pub mod parsers;
pub mod json_utils;
pub mod error;
pub mod compat;
pub mod end_condition;
pub use parsers::{JsonStreamProcessor, process_complete_text};
use schemars::JsonSchema;
use serde::de::DeserializeOwned;
use serde::Deserialize;
use tracing::instrument;
use tokio::io::{AsyncRead, AsyncReadExt};
use async_stream::stream;
use futures_core::stream::Stream as FuturesStream;
use futures_util::StreamExt;
use bytes::Bytes;
use std::pin::Pin;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SSEFormat {
OpenAI,
Claude,
}
impl SSEFormat {
pub fn create_provider<T>(&self) -> Box<dyn SSEProvider<T>>
where
T: DeserializeOwned + JsonSchema + Send + 'static,
{
match self {
SSEFormat::OpenAI => Box::new(sse::openai_provider()),
SSEFormat::Claude => Box::new(sse::claude_provider()),
}
}
}
pub trait SSEProvider<T>: Send + Sync
where
T: DeserializeOwned + JsonSchema + Send + 'static,
{
fn parse_sse_stream(
&self,
byte_stream: Pin<Box<dyn FuturesStream<Item = Result<Bytes, error::AIError>> + Send>>
) -> Pin<Box<dyn FuturesStream<Item = Result<StreamItem<T>, error::QueryResolverError>> + Send>>;
}
#[derive(Debug, Clone, Deserialize, JsonSchema, PartialEq)]
pub struct TextContent {
pub text: String,
}
#[derive(Debug, Clone, Deserialize, JsonSchema)]
#[serde(tag = "kind", content = "content")]
pub enum StreamItem<T>
where
T: JsonSchema,
{
#[serde(skip)]
Token(String),
Text(TextContent),
Data(T),
}
pub type ParsedStream<T> = Vec<StreamItem<T>>;
#[instrument(skip(raw))]
pub fn build_parsed_stream<T>(raw: &str) -> ParsedStream<T>
where
T: DeserializeOwned + JsonSchema,
{
process_complete_text(raw)
}
pub fn stream_from_async_read<R, T>(mut reader: R, buf_size: usize) -> impl FuturesStream<Item = StreamItem<T>>
where
R: AsyncRead + Unpin + Send + 'static,
T: DeserializeOwned + JsonSchema + Send + 'static,
{
stream! {
let mut processor = JsonStreamProcessor::<T>::new();
let mut buf = vec![0u8; buf_size.max(1024)];
loop {
match reader.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if let Ok(s) = std::str::from_utf8(&buf[..n]) {
let items = processor.process_chunk(s);
for item in items {
yield item;
}
}
}
Err(_) => break,
}
}
let final_items = processor.finalize();
for item in final_items {
yield item;
}
}
}
pub fn stream_from_bytes<T>(
byte_stream: Pin<Box<dyn FuturesStream<Item = Result<Bytes, error::AIError>> + Send>>
) -> impl FuturesStream<Item = Result<StreamItem<T>, error::QueryResolverError>>
where
T: DeserializeOwned + JsonSchema + Send + 'static,
{
stream! {
let mut processor = JsonStreamProcessor::<T>::new();
let mut byte_stream = byte_stream;
while let Some(chunk_result) = byte_stream.next().await {
match chunk_result {
Ok(bytes) => {
match std::str::from_utf8(&bytes) {
Ok(s) => {
let items = processor.process_chunk(s);
for item in items {
yield Ok(item);
}
}
Err(utf8_err) => {
yield Err(error::QueryResolverError::Ai(
error::AIError::Mock(format!("UTF-8 decode error: {}", utf8_err))
));
break;
}
}
}
Err(ai_error) => {
yield Err(error::QueryResolverError::Ai(ai_error));
break;
}
}
}
let final_items = processor.finalize();
for item in final_items {
yield Ok(item);
}
}
}
pub mod sse;
pub use sse::{
ConfigurableSSEProvider, OpenAISSEProvider, ClaudeSSEProvider,
openai_provider, claude_provider, stream_from_sse_bytes
};
pub use compat::{StreamEvent, StreamProcessor, StreamEventExt};
pub type Stream = Pin<Box<dyn FuturesStream<Item = crate::error::Result<StreamEvent>> + Send>>;