#[derive(Debug)]
pub enum WireEvent<T> {
Known(T),
Unknown {
event_type: String,
value: crate::streaming::UnknownPayload,
},
Corrupt(serde_json::Error),
}
impl<T> WireEvent<T> {
pub fn map<U>(self, f: impl FnOnce(T) -> U) -> WireEvent<U> {
match self {
Self::Known(event) => WireEvent::Known(f(event)),
Self::Unknown { event_type, value } => WireEvent::Unknown { event_type, value },
Self::Corrupt(error) => WireEvent::Corrupt(error),
}
}
}
pub fn classify_tagged_frame<T>(
data: &str,
tag: &str,
is_known_event_type: impl Fn(&str) -> bool,
) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
let scanned = match scan_discriminators(data, &[tag], true) {
Ok(scanned) => scanned,
Err(error) => return WireEvent::Corrupt(error),
};
match scanned {
DiscriminatorScan::Object(found) => {
match found.first().and_then(|key| key.string_value.as_deref()) {
Some(event_type) if !is_known_event_type(event_type) => {
unknown_with_value(data, event_type.to_owned())
}
_ => decode_known(data),
}
}
DiscriminatorScan::NotObject => unknown_with_value(data, String::new()),
}
}
fn unknown_with_value<T>(data: &str, event_type: String) -> WireEvent<T> {
match serde_json::from_str::<serde_json::Value>(data) {
Ok(value) => WireEvent::Unknown {
event_type,
value: value.into(),
},
Err(error) => WireEvent::Corrupt(error),
}
}
pub fn classify_chat_completions_frame<T>(data: &str) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
let scanned = match scan_discriminators(data, &["object", "choices"], true) {
Ok(scanned) => scanned,
Err(error) => return WireEvent::Corrupt(error),
};
let found = match scanned {
DiscriminatorScan::Object(found) => found,
DiscriminatorScan::NotObject => return unknown_with_value(data, String::new()),
};
let object_value = found.first().and_then(|key| key.string_value.as_deref());
let has_choices = found.get(1).is_some_and(|key| key.present);
let is_chat_chunk =
object_value.is_some_and(|object| object == "chat.completion.chunk") || has_choices;
if !is_chat_chunk {
return unknown_with_value(data, object_value.unwrap_or_default().to_owned());
}
decode_known(data)
}
pub fn classify_marker_keyed_frame<T>(data: &str, marker_keys: &[&str]) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
let scanned = match scan_discriminators(data, marker_keys, false) {
Ok(scanned) => scanned,
Err(error) => return WireEvent::Corrupt(error),
};
let recognizable = match &scanned {
DiscriminatorScan::Object(found) => found.iter().any(|key| key.present),
DiscriminatorScan::NotObject => false,
};
if !recognizable {
let value = match serde_json::from_str::<serde_json::Value>(data) {
Ok(value) => value,
Err(error) => return WireEvent::Corrupt(error),
};
let event_type = value
.as_object()
.map(|object| object.keys().cloned().collect::<Vec<_>>().join(","))
.unwrap_or_default();
return WireEvent::Unknown {
event_type,
value: value.into(),
};
}
decode_known(data)
}
pub fn classify_untyped_line<T>(line: &[u8]) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
match serde_json::from_slice::<T>(line) {
Ok(event) => WireEvent::Known(event),
Err(error) => WireEvent::Corrupt(error),
}
}
#[derive(Debug)]
pub enum TypedEvent<T> {
Modeled(T),
Unrecognized {
event_type: String,
detail: String,
},
Malformed(String),
}
pub fn classify_typed_event<T>(event: TypedEvent<T>) -> WireEvent<T> {
match event {
TypedEvent::Modeled(event) => WireEvent::Known(event),
TypedEvent::Unrecognized { event_type, detail } => WireEvent::Unknown {
event_type,
value: serde_json::Value::String(detail).into(),
},
TypedEvent::Malformed(message) => {
WireEvent::Corrupt(<serde_json::Error as serde::de::Error>::custom(message))
}
}
}
pub fn classify_with_repair<T>(
data: &str,
classify: impl Fn(&str) -> WireEvent<T>,
repair: impl FnOnce(&str) -> Option<String>,
on_unrepairable: impl FnOnce(&serde_json::Error) -> serde_json::Error,
on_still_corrupt: impl FnOnce() -> serde_json::Error,
) -> WireEvent<T> {
match classify(data) {
WireEvent::Corrupt(corrupt) => match repair(data) {
None => WireEvent::Corrupt(on_unrepairable(&corrupt)),
Some(repaired) => match classify(&repaired) {
WireEvent::Known(event) => WireEvent::Known(event),
WireEvent::Unknown { .. } | WireEvent::Corrupt(_) => {
WireEvent::Corrupt(on_still_corrupt())
}
},
},
event => event,
}
}
enum DiscriminatorScan {
Object(Vec<KeyScan>),
NotObject,
}
#[derive(Default, Clone)]
struct KeyScan {
present: bool,
string_value: Option<String>,
}
fn scan_discriminators(
data: &str,
keys: &[&str],
reject_duplicates: bool,
) -> Result<DiscriminatorScan, serde_json::Error> {
struct Scan<'a> {
keys: &'a [&'a str],
reject_duplicates: bool,
}
impl<'de> serde::de::Visitor<'de> for Scan<'_> {
type Value = DiscriminatorScan;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a JSON value")
}
fn visit_map<A>(self, mut map: A) -> Result<DiscriminatorScan, A::Error>
where
A: serde::de::MapAccess<'de>,
{
let mut found = vec![KeyScan::default(); self.keys.len()];
while let Some(key) = map.next_key::<String>()? {
match self.keys.iter().position(|candidate| *candidate == key) {
Some(index) => {
let entry = found.get_mut(index).ok_or_else(|| {
serde::de::Error::custom("discriminator index out of range")
})?;
if entry.present {
if self.reject_duplicates {
return Err(serde::de::Error::custom(format!(
"duplicate `{key}` discriminator key in stream frame"
)));
}
map.next_value::<serde::de::IgnoredAny>()?;
continue;
}
entry.present = true;
entry.string_value = match map.next_value::<StringOrIgnored>()? {
StringOrIgnored::String(value) => Some(value),
StringOrIgnored::Ignored => None,
};
}
None => {
map.next_value::<serde::de::IgnoredAny>()?;
}
}
}
Ok(DiscriminatorScan::Object(found))
}
fn visit_bool<E>(self, _: bool) -> Result<DiscriminatorScan, E> {
Ok(DiscriminatorScan::NotObject)
}
fn visit_i64<E>(self, _: i64) -> Result<DiscriminatorScan, E> {
Ok(DiscriminatorScan::NotObject)
}
fn visit_u64<E>(self, _: u64) -> Result<DiscriminatorScan, E> {
Ok(DiscriminatorScan::NotObject)
}
fn visit_f64<E>(self, _: f64) -> Result<DiscriminatorScan, E> {
Ok(DiscriminatorScan::NotObject)
}
fn visit_str<E>(self, _: &str) -> Result<DiscriminatorScan, E> {
Ok(DiscriminatorScan::NotObject)
}
fn visit_unit<E>(self) -> Result<DiscriminatorScan, E> {
Ok(DiscriminatorScan::NotObject)
}
fn visit_seq<A>(self, mut seq: A) -> Result<DiscriminatorScan, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
while seq.next_element::<serde::de::IgnoredAny>()?.is_some() {}
Ok(DiscriminatorScan::NotObject)
}
}
enum StringOrIgnored {
String(String),
Ignored,
}
impl<'de> serde::Deserialize<'de> for StringOrIgnored {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
struct V;
impl<'de> serde::de::Visitor<'de> for V {
type Value = StringOrIgnored;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("any JSON value")
}
fn visit_str<E>(self, value: &str) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::String(value.to_owned()))
}
fn visit_string<E>(self, value: String) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::String(value))
}
fn visit_bool<E>(self, _: bool) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::Ignored)
}
fn visit_i64<E>(self, _: i64) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::Ignored)
}
fn visit_u64<E>(self, _: u64) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::Ignored)
}
fn visit_f64<E>(self, _: f64) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::Ignored)
}
fn visit_unit<E>(self) -> Result<StringOrIgnored, E> {
Ok(StringOrIgnored::Ignored)
}
fn visit_map<A>(self, mut map: A) -> Result<StringOrIgnored, A::Error>
where
A: serde::de::MapAccess<'de>,
{
while map
.next_entry::<serde::de::IgnoredAny, serde::de::IgnoredAny>()?
.is_some()
{}
Ok(StringOrIgnored::Ignored)
}
fn visit_seq<A>(self, mut seq: A) -> Result<StringOrIgnored, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
while seq.next_element::<serde::de::IgnoredAny>()?.is_some() {}
Ok(StringOrIgnored::Ignored)
}
}
deserializer.deserialize_any(V)
}
}
let mut deserializer = serde_json::Deserializer::from_str(data);
let scanned = serde::Deserializer::deserialize_any(
&mut deserializer,
Scan {
keys,
reject_duplicates,
},
)?;
deserializer.end()?;
Ok(scanned)
}
fn decode_known<T>(data: &str) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
match serde_json::from_str::<T>(data) {
Ok(event) => WireEvent::Known(event),
Err(error) => WireEvent::Corrupt(error),
}
}
#[cfg(test)]
mod tests {
use super::{WireEvent, classify_chat_completions_frame, classify_tagged_frame};
#[derive(Debug, serde::Deserialize)]
#[serde(tag = "type")]
enum TestEvent {
#[serde(rename = "text.delta")]
TextDelta { delta: String },
}
fn known(event_type: &str) -> bool {
event_type == "text.delta"
}
#[derive(Debug, serde::Deserialize)]
struct TestChunk {
#[allow(dead_code)]
choices: Vec<serde_json::Value>,
}
#[test]
fn tagged_known_frame_decodes() {
let event = classify_tagged_frame::<TestEvent>(
r#"{"type":"text.delta","delta":"hi"}"#,
"type",
known,
);
assert!(matches!(event, WireEvent::Known(TestEvent::TextDelta { delta }) if delta == "hi"));
}
#[test]
fn tagged_unknown_type_is_unknown() {
let event = classify_tagged_frame::<TestEvent>(r#"{"type":"future.event"}"#, "type", known);
assert!(matches!(
event,
WireEvent::Unknown { event_type, .. } if event_type == "future.event"
));
}
#[test]
fn tagged_invalid_json_is_corrupt() {
let event = classify_tagged_frame::<TestEvent>("{not json", "type", known);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn tagged_known_type_with_defective_payload_is_corrupt() {
let event = classify_tagged_frame::<TestEvent>(
r#"{"type":"text.delta","delta":42}"#,
"type",
known,
);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn tagged_typeless_frame_is_corrupt() {
let event = classify_tagged_frame::<TestEvent>("{}", "type", known);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn tagged_duplicate_discriminator_is_corrupt() {
let event = classify_tagged_frame::<TestEvent>(
r#"{"type":"text.delta","type":"future.event","delta":"hi"}"#,
"type",
known,
);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn chat_duplicate_object_discriminator_is_corrupt() {
let event = classify_chat_completions_frame::<TestChunk>(
r#"{"object":"chat.completion.chunk","object":"future.thing","data":1}"#,
);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn chat_duplicate_choices_key_is_corrupt() {
let event = classify_chat_completions_frame::<TestChunk>(r#"{"choices":[],"choices":42}"#);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn tagged_duplicate_non_discriminator_key_still_classifies() {
let event = classify_tagged_frame::<TestEvent>(
r#"{"type":"text.delta","ignored":1,"ignored":2,"delta":"hi"}"#,
"type",
known,
);
assert!(matches!(event, WireEvent::Known(TestEvent::TextDelta { delta }) if delta == "hi"));
}
#[test]
fn chat_recognizable_chunk_decodes() {
let event = classify_chat_completions_frame::<TestChunk>(r#"{"choices":[]}"#);
assert!(matches!(event, WireEvent::Known(_)));
}
#[test]
fn chat_unrecognizable_json_is_unknown() {
let event = classify_chat_completions_frame::<TestChunk>(r#"{"object":"ping"}"#);
assert!(matches!(
event,
WireEvent::Unknown { event_type, .. } if event_type == "ping"
));
}
#[test]
fn chat_recognizable_chunk_with_defective_payload_is_corrupt() {
let event = classify_chat_completions_frame::<TestChunk>(r#"{"choices":42}"#);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn chat_invalid_json_is_corrupt() {
let event = classify_chat_completions_frame::<TestChunk>("{not json");
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn non_object_json_is_unknown_never_corrupt() {
for frame in ["null", "[]", "42", r#""ping""#] {
let event = classify_chat_completions_frame::<TestChunk>(frame);
assert!(
matches!(event, WireEvent::Unknown { .. }),
"chat classifier must skip {frame}, got {event:?}"
);
let event = classify_tagged_frame::<TestEvent>(frame, "type", known);
assert!(
matches!(event, WireEvent::Unknown { .. }),
"tagged classifier must skip {frame}, got {event:?}"
);
}
}
#[test]
fn tagged_dispatch_honors_a_non_type_tag_name() {
#[derive(Debug, serde::Deserialize)]
#[serde(tag = "event_type")]
enum EventTypeTagged {
#[serde(rename = "step.delta")]
StepDelta { delta: String },
}
let event = classify_tagged_frame::<EventTypeTagged>(
r#"{"event_type":"step.delta","delta":"hi"}"#,
"event_type",
|event_type| event_type == "step.delta",
);
assert!(matches!(
event,
WireEvent::Known(EventTypeTagged::StepDelta { delta }) if delta == "hi"
));
let event = classify_tagged_frame::<EventTypeTagged>(
r#"{"event_type":"future.event"}"#,
"event_type",
|event_type| event_type == "step.delta",
);
assert!(matches!(
event,
WireEvent::Unknown { event_type, .. } if event_type == "future.event"
));
}
#[test]
fn marker_keyed_recognizable_chunk_decodes() {
let event = super::classify_marker_keyed_frame::<TestChunk>(
r#"{"choices":[]}"#,
&["choices", "usage"],
);
assert!(matches!(event, WireEvent::Known(_)));
}
#[test]
fn marker_keyed_unrecognizable_json_is_unknown() {
let event = super::classify_marker_keyed_frame::<TestChunk>(
r#"{"noise":true,"other":1}"#,
&["choices", "usage"],
);
assert!(matches!(
event,
WireEvent::Unknown { event_type, .. } if event_type == "noise,other"
));
}
#[test]
fn marker_keyed_recognizable_chunk_with_defective_payload_is_corrupt() {
let event =
super::classify_marker_keyed_frame::<TestChunk>(r#"{"choices":42}"#, &["choices"]);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn marker_keyed_invalid_json_is_corrupt() {
let event = super::classify_marker_keyed_frame::<TestChunk>("{not json", &["choices"]);
assert!(matches!(event, WireEvent::Corrupt(_)));
}
#[test]
fn typed_event_triage_maps_onto_the_shared_policy() {
let event = super::classify_typed_event(super::TypedEvent::Modeled(7u8));
assert!(matches!(event, WireEvent::Known(7)));
let event = super::classify_typed_event::<u8>(super::TypedEvent::Unrecognized {
event_type: "unknown".to_string(),
detail: "FutureEvent".to_string(),
});
assert!(matches!(
event,
WireEvent::Unknown { event_type, value }
if event_type == "unknown" && value.value() == &serde_json::Value::String("FutureEvent".into())
));
let event =
super::classify_typed_event::<u8>(super::TypedEvent::Malformed("bad frame".into()));
assert!(
matches!(event, WireEvent::Corrupt(error) if error.to_string().contains("bad frame"))
);
}
#[test]
fn untyped_line_is_known_or_corrupt() {
assert!(matches!(
super::classify_untyped_line::<TestChunk>(br#"{"choices":[]}"#),
WireEvent::Known(_)
));
assert!(matches!(
super::classify_untyped_line::<TestChunk>(b"{not json"),
WireEvent::Corrupt(_)
));
}
}