Skip to main content

mermaid_cli/providers/model/
openai_compat.rs

1//! OpenAI-compatible provider — wraps
2//! `models::adapters::openai_compat::OpenAICompatAdapter`.
3//!
4//! This provider covers the OpenAI long-tail: OpenRouter, Groq,
5//! Fireworks, Together, custom vLLM endpoints, plus the user-defined
6//! entries in `[providers.*]`. The adapter looks up a
7//! `ProviderProfile` (registry entry) and applies per-provider
8//! reasoning shapes (flat `reasoning_effort` vs nested `reasoning:
9//! {effort}`). This wrapper just forwards.
10
11use std::collections::HashMap;
12use std::sync::Arc;
13
14use async_trait::async_trait;
15
16use crate::domain::ChatRequest;
17use crate::models::adapters::ModelLimits;
18use crate::models::adapters::openai_compat::OpenAICompatAdapter;
19use crate::models::{
20    Model, ModelConfig, ModelError, ProviderProfile, ReasoningChunk, Result, StreamCallback,
21    StreamEvent as ModelStreamEvent,
22};
23
24use super::super::capabilities::Capabilities;
25use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
26use super::{
27    ContextSizing, ModelProvider, learn_output_cap, output_cap_from_error, resolve_limits_cached,
28    retry_cap,
29};
30
31pub struct OpenAICompatProvider {
32    adapter: OpenAICompatAdapter,
33    capabilities: Capabilities,
34}
35
36impl OpenAICompatProvider {
37    pub fn new(
38        profile: &'static ProviderProfile,
39        base_url: String,
40        api_key: Option<String>,
41        model_name: String,
42        extra_headers: HashMap<String, String>,
43    ) -> Result<Self> {
44        let adapter =
45            OpenAICompatAdapter::new(profile, base_url, api_key, model_name, extra_headers)?;
46        let capabilities = Capabilities::from_legacy(adapter.capabilities());
47        Ok(Self {
48            adapter,
49            capabilities,
50        })
51    }
52}
53
54#[async_trait]
55impl ModelProvider for OpenAICompatProvider {
56    fn capabilities(&self) -> &Capabilities {
57        &self.capabilities
58    }
59
60    /// Live limit discovery: most OpenAI-compatible providers attach the
61    /// model's context window / output ceiling to their `/models` metadata
62    /// (OpenRouter et al); Cloudflare exposes it on its account-level
63    /// `models/search` endpoint instead (`list_models_for_limits` routes
64    /// there). Cache-first via `provider_probes` (TTL-bounded), one live
65    /// fetch on a miss, static fallback (all `None`) when the provider
66    /// exposes nothing or the fetch fails.
67    async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
68        let _ = request;
69        let provider = self.adapter.provider_name().to_string();
70        let model = Model::name(&self.adapter).to_string();
71        let limits = resolve_limits_cached(&provider, &model, || async {
72            let listings = self.adapter.list_models_for_limits().await?;
73            let found = listings.into_iter().find(|m| m.id == model);
74            Ok(ModelLimits {
75                max_context_tokens: found.as_ref().and_then(|m| m.max_context_tokens),
76                max_output_tokens: found.as_ref().and_then(|m| m.max_output_tokens),
77            })
78        })
79        .await;
80        let window = limits.as_ref().and_then(|l| l.max_context_tokens);
81        ContextSizing {
82            model_max: window,
83            effective: window,
84            source: None,
85            max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
86        }
87    }
88
89    async fn supports_vision(&self) -> Option<bool> {
90        // Report the model-driven capability (derived from the model id) so the
91        // no-vision-model warning fires for text-only models and stays quiet for
92        // genuine vision models — instead of the default `None` (never warns).
93        Some(self.capabilities.supports_vision)
94    }
95
96    async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
97        let config = build_model_config(&request);
98        let (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
99        let callback = forward_callback(relay_tx.clone());
100        let chat_fut = async {
101            match self
102                .adapter
103                .chat(&request.messages, &config, Some(callback.clone()))
104                .await
105            {
106                Ok(response) => Ok(response),
107                Err(err) => {
108                    // Learn-from-400 parity with the Ollama wrapper. AUTO
109                    // omits max_tokens here, so this fires mainly for
110                    // explicit user caps above the model's real ceiling:
111                    // learn the cap (persisted for later sizing), clamp,
112                    // retry ONCE. A 400 streamed no events, so the relay is
113                    // untouched and reusable.
114                    let Some(cap) = output_cap_from_error(&err) else {
115                        return Err(err);
116                    };
117                    let Some(clamped) = retry_cap(config.max_tokens, cap) else {
118                        return Err(err);
119                    };
120                    let provider = self.adapter.provider_name().to_string();
121                    let model = Model::name(&self.adapter).to_string();
122                    learn_output_cap(provider, model.clone(), cap).await;
123                    let _ = relay_tx.send(StreamEvent::Status(format!(
124                        "{model} rejected the output budget; learned its {cap}-token cap and retrying"
125                    )));
126                    let retry_config = ModelConfig {
127                        max_tokens: clamped,
128                        ..config.clone()
129                    };
130                    self.adapter
131                        .chat(&request.messages, &retry_config, Some(callback.clone()))
132                        .await
133                },
134            }
135        };
136
137        let response = tokio::select! {
138            biased;
139            _ = ctx.token.cancelled() => {
140                return Err(ModelError::Cancelled);
141            },
142            r = chat_fut => r?,
143        };
144
145        let usage = response.usage.clone();
146        let stop_reason = response.stop_reason.clone();
147        // Route the terminal Done through the SAME ordered relay (not directly
148        // on the bounded sink) so it can't overtake a still-buffered ToolCall,
149        // then await the relay drain before returning.
150        let _ = relay_tx.send(StreamEvent::Done {
151            usage: usage.clone(),
152            provider_continuation: None,
153            stop_reason: stop_reason.clone(),
154        });
155        drop(relay_tx);
156        crate::utils::join_logged(relay_handle.take(), "stream_relay").await;
157
158        Ok(FinalResponse {
159            usage,
160            provider_continuation: None,
161            tool_calls: response.tool_calls.unwrap_or_default(),
162            stop_reason,
163        })
164    }
165}
166
167fn build_model_config(request: &ChatRequest) -> ModelConfig {
168    ModelConfig {
169        model: request.model_id.clone(),
170        temperature: request.temperature,
171        max_tokens: request.max_tokens,
172        reasoning: request.reasoning,
173        system_prompt: Some(request.system_prompt.clone()),
174        dynamic_system_suffix: request.instructions.clone(),
175        tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
176        output_schema: request.output_schema.clone(),
177        ..Default::default()
178    }
179}
180
181fn forward_callback(sink: tokio::sync::mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
182    Arc::new(move |event: ModelStreamEvent| {
183        let mapped = match event {
184            ModelStreamEvent::Text(s) => StreamEvent::Text(s),
185            ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
186                text: chunk.text,
187                signature: chunk.signature,
188            }),
189            ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
190            ModelStreamEvent::Status(s) => StreamEvent::Status(s),
191            // No adapter emits `Done` through this callback — the wrapper
192            // sends the authoritative terminal `Done` built from the
193            // returned `ModelResponse` (F3). Map defensively without
194            // inventing usage (the old placeholder misfiled everything
195            // as completion tokens).
196            ModelStreamEvent::Done { .. } => StreamEvent::Done {
197                usage: None,
198                provider_continuation: None,
199                stop_reason: None,
200            },
201        };
202        let _ = sink.send(mapped);
203    })
204}
205
206#[cfg(test)]
207mod tests {
208    use super::*;
209
210    #[test]
211    fn build_model_config_maps_fields() {
212        let req = ChatRequest {
213            model_id: "groq/llama-3.3-70b-versatile".to_string(),
214            messages: vec![],
215            system_prompt: "sys".to_string(),
216            instructions: None,
217            reasoning: crate::models::ReasoningLevel::Medium,
218            temperature: 0.7,
219            max_tokens: 4096,
220            tools: vec![],
221
222            ollama_num_ctx: None,
223            ollama_allow_ram_offload: None,
224            resolved_context_window: None,
225            resolved_max_output: None,
226            output_schema: None,
227            suppress_auto_compact: false,
228        };
229        let cfg = build_model_config(&req);
230        assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
231    }
232}