Complete local HTTP endpoint parity
This commit is contained in:
186
src/server.rs
186
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<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<|DSML|tool_calls>\n<|DSML|invoke name=\"shell\">\n<|DSML|parameter name=\"command\" string=\"true\">pwd</|DSML|parameter>\n</|DSML|invoke>\n</|DSML|tool_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 = "<|DSML|tool_calls><|DSML|invoke name=\"echo\"><|DSML|parameter name=\"text\" string=\"true\">hi</|DSML|parameter></|DSML|invoke></|DSML|tool_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<DSML|tool_calls><DSML|invoke name=\"run\"><DSML|parameter name=\"options\"><DSML|parameter name=\"path\">/tmp</DSML|parameter></DSML|parameter></DSML|invoke></DSML|tool_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<|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 = "<think>example <|DSML|tool_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-")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user