Complete local HTTP endpoint parity
This commit is contained in:
2
migrations/20260725130000_add_endpoint_controls/down.sql
Normal file
2
migrations/20260725130000_add_endpoint_controls/down.sql
Normal file
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE preferences DROP COLUMN endpoint_cors;
|
||||
ALTER TABLE preferences DROP COLUMN endpoint_enabled;
|
||||
4
migrations/20260725130000_add_endpoint_controls/up.sql
Normal file
4
migrations/20260725130000_add_endpoint_controls/up.sql
Normal 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));
|
||||
34
src/app.rs
34
src/app.rs
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
));
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
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-")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user