1use crate::error::ProviderError;
23use crate::ollama::OllamaChat;
24use crate::ollama::OllamaConfig;
25use crate::openai::OpenAIChat;
26use crate::openai::OpenAIConfig;
27use crate::providers::anthropic::AnthropicChat;
28use crate::providers::anthropic::AnthropicConfig;
29use crate::providers::azure::AzureOpenAIChat;
30use crate::providers::azure::AzureOpenAIConfig;
31use crate::providers::cohere::CohereChat;
32use crate::providers::cohere::CohereConfig;
33use crate::providers::deepseek::DeepSeekChat;
34use crate::providers::deepseek::DeepSeekConfig;
35use crate::providers::gemini::GeminiChat;
36use crate::providers::gemini::GeminiConfig;
37use crate::providers::mistral::MistralChat;
38use crate::providers::mistral::MistralConfig;
39use crate::providers::moonshot::MoonshotChat;
40use crate::providers::moonshot::MoonshotConfig;
41use crate::providers::qwen::QwenChat;
42use crate::providers::qwen::QwenConfig;
43use crate::providers::zhipu::ZhipuChat;
44use crate::providers::zhipu::ZhipuConfig;
45use crate::wrap_chat_model;
46use async_trait::async_trait;
47use futures_util::Stream;
48use lc_core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult, StreamChunk};
49use lc_core::runnables::Runnable;
50use lc_core::tools::ToolDefinition;
51use lc_core::RunnableConfig;
52use lc_schema::Message;
53use std::sync::{Arc, Mutex};
54
55#[derive(Debug, Default)]
62struct ClientOverrides {
63 temperature: Option<f32>,
64 max_tokens: Option<usize>,
65}
66
67pub struct LLMClient {
74 inner: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
75 overrides: Mutex<ClientOverrides>,
76}
77
78impl LLMClient {
79 fn from_inner(inner: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>) -> Self {
81 Self {
82 inner,
83 overrides: Mutex::new(ClientOverrides::default()),
84 }
85 }
86
87 fn apply_overrides(&self, config: Option<RunnableConfig>) -> Option<RunnableConfig> {
89 let overrides = self.overrides.lock().unwrap_or_else(|e| e.into_inner());
90 if overrides.temperature.is_none() && overrides.max_tokens.is_none() {
91 return config;
92 }
93 let mut cfg = config.unwrap_or_default();
94 cfg.temperature = overrides.temperature.or(cfg.temperature);
95 cfg.max_tokens = overrides.max_tokens.or(cfg.max_tokens);
96 Some(cfg)
97 }
98}
99
100impl std::fmt::Debug for LLMClient {
101 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
102 f.debug_struct("LLMClient")
103 .field("model_name", &self.inner.model_name())
104 .finish()
105 }
106}
107
108impl LLMClient {
109 pub fn from_env() -> Result<Self, ProviderError> {
132 if std::env::var("OPENAI_API_KEY").is_ok() {
134 return Ok(Self::openai(OpenAIConfig::from_env_result()?));
135 }
136
137 if std::env::var("ANTHROPIC_API_KEY").is_ok() {
139 return Ok(Self::anthropic(AnthropicConfig::from_env_result()?));
140 }
141
142 if std::env::var("AZURE_OPENAI_API_KEY").is_ok() {
144 return Ok(Self::azure(AzureOpenAIConfig::from_env_result()?));
145 }
146
147 if std::env::var("DEEPSEEK_API_KEY").is_ok() {
149 return Ok(Self::deepseek(DeepSeekConfig::from_env_result()?));
150 }
151
152 if std::env::var("QWEN_API_KEY").is_ok() {
154 return Ok(Self::qwen(QwenConfig::from_env_result()?));
155 }
156
157 if std::env::var("MOONSHOT_API_KEY").is_ok() {
159 return Ok(Self::moonshot(MoonshotConfig::from_env_result()?));
160 }
161
162 if std::env::var("ZHIPU_API_KEY").is_ok() {
164 return Ok(Self::zhipu(ZhipuConfig::from_env_result()?));
165 }
166
167 if std::env::var("MISTRAL_API_KEY").is_ok() {
169 return Ok(Self::mistral(MistralConfig::from_env_result()?));
170 }
171
172 if std::env::var("COHERE_API_KEY").is_ok() {
174 return Ok(Self::cohere(CohereConfig::from_env_result()?));
175 }
176
177 if std::env::var("GEMINI_API_KEY").is_ok() || std::env::var("GOOGLE_API_KEY").is_ok() {
179 return Ok(Self::gemini(GeminiConfig::from_env_result()?));
180 }
181
182 if std::env::var("OLLAMA_BASE_URL").is_ok() {
184 return Ok(Self::ollama(OllamaConfig::from_env_result()?));
185 }
186
187 Err(ProviderError::Config(
188 "No LLM provider detected. Set one of: OPENAI_API_KEY, ANTHROPIC_API_KEY, \
189 AZURE_OPENAI_API_KEY, DEEPSEEK_API_KEY, QWEN_API_KEY, MOONSHOT_API_KEY, \
190 ZHIPU_API_KEY, MISTRAL_API_KEY, COHERE_API_KEY, GEMINI_API_KEY, OLLAMA_BASE_URL"
191 .to_string(),
192 ))
193 }
194
195 pub fn openai(config: OpenAIConfig) -> Self {
201 let llm = OpenAIChat::new(config);
202 Self::from_inner(wrap_chat_model(llm))
203 }
204
205 pub fn anthropic(config: AnthropicConfig) -> Self {
207 let llm = AnthropicChat::new(config);
208 Self::from_inner(wrap_chat_model(llm))
209 }
210
211 pub fn ollama(config: OllamaConfig) -> Self {
213 let llm = OllamaChat::with_config(config);
214 Self::from_inner(wrap_chat_model(llm))
215 }
216
217 pub fn gemini(config: GeminiConfig) -> Self {
219 let llm = GeminiChat::new(config);
220 Self::from_inner(wrap_chat_model(llm))
221 }
222
223 pub fn deepseek(config: DeepSeekConfig) -> Self {
225 let llm = DeepSeekChat::new(config);
226 Self::from_inner(wrap_chat_model(llm))
227 }
228
229 pub fn qwen(config: QwenConfig) -> Self {
231 let llm = QwenChat::new(config);
232 Self::from_inner(wrap_chat_model(llm))
233 }
234
235 pub fn moonshot(config: MoonshotConfig) -> Self {
237 let llm = MoonshotChat::new(config);
238 Self::from_inner(wrap_chat_model(llm))
239 }
240
241 pub fn zhipu(config: ZhipuConfig) -> Self {
243 let llm = ZhipuChat::new(config);
244 Self::from_inner(wrap_chat_model(llm))
245 }
246
247 pub fn mistral(config: MistralConfig) -> Self {
249 let llm = MistralChat::new(config);
250 Self::from_inner(wrap_chat_model(llm))
251 }
252
253 pub fn azure(config: AzureOpenAIConfig) -> Self {
255 let llm = AzureOpenAIChat::new(config);
256 Self::from_inner(wrap_chat_model(llm))
257 }
258
259 pub fn cohere(config: CohereConfig) -> Self {
261 let llm = CohereChat::new(config);
262 Self::from_inner(wrap_chat_model(llm))
263 }
264
265 pub fn from_llm<L>(llm: L) -> Self
271 where
272 L: BaseChatModel + Send + Sync + 'static,
273 L::Error: Into<ProviderError>,
274 {
275 Self::from_inner(wrap_chat_model(llm))
276 }
277
278 pub fn from_arc(llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>) -> Self {
280 Self::from_inner(llm)
281 }
282
283 pub fn into_inner(self) -> Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync> {
289 self.inner
290 }
291
292 pub fn inner(&self) -> &Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync> {
294 &self.inner
295 }
296}
297
298#[async_trait]
301impl Runnable<Vec<Message>, LLMResult> for LLMClient {
302 type Error = ProviderError;
303
304 async fn invoke(
305 &self,
306 input: Vec<Message>,
307 config: Option<RunnableConfig>,
308 ) -> Result<LLMResult, ProviderError> {
309 self.inner.invoke(input, self.apply_overrides(config)).await
310 }
311
312 async fn batch(
313 &self,
314 inputs: Vec<Vec<Message>>,
315 config: Option<RunnableConfig>,
316 ) -> Result<Vec<LLMResult>, ProviderError> {
317 self.inner.batch(inputs, self.apply_overrides(config)).await
318 }
319
320 async fn stream(
321 &self,
322 input: Vec<Message>,
323 config: Option<RunnableConfig>,
324 ) -> Result<
325 std::pin::Pin<Box<dyn Stream<Item = Result<LLMResult, ProviderError>> + Send>>,
326 ProviderError,
327 > {
328 self.inner.stream(input, self.apply_overrides(config)).await
329 }
330}
331
332impl BaseLanguageModel<Vec<Message>, LLMResult> for LLMClient {
333 fn model_name(&self) -> &str {
334 self.inner.model_name()
335 }
336
337 fn get_num_tokens(&self, text: &str) -> usize {
338 self.inner.get_num_tokens(text)
339 }
340
341 fn temperature(&self) -> Option<f32> {
342 self.overrides
344 .lock()
345 .unwrap_or_else(|e| e.into_inner())
346 .temperature
347 .or_else(|| self.inner.temperature())
348 }
349
350 fn max_tokens(&self) -> Option<usize> {
351 self.overrides
352 .lock()
353 .unwrap_or_else(|e| e.into_inner())
354 .max_tokens
355 .or_else(|| self.inner.max_tokens())
356 }
357
358 fn with_temperature(self, temp: f32) -> Self
359 where
360 Self: Sized,
361 {
362 self.overrides
365 .lock()
366 .unwrap_or_else(|e| e.into_inner())
367 .temperature = Some(temp);
368 self
369 }
370
371 fn with_max_tokens(self, max: usize) -> Self
372 where
373 Self: Sized,
374 {
375 self.overrides
376 .lock()
377 .unwrap_or_else(|e| e.into_inner())
378 .max_tokens = Some(max);
379 self
380 }
381}
382
383#[async_trait]
384impl BaseChatModel for LLMClient {
385 async fn chat(
386 &self,
387 messages: Vec<Message>,
388 config: Option<RunnableConfig>,
389 ) -> Result<LLMResult, ProviderError> {
390 self.inner
391 .chat(messages, self.apply_overrides(config))
392 .await
393 }
394
395 async fn stream_chat(
396 &self,
397 messages: Vec<Message>,
398 config: Option<RunnableConfig>,
399 ) -> Result<
400 std::pin::Pin<Box<dyn Stream<Item = Result<StreamChunk, ProviderError>> + Send>>,
401 ProviderError,
402 > {
403 self.inner
404 .stream_chat(messages, self.apply_overrides(config))
405 .await
406 }
407
408 fn bind_tools(
409 &self,
410 tools: Vec<ToolDefinition>,
411 ) -> Option<Box<dyn BaseChatModel<Error = ProviderError> + Send + Sync>> {
412 self.inner.bind_tools(tools)
413 }
414}
415
416impl std::ops::Deref for LLMClient {
417 type Target = dyn BaseChatModel<Error = ProviderError> + Send + Sync;
418
419 fn deref(&self) -> &Self::Target {
420 &*self.inner
421 }
422}
423
424#[cfg(test)]
425mod tests {
426 use super::*;
427 use crate::openai::{OpenAIChat, OpenAIConfig};
428 use crate::ENV_TEST_LOCK;
429
430 const DETECTION_ENV_VARS: [&str; 11] = [
432 "OPENAI_API_KEY",
433 "ANTHROPIC_API_KEY",
434 "AZURE_OPENAI_API_KEY",
435 "DEEPSEEK_API_KEY",
436 "QWEN_API_KEY",
437 "MOONSHOT_API_KEY",
438 "ZHIPU_API_KEY",
439 "MISTRAL_API_KEY",
440 "COHERE_API_KEY",
441 "GEMINI_API_KEY",
442 "OLLAMA_BASE_URL",
443 ];
444
445 fn save_and_set(key: &str, value: &str) -> Option<String> {
446 let old = std::env::var(key).ok();
447 std::env::set_var(key, value);
448 old
449 }
450
451 fn restore(key: &str, old: Option<String>) {
452 match old {
453 Some(v) => std::env::set_var(key, v),
454 None => std::env::remove_var(key),
455 }
456 }
457
458 #[test]
459 fn test_from_llm_openai() {
460 let config = OpenAIConfig::new("test_key");
461 let _client = LLMClient::from_llm(OpenAIChat::new(config));
462 }
463
464 #[test]
465 fn test_openai_constructor() {
466 let config = OpenAIConfig::new("test_key");
467 let _client = LLMClient::openai(config);
468 }
469
470 #[test]
471 fn test_from_arc() {
472 let config = OpenAIConfig::new("test_key");
473 let arc = wrap_chat_model(OpenAIChat::new(config));
474 let _client = LLMClient::from_arc(arc);
475 }
476
477 #[test]
478 fn test_into_inner() {
479 let config = OpenAIConfig::new("test_key");
480 let client = LLMClient::openai(config);
481 let _arc = client.into_inner();
482 }
483
484 #[test]
485 fn test_from_env_no_keys() {
486 let _lock = ENV_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
487 let saved: Vec<(&str, Option<String>)> = DETECTION_ENV_VARS
489 .iter()
490 .map(|k| {
491 let old = std::env::var(k).ok();
492 std::env::remove_var(k);
493 (*k, old)
494 })
495 .collect();
496
497 let result = LLMClient::from_env();
498 assert!(result.is_err());
499 assert!(result
500 .unwrap_err()
501 .to_string()
502 .contains("No LLM provider detected"));
503
504 for (k, old) in saved {
505 restore(k, old);
506 }
507 }
508
509 #[test]
510 fn test_from_env_detects_each_provider() {
511 for key in DETECTION_ENV_VARS {
512 let _lock = ENV_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
513 let saved: Vec<(&str, Option<String>)> = DETECTION_ENV_VARS
515 .iter()
516 .map(|k| {
517 let old = std::env::var(k).ok();
518 std::env::remove_var(k);
519 (*k, old)
520 })
521 .collect();
522
523 let old = save_and_set(key, "test-value");
524 let azure_extra: Vec<(&str, Option<String>)> = if key == "AZURE_OPENAI_API_KEY" {
526 vec![
527 (
528 "AZURE_OPENAI_ENDPOINT",
529 save_and_set("AZURE_OPENAI_ENDPOINT", "https://test.openai.azure.com"),
530 ),
531 (
532 "AZURE_OPENAI_DEPLOYMENT_NAME",
533 save_and_set("AZURE_OPENAI_DEPLOYMENT_NAME", "test-deployment"),
534 ),
535 (
536 "AZURE_OPENAI_API_VERSION",
537 save_and_set("AZURE_OPENAI_API_VERSION", "2024-02-01"),
538 ),
539 ]
540 } else {
541 vec![]
542 };
543
544 let result = LLMClient::from_env();
545 assert!(result.is_ok(), "expected detection via {key}");
546
547 restore(key, old);
548 for (k, v) in azure_extra {
549 restore(k, v);
550 }
551 for (k, old) in saved {
552 restore(k, old);
553 }
554 }
555 }
556
557 #[test]
558 fn test_with_temperature_override_applies_to_config() {
559 let config = OpenAIConfig::new("test_key");
560 let client = LLMClient::openai(config)
561 .with_temperature(0.7)
562 .with_max_tokens(128);
563
564 assert_eq!(client.temperature(), Some(0.7));
566 assert_eq!(client.max_tokens(), Some(128));
567
568 let merged = client.apply_overrides(None).unwrap();
570 assert_eq!(merged.temperature, Some(0.7));
571 assert_eq!(merged.max_tokens, Some(128));
572
573 let cfg = RunnableConfig::default().with_temperature(0.2);
575 let merged = client.apply_overrides(Some(cfg)).unwrap();
576 assert_eq!(merged.temperature, Some(0.7));
577 assert_eq!(merged.max_tokens, Some(128));
578 }
579
580 #[test]
581 fn test_no_overrides_passes_config_through() {
582 let config = OpenAIConfig::new("test_key");
583 let client = LLMClient::openai(config);
584
585 assert_eq!(client.temperature(), None);
586 assert_eq!(client.max_tokens(), None);
587
588 let cfg = RunnableConfig::default().with_temperature(0.5);
590 let merged = client.apply_overrides(Some(cfg.clone())).unwrap();
591 assert_eq!(merged.temperature, Some(0.5));
592 assert!(client.apply_overrides(None).is_none());
593 }
594
595 #[test]
596 fn test_deref_works() {
597 let config = OpenAIConfig::new("test_key");
598 let client = LLMClient::openai(config);
599 let _name = client.model_name();
601 }
602}