use serde_json::Value;
use crate::wire::WireEvent;
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,
{
match scan(data, &[tag], true) {
Err(error) => WireEvent::Corrupt(error),
Ok((value, None)) => unknown(value, String::new()),
Ok((value, Some(found))) => match found.first().and_then(|tag| tag.as_ref()?.as_str()) {
Some(event_type) if !is_known_event_type(event_type) => {
let event_type = event_type.to_owned();
unknown(value, event_type)
}
_ => decode_known(data),
},
}
}
fn unknown<T>(value: Value, event_type: String) -> WireEvent<T> {
WireEvent::Unknown {
event_type,
value: value.into(),
}
}
pub fn classify_chat_completions_frame<T>(data: &str) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
let (value, found) = match scan(data, &["object", "choices"], true) {
Ok((value, Some(found))) => (value, found),
Ok((value, None)) => return unknown(value, String::new()),
Err(error) => return WireEvent::Corrupt(error),
};
let object = found
.first()
.and_then(|object| object.as_ref()?.as_str())
.map(str::to_owned);
let has_choices = found.get(1).is_some_and(Option::is_some);
if object.as_deref() == Some("chat.completion.chunk") || has_choices {
decode_known(data)
} else {
unknown(value, object.unwrap_or_default())
}
}
pub fn classify_marker_keyed_frame<T>(data: &str, marker_keys: &[&str]) -> WireEvent<T>
where
T: serde::de::DeserializeOwned,
{
match scan(data, marker_keys, false) {
Err(error) => WireEvent::Corrupt(error),
Ok((_, Some(found))) if found.iter().any(Option::is_some) => decode_known(data),
Ok((value, _)) => {
let event_type = value
.as_object()
.map(|object| object.keys().cloned().collect::<Vec<_>>().join(","))
.unwrap_or_default();
unknown(value, event_type)
}
}
}
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),
}
}
pub fn classify_or<T>(
data: &str,
first: impl Fn(&str) -> WireEvent<T>,
then: impl Fn(&str) -> WireEvent<T>,
) -> WireEvent<T> {
match first(data) {
WireEvent::Corrupt(first_error) => match then(data) {
WireEvent::Known(event) => WireEvent::Known(event),
WireEvent::Corrupt(error) => WireEvent::Corrupt(error),
WireEvent::Unknown { .. } => WireEvent::Corrupt(first_error),
},
event => event,
}
}
#[derive(serde::Deserialize)]
struct MessageEnvelope {
#[allow(dead_code)]
message: String,
}
pub fn classify_reply_or_message_envelope<T>(
data: &str,
reply_marker: &str,
) -> WireEvent<Result<T, String>>
where
T: serde::de::DeserializeOwned,
{
classify_or(
data,
|data| classify_marker_keyed_frame::<T>(data, &[reply_marker, "message"]).map(Ok),
|data| {
classify_marker_keyed_frame::<MessageEnvelope>(data, &["message"])
.map(|_| Err(data.to_owned()))
},
)
}
pub fn classify_or_untagged<T>(
data: &str,
tag: &str,
first: impl Fn(&str) -> WireEvent<T>,
then: impl Fn(&str) -> WireEvent<T>,
) -> WireEvent<T> {
classify_or(data, first, |data| match scan(data, &[tag], false) {
Ok((value, Some(found))) if found.iter().any(Option::is_some) => {
unknown(value, tag.to_owned())
}
_ => then(data),
})
}
struct Entries(Vec<(String, Value)>);
impl<'de> serde::Deserialize<'de> for Entries {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct Visitor;
impl<'de> serde::de::Visitor<'de> for Visitor {
type Value = Entries;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a JSON object")
}
fn visit_map<A: serde::de::MapAccess<'de>>(
self,
mut map: A,
) -> Result<Entries, A::Error> {
let mut entries = Vec::new();
while let Some(entry) = map.next_entry()? {
entries.push(entry);
}
Ok(Entries(entries))
}
}
deserializer.deserialize_map(Visitor)
}
}
fn scan(
data: &str,
keys: &[&str],
unique: bool,
) -> Result<(Value, Option<Vec<Option<Value>>>), serde_json::Error> {
let value: Value = serde_json::from_str(data)?;
if !value.is_object() {
return Ok((value, None));
}
let Entries(entries) = serde_json::from_str(data)?;
let mut found = vec![None; keys.len()];
for (key, field) in entries {
match keys
.iter()
.position(|candidate| *candidate == key)
.and_then(|index| found.get_mut(index))
{
Some(slot @ None) => *slot = Some(field),
Some(Some(_)) if unique => {
return Err(serde::de::Error::custom(format!(
"duplicate `{key}` discriminator key in stream frame"
)));
}
Some(Some(_)) | None => {}
}
}
Ok((value, Some(found)))
}
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;