Skip to main content

zeph_llm/
compatible.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! `OpenAI`-compatible provider adapter.
5//!
6//! [`CompatibleProvider`] wraps [`crate::openai::OpenAiProvider`] and adds a named
7//! provider label for logging. Use it for any endpoint that exposes the `OpenAI` Chat
8//! Completions and Embeddings API (Together AI, Fireworks, Anyscale, local vLLM, etc.).
9//!
10//! # Configuration
11//!
12//! ```toml
13//! [[llm.providers]]
14//! name = "together"
15//! type = "compatible"
16//! provider_name = "together-ai"
17//! base_url = "https://api.together.xyz/v1"
18//! model = "meta-llama/Llama-3.3-70B-Instruct-Turbo"
19//! max_tokens = 4096
20//! api_key_vault = "ZEPH_TOGETHER_API_KEY"
21//! ```
22
23use std::fmt;
24
25use crate::error::LlmError;
26use crate::openai::{CompletionTokensParam, OpenAiConfig, OpenAiProvider};
27use crate::provider::{
28    ChatExtras, ChatResponse, ChatStream, GenerationOverrides, LlmProvider, Message, StatusTx,
29    ToolDefinition,
30};
31
32/// Configuration for [`CompatibleProvider`].
33///
34/// Pass to [`CompatibleProvider::new`] instead of individual positional arguments to avoid
35/// silent parameter transposition.
36///
37/// # Examples
38///
39/// ```
40/// use zeph_llm::compatible::{CompatibleConfig, CompatibleProvider};
41///
42/// let cfg = CompatibleConfig {
43///     provider_name: "together-ai".into(),
44///     api_key: "key".into(),
45///     base_url: "https://api.together.xyz/v1".into(),
46///     model: "meta-llama/Llama-3.3-70B-Instruct-Turbo".into(),
47///     max_tokens: 4096,
48///     embedding_model: None,
49///     completion_tokens_param: None,
50///     vision: None,
51/// };
52/// let provider = CompatibleProvider::new(cfg);
53/// ```
54#[derive(Clone)]
55pub struct CompatibleConfig {
56    /// Human-readable provider name used in logs and [`LlmProvider::name`].
57    pub provider_name: String,
58    /// Secret API key sent in the `Authorization: Bearer` header.
59    pub api_key: String,
60    /// Base URL of the endpoint, e.g. `"https://api.together.xyz/v1"`.
61    pub base_url: String,
62    /// Chat model identifier.
63    pub model: String,
64    /// Upper bound on completion tokens returned by the model.
65    pub max_tokens: u32,
66    /// Embedding model identifier. Set to `None` when the endpoint does not support embeddings.
67    pub embedding_model: Option<String>,
68    /// Override which token-limit parameter is used in API requests.
69    ///
70    /// When `None`, the provider infers the correct field from the model name via the built-in
71    /// prefix table. Set explicitly for models the table does not recognise (e.g. fine-tuned
72    /// reasoning models whose names do not start with `o` + digit).
73    pub completion_tokens_param: Option<CompletionTokensParam>,
74    /// Explicit vision-capability override.
75    ///
76    /// `OpenAiProvider`'s built-in model-name prefix table cannot recognise arbitrary
77    /// `compatible` endpoint model names (Together AI, local vLLM, etc.), so it fails safe
78    /// to `false` for all of them. Set this when the endpoint's configured model actually
79    /// accepts image input.
80    pub vision: Option<bool>,
81}
82
83impl fmt::Debug for CompatibleConfig {
84    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
85        f.debug_struct("CompatibleConfig")
86            .field("provider_name", &self.provider_name)
87            .field("api_key", &"<redacted>")
88            .field("base_url", &self.base_url)
89            .field("model", &self.model)
90            .field("max_tokens", &self.max_tokens)
91            .field("embedding_model", &self.embedding_model)
92            .field("completion_tokens_param", &self.completion_tokens_param)
93            .field("vision", &self.vision)
94            .finish()
95    }
96}
97
98/// [`LlmProvider`] adapter for OpenAI-compatible REST endpoints.
99///
100/// Delegates all operations to an inner [`OpenAiProvider`] while exposing a
101/// configurable `provider_name` for logging and routing identification.
102pub struct CompatibleProvider {
103    inner: OpenAiProvider,
104    /// Human-readable name used in logs and [`LlmProvider::name`].
105    provider_name: String,
106}
107
108impl CompatibleProvider {
109    /// Create a new provider from a [`CompatibleConfig`].
110    #[must_use]
111    pub fn new(cfg: CompatibleConfig) -> Self {
112        let provider_name = cfg.provider_name;
113        let inner = OpenAiProvider::new(OpenAiConfig {
114            api_key: cfg.api_key,
115            base_url: cfg.base_url,
116            model: cfg.model,
117            max_tokens: cfg.max_tokens,
118            embedding_model: cfg.embedding_model,
119            reasoning_effort: None,
120            context_window: None,
121            completion_tokens_param: cfg.completion_tokens_param,
122            vision: cfg.vision,
123        });
124        Self {
125            inner,
126            provider_name,
127        }
128    }
129}
130
131impl fmt::Debug for CompatibleProvider {
132    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
133        f.debug_struct("CompatibleProvider")
134            .field("provider_name", &self.provider_name)
135            .field("inner", &self.inner)
136            .finish_non_exhaustive()
137    }
138}
139
140impl Clone for CompatibleProvider {
141    fn clone(&self) -> Self {
142        Self {
143            inner: self.inner.clone(),
144            provider_name: self.provider_name.clone(),
145        }
146    }
147}
148
149impl CompatibleProvider {
150    /// Fetch models via the inner `OpenAiProvider`. Cache slug is derived from base URL.
151    ///
152    /// # Errors
153    ///
154    /// Returns an error if the API request fails.
155    pub async fn list_models_remote(
156        &self,
157    ) -> Result<Vec<crate::model_cache::RemoteModelInfo>, LlmError> {
158        self.inner.list_models_remote().await
159    }
160}
161
162impl CompatibleProvider {
163    /// Attach a status channel for streaming progress events to the TUI.
164    pub fn set_status_tx(&mut self, tx: StatusTx) {
165        self.inner.status_tx = Some(tx);
166    }
167
168    /// Override generation parameters (temperature, top-p, etc.) for all subsequent calls.
169    #[must_use]
170    pub fn with_generation_overrides(mut self, overrides: GenerationOverrides) -> Self {
171        self.inner = self.inner.with_generation_overrides(overrides);
172        self
173    }
174
175    /// Override which token-limit parameter is sent in API requests.
176    ///
177    /// Delegates to the inner [`OpenAiProvider`]. Use this when the model name is not covered
178    /// by the built-in prefix table and the inferred field would produce a 400 error.
179    ///
180    /// # Examples
181    ///
182    /// ```
183    /// use zeph_llm::compatible::{CompatibleConfig, CompatibleProvider};
184    /// use zeph_llm::openai::CompletionTokensParam;
185    ///
186    /// let provider = CompatibleProvider::new(CompatibleConfig {
187    ///     provider_name: "my-provider".into(),
188    ///     api_key: "key".into(),
189    ///     base_url: "https://api.example.com/v1".into(),
190    ///     model: "my-ft-reasoner-v1".into(),
191    ///     max_tokens: 4096,
192    ///     embedding_model: None,
193    ///     completion_tokens_param: None,
194    ///     vision: None,
195    /// })
196    /// .with_completion_tokens_param(CompletionTokensParam::MaxCompletionTokens);
197    /// ```
198    #[must_use]
199    pub fn with_completion_tokens_param(mut self, param: CompletionTokensParam) -> Self {
200        self.inner = self.inner.with_completion_tokens_param(param);
201        self
202    }
203
204    /// Override the vision-capability value reported by [`LlmProvider::supports_vision`].
205    ///
206    /// Delegates to the inner [`OpenAiProvider`]. Use this for `compatible` endpoints whose
207    /// model name is not covered by `OpenAiProvider`'s built-in prefix table — which is the
208    /// common case, since third-party model names carry no `OpenAI` naming convention and the
209    /// table therefore fails safe to `false` for them.
210    ///
211    /// # Examples
212    ///
213    /// ```
214    /// use zeph_llm::compatible::{CompatibleConfig, CompatibleProvider};
215    /// use zeph_llm::provider::LlmProvider;
216    ///
217    /// let provider = CompatibleProvider::new(CompatibleConfig {
218    ///     provider_name: "local-vlm".into(),
219    ///     api_key: "key".into(),
220    ///     base_url: "http://localhost:8000/v1".into(),
221    ///     model: "llava-onevision".into(),
222    ///     max_tokens: 4096,
223    ///     embedding_model: None,
224    ///     completion_tokens_param: None,
225    ///     vision: None,
226    /// })
227    /// .with_vision(true);
228    /// assert!(provider.supports_vision());
229    /// ```
230    #[must_use]
231    pub fn with_vision(mut self, supported: bool) -> Self {
232        self.inner = self.inner.with_vision(supported);
233        self
234    }
235
236    /// Forward MCP tool output schemas as JSON hints appended to tool descriptions.
237    ///
238    /// Delegates to the inner [`OpenAiProvider`]. When `enabled` is `false` the call is a no-op.
239    /// `hint_bytes` caps the JSON representation; `max_description_bytes` caps the combined
240    /// description string.
241    #[must_use]
242    pub fn with_output_schema_forwarding(
243        mut self,
244        enabled: bool,
245        hint_bytes: usize,
246        max_description_bytes: usize,
247    ) -> Self {
248        self.inner =
249            self.inner
250                .with_output_schema_forwarding(enabled, hint_bytes, max_description_bytes);
251        self
252    }
253
254    /// Apply a `reasoning_effort` override to the inner [`OpenAiProvider`].
255    ///
256    /// Delegates to [`OpenAiProvider::set_reasoning_effort`], which validates the value
257    /// (`"low"`, `"medium"`, or `"high"`) and logs a warning for any unknown value.
258    /// Pass `None` to clear a previously-set effort level.
259    pub fn set_reasoning_effort(&mut self, effort: Option<String>) {
260        self.inner.set_reasoning_effort(effort);
261    }
262
263    /// Return the currently configured `reasoning_effort` value on the inner
264    /// [`OpenAiProvider`], if any.
265    #[must_use]
266    pub fn current_reasoning_effort(&self) -> Option<String> {
267        self.inner.reasoning_effort.clone()
268    }
269}
270
271impl LlmProvider for CompatibleProvider {
272    fn context_window(&self) -> Option<usize> {
273        self.inner.context_window()
274    }
275
276    #[tracing::instrument(
277        name = "llm.chat",
278        skip_all,
279        fields(provider = self.name(), model = self.model_identifier())
280    )]
281    async fn chat(&self, messages: &[Message]) -> Result<String, LlmError> {
282        self.inner.chat(messages).await
283    }
284
285    async fn chat_with_extras(
286        &self,
287        messages: &[Message],
288    ) -> Result<(String, ChatExtras), LlmError> {
289        self.inner.chat_with_extras(messages).await
290    }
291
292    #[tracing::instrument(
293        name = "llm.chat_stream",
294        skip_all,
295        fields(provider = self.name(), model = self.model_identifier())
296    )]
297    async fn chat_stream(&self, messages: &[Message]) -> Result<ChatStream, LlmError> {
298        self.inner.chat_stream(messages).await
299    }
300
301    fn supports_streaming(&self) -> bool {
302        self.inner.supports_streaming()
303    }
304
305    #[tracing::instrument(
306        name = "llm.embed",
307        skip_all,
308        fields(provider = self.name(), model = self.model_identifier())
309    )]
310    async fn embed(&self, text: &str) -> Result<Vec<f32>, LlmError> {
311        self.inner.embed(text).await
312    }
313
314    #[tracing::instrument(
315        name = "llm.embed_batch",
316        skip_all,
317        fields(provider = self.name(), model = self.model_identifier())
318    )]
319    async fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, LlmError> {
320        self.inner.embed_batch(texts).await
321    }
322
323    fn supports_embeddings(&self) -> bool {
324        self.inner.supports_embeddings()
325    }
326
327    fn name(&self) -> &str {
328        &self.provider_name
329    }
330
331    fn model_identifier(&self) -> &str {
332        self.inner.model_identifier()
333    }
334
335    fn list_models(&self) -> Vec<String> {
336        self.inner.list_models()
337    }
338
339    fn supports_structured_output(&self) -> bool {
340        self.inner.supports_structured_output()
341    }
342
343    async fn chat_typed<T>(&self, messages: &[Message]) -> Result<T, LlmError>
344    where
345        T: serde::de::DeserializeOwned + schemars::JsonSchema + 'static,
346        Self: Sized,
347    {
348        self.inner.chat_typed(messages).await
349    }
350
351    #[tracing::instrument(
352        name = "llm.chat_with_tools",
353        skip_all,
354        fields(provider = self.name(), model = self.model_identifier(), tool_count = tools.len())
355    )]
356    async fn chat_with_tools(
357        &self,
358        messages: &[Message],
359        tools: &[ToolDefinition],
360    ) -> Result<ChatResponse, LlmError> {
361        self.inner.chat_with_tools(messages, tools).await
362    }
363
364    fn last_cache_usage(&self) -> Option<(u64, u64)> {
365        self.inner.last_cache_usage()
366    }
367
368    fn last_usage(&self) -> Option<(u64, u64)> {
369        self.inner.last_usage()
370    }
371
372    fn last_reasoning_tokens(&self) -> Option<u64> {
373        self.inner.last_reasoning_tokens()
374    }
375
376    fn supports_vision(&self) -> bool {
377        self.inner.supports_vision()
378    }
379
380    fn supports_tool_use(&self) -> bool {
381        self.inner.supports_tool_use()
382    }
383
384    fn debug_request_json(
385        &self,
386        messages: &[Message],
387        tools: &[ToolDefinition],
388        stream: bool,
389    ) -> serde_json::Value {
390        self.inner.debug_request_json(messages, tools, stream)
391    }
392}
393
394#[cfg(test)]
395mod tests {
396    use super::*;
397
398    fn test_provider() -> CompatibleProvider {
399        CompatibleProvider::new(CompatibleConfig {
400            provider_name: "groq".into(),
401            api_key: "key".into(),
402            base_url: "https://api.groq.com/openai/v1".into(),
403            model: "llama-3.3-70b".into(),
404            max_tokens: 4096,
405            embedding_model: None,
406            completion_tokens_param: None,
407            vision: None,
408        })
409    }
410
411    #[test]
412    fn name_returns_custom_provider_name() {
413        let p = test_provider();
414        assert_eq!(p.name(), "groq");
415    }
416
417    #[test]
418    fn context_window_delegates_to_inner() {
419        // "gpt-4o" is in the prefix table → Some(128_000)
420        let p = CompatibleProvider::new(CompatibleConfig {
421            provider_name: "openai".into(),
422            api_key: "key".into(),
423            base_url: "https://api.openai.com/v1".into(),
424            model: "gpt-4o".into(),
425            max_tokens: 4096,
426            embedding_model: None,
427            completion_tokens_param: None,
428            vision: None,
429        });
430        assert_eq!(p.context_window(), Some(128_000));
431    }
432
433    #[test]
434    fn context_window_unknown_model_returns_some_fallback() {
435        // Unknown model falls back to 128_000 default in OpenAiProvider.
436        let p = CompatibleProvider::new(CompatibleConfig {
437            provider_name: "local".into(),
438            api_key: "key".into(),
439            base_url: "http://localhost/v1".into(),
440            model: "unknown-custom-model".into(),
441            max_tokens: 4096,
442            embedding_model: None,
443            completion_tokens_param: None,
444            vision: None,
445        });
446        // OpenAiProvider returns Some(128_000) as fallback for unrecognised models.
447        assert!(p.context_window().is_some());
448    }
449
450    #[test]
451    fn supports_streaming_delegates() {
452        assert!(test_provider().supports_streaming());
453    }
454
455    #[test]
456    fn supports_embeddings_without_model() {
457        assert!(!test_provider().supports_embeddings());
458    }
459
460    #[test]
461    fn supports_embeddings_with_model() {
462        let p = CompatibleProvider::new(CompatibleConfig {
463            provider_name: "test".into(),
464            api_key: "key".into(),
465            base_url: "http://localhost".into(),
466            model: "m".into(),
467            max_tokens: 100,
468            embedding_model: Some("embed-model".into()),
469            completion_tokens_param: None,
470            vision: None,
471        });
472        assert!(p.supports_embeddings());
473    }
474
475    #[test]
476    fn clone_preserves_name() {
477        let p = test_provider();
478        let c = p.clone();
479        assert_eq!(c.name(), "groq");
480    }
481
482    #[test]
483    fn debug_contains_provider_name() {
484        let debug = format!("{:?}", test_provider());
485        assert!(debug.contains("groq"));
486        assert!(debug.contains("CompatibleProvider"));
487    }
488
489    #[tokio::test]
490    async fn chat_unreachable_errors() {
491        let p = CompatibleProvider::new(CompatibleConfig {
492            provider_name: "test".into(),
493            api_key: "key".into(),
494            base_url: "http://127.0.0.1:1".into(),
495            model: "m".into(),
496            max_tokens: 100,
497            embedding_model: None,
498            completion_tokens_param: None,
499            vision: None,
500        });
501        let msgs = vec![Message::from_legacy(crate::provider::Role::User, "hello")];
502        assert!(p.chat(&msgs).await.is_err());
503    }
504
505    #[tokio::test]
506    async fn embed_without_model_errors() {
507        let p = test_provider();
508        let result = p.embed("test").await;
509        assert!(result.is_err());
510    }
511
512    #[test]
513    fn last_usage_initially_none() {
514        assert!(test_provider().last_usage().is_none());
515    }
516
517    #[test]
518    fn with_output_schema_forwarding_does_not_panic() {
519        // Smoke-test that the builder compiles and returns self without panicking.
520        let p = test_provider().with_output_schema_forwarding(true, 512, usize::MAX);
521        assert_eq!(p.name(), "groq");
522    }
523
524    // ── reasoning_effort restore path (#5007 Phase 2) ────────────────────────
525
526    #[test]
527    fn set_reasoning_effort_applies_via_compatible() {
528        let mut p = test_provider();
529        p.set_reasoning_effort(Some("high".into()));
530        assert_eq!(p.inner.reasoning_effort.as_deref(), Some("high"));
531    }
532
533    #[test]
534    fn any_provider_set_reasoning_effort_delegates_to_compatible() {
535        use crate::any::AnyProvider;
536        let mut any = AnyProvider::Compatible(test_provider());
537        any.set_reasoning_effort(Some("high".into()));
538        let AnyProvider::Compatible(ref p) = any else {
539            panic!("variant must remain Compatible");
540        };
541        assert_eq!(
542            p.inner.reasoning_effort.as_deref(),
543            Some("high"),
544            "Compatible inner OpenAiProvider must have reasoning_effort applied"
545        );
546    }
547
548    #[test]
549    fn supports_vision_delegates_to_inner() {
550        // #6411: test_provider() uses model "llama-3.3-70b", which matches no OpenAI
551        // vision prefix — OpenAiProvider's fail-safe default must be false, not true.
552        assert!(!test_provider().supports_vision());
553    }
554
555    #[test]
556    fn supports_vision_with_vision_override_delegates_to_inner() {
557        // with_vision(true) must override the prefix-table default even for a model
558        // name the table cannot recognise.
559        let p = test_provider().with_vision(true);
560        assert!(p.supports_vision());
561    }
562
563    #[test]
564    fn supports_vision_config_field_true_forwards_to_inner() {
565        // Distinct code path from with_vision(): CompatibleConfig::vision must be forwarded
566        // into the inner OpenAiConfig by CompatibleProvider::new, not just by the builder.
567        let p = CompatibleProvider::new(CompatibleConfig {
568            provider_name: "test".into(),
569            api_key: "key".into(),
570            base_url: "http://localhost".into(),
571            model: "llama-3.3-70b".into(),
572            max_tokens: 100,
573            embedding_model: None,
574            completion_tokens_param: None,
575            vision: Some(true),
576        });
577        assert!(p.supports_vision());
578    }
579
580    #[test]
581    fn supports_vision_config_field_false_forwards_to_inner() {
582        let p = CompatibleProvider::new(CompatibleConfig {
583            provider_name: "test".into(),
584            api_key: "key".into(),
585            base_url: "http://localhost".into(),
586            model: "llama-3.3-70b".into(),
587            max_tokens: 100,
588            embedding_model: None,
589            completion_tokens_param: None,
590            vision: Some(false),
591        });
592        assert!(!p.supports_vision());
593    }
594
595    #[test]
596    fn supports_tool_use_delegates_to_inner() {
597        // OpenAiProvider always returns true for supports_tool_use.
598        assert!(test_provider().supports_tool_use());
599    }
600
601    #[test]
602    fn last_reasoning_tokens_initially_none() {
603        assert!(test_provider().last_reasoning_tokens().is_none());
604    }
605
606    #[test]
607    fn compatible_config_debug_redacts_api_key() {
608        let cfg = CompatibleConfig {
609            provider_name: "together-ai".into(),
610            api_key: "sk-SUPERSECRET".into(),
611            base_url: "https://api.together.xyz/v1".into(),
612            model: "meta-llama/Llama-3.3-70B-Instruct-Turbo".into(),
613            max_tokens: 4096,
614            embedding_model: None,
615            completion_tokens_param: None,
616            vision: None,
617        };
618        let dbg = format!("{cfg:?}");
619        assert!(!dbg.contains("sk-SUPERSECRET"));
620        assert!(dbg.contains("<redacted>"));
621    }
622}