mermaid_cli/providers/model/
openai_compat.rs1use 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 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 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 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 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 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}