From 668d8b787ee0605a4f053e069026c6b934b40093 Mon Sep 17 00:00:00 2001 From: Georg Bauer Date: Sat, 25 Jul 2026 14:18:57 +0200 Subject: [PATCH] Complete local HTTP endpoint parity --- .../down.sql | 2 + .../up.sql | 4 + src/app.rs | 34 +++- src/app/preferences.rs | 27 ++- src/app/view/preferences.rs | 12 +- src/database.rs | 30 ++- src/engine.rs | 2 +- src/runtime.rs | 7 + src/schema.rs | 2 + src/server.rs | 186 +++++++++++++++--- src/server/http.rs | 43 +++- src/server/response.rs | 106 ++++++++-- src/server/tools.rs | 107 +++++++++- 13 files changed, 493 insertions(+), 69 deletions(-) create mode 100644 migrations/20260725130000_add_endpoint_controls/down.sql create mode 100644 migrations/20260725130000_add_endpoint_controls/up.sql diff --git a/migrations/20260725130000_add_endpoint_controls/down.sql b/migrations/20260725130000_add_endpoint_controls/down.sql new file mode 100644 index 0000000..1076b13 --- /dev/null +++ b/migrations/20260725130000_add_endpoint_controls/down.sql @@ -0,0 +1,2 @@ +ALTER TABLE preferences DROP COLUMN endpoint_cors; +ALTER TABLE preferences DROP COLUMN endpoint_enabled; diff --git a/migrations/20260725130000_add_endpoint_controls/up.sql b/migrations/20260725130000_add_endpoint_controls/up.sql new file mode 100644 index 0000000..9859036 --- /dev/null +++ b/migrations/20260725130000_add_endpoint_controls/up.sql @@ -0,0 +1,4 @@ +ALTER TABLE preferences ADD COLUMN endpoint_enabled BOOLEAN NOT NULL DEFAULT TRUE + CHECK (endpoint_enabled IN (FALSE, TRUE)); +ALTER TABLE preferences ADD COLUMN endpoint_cors BOOLEAN NOT NULL DEFAULT FALSE + CHECK (endpoint_cors IN (FALSE, TRUE)); diff --git a/src/app.rs b/src/app.rs index 3c2f603..0a9ddbc 100644 --- a/src/app.rs +++ b/src/app.rs @@ -123,6 +123,8 @@ pub(crate) enum Message { PreferenceDsparkChanged(bool), PreferenceTimeoutChanged(String), PreferenceEndpointPortChanged(String), + PreferenceEndpointEnabledChanged(bool), + PreferenceEndpointCorsChanged(bool), PreferenceContextChanged(String), PreferenceMaxTokensChanged(String), PreferenceSystemPromptChanged(String), @@ -429,6 +431,14 @@ impl App { self.preference_draft.endpoint_port = value; self.preference_error = None; } + Message::PreferenceEndpointEnabledChanged(enabled) => { + self.preference_draft.endpoint_enabled = enabled; + self.preference_error = None; + } + Message::PreferenceEndpointCorsChanged(enabled) => { + self.preference_draft.endpoint_cors = enabled; + self.preference_error = None; + } Message::PreferenceContextChanged(value) => { self.preference_draft.context_tokens = value; self.preference_error = None; @@ -970,17 +980,21 @@ fn spawn_services( Ok(generation) => generation, Err(error) => return (runtime_preferences, None, None, Some(error)), }; - let endpoint = crate::server::ServerHandle::spawn( - generation.clone(), - Arc::clone(&runtime_preferences), - models_path(), - transient_cache_path(), - u16::try_from(preferences.endpoint_port).unwrap_or(4000), - metrics, - ); + let endpoint = preferences.endpoint_enabled.then(|| { + crate::server::ServerHandle::spawn( + generation.clone(), + Arc::clone(&runtime_preferences), + models_path(), + transient_cache_path(), + u16::try_from(preferences.endpoint_port).unwrap_or(4000), + preferences.endpoint_cors, + metrics, + ) + }); match endpoint { - Ok(endpoint) => (runtime_preferences, Some(generation), Some(endpoint), None), - Err(error) => (runtime_preferences, Some(generation), None, Some(error)), + Some(Ok(endpoint)) => (runtime_preferences, Some(generation), Some(endpoint), None), + Some(Err(error)) => (runtime_preferences, Some(generation), None, Some(error)), + None => (runtime_preferences, Some(generation), None, None), } } diff --git a/src/app/preferences.rs b/src/app/preferences.rs index f049a7a..9ed3721 100644 --- a/src/app/preferences.rs +++ b/src/app/preferences.rs @@ -6,6 +6,8 @@ pub(super) struct PreferenceDraft { pub(super) dspark_enabled: bool, pub(super) idle_timeout_minutes: String, pub(super) endpoint_port: String, + pub(super) endpoint_enabled: bool, + pub(super) endpoint_cors: bool, pub(super) context_tokens: String, pub(super) max_generated_tokens: String, pub(super) system_prompt: String, @@ -51,6 +53,8 @@ impl PreferenceDraft { dspark_enabled: speculative.dspark_enabled, idle_timeout_minutes: preferences.idle_timeout_minutes.to_string(), endpoint_port: preferences.endpoint_port.to_string(), + endpoint_enabled: preferences.endpoint_enabled, + endpoint_cors: preferences.endpoint_cors, context_tokens: generation.context_tokens.to_string(), max_generated_tokens: generation.max_generated_tokens.to_string(), system_prompt: generation.system_prompt, @@ -147,6 +151,8 @@ impl PreferenceDraft { self.dspark_enabled = false; self.idle_timeout_minutes = "10".into(); self.endpoint_port = "4000".into(); + self.endpoint_enabled = true; + self.endpoint_cors = false; self.context_tokens = defaults.context_tokens.to_string(); self.max_generated_tokens = defaults.max_generated_tokens.to_string(); self.system_prompt = defaults.system_prompt; @@ -429,8 +435,18 @@ impl App { return; } #[cfg(target_os = "macos")] - let pending_endpoint = if self.preferences.endpoint_port != i32::from(endpoint_port) - || self._endpoint.is_none() + let endpoint_changed = self.preferences.endpoint_port != i32::from(endpoint_port) + || self.preferences.endpoint_enabled != self.preference_draft.endpoint_enabled + || self.preferences.endpoint_cors != self.preference_draft.endpoint_cors; + #[cfg(target_os = "macos")] + if endpoint_changed + && self.preferences.endpoint_port == i32::from(endpoint_port) + && self._endpoint.is_some() + { + self._endpoint = None; + } + let pending_endpoint = if self.preference_draft.endpoint_enabled + && (endpoint_changed || self._endpoint.is_none()) { let Some(generation) = &self.generation_service else { self.preference_error = Some("The model runtime is unavailable.".into()); @@ -442,6 +458,7 @@ impl App { models_path(), application_support_path().join("kv-cache").join("http"), endpoint_port, + self.preference_draft.endpoint_cors, Arc::clone(&self.metrics), ) { Ok(endpoint) => Some(endpoint), @@ -460,6 +477,8 @@ impl App { model.id(), idle_timeout_minutes, i32::from(endpoint_port), + self.preference_draft.endpoint_enabled, + self.preference_draft.endpoint_cors, &generation, &runtime, ) { @@ -468,7 +487,9 @@ impl App { #[cfg(target_os = "macos")] update_runtime_preferences(&self.runtime_preferences, &self.preferences); #[cfg(target_os = "macos")] - if let Some(endpoint) = pending_endpoint { + if endpoint_changed { + self._endpoint = pending_endpoint; + } else if let Some(endpoint) = pending_endpoint { self._endpoint = Some(endpoint); } self.preference_draft = PreferenceDraft::from_saved(&self.preferences) diff --git a/src/app/view/preferences.rs b/src/app/view/preferences.rs index 95fd4e4..e1c15d0 100644 --- a/src/app/view/preferences.rs +++ b/src/app/view/preferences.rs @@ -107,12 +107,22 @@ impl App { let endpoint_group = preference_group( "LOCAL ENDPOINT", column![ + checkbox( + "Enable OpenAI-compatible endpoint", + self.preference_draft.endpoint_enabled, + ) + .on_toggle(Message::PreferenceEndpointEnabledChanged), preference_input_row( "Port", text_input("4000", &self.preference_draft.endpoint_port) .on_input(Message::PreferenceEndpointPortChanged), ), - text("Listens on 127.0.0.1. Saving a changed port restarts the local endpoint.") + checkbox( + "Allow browser clients (CORS)", + self.preference_draft.endpoint_cors, + ) + .on_toggle(Message::PreferenceEndpointCorsChanged), + text("Listens only on 127.0.0.1. Saving changed endpoint settings restarts it.") .size(12), ] .spacing(10), diff --git a/src/database.rs b/src/database.rs index 9dabd2c..cf9471d 100644 --- a/src/database.rs +++ b/src/database.rs @@ -52,6 +52,8 @@ pub struct AppPreferences { pub simulated_used_memory_gib: Option, pub expert_profile_path: Option, pub endpoint_port: i32, + pub endpoint_enabled: bool, + pub endpoint_cors: bool, } impl Default for AppPreferences { @@ -92,6 +94,8 @@ impl Default for AppPreferences { simulated_used_memory_gib: None, expert_profile_path: None, endpoint_port: 4000, + endpoint_enabled: true, + endpoint_cors: false, } } } @@ -237,6 +241,8 @@ struct PreferenceChanges<'a> { simulated_used_memory_gib: Option, expert_profile_path: Option<&'a str>, endpoint_port: i32, + endpoint_enabled: bool, + endpoint_cors: bool, } #[derive(Clone, Debug, Identifiable, Queryable, Selectable)] @@ -423,11 +429,14 @@ impl Database { .map_err(|error| error.to_string()) } + #[allow(clippy::too_many_arguments)] pub fn update_preferences( &mut self, selected_model: &str, idle_timeout_minutes: i32, endpoint_port: i32, + endpoint_enabled: bool, + endpoint_cors: bool, generation: &GenerationPreferences, runtime: &RuntimePreferences, ) -> Result { @@ -488,6 +497,8 @@ impl Database { .map(|value| value as i64), expert_profile_path: runtime.diagnostics.expert_profile_path.as_deref(), endpoint_port, + endpoint_enabled, + endpoint_cors, }) .returning(AppPreferences::as_returning()) .get_result(&mut self.connection) @@ -669,6 +680,8 @@ mod tests { assert!(!preferences.dspark_enabled); assert_eq!(preferences.idle_timeout_minutes, 10); assert_eq!(preferences.endpoint_port, 4000); + assert!(preferences.endpoint_enabled); + assert!(!preferences.endpoint_cors); let generation = GenerationPreferences::default(); let runtime = RuntimePreferences::default(); assert!( @@ -677,6 +690,8 @@ mod tests { "glm-5.2", 30, 4000, + true, + false, &generation, &RuntimePreferences { speculative: SpeculativePreferences { @@ -690,7 +705,15 @@ mod tests { ); assert!( database - .update_preferences("deepseek-v4-flash", 0, 4000, &generation, &runtime,) + .update_preferences( + "deepseek-v4-flash", + 0, + 4000, + true, + false, + &generation, + &runtime, + ) .is_err() ); let generation = GenerationPreferences { @@ -728,8 +751,11 @@ mod tests { ..RuntimePreferences::default() }; database - .update_preferences("glm-5.2", 30, 4567, &generation, &runtime) + .update_preferences("glm-5.2", 30, 4567, false, true, &generation, &runtime) .unwrap(); + let preferences = database.load_preferences().unwrap(); + assert!(!preferences.endpoint_enabled); + assert!(preferences.endpoint_cors); let project = database.create_project("DS4", "/tmp/ds4").unwrap(); let first = database diff --git a/src/engine.rs b/src/engine.rs index 9ac0811..ad01fc1 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -560,7 +560,7 @@ impl Generator { let max_context = self.executor.context() as usize; if tokens.len() >= max_context { return Err(format!( - "the conversation uses {} tokens; the configured context holds fewer than {max_context}", + "Prompt has {} tokens, but the configured context size is {max_context} tokens", tokens.len() )); } diff --git a/src/runtime.rs b/src/runtime.rs index 341d27a..6d89d56 100644 --- a/src/runtime.rs +++ b/src/runtime.rs @@ -165,6 +165,13 @@ fn run_command( idle_timeout: &mut Duration, ) { *idle_timeout = command.idle_timeout; + if command.cancel.load(Ordering::Relaxed) { + metrics.request_failed(request_started.elapsed()); + let _ = command.events.send(GenerationEvent::Finished( + Err("generation cancelled".into()), + )); + return; + } if loaded .as_ref() .is_none_or(|(settings, _)| settings != &command.engine) diff --git a/src/schema.rs b/src/schema.rs index ec0c124..de135a0 100644 --- a/src/schema.rs +++ b/src/schema.rs @@ -54,6 +54,8 @@ diesel::table! { simulated_used_memory_gib -> Nullable, expert_profile_path -> Nullable, endpoint_port -> Integer, + endpoint_enabled -> Bool, + endpoint_cors -> Bool, } } diff --git a/src/server.rs b/src/server.rs index 70d0ebd..c844322 100644 --- a/src/server.rs +++ b/src/server.rs @@ -4,8 +4,8 @@ mod response; mod tools; use http::{ - random_id, random_tool_id, read_request, send_error, send_json, send_response, send_sse, - send_sse_error, send_sse_headers, unix_time, + random_id, random_tool_id, read_request, send_error_with_cors, send_json_with_cors, + send_response_with_cors, send_sse, send_sse_error, send_sse_headers_with_cors, unix_time, }; use request::{ anthropic_request, completion_request, parse_chat_request, raw_tool_schemas, responses_request, @@ -15,7 +15,7 @@ use response::{final_response, stream_response}; use response::{receive_stream_event, send_stream_start}; use tools::{ TOOL_SYNTAXES, ToolProjectionEvent, ToolProjector, content_text, parse_generated_tools, - parse_generated_tools_with_ids, render_messages, tool_calls_json, + parse_generated_tools_with_ids, read_tool_memory, render_messages, tool_calls_json, }; #[cfg(test)] use tools::{canonical_tools, validate_tool_results}; @@ -77,6 +77,7 @@ struct State { connections: Arc, tool_memory: Mutex>, metrics: Arc, + cors: bool, } #[derive(Deserialize)] @@ -179,6 +180,7 @@ struct ResponseOptions { stream: bool, include_usage: bool, has_tools: bool, + cors: bool, } #[derive(Clone, Copy, PartialEq, Eq)] @@ -214,6 +216,7 @@ impl ServerHandle { models_path: PathBuf, cache_path: PathBuf, port: u16, + cors: bool, metrics: Arc, ) -> Result { let address = format!("127.0.0.1:{port}"); @@ -224,6 +227,7 @@ impl ServerHandle { .map_err(|error| format!("Could not configure http://{address}: {error}"))?; let stop = Arc::new(AtomicBool::new(false)); let worker_stop = Arc::clone(&stop); + let tool_memory = read_tool_memory(&cache_path); let state = Arc::new(State { generation, preferences, @@ -231,8 +235,9 @@ impl ServerHandle { cache_path, sequence: AtomicU64::new(0), connections: Arc::new(AtomicUsize::new(0)), - tool_memory: Mutex::new(HashMap::new()), + tool_memory: Mutex::new(tool_memory), metrics: Arc::clone(&metrics), + cors, }); let thread = thread::Builder::new() .name("local-http".into()) @@ -314,7 +319,7 @@ fn handle(mut stream: TcpStream, state: &State) { Ok(request) => request, Err(error) => { state.metrics.http_started("", 0); - let _ = send_error(&mut stream, 400, &error); + let _ = send_error_with_cors(&mut stream, 400, &error, state.cors); state.metrics.http_finished(started.elapsed(), true); return; } @@ -323,16 +328,16 @@ fn handle(mut stream: TcpStream, state: &State) { .metrics .http_started(&request.path, request.body.len()); let result = match (request.method.as_str(), request.path.as_str()) { - ("OPTIONS", _) => send_response(&mut stream, 204, None, &[]), + ("OPTIONS", _) => send_response_with_cors(&mut stream, 204, None, &[], state.cors), ("GET", "/v1/models") => { let body = models_json(state); - send_json(&mut stream, 200, &body) + send_json_with_cors(&mut stream, 200, &body, state.cors) } ("POST", "/v1/chat/completions") => { match chat_completion(&mut stream, state, &request.body) { Ok(()) => Ok(()), Err(error) => { - let _ = send_error(&mut stream, error.0, &error.1); + let _ = send_error_with_cors(&mut stream, error.0, &error.1, state.cors); Err(error.1) } } @@ -341,7 +346,7 @@ fn handle(mut stream: TcpStream, state: &State) { match compatible_completion(&mut stream, state, &request.body, Protocol::Completion) { Ok(()) => Ok(()), Err(error) => { - let _ = send_error(&mut stream, error.0, &error.1); + let _ = send_error_with_cors(&mut stream, error.0, &error.1, state.cors); Err(error.1) } } @@ -350,7 +355,7 @@ fn handle(mut stream: TcpStream, state: &State) { match compatible_completion(&mut stream, state, &request.body, Protocol::Anthropic) { Ok(()) => Ok(()), Err(error) => { - let _ = send_error(&mut stream, error.0, &error.1); + let _ = send_error_with_cors(&mut stream, error.0, &error.1, state.cors); Err(error.1) } } @@ -359,29 +364,35 @@ fn handle(mut stream: TcpStream, state: &State) { match compatible_completion(&mut stream, state, &request.body, Protocol::Responses) { Ok(()) => Ok(()), Err(error) => { - let _ = send_error(&mut stream, error.0, &error.1); + let _ = send_error_with_cors(&mut stream, error.0, &error.1, state.cors); Err(error.1) } } } ("GET", path) if path.starts_with("/v1/models/") => { let id = &path[11..]; - let model = installed_endpoint_models(&state.models_path) - .into_iter() - .find(|model| model.id() == id); + let installed = installed_endpoint_models(&state.models_path); + let model = model_alias(id).filter(|model| installed.contains(model)); if let Some(model) = model { - let context = state + let (context, tokens) = state .preferences .read() - .map_or(32_768, |preferences| preferences.context_tokens); - send_json(&mut stream, 200, &model_json(model, context)) + .map_or((32_768, 32_768), |preferences| { + (preferences.context_tokens, preferences.max_generated_tokens) + }); + send_json_with_cors( + &mut stream, + 200, + &model_json(id, model, context, tokens), + state.cors, + ) } else { - let _ = send_error(&mut stream, 404, "unknown model"); + let _ = send_error_with_cors(&mut stream, 404, "unknown model", state.cors); Err("unknown model".into()) } } _ => { - let _ = send_error(&mut stream, 404, "unknown endpoint"); + let _ = send_error_with_cors(&mut stream, 404, "unknown endpoint", state.cors); Err("unknown endpoint".into()) } }; @@ -409,6 +420,7 @@ fn chat_completion( stream: parsed.stream, include_usage: parsed.include_usage, has_tools: parsed.has_tools, + cors: state.cors, }; let active = state .generation @@ -460,6 +472,7 @@ fn compatible_completion( stream: parsed.stream, include_usage: parsed.include_usage, has_tools: parsed.has_tools, + cors: state.cors, }; let active = state .generation @@ -493,22 +506,24 @@ fn installed_endpoint_models(models_path: &std::path::Path) -> Vec } fn models_json(state: &State) -> Value { - let context = state + let (context, tokens) = state .preferences .read() - .map_or(32_768, |preferences| preferences.context_tokens); + .map_or((32_768, 32_768), |preferences| { + (preferences.context_tokens, preferences.max_generated_tokens) + }); json!({ "object": "list", "data": installed_endpoint_models(&state.models_path) .into_iter() - .map(|model| model_json(model, context)) + .map(|model| model_json(model.id(), model, context, tokens)) .collect::>() }) } -fn model_json(model: ModelChoice, context: i32) -> Value { +fn model_json(id: &str, model: ModelChoice, context: i32, default_tokens: i32) -> Value { json!({ - "id": model.id(), + "id": id, "object": "model", "created": 1767225600_i64, "owned_by": "ds4.c", @@ -516,7 +531,7 @@ fn model_json(model: ModelChoice, context: i32) -> Value { "context_length": context, "top_provider": { "context_length": context, - "max_completion_tokens": context, + "max_completion_tokens": default_tokens.min(context), "is_moderated": false }, "supported_parameters": [ @@ -684,6 +699,7 @@ mod tests { connections: Arc::new(AtomicUsize::new(0)), tool_memory: Mutex::new(HashMap::new()), metrics, + cors: false, }; let text = "done\n\n<|DSML|tool_calls>\n<|DSML|invoke name=\"shell\">\n<|DSML|parameter name=\"command\" string=\"true\">pwd\n\n"; let (content, calls) = parse_generated_tools(&state, text, Protocol::Chat); @@ -708,6 +724,28 @@ mod tests { assert_eq!(calls[0].id.len(), 38); } + #[test] + fn sampled_tool_replay_survives_server_restart() { + let directory = std::env::temp_dir().join(random_id("ds4-tool-replay-")); + let metrics = Arc::new(Metrics::new(&directory)); + let state = State { + generation: GenerationService::spawn(Arc::clone(&metrics)).unwrap(), + preferences: Arc::new(RwLock::new(AppPreferences::default())), + models_path: PathBuf::new(), + cache_path: directory.clone(), + sequence: AtomicU64::new(0), + connections: Arc::new(AtomicUsize::new(0)), + tool_memory: Mutex::new(HashMap::new()), + metrics, + cors: false, + }; + let raw = "<|DSML|tool_calls><|DSML|invoke name=\"echo\"><|DSML|parameter name=\"text\" string=\"true\">hi"; + let (_, calls) = parse_generated_tools(&state, raw, Protocol::Chat); + let restored = read_tool_memory(&directory); + assert_eq!(restored.get(&calls[0].id).map(String::as_str), Some(raw)); + std::fs::remove_dir_all(directory).unwrap(); + } + #[test] fn generated_tool_parser_accepts_ds4_recovery_syntaxes() { let metrics = Arc::new(Metrics::new(PathBuf::new().as_path())); @@ -720,6 +758,7 @@ mod tests { connections: Arc::new(AtomicUsize::new(0)), tool_memory: Mutex::new(HashMap::new()), metrics, + cors: false, }; let short = "done\n\n/tmp"; let (content, calls) = parse_generated_tools(&state, short, Protocol::Chat); @@ -733,6 +772,31 @@ mod tests { let plain = "invalid"; let (_, calls) = parse_generated_tools(&state, plain, Protocol::Chat); assert_eq!(calls[0].function.arguments, r#"{"count":null}"#); + + let truncated = "before<|DSML|tool_calls><|DSML|invoke name=\"echo\"><|DSML|parameter name=\"text\" string=\"true\">hi"; + let (content, calls) = parse_generated_tools(&state, truncated, Protocol::Chat); + assert_eq!(content, "before"); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].function.arguments, r#"{"text":"hi"}"#); + + let quoted_in_reasoning = "example <|DSML|tool_calls>answer"; + assert!( + parse_generated_tools(&state, quoted_in_reasoning, Protocol::Chat) + .1 + .is_empty() + ); + + let mut projector = ToolProjector::new(); + projector.push(truncated, false, "call_"); + let repaired_arguments = projector + .finish("call_") + .into_iter() + .filter_map(|event| match event { + ToolProjectionEvent::Arguments { fragment, .. } => Some(fragment), + _ => None, + }) + .collect::(); + assert_eq!(repaired_arguments, "\"}"); } #[test] @@ -747,6 +811,7 @@ mod tests { connections: Arc::new(AtomicUsize::new(0)), tool_memory: Mutex::new(HashMap::new()), metrics, + cors: false, }; let tool_result: ApiMessage = serde_json::from_value(json!({ "role": "tool", "content": "ok", "tool_call_id": "call_missing" @@ -865,9 +930,10 @@ mod tests { stream: true, include_usage: false, has_tools: true, + cors: false, }; let mut response = Vec::new(); - send_sse_headers(&mut response).unwrap(); + send_sse_headers_with_cors(&mut response, false).unwrap(); response.write_all(b": prefill\n\n").unwrap(); let mut headers_sent = true; let mut role_sent = false; @@ -892,4 +958,72 @@ mod tests { "event: error\ndata: {\"error\":{\"message\":\"prefill failed\",\"type\":\"server_error\"}}\n\n" ); } + + #[test] + fn context_limit_errors_use_protocol_standard_fields() { + assert_eq!( + response::context_error_dimensions( + "Prompt has 32768 tokens, but the configured context size is 32768 tokens" + ), + Some((32_768, 32_768)) + ); + + let listener = match TcpListener::bind("127.0.0.1:0") { + Ok(listener) => listener, + Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => return, + Err(error) => panic!("could not start test server: {error}"), + }; + let address = listener.local_addr().unwrap(); + let client = thread::spawn(move || { + let mut stream = TcpStream::connect(address).unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).unwrap(); + response + }); + let (mut stream, _) = listener.accept().unwrap(); + response::send_generation_error( + &mut stream, + Protocol::Responses, + false, + "Prompt has 16 tokens, but the configured context size is 16 tokens".into(), + ) + .unwrap(); + drop(stream); + let response = client.join().unwrap(); + assert!(response.starts_with("HTTP/1.1 400 Bad Request")); + assert!(response.contains("\"code\":\"context_length_exceeded\"")); + assert!(response.contains("\"param\":\"input\"")); + assert!(response.contains("\"n_prompt_tokens\":16")); + assert!(response.contains("\"n_ctx\":16")); + } + + #[test] + fn model_metadata_uses_the_requested_alias_and_default_token_limit() { + let model = model_json( + "deepseek-reasoner", + ModelChoice::DeepSeekV4Flash, + 32_768, + 50_000, + ); + assert_eq!(model["id"], "deepseek-reasoner"); + assert_eq!(model["top_provider"]["max_completion_tokens"], 32_768); + } + + #[test] + fn cors_headers_are_opt_in_and_cover_preflight() { + let mut response = Vec::new(); + send_response_with_cors(&mut response, 204, None, &[], true).unwrap(); + let response = String::from_utf8(response).unwrap(); + assert!(response.contains("Access-Control-Allow-Origin: *\r\n")); + assert!(response.contains("Access-Control-Allow-Methods: GET, POST, OPTIONS\r\n")); + assert!(response.contains("Access-Control-Allow-Headers: *\r\n")); + + let mut response = Vec::new(); + send_sse_headers_with_cors(&mut response, false).unwrap(); + assert!( + !String::from_utf8(response) + .unwrap() + .contains("Access-Control-Allow-") + ); + } } diff --git a/src/server/http.rs b/src/server/http.rs index a80b3e1..85ddf2d 100644 --- a/src/server/http.rs +++ b/src/server/http.rs @@ -1,10 +1,16 @@ use super::*; -pub(super) fn send_sse_headers(stream: &mut impl Write) -> Result<(), String> { +pub(super) fn send_sse_headers_with_cors( + stream: &mut impl Write, + cors: bool, +) -> Result<(), String> { + let cors = if cors { + "Access-Control-Allow-Origin: *\r\nAccess-Control-Allow-Methods: GET, POST, OPTIONS\r\nAccess-Control-Allow-Headers: *\r\n" + } else { + "" + }; stream - .write_all( - b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\nConnection: close\r\n\r\n", - ) + .write_all(format!("HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\n{cors}Connection: close\r\n\r\n").as_bytes()) .map_err(|error| error.to_string()) } @@ -27,25 +33,37 @@ pub(super) fn send_sse_error(stream: &mut impl Write, message: &str) -> Result<( stream.write_all(b"\n\n").map_err(|error| error.to_string()) } -pub(super) fn send_json(stream: &mut TcpStream, code: u16, value: &Value) -> Result<(), String> { +pub(super) fn send_json_with_cors( + stream: &mut TcpStream, + code: u16, + value: &Value, + cors: bool, +) -> Result<(), String> { let mut body = serde_json::to_vec(value).map_err(|error| error.to_string())?; body.push(b'\n'); - send_response(stream, code, Some("application/json"), &body) + send_response_with_cors(stream, code, Some("application/json"), &body, cors) } -pub(super) fn send_error(stream: &mut TcpStream, code: u16, message: &str) -> Result<(), String> { - send_json( +pub(super) fn send_error_with_cors( + stream: &mut TcpStream, + code: u16, + message: &str, + cors: bool, +) -> Result<(), String> { + send_json_with_cors( stream, code, &json!({"error": {"message": message, "type": "invalid_request_error"}}), + cors, ) } -pub(super) fn send_response( - stream: &mut TcpStream, +pub(super) fn send_response_with_cors( + stream: &mut impl Write, code: u16, content_type: Option<&str>, body: &[u8], + cors: bool, ) -> Result<(), String> { let reason = match code { 200 => "OK", @@ -65,6 +83,11 @@ pub(super) fn send_response( header.push_str(content_type); header.push_str("\r\n"); } + if cors { + header.push_str("Access-Control-Allow-Origin: *\r\n"); + header.push_str("Access-Control-Allow-Methods: GET, POST, OPTIONS\r\n"); + header.push_str("Access-Control-Allow-Headers: *\r\n"); + } header.push_str("Connection: close\r\n\r\n"); stream .write_all(header.as_bytes()) diff --git a/src/server/response.rs b/src/server/response.rs index 47ca807..217bc56 100644 --- a/src/server/response.rs +++ b/src/server/response.rs @@ -7,7 +7,11 @@ pub(super) fn final_response( active: crate::runtime::ActiveGeneration, id: &str, ) -> Result<(), (u16, String)> { - let output = wait_for_output(active)?; + let output = match wait_for_output(stream, active) { + Ok(output) => output, + Err(error) if error == "client disconnected" => return Ok(()), + Err(error) => return send_generation_error(stream, request.protocol, request.cors, error), + }; let (content, calls) = parse_generated_tools(state, &output.message.content, request.protocol); let finish = if calls.is_empty() { output.finish_reason @@ -114,7 +118,7 @@ pub(super) fn final_response( }) } }; - send_json(stream, 200, &body).map_err(|error| (500, error)) + send_json_with_cors(stream, 200, &body, request.cors).map_err(|error| (500, error)) } pub(super) fn stream_response( @@ -158,7 +162,7 @@ fn completion_stream_response( active: crate::runtime::ActiveGeneration, id: &str, ) -> Result<(), (u16, String)> { - send_sse_headers(stream).map_err(|error| (500, error))?; + send_sse_headers_with_cors(stream, request.cors).map_err(|error| (500, error))?; while let Ok(event) = active.events.recv() { match event { GenerationEvent::Chunk { @@ -215,7 +219,7 @@ fn anthropic_stream_response( active: crate::runtime::ActiveGeneration, id: &str, ) -> Result<(), (u16, String)> { - send_sse_headers(stream).map_err(|error| (500, error))?; + send_sse_headers_with_cors(stream, request.cors).map_err(|error| (500, error))?; let mut prompt_tokens = 0; let mut started = false; let mut block = None::<(usize, bool)>; @@ -298,7 +302,7 @@ fn anthropic_stream_response( )?; } if request.has_tools { - let events = projector.push("", true, "toolu_"); + let events = projector.finish("toolu_"); send_anthropic_projection_events( stream, events, @@ -468,7 +472,7 @@ fn responses_stream_response( active: crate::runtime::ActiveGeneration, id: &str, ) -> Result<(), (u16, String)> { - send_sse_headers(stream).map_err(|error| (500, error))?; + send_sse_headers_with_cors(stream, request.cors).map_err(|error| (500, error))?; let created = unix_time(); let message_id = random_id("msg_"); let reasoning_id = random_id("rs_"); @@ -739,7 +743,8 @@ fn stream_response_with_keepalive( } => { prefilling = tokens_per_second.is_none(); if prefilling && !headers_sent { - send_sse_headers(stream).map_err(|error| (500, error))?; + send_sse_headers_with_cors(stream, request.cors) + .map_err(|error| (500, error))?; headers_sent = true; last_keepalive = Instant::now(); } else if !prefilling { @@ -785,7 +790,12 @@ fn stream_response_with_keepalive( let _ = send_sse_error(stream, &error); return Ok(()); } - return Err((500, error)); + return send_generation_error( + stream, + request.protocol, + request.cors, + error, + ); } } break; @@ -801,7 +811,7 @@ fn stream_response_with_keepalive( None => return Err((500, "The model runtime stopped unexpectedly.".into())), }; if request.has_tools { - let events = projector.push("", true, "call_"); + let events = projector.finish("call_"); send_chat_projection_events(stream, &request, id, events)?; } let (_, calls) = parse_generated_tools_with_ids( @@ -885,7 +895,7 @@ pub(super) fn send_stream_start( role_sent: &mut bool, ) -> Result<(), (u16, String)> { if !*headers_sent { - send_sse_headers(stream).map_err(|error| (500, error))?; + send_sse_headers_with_cors(stream, request.cors).map_err(|error| (500, error))?; *headers_sent = true; } if !*role_sent { @@ -924,10 +934,28 @@ pub(super) fn receive_stream_event( } fn wait_for_output( + stream: &mut TcpStream, active: crate::runtime::ActiveGeneration, -) -> Result { +) -> Result { let mut content = String::new(); - while let Ok(event) = active.events.recv() { + stream + .set_nonblocking(true) + .map_err(|error| error.to_string())?; + loop { + let event = match active.events.recv_timeout(Duration::from_millis(100)) { + Ok(event) => event, + Err(std::sync::mpsc::RecvTimeoutError::Timeout) => { + let mut byte = [0]; + match stream.peek(&mut byte) { + Ok(0) => return Err("client disconnected".into()), + Ok(_) => {} + Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {} + Err(error) => return Err(error.to_string()), + } + continue; + } + Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => break, + }; match event { GenerationEvent::Chunk { reasoning: false, @@ -942,12 +970,62 @@ fn wait_for_output( } } GenerationEvent::Finished(result) => { - return result.map_err(|error| (500, error)); + stream + .set_nonblocking(false) + .map_err(|error| error.to_string())?; + return result; } _ => {} } } - Err((500, "The model runtime stopped unexpectedly.".into())) + let _ = stream.set_nonblocking(false); + Err("The model runtime stopped unexpectedly.".into()) +} + +pub(super) fn send_generation_error( + stream: &mut TcpStream, + protocol: Protocol, + cors: bool, + error: String, +) -> Result<(), (u16, String)> { + let Some((prompt_tokens, context)) = context_error_dimensions(&error) else { + return Err((500, error)); + }; + let body = if protocol == Protocol::Anthropic { + json!({ + "type": "error", + "error": { + "type": "invalid_request_error", + "message": error, + "n_prompt_tokens": prompt_tokens, + "n_ctx": context + } + }) + } else { + let parameter = match protocol { + Protocol::Completion => "prompt", + Protocol::Responses => "input", + Protocol::Chat => "messages", + Protocol::Anthropic => unreachable!(), + }; + json!({"error": { + "message": error, + "type": "invalid_request_error", + "param": parameter, + "code": "context_length_exceeded", + "n_prompt_tokens": prompt_tokens, + "n_ctx": context + }}) + }; + send_json_with_cors(stream, 400, &body, cors).map_err(|error| (500, error)) +} + +pub(super) fn context_error_dimensions(error: &str) -> Option<(u32, u32)> { + let values = error + .strip_prefix("Prompt has ")? + .strip_suffix(" tokens")? + .split_once(" tokens, but the configured context size is ")?; + Some((values.0.parse().ok()?, values.1.parse().ok()?)) } pub(super) fn usage_json(prompt: u32, cached: u32, completion: u32) -> Value { diff --git a/src/server/tools.rs b/src/server/tools.rs index 4a556cd..d1368de 100644 --- a/src/server/tools.rs +++ b/src/server/tools.rs @@ -1,5 +1,9 @@ use super::*; +const TOOL_MEMORY_FILE: &str = "tool-replay.json"; +const TOOL_MEMORY_MAX_IDS: usize = 100_000; +const TOOL_MEMORY_MAX_BYTES: u64 = 512 * 1024 * 1024; + enum ToolProjectionState { Seeking, Invokes, @@ -226,6 +230,15 @@ impl ToolProjector { events } + pub(super) fn finish(&mut self, prefix: &str) -> Vec { + if let Some(repaired) = repair_generated_tools(&self.raw) { + let suffix = repaired[self.raw.len()..].to_owned(); + self.push(&suffix, true, prefix) + } else { + self.push("", true, prefix) + } + } + fn emit_value(&mut self, end: usize, events: &mut Vec) { if end <= self.position { return; @@ -506,6 +519,17 @@ pub(super) fn parse_generated_tools_with_ids( text: &str, protocol: Protocol, streamed_ids: &[String], +) -> (String, Vec) { + let repaired = repair_generated_tools(text); + let text = repaired.as_deref().unwrap_or(text); + parse_generated_tools_once(state, text, protocol, streamed_ids) +} + +fn parse_generated_tools_once( + state: &State, + text: &str, + protocol: Protocol, + streamed_ids: &[String], ) -> (String, Vec) { let Some((start, syntax)) = TOOL_SYNTAXES .iter() @@ -575,16 +599,95 @@ pub(super) fn parse_generated_tools_with_ids( .tool_memory .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); - // ponytail: one process-local replay table; add LRU eviction if 100k live tool ids is measured insufficient. - if memory.len() >= 100_000 { + // ponytail: clear-at-cap is cheaper than an LRU; add ordered eviction if real clients retain 100k live ids. + if memory.len() >= TOOL_MEMORY_MAX_IDS { memory.clear(); } for call in &calls { memory.insert(call.id.clone(), raw.to_owned()); } + write_tool_memory(&state.cache_path, &memory); (content.to_owned(), calls) } +pub(super) fn read_tool_memory(directory: &std::path::Path) -> HashMap { + if directory.as_os_str().is_empty() { + return HashMap::new(); + } + let path = directory.join(TOOL_MEMORY_FILE); + if std::fs::metadata(&path).is_ok_and(|metadata| metadata.len() > TOOL_MEMORY_MAX_BYTES) { + return HashMap::new(); + } + let Ok(file) = File::open(path) else { + return HashMap::new(); + }; + let Ok(memory) = serde_json::from_reader::<_, HashMap>(file) else { + return HashMap::new(); + }; + if memory.len() <= TOOL_MEMORY_MAX_IDS { + memory + } else { + HashMap::new() + } +} + +fn write_tool_memory(directory: &std::path::Path, memory: &HashMap) { + if directory.as_os_str().is_empty() || std::fs::create_dir_all(directory).is_err() { + return; + } + let path = directory.join(TOOL_MEMORY_FILE); + let temporary = path.with_extension("json.tmp"); + let Ok(mut file) = File::create(&temporary) else { + return; + }; + if serde_json::to_writer(&mut file, memory).is_ok() + && file.sync_all().is_ok() + && file + .metadata() + .is_ok_and(|metadata| metadata.len() <= TOOL_MEMORY_MAX_BYTES) + { + let _ = std::fs::rename(&temporary, path); + } else { + let _ = std::fs::remove_file(temporary); + } +} + +fn repair_generated_tools(text: &str) -> Option { + let scan_start = text + .rfind("") + .map_or(0, |position| position + "".len()); + let scan = &text[scan_start..]; + let syntax = TOOL_SYNTAXES + .iter() + .filter_map(|syntax| scan.find(syntax.tool_start).map(|start| (start, *syntax))) + .min_by_key(|(start, _)| *start)? + .1; + let tool_open = scan.matches(syntax.tool_start).count(); + let tool_close = scan.matches(syntax.tool_end).count(); + let invoke_open = scan.matches(syntax.invoke_start).count(); + let invoke_close = scan.matches(syntax.invoke_end).count(); + let parameter_open = scan.matches(syntax.parameter_start).count(); + let parameter_close = scan.matches(syntax.parameter_end).count(); + if (tool_open, invoke_open, parameter_open) == (tool_close, invoke_close, parameter_close) + || tool_close > tool_open + || invoke_close > invoke_open + || parameter_close > parameter_open + { + return None; + } + let mut repaired = text.to_owned(); + for _ in parameter_close..parameter_open { + repaired.push_str(syntax.parameter_end); + } + for _ in invoke_close..invoke_open { + repaired.push_str(syntax.invoke_end); + } + for _ in tool_close..tool_open { + repaired.push_str(syntax.tool_end); + } + Some(repaired) +} + #[derive(Clone, Copy)] pub(super) struct ToolSyntax { pub(super) tool_start: &'static str,