Skip to main content

llm/providers/openai_compatible/
types.rs

1use async_openai::types::chat::{
2    ChatCompletionMessageToolCall, ChatCompletionMessageToolCalls, ChatCompletionStreamOptions, ChatCompletionTools,
3    FunctionCall, Role,
4};
5use serde::{Deserialize, Serialize};
6
7use crate::{ChatMessage, ContentBlock, TokenUsage};
8
9/// Unified custom types for OpenAI-compatible APIs that deviate slightly from the standard.
10/// This handles quirks from providers like `OpenRouter`, Z.ai, and potentially others.
11///
12/// Common deviations handled:
13/// - Missing 'object' field (z.ai)
14/// - Negative token counts (openrouter)
15/// - Additional finish reasons like 'error' (openrouter)
16/// - Optional `system_fingerprint` and usage fields
17
18#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
19#[serde(rename_all = "snake_case")]
20pub enum FinishReason {
21    Stop,
22    Length,
23    ToolCalls,
24    ContentFilter,
25    FunctionCall,
26    Error,
27    NetworkError,
28    ModelContextWindowExceeded,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize)]
32pub struct ChatCompletionStreamResponse {
33    pub id: String,
34    pub choices: Vec<ChatCompletionStreamChoice>,
35    pub created: u64,
36    pub model: String,
37    #[serde(default)]
38    pub system_fingerprint: Option<String>,
39    #[serde(default = "default_object")]
40    pub object: String,
41    #[serde(default)]
42    pub usage: Option<Usage>,
43}
44
45fn default_object() -> String {
46    "chat.completion.chunk".to_string()
47}
48
49#[derive(Debug, Clone, Serialize)]
50#[serde(untagged)]
51pub(crate) enum UserContent {
52    Text(String),
53    Parts(Vec<UserContentPart>),
54}
55
56#[derive(Debug, Clone, Serialize)]
57#[serde(tag = "type", rename_all = "snake_case")]
58pub(crate) enum UserContentPart {
59    Text { text: String },
60    ImageUrl { image_url: ImageUrlContent },
61}
62
63#[derive(Debug, Clone, Serialize)]
64pub(crate) struct ImageUrlContent {
65    pub url: String,
66}
67
68#[derive(Debug, Clone, Serialize)]
69#[serde(tag = "role", rename_all = "lowercase")]
70pub(crate) enum CompatibleChatMessage {
71    System {
72        content: String,
73    },
74    User {
75        content: UserContent,
76    },
77    Assistant {
78        content: String,
79        #[serde(skip_serializing_if = "Option::is_none")]
80        reasoning_content: Option<String>,
81        #[serde(skip_serializing_if = "Option::is_none")]
82        tool_calls: Option<Vec<ChatCompletionMessageToolCalls>>,
83    },
84    Tool {
85        content: String,
86        tool_call_id: String,
87    },
88}
89
90#[derive(Debug, Clone, Serialize)]
91pub(crate) struct CompatibleChatRequest {
92    pub model: String,
93    pub messages: Vec<CompatibleChatMessage>,
94    #[serde(skip_serializing_if = "Option::is_none")]
95    pub stream: Option<bool>,
96    #[serde(skip_serializing_if = "Option::is_none")]
97    pub tools: Option<Vec<ChatCompletionTools>>,
98    #[serde(skip_serializing_if = "Option::is_none")]
99    pub stream_options: Option<ChatCompletionStreamOptions>,
100    #[serde(skip_serializing_if = "Option::is_none")]
101    pub reasoning_effort: Option<crate::ReasoningEffort>,
102    #[serde(skip_serializing_if = "Option::is_none")]
103    pub temperature: Option<f32>,
104    #[serde(skip_serializing_if = "Option::is_none")]
105    pub top_p: Option<f32>,
106    #[serde(skip_serializing_if = "Option::is_none")]
107    pub max_tokens: Option<u32>,
108    /// OpenAI-style prompt-cache affinity. Only set for providers that document support.
109    #[serde(skip_serializing_if = "Option::is_none")]
110    pub prompt_cache_key: Option<String>,
111}
112
113pub(crate) fn map_messages(messages: &[ChatMessage]) -> crate::Result<Vec<CompatibleChatMessage>> {
114    let mut result = Vec::new();
115
116    for message in messages {
117        let mapped = match message {
118            ChatMessage::System { content, .. } => Some(CompatibleChatMessage::System { content: content.clone() }),
119            ChatMessage::User { content, .. } => {
120                Some(CompatibleChatMessage::User { content: map_user_content(content)? })
121            }
122            ChatMessage::Assistant { content, reasoning, tool_calls, .. } => {
123                let openai_tool_calls: Vec<_> = tool_calls
124                    .iter()
125                    .map(|call| {
126                        ChatCompletionMessageToolCalls::Function(ChatCompletionMessageToolCall {
127                            id: call.id.clone(),
128                            function: FunctionCall { name: call.name.clone(), arguments: call.arguments.clone() },
129                        })
130                    })
131                    .collect();
132
133                let has_tool_calls = !openai_tool_calls.is_empty();
134                let tool_calls = has_tool_calls.then_some(openai_tool_calls);
135
136                let reasoning_content = if reasoning.summary_text.is_some() {
137                    reasoning.summary_text.clone()
138                } else if has_tool_calls {
139                    Some(".".to_string())
140                } else {
141                    None
142                };
143
144                Some(CompatibleChatMessage::Assistant { content: content.clone(), reasoning_content, tool_calls })
145            }
146            ChatMessage::ToolCallResult(r) => {
147                let (content, tool_call_id) = match r {
148                    Ok(tool_result) => (tool_result.result.clone(), tool_result.id.clone()),
149                    Err(tool_error) => (tool_error.error.clone(), tool_error.id.clone()),
150                };
151
152                Some(CompatibleChatMessage::Tool { content, tool_call_id })
153            }
154            ChatMessage::Summary { content, .. } => Some(CompatibleChatMessage::User {
155                content: UserContent::Text(format!("[Previous conversation handoff]\n\n{content}")),
156            }),
157            ChatMessage::Error { .. } => None,
158        };
159
160        if let Some(msg) = mapped {
161            result.push(msg);
162        }
163    }
164
165    Ok(result)
166}
167
168fn map_user_content(parts: &[ContentBlock]) -> crate::Result<UserContent> {
169    let has_non_text = parts.iter().any(|p| !matches!(p, ContentBlock::Text { .. }));
170
171    if !has_non_text {
172        return Ok(UserContent::Text(ContentBlock::join_text(parts)));
173    }
174
175    let mut items = Vec::with_capacity(parts.len());
176    for p in parts {
177        match p {
178            ContentBlock::Text { text } => items.push(UserContentPart::Text { text: text.clone() }),
179            ContentBlock::Image { .. } => {
180                items.push(UserContentPart::ImageUrl { image_url: ImageUrlContent { url: p.as_data_uri().unwrap() } });
181            }
182            ContentBlock::Audio { .. } => {
183                return Err(crate::LlmError::UnsupportedContent("This provider does not support audio input".into()));
184            }
185        }
186    }
187
188    Ok(UserContent::Parts(items))
189}
190
191#[derive(Debug, Clone, Serialize, Deserialize)]
192pub struct ChatCompletionStreamChoice {
193    pub index: i32,
194    pub delta: ChatCompletionStreamResponseDelta,
195    pub finish_reason: Option<FinishReason>,
196    #[serde(default)]
197    pub logprobs: Option<serde_json::Value>,
198}
199
200#[derive(Debug, Clone, Default, Serialize, Deserialize)]
201pub struct ChatCompletionStreamResponseDelta {
202    pub role: Option<Role>,
203    pub content: Option<String>,
204    #[serde(default)]
205    pub reasoning_content: Option<String>,
206    pub tool_calls: Option<Vec<ToolCallDelta>>,
207}
208
209#[derive(Debug, Clone, Serialize, Deserialize)]
210pub struct ToolCallDelta {
211    pub index: i32,
212    pub id: Option<String>,
213    #[serde(rename = "type")]
214    pub tool_type: Option<String>,
215    pub function: Option<FunctionCallDelta>,
216}
217
218#[derive(Debug, Clone, Serialize, Deserialize)]
219pub struct FunctionCallDelta {
220    pub name: Option<String>,
221    pub arguments: Option<String>,
222}
223
224#[derive(Debug, Clone, Default, Serialize, Deserialize)]
225pub struct PromptTokensDetails {
226    #[serde(default)]
227    pub cached_tokens: Option<u32>,
228    /// `OpenRouter`-specific: tokens written to cache (cache creation).
229    /// Only returned for models with explicit caching and cache write pricing.
230    #[serde(default)]
231    pub cache_write_tokens: Option<u32>,
232    /// `OpenAI` + `OpenRouter`: input audio tokens.
233    #[serde(default)]
234    pub audio_tokens: Option<u32>,
235    /// `OpenRouter`-specific: input video tokens.
236    #[serde(default)]
237    pub video_tokens: Option<u32>,
238}
239
240#[derive(Debug, Clone, Default, Serialize, Deserialize)]
241pub struct CompletionTokensDetails {
242    #[serde(default)]
243    pub reasoning_tokens: Option<u32>,
244    #[serde(default)]
245    pub audio_tokens: Option<u32>,
246    #[serde(default)]
247    pub accepted_prediction_tokens: Option<u32>,
248    #[serde(default)]
249    pub rejected_prediction_tokens: Option<u32>,
250}
251
252#[derive(Debug, Clone, Default, Serialize, Deserialize)]
253pub struct Usage {
254    pub prompt_tokens: i64,
255    pub completion_tokens: i64,
256    pub total_tokens: i64,
257    #[serde(default)]
258    pub prompt_tokens_details: Option<PromptTokensDetails>,
259    #[serde(default)]
260    pub completion_tokens_details: Option<CompletionTokensDetails>,
261}
262
263impl From<Usage> for TokenUsage {
264    fn from(usage: Usage) -> Self {
265        let prompt = usage.prompt_tokens_details.unwrap_or_default();
266        let completion = usage.completion_tokens_details.unwrap_or_default();
267        TokenUsage {
268            input_tokens: u32::try_from(usage.prompt_tokens.max(0)).unwrap_or(0).into(),
269            output_tokens: u32::try_from(usage.completion_tokens.max(0)).unwrap_or(0).into(),
270            cache_read_tokens: prompt.cached_tokens.map(Into::into),
271            cache_creation_tokens: prompt.cache_write_tokens.map(Into::into),
272            input_audio_tokens: prompt.audio_tokens.map(Into::into),
273            input_video_tokens: prompt.video_tokens.map(Into::into),
274            reasoning_tokens: completion.reasoning_tokens.map(Into::into),
275            output_audio_tokens: completion.audio_tokens.map(Into::into),
276            accepted_prediction_tokens: completion.accepted_prediction_tokens.map(Into::into),
277            rejected_prediction_tokens: completion.rejected_prediction_tokens.map(Into::into),
278        }
279    }
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285    use crate::providers::openai_compatible::build_chat_request;
286    use crate::types::IsoString;
287    use crate::{Context, ModelSettings, ToolCallRequest, ToolDefinition};
288
289    fn assistant_with_tool_call(reasoning_content: Option<&str>) -> ChatMessage {
290        ChatMessage::Assistant {
291            content: String::new(),
292            reasoning: crate::AssistantReasoning {
293                summary_text: reasoning_content.map(ToString::to_string),
294                encrypted_content: None,
295            },
296            timestamp: IsoString::now(),
297            tool_calls: vec![ToolCallRequest {
298                id: "call_1".to_string(),
299                name: "test__tool".to_string(),
300                arguments: "{\"path\":\"src/main.rs\"}".to_string(),
301            }],
302        }
303    }
304
305    fn context_with_assistant_message(message: ChatMessage) -> crate::Context {
306        crate::Context::new(
307            vec![ChatMessage::user("run a tool"), message],
308            vec![ToolDefinition::new("test__tool", "test", serde_json::json!({ "type": "object" }))],
309        )
310    }
311
312    #[test]
313    fn test_build_request_includes_reasoning_content_on_assistant_tool_message() {
314        let context = context_with_assistant_message(assistant_with_tool_call(Some("trace chunk")));
315        let request = build_chat_request("test-model", &context, None).unwrap();
316
317        let json = serde_json::to_value(&request).unwrap();
318        assert_eq!(json["messages"][1]["role"], "assistant");
319        assert_eq!(json["messages"][1]["reasoning_content"], "trace chunk");
320    }
321
322    #[test]
323    fn test_build_request_maps_model_settings_and_omits_when_unset() {
324        let user = || ChatMessage::user("hello");
325
326        let mut context = Context::new(vec![user()], vec![]);
327        context.set_model_settings(ModelSettings { temperature: Some(0.0), top_p: Some(0.5), max_tokens: Some(64) });
328        let json = serde_json::to_value(build_chat_request("test-model", &context, None).unwrap()).unwrap();
329        assert_eq!(json["temperature"], 0.0);
330        assert_eq!(json["top_p"], 0.5);
331        assert_eq!(json["max_tokens"], 64);
332
333        let unset = Context::new(vec![user()], vec![]);
334        let json = serde_json::to_value(build_chat_request("test-model", &unset, None).unwrap()).unwrap();
335        assert!(json.get("temperature").is_none());
336        assert!(json.get("top_p").is_none());
337        assert!(json.get("max_tokens").is_none());
338    }
339
340    #[test]
341    fn test_build_request_includes_stream_options_with_usage() {
342        let context = crate::Context::new(vec![ChatMessage::user("hello")], vec![]);
343        let request = build_chat_request("test-model", &context, None).unwrap();
344
345        let json = serde_json::to_value(&request).unwrap();
346        assert_eq!(json["stream_options"]["include_usage"], true);
347    }
348
349    #[test]
350    fn test_build_request_sends_empty_reasoning_content_on_tool_call_when_none() {
351        let context = context_with_assistant_message(assistant_with_tool_call(None));
352        let request = build_chat_request("test-model", &context, None).unwrap();
353
354        let json = serde_json::to_value(&request).unwrap();
355        assert_eq!(json["messages"][1]["role"], "assistant");
356        assert_eq!(json["messages"][1]["reasoning_content"], ".");
357    }
358
359    #[test]
360    fn test_user_message_text_only_serializes_as_string() {
361        let content = map_user_content(&[ContentBlock::text("Hello")]).unwrap();
362        let json = serde_json::to_value(&content).unwrap();
363        assert_eq!(json, "Hello");
364    }
365
366    #[test]
367    fn test_user_message_with_image_serializes_as_array() {
368        let content = map_user_content(&[
369            ContentBlock::text("Look:"),
370            ContentBlock::Image { data: "aW1n".to_string(), mime_type: "image/png".to_string() },
371        ])
372        .unwrap();
373        let json = serde_json::to_value(&content).unwrap();
374        let parts = json.as_array().expect("Expected array");
375        assert_eq!(parts.len(), 2);
376        assert_eq!(parts[0]["type"], "text");
377        assert_eq!(parts[0]["text"], "Look:");
378        assert_eq!(parts[1]["type"], "image_url");
379        assert!(parts[1]["image_url"]["url"].as_str().unwrap().starts_with("data:image/png;base64,"));
380    }
381
382    #[test]
383    fn test_user_message_audio_only_errors() {
384        let result = map_user_content(&[ContentBlock::Audio {
385            data: "YXVkaW8=".to_string(),
386            mime_type: "audio/wav".to_string(),
387        }]);
388        assert!(matches!(result, Err(crate::LlmError::UnsupportedContent(_))));
389    }
390
391    #[test]
392    fn test_user_message_audio_with_text_errors() {
393        let result = map_user_content(&[
394            ContentBlock::text("Listen:"),
395            ContentBlock::Audio { data: "YXVkaW8=".to_string(), mime_type: "audio/wav".to_string() },
396        ]);
397        assert!(matches!(result, Err(crate::LlmError::UnsupportedContent(_))));
398    }
399}