diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4041621..5dfec37 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -26,6 +26,18 @@ jobs: run: cargo clippy --workspace --locked --all-targets -- -D warnings - name: Test run: cargo test --workspace --locked + - name: Official Laya packing parity + env: + LAYA_TOKENIZER: ${{ runner.temp }}/laya-tokenizer.json + LAYA_PACKING_ORACLE: ${{ runner.temp }}/laya-packing.json + run: | + curl --fail --location --retry 3 \ + https://huggingface.co/convaiinnovations/laya/resolve/55cf4c4ebb4ebe31b2550e8bdf3bd21b99753851/tokenizer/tokenizer.json \ + --output "$LAYA_TOKENIZER" + curl --fail --location --retry 3 \ + https://raw.githubusercontent.com/linear3735/system1-omni/5e4dd4215c925ebd93bb9ce4097b27bd6375f7c0/recipe/laya/native/packing-golden.json \ + --output "$LAYA_PACKING_ORACLE" + cargo test --locked -p omni-laya --test packing -- --ignored - name: Build run: cargo build --workspace --release --locked diff --git a/Cargo.lock b/Cargo.lock index 4b9adda..06f296f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -113,6 +113,21 @@ version = "0.23.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ac07cdecf99051d9a5238b80f35af32cdeba5b336e55d957b318b50137e18da5" +[[package]] +name = "bit-set" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" +dependencies = [ + "bit-vec", +] + +[[package]] +name = "bit-vec" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" + [[package]] name = "bitflags" version = "2.13.2" @@ -256,6 +271,12 @@ dependencies = [ "typenum", ] +[[package]] +name = "daachorse" +version = "3.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5614204febbc33cc07a2806aa6440b904ac012b68eecc37f4493ea4a76455a3d" + [[package]] name = "darling" version = "0.20.11" @@ -380,6 +401,17 @@ version = "0.1.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" +[[package]] +name = "fancy-regex" +version = "0.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72cf461f865c862bb7dc573f643dd6a2b6842f7c30b07882b56bd148cc2761b8" +dependencies = [ + "bit-set", + "regex-automata", + "regex-syntax", +] + [[package]] name = "fastrand" version = "2.5.0" @@ -954,7 +986,7 @@ dependencies = [ "safetensors 0.8.0", "serde", "serde_json", - "tokenizers", + "tokenizers 0.22.2", "tokio", ] @@ -980,6 +1012,7 @@ dependencies = [ "serde_json", "sha2", "tempfile", + "tokenizers 0.23.2", ] [[package]] @@ -1679,6 +1712,39 @@ dependencies = [ "unicode_categories", ] +[[package]] +name = "tokenizers" +version = "0.23.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7afbf6e88718afcc138bad01d6ccc3051dbbc3b2ce9793d8b8a3aeb610969cfc" +dependencies = [ + "ahash", + "compact_str", + "daachorse", + "dary_heap", + "derive_builder", + "esaxx-rs", + "fancy-regex", + "getrandom 0.3.4", + "itertools", + "log", + "macro_rules_attribute", + "monostate", + "paste", + "rand 0.9.5", + "rayon", + "rayon-cond", + "regex", + "regex-syntax", + "serde", + "serde_json", + "spm_precompiled", + "thiserror", + "unicode-normalization-alignments", + "unicode-segmentation", + "unicode_categories", +] + [[package]] name = "tokio" version = "1.53.1" diff --git a/recipe/laya/README.md b/recipe/laya/README.md index 254c12c..609cd7c 100644 --- a/recipe/laya/README.md +++ b/recipe/laya/README.md @@ -56,3 +56,47 @@ if the worker requires a bearer token. See the [frontend documentation](../../src/frontend/README.md) for configuration and transport behavior. + +## Native CPU packing check + +The `omni-laya` preprocessor packs English Laya 0.3.20 requests without weights +or a GPU. + +```sh +cargo test --locked -p omni-laya --test preprocess +``` + +Native callers use `Request::from_json(&str)` for a single top-level JSON request, +or `Request::from_value(Value)` for an existing structured value. `Request` +retains its public fields and `Serialize`; it does not implement generic +`Deserialize`. The JSON entry checks the complete request against serde_json's +default nesting limit. The value entry preserves existing nested values without +reparsing. Both preserve literal private Number/RawValue object keys and reject +unknown request fields. + +The normal tests cover validation, JSON rendering, question and option order, +and truncation with a small tokenizer: +Pass raw JSON directly to `from_json`. + +For the official 17-case comparison, use the same pinned inputs as CPU CI. +The test checks both files by SHA-256 before comparing: + +```sh +LAYA_PACKING_DIR=$(mktemp -d) +export LAYA_TOKENIZER="$LAYA_PACKING_DIR/tokenizer.json" +export LAYA_PACKING_ORACLE="$LAYA_PACKING_DIR/packing-golden.json" +curl --fail --location --retry 3 \ + https://huggingface.co/convaiinnovations/laya/resolve/55cf4c4ebb4ebe31b2550e8bdf3bd21b99753851/tokenizer/tokenizer.json \ + --output "$LAYA_TOKENIZER" +curl --fail --location --retry 3 \ + https://raw.githubusercontent.com/linear3735/system1-omni/5e4dd4215c925ebd93bb9ce4097b27bd6375f7c0/recipe/laya/native/packing-golden.json \ + --output "$LAYA_PACKING_ORACLE" +cargo test --locked -p omni-laya --test packing -- --ignored +``` + +Existing copies of these pinned files can be supplied through `LAYA_TOKENIZER` +and `LAYA_PACKING_ORACLE` instead. The comparison covers every token, marker, +question type, row length, question order and usage count; it excludes backend +padding and bucket dimensions. The [reference generator and inputs](https://github.com/linear3735/system1-omni/tree/5e4dd4215c925ebd93bb9ce4097b27bd6375f7c0/recipe/laya/native) +use `laya==0.3.20`. Packing parity does not measure model quality or execute +native model inference. diff --git a/src/models/cua_s1/native/Cargo.toml b/src/models/cua_s1/native/Cargo.toml index 2e8caf6..998f7ef 100644 --- a/src/models/cua_s1/native/Cargo.toml +++ b/src/models/cua_s1/native/Cargo.toml @@ -18,7 +18,7 @@ libloading = "0.8" memmap2 = "0.9.9" safetensors = "0.8.0" serde = "1" -serde_json = { version = "1.0.149", features = ["float_roundtrip", "preserve_order"] } +serde_json = { version = "1.0.149", features = ["float_roundtrip", "preserve_order", "raw_value"] } # the onig regex backend, as in the Python tokenizers wheel tokenizers = { version = "=0.22.2", default-features = false, features = ["onig"] } tokio = { version = "1.49.0", features = ["macros", "net", "rt-multi-thread", "sync"] } diff --git a/src/models/cua_s1/native/src/json.rs b/src/models/cua_s1/native/src/json.rs index a0e2ed5..59aaeb0 100644 --- a/src/models/cua_s1/native/src/json.rs +++ b/src/models/cua_s1/native/src/json.rs @@ -10,14 +10,14 @@ use std::fmt::Write as _; use std::io; use serde::Serialize; -use serde::de::{self, Deserializer, MapAccess, SeqAccess, Visitor}; +use serde::de::{self, DeserializeSeed, Deserializer, MapAccess, SeqAccess, Visitor}; use serde_json::{Map, Number, Value}; /// Decode a request body into its top-level object; the error is the 400 message. pub fn parse(raw: &[u8]) -> Result, String> { let mut de = serde_json::Deserializer::from_slice(raw); - let value = de - .deserialize_any(NoDuplicates) + let value = NoDuplicates(0) + .deserialize(&mut de) .and_then(|v| de.end().map(|()| v)) .map_err(|e| format!("request body is not valid JSON: {e}"))?; match value { @@ -27,16 +27,37 @@ pub fn parse(raw: &[u8]) -> Result, String> { } /// Builds a `Value` like serde_json does, but fails on a repeated key. -struct NoDuplicates; +struct NoDuplicates(usize); -impl<'de> de::Deserialize<'de> for Wrapped { - fn deserialize>(d: D) -> Result { - d.deserialize_any(NoDuplicates).map(Wrapped) +impl<'de> DeserializeSeed<'de> for NoDuplicates { + type Value = Value; + + fn deserialize>(self, d: D) -> Result { + let raw = <&serde_json::value::RawValue as de::Deserialize>::deserialize(d)?; + let text = raw.get(); + // Another workspace member may enable arbitrary_precision. Read number + // tokens directly so its private map encoding cannot become user data. + if matches!(text.as_bytes()[0], b'-' | b'0'..=b'9') { + if text != "-0" { + if let Ok(n) = text.parse::() { + return self.visit_i64(n); + } + if let Ok(n) = text.parse::() { + return self.visit_u64(n); + } + } + let n = serde_json::from_str::(text).map_err(de::Error::custom)?; + return self.visit_f64(n); + } + if self.0 >= 127 && matches!(text.as_bytes()[0], b'{' | b'[') { + return Err(de::Error::custom("recursion limit exceeded")); + } + serde_json::Deserializer::from_str(text) + .deserialize_any(self) + .map_err(de::Error::custom) } } -struct Wrapped(Value); - impl<'de> Visitor<'de> for NoDuplicates { type Value = Value; @@ -68,7 +89,7 @@ impl<'de> Visitor<'de> for NoDuplicates { } fn visit_seq>(self, mut seq: A) -> Result { let mut items = Vec::new(); - while let Some(Wrapped(v)) = seq.next_element()? { + while let Some(v) = seq.next_element_seed(NoDuplicates(self.0 + 1))? { items.push(v); } Ok(Value::Array(items)) @@ -76,7 +97,7 @@ impl<'de> Visitor<'de> for NoDuplicates { fn visit_map>(self, mut map: A) -> Result { let mut obj = Map::new(); while let Some(key) = map.next_key::()? { - let Wrapped(v) = map.next_value()?; + let v = map.next_value_seed(NoDuplicates(self.0 + 1))?; if obj.contains_key(&key) { return Err(de::Error::custom(format_args!( "duplicate key {}", @@ -107,6 +128,14 @@ pub fn dumps(value: &Value) -> String { struct PyFormatter; impl serde_json::ser::Formatter for PyFormatter { + fn write_number_str(&mut self, w: &mut W, n: &str) -> io::Result<()> { + if n.contains(['.', 'e', 'E']) { + let x = serde_json::from_str::(n).map_err(io::Error::other)?; + self.write_f64(w, x) + } else { + w.write_all(n.as_bytes()) + } + } fn begin_array_value( &mut self, w: &mut W, @@ -243,3 +272,7 @@ mod tests { assert!(parse(b"{\"a\": \"\xff\"}").is_err()); } } + +#[cfg(test)] +#[path = "../../../../../tests/cua_s1/json.rs"] +mod json_regression_tests; diff --git a/src/models/laya/Cargo.toml b/src/models/laya/Cargo.toml index 0eba920..3ba4850 100644 --- a/src/models/laya/Cargo.toml +++ b/src/models/laya/Cargo.toml @@ -10,7 +10,8 @@ half = "2" memmap2 = "0.9" safetensors = "0.6" serde = { version = "1", features = ["derive"] } -serde_json = "1" +serde_json = { version = "1", features = ["preserve_order", "arbitrary_precision", "raw_value"] } +tokenizers = { version = "0.23.2", default-features = false, features = ["fancy-regex"] } [dev-dependencies] sha2 = "0.10" @@ -23,3 +24,11 @@ path = "../../../tests/laya/checkpoint.rs" [[test]] name = "weights" path = "../../../tests/laya/weights.rs" + +[[test]] +name = "preprocess" +path = "../../../tests/laya/preprocess.rs" + +[[test]] +name = "packing" +path = "../../../tests/laya/packing.rs" diff --git a/src/models/laya/README.md b/src/models/laya/README.md index 3cf9ec0..f8b90ed 100644 --- a/src/models/laya/README.md +++ b/src/models/laya/README.md @@ -8,6 +8,18 @@ The `omni-laya` crate currently reads and checks the English Laya 0.3.20 checkpo Keep checkpoint files unchanged while `Weights` holds a read-only memory mapping. This crate does not yet execute inference. +`Preprocessor::load` reads a tokenizer JSON file. `prepare` packs English `choice`, `score` and `noul` questions into ordered token rows, option-marker positions and type IDs. Rows follow Laya 0.3.20's 512-token limit and 192-token head budget. Conversation lists keep the newest state tokens; other state values keep the beginning. The result includes normalized criteria for later decoding and the total input-token usage. Backends own padding, batching and resource limits. + +Build a request with `Request::from_json(&str)` for one top-level JSON object, or +`Request::from_value(Value)` for an already constructed value. The JSON entry +preserves object order and arbitrary-size integers, treats serde_json's private +Number/RawValue keys as ordinary user keys, and applies its default nesting limit +to the complete request. The value entry moves state and questions without +reparsing or adding a depth limit. Both reject unknown request fields and require +state and an object of questions. Public fields and `Serialize` remain available; +`Request` does not implement generic `Deserialize`, so use these explicit entries +instead of `serde_json::from_str::` or `serde_json::from_value::`. + ## CPU checks The normal workspace tests cover configuration errors, malformed tensors, inventory mismatches and conversion boundaries without downloading weights. @@ -23,6 +35,8 @@ cargo test --release --locked -p omni-laya --test weights -- --ignored These two CPU tests check all 206 tensor names and shapes, 618 conversion hashes, and the legacy temperature buffer. The normal CI job skips them because it does not download the full checkpoint. +The normal tests also check input validation, question and option order, truncation and JSON rendering with a small test tokenizer. The [Laya recipe](../../../recipe/laya/README.md#native-cpu-packing-check) provides the CPU packing validation commands and pinned inputs for the official 17-case comparison. No weights or GPU are needed; packing parity does not measure model quality. + ## Python worker The Python worker serves LAYA through laya-serve on CPU and Apple Silicon (PyTorch MPS, diff --git a/src/models/laya/src/lib.rs b/src/models/laya/src/lib.rs index fbc136f..ddda161 100644 --- a/src/models/laya/src/lib.rs +++ b/src/models/laya/src/lib.rs @@ -1,2 +1,3 @@ pub mod config; +pub mod preprocess; pub mod weights; diff --git a/src/models/laya/src/preprocess.rs b/src/models/laya/src/preprocess.rs new file mode 100644 index 0000000..cbdf7e2 --- /dev/null +++ b/src/models/laya/src/preprocess.rs @@ -0,0 +1,415 @@ +//! English Laya 0.3.20 token packing, before backend padding or batching. +use anyhow::{Result, anyhow, bail, ensure}; +use serde::de::{self, MapAccess, SeqAccess, Visitor}; +use serde::{Deserializer, Serialize}; +use serde_json::value::RawValue; +use serde_json::{Map, Value}; +use std::path::Path; +use tokenizers::Tokenizer; + +/// Use `from_json` for a top-level JSON request or `from_value` for an existing value. +#[derive(Debug, Serialize)] +pub struct Request { + pub state: Value, + pub model: Option, + pub questions: Map, + pub lang: Option, +} + +impl Request { + /// Decode one JSON request with serde_json's default nesting limit. + pub fn from_json(raw: &str) -> Result { + let raw: Box = serde_json::from_str(raw)?; + Self::from_value(parse_value(&raw, 0)?) + } + + /// Keep existing structured values without reparsing or imposing a JSON depth limit. + pub fn from_value(value: Value) -> Result { + let Value::Object(mut fields) = value else { + bail!("request must be an object"); + }; + let state = fields + .remove("state") + .ok_or_else(|| anyhow!("missing state"))?; + let questions = fields + .remove("questions") + .ok_or_else(|| anyhow!("missing questions"))?; + let Value::Object(questions) = questions else { + bail!("questions must be an object"); + }; + let model = serde_json::from_value(fields.remove("model").unwrap_or(Value::Null))?; + let lang = serde_json::from_value(fields.remove("lang").unwrap_or(Value::Null))?; + ensure!(fields.is_empty(), "unknown request fields"); + Ok(Self { + state, + model, + questions, + lang, + }) + } +} + +fn parse_value(raw: &RawValue, depth: usize) -> serde_json::Result { + let text = raw.get(); + if !matches!(text.as_bytes()[0], b'{' | b'[') { + return serde_json::from_str(text); + } + if depth >= 127 { + return Err(de::Error::custom("recursion limit exceeded")); + } + // Construct containers explicitly: serde_json's private Number/RawValue + // map encodings must not interpret legitimate user object keys. + struct Containers(usize); + impl<'de> Visitor<'de> for Containers { + type Value = Value; + fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + f.write_str("a JSON object or array") + } + fn visit_seq>(self, mut seq: A) -> Result { + let mut values = Vec::new(); + while let Some(raw) = seq.next_element::>()? { + values.push(parse_value(&raw, self.0 + 1).map_err(de::Error::custom)?); + } + Ok(Value::Array(values)) + } + fn visit_map>(self, mut map: A) -> Result { + let mut values = Map::new(); + while let Some((key, raw)) = map.next_entry::>()? { + if self.0 == 0 && values.contains_key(&key) { + return Err(de::Error::custom(format!("duplicate field {key}"))); + } + values.insert( + key, + parse_value(&raw, self.0 + 1).map_err(de::Error::custom)?, + ); + } + Ok(Value::Object(values)) + } + } + serde_json::Deserializer::from_str(text).deserialize_any(Containers(depth)) +} + +#[derive(Debug, Serialize)] +pub struct Question { + pub id: String, + pub kind: String, + pub criteria: Value, + pub ids: Vec, + pub markers: Vec, + pub qtype: i64, +} + +#[derive(Debug, Serialize)] +pub struct Prepared { + pub questions: Vec, + pub usage: usize, +} + +pub struct Preprocessor { + tokenizer: Tokenizer, + cls: u32, + sep: u32, + mask: u32, +} + +// Python json.dumps(..., ensure_ascii=False) uses spaces after commas/colons. +// Walk values so punctuation inside strings remains untouched. +pub fn render(value: &Value) -> String { + match value { + Value::String(s) => s.clone(), + _ => spaced_json(value), + } +} +fn spaced_json(value: &Value) -> String { + match value { + Value::Array(a) => format!( + "[{}]", + a.iter().map(spaced_json).collect::>().join(", ") + ), + Value::Object(o) => format!( + "{{{}}}", + o.iter() + .map(|(k, v)| format!("{}: {}", serde_json::to_string(k).unwrap(), spaced_json(v))) + .collect::>() + .join(", ") + ), + Value::Number(n) => python_number(n), + _ => serde_json::to_string(value).unwrap(), + } +} + +impl Preprocessor { + pub fn load(tokenizer_path: &Path) -> Result { + let mut tokenizer = Tokenizer::from_file(tokenizer_path).map_err(|e| anyhow!("{e}"))?; + tokenizer.with_padding(None); + tokenizer + .with_truncation(None) + .map_err(|e| anyhow!("{e}"))?; + let token = |s| { + tokenizer + .token_to_id(s) + .ok_or_else(|| anyhow!("missing token {s}")) + }; + Ok(Self { + cls: token("[CLS]")?, + sep: token("[SEP]")?, + mask: token("[MASK]")?, + tokenizer, + }) + } + fn encode(&self, text: &str) -> Result> { + Ok(self + .tokenizer + .encode(text.replace("[MASK]", " "), false) + .map_err(|e| anyhow!("{e}"))? + .get_ids() + .to_vec()) + } + pub fn prepare(&self, request: &Request) -> Result { + validate_numbers(&request.state)?; + for question in request.questions.values() { + validate_numbers(question)?; + } + ensure!( + request.model.as_deref().is_none_or(|s| s == "english"), + "only model=english is supported" + ); + ensure!( + request + .lang + .as_deref() + .is_none_or(|s| s == "en" || s == "english"), + "only English is supported; use lang=en" + ); + let state = self.encode(&render(&request.state))?; + let mut questions = Vec::new(); + for (id, definition) in &request.questions { + let (kind, criteria, opts) = + options(definition).map_err(|e| anyhow!("question {id:?}: {e}"))?; + let ins = definition + .get("instructions") + .ok_or_else(|| anyhow!("question {id:?}: missing instructions"))?; + let mut head = self.encode(&format!("{kind} question: {}", render(ins)))?; + let mut opt_ids = Vec::new(); + for opt in &opts { + let mut ids = vec![self.mask]; + ids.extend(self.encode(&format!(" {opt}"))?.into_iter().take(48)); + opt_ids.push(ids); + } + let mut budget = 192isize - opt_ids.iter().map(Vec::len).sum::() as isize; + if budget < 16 { + let per = (176 / opt_ids.len()).max(4); + for o in &mut opt_ids { + o.truncate(per); + } + budget = 192 - opt_ids.iter().map(Vec::len).sum::() as isize; + } + head.truncate(budget.max(8) as usize); + let mut ids = vec![self.cls]; + ids.extend(head); + ids.push(self.sep); + let mut markers = Vec::new(); + for opt in opt_ids { + markers.push(ids.len()); + ids.extend(opt); + } + ids.push(self.sep); + let room = 512usize.saturating_sub(ids.len() + 1); + if request.state.is_array() { + ids.extend_from_slice(&state[state.len().saturating_sub(room)..]); + } else { + ids.extend_from_slice(&state[..room.min(state.len())]); + } + ids.push(self.sep); + ids.truncate(512); + ensure!( + markers.iter().all(|&m| m < 512), + "question {id:?}: options exceed head_max_len=192" + ); + let qtype = match kind.as_str() { + "choice" => 0, + "score" => 1, + _ => 2, + }; + questions.push(Question { + id: id.clone(), + kind, + criteria, + ids, + markers, + qtype, + }); + } + let usage = questions.iter().map(|q| q.ids.len()).sum(); + Ok(Prepared { questions, usage }) + } +} + +fn options(q: &Value) -> Result<(String, Value, Vec)> { + ensure!(q.is_object(), "definition must be an object"); + let kind = q["type"].as_str().ok_or_else(|| anyhow!("missing type"))?; + ensure!( + ["choice", "score", "noul"].contains(&kind), + "unknown type {kind}" + ); + ensure!( + kind == "noul" || q.get("labels").is_none(), + "labels only apply to noul" + ); + let mut criteria = q.get("criteria").cloned().unwrap_or(Value::Null); + let opts = match kind { + "choice" => { + if let Some(a) = criteria.as_array() { + let mut o = Map::new(); + for key in a { + o.insert( + key.as_str() + .ok_or_else(|| anyhow!("choice labels must be strings"))? + .to_owned(), + Value::Null, + ); + } + criteria = Value::Object(o); + } + let o = criteria + .as_object() + .ok_or_else(|| anyhow!("choice criteria must be object or list"))?; + ensure!(!o.is_empty(), "at least one choice required"); + o.iter() + .map(|(k, v)| { + if v.is_null() || v.as_str() == Some("") { + k.clone() + } else { + format!("{k}: {}", render(v)) + } + }) + .collect() + } + "score" => { + let a = criteria + .as_array() + .ok_or_else(|| anyhow!("score criteria must be a list"))?; + ensure!(!a.is_empty(), "at least one level required"); + a.iter() + .enumerate() + .map(|(i, v)| format!("level {i}: {}", render(v))) + .collect() + } + _ => { + if criteria.is_null() { + criteria = Value::Object(Map::new()); + } + let o = criteria + .as_object() + .ok_or_else(|| anyhow!("noul criteria must be object"))?; + let mut normalized = Map::new(); + for (k, v) in o { + let key = k.to_lowercase(); + ensure!( + key == "true" || key == "false", + "noul criteria keys must be true/false" + ); + normalized.insert(key, v.clone()); + } + criteria = Value::Object(normalized); + let labels = match q.get("labels") { + None | Some(Value::Null) => ["false", "true"], + Some(Value::Object(o)) if o.len() == 2 => [ + python_strip(o.get("false").and_then(Value::as_str).unwrap_or("")), + python_strip(o.get("true").and_then(Value::as_str).unwrap_or("")), + ], + _ => bail!("noul labels must map true/false to distinct non-empty strings"), + }; + ensure!( + !labels[0].is_empty() && !labels[1].is_empty() && labels[0] != labels[1], + "invalid noul labels" + ); + ["false", "true"] + .iter() + .enumerate() + .map(|(i, k)| { + let v = &criteria[*k]; + let desc = if v.is_null() || v.as_str() == Some("") { + if i == 0 { + "no, the statement does not hold".to_owned() + } else { + "yes, the statement holds".to_owned() + } + } else { + render(v) + }; + format!("{}: {desc}", labels[i]) + }) + .collect() + } + }; + Ok((kind.to_owned(), criteria, opts)) +} + +fn python_strip(value: &str) -> &str { + value.trim_matches(|c: char| c.is_whitespace() || ('\u{1c}'..='\u{1f}').contains(&c)) +} + +// Preserve arbitrary-size integers; floats use Python's shortest decimal and notation rules. +fn python_number(n: &serde_json::Number) -> String { + let raw = n.to_string(); + if !raw.contains(['.', 'e', 'E']) { + return raw; + } + let Some(value) = n.as_f64().filter(|v| v.is_finite()) else { + return raw; + }; + let shortest = serde_json::Number::from_f64(value).unwrap().to_string(); + let (mantissa, exponent) = shortest.split_once('e').unwrap_or((&shortest, "0")); + let sign = if value.is_sign_negative() { "-" } else { "" }; + let mantissa = mantissa.trim_start_matches('-'); + let point = mantissa.find('.').unwrap_or(mantissa.len()); + let digits = mantissa.replace('.', ""); + let significant = digits.trim_start_matches('0'); + if significant.is_empty() { + return format!("{sign}0.0"); + } + let e = exponent.parse::().unwrap() + point as i32 + - (digits.len() - significant.len()) as i32 + - 1; + let significant = significant.trim_end_matches('0'); + if !(-4..16).contains(&e) { + let mut mantissa = significant.to_owned(); + if mantissa.len() > 1 { + mantissa.insert(1, '.'); + } + return format!("{sign}{mantissa}e{e:+03}"); + } + let point = e + 1; + let mut plain = significant.to_owned(); + if point <= 0 { + plain = format!("0.{}{plain}", "0".repeat((-point) as usize)); + } else if point as usize >= plain.len() { + plain.push_str(&"0".repeat(point as usize - plain.len())); + plain.push_str(".0"); + } else { + plain.insert(point as usize, '.'); + } + format!("{sign}{plain}") +} + +fn validate_numbers(value: &Value) -> Result<()> { + match value { + Value::Number(n) if n.to_string().contains(['.', 'e', 'E']) => ensure!( + n.as_f64().is_some_and(f64::is_finite), + "floating point value outside supported finite range" + ), + Value::Array(a) => { + for v in a { + validate_numbers(v)?; + } + } + Value::Object(o) => { + for v in o.values() { + validate_numbers(v)?; + } + } + _ => {} + } + Ok(()) +} diff --git a/tests/cua_s1/json.rs b/tests/cua_s1/json.rs new file mode 100644 index 0000000..bb0e070 --- /dev/null +++ b/tests/cua_s1/json.rs @@ -0,0 +1,18 @@ +use super::*; + +#[test] +fn number_encoding_does_not_change_user_objects() { + let input = + r#"{"a": {"$serde_json::private::Number": "1.5"}, "b": -0, "c": 18446744073709551616}"#; + let value = Value::Object(parse(input.as_bytes()).unwrap()); + assert_eq!(value["a"]["$serde_json::private::Number"], "1.5"); + assert_eq!( + dumps(&value), + r#"{"a": {"$serde_json::private::Number": "1.5"}, "b": -0.0, "c": 1.8446744073709552e+19}"# + ); + assert!( + parse(br#"{"a": {"x": 1, "\u0078": 2}}"#) + .unwrap_err() + .contains("duplicate key") + ); +} diff --git a/tests/laya/packing.rs b/tests/laya/packing.rs new file mode 100644 index 0000000..bd753d8 --- /dev/null +++ b/tests/laya/packing.rs @@ -0,0 +1,62 @@ +use omni_laya::preprocess::{Preprocessor, Request}; +use serde_json::Value; +use sha2::{Digest, Sha256}; +use std::path::PathBuf; + +fn checked_file(variable: &str, sha256: &str) -> PathBuf { + let path = + PathBuf::from(std::env::var_os(variable).unwrap_or_else(|| panic!("set {variable}"))); + let bytes = std::fs::read(&path).unwrap(); + assert_eq!( + format!("{:x}", Sha256::digest(&bytes)), + sha256, + "{variable}" + ); + path +} + +#[test] +#[ignore = "requires pinned tokenizer and packing oracle; run by CPU CI, no GPU"] +fn official_packing_parity() { + let tokenizer = checked_file( + "LAYA_TOKENIZER", + "6c8aaa9a542084f2457eab775d4eeb51f92a70c0fd9de28d5edb0ddec3c08d30", + ); + let oracle = checked_file( + "LAYA_PACKING_ORACLE", + "8cbaa311a59924ac9f9f2ba6f8438f82dafbd08476f64e86868289c2ede67f60", + ); + let pre = Preprocessor::load(&tokenizer).unwrap(); + let cases: Vec = serde_json::from_slice(&std::fs::read(oracle).unwrap()).unwrap(); + assert_eq!(cases.len(), 17); + for case in cases { + let request: Request = Request::from_value(case["request"].clone()).unwrap(); + let got = pre.prepare(&request).unwrap(); + let expected = &case["expected"]; + let items = expected["items"].as_array().unwrap(); + assert_eq!( + got.questions.len(), + items.len(), + "{} row count", + case["name"] + ); + assert_eq!(got.usage, expected["usage"].as_u64().unwrap() as usize); + for (i, ((question, item), id)) in got + .questions + .iter() + .zip(items) + .zip(request.questions.keys()) + .enumerate() + { + assert_eq!(&question.id, id, "{} row order", case["name"]); + let actual = serde_json::to_value(question).unwrap(); + for key in ["ids", "markers", "qtype"] { + assert_eq!(actual[key], item[key], "{} row {i} {key}", case["name"]); + } + assert_eq!( + question.ids.len(), + expected["lens"][i].as_u64().unwrap() as usize + ); + } + } +} diff --git a/tests/laya/preprocess.rs b/tests/laya/preprocess.rs new file mode 100644 index 0000000..0cbcc05 --- /dev/null +++ b/tests/laya/preprocess.rs @@ -0,0 +1,318 @@ +use omni_laya::preprocess::{Preprocessor, Request, render}; +use serde::Deserialize; +use serde_json::{Map, Value, json}; +use tokenizers::{Tokenizer, models::wordlevel::WordLevel, pre_tokenizers::whitespace::Whitespace}; + +// Small tokenizer for validation and packing boundaries; official parity is in packing.rs. +fn preprocessor() -> (tempfile::TempDir, Preprocessor) { + let dir = tempfile::tempdir().unwrap(); + let mut tokenizer = Tokenizer::new( + WordLevel::builder() + .vocab( + ["[UNK]", "[CLS]", "[SEP]", "[MASK]", "old", "new"] + .into_iter() + .enumerate() + .map(|(i, token)| (token.to_owned(), i as u32)) + .collect(), + ) + .unk_token("[UNK]".to_owned()) + .build() + .unwrap(), + ); + tokenizer.with_pre_tokenizer(Some(Whitespace)); + let path = dir.path().join("tokenizer.json"); + tokenizer.save(&path, false).unwrap(); + let pre = Preprocessor::load(&path).unwrap(); + (dir, pre) +} + +#[test] +fn python_json_numbers_and_order() { + for (input, expected) in [ + ("1e-6", "1e-06"), + ("1e20", "1e+20"), + ("1e16", "1e+16"), + ("1e-4", "0.0001"), + ("1.0", "1.0"), + ("-0.0", "-0.0"), + ("1331752170181752.2", "1331752170181752.2"), + ("-243915020125850.12", "-243915020125850.12"), + ("1e-5", "1e-05"), + ("5e-324", "5e-324"), + ("18446744073709551616000", "18446744073709551616000"), + ] { + assert_eq!(render(&serde_json::from_str(input).unwrap()), expected); + } + let v = serde_json::from_str(r#"{"z":1e-6,"a":"你好,x:y"}"#).unwrap(); + assert_eq!(render(&v), r#"{"z": 1e-06, "a": "你好,x:y"}"#); +} + +#[test] +fn strips_python_whitespace_from_noul_labels() { + let (_dir, pre) = preprocessor(); + let packed = |label| { + let request: Request = Request::from_value(json!({ + "state":"", "questions":{"q":{"type":"noul","instructions":"New?", + "labels":{"false":label,"true":"new"}}} + })) + .unwrap(); + pre.prepare(&request).unwrap().questions.remove(0).ids + }; + assert_eq!(packed("\u{1c}old\u{1f}"), packed("old")); + let invalid: Request = Request::from_value(json!({ + "state":"", "questions":{"q":{"type":"noul","instructions":"New?", + "labels":{"false":"\u{1c}\u{1f}","true":"new"}}} + })) + .unwrap(); + assert!( + pre.prepare(&invalid) + .unwrap_err() + .to_string() + .contains("invalid noul labels") + ); +} + +#[test] +fn validation_names_the_question_and_leaves_preprocessor_usable() { + let (_dir, pre) = preprocessor(); + for definition in [ + json!({"type":"choice","instructions":"Pick","criteria":[]}), + json!({"type":"noul","instructions":"Pick","criteria":{"yes":"ok"}}), + json!({"type":"noul","instructions":"Pick","labels":{"false":"x","true":"x"}}), + json!({"type":"score","criteria":["low","high"]}), + ] { + let request: Request = Request::from_value(json!({ + "state":"new", "questions":{"broken":definition} + })) + .unwrap(); + assert!( + pre.prepare(&request) + .unwrap_err() + .to_string() + .contains("broken") + ); + } + let request: Request = Request::from_value(json!({ + "state":"new", "questions":{"valid":{"type":"noul","instructions":"New?"}} + })) + .unwrap(); + assert_eq!(pre.prepare(&request).unwrap().questions.len(), 1); +} + +#[test] +fn preserves_question_and_choice_order() { + let (_dir, pre) = preprocessor(); + let request: Request = Request::from_json( + r#"{"state":"","questions":{ + "z":{"type":"choice","instructions":"Pick","criteria":["new","old","new"]}, + "a":{"type":"noul","instructions":"New?"} + }}"#, + ) + .unwrap(); + let prepared = pre.prepare(&request).unwrap(); + assert_eq!(prepared.questions[0].id, "z"); + assert_eq!(prepared.questions[1].id, "a"); + let choice = &prepared.questions[0]; + assert_eq!(choice.markers.len(), 2); + assert_eq!(choice.ids[choice.markers[0] + 1], 5); + assert_eq!(choice.ids[choice.markers[1] + 1], 4); + assert_eq!( + prepared.usage, + prepared + .questions + .iter() + .map(|q| q.ids.len()) + .sum::() + ); +} + +#[test] +fn keeps_newest_conversation_and_start_of_plain_text() { + let (_dir, pre) = preprocessor(); + for state in [ + json!(format!("{} new", "old ".repeat(1000))), + json!(["old ".repeat(1000), "new"]), + ] { + let is_conversation = state.is_array(); + let request: Request = Request::from_value(json!({ + "state":state, "questions":{"q":{"type":"noul","instructions":"New?"}} + })) + .unwrap(); + let prepared = pre.prepare(&request).unwrap(); + let ids = &prepared.questions[0].ids; + assert_eq!(ids.len(), 512); + assert_eq!(ids.contains(&5), is_conversation); + } +} + +#[test] +fn rejects_truncated_option_markers() { + let (_dir, pre) = preprocessor(); + let criteria: Vec<_> = (0..300).map(|i| i.to_string()).collect(); + let request: Request = Request::from_value(json!({ + "state":"", "questions":{"q":{"type":"choice","instructions":"Pick","criteria":criteria}} + })) + .unwrap(); + assert!( + pre.prepare(&request) + .unwrap_err() + .to_string() + .contains("options exceed") + ); +} + +#[test] +fn rejects_unsupported_language_and_nonfinite_numbers() { + let (_dir, pre) = preprocessor(); + for input in [ + r#"{"state":"","lang":"de","questions":{}}"#, + r#"{"state":"","model":"multilingual","questions":{}}"#, + r#"{"state":1e400,"questions":{}}"#, + ] { + assert!(pre.prepare(&Request::from_json(input).unwrap()).is_err()); + } +} + +#[test] +fn private_json_keys_stay_objects_in_raw_requests() { + let (_dir, pre) = preprocessor(); + for key in [ + "$serde_json::private::Number", + "$serde_json::private::RawValue", + ] { + let state = json!({"outer": [{(key): "1.5"}]}); + let questions = json!({ + "z": {"type": "choice", "instructions": {(key): "2"}, + "criteria": {"new": [{(key): "3"}], "old": null}}, + "a": {"type": "noul", "instructions": "New?"} + }) + .as_object() + .unwrap() + .clone(); + let expected = Request { + state, + model: None, + questions, + lang: None, + }; + let encoded = serde_json::to_string(&expected).unwrap(); + let request: Request = Request::from_json(&encoded).unwrap(); + assert_eq!(request.state, expected.state, "{key}"); + assert_eq!(request.questions, expected.questions, "{key}"); + assert_eq!(serde_json::to_string(&request).unwrap(), encoded); + assert_eq!( + render(&request.state), + format!(r#"{{"outer": [{{"{key}": "1.5"}}]}}"#) + ); + let packed = pre.prepare(&request).unwrap(); + let expected_packed = pre.prepare(&expected).unwrap(); + assert_eq!( + packed.questions[0].criteria, + expected_packed.questions[0].criteria + ); + assert_eq!(packed.questions[0].ids, expected_packed.questions[0].ids); + assert_eq!(packed.questions[0].id, "z"); + assert_eq!(packed.questions[1].id, "a"); + } +} + +#[test] +fn private_json_keys_stay_objects_from_value() { + for key in [ + "$serde_json::private::RawValue", + "$serde_json::private::Number", + ] { + let value = json!({"state": {(key): "1.5"}, "questions": { + "q": {"type": "score", "instructions": "New?", + "criteria": [{"outer": {(key): "2"}}]} + }}); + let request: Request = Request::from_value(value.clone()).unwrap(); + assert_eq!(request.state, value["state"]); + assert_eq!(Value::Object(request.questions), value["questions"]); + } +} + +#[test] +fn request_numbers_keep_arbitrary_precision_and_syntax() { + for number in [ + "18446744073709551616000", + "-18446744073709551616000", + "1e+03", + "1.2300", + "-0", + "1e400", + ] { + let encoded = + format!(r#"{{"state":[{number}],"questions":{{"q":{{"criteria":[{number}]}}}}}}"#); + let request: Request = Request::from_json(&encoded).unwrap(); + let scalar: Value = serde_json::from_str(number).unwrap(); + assert_eq!(request.state[0], scalar); + assert_eq!(request.questions["q"]["criteria"][0], scalar); + let roundtrip: Request = + Request::from_value(serde_json::to_value(&request).unwrap()).unwrap(); + assert_eq!(roundtrip.state, request.state); + assert_eq!(roundtrip.questions, request.questions); + } +} + +#[test] +fn request_keeps_default_json_recursion_limit() { + #[derive(Deserialize)] + struct Reference { + #[serde(rename = "state")] + _state: Value, + #[serde(rename = "questions")] + _questions: Map, + } + for depth in [124, 125, 126, 127, 128, 200] { + let nested = format!("{}0{}", "[".repeat(depth), "]".repeat(depth)); + for encoded in [ + format!(r#"{{"state":{nested},"questions":{{}}}}"#), + format!(r#"{{"state":null,"questions":{{"q":{{"criteria":{nested}}}}}}}"#), + ] { + let expected = serde_json::from_str::(&encoded).is_ok(); + let actual = Request::from_json(&encoded); + assert_eq!(actual.is_ok(), expected, "depth {depth}: {encoded}"); + if !expected { + assert!(actual.unwrap_err().to_string().contains("recursion limit")); + } + } + } +} + +#[test] +fn request_from_value_preserves_deep_values() { + let mut nested = json!({"$serde_json::private::Number": "1.5"}); + for _ in 0..200 { + nested = Value::Array(vec![nested]); + } + let value = json!({"state": nested.clone(), "questions": { + "q": {"criteria": nested.clone()} + }}); + let request = Request::from_value(value).unwrap(); + assert_eq!(request.state, nested); + assert_eq!(request.questions["q"]["criteria"], nested); + let request = Request::from_value(serde_json::to_value(&request).unwrap()).unwrap(); + assert_eq!(request.state, nested); +} + +#[test] +fn request_rejects_invalid_json_values() { + for state in ["NaN", "01", "[1,]", r#""\ud800""#] { + let encoded = format!(r#"{{"state":{state},"questions":{{}}}}"#); + assert!(Request::from_json(&encoded).is_err(), "{state}"); + } + assert!(Request::from_json(r#"{"state":null,"questions":[]}"#).is_err()); + for encoded in [ + r#"{"state":null,"questions":{},"extra":1}"#, + r#"{"state":null,"state":1,"questions":{}}"#, + r#"{"state":null}"#, + r#"{"questions":{}}"#, + r#"{"state":null,"questions":{},"lang":1}"#, + r#"[{"state":null,"questions":{}}]"#, + ] { + assert!(Request::from_json(encoded).is_err(), "{encoded}"); + } + let request = Request::from_json(r#"{"state":{"x":1,"x":2},"questions":{}}"#).unwrap(); + assert_eq!(request.state["x"], 2); +}