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::openai_compat::OpenAICompatAdapter;
18use crate::models::{
19 Model, ModelConfig, ModelError, ProviderProfile, ReasoningChunk, Result, StreamCallback,
20 StreamEvent as ModelStreamEvent,
21};
22
23use super::super::capabilities::Capabilities;
24use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
25use super::ModelProvider;
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: 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 chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
57 let config = build_model_config(&request);
58 let (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
59 let callback = forward_callback(relay_tx.clone());
60 let chat_fut = self
61 .adapter
62 .chat(&request.messages, &config, Some(callback));
63
64 let response = tokio::select! {
65 biased;
66 _ = ctx.token.cancelled() => {
67 return Err(ModelError::Cancelled);
68 },
69 r = chat_fut => r?,
70 };
71
72 let usage = response.usage.clone();
73 let stop_reason = response.stop_reason.clone();
74 let _ = relay_tx.send(StreamEvent::Done {
78 usage: usage.clone(),
79 thinking_signature: None,
80 stop_reason: stop_reason.clone(),
81 });
82 drop(relay_tx);
83 let _ = relay_handle.await;
84
85 Ok(FinalResponse {
86 usage,
87 thinking_signature: None,
88 tool_calls: response.tool_calls.unwrap_or_default(),
89 stop_reason,
90 })
91 }
92}
93
94fn build_model_config(request: &ChatRequest) -> ModelConfig {
95 ModelConfig {
96 model: request.model_id.clone(),
97 temperature: request.temperature,
98 max_tokens: request.max_tokens,
99 reasoning: request.reasoning,
100 system_prompt: Some(request.system_prompt.clone()),
101 dynamic_system_suffix: request.instructions.clone(),
102 tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
103 ..Default::default()
104 }
105}
106
107fn forward_callback(sink: tokio::sync::mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
108 Arc::new(move |event: ModelStreamEvent| {
109 let mapped = match event {
110 ModelStreamEvent::Text(s) => StreamEvent::Text(s),
111 ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
112 text: chunk.text,
113 signature: chunk.signature,
114 }),
115 ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
116 ModelStreamEvent::Done { tokens } => StreamEvent::Done {
117 usage: if tokens > 0 {
118 Some(crate::models::TokenUsage::provider(0, tokens, tokens))
119 } else {
120 None
121 },
122 thinking_signature: None,
123 stop_reason: None,
124 },
125 };
126 let _ = sink.send(mapped);
127 })
128}
129
130#[cfg(test)]
131mod tests {
132 use super::*;
133
134 #[test]
135 fn build_model_config_maps_fields() {
136 let req = ChatRequest {
137 model_id: "groq/llama-3.3-70b-versatile".to_string(),
138 messages: vec![],
139 system_prompt: "sys".to_string(),
140 instructions: None,
141 reasoning: crate::models::ReasoningLevel::Medium,
142 temperature: 0.7,
143 max_tokens: 4096,
144 tools: vec![],
145
146 ollama_num_ctx: None,
147 ollama_allow_ram_offload: None,
148 };
149 let cfg = build_model_config(&req);
150 assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
151 }
152}