use crate::error::ProviderError;
use crate::ollama::OllamaChat;
use crate::ollama::OllamaConfig;
use crate::openai::OpenAIChat;
use crate::openai::OpenAIConfig;
use crate::providers::anthropic::AnthropicChat;
use crate::providers::anthropic::AnthropicConfig;
use crate::providers::azure::AzureOpenAIChat;
use crate::providers::azure::AzureOpenAIConfig;
use crate::providers::cohere::CohereChat;
use crate::providers::cohere::CohereConfig;
use crate::providers::deepseek::DeepSeekChat;
use crate::providers::deepseek::DeepSeekConfig;
use crate::providers::gemini::GeminiChat;
use crate::providers::gemini::GeminiConfig;
use crate::providers::mistral::MistralChat;
use crate::providers::mistral::MistralConfig;
use crate::providers::moonshot::MoonshotChat;
use crate::providers::moonshot::MoonshotConfig;
use crate::providers::qwen::QwenChat;
use crate::providers::qwen::QwenConfig;
use crate::providers::zhipu::ZhipuChat;
use crate::providers::zhipu::ZhipuConfig;
use crate::wrap_chat_model;
use async_trait::async_trait;
use futures_util::Stream;
use lc_core::language_models::{BaseChatModel, BaseLanguageModel, LLMResult};
use lc_core::runnables::Runnable;
use lc_core::tools::ToolDefinition;
use lc_core::RunnableConfig;
use lc_schema::Message;
use std::sync::{Arc, Mutex};
#[derive(Debug, Default)]
struct ClientOverrides {
temperature: Option<f32>,
max_tokens: Option<usize>,
}
pub struct LLMClient {
inner: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
overrides: Mutex<ClientOverrides>,
}
impl LLMClient {
fn from_inner(inner: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>) -> Self {
Self {
inner,
overrides: Mutex::new(ClientOverrides::default()),
}
}
fn apply_overrides(&self, config: Option<RunnableConfig>) -> Option<RunnableConfig> {
let overrides = self.overrides.lock().unwrap_or_else(|e| e.into_inner());
if overrides.temperature.is_none() && overrides.max_tokens.is_none() {
return config;
}
let mut cfg = config.unwrap_or_default();
cfg.temperature = overrides.temperature.or(cfg.temperature);
cfg.max_tokens = overrides.max_tokens.or(cfg.max_tokens);
Some(cfg)
}
}
impl std::fmt::Debug for LLMClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LLMClient")
.field("model_name", &self.inner.model_name())
.finish()
}
}
impl LLMClient {
pub fn from_env() -> Result<Self, ProviderError> {
if std::env::var("OPENAI_API_KEY").is_ok() {
return Ok(Self::openai(OpenAIConfig::from_env_result()?));
}
if std::env::var("ANTHROPIC_API_KEY").is_ok() {
return Ok(Self::anthropic(AnthropicConfig::from_env_result()?));
}
if std::env::var("AZURE_OPENAI_API_KEY").is_ok() {
return Ok(Self::azure(AzureOpenAIConfig::from_env_result()?));
}
if std::env::var("DEEPSEEK_API_KEY").is_ok() {
return Ok(Self::deepseek(DeepSeekConfig::from_env_result()?));
}
if std::env::var("QWEN_API_KEY").is_ok() {
return Ok(Self::qwen(QwenConfig::from_env_result()?));
}
if std::env::var("MOONSHOT_API_KEY").is_ok() {
return Ok(Self::moonshot(MoonshotConfig::from_env_result()?));
}
if std::env::var("ZHIPU_API_KEY").is_ok() {
return Ok(Self::zhipu(ZhipuConfig::from_env_result()?));
}
if std::env::var("MISTRAL_API_KEY").is_ok() {
return Ok(Self::mistral(MistralConfig::from_env_result()?));
}
if std::env::var("COHERE_API_KEY").is_ok() {
return Ok(Self::cohere(CohereConfig::from_env_result()?));
}
if std::env::var("GEMINI_API_KEY").is_ok() || std::env::var("GOOGLE_API_KEY").is_ok() {
return Ok(Self::gemini(GeminiConfig::from_env_result()?));
}
if std::env::var("OLLAMA_BASE_URL").is_ok() {
return Ok(Self::ollama(OllamaConfig::from_env_result()?));
}
Err(ProviderError::Config(
"No LLM provider detected. Set one of: OPENAI_API_KEY, ANTHROPIC_API_KEY, \
AZURE_OPENAI_API_KEY, DEEPSEEK_API_KEY, QWEN_API_KEY, MOONSHOT_API_KEY, \
ZHIPU_API_KEY, MISTRAL_API_KEY, COHERE_API_KEY, GEMINI_API_KEY, OLLAMA_BASE_URL"
.to_string(),
))
}
pub fn openai(config: OpenAIConfig) -> Self {
let llm = OpenAIChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn anthropic(config: AnthropicConfig) -> Self {
let llm = AnthropicChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn ollama(config: OllamaConfig) -> Self {
let llm = OllamaChat::with_config(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn gemini(config: GeminiConfig) -> Self {
let llm = GeminiChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn deepseek(config: DeepSeekConfig) -> Self {
let llm = DeepSeekChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn qwen(config: QwenConfig) -> Self {
let llm = QwenChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn moonshot(config: MoonshotConfig) -> Self {
let llm = MoonshotChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn zhipu(config: ZhipuConfig) -> Self {
let llm = ZhipuChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn mistral(config: MistralConfig) -> Self {
let llm = MistralChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn azure(config: AzureOpenAIConfig) -> Self {
let llm = AzureOpenAIChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn cohere(config: CohereConfig) -> Self {
let llm = CohereChat::new(config);
Self::from_inner(wrap_chat_model(llm))
}
pub fn from_llm<L>(llm: L) -> Self
where
L: BaseChatModel + Send + Sync + 'static,
L::Error: Into<ProviderError>,
{
Self::from_inner(wrap_chat_model(llm))
}
pub fn from_arc(llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>) -> Self {
Self::from_inner(llm)
}
pub fn into_inner(self) -> Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync> {
self.inner
}
pub fn inner(&self) -> &Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync> {
&self.inner
}
}
#[async_trait]
impl Runnable<Vec<Message>, LLMResult> for LLMClient {
type Error = ProviderError;
async fn invoke(
&self,
input: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, ProviderError> {
self.inner.invoke(input, self.apply_overrides(config)).await
}
async fn batch(
&self,
inputs: Vec<Vec<Message>>,
config: Option<RunnableConfig>,
) -> Result<Vec<LLMResult>, ProviderError> {
self.inner.batch(inputs, self.apply_overrides(config)).await
}
async fn stream(
&self,
input: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<
std::pin::Pin<Box<dyn Stream<Item = Result<LLMResult, ProviderError>> + Send>>,
ProviderError,
> {
self.inner.stream(input, self.apply_overrides(config)).await
}
}
impl BaseLanguageModel<Vec<Message>, LLMResult> for LLMClient {
fn model_name(&self) -> &str {
self.inner.model_name()
}
fn get_num_tokens(&self, text: &str) -> usize {
self.inner.get_num_tokens(text)
}
fn temperature(&self) -> Option<f32> {
self.overrides
.lock()
.unwrap_or_else(|e| e.into_inner())
.temperature
.or_else(|| self.inner.temperature())
}
fn max_tokens(&self) -> Option<usize> {
self.overrides
.lock()
.unwrap_or_else(|e| e.into_inner())
.max_tokens
.or_else(|| self.inner.max_tokens())
}
fn with_temperature(self, temp: f32) -> Self
where
Self: Sized,
{
self.overrides
.lock()
.unwrap_or_else(|e| e.into_inner())
.temperature = Some(temp);
self
}
fn with_max_tokens(self, max: usize) -> Self
where
Self: Sized,
{
self.overrides
.lock()
.unwrap_or_else(|e| e.into_inner())
.max_tokens = Some(max);
self
}
}
#[async_trait]
impl BaseChatModel for LLMClient {
async fn chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, ProviderError> {
self.inner
.chat(messages, self.apply_overrides(config))
.await
}
async fn stream_chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<
std::pin::Pin<Box<dyn Stream<Item = Result<String, ProviderError>> + Send>>,
ProviderError,
> {
self.inner
.stream_chat(messages, self.apply_overrides(config))
.await
}
fn bind_tools(
&self,
tools: Vec<ToolDefinition>,
) -> Option<Box<dyn BaseChatModel<Error = ProviderError> + Send + Sync>> {
self.inner.bind_tools(tools)
}
}
impl std::ops::Deref for LLMClient {
type Target = dyn BaseChatModel<Error = ProviderError> + Send + Sync;
fn deref(&self) -> &Self::Target {
&*self.inner
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::openai::{OpenAIChat, OpenAIConfig};
use crate::ENV_TEST_LOCK;
const DETECTION_ENV_VARS: [&str; 11] = [
"OPENAI_API_KEY",
"ANTHROPIC_API_KEY",
"AZURE_OPENAI_API_KEY",
"DEEPSEEK_API_KEY",
"QWEN_API_KEY",
"MOONSHOT_API_KEY",
"ZHIPU_API_KEY",
"MISTRAL_API_KEY",
"COHERE_API_KEY",
"GEMINI_API_KEY",
"OLLAMA_BASE_URL",
];
fn save_and_set(key: &str, value: &str) -> Option<String> {
let old = std::env::var(key).ok();
std::env::set_var(key, value);
old
}
fn restore(key: &str, old: Option<String>) {
match old {
Some(v) => std::env::set_var(key, v),
None => std::env::remove_var(key),
}
}
#[test]
fn test_from_llm_openai() {
let config = OpenAIConfig::new("test_key");
let _client = LLMClient::from_llm(OpenAIChat::new(config));
}
#[test]
fn test_openai_constructor() {
let config = OpenAIConfig::new("test_key");
let _client = LLMClient::openai(config);
}
#[test]
fn test_from_arc() {
let config = OpenAIConfig::new("test_key");
let arc = wrap_chat_model(OpenAIChat::new(config));
let _client = LLMClient::from_arc(arc);
}
#[test]
fn test_into_inner() {
let config = OpenAIConfig::new("test_key");
let client = LLMClient::openai(config);
let _arc = client.into_inner();
}
#[test]
fn test_from_env_no_keys() {
let _lock = ENV_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let saved: Vec<(&str, Option<String>)> = DETECTION_ENV_VARS
.iter()
.map(|k| {
let old = std::env::var(k).ok();
std::env::remove_var(k);
(*k, old)
})
.collect();
let result = LLMClient::from_env();
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("No LLM provider detected"));
for (k, old) in saved {
restore(k, old);
}
}
#[test]
fn test_from_env_detects_each_provider() {
for key in DETECTION_ENV_VARS {
let _lock = ENV_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let saved: Vec<(&str, Option<String>)> = DETECTION_ENV_VARS
.iter()
.map(|k| {
let old = std::env::var(k).ok();
std::env::remove_var(k);
(*k, old)
})
.collect();
let old = save_and_set(key, "test-value");
let azure_extra: Vec<(&str, Option<String>)> = if key == "AZURE_OPENAI_API_KEY" {
vec![
(
"AZURE_OPENAI_ENDPOINT",
save_and_set("AZURE_OPENAI_ENDPOINT", "https://test.openai.azure.com"),
),
(
"AZURE_OPENAI_DEPLOYMENT_NAME",
save_and_set("AZURE_OPENAI_DEPLOYMENT_NAME", "test-deployment"),
),
(
"AZURE_OPENAI_API_VERSION",
save_and_set("AZURE_OPENAI_API_VERSION", "2024-02-01"),
),
]
} else {
vec![]
};
let result = LLMClient::from_env();
assert!(result.is_ok(), "expected detection via {key}");
restore(key, old);
for (k, v) in azure_extra {
restore(k, v);
}
for (k, old) in saved {
restore(k, old);
}
}
}
#[test]
fn test_with_temperature_override_applies_to_config() {
let config = OpenAIConfig::new("test_key");
let client = LLMClient::openai(config)
.with_temperature(0.7)
.with_max_tokens(128);
assert_eq!(client.temperature(), Some(0.7));
assert_eq!(client.max_tokens(), Some(128));
let merged = client.apply_overrides(None).unwrap();
assert_eq!(merged.temperature, Some(0.7));
assert_eq!(merged.max_tokens, Some(128));
let cfg = RunnableConfig::default().with_temperature(0.2);
let merged = client.apply_overrides(Some(cfg)).unwrap();
assert_eq!(merged.temperature, Some(0.7));
assert_eq!(merged.max_tokens, Some(128));
}
#[test]
fn test_no_overrides_passes_config_through() {
let config = OpenAIConfig::new("test_key");
let client = LLMClient::openai(config);
assert_eq!(client.temperature(), None);
assert_eq!(client.max_tokens(), None);
let cfg = RunnableConfig::default().with_temperature(0.5);
let merged = client.apply_overrides(Some(cfg.clone())).unwrap();
assert_eq!(merged.temperature, Some(0.5));
assert!(client.apply_overrides(None).is_none());
}
#[test]
fn test_deref_works() {
let config = OpenAIConfig::new("test_key");
let client = LLMClient::openai(config);
let _name = client.model_name();
}
}