Skip to main content

mermaid_cli/providers/model/
mod.rs

1//! Model adapters wrapped as `ModelProvider` implementations.
2//!
3//! Five providers today: Ollama, Anthropic, Gemini, Meta, and OpenAI-
4//! compat (covering OpenAI, OpenRouter, Groq, Cerebras, DeepInfra,
5//! Together, and user-defined endpoints). Each wraps the
6//! corresponding adapter in `mermaid_model::models::adapters`; the adapter
7//! owns the wire format and the wrapper owns the trait shape.
8
9pub mod anthropic;
10pub mod gemini;
11pub mod meta;
12pub mod ollama;
13pub mod openai_compat;
14
15use std::sync::Arc;
16
17use async_trait::async_trait;
18
19use mermaid_domain::{ChatRequest, TurnId};
20use mermaid_model::models::adapters::ModelLimits;
21use mermaid_model::models::adapters::ollama_sizing::NumCtxSource;
22use mermaid_model::models::{FinishReason, ModelError, Result, TokenUsage};
23use mermaid_runtime::NewProviderProbe;
24
25use super::ctx::{FinalResponse, StreamContext, StreamEvent};
26use mermaid_model::models::ModelCapabilities;
27
28/// Resolved context sizing for a turn. For most providers `model_max ==
29/// effective` (the static advertised window). For Ollama they differ:
30/// `model_max` is the probed architectural window, while `effective` is what we
31/// actually enforce as `num_ctx` (auto-fitted to memory, capped, or an
32/// override). Compaction and the status bar use `effective`; "model supports up
33/// to X" uses `model_max`.
34#[derive(Debug, Clone, Copy, Default)]
35pub struct ContextSizing {
36    pub model_max: Option<usize>,
37    pub effective: Option<usize>,
38    /// How `effective` was chosen (Ollama only). `None` for static/advertised.
39    pub source: Option<NumCtxSource>,
40    /// The model's per-response output ceiling, when the provider exposes one
41    /// (`/models` metadata, or a documented static table). Rides the same
42    /// resolve→reducer pipeline as the window so `provider_capabilities` can be
43    /// refreshed live.
44    pub max_output: Option<usize>,
45}
46
47/// Where a loaded model actually sits in memory, from a post-turn probe (Ollama
48/// `/api/ps`). `total_bytes` is weights + KV + buffers; `size_vram_bytes` is the
49/// part resident in VRAM. `size_vram_bytes < total_bytes` means the model spilled
50/// to CPU/RAM (partial offload → slow). Only Ollama reports this.
51#[derive(Debug, Clone, Copy)]
52pub struct ModelPlacement {
53    pub size_vram_bytes: u64,
54    pub total_bytes: u64,
55    /// Auto-converge target: when the model spilled, the largest `num_ctx` that
56    /// would fit instead — or `None` if it already fits or shrinking can't help
57    /// (weights-bound). Computed against the *measured* footprint.
58    pub suggested_num_ctx: Option<u32>,
59}
60
61/// Provider-facing interface. A `ModelProvider` impl owns whatever
62/// HTTP client / state it needs and exposes `chat()` — that's the
63/// whole surface.
64#[async_trait]
65pub trait ModelProvider: Send + Sync {
66    /// `ModelCapabilities` the provider advertises. The reducer reads this
67    /// when building the outgoing `ChatRequest` (e.g. whether to
68    /// attach reasoning controls).
69    fn capabilities(&self) -> &ModelCapabilities;
70
71    /// Resolve the *effective* context window for a turn (what the model will
72    /// actually enforce). The default returns the static advertised window;
73    /// Ollama overrides this to probe the model's real window and auto-fit
74    /// `num_ctx` to host memory, honoring the request's per-model
75    /// `ollama_num_ctx` override. `None` means "let the backend decide".
76    ///
77    /// Awaited only on the effect runtime (never the reducer), so a probe never
78    /// blocks the UI.
79    async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
80        let _ = request;
81        let max = self.capabilities().max_context_tokens;
82        ContextSizing {
83            model_max: max,
84            effective: max,
85            source: None,
86            max_output: self.capabilities().max_output_tokens,
87        }
88    }
89
90    /// Best-effort: where the loaded model currently sits in memory. The default
91    /// returns `None` (unknown / not applicable); Ollama overrides it to probe
92    /// `/api/ps`. Awaited only on the effect runtime, *after* a turn (when the
93    /// model is resident), so it never blocks the UI.
94    async fn verify_placement(&self, current_num_ctx: Option<usize>) -> Option<ModelPlacement> {
95        let _ = current_num_ctx;
96        None
97    }
98
99    /// Best-effort: whether the active model can actually see images. `None`
100    /// means "unknown / not applicable" — the default for providers that don't
101    /// probe, and for cloud providers whose vision support is already known
102    /// good; `Some(false)` is what drives the no-vision-model warning. Awaited
103    /// only on the effect runtime, so a probe never blocks the UI.
104    async fn supports_vision(&self) -> Option<bool> {
105        None
106    }
107
108    /// Stream a chat turn. Typed events flow through
109    /// `ctx.sink`; the returned `FinalResponse` carries token usage
110    /// and the Anthropic thinking-signature (opaque blob required to
111    /// continue extended thinking across turns).
112    ///
113    /// Cancellation: the provider MUST select! on `ctx.token.
114    /// cancelled()` inside any await that could block for more than
115    /// a few hundred ms. This is the contract that replaces the old
116    /// `check_interrupt` polling pattern.
117    async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse>;
118}
119
120/// Output collected from a one-shot, non-interactive streaming model call.
121#[derive(Debug, Clone, Default)]
122pub struct CollectedText {
123    pub text: String,
124    pub usage: Option<TokenUsage>,
125    pub stop_reason: Option<FinishReason>,
126    pub reasoning: Option<String>,
127}
128
129/// Run a one-shot, non-interactive model call and collect its streamed text
130/// into a `CollectedText`. For internal calls whose output is NOT shown to the user
131/// as it streams — context compaction, memory consolidation, and the Auto-mode safety classifier.
132/// Drains a private event channel, capturing plain text, reasoning trace, final token
133/// usage, and stop reasons (e.g. length limits / content filters). The `token` lets the
134/// caller cancel the call (e.g. on Ctrl+C) like any other turn work.
135pub(crate) async fn collect_text(
136    provider: Arc<dyn ModelProvider>,
137    turn: TurnId,
138    request: ChatRequest,
139    token: tokio_util::sync::CancellationToken,
140) -> Result<CollectedText> {
141    let (stream_tx, mut stream_rx) = tokio::sync::mpsc::channel::<StreamEvent>(128);
142    let ctx = StreamContext::new(token, stream_tx, turn);
143    let collector = tokio::task::spawn(async move {
144        let mut text = String::new();
145        let mut reasoning = String::new();
146        let mut usage = None;
147        let mut stop_reason = None;
148        while let Some(event) = stream_rx.recv().await {
149            match event {
150                StreamEvent::Text(chunk) => text.push_str(&chunk),
151                StreamEvent::Reasoning(chunk) => reasoning.push_str(&chunk.text),
152                StreamEvent::Done {
153                    usage: done_usage,
154                    stop_reason: done_stop,
155                    ..
156                } => {
157                    usage = done_usage;
158                    stop_reason = done_stop;
159                },
160                // Status is a user-facing plumbing notice, not content —
161                // a text collector has nowhere to surface it.
162                StreamEvent::ToolCall(_) | StreamEvent::Status(_) => {},
163            }
164        }
165        (
166            text,
167            (!reasoning.is_empty()).then_some(reasoning),
168            usage,
169            stop_reason,
170        )
171    });
172
173    let response = provider.chat(request, ctx).await;
174    let (text, reasoning, stream_usage, stream_stop_reason) = collector
175        .await
176        .map_err(|err| ModelError::StreamError(format!("collect_text collector failed: {err}")))?;
177    match response {
178        Ok(final_response) => Ok(CollectedText {
179            text,
180            usage: final_response.usage.or(stream_usage),
181            stop_reason: final_response.stop_reason.or(stream_stop_reason),
182            reasoning,
183        }),
184        Err(err) => Err(err),
185    }
186}
187
188pub use anthropic::AnthropicProvider;
189pub use gemini::GeminiProvider;
190pub use meta::MetaProvider;
191pub use ollama::OllamaProvider;
192pub use openai_compat::OpenAICompatProvider;
193
194/// True when a `provider_probes` row is older than the probe TTL (shared by
195/// the Ollama context probe and the per-provider limits probes). An
196/// unparseable timestamp is treated as stale so it re-probes.
197pub(crate) fn probe_is_stale(probed_at: &str) -> bool {
198    use chrono::{DateTime, Utc};
199    match DateTime::parse_from_rfc3339(probed_at) {
200        Ok(t) => {
201            Utc::now()
202                .signed_duration_since(t.with_timezone(&Utc))
203                .num_days()
204                >= mermaid_model::constants::PROVIDER_PROBE_TTL_DAYS
205        },
206        // Unparseable timestamp → treat as stale and re-probe.
207        Err(_) => true,
208    }
209}
210
211/// Limits learned from one live probe of a provider's models endpoint, cached
212/// per (provider, model) in `provider_probes`. A successful fetch WITHOUT
213/// limit metadata (or a definitive "model not listed") is cached too — as
214/// `None`s — so providers that don't expose limits aren't re-fetched every
215/// turn. Fetch *failures* are never cached.
216#[derive(serde::Serialize, serde::Deserialize)]
217pub(crate) struct CachedLimits {
218    pub(crate) max_context_tokens: Option<usize>,
219    pub(crate) max_output_tokens: Option<usize>,
220}
221
222pub(crate) const LIMITS_PROBE_KEY: &str = "limits_probe";
223
224/// Load fresh cached limits, off the async runtime. Best-effort → `None`.
225pub(crate) async fn load_limits_from_db(provider: String, model: String) -> Option<CachedLimits> {
226    tokio::task::spawn_blocking(move || {
227        let rec = mermaid_runtime::with_shared_store(|store| {
228            store
229                .provider_probes()
230                .get(&provider, &model, LIMITS_PROBE_KEY)
231        })
232        .ok()??;
233        if probe_is_stale(&rec.probed_at) {
234            return None;
235        }
236        serde_json::from_str::<CachedLimits>(&rec.capability_value).ok()
237    })
238    .await
239    .ok()
240    .flatten()
241}
242
243/// Persist probed limits for subsequent sessions. Best-effort.
244pub(crate) async fn save_limits_to_db(provider: String, model: String, limits: &CachedLimits) {
245    let value = match serde_json::to_string(limits) {
246        Ok(v) => v,
247        Err(_) => return,
248    };
249    let _ = tokio::task::spawn_blocking(move || -> Option<()> {
250        mermaid_runtime::with_shared_store(|store| {
251            store.provider_probes().upsert(NewProviderProbe {
252                provider,
253                model_id: model,
254                capability_key: LIMITS_PROBE_KEY.into(),
255                capability_value: value,
256                confidence: "probed".into(),
257                error: None,
258            })
259        })
260        .ok()?;
261        Some(())
262    })
263    .await;
264}
265
266/// Cache-first limit resolution: return fresh cached limits when present,
267/// otherwise run `fetch` against the provider's models endpoint. A successful
268/// fetch is cached even when all-`None` (definitive "provider exposes
269/// nothing"); a failed fetch is NOT cached and resolves to `None` so the next
270/// turn retries.
271pub(crate) async fn resolve_limits_cached<F, Fut>(
272    provider: &str,
273    model: &str,
274    fetch: F,
275) -> Option<CachedLimits>
276where
277    F: FnOnce() -> Fut,
278    Fut: std::future::Future<Output = Result<ModelLimits>>,
279{
280    if let Some(cached) = load_limits_from_db(provider.to_string(), model.to_string()).await {
281        return Some(cached);
282    }
283    match fetch().await {
284        Ok(limits) => {
285            let cached = CachedLimits {
286                max_context_tokens: limits.max_context_tokens,
287                max_output_tokens: limits.max_output_tokens,
288            };
289            save_limits_to_db(provider.to_string(), model.to_string(), &cached).await;
290            Some(cached)
291        },
292        // Network/parse failure: don't cache; caller falls back to static.
293        Err(_) => None,
294    }
295}
296
297/// Extract a model's real per-response output ceiling from a provider's 400
298/// rejection body. Fires ONLY on unambiguous output-cap wordings:
299///
300/// - Ollama Cloud / MiniMax: `max_tokens (521276) exceeds model's maximum
301///   output tokens (131072) for model …`
302/// - OpenAI-compat: `max_tokens is too large: … This model supports at most
303///   16384 completion tokens …`
304///
305/// Anything else — in particular context-limit wordings ("prompt is too
306/// long", "maximum context length") — returns `None`: learning a window as
307/// an output cap would poison the cache, and a missed match just means
308/// today's behavior (the error surfaces). Values outside a sanity range are
309/// rejected as parser noise.
310pub(crate) fn parse_output_cap_message(body: &str) -> Option<usize> {
311    let cap = if let Some(rest) = text_after(body, "exceeds model's maximum output tokens") {
312        leading_integer(rest)
313    } else if body.contains("max_tokens is too large") {
314        text_after(body, "supports at most").and_then(leading_integer)
315    } else {
316        None
317    }?;
318    (1_024..10_000_000).contains(&cap).then_some(cap)
319}
320
321/// The slice of `haystack` after the first occurrence of `marker`.
322fn text_after<'a>(haystack: &'a str, marker: &str) -> Option<&'a str> {
323    haystack.find(marker).map(|i| &haystack[i + marker.len()..])
324}
325
326/// The first integer in `s`, required to start within a few characters —
327/// both documented wordings put the number right after the marker (`" ("` /
328/// `" "`), and a distant number would belong to something else.
329fn leading_integer(s: &str) -> Option<usize> {
330    let start = s.find(|c: char| c.is_ascii_digit()).filter(|&i| i <= 8)?;
331    s[start..]
332        .chars()
333        .take_while(char::is_ascii_digit)
334        .collect::<String>()
335        .parse()
336        .ok()
337}
338
339/// Decide the output cap for a one-shot retry after learning `learned` from
340/// a 400. AUTO (`requested == 0`, including "field omitted") retries at the
341/// learned cap; an explicit ask above the cap retries clamped to it; an ask
342/// already within the cap returns `None` — the 400 was about something else,
343/// so retrying the same request would loop.
344pub(crate) fn retry_cap(requested: usize, learned: usize) -> Option<usize> {
345    (requested == 0 || requested > learned).then_some(learned)
346}
347
348/// `parse_output_cap_message` gated to actual HTTP 400s — the only status
349/// where the body names a rejected parameter rather than a transient fault.
350pub(crate) fn output_cap_from_error(err: &ModelError) -> Option<usize> {
351    match err {
352        ModelError::Backend(mermaid_model::models::BackendError::HttpError {
353            status: 400,
354            message,
355            ..
356        }) => parse_output_cap_message(message),
357        _ => None,
358    }
359}
360
361/// Persist an output cap learned from a provider's 400 rejection (the error
362/// body names the model's real ceiling). Merges into any existing cached
363/// row — reading it raw, ignoring the TTL, since a stale window is still
364/// better than dropping it — and upserts as "probed". Best-effort.
365pub(crate) async fn learn_output_cap(provider: String, model: String, cap: usize) {
366    let _ = tokio::task::spawn_blocking(move || -> Option<()> {
367        let existing = mermaid_runtime::with_shared_store(|store| {
368            store
369                .provider_probes()
370                .get(&provider, &model, LIMITS_PROBE_KEY)
371        })
372        .ok()
373        .flatten()
374        .and_then(|rec| serde_json::from_str::<CachedLimits>(&rec.capability_value).ok());
375        let merged = CachedLimits {
376            max_context_tokens: existing.and_then(|l| l.max_context_tokens),
377            max_output_tokens: Some(cap),
378        };
379        let value = serde_json::to_string(&merged).ok()?;
380        mermaid_runtime::with_shared_store(|store| {
381            store.provider_probes().upsert(NewProviderProbe {
382                provider,
383                model_id: model,
384                capability_key: LIMITS_PROBE_KEY.into(),
385                capability_value: value,
386                confidence: "probed".into(),
387                error: None,
388            })
389        })
390        .ok()?;
391        Some(())
392    })
393    .await;
394}
395
396#[cfg(test)]
397mod tests {
398    use super::*;
399
400    // The incident wording (Ollama Cloud, minimax-m3), raw and JSON-wrapped.
401    const MINIMAX_RAW: &str =
402        "max_tokens (521276) exceeds model's maximum output tokens (131072) for model minimax-m3";
403    const MINIMAX_JSON: &str = r#"{"error":"max_tokens (521276) exceeds model's maximum output tokens (131072) for model minimax-m3 (ref: a05c9ffb-168f)"}"#;
404    const OPENAI_STYLE: &str = r#"{"error":{"message":"max_tokens is too large: 200000. This model supports at most 16384 completion tokens, whereas you provided 200000.","type":"invalid_request_error"}}"#;
405
406    #[test]
407    fn parse_output_cap_matches_documented_wordings() {
408        assert_eq!(parse_output_cap_message(MINIMAX_RAW), Some(131_072));
409        assert_eq!(parse_output_cap_message(MINIMAX_JSON), Some(131_072));
410        assert_eq!(parse_output_cap_message(OPENAI_STYLE), Some(16_384));
411    }
412
413    #[test]
414    fn parse_output_cap_never_matches_context_limit_wordings() {
415        // Learning a context window as an output cap would poison the cache —
416        // these must all be None even though they mention token limits.
417        for body in [
418            "prompt is too long: 210000 tokens > 200000 maximum",
419            "This model's maximum context length is 128000 tokens",
420            "input length and max_tokens exceed context limit: 190000 + 20000 > 200000",
421            "the request exceeds the maximum context window of 131072 tokens",
422            "rate limit exceeded, try again in 20s",
423            "",
424        ] {
425            assert_eq!(parse_output_cap_message(body), None, "matched: {body}");
426        }
427    }
428
429    #[test]
430    fn parse_output_cap_rejects_nonsense_values() {
431        // Sub-1024 and absurd values are parser noise, not real ceilings.
432        assert_eq!(
433            parse_output_cap_message("exceeds model's maximum output tokens (512)"),
434            None
435        );
436        assert_eq!(
437            parse_output_cap_message("exceeds model's maximum output tokens (99999999999)"),
438            None
439        );
440        // A number too far from the marker belongs to something else.
441        assert_eq!(
442            parse_output_cap_message(
443                "exceeds model's maximum output tokens for this deployment tier which is 131072"
444            ),
445            None
446        );
447    }
448
449    #[test]
450    fn retry_cap_triple() {
451        // AUTO (0 / omitted) → retry at the learned cap.
452        assert_eq!(retry_cap(0, 131_072), Some(131_072));
453        // Explicit ask above the cap → clamp.
454        assert_eq!(retry_cap(521_276, 131_072), Some(131_072));
455        // Ask already within the cap → the 400 was about something else.
456        assert_eq!(retry_cap(4_096, 131_072), None);
457    }
458
459    #[test]
460    fn output_cap_from_error_gates_on_http_400() {
461        let err_400 = ModelError::Backend(mermaid_model::models::BackendError::HttpError {
462            status: 400,
463            message: MINIMAX_JSON.to_string(),
464            debug: Default::default(),
465        });
466        assert_eq!(output_cap_from_error(&err_400), Some(131_072));
467        // Same body on a 500 is a transient fault, not a learned limit.
468        let err_500 = ModelError::Backend(mermaid_model::models::BackendError::HttpError {
469            status: 500,
470            message: MINIMAX_JSON.to_string(),
471            debug: Default::default(),
472        });
473        assert_eq!(output_cap_from_error(&err_500), None);
474        assert_eq!(output_cap_from_error(&ModelError::Cancelled), None);
475    }
476}