diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index 6dce1f94..12020f10 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, per-request `model` override | | Canvas hosting | ✅ | ❌ | Agent-driven UI | | Gateway lock (PID-based) | ✅ | ❌ | | | launchd/systemd integration | ✅ | ❌ | | diff --git a/src/channels/web/openai_compat.rs b/src/channels/web/openai_compat.rs index b2dfa007..c493ef5c 100644 --- a/src/channels/web/openai_compat.rs +++ b/src/channels/web/openai_compat.rs @@ -24,6 +24,8 @@ use crate::llm::{ use super::server::GatewayState; +const MAX_MODEL_NAME_BYTES: usize = 256; + // --------------------------------------------------------------------------- // OpenAI request types // --------------------------------------------------------------------------- @@ -380,6 +382,27 @@ fn unix_timestamp() -> u64 { .as_secs() } +fn validate_model_name(model: &str) -> Result<(), String> { + let trimmed = model.trim(); + + if trimmed.is_empty() { + return Err("model must not be empty".to_string()); + } + if trimmed != model { + return Err("model must not have leading or trailing whitespace".to_string()); + } + if model.len() > MAX_MODEL_NAME_BYTES { + return Err(format!( + "model must be at most {} bytes", + MAX_MODEL_NAME_BYTES + )); + } + if model.chars().any(char::is_control) { + return Err("model contains control characters".to_string()); + } + Ok(()) +} + /// Extract stop sequences from the flexible `stop` field. fn parse_stop(val: &serde_json::Value) -> Option> { match val { @@ -426,29 +449,17 @@ pub async fn chat_completions_handler( "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()), - }, - }), + if let Err(e) = validate_model_name(&req.model) { + return Err(openai_error( + StatusCode::BAD_REQUEST, + e, + "invalid_request_error", )); } let has_tools = req.tools.as_ref().is_some_and(|t| !t.is_empty()); let stream = req.stream.unwrap_or(false); + let requested_model = req.model.clone(); if stream { return handle_streaming(llm.clone(), req, has_tools) @@ -460,13 +471,12 @@ pub async fn chat_completions_handler( 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); + let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model); if let Some(t) = req.temperature { tool_req = tool_req.with_temperature(t); } @@ -483,6 +493,7 @@ pub async fn chat_completions_handler( .complete_with_tools(tool_req) .await .map_err(map_llm_error)?; + let model_name = llm.effective_model_name(Some(requested_model.as_str())); let tool_calls_openai = if resp.tool_calls.is_empty() { None @@ -515,7 +526,7 @@ pub async fn chat_completions_handler( Ok(Json(response).into_response()) } else { - let mut comp_req = CompletionRequest::new(messages); + let mut comp_req = CompletionRequest::new(messages).with_model(req.model); if let Some(t) = req.temperature { comp_req = comp_req.with_temperature(t); } @@ -527,6 +538,7 @@ pub async fn chat_completions_handler( } let resp = llm.complete(comp_req).await.map_err(map_llm_error)?; + let model_name = llm.effective_model_name(Some(requested_model.as_str())); let response = OpenAiChatResponse { id, @@ -570,7 +582,7 @@ async fn handle_streaming( 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 requested_model = req.model.clone(); let id = chat_completion_id(); let created = unix_timestamp(); @@ -584,7 +596,7 @@ async fn handle_streaming( let llm_result = if has_tools { let tools = convert_tools(req.tools.as_deref().unwrap_or(&[])); - let mut tool_req = ToolCompletionRequest::new(messages, tools); + let mut tool_req = ToolCompletionRequest::new(messages, tools).with_model(req.model); if let Some(t) = req.temperature { tool_req = tool_req.with_temperature(t); } @@ -602,7 +614,7 @@ async fn handle_streaming( .map_err(map_llm_error)?, ) } else { - let mut comp_req = CompletionRequest::new(messages); + let mut comp_req = CompletionRequest::new(messages).with_model(req.model); if let Some(t) = req.temperature { comp_req = comp_req.with_temperature(t); } @@ -614,6 +626,7 @@ async fn handle_streaming( } LlmResult::Simple(llm.complete(comp_req).await.map_err(map_llm_error)?) }; + let model_name = llm.effective_model_name(Some(requested_model.as_str())); // LLM succeeded — emit the response as SSE chunks let (tx, rx) = tokio::sync::mpsc::channel::>(64); @@ -1091,4 +1104,18 @@ mod tests { let v = serde_json::Value::Null; assert_eq!(parse_stop(&v), None); } + + #[test] + fn test_validate_model_name_rejects_leading_or_trailing_whitespace() { + let err = validate_model_name(" gpt-4").unwrap_err(); + assert!(err.contains("leading or trailing whitespace")); + + let err = validate_model_name("gpt-4 ").unwrap_err(); + assert!(err.contains("leading or trailing whitespace")); + } + + #[test] + fn test_validate_model_name_accepts_normal_name() { + assert!(validate_model_name("gpt-4").is_ok()); + } } diff --git a/src/llm/circuit_breaker.rs b/src/llm/circuit_breaker.rs index 8f7718fd..12e46b30 100644 --- a/src/llm/circuit_breaker.rs +++ b/src/llm/circuit_breaker.rs @@ -273,6 +273,10 @@ impl LlmProvider for CircuitBreakerProvider { self.inner.model_metadata().await } + fn effective_model_name(&self, requested_model: Option<&str>) -> String { + self.inner.effective_model_name(requested_model) + } + fn active_model_name(&self) -> String { self.inner.active_model_name() } diff --git a/src/llm/failover.rs b/src/llm/failover.rs index a9cb9ed2..57836a3f 100644 --- a/src/llm/failover.rs +++ b/src/llm/failover.rs @@ -7,8 +7,10 @@ //! so subsequent requests skip them, reducing latency when a provider //! is known to be down. Cooldown state is lock-free (atomics only). +use std::collections::HashMap; use std::future::Future; use std::sync::Arc; +use std::sync::Mutex; use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering}; use std::time::{Duration, Instant}; @@ -139,6 +141,12 @@ pub struct FailoverProvider { epoch: Instant, /// Cooldown configuration. cooldown_config: CooldownConfig, + /// Request-scoped provider index keyed by Tokio task ID. + /// + /// This allows `effective_model_name()` to report the provider that handled + /// the *current* request, even when other concurrent requests update + /// `last_used`. + provider_for_task: Mutex>, } impl FailoverProvider { @@ -171,6 +179,7 @@ impl FailoverProvider { cooldowns, epoch: Instant::now(), cooldown_config, + provider_for_task: Mutex::new(HashMap::new()), }) } @@ -182,12 +191,36 @@ impl FailoverProvider { self.epoch.elapsed().as_nanos() as u64 } + /// Current Tokio task ID if available. + fn current_task_id() -> Option { + tokio::task::try_id() + } + + /// Bind the selected provider index to the current task. + fn bind_provider_to_current_task(&self, provider_idx: usize) { + let Some(task_id) = Self::current_task_id() else { + return; + }; + if let Ok(mut guard) = self.provider_for_task.lock() { + guard.insert(task_id, provider_idx); + } + } + + /// Take and remove the provider index bound to the current task. + fn take_bound_provider_for_current_task(&self) -> Option { + let task_id = Self::current_task_id()?; + self.provider_for_task + .lock() + .ok() + .and_then(|mut guard| guard.remove(&task_id)) + } + /// Try each provider in sequence until one succeeds or all fail. /// /// Providers in cooldown are skipped unless *all* providers are in /// cooldown, in which case the one with the oldest cooldown timestamp /// (most likely to have recovered) is tried. - async fn try_providers(&self, mut call: F) -> Result + async fn try_providers(&self, mut call: F) -> Result<(usize, T), LlmError> where F: FnMut(Arc) -> Fut, Fut: Future>, @@ -236,7 +269,7 @@ impl FailoverProvider { Ok(response) => { self.last_used.store(i, Ordering::Relaxed); self.cooldowns[i].reset(); - return Ok(response); + return Ok((i, response)); } Err(err) => { if !is_retryable(&err) { @@ -287,22 +320,28 @@ impl LlmProvider for FailoverProvider { } async fn complete(&self, request: CompletionRequest) -> Result { - self.try_providers(|provider| { - let req = request.clone(); - async move { provider.complete(req).await } - }) - .await + let (provider_idx, response) = self + .try_providers(|provider| { + let req = request.clone(); + async move { provider.complete(req).await } + }) + .await?; + self.bind_provider_to_current_task(provider_idx); + Ok(response) } async fn complete_with_tools( &self, request: ToolCompletionRequest, ) -> Result { - self.try_providers(|provider| { - let req = request.clone(); - async move { provider.complete_with_tools(req).await } - }) - .await + let (provider_idx, response) = self + .try_providers(|provider| { + let req = request.clone(); + async move { provider.complete_with_tools(req).await } + }) + .await?; + self.bind_provider_to_current_task(provider_idx); + Ok(response) } fn active_model_name(&self) -> String { @@ -336,6 +375,14 @@ impl LlmProvider for FailoverProvider { all_models.dedup(); Ok(all_models) } + + fn effective_model_name(&self, requested_model: Option<&str>) -> String { + if let Some(provider_idx) = self.take_bound_provider_for_current_task() { + return self.providers[provider_idx].effective_model_name(requested_model); + } + + self.providers[self.last_used.load(Ordering::Relaxed)].effective_model_name(requested_model) + } } #[cfg(test)] @@ -610,6 +657,49 @@ mod tests { assert_eq!(failover.cost_per_token(), (fallback_cost, fallback_cost)); } + // Test: model reporting is request-scoped under concurrent requests. + #[tokio::test] + async fn effective_model_name_is_request_scoped_under_concurrency() { + let config = CooldownConfig { + cooldown_duration: Duration::from_secs(60), + failure_threshold: 3, + }; + let primary = Arc::new(MultiCallMockProvider::fail_then_ok("primary", 1)); + let fallback = Arc::new(MultiCallMockProvider::always_ok("fallback")); + let failover = + Arc::new(FailoverProvider::with_cooldown(vec![primary, fallback], config).unwrap()); + + let (first_done_tx, first_done_rx) = tokio::sync::oneshot::channel::<()>(); + let (second_done_tx, second_done_rx) = tokio::sync::oneshot::channel::<()>(); + + let failover_a = Arc::clone(&failover); + let task_a = tokio::spawn(async move { + // First request: primary fails once, fallback serves. + let _ = failover_a.complete(make_request()).await.unwrap(); + let _ = first_done_tx.send(()); + + // Wait until the second request finishes and updates global state. + let _ = second_done_rx.await; + failover_a.effective_model_name(None) + }); + + let failover_b = Arc::clone(&failover); + let task_b = tokio::spawn(async move { + let _ = first_done_rx.await; + // Second request: primary now succeeds. + let _ = failover_b.complete(make_request()).await.unwrap(); + let model = failover_b.effective_model_name(None); + let _ = second_done_tx.send(()); + model + }); + + let model_b = task_b.await.unwrap(); + let model_a = task_a.await.unwrap(); + + assert_eq!(model_a, "fallback"); + assert_eq!(model_b, "primary"); + } + // Test: list_models aggregates from all providers. #[tokio::test] async fn list_models_aggregates_all() { diff --git a/src/llm/nearai.rs b/src/llm/nearai.rs index 2a995805..de58a8be 100644 --- a/src/llm/nearai.rs +++ b/src/llm/nearai.rs @@ -462,11 +462,12 @@ fn split_messages( #[async_trait] impl LlmProvider for NearAiProvider { async fn complete(&self, req: CompletionRequest) -> Result { + let model = req.model.unwrap_or_else(|| self.active_model_name()); let thread_id = req.metadata.get("thread_id").cloned(); let (instructions, input) = split_messages(req.messages, false); let request = NearAiRequest { - model: self.active_model_name(), + model, instructions, input, previous_response_id: None, @@ -579,6 +580,7 @@ impl LlmProvider for NearAiProvider { &self, req: ToolCompletionRequest, ) -> Result { + let model = req.model.unwrap_or_else(|| self.active_model_name()); let thread_id = req.metadata.get("thread_id").cloned(); // Look up chaining state for this thread @@ -619,7 +621,7 @@ impl LlmProvider for NearAiProvider { .collect(); let request = NearAiRequest { - model: self.active_model_name(), + model: model.clone(), instructions: if chaining { None } else { instructions.clone() }, input, previous_response_id: previous_response_id.clone(), @@ -660,7 +662,7 @@ impl LlmProvider for NearAiProvider { false, ); let retry_request = NearAiRequest { - model: self.active_model_name(), + model, instructions: instructions_full, input: input_full, previous_response_id: None, diff --git a/src/llm/nearai_chat.rs b/src/llm/nearai_chat.rs index dbe51bc7..2f97af36 100644 --- a/src/llm/nearai_chat.rs +++ b/src/llm/nearai_chat.rs @@ -227,11 +227,12 @@ struct ApiModelEntry { #[async_trait] impl LlmProvider for NearAiChatProvider { async fn complete(&self, req: CompletionRequest) -> Result { + let model = req.model.unwrap_or_else(|| self.active_model_name()); let messages: Vec = req.messages.into_iter().map(|m| m.into()).collect(); let request = ChatCompletionRequest { - model: self.active_model_name(), + model, messages, temperature: req.temperature, max_tokens: req.max_tokens, @@ -273,6 +274,7 @@ impl LlmProvider for NearAiChatProvider { &self, req: ToolCompletionRequest, ) -> Result { + let model = req.model.unwrap_or_else(|| self.active_model_name()); let messages: Vec = req.messages.into_iter().map(|m| m.into()).collect(); @@ -296,7 +298,7 @@ impl LlmProvider for NearAiChatProvider { .collect(); let request = ChatCompletionRequest { - model: self.active_model_name(), + model, messages, temperature: req.temperature, max_tokens: req.max_tokens, diff --git a/src/llm/provider.rs b/src/llm/provider.rs index e06d8b77..1c4e8510 100644 --- a/src/llm/provider.rs +++ b/src/llm/provider.rs @@ -105,6 +105,8 @@ impl ChatMessage { #[derive(Debug, Clone)] pub struct CompletionRequest { pub messages: Vec, + /// Optional per-request model override. + pub model: Option, pub max_tokens: Option, pub temperature: Option, pub stop_sequences: Option>, @@ -117,6 +119,7 @@ impl CompletionRequest { pub fn new(messages: Vec) -> Self { Self { messages, + model: None, max_tokens: None, temperature: None, stop_sequences: None, @@ -124,6 +127,12 @@ impl CompletionRequest { } } + /// Set model override. + pub fn with_model(mut self, model: impl Into) -> Self { + self.model = Some(model.into()); + self + } + /// Set max tokens. pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { self.max_tokens = Some(max_tokens); @@ -188,6 +197,8 @@ pub struct ToolResult { pub struct ToolCompletionRequest { pub messages: Vec, pub tools: Vec, + /// Optional per-request model override. + pub model: Option, pub max_tokens: Option, pub temperature: Option, /// How to handle tool use: "auto", "required", or "none". @@ -202,6 +213,7 @@ impl ToolCompletionRequest { Self { messages, tools, + model: None, max_tokens: None, temperature: None, tool_choice: None, @@ -209,6 +221,12 @@ impl ToolCompletionRequest { } } + /// Set model override. + pub fn with_model(mut self, model: impl Into) -> Self { + self.model = Some(model.into()); + self + } + /// Set max tokens. pub fn with_max_tokens(mut self, max_tokens: u32) -> Self { self.max_tokens = Some(max_tokens); @@ -283,6 +301,16 @@ pub trait LlmProvider: Send + Sync { }) } + /// Resolve which model should be reported for a given request. + /// + /// Providers that ignore per-request model overrides should override this + /// and return `active_model_name()`. + fn effective_model_name(&self, requested_model: Option<&str>) -> String { + requested_model + .map(std::borrow::ToOwned::to_owned) + .unwrap_or_else(|| self.active_model_name()) + } + /// Get the currently active model name. /// /// May differ from `model_name()` if the model was switched at runtime diff --git a/src/llm/response_cache.rs b/src/llm/response_cache.rs index 0f8468df..61437e6f 100644 --- a/src/llm/response_cache.rs +++ b/src/llm/response_cache.rs @@ -144,7 +144,8 @@ impl LlmProvider for CachedProvider { } async fn complete(&self, request: CompletionRequest) -> Result { - let key = cache_key(self.inner.model_name(), &request); + let effective_model = self.inner.effective_model_name(request.model.as_deref()); + let key = cache_key(&effective_model, &request); let now = Instant::now(); // Check cache @@ -216,6 +217,10 @@ impl LlmProvider for CachedProvider { self.inner.model_metadata().await } + fn effective_model_name(&self, requested_model: Option<&str>) -> String { + self.inner.effective_model_name(requested_model) + } + fn active_model_name(&self) -> String { self.inner.active_model_name() } @@ -242,6 +247,7 @@ mod tests { fn simple_request() -> CompletionRequest { CompletionRequest { messages: vec![ChatMessage::user("hello")], + model: None, max_tokens: None, temperature: None, stop_sequences: None, @@ -252,6 +258,7 @@ mod tests { fn different_request() -> CompletionRequest { CompletionRequest { messages: vec![ChatMessage::user("goodbye")], + model: None, max_tokens: None, temperature: None, stop_sequences: None, @@ -378,6 +385,7 @@ mod tests { // Add a third: should evict the oldest let third = CompletionRequest { messages: vec![ChatMessage::user("third")], + model: None, max_tokens: None, temperature: None, stop_sequences: None, @@ -396,6 +404,7 @@ mod tests { let req = ToolCompletionRequest { messages: vec![ChatMessage::user("use tool")], tools: vec![], + model: None, max_tokens: None, temperature: None, tool_choice: None, @@ -444,6 +453,23 @@ mod tests { assert!(cached.is_empty().await); } + #[tokio::test] + async fn model_override_gets_distinct_cache_entries() { + let stub = Arc::new(StubLlm::new("cached response")); + let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default()); + + let mut req_a = simple_request(); + req_a.model = Some("model-a".to_string()); + let mut req_b = simple_request(); + req_b.model = Some("model-b".to_string()); + + cached.complete(req_a).await.unwrap(); + cached.complete(req_b).await.unwrap(); + + assert_eq!(stub.calls(), 2); + assert_eq!(cached.len().await, 2); + } + #[test] fn default_config_is_reasonable() { let cfg = ResponseCacheConfig::default(); diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index 7a520989..ce2b84af 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -404,6 +404,16 @@ where } async fn complete(&self, request: CompletionRequest) -> Result { + if let Some(requested_model) = request.model.as_deref() + && requested_model != self.model_name.as_str() + { + tracing::warn!( + requested_model = requested_model, + active_model = %self.model_name, + "Per-request model override is not supported for this provider; using configured model" + ); + } + let (preamble, history) = convert_messages(&request.messages); let rig_req = build_rig_request( @@ -439,6 +449,16 @@ where &self, request: ToolCompletionRequest, ) -> Result { + if let Some(requested_model) = request.model.as_deref() + && requested_model != self.model_name.as_str() + { + tracing::warn!( + requested_model = requested_model, + active_model = %self.model_name, + "Per-request model override is not supported for this provider; using configured model" + ); + } + let (preamble, history) = convert_messages(&request.messages); let tools = convert_tools(&request.tools); let tool_choice = convert_tool_choice(request.tool_choice.as_deref()); @@ -477,6 +497,10 @@ where self.model_name.clone() } + fn effective_model_name(&self, _requested_model: Option<&str>) -> String { + self.active_model_name() + } + fn set_model(&self, _model: &str) -> Result<(), LlmError> { // rig-core models are baked at construction time. // Switching requires creating a new adapter. diff --git a/src/main.rs b/src/main.rs index a7b3e2af..dc306bce 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1305,7 +1305,7 @@ async fn main() -> anyhow::Result<()> { // Add web gateway channel if configured let mut gateway_url: Option = None; if let Some(ref gw_config) = config.channels.gateway { - let mut gw = GatewayChannel::new(gw_config.clone()); + let mut gw = GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&llm)); if let Some(ref ws) = workspace { gw = gw.with_workspace(Arc::clone(ws)); } diff --git a/src/orchestrator/api.rs b/src/orchestrator/api.rs index 8138dc1c..45803e4e 100644 --- a/src/orchestrator/api.rs +++ b/src/orchestrator/api.rs @@ -143,6 +143,7 @@ async fn llm_complete( ) -> Result, StatusCode> { let completion_req = CompletionRequest { messages: req.messages, + model: req.model, max_tokens: req.max_tokens, temperature: req.temperature, stop_sequences: req.stop_sequences, @@ -170,6 +171,7 @@ async fn llm_complete_with_tools( let tool_req = ToolCompletionRequest { messages: req.messages, tools: req.tools, + model: req.model, max_tokens: req.max_tokens, temperature: req.temperature, tool_choice: req.tool_choice, diff --git a/src/worker/api.rs b/src/worker/api.rs index 1765d66b..292453fe 100644 --- a/src/worker/api.rs +++ b/src/worker/api.rs @@ -40,6 +40,7 @@ pub struct JobDescription { #[derive(Debug, Serialize, Deserialize)] pub struct ProxyCompletionRequest { pub messages: Vec, + pub model: Option, pub max_tokens: Option, pub temperature: Option, pub stop_sequences: Option>, @@ -57,6 +58,7 @@ pub struct ProxyCompletionResponse { pub struct ProxyToolCompletionRequest { pub messages: Vec, pub tools: Vec, + pub model: Option, pub max_tokens: Option, pub temperature: Option, pub tool_choice: Option, @@ -210,6 +212,7 @@ impl WorkerHttpClient { ) -> Result { let proxy_req = ProxyCompletionRequest { messages: request.messages.clone(), + model: request.model.clone(), max_tokens: request.max_tokens, temperature: request.temperature, stop_sequences: request.stop_sequences.clone(), @@ -236,6 +239,7 @@ impl WorkerHttpClient { let proxy_req = ProxyToolCompletionRequest { messages: request.messages.clone(), tools: request.tools.clone(), + model: request.model.clone(), max_tokens: request.max_tokens, temperature: request.temperature, tool_choice: request.tool_choice.clone(), diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index 67cbe7b4..a4636c4f 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -24,7 +24,21 @@ const AUTH_TOKEN: &str = "test-openai-token"; // Mock LLM provider // --------------------------------------------------------------------------- -struct MockLlmProvider; +#[derive(Default)] +struct MockLlmState { + completion_models: tokio::sync::Mutex>>, + tool_completion_models: tokio::sync::Mutex>>, +} + +struct MockLlmProvider { + state: Arc, +} + +impl MockLlmProvider { + fn new(state: Arc) -> Self { + Self { state } + } +} #[async_trait] impl LlmProvider for MockLlmProvider { @@ -37,6 +51,12 @@ impl LlmProvider for MockLlmProvider { } async fn complete(&self, req: CompletionRequest) -> Result { + self.state + .completion_models + .lock() + .await + .push(req.model.clone()); + // Echo the last user message back let user_msg = req .messages @@ -59,6 +79,12 @@ impl LlmProvider for MockLlmProvider { &self, req: ToolCompletionRequest, ) -> Result { + self.state + .tool_completion_models + .lock() + .await + .push(req.model.clone()); + // If tools are provided, return a tool call if let Some(tool) = req.tools.first() { Ok(ToolCompletionResponse { @@ -93,11 +119,71 @@ impl LlmProvider for MockLlmProvider { } } +struct FixedModelProvider { + model: &'static str, +} + +impl FixedModelProvider { + fn new(model: &'static str) -> Self { + Self { model } + } +} + +#[async_trait] +impl LlmProvider for FixedModelProvider { + fn model_name(&self) -> &str { + self.model + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (Decimal::ZERO, Decimal::ZERO) + } + + async fn complete(&self, _req: CompletionRequest) -> Result { + Ok(CompletionResponse { + content: "fixed response".to_string(), + input_tokens: 10, + output_tokens: 5, + finish_reason: FinishReason::Stop, + response_id: None, + }) + } + + async fn complete_with_tools( + &self, + _req: ToolCompletionRequest, + ) -> Result { + Ok(ToolCompletionResponse { + content: Some("fixed response".to_string()), + tool_calls: vec![], + input_tokens: 10, + output_tokens: 5, + finish_reason: FinishReason::Stop, + response_id: None, + }) + } + + fn effective_model_name(&self, _requested_model: Option<&str>) -> String { + self.model.to_string() + } +} + // --------------------------------------------------------------------------- // Test helpers // --------------------------------------------------------------------------- -async fn start_test_server() -> (SocketAddr, Arc) { +async fn start_test_server() -> (SocketAddr, Arc, Arc) { + let mock_state = Arc::new(MockLlmState::default()); + + let llm_provider: Arc = Arc::new(MockLlmProvider::new(mock_state.clone())); + let (bound_addr, state) = start_test_server_with_provider(llm_provider).await; + + (bound_addr, state, mock_state) +} + +async fn start_test_server_with_provider( + llm_provider: Arc, +) -> (SocketAddr, Arc) { let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), sse: SseManager::new(), @@ -112,7 +198,7 @@ async fn start_test_server() -> (SocketAddr, Arc) { 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)), + llm_provider: Some(llm_provider), skill_registry: None, skill_catalog: None, chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), @@ -139,7 +225,7 @@ fn client() -> reqwest::Client { #[tokio::test] async fn test_chat_completions_basic() { - let (addr, _state) = start_test_server().await; + let (addr, _state, mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -173,11 +259,14 @@ async fn test_chat_completions_basic() { assert_eq!(body["usage"]["prompt_tokens"], 10); assert_eq!(body["usage"]["completion_tokens"], 5); assert_eq!(body["usage"]["total_tokens"], 15); + + let models = mock_state.completion_models.lock().await; + assert_eq!(*models, vec![Some("mock-model-v1".to_string())]); } #[tokio::test] async fn test_chat_completions_with_system_message() { - let (addr, _state) = start_test_server().await; + let (addr, _state, _mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -204,7 +293,7 @@ async fn test_chat_completions_with_system_message() { #[tokio::test] async fn test_chat_completions_with_tools() { - let (addr, _state) = start_test_server().await; + let (addr, _state, mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -243,11 +332,14 @@ async fn test_chat_completions_with_tools() { 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"); + + let models = mock_state.tool_completion_models.lock().await; + assert_eq!(*models, vec![Some("mock-model-v1".to_string())]); } #[tokio::test] async fn test_chat_completions_streaming() { - let (addr, _state) = start_test_server().await; + let (addr, _state, mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -316,11 +408,14 @@ async fn test_chat_completions_streaming() { "Expected reassembled content to contain 'Stream test', got: '{}'", full_content ); + + let models = mock_state.completion_models.lock().await; + assert_eq!(*models, vec![Some("mock-model-v1".to_string())]); } #[tokio::test] async fn test_chat_completions_empty_messages() { - let (addr, _state) = start_test_server().await; + let (addr, _state, _mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -340,8 +435,8 @@ async fn test_chat_completions_empty_messages() { } #[tokio::test] -async fn test_chat_completions_model_mismatch() { - let (addr, _state) = start_test_server().await; +async fn test_chat_completions_model_override() { + let (addr, _state, mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -355,20 +450,173 @@ async fn test_chat_completions_model_mismatch() { .await .unwrap(); - assert_eq!(resp.status(), 404); + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["model"], "gpt-4"); + + let models = mock_state.completion_models.lock().await; + assert_eq!(*models, vec![Some("gpt-4".to_string())]); +} + +#[tokio::test] +async fn test_chat_completions_uses_effective_model_when_override_ignored() { + let provider: Arc = Arc::new(FixedModelProvider::new("configured-model")); + let (addr, _state) = start_test_server_with_provider(provider).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(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["model"], "configured-model"); +} + +#[tokio::test] +async fn test_chat_completions_streaming_uses_effective_model_when_override_ignored() { + let provider: Arc = Arc::new(FixedModelProvider::new("configured-model")); + let (addr, _state) = start_test_server_with_provider(provider).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"}], + "stream": true + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let text = resp.text().await.unwrap(); + assert!( + text.contains("\"model\":\"configured-model\""), + "Expected streaming chunks to report configured model, got: {}", + text + ); +} + +#[tokio::test] +async fn test_chat_completions_model_too_long() { + let (addr, _state, mock_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": "m".repeat(300), + "messages": [{"role": "user", "content": "Hi"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 400); 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") + .unwrap_or("") + .contains("model"), + "Expected model validation error, got: {}", + body + ); + + // Validation should fail before provider invocation. + let models = mock_state.completion_models.lock().await; + assert!( + models.is_empty(), + "provider should not be called: {:?}", + *models + ); +} + +#[tokio::test] +async fn test_chat_completions_model_with_control_chars() { + let (addr, _state, mock_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\noops", + "messages": [{"role": "user", "content": "Hi"}] + })) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 400); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"]["message"] + .as_str() + .unwrap_or("") + .contains("control"), + "Expected model validation error, got: {}", + body + ); + + // Validation should fail before provider invocation. + let models = mock_state.completion_models.lock().await; + assert!( + models.is_empty(), + "provider should not be called: {:?}", + *models + ); +} + +#[tokio::test] +async fn test_chat_completions_model_with_surrounding_whitespace() { + let (addr, _state, mock_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(), 400); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"]["message"] + .as_str() + .unwrap_or("") + .contains("leading or trailing whitespace"), + "Expected model validation error, got: {}", + body + ); + + let models = mock_state.completion_models.lock().await; + assert!( + models.is_empty(), + "provider should not be called: {:?}", + *models ); } #[tokio::test] async fn test_chat_completions_no_auth() { - let (addr, _state) = start_test_server().await; + let (addr, _state, _mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); let resp = client() @@ -387,7 +635,7 @@ async fn test_chat_completions_no_auth() { #[tokio::test] async fn test_models_endpoint() { - let (addr, _state) = start_test_server().await; + let (addr, _state, _mock_state) = start_test_server().await; let url = format!("http://{}/v1/models", addr); let resp = client() @@ -410,7 +658,7 @@ async fn test_models_endpoint() { #[tokio::test] async fn test_models_no_auth() { - let (addr, _state) = start_test_server().await; + let (addr, _state, _mock_state) = start_test_server().await; let url = format!("http://{}/v1/models", addr); let resp = client().get(&url).send().await.unwrap(); @@ -462,7 +710,7 @@ async fn test_no_llm_provider_returns_503() { #[tokio::test] async fn test_chat_completions_body_too_large() { - let (addr, _state) = start_test_server().await; + let (addr, _state, _mock_state) = start_test_server().await; let url = format!("http://{}/v1/chat/completions", addr); // Build a payload over 1 MB (the gateway's DefaultBodyLimit)