Skip to main content

lc_providers/
client.rs

1// lc-providers/src/client.rs
2//! LLMClient — zero-config unified entry point for switching providers
3//!
4//! Provides three creation modes:
5//! 1. `from_env()` — auto-detect environment variables
6//! 2. `openai(config)` / `anthropic(config)` etc. — explicit Config
7//! 3. `openai(OpenAIConfig::from_env_result()?)` — read config from env, then override
8//!
9//! # Example
10//!
11//! ```ignore
12//! // Mode 1: auto-detect
13//! let llm = LLMClient::from_env()?;
14//!
15//! // Mode 2: explicit Config
16//! let llm = LLMClient::openai(OpenAIConfig::new("sk-...").with_model("gpt-4o"));
17//!
18//! // Mode 3: from env + override
19//! let llm = LLMClient::openai(OpenAIConfig::from_env_result()?.with_model("gpt-4o"));
20//! ```
21
22use 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/// Per-client sampling overrides (providers Q2).
56///
57/// Stored on `LLMClient` so the consuming `with_temperature` /
58/// `with_max_tokens` builders can affect later `chat` / `stream_chat` calls
59/// even though the wrapped model sits behind a trait object. The overrides
60/// are merged into the `RunnableConfig` per call.
61#[derive(Debug, Default)]
62struct ClientOverrides {
63    temperature: Option<f32>,
64    max_tokens: Option<usize>,
65}
66
67/// LLM Client unified entry point
68///
69/// Wraps any `BaseChatModel` as `Arc<dyn BaseChatModel<Error = ProviderError>>`,
70/// providing zero-config auto-detection and explicit construction.
71///
72/// Implements `Deref<Target = dyn BaseChatModel>`, so you can call `.chat()` etc. directly.
73pub struct LLMClient {
74    inner: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
75    overrides: Mutex<ClientOverrides>,
76}
77
78impl LLMClient {
79    /// Wrap an already-normalized `Arc<dyn BaseChatModel>` with fresh overrides.
80    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    /// Merge per-client sampling overrides into the invocation config.
88    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    // -----------------------------------------------------------------------
110    // Auto-detect
111    // -----------------------------------------------------------------------
112
113    /// Create LLM Client from environment variables (auto-detect)
114    ///
115    /// Detection priority:
116    /// 1. `OPENAI_API_KEY` -> OpenAIChat
117    /// 2. `ANTHROPIC_API_KEY` -> AnthropicChat
118    /// 3. `AZURE_OPENAI_API_KEY` -> AzureOpenAIChat
119    /// 4. `DEEPSEEK_API_KEY` -> DeepSeekChat
120    /// 5. `QWEN_API_KEY` -> QwenChat
121    /// 6. `MOONSHOT_API_KEY` -> MoonshotChat
122    /// 7. `ZHIPU_API_KEY` -> ZhipuChat
123    /// 8. `MISTRAL_API_KEY` -> MistralChat
124    /// 9. `COHERE_API_KEY` -> CohereChat
125    /// 10. `GEMINI_API_KEY` (or `GOOGLE_API_KEY`) -> GeminiChat
126    /// 11. `OLLAMA_BASE_URL` -> OllamaChat
127    ///
128    /// # Errors
129    ///
130    /// Returns error if no known environment variable is set.
131    pub fn from_env() -> Result<Self, ProviderError> {
132        // Priority 1: OpenAI
133        if std::env::var("OPENAI_API_KEY").is_ok() {
134            return Ok(Self::openai(OpenAIConfig::from_env_result()?));
135        }
136
137        // Priority 2: Anthropic
138        if std::env::var("ANTHROPIC_API_KEY").is_ok() {
139            return Ok(Self::anthropic(AnthropicConfig::from_env_result()?));
140        }
141
142        // Priority 3: Azure OpenAI (endpoint + deployment based)
143        if std::env::var("AZURE_OPENAI_API_KEY").is_ok() {
144            return Ok(Self::azure(AzureOpenAIConfig::from_env_result()?));
145        }
146
147        // Priority 4: DeepSeek
148        if std::env::var("DEEPSEEK_API_KEY").is_ok() {
149            return Ok(Self::deepseek(DeepSeekConfig::from_env_result()?));
150        }
151
152        // Priority 5: Qwen
153        if std::env::var("QWEN_API_KEY").is_ok() {
154            return Ok(Self::qwen(QwenConfig::from_env_result()?));
155        }
156
157        // Priority 6: Moonshot
158        if std::env::var("MOONSHOT_API_KEY").is_ok() {
159            return Ok(Self::moonshot(MoonshotConfig::from_env_result()?));
160        }
161
162        // Priority 7: Zhipu
163        if std::env::var("ZHIPU_API_KEY").is_ok() {
164            return Ok(Self::zhipu(ZhipuConfig::from_env_result()?));
165        }
166
167        // Priority 8: Mistral
168        if std::env::var("MISTRAL_API_KEY").is_ok() {
169            return Ok(Self::mistral(MistralConfig::from_env_result()?));
170        }
171
172        // Priority 9: Cohere
173        if std::env::var("COHERE_API_KEY").is_ok() {
174            return Ok(Self::cohere(CohereConfig::from_env_result()?));
175        }
176
177        // Priority 10: Gemini
178        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        // Priority 11: Ollama (local, no API key required)
183        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    // -----------------------------------------------------------------------
196    // Explicit construction
197    // -----------------------------------------------------------------------
198
199    /// Create OpenAI Client
200    pub fn openai(config: OpenAIConfig) -> Self {
201        let llm = OpenAIChat::new(config);
202        Self::from_inner(wrap_chat_model(llm))
203    }
204
205    /// Create Anthropic Client
206    pub fn anthropic(config: AnthropicConfig) -> Self {
207        let llm = AnthropicChat::new(config);
208        Self::from_inner(wrap_chat_model(llm))
209    }
210
211    /// Create Ollama Client
212    pub fn ollama(config: OllamaConfig) -> Self {
213        let llm = OllamaChat::with_config(config);
214        Self::from_inner(wrap_chat_model(llm))
215    }
216
217    /// Create Gemini Client
218    pub fn gemini(config: GeminiConfig) -> Self {
219        let llm = GeminiChat::new(config);
220        Self::from_inner(wrap_chat_model(llm))
221    }
222
223    /// Create DeepSeek Client
224    pub fn deepseek(config: DeepSeekConfig) -> Self {
225        let llm = DeepSeekChat::new(config);
226        Self::from_inner(wrap_chat_model(llm))
227    }
228
229    /// Create Qwen Client
230    pub fn qwen(config: QwenConfig) -> Self {
231        let llm = QwenChat::new(config);
232        Self::from_inner(wrap_chat_model(llm))
233    }
234
235    /// Create Moonshot Client
236    pub fn moonshot(config: MoonshotConfig) -> Self {
237        let llm = MoonshotChat::new(config);
238        Self::from_inner(wrap_chat_model(llm))
239    }
240
241    /// Create Zhipu Client
242    pub fn zhipu(config: ZhipuConfig) -> Self {
243        let llm = ZhipuChat::new(config);
244        Self::from_inner(wrap_chat_model(llm))
245    }
246
247    /// Create Mistral Client
248    pub fn mistral(config: MistralConfig) -> Self {
249        let llm = MistralChat::new(config);
250        Self::from_inner(wrap_chat_model(llm))
251    }
252
253    /// Create Azure OpenAI Client
254    pub fn azure(config: AzureOpenAIConfig) -> Self {
255        let llm = AzureOpenAIChat::new(config);
256        Self::from_inner(wrap_chat_model(llm))
257    }
258
259    /// Create Cohere Client
260    pub fn cohere(config: CohereConfig) -> Self {
261        let llm = CohereChat::new(config);
262        Self::from_inner(wrap_chat_model(llm))
263    }
264
265    // -----------------------------------------------------------------------
266    // Generic construction
267    // -----------------------------------------------------------------------
268
269    /// Create Client from any `BaseChatModel`
270    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    /// Create Client from `Arc<dyn BaseChatModel>`
279    pub fn from_arc(llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>) -> Self {
280        Self::from_inner(llm)
281    }
282
283    // -----------------------------------------------------------------------
284    // Access inner
285    // -----------------------------------------------------------------------
286
287    /// Get the inner `Arc<dyn BaseChatModel>`, can be passed directly to Agent
288    pub fn into_inner(self) -> Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync> {
289        self.inner
290    }
291
292    /// Get inner reference
293    pub fn inner(&self) -> &Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync> {
294        &self.inner
295    }
296}
297
298// LLMClient implements the full trait hierarchy: Runnable -> BaseLanguageModel -> BaseChatModel
299
300#[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        // Per-client override takes precedence over the wrapped model's own value.
343        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        // Store the override; `chat`/`stream_chat` merge it into the
363        // `RunnableConfig` for the wrapped model (providers Q2).
364        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    /// Env vars that `LLMClient::from_env` checks, in detection order.
431    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        // Clear all detection vars, ensure from_env errors
488        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            // Clear every detection var, then set exactly one.
514            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            // Azure also needs endpoint + deployment to build its config.
525            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        // Getter reflects the per-client override (providers Q2).
565        assert_eq!(client.temperature(), Some(0.7));
566        assert_eq!(client.max_tokens(), Some(128));
567
568        // The merged config carries the overrides to the wrapped model.
569        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        // Per-client override takes precedence over a per-call config value.
574        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        // No overrides: config is returned as-is (not cloned).
589        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        // Can directly call BaseChatModel methods
600        let _name = client.model_name();
601    }
602}