diff --git a/src/a2ui.rs b/src/a2ui.rs index 3196322..53580cc 100644 --- a/src/a2ui.rs +++ b/src/a2ui.rs @@ -1,3 +1,14 @@ +mod evaluation; +mod validation; + +use evaluation::{evaluate, resolve, resolve_at}; +use validation::{validate_component, validate_function}; + +#[cfg(test)] +use evaluation::format_date; +#[cfg(test)] +use validation::ds4_catalog; + use regex::Regex; use serde_json::{Map, Value, json}; use std::cell::Cell; @@ -749,1106 +760,6 @@ pub(crate) fn first_failed_check(component: &Map, data: &Value) - }) } -fn validate_component(component: &Map, catalog_id: &str) -> Result<(), String> { - let id = required_string(component, "id")?; - let kind = required_string(component, "component")?; - let allowed = if catalog_id == CATALOG_ID { - ds4_catalog() - .pointer("/components") - .and_then(Value::as_object) - .is_some_and(|components| components.contains_key(kind)) - } else { - BASIC_COMPONENTS.contains(&kind) - }; - if !allowed { - return Err(format!( - "component `{id}` uses `{kind}`, which is not in catalog `{catalog_id}`" - )); - } - let required: &[&str] = match kind { - "Text" => &["text"], - "Image" => &["url"], - "Icon" => &["name"], - "Video" | "AudioPlayer" => &["url"], - "Row" | "Column" | "List" => &["children"], - "Card" => &["child"], - "Modal" => &["trigger", "content"], - "Tabs" => &["tabs"], - "Button" => &["child", "action"], - "TextField" => &["label"], - "CheckBox" => &["label", "value"], - "Slider" => &["value", "max"], - "DateTimeInput" => &["value"], - "ChoicePicker" => &["options", "value"], - "Chart" => &["chartType", "series"], - "Table" => &["columns", "rows"], - "Metric" => &["label", "value"], - "Timeline" => &["events"], - "Map" => &["locations"], - "MindMap" => &["nodes"], - "Form" => &["children"], - _ => &[], - }; - if let Some(field) = required - .iter() - .find(|field| !component.contains_key(**field)) - { - return Err(format!("component `{id}` ({kind}) requires `{field}`")); - } - let common = ["id", "component", "accessibility", "weight", "checks"]; - let specific: &[&str] = match kind { - "Text" => &["text", "variant"], - "Image" => &["url", "description", "fit", "variant"], - "Icon" => &["name"], - "Video" => &["url", "posterUrl"], - "AudioPlayer" => &["url", "description"], - "Divider" => &["axis"], - "Row" | "Column" => &["children", "justify", "align"], - "List" => &["children", "direction", "align"], - "Card" => &["child"], - "Modal" => &["trigger", "content"], - "Tabs" => &["tabs"], - "Button" => &["child", "variant", "action"], - "TextField" => &["label", "value", "placeholder", "variant"], - "CheckBox" => &["label", "value"], - "Slider" => &["label", "min", "max", "value", "steps"], - "DateTimeInput" => &["label", "value", "enableDate", "enableTime", "min", "max"], - "ChoicePicker" => &[ - "label", - "variant", - "options", - "value", - "displayStyle", - "filterable", - ], - "Chart" => &["title", "chartType", "series"], - "Table" => &["title", "columns", "rows"], - "Metric" => &["label", "value", "detail", "trend"], - "Timeline" => &["title", "events"], - "Map" => &["title", "locations"], - "MindMap" => &["title", "nodes"], - "Form" => &["title", "children", "submitLabel", "action"], - _ => &[], - }; - if let Some(field) = component - .keys() - .find(|field| !common.contains(&field.as_str()) && !specific.contains(&field.as_str())) - { - return Err(format!( - "component `{id}` ({kind}) contains unknown property `{field}`" - )); - } - if let Some(accessibility) = component.get("accessibility") { - let accessibility = object(Some(accessibility), "accessibility")?; - for field in ["label", "description"] { - validate_optional_dynamic(accessibility, field, Value::is_string, id)?; - } - } - if component - .get("weight") - .is_some_and(|value| !value.is_number()) - { - return Err(format!("component `{id}` weight must be a number")); - } - if component - .get("title") - .is_some_and(|value| !value.is_string()) - { - return Err(format!("component `{id}` title must be a string")); - } - if catalog_id != CATALOG_ID - && component.contains_key("checks") - && !matches!( - kind, - "Button" | "TextField" | "CheckBox" | "ChoicePicker" | "Slider" | "DateTimeInput" - ) - { - return Err(format!("component `{id}` ({kind}) does not support checks")); - } - match kind { - "Row" | "Column" | "List" | "Form" => { - validate_children(component.get("children").unwrap(), id)? - } - "Card" | "Button" => expect_string(component, "child", id)?, - "Modal" => { - expect_string(component, "trigger", id)?; - expect_string(component, "content", id)?; - } - "Tabs" => validate_tabs(component.get("tabs").unwrap(), id)?, - "ChoicePicker" => validate_options(component.get("options").unwrap(), id)?, - "Slider" => { - expect_number(component, "max", id)?; - if let Some(min) = component.get("min") - && !min.is_number() - { - return Err(format!("component `{id}` min must be a number")); - } - if let Some(steps) = component.get("steps") - && !steps.as_u64().is_some_and(|steps| steps > 0) - { - return Err(format!("component `{id}` steps must be a positive integer")); - } - } - _ => {} - } - match kind { - "Text" => validate_dynamic_field(component, "text", Value::is_string, id)?, - "Image" | "Video" | "AudioPlayer" => { - validate_dynamic_field(component, "url", Value::is_string, id)?; - validate_optional_dynamic(component, "description", Value::is_string, id)?; - validate_optional_dynamic(component, "posterUrl", Value::is_string, id)?; - } - "Icon" => validate_dynamic_field(component, "name", Value::is_string, id)?, - "TextField" | "DateTimeInput" => { - validate_optional_dynamic(component, "label", Value::is_string, id)?; - validate_optional_dynamic(component, "value", Value::is_string, id)?; - } - "CheckBox" => { - validate_dynamic_field(component, "label", Value::is_string, id)?; - validate_dynamic_field(component, "value", Value::is_boolean, id)?; - } - "Slider" => validate_dynamic_field(component, "value", Value::is_number, id)?, - "ChoicePicker" => { - validate_optional_dynamic(component, "label", Value::is_string, id)?; - validate_dynamic_field( - component, - "value", - |value| { - value - .as_array() - .is_some_and(|values| values.iter().all(Value::is_string)) - }, - id, - )?; - } - "Chart" | "Table" | "Timeline" | "Map" | "MindMap" => { - for field in required { - if *field != "chartType" { - validate_dynamic_field(component, field, Value::is_array, id)?; - } - } - } - _ => {} - } - match kind { - "Text" => validate_enum(component, "variant", &["caption", "body"], id)?, - "Image" => { - validate_enum( - component, - "fit", - &["contain", "cover", "fill", "none", "scaleDown"], - id, - )?; - validate_enum( - component, - "variant", - &[ - "icon", - "avatar", - "smallFeature", - "mediumFeature", - "largeFeature", - "header", - ], - id, - )?; - } - "Icon" if component.get("name").is_some_and(Value::is_string) => validate_enum( - component, - "name", - &[ - "accountCircle", - "add", - "arrowBack", - "arrowForward", - "attachFile", - "calendarToday", - "call", - "camera", - "check", - "close", - "delete", - "download", - "edit", - "event", - "error", - "fastForward", - "favorite", - "favoriteOff", - "folder", - "help", - "home", - "info", - "locationOn", - "lock", - "lockOpen", - "mail", - "menu", - "moreVert", - "moreHoriz", - "notificationsOff", - "notifications", - "pause", - "payment", - "person", - "phone", - "photo", - "play", - "print", - "refresh", - "rewind", - "search", - "send", - "settings", - "share", - "shoppingCart", - "skipNext", - "skipPrevious", - "star", - "starHalf", - "starOff", - "stop", - "upload", - "visibility", - "visibilityOff", - "volumeDown", - "volumeMute", - "volumeOff", - "volumeUp", - "warning", - ], - id, - )?, - "Divider" => validate_enum(component, "axis", &["horizontal", "vertical"], id)?, - "Row" | "Column" => { - validate_enum( - component, - "justify", - &[ - "start", - "center", - "end", - "spaceBetween", - "spaceAround", - "spaceEvenly", - "stretch", - ], - id, - )?; - validate_enum( - component, - "align", - &["start", "center", "end", "stretch"], - id, - )?; - } - "List" => { - validate_enum(component, "direction", &["vertical", "horizontal"], id)?; - validate_enum( - component, - "align", - &["start", "center", "end", "stretch"], - id, - )?; - } - "Button" => validate_enum( - component, - "variant", - &["default", "primary", "borderless"], - id, - )?, - "TextField" => validate_enum( - component, - "variant", - &["longText", "number", "shortText", "obscured"], - id, - )?, - "ChoicePicker" => { - validate_enum( - component, - "variant", - &["multipleSelection", "mutuallyExclusive"], - id, - )?; - validate_enum(component, "displayStyle", &["checkbox", "chips"], id)?; - optional_bool(component, "filterable")?; - } - "DateTimeInput" => { - optional_bool(component, "enableDate")?; - optional_bool(component, "enableTime")?; - validate_optional_dynamic(component, "min", Value::is_string, id)?; - validate_optional_dynamic(component, "max", Value::is_string, id)?; - } - "Chart" => validate_enum( - component, - "chartType", - &[ - "bar", - "line", - "area", - "stackedBar", - "pie", - "donut", - "heatmap", - ], - id, - )?, - _ => {} - } - if let Some(action) = component.get("action") { - validate_action(action, id)?; - } - if let Some(checks) = component.get("checks") { - let checks = checks - .as_array() - .ok_or_else(|| format!("component `{id}` checks must be an array"))?; - for check in checks { - let check = object(Some(check), "check")?; - reject_unknown(check, &["condition", "message"], "check")?; - let condition = check - .get("condition") - .ok_or_else(|| format!("component `{id}` check requires condition"))?; - required_string(check, "message")?; - validate_function(condition)?; - } - } - validate_dynamic_values(component)?; - Ok(()) -} - -fn validate_dynamic_field( - object: &Map, - field: &str, - literal: impl Fn(&Value) -> bool, - id: &str, -) -> Result<(), String> { - let value = object - .get(field) - .ok_or_else(|| format!("component `{id}` requires `{field}`"))?; - if literal(value) { - return Ok(()); - } - let dynamic = value - .as_object() - .ok_or_else(|| format!("component `{id}` {field} has the wrong literal or dynamic type"))?; - if let Some(path) = dynamic.get("path") { - if dynamic.len() == 1 && path.is_string() { - return Ok(()); - } - return Err(format!( - "component `{id}` {field} has an invalid data binding" - )); - } - validate_function(value).map_err(|error| format!("component `{id}` {field}: {error}")) -} - -fn validate_optional_dynamic( - object: &Map, - field: &str, - literal: impl Fn(&Value) -> bool, - id: &str, -) -> Result<(), String> { - if object.contains_key(field) { - validate_dynamic_field(object, field, literal, id) - } else { - Ok(()) - } -} - -fn ds4_catalog() -> &'static Value { - static CATALOG: OnceLock = OnceLock::new(); - CATALOG.get_or_init(|| { - serde_json::from_str(CATALOG_JSON).expect("embedded DS4Server A2UI catalog must be valid") - }) -} - -fn expect_string(object: &Map, field: &str, id: &str) -> Result<(), String> { - if !object.get(field).is_some_and(Value::is_string) { - return Err(format!("component `{id}` {field} must be a string")); - } - Ok(()) -} - -fn expect_number(object: &Map, field: &str, id: &str) -> Result<(), String> { - if !object.get(field).is_some_and(Value::is_number) { - return Err(format!("component `{id}` {field} must be a number")); - } - Ok(()) -} - -fn validate_enum( - object: &Map, - field: &str, - allowed: &[&str], - id: &str, -) -> Result<(), String> { - if let Some(value) = object.get(field) { - let value = value - .as_str() - .ok_or_else(|| format!("component `{id}` {field} must be a string"))?; - if !allowed.contains(&value) { - return Err(format!( - "component `{id}` has invalid {field} `{value}`; expected one of: {}", - allowed.join(", ") - )); - } - } - Ok(()) -} - -fn validate_children(value: &Value, id: &str) -> Result<(), String> { - if value - .as_array() - .is_some_and(|children| children.iter().all(Value::is_string)) - { - return Ok(()); - } - let template = object(Some(value), "children")?; - reject_unknown(template, &["componentId", "path"], "children template")?; - required_string(template, "componentId")?; - required_string(template, "path")?; - if !required_string(template, "path")?.starts_with('/') { - return Err(format!( - "component `{id}` child template path must be a JSON Pointer" - )); - } - Ok(()) -} - -fn validate_tabs(value: &Value, id: &str) -> Result<(), String> { - let tabs = value - .as_array() - .filter(|tabs| !tabs.is_empty()) - .ok_or_else(|| format!("component `{id}` tabs must be a non-empty array"))?; - for tab in tabs { - let tab = object(Some(tab), "tab")?; - reject_unknown(tab, &["title", "child"], "tab")?; - validate_dynamic_field(tab, "title", Value::is_string, id)?; - required_string(tab, "child")?; - } - Ok(()) -} - -fn validate_options(value: &Value, id: &str) -> Result<(), String> { - let options = value - .as_array() - .ok_or_else(|| format!("component `{id}` options must be an array"))?; - for option in options { - let option = object(Some(option), "choice option")?; - reject_unknown(option, &["label", "value"], "choice option")?; - validate_dynamic_field(option, "label", Value::is_string, id)?; - required_string(option, "value")?; - } - Ok(()) -} - -fn validate_action(value: &Value, id: &str) -> Result<(), String> { - let action = object(Some(value), "action")?; - if let Some(event) = action.get("event") { - reject_unknown(action, &["event"], "action")?; - let event = object(Some(event), "action.event")?; - reject_unknown( - event, - &["name", "context", "wantResponse", "responsePath"], - "action.event", - )?; - required_string(event, "name")?; - if let Some(context) = event.get("context") - && !context.is_object() - { - return Err(format!("component `{id}` action context must be an object")); - } - optional_bool(event, "wantResponse")?; - if let Some(path) = event.get("responsePath") - && !path.as_str().is_some_and(|path| path.starts_with('/')) - { - return Err(format!( - "component `{id}` responsePath must be a JSON Pointer" - )); - } - return Ok(()); - } - reject_unknown(action, &["functionCall"], "action")?; - validate_function( - action - .get("functionCall") - .ok_or_else(|| format!("component `{id}` action requires event or functionCall"))?, - ) -} - -fn validate_dynamic_values(value: &Map) -> Result<(), String> { - for value in value.values() { - if value.get("call").is_some() { - validate_function(value)?; - } - match value { - Value::Object(object) => validate_dynamic_values(object)?, - Value::Array(values) => { - for value in values { - if let Value::Object(object) = value { - validate_dynamic_values(object)?; - } - } - } - _ => {} - } - } - Ok(()) -} - -fn validate_function(value: &Value) -> Result<(), String> { - let function = object(Some(value), "function call")?; - reject_unknown(function, &["call", "args"], "function call")?; - let call = value - .get("call") - .and_then(Value::as_str) - .ok_or_else(|| "function call requires `call`".to_owned())?; - if !FUNCTIONS.contains(&call) { - return Err(format!("function `{call}` is not declared by the catalog")); - } - if value.get("args").is_some_and(|args| !args.is_object()) { - return Err(format!("function `{call}` requires object `args`")); - } - let args = value - .get("args") - .and_then(Value::as_object) - .cloned() - .unwrap_or_default(); - let allowed: &[&str] = match call { - "required" | "email" | "formatString" | "not" => &["value"], - "regex" => &["value", "pattern"], - "length" | "numeric" => &["value", "min", "max"], - "formatNumber" => &["value", "decimals", "grouping"], - "formatCurrency" => &["value", "currency", "decimals", "grouping"], - "formatDate" => &["value", "format"], - "pluralize" => &["value", "zero", "one", "two", "few", "many", "other"], - "openUrl" => &["url"], - "and" | "or" => &["values"], - "@index" => &["offset"], - _ => &[], - }; - reject_unknown(&args, allowed, &format!("function `{call}` args"))?; - let required: &[&str] = match call { - "required" | "regex" | "length" | "numeric" | "email" | "formatString" | "formatNumber" - | "not" | "pluralize" => &["value"], - "formatCurrency" => &["value", "currency"], - "formatDate" => &["value", "format"], - "openUrl" => &["url"], - "and" | "or" => &["values"], - "@index" => &[], - _ => &[], - }; - if let Some(field) = required.iter().find(|field| !args.contains_key(**field)) { - return Err(format!("function `{call}` requires argument `{field}`")); - } - if call == "regex" && !args.get("pattern").is_some_and(Value::is_string) { - return Err("function `regex` requires string argument `pattern`".into()); - } - if matches!(call, "length" | "numeric") - && !args.contains_key("min") - && !args.contains_key("max") - { - return Err(format!("function `{call}` requires `min` or `max`")); - } - if call == "pluralize" && !args.contains_key("other") { - return Err("function `pluralize` requires argument `other`".into()); - } - if matches!(call, "and" | "or") - && !args - .get("values") - .and_then(Value::as_array) - .is_some_and(|values| values.len() >= 2) - { - return Err(format!("function `{call}` requires at least two values")); - } - match call { - "regex" | "length" | "email" | "formatString" => { - validate_dynamic_field(&args, "value", Value::is_string, call)?; - } - "numeric" | "formatNumber" | "formatCurrency" | "pluralize" => { - validate_dynamic_field(&args, "value", Value::is_number, call)?; - } - "not" => validate_dynamic_field(&args, "value", Value::is_boolean, call)?, - "formatDate" => validate_dynamic_field(&args, "format", Value::is_string, call)?, - "openUrl" => { - let url = required_string(&args, "url")?; - url::Url::parse(url).map_err(|error| format!("invalid openUrl URL: {error}"))?; - } - "and" | "or" => { - for item in args["values"].as_array().unwrap() { - if !item.is_boolean() { - let item = item.as_object().ok_or_else(|| { - format!("function `{call}` values must be dynamic booleans") - })?; - if !item.contains_key("path") && !item.contains_key("call") { - return Err(format!("function `{call}` values must be dynamic booleans")); - } - } - } - } - _ => {} - } - for field in ["min", "max", "decimals", "offset"] { - if args.contains_key(field) { - validate_dynamic_field(&args, field, Value::is_number, call)?; - } - } - if args.contains_key("grouping") { - validate_dynamic_field(&args, "grouping", Value::is_boolean, call)?; - } - for field in ["currency", "zero", "one", "two", "few", "many", "other"] { - if args.contains_key(field) { - validate_dynamic_field(&args, field, Value::is_string, call)?; - } - } - Ok(()) -} - -fn resolve(value: &Value, data: &Value) -> Result { - resolve_at(value, data, data) -} - -fn resolve_at(value: &Value, data: &Value, context: &Value) -> Result { - if let Some(path) = value - .as_object() - .and_then(|value| value.get("path")) - .and_then(Value::as_str) - { - let source = if path.starts_with('/') { data } else { context }; - let pointer = if path.starts_with('/') { - normalize_pointer(path).to_owned() - } else { - format!("/{path}") - }; - return Ok(source.pointer(&pointer).cloned().unwrap_or(Value::Null)); - } - if let Some(call) = value - .as_object() - .and_then(|value| value.get("call")) - .and_then(Value::as_str) - { - return evaluate( - call, - value.get("args").unwrap_or(&Value::Null), - data, - context, - ); - } - match value { - Value::Array(values) => values - .iter() - .map(|value| resolve_at(value, data, context)) - .collect::, _>>() - .map(Value::Array), - Value::Object(values) => values - .iter() - .map(|(key, value)| Ok((key.clone(), resolve_at(value, data, context)?))) - .collect::, String>>() - .map(Value::Object), - _ => Ok(value.clone()), - } -} - -fn evaluate(call: &str, args: &Value, data: &Value, context: &Value) -> Result { - if !FUNCTIONS.contains(&call) { - return Err(format!("function `{call}` is not declared by the catalog")); - } - let args = args - .as_object() - .ok_or_else(|| format!("function `{call}` requires object args"))?; - let value = || resolve_at(args.get("value").unwrap_or(&Value::Null), data, context); - match call { - "required" => Ok(Value::Bool(match value()? { - Value::Null => false, - Value::String(value) => !value.is_empty(), - Value::Array(value) => !value.is_empty(), - Value::Object(value) => !value.is_empty(), - _ => true, - })), - "regex" => { - let pattern = args - .get("pattern") - .and_then(Value::as_str) - .ok_or_else(|| "regex requires a pattern".to_owned())?; - let regex = Regex::new(pattern).map_err(|error| format!("invalid regex: {error}"))?; - Ok(Value::Bool(regex.is_match(&display_value(&value()?)))) - } - "length" => { - let length = display_value(&value()?).chars().count() as u64; - let min = args.get("min").and_then(Value::as_u64).unwrap_or(0); - let max = args.get("max").and_then(Value::as_u64).unwrap_or(u64::MAX); - Ok(Value::Bool((min..=max).contains(&length))) - } - "numeric" => { - let resolved = value()?; - let number = resolved - .as_f64() - .or_else(|| resolved.as_str().and_then(|value| value.parse().ok())); - let min = args - .get("min") - .and_then(Value::as_f64) - .unwrap_or(f64::NEG_INFINITY); - let max = args - .get("max") - .and_then(Value::as_f64) - .unwrap_or(f64::INFINITY); - Ok(Value::Bool( - number.is_some_and(|number| (min..=max).contains(&number)), - )) - } - "email" => { - let email = display_value(&value()?); - let valid = Regex::new(r"^[^\s@]+@[^\s@]+\.[^\s@]+$") - .unwrap() - .is_match(&email); - Ok(Value::Bool(valid)) - } - "formatString" => Ok(Value::String(interpolate( - &display_value(&value()?), - data, - context, - )?)), - "formatNumber" => { - let number = value()?.as_f64().unwrap_or(0.0); - let digits = resolve_at(args.get("decimals").unwrap_or(&json!(2)), data, context)? - .as_f64() - .map(|value| value.max(0.0) as u64) - .unwrap_or(2) - .min(12) as usize; - let formatted = format!("{number:.digits$}"); - Ok(Value::String( - if resolve_at( - args.get("grouping").unwrap_or(&Value::Bool(true)), - data, - context, - )? - .as_bool() - .unwrap_or(true) - { - group_number(&formatted) - } else { - formatted - }, - )) - } - "formatCurrency" => { - let number = value()?.as_f64().unwrap_or(0.0); - let currency = display_value(&resolve_at( - args.get("currency").unwrap_or(&Value::Null), - data, - context, - )?); - let digits = resolve_at(args.get("decimals").unwrap_or(&json!(2)), data, context)? - .as_f64() - .map(|value| value.max(0.0) as u64) - .unwrap_or(2) - .min(12) as usize; - let formatted = format!("{number:.digits$}"); - let number = if resolve_at( - args.get("grouping").unwrap_or(&Value::Bool(true)), - data, - context, - )? - .as_bool() - .unwrap_or(true) - { - group_number(&formatted) - } else { - formatted - }; - Ok(Value::String(format!("{currency} {number}"))) - } - "formatDate" => { - let input = display_value(&value()?); - let date = - time::OffsetDateTime::parse(&input, &time::format_description::well_known::Rfc3339) - .map_err(|error| format!("formatDate requires an RFC 3339 value: {error}"))?; - let pattern = display_value(&resolve_at( - args.get("format").unwrap_or(&Value::Null), - data, - context, - )?); - Ok(Value::String(format_date(date, &pattern))) - } - "pluralize" => { - let number = value()?.as_f64().unwrap_or(0.0); - let key = if number == 0.0 && args.contains_key("zero") { - "zero" - } else if number == 1.0 && args.contains_key("one") { - "one" - } else if number == 2.0 && args.contains_key("two") { - "two" - } else { - "other" - }; - resolve_at(args.get(key).unwrap_or(&Value::Null), data, context) - } - "and" | "or" => { - let values = args - .get("values") - .and_then(Value::as_array) - .ok_or_else(|| format!("{call} requires array `values`"))?; - let values = values - .iter() - .map(|value| { - resolve_at(value, data, context)? - .as_bool() - .ok_or_else(|| format!("{call} requires boolean values")) - }) - .collect::, _>>()?; - Ok(Value::Bool(if call == "and" { - values.into_iter().all(|value| value) - } else { - values.into_iter().any(|value| value) - })) - } - "not" => { - Ok(Value::Bool(!value()?.as_bool().ok_or_else(|| { - "not requires a boolean value".to_owned() - })?)) - } - "openUrl" => Ok(Value::Null), - "@index" => { - let index = TEMPLATE_INDEX - .with(Cell::get) - .ok_or_else(|| "@index is only available in a list template".to_owned())?; - let offset = args - .get("offset") - .map(|value| resolve_at(value, data, context)) - .transpose()? - .and_then(|value| value.as_i64()) - .unwrap_or(0); - Ok(json!(index as i64 + offset)) - } - _ => unreachable!(), - } -} - -fn interpolate(template: &str, data: &Value, context: &Value) -> Result { - let mut output = String::new(); - let mut offset = 0; - while let Some(relative) = template[offset..].find("${") { - let start = offset + relative; - if start > offset && template.as_bytes()[start - 1] == b'\\' { - output.push_str(&template[offset..start - 1]); - output.push_str("${"); - offset = start + 2; - continue; - } - output.push_str(&template[offset..start]); - let end = expression_end(template, start + 2) - .ok_or_else(|| "formatString contains an unclosed expression".to_owned())?; - output.push_str(&display_value(&evaluate_expression( - &template[start + 2..end], - data, - context, - )?)); - offset = end + 1; - } - output.push_str(&template[offset..]); - Ok(output) -} - -fn expression_end(text: &str, mut offset: usize) -> Option { - let mut depth = 1; - let mut quote = None; - while offset < text.len() { - let character = text[offset..].chars().next()?; - if let Some(current) = quote { - if character == current && text.as_bytes().get(offset.wrapping_sub(1)) != Some(&b'\\') { - quote = None; - } - } else if matches!(character, '\'' | '"') { - quote = Some(character); - } else if text[offset..].starts_with("${") { - depth += 1; - offset += 2; - continue; - } else if character == '}' { - depth -= 1; - if depth == 0 { - return Some(offset); - } - } - offset += character.len_utf8(); - } - None -} - -fn evaluate_expression(expression: &str, data: &Value, context: &Value) -> Result { - let expression = expression.trim(); - if let Some(open) = expression.find('(') - && expression.ends_with(')') - { - let call = expression[..open].trim(); - let mut args = Map::new(); - for argument in split_expression_args(&expression[open + 1..expression.len() - 1]) { - let colon = top_level_separator(argument, ':') - .ok_or_else(|| format!("formatString argument `{argument}` must be named"))?; - let name = argument[..colon].trim(); - if name.is_empty() { - return Err("formatString contains an empty argument name".into()); - } - args.insert( - name.to_owned(), - expression_value(argument[colon + 1..].trim(), data, context)?, - ); - } - return evaluate(call, &Value::Object(args), data, context); - } - let (source, pointer) = if expression.starts_with('/') { - (data, normalize_pointer(expression).to_owned()) - } else { - (context, format!("/{expression}")) - }; - Ok(source.pointer(&pointer).cloned().unwrap_or(Value::Null)) -} - -fn split_expression_args(input: &str) -> Vec<&str> { - let mut arguments = Vec::new(); - let mut start = 0; - while let Some(relative) = top_level_separator(&input[start..], ',') { - arguments.push(input[start..start + relative].trim()); - start += relative + 1; - } - if !input[start..].trim().is_empty() { - arguments.push(input[start..].trim()); - } - arguments -} - -fn top_level_separator(input: &str, separator: char) -> Option { - let mut round = 0; - let mut braces = 0; - let mut quote = None; - for (index, character) in input.char_indices() { - if let Some(current) = quote { - if character == current && input.as_bytes().get(index.wrapping_sub(1)) != Some(&b'\\') { - quote = None; - } - continue; - } - match character { - '\'' | '"' => quote = Some(character), - '(' => round += 1, - ')' => round -= 1, - '{' => braces += 1, - '}' => braces -= 1, - _ if character == separator && round == 0 && braces == 0 => return Some(index), - _ => {} - } - } - None -} - -fn expression_value(expression: &str, data: &Value, context: &Value) -> Result { - if expression.starts_with("${") && expression.ends_with('}') { - return evaluate_expression(&expression[2..expression.len() - 1], data, context); - } - if expression.starts_with('\'') && expression.ends_with('\'') && expression.len() >= 2 { - return Ok(Value::String( - expression[1..expression.len() - 1].to_owned(), - )); - } - if let Ok(value) = serde_json::from_str(expression) { - return Ok(value); - } - if expression.starts_with('/') { - return Ok(data - .pointer(normalize_pointer(expression)) - .cloned() - .unwrap_or(Value::Null)); - } - Ok(Value::String(expression.to_owned())) -} - -fn group_number(number: &str) -> String { - let (whole, fraction) = number.split_once('.').unwrap_or((number, "")); - let (sign, digits) = whole - .strip_prefix('-') - .map_or(("", whole), |digits| ("-", digits)); - let mut grouped = String::with_capacity(number.len() + number.len() / 3); - grouped.push_str(sign); - for (index, digit) in digits.chars().enumerate() { - if index > 0 && (digits.len() - index).is_multiple_of(3) { - grouped.push(','); - } - grouped.push(digit); - } - if !fraction.is_empty() { - grouped.push('.'); - grouped.push_str(fraction); - } - grouped -} - -fn format_date(date: time::OffsetDateTime, pattern: &str) -> String { - const MONTHS: [&str; 12] = [ - "January", - "February", - "March", - "April", - "May", - "June", - "July", - "August", - "September", - "October", - "November", - "December", - ]; - const DAYS: [&str; 7] = [ - "Monday", - "Tuesday", - "Wednesday", - "Thursday", - "Friday", - "Saturday", - "Sunday", - ]; - let month = MONTHS[date.month() as usize - 1]; - let day = DAYS[date.weekday().number_days_from_monday() as usize]; - let hour_12 = match date.hour() % 12 { - 0 => 12, - hour => hour, - }; - let replacements = [ - ("EEEE", day.to_owned()), - ("MMMM", month.to_owned()), - ("yyyy", format!("{:04}", date.year())), - ("MMM", month[..3].to_owned()), - ("yy", format!("{:02}", date.year().rem_euclid(100))), - ("MM", format!("{:02}", date.month() as u8)), - ("dd", format!("{:02}", date.day())), - ("HH", format!("{:02}", date.hour())), - ("hh", format!("{hour_12:02}")), - ("mm", format!("{:02}", date.minute())), - ("ss", format!("{:02}", date.second())), - ("E", day[..3].to_owned()), - ("M", (date.month() as u8).to_string()), - ("d", date.day().to_string()), - ("H", date.hour().to_string()), - ("h", hour_12.to_string()), - ("a", if date.hour() < 12 { "AM" } else { "PM" }.to_owned()), - ]; - let mut output = String::new(); - let mut remaining = pattern; - while !remaining.is_empty() { - if let Some((token, replacement)) = replacements - .iter() - .find(|(token, _)| remaining.starts_with(token)) - { - output.push_str(replacement); - remaining = &remaining[token.len()..]; - } else { - let character = remaining.chars().next().unwrap(); - output.push(character); - remaining = &remaining[character.len_utf8()..]; - } - } - output -} - fn set_pointer(root: &mut Value, path: &str, value: Option) -> Result<(), String> { if path.is_empty() || path == "/" { *root = value.unwrap_or(Value::Null); diff --git a/src/a2ui/evaluation.rs b/src/a2ui/evaluation.rs new file mode 100644 index 0000000..e978266 --- /dev/null +++ b/src/a2ui/evaluation.rs @@ -0,0 +1,451 @@ +use super::*; + +pub(super) fn resolve(value: &Value, data: &Value) -> Result { + resolve_at(value, data, data) +} + +pub(super) fn resolve_at(value: &Value, data: &Value, context: &Value) -> Result { + if let Some(path) = value + .as_object() + .and_then(|value| value.get("path")) + .and_then(Value::as_str) + { + let source = if path.starts_with('/') { data } else { context }; + let pointer = if path.starts_with('/') { + normalize_pointer(path).to_owned() + } else { + format!("/{path}") + }; + return Ok(source.pointer(&pointer).cloned().unwrap_or(Value::Null)); + } + if let Some(call) = value + .as_object() + .and_then(|value| value.get("call")) + .and_then(Value::as_str) + { + return evaluate( + call, + value.get("args").unwrap_or(&Value::Null), + data, + context, + ); + } + match value { + Value::Array(values) => values + .iter() + .map(|value| resolve_at(value, data, context)) + .collect::, _>>() + .map(Value::Array), + Value::Object(values) => values + .iter() + .map(|(key, value)| Ok((key.clone(), resolve_at(value, data, context)?))) + .collect::, String>>() + .map(Value::Object), + _ => Ok(value.clone()), + } +} + +pub(super) fn evaluate( + call: &str, + args: &Value, + data: &Value, + context: &Value, +) -> Result { + if !FUNCTIONS.contains(&call) { + return Err(format!("function `{call}` is not declared by the catalog")); + } + let args = args + .as_object() + .ok_or_else(|| format!("function `{call}` requires object args"))?; + let value = || resolve_at(args.get("value").unwrap_or(&Value::Null), data, context); + match call { + "required" => Ok(Value::Bool(match value()? { + Value::Null => false, + Value::String(value) => !value.is_empty(), + Value::Array(value) => !value.is_empty(), + Value::Object(value) => !value.is_empty(), + _ => true, + })), + "regex" => { + let pattern = args + .get("pattern") + .and_then(Value::as_str) + .ok_or_else(|| "regex requires a pattern".to_owned())?; + let regex = Regex::new(pattern).map_err(|error| format!("invalid regex: {error}"))?; + Ok(Value::Bool(regex.is_match(&display_value(&value()?)))) + } + "length" => { + let length = display_value(&value()?).chars().count() as u64; + let min = args.get("min").and_then(Value::as_u64).unwrap_or(0); + let max = args.get("max").and_then(Value::as_u64).unwrap_or(u64::MAX); + Ok(Value::Bool((min..=max).contains(&length))) + } + "numeric" => { + let resolved = value()?; + let number = resolved + .as_f64() + .or_else(|| resolved.as_str().and_then(|value| value.parse().ok())); + let min = args + .get("min") + .and_then(Value::as_f64) + .unwrap_or(f64::NEG_INFINITY); + let max = args + .get("max") + .and_then(Value::as_f64) + .unwrap_or(f64::INFINITY); + Ok(Value::Bool( + number.is_some_and(|number| (min..=max).contains(&number)), + )) + } + "email" => { + let email = display_value(&value()?); + let valid = Regex::new(r"^[^\s@]+@[^\s@]+\.[^\s@]+$") + .unwrap() + .is_match(&email); + Ok(Value::Bool(valid)) + } + "formatString" => Ok(Value::String(interpolate( + &display_value(&value()?), + data, + context, + )?)), + "formatNumber" => { + let number = value()?.as_f64().unwrap_or(0.0); + let digits = resolve_at(args.get("decimals").unwrap_or(&json!(2)), data, context)? + .as_f64() + .map(|value| value.max(0.0) as u64) + .unwrap_or(2) + .min(12) as usize; + let formatted = format!("{number:.digits$}"); + Ok(Value::String( + if resolve_at( + args.get("grouping").unwrap_or(&Value::Bool(true)), + data, + context, + )? + .as_bool() + .unwrap_or(true) + { + group_number(&formatted) + } else { + formatted + }, + )) + } + "formatCurrency" => { + let number = value()?.as_f64().unwrap_or(0.0); + let currency = display_value(&resolve_at( + args.get("currency").unwrap_or(&Value::Null), + data, + context, + )?); + let digits = resolve_at(args.get("decimals").unwrap_or(&json!(2)), data, context)? + .as_f64() + .map(|value| value.max(0.0) as u64) + .unwrap_or(2) + .min(12) as usize; + let formatted = format!("{number:.digits$}"); + let number = if resolve_at( + args.get("grouping").unwrap_or(&Value::Bool(true)), + data, + context, + )? + .as_bool() + .unwrap_or(true) + { + group_number(&formatted) + } else { + formatted + }; + Ok(Value::String(format!("{currency} {number}"))) + } + "formatDate" => { + let input = display_value(&value()?); + let date = + time::OffsetDateTime::parse(&input, &time::format_description::well_known::Rfc3339) + .map_err(|error| format!("formatDate requires an RFC 3339 value: {error}"))?; + let pattern = display_value(&resolve_at( + args.get("format").unwrap_or(&Value::Null), + data, + context, + )?); + Ok(Value::String(format_date(date, &pattern))) + } + "pluralize" => { + let number = value()?.as_f64().unwrap_or(0.0); + let key = if number == 0.0 && args.contains_key("zero") { + "zero" + } else if number == 1.0 && args.contains_key("one") { + "one" + } else if number == 2.0 && args.contains_key("two") { + "two" + } else { + "other" + }; + resolve_at(args.get(key).unwrap_or(&Value::Null), data, context) + } + "and" | "or" => { + let values = args + .get("values") + .and_then(Value::as_array) + .ok_or_else(|| format!("{call} requires array `values`"))?; + let values = values + .iter() + .map(|value| { + resolve_at(value, data, context)? + .as_bool() + .ok_or_else(|| format!("{call} requires boolean values")) + }) + .collect::, _>>()?; + Ok(Value::Bool(if call == "and" { + values.into_iter().all(|value| value) + } else { + values.into_iter().any(|value| value) + })) + } + "not" => { + Ok(Value::Bool(!value()?.as_bool().ok_or_else(|| { + "not requires a boolean value".to_owned() + })?)) + } + "openUrl" => Ok(Value::Null), + "@index" => { + let index = TEMPLATE_INDEX + .with(Cell::get) + .ok_or_else(|| "@index is only available in a list template".to_owned())?; + let offset = args + .get("offset") + .map(|value| resolve_at(value, data, context)) + .transpose()? + .and_then(|value| value.as_i64()) + .unwrap_or(0); + Ok(json!(index as i64 + offset)) + } + _ => unreachable!(), + } +} + +fn interpolate(template: &str, data: &Value, context: &Value) -> Result { + let mut output = String::new(); + let mut offset = 0; + while let Some(relative) = template[offset..].find("${") { + let start = offset + relative; + if start > offset && template.as_bytes()[start - 1] == b'\\' { + output.push_str(&template[offset..start - 1]); + output.push_str("${"); + offset = start + 2; + continue; + } + output.push_str(&template[offset..start]); + let end = expression_end(template, start + 2) + .ok_or_else(|| "formatString contains an unclosed expression".to_owned())?; + output.push_str(&display_value(&evaluate_expression( + &template[start + 2..end], + data, + context, + )?)); + offset = end + 1; + } + output.push_str(&template[offset..]); + Ok(output) +} + +fn expression_end(text: &str, mut offset: usize) -> Option { + let mut depth = 1; + let mut quote = None; + while offset < text.len() { + let character = text[offset..].chars().next()?; + if let Some(current) = quote { + if character == current && text.as_bytes().get(offset.wrapping_sub(1)) != Some(&b'\\') { + quote = None; + } + } else if matches!(character, '\'' | '"') { + quote = Some(character); + } else if text[offset..].starts_with("${") { + depth += 1; + offset += 2; + continue; + } else if character == '}' { + depth -= 1; + if depth == 0 { + return Some(offset); + } + } + offset += character.len_utf8(); + } + None +} + +fn evaluate_expression(expression: &str, data: &Value, context: &Value) -> Result { + let expression = expression.trim(); + if let Some(open) = expression.find('(') + && expression.ends_with(')') + { + let call = expression[..open].trim(); + let mut args = Map::new(); + for argument in split_expression_args(&expression[open + 1..expression.len() - 1]) { + let colon = top_level_separator(argument, ':') + .ok_or_else(|| format!("formatString argument `{argument}` must be named"))?; + let name = argument[..colon].trim(); + if name.is_empty() { + return Err("formatString contains an empty argument name".into()); + } + args.insert( + name.to_owned(), + expression_value(argument[colon + 1..].trim(), data, context)?, + ); + } + return evaluate(call, &Value::Object(args), data, context); + } + let (source, pointer) = if expression.starts_with('/') { + (data, normalize_pointer(expression).to_owned()) + } else { + (context, format!("/{expression}")) + }; + Ok(source.pointer(&pointer).cloned().unwrap_or(Value::Null)) +} + +fn split_expression_args(input: &str) -> Vec<&str> { + let mut arguments = Vec::new(); + let mut start = 0; + while let Some(relative) = top_level_separator(&input[start..], ',') { + arguments.push(input[start..start + relative].trim()); + start += relative + 1; + } + if !input[start..].trim().is_empty() { + arguments.push(input[start..].trim()); + } + arguments +} + +fn top_level_separator(input: &str, separator: char) -> Option { + let mut round = 0; + let mut braces = 0; + let mut quote = None; + for (index, character) in input.char_indices() { + if let Some(current) = quote { + if character == current && input.as_bytes().get(index.wrapping_sub(1)) != Some(&b'\\') { + quote = None; + } + continue; + } + match character { + '\'' | '"' => quote = Some(character), + '(' => round += 1, + ')' => round -= 1, + '{' => braces += 1, + '}' => braces -= 1, + _ if character == separator && round == 0 && braces == 0 => return Some(index), + _ => {} + } + } + None +} + +fn expression_value(expression: &str, data: &Value, context: &Value) -> Result { + if expression.starts_with("${") && expression.ends_with('}') { + return evaluate_expression(&expression[2..expression.len() - 1], data, context); + } + if expression.starts_with('\'') && expression.ends_with('\'') && expression.len() >= 2 { + return Ok(Value::String( + expression[1..expression.len() - 1].to_owned(), + )); + } + if let Ok(value) = serde_json::from_str(expression) { + return Ok(value); + } + if expression.starts_with('/') { + return Ok(data + .pointer(normalize_pointer(expression)) + .cloned() + .unwrap_or(Value::Null)); + } + Ok(Value::String(expression.to_owned())) +} + +fn group_number(number: &str) -> String { + let (whole, fraction) = number.split_once('.').unwrap_or((number, "")); + let (sign, digits) = whole + .strip_prefix('-') + .map_or(("", whole), |digits| ("-", digits)); + let mut grouped = String::with_capacity(number.len() + number.len() / 3); + grouped.push_str(sign); + for (index, digit) in digits.chars().enumerate() { + if index > 0 && (digits.len() - index).is_multiple_of(3) { + grouped.push(','); + } + grouped.push(digit); + } + if !fraction.is_empty() { + grouped.push('.'); + grouped.push_str(fraction); + } + grouped +} + +pub(super) fn format_date(date: time::OffsetDateTime, pattern: &str) -> String { + const MONTHS: [&str; 12] = [ + "January", + "February", + "March", + "April", + "May", + "June", + "July", + "August", + "September", + "October", + "November", + "December", + ]; + const DAYS: [&str; 7] = [ + "Monday", + "Tuesday", + "Wednesday", + "Thursday", + "Friday", + "Saturday", + "Sunday", + ]; + let month = MONTHS[date.month() as usize - 1]; + let day = DAYS[date.weekday().number_days_from_monday() as usize]; + let hour_12 = match date.hour() % 12 { + 0 => 12, + hour => hour, + }; + let replacements = [ + ("EEEE", day.to_owned()), + ("MMMM", month.to_owned()), + ("yyyy", format!("{:04}", date.year())), + ("MMM", month[..3].to_owned()), + ("yy", format!("{:02}", date.year().rem_euclid(100))), + ("MM", format!("{:02}", date.month() as u8)), + ("dd", format!("{:02}", date.day())), + ("HH", format!("{:02}", date.hour())), + ("hh", format!("{hour_12:02}")), + ("mm", format!("{:02}", date.minute())), + ("ss", format!("{:02}", date.second())), + ("E", day[..3].to_owned()), + ("M", (date.month() as u8).to_string()), + ("d", date.day().to_string()), + ("H", date.hour().to_string()), + ("h", hour_12.to_string()), + ("a", if date.hour() < 12 { "AM" } else { "PM" }.to_owned()), + ]; + let mut output = String::new(); + let mut remaining = pattern; + while !remaining.is_empty() { + if let Some((token, replacement)) = replacements + .iter() + .find(|(token, _)| remaining.starts_with(token)) + { + output.push_str(replacement); + remaining = &remaining[token.len()..]; + } else { + let character = remaining.chars().next().unwrap(); + output.push(character); + remaining = &remaining[character.len_utf8()..]; + } + } + output +} diff --git a/src/a2ui/validation.rs b/src/a2ui/validation.rs new file mode 100644 index 0000000..9f30379 --- /dev/null +++ b/src/a2ui/validation.rs @@ -0,0 +1,659 @@ +use super::*; + +pub(super) fn validate_component( + component: &Map, + catalog_id: &str, +) -> Result<(), String> { + let id = required_string(component, "id")?; + let kind = required_string(component, "component")?; + let allowed = if catalog_id == CATALOG_ID { + ds4_catalog() + .pointer("/components") + .and_then(Value::as_object) + .is_some_and(|components| components.contains_key(kind)) + } else { + BASIC_COMPONENTS.contains(&kind) + }; + if !allowed { + return Err(format!( + "component `{id}` uses `{kind}`, which is not in catalog `{catalog_id}`" + )); + } + let required: &[&str] = match kind { + "Text" => &["text"], + "Image" => &["url"], + "Icon" => &["name"], + "Video" | "AudioPlayer" => &["url"], + "Row" | "Column" | "List" => &["children"], + "Card" => &["child"], + "Modal" => &["trigger", "content"], + "Tabs" => &["tabs"], + "Button" => &["child", "action"], + "TextField" => &["label"], + "CheckBox" => &["label", "value"], + "Slider" => &["value", "max"], + "DateTimeInput" => &["value"], + "ChoicePicker" => &["options", "value"], + "Chart" => &["chartType", "series"], + "Table" => &["columns", "rows"], + "Metric" => &["label", "value"], + "Timeline" => &["events"], + "Map" => &["locations"], + "MindMap" => &["nodes"], + "Form" => &["children"], + _ => &[], + }; + if let Some(field) = required + .iter() + .find(|field| !component.contains_key(**field)) + { + return Err(format!("component `{id}` ({kind}) requires `{field}`")); + } + let common = ["id", "component", "accessibility", "weight", "checks"]; + let specific: &[&str] = match kind { + "Text" => &["text", "variant"], + "Image" => &["url", "description", "fit", "variant"], + "Icon" => &["name"], + "Video" => &["url", "posterUrl"], + "AudioPlayer" => &["url", "description"], + "Divider" => &["axis"], + "Row" | "Column" => &["children", "justify", "align"], + "List" => &["children", "direction", "align"], + "Card" => &["child"], + "Modal" => &["trigger", "content"], + "Tabs" => &["tabs"], + "Button" => &["child", "variant", "action"], + "TextField" => &["label", "value", "placeholder", "variant"], + "CheckBox" => &["label", "value"], + "Slider" => &["label", "min", "max", "value", "steps"], + "DateTimeInput" => &["label", "value", "enableDate", "enableTime", "min", "max"], + "ChoicePicker" => &[ + "label", + "variant", + "options", + "value", + "displayStyle", + "filterable", + ], + "Chart" => &["title", "chartType", "series"], + "Table" => &["title", "columns", "rows"], + "Metric" => &["label", "value", "detail", "trend"], + "Timeline" => &["title", "events"], + "Map" => &["title", "locations"], + "MindMap" => &["title", "nodes"], + "Form" => &["title", "children", "submitLabel", "action"], + _ => &[], + }; + if let Some(field) = component + .keys() + .find(|field| !common.contains(&field.as_str()) && !specific.contains(&field.as_str())) + { + return Err(format!( + "component `{id}` ({kind}) contains unknown property `{field}`" + )); + } + if let Some(accessibility) = component.get("accessibility") { + let accessibility = object(Some(accessibility), "accessibility")?; + for field in ["label", "description"] { + validate_optional_dynamic(accessibility, field, Value::is_string, id)?; + } + } + if component + .get("weight") + .is_some_and(|value| !value.is_number()) + { + return Err(format!("component `{id}` weight must be a number")); + } + if component + .get("title") + .is_some_and(|value| !value.is_string()) + { + return Err(format!("component `{id}` title must be a string")); + } + if catalog_id != CATALOG_ID + && component.contains_key("checks") + && !matches!( + kind, + "Button" | "TextField" | "CheckBox" | "ChoicePicker" | "Slider" | "DateTimeInput" + ) + { + return Err(format!("component `{id}` ({kind}) does not support checks")); + } + match kind { + "Row" | "Column" | "List" | "Form" => { + validate_children(component.get("children").unwrap(), id)? + } + "Card" | "Button" => expect_string(component, "child", id)?, + "Modal" => { + expect_string(component, "trigger", id)?; + expect_string(component, "content", id)?; + } + "Tabs" => validate_tabs(component.get("tabs").unwrap(), id)?, + "ChoicePicker" => validate_options(component.get("options").unwrap(), id)?, + "Slider" => { + expect_number(component, "max", id)?; + if let Some(min) = component.get("min") + && !min.is_number() + { + return Err(format!("component `{id}` min must be a number")); + } + if let Some(steps) = component.get("steps") + && !steps.as_u64().is_some_and(|steps| steps > 0) + { + return Err(format!("component `{id}` steps must be a positive integer")); + } + } + _ => {} + } + match kind { + "Text" => validate_dynamic_field(component, "text", Value::is_string, id)?, + "Image" | "Video" | "AudioPlayer" => { + validate_dynamic_field(component, "url", Value::is_string, id)?; + validate_optional_dynamic(component, "description", Value::is_string, id)?; + validate_optional_dynamic(component, "posterUrl", Value::is_string, id)?; + } + "Icon" => validate_dynamic_field(component, "name", Value::is_string, id)?, + "TextField" | "DateTimeInput" => { + validate_optional_dynamic(component, "label", Value::is_string, id)?; + validate_optional_dynamic(component, "value", Value::is_string, id)?; + } + "CheckBox" => { + validate_dynamic_field(component, "label", Value::is_string, id)?; + validate_dynamic_field(component, "value", Value::is_boolean, id)?; + } + "Slider" => validate_dynamic_field(component, "value", Value::is_number, id)?, + "ChoicePicker" => { + validate_optional_dynamic(component, "label", Value::is_string, id)?; + validate_dynamic_field( + component, + "value", + |value| { + value + .as_array() + .is_some_and(|values| values.iter().all(Value::is_string)) + }, + id, + )?; + } + "Chart" | "Table" | "Timeline" | "Map" | "MindMap" => { + for field in required { + if *field != "chartType" { + validate_dynamic_field(component, field, Value::is_array, id)?; + } + } + } + _ => {} + } + match kind { + "Text" => validate_enum(component, "variant", &["caption", "body"], id)?, + "Image" => { + validate_enum( + component, + "fit", + &["contain", "cover", "fill", "none", "scaleDown"], + id, + )?; + validate_enum( + component, + "variant", + &[ + "icon", + "avatar", + "smallFeature", + "mediumFeature", + "largeFeature", + "header", + ], + id, + )?; + } + "Icon" if component.get("name").is_some_and(Value::is_string) => validate_enum( + component, + "name", + &[ + "accountCircle", + "add", + "arrowBack", + "arrowForward", + "attachFile", + "calendarToday", + "call", + "camera", + "check", + "close", + "delete", + "download", + "edit", + "event", + "error", + "fastForward", + "favorite", + "favoriteOff", + "folder", + "help", + "home", + "info", + "locationOn", + "lock", + "lockOpen", + "mail", + "menu", + "moreVert", + "moreHoriz", + "notificationsOff", + "notifications", + "pause", + "payment", + "person", + "phone", + "photo", + "play", + "print", + "refresh", + "rewind", + "search", + "send", + "settings", + "share", + "shoppingCart", + "skipNext", + "skipPrevious", + "star", + "starHalf", + "starOff", + "stop", + "upload", + "visibility", + "visibilityOff", + "volumeDown", + "volumeMute", + "volumeOff", + "volumeUp", + "warning", + ], + id, + )?, + "Divider" => validate_enum(component, "axis", &["horizontal", "vertical"], id)?, + "Row" | "Column" => { + validate_enum( + component, + "justify", + &[ + "start", + "center", + "end", + "spaceBetween", + "spaceAround", + "spaceEvenly", + "stretch", + ], + id, + )?; + validate_enum( + component, + "align", + &["start", "center", "end", "stretch"], + id, + )?; + } + "List" => { + validate_enum(component, "direction", &["vertical", "horizontal"], id)?; + validate_enum( + component, + "align", + &["start", "center", "end", "stretch"], + id, + )?; + } + "Button" => validate_enum( + component, + "variant", + &["default", "primary", "borderless"], + id, + )?, + "TextField" => validate_enum( + component, + "variant", + &["longText", "number", "shortText", "obscured"], + id, + )?, + "ChoicePicker" => { + validate_enum( + component, + "variant", + &["multipleSelection", "mutuallyExclusive"], + id, + )?; + validate_enum(component, "displayStyle", &["checkbox", "chips"], id)?; + optional_bool(component, "filterable")?; + } + "DateTimeInput" => { + optional_bool(component, "enableDate")?; + optional_bool(component, "enableTime")?; + validate_optional_dynamic(component, "min", Value::is_string, id)?; + validate_optional_dynamic(component, "max", Value::is_string, id)?; + } + "Chart" => validate_enum( + component, + "chartType", + &[ + "bar", + "line", + "area", + "stackedBar", + "pie", + "donut", + "heatmap", + ], + id, + )?, + _ => {} + } + if let Some(action) = component.get("action") { + validate_action(action, id)?; + } + if let Some(checks) = component.get("checks") { + let checks = checks + .as_array() + .ok_or_else(|| format!("component `{id}` checks must be an array"))?; + for check in checks { + let check = object(Some(check), "check")?; + reject_unknown(check, &["condition", "message"], "check")?; + let condition = check + .get("condition") + .ok_or_else(|| format!("component `{id}` check requires condition"))?; + required_string(check, "message")?; + validate_function(condition)?; + } + } + validate_dynamic_values(component)?; + Ok(()) +} + +fn validate_dynamic_field( + object: &Map, + field: &str, + literal: impl Fn(&Value) -> bool, + id: &str, +) -> Result<(), String> { + let value = object + .get(field) + .ok_or_else(|| format!("component `{id}` requires `{field}`"))?; + if literal(value) { + return Ok(()); + } + let dynamic = value + .as_object() + .ok_or_else(|| format!("component `{id}` {field} has the wrong literal or dynamic type"))?; + if let Some(path) = dynamic.get("path") { + if dynamic.len() == 1 && path.is_string() { + return Ok(()); + } + return Err(format!( + "component `{id}` {field} has an invalid data binding" + )); + } + validate_function(value).map_err(|error| format!("component `{id}` {field}: {error}")) +} + +fn validate_optional_dynamic( + object: &Map, + field: &str, + literal: impl Fn(&Value) -> bool, + id: &str, +) -> Result<(), String> { + if object.contains_key(field) { + validate_dynamic_field(object, field, literal, id) + } else { + Ok(()) + } +} + +pub(super) fn ds4_catalog() -> &'static Value { + static CATALOG: OnceLock = OnceLock::new(); + CATALOG.get_or_init(|| { + serde_json::from_str(CATALOG_JSON).expect("embedded DS4Server A2UI catalog must be valid") + }) +} + +fn expect_string(object: &Map, field: &str, id: &str) -> Result<(), String> { + if !object.get(field).is_some_and(Value::is_string) { + return Err(format!("component `{id}` {field} must be a string")); + } + Ok(()) +} + +fn expect_number(object: &Map, field: &str, id: &str) -> Result<(), String> { + if !object.get(field).is_some_and(Value::is_number) { + return Err(format!("component `{id}` {field} must be a number")); + } + Ok(()) +} + +fn validate_enum( + object: &Map, + field: &str, + allowed: &[&str], + id: &str, +) -> Result<(), String> { + if let Some(value) = object.get(field) { + let value = value + .as_str() + .ok_or_else(|| format!("component `{id}` {field} must be a string"))?; + if !allowed.contains(&value) { + return Err(format!( + "component `{id}` has invalid {field} `{value}`; expected one of: {}", + allowed.join(", ") + )); + } + } + Ok(()) +} + +fn validate_children(value: &Value, id: &str) -> Result<(), String> { + if value + .as_array() + .is_some_and(|children| children.iter().all(Value::is_string)) + { + return Ok(()); + } + let template = object(Some(value), "children")?; + reject_unknown(template, &["componentId", "path"], "children template")?; + required_string(template, "componentId")?; + required_string(template, "path")?; + if !required_string(template, "path")?.starts_with('/') { + return Err(format!( + "component `{id}` child template path must be a JSON Pointer" + )); + } + Ok(()) +} + +fn validate_tabs(value: &Value, id: &str) -> Result<(), String> { + let tabs = value + .as_array() + .filter(|tabs| !tabs.is_empty()) + .ok_or_else(|| format!("component `{id}` tabs must be a non-empty array"))?; + for tab in tabs { + let tab = object(Some(tab), "tab")?; + reject_unknown(tab, &["title", "child"], "tab")?; + validate_dynamic_field(tab, "title", Value::is_string, id)?; + required_string(tab, "child")?; + } + Ok(()) +} + +fn validate_options(value: &Value, id: &str) -> Result<(), String> { + let options = value + .as_array() + .ok_or_else(|| format!("component `{id}` options must be an array"))?; + for option in options { + let option = object(Some(option), "choice option")?; + reject_unknown(option, &["label", "value"], "choice option")?; + validate_dynamic_field(option, "label", Value::is_string, id)?; + required_string(option, "value")?; + } + Ok(()) +} + +fn validate_action(value: &Value, id: &str) -> Result<(), String> { + let action = object(Some(value), "action")?; + if let Some(event) = action.get("event") { + reject_unknown(action, &["event"], "action")?; + let event = object(Some(event), "action.event")?; + reject_unknown( + event, + &["name", "context", "wantResponse", "responsePath"], + "action.event", + )?; + required_string(event, "name")?; + if let Some(context) = event.get("context") + && !context.is_object() + { + return Err(format!("component `{id}` action context must be an object")); + } + optional_bool(event, "wantResponse")?; + if let Some(path) = event.get("responsePath") + && !path.as_str().is_some_and(|path| path.starts_with('/')) + { + return Err(format!( + "component `{id}` responsePath must be a JSON Pointer" + )); + } + return Ok(()); + } + reject_unknown(action, &["functionCall"], "action")?; + validate_function( + action + .get("functionCall") + .ok_or_else(|| format!("component `{id}` action requires event or functionCall"))?, + ) +} + +fn validate_dynamic_values(value: &Map) -> Result<(), String> { + for value in value.values() { + if value.get("call").is_some() { + validate_function(value)?; + } + match value { + Value::Object(object) => validate_dynamic_values(object)?, + Value::Array(values) => { + for value in values { + if let Value::Object(object) = value { + validate_dynamic_values(object)?; + } + } + } + _ => {} + } + } + Ok(()) +} + +pub(super) fn validate_function(value: &Value) -> Result<(), String> { + let function = object(Some(value), "function call")?; + reject_unknown(function, &["call", "args"], "function call")?; + let call = value + .get("call") + .and_then(Value::as_str) + .ok_or_else(|| "function call requires `call`".to_owned())?; + if !FUNCTIONS.contains(&call) { + return Err(format!("function `{call}` is not declared by the catalog")); + } + if value.get("args").is_some_and(|args| !args.is_object()) { + return Err(format!("function `{call}` requires object `args`")); + } + let args = value + .get("args") + .and_then(Value::as_object) + .cloned() + .unwrap_or_default(); + let allowed: &[&str] = match call { + "required" | "email" | "formatString" | "not" => &["value"], + "regex" => &["value", "pattern"], + "length" | "numeric" => &["value", "min", "max"], + "formatNumber" => &["value", "decimals", "grouping"], + "formatCurrency" => &["value", "currency", "decimals", "grouping"], + "formatDate" => &["value", "format"], + "pluralize" => &["value", "zero", "one", "two", "few", "many", "other"], + "openUrl" => &["url"], + "and" | "or" => &["values"], + "@index" => &["offset"], + _ => &[], + }; + reject_unknown(&args, allowed, &format!("function `{call}` args"))?; + let required: &[&str] = match call { + "required" | "regex" | "length" | "numeric" | "email" | "formatString" | "formatNumber" + | "not" | "pluralize" => &["value"], + "formatCurrency" => &["value", "currency"], + "formatDate" => &["value", "format"], + "openUrl" => &["url"], + "and" | "or" => &["values"], + "@index" => &[], + _ => &[], + }; + if let Some(field) = required.iter().find(|field| !args.contains_key(**field)) { + return Err(format!("function `{call}` requires argument `{field}`")); + } + if call == "regex" && !args.get("pattern").is_some_and(Value::is_string) { + return Err("function `regex` requires string argument `pattern`".into()); + } + if matches!(call, "length" | "numeric") + && !args.contains_key("min") + && !args.contains_key("max") + { + return Err(format!("function `{call}` requires `min` or `max`")); + } + if call == "pluralize" && !args.contains_key("other") { + return Err("function `pluralize` requires argument `other`".into()); + } + if matches!(call, "and" | "or") + && !args + .get("values") + .and_then(Value::as_array) + .is_some_and(|values| values.len() >= 2) + { + return Err(format!("function `{call}` requires at least two values")); + } + match call { + "regex" | "length" | "email" | "formatString" => { + validate_dynamic_field(&args, "value", Value::is_string, call)?; + } + "numeric" | "formatNumber" | "formatCurrency" | "pluralize" => { + validate_dynamic_field(&args, "value", Value::is_number, call)?; + } + "not" => validate_dynamic_field(&args, "value", Value::is_boolean, call)?, + "formatDate" => validate_dynamic_field(&args, "format", Value::is_string, call)?, + "openUrl" => { + let url = required_string(&args, "url")?; + url::Url::parse(url).map_err(|error| format!("invalid openUrl URL: {error}"))?; + } + "and" | "or" => { + for item in args["values"].as_array().unwrap() { + if !item.is_boolean() { + let item = item.as_object().ok_or_else(|| { + format!("function `{call}` values must be dynamic booleans") + })?; + if !item.contains_key("path") && !item.contains_key("call") { + return Err(format!("function `{call}` values must be dynamic booleans")); + } + } + } + } + _ => {} + } + for field in ["min", "max", "decimals", "offset"] { + if args.contains_key(field) { + validate_dynamic_field(&args, field, Value::is_number, call)?; + } + } + if args.contains_key("grouping") { + validate_dynamic_field(&args, "grouping", Value::is_boolean, call)?; + } + for field in ["currency", "zero", "one", "two", "few", "many", "other"] { + if args.contains_key(field) { + validate_dynamic_field(&args, field, Value::is_string, call)?; + } + } + Ok(()) +} diff --git a/src/agent.rs b/src/agent.rs index 4db1ce7..7923542 100644 --- a/src/agent.rs +++ b/src/agent.rs @@ -1173,7 +1173,7 @@ pub(crate) fn parse_tool_calls( let (content, calls) = if model == ModelChoice::Glm52 { parse_glm_calls(text)? } else { - crate::server::parse_dsml_tool_calls(text)? + crate::dsml::parse_tool_calls(text)? }; calls .into_iter() diff --git a/src/app.rs b/src/app.rs index 032f421..1f46f84 100644 --- a/src/app.rs +++ b/src/app.rs @@ -19,12 +19,16 @@ use crate::config::{ GitDiffWhitespace, }; use crate::database::{Database, ProjectWithSessions, SessionState, StoredMessage}; -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] use crate::engine::ChatTurn; +#[cfg(target_os = "macos")] +use crate::metrics::WorkSource; use crate::metrics::{KvCacheReport, Metrics, MetricsSnapshot}; use crate::model::{self, DownloadOutcome, DownloadProgress, ManagedArtifactId, ModelChoice}; #[cfg(target_os = "macos")] -use crate::runtime::{ActiveGeneration, CheckpointTarget, GenerationEvent, GenerationService}; +use crate::runtime::{ + ActiveGeneration, CheckpointTarget, CompactionInput, GenerationEvent, GenerationService, +}; use crate::settings::{ DiagnosticPreferences, ExecutionPreferences, GIB, GenerationPreferences, KvCachePreferences, ReasoningMode, RuntimePreferences, SpeculativePreferences, SsdPreferences, SteeringPreferences, @@ -36,9 +40,11 @@ use rfd::AsyncFileDialog; use std::collections::{HashMap, HashSet, VecDeque}; use std::fs; use std::path::{Path, PathBuf}; +use std::sync::Arc; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::mpsc::{self, TryRecvError}; -use std::sync::{Arc, Mutex, RwLock}; +#[cfg(target_os = "macos")] +use std::sync::{Mutex, RwLock}; use std::thread; use std::time::{Duration, Instant}; @@ -595,6 +601,7 @@ impl App { active_generation: None, #[cfg(target_os = "macos")] active_compaction: None, + #[cfg(target_os = "macos")] active_tool_check: None, #[cfg(target_os = "macos")] agent_tools: None, @@ -737,6 +744,7 @@ impl App { active_generation: None, #[cfg(target_os = "macos")] active_compaction: None, + #[cfg(target_os = "macos")] active_tool_check: None, #[cfg(target_os = "macos")] agent_tools: None, @@ -938,6 +946,10 @@ impl App { } pub(crate) fn update(&mut self, message: Message) -> Task { + let message = match self.update_preference_message(message) { + Ok(message) => message, + Err(task) => return task, + }; match message { Message::Noop => { #[cfg(target_os = "macos")] @@ -1132,312 +1144,6 @@ impl App { self.sync_native_menu(); } Message::MetricsTick => self.sample_metrics(), - Message::PreferenceModelChanged(model) => { - self.preference_draft.model = model; - if !model.supports_dspark() { - self.preference_draft.legacy_mtp_enabled = false; - self.preference_draft.dspark_enabled = false; - self.preference_draft.dspark_confidence_threshold.clear(); - self.preference_draft.dspark_strict = false; - } - if model == ModelChoice::Glm52 { - self.preference_draft.power_percent.clear(); - self.preference_draft.prefill_chunk.clear(); - self.preference_draft.directional_steering_file.clear(); - self.preference_draft.directional_steering_ffn.clear(); - self.preference_draft.directional_steering_attn.clear(); - } else { - self.preference_draft.glm_mtp = false; - self.preference_draft.glm_mtp_timing = false; - self.preference_draft.ssd_full_layers.clear(); - } - self.preference_error = None; - } - Message::PreferenceLegacyMtpChanged(enabled) => { - self.preference_draft.legacy_mtp_enabled = - self.preference_draft.model.supports_dspark() && enabled; - if self.preference_draft.legacy_mtp_enabled { - self.preference_draft.dspark_enabled = false; - self.preference_draft.dspark_confidence_threshold.clear(); - self.preference_draft.dspark_strict = false; - } - self.preference_error = None; - } - Message::PreferenceDsparkChanged(enabled) => { - self.preference_draft.dspark_enabled = - self.preference_draft.model.supports_dspark() && enabled; - if !self.preference_draft.dspark_enabled { - self.preference_draft.dspark_confidence_threshold.clear(); - self.preference_draft.dspark_strict = false; - } else { - self.preference_draft.legacy_mtp_enabled = false; - } - self.preference_error = None; - } - Message::PreferenceTimeoutChanged(value) => { - self.preference_draft.idle_timeout_minutes = value; - self.preference_error = None; - } - Message::PreferenceA2uiChanged(enabled) => { - self.preference_draft.a2ui_enabled = enabled; - self.preference_error = None; - } - Message::PreferenceEndpointPortChanged(value) => { - 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::PreferenceDevBrainEnabledChanged(enabled) => { - self.preference_draft.dev_brain_enabled = enabled; - self.preference_error = None; - } - Message::PreferenceDevBrainVaultChanged(value) => { - self.preference_draft.dev_brain_vault_path = value; - self.preference_error = None; - } - Message::PreferenceGitDiffLayoutChanged(layout) => { - self.preference_draft.git_diff_layout = layout; - self.preference_error = None; - } - Message::PreferenceGitDiffAlgorithmChanged(algorithm) => { - self.preference_draft.git_diff_algorithm = algorithm; - self.preference_error = None; - } - Message::PreferenceGitContextLinesChanged(value) => { - self.preference_draft.git_context_lines = value; - self.preference_error = None; - } - Message::PreferenceGitInterhunkLinesChanged(value) => { - self.preference_draft.git_interhunk_lines = value; - self.preference_error = None; - } - Message::PreferenceGitIndentHeuristicChanged(enabled) => { - self.preference_draft.git_indent_heuristic = enabled; - self.preference_error = None; - } - Message::PreferenceGitWhitespaceChanged(whitespace) => { - self.preference_draft.git_whitespace = whitespace; - self.preference_error = None; - } - Message::PreferenceGitIgnoreBlankLinesChanged(enabled) => { - self.preference_draft.git_ignore_blank_lines = enabled; - self.preference_error = None; - } - Message::ChooseDevBrainVault => { - return Task::perform( - async { - AsyncFileDialog::new() - .set_title("Choose an Obsidian vault") - .pick_folder() - .await - .map(|folder| folder.path().to_path_buf()) - }, - Message::DevBrainVaultPicked, - ); - } - Message::DevBrainVaultPicked(path) => { - if let Some(path) = path { - self.preference_draft.dev_brain_vault_path = - path.to_string_lossy().into_owned(); - self.preference_error = None; - } - } - Message::RestoreDevBrainDefaultGuides => { - self.preference_error = None; - self.restore_dev_brain_confirmation = true; - } - Message::ConfirmRestoreDevBrainDefaultGuides => { - self.restore_dev_brain_confirmation = false; - self.preference_error = crate::dev_brain::restore_default_guides(Path::new( - &self.preference_draft.dev_brain_vault_path, - )) - .err(); - } - Message::CancelRestoreDevBrainDefaultGuides => { - self.restore_dev_brain_confirmation = false; - } - Message::PreferenceContextChanged(value) => { - self.preference_draft.context_tokens = value; - self.preference_error = None; - } - Message::PreferenceMaxTokensChanged(value) => { - self.preference_draft.max_generated_tokens = value; - self.preference_error = None; - } - Message::PreferenceSystemPromptAction(action) => { - self.preference_draft.system_prompt.perform(action); - self.preference_error = None; - } - Message::PreferenceTemperatureChanged(value) => { - self.preference_draft.temperature = value; - self.preference_error = None; - } - Message::PreferenceTopPChanged(value) => { - self.preference_draft.top_p = value; - self.preference_error = None; - } - Message::PreferenceMinPChanged(value) => { - self.preference_draft.min_p = value; - self.preference_error = None; - } - Message::PreferenceSeedChanged(value) => { - self.preference_draft.seed = value; - self.preference_error = None; - } - Message::PreferenceReasoningChanged(value) => { - self.preference_draft.reasoning_mode = value; - self.preference_error = None; - } - Message::PreferenceCpuThreadsChanged(value) => { - self.preference_draft.cpu_threads = value; - self.preference_error = None; - } - Message::PreferencePowerChanged(value) => { - self.preference_draft.power_percent = value; - self.preference_error = None; - } - Message::PreferencePrefillChunkChanged(value) => { - self.preference_draft.prefill_chunk = value; - self.preference_error = None; - } - Message::PreferenceQualityChanged(value) => { - self.preference_draft.quality = value; - self.preference_error = None; - } - Message::PreferenceWarmWeightsChanged(value) => { - self.preference_draft.warm_weights = value; - self.preference_error = None; - } - Message::PreferenceMtpDraftChanged(value) => { - self.preference_draft.mtp_draft_tokens = value; - self.preference_error = None; - } - Message::PreferenceMtpMarginChanged(value) => { - self.preference_draft.mtp_margin = value; - self.preference_error = None; - } - Message::PreferenceGlmMtpChanged(value) => { - self.preference_draft.glm_mtp = - self.preference_draft.model == ModelChoice::Glm52 && value; - if !self.preference_draft.glm_mtp { - self.preference_draft.glm_mtp_timing = false; - } - self.preference_error = None; - } - Message::PreferenceGlmMtpTimingChanged(value) => { - self.preference_draft.glm_mtp_timing = - self.preference_draft.model == ModelChoice::Glm52 && value; - if self.preference_draft.glm_mtp_timing { - self.preference_draft.glm_mtp = true; - } - self.preference_error = None; - } - Message::PreferenceDsparkConfidenceChanged(value) => { - self.preference_draft.dspark_confidence_threshold = value; - if self.preference_draft.model.supports_dspark() - && !self - .preference_draft - .dspark_confidence_threshold - .trim() - .is_empty() - { - self.preference_draft.dspark_enabled = true; - self.preference_draft.legacy_mtp_enabled = false; - } - self.preference_error = None; - } - Message::PreferenceDsparkStrictChanged(value) => { - self.preference_draft.dspark_strict = - self.preference_draft.model.supports_dspark() && value; - if self.preference_draft.dspark_strict { - self.preference_draft.dspark_enabled = true; - self.preference_draft.legacy_mtp_enabled = false; - } - self.preference_error = None; - } - Message::PreferenceSsdChanged(value) => { - self.preference_draft.ssd_streaming = value; - self.preference_error = None; - } - Message::PreferenceSsdColdChanged(value) => { - self.preference_draft.ssd_streaming_cold = value; - self.preference_error = None; - } - Message::PreferenceSsdCacheChanged(value) => { - self.preference_draft.ssd_cache = value; - self.preference_error = None; - } - Message::PreferenceSsdFullLayersChanged(value) => { - if self.preference_draft.model == ModelChoice::Glm52 { - self.preference_draft.ssd_full_layers = value; - } - self.preference_error = None; - } - Message::PreferenceSsdPreloadChanged(value) => { - self.preference_draft.ssd_preload_experts = value; - self.preference_error = None; - } - Message::PreferenceSteeringFileChanged(value) => { - if self.preference_draft.model != ModelChoice::Glm52 { - self.preference_draft.directional_steering_file = value; - } - self.preference_error = None; - } - Message::PreferenceSteeringFfnChanged(value) => { - if self.preference_draft.model != ModelChoice::Glm52 { - self.preference_draft.directional_steering_ffn = value; - } - self.preference_error = None; - } - Message::PreferenceSteeringAttnChanged(value) => { - if self.preference_draft.model != ModelChoice::Glm52 { - self.preference_draft.directional_steering_attn = value; - } - self.preference_error = None; - } - Message::PreferenceSimulatedMemoryChanged(value) => { - self.preference_draft.simulated_used_memory_gib = value; - self.preference_error = None; - } - Message::PreferenceExpertProfileChanged(value) => { - self.preference_draft.expert_profile_path = value; - self.preference_error = None; - } - Message::PreferenceKvBudgetChanged(value) => { - self.preference_draft.kv_budget_gib = value; - self.preference_error = None; - } - Message::PreferenceKvMinTokensChanged(value) => { - self.preference_draft.kv_min_tokens = value; - self.preference_error = None; - } - Message::PreferenceKvColdMaxChanged(value) => { - self.preference_draft.kv_cold_max_tokens = value; - self.preference_error = None; - } - Message::PreferenceKvContinuedIntervalChanged(value) => { - self.preference_draft.kv_continued_interval_tokens = value; - self.preference_error = None; - } - Message::ResetPreferences => { - self.preference_draft.reset(); - self.preference_error = None; - } - Message::SavePreferences => { - self.save_preferences(); - if self.preference_error.is_none() - && let Some(id) = self.preferences_window - { - return window::close(id); - } - } Message::DownloadArtifact(artifact) => { self.start_model_operation(artifact, ModelOperation::Download) } @@ -1567,8 +1273,11 @@ impl App { self.error = Some(format!("Could not play media: {error}")); } #[cfg(not(target_os = "macos"))] - if let Err(error) = std::process::Command::new("open").arg(url).spawn() { - self.error = Some(format!("Could not open media: {error}")); + { + let _ = (title, video); + if let Err(error) = std::process::Command::new("open").arg(url).spawn() { + self.error = Some(format!("Could not open media: {error}")); + } } } Message::RequestA2uiDismiss(surface_id) => { @@ -1788,8 +1497,10 @@ impl App { if let Some(database) = &mut self.database { match database.delete_project(project_id) { Ok(()) => { + #[cfg(not(target_os = "macos"))] + let _ = checkpoint_ids; + #[cfg(target_os = "macos")] for session_id in checkpoint_ids { - #[cfg(target_os = "macos")] self.background_chats.remove(&session_id); } self.drafts.remove(&project_id); @@ -2122,6 +1833,7 @@ impl App { } } } + _ => unreachable!("preference messages are dispatched before the main update"), } Task::none() } diff --git a/src/app/generation.rs b/src/app/generation.rs index d7fc437..f64258a 100644 --- a/src/app/generation.rs +++ b/src/app/generation.rs @@ -293,7 +293,7 @@ fn sync_a2ui_message( renderable_surface_updated } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn chat_turn(message: &ChatMessage) -> ChatTurn { ChatTurn { user: message.user, @@ -368,7 +368,7 @@ fn project_agents(path: &Path) -> Result, String> { } } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn title_context(messages: impl IntoIterator) -> Vec { let mut started = false; messages @@ -497,6 +497,8 @@ impl App { #[cfg(target_os = "macos")] let agents = agents_prompt_for_turn(opening_turn, opening_agents, self.session_agents_prompt()); + #[cfg(not(target_os = "macos"))] + let agents = None::; self.a2ui_auto_switch_pending = true; self.context_notice = None; #[cfg(target_os = "macos")] @@ -535,6 +537,7 @@ impl App { &effective.turn.system_prompt, self.compaction_summary(), ); + #[cfg(target_os = "macos")] let assistant_reasoning = effective.turn.reasoning_mode != ReasoningMode::Direct; #[cfg(target_os = "macos")] let model_prompt = if self.config.a2ui_enabled { @@ -667,6 +670,7 @@ impl App { effective.turn, messages, checkpoint, + WorkSource::LocalChat, idle_timeout, ) { Ok(active) => Some(active), @@ -700,7 +704,6 @@ impl App { { let _ = effective; self.error = Some("Local Metal generation requires macOS.".into()); - return; } } @@ -1304,6 +1307,7 @@ impl App { checkpoint: session_checkpoint_path(session_id), bootstrap: None, }, + WorkSource::LocalChat, idle_timeout, )?, ); @@ -1520,18 +1524,18 @@ impl App { .generation_service .as_ref() .ok_or_else(|| "The model runtime is unavailable.".to_owned())? - .compact( - effective.engine, - effective.turn, + .compact(CompactionInput { + engine: effective.engine, + turn: effective.turn, messages, - reason, + reason: reason.to_owned(), rebuild_system_prompt, - session_compaction_checkpoint_path( + checkpoint: session_compaction_checkpoint_path( self.selected_session .ok_or_else(|| "The active session is unavailable.".to_owned())?, ), idle_timeout, - )?; + })?; self.active_compaction = Some(CompactionRequest { active, pending, @@ -1789,7 +1793,8 @@ impl App { effective.engine, effective.turn, messages, - CheckpointTarget::OneShot(transient_cache_path()), + CheckpointTarget::Transient(transient_cache_path()), + WorkSource::LocalChat, idle_timeout, ) { Ok(active) => { diff --git a/src/app/preferences.rs b/src/app/preferences.rs index 5421806..8162561 100644 --- a/src/app/preferences.rs +++ b/src/app/preferences.rs @@ -1,4 +1,5 @@ use super::*; +use std::sync::RwLock; #[derive(Clone)] pub(super) struct PreferenceDraft { @@ -437,6 +438,7 @@ impl App { return; } } + #[cfg(target_os = "macos")] let dev_brain_changed = self.config.dev_brain != config.dev_brain; #[cfg(target_os = "macos")] let endpoint_changed = self.config.endpoint != config.endpoint; @@ -447,6 +449,7 @@ impl App { { self._endpoint = None; } + #[cfg(target_os = "macos")] let pending_endpoint = if config.endpoint.enabled && (endpoint_changed || self._endpoint.is_none()) { let Some(generation) = &self.generation_service else { @@ -494,6 +497,324 @@ impl App { } } +impl App { + pub(super) fn update_preference_message( + &mut self, + message: Message, + ) -> Result> { + match message { + Message::PreferenceModelChanged(model) => { + self.preference_draft.model = model; + if !model.supports_dspark() { + self.preference_draft.legacy_mtp_enabled = false; + self.preference_draft.dspark_enabled = false; + self.preference_draft.dspark_confidence_threshold.clear(); + self.preference_draft.dspark_strict = false; + } + if model == ModelChoice::Glm52 { + self.preference_draft.power_percent.clear(); + self.preference_draft.prefill_chunk.clear(); + self.preference_draft.directional_steering_file.clear(); + self.preference_draft.directional_steering_ffn.clear(); + self.preference_draft.directional_steering_attn.clear(); + } else { + self.preference_draft.glm_mtp = false; + self.preference_draft.glm_mtp_timing = false; + self.preference_draft.ssd_full_layers.clear(); + } + self.preference_error = None; + } + Message::PreferenceLegacyMtpChanged(enabled) => { + self.preference_draft.legacy_mtp_enabled = + self.preference_draft.model.supports_dspark() && enabled; + if self.preference_draft.legacy_mtp_enabled { + self.preference_draft.dspark_enabled = false; + self.preference_draft.dspark_confidence_threshold.clear(); + self.preference_draft.dspark_strict = false; + } + self.preference_error = None; + } + Message::PreferenceDsparkChanged(enabled) => { + self.preference_draft.dspark_enabled = + self.preference_draft.model.supports_dspark() && enabled; + if !self.preference_draft.dspark_enabled { + self.preference_draft.dspark_confidence_threshold.clear(); + self.preference_draft.dspark_strict = false; + } else { + self.preference_draft.legacy_mtp_enabled = false; + } + self.preference_error = None; + } + Message::PreferenceTimeoutChanged(value) => { + self.preference_draft.idle_timeout_minutes = value; + self.preference_error = None; + } + Message::PreferenceA2uiChanged(enabled) => { + self.preference_draft.a2ui_enabled = enabled; + self.preference_error = None; + } + Message::PreferenceEndpointPortChanged(value) => { + 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::PreferenceDevBrainEnabledChanged(enabled) => { + self.preference_draft.dev_brain_enabled = enabled; + self.preference_error = None; + } + Message::PreferenceDevBrainVaultChanged(value) => { + self.preference_draft.dev_brain_vault_path = value; + self.preference_error = None; + } + Message::PreferenceGitDiffLayoutChanged(layout) => { + self.preference_draft.git_diff_layout = layout; + self.preference_error = None; + } + Message::PreferenceGitDiffAlgorithmChanged(algorithm) => { + self.preference_draft.git_diff_algorithm = algorithm; + self.preference_error = None; + } + Message::PreferenceGitContextLinesChanged(value) => { + self.preference_draft.git_context_lines = value; + self.preference_error = None; + } + Message::PreferenceGitInterhunkLinesChanged(value) => { + self.preference_draft.git_interhunk_lines = value; + self.preference_error = None; + } + Message::PreferenceGitIndentHeuristicChanged(enabled) => { + self.preference_draft.git_indent_heuristic = enabled; + self.preference_error = None; + } + Message::PreferenceGitWhitespaceChanged(whitespace) => { + self.preference_draft.git_whitespace = whitespace; + self.preference_error = None; + } + Message::PreferenceGitIgnoreBlankLinesChanged(enabled) => { + self.preference_draft.git_ignore_blank_lines = enabled; + self.preference_error = None; + } + Message::ChooseDevBrainVault => { + return Err(Task::perform( + async { + AsyncFileDialog::new() + .set_title("Choose an Obsidian vault") + .pick_folder() + .await + .map(|folder| folder.path().to_path_buf()) + }, + Message::DevBrainVaultPicked, + )); + } + Message::DevBrainVaultPicked(path) => { + if let Some(path) = path { + self.preference_draft.dev_brain_vault_path = + path.to_string_lossy().into_owned(); + self.preference_error = None; + } + } + Message::RestoreDevBrainDefaultGuides => { + self.preference_error = None; + self.restore_dev_brain_confirmation = true; + } + Message::ConfirmRestoreDevBrainDefaultGuides => { + self.restore_dev_brain_confirmation = false; + self.preference_error = crate::dev_brain::restore_default_guides(Path::new( + &self.preference_draft.dev_brain_vault_path, + )) + .err(); + } + Message::CancelRestoreDevBrainDefaultGuides => { + self.restore_dev_brain_confirmation = false; + } + Message::PreferenceContextChanged(value) => { + self.preference_draft.context_tokens = value; + self.preference_error = None; + } + Message::PreferenceMaxTokensChanged(value) => { + self.preference_draft.max_generated_tokens = value; + self.preference_error = None; + } + Message::PreferenceSystemPromptAction(action) => { + self.preference_draft.system_prompt.perform(action); + self.preference_error = None; + } + Message::PreferenceTemperatureChanged(value) => { + self.preference_draft.temperature = value; + self.preference_error = None; + } + Message::PreferenceTopPChanged(value) => { + self.preference_draft.top_p = value; + self.preference_error = None; + } + Message::PreferenceMinPChanged(value) => { + self.preference_draft.min_p = value; + self.preference_error = None; + } + Message::PreferenceSeedChanged(value) => { + self.preference_draft.seed = value; + self.preference_error = None; + } + Message::PreferenceReasoningChanged(value) => { + self.preference_draft.reasoning_mode = value; + self.preference_error = None; + } + Message::PreferenceCpuThreadsChanged(value) => { + self.preference_draft.cpu_threads = value; + self.preference_error = None; + } + Message::PreferencePowerChanged(value) => { + self.preference_draft.power_percent = value; + self.preference_error = None; + } + Message::PreferencePrefillChunkChanged(value) => { + self.preference_draft.prefill_chunk = value; + self.preference_error = None; + } + Message::PreferenceQualityChanged(value) => { + self.preference_draft.quality = value; + self.preference_error = None; + } + Message::PreferenceWarmWeightsChanged(value) => { + self.preference_draft.warm_weights = value; + self.preference_error = None; + } + Message::PreferenceMtpDraftChanged(value) => { + self.preference_draft.mtp_draft_tokens = value; + self.preference_error = None; + } + Message::PreferenceMtpMarginChanged(value) => { + self.preference_draft.mtp_margin = value; + self.preference_error = None; + } + Message::PreferenceGlmMtpChanged(value) => { + self.preference_draft.glm_mtp = + self.preference_draft.model == ModelChoice::Glm52 && value; + if !self.preference_draft.glm_mtp { + self.preference_draft.glm_mtp_timing = false; + } + self.preference_error = None; + } + Message::PreferenceGlmMtpTimingChanged(value) => { + self.preference_draft.glm_mtp_timing = + self.preference_draft.model == ModelChoice::Glm52 && value; + if self.preference_draft.glm_mtp_timing { + self.preference_draft.glm_mtp = true; + } + self.preference_error = None; + } + Message::PreferenceDsparkConfidenceChanged(value) => { + self.preference_draft.dspark_confidence_threshold = value; + if self.preference_draft.model.supports_dspark() + && !self + .preference_draft + .dspark_confidence_threshold + .trim() + .is_empty() + { + self.preference_draft.dspark_enabled = true; + self.preference_draft.legacy_mtp_enabled = false; + } + self.preference_error = None; + } + Message::PreferenceDsparkStrictChanged(value) => { + self.preference_draft.dspark_strict = + self.preference_draft.model.supports_dspark() && value; + if self.preference_draft.dspark_strict { + self.preference_draft.dspark_enabled = true; + self.preference_draft.legacy_mtp_enabled = false; + } + self.preference_error = None; + } + Message::PreferenceSsdChanged(value) => { + self.preference_draft.ssd_streaming = value; + self.preference_error = None; + } + Message::PreferenceSsdColdChanged(value) => { + self.preference_draft.ssd_streaming_cold = value; + self.preference_error = None; + } + Message::PreferenceSsdCacheChanged(value) => { + self.preference_draft.ssd_cache = value; + self.preference_error = None; + } + Message::PreferenceSsdFullLayersChanged(value) => { + if self.preference_draft.model == ModelChoice::Glm52 { + self.preference_draft.ssd_full_layers = value; + } + self.preference_error = None; + } + Message::PreferenceSsdPreloadChanged(value) => { + self.preference_draft.ssd_preload_experts = value; + self.preference_error = None; + } + Message::PreferenceSteeringFileChanged(value) => { + if self.preference_draft.model != ModelChoice::Glm52 { + self.preference_draft.directional_steering_file = value; + } + self.preference_error = None; + } + Message::PreferenceSteeringFfnChanged(value) => { + if self.preference_draft.model != ModelChoice::Glm52 { + self.preference_draft.directional_steering_ffn = value; + } + self.preference_error = None; + } + Message::PreferenceSteeringAttnChanged(value) => { + if self.preference_draft.model != ModelChoice::Glm52 { + self.preference_draft.directional_steering_attn = value; + } + self.preference_error = None; + } + Message::PreferenceSimulatedMemoryChanged(value) => { + self.preference_draft.simulated_used_memory_gib = value; + self.preference_error = None; + } + Message::PreferenceExpertProfileChanged(value) => { + self.preference_draft.expert_profile_path = value; + self.preference_error = None; + } + Message::PreferenceKvBudgetChanged(value) => { + self.preference_draft.kv_budget_gib = value; + self.preference_error = None; + } + Message::PreferenceKvMinTokensChanged(value) => { + self.preference_draft.kv_min_tokens = value; + self.preference_error = None; + } + Message::PreferenceKvColdMaxChanged(value) => { + self.preference_draft.kv_cold_max_tokens = value; + self.preference_error = None; + } + Message::PreferenceKvContinuedIntervalChanged(value) => { + self.preference_draft.kv_continued_interval_tokens = value; + self.preference_error = None; + } + Message::ResetPreferences => { + self.preference_draft.reset(); + self.preference_error = None; + } + Message::SavePreferences => { + self.save_preferences(); + if self.preference_error.is_none() + && let Some(id) = self.preferences_window + { + return Err(window::close(id)); + } + } + message => return Ok(message), + } + Err(Task::none()) + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/src/database.rs b/src/database.rs index d8503fa..fd34270 100644 --- a/src/database.rs +++ b/src/database.rs @@ -147,6 +147,54 @@ struct NewMessage<'a> { compaction_tail_start: Option, } +impl<'a> NewMessage<'a> { + fn system(session_id: i32, content: &'a str) -> Self { + Self::new(session_id, content, false, false, true) + } + + fn user(session_id: i32, content: &'a str, model_content: Option<&'a str>) -> Self { + Self { + model_content, + ..Self::new(session_id, content, true, false, false) + } + } + + fn tool_result(session_id: i32, content: &'a str) -> Self { + Self::new(session_id, content, false, true, false) + } + + fn assistant(session_id: i32, reasoning: bool) -> Self { + Self { + reasoning: reasoning.then_some(""), + reasoning_complete: !reasoning, + ..Self::new(session_id, "", false, false, false) + } + } + + fn compaction(session_id: i32, content: &'a str, tail_start: Option) -> Self { + Self { + compaction: true, + compaction_tail_start: tail_start, + ..Self::system(session_id, content) + } + } + + fn new(session_id: i32, content: &'a str, user: bool, tool: bool, system: bool) -> Self { + Self { + session_id, + user, + tool, + reasoning: None, + reasoning_complete: true, + content, + model_content: None, + system, + compaction: false, + compaction_tail_start: None, + } + } +} + #[derive(Clone, Debug, Identifiable, Queryable, Selectable)] #[diesel(table_name = a2ui_messages)] #[diesel(check_for_backend(diesel::sqlite::Sqlite))] @@ -388,18 +436,7 @@ impl Database { self.connection .transaction(|connection| { let message = diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: false, - reasoning: None, - reasoning_complete: true, - content: &content, - model_content: None, - system: true, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::system(session_id, &content)) .returning(StoredMessage::as_returning()) .get_result(connection)?; diesel::insert_into(a2ui_messages::table) @@ -454,50 +491,17 @@ impl Database { for content in system_messages { stored.push( diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: false, - reasoning: None, - reasoning_complete: true, - content, - model_content: None, - system: true, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::system(session_id, content)) .returning(StoredMessage::as_returning()) .get_result(connection)?, ); } let user = diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: true, - tool: false, - reasoning: None, - reasoning_complete: true, - content: prompt, - model_content: model_prompt, - system: false, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::user(session_id, prompt, model_prompt)) .returning(StoredMessage::as_returning()) .get_result(connection)?; let assistant = diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: false, - reasoning: reasoning.then_some(""), - reasoning_complete: !reasoning, - content: "", - model_content: None, - system: false, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::assistant(session_id, reasoning)) .returning(StoredMessage::as_returning()) .get_result(connection)?; stored.extend([user, assistant]); @@ -520,36 +524,14 @@ impl Database { let mut stored = Vec::with_capacity(system_messages.len() + 3); stored.push( diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: true, - reasoning: None, - reasoning_complete: true, - content: result, - model_content: None, - system: false, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::tool_result(session_id, result)) .returning(StoredMessage::as_returning()) .get_result(connection)?, ); if let Some(content) = queued_user { stored.push( diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: true, - tool: false, - reasoning: None, - reasoning_complete: true, - content, - model_content: None, - system: false, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::user(session_id, content, None)) .returning(StoredMessage::as_returning()) .get_result(connection)?, ); @@ -557,35 +539,13 @@ impl Database { for content in system_messages { stored.push( diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: false, - reasoning: None, - reasoning_complete: true, - content, - model_content: None, - system: true, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::system(session_id, content)) .returning(StoredMessage::as_returning()) .get_result(connection)?, ); } let assistant = diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: false, - reasoning: reasoning.then_some(""), - reasoning_complete: !reasoning, - content: "", - model_content: None, - system: false, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::assistant(session_id, reasoning)) .returning(StoredMessage::as_returning()) .get_result(connection)?; stored.push(assistant); @@ -667,36 +627,14 @@ impl Database { let mut stored = Vec::with_capacity(2); stored.push( diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: false, - reasoning: None, - reasoning_complete: true, - content: summary, - model_content: None, - system: true, - compaction: true, - compaction_tail_start: tail_start, - }) + .values(NewMessage::compaction(session_id, summary, tail_start)) .returning(StoredMessage::as_returning()) .get_result(connection)?, ); if let Some(content) = running_jobs { stored.push( diesel::insert_into(messages::table) - .values(NewMessage { - session_id, - user: false, - tool: true, - reasoning: None, - reasoning_complete: true, - content, - model_content: None, - system: false, - compaction: false, - compaction_tail_start: None, - }) + .values(NewMessage::tool_result(session_id, content)) .returning(StoredMessage::as_returning()) .get_result(connection)?, ); diff --git a/src/dsml.rs b/src/dsml.rs new file mode 100644 index 0000000..9973c10 --- /dev/null +++ b/src/dsml.rs @@ -0,0 +1,256 @@ +use serde_json::{Map, Value}; + +const INCOMPLETE_TOOL_CALL: &str = "invalid or incomplete DSML tool call"; + +#[derive(Clone, Copy)] +pub(crate) struct Syntax { + pub(crate) tool_start: &'static str, + pub(crate) tool_end: &'static str, + pub(crate) invoke_start: &'static str, + pub(crate) invoke_end: &'static str, + pub(crate) parameter_start: &'static str, + pub(crate) parameter_end: &'static str, +} + +pub(crate) const SYNTAXES: [Syntax; 3] = [ + Syntax { + tool_start: "<|DSML|tool_calls>", + tool_end: "", + invoke_start: "<|DSML|invoke", + invoke_end: "", + parameter_start: "<|DSML|parameter", + parameter_end: "", + }, + Syntax { + tool_start: "", + tool_end: "", + invoke_start: "", + parameter_start: "", + }, + Syntax { + tool_start: "", + tool_end: "", + invoke_start: "", + parameter_start: "", + }, +]; + +pub(crate) fn parse_tool_calls(text: &str) -> Result<(String, Vec<(String, Value)>), String> { + let Some((start, syntax)) = SYNTAXES + .iter() + .filter_map(|syntax| text.find(syntax.tool_start).map(|start| (start, *syntax))) + .min_by_key(|(start, _)| *start) + else { + return Ok((text.to_owned(), Vec::new())); + }; + let relative_end = text[start..].find(syntax.tool_end); + let end = relative_end + .map(|relative_end| start + relative_end + syntax.tool_end.len()) + .unwrap_or(text.len()); + let outer_call_complete = relative_end.is_some(); + let content = text[..start].trim_end().to_owned(); + let raw = &text[start..end]; + let mut cursor = syntax.tool_start.len(); + let mut calls = Vec::new(); + loop { + skip_whitespace(raw, &mut cursor); + if raw[cursor..].starts_with(syntax.tool_end) { + break; + } + // DS4 recovers calls whose invokes are complete even if generation + // stopped before emitting the outer tool-calls closing tag. + if !outer_call_complete && cursor == raw.len() && !calls.is_empty() { + break; + } + let call = (|| { + if !raw[cursor..].starts_with(syntax.invoke_start) { + return Err(INCOMPLETE_TOOL_CALL.into()); + } + let tag_end = raw[cursor..] + .find('>') + .map(|offset| cursor + offset + 1) + .ok_or_else(|| INCOMPLETE_TOOL_CALL.to_owned())?; + let name = attribute(&raw[cursor..tag_end], "name") + .ok_or_else(|| INCOMPLETE_TOOL_CALL.to_owned())?; + cursor = tag_end; + let mut arguments = Map::new(); + loop { + skip_whitespace(raw, &mut cursor); + if raw[cursor..].starts_with(syntax.invoke_end) { + cursor += syntax.invoke_end.len(); + break; + } + let (name, value) = parse_parameter(raw, &mut cursor, syntax)? + .ok_or_else(|| INCOMPLETE_TOOL_CALL.to_owned())?; + arguments.insert(name, value); + } + Ok((name, Value::Object(arguments))) + })(); + match call { + Ok(call) => calls.push(call), + Err(error) if !calls.is_empty() && error == INCOMPLETE_TOOL_CALL => { + break; + } + Err(error) => return Err(error), + } + } + if calls.is_empty() { + return Err(INCOMPLETE_TOOL_CALL.into()); + } + Ok((content, calls)) +} + +fn parse_parameter( + text: &str, + cursor: &mut usize, + syntax: Syntax, +) -> Result, String> { + if !text[*cursor..].starts_with(syntax.parameter_start) { + return Ok(None); + } + let Some(tag_end) = text[*cursor..].find('>').map(|end| *cursor + end + 1) else { + return Ok(None); + }; + let tag = &text[*cursor..tag_end]; + let Some(name) = attribute(tag, "name") else { + return Ok(None); + }; + let is_string = attribute(tag, "string"); + *cursor = tag_end; + let mut nested_start = *cursor; + skip_whitespace(text, &mut nested_start); + if is_string.is_none() && text[nested_start..].starts_with(syntax.parameter_start) { + *cursor = nested_start; + let mut nested = Map::new(); + loop { + skip_whitespace(text, cursor); + if !text[*cursor..].starts_with(syntax.parameter_start) { + break; + } + let Some((name, value)) = parse_parameter(text, cursor, syntax)? else { + return Ok(None); + }; + nested.insert(name, value); + } + skip_whitespace(text, cursor); + if !text[*cursor..].starts_with(syntax.parameter_end) { + return Ok(None); + } + *cursor += syntax.parameter_end.len(); + return Ok(Some((name, Value::Object(nested)))); + } + let Some(value_end) = text[*cursor..] + .find(syntax.parameter_end) + .map(|end| *cursor + end) + else { + return Ok(None); + }; + let raw = &text[*cursor..value_end]; + *cursor = value_end + syntax.parameter_end.len(); + let value = if is_string.as_deref().unwrap_or("true") == "true" { + Value::String(unescape(raw)) + } else { + serde_json::from_str(raw) + .map_err(|error| format!("invalid DSML tool arguments: {error}"))? + }; + Ok(Some((name, value))) +} + +fn attribute(tag: &str, name: &str) -> Option { + let start = tag.find(&format!("{name}=\""))? + name.len() + 2; + let end = start + tag[start..].find('"')?; + Some(unescape(&tag[start..end])) +} + +fn skip_whitespace(text: &str, cursor: &mut usize) { + while text[*cursor..] + .chars() + .next() + .is_some_and(char::is_whitespace) + { + *cursor += text[*cursor..].chars().next().unwrap().len_utf8(); + } +} + +fn unescape(text: &str) -> String { + text.replace(""", "\"") + .replace(">", ">") + .replace("<", "<") + .replace("&", "&") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_reference_and_legacy_dsml_syntaxes() { + for (start, end, invoke, invoke_end, parameter, parameter_end) in [ + ( + "<|DSML|tool_calls>", + "", + "<|DSML|invoke", + "", + "<|DSML|parameter", + "", + ), + ( + "", + "", + "", + "", + ), + ( + "", + "", + "", + "", + ), + ] { + let raw = format!( + "done{start}{invoke} name=\"read\">{parameter} name=\"path\" string=\"true\">src/main.rs{parameter_end}{invoke_end}{end}" + ); + let (content, calls) = parse_tool_calls(&raw).unwrap(); + assert_eq!(content, "done"); + assert_eq!(calls[0].0, "read"); + assert_eq!(calls[0].1["path"], "src/main.rs"); + } + } + + #[test] + fn recovers_complete_invokes_without_the_outer_closing_tag() { + let raw = "before<|DSML|tool_calls><|DSML|invoke name=\"read\"><|DSML|parameter name=\"path\" string=\"true\">src/main.rs"; + + let (content, calls) = parse_tool_calls(raw).unwrap(); + + assert_eq!(content, "before"); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].0, "read"); + assert_eq!(calls[0].1["path"], "src/main.rs"); + } + + #[test] + fn recovers_complete_invokes_before_an_incomplete_one() { + let raw = "truncated"; + + let (_, calls) = parse_tool_calls(raw).unwrap(); + + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].0, "first"); + + let invalid = "invalid"; + assert!( + parse_tool_calls(invalid) + .unwrap_err() + .starts_with("invalid DSML tool arguments") + ); + } +} diff --git a/src/engine.rs b/src/engine.rs index 4cff18e..00f3bf1 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -1,5 +1,5 @@ mod gguf; -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] mod kvstore; #[cfg(target_os = "macos")] mod metal; @@ -9,6 +9,7 @@ mod validation; #[cfg(target_os = "macos")] use crate::metrics::{KvLookup, Metrics}; use crate::model::ModelChoice; +#[cfg(target_os = "macos")] use crate::settings::TurnSettings; use crate::settings::{EngineSettings, ReasoningMode}; use gguf::{F16, F32, Gguf, I32, IQ2_XXS, Q2_K, Q4_0, Q4_K, Q5_K, Q6_K, Q8_0, Tensor, Value}; @@ -17,9 +18,12 @@ use kvstore::{KvStore, StoreReason}; use sha2::{Digest, Sha256}; #[cfg(target_os = "macos")] use std::collections::HashMap; -use std::path::{Path, PathBuf}; +use std::path::Path; +#[cfg(target_os = "macos")] +use std::path::PathBuf; #[cfg(target_os = "macos")] use std::sync::Arc; +#[cfg(target_os = "macos")] use std::sync::atomic::{AtomicBool, Ordering}; #[cfg(target_os = "macos")] use std::time::Instant; @@ -1386,7 +1390,7 @@ impl Generator { } } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn append_generated_bytes( generated: &mut ChatTurn, reasoning: bool, @@ -1449,7 +1453,7 @@ fn flush_generated( ); } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn emit_safe_text( text: &mut String, emitted: &mut usize, @@ -1494,7 +1498,7 @@ fn emit_safe_text( false } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn conversation_key(system: &str, reasoning: ReasoningMode, messages: &[ChatTurn]) -> Vec { fn text(output: &mut Vec, value: &str) { output.extend_from_slice(&(value.len() as u64).to_le_bytes()); @@ -1526,12 +1530,12 @@ fn conversation_key(system: &str, reasoning: ReasoningMode, messages: &[ChatTurn output } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn conversation_tag(system: &str, reasoning: ReasoningMode, messages: &[ChatTurn]) -> [u8; 32] { Sha256::digest(conversation_key(system, reasoning, messages)).into() } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn checkpoint_matches_prefix( checkpoint: [u8; 32], system: &str, @@ -1555,7 +1559,7 @@ fn resident_key(directory: &Path, tag: [u8; 32]) -> PathBuf { directory.join("resident").join(name) } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] fn sample( logits: &[f32], temperature: f32, @@ -1631,10 +1635,10 @@ fn sample( probabilities.last().map_or(0, |(token, _)| *token as i32) } -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] struct Rng(u64); -#[cfg(target_os = "macos")] +#[cfg(any(target_os = "macos", test))] impl Rng { fn new(seed: u64) -> Self { Self(seed.max(1)) @@ -1654,7 +1658,7 @@ impl Rng { } } -#[cfg(all(test, target_os = "macos"))] +#[cfg(test)] mod sampling_tests { use super::*; @@ -1835,6 +1839,7 @@ mod sampling_tests { } #[test] + #[cfg(target_os = "macos")] #[ignore = "requires the 80 GiB Flash checkpoint and Apple Metal"] fn metal_executes_real_flash_token() { configure_metal_sources().unwrap(); diff --git a/src/main.rs b/src/main.rs index 09209ee..e495110 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,3 +1,5 @@ +#![cfg_attr(not(target_os = "macos"), allow(dead_code))] + mod a2ui; mod a2ui_validation; mod agent; @@ -6,6 +8,7 @@ mod compaction; mod config; mod database; mod dev_brain; +mod dsml; mod engine; mod metrics; mod model; diff --git a/src/runtime.rs b/src/runtime.rs index 26648b4..04caab0 100644 --- a/src/runtime.rs +++ b/src/runtime.rs @@ -34,18 +34,6 @@ pub(crate) enum CheckpointTarget { bootstrap: Option, }, Transient(PathBuf), - /// Same transient KV handling as [`CheckpointTarget::Transient`], but asked - /// for by the app itself (session titling) rather than by an HTTP client. - OneShot(PathBuf), -} - -impl CheckpointTarget { - fn source(&self) -> WorkSource { - match self { - Self::Local { .. } | Self::OneShot(_) => WorkSource::LocalChat, - Self::Transient(_) => WorkSource::Http, - } - } } pub(crate) enum GenerationEvent { @@ -74,24 +62,48 @@ enum Operation { Measure, } -#[derive(Clone, Copy)] -enum ResponseKind { - Generation, - Compaction, - Measurement, +impl Operation { + fn tracks_metrics(&self) -> bool { + !matches!(self, Self::Measure) + } + + fn error_handler(&self) -> fn(String) -> GenerationEvent { + match self { + Self::Generate => |error| GenerationEvent::Finished(Err(error)), + Self::Compact { .. } => |error| GenerationEvent::Compacted(Err(error)), + Self::Measure => |error| GenerationEvent::Measured(Err(error)), + } + } } struct Command { engine: EngineSettings, turn: TurnSettings, messages: Vec, - checkpoint: CheckpointTarget, + checkpoint: Option, + source: WorkSource, operation: Operation, idle_timeout: Duration, cancel: Arc, events: Sender, } +pub(crate) struct CompactionInput { + pub(crate) engine: EngineSettings, + pub(crate) turn: TurnSettings, + pub(crate) messages: Vec, + pub(crate) reason: String, + pub(crate) rebuild_system_prompt: String, + pub(crate) checkpoint: PathBuf, + pub(crate) idle_timeout: Duration, +} + +struct RuntimeState { + loaded: Option<(EngineSettings, Generator)>, + last_used: Instant, + idle_timeout: Duration, +} + impl GenerationService { pub(crate) fn spawn(metrics: Arc) -> Result { let (commands, receiver) = mpsc::channel::(); @@ -109,69 +121,35 @@ impl GenerationService { turn: TurnSettings, messages: Vec, checkpoint: CheckpointTarget, + source: WorkSource, idle_timeout: Duration, ) -> Result { - let source = checkpoint.source(); - let cancel = Arc::new(AtomicBool::new(false)); - let (events, receiver) = mpsc::channel(); - self.metrics.request_queued(source); - if self - .commands - .send(Command { - engine, - turn, - messages, - checkpoint, - operation: Operation::Generate, - idle_timeout, - cancel: Arc::clone(&cancel), - events, - }) - .is_err() - { - self.metrics.request_rejected(); - return Err("The model runtime stopped unexpectedly.".to_owned()); - } - Ok(ActiveGeneration { - events: receiver, - cancel, + self.submit(CommandRequest { + engine, + turn, + messages, + checkpoint: Some(checkpoint), + source, + operation: Operation::Generate, + idle_timeout, }) } - #[allow(clippy::too_many_arguments)] - pub(crate) fn compact( - &self, - engine: EngineSettings, - turn: TurnSettings, - messages: Vec, - reason: &str, - rebuild_system_prompt: String, - checkpoint: PathBuf, - idle_timeout: Duration, - ) -> Result { - let cancel = Arc::new(AtomicBool::new(false)); - let (events, receiver) = mpsc::channel(); - self.commands - .send(Command { - engine, - turn, - messages, - checkpoint: CheckpointTarget::Local { - checkpoint, - bootstrap: None, - }, - operation: Operation::Compact { - reason: reason.to_owned(), - rebuild_system_prompt, - }, - idle_timeout, - cancel: Arc::clone(&cancel), - events, - }) - .map_err(|_| "The model runtime stopped unexpectedly.".to_owned())?; - Ok(ActiveGeneration { - events: receiver, - cancel, + pub(crate) fn compact(&self, input: CompactionInput) -> Result { + self.submit(CommandRequest { + engine: input.engine, + turn: input.turn, + messages: input.messages, + checkpoint: Some(CheckpointTarget::Local { + checkpoint: input.checkpoint, + bootstrap: None, + }), + source: WorkSource::LocalChat, + operation: Operation::Compact { + reason: input.reason, + rebuild_system_prompt: input.rebuild_system_prompt, + }, + idle_timeout: input.idle_timeout, }) } @@ -182,21 +160,41 @@ impl GenerationService { messages: Vec, idle_timeout: Duration, ) -> Result { + self.submit(CommandRequest { + engine, + turn, + messages, + checkpoint: None, + source: WorkSource::LocalChat, + operation: Operation::Measure, + idle_timeout, + }) + } + + fn submit(&self, request: CommandRequest) -> Result { let cancel = Arc::new(AtomicBool::new(false)); let (events, receiver) = mpsc::channel(); - self.metrics.request_queued(WorkSource::LocalChat); - self.commands - .send(Command { - engine, - turn, - messages, - checkpoint: CheckpointTarget::OneShot(PathBuf::new()), - operation: Operation::Measure, - idle_timeout, - cancel: Arc::clone(&cancel), - events, - }) - .map_err(|_| "The model runtime stopped unexpectedly.".to_owned())?; + let tracked = request.operation.tracks_metrics(); + if tracked { + self.metrics.request_queued(request.source); + } + let command = Command { + engine: request.engine, + turn: request.turn, + messages: request.messages, + checkpoint: request.checkpoint, + source: request.source, + operation: request.operation, + idle_timeout: request.idle_timeout, + cancel: Arc::clone(&cancel), + events, + }; + if self.commands.send(command).is_err() { + if tracked { + self.metrics.request_rejected(); + } + return Err("The model runtime stopped unexpectedly.".to_owned()); + } Ok(ActiveGeneration { events: receiver, cancel, @@ -204,44 +202,48 @@ impl GenerationService { } } +struct CommandRequest { + engine: EngineSettings, + turn: TurnSettings, + messages: Vec, + checkpoint: Option, + source: WorkSource, + operation: Operation, + idle_timeout: Duration, +} + fn run(commands: Receiver, metrics: Arc) { - let mut loaded = None::<(EngineSettings, Generator)>; - let mut last_used = Instant::now(); - let mut idle_timeout = Duration::from_secs(15 * 60); + let mut state = RuntimeState { + loaded: None, + last_used: Instant::now(), + idle_timeout: Duration::from_secs(15 * 60), + }; loop { match commands.recv_timeout(Duration::from_secs(1)) { Ok(command) => { let request_started = Instant::now(); - let source = command.checkpoint.source(); + let source = command.source; let events = command.events.clone(); - let response = response_kind(&command.operation); - let tracked = !matches!(command.operation, Operation::Measure); + let error_event = command.operation.error_handler(); + let tracked = command.operation.tracks_metrics(); if tracked { metrics.request_started(source); } if let Err(error) = catch_runtime_panic(|| { - run_command( - command, - &mut loaded, - &metrics, - request_started, - source, - &mut last_used, - &mut idle_timeout, - ); + run_command(command, &mut state, &metrics, request_started, source); }) { if tracked { metrics.request_failed(request_started.elapsed()); } - if loaded.take().is_some() { + if state.loaded.take().is_some() { metrics.unloaded(); } - let _ = events.send(error_event(response, error)); + let _ = events.send(error_event(error)); } } Err(mpsc::RecvTimeoutError::Timeout) => { - if loaded.is_some() && last_used.elapsed() >= idle_timeout { - loaded = None; + if state.loaded.is_some() && state.last_used.elapsed() >= state.idle_timeout { + state.loaded = None; metrics.unloaded(); } } @@ -250,39 +252,37 @@ fn run(commands: Receiver, metrics: Arc) { } } -#[allow(clippy::too_many_arguments)] fn run_command( command: Command, - loaded: &mut Option<(EngineSettings, Generator)>, + state: &mut RuntimeState, metrics: &Arc, request_started: Instant, source: WorkSource, - last_used: &mut Instant, - idle_timeout: &mut Duration, ) { - let response = response_kind(&command.operation); - let tracked = !matches!(command.operation, Operation::Measure); - *idle_timeout = command.idle_timeout; + let error_event = command.operation.error_handler(); + let tracked = command.operation.tracks_metrics(); + state.idle_timeout = command.idle_timeout; if command.cancel.load(Ordering::Relaxed) { if tracked { metrics.request_failed(request_started.elapsed()); } let _ = command .events - .send(error_event(response, "generation cancelled".into())); + .send(error_event("generation cancelled".into())); return; } - if loaded + if state + .loaded .as_ref() .is_none_or(|(settings, _)| settings != &command.engine) { - if loaded.take().is_some() { + if state.loaded.take().is_some() { metrics.unloaded(); } let _ = command.events.send(GenerationEvent::Loading); metrics.loading(); let load_started = Instant::now(); - *loaded = match Generator::open(&command.engine, Arc::clone(metrics)) { + state.loaded = match Generator::open(&command.engine, Arc::clone(metrics)) { Ok(generator) => { let summary = generator.summary(); metrics.loaded( @@ -298,12 +298,12 @@ fn run_command( if tracked { metrics.request_failed(request_started.elapsed()); } - let _ = command.events.send(error_event(response, error)); + let _ = command.events.send(error_event(error)); None } }; } - if let Some((_, generator)) = loaded { + if let Some((_, generator)) = &mut state.loaded { let mut prefill_started = None::<(Instant, u32)>; let mut emit = |reasoning, content| { let _ = command @@ -335,7 +335,7 @@ fn run_command( rebuild_system_prompt, } = &command.operation { - let CheckpointTarget::Local { checkpoint, .. } = &command.checkpoint else { + let Some(CheckpointTarget::Local { checkpoint, .. }) = &command.checkpoint else { unreachable!("compaction checkpoints are local") }; let result = generator.compact( @@ -357,16 +357,19 @@ fn run_command( Err(_) => metrics.request_failed(request_started.elapsed()), } let _ = command.events.send(GenerationEvent::Compacted(result)); - *last_used = Instant::now(); + state.last_used = Instant::now(); return; } if matches!(command.operation, Operation::Measure) { let result = generator.rendered_history_tokens(&command.messages, &command.turn); let _ = command.events.send(GenerationEvent::Measured(result)); - *last_used = Instant::now(); + state.last_used = Instant::now(); return; } - let result = match command.checkpoint { + let result = match command + .checkpoint + .expect("generation requires a checkpoint target") + { CheckpointTarget::Local { checkpoint, bootstrap, @@ -382,16 +385,14 @@ fn run_command( let _ = command.events.send(GenerationEvent::Activity(activity)); }, ), - CheckpointTarget::Transient(directory) | CheckpointTarget::OneShot(directory) => { - generator.generate_transient( - &directory, - &command.messages, - &command.turn, - &command.cancel, - &mut emit, - &mut progress, - ) - } + CheckpointTarget::Transient(directory) => generator.generate_transient( + &directory, + &command.messages, + &command.turn, + &command.cancel, + &mut emit, + &mut progress, + ), }; match &result { Ok(output) => metrics.request_finished( @@ -406,23 +407,7 @@ fn run_command( Err(_) => metrics.request_failed(request_started.elapsed()), } let _ = command.events.send(GenerationEvent::Finished(result)); - *last_used = Instant::now(); - } -} - -fn response_kind(operation: &Operation) -> ResponseKind { - match operation { - Operation::Generate => ResponseKind::Generation, - Operation::Compact { .. } => ResponseKind::Compaction, - Operation::Measure => ResponseKind::Measurement, - } -} - -fn error_event(kind: ResponseKind, error: String) -> GenerationEvent { - match kind { - ResponseKind::Generation => GenerationEvent::Finished(Err(error)), - ResponseKind::Compaction => GenerationEvent::Compacted(Err(error)), - ResponseKind::Measurement => GenerationEvent::Measured(Err(error)), + state.last_used = Instant::now(); } } @@ -461,12 +446,29 @@ mod tests { #[test] fn operation_failures_use_the_matching_event() { assert!(matches!( - error_event(ResponseKind::Compaction, "stopped".into()), + Operation::Compact { + reason: String::new(), + rebuild_system_prompt: String::new(), + } + .error_handler()("stopped".into()), GenerationEvent::Compacted(Err(error)) if error == "stopped" )); assert!(matches!( - error_event(ResponseKind::Measurement, "stopped".into()), + Operation::Measure.error_handler()("stopped".into()), GenerationEvent::Measured(Err(error)) if error == "stopped" )); } + + #[test] + fn operation_tracking_is_consistent() { + assert!(Operation::Generate.tracks_metrics()); + assert!( + Operation::Compact { + reason: String::new(), + rebuild_system_prompt: String::new(), + } + .tracks_metrics() + ); + assert!(!Operation::Measure.tracks_metrics()); + } } diff --git a/src/server.rs b/src/server.rs index 9c39996..84c6f13 100644 --- a/src/server.rs +++ b/src/server.rs @@ -22,7 +22,7 @@ use tools::{canonical_tools, validate_tool_results}; use crate::config::Config; use crate::engine::ChatTurn; -use crate::metrics::Metrics; +use crate::metrics::{Metrics, WorkSource}; use crate::model::{self, ModelChoice}; use crate::runtime::{CheckpointTarget, GenerationEvent, GenerationService}; use crate::settings::{ReasoningMode, effective_settings}; @@ -312,49 +312,6 @@ impl Drop for ConnectionSlot { } } -pub(crate) fn parse_dsml_tool_calls(text: &str) -> Result<(String, Vec<(String, Value)>), String> { - let has_tool_marker = TOOL_SYNTAXES - .iter() - .any(|syntax| text.contains(syntax.tool_start)); - let mut projector = ToolProjector::new(); - let events = projector.push(text, true, ""); - let mut content = String::new(); - let mut calls = Vec::<(String, String, bool)>::new(); - for event in events { - match event { - ToolProjectionEvent::Text(text) => content.push_str(&text), - ToolProjectionEvent::Start { index, name, .. } => { - if calls.len() == index { - calls.push((name, String::new(), false)); - } - } - ToolProjectionEvent::Arguments { index, fragment } => { - if let Some((_, arguments, _)) = calls.get_mut(index) { - arguments.push_str(&fragment); - } - } - ToolProjectionEvent::End { index } => { - if let Some((_, _, complete)) = calls.get_mut(index) { - *complete = true; - } - } - } - } - let calls = calls - .into_iter() - .filter(|(_, _, complete)| *complete) - .map(|(name, arguments, _)| { - serde_json::from_str(&arguments) - .map(|arguments| (name, arguments)) - .map_err(|error| format!("invalid DSML tool arguments: {error}")) - }) - .collect::, _>>()?; - if has_tool_marker && calls.is_empty() { - return Err("invalid or incomplete DSML tool call".into()); - } - Ok((content, calls)) -} - fn handle(mut stream: TcpStream, state: &State) { let _ = stream.set_write_timeout(Some(HTTP_IO_TIMEOUT)); let started = Instant::now(); @@ -472,6 +429,7 @@ fn chat_completion( parsed.turn, parsed.messages, CheckpointTarget::Transient(state.cache_path.clone()), + WorkSource::Http, parsed.idle_timeout, ) .map_err(|error| (500, error))?; @@ -524,6 +482,7 @@ fn compatible_completion( parsed.turn, parsed.messages, CheckpointTarget::Transient(state.cache_path.clone()), + WorkSource::Http, parsed.idle_timeout, ) .map_err(|error| (500, error))?; diff --git a/src/server/response.rs b/src/server/response.rs index fd786fe..44d110c 100644 --- a/src/server/response.rs +++ b/src/server/response.rs @@ -682,7 +682,6 @@ fn responses_stream_response( Err((500, "The model runtime stopped unexpectedly.".into())) } -#[allow(clippy::too_many_arguments)] fn send_named_sse( stream: &mut impl Write, event: &str, diff --git a/src/server/tools.rs b/src/server/tools.rs index 199bfa8..a5f4143 100644 --- a/src/server/tools.rs +++ b/src/server/tools.rs @@ -1,4 +1,5 @@ use super::*; +pub(super) use crate::dsml::{SYNTAXES as TOOL_SYNTAXES, Syntax as ToolSyntax}; const TOOL_MEMORY_FILE: &str = "tool-replay.json"; const TOOL_MEMORY_MAX_IDS: usize = 100_000; @@ -698,43 +699,6 @@ fn repair_generated_tools(text: &str) -> Option { Some(repaired) } -#[derive(Clone, Copy)] -pub(super) struct ToolSyntax { - pub(super) tool_start: &'static str, - pub(super) tool_end: &'static str, - pub(super) invoke_start: &'static str, - pub(super) invoke_end: &'static str, - pub(super) parameter_start: &'static str, - pub(super) parameter_end: &'static str, -} - -pub(super) const TOOL_SYNTAXES: [ToolSyntax; 3] = [ - ToolSyntax { - tool_start: "<|DSML|tool_calls>", - tool_end: "", - invoke_start: "<|DSML|invoke", - invoke_end: "", - parameter_start: "<|DSML|parameter", - parameter_end: "", - }, - ToolSyntax { - tool_start: "", - tool_end: "", - invoke_start: "", - parameter_start: "", - }, - ToolSyntax { - tool_start: "", - tool_end: "", - invoke_start: "", - parameter_start: "", - }, -]; - fn parse_tool_parameter( text: &str, cursor: &mut usize,