Skip to main content

embacle_server/
completions.rs

1// ABOUTME: POST /v1/chat/completions handler for OpenAI-compatible chat completion
2// ABOUTME: Routes to single provider or multiplex, supports both streaming and non-streaming
3//
4// SPDX-License-Identifier: Apache-2.0
5// Copyright (c) 2026 dravr.ai
6
7use std::env;
8use std::sync::atomic::{AtomicU64, Ordering};
9use std::time::{SystemTime, UNIX_EPOCH};
10
11use axum::extract::State;
12use axum::http::StatusCode;
13use axum::response::{IntoResponse, Response};
14use axum::Json;
15use embacle::config::CliRunnerType;
16use embacle::types::{
17    ChatMessage, ChatRequest, ErrorKind, LlmCapabilities, LlmProvider, MessageRole, ResponseFormat,
18    RunnerError,
19};
20use embacle::FunctionDeclaration;
21use tracing::{debug, error, warn};
22
23use crate::openai_types::{
24    ChatCompletionMessage, ChatCompletionRequest, ChatCompletionResponse, Choice, ContentPart,
25    ErrorResponse, MessageContent, ModelField, MultiplexProviderResult, MultiplexResponse,
26    ResponseFormatRequest, ResponseMessage, StopField, ToolCall, ToolCallFunction, ToolChoice,
27    ToolDefinition as OpenAiToolDefinition, Usage,
28};
29use crate::provider_resolver::resolve_model;
30use crate::runner::multiplex::{MultiplexEngine, MultiplexParams};
31use crate::state::SharedState;
32use crate::streaming;
33
34/// OpenAI-specified upper bound for temperature
35const MAX_TEMPERATURE: f32 = 2.0;
36
37/// Handle POST /v1/chat/completions
38///
39/// Dispatches to single-provider or multiplex mode based on the model field.
40/// Supports both streaming (SSE) and non-streaming (JSON) responses.
41pub async fn handle(
42    State(state): State<SharedState>,
43    Json(request): Json<ChatCompletionRequest>,
44) -> Response {
45    if let Some(temp) = request.temperature {
46        if !(0.0..=MAX_TEMPERATURE).contains(&temp) {
47            return error_response(
48                StatusCode::BAD_REQUEST,
49                &format!("temperature must be between 0.0 and {MAX_TEMPERATURE}"),
50            );
51        }
52    }
53    if let Some(max) = request.max_tokens {
54        if max == 0 {
55            return error_response(StatusCode::BAD_REQUEST, "max_tokens must be greater than 0");
56        }
57    }
58    if let Some(top_p) = request.top_p {
59        if !(0.0..=1.0).contains(&top_p) {
60            return error_response(StatusCode::BAD_REQUEST, "top_p must be between 0.0 and 1.0");
61        }
62    }
63    if let Some(ref stop) = request.stop {
64        if stop.len() > 4 {
65            return error_response(
66                StatusCode::BAD_REQUEST,
67                "stop must have at most 4 sequences",
68            );
69        }
70    }
71
72    match request.model {
73        ModelField::Multiple(ref models) if models.len() > 1 => {
74            handle_multiplex(&state, &request, models).await
75        }
76        ModelField::Multiple(ref models) if models.len() == 1 => {
77            handle_single(&state, &request, &models[0]).await
78        }
79        ModelField::Multiple(_) => {
80            error_response(StatusCode::BAD_REQUEST, "Model array must not be empty")
81        }
82        ModelField::Single(ref model) => handle_single(&state, &request, model).await,
83    }
84}
85
86/// Handle a single-provider request (standard case)
87async fn handle_single(
88    state: &SharedState,
89    request: &ChatCompletionRequest,
90    model_str: &str,
91) -> Response {
92    let has_tools = request
93        .tools
94        .as_ref()
95        .is_some_and(|t| !t.is_empty() && !is_tool_choice_none(request.tool_choice.as_ref()));
96
97    let state_guard = state.read().await;
98    let resolved = resolve_model(model_str, state_guard.active_provider());
99    debug!(
100        provider = %resolved.runner_type,
101        model = ?resolved.model,
102        stream = request.stream,
103        has_tools,
104        "Dispatching completion"
105    );
106
107    let runner = match state_guard.get_runner(resolved.runner_type).await {
108        Ok(r) => r,
109        Err(e) => return runner_error_to_response(&e),
110    };
111    drop(state_guard);
112
113    let strict = request
114        .strict_capabilities
115        .unwrap_or_else(|| env::var("EMBACLE_STRICT_CAPS").is_ok_and(|v| v == "true" || v == "1"));
116
117    let mut messages = convert_messages(&request.messages);
118
119    // Inject tool catalog using the most effective strategy for this provider
120    if has_tools {
121        let declarations = tools_to_declarations(request.tools.as_deref().unwrap_or_default());
122        let catalog = embacle::generate_tool_catalog(&declarations);
123
124        if runner
125            .capabilities()
126            .contains(LlmCapabilities::SYSTEM_MESSAGES)
127        {
128            embacle::inject_tool_catalog(&mut messages, &catalog);
129        } else {
130            inject_tool_catalog_as_user_message(&mut messages, &catalog);
131        }
132    }
133
134    let mut chat_request = ChatRequest::new(messages);
135    chat_request.model = resolved.model;
136    chat_request.temperature = request.temperature;
137    chat_request.max_tokens = request.max_tokens;
138    chat_request.top_p = request.top_p;
139    chat_request.stop = request.stop.as_ref().map(StopField::to_bounded_vec);
140    chat_request.response_format = request.response_format.as_ref().map(server_format_to_core);
141    chat_request.tools = request
142        .tools
143        .as_ref()
144        .map(|tools| tools.iter().map(server_tool_to_core).collect());
145    chat_request.tool_choice = request.tool_choice.as_ref().map(server_choice_to_core);
146
147    let warnings = match embacle::validate_capabilities(
148        runner.name(),
149        runner.capabilities(),
150        &chat_request,
151        strict,
152    ) {
153        Ok(w) => w,
154        Err(e) => return runner_error_to_response(&e),
155    };
156    let warnings_for_response = if warnings.is_empty() {
157        None
158    } else {
159        Some(warnings)
160    };
161
162    let supports_streaming = runner.capabilities().contains(LlmCapabilities::STREAMING);
163
164    dispatch_completion(
165        runner.as_ref(),
166        resolved.runner_type,
167        chat_request,
168        request.stream,
169        has_tools,
170        supports_streaming,
171        warnings_for_response,
172    )
173    .await
174}
175
176/// Check if the response format requests JSON output
177fn wants_json(format: Option<&ResponseFormat>) -> bool {
178    matches!(
179        format,
180        Some(ResponseFormat::JsonObject | ResponseFormat::JsonSchema { .. })
181    )
182}
183
184/// Strip markdown code fences from content when JSON output is requested.
185///
186/// CLI runners often wrap JSON in `` ```json ... ``` `` fences. When the client
187/// has requested `response_format: json_object` or `json_schema`, extract the
188/// raw JSON so the response is directly parseable.
189fn strip_json_fences(content: String, json_mode: bool) -> String {
190    if json_mode {
191        embacle::extract_json_from_response(&content)
192    } else {
193        content
194    }
195}
196
197/// Dispatch the completion request to the appropriate execution path
198///
199/// Routes between four modes:
200/// 1. Streaming with tools: downgrade to `complete()`, emit as SSE
201/// 2. Streaming without provider support: downgrade to `complete()`, emit as SSE
202/// 3. Pure streaming: use `complete_stream()`
203/// 4. Non-streaming: use `complete()`, return JSON
204async fn dispatch_completion(
205    runner: &dyn LlmProvider,
206    runner_type: CliRunnerType,
207    mut chat_request: ChatRequest,
208    stream: bool,
209    has_tools: bool,
210    supports_streaming: bool,
211    warnings: Option<Vec<String>>,
212) -> Response {
213    let json_mode = wants_json(chat_request.response_format.as_ref());
214
215    if stream && (has_tools || !supports_streaming) {
216        // Downgrade to non-streaming complete(), emit result as SSE
217        if has_tools {
218            debug!("Downgrading stream+tools to non-streaming complete");
219        } else {
220            debug!(
221                provider = runner.name(),
222                "Provider does not support streaming; downgrading to non-streaming complete"
223            );
224        }
225        match runner.complete(&chat_request).await {
226            Ok(response) => {
227                let model_name = format!("{runner_type}:{}", response.model);
228                let content = strip_json_fences(response.content, json_mode);
229                let (message, finish_reason) = build_response_message(
230                    has_tools,
231                    content,
232                    response.finish_reason,
233                    response.tool_calls.as_ref(),
234                );
235                let reason = finish_reason.as_deref().unwrap_or("stop");
236                streaming::sse_single_response(message, reason, &model_name)
237            }
238            Err(e) => runner_error_to_response(&e),
239        }
240    } else if stream {
241        chat_request.stream = true;
242        match runner.complete_stream(&chat_request).await {
243            Ok(s) => {
244                let model_name = format!("{runner_type}:{}", runner.default_model());
245                if json_mode {
246                    streaming::sse_response_strip_fences(s, &model_name)
247                } else {
248                    streaming::sse_response(s, &model_name)
249                }
250            }
251            Err(e) => runner_error_to_response(&e),
252        }
253    } else {
254        match runner.complete(&chat_request).await {
255            Ok(response) => {
256                let model_name = format!("{runner_type}:{}", response.model);
257                let usage = response.usage.map(|u| Usage {
258                    prompt: u.prompt_tokens,
259                    completion: u.completion_tokens,
260                    total: u.total_tokens,
261                });
262
263                let content = strip_json_fences(response.content, json_mode);
264                let (message, finish_reason) = build_response_message(
265                    has_tools,
266                    content,
267                    response.finish_reason,
268                    response.tool_calls.as_ref(),
269                );
270
271                let resp = ChatCompletionResponse {
272                    id: generate_id(),
273                    object: "chat.completion",
274                    created: unix_timestamp(),
275                    model: model_name,
276                    choices: vec![Choice {
277                        index: 0,
278                        message,
279                        finish_reason,
280                    }],
281                    usage,
282                    warnings,
283                };
284
285                (StatusCode::OK, Json(resp)).into_response()
286            }
287            Err(e) => runner_error_to_response(&e),
288        }
289    }
290}
291
292/// Handle a multiplex request (multiple providers)
293async fn handle_multiplex(
294    state: &SharedState,
295    request: &ChatCompletionRequest,
296    models: &[String],
297) -> Response {
298    if request.stream {
299        return error_response(
300            StatusCode::BAD_REQUEST,
301            "Streaming is not supported for multiplex requests",
302        );
303    }
304
305    let strict = request
306        .strict_capabilities
307        .unwrap_or_else(|| env::var("EMBACLE_STRICT_CAPS").is_ok_and(|v| v == "true" || v == "1"));
308
309    let state_guard = state.read().await;
310    let default_provider = state_guard.active_provider();
311    let resolved: Vec<_> = models
312        .iter()
313        .map(|m| resolve_model(m, default_provider))
314        .collect();
315
316    let providers: Vec<_> = resolved.iter().map(|r| r.runner_type).collect();
317    let messages = convert_messages(&request.messages);
318
319    // Build a temporary ChatRequest for capability validation
320    let mut validation_request = ChatRequest::new(messages.clone());
321    validation_request.temperature = request.temperature;
322    validation_request.max_tokens = request.max_tokens;
323    validation_request.top_p = request.top_p;
324    validation_request.stop = request.stop.as_ref().map(StopField::to_bounded_vec);
325    validation_request.response_format =
326        request.response_format.as_ref().map(server_format_to_core);
327
328    for &provider_type in &providers {
329        let runner = match state_guard.get_runner(provider_type).await {
330            Ok(r) => r,
331            Err(e) => return runner_error_to_response(&e),
332        };
333        match embacle::validate_capabilities(
334            runner.name(),
335            runner.capabilities(),
336            &validation_request,
337            strict,
338        ) {
339            Ok(w) => {
340                for warning in &w {
341                    warn!(provider = runner.name(), warning = %warning, "Capability warning");
342                }
343            }
344            Err(e) => return runner_error_to_response(&e),
345        }
346    }
347
348    drop(state_guard);
349    let engine = MultiplexEngine::new(state);
350    let params = MultiplexParams {
351        temperature: request.temperature,
352        max_tokens: request.max_tokens,
353        top_p: request.top_p,
354        stop: request.stop.as_ref().map(StopField::to_bounded_vec),
355        response_format: request.response_format.as_ref().map(server_format_to_core),
356    };
357    match engine.execute(&messages, &providers, &params).await {
358        Ok(result) => {
359            let results = result
360                .responses
361                .into_iter()
362                .map(|r| MultiplexProviderResult {
363                    provider: r.provider,
364                    model: r.model,
365                    content: r.content,
366                    error: r.error,
367                    duration_ms: r.duration_ms,
368                })
369                .collect();
370
371            let resp = MultiplexResponse {
372                id: generate_id(),
373                object: "chat.completion.multiplex",
374                created: unix_timestamp(),
375                results,
376                summary: result.summary,
377            };
378
379            (StatusCode::OK, Json(resp)).into_response()
380        }
381        Err(e) => runner_error_to_response(&e),
382    }
383}
384
385/// Build a `ResponseMessage` from LLM output, using native tool calls if available
386/// or falling back to XML parsing if tools were requested
387fn build_response_message(
388    has_tools: bool,
389    content: String,
390    finish_reason: Option<String>,
391    native_tool_calls: Option<&Vec<embacle::ToolCallRequest>>,
392) -> (ResponseMessage, Option<String>) {
393    // If the provider returned native tool calls, use them directly
394    if let Some(calls) = native_tool_calls {
395        if !calls.is_empty() {
396            let tool_calls: Vec<ToolCall> = calls
397                .iter()
398                .enumerate()
399                .map(|(i, tc)| ToolCall {
400                    index: i,
401                    id: tc.id.clone(),
402                    tool_type: "function".to_owned(),
403                    function: ToolCallFunction {
404                        name: tc.function_name.clone(),
405                        arguments: serde_json::to_string(&tc.arguments)
406                            .unwrap_or_else(|_| "{}".to_owned()),
407                    },
408                })
409                .collect();
410            let text_content = if content.is_empty() {
411                None
412            } else {
413                Some(content)
414            };
415            return (
416                ResponseMessage {
417                    role: "assistant",
418                    content: text_content,
419                    tool_calls: Some(tool_calls),
420                },
421                Some("tool_calls".to_owned()),
422            );
423        }
424    }
425
426    // Fall back to XML parsing for text-based tool simulation
427    if has_tools {
428        let parsed_calls = embacle::parse_tool_call_blocks(&content);
429        if parsed_calls.is_empty() {
430            (
431                ResponseMessage {
432                    role: "assistant",
433                    content: Some(content),
434                    tool_calls: None,
435                },
436                finish_reason.or_else(|| Some("stop".to_owned())),
437            )
438        } else {
439            let remaining_text = embacle::strip_tool_call_blocks(&content);
440            let text_content = if remaining_text.is_empty() {
441                None
442            } else {
443                Some(remaining_text)
444            };
445            let tool_calls: Vec<ToolCall> = parsed_calls
446                .iter()
447                .enumerate()
448                .map(|(i, fc)| ToolCall {
449                    index: i,
450                    id: generate_tool_call_id(&fc.name, i),
451                    tool_type: "function".to_owned(),
452                    function: ToolCallFunction {
453                        name: fc.name.clone(),
454                        arguments: serde_json::to_string(&fc.args)
455                            .unwrap_or_else(|_| "{}".to_owned()),
456                    },
457                })
458                .collect();
459            (
460                ResponseMessage {
461                    role: "assistant",
462                    content: text_content,
463                    tool_calls: Some(tool_calls),
464                },
465                Some("tool_calls".to_owned()),
466            )
467        }
468    } else {
469        (
470            ResponseMessage {
471                role: "assistant",
472                content: Some(content),
473                tool_calls: None,
474            },
475            finish_reason.or_else(|| Some("stop".to_owned())),
476        )
477    }
478}
479
480/// Extract text content from a `MessageContent`, returning an empty string for None
481fn content_as_text(content: Option<&MessageContent>) -> String {
482    content.map(MessageContent::as_text).unwrap_or_default()
483}
484
485/// Parse a `data:` URI into an `ImagePart`
486///
487/// Expected format: `data:<mime_type>;base64,<data>`
488fn parse_data_uri(url: &str) -> Option<embacle::ImagePart> {
489    let rest = url.strip_prefix("data:")?;
490    let (mime_type, data) = rest.split_once(";base64,")?;
491    embacle::ImagePart::new(data, mime_type).ok()
492}
493
494/// Extract images from a `MessageContent::Parts` variant
495fn extract_images(content: Option<&MessageContent>) -> Option<Vec<embacle::ImagePart>> {
496    let Some(MessageContent::Parts(parts)) = content else {
497        return None;
498    };
499
500    let images: Vec<embacle::ImagePart> = parts
501        .iter()
502        .filter_map(|p| match p {
503            ContentPart::ImageUrl { image_url } => parse_data_uri(&image_url.url),
504            ContentPart::Text { .. } => None,
505        })
506        .collect();
507
508    if images.is_empty() {
509        None
510    } else {
511        Some(images)
512    }
513}
514
515/// Convert `OpenAI` message format to embacle `ChatMessage`
516///
517/// Handles all `OpenAI` roles including "tool" messages and assistant messages
518/// with `tool_calls`. Tool messages are collected and formatted as `<tool_result>`
519/// blocks. Assistant messages with `tool_calls` are reconstructed as `<tool_call>` blocks.
520/// User messages with multipart content (text + images) are converted to `ChatMessage`
521/// with attached `ImagePart` entries.
522fn convert_messages(messages: &[ChatCompletionMessage]) -> Vec<ChatMessage> {
523    let mut result = Vec::with_capacity(messages.len());
524    let mut i = 0;
525
526    while i < messages.len() {
527        let m = &messages[i];
528        match m.role.as_str() {
529            "system" => {
530                result.push(ChatMessage::system(content_as_text(m.content.as_ref())));
531                i += 1;
532            }
533            "user" => {
534                let text = content_as_text(m.content.as_ref());
535                let images = extract_images(m.content.as_ref());
536                if let Some(imgs) = images {
537                    result.push(ChatMessage::user_with_images(text, imgs));
538                } else {
539                    result.push(ChatMessage::user(text));
540                }
541                i += 1;
542            }
543            "assistant" => {
544                if let Some(ref tool_calls) = m.tool_calls {
545                    // Reconstruct <tool_call> XML blocks from stored tool calls
546                    let mut text = content_as_text(m.content.as_ref());
547                    for tc in tool_calls {
548                        text.push_str("\n<tool_call>\n");
549                        let payload = serde_json::json!({
550                            "name": tc.function.name,
551                            "arguments": serde_json::from_str::<serde_json::Value>(&tc.function.arguments)
552                                .unwrap_or_else(|_| serde_json::Value::Object(serde_json::Map::new()))
553                        });
554                        text.push_str(
555                            &serde_json::to_string(&payload).unwrap_or_else(|_| "{}".to_owned()),
556                        );
557                        text.push_str("\n</tool_call>");
558                    }
559                    result.push(ChatMessage::assistant(text));
560                } else {
561                    result.push(ChatMessage::assistant(content_as_text(m.content.as_ref())));
562                }
563                i += 1;
564            }
565            "tool" => {
566                // Collect consecutive tool messages into a single user message
567                let mut tool_responses = Vec::new();
568                while i < messages.len() && messages[i].role == "tool" {
569                    let tool_msg = &messages[i];
570                    let name = tool_msg.name.as_deref().unwrap_or("unknown");
571                    let content_text = content_as_text(tool_msg.content.as_ref());
572                    let response_value: serde_json::Value = if content_text.is_empty() {
573                        serde_json::Value::Null
574                    } else {
575                        serde_json::from_str(&content_text)
576                            .unwrap_or(serde_json::Value::String(content_text))
577                    };
578                    tool_responses.push(embacle::FunctionResponse {
579                        name: name.to_owned(),
580                        response: response_value,
581                    });
582                    i += 1;
583                }
584                let text = embacle::format_tool_results_as_text(&tool_responses);
585                result.push(ChatMessage::user(text));
586            }
587            other => {
588                warn!(role = other, "Unknown message role, mapping to user");
589                result.push(ChatMessage::user(content_as_text(m.content.as_ref())));
590                i += 1;
591            }
592        }
593    }
594
595    result
596}
597
598/// Convert a server `ToolDefinition` to core `ToolDefinition`
599fn server_tool_to_core(tool: &OpenAiToolDefinition) -> embacle::ToolDefinition {
600    embacle::ToolDefinition {
601        name: tool.function.name.clone(),
602        description: tool.function.description.clone().unwrap_or_default(),
603        parameters: tool.function.parameters.clone(),
604    }
605}
606
607/// Convert a server `ToolChoice` to core `ToolChoice`
608fn server_choice_to_core(choice: &ToolChoice) -> embacle::ToolChoice {
609    match choice {
610        ToolChoice::Mode(m) => match m.as_str() {
611            "none" => embacle::ToolChoice::None,
612            "required" => embacle::ToolChoice::Required,
613            _ => embacle::ToolChoice::Auto,
614        },
615        ToolChoice::Specific(s) => embacle::ToolChoice::Specific {
616            name: s.function.name.clone(),
617        },
618    }
619}
620
621/// Convert a server `ResponseFormatRequest` to core `ResponseFormat`
622fn server_format_to_core(format: &ResponseFormatRequest) -> embacle::ResponseFormat {
623    match format {
624        ResponseFormatRequest::Text => embacle::ResponseFormat::Text,
625        ResponseFormatRequest::JsonObject => embacle::ResponseFormat::JsonObject,
626        ResponseFormatRequest::JsonSchema { json_schema } => embacle::ResponseFormat::JsonSchema {
627            name: json_schema.name.clone(),
628            schema: json_schema.schema.clone(),
629        },
630    }
631}
632
633/// Convert `OpenAI` tool definitions to embacle `FunctionDeclaration` format
634fn tools_to_declarations(tools: &[OpenAiToolDefinition]) -> Vec<FunctionDeclaration> {
635    tools
636        .iter()
637        .map(|t| FunctionDeclaration {
638            name: t.function.name.clone(),
639            description: t.function.description.clone().unwrap_or_default(),
640            parameters: t.function.parameters.clone(),
641        })
642        .collect()
643}
644
645/// Inject tool catalog into the last user message content
646///
647/// Used for providers that do not support system messages (e.g. Copilot CLI).
648/// The catalog is prepended to the last user message so the LLM sees it in
649/// the conversational flow rather than in a system prompt it cannot parse.
650fn inject_tool_catalog_as_user_message(messages: &mut [ChatMessage], catalog: &str) {
651    if let Some(last_user) = messages
652        .iter_mut()
653        .rev()
654        .find(|m| m.role == MessageRole::User)
655    {
656        let augmented = format!("{catalog}\n\n{}", last_user.content);
657        *last_user = ChatMessage::user(augmented);
658    } else {
659        // No user message found; this shouldn't happen in practice but
660        // handle gracefully by appending a user message with the catalog.
661        warn!("No user message found for tool catalog injection");
662    }
663}
664
665/// Check if `tool_choice` is explicitly "none"
666fn is_tool_choice_none(tool_choice: Option<&ToolChoice>) -> bool {
667    matches!(tool_choice, Some(ToolChoice::Mode(ref m)) if m == "none")
668}
669
670/// Generate a deterministic tool call ID from function name and index
671fn generate_tool_call_id(name: &str, index: usize) -> String {
672    format!("call_{name}_{index}")
673}
674
675/// Map a `RunnerError` to an appropriate HTTP status code and `OpenAI` error response
676fn runner_error_to_response(err: &RunnerError) -> Response {
677    let (status, error_type) = match err.kind {
678        ErrorKind::BinaryNotFound => (StatusCode::SERVICE_UNAVAILABLE, "provider_not_available"),
679        ErrorKind::AuthFailure => (StatusCode::UNAUTHORIZED, "authentication_error"),
680        ErrorKind::Timeout => (StatusCode::GATEWAY_TIMEOUT, "timeout_error"),
681        ErrorKind::ExternalService => (StatusCode::BAD_GATEWAY, "external_service_error"),
682        ErrorKind::Config => (StatusCode::BAD_REQUEST, "invalid_request_error"),
683        ErrorKind::Guardrail => (StatusCode::BAD_REQUEST, "guardrail_error"),
684        ErrorKind::ModelUnavailable => (StatusCode::NOT_FOUND, "model_not_found"),
685        ErrorKind::Internal => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
686    };
687
688    error!(kind = ?err.kind, message = %err.message, "Runner error");
689    let body = ErrorResponse::new(error_type, &err.message);
690    (status, Json(body)).into_response()
691}
692
693/// Build an error response with a given status and message
694fn error_response(status: StatusCode, message: &str) -> Response {
695    let body = ErrorResponse::new("invalid_request_error", message);
696    (status, Json(body)).into_response()
697}
698
699/// Monotonic counter ensuring unique IDs even for requests within the same second
700static ID_COUNTER: AtomicU64 = AtomicU64::new(0);
701
702/// Generate a unique completion ID
703///
704/// Combines the unix timestamp with a monotonically increasing counter
705/// to guarantee uniqueness across concurrent and rapid-fire requests.
706pub fn generate_id() -> String {
707    let ts = unix_timestamp();
708    let seq = ID_COUNTER.fetch_add(1, Ordering::Relaxed);
709    format!("chatcmpl-{ts:x}{seq:08x}")
710}
711
712/// Get current unix timestamp in seconds
713pub fn unix_timestamp() -> u64 {
714    SystemTime::now()
715        .duration_since(UNIX_EPOCH)
716        .map_or(0, |d| d.as_secs())
717}
718
719#[cfg(test)]
720mod tests {
721    use super::*;
722    use crate::openai_types::{
723        ContentPart, FunctionObject, ImageUrlDetail, ToolCall, ToolCallFunction, ToolDefinition,
724    };
725    use MessageRole;
726
727    /// Helper to create a `ChatCompletionMessage` with plain text content
728    fn text_msg(role: &str, content: Option<&str>) -> ChatCompletionMessage {
729        ChatCompletionMessage {
730            role: role.to_owned(),
731            content: content.map(|c| MessageContent::Text(c.to_owned())),
732            tool_calls: None,
733            tool_call_id: None,
734            name: None,
735        }
736    }
737
738    #[test]
739    fn convert_messages_maps_roles() {
740        let openai_msgs = vec![
741            text_msg("system", Some("You are helpful")),
742            text_msg("user", Some("Hello")),
743            text_msg("assistant", Some("Hi there")),
744        ];
745
746        let messages = convert_messages(&openai_msgs);
747        assert_eq!(messages.len(), 3);
748        assert_eq!(messages[0].role, MessageRole::System);
749        assert_eq!(messages[1].role, MessageRole::User);
750        assert_eq!(messages[2].role, MessageRole::Assistant);
751    }
752
753    #[test]
754    fn convert_unknown_role_defaults_to_user() {
755        let openai_msgs = vec![text_msg("function", Some("result"))];
756
757        let messages = convert_messages(&openai_msgs);
758        assert_eq!(messages[0].role, MessageRole::User);
759    }
760
761    #[test]
762    fn convert_assistant_with_tool_calls() {
763        let openai_msgs = vec![ChatCompletionMessage {
764            role: "assistant".to_owned(),
765            content: None,
766            tool_calls: Some(vec![ToolCall {
767                index: 0,
768                id: "call_1".to_owned(),
769                tool_type: "function".to_owned(),
770                function: ToolCallFunction {
771                    name: "get_weather".to_owned(),
772                    arguments: r#"{"city":"Paris"}"#.to_owned(),
773                },
774            }]),
775            tool_call_id: None,
776            name: None,
777        }];
778
779        let messages = convert_messages(&openai_msgs);
780        assert_eq!(messages.len(), 1);
781        assert_eq!(messages[0].role, MessageRole::Assistant);
782        assert!(messages[0].content.contains("<tool_call>"));
783        assert!(messages[0].content.contains("get_weather"));
784        assert!(messages[0].content.contains("</tool_call>"));
785    }
786
787    #[test]
788    fn convert_tool_messages_to_user() {
789        let openai_msgs = vec![
790            ChatCompletionMessage {
791                role: "tool".to_owned(),
792                content: Some(MessageContent::Text(r#"{"temp":72}"#.to_owned())),
793                tool_calls: None,
794                tool_call_id: Some("call_1".to_owned()),
795                name: Some("get_weather".to_owned()),
796            },
797            ChatCompletionMessage {
798                role: "tool".to_owned(),
799                content: Some(MessageContent::Text(r#"{"time":"14:30"}"#.to_owned())),
800                tool_calls: None,
801                tool_call_id: Some("call_2".to_owned()),
802                name: Some("get_time".to_owned()),
803            },
804        ];
805
806        let messages = convert_messages(&openai_msgs);
807        // Consecutive tool messages should be merged into one user message
808        assert_eq!(messages.len(), 1);
809        assert_eq!(messages[0].role, MessageRole::User);
810        assert!(messages[0].content.contains("tool_result"));
811        assert!(messages[0].content.contains("get_weather"));
812        assert!(messages[0].content.contains("get_time"));
813    }
814
815    #[test]
816    fn convert_messages_none_content() {
817        let openai_msgs = vec![text_msg("user", None)];
818
819        let messages = convert_messages(&openai_msgs);
820        assert_eq!(messages[0].content, "");
821    }
822
823    #[test]
824    fn convert_multipart_user_message_extracts_images() {
825        let openai_msgs = vec![ChatCompletionMessage {
826            role: "user".to_owned(),
827            content: Some(MessageContent::Parts(vec![
828                ContentPart::Text {
829                    text: "What is this?".to_owned(),
830                },
831                ContentPart::ImageUrl {
832                    image_url: ImageUrlDetail {
833                        url: "data:image/png;base64,aGVsbG8=".to_owned(),
834                    },
835                },
836            ])),
837            tool_calls: None,
838            tool_call_id: None,
839            name: None,
840        }];
841
842        let messages = convert_messages(&openai_msgs);
843        assert_eq!(messages.len(), 1);
844        assert_eq!(messages[0].content, "What is this?");
845        let images = messages[0].images.as_ref().expect("images present"); // Safe: test assertion
846        assert_eq!(images.len(), 1);
847        assert_eq!(images[0].mime_type, "image/png");
848        assert_eq!(images[0].data, "aGVsbG8=");
849    }
850
851    #[test]
852    fn parse_data_uri_valid() {
853        let img = parse_data_uri("data:image/jpeg;base64,AAAA").expect("should parse"); // Safe: test assertion
854        assert_eq!(img.mime_type, "image/jpeg");
855        assert_eq!(img.data, "AAAA");
856    }
857
858    #[test]
859    fn parse_data_uri_invalid_format() {
860        assert!(parse_data_uri("https://example.com/image.png").is_none());
861        assert!(parse_data_uri("data:text/plain;base64,abc").is_none());
862        assert!(parse_data_uri("data:image/png;abc").is_none());
863    }
864
865    #[test]
866    fn convert_plain_string_content_backward_compat() {
867        let openai_msgs = vec![text_msg("user", Some("hello"))];
868        let messages = convert_messages(&openai_msgs);
869        assert_eq!(messages[0].content, "hello");
870        assert!(messages[0].images.is_none());
871    }
872
873    #[test]
874    fn tools_to_declarations_converts() {
875        let tools = vec![ToolDefinition {
876            tool_type: "function".to_owned(),
877            function: FunctionObject {
878                name: "search".to_owned(),
879                description: Some("Search the web".to_owned()),
880                parameters: Some(serde_json::json!({
881                    "type": "object",
882                    "properties": {"q": {"type": "string"}},
883                    "required": ["q"]
884                })),
885            },
886        }];
887
888        let decls = tools_to_declarations(&tools);
889        assert_eq!(decls.len(), 1);
890        assert_eq!(decls[0].name, "search");
891        assert_eq!(decls[0].description, "Search the web");
892        assert!(decls[0].parameters.is_some());
893    }
894
895    #[test]
896    fn tool_choice_none_detection() {
897        let none_choice = ToolChoice::Mode("none".to_owned());
898        assert!(is_tool_choice_none(Some(&none_choice)));
899        let auto_choice = ToolChoice::Mode("auto".to_owned());
900        assert!(!is_tool_choice_none(Some(&auto_choice)));
901        assert!(!is_tool_choice_none(None));
902    }
903
904    #[test]
905    fn content_as_text_none() {
906        assert_eq!(content_as_text(None), "");
907    }
908
909    #[test]
910    fn content_as_text_plain() {
911        let content = MessageContent::Text("hello".to_owned());
912        assert_eq!(content_as_text(Some(&content)), "hello");
913    }
914
915    #[test]
916    fn generate_tool_call_id_format() {
917        let id = generate_tool_call_id("get_weather", 0);
918        assert_eq!(id, "call_get_weather_0");
919    }
920
921    #[test]
922    fn generate_id_has_prefix() {
923        let id = generate_id();
924        assert!(id.starts_with("chatcmpl-"));
925    }
926
927    #[test]
928    fn error_maps_binary_not_found_to_503() {
929        let err = RunnerError::binary_not_found("claude");
930        let (status, _) = match err.kind {
931            ErrorKind::BinaryNotFound => {
932                (StatusCode::SERVICE_UNAVAILABLE, "provider_not_available")
933            }
934            _ => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
935        };
936        assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE);
937    }
938
939    #[test]
940    fn error_maps_auth_to_401() {
941        let err = RunnerError::auth_failure("bad token");
942        let (status, _) = match err.kind {
943            ErrorKind::AuthFailure => (StatusCode::UNAUTHORIZED, "authentication_error"),
944            _ => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
945        };
946        assert_eq!(status, StatusCode::UNAUTHORIZED);
947    }
948
949    #[test]
950    fn error_maps_timeout_to_504() {
951        let err = RunnerError::timeout("too slow");
952        let (status, _) = match err.kind {
953            ErrorKind::Timeout => (StatusCode::GATEWAY_TIMEOUT, "timeout_error"),
954            _ => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
955        };
956        assert_eq!(status, StatusCode::GATEWAY_TIMEOUT);
957    }
958
959    #[test]
960    fn inject_tool_catalog_as_user_message_prepends_to_last_user() {
961        let mut messages = vec![
962            ChatMessage::user("First question"),
963            ChatMessage::assistant("Some answer"),
964            ChatMessage::user("What is the weather?"),
965        ];
966        let catalog = "## Available Tools\n- get_weather: Get the weather";
967
968        inject_tool_catalog_as_user_message(&mut messages, catalog);
969
970        assert_eq!(messages.len(), 3);
971        assert!(messages[2].content.starts_with("## Available Tools"));
972        assert!(messages[2].content.contains("What is the weather?"));
973        // First user message should be untouched
974        assert_eq!(messages[0].content, "First question");
975    }
976
977    #[test]
978    fn inject_tool_catalog_as_user_message_single_user() {
979        let mut messages = vec![
980            ChatMessage::system("You are helpful"),
981            ChatMessage::user("Hello"),
982        ];
983        let catalog = "## Tools\nsome tools";
984
985        inject_tool_catalog_as_user_message(&mut messages, catalog);
986
987        assert!(messages[1].content.starts_with("## Tools"));
988        assert!(messages[1].content.contains("Hello"));
989    }
990
991    #[test]
992    fn wants_json_matches_json_formats() {
993        use embacle::types::ResponseFormat;
994
995        assert!(!wants_json(None));
996        assert!(!wants_json(Some(&ResponseFormat::Text)));
997        assert!(wants_json(Some(&ResponseFormat::JsonObject)));
998        assert!(wants_json(Some(&ResponseFormat::JsonSchema {
999            name: "test".to_owned(),
1000            schema: serde_json::json!({}),
1001        })));
1002    }
1003
1004    #[test]
1005    fn strip_json_fences_removes_markdown_wrapper() {
1006        let fenced = "```json\n{\"key\":\"value\"}\n```".to_owned();
1007        assert_eq!(strip_json_fences(fenced, true), "{\"key\":\"value\"}");
1008    }
1009
1010    #[test]
1011    fn strip_json_fences_passes_through_in_text_mode() {
1012        let fenced = "```json\n{\"key\":\"value\"}\n```".to_owned();
1013        assert_eq!(strip_json_fences(fenced.clone(), false), fenced);
1014    }
1015
1016    #[test]
1017    fn strip_json_fences_leaves_clean_json_unchanged() {
1018        let clean = "{\"key\":\"value\"}".to_owned();
1019        assert_eq!(strip_json_fences(clean.clone(), true), clean);
1020    }
1021}