use std::fmt;
use crate::error::LlmError;
use crate::openai::{CompletionTokensParam, OpenAiConfig, OpenAiProvider};
use crate::provider::{
ChatExtras, ChatResponse, ChatStream, GenerationOverrides, LlmProvider, Message, StatusTx,
ToolDefinition,
};
#[derive(Clone)]
pub struct CompatibleConfig {
pub provider_name: String,
pub api_key: String,
pub base_url: String,
pub model: String,
pub max_tokens: u32,
pub embedding_model: Option<String>,
pub completion_tokens_param: Option<CompletionTokensParam>,
pub vision: Option<bool>,
}
impl fmt::Debug for CompatibleConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CompatibleConfig")
.field("provider_name", &self.provider_name)
.field("api_key", &"<redacted>")
.field("base_url", &self.base_url)
.field("model", &self.model)
.field("max_tokens", &self.max_tokens)
.field("embedding_model", &self.embedding_model)
.field("completion_tokens_param", &self.completion_tokens_param)
.field("vision", &self.vision)
.finish()
}
}
pub struct CompatibleProvider {
inner: OpenAiProvider,
provider_name: String,
}
impl CompatibleProvider {
#[must_use]
pub fn new(cfg: CompatibleConfig) -> Self {
let provider_name = cfg.provider_name;
let inner = OpenAiProvider::new(OpenAiConfig {
api_key: cfg.api_key,
base_url: cfg.base_url,
model: cfg.model,
max_tokens: cfg.max_tokens,
embedding_model: cfg.embedding_model,
reasoning_effort: None,
context_window: None,
completion_tokens_param: cfg.completion_tokens_param,
vision: cfg.vision,
});
Self {
inner,
provider_name,
}
}
}
impl fmt::Debug for CompatibleProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CompatibleProvider")
.field("provider_name", &self.provider_name)
.field("inner", &self.inner)
.finish_non_exhaustive()
}
}
impl Clone for CompatibleProvider {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
provider_name: self.provider_name.clone(),
}
}
}
impl CompatibleProvider {
pub async fn list_models_remote(
&self,
) -> Result<Vec<crate::model_cache::RemoteModelInfo>, LlmError> {
self.inner.list_models_remote().await
}
}
impl CompatibleProvider {
pub fn set_status_tx(&mut self, tx: StatusTx) {
self.inner.status_tx = Some(tx);
}
#[must_use]
pub fn with_generation_overrides(mut self, overrides: GenerationOverrides) -> Self {
self.inner = self.inner.with_generation_overrides(overrides);
self
}
#[must_use]
pub fn with_completion_tokens_param(mut self, param: CompletionTokensParam) -> Self {
self.inner = self.inner.with_completion_tokens_param(param);
self
}
#[must_use]
pub fn with_vision(mut self, supported: bool) -> Self {
self.inner = self.inner.with_vision(supported);
self
}
#[must_use]
pub fn with_output_schema_forwarding(
mut self,
enabled: bool,
hint_bytes: usize,
max_description_bytes: usize,
) -> Self {
self.inner =
self.inner
.with_output_schema_forwarding(enabled, hint_bytes, max_description_bytes);
self
}
pub fn set_reasoning_effort(&mut self, effort: Option<String>) {
self.inner.set_reasoning_effort(effort);
}
#[must_use]
pub fn current_reasoning_effort(&self) -> Option<String> {
self.inner.reasoning_effort.clone()
}
}
impl LlmProvider for CompatibleProvider {
fn context_window(&self) -> Option<usize> {
self.inner.context_window()
}
#[tracing::instrument(
name = "llm.chat",
skip_all,
fields(provider = self.name(), model = self.model_identifier())
)]
async fn chat(&self, messages: &[Message]) -> Result<String, LlmError> {
self.inner.chat(messages).await
}
async fn chat_with_extras(
&self,
messages: &[Message],
) -> Result<(String, ChatExtras), LlmError> {
self.inner.chat_with_extras(messages).await
}
#[tracing::instrument(
name = "llm.chat_stream",
skip_all,
fields(provider = self.name(), model = self.model_identifier())
)]
async fn chat_stream(&self, messages: &[Message]) -> Result<ChatStream, LlmError> {
self.inner.chat_stream(messages).await
}
fn supports_streaming(&self) -> bool {
self.inner.supports_streaming()
}
#[tracing::instrument(
name = "llm.embed",
skip_all,
fields(provider = self.name(), model = self.model_identifier())
)]
async fn embed(&self, text: &str) -> Result<Vec<f32>, LlmError> {
self.inner.embed(text).await
}
#[tracing::instrument(
name = "llm.embed_batch",
skip_all,
fields(provider = self.name(), model = self.model_identifier())
)]
async fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, LlmError> {
self.inner.embed_batch(texts).await
}
fn supports_embeddings(&self) -> bool {
self.inner.supports_embeddings()
}
fn name(&self) -> &str {
&self.provider_name
}
fn model_identifier(&self) -> &str {
self.inner.model_identifier()
}
fn list_models(&self) -> Vec<String> {
self.inner.list_models()
}
fn supports_structured_output(&self) -> bool {
self.inner.supports_structured_output()
}
async fn chat_typed<T>(&self, messages: &[Message]) -> Result<T, LlmError>
where
T: serde::de::DeserializeOwned + schemars::JsonSchema + 'static,
Self: Sized,
{
self.inner.chat_typed(messages).await
}
#[tracing::instrument(
name = "llm.chat_with_tools",
skip_all,
fields(provider = self.name(), model = self.model_identifier(), tool_count = tools.len())
)]
async fn chat_with_tools(
&self,
messages: &[Message],
tools: &[ToolDefinition],
) -> Result<ChatResponse, LlmError> {
self.inner.chat_with_tools(messages, tools).await
}
fn last_cache_usage(&self) -> Option<(u64, u64)> {
self.inner.last_cache_usage()
}
fn last_usage(&self) -> Option<(u64, u64)> {
self.inner.last_usage()
}
fn last_reasoning_tokens(&self) -> Option<u64> {
self.inner.last_reasoning_tokens()
}
fn supports_vision(&self) -> bool {
self.inner.supports_vision()
}
fn supports_tool_use(&self) -> bool {
self.inner.supports_tool_use()
}
fn debug_request_json(
&self,
messages: &[Message],
tools: &[ToolDefinition],
stream: bool,
) -> serde_json::Value {
self.inner.debug_request_json(messages, tools, stream)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_provider() -> CompatibleProvider {
CompatibleProvider::new(CompatibleConfig {
provider_name: "groq".into(),
api_key: "key".into(),
base_url: "https://api.groq.com/openai/v1".into(),
model: "llama-3.3-70b".into(),
max_tokens: 4096,
embedding_model: None,
completion_tokens_param: None,
vision: None,
})
}
#[test]
fn name_returns_custom_provider_name() {
let p = test_provider();
assert_eq!(p.name(), "groq");
}
#[test]
fn context_window_delegates_to_inner() {
let p = CompatibleProvider::new(CompatibleConfig {
provider_name: "openai".into(),
api_key: "key".into(),
base_url: "https://api.openai.com/v1".into(),
model: "gpt-4o".into(),
max_tokens: 4096,
embedding_model: None,
completion_tokens_param: None,
vision: None,
});
assert_eq!(p.context_window(), Some(128_000));
}
#[test]
fn context_window_unknown_model_returns_some_fallback() {
let p = CompatibleProvider::new(CompatibleConfig {
provider_name: "local".into(),
api_key: "key".into(),
base_url: "http://localhost/v1".into(),
model: "unknown-custom-model".into(),
max_tokens: 4096,
embedding_model: None,
completion_tokens_param: None,
vision: None,
});
assert!(p.context_window().is_some());
}
#[test]
fn supports_streaming_delegates() {
assert!(test_provider().supports_streaming());
}
#[test]
fn supports_embeddings_without_model() {
assert!(!test_provider().supports_embeddings());
}
#[test]
fn supports_embeddings_with_model() {
let p = CompatibleProvider::new(CompatibleConfig {
provider_name: "test".into(),
api_key: "key".into(),
base_url: "http://localhost".into(),
model: "m".into(),
max_tokens: 100,
embedding_model: Some("embed-model".into()),
completion_tokens_param: None,
vision: None,
});
assert!(p.supports_embeddings());
}
#[test]
fn clone_preserves_name() {
let p = test_provider();
let c = p.clone();
assert_eq!(c.name(), "groq");
}
#[test]
fn debug_contains_provider_name() {
let debug = format!("{:?}", test_provider());
assert!(debug.contains("groq"));
assert!(debug.contains("CompatibleProvider"));
}
#[tokio::test]
async fn chat_unreachable_errors() {
let p = CompatibleProvider::new(CompatibleConfig {
provider_name: "test".into(),
api_key: "key".into(),
base_url: "http://127.0.0.1:1".into(),
model: "m".into(),
max_tokens: 100,
embedding_model: None,
completion_tokens_param: None,
vision: None,
});
let msgs = vec![Message::from_legacy(crate::provider::Role::User, "hello")];
assert!(p.chat(&msgs).await.is_err());
}
#[tokio::test]
async fn embed_without_model_errors() {
let p = test_provider();
let result = p.embed("test").await;
assert!(result.is_err());
}
#[test]
fn last_usage_initially_none() {
assert!(test_provider().last_usage().is_none());
}
#[test]
fn with_output_schema_forwarding_does_not_panic() {
let p = test_provider().with_output_schema_forwarding(true, 512, usize::MAX);
assert_eq!(p.name(), "groq");
}
#[test]
fn set_reasoning_effort_applies_via_compatible() {
let mut p = test_provider();
p.set_reasoning_effort(Some("high".into()));
assert_eq!(p.inner.reasoning_effort.as_deref(), Some("high"));
}
#[test]
fn any_provider_set_reasoning_effort_delegates_to_compatible() {
use crate::any::AnyProvider;
let mut any = AnyProvider::Compatible(test_provider());
any.set_reasoning_effort(Some("high".into()));
let AnyProvider::Compatible(ref p) = any else {
panic!("variant must remain Compatible");
};
assert_eq!(
p.inner.reasoning_effort.as_deref(),
Some("high"),
"Compatible inner OpenAiProvider must have reasoning_effort applied"
);
}
#[test]
fn supports_vision_delegates_to_inner() {
assert!(!test_provider().supports_vision());
}
#[test]
fn supports_vision_with_vision_override_delegates_to_inner() {
let p = test_provider().with_vision(true);
assert!(p.supports_vision());
}
#[test]
fn supports_vision_config_field_true_forwards_to_inner() {
let p = CompatibleProvider::new(CompatibleConfig {
provider_name: "test".into(),
api_key: "key".into(),
base_url: "http://localhost".into(),
model: "llama-3.3-70b".into(),
max_tokens: 100,
embedding_model: None,
completion_tokens_param: None,
vision: Some(true),
});
assert!(p.supports_vision());
}
#[test]
fn supports_vision_config_field_false_forwards_to_inner() {
let p = CompatibleProvider::new(CompatibleConfig {
provider_name: "test".into(),
api_key: "key".into(),
base_url: "http://localhost".into(),
model: "llama-3.3-70b".into(),
max_tokens: 100,
embedding_model: None,
completion_tokens_param: None,
vision: Some(false),
});
assert!(!p.supports_vision());
}
#[test]
fn supports_tool_use_delegates_to_inner() {
assert!(test_provider().supports_tool_use());
}
#[test]
fn last_reasoning_tokens_initially_none() {
assert!(test_provider().last_reasoning_tokens().is_none());
}
#[test]
fn compatible_config_debug_redacts_api_key() {
let cfg = CompatibleConfig {
provider_name: "together-ai".into(),
api_key: "sk-SUPERSECRET".into(),
base_url: "https://api.together.xyz/v1".into(),
model: "meta-llama/Llama-3.3-70B-Instruct-Turbo".into(),
max_tokens: 4096,
embedding_model: None,
completion_tokens_param: None,
vision: None,
};
let dbg = format!("{cfg:?}");
assert!(!dbg.contains("sk-SUPERSECRET"));
assert!(dbg.contains("<redacted>"));
}
}