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