mermaid_cli/providers/model/
openai_compat.rs1use std::collections::HashMap;
12use std::sync::Arc;
13
14use async_trait::async_trait;
15
16use crate::domain::ChatRequest;
17use crate::models::adapters::ModelLimits;
18use crate::models::adapters::openai_compat::OpenAICompatAdapter;
19use crate::models::{
20 Model, ModelConfig, ModelError, ProviderProfile, ReasoningChunk, Result, StreamCallback,
21 StreamEvent as ModelStreamEvent,
22};
23
24use super::super::capabilities::Capabilities;
25use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
26use super::{
27 ContextSizing, ModelProvider, learn_output_cap, output_cap_from_error, resolve_limits_cached,
28 retry_cap,
29};
30
31pub struct OpenAICompatProvider {
32 adapter: OpenAICompatAdapter,
33 capabilities: Capabilities,
34}
35
36impl OpenAICompatProvider {
37 pub fn new(
38 profile: &'static ProviderProfile,
39 base_url: String,
40 api_key: Option<String>,
41 model_name: String,
42 extra_headers: HashMap<String, String>,
43 ) -> Result<Self> {
44 let adapter =
45 OpenAICompatAdapter::new(profile, base_url, api_key, model_name, extra_headers)?;
46 let capabilities = Capabilities::from_legacy(adapter.capabilities());
47 Ok(Self {
48 adapter,
49 capabilities,
50 })
51 }
52}
53
54#[async_trait]
55impl ModelProvider for OpenAICompatProvider {
56 fn capabilities(&self) -> &Capabilities {
57 &self.capabilities
58 }
59
60 async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
68 let _ = request;
69 let provider = self.adapter.provider_name().to_string();
70 let model = Model::name(&self.adapter).to_string();
71 let limits = resolve_limits_cached(&provider, &model, || async {
72 let listings = self.adapter.list_models_for_limits().await?;
73 let found = listings.into_iter().find(|m| m.id == model);
74 Ok(ModelLimits {
75 max_context_tokens: found.as_ref().and_then(|m| m.max_context_tokens),
76 max_output_tokens: found.as_ref().and_then(|m| m.max_output_tokens),
77 })
78 })
79 .await;
80 let window = limits.as_ref().and_then(|l| l.max_context_tokens);
81 ContextSizing {
82 model_max: window,
83 effective: window,
84 source: None,
85 max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
86 }
87 }
88
89 async fn supports_vision(&self) -> Option<bool> {
90 Some(self.capabilities.supports_vision)
94 }
95
96 async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
97 let config = build_model_config(&request);
98 let (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
99 let callback = forward_callback(relay_tx.clone());
100 let chat_fut = async {
101 match self
102 .adapter
103 .chat(&request.messages, &config, Some(callback.clone()))
104 .await
105 {
106 Ok(response) => Ok(response),
107 Err(err) => {
108 let Some(cap) = output_cap_from_error(&err) else {
115 return Err(err);
116 };
117 let Some(clamped) = retry_cap(config.max_tokens, cap) else {
118 return Err(err);
119 };
120 let provider = self.adapter.provider_name().to_string();
121 let model = Model::name(&self.adapter).to_string();
122 learn_output_cap(provider, model.clone(), cap).await;
123 let _ = relay_tx.send(StreamEvent::Status(format!(
124 "{model} rejected the output budget; learned its {cap}-token cap and retrying"
125 )));
126 let retry_config = ModelConfig {
127 max_tokens: clamped,
128 ..config.clone()
129 };
130 self.adapter
131 .chat(&request.messages, &retry_config, Some(callback.clone()))
132 .await
133 },
134 }
135 };
136
137 let response = tokio::select! {
138 biased;
139 _ = ctx.token.cancelled() => {
140 return Err(ModelError::Cancelled);
141 },
142 r = chat_fut => r?,
143 };
144
145 let usage = response.usage.clone();
146 let stop_reason = response.stop_reason.clone();
147 let _ = relay_tx.send(StreamEvent::Done {
151 usage: usage.clone(),
152 provider_continuation: None,
153 stop_reason: stop_reason.clone(),
154 });
155 drop(relay_tx);
156 crate::utils::join_logged(relay_handle.take(), "stream_relay").await;
157
158 Ok(FinalResponse {
159 usage,
160 provider_continuation: None,
161 tool_calls: response.tool_calls.unwrap_or_default(),
162 stop_reason,
163 })
164 }
165}
166
167fn build_model_config(request: &ChatRequest) -> ModelConfig {
168 ModelConfig {
169 model: request.model_id.clone(),
170 temperature: request.temperature,
171 max_tokens: request.max_tokens,
172 reasoning: request.reasoning,
173 system_prompt: Some(request.system_prompt.clone()),
174 dynamic_system_suffix: request.instructions.clone(),
175 tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
176 output_schema: request.output_schema.clone(),
177 ..Default::default()
178 }
179}
180
181fn forward_callback(sink: tokio::sync::mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
182 Arc::new(move |event: ModelStreamEvent| {
183 let mapped = match event {
184 ModelStreamEvent::Text(s) => StreamEvent::Text(s),
185 ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
186 text: chunk.text,
187 signature: chunk.signature,
188 }),
189 ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
190 ModelStreamEvent::Status(s) => StreamEvent::Status(s),
191 ModelStreamEvent::Done { .. } => StreamEvent::Done {
197 usage: None,
198 provider_continuation: None,
199 stop_reason: None,
200 },
201 };
202 let _ = sink.send(mapped);
203 })
204}
205
206#[cfg(test)]
207mod tests {
208 use super::*;
209
210 #[test]
211 fn build_model_config_maps_fields() {
212 let req = ChatRequest {
213 model_id: "groq/llama-3.3-70b-versatile".to_string(),
214 messages: vec![],
215 system_prompt: "sys".to_string(),
216 instructions: None,
217 reasoning: crate::models::ReasoningLevel::Medium,
218 temperature: 0.7,
219 max_tokens: 4096,
220 tools: vec![],
221
222 ollama_num_ctx: None,
223 ollama_allow_ram_offload: None,
224 resolved_context_window: None,
225 resolved_max_output: None,
226 output_schema: None,
227 suppress_auto_compact: false,
228 };
229 let cfg = build_model_config(&req);
230 assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
231 }
232}