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;
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 (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
102        let callback = super::stream_bridge::forward_callback(relay_tx.clone());
103        let chat_fut = async {
104            match self
105                .adapter
106                .chat(&request.messages, &config, Some(callback.clone()))
107                .await
108            {
109                Ok(response) => Ok(response),
110                Err(err) => {
111                    // Learn-from-400 parity with the Ollama wrapper. AUTO
112                    // omits max_tokens here, so this fires mainly for
113                    // explicit user caps above the model's real ceiling:
114                    // learn the cap (persisted for later sizing), clamp,
115                    // retry ONCE. A 400 streamed no events, so the relay is
116                    // untouched and reusable.
117                    let Some(cap) = output_cap_from_error(&err) else {
118                        return Err(err);
119                    };
120                    let Some(clamped) = retry_cap(config.max_tokens, cap) else {
121                        return Err(err);
122                    };
123                    let provider = self.adapter.provider_name().to_string();
124                    let model = Model::name(&self.adapter).to_string();
125                    learn_output_cap(provider, model.clone(), cap).await;
126                    let _ = relay_tx.send(StreamEvent::Status(format!(
127                        "{model} rejected the output budget; learned its {cap}-token cap and retrying"
128                    )));
129                    let retry_config = ModelConfig {
130                        max_tokens: clamped,
131                        ..config.clone()
132                    };
133                    self.adapter
134                        .chat(&request.messages, &retry_config, Some(callback.clone()))
135                        .await
136                },
137            }
138        };
139
140        let response = tokio::select! {
141            biased;
142            _ = ctx.token.cancelled() => {
143                return Err(ModelError::Cancelled);
144            },
145            r = chat_fut => r?,
146        };
147
148        let usage = response.usage.clone();
149        let stop_reason = response.stop_reason.clone();
150        // Route the terminal Done through the SAME ordered relay (not directly
151        // on the bounded sink) so it can't overtake a still-buffered ToolCall,
152        // then await the relay drain before returning.
153        let _ = relay_tx.send(StreamEvent::Done {
154            usage: usage.clone(),
155            provider_continuation: None,
156            stop_reason: stop_reason.clone(),
157        });
158        drop(relay_tx);
159        mermaid_model::utils::join_logged(relay_handle.take(), "stream_relay").await;
160
161        Ok(FinalResponse {
162            usage,
163            provider_continuation: None,
164            tool_calls: response.tool_calls.unwrap_or_default(),
165            stop_reason,
166        })
167    }
168}
169
170fn build_model_config(request: &ChatRequest) -> ModelConfig {
171    ModelConfig {
172        model: request.model_id.clone(),
173        temperature: request.temperature,
174        max_tokens: request.max_tokens,
175        reasoning: request.reasoning,
176        system_prompt: Some(request.system_prompt.clone()),
177        dynamic_system_suffix: request.instructions.clone(),
178        tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
179        output_schema: request.output_schema.clone(),
180        ..Default::default()
181    }
182}
183
184#[cfg(test)]
185mod tests {
186    use super::*;
187
188    #[test]
189    fn build_model_config_maps_fields() {
190        let req = ChatRequest {
191            model_id: "groq/llama-3.3-70b-versatile".to_string(),
192            messages: vec![],
193            system_prompt: "sys".to_string(),
194            instructions: None,
195            reasoning: mermaid_model::models::ReasoningLevel::Medium,
196            temperature: 0.7,
197            max_tokens: 4096,
198            tools: vec![],
199
200            ollama_num_ctx: None,
201            ollama_allow_ram_offload: None,
202            resolved_context_window: None,
203            resolved_max_output: None,
204            output_schema: None,
205            suppress_auto_compact: false,
206            suppressed_builtin_tools: Vec::new(),
207        };
208        let cfg = build_model_config(&req);
209        assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
210    }
211}