use std::collections::VecDeque;
use std::time::Duration;
use ferrin_spec::BoxStream;
use ferrin_spec::JsonValue;
use ferrin_spec::PartId;
use ferrin_spec::ToolCallId;
use ferrin_spec::error::ProviderError;
use ferrin_spec::language_model::StreamPart;
use futures_util::StreamExt;
use futures_util::stream;
use crate::http::ParseResult;
pub const ACCEPTED_GRACE: Duration = Duration::from_millis(50);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EarlyChunk {
Error,
Output,
Accepted,
Other,
}
pub async fn fail_on_early_error<T: Send + 'static>(
stream: BoxStream<'static, ParseResult<T>>,
classify: impl Fn(&T) -> EarlyChunk,
to_error: impl FnOnce(&T, &JsonValue) -> ProviderError,
) -> Result<BoxStream<'static, ParseResult<T>>, ProviderError> {
let mut stream = stream.fuse();
let mut buffered: Vec<ParseResult<T>> = Vec::new();
let mut accepted = false;
loop {
let next = if accepted {
match tokio::time::timeout(ACCEPTED_GRACE, stream.next()).await {
Ok(next) => next,
Err(_) => break,
}
} else {
stream.next().await
};
let Some(chunk) = next else {
break;
};
let ParseResult::Ok { value, raw } = &chunk else {
buffered.push(chunk);
break;
};
match classify(value) {
EarlyChunk::Error => return Err(to_error(value, raw)),
EarlyChunk::Output => {
buffered.push(chunk);
break;
}
EarlyChunk::Accepted => {
accepted = true;
buffered.push(chunk);
}
EarlyChunk::Other => buffered.push(chunk),
}
}
Ok(Box::pin(stream::iter(buffered).chain(stream)))
}
pub trait StreamMachine: Send + 'static {
type Chunk: Send + 'static;
fn handle(&mut self, chunk: ParseResult<Self::Chunk>, include_raw: bool) -> Vec<StreamPart>;
fn finish(self) -> Vec<StreamPart>;
}
#[derive(Debug, Default)]
struct OpenParts {
text: Vec<PartId>,
reasoning: Vec<PartId>,
tool_inputs: Vec<ToolCallId>,
}
impl OpenParts {
fn observe(&mut self, part: &StreamPart) {
match part {
StreamPart::TextStart { id, .. } => self.text.push(id.clone()),
StreamPart::TextEnd { id, .. } => self.text.retain(|open| open != id),
StreamPart::ReasoningStart { id, .. } => self.reasoning.push(id.clone()),
StreamPart::ReasoningEnd { id, .. } => self.reasoning.retain(|open| open != id),
StreamPart::ToolInputStart { id, .. } => self.tool_inputs.push(id.clone()),
StreamPart::ToolInputEnd { id, .. } => self.tool_inputs.retain(|open| open != id),
_ => {}
}
}
fn close(&mut self) -> Vec<StreamPart> {
let mut parts = Vec::new();
parts.extend(std::mem::take(&mut self.tool_inputs).into_iter().map(|id| {
StreamPart::ToolInputEnd {
id,
provider_metadata: None,
}
}));
parts.extend(std::mem::take(&mut self.reasoning).into_iter().map(|id| {
StreamPart::ReasoningEnd {
id,
provider_metadata: None,
}
}));
parts.extend(
std::mem::take(&mut self.text)
.into_iter()
.map(|id| StreamPart::TextEnd {
id,
provider_metadata: None,
}),
);
parts
}
}
#[must_use]
pub fn drive_stream<M: StreamMachine>(
start: StreamPart,
chunks: BoxStream<'static, ParseResult<M::Chunk>>,
machine: M,
include_raw: bool,
) -> BoxStream<'static, StreamPart> {
struct Driver<M: StreamMachine> {
chunks: BoxStream<'static, ParseResult<M::Chunk>>,
machine: Option<M>,
pending: VecDeque<StreamPart>,
open: OpenParts,
include_raw: bool,
}
impl<M: StreamMachine> Driver<M> {
fn enqueue(&mut self, parts: Vec<StreamPart>) {
for part in parts {
if matches!(part, StreamPart::Error { .. }) {
self.pending.extend(self.open.close());
self.pending.push_back(part);
self.machine = None;
return;
}
self.open.observe(&part);
self.pending.push_back(part);
}
}
}
let driver = Driver {
chunks,
machine: Some(machine),
pending: VecDeque::from([start]),
open: OpenParts::default(),
include_raw,
};
Box::pin(stream::unfold(driver, |mut driver| async move {
loop {
if let Some(part) = driver.pending.pop_front() {
return Some((part, driver));
}
driver.machine.as_ref()?;
match driver.chunks.next().await {
Some(chunk) => {
let include_raw = driver.include_raw;
let parts = driver.machine.as_mut()?.handle(chunk, include_raw);
driver.enqueue(parts);
}
None => {
let parts = driver.machine.take()?.finish();
driver.enqueue(parts);
}
}
}
}))
}