mermaid_cli/providers/model/
openai_compat.rs1use std::collections::HashMap;
12
13use async_trait::async_trait;
14
15use mermaid_domain::{ChatRequest, ToolDefinition};
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 chat_fut = async {
102 match self
103 .adapter
104 .chat(&request.messages, &config, Some(ctx.sink.clone()))
105 .await
106 {
107 Ok(response) => Ok(response),
108 Err(err) => {
109 let Some(cap) = output_cap_from_error(&err) else {
116 return Err(err);
117 };
118 let Some(clamped) = retry_cap(config.max_tokens, cap) else {
119 return Err(err);
120 };
121 let provider = self.adapter.provider_name().to_string();
122 let model = Model::name(&self.adapter).to_string();
123 learn_output_cap(provider, model.clone(), cap).await;
124 let _ = ctx.sink.send(StreamEvent::Status(format!(
125 "{model} rejected the output budget; learned its {cap}-token cap and retrying"
126 ))).await;
127 let retry_config = ModelConfig {
128 max_tokens: clamped,
129 ..config.clone()
130 };
131 self.adapter
132 .chat(&request.messages, &retry_config, Some(ctx.sink.clone()))
133 .await
134 },
135 }
136 };
137
138 let response = tokio::select! {
139 biased;
140 _ = ctx.token.cancelled() => {
141 return Err(ModelError::Cancelled);
142 },
143 r = chat_fut => r?,
144 };
145
146 let usage = response.usage.clone();
147 let stop_reason = response.stop_reason.clone();
148 let _ = ctx
151 .sink
152 .send(StreamEvent::Done {
153 usage: usage.clone(),
154 provider_continuation: None,
155 stop_reason: stop_reason.clone(),
156 })
157 .await;
158
159 Ok(FinalResponse {
160 usage,
161 provider_continuation: None,
162 tool_calls: response.tool_calls.unwrap_or_default(),
163 stop_reason,
164 })
165 }
166}
167
168fn build_model_config(request: &ChatRequest) -> ModelConfig {
169 ModelConfig {
170 model: request.model_id.clone(),
171 temperature: request.temperature,
172 max_tokens: request.max_tokens,
173 reasoning: request.reasoning,
174 system_prompt: Some(request.system_prompt.clone()),
175 dynamic_system_suffix: request.instructions.clone(),
176 tools: request
177 .tools
178 .iter()
179 .map(ToolDefinition::to_openai_json)
180 .collect(),
181 output_schema: request.output_schema.clone(),
182 ..Default::default()
183 }
184}
185
186#[cfg(test)]
187mod tests {
188 use super::*;
189
190 #[test]
191 fn build_model_config_maps_fields() {
192 let req = ChatRequest {
193 model_id: "groq/llama-3.3-70b-versatile".to_string(),
194 messages: vec![],
195 system_prompt: "sys".to_string(),
196 instructions: None,
197 reasoning: mermaid_model::models::ReasoningLevel::Medium,
198 temperature: 0.7,
199 max_tokens: 4096,
200 tools: vec![],
201
202 ollama_num_ctx: None,
203 ollama_allow_ram_offload: None,
204 resolved_context_window: None,
205 resolved_max_output: None,
206 output_schema: None,
207 suppress_auto_compact: false,
208 suppressed_builtin_tools: Vec::new(),
209 };
210 let cfg = build_model_config(&req);
211 assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
212 }
213}