#![deny(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
use std::collections::BTreeMap;
use proptest::prelude::*;
use typesayer::{Adapter, ChatAdapter};
use typesayer_parser::{
ChatStreamParser, ParseEvent,
proptest_strategies::{
DriftPerturbation, arb_chunking_positions, arb_drift_perturbation,
arb_supported_type_and_value, arb_tagged_oneof_type_and_value, arb_unsupported_field_type,
arb_variant_type_and_value, chunk_at_positions, serialize_completion,
serialize_completion_with_drift, single_output_signature,
},
};
use typesayer_types::field::{FieldType, FieldValue};
fn parse_value_through_chunks(
field_type: &FieldType,
value: &FieldValue,
positions: &[u8],
) -> Vec<ParseEvent> {
let signature = single_output_signature("answer", field_type.clone());
let completion = serialize_completion("answer", field_type, value);
let chunks = chunk_at_positions(&completion, positions);
let mut parser = ChatStreamParser::new(&signature);
let mut events = Vec::new();
for chunk in chunks {
events.extend(parser.push(chunk));
}
events.extend(parser.finish());
events
}
proptest! {
#[test]
fn chat_stream_parser_roundtrips_any_supported_value(
(field_type, value) in arb_supported_type_and_value(),
positions in arb_chunking_positions(),
) {
let events = parse_value_through_chunks(&field_type, &value, &positions);
let expected = match (&field_type, &value) {
(FieldType::String, FieldValue::Str(s)) => FieldValue::Str(s.trim().to_owned()),
_ => value.clone(),
};
match events.last() {
Some(ParseEvent::StreamComplete { output: FieldValue::Object(map) }) => {
prop_assert_eq!(map.get("answer"), Some(&expected));
}
Some(other) => {
prop_assert!(false, "expected StreamComplete, got {other:?}");
}
None => {
prop_assert!(false, "no events emitted");
}
}
}
}
proptest! {
#[test]
fn unsupported_field_type_emits_stream_error_without_panic(
field_type in arb_unsupported_field_type(),
positions in arb_chunking_positions(),
) {
let signature = single_output_signature("answer", field_type.clone());
let completion = "[[ ## answer ## ]]\n \n[[ ## completed ## ]]";
let chunks = chunk_at_positions(completion, &positions);
let mut parser = ChatStreamParser::new(&signature);
let mut events = Vec::new();
for chunk in chunks {
events.extend(parser.push(chunk));
}
events.extend(parser.finish());
let has_stream_error = events.iter().any(|e| matches!(e, ParseEvent::StreamError { .. }));
prop_assert!(
has_stream_error,
"expected StreamError for unsupported type {field_type:?}, events: {events:?}"
);
}
}
proptest! {
#[test]
fn buffered_and_streaming_parsers_agree(
(field_type, value) in arb_supported_type_and_value(),
positions in arb_chunking_positions(),
perturb in arb_drift_perturbation(),
) {
let signature = single_output_signature("answer", field_type.clone());
let completion = serialize_completion_with_drift(
"answer", &field_type, &value, perturb,
);
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|error| TestCaseError::fail(error.to_string()))?;
let buffered_result: Result<BTreeMap<String, typesayer_types::FieldValue>, _> =
runtime.block_on(async {
ChatAdapter::default().parse(&signature, &completion).await
});
let Ok(buffered_map) = buffered_result else {
return Ok(());
};
let Some(buffered_value) = buffered_map.get("answer") else {
return Ok(());
};
let chunks = chunk_at_positions(&completion, &positions);
let mut parser = ChatStreamParser::new(&signature);
let mut events = Vec::new();
for chunk in chunks {
events.extend(parser.push(chunk));
}
events.extend(parser.finish());
let streamed_value = match events.last() {
Some(ParseEvent::StreamComplete { output: FieldValue::Object(map) }) => {
map.get("answer").cloned()
}
_ => None,
};
let expected = match (&field_type, &value) {
(FieldType::String, FieldValue::Str(s)) => FieldValue::Str(s.trim().to_owned()),
_ => buffered_value.clone(),
};
prop_assert_eq!(
streamed_value.as_ref(),
Some(&expected),
"buffered/streaming divergence on {:?} with perturbation {:?}",
field_type,
perturb
);
if perturb == DriftPerturbation::None {
let has_stream_error =
events.iter().any(|e| matches!(e, ParseEvent::StreamError { .. }));
prop_assert!(
!has_stream_error,
"streaming emitted StreamError on a buffered-accepted value: {events:?}"
);
}
}
}
proptest! {
#[test]
fn variant_buffered_and_streaming_agree(
(field_type, value) in arb_variant_type_and_value(),
positions in arb_chunking_positions(),
perturb in arb_drift_perturbation(),
) {
let signature = single_output_signature("answer", field_type.clone());
let completion = serialize_completion_with_drift(
"answer", &field_type, &value, perturb,
);
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|error| TestCaseError::fail(error.to_string()))?;
let buffered_result: Result<BTreeMap<String, typesayer_types::FieldValue>, _> =
runtime.block_on(async {
ChatAdapter::default().parse(&signature, &completion).await
});
let Ok(buffered_map) = buffered_result else {
return Ok(());
};
let Some(buffered_value) = buffered_map.get("answer") else {
return Ok(());
};
let chunks = chunk_at_positions(&completion, &positions);
let mut parser = ChatStreamParser::new(&signature);
let mut events = Vec::new();
for chunk in chunks {
events.extend(parser.push(chunk));
}
events.extend(parser.finish());
let streamed_value = match events.last() {
Some(ParseEvent::StreamComplete { output: FieldValue::Object(map) }) => {
map.get("answer").cloned()
}
_ => None,
};
prop_assert_eq!(
streamed_value.as_ref(),
Some(buffered_value),
"variant buffered/streaming divergence on {:?}",
field_type
);
}
}
proptest! {
#[test]
fn tagged_variant_buffered_and_streaming_agree(
(field_type, value) in arb_tagged_oneof_type_and_value(),
positions in arb_chunking_positions(),
perturb in arb_drift_perturbation(),
) {
let signature = single_output_signature("answer", field_type.clone());
let completion = serialize_completion_with_drift(
"answer", &field_type, &value, perturb,
);
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|error| TestCaseError::fail(error.to_string()))?;
let buffered_result: Result<BTreeMap<String, typesayer_types::FieldValue>, _> =
runtime.block_on(async {
ChatAdapter::default().parse(&signature, &completion).await
});
let Ok(buffered_map) = buffered_result else {
return Ok(());
};
let Some(buffered_value) = buffered_map.get("answer") else {
return Ok(());
};
let chunks = chunk_at_positions(&completion, &positions);
let mut parser = ChatStreamParser::new(&signature);
let mut events = Vec::new();
for chunk in chunks {
events.extend(parser.push(chunk));
}
events.extend(parser.finish());
let streamed_value = match events.last() {
Some(ParseEvent::StreamComplete { output: FieldValue::Object(map) }) => {
map.get("answer").cloned()
}
_ => None,
};
prop_assert_eq!(
streamed_value.as_ref(),
Some(buffered_value),
"tagged-variant buffered/streaming divergence on {:?}",
field_type
);
}
}