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 _ = relay_tx.send(StreamEvent::Done {
77 usage: usage.clone(),
78 thinking_signature: None,
79 });
80 drop(relay_tx);
81 let _ = relay_handle.await;
82
83 Ok(FinalResponse {
84 usage,
85 thinking_signature: None,
86 tool_calls: response.tool_calls.unwrap_or_default(),
87 })
88 }
89}
90
91fn build_model_config(request: &ChatRequest) -> ModelConfig {
92 ModelConfig {
93 model: request.model_id.clone(),
94 temperature: request.temperature,
95 max_tokens: request.max_tokens,
96 reasoning: request.reasoning,
97 system_prompt: Some(request.system_prompt.clone()),
98 dynamic_system_suffix: request.instructions.clone(),
99 tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
100 ..Default::default()
101 }
102}
103
104fn forward_callback(sink: tokio::sync::mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
105 Arc::new(move |event: ModelStreamEvent| {
106 let mapped = match event {
107 ModelStreamEvent::Text(s) => StreamEvent::Text(s),
108 ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
109 text: chunk.text,
110 signature: chunk.signature,
111 }),
112 ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
113 ModelStreamEvent::Done { tokens } => StreamEvent::Done {
114 usage: if tokens > 0 {
115 Some(crate::models::TokenUsage::provider(0, tokens, tokens))
116 } else {
117 None
118 },
119 thinking_signature: None,
120 },
121 };
122 let _ = sink.send(mapped);
123 })
124}
125
126#[cfg(test)]
127mod tests {
128 use super::*;
129
130 #[test]
131 fn build_model_config_maps_fields() {
132 let req = ChatRequest {
133 model_id: "groq/llama-3.3-70b-versatile".to_string(),
134 messages: vec![],
135 system_prompt: "sys".to_string(),
136 instructions: None,
137 reasoning: crate::models::ReasoningLevel::Medium,
138 temperature: 0.7,
139 max_tokens: 4096,
140 tools: vec![],
141 };
142 let cfg = build_model_config(&req);
143 assert_eq!(cfg.model, "groq/llama-3.3-70b-versatile");
144 }
145}