use std::borrow::Cow;
use futures::{Stream, StreamExt};
use super::wire::WireEvent;
use crate::completion::CompletionError;
use crate::streaming::{RawStreamingChoice, RawStreamingResult};
use crate::wasm_compat::WasmCompatSend;
#[derive(Debug, Clone)]
pub enum WireFrame {
Text(String),
Bytes(Vec<u8>),
}
impl WireFrame {
pub fn as_str(&self) -> Cow<'_, str> {
match self {
Self::Text(text) => Cow::Borrowed(text),
Self::Bytes(bytes) => String::from_utf8_lossy(bytes),
}
}
}
pub type AdapterOutput<R> = Vec<Result<RawStreamingChoice<R>, CompletionError>>;
pub trait WireAdapter {
type Frame;
type Event;
type Response;
fn classify(&self, frame: Self::Frame) -> WireEvent<Self::Event>;
fn interpret(&mut self, event: Self::Event, out: &mut AdapterOutput<Self::Response>);
fn finish(&mut self, out: &mut AdapterOutput<Self::Response>);
fn flush_before_terminal_error(&mut self, _out: &mut AdapterOutput<Self::Response>) {}
fn is_finished(&self) -> bool {
false
}
}
#[derive(Debug)]
pub enum TriagedFrame<T> {
Event(T),
Unknown(crate::streaming::UnknownPayload),
}
pub fn triage_frame<T>(event: WireEvent<T>) -> Result<TriagedFrame<T>, CompletionError> {
match event {
WireEvent::Known(event) => Ok(TriagedFrame::Event(event)),
WireEvent::Unknown { event_type, value } => {
warn_unmodeled(&event_type, &value);
Ok(TriagedFrame::Unknown(value))
}
WireEvent::Corrupt(error) => Err(CompletionError::JsonError(error)),
}
}
pub fn warn_unmodeled(kind: &str, payload: &impl serde::Serialize) {
tracing::warn!(
kind,
payload_bytes = unknown_payload_bytes(payload),
"skipping unmodeled wire payload"
);
}
fn unknown_payload_bytes(value: &impl serde::Serialize) -> u64 {
struct CountingWriter(u64);
impl std::io::Write for CountingWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0 += buf.len() as u64;
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
let mut counter = CountingWriter(0);
let _ = serde_json::to_writer(&mut counter, value);
counter.0
}
pub fn run_wire_stream<A, S>(transport: S, mut adapter: A) -> RawStreamingResult<A::Response>
where
A: WireAdapter + WasmCompatSend + 'static,
A::Frame: WasmCompatSend,
A::Event: WasmCompatSend,
A::Response: WasmCompatSend + 'static,
S: Stream<Item = Result<A::Frame, CompletionError>> + WasmCompatSend + 'static,
{
Box::pin(async_stream::stream! {
let mut transport = Box::pin(transport);
let mut out: AdapterOutput<A::Response> = Vec::new();
#[cfg(any(test, debug_assertions))]
let mut sequence_laws = super::sequence_law::SequenceLaws::default();
while let Some(frame) = transport.next().await {
let frame = match frame {
Ok(frame) => frame,
Err(error) => {
adapter.flush_before_terminal_error(&mut out);
for item in out.drain(..) {
yield item;
}
yield Err(error);
return;
}
};
match triage_frame(adapter.classify(frame)) {
Ok(TriagedFrame::Event(event)) => adapter.interpret(event, &mut out),
Ok(TriagedFrame::Unknown(value)) => {
out.push(Ok(RawStreamingChoice::Unknown(value)));
}
Err(error) => {
yield Err(error);
}
}
#[cfg(any(test, debug_assertions))]
sequence_laws.check_batch(&out);
let saw_terminal = out
.iter()
.any(|item| matches!(item, Ok(RawStreamingChoice::FinalResponse(_))));
for item in out.drain(..) {
yield item;
}
if saw_terminal || adapter.is_finished() {
return;
}
}
adapter.finish(&mut out);
#[cfg(any(test, debug_assertions))]
sequence_laws.check_batch(&out);
for item in out.drain(..) {
yield item;
}
})
}
pub fn run_wire_buffered<A>(
frames: impl IntoIterator<Item = A::Frame>,
mut adapter: A,
) -> Result<Vec<RawStreamingChoice<A::Response>>, CompletionError>
where
A: WireAdapter,
{
let mut out: AdapterOutput<A::Response> = Vec::new();
let mut choices = Vec::new();
#[cfg(any(test, debug_assertions))]
let mut sequence_laws = super::sequence_law::SequenceLaws::default();
for frame in frames {
match adapter.classify(frame) {
WireEvent::Known(event) => adapter.interpret(event, &mut out),
WireEvent::Unknown { event_type, value } => {
tracing::warn!(
event_type,
payload_bytes = unknown_payload_bytes(&value),
"skipping unrecognized stream event"
);
}
WireEvent::Corrupt(error) => {
return Err(CompletionError::ResponseError(error.to_string()));
}
}
#[cfg(any(test, debug_assertions))]
sequence_laws.check_batch(&out);
let saw_terminal = drain_buffered(&mut out, &mut choices)?;
if saw_terminal || adapter.is_finished() {
return Ok(choices);
}
}
adapter.finish(&mut out);
#[cfg(any(test, debug_assertions))]
sequence_laws.check_batch(&out);
drain_buffered(&mut out, &mut choices)?;
Ok(choices)
}
fn drain_buffered<R>(
out: &mut AdapterOutput<R>,
choices: &mut Vec<RawStreamingChoice<R>>,
) -> Result<bool, CompletionError> {
let mut saw_terminal = false;
for item in out.drain(..) {
let choice = item?;
saw_terminal |= matches!(choice, RawStreamingChoice::FinalResponse(_));
choices.push(choice);
}
Ok(saw_terminal)
}
pub use crate::streaming::SyntheticIds;