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;
12
13use async_trait::async_trait;
14
15use mermaid_domain::{ChatRequest, ToolDefinition};
16use mermaid_model::models::adapters::ModelLimits;
17use mermaid_model::models::adapters::openai_compat::OpenAICompatAdapter;
18use mermaid_model::models::{Model, ModelConfig, ModelError, ProviderProfile, Result};
19
20use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
21use super::{
22    ContextSizing, ModelProvider, learn_output_cap, output_cap_from_error, resolve_limits_cached,
23    retry_cap,
24};
25use mermaid_model::models::ModelCapabilities;
26
27pub struct OpenAICompatProvider {
28    adapter: OpenAICompatAdapter,
29    capabilities: ModelCapabilities,
30}
31
32impl OpenAICompatProvider {
33    /// Wrap a fresh [`OpenAICompatAdapter`] as a `ModelProvider`.
34    ///
35    /// # Errors
36    ///
37    /// Only [`OpenAICompatAdapter::new`]'s — the HTTP client build. The
38    /// endpoint is not contacted here, so a wrong `base_url` or missing key
39    /// still constructs and fails on the first request.
40    pub fn new(
41        profile: &'static ProviderProfile,
42        base_url: String,
43        api_key: Option<String>,
44        model_name: String,
45        extra_headers: HashMap<String, String>,
46    ) -> Result<Self> {
47        let adapter =
48            OpenAICompatAdapter::new(profile, base_url, api_key, model_name, extra_headers)?;
49        let capabilities = adapter.capabilities().clone();
50        Ok(Self {
51            adapter,
52            capabilities,
53        })
54    }
55}
56
57#[async_trait]
58impl ModelProvider for OpenAICompatProvider {
59    fn capabilities(&self) -> &ModelCapabilities {
60        &self.capabilities
61    }
62
63    /// Live limit discovery: most OpenAI-compatible providers attach the
64    /// model's context window / output ceiling to their `/models` metadata
65    /// (OpenRouter et al); Cloudflare exposes it on its account-level
66    /// `models/search` endpoint instead (`list_models_for_limits` routes
67    /// there). Cache-first via `provider_probes` (TTL-bounded), one live
68    /// fetch on a miss, static fallback (all `None`) when the provider
69    /// exposes nothing or the fetch fails.
70    async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
71        let _ = request;
72        let provider = self.adapter.provider_name().to_string();
73        let model = Model::name(&self.adapter).to_string();
74        let limits = resolve_limits_cached(&provider, &model, || async {
75            let listings = self.adapter.list_models_for_limits().await?;
76            let found = listings.into_iter().find(|m| m.id == model);
77            Ok(ModelLimits {
78                max_context_tokens: found.as_ref().and_then(|m| m.max_context_tokens),
79                max_output_tokens: found.as_ref().and_then(|m| m.max_output_tokens),
80            })
81        })
82        .await;
83        let window = limits.as_ref().and_then(|l| l.max_context_tokens);
84        ContextSizing {
85            model_max: window,
86            effective: window,
87            source: None,
88            max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
89        }
90    }
91
92    async fn supports_vision(&self) -> Option<bool> {
93        // Report the model-driven capability (derived from the model id) so the
94        // no-vision-model warning fires for text-only models and stays quiet for
95        // genuine vision models — instead of the default `None` (never warns).
96        Some(self.capabilities.supports_vision)
97    }
98
99    async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
100        let config = build_model_config(&request);
101        let chat_fut = async {
102            match self
103                .adapter
104                .chat(&request.messages, &config, Some(ctx.sink.clone()))
105                .await
106            {
107                Ok(response) => Ok(response),
108                Err(err) => {
109                    // Learn-from-400 parity with the Ollama wrapper. AUTO
110                    // omits max_tokens here, so this fires mainly for
111                    // explicit user caps above the model's real ceiling:
112                    // learn the cap (persisted for later sizing), clamp,
113                    // retry ONCE. A 400 streamed no events, so the sink is
114                    // untouched and the retry starts from a clean stream.
115                    let Some(cap) = output_cap_from_error(&err) else {
116                        return Err(err);
117                    };
118                    let Some(clamped) = retry_cap(config.max_tokens, cap) else {
119                        return Err(err);
120                    };
121                    let provider = self.adapter.provider_name().to_string();
122                    let model = Model::name(&self.adapter).to_string();
123                    learn_output_cap(provider, model.clone(), cap).await;
124                    let _ = ctx.sink.send(StreamEvent::Status(format!(
125                        "{model} rejected the output budget; learned its {cap}-token cap and retrying"
126                    ))).await;
127                    let retry_config = ModelConfig {
128                        max_tokens: clamped,
129                        ..config.clone()
130                    };
131                    self.adapter
132                        .chat(&request.messages, &retry_config, Some(ctx.sink.clone()))
133                        .await
134                },
135            }
136        };
137
138        let response = tokio::select! {
139            biased;
140            _ = ctx.token.cancelled() => {
141                return Err(ModelError::Cancelled);
142            },
143            r = chat_fut => r?,
144        };
145
146        let usage = response.usage.clone();
147        let stop_reason = response.stop_reason.clone();
148        // The terminal Done goes on the same sink the adapter just finished
149        // writing to, so it cannot overtake a still-queued ToolCall.
150        let _ = ctx
151            .sink
152            .send(StreamEvent::Done {
153                usage: usage.clone(),
154                provider_continuation: None,
155                stop_reason: stop_reason.clone(),
156            })
157            .await;
158
159        Ok(FinalResponse {
160            usage,
161            provider_continuation: None,
162            tool_calls: response.tool_calls.unwrap_or_default(),
163            stop_reason,
164        })
165    }
166}
167
168fn build_model_config(request: &ChatRequest) -> ModelConfig {
169    ModelConfig {
170        model: request.model_id.clone(),
171        temperature: request.temperature,
172        max_tokens: request.max_tokens,
173        reasoning: request.reasoning,
174        system_prompt: Some(request.system_prompt.clone()),
175        dynamic_system_suffix: request.instructions.clone(),
176        tools: request
177            .tools
178            .iter()
179            .map(ToolDefinition::to_openai_json)
180            .collect(),
181        output_schema: request.output_schema.clone(),
182        ..Default::default()
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189
190    #[test]
191    fn build_model_config_maps_fields() {
192        let req = ChatRequest {
193            model_id: "groq/llama-3.3-70b-versatile".to_string(),
194            messages: vec![],
195            system_prompt: "sys".to_string(),
196            instructions: None,
197            reasoning: mermaid_model::models::ReasoningLevel::Medium,
198            temperature: 0.7,
199            max_tokens: 4096,
200            tools: vec![],
201
202            ollama_num_ctx: None,
203            ollama_allow_ram_offload: None,
204            resolved_context_window: None,
205            resolved_max_output: None,
206            output_schema: None,
207            suppress_auto_compact: false,
208            suppressed_builtin_tools: Vec::new(),
209        };
210        let cfg = build_model_config(&req);
211        assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
212    }
213}