From bbb68f74908192f71f5cbb68c9bedb2a606882b6 Mon Sep 17 00:00:00 2001 From: Jaswinder Date: Fri, 13 Feb 2026 15:26:49 +0100 Subject: [PATCH] Add OpenAI-compatible HTTP API (/v1/chat/completions, /v1/models) (#31) * feat: add OpenAI-compatible HTTP API (/v1/chat/completions, /v1/models) * - Reject model mismatches: validate req.model against the active model and return 404 model_not_found instead of silently ignoring it - Add x-ironclaw-streaming: simulated response header so clients know streaming is not true token-by-token delivery - Use SSE event type "error" for mid-stream LLM failures so clients can distinguish errors from content chunks - Mark docker-compose credentials as dev-only - Add integration tests for model mismatch, streaming header, and body size limit (axum's default 2MB) * fix: address Copilot review feedback on OpenAI-compat API - Wire chat_rate_limiter into /v1/chat/completions handler - Execute LLM before starting SSE stream so failures return proper HTTP errors instead of SSE error events - Validate tool-role messages require tool_call_id and name fields - Surface list_models() errors in models_handler via map_llm_error - Reject unknown roles with 400 instead of defaulting to User Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: firat.sertgoz Co-authored-by: Claude Opus 4.6 --- FEATURE_PARITY.md | 2 +- docker-compose.yml | 20 + src/channels/web/mod.rs | 9 + src/channels/web/openai_compat.rs | 1094 ++++++++++++++++++++++++++++ src/channels/web/server.rs | 8 + src/channels/web/ws.rs | 1 + src/main.rs | 1 + tests/openai_compat_integration.rs | 478 ++++++++++++ tests/ws_gateway_integration.rs | 1 + 9 files changed, 1613 insertions(+), 1 deletion(-) create mode 100644 docker-compose.yml create mode 100644 src/channels/web/openai_compat.rs create mode 100644 tests/openai_compat_integration.rs diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index a5ce7092..b6265291 100644 --- a/FEATURE_PARITY.md +++ b/FEATURE_PARITY.md @@ -37,7 +37,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O | Session management/routing | ✅ | ✅ | SessionManager exists | | Configuration hot-reload | ✅ | ❌ | | | Network modes (loopback/LAN/remote) | ✅ | 🚧 | HTTP only | -| OpenAI-compatible HTTP API | ✅ | ❌ | /v1/chat/completions | +| OpenAI-compatible HTTP API | ✅ | ✅ | /v1/chat/completions | | Canvas hosting | ✅ | ❌ | Agent-driven UI | | Gateway lock (PID-based) | ✅ | ❌ | | | launchd/systemd integration | ✅ | ❌ | | diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 00000000..9b82fcdb --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,20 @@ +# Local development only — do NOT use these credentials in production. +services: + postgres: + image: pgvector/pgvector:pg16 + ports: + - "5432:5432" + environment: + POSTGRES_DB: ironclaw + POSTGRES_USER: ironclaw + POSTGRES_PASSWORD: ironclaw # dev-only, change for any non-local deployment + volumes: + - pgdata:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U ironclaw"] + interval: 5s + timeout: 3s + retries: 5 + +volumes: + pgdata: diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index 80a5242b..ac6ef2f6 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -16,6 +16,7 @@ pub mod auth; pub mod log_layer; +pub mod openai_compat; pub mod server; pub mod sse; pub mod types; @@ -81,6 +82,7 @@ impl GatewayChannel { user_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), + llm_provider: None, chat_rate_limiter: server::RateLimiter::new(30, 60), }); @@ -107,6 +109,7 @@ impl GatewayChannel { user_id: self.state.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: self.state.ws_tracker.clone(), + llm_provider: self.state.llm_provider.clone(), chat_rate_limiter: server::RateLimiter::new(30, 60), }; mutate(&mut new_state); @@ -171,6 +174,12 @@ impl GatewayChannel { self } + /// Inject the LLM provider for OpenAI-compatible API proxy. + pub fn with_llm_provider(mut self, llm: Arc) -> Self { + self.rebuild_state(|s| s.llm_provider = Some(llm)); + self + } + /// Get the auth token (for printing to console on startup). pub fn auth_token(&self) -> &str { &self.auth_token diff --git a/src/channels/web/openai_compat.rs b/src/channels/web/openai_compat.rs new file mode 100644 index 00000000..7185ca7b --- /dev/null +++ b/src/channels/web/openai_compat.rs @@ -0,0 +1,1094 @@ +//! OpenAI-compatible HTTP API (`/v1/chat/completions`, `/v1/models`). +//! +//! This module provides a direct LLM proxy through the web gateway so any +//! standard OpenAI client library can use IronClaw as a backend by simply +//! changing the `base_url`. + +use std::sync::Arc; + +use axum::{ + Json, + extract::State, + http::{HeaderValue, StatusCode}, + response::{ + IntoResponse, Response, + sse::{Event, KeepAlive, Sse}, + }, +}; +use serde::{Deserialize, Serialize}; + +use crate::llm::{ + ChatMessage, CompletionRequest, FinishReason, Role, ToolCall, ToolCompletionRequest, + ToolDefinition, +}; + +use super::server::GatewayState; + +// --------------------------------------------------------------------------- +// OpenAI request types +// --------------------------------------------------------------------------- + +#[derive(Debug, Deserialize)] +pub struct OpenAiChatRequest { + pub model: String, + pub messages: Vec, + #[serde(default)] + pub temperature: Option, + #[serde(default)] + pub max_tokens: Option, + #[serde(default)] + pub stream: Option, + #[serde(default)] + pub tools: Option>, + #[serde(default)] + pub tool_choice: Option, + #[serde(default)] + pub stop: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAiMessage { + pub role: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAiTool { + #[serde(rename = "type")] + pub tool_type: String, + pub function: OpenAiFunction, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAiFunction { + pub name: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub description: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub parameters: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAiToolCall { + pub id: String, + #[serde(rename = "type")] + pub call_type: String, + pub function: OpenAiToolCallFunction, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAiToolCallFunction { + pub name: String, + pub arguments: String, +} + +// --------------------------------------------------------------------------- +// OpenAI response types (non-streaming) +// --------------------------------------------------------------------------- + +#[derive(Debug, Serialize)] +pub struct OpenAiChatResponse { + pub id: String, + pub object: &'static str, + pub created: u64, + pub model: String, + pub choices: Vec, + pub usage: OpenAiUsage, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiChoice { + pub index: u32, + pub message: OpenAiMessage, + pub finish_reason: String, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiUsage { + pub prompt_tokens: u32, + pub completion_tokens: u32, + pub total_tokens: u32, +} + +// --------------------------------------------------------------------------- +// OpenAI response types (streaming) +// --------------------------------------------------------------------------- + +#[derive(Debug, Serialize)] +pub struct OpenAiChatChunk { + pub id: String, + pub object: &'static str, + pub created: u64, + pub model: String, + pub choices: Vec, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiChunkChoice { + pub index: u32, + pub delta: OpenAiDelta, + #[serde(skip_serializing_if = "Option::is_none")] + pub finish_reason: Option, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiDelta { + #[serde(skip_serializing_if = "Option::is_none")] + pub role: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiToolCallDelta { + pub index: u32, + #[serde(skip_serializing_if = "Option::is_none")] + pub id: Option, + #[serde(rename = "type", skip_serializing_if = "Option::is_none")] + pub call_type: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub function: Option, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiToolCallFunctionDelta { + #[serde(skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub arguments: Option, +} + +// --------------------------------------------------------------------------- +// Error response +// --------------------------------------------------------------------------- + +#[derive(Debug, Serialize)] +pub struct OpenAiErrorResponse { + pub error: OpenAiErrorDetail, +} + +#[derive(Debug, Serialize)] +pub struct OpenAiErrorDetail { + pub message: String, + #[serde(rename = "type")] + pub error_type: String, + pub param: Option, + pub code: Option, +} + +// --------------------------------------------------------------------------- +// Conversion functions +// --------------------------------------------------------------------------- + +fn parse_role(s: &str) -> Result { + match s { + "system" => Ok(Role::System), + "user" => Ok(Role::User), + "assistant" => Ok(Role::Assistant), + "tool" => Ok(Role::Tool), + _ => Err(format!("Unknown role: '{}'", s)), + } +} + +pub fn convert_messages(messages: &[OpenAiMessage]) -> Result, String> { + messages + .iter() + .enumerate() + .map(|(i, m)| { + let role = parse_role(&m.role).map_err(|e| format!("messages[{}]: {}", i, e))?; + match role { + Role::Tool => { + let tool_call_id = m.tool_call_id.as_deref().ok_or_else(|| { + format!("messages[{}]: tool message requires 'tool_call_id'", i) + })?; + let name = m + .name + .as_deref() + .ok_or_else(|| format!("messages[{}]: tool message requires 'name'", i))?; + Ok(ChatMessage::tool_result( + tool_call_id, + name, + m.content.as_deref().unwrap_or(""), + )) + } + Role::Assistant => { + if let Some(ref tcs) = m.tool_calls { + let calls: Vec = tcs + .iter() + .map(|tc| ToolCall { + id: tc.id.clone(), + name: tc.function.name.clone(), + arguments: serde_json::from_str(&tc.function.arguments) + .unwrap_or(serde_json::Value::Object(Default::default())), + }) + .collect(); + Ok(ChatMessage::assistant_with_tool_calls( + m.content.clone(), + calls, + )) + } else { + Ok(ChatMessage::assistant(m.content.as_deref().unwrap_or(""))) + } + } + _ => Ok(ChatMessage { + role, + content: m.content.as_deref().unwrap_or("").to_string(), + tool_call_id: None, + name: m.name.clone(), + tool_calls: None, + }), + } + }) + .collect() +} + +pub fn convert_tools(tools: &[OpenAiTool]) -> Vec { + tools + .iter() + .filter(|t| t.tool_type == "function") + .map(|t| ToolDefinition { + name: t.function.name.clone(), + description: t.function.description.clone().unwrap_or_default(), + parameters: t + .function + .parameters + .clone() + .unwrap_or(serde_json::json!({"type": "object", "properties": {}})), + }) + .collect() +} + +fn convert_tool_calls_to_openai(calls: &[ToolCall]) -> Vec { + calls + .iter() + .map(|tc| OpenAiToolCall { + id: tc.id.clone(), + call_type: "function".to_string(), + function: OpenAiToolCallFunction { + name: tc.name.clone(), + arguments: serde_json::to_string(&tc.arguments).unwrap_or_default(), + }, + }) + .collect() +} + +pub fn finish_reason_str(reason: FinishReason) -> String { + match reason { + FinishReason::Stop => "stop".to_string(), + FinishReason::Length => "length".to_string(), + FinishReason::ToolUse => "tool_calls".to_string(), + FinishReason::ContentFilter => "content_filter".to_string(), + FinishReason::Unknown => "stop".to_string(), + } +} + +fn normalize_tool_choice(val: &serde_json::Value) -> Option { + match val { + serde_json::Value::String(s) => Some(s.clone()), + serde_json::Value::Object(obj) => { + // { "type": "function", "function": { "name": "foo" } } → "required" + if obj.contains_key("function") { + Some("required".to_string()) + } else { + obj.get("type") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()) + } + } + _ => None, + } +} + +fn map_llm_error(err: crate::error::LlmError) -> (StatusCode, Json) { + let (status, error_type, code) = match &err { + crate::error::LlmError::AuthFailed { .. } + | crate::error::LlmError::SessionExpired { .. } => ( + StatusCode::UNAUTHORIZED, + "authentication_error", + "auth_error", + ), + crate::error::LlmError::RateLimited { .. } => ( + StatusCode::TOO_MANY_REQUESTS, + "rate_limit_error", + "rate_limit", + ), + crate::error::LlmError::ContextLengthExceeded { .. } => ( + StatusCode::BAD_REQUEST, + "invalid_request_error", + "context_length_exceeded", + ), + crate::error::LlmError::ModelNotAvailable { .. } => ( + StatusCode::NOT_FOUND, + "invalid_request_error", + "model_not_found", + ), + _ => ( + StatusCode::INTERNAL_SERVER_ERROR, + "server_error", + "internal_error", + ), + }; + + ( + status, + Json(OpenAiErrorResponse { + error: OpenAiErrorDetail { + message: err.to_string(), + error_type: error_type.to_string(), + param: None, + code: Some(code.to_string()), + }, + }), + ) +} + +fn openai_error( + status: StatusCode, + message: impl Into, + error_type: &str, +) -> (StatusCode, Json) { + ( + status, + Json(OpenAiErrorResponse { + error: OpenAiErrorDetail { + message: message.into(), + error_type: error_type.to_string(), + param: None, + code: None, + }, + }), + ) +} + +fn chat_completion_id() -> String { + format!("chatcmpl-{}", uuid::Uuid::new_v4().simple()) +} + +fn unix_timestamp() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +/// Extract stop sequences from the flexible `stop` field. +fn parse_stop(val: &serde_json::Value) -> Option> { + match val { + serde_json::Value::String(s) => Some(vec![s.clone()]), + serde_json::Value::Array(arr) => { + let strs: Vec = arr + .iter() + .filter_map(|v| v.as_str().map(String::from)) + .collect(); + if strs.is_empty() { None } else { Some(strs) } + } + _ => None, + } +} + +// --------------------------------------------------------------------------- +// Handlers +// --------------------------------------------------------------------------- + +pub async fn chat_completions_handler( + State(state): State>, + Json(req): Json, +) -> Result)> { + if !state.chat_rate_limiter.check() { + return Err(openai_error( + StatusCode::TOO_MANY_REQUESTS, + "Rate limit exceeded. Please try again later.", + "rate_limit_error", + )); + } + + let llm = state.llm_provider.as_ref().ok_or_else(|| { + openai_error( + StatusCode::SERVICE_UNAVAILABLE, + "LLM provider not configured", + "server_error", + ) + })?; + + if req.messages.is_empty() { + return Err(openai_error( + StatusCode::BAD_REQUEST, + "messages must not be empty", + "invalid_request_error", + )); + } + + // Validate the requested model matches the active model. + // Per-request model switching is not yet supported (see GH issue). + let active_model = llm.active_model_name(); + if req.model != active_model { + return Err(( + StatusCode::NOT_FOUND, + Json(OpenAiErrorResponse { + error: OpenAiErrorDetail { + message: format!( + "Model '{}' not found. The active model is '{}'.", + req.model, active_model + ), + error_type: "invalid_request_error".to_string(), + param: Some("model".to_string()), + code: Some("model_not_found".to_string()), + }, + }), + )); + } + + let has_tools = req.tools.as_ref().is_some_and(|t| !t.is_empty()); + let stream = req.stream.unwrap_or(false); + + if stream { + return handle_streaming(llm.clone(), req, has_tools) + .await + .map(IntoResponse::into_response); + } + + // --- Non-streaming path --- + + let messages = convert_messages(&req.messages) + .map_err(|e| openai_error(StatusCode::BAD_REQUEST, e, "invalid_request_error"))?; + let model_name = llm.active_model_name(); + let id = chat_completion_id(); + let created = unix_timestamp(); + + if has_tools { + let tools = convert_tools(req.tools.as_deref().unwrap_or(&[])); + let mut tool_req = ToolCompletionRequest::new(messages, tools); + if let Some(t) = req.temperature { + tool_req = tool_req.with_temperature(t); + } + if let Some(mt) = req.max_tokens { + tool_req = tool_req.with_max_tokens(mt); + } + if let Some(ref tc) = req.tool_choice { + if let Some(choice) = normalize_tool_choice(tc) { + tool_req = tool_req.with_tool_choice(choice); + } + } + + let resp = llm + .complete_with_tools(tool_req) + .await + .map_err(map_llm_error)?; + + let tool_calls_openai = if resp.tool_calls.is_empty() { + None + } else { + Some(convert_tool_calls_to_openai(&resp.tool_calls)) + }; + + let response = OpenAiChatResponse { + id, + object: "chat.completion", + created, + model: model_name, + choices: vec![OpenAiChoice { + index: 0, + message: OpenAiMessage { + role: "assistant".to_string(), + content: resp.content.clone(), + name: None, + tool_call_id: None, + tool_calls: tool_calls_openai, + }, + finish_reason: finish_reason_str(resp.finish_reason), + }], + usage: OpenAiUsage { + prompt_tokens: resp.input_tokens, + completion_tokens: resp.output_tokens, + total_tokens: resp.input_tokens + resp.output_tokens, + }, + }; + + Ok(Json(response).into_response()) + } else { + let mut comp_req = CompletionRequest::new(messages); + if let Some(t) = req.temperature { + comp_req = comp_req.with_temperature(t); + } + if let Some(mt) = req.max_tokens { + comp_req = comp_req.with_max_tokens(mt); + } + if let Some(ref stop_val) = req.stop { + comp_req.stop_sequences = parse_stop(stop_val); + } + + let resp = llm.complete(comp_req).await.map_err(map_llm_error)?; + + let response = OpenAiChatResponse { + id, + object: "chat.completion", + created, + model: model_name, + choices: vec![OpenAiChoice { + index: 0, + message: OpenAiMessage { + role: "assistant".to_string(), + content: Some(resp.content), + name: None, + tool_call_id: None, + tool_calls: None, + }, + finish_reason: finish_reason_str(resp.finish_reason), + }], + usage: OpenAiUsage { + prompt_tokens: resp.input_tokens, + completion_tokens: resp.output_tokens, + total_tokens: resp.input_tokens + resp.output_tokens, + }, + }; + + Ok(Json(response).into_response()) + } +} + +/// Handle streaming responses. +/// +/// The current `LlmProvider` returns complete responses (no streaming method). +/// We execute the LLM call first, then simulate chunked delivery by splitting +/// the response into word-boundary chunks. This ensures LLM failures return +/// proper HTTP errors instead of SSE error events. True token streaming can be +/// added later by extending `LlmProvider` with a `complete_stream()` method. +async fn handle_streaming( + llm: Arc, + req: OpenAiChatRequest, + has_tools: bool, +) -> Result)> { + let messages = convert_messages(&req.messages) + .map_err(|e| openai_error(StatusCode::BAD_REQUEST, e, "invalid_request_error"))?; + + let model_name = llm.active_model_name(); + let id = chat_completion_id(); + let created = unix_timestamp(); + + // Execute the LLM call before starting the SSE stream. + // Since streaming is simulated (LlmProvider returns complete responses), + // this lets us return proper HTTP errors on failure. + enum LlmResult { + Simple(crate::llm::CompletionResponse), + WithTools(crate::llm::ToolCompletionResponse), + } + + let llm_result = if has_tools { + let tools = convert_tools(req.tools.as_deref().unwrap_or(&[])); + let mut tool_req = ToolCompletionRequest::new(messages, tools); + if let Some(t) = req.temperature { + tool_req = tool_req.with_temperature(t); + } + if let Some(mt) = req.max_tokens { + tool_req = tool_req.with_max_tokens(mt); + } + if let Some(ref tc) = req.tool_choice { + if let Some(choice) = normalize_tool_choice(tc) { + tool_req = tool_req.with_tool_choice(choice); + } + } + LlmResult::WithTools( + llm.complete_with_tools(tool_req) + .await + .map_err(map_llm_error)?, + ) + } else { + let mut comp_req = CompletionRequest::new(messages); + if let Some(t) = req.temperature { + comp_req = comp_req.with_temperature(t); + } + if let Some(mt) = req.max_tokens { + comp_req = comp_req.with_max_tokens(mt); + } + if let Some(ref stop_val) = req.stop { + comp_req.stop_sequences = parse_stop(stop_val); + } + LlmResult::Simple(llm.complete(comp_req).await.map_err(map_llm_error)?) + }; + + // LLM succeeded — emit the response as SSE chunks + let (tx, rx) = tokio::sync::mpsc::channel::>(64); + + tokio::spawn(async move { + // Send initial chunk with role + let role_chunk = OpenAiChatChunk { + id: id.clone(), + object: "chat.completion.chunk", + created, + model: model_name.clone(), + choices: vec![OpenAiChunkChoice { + index: 0, + delta: OpenAiDelta { + role: Some("assistant".to_string()), + content: None, + tool_calls: None, + }, + finish_reason: None, + }], + }; + let data = serde_json::to_string(&role_chunk).unwrap_or_default(); + let _ = tx.send(Ok(Event::default().data(data))).await; + + match llm_result { + LlmResult::WithTools(resp) => { + // Stream content chunks + if let Some(ref content) = resp.content { + stream_content_chunks(&tx, &id, created, &model_name, content).await; + } + + // Stream tool calls + if !resp.tool_calls.is_empty() { + let deltas: Vec = resp + .tool_calls + .iter() + .enumerate() + .map(|(i, tc)| OpenAiToolCallDelta { + index: i as u32, + id: Some(tc.id.clone()), + call_type: Some("function".to_string()), + function: Some(OpenAiToolCallFunctionDelta { + name: Some(tc.name.clone()), + arguments: Some( + serde_json::to_string(&tc.arguments).unwrap_or_default(), + ), + }), + }) + .collect(); + + let chunk = OpenAiChatChunk { + id: id.clone(), + object: "chat.completion.chunk", + created, + model: model_name.clone(), + choices: vec![OpenAiChunkChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: None, + tool_calls: Some(deltas), + }, + finish_reason: None, + }], + }; + let data = serde_json::to_string(&chunk).unwrap_or_default(); + let _ = tx.send(Ok(Event::default().data(data))).await; + } + + // Final chunk with finish_reason + send_finish_chunk(&tx, &id, created, &model_name, resp.finish_reason).await; + } + LlmResult::Simple(resp) => { + stream_content_chunks(&tx, &id, created, &model_name, &resp.content).await; + send_finish_chunk(&tx, &id, created, &model_name, resp.finish_reason).await; + } + } + + // Send [DONE] sentinel + let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; + }); + + let stream = tokio_stream::wrappers::ReceiverStream::new(rx); + let sse = Sse::new(stream).keep_alive(KeepAlive::new().text("")); + let mut response = sse.into_response(); + response.headers_mut().insert( + "x-ironclaw-streaming", + HeaderValue::from_static("simulated"), + ); + Ok(response) +} + +/// Split content into word-boundary chunks and send as SSE events. +async fn stream_content_chunks( + tx: &tokio::sync::mpsc::Sender>, + id: &str, + created: u64, + model: &str, + content: &str, +) { + // Split on word boundaries, grouping ~20 chars per chunk + let mut buf = String::new(); + for word in content.split_inclusive(char::is_whitespace) { + buf.push_str(word); + if buf.len() >= 20 { + let chunk = OpenAiChatChunk { + id: id.to_string(), + object: "chat.completion.chunk", + created, + model: model.to_string(), + choices: vec![OpenAiChunkChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: Some(buf.clone()), + tool_calls: None, + }, + finish_reason: None, + }], + }; + let data = serde_json::to_string(&chunk).unwrap_or_default(); + if tx.send(Ok(Event::default().data(data))).await.is_err() { + return; + } + buf.clear(); + } + } + // Flush remaining + if !buf.is_empty() { + let chunk = OpenAiChatChunk { + id: id.to_string(), + object: "chat.completion.chunk", + created, + model: model.to_string(), + choices: vec![OpenAiChunkChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: Some(buf), + tool_calls: None, + }, + finish_reason: None, + }], + }; + let data = serde_json::to_string(&chunk).unwrap_or_default(); + let _ = tx.send(Ok(Event::default().data(data))).await; + } +} + +async fn send_finish_chunk( + tx: &tokio::sync::mpsc::Sender>, + id: &str, + created: u64, + model: &str, + reason: FinishReason, +) { + let chunk = OpenAiChatChunk { + id: id.to_string(), + object: "chat.completion.chunk", + created, + model: model.to_string(), + choices: vec![OpenAiChunkChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: None, + tool_calls: None, + }, + finish_reason: Some(finish_reason_str(reason)), + }], + }; + let data = serde_json::to_string(&chunk).unwrap_or_default(); + let _ = tx.send(Ok(Event::default().data(data))).await; +} + +pub async fn models_handler( + State(state): State>, +) -> Result, (StatusCode, Json)> { + let llm = state.llm_provider.as_ref().ok_or_else(|| { + openai_error( + StatusCode::SERVICE_UNAVAILABLE, + "LLM provider not configured", + "server_error", + ) + })?; + + let model_name = llm.active_model_name(); + let created = unix_timestamp(); + + // Try to fetch available models from the provider + let models = match llm.list_models().await { + Ok(names) if !names.is_empty() => names + .into_iter() + .map(|name| { + serde_json::json!({ + "id": name, + "object": "model", + "created": created, + "owned_by": "ironclaw" + }) + }) + .collect(), + Ok(_) => { + // Empty list: fall back to active model + vec![serde_json::json!({ + "id": model_name, + "object": "model", + "created": created, + "owned_by": "ironclaw" + })] + } + Err(e) => return Err(map_llm_error(e)), + }; + + Ok(Json(serde_json::json!({ + "object": "list", + "data": models + }))) +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_role() { + assert_eq!(parse_role("system").unwrap(), Role::System); + assert_eq!(parse_role("user").unwrap(), Role::User); + assert_eq!(parse_role("assistant").unwrap(), Role::Assistant); + assert_eq!(parse_role("tool").unwrap(), Role::Tool); + } + + #[test] + fn test_parse_role_unknown_rejected() { + let err = parse_role("unknown").unwrap_err(); + assert!(err.contains("Unknown role")); + assert!(err.contains("unknown")); + } + + #[test] + fn test_finish_reason_str() { + assert_eq!(finish_reason_str(FinishReason::Stop), "stop"); + assert_eq!(finish_reason_str(FinishReason::Length), "length"); + assert_eq!(finish_reason_str(FinishReason::ToolUse), "tool_calls"); + assert_eq!( + finish_reason_str(FinishReason::ContentFilter), + "content_filter" + ); + assert_eq!(finish_reason_str(FinishReason::Unknown), "stop"); + } + + #[test] + fn test_convert_messages_basic() { + let msgs = vec![ + OpenAiMessage { + role: "system".to_string(), + content: Some("You are helpful.".to_string()), + name: None, + tool_call_id: None, + tool_calls: None, + }, + OpenAiMessage { + role: "user".to_string(), + content: Some("Hello".to_string()), + name: None, + tool_call_id: None, + tool_calls: None, + }, + ]; + + let converted = convert_messages(&msgs).unwrap(); + assert_eq!(converted.len(), 2); + assert_eq!(converted[0].role, Role::System); + assert_eq!(converted[0].content, "You are helpful."); + assert_eq!(converted[1].role, Role::User); + assert_eq!(converted[1].content, "Hello"); + } + + #[test] + fn test_convert_messages_with_tool_results() { + let msgs = vec![OpenAiMessage { + role: "tool".to_string(), + content: Some("42".to_string()), + name: Some("calculator".to_string()), + tool_call_id: Some("call_123".to_string()), + tool_calls: None, + }]; + + let converted = convert_messages(&msgs).unwrap(); + assert_eq!(converted.len(), 1); + assert_eq!(converted[0].role, Role::Tool); + assert_eq!(converted[0].content, "42"); + assert_eq!(converted[0].tool_call_id.as_deref(), Some("call_123")); + assert_eq!(converted[0].name.as_deref(), Some("calculator")); + } + + #[test] + fn test_convert_tools() { + let tools = vec![OpenAiTool { + tool_type: "function".to_string(), + function: OpenAiFunction { + name: "get_weather".to_string(), + description: Some("Get weather for a location".to_string()), + parameters: Some(serde_json::json!({ + "type": "object", + "properties": { + "location": { "type": "string" } + }, + "required": ["location"] + })), + }, + }]; + + let converted = convert_tools(&tools); + assert_eq!(converted.len(), 1); + assert_eq!(converted[0].name, "get_weather"); + assert_eq!(converted[0].description, "Get weather for a location"); + } + + #[test] + fn test_convert_tool_calls_to_openai() { + let calls = vec![ToolCall { + id: "call_abc".to_string(), + name: "search".to_string(), + arguments: serde_json::json!({"query": "rust"}), + }]; + + let converted = convert_tool_calls_to_openai(&calls); + assert_eq!(converted.len(), 1); + assert_eq!(converted[0].id, "call_abc"); + assert_eq!(converted[0].call_type, "function"); + assert_eq!(converted[0].function.name, "search"); + assert!(converted[0].function.arguments.contains("rust")); + } + + #[test] + fn test_normalize_tool_choice() { + // String variant + let v = serde_json::json!("auto"); + assert_eq!(normalize_tool_choice(&v), Some("auto".to_string())); + + // Object with function + let v = serde_json::json!({"type": "function", "function": {"name": "foo"}}); + assert_eq!(normalize_tool_choice(&v), Some("required".to_string())); + + // Object with type only + let v = serde_json::json!({"type": "none"}); + assert_eq!(normalize_tool_choice(&v), Some("none".to_string())); + + // Null + let v = serde_json::Value::Null; + assert_eq!(normalize_tool_choice(&v), None); + } + + #[test] + fn test_openai_request_deserialize_minimal() { + let json = r#"{"model":"gpt-4","messages":[{"role":"user","content":"Hi"}]}"#; + let req: OpenAiChatRequest = serde_json::from_str(json).unwrap(); + assert_eq!(req.model, "gpt-4"); + assert_eq!(req.messages.len(), 1); + assert_eq!(req.stream, None); + assert_eq!(req.temperature, None); + } + + #[test] + fn test_openai_request_deserialize_streaming() { + let json = r#"{"model":"gpt-4","messages":[{"role":"user","content":"Hi"}],"stream":true,"temperature":0.7}"#; + let req: OpenAiChatRequest = serde_json::from_str(json).unwrap(); + assert_eq!(req.stream, Some(true)); + assert_eq!(req.temperature, Some(0.7)); + } + + #[test] + fn test_openai_response_serialize() { + let resp = OpenAiChatResponse { + id: "chatcmpl-test".to_string(), + object: "chat.completion", + created: 1234567890, + model: "test-model".to_string(), + choices: vec![OpenAiChoice { + index: 0, + message: OpenAiMessage { + role: "assistant".to_string(), + content: Some("Hello!".to_string()), + name: None, + tool_call_id: None, + tool_calls: None, + }, + finish_reason: "stop".to_string(), + }], + usage: OpenAiUsage { + prompt_tokens: 10, + completion_tokens: 5, + total_tokens: 15, + }, + }; + + let json = serde_json::to_value(&resp).unwrap(); + assert_eq!(json["object"], "chat.completion"); + assert_eq!(json["choices"][0]["finish_reason"], "stop"); + assert_eq!(json["choices"][0]["message"]["content"], "Hello!"); + assert_eq!(json["usage"]["total_tokens"], 15); + } + + #[test] + fn test_openai_message_with_null_content() { + let json = r#"{"role":"assistant","content":null,"tool_calls":[{"id":"call_1","type":"function","function":{"name":"search","arguments":"{\"q\":\"test\"}"}}]}"#; + let msg: OpenAiMessage = serde_json::from_str(json).unwrap(); + assert_eq!(msg.role, "assistant"); + assert!(msg.content.is_none()); + assert!(msg.tool_calls.is_some()); + assert_eq!(msg.tool_calls.as_ref().unwrap().len(), 1); + } + + #[test] + fn test_convert_messages_unknown_role_rejected() { + let msgs = vec![OpenAiMessage { + role: "moderator".to_string(), + content: Some("Hi".to_string()), + name: None, + tool_call_id: None, + tool_calls: None, + }]; + let err = convert_messages(&msgs).unwrap_err(); + assert!(err.contains("messages[0]")); + assert!(err.contains("Unknown role")); + } + + #[test] + fn test_convert_messages_tool_missing_fields() { + // Missing tool_call_id + let msgs = vec![OpenAiMessage { + role: "tool".to_string(), + content: Some("result".to_string()), + name: Some("calc".to_string()), + tool_call_id: None, + tool_calls: None, + }]; + let err = convert_messages(&msgs).unwrap_err(); + assert!(err.contains("tool_call_id")); + + // Missing name + let msgs = vec![OpenAiMessage { + role: "tool".to_string(), + content: Some("result".to_string()), + name: None, + tool_call_id: Some("call_1".to_string()), + tool_calls: None, + }]; + let err = convert_messages(&msgs).unwrap_err(); + assert!(err.contains("'name'")); + } + + #[test] + fn test_parse_stop_string() { + let v = serde_json::json!("STOP"); + assert_eq!(parse_stop(&v), Some(vec!["STOP".to_string()])); + } + + #[test] + fn test_parse_stop_array() { + let v = serde_json::json!(["STOP", "END"]); + assert_eq!( + parse_stop(&v), + Some(vec!["STOP".to_string(), "END".to_string()]) + ); + } + + #[test] + fn test_parse_stop_null() { + let v = serde_json::Value::Null; + assert_eq!(parse_stop(&v), None); + } +} diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 8298a0bb..540b8174 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -137,6 +137,8 @@ pub struct GatewayState { pub shutdown_tx: tokio::sync::RwLock>>, /// WebSocket connection tracker. pub ws_tracker: Option>, + /// LLM provider for OpenAI-compatible API proxy. + pub llm_provider: Option>, /// Rate limiter for chat endpoints (30 messages per 60 seconds). pub chat_rate_limiter: RateLimiter, } @@ -235,6 +237,12 @@ pub async fn start_server( ) // Gateway control plane .route("/api/gateway/status", get(gateway_status_handler)) + // OpenAI-compatible API + .route( + "/v1/chat/completions", + post(super::openai_compat::chat_completions_handler), + ) + .route("/v1/models", get(super::openai_compat::models_handler)) .route_layer(middleware::from_fn_with_state( auth_state.clone(), auth_middleware, diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 9dde5a15..d6ebc0f0 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -485,6 +485,7 @@ mod tests { user_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60), } } diff --git a/src/main.rs b/src/main.rs index 4d8e08c1..7b981979 100644 --- a/src/main.rs +++ b/src/main.rs @@ -997,6 +997,7 @@ async fn main() -> anyhow::Result<()> { if let Some(ref jm) = container_job_manager { gw = gw.with_job_manager(Arc::clone(jm)); } + gw = gw.with_llm_provider(Arc::clone(&llm)); if config.sandbox.enabled { gw = gw.with_prompt_queue(Arc::clone(&prompt_queue)); diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs new file mode 100644 index 00000000..d0927a2c --- /dev/null +++ b/tests/openai_compat_integration.rs @@ -0,0 +1,478 @@ +//! Integration tests for the OpenAI-compatible API endpoints. +//! +//! Uses a mock LLM provider so no real API key is needed. + +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use async_trait::async_trait; +use rust_decimal::Decimal; + +use ironclaw::channels::web::server::{GatewayState, start_server}; +use ironclaw::channels::web::sse::SseManager; +use ironclaw::channels::web::ws::WsConnectionTracker; +use ironclaw::error::LlmError; +use ironclaw::llm::{ + CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ToolCompletionRequest, + ToolCompletionResponse, +}; + +const AUTH_TOKEN: &str = "test-openai-token"; + +// --------------------------------------------------------------------------- +// Mock LLM provider +// --------------------------------------------------------------------------- + +struct MockLlmProvider; + +#[async_trait] +impl LlmProvider for MockLlmProvider { + fn model_name(&self) -> &str { + "mock-model-v1" + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (Decimal::ZERO, Decimal::ZERO) + } + + async fn complete(&self, req: CompletionRequest) -> Result { + // Echo the last user message back + let user_msg = req + .messages + .iter() + .rev() + .find(|m| m.role == ironclaw::llm::Role::User) + .map(|m| m.content.clone()) + .unwrap_or_else(|| "no user message".to_string()); + + Ok(CompletionResponse { + content: format!("Mock response to: {}", user_msg), + input_tokens: 10, + output_tokens: 5, + finish_reason: FinishReason::Stop, + response_id: None, + }) + } + + async fn complete_with_tools( + &self, + req: ToolCompletionRequest, + ) -> Result { + // If tools are provided, return a tool call + if let Some(tool) = req.tools.first() { + Ok(ToolCompletionResponse { + content: None, + tool_calls: vec![ironclaw::llm::ToolCall { + id: "call_mock_001".to_string(), + name: tool.name.clone(), + arguments: serde_json::json!({"test": true}), + }], + input_tokens: 15, + output_tokens: 8, + finish_reason: FinishReason::ToolUse, + response_id: None, + }) + } else { + Ok(ToolCompletionResponse { + content: Some("No tools available".to_string()), + tool_calls: vec![], + input_tokens: 10, + output_tokens: 4, + finish_reason: FinishReason::Stop, + response_id: None, + }) + } + } + + async fn list_models(&self) -> Result, LlmError> { + Ok(vec![ + "mock-model-v1".to_string(), + "mock-model-v2".to_string(), + ]) + } +} + +// --------------------------------------------------------------------------- +// Test helpers +// --------------------------------------------------------------------------- + +async fn start_test_server() -> (SocketAddr, Arc) { + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: SseManager::new(), + workspace: None, + session_manager: None, + log_broadcaster: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + user_id: "test-user".to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: Some(Arc::new(MockLlmProvider)), + chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + }); + + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + .await + .expect("Failed to start test server"); + + (bound_addr, state) +} + +fn client() -> reqwest::Client { + reqwest::Client::builder() + .timeout(Duration::from_secs(10)) + .build() + .unwrap() +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[tokio::test] +async fn test_chat_completions_basic() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [ + {"role": "user", "content": "Hello world"} + ] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["object"], "chat.completion"); + assert_eq!(body["model"], "mock-model-v1"); + assert_eq!(body["choices"][0]["finish_reason"], "stop"); + + let content = body["choices"][0]["message"]["content"].as_str().unwrap(); + assert!( + content.contains("Hello world"), + "Expected echo, got: {}", + content + ); + + // Check usage + assert_eq!(body["usage"]["prompt_tokens"], 10); + assert_eq!(body["usage"]["completion_tokens"], 5); + assert_eq!(body["usage"]["total_tokens"], 15); +} + +#[tokio::test] +async fn test_chat_completions_with_system_message() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "What is 2+2?"} + ], + "temperature": 0.5, + "max_tokens": 100 + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + let content = body["choices"][0]["message"]["content"].as_str().unwrap(); + assert!(content.contains("2+2")); +} + +#[tokio::test] +async fn test_chat_completions_with_tools() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [ + {"role": "user", "content": "What's the weather?"} + ], + "tools": [{ + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the weather", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"} + } + } + } + }] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + + assert_eq!(body["choices"][0]["finish_reason"], "tool_calls"); + + let tool_calls = &body["choices"][0]["message"]["tool_calls"]; + assert!(tool_calls.is_array()); + assert_eq!(tool_calls[0]["id"], "call_mock_001"); + assert_eq!(tool_calls[0]["type"], "function"); + assert_eq!(tool_calls[0]["function"]["name"], "get_weather"); +} + +#[tokio::test] +async fn test_chat_completions_streaming() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [ + {"role": "user", "content": "Stream test"} + ], + "stream": true + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + + // Check simulated streaming header + assert_eq!( + resp.headers() + .get("x-ironclaw-streaming") + .and_then(|v| v.to_str().ok()), + Some("simulated"), + "Expected x-ironclaw-streaming: simulated header" + ); + + let text = resp.text().await.unwrap(); + + // Should contain SSE data lines + assert!( + text.contains("data:"), + "Expected SSE data lines, got: {}", + text + ); + // Should end with [DONE] + assert!( + text.contains("[DONE]"), + "Expected [DONE] sentinel, got: {}", + text + ); + // Should contain the role chunk + assert!( + text.contains("\"role\":\"assistant\""), + "Expected role chunk, got: {}", + text + ); + + // Collect all content from the chunks + let mut full_content = String::new(); + for line in text.lines() { + if let Some(data) = line.strip_prefix("data:") { + let data = data.trim(); + if data == "[DONE]" { + continue; + } + if let Ok(chunk) = serde_json::from_str::(data) { + if let Some(content) = chunk["choices"][0]["delta"]["content"].as_str() { + full_content.push_str(content); + } + } + } + } + assert!( + full_content.contains("Stream test"), + "Expected reassembled content to contain 'Stream test', got: '{}'", + full_content + ); +} + +#[tokio::test] +async fn test_chat_completions_empty_messages() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 400); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!(body["error"]["message"].as_str().unwrap().contains("empty")); +} + +#[tokio::test] +async fn test_chat_completions_model_mismatch() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hi"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 404); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["error"]["code"], "model_not_found"); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("mock-model-v1") + ); +} + +#[tokio::test] +async fn test_chat_completions_no_auth() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + let resp = client() + .post(&url) + // No auth header + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [{"role": "user", "content": "Hi"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn test_models_endpoint() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/models", addr); + + let resp = client() + .get(&url) + .bearer_auth(AUTH_TOKEN) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + + assert_eq!(body["object"], "list"); + let data = body["data"].as_array().unwrap(); + assert_eq!(data.len(), 2); + assert_eq!(data[0]["id"], "mock-model-v1"); + assert_eq!(data[1]["id"], "mock-model-v2"); + assert_eq!(data[0]["object"], "model"); +} + +#[tokio::test] +async fn test_models_no_auth() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/models", addr); + + let resp = client().get(&url).send().await.unwrap(); + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn test_no_llm_provider_returns_503() { + // Create state WITHOUT llm_provider + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: SseManager::new(), + workspace: None, + session_manager: None, + log_broadcaster: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + user_id: "test-user".to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, // No LLM! + chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + }); + + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string()) + .await + .unwrap(); + + let url = format!("http://{}/v1/chat/completions", bound_addr); + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [{"role": "user", "content": "Hi"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 503); +} + +#[tokio::test] +async fn test_chat_completions_body_too_large() { + let (addr, _state) = start_test_server().await; + let url = format!("http://{}/v1/chat/completions", addr); + + // Build a payload over 1 MB (the gateway's DefaultBodyLimit) + let big_content = "x".repeat(2 * 1024 * 1024); + let resp = client() + .post(&url) + .bearer_auth(AUTH_TOKEN) + .json(&serde_json::json!({ + "model": "mock-model-v1", + "messages": [{"role": "user", "content": big_content}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 413); +} diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 6888a7f3..d2dbd756 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -51,6 +51,7 @@ async fn start_test_server() -> ( user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), });