Skip to main content

mermaid_cli/providers/model/
openai_compat.rs

1//! OpenAI-compatible provider — wraps
2//! `models::adapters::openai_compat::OpenAICompatAdapter`.
3//!
4//! This provider covers the OpenAI long-tail: OpenRouter, Groq,
5//! Fireworks, Together, custom vLLM endpoints, plus the user-defined
6//! entries in `[providers.*]`. The adapter looks up a
7//! `ProviderProfile` (registry entry) and applies per-provider
8//! reasoning shapes (flat `reasoning_effort` vs nested `reasoning:
9//! {effort}`). This wrapper just forwards.
10
11use 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        // Route the terminal Done through the SAME ordered relay (not directly
75        // on the bounded sink) so it can't overtake a still-buffered ToolCall,
76        // then await the relay drain before returning.
77        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}