Skip to main content

xz_provider/
key_source.rs

1use async_trait::async_trait;
2
3use crate::error::ProviderError;
4
5/// 用于获取 LLM Provider API Key 的可插拔 trait。
6///
7/// 不同场景下 API Key 来源不同:
8/// - 配置文件(`ConfigKeySource`)
9/// - 用户自己提供(`UserKeySource`)
10///
11/// 所有实现均通过本 trait 统一接口接入 `ProviderRouter`。
12#[async_trait]
13pub trait KeySource: Send + Sync {
14    /// 获取 API Key。
15    ///
16    /// 对于 ConfigKeySource:直接从配置中读取。
17    async fn get_api_key(&self) -> Result<String, ProviderError>;
18}
19
20#[cfg(test)]
21mod tests {
22    use super::*;
23
24    struct TestSource {
25        key: String,
26    }
27
28    #[async_trait]
29    impl KeySource for TestSource {
30        async fn get_api_key(&self) -> Result<String, ProviderError> {
31            Ok(self.key.clone())
32        }
33    }
34
35    #[tokio::test]
36    async fn trait_is_object_safe() {
37        use std::sync::Arc;
38
39        let source: Arc<dyn KeySource> = Arc::new(TestSource { key: "sk-test".into() });
40        let api_key = source.get_api_key().await.ok();
41
42        assert_eq!(api_key.as_deref(), Some("sk-test"));
43    }
44
45    struct ErrSource;
46
47    #[async_trait]
48    impl KeySource for ErrSource {
49        async fn get_api_key(&self) -> Result<String, ProviderError> {
50            Err(ProviderError::KeySource("test error".into()))
51        }
52    }
53
54    #[tokio::test]
55    async fn trait_returns_error() {
56        let source = ErrSource;
57
58        assert!(source.get_api_key().await.is_err());
59    }
60}