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

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