use std::cell::Cell;
use std::collections::HashSet;
use std::fmt;
use std::io::Read;
use serde::de::{
self, DeserializeSeed, Deserializer, Error as _, IgnoredAny, MapAccess, SeqAccess, Visitor,
};
use super::limits::MAX_JSON_DEPTH;
use super::HeadlessError;
const DEPTH_EXCEEDED: &str = "JSON nesting too deep";
const DUPLICATE_KEY: &str = "duplicate top-level key";
const UNKNOWN_FIELD: &str = "unknown field alongside prompt";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum InputFormat {
Text,
Json,
}
#[derive(Debug, Clone)]
pub struct Envelope {
pub prompt: String,
pub system: Option<String>,
pub model: Option<String>,
pub provider: Option<String>,
pub max_tool_calls: Option<u32>,
pub consult: Option<bool>,
}
enum MapOutcome {
Envelope(Envelope),
NoPrompt,
}
fn text_envelope(text: &str) -> Envelope {
Envelope {
prompt: text.to_string(),
system: None,
model: None,
provider: None,
max_tool_calls: None,
consult: None,
}
}
fn is_json_whitespace(byte: u8) -> bool {
matches!(byte, b' ' | b'\t' | b'\n' | b'\r')
}
fn enter_depth<E: de::Error>(depth: &Cell<u32>) -> Result<(), E> {
let next = depth.get().saturating_add(1);
if next > MAX_JSON_DEPTH {
return Err(E::custom(DEPTH_EXCEEDED));
}
depth.set(next);
Ok(())
}
fn leave_depth(depth: &Cell<u32>) {
depth.set(depth.get().saturating_sub(1));
}
struct DepthLimitedIgnoredAny<'a> {
depth: &'a Cell<u32>,
}
impl<'de, 'a> DeserializeSeed<'de> for DepthLimitedIgnoredAny<'a> {
type Value = ();
fn deserialize<D>(self, deserializer: D) -> Result<(), D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_any(DepthLimitedVisitor { depth: self.depth })
}
}
struct DepthLimitedVisitor<'a> {
depth: &'a Cell<u32>,
}
impl<'de, 'a> Visitor<'de> for DepthLimitedVisitor<'a> {
type Value = ();
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("any JSON value within the depth limit")
}
fn visit_bool<E>(self, _v: bool) -> Result<(), E> {
Ok(())
}
fn visit_i64<E>(self, _v: i64) -> Result<(), E> {
Ok(())
}
fn visit_u64<E>(self, _v: u64) -> Result<(), E> {
Ok(())
}
fn visit_f64<E>(self, _v: f64) -> Result<(), E> {
Ok(())
}
fn visit_str<E>(self, _v: &str) -> Result<(), E> {
Ok(())
}
fn visit_none<E>(self) -> Result<(), E> {
Ok(())
}
fn visit_unit<E>(self) -> Result<(), E> {
Ok(())
}
fn visit_seq<A>(self, mut seq: A) -> Result<(), A::Error>
where
A: SeqAccess<'de>,
{
enter_depth::<A::Error>(self.depth)?;
while seq
.next_element_seed(DepthLimitedIgnoredAny { depth: self.depth })?
.is_some()
{}
leave_depth(self.depth);
Ok(())
}
fn visit_map<A>(self, mut map: A) -> Result<(), A::Error>
where
A: MapAccess<'de>,
{
enter_depth::<A::Error>(self.depth)?;
while map.next_key::<IgnoredAny>()?.is_some() {
map.next_value_seed(DepthLimitedIgnoredAny { depth: self.depth })?;
}
leave_depth(self.depth);
Ok(())
}
}
struct EnvelopeVisitor {
depth: Cell<u32>,
}
impl<'de> Visitor<'de> for EnvelopeVisitor {
type Value = MapOutcome;
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("a JSON envelope object")
}
fn visit_map<A>(self, mut map: A) -> Result<MapOutcome, A::Error>
where
A: MapAccess<'de>,
{
enter_depth::<A::Error>(&self.depth)?;
let mut seen: HashSet<String> = HashSet::new();
let mut prompt: Option<String> = None;
let mut system: Option<String> = None;
let mut model: Option<String> = None;
let mut provider: Option<String> = None;
let mut max_tool_calls: Option<u32> = None;
let mut consult: Option<bool> = None;
let mut unknown_seen = false;
while let Some(key) = map.next_key::<String>()? {
if !seen.insert(key.clone()) {
return Err(A::Error::custom(DUPLICATE_KEY));
}
match key.as_str() {
"prompt" => prompt = Some(map.next_value::<String>()?),
"system" => system = map.next_value::<Option<String>>()?,
"model" => model = map.next_value::<Option<String>>()?,
"provider" => provider = map.next_value::<Option<String>>()?,
"max_tool_calls" => max_tool_calls = map.next_value::<Option<u32>>()?,
"consult" => consult = map.next_value::<Option<bool>>()?,
_ => {
unknown_seen = true;
map.next_value_seed(DepthLimitedIgnoredAny { depth: &self.depth })?;
}
}
}
leave_depth(&self.depth);
match prompt {
None => Ok(MapOutcome::NoPrompt),
Some(prompt) => {
if unknown_seen {
Err(A::Error::custom(UNKNOWN_FIELD))
} else {
Ok(MapOutcome::Envelope(Envelope {
prompt,
system,
model,
provider,
max_tool_calls,
consult,
}))
}
}
}
}
}
pub fn read_input_bounded(
reader: impl Read,
max_input_bytes: usize,
) -> Result<Vec<u8>, HeadlessError> {
let mut buf = Vec::new();
reader
.take(max_input_bytes as u64 + 1)
.read_to_end(&mut buf)
.map_err(|e| HeadlessError::Io(e.to_string()))?;
if buf.len() > max_input_bytes {
return Err(HeadlessError::InputTooLarge(max_input_bytes));
}
Ok(buf)
}
pub fn parse_input(
bytes: &[u8],
forced_fmt: Option<InputFormat>,
) -> Result<Envelope, HeadlessError> {
let text = std::str::from_utf8(bytes)
.map_err(|_| HeadlessError::InputInvalid("input is not valid UTF-8".to_string()))?;
if forced_fmt == Some(InputFormat::Text) {
return Ok(text_envelope(text));
}
let looks_like_object = text.bytes().find(|&b| !is_json_whitespace(b)) == Some(b'{');
if !looks_like_object {
if forced_fmt == Some(InputFormat::Json) {
return Err(HeadlessError::InputInvalid(
"expected a JSON object under --input-format json".to_string(),
));
}
return Ok(text_envelope(text));
}
let mut de = serde_json::Deserializer::from_slice(bytes);
let outcome = (&mut de)
.deserialize_map(EnvelopeVisitor {
depth: Cell::new(0),
})
.map_err(|_| HeadlessError::InputInvalid("malformed JSON envelope".to_string()))?;
de.end().map_err(|_| {
HeadlessError::InputInvalid("trailing data after JSON envelope".to_string())
})?;
match outcome {
MapOutcome::Envelope(envelope) => Ok(envelope),
MapOutcome::NoPrompt => {
if forced_fmt == Some(InputFormat::Json) {
Err(HeadlessError::InputInvalid(
"JSON object has no `prompt` field under --input-format json".to_string(),
))
} else {
Ok(text_envelope(text))
}
}
}
}
#[cfg(test)]
mod tests {
use std::io::Cursor;
use super::super::limits::MAX_INPUT_BYTES;
use super::*;
#[test]
fn test_read_input_rejects_oversized_without_buffering_all() {
let r = std::io::repeat(b'a');
assert!(matches!(
read_input_bounded(r, MAX_INPUT_BYTES),
Err(HeadlessError::InputTooLarge(_))
));
}
#[test]
fn test_read_input_empty_reader_returns_empty_vec() {
let r = Cursor::new(Vec::new());
assert_eq!(
read_input_bounded(r, MAX_INPUT_BYTES).unwrap(),
Vec::<u8>::new()
);
}
#[test]
fn test_read_input_accepts_exactly_max_input_bytes() {
let r = std::io::repeat(b'x').take(MAX_INPUT_BYTES as u64);
let out = read_input_bounded(r, MAX_INPUT_BYTES).expect("exactly the cap must be accepted");
assert_eq!(out.len(), MAX_INPUT_BYTES);
}
#[test]
fn test_read_input_rejects_max_input_bytes_plus_one() {
let r = std::io::repeat(b'x').take(MAX_INPUT_BYTES as u64 + 1);
assert!(matches!(
read_input_bounded(r, MAX_INPUT_BYTES),
Err(HeadlessError::InputTooLarge(limit)) if limit == MAX_INPUT_BYTES
));
}
#[test]
fn test_read_input_bounded_respects_custom_effective_cap() {
let small_cap = 10usize;
let r = Cursor::new(vec![b'x'; small_cap + 1]);
assert!(
matches!(
read_input_bounded(r, small_cap),
Err(HeadlessError::InputTooLarge(limit)) if limit == small_cap
),
"a custom (smaller) effective cap must be enforced, not the module constant"
);
let r_ok = Cursor::new(vec![b'x'; small_cap]);
let out =
read_input_bounded(r_ok, small_cap).expect("exactly the custom cap must be accepted");
assert_eq!(out.len(), small_cap);
}
#[test]
fn test_parse_input_autodetect() {
let e = parse_input(br#"{"prompt":"hi","consult":true}"#, None).unwrap();
assert_eq!(e.prompt, "hi");
assert_eq!(e.consult, Some(true));
let t = parse_input(b"just text", None).unwrap();
assert_eq!(t.prompt, "just text");
let j = parse_input(br#"{"foo":1}"#, None).unwrap(); assert_eq!(j.prompt, r#"{"foo":1}"#);
}
#[test]
fn test_parse_input_unknown_field_priority_depends_on_prompt_presence() {
let t = parse_input(br#"{"foo":1}"#, None).unwrap();
assert_eq!(t.prompt, r#"{"foo":1}"#);
assert!(matches!(
parse_input(br#"{"foo":1,"prompt":"x"}"#, None),
Err(HeadlessError::InputInvalid(_))
));
assert_eq!(
parse_input(br#"{"prompt":"x","consult":true}"#, None)
.unwrap()
.prompt,
"x"
);
}
#[test]
fn test_parse_input_rejects_nonstring_prompt_dupkey_and_deep() {
assert!(matches!(
parse_input(br#"{"prompt":123}"#, Some(InputFormat::Json)),
Err(HeadlessError::InputInvalid(_))
));
assert!(matches!(
parse_input(br#"{"prompt":"a","prompt":"b"}"#, None),
Err(HeadlessError::InputInvalid(_))
));
let deep = format!("{}{}{}", "[".repeat(100), "1", "]".repeat(100));
assert!(matches!(
parse_input(deep.as_bytes(), Some(InputFormat::Json)),
Err(HeadlessError::InputInvalid(_))
));
}
fn nested_object(levels: u32) -> String {
let mut s = String::from("1");
for _ in 0..levels {
s = format!(r#"{{"a":{s}}}"#);
}
s
}
#[test]
fn test_parse_input_depth_boundary_64_ok_65_rejected() {
assert!(parse_input(nested_object(64).as_bytes(), None).is_ok());
assert!(matches!(
parse_input(nested_object(65).as_bytes(), None),
Err(HeadlessError::InputInvalid(_))
));
}
#[test]
fn test_parse_input_depth_inside_unknown_field_is_bounded() {
let deep_value = format!("{}{}{}", "[".repeat(100), "1", "]".repeat(100));
let input = format!(r#"{{"prompt":"x","foo":{deep_value}}}"#);
assert!(matches!(
parse_input(input.as_bytes(), None),
Err(HeadlessError::InputInvalid(_))
));
}
#[test]
fn test_parse_input_deep_object_without_prompt_rejected_by_depth() {
let deep_value = format!("{}{}{}", "[".repeat(100), "1", "]".repeat(100));
let input = format!(r#"{{"foo":{deep_value}}}"#);
assert!(matches!(
parse_input(input.as_bytes(), None),
Err(HeadlessError::InputInvalid(_))
));
}
#[test]
fn test_parse_input_forced_text_never_parses_deep_object() {
let deep_value = format!("{}{}{}", "[".repeat(100), "1", "]".repeat(100));
let input = format!(r#"{{"foo":{deep_value}}}"#);
let e = parse_input(input.as_bytes(), Some(InputFormat::Text)).unwrap();
assert_eq!(e.prompt, input);
}
#[test]
fn test_parse_input_format_forcing() {
assert!(matches!(
parse_input(b"just text", Some(InputFormat::Json)),
Err(HeadlessError::InputInvalid(_))
));
let e = parse_input(br#"{"prompt":"x"}"#, Some(InputFormat::Text)).unwrap();
assert_eq!(e.prompt, r#"{"prompt":"x"}"#);
}
#[test]
fn test_parse_input_rejects_non_utf8() {
assert!(matches!(
parse_input(&[0xff, 0xfe, 0x00], None),
Err(HeadlessError::InputInvalid(_))
));
}
#[test]
fn test_parse_input_smoke_never_panics_on_degenerate_bytes() {
let deep = format!("{}1{}", "[".repeat(200), "]".repeat(200));
let cases: Vec<Vec<u8>> = vec![
Vec::new(),
vec![0xff, 0xfe, 0x00, 0x80],
deep.into_bytes(),
br#"{"prompt":"a","prompt":"b"}"#.to_vec(),
br#"{"prompt":123}"#.to_vec(),
b"{[not valid json".to_vec(),
br#"["array","not","object"]"#.to_vec(),
b"{".to_vec(),
b"plain text with { and [ chars".to_vec(),
br#"{"prompt":"x","unknown":{"nested":[1,2,3]}}"#.to_vec(),
];
for bytes in &cases {
for fmt in [None, Some(InputFormat::Json), Some(InputFormat::Text)] {
let _ = parse_input(bytes, fmt);
}
let _ = read_input_bounded(Cursor::new(bytes.clone()), MAX_INPUT_BYTES);
}
}
}