use crate::error::{MiniLLMError, Result};
#[derive(Debug)]
pub struct PayloadExtractor {
field: String,
lenient: bool,
state: State,
esc: Esc,
raw: String,
}
#[derive(Debug)]
enum State {
Prelude(Prelude),
Body,
MaybeClosed {
tail: String,
},
Closed {
saw_brace: bool,
},
Failed,
}
#[derive(Debug)]
enum Prelude {
BeforeBrace,
BeforeKey,
InKey {
key: String,
escaped: bool,
},
AfterKey {
matched: bool,
},
BeforeValue {
matched: bool,
},
SkipValue {
depth: u32,
in_string: bool,
escaped: bool,
},
AfterValue,
}
#[derive(Debug, Default, PartialEq)]
enum Esc {
#[default]
None,
Slash,
Unicode(String),
AwaitLowSlash { high: u16 },
AwaitLowU { high: u16 },
AwaitLowHex { high: u16, buf: String },
}
impl PayloadExtractor {
pub fn strict(field: impl Into<String>) -> Self {
Self::new(field, false)
}
pub fn lenient(field: impl Into<String>) -> Self {
Self::new(field, true)
}
fn new(field: impl Into<String>, lenient: bool) -> Self {
Self {
field: field.into(),
lenient,
state: State::Prelude(Prelude::BeforeBrace),
esc: Esc::default(),
raw: String::new(),
}
}
pub fn feed(&mut self, fragment: &str) -> Result<String> {
if matches!(self.state, State::Failed) {
return Err(self.error("extractor already failed"));
}
self.raw.push_str(fragment);
let mut out = String::new();
for c in fragment.chars() {
if let Err(e) = self.step(c, &mut out) {
self.state = State::Failed;
return Err(e);
}
}
Ok(out)
}
pub fn finish(mut self) -> Result<String> {
match std::mem::replace(&mut self.state, State::Failed) {
State::Failed => Err(self.error("extractor already failed")),
State::Prelude(_) => Err(self.error(&format!(
"payload field '{}' never started (arguments incomplete or field missing)",
self.field
))),
State::Body => {
if self.lenient {
let mut out = String::new();
match std::mem::take(&mut self.esc) {
Esc::None => {}
Esc::Slash => out.push('\\'),
Esc::Unicode(buf) => {
out.push_str("\\u");
out.push_str(&buf);
}
Esc::AwaitLowSlash { .. } => out.push('\u{FFFD}'),
Esc::AwaitLowU { .. } => {
out.push('\u{FFFD}');
out.push('\\');
}
Esc::AwaitLowHex { buf, .. } => {
out.push('\u{FFFD}');
out.push_str("\\u");
out.push_str(&buf);
}
}
Ok(out)
} else {
Err(self.error("payload string never closed"))
}
}
State::MaybeClosed { .. } => Ok(String::new()),
State::Closed { saw_brace } => {
if saw_brace {
Ok(String::new())
} else {
Err(self.error("arguments object never closed after the payload string"))
}
}
}
}
fn error(&self, what: &str) -> MiniLLMError {
MiniLLMError::MalformedResponse(format!(
"streamed tool arguments: {} (field '{}', raw: {})",
what, self.field, self.raw
))
}
fn step(&mut self, c: char, out: &mut String) -> Result<()> {
match &mut self.state {
State::Prelude(_) => self.step_prelude(c),
State::Body => self.step_body(c, out),
State::MaybeClosed { tail } => {
let fits = c.is_whitespace() || (c == '}' && !tail.contains('}'));
if fits {
tail.push(c);
return Ok(());
}
let tail = std::mem::take(tail);
self.state = State::Body;
out.push('"');
for tc in tail.chars() {
self.step_body(tc, out)?;
}
self.step_body(c, out)
}
State::Closed { saw_brace } => {
if c.is_whitespace() {
Ok(())
} else if c == '}' && !*saw_brace {
*saw_brace = true;
Ok(())
} else {
Err(self.error(&format!("unexpected '{c}' after the payload string closed")))
}
}
State::Failed => unreachable!("feed guards Failed"),
}
}
fn step_prelude(&mut self, c: char) -> Result<()> {
let State::Prelude(p) = &mut self.state else {
unreachable!()
};
match p {
Prelude::BeforeBrace => {
if c.is_whitespace() {
} else if c == '{' {
*p = Prelude::BeforeKey;
} else {
return Err(
self.error(&format!("arguments do not start with '{{' (got '{c}')"))
);
}
}
Prelude::BeforeKey => {
if c.is_whitespace() {
} else if c == '"' {
*p = Prelude::InKey {
key: String::new(),
escaped: false,
};
} else if c == '}' {
return Err(self.error(&format!(
"arguments object closed before the payload field '{}'",
self.field
)));
} else {
return Err(self.error(&format!("expected a field name, got '{c}'")));
}
}
Prelude::InKey { key, escaped } => {
if *escaped {
key.push(c);
*escaped = false;
} else if c == '\\' {
*escaped = true;
} else if c == '"' {
let matched = *key == self.field;
*p = Prelude::AfterKey { matched };
} else {
key.push(c);
}
}
Prelude::AfterKey { matched } => {
if c.is_whitespace() {
} else if c == ':' {
*p = Prelude::BeforeValue { matched: *matched };
} else {
return Err(self.error(&format!("expected ':' after a field name, got '{c}'")));
}
}
Prelude::BeforeValue { matched } => {
if c.is_whitespace() {
} else if *matched {
if c == '"' {
self.state = State::Body;
} else {
return Err(self.error(&format!(
"payload field '{}' is not a string (starts with '{c}')",
self.field
)));
}
} else if c == '"' {
*p = Prelude::SkipValue {
depth: 0,
in_string: true,
escaped: false,
};
} else if c == '{' || c == '[' {
*p = Prelude::SkipValue {
depth: 1,
in_string: false,
escaped: false,
};
} else {
*p = Prelude::SkipValue {
depth: 0,
in_string: false,
escaped: false,
};
}
}
Prelude::SkipValue {
depth,
in_string,
escaped,
} => {
if *in_string {
if *escaped {
*escaped = false;
} else if c == '\\' {
*escaped = true;
} else if c == '"' {
*in_string = false;
if *depth == 0 {
*p = Prelude::AfterValue;
}
}
} else if *depth > 0 {
match c {
'"' => *in_string = true,
'{' | '[' => *depth += 1,
'}' | ']' => {
*depth -= 1;
if *depth == 0 {
*p = Prelude::AfterValue;
}
}
_ => {}
}
} else {
if c == ',' {
*p = Prelude::BeforeKey;
} else if c == '}' {
return Err(self.error(&format!(
"arguments object closed before the payload field '{}'",
self.field
)));
} else if c.is_whitespace() {
*p = Prelude::AfterValue;
}
}
}
Prelude::AfterValue => {
if c.is_whitespace() {
} else if c == ',' {
*p = Prelude::BeforeKey;
} else if c == '}' {
return Err(self.error(&format!(
"arguments object closed before the payload field '{}'",
self.field
)));
} else {
return Err(self.error(&format!("expected ',' or '}}', got '{c}'")));
}
}
}
Ok(())
}
fn step_body(&mut self, c: char, out: &mut String) -> Result<()> {
match std::mem::take(&mut self.esc) {
Esc::None => match c {
'\\' => self.esc = Esc::Slash,
'"' => {
self.state = if self.lenient {
State::MaybeClosed {
tail: String::new(),
}
} else {
State::Closed { saw_brace: false }
};
}
c if (c as u32) < 0x20 => {
if self.lenient {
out.push(c); } else {
return Err(self.error(&format!(
"unescaped control character U+{:04X} in payload string",
c as u32
)));
}
}
c => out.push(c),
},
Esc::Slash => match c {
'"' => out.push('"'),
'\\' => out.push('\\'),
'/' => out.push('/'),
'b' => out.push('\u{8}'),
'f' => out.push('\u{c}'),
'n' => out.push('\n'),
'r' => out.push('\r'),
't' => out.push('\t'),
'u' => self.esc = Esc::Unicode(String::new()),
other => {
if self.lenient {
out.push('\\');
return self.step_body(other, out);
}
return Err(self.error(&format!("invalid escape '\\{other}'")));
}
},
Esc::Unicode(mut buf) => {
if c.is_ascii_hexdigit() {
buf.push(c);
if buf.len() == 4 {
let code = u16::from_str_radix(&buf, 16).expect("4 hex digits");
self.take_unit(code, out)?;
} else {
self.esc = Esc::Unicode(buf);
}
} else {
if self.lenient {
out.push_str("\\u");
out.push_str(&buf);
return self.step_body(c, out);
}
return Err(self.error(&format!("invalid \\u escape '\\u{buf}{c}'")));
}
}
Esc::AwaitLowSlash { high } => {
if c == '\\' {
self.esc = Esc::AwaitLowU { high };
} else {
if self.lenient {
out.push('\u{FFFD}'); return self.step_body(c, out);
}
return Err(self.error("unpaired \\u surrogate in payload string"));
}
}
Esc::AwaitLowU { high } => {
if c == 'u' {
self.esc = Esc::AwaitLowHex {
high,
buf: String::new(),
};
} else {
if self.lenient {
out.push('\u{FFFD}');
self.esc = Esc::Slash;
return self.step_body(c, out);
}
return Err(self.error("unpaired \\u surrogate in payload string"));
}
}
Esc::AwaitLowHex { high, mut buf } => {
if c.is_ascii_hexdigit() {
buf.push(c);
if buf.len() == 4 {
let low = u16::from_str_radix(&buf, 16).expect("4 hex digits");
if (0xDC00..=0xDFFF).contains(&low) {
let combined =
0x10000 + ((high as u32 - 0xD800) << 10) + (low as u32 - 0xDC00);
out.push(char::from_u32(combined).expect("valid surrogate pair"));
} else {
if !self.lenient {
return Err(self.error("unpaired \\u surrogate in payload string"));
}
out.push('\u{FFFD}');
self.take_unit(low, out)?;
}
} else {
self.esc = Esc::AwaitLowHex { high, buf };
}
} else {
if self.lenient {
out.push('\u{FFFD}');
out.push_str("\\u");
out.push_str(&buf);
return self.step_body(c, out);
}
return Err(self.error("unpaired \\u surrogate in payload string"));
}
}
}
Ok(())
}
fn take_unit(&mut self, code: u16, out: &mut String) -> Result<()> {
if (0xD800..=0xDBFF).contains(&code) {
self.esc = Esc::AwaitLowSlash { high: code };
Ok(())
} else if (0xDC00..=0xDFFF).contains(&code) {
if self.lenient {
out.push('\u{FFFD}');
Ok(())
} else {
Err(self.error("unpaired \\u surrogate in payload string"))
}
} else {
out.push(char::from_u32(code as u32).expect("non-surrogate BMP code point"));
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn all_splits(field: &str, lenient: bool, raw: &str) -> Result<String> {
let make = || {
if lenient {
PayloadExtractor::lenient(field)
} else {
PayloadExtractor::strict(field)
}
};
let run = |fragments: Vec<&str>| -> Result<String> {
let mut ex = make();
let mut out = String::new();
for f in fragments {
out.push_str(&ex.feed(f)?);
}
out.push_str(&ex.finish()?);
Ok(out)
};
let reference = run(vec![raw]);
for i in 0..=raw.len() {
if !raw.is_char_boundary(i) {
continue;
}
let split = run(vec![&raw[..i], &raw[i..]]);
match (&reference, &split) {
(Ok(a), Ok(b)) => assert_eq!(a, b, "split at {i} diverged for {raw:?}"),
(Err(_), Err(_)) => {}
_ => panic!("split at {i} changed ok/err for {raw:?}"),
}
}
let mut ex = make();
let mut out = String::new();
let mut failed = false;
for c in raw.chars() {
match ex.feed(&c.to_string()) {
Ok(s) => out.push_str(&s),
Err(_) => {
failed = true;
break;
}
}
}
if !failed {
match (ex.finish(), &reference) {
(Ok(s), Ok(r)) => {
out.push_str(&s);
assert_eq!(&out, r, "char-by-char diverged for {raw:?}");
}
(Err(_), Err(_)) => {}
(a, b) => panic!("char-by-char changed ok/err for {raw:?}: {a:?} vs {b:?}"),
}
} else {
assert!(reference.is_err(), "char-by-char failed but whole ok");
}
reference
}
fn strict(raw: &str) -> Result<String> {
all_splits("content", false, raw)
}
fn lenient(raw: &str) -> Result<String> {
all_splits("content", true, raw)
}
#[test]
fn strict_decodes_simple_payload() {
assert_eq!(
strict(r#"{"content": "hello world"}"#).unwrap(),
"hello world"
);
}
#[test]
fn strict_decodes_every_escape() {
assert_eq!(
strict(r#"{"content":"a\"b\\c\/d\be\ff\ng\rh\tiéj"}"#).unwrap(),
"a\"b\\c/d\u{8}e\u{c}f\ng\rh\ti\u{e9}j"
);
}
#[test]
fn strict_decodes_surrogate_pairs() {
assert_eq!(strict(r#"{"content":"ok 😀!"}"#).unwrap(), "ok 😀!");
}
#[test]
fn strict_skips_preceding_fields() {
assert_eq!(
strict(
r#"{"lang": "py\"x", "n": 4.5e2, "flag": true,
"meta": {"a": ["b", {"c": 1}]}, "content": "payload"}"#
)
.unwrap(),
"payload"
);
}
#[test]
fn strict_handles_empty_payload_and_whitespace_envelope() {
assert_eq!(strict(" { \"content\" : \"\" } ").unwrap(), "");
}
#[test]
fn strict_rejects_malformed() {
for bad in [
r#"{"content": "unterminated"#, r#"{"content": "no brace""#, r#"{"content": "bad \q escape"}"#, r#"{"content": "lone \ud800 high"}"#, r#"{"content": 42}"#, r#"{"other": "x"}"#, r#"{"content": "a"} trailing"#, "{\"content\": \"raw\nnewline\"}", r#"["content"]"#, ] {
assert!(strict(bad).is_err(), "must reject: {bad:?}");
}
}
#[test]
fn error_carries_the_raw_text() {
let mut ex = PayloadExtractor::strict("content");
let err = ex.feed(r#"{"content": 42"#).unwrap_err().to_string();
assert!(
err.contains(r#"{"content": 42"#),
"raw text in error: {err}"
);
assert!(ex.feed("x").is_err());
}
#[test]
fn lenient_decodes_well_formed_identically() {
let raw = r#"{"content":"a\"b\\c\nd 😀"}"#;
assert_eq!(lenient(raw).unwrap(), strict(raw).unwrap());
}
#[test]
fn lenient_unescaped_quote_mid_content_is_literal() {
assert_eq!(
lenient(r#"{"content": "say "hi" ok"}"#).unwrap(),
r#"say "hi" ok"#
);
}
#[test]
fn lenient_raw_newline_is_itself() {
assert_eq!(
lenient("{\"content\": \"line1\nline2\"}").unwrap(),
"line1\nline2"
);
}
#[test]
fn lenient_backslash_before_non_escape_is_literal() {
assert_eq!(
lenient(r#"{"content": "C:\path \q \ x"}"#).unwrap(),
r#"C:\path \q \ x"#
);
}
#[test]
fn lenient_model_forgot_the_closing_quote_and_brace() {
assert_eq!(
lenient(r#"{"content": "the model stopped here"#).unwrap(),
"the model stopped here"
);
}
#[test]
fn lenient_model_closed_string_but_not_object() {
assert_eq!(lenient(r#"{"content": "done""#).unwrap(), "done");
}
#[test]
fn lenient_trailing_partial_escape_is_flushed_literally() {
assert_eq!(
lenient(r#"{"content": "ends with \"#).unwrap(),
"ends with \\"
);
assert_eq!(
lenient(r#"{"content": "ends with \u12"#).unwrap(),
"ends with \\u12"
);
}
#[test]
fn lenient_quote_then_more_content_after_whitespace_and_brace() {
assert_eq!(lenient(r#"{"content": "a "} b"}"#).unwrap(), r#"a "} b"#);
}
#[test]
fn lenient_unpaired_surrogate_is_replacement_not_dropped() {
assert_eq!(
lenient(r#"{"content": "x \ud800 y"}"#).unwrap(),
"x \u{FFFD} y"
);
}
#[test]
fn lenient_still_rejects_a_broken_prelude() {
assert!(lenient(r#"{"other": "x"}"#).is_err());
assert!(lenient(r#"not json at all"#).is_err());
}
#[test]
fn code_file_payload_streams_decoded() {
let code = "fn main() {\n println!(\"hi \\\\ there\");\n}\n";
let raw = format!(r#"{{"content": {}}}"#, serde_json::to_string(code).unwrap());
assert_eq!(strict(&raw).unwrap(), code);
assert_eq!(lenient(&raw).unwrap(), code);
}
}