1use 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
34const MAX_TEMPERATURE: f32 = 2.0;
36
37pub 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
86async 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 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
176fn wants_json(format: Option<&ResponseFormat>) -> bool {
178 matches!(
179 format,
180 Some(ResponseFormat::JsonObject | ResponseFormat::JsonSchema { .. })
181 )
182}
183
184fn 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
197async 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 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
292async 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 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, ¶ms).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
385fn 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 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 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
480fn content_as_text(content: Option<&MessageContent>) -> String {
482 content.map(MessageContent::as_text).unwrap_or_default()
483}
484
485fn 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
494fn 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
515fn 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 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 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
598fn 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
607fn 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
621fn 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
633fn 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
645fn 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 warn!("No user message found for tool catalog injection");
662 }
663}
664
665fn is_tool_choice_none(tool_choice: Option<&ToolChoice>) -> bool {
667 matches!(tool_choice, Some(ToolChoice::Mode(ref m)) if m == "none")
668}
669
670fn generate_tool_call_id(name: &str, index: usize) -> String {
672 format!("call_{name}_{index}")
673}
674
675fn 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
693fn error_response(status: StatusCode, message: &str) -> Response {
695 let body = ErrorResponse::new("invalid_request_error", message);
696 (status, Json(body)).into_response()
697}
698
699static ID_COUNTER: AtomicU64 = AtomicU64::new(0);
701
702pub 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
712pub 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 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 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"); 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"); 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 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}