如何减少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
相关产品推荐
相关产品推荐

