Complete local HTTP endpoint parity

This commit is contained in:
Georg Bauer
2026-07-25 14:18:57 +02:00
parent a1761fa731
commit 668d8b787e
13 changed files with 493 additions and 69 deletions

View File

@@ -0,0 +1,2 @@
ALTER TABLE preferences DROP COLUMN endpoint_cors;
ALTER TABLE preferences DROP COLUMN endpoint_enabled;

View File

@@ -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));

View File

@@ -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),
}
}

View File

@@ -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)

View File

@@ -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),

View File

@@ -52,6 +52,8 @@ pub struct AppPreferences {
pub simulated_used_memory_gib: Option<i64>,
pub expert_profile_path: Option<String>,
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<i64>,
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<AppPreferences, String> {
@@ -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

View File

@@ -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()
));
}

View File

@@ -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)

View File

@@ -54,6 +54,8 @@ diesel::table! {
simulated_used_memory_gib -> Nullable<BigInt>,
expert_profile_path -> Nullable<Text>,
endpoint_port -> Integer,
endpoint_enabled -> Bool,
endpoint_cors -> Bool,
}
}

View File

@@ -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<AtomicUsize>,
tool_memory: Mutex<HashMap<String, String>>,
metrics: Arc<Metrics>,
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<Metrics>,
) -> Result<Self, String> {
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<ModelChoice>
}
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::<Vec<_>>()
})
}
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<DSMLtool_calls>\n<DSMLinvoke name=\"shell\">\n<DSMLparameter name=\"command\" string=\"true\">pwd</DSMLparameter>\n</DSMLinvoke>\n</DSMLtool_calls>";
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 = "<DSMLtool_calls><DSMLinvoke name=\"echo\"><DSMLparameter name=\"text\" string=\"true\">hi</DSMLparameter></DSMLinvoke></DSMLtool_calls>";
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<DSMLtool_calls><DSMLinvoke name=\"run\"><DSMLparameter name=\"options\"><DSMLparameter name=\"path\">/tmp</DSMLparameter></DSMLparameter></DSMLinvoke></DSMLtool_calls>";
let (content, calls) = parse_generated_tools(&state, short, Protocol::Chat);
@@ -733,6 +772,31 @@ mod tests {
let plain = "<tool_calls><invoke name=\"echo\"><parameter name=\"count\" string=\"false\">invalid</parameter></invoke></tool_calls>";
let (_, calls) = parse_generated_tools(&state, plain, Protocol::Chat);
assert_eq!(calls[0].function.arguments, r#"{"count":null}"#);
let truncated = "before<DSMLtool_calls><DSMLinvoke name=\"echo\"><DSMLparameter 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 = "<think>example <DSMLtool_calls></think>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::<String>();
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-")
);
}
}

View File

@@ -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())

View File

@@ -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<crate::engine::GenerationOutput, (u16, String)> {
) -> Result<crate::engine::GenerationOutput, String> {
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 {

View File

@@ -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<ToolProjectionEvent> {
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<ToolProjectionEvent>) {
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<ApiToolCall>) {
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<ApiToolCall>) {
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<String, String> {
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<String, String>>(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<String, String>) {
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<String> {
let scan_start = text
.rfind("</think>")
.map_or(0, |position| position + "</think>".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,