use std::collections::HashSet;
use std::env;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use axum::extract::State;
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::Json;
use embacle::config::CliRunnerType;
use embacle::mcp_tool_bridge::create_mcp_tool_handler;
use embacle::types::{
ChatMessage, ChatRequest, ErrorKind, LlmCapabilities, LlmProvider, MessageRole, ResponseFormat,
RunnerError,
};
use embacle::{AgentExecutor, FunctionDeclaration};
use tracing::{debug, error, warn};
use crate::openai_types::{
ChatCompletionMessage, ChatCompletionRequest, ChatCompletionResponse, Choice, ContentPart,
ErrorResponse, MessageContent, ModelField, MultiplexProviderResult, MultiplexResponse,
ResponseFormatRequest, ResponseMessage, StopField, ToolCall, ToolCallFunction, ToolChoice,
ToolDefinition as OpenAiToolDefinition, ToolExecutionMode, Usage,
};
use crate::provider_resolver::resolve_model;
use crate::runner::multiplex::{MultiplexEngine, MultiplexParams};
use crate::state::{AppState, SharedState};
use crate::streaming;
const MAX_TEMPERATURE: f32 = 2.0;
pub async fn handle(
State(app): State<AppState>,
Json(request): Json<ChatCompletionRequest>,
) -> Response {
if let Some(temp) = request.temperature {
if !(0.0..=MAX_TEMPERATURE).contains(&temp) {
return error_response(
StatusCode::BAD_REQUEST,
&format!("temperature must be between 0.0 and {MAX_TEMPERATURE}"),
);
}
}
if let Some(max) = request.max_tokens {
if max == 0 {
return error_response(StatusCode::BAD_REQUEST, "max_tokens must be greater than 0");
}
}
if let Some(top_p) = request.top_p {
if !(0.0..=1.0).contains(&top_p) {
return error_response(StatusCode::BAD_REQUEST, "top_p must be between 0.0 and 1.0");
}
}
if let Some(ref stop) = request.stop {
if stop.len() > 4 {
return error_response(
StatusCode::BAD_REQUEST,
"stop must have at most 4 sequences",
);
}
}
match request.model {
ModelField::Multiple(ref models) if models.len() > 1 => {
handle_multiplex(&app.shared, &request, models).await
}
ModelField::Multiple(ref models) if models.len() == 1 => {
handle_single(&app, &request, &models[0]).await
}
ModelField::Multiple(_) => {
error_response(StatusCode::BAD_REQUEST, "Model array must not be empty")
}
ModelField::Single(ref model) => handle_single(&app, &request, model).await,
}
}
async fn handle_single(
app: &AppState,
request: &ChatCompletionRequest,
model_str: &str,
) -> Response {
let state = &app.shared;
let has_tools = request
.tools
.as_ref()
.is_some_and(|t| !t.is_empty() && !is_tool_choice_none(request.tool_choice.as_ref()));
if request.tool_execution == Some(ToolExecutionMode::Server) {
return handle_server_side_tools(app, request, model_str).await;
}
let state_guard = state.read().await;
let resolved = resolve_model(model_str, state_guard.active_provider());
debug!(
provider = %resolved.runner_type,
model = ?resolved.model,
stream = request.stream,
has_tools,
"Dispatching completion"
);
let runner = match state_guard.get_runner(resolved.runner_type).await {
Ok(r) => r,
Err(e) => return runner_error_to_response(&e),
};
drop(state_guard);
let strict = request
.strict_capabilities
.unwrap_or_else(|| env::var("EMBACLE_STRICT_CAPS").is_ok_and(|v| v == "true" || v == "1"));
let mut messages = convert_messages(&request.messages);
if has_tools {
let declarations = tools_to_declarations(request.tools.as_deref().unwrap_or_default());
let catalog = embacle::generate_tool_catalog(&declarations);
if runner
.capabilities()
.contains(LlmCapabilities::SYSTEM_MESSAGES)
{
embacle::inject_tool_catalog(&mut messages, &catalog);
} else {
inject_tool_catalog_as_user_message(&mut messages, &catalog);
}
}
let mut chat_request = ChatRequest::new(messages);
chat_request.model = resolved.model;
chat_request.temperature = request.temperature;
chat_request.max_tokens = request.max_tokens;
chat_request.top_p = request.top_p;
chat_request.stop = request.stop.as_ref().map(StopField::to_bounded_vec);
chat_request.response_format = request.response_format.as_ref().map(server_format_to_core);
chat_request.tools = request
.tools
.as_ref()
.map(|tools| tools.iter().map(server_tool_to_core).collect());
chat_request.tool_choice = request.tool_choice.as_ref().map(server_choice_to_core);
let warnings = match embacle::validate_capabilities(
runner.name(),
runner.capabilities(),
&chat_request,
strict,
) {
Ok(w) => w,
Err(e) => return runner_error_to_response(&e),
};
let warnings_for_response = if warnings.is_empty() {
None
} else {
Some(warnings)
};
let supports_streaming = runner.capabilities().contains(LlmCapabilities::STREAMING);
let native_tools = runner
.capabilities()
.contains(LlmCapabilities::FUNCTION_CALLING);
let mode = if request.stream && has_tools && supports_streaming && !native_tools {
DispatchMode::StreamToolCalls
} else if request.stream && (has_tools || !supports_streaming) {
DispatchMode::StreamDowngrade
} else if request.stream {
DispatchMode::PureStream
} else {
DispatchMode::NonStreaming
};
dispatch_completion(
runner.as_ref(),
resolved.runner_type,
chat_request,
mode,
has_tools,
warnings_for_response,
)
.await
}
async fn handle_server_side_tools(
app: &AppState,
request: &ChatCompletionRequest,
model_str: &str,
) -> Response {
let Some(server_tools) = app.server_tools.as_ref() else {
return error_response(
StatusCode::BAD_REQUEST,
"tool_execution=server requires the server to be started with configured MCP tool servers ([[mcp_servers]])",
);
};
let state_guard = app.shared.read().await;
let resolved = resolve_model(model_str, state_guard.active_provider());
debug!(
provider = %resolved.runner_type,
model = ?resolved.model,
"Dispatching server-side tool execution"
);
let runner = match state_guard.get_runner(resolved.runner_type).await {
Ok(r) => r,
Err(e) => return runner_error_to_response(&e),
};
drop(state_guard);
let declarations =
select_server_declarations(&server_tools.declarations, request.tools.as_deref());
if declarations.is_empty() {
return error_response(
StatusCode::BAD_REQUEST,
"None of the requested tools are available on the server's configured MCP tool servers",
);
}
let messages = convert_messages(&request.messages);
let mut template = ChatRequest::new(Vec::new());
template.model = resolved.model;
template.temperature = request.temperature;
template.max_tokens = request.max_tokens;
template.top_p = request.top_p;
template.stop = request.stop.as_ref().map(StopField::to_bounded_vec);
let handler = create_mcp_tool_handler(Arc::clone(&server_tools.executor));
let agent =
AgentExecutor::new(runner.as_ref(), declarations, handler).with_request_template(template);
let result = match agent.run(messages).await {
Ok(r) => r,
Err(e) => return runner_error_to_response(&e),
};
let model_name = format!("{}:{}", resolved.runner_type, runner.default_model());
let usage = Some(Usage {
prompt: result.total_usage.prompt_tokens,
completion: result.total_usage.completion_tokens,
total: result.total_usage.total_tokens,
});
let finish_reason = result.finish_reason.unwrap_or_else(|| "stop".to_owned());
let warnings = if finish_reason == "max_turns" {
Some(vec![
"Agent reached the maximum number of tool-calling turns before completing".to_owned(),
])
} else {
None
};
let resp = ChatCompletionResponse {
id: generate_id(),
object: "chat.completion",
created: unix_timestamp(),
model: model_name,
choices: vec![Choice {
index: 0,
message: ResponseMessage {
role: "assistant",
content: Some(result.content),
tool_calls: None,
},
finish_reason: Some(finish_reason),
}],
usage,
warnings,
};
(StatusCode::OK, Json(resp)).into_response()
}
fn select_server_declarations(
available: &[FunctionDeclaration],
requested: Option<&[OpenAiToolDefinition]>,
) -> Vec<FunctionDeclaration> {
match requested {
Some(tools) if !tools.is_empty() => {
let names: HashSet<&str> = tools.iter().map(|t| t.function.name.as_str()).collect();
available
.iter()
.filter(|d| names.contains(d.name.as_str()))
.cloned()
.collect()
}
_ => available.to_vec(),
}
}
fn wants_json(format: Option<&ResponseFormat>) -> bool {
matches!(
format,
Some(ResponseFormat::JsonObject | ResponseFormat::JsonSchema { .. })
)
}
fn strip_json_fences(content: String, json_mode: bool) -> String {
if json_mode {
embacle::extract_json_from_response(&content)
} else {
content
}
}
enum DispatchMode {
StreamToolCalls,
StreamDowngrade,
PureStream,
NonStreaming,
}
async fn dispatch_completion(
runner: &dyn LlmProvider,
runner_type: CliRunnerType,
mut chat_request: ChatRequest,
mode: DispatchMode,
has_tools: bool,
warnings: Option<Vec<String>>,
) -> Response {
let json_mode = wants_json(chat_request.response_format.as_ref());
match mode {
DispatchMode::StreamToolCalls => {
chat_request.stream = true;
match runner.complete_stream(&chat_request).await {
Ok(s) => {
let model_name = format!("{runner_type}:{}", runner.default_model());
streaming::sse_response_with_tool_calls(s, &model_name)
}
Err(e) => runner_error_to_response(&e),
}
}
DispatchMode::StreamDowngrade => {
if has_tools {
debug!("Downgrading stream+tools to non-streaming complete");
} else {
debug!(
provider = runner.name(),
"Provider does not support streaming; downgrading to non-streaming complete"
);
}
match runner.complete(&chat_request).await {
Ok(response) => {
let model_name = format!("{runner_type}:{}", response.model);
let content = strip_json_fences(response.content, json_mode);
let (message, finish_reason) = build_response_message(
has_tools,
content,
response.finish_reason,
response.tool_calls.as_ref(),
);
let reason = finish_reason.as_deref().unwrap_or("stop");
streaming::sse_single_response(message, reason, &model_name)
}
Err(e) => runner_error_to_response(&e),
}
}
DispatchMode::PureStream => {
chat_request.stream = true;
match runner.complete_stream(&chat_request).await {
Ok(s) => {
let model_name = format!("{runner_type}:{}", runner.default_model());
if json_mode {
streaming::sse_response_strip_fences(s, &model_name)
} else {
streaming::sse_response(s, &model_name)
}
}
Err(e) => runner_error_to_response(&e),
}
}
DispatchMode::NonStreaming => match runner.complete(&chat_request).await {
Ok(response) => {
let model_name = format!("{runner_type}:{}", response.model);
let usage = response.usage.map(|u| Usage {
prompt: u.prompt_tokens,
completion: u.completion_tokens,
total: u.total_tokens,
});
let content = strip_json_fences(response.content, json_mode);
let (message, finish_reason) = build_response_message(
has_tools,
content,
response.finish_reason,
response.tool_calls.as_ref(),
);
let resp = ChatCompletionResponse {
id: generate_id(),
object: "chat.completion",
created: unix_timestamp(),
model: model_name,
choices: vec![Choice {
index: 0,
message,
finish_reason,
}],
usage,
warnings,
};
(StatusCode::OK, Json(resp)).into_response()
}
Err(e) => runner_error_to_response(&e),
},
}
}
async fn handle_multiplex(
state: &SharedState,
request: &ChatCompletionRequest,
models: &[String],
) -> Response {
if request.stream {
return error_response(
StatusCode::BAD_REQUEST,
"Streaming is not supported for multiplex requests",
);
}
let strict = request
.strict_capabilities
.unwrap_or_else(|| env::var("EMBACLE_STRICT_CAPS").is_ok_and(|v| v == "true" || v == "1"));
let state_guard = state.read().await;
let default_provider = state_guard.active_provider();
let resolved: Vec<_> = models
.iter()
.map(|m| resolve_model(m, default_provider))
.collect();
let providers: Vec<_> = resolved.iter().map(|r| r.runner_type).collect();
let messages = convert_messages(&request.messages);
let mut validation_request = ChatRequest::new(messages.clone());
validation_request.temperature = request.temperature;
validation_request.max_tokens = request.max_tokens;
validation_request.top_p = request.top_p;
validation_request.stop = request.stop.as_ref().map(StopField::to_bounded_vec);
validation_request.response_format =
request.response_format.as_ref().map(server_format_to_core);
for &provider_type in &providers {
let runner = match state_guard.get_runner(provider_type).await {
Ok(r) => r,
Err(e) => return runner_error_to_response(&e),
};
match embacle::validate_capabilities(
runner.name(),
runner.capabilities(),
&validation_request,
strict,
) {
Ok(w) => {
for warning in &w {
warn!(provider = runner.name(), warning = %warning, "Capability warning");
}
}
Err(e) => return runner_error_to_response(&e),
}
}
drop(state_guard);
let engine = MultiplexEngine::new(state);
let params = MultiplexParams {
temperature: request.temperature,
max_tokens: request.max_tokens,
top_p: request.top_p,
stop: request.stop.as_ref().map(StopField::to_bounded_vec),
response_format: request.response_format.as_ref().map(server_format_to_core),
};
match engine.execute(&messages, &providers, ¶ms).await {
Ok(result) => {
let results = result
.responses
.into_iter()
.map(|r| MultiplexProviderResult {
provider: r.provider,
model: r.model,
content: r.content,
error: r.error,
duration_ms: r.duration_ms,
})
.collect();
let resp = MultiplexResponse {
id: generate_id(),
object: "chat.completion.multiplex",
created: unix_timestamp(),
results,
summary: result.summary,
};
(StatusCode::OK, Json(resp)).into_response()
}
Err(e) => runner_error_to_response(&e),
}
}
fn build_response_message(
has_tools: bool,
content: String,
finish_reason: Option<String>,
native_tool_calls: Option<&Vec<embacle::ToolCallRequest>>,
) -> (ResponseMessage, Option<String>) {
if let Some(calls) = native_tool_calls {
if !calls.is_empty() {
let tool_calls: Vec<ToolCall> = calls
.iter()
.enumerate()
.map(|(i, tc)| ToolCall {
index: i,
id: tc.id.clone(),
tool_type: "function".to_owned(),
function: ToolCallFunction {
name: tc.function_name.clone(),
arguments: serde_json::to_string(&tc.arguments)
.unwrap_or_else(|_| "{}".to_owned()),
},
})
.collect();
let text_content = if content.is_empty() {
None
} else {
Some(content)
};
return (
ResponseMessage {
role: "assistant",
content: text_content,
tool_calls: Some(tool_calls),
},
Some("tool_calls".to_owned()),
);
}
}
if has_tools {
let parsed_calls = embacle::parse_tool_call_blocks(&content);
if parsed_calls.is_empty() {
(
ResponseMessage {
role: "assistant",
content: Some(content),
tool_calls: None,
},
finish_reason.or_else(|| Some("stop".to_owned())),
)
} else {
let remaining_text = embacle::strip_tool_call_blocks(&content);
let text_content = if remaining_text.is_empty() {
None
} else {
Some(remaining_text)
};
let tool_calls: Vec<ToolCall> = parsed_calls
.iter()
.enumerate()
.map(|(i, fc)| ToolCall {
index: i,
id: generate_tool_call_id(&fc.name, i),
tool_type: "function".to_owned(),
function: ToolCallFunction {
name: fc.name.clone(),
arguments: serde_json::to_string(&fc.args)
.unwrap_or_else(|_| "{}".to_owned()),
},
})
.collect();
(
ResponseMessage {
role: "assistant",
content: text_content,
tool_calls: Some(tool_calls),
},
Some("tool_calls".to_owned()),
)
}
} else {
(
ResponseMessage {
role: "assistant",
content: Some(content),
tool_calls: None,
},
finish_reason.or_else(|| Some("stop".to_owned())),
)
}
}
fn content_as_text(content: Option<&MessageContent>) -> String {
content.map(MessageContent::as_text).unwrap_or_default()
}
fn parse_data_uri(url: &str) -> Option<embacle::ImagePart> {
let rest = url.strip_prefix("data:")?;
let (mime_type, data) = rest.split_once(";base64,")?;
embacle::ImagePart::new(data, mime_type).ok()
}
fn extract_images(content: Option<&MessageContent>) -> Option<Vec<embacle::ImagePart>> {
let Some(MessageContent::Parts(parts)) = content else {
return None;
};
let images: Vec<embacle::ImagePart> = parts
.iter()
.filter_map(|p| match p {
ContentPart::ImageUrl { image_url } => parse_data_uri(&image_url.url),
ContentPart::Text { .. } => None,
})
.collect();
if images.is_empty() {
None
} else {
Some(images)
}
}
fn convert_messages(messages: &[ChatCompletionMessage]) -> Vec<ChatMessage> {
let mut result = Vec::with_capacity(messages.len());
let mut i = 0;
while i < messages.len() {
let m = &messages[i];
match m.role.as_str() {
"system" => {
result.push(ChatMessage::system(content_as_text(m.content.as_ref())));
i += 1;
}
"user" => {
let text = content_as_text(m.content.as_ref());
let images = extract_images(m.content.as_ref());
if let Some(imgs) = images {
result.push(ChatMessage::user_with_images(text, imgs));
} else {
result.push(ChatMessage::user(text));
}
i += 1;
}
"assistant" => {
if let Some(ref tool_calls) = m.tool_calls {
let mut text = content_as_text(m.content.as_ref());
for tc in tool_calls {
text.push_str("\n<tool_call>\n");
let payload = serde_json::json!({
"name": tc.function.name,
"arguments": serde_json::from_str::<serde_json::Value>(&tc.function.arguments)
.unwrap_or_else(|_| serde_json::Value::Object(serde_json::Map::new()))
});
text.push_str(
&serde_json::to_string(&payload).unwrap_or_else(|_| "{}".to_owned()),
);
text.push_str("\n</tool_call>");
}
result.push(ChatMessage::assistant(text));
} else {
result.push(ChatMessage::assistant(content_as_text(m.content.as_ref())));
}
i += 1;
}
"tool" => {
let mut tool_responses = Vec::new();
while i < messages.len() && messages[i].role == "tool" {
let tool_msg = &messages[i];
let name = tool_msg.name.as_deref().unwrap_or("unknown");
let content_text = content_as_text(tool_msg.content.as_ref());
let response_value: serde_json::Value = if content_text.is_empty() {
serde_json::Value::Null
} else {
serde_json::from_str(&content_text)
.unwrap_or(serde_json::Value::String(content_text))
};
tool_responses.push(embacle::FunctionResponse {
name: name.to_owned(),
response: response_value,
});
i += 1;
}
let text = embacle::format_tool_results_as_text(&tool_responses);
result.push(ChatMessage::user(text));
}
other => {
warn!(role = other, "Unknown message role, mapping to user");
result.push(ChatMessage::user(content_as_text(m.content.as_ref())));
i += 1;
}
}
}
result
}
fn server_tool_to_core(tool: &OpenAiToolDefinition) -> embacle::ToolDefinition {
embacle::ToolDefinition {
name: tool.function.name.clone(),
description: tool.function.description.clone().unwrap_or_default(),
parameters: tool.function.parameters.clone(),
}
}
fn server_choice_to_core(choice: &ToolChoice) -> embacle::ToolChoice {
match choice {
ToolChoice::Mode(m) => match m.as_str() {
"none" => embacle::ToolChoice::None,
"required" => embacle::ToolChoice::Required,
_ => embacle::ToolChoice::Auto,
},
ToolChoice::Specific(s) => embacle::ToolChoice::Specific {
name: s.function.name.clone(),
},
}
}
fn server_format_to_core(format: &ResponseFormatRequest) -> embacle::ResponseFormat {
match format {
ResponseFormatRequest::Text => embacle::ResponseFormat::Text,
ResponseFormatRequest::JsonObject => embacle::ResponseFormat::JsonObject,
ResponseFormatRequest::JsonSchema { json_schema } => embacle::ResponseFormat::JsonSchema {
name: json_schema.name.clone(),
schema: json_schema.schema.clone(),
},
}
}
fn tools_to_declarations(tools: &[OpenAiToolDefinition]) -> Vec<FunctionDeclaration> {
tools
.iter()
.map(|t| FunctionDeclaration {
name: t.function.name.clone(),
description: t.function.description.clone().unwrap_or_default(),
parameters: t.function.parameters.clone(),
})
.collect()
}
fn inject_tool_catalog_as_user_message(messages: &mut [ChatMessage], catalog: &str) {
if let Some(last_user) = messages
.iter_mut()
.rev()
.find(|m| m.role == MessageRole::User)
{
let augmented = format!("{catalog}\n\n{}", last_user.content);
*last_user = ChatMessage::user(augmented);
} else {
warn!("No user message found for tool catalog injection");
}
}
fn is_tool_choice_none(tool_choice: Option<&ToolChoice>) -> bool {
matches!(tool_choice, Some(ToolChoice::Mode(ref m)) if m == "none")
}
pub(crate) fn generate_tool_call_id(name: &str, index: usize) -> String {
format!("call_{name}_{index}")
}
fn runner_error_to_response(err: &RunnerError) -> Response {
let (status, error_type) = match err.kind {
ErrorKind::BinaryNotFound => (StatusCode::SERVICE_UNAVAILABLE, "provider_not_available"),
ErrorKind::AuthFailure => (StatusCode::UNAUTHORIZED, "authentication_error"),
ErrorKind::Timeout => (StatusCode::GATEWAY_TIMEOUT, "timeout_error"),
ErrorKind::ExternalService => (StatusCode::BAD_GATEWAY, "external_service_error"),
ErrorKind::Config => (StatusCode::BAD_REQUEST, "invalid_request_error"),
ErrorKind::Guardrail => (StatusCode::BAD_REQUEST, "guardrail_error"),
ErrorKind::ModelUnavailable => (StatusCode::NOT_FOUND, "model_not_found"),
ErrorKind::Internal => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
};
error!(kind = ?err.kind, message = %err.message, "Runner error");
let body = ErrorResponse::new(error_type, &err.message);
(status, Json(body)).into_response()
}
fn error_response(status: StatusCode, message: &str) -> Response {
let body = ErrorResponse::new("invalid_request_error", message);
(status, Json(body)).into_response()
}
static ID_COUNTER: AtomicU64 = AtomicU64::new(0);
pub fn generate_id() -> String {
let ts = unix_timestamp();
let seq = ID_COUNTER.fetch_add(1, Ordering::Relaxed);
format!("chatcmpl-{ts:x}{seq:08x}")
}
pub fn unix_timestamp() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |d| d.as_secs())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::openai_types::{
ContentPart, FunctionObject, ImageUrlDetail, ToolCall, ToolCallFunction, ToolDefinition,
};
use MessageRole;
fn text_msg(role: &str, content: Option<&str>) -> ChatCompletionMessage {
ChatCompletionMessage {
role: role.to_owned(),
content: content.map(|c| MessageContent::Text(c.to_owned())),
tool_calls: None,
tool_call_id: None,
name: None,
}
}
#[test]
fn convert_messages_maps_roles() {
let openai_msgs = vec![
text_msg("system", Some("You are helpful")),
text_msg("user", Some("Hello")),
text_msg("assistant", Some("Hi there")),
];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages.len(), 3);
assert_eq!(messages[0].role, MessageRole::System);
assert_eq!(messages[1].role, MessageRole::User);
assert_eq!(messages[2].role, MessageRole::Assistant);
}
#[test]
fn convert_unknown_role_defaults_to_user() {
let openai_msgs = vec![text_msg("function", Some("result"))];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages[0].role, MessageRole::User);
}
#[test]
fn convert_assistant_with_tool_calls() {
let openai_msgs = vec![ChatCompletionMessage {
role: "assistant".to_owned(),
content: None,
tool_calls: Some(vec![ToolCall {
index: 0,
id: "call_1".to_owned(),
tool_type: "function".to_owned(),
function: ToolCallFunction {
name: "get_weather".to_owned(),
arguments: r#"{"city":"Paris"}"#.to_owned(),
},
}]),
tool_call_id: None,
name: None,
}];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, MessageRole::Assistant);
assert!(messages[0].content.contains("<tool_call>"));
assert!(messages[0].content.contains("get_weather"));
assert!(messages[0].content.contains("</tool_call>"));
}
#[test]
fn convert_tool_messages_to_user() {
let openai_msgs = vec![
ChatCompletionMessage {
role: "tool".to_owned(),
content: Some(MessageContent::Text(r#"{"temp":72}"#.to_owned())),
tool_calls: None,
tool_call_id: Some("call_1".to_owned()),
name: Some("get_weather".to_owned()),
},
ChatCompletionMessage {
role: "tool".to_owned(),
content: Some(MessageContent::Text(r#"{"time":"14:30"}"#.to_owned())),
tool_calls: None,
tool_call_id: Some("call_2".to_owned()),
name: Some("get_time".to_owned()),
},
];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, MessageRole::User);
assert!(messages[0].content.contains("tool_result"));
assert!(messages[0].content.contains("get_weather"));
assert!(messages[0].content.contains("get_time"));
}
#[test]
fn convert_messages_none_content() {
let openai_msgs = vec![text_msg("user", None)];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages[0].content, "");
}
#[test]
fn convert_multipart_user_message_extracts_images() {
let openai_msgs = vec![ChatCompletionMessage {
role: "user".to_owned(),
content: Some(MessageContent::Parts(vec![
ContentPart::Text {
text: "What is this?".to_owned(),
},
ContentPart::ImageUrl {
image_url: ImageUrlDetail {
url: "data:image/png;base64,aGVsbG8=".to_owned(),
},
},
])),
tool_calls: None,
tool_call_id: None,
name: None,
}];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].content, "What is this?");
let images = messages[0].images.as_ref().expect("images present"); assert_eq!(images.len(), 1);
assert_eq!(images[0].mime_type, "image/png");
assert_eq!(images[0].data, "aGVsbG8=");
}
#[test]
fn parse_data_uri_valid() {
let img = parse_data_uri("data:image/jpeg;base64,AAAA").expect("should parse"); assert_eq!(img.mime_type, "image/jpeg");
assert_eq!(img.data, "AAAA");
}
#[test]
fn parse_data_uri_invalid_format() {
assert!(parse_data_uri("https://example.com/image.png").is_none());
assert!(parse_data_uri("data:text/plain;base64,abc").is_none());
assert!(parse_data_uri("data:image/png;abc").is_none());
}
#[test]
fn convert_plain_string_content_backward_compat() {
let openai_msgs = vec![text_msg("user", Some("hello"))];
let messages = convert_messages(&openai_msgs);
assert_eq!(messages[0].content, "hello");
assert!(messages[0].images.is_none());
}
#[test]
fn tools_to_declarations_converts() {
let tools = vec![ToolDefinition {
tool_type: "function".to_owned(),
function: FunctionObject {
name: "search".to_owned(),
description: Some("Search the web".to_owned()),
parameters: Some(serde_json::json!({
"type": "object",
"properties": {"q": {"type": "string"}},
"required": ["q"]
})),
},
}];
let decls = tools_to_declarations(&tools);
assert_eq!(decls.len(), 1);
assert_eq!(decls[0].name, "search");
assert_eq!(decls[0].description, "Search the web");
assert!(decls[0].parameters.is_some());
}
#[test]
fn tool_choice_none_detection() {
let none_choice = ToolChoice::Mode("none".to_owned());
assert!(is_tool_choice_none(Some(&none_choice)));
let auto_choice = ToolChoice::Mode("auto".to_owned());
assert!(!is_tool_choice_none(Some(&auto_choice)));
assert!(!is_tool_choice_none(None));
}
#[test]
fn content_as_text_none() {
assert_eq!(content_as_text(None), "");
}
#[test]
fn content_as_text_plain() {
let content = MessageContent::Text("hello".to_owned());
assert_eq!(content_as_text(Some(&content)), "hello");
}
#[test]
fn generate_tool_call_id_format() {
let id = generate_tool_call_id("get_weather", 0);
assert_eq!(id, "call_get_weather_0");
}
fn named_tool(name: &str) -> ToolDefinition {
ToolDefinition {
tool_type: "function".to_owned(),
function: FunctionObject {
name: name.to_owned(),
description: None,
parameters: None,
},
}
}
fn named_decl(name: &str) -> FunctionDeclaration {
FunctionDeclaration {
name: name.to_owned(),
description: String::new(),
parameters: None,
}
}
#[test]
fn select_server_declarations_returns_all_when_unrestricted() {
let available = vec![named_decl("a"), named_decl("b")];
let selected = select_server_declarations(&available, None);
assert_eq!(selected.len(), 2);
}
#[test]
fn select_server_declarations_restricts_to_requested_names() {
let available = vec![named_decl("a"), named_decl("b"), named_decl("c")];
let requested = vec![named_tool("b"), named_tool("nonexistent")];
let selected = select_server_declarations(&available, Some(&requested));
assert_eq!(selected.len(), 1);
assert_eq!(selected[0].name, "b");
}
#[test]
fn select_server_declarations_empty_request_offers_all() {
let available = vec![named_decl("a")];
let requested: Vec<ToolDefinition> = vec![];
let selected = select_server_declarations(&available, Some(&requested));
assert_eq!(selected.len(), 1);
}
#[test]
fn generate_id_has_prefix() {
let id = generate_id();
assert!(id.starts_with("chatcmpl-"));
}
#[test]
fn error_maps_binary_not_found_to_503() {
let err = RunnerError::binary_not_found("claude");
let (status, _) = match err.kind {
ErrorKind::BinaryNotFound => {
(StatusCode::SERVICE_UNAVAILABLE, "provider_not_available")
}
_ => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
};
assert_eq!(status, StatusCode::SERVICE_UNAVAILABLE);
}
#[test]
fn error_maps_auth_to_401() {
let err = RunnerError::auth_failure("bad token");
let (status, _) = match err.kind {
ErrorKind::AuthFailure => (StatusCode::UNAUTHORIZED, "authentication_error"),
_ => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
};
assert_eq!(status, StatusCode::UNAUTHORIZED);
}
#[test]
fn error_maps_timeout_to_504() {
let err = RunnerError::timeout("too slow");
let (status, _) = match err.kind {
ErrorKind::Timeout => (StatusCode::GATEWAY_TIMEOUT, "timeout_error"),
_ => (StatusCode::INTERNAL_SERVER_ERROR, "server_error"),
};
assert_eq!(status, StatusCode::GATEWAY_TIMEOUT);
}
#[test]
fn inject_tool_catalog_as_user_message_prepends_to_last_user() {
let mut messages = vec![
ChatMessage::user("First question"),
ChatMessage::assistant("Some answer"),
ChatMessage::user("What is the weather?"),
];
let catalog = "## Available Tools\n- get_weather: Get the weather";
inject_tool_catalog_as_user_message(&mut messages, catalog);
assert_eq!(messages.len(), 3);
assert!(messages[2].content.starts_with("## Available Tools"));
assert!(messages[2].content.contains("What is the weather?"));
assert_eq!(messages[0].content, "First question");
}
#[test]
fn inject_tool_catalog_as_user_message_single_user() {
let mut messages = vec![
ChatMessage::system("You are helpful"),
ChatMessage::user("Hello"),
];
let catalog = "## Tools\nsome tools";
inject_tool_catalog_as_user_message(&mut messages, catalog);
assert!(messages[1].content.starts_with("## Tools"));
assert!(messages[1].content.contains("Hello"));
}
#[test]
fn wants_json_matches_json_formats() {
use embacle::types::ResponseFormat;
assert!(!wants_json(None));
assert!(!wants_json(Some(&ResponseFormat::Text)));
assert!(wants_json(Some(&ResponseFormat::JsonObject)));
assert!(wants_json(Some(&ResponseFormat::JsonSchema {
name: "test".to_owned(),
schema: serde_json::json!({}),
})));
}
#[test]
fn strip_json_fences_removes_markdown_wrapper() {
let fenced = "```json\n{\"key\":\"value\"}\n```".to_owned();
assert_eq!(strip_json_fences(fenced, true), "{\"key\":\"value\"}");
}
#[test]
fn strip_json_fences_passes_through_in_text_mode() {
let fenced = "```json\n{\"key\":\"value\"}\n```".to_owned();
assert_eq!(strip_json_fences(fenced.clone(), false), fenced);
}
#[test]
fn strip_json_fences_leaves_clean_json_unchanged() {
let clean = "{\"key\":\"value\"}".to_owned();
assert_eq!(strip_json_fences(clean.clone(), true), clean);
}
}