基于Tokio的Rust递归API调用高效实现方案问询
高效实现递归函数f(n)的Rust方案
问题描述
给定数字n和API端点,实现函数f(n)并输出结果,规则为:
f(0) = 1
f(2) = 2
当n为偶数时,f(n) = f(n-1)+f(n-2)+f(n-3)
当n为奇数时,f(n) = CallAPI(n)
CallAPI返回0-100的数字,相同n返回值固定,需尽可能少调用。
原实现采用memoization+Mutex,但因频繁锁竞争导致性能不佳,求高效优化方案。
原代码实现:
use std::{ collections::HashMap, sync::{Arc, Mutex}, }; use async_recursion::async_recursion; type Memo = Arc<Mutex<HashMap<usize, usize>>>; #[tokio::main] pub async fn main() -> Result<(), Box<dyn std::error::Error>> { // 实际场景中n会作为输入参数传入 let n = "60"; let mut temp = HashMap::new(); temp.insert(0, 1); temp.insert(2, 2); let memo: Memo = Arc::new(Mutex::new(temp)); let ans = f(n.parse::<usize>().expect("error"), memo).await; println!("{}", ans); Ok(()) } #[async_recursion] async fn f(n: usize, memo: Memo) -> usize { if let Some(&ans) = memo.lock().unwrap().get(&n) { return ans; } if n % 2 == 0 { let memo1 = memo.clone(); let memo2 = memo.clone(); let memo3 = memo.clone(); let (one, two, three) = tokio::join!(f(n - 1, memo1), f(n - 2, memo2), f(n - 3, memo3)); let ans = one + two + three; let mut mut_memo = memo.lock().unwrap(); mut_memo.insert(n, ans); ans } else { let n_str = n.to_string(); let ans = ask_server(&n_str).await; let mut mut_memo = memo.lock().unwrap(); mut_memo.insert(n, ans); ans } } async fn ask_server(n: &String) -> usize { // 实际场景中会发送HTTP GET请求并返回响应中的数字,此处简化返回100 100 }
优化方案
1. 替换Mutex为高效并发数据结构
原代码用Arc<Mutex<HashMap>>,每次读写都需要独占锁,并发场景下竞争激烈。推荐使用**DashMap**——一个基于分片锁实现的高性能并发哈希表,读写操作只会锁住对应数据分片,大幅降低锁冲突概率。
2. 移除异步递归,改用迭代计算
异步递归会产生大量Future对象,且async_recursion宏带来额外开销。观察递推规则可知,偶数n的f(n)仅依赖更小的n值,完全可以从低到高迭代计算,避免递归开销,同时确保每个值只计算一次。
3. 最小化API调用次数
每个奇数n的API结果只会被调用一次并缓存,迭代过程中直接读取缓存值,完全满足"尽可能少调用"的要求。
优化后代码
use dashmap::DashMap; use std::sync::Arc; #[tokio::main] pub async fn main() -> Result<(), Box<dyn std::error::Error>> { let n = "60".parse::<usize>().expect("invalid number"); let memo = Arc::new(DashMap::new()); // 初始化已知固定值 memo.insert(0, 1); memo.insert(2, 2); let ans = compute_f(n, memo).await; println!("{}", ans); Ok(()) } async fn compute_f(n: usize, memo: Arc<DashMap<usize, usize>>) -> usize { // 直接返回已缓存的值 if let Some(&val) = memo.get(&n) { return val; } // 从3开始迭代计算到n的所有值 for i in 3..=n { if memo.contains_key(&i) { continue; } if i % 2 == 1 { // 奇数调用API并缓存结果 let ans = ask_server(&i.to_string()).await; memo.insert(i, ans); } else { // 偶数直接取前三个缓存值计算 let a = memo.get(&(i-1)).unwrap().clone(); let b = memo.get(&(i-2)).unwrap().clone(); let c = memo.get(&(i-3)).unwrap().clone(); let ans = a + b + c; memo.insert(i, ans); } } memo.get(&n).unwrap().clone() } async fn ask_server(n: &str) -> usize { // 实际场景为HTTP请求,此处简化返回100 100 }
备选方案:RwLock替代Mutex
若不需要极致并发性能,也可以用Arc<RwLock<HashMap>>替代Mutex,读操作使用共享锁,写操作使用独占锁,提升读多写少场景的性能:
use std::collections::HashMap; use std::sync::{Arc, RwLock}; type Memo = Arc<RwLock<HashMap<usize, usize>>>; #[tokio::main] pub async fn main() -> Result<(), Box<dyn std::error::Error>> { let n = "60".parse::<usize>().expect("invalid number"); let memo: Memo = Arc::new(RwLock::new({ let mut m = HashMap::new(); m.insert(0, 1); m.insert(2, 2); m })); let ans = compute_f(n, memo).await; println!("{}", ans); Ok(()) } async fn compute_f(n: usize, memo: Memo) -> usize { // 读锁查询缓存 if let Ok(memo_read) = memo.read() { if let Some(&val) = memo_read.get(&n) { return val; } } for i in 3..=n { if let Ok(memo_read) = memo.read() { if memo_read.contains_key(&i) { continue; } } if i % 2 == 1 { let ans = ask_server(&i.to_string()).await; memo.write().unwrap().insert(i, ans); } else { let memo_read = memo.read().unwrap(); let a = memo_read.get(&(i-1)).unwrap().clone(); let b = memo_read.get(&(i-2)).unwrap().clone(); let c = memo_read.get(&(i-3)).unwrap().clone(); drop(memo_read); // 提前释放读锁 let ans = a + b + c; memo.write().unwrap().insert(i, ans); } } memo.read().unwrap().get(&n).unwrap().clone() } async fn ask_server(n: &str) -> usize { 100 }
内容的提问来源于stack exchange,提问作者S0i
相关产品推荐
相关产品推荐

