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