use serde::de::DeserializeOwned;
use schemars::JsonSchema;
use super::super::super::{StreamItem, TextContent};
use super::types::{TokenExtractor, TokenExtractionResult};
use super::super::super::json_utils::find_json_structures;
use tokio::io::{AsyncBufReadExt, BufReader};
use async_stream::stream;
use futures_core::stream::Stream;
use futures_util::StreamExt;
use bytes::Bytes;
use std::pin::Pin;
pub struct SSEStreamProcessor;
impl SSEStreamProcessor {
pub fn process_sse_stream<T, E>(
byte_stream: Pin<Box<dyn Stream<Item = Result<Bytes, super::super::super::error::AIError>> + Send>>,
extractor: E,
) -> Pin<Box<dyn Stream<Item = Result<StreamItem<T>, super::super::super::error::QueryResolverError>> + Send>>
where
T: DeserializeOwned + JsonSchema + Send + 'static,
E: TokenExtractor + 'static,
{
Box::pin(stream! {
use tokio_util::io::StreamReader;
let io_stream = byte_stream.map(|res| match res {
Ok(bytes) => Ok::<Bytes, std::io::Error>(bytes),
Err(e) => Err(std::io::Error::other(e.to_string())),
});
let reader = StreamReader::new(io_stream);
let mut br = BufReader::new(reader).lines();
let mut sse_event = String::new();
let mut text_buf = String::new();
while let Ok(Some(line)) = br.next_line().await {
if line.is_empty() {
if let Some(payload) = sse_event.strip_prefix("data: ") {
if extractor.is_end_payload(payload) {
let tail = text_buf.trim();
if !tail.is_empty() {
yield Ok(StreamItem::Text(TextContent { text: tail.to_string() }));
}
break;
}
if let Ok(v) = serde_json::from_str::<serde_json::Value>(payload) {
match extractor.extract_token(&v) {
TokenExtractionResult::Token(token) => {
yield Ok(StreamItem::Token(token.clone()));
text_buf.push_str(&token);
let coords = find_json_structures(&text_buf);
let mut consumed_up_to = 0usize;
for node in coords {
let end = node.end.saturating_add(1);
let slice = &text_buf[node.start..end];
if let Ok(item) = serde_json::from_str::<T>(slice) {
if node.start > 0 {
let chunk = text_buf[..node.start].trim();
if !chunk.is_empty() {
yield Ok(StreamItem::Text(TextContent { text: chunk.to_string() }));
}
}
yield Ok(StreamItem::Data(item));
consumed_up_to = consumed_up_to.max(end);
}
}
if consumed_up_to > 0 {
text_buf.drain(..consumed_up_to);
}
if let Some(idx) = text_buf.find("\n\n") {
let (chunk, rest) = text_buf.split_at(idx);
let chunk = chunk.trim();
if !chunk.is_empty() {
yield Ok(StreamItem::Text(TextContent { text: chunk.to_string() }));
}
text_buf = rest[2..].to_string();
}
if let TokenExtractionResult::EndStream = extractor.extract_token(&v) {
let tail = text_buf.trim();
if !tail.is_empty() {
yield Ok(StreamItem::Text(TextContent { text: tail.to_string() }));
}
text_buf.clear();
}
}
TokenExtractionResult::EndStream => {
let tail = text_buf.trim();
if !tail.is_empty() {
yield Ok(StreamItem::Text(TextContent { text: tail.to_string() }));
}
break;
}
TokenExtractionResult::NoToken => {
}
}
}
}
sse_event.clear();
} else {
if line.starts_with("data: ") {
sse_event = line.to_string();
}
}
}
let tail = text_buf.trim();
if !tail.is_empty() {
yield Ok(StreamItem::Text(TextContent { text: tail.to_string() }));
}
})
}
}