Skip to main content

agent_base/llm/
registry.rs

1use std::sync::Arc;
2
3use super::{AnthropicClient, LlmClient, OpenAiClient};
4
5#[derive(Clone, Debug)]
6pub enum LlmProvider {
7    OpenAi,
8    Anthropic,
9    Custom(String),
10}
11
12impl LlmProvider {
13    #[allow(clippy::should_implement_trait)]
14    pub fn from_str(s: &str) -> Self {
15        match s.to_lowercase().as_str() {
16            "openai" => Self::OpenAi,
17            "anthropic" => Self::Anthropic,
18            other => Self::Custom(other.to_string()),
19        }
20    }
21}
22
23pub struct LlmClientBuilder {
24    provider: LlmProvider,
25    api_key: String,
26    model: String,
27    base_url: Option<String>,
28}
29
30impl LlmClientBuilder {
31    pub fn new(
32        provider: LlmProvider,
33        api_key: impl Into<String>,
34        model: impl Into<String>,
35    ) -> Self {
36        Self {
37            provider,
38            api_key: api_key.into(),
39            model: model.into(),
40            base_url: None,
41        }
42    }
43
44    pub fn from_env() -> Option<Self> {
45        let api_key = std::env::var("LLM_API_KEY").ok()?;
46        let model = std::env::var("LLM_MODEL").unwrap_or_else(|_| "gpt-4o".to_string());
47        let base_url = std::env::var("LLM_BASE_URL").ok();
48        let provider_str = std::env::var("LLM_PROVIDER").unwrap_or_else(|_| "openai".to_string());
49
50        Some(Self {
51            provider: LlmProvider::from_str(&provider_str),
52            api_key,
53            model,
54            base_url,
55        })
56    }
57
58    pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
59        self.base_url = Some(base_url.into());
60        self
61    }
62
63    pub fn build(self) -> Arc<dyn LlmClient> {
64        let base_url = self.base_url;
65        match self.provider {
66            LlmProvider::OpenAi => {
67                let url = base_url.unwrap_or_else(|| "https://api.openai.com/v1".to_string());
68                Arc::new(OpenAiClient::new(self.api_key, self.model, Some(url)))
69            }
70            LlmProvider::Anthropic => {
71                let url = base_url.unwrap_or_else(|| "https://api.anthropic.com".to_string());
72                Arc::new(AnthropicClient::new(self.api_key, self.model, Some(url)))
73            }
74            LlmProvider::Custom(_) => {
75                let url = base_url.unwrap_or_else(|| "https://api.openai.com/v1".to_string());
76                Arc::new(OpenAiClient::new(self.api_key, self.model, Some(url)))
77            }
78        }
79    }
80}
81
82#[cfg(test)]
83mod tests {
84    use super::*;
85
86    #[test]
87    fn provider_from_str_is_case_insensitive() {
88        assert!(matches!(
89            LlmProvider::from_str("openai"),
90            LlmProvider::OpenAi
91        ));
92        assert!(matches!(
93            LlmProvider::from_str("OpenAI"),
94            LlmProvider::OpenAi
95        ));
96        assert!(matches!(
97            LlmProvider::from_str("OPENAI"),
98            LlmProvider::OpenAi
99        ));
100        assert!(matches!(
101            LlmProvider::from_str("anthropic"),
102            LlmProvider::Anthropic
103        ));
104        assert!(matches!(
105            LlmProvider::from_str("Anthropic"),
106            LlmProvider::Anthropic
107        ));
108    }
109
110    #[test]
111    fn provider_from_str_unknown_becomes_custom() {
112        assert!(matches!(
113            LlmProvider::from_str("ollama"),
114            LlmProvider::Custom(ref s) if s == "ollama"
115        ));
116        assert!(matches!(
117            LlmProvider::from_str(""),
118            LlmProvider::Custom(ref s) if s.is_empty()
119        ));
120    }
121
122    #[test]
123    fn build_routes_openai() {
124        let client = LlmClientBuilder::new(LlmProvider::OpenAi, "sk-test", "gpt-4o").build();
125        assert_eq!(client.model_name(), "gpt-4o");
126        assert_eq!(client.capabilities().max_context_tokens, Some(128_000));
127    }
128
129    #[test]
130    fn build_routes_anthropic() {
131        let client = LlmClientBuilder::new(LlmProvider::Anthropic, "sk-ant", "claude").build();
132        assert_eq!(client.model_name(), "claude");
133        assert_eq!(client.capabilities().max_context_tokens, Some(200_000));
134    }
135
136    #[test]
137    fn build_custom_defaults_to_openai() {
138        let client =
139            LlmClientBuilder::new(LlmProvider::Custom("ollama".into()), "sk", "llama").build();
140        assert_eq!(client.model_name(), "llama");
141        assert_eq!(client.capabilities().max_context_tokens, Some(128_000));
142    }
143
144    #[test]
145    fn base_url_is_chainable() {
146        let client = LlmClientBuilder::new(LlmProvider::OpenAi, "sk", "gpt-4o")
147            .base_url("http://localhost:9999/v1")
148            .build();
149        assert_eq!(client.model_name(), "gpt-4o");
150    }
151
152    #[test]
153    fn builder_from_env_reads_vars_and_requires_key() {
154        // Single test (no intra-module parallelism) so env mutations can't race.
155        unsafe {
156            std::env::remove_var("LLM_API_KEY");
157            std::env::remove_var("LLM_MODEL");
158            std::env::remove_var("LLM_BASE_URL");
159            std::env::remove_var("LLM_PROVIDER");
160        }
161        assert!(LlmClientBuilder::from_env().is_none());
162
163        unsafe {
164            std::env::set_var("LLM_API_KEY", "env-key");
165            std::env::set_var("LLM_MODEL", "env-model");
166            std::env::set_var("LLM_BASE_URL", "http://env.test/v1");
167            std::env::set_var("LLM_PROVIDER", "anthropic");
168        }
169        let client = LlmClientBuilder::from_env().expect("all vars set").build();
170        assert_eq!(client.model_name(), "env-model");
171        assert_eq!(client.capabilities().max_context_tokens, Some(200_000));
172
173        unsafe {
174            std::env::remove_var("LLM_API_KEY");
175            std::env::remove_var("LLM_MODEL");
176            std::env::remove_var("LLM_BASE_URL");
177            std::env::remove_var("LLM_PROVIDER");
178        }
179    }
180}