You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 18:55:16