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

如何减少Rust中仅因.await差异导致的同步异步代码重复?

问题描述

我正在开发一个Rust代码库,计划通过feature flag提供同步与异步版本,确保整个生态一致性,避免异步与同步代码转换的复杂操作。但实现中出现大量代码重复,仅差异在于异步版本需对异步函数调用添加.await,每次修改都要维护两处,操作繁琐。

重复代码示例

#[cfg(feature = "async")]
pub async fn new_from_env(realm_id: &str, environment: Environment) -> Result<Self, AuthError> {
    let discovery_doc = Self::get_discovery_doc(&environment).await?; // 唯一差异
    let client_id = ClientId::new(std::env::var("INTUIT_CLIENT_ID")?);
    let client_secret = ClientSecret::new(std::env::var("INTUIT_CLIENT_SECRET")?);
    let redirect_uri = RedirectUrl::new(std::env::var("INTUIT_REDIRECT_URI")?)?;
    log::info!("Got Discovery Doc and Intuit Credentials Successfully");
    Ok(Self {
        redirect_uri,
        realm_id: realm_id.to_string(),
        environment,
        data: Unauthorized {
            client_id,
            client_secret,
            discovery_doc,
        },
    })
}

#[cfg(not(feature = "async"))]
pub fn new_from_env(realm_id: &str, environment: Environment) -> Result<Self, AuthError> {
    let discovery_doc = Self::get_discovery_doc(&environment)?; // 无await
    let client_id = ClientId::new(std::env::var("INTUIT_CLIENT_ID")?);
    let client_secret = ClientSecret::new(std::env::var("INTUIT_CLIENT_SECRET")?);
    let redirect_uri = RedirectUrl::new(std::env::var("INTUIT_REDIRECT_URI")?)?;
    log::info!("Got Discovery Doc and Intuit Credentials Successfully");
    Ok(Self {
        redirect_uri,
        realm_id: realm_id.to_string(),
        environment,
        data: Unauthorized {
            client_id,
            client_secret,
            discovery_doc,
        },
    })
}

#[cfg(feature = "async")]
async fn default_grab_token_session(
    client_ref: &BasicClient,
    scopes: Option<&[IntuitScope]>,
) -> Result<TokenSession, AuthError> {
    let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
    let (auth_url, csrf_state) = Self::get_auth_url(client_ref, pkce_challenge, scopes);
    let listener = TcpListener::bind("127.0.0.1:3320")
        .await // 此处有await
        .expect("Error starting localhost callback listener! (async)");
    open::that_detached(auth_url.as_str())?;
    log::info!("Opened Auth URL: {}", auth_url);
    Self::handle_oauth_callback(client_ref, listener, csrf_state, pkce_verifier).await
}

#[cfg(not(feature = "async"))]
fn default_grab_token_session(
    client_ref: &BasicClient,
    scopes: Option<&[IntuitScope]>,
) -> Result<TokenSession, AuthError> {
    let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
    let (auth_url, csrf_state) = Self::get_auth_url(client_ref, pkce_challenge, scopes);
    let listener = TcpListener::bind("127.0.0.1:3320")
        .expect("Error starting localhost callback listener!"); // 无await
    open::that_detached(auth_url.as_str())?;
    log::info!("Opened Auth URL: {}", auth_url);
    Self::handle_oauth_callback(client_ref, listener, csrf_state, pkce_verifier)
}

已尝试的方案

曾使用duplicate crate及自定义macro_rules!宏解决,但因无法简单判断类型是否可被await,效果不佳。示例宏代码如下:

#[cfg(feature = "async")]
use tokio::fs;
#[cfg(not(feature = "async"))]
use std::fs;

macro_rules! cfg_async {
    ($func_name:ident ($($args:ident: $state_ty:ty),*) -> $output:ident { $body:expr } ) => {
      #[cfg(feature = "async")]
      async fn $func_name($($args: $state_ty),*) -> $output {
         // 无法自动判断是否需要await
         $body
      }
  
      #[cfg(not(feature = "async"))]
      fn $func_name($($args: $state_ty),*) -> $output {
         $body
      }
    };
}

cfg_async!(foo (path: &str) -> String {
    fs::read_to_string(path).unwrap()
    // 异步场景下需要改为fs::read_to_string("foo.txt").await.unwrap()
});

求简洁高效的解决方法?


解决方案

1. 自定义宏注入条件式await

改进宏逻辑,通过显式标记需要处理的调用,让宏在编译时根据feature flag自动添加或省略.await:

// 处理单个表达式的await条件注入
macro_rules! await_if_async {
    ($expr:expr) => {
        #[cfg(feature = "async")]
        { $expr.await }
        #[cfg(not(feature = "async"))]
        { $expr }
    };
}

// 定义同步/异步双版本函数的宏
macro_rules! cfg_async_fn {
    ($vis:vis $func_name:ident ($($args:ident: $state_ty:ty),*) -> $output:ty { $body:block } ) => {
        #[cfg(feature = "async")]
        $vis async fn $func_name($($args: $state_ty),*) -> $output {
            $body
        }

        #[cfg(not(feature = "async"))]
        $vis fn $func_name($($args: $state_ty),*) -> $output {
            $body
        }
    };
}

使用示例:

cfg_async_fn!(pub new_from_env (realm_id: &str, environment: Environment) -> Result<Self, AuthError> {
    let discovery_doc = await_if_async!(Self::get_discovery_doc(&environment))?;
    let client_id = ClientId::new(std::env::var("INTUIT_CLIENT_ID")?);
    let client_secret = ClientSecret::new(std::env::var("INTUIT_CLIENT_SECRET")?);
    let redirect_uri = RedirectUrl::new(std::env::var("INTUIT_REDIRECT_URI")?)?;
    log::info!("Got Discovery Doc and Intuit Credentials Successfully");
    Ok(Self {
        redirect_uri,
        realm_id: realm_id.to_string(),
        environment,
        data: Unauthorized {
            client_id,
            client_secret,
            discovery_doc,
        },
    })
});

2. 封装Awaitable trait抽象差异

定义一个统一处理同步结果/异步Future的trait,让核心逻辑无需区分同步异步:

#[cfg(feature = "async")]
use std::future::Future;

trait Awaitable<T> {
    fn resolve(self) -> Result<T, AuthError>;
}

// 异步场景下,Future转为阻塞调用
#[cfg(feature = "async")]
impl<T, E: Into<AuthError>> Awaitable<T> for impl Future<Output = Result<T, E>> {
    fn resolve(self) -> Result<T, AuthError> {
        futures::executor::block_on(self).map_err(Into::into)
    }
}

// 同步场景下,直接返回结果
#[cfg(not(feature = "async"))]
impl<T, E: Into<AuthError>> Awaitable<T> for Result<T, E> {
    fn resolve(self) -> Result<T, AuthError> {
        self.map_err(Into::into)
    }
}

核心逻辑写一次,再通过条件编译导出同步/异步接口:

fn new_from_env_core(realm_id: &str, environment: Environment) -> Result<Self, AuthError> {
    let discovery_doc = Self::get_discovery_doc(&environment).resolve()?;
    // 其余逻辑完全不变
    let client_id = ClientId::new(std::env::var("INTUIT_CLIENT_ID")?);
    let client_secret = ClientSecret::new(std::env::var("INTUIT_CLIENT_SECRET")?);
    let redirect_uri = RedirectUrl::new(std::env::var("INTUIT_REDIRECT_URI")?)?;
    log::info!("Got Discovery Doc and Intuit Credentials Successfully");
    Ok(Self {
        redirect_uri,
        realm_id: realm_id.to_string(),
        environment,
        data: Unauthorized {
            client_id,
            client_secret,
            discovery_doc,
        },
    })
}

#[cfg(feature = "async")]
pub async fn new_from_env(realm_id: &str, environment: Environment) -> Result<Self, AuthError> {
    new_from_env_core(realm_id, environment)
}

#[cfg(not(feature = "async"))]
pub fn new_from_env(realm_id: &str, environment: Environment) -> Result<Self, AuthError> {
    new_from_env_core(realm_id, environment)
}

3. 抽象同步/异步组件类型

将依赖的异步/同步组件(如TcpListener、HTTP客户端)抽象为trait,通过feature flag选择具体实现,核心逻辑基于trait编写:

trait Listener {
    type Stream;
    fn bind(addr: &str) -> Result<Self, std::io::Error>;
    #[cfg(feature = "async")]
    async fn accept(&self) -> Result<(Self::Stream, std::net::SocketAddr), std::io::Error>;
    #[cfg(not(feature = "async"))]
    fn accept(&self) -> Result<(Self::Stream, std::net::SocketAddr), std::io::Error>;
}

// 异步场景绑定tokio的TcpListener
#[cfg(feature = "async")]
impl Listener for tokio::net::TcpListener {
    type Stream = tokio::net::TcpStream;
    fn bind(addr: &str) -> Result<Self, std::io::Error> {
        tokio::net::TcpListener::bind(addr).map_err(Into::into)
    }
    async fn accept(&self) -> Result<(Self::Stream, std::net::SocketAddr), std::io::Error> {
        self.accept().await
    }
}

// 同步场景绑定标准库的TcpListener
#[cfg(not(feature = "async"))]
impl Listener for std::net::TcpListener {
    type Stream = std::net::TcpStream;
    fn bind(addr: &str) -> Result<Self, std::io::Error> {
        std::net::TcpListener::bind(addr)
    }
    fn accept(&self) -> Result<(Self::Stream, std::net::SocketAddr), std::io::Error> {
        self.accept()
    }
}

配合类型别名和之前的宏使用:

#[cfg(feature = "async")]
type AppListener = tokio::net::TcpListener;
#[cfg(not(feature = "async"))]
type AppListener = std::net::TcpListener;

cfg_async_fn!(async fn default_grab_token_session (
    client_ref: &BasicClient,
    scopes: Option<&[IntuitScope]>,
) -> Result<TokenSession, AuthError> {
    let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
    let (auth_url, csrf_state) = Self::get_auth_url(client_ref, pkce_challenge, scopes);
    let listener = AppListener::bind("127.0.0.1:3320")
        .expect("Error starting localhost callback listener!");
    open::that_detached(auth_url.as_str())?;
    log::info!("Opened Auth URL: {}", auth_url);
    await_if_async!(Self::handle_oauth_callback(client_ref, listener, csrf_state, pkce_verifier))
});

内容的提问来源于stack exchange,提问作者Tristen Gordon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 15:12:31