mod http;
mod request;
mod response;
mod tools;
use http::{
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,
};
use response::{final_response, stream_response};
#[cfg(test)]
use response::{receive_stream_event, send_stream_start};
use tools::{
TOOL_SYNTAXES, ToolProjectionEvent, ToolProjector, content_text, parse_generated_tools,
parse_generated_tools_with_ids, read_tool_memory, render_messages, tool_calls_json,
};
#[cfg(test)]
use tools::{canonical_tools, validate_tool_results};
use crate::database::AppPreferences;
use crate::engine::ChatTurn;
use crate::metrics::Metrics;
use crate::model::{self, ModelChoice};
use crate::runtime::{CheckpointTarget, GenerationEvent, GenerationService};
use crate::settings::{ReasoningMode, effective_settings};
use serde::Deserialize;
use serde_json::value::RawValue;
use serde_json::{Map, Value, json};
use std::collections::HashMap;
use std::fs::File;
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, RwLock};
use std::thread::{self, JoinHandle};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
const MAX_HEADER_BYTES: usize = 64 * 1024;
const MAX_BODY_BYTES: usize = 64 * 1024 * 1024;
const MAX_CONNECTIONS: usize = 64;
const HTTP_IO_TIMEOUT: Duration = Duration::from_secs(10);
const HTTP_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const PREFILL_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(5);
const TOOLS_PROMPT: &str = "## Tools\n\n\
You have access to a set of tools to help answer the user question. You can invoke tools by writing a \"<|DSML|tool_calls>\" block like the following:\n\n\
<|DSML|tool_calls>\n\
<|DSML|invoke name=\"$TOOL_NAME\">\n\
<|DSML|parameter name=\"$PARAMETER_NAME\" string=\"true|false\">$PARAMETER_VALUE|DSML|parameter>\n\
...\n\
|DSML|invoke>\n\
<|DSML|invoke name=\"$TOOL_NAME2\">\n\
...\n\
|DSML|invoke>\n\
|DSML|tool_calls>\n\n\
String parameters should be specified as raw text and set `string=\"true\"`. Preserve characters such as `>`, `&`, and `&&` exactly; never replace normal string characters with XML or HTML entity escapes. Only if a string value itself contains the exact closing parameter tag `|DSML|parameter>`, write that tag as `</|DSML|parameter>` inside the value. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string=\"false\"`.\n\n\
If thinking_mode is enabled (triggered by ), you MUST output your complete reasoning inside ... BEFORE any tool calls or final response.\n\n\
Otherwise, output directly after with tool calls or final response.\n\n\
### Available Tool Schemas\n\n";
pub(crate) struct ServerHandle {
stop: Arc,
thread: Option>,
metrics: Arc,
port: u16,
}
struct State {
generation: GenerationService,
preferences: Arc>,
models_path: PathBuf,
cache_path: PathBuf,
sequence: AtomicU64,
connections: Arc,
tool_memory: Mutex>,
metrics: Arc,
cors: bool,
}
#[derive(Deserialize)]
struct ChatRequest {
#[serde(default)]
model: Option,
#[serde(default)]
messages: Vec,
#[serde(default)]
tools: Vec,
#[serde(skip)]
tool_schemas: Vec,
#[serde(default)]
tool_choice: Option,
#[serde(default)]
max_tokens: Option,
#[serde(default)]
max_completion_tokens: Option,
#[serde(default)]
temperature: Option,
#[serde(default)]
top_p: Option,
#[serde(default)]
min_p: Option,
#[serde(default)]
top_k: Option,
#[serde(default)]
seed: Option,
#[serde(default)]
stream: bool,
#[serde(default)]
stream_options: Option,
#[serde(default)]
thinking: Option,
#[serde(default)]
think: Option,
#[serde(default)]
reasoning_effort: Option,
#[serde(default)]
stop: Option,
}
#[derive(Deserialize)]
struct ApiMessage {
#[serde(default = "default_role")]
role: String,
#[serde(default)]
content: Value,
#[serde(default)]
reasoning_content: Value,
#[serde(default)]
tool_calls: Vec,
#[serde(default)]
tool_call_id: String,
}
#[derive(Clone, Deserialize)]
struct ApiToolCall {
#[serde(default)]
id: String,
function: ApiFunction,
}
#[derive(Clone, Deserialize)]
struct ApiFunction {
name: String,
#[serde(default = "empty_arguments")]
arguments: String,
}
#[derive(Deserialize)]
struct StreamOptions {
#[serde(default)]
include_usage: bool,
}
#[derive(Deserialize)]
#[serde(untagged)]
enum OneOrMany {
One(String),
Many(Vec),
}
struct ParsedRequest {
protocol: Protocol,
model_id: String,
messages: Vec,
turn: crate::settings::TurnSettings,
engine: crate::settings::EngineSettings,
idle_timeout: Duration,
stream: bool,
include_usage: bool,
has_tools: bool,
}
struct ResponseOptions {
protocol: Protocol,
reasoning_summary: bool,
model_id: String,
stream: bool,
include_usage: bool,
has_tools: bool,
cors: bool,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Protocol {
Chat,
Completion,
Anthropic,
Responses,
}
struct HttpRequest {
method: String,
path: String,
body: Vec,
}
#[derive(Deserialize)]
struct RawToolsRequest<'a> {
#[serde(default, borrow)]
tools: Vec<&'a RawValue>,
}
#[derive(Deserialize)]
struct RawOpenAiTool<'a> {
#[serde(default, borrow)]
function: Option<&'a RawValue>,
}
impl ServerHandle {
pub(crate) fn spawn(
generation: GenerationService,
preferences: Arc>,
models_path: PathBuf,
cache_path: PathBuf,
port: u16,
cors: bool,
metrics: Arc,
) -> Result {
let address = format!("127.0.0.1:{port}");
let listener = TcpListener::bind(("127.0.0.1", port))
.map_err(|error| format!("Could not listen on http://{address}: {error}"))?;
listener
.set_nonblocking(true)
.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,
models_path,
cache_path,
sequence: AtomicU64::new(0),
connections: Arc::new(AtomicUsize::new(0)),
tool_memory: Mutex::new(tool_memory),
metrics: Arc::clone(&metrics),
cors,
});
let thread = thread::Builder::new()
.name("local-http".into())
.spawn(move || serve(listener, state, worker_stop))
.map_err(|error| format!("Could not start the local HTTP service: {error}"))?;
metrics.server_listening(port);
Ok(Self {
stop,
thread: Some(thread),
metrics,
port,
})
}
}
impl Drop for ServerHandle {
fn drop(&mut self) {
self.stop.store(true, Ordering::Relaxed);
if let Some(thread) = self.thread.take() {
let _ = thread.join();
}
self.metrics.server_stopped(self.port);
}
}
fn serve(listener: TcpListener, state: Arc, stop: Arc) {
while !stop.load(Ordering::Relaxed) {
match listener.accept() {
Ok((stream, _)) => {
let Some(slot) = ConnectionSlot::acquire(&state.connections) else {
continue;
};
let state = Arc::clone(&state);
if let Err(error) =
thread::Builder::new()
.name("http-request".into())
.spawn(move || {
let _slot = slot;
handle(stream, &state);
})
{
eprintln!("DS4Server: endpoint request thread failed: {error}");
}
}
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(50));
}
Err(error) => {
eprintln!("DS4Server: endpoint accept failed: {error}");
thread::sleep(Duration::from_millis(100));
}
}
}
}
struct ConnectionSlot(Arc);
impl ConnectionSlot {
fn acquire(active: &Arc) -> Option {
active
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| {
(count < MAX_CONNECTIONS).then_some(count + 1)
})
.ok()?;
Some(Self(Arc::clone(active)))
}
}
impl Drop for ConnectionSlot {
fn drop(&mut self) {
self.0.fetch_sub(1, Ordering::Relaxed);
}
}
fn handle(mut stream: TcpStream, state: &State) {
let _ = stream.set_write_timeout(Some(HTTP_IO_TIMEOUT));
let started = Instant::now();
let request = match read_request(&mut stream, started + HTTP_REQUEST_TIMEOUT) {
Ok(request) => request,
Err(error) => {
state.metrics.http_started("", 0);
let _ = send_error_with_cors(&mut stream, 400, &error, state.cors);
state.metrics.http_finished(started.elapsed(), true);
return;
}
};
state
.metrics
.http_started(&request.path, request.body.len());
let result = match (request.method.as_str(), request.path.as_str()) {
("OPTIONS", _) => send_response_with_cors(&mut stream, 204, None, &[], state.cors),
("GET", "/v1/models") => {
let body = models_json(state);
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_with_cors(&mut stream, error.0, &error.1, state.cors);
Err(error.1)
}
}
}
("POST", "/v1/completions") => {
match compatible_completion(&mut stream, state, &request.body, Protocol::Completion) {
Ok(()) => Ok(()),
Err(error) => {
let _ = send_error_with_cors(&mut stream, error.0, &error.1, state.cors);
Err(error.1)
}
}
}
("POST", "/v1/messages") => {
match compatible_completion(&mut stream, state, &request.body, Protocol::Anthropic) {
Ok(()) => Ok(()),
Err(error) => {
let _ = send_error_with_cors(&mut stream, error.0, &error.1, state.cors);
Err(error.1)
}
}
}
("POST", "/v1/responses") => {
match compatible_completion(&mut stream, state, &request.body, Protocol::Responses) {
Ok(()) => Ok(()),
Err(error) => {
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 installed = installed_endpoint_models(&state.models_path);
let model = model_alias(id).filter(|model| installed.contains(model));
if let Some(model) = model {
let (context, tokens) = state
.preferences
.read()
.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_with_cors(&mut stream, 404, "unknown model", state.cors);
Err("unknown model".into())
}
}
_ => {
let _ = send_error_with_cors(&mut stream, 404, "unknown endpoint", state.cors);
Err("unknown endpoint".into())
}
};
state
.metrics
.http_finished(started.elapsed(), result.is_err());
}
fn chat_completion(
stream: &mut TcpStream,
state: &State,
body: &[u8],
) -> Result<(), (u16, String)> {
let mut request: ChatRequest =
serde_json::from_slice(body).map_err(|_| (400, "invalid JSON request".to_owned()))?;
request.tool_schemas = raw_tool_schemas(body, Protocol::Chat)?;
let parsed = parse_chat_request(state, request, Protocol::Chat)?;
if parsed.stream {
state.metrics.http_streaming();
}
let response = ResponseOptions {
protocol: parsed.protocol,
reasoning_summary: false,
model_id: parsed.model_id,
stream: parsed.stream,
include_usage: parsed.include_usage,
has_tools: parsed.has_tools,
cors: state.cors,
};
let active = state
.generation
.generate(
parsed.engine,
parsed.turn,
parsed.messages,
CheckpointTarget::Transient(state.cache_path.clone()),
parsed.idle_timeout,
)
.map_err(|error| (500, error))?;
let sequence = state.sequence.fetch_add(1, Ordering::Relaxed) + 1;
let id = format!("chatcmpl-{sequence}");
if response.stream {
stream_response(stream, state, response, active, &id)
} else {
final_response(stream, state, response, active, &id)
}
}
fn compatible_completion(
stream: &mut TcpStream,
state: &State,
body: &[u8],
protocol: Protocol,
) -> Result<(), (u16, String)> {
let value: Value =
serde_json::from_slice(body).map_err(|_| (400, "invalid JSON request".to_owned()))?;
let reasoning_summary = protocol == Protocol::Responses
&& matches!(
value.pointer("/reasoning/summary").and_then(Value::as_str),
Some("auto" | "concise" | "detailed")
);
let mut request = match protocol {
Protocol::Completion => completion_request(value)?,
Protocol::Anthropic => anthropic_request(value)?,
Protocol::Responses => responses_request(value)?,
Protocol::Chat => unreachable!(),
};
request.tool_schemas = raw_tool_schemas(body, protocol)?;
let parsed = parse_chat_request(state, request, protocol)?;
if parsed.stream {
state.metrics.http_streaming();
}
let response = ResponseOptions {
protocol,
reasoning_summary,
model_id: parsed.model_id,
stream: parsed.stream,
include_usage: parsed.include_usage,
has_tools: parsed.has_tools,
cors: state.cors,
};
let active = state
.generation
.generate(
parsed.engine,
parsed.turn,
parsed.messages,
CheckpointTarget::Transient(state.cache_path.clone()),
parsed.idle_timeout,
)
.map_err(|error| (500, error))?;
let sequence = state.sequence.fetch_add(1, Ordering::Relaxed) + 1;
let id = match protocol {
Protocol::Anthropic => format!("chatcmpl-{sequence}"),
Protocol::Responses => random_id("resp_"),
Protocol::Completion => format!("cmpl-{sequence}"),
Protocol::Chat => unreachable!(),
};
if response.stream {
stream_response(stream, state, response, active, &id)
} else {
final_response(stream, state, response, active, &id)
}
}
fn installed_endpoint_models(models_path: &std::path::Path) -> Vec {
model::installed_models(models_path)
.into_iter()
.filter(|model| *model == ModelChoice::DeepSeekV4Flash)
.collect()
}
fn models_json(state: &State) -> Value {
let (context, tokens) = state
.preferences
.read()
.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.id(), model, context, tokens))
.collect::>()
})
}
fn model_json(id: &str, model: ModelChoice, context: i32, default_tokens: i32) -> Value {
json!({
"id": id,
"object": "model",
"created": 1767225600_i64,
"owned_by": "ds4.c",
"name": model.to_string(),
"context_length": context,
"top_provider": {
"context_length": context,
"max_completion_tokens": default_tokens.min(context),
"is_moderated": false
},
"supported_parameters": [
"tools", "tool_choice", "max_tokens", "temperature", "top_p",
"top_k", "min_p", "stop", "seed", "stream", "reasoning_effort"
]
})
}
fn model_alias(id: &str) -> Option {
match id {
"deepseek-chat" | "deepseek-reasoner" => Some(ModelChoice::DeepSeekV4Flash),
"glm-5.2-chat"
| "glm-5.2-no-think"
| "glm-5.2-nothink"
| "glm-5.2-reasoner"
| "zai/glm-5.2"
| "zai/glm-5.2-chat"
| "zai/glm-5.2-reasoner" => Some(ModelChoice::Glm52),
_ => ModelChoice::from_id(id),
}
}
fn default_role() -> String {
"user".into()
}
fn empty_arguments() -> String {
"{}".into()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn connection_slots_enforce_the_limit_and_release_on_drop() {
let active = Arc::new(AtomicUsize::new(0));
let mut slots = (0..MAX_CONNECTIONS)
.map(|_| ConnectionSlot::acquire(&active).unwrap())
.collect::>();
assert!(ConnectionSlot::acquire(&active).is_none());
slots.pop();
assert!(ConnectionSlot::acquire(&active).is_some());
}
#[test]
fn request_deadline_stops_slow_clients() {
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();
stream.write_all(b"GET /v1/models HTTP/1.1\r\n").unwrap();
thread::sleep(Duration::from_millis(200));
});
let (mut stream, _) = listener.accept().unwrap();
let error = read_request(&mut stream, Instant::now() + Duration::from_millis(50))
.err()
.unwrap();
client.join().unwrap();
assert_eq!(error, "HTTP request timed out");
}
#[test]
fn tool_prompt_matches_ds4_server_text() {
assert_eq!(TOOLS_PROMPT.len(), 1_183);
}
#[test]
fn dsml_tool_stream_projects_incremental_json_arguments() {
let mut projector = ToolProjector::new();
let mut fragments = Vec::new();
for (chunk, final_chunk) in [
("\n\n", false),
(
"<|DSML|tool_calls>\n<|DSML|invoke name=\"echo\">\n<|DSML|parameter name=\"text\" string=\"true\">h&am",
false,
),
(
"p;i|DSML|parameter>\n|DSML|invoke>\n|DSML|tool_calls>",
true,
),
] {
for event in projector.push(chunk, final_chunk, "call_") {
match event {
ToolProjectionEvent::Start { name, .. } => {
fragments.push(format!("start:{name}"));
}
ToolProjectionEvent::Arguments { fragment, .. } => fragments.push(fragment),
ToolProjectionEvent::End { .. } => fragments.push("end".into()),
ToolProjectionEvent::Text(_) => {}
}
}
}
assert_eq!(
fragments,
[
"start:echo",
"{",
"\"text\":\"",
"h",
"&i",
"\"",
"}",
"end"
]
);
assert_eq!(projector.ids.len(), 1);
}
#[test]
fn tool_stream_projects_ds4_recovery_syntaxes() {
for text in [
"hi",
"hi",
] {
let mut projector = ToolProjector::new();
let events = projector.push(text, true, "call_");
assert!(events.iter().any(|event| matches!(
event,
ToolProjectionEvent::Start { name, .. } if name == "echo"
)));
let arguments = events
.into_iter()
.filter_map(|event| match event {
ToolProjectionEvent::Arguments { fragment, .. } => Some(fragment),
_ => None,
})
.collect::();
assert_eq!(arguments, r#"{"text":"hi"}"#);
}
}
#[test]
fn canonical_tool_replay_preserves_argument_types() {
let calls = vec![ApiToolCall {
id: "call_1".into(),
function: ApiFunction {
name: "shell".into(),
arguments: r#"{"command":"printf hi","timeout":30,"login":false}"#.into(),
},
}];
let rendered = canonical_tools(&calls);
assert!(rendered.contains("name=\"command\" string=\"true\">printf hi"));
assert!(rendered.contains("name=\"timeout\" string=\"false\">30"));
assert!(rendered.contains("name=\"login\" string=\"false\">false"));
}
#[test]
fn generated_dsml_becomes_openai_tool_calls() {
let metrics = Arc::new(Metrics::new(PathBuf::new().as_path()));
let state = State {
generation: GenerationService::spawn(Arc::clone(&metrics)).unwrap(),
preferences: Arc::new(RwLock::new(AppPreferences::default())),
models_path: PathBuf::new(),
cache_path: PathBuf::new(),
sequence: AtomicU64::new(0),
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);
assert_eq!(content, "done");
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].function.name, "shell");
assert_eq!(calls[0].function.arguments, r#"{"command":"pwd"}"#);
assert_eq!(calls[0].id.len(), 37);
assert_eq!(state.sequence.load(Ordering::Relaxed), 0);
assert!(
state
.tool_memory
.lock()
.unwrap()
.get(&calls[0].id)
.unwrap()
.starts_with("\n\n<|DSML|tool_calls>")
);
let (_, calls) = parse_generated_tools(&state, text, Protocol::Anthropic);
assert!(calls[0].id.starts_with("toolu_"));
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()));
let state = State {
generation: GenerationService::spawn(Arc::clone(&metrics)).unwrap(),
preferences: Arc::new(RwLock::new(AppPreferences::default())),
models_path: PathBuf::new(),
cache_path: PathBuf::new(),
sequence: AtomicU64::new(0),
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);
assert_eq!(content, "done");
assert_eq!(calls[0].function.name, "run");
assert_eq!(
calls[0].function.arguments,
r#"{"options":{"path":"/tmp"}}"#
);
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]
fn protocol_tool_results_require_live_or_replayed_call_ids() {
let metrics = Arc::new(Metrics::new(PathBuf::new().as_path()));
let state = State {
generation: GenerationService::spawn(Arc::clone(&metrics)).unwrap(),
preferences: Arc::new(RwLock::new(AppPreferences::default())),
models_path: PathBuf::new(),
cache_path: PathBuf::new(),
sequence: AtomicU64::new(0),
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"
}))
.unwrap();
let error = validate_tool_results(&state, &[tool_result], Protocol::Responses).unwrap_err();
assert_eq!(error.0, 400);
assert!(
error
.1
.contains("Responses continuation state is not available")
);
let replay: Vec = serde_json::from_value(json!([
{"role": "assistant", "tool_calls": [{"id": "call_replay", "function": {"name": "echo", "arguments": "{}"}}]},
{"role": "tool", "content": "ok", "tool_call_id": "call_replay"}
]))
.unwrap();
assert!(validate_tool_results(&state, &replay, Protocol::Responses).is_ok());
state
.tool_memory
.lock()
.unwrap()
.insert("toolu_live".into(), String::new());
let live: ApiMessage = serde_json::from_value(json!({
"role": "tool", "content": "ok", "tool_call_id": "toolu_live"
}))
.unwrap();
assert!(validate_tool_results(&state, &[live], Protocol::Anthropic).is_ok());
}
#[test]
fn responses_parser_rejects_dropped_state_and_preserves_hosted_calls() {
let error = responses_request(json!({
"input": "hello", "previous_response_id": "resp_old"
}))
.err()
.unwrap();
assert_eq!(
error.1,
"previous_response_id is not supported; replay full input instead"
);
assert_eq!(
responses_request(json!({"input": "hello", "tool_choice": "required"}))
.err()
.unwrap()
.1,
"tool_choice=required not supported"
);
assert!(responses_request(json!({"input": [{"type": "unknown"}]})).is_err());
assert!(
responses_request(json!({
"input": [{"type": "message", "role": "user", "content": [{"type": "image", "url": "x"}]}]
}))
.is_err()
);
let request = responses_request(json!({
"input": [
{"type": "reasoning", "summary": [{"type": "summary_text", "text": "why"}]},
{"type": "local_shell_call", "id": "call_1", "action": {"command": "pwd"}, "status": "completed"},
{"type": "local_shell_call_output", "call_id": "call_1", "output": "/tmp"}
]
}))
.unwrap();
assert_eq!(request.messages[0].reasoning_content, json!("why"));
assert_eq!(
request.messages[0].tool_calls[0].function.name,
"local_shell"
);
assert_eq!(
request.messages[0].tool_calls[0].function.arguments,
r#"{"command":"pwd"}"#
);
assert_eq!(request.messages[1].tool_call_id, "call_1");
assert_eq!(request.messages[1].content, json!("/tmp"));
}
#[test]
fn content_arrays_match_ds4_text_projection() {
assert_eq!(
content_text(&json!([{"type": "text", "text": "one"}, " two"])),
"one two"
);
}
#[test]
fn streaming_prefill_keeps_the_client_connection_alive() {
let (events, receiver) = std::sync::mpsc::channel();
let active = crate::runtime::ActiveGeneration {
events: receiver,
cancel: Arc::new(AtomicBool::new(false)),
};
let producer = thread::spawn(move || {
thread::sleep(Duration::from_millis(25));
events.send(GenerationEvent::Loading).unwrap();
});
let mut response = Vec::new();
let event = receive_stream_event(
&mut response,
&active,
true,
&mut Instant::now(),
Duration::from_millis(10),
)
.unwrap();
producer.join().unwrap();
assert!(matches!(event, Some(GenerationEvent::Loading)));
assert!(response.starts_with(b": prefill\n\n"));
let request = ResponseOptions {
protocol: Protocol::Chat,
reasoning_summary: false,
model_id: "deepseek-v4-flash".into(),
stream: true,
include_usage: false,
has_tools: true,
cors: false,
};
let mut response = Vec::new();
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;
send_stream_start(
&mut response,
&request,
"chatcmpl-test",
&mut headers_sent,
&mut role_sent,
)
.unwrap();
let response = String::from_utf8(response).unwrap();
assert!(response.find(": prefill").unwrap() < response.find("assistant").unwrap());
}
#[test]
fn streaming_errors_match_ds4_server_framing() {
let mut response = Vec::new();
send_sse_error(&mut response, "prefill failed").unwrap();
assert_eq!(
String::from_utf8(response).unwrap(),
"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-")
);
}
}