mermaid_cli/providers/model/
mod.rs1pub 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#[derive(Debug, Clone, Copy, Default)]
35pub struct ContextSizing {
36 pub model_max: Option<usize>,
37 pub effective: Option<usize>,
38 pub source: Option<NumCtxSource>,
40 pub max_output: Option<usize>,
45}
46
47#[derive(Debug, Clone, Copy)]
52pub struct ModelPlacement {
53 pub size_vram_bytes: u64,
54 pub total_bytes: u64,
55 pub suggested_num_ctx: Option<u32>,
59}
60
61#[async_trait]
65pub trait ModelProvider: Send + Sync {
66 fn capabilities(&self) -> &ModelCapabilities;
70
71 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 async fn verify_placement(&self, current_num_ctx: Option<usize>) -> Option<ModelPlacement> {
95 let _ = current_num_ctx;
96 None
97 }
98
99 async fn supports_vision(&self) -> Option<bool> {
105 None
106 }
107
108 async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse>;
118}
119
120#[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
129pub(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 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
194pub(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 Err(_) => true,
208 }
209}
210
211#[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
224pub(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
243pub(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
266pub(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 Err(_) => None,
294 }
295}
296
297pub(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
321fn text_after<'a>(haystack: &'a str, marker: &str) -> Option<&'a str> {
323 haystack.find(marker).map(|i| &haystack[i + marker.len()..])
324}
325
326fn 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
339pub(crate) fn retry_cap(requested: usize, learned: usize) -> Option<usize> {
345 (requested == 0 || requested > learned).then_some(learned)
346}
347
348pub(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
361pub(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 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 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 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 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 assert_eq!(retry_cap(0, 131_072), Some(131_072));
453 assert_eq!(retry_cap(521_276, 131_072), Some(131_072));
455 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 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}