agent_base/llm/
registry.rs1use 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 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 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 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}