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