Skip to main content

agent_base/llm/
registry.rs

1use std::sync::Arc;
2
3use super::{AnthropicClient, LlmClient, OpenAiClient, OpenAiResponsesClient, StreamClient};
4
5#[derive(Clone, Debug)]
6pub enum LlmProvider {
7    OpenAi,
8    OpenAiResponses,
9    Anthropic,
10    Custom(String),
11}
12
13impl LlmProvider {
14    #[allow(clippy::should_implement_trait)]
15    pub fn from_str(s: &str) -> Self {
16        match s.to_lowercase().as_str() {
17            "openai" => Self::OpenAi,
18            "openai-responses" | "responses" => Self::OpenAiResponses,
19            "anthropic" => Self::Anthropic,
20            other => Self::Custom(other.to_string()),
21        }
22    }
23}
24
25pub struct LlmClientBuilder {
26    provider: LlmProvider,
27    api_key: String,
28    model: String,
29    base_url: Option<String>,
30}
31
32impl LlmClientBuilder {
33    pub fn new(
34        provider: LlmProvider,
35        api_key: impl Into<String>,
36        model: impl Into<String>,
37    ) -> Self {
38        Self {
39            provider,
40            api_key: api_key.into(),
41            model: model.into(),
42            base_url: None,
43        }
44    }
45
46    pub fn from_env() -> Option<Self> {
47        let api_key = std::env::var("LLM_API_KEY").ok()?;
48        let model = std::env::var("LLM_MODEL").unwrap_or_else(|_| "gpt-4o".to_string());
49        let base_url = std::env::var("LLM_BASE_URL").ok();
50        let provider_str = std::env::var("LLM_PROVIDER").unwrap_or_else(|_| "openai".to_string());
51
52        Some(Self {
53            provider: LlmProvider::from_str(&provider_str),
54            api_key,
55            model,
56            base_url,
57        })
58    }
59
60    pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
61        self.base_url = Some(base_url.into());
62        self
63    }
64
65    pub fn build(self) -> Arc<dyn LlmClient> {
66        let base_url = self.base_url;
67        match self.provider {
68            LlmProvider::OpenAi => {
69                let url = base_url.unwrap_or_else(|| "https://api.openai.com/v1".to_string());
70                Arc::new(OpenAiClient::new(self.api_key, self.model, Some(url)))
71            }
72            LlmProvider::Anthropic => {
73                let url = base_url.unwrap_or_else(|| "https://api.anthropic.com".to_string());
74                Arc::new(AnthropicClient::new(self.api_key, self.model, Some(url)))
75            }
76            LlmProvider::OpenAiResponses => {
77                panic!(
78                    "OpenAiResponsesClient implements StreamClient, not LlmClient. \
79                     Use build_stream_client() instead of build()."
80                )
81            }
82            LlmProvider::Custom(_) => {
83                let url = base_url.unwrap_or_else(|| "https://api.openai.com/v1".to_string());
84                Arc::new(OpenAiClient::new(self.api_key, self.model, Some(url)))
85            }
86        }
87    }
88
89    /// Build an `Arc<dyn StreamClient>`.
90    ///
91    /// Use this for providers that implement [`StreamClient`] directly
92    /// (e.g. [`OpenAiResponsesClient`]) rather than the legacy [`LlmClient`] trait.
93    pub fn build_stream_client(self) -> Arc<dyn StreamClient> {
94        match &self.provider {
95            LlmProvider::OpenAiResponses => {
96                let url = self
97                    .base_url
98                    .clone()
99                    .unwrap_or_else(|| "https://api.openai.com/v1".to_string());
100                Arc::new(OpenAiResponsesClient::new(
101                    self.api_key,
102                    self.model,
103                    Some(url),
104                ))
105            }
106            _ => {
107                // For other providers, wrap the LlmClient in an adapter.
108                super::adapt(self.build())
109            }
110        }
111    }
112}
113
114#[cfg(test)]
115mod tests {
116    use super::*;
117
118    #[test]
119    fn provider_from_str_is_case_insensitive() {
120        assert!(matches!(
121            LlmProvider::from_str("openai"),
122            LlmProvider::OpenAi
123        ));
124        assert!(matches!(
125            LlmProvider::from_str("OpenAI"),
126            LlmProvider::OpenAi
127        ));
128        assert!(matches!(
129            LlmProvider::from_str("OPENAI"),
130            LlmProvider::OpenAi
131        ));
132        assert!(matches!(
133            LlmProvider::from_str("anthropic"),
134            LlmProvider::Anthropic
135        ));
136        assert!(matches!(
137            LlmProvider::from_str("Anthropic"),
138            LlmProvider::Anthropic
139        ));
140    }
141
142    #[test]
143    fn provider_from_str_openai_responses() {
144        assert!(matches!(
145            LlmProvider::from_str("openai-responses"),
146            LlmProvider::OpenAiResponses
147        ));
148        assert!(matches!(
149            LlmProvider::from_str("OpenAI-Responses"),
150            LlmProvider::OpenAiResponses
151        ));
152        assert!(matches!(
153            LlmProvider::from_str("responses"),
154            LlmProvider::OpenAiResponses
155        ));
156    }
157
158    #[test]
159    fn provider_from_str_unknown_becomes_custom() {
160        assert!(matches!(
161            LlmProvider::from_str("ollama"),
162            LlmProvider::Custom(ref s) if s == "ollama"
163        ));
164        assert!(matches!(
165            LlmProvider::from_str(""),
166            LlmProvider::Custom(ref s) if s.is_empty()
167        ));
168    }
169
170    #[test]
171    fn build_routes_openai() {
172        let client = LlmClientBuilder::new(LlmProvider::OpenAi, "sk-test", "gpt-4o").build();
173        assert_eq!(client.model_name(), "gpt-4o");
174        assert_eq!(client.capabilities().max_context_tokens, Some(128_000));
175    }
176
177    #[test]
178    fn build_routes_anthropic() {
179        let client = LlmClientBuilder::new(LlmProvider::Anthropic, "sk-ant", "claude").build();
180        assert_eq!(client.model_name(), "claude");
181        assert_eq!(client.capabilities().max_context_tokens, Some(200_000));
182    }
183
184    #[test]
185    fn build_custom_defaults_to_openai() {
186        let client =
187            LlmClientBuilder::new(LlmProvider::Custom("ollama".into()), "sk", "llama").build();
188        assert_eq!(client.model_name(), "llama");
189        assert_eq!(client.capabilities().max_context_tokens, Some(128_000));
190    }
191
192    #[test]
193    fn build_stream_client_responses() {
194        let client = LlmClientBuilder::new(LlmProvider::OpenAiResponses, "sk", "gpt-4o")
195            .build_stream_client();
196        assert_eq!(client.model_name(), "gpt-4o");
197        assert_eq!(client.capabilities().max_context_tokens, Some(128_000));
198    }
199
200    #[test]
201    #[should_panic(expected = "build_stream_client")]
202    fn build_panics_for_responses_provider() {
203        LlmClientBuilder::new(LlmProvider::OpenAiResponses, "sk", "gpt-4o").build();
204    }
205
206    #[test]
207    fn build_stream_client_falls_back_to_adapter() {
208        let client =
209            LlmClientBuilder::new(LlmProvider::OpenAi, "sk", "gpt-4o").build_stream_client();
210        assert_eq!(client.model_name(), "gpt-4o");
211    }
212
213    #[test]
214    fn base_url_is_chainable() {
215        let client = LlmClientBuilder::new(LlmProvider::OpenAi, "sk", "gpt-4o")
216            .base_url("http://localhost:9999/v1")
217            .build();
218        assert_eq!(client.model_name(), "gpt-4o");
219    }
220
221    #[test]
222    fn builder_from_env_reads_vars_and_requires_key() {
223        // Single test (no intra-module parallelism) so env mutations can't race.
224        unsafe {
225            std::env::remove_var("LLM_API_KEY");
226            std::env::remove_var("LLM_MODEL");
227            std::env::remove_var("LLM_BASE_URL");
228            std::env::remove_var("LLM_PROVIDER");
229        }
230        assert!(LlmClientBuilder::from_env().is_none());
231
232        unsafe {
233            std::env::set_var("LLM_API_KEY", "env-key");
234            std::env::set_var("LLM_MODEL", "env-model");
235            std::env::set_var("LLM_BASE_URL", "http://env.test/v1");
236            std::env::set_var("LLM_PROVIDER", "anthropic");
237        }
238        let client = LlmClientBuilder::from_env().expect("all vars set").build();
239        assert_eq!(client.model_name(), "env-model");
240        assert_eq!(client.capabilities().max_context_tokens, Some(200_000));
241
242        unsafe {
243            std::env::remove_var("LLM_API_KEY");
244            std::env::remove_var("LLM_MODEL");
245            std::env::remove_var("LLM_BASE_URL");
246            std::env::remove_var("LLM_PROVIDER");
247        }
248    }
249}