1030 lines
36 KiB
Rust
1030 lines
36 KiB
Rust
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 <think>), you MUST output your complete reasoning inside <think>...</think> BEFORE any tool calls or final response.\n\n\
|
||
Otherwise, output directly after </think> with tool calls or final response.\n\n\
|
||
### Available Tool Schemas\n\n";
|
||
|
||
pub(crate) struct ServerHandle {
|
||
stop: Arc<AtomicBool>,
|
||
thread: Option<JoinHandle<()>>,
|
||
metrics: Arc<Metrics>,
|
||
port: u16,
|
||
}
|
||
|
||
struct State {
|
||
generation: GenerationService,
|
||
preferences: Arc<RwLock<AppPreferences>>,
|
||
models_path: PathBuf,
|
||
cache_path: PathBuf,
|
||
sequence: AtomicU64,
|
||
connections: Arc<AtomicUsize>,
|
||
tool_memory: Mutex<HashMap<String, String>>,
|
||
metrics: Arc<Metrics>,
|
||
cors: bool,
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
struct ChatRequest {
|
||
#[serde(default)]
|
||
model: Option<String>,
|
||
#[serde(default)]
|
||
messages: Vec<ApiMessage>,
|
||
#[serde(default)]
|
||
tools: Vec<Value>,
|
||
#[serde(skip)]
|
||
tool_schemas: Vec<String>,
|
||
#[serde(default)]
|
||
tool_choice: Option<Value>,
|
||
#[serde(default)]
|
||
max_tokens: Option<i32>,
|
||
#[serde(default)]
|
||
max_completion_tokens: Option<i32>,
|
||
#[serde(default)]
|
||
temperature: Option<f32>,
|
||
#[serde(default)]
|
||
top_p: Option<f32>,
|
||
#[serde(default)]
|
||
min_p: Option<f32>,
|
||
#[serde(default)]
|
||
top_k: Option<i32>,
|
||
#[serde(default)]
|
||
seed: Option<u64>,
|
||
#[serde(default)]
|
||
stream: bool,
|
||
#[serde(default)]
|
||
stream_options: Option<StreamOptions>,
|
||
#[serde(default)]
|
||
thinking: Option<Value>,
|
||
#[serde(default)]
|
||
think: Option<bool>,
|
||
#[serde(default)]
|
||
reasoning_effort: Option<String>,
|
||
#[serde(default)]
|
||
stop: Option<OneOrMany>,
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
struct ApiMessage {
|
||
#[serde(default = "default_role")]
|
||
role: String,
|
||
#[serde(default)]
|
||
content: Value,
|
||
#[serde(default)]
|
||
reasoning_content: Value,
|
||
#[serde(default)]
|
||
tool_calls: Vec<ApiToolCall>,
|
||
#[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<String>),
|
||
}
|
||
|
||
struct ParsedRequest {
|
||
protocol: Protocol,
|
||
model_id: String,
|
||
messages: Vec<ChatTurn>,
|
||
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<u8>,
|
||
}
|
||
|
||
#[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<RwLock<AppPreferences>>,
|
||
models_path: PathBuf,
|
||
cache_path: PathBuf,
|
||
port: u16,
|
||
cors: bool,
|
||
metrics: Arc<Metrics>,
|
||
) -> Result<Self, String> {
|
||
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<State>, stop: Arc<AtomicBool>) {
|
||
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<AtomicUsize>);
|
||
|
||
impl ConnectionSlot {
|
||
fn acquire(active: &Arc<AtomicUsize>) -> Option<Self> {
|
||
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<ModelChoice> {
|
||
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::<Vec<_>>()
|
||
})
|
||
}
|
||
|
||
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<ModelChoice> {
|
||
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::<Vec<_>>();
|
||
|
||
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 [
|
||
"<DSML|tool_calls><DSML|invoke name=\"echo\"><DSML|parameter name=\"text\">hi</DSML|parameter></DSML|invoke></DSML|tool_calls>",
|
||
"<tool_calls><invoke name=\"echo\"><parameter name=\"text\">hi</parameter></invoke></tool_calls>",
|
||
] {
|
||
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::<String>();
|
||
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<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);
|
||
assert_eq!(content, "done");
|
||
assert_eq!(calls[0].function.name, "run");
|
||
assert_eq!(
|
||
calls[0].function.arguments,
|
||
r#"{"options":{"path":"/tmp"}}"#
|
||
);
|
||
|
||
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]
|
||
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<ApiMessage> = 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-")
|
||
);
|
||
}
|
||
}
|