use async_trait::async_trait;
use futures_util::Stream;
use futures_util::StreamExt;
use serde::de::DeserializeOwned;
use serde::Serialize;
use serde_json::Value;
use std::pin::Pin;
use crate::core::language_models::BaseChatModel;
use crate::schema::Message;
use super::extract::{build_structured_system_prompt, StructuredOutputError};
use super::parser::{PartialJsonError, PartialJsonParser};
#[async_trait]
pub trait StreamingStructuredOutputExt: BaseChatModel {
async fn stream_structured_output<T>(
&self,
schema: Value,
prompt: &str,
) -> Result<
Pin<Box<dyn Stream<Item = Result<T, StructuredOutputError>> + Send>>,
StructuredOutputError,
>
where
T: DeserializeOwned + Serialize + Clone + PartialEq + Unpin + Send + Sync + 'static,
{
stream_structured_output(self, schema, prompt).await
}
}
impl<M: BaseChatModel> StreamingStructuredOutputExt for M {}
pub async fn stream_structured_output<T, M>(
llm: &M,
schema: Value,
prompt: &str,
) -> Result<
Pin<Box<dyn Stream<Item = Result<T, StructuredOutputError>> + Send>>,
StructuredOutputError,
>
where
T: DeserializeOwned + Serialize + Clone + PartialEq + Unpin + Send + Sync + 'static,
M: BaseChatModel + ?Sized,
{
if !schema.is_object() {
return Err(StructuredOutputError::SchemaError(format!(
"Schema must be a JSON object, got: {}",
schema
)));
}
let system_prompt = build_structured_system_prompt(&schema);
let messages = vec![Message::system(system_prompt), Message::human(prompt)];
let token_stream = llm
.stream_chat(messages, None)
.await
.map_err(|e| StructuredOutputError::LLMError(e.to_string()))?;
let mapped_stream =
token_stream.map(|item| item.map_err(|e| StructuredOutputError::LLMError(e.to_string())));
let boxed: Pin<Box<dyn Stream<Item = Result<String, StructuredOutputError>> + Send>> =
Box::pin(mapped_stream);
let output_stream = StructuredStreamProcessor::<T>::new(boxed);
Ok(Box::pin(output_stream))
}
struct StructuredStreamProcessor<T> {
inner: Pin<Box<dyn Stream<Item = Result<String, StructuredOutputError>> + Send>>,
parser: PartialJsonParser,
last_value: Option<T>,
done: bool,
}
impl<T: Unpin> Unpin for StructuredStreamProcessor<T> {}
impl<T> StructuredStreamProcessor<T>
where
T: DeserializeOwned + Serialize + Clone + PartialEq + Unpin + Send + Sync + 'static,
{
fn new(
inner: Pin<Box<dyn Stream<Item = Result<String, StructuredOutputError>> + Send>>,
) -> Self {
Self {
inner,
parser: PartialJsonParser::new(),
last_value: None,
done: false,
}
}
}
impl<T> Stream for StructuredStreamProcessor<T>
where
T: DeserializeOwned + Serialize + Clone + PartialEq + Send + Sync + Unpin + 'static,
{
type Item = Result<T, StructuredOutputError>;
fn poll_next(
self: Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.done {
return std::task::Poll::Ready(None);
}
loop {
match this.inner.as_mut().poll_next(cx) {
std::task::Poll::Ready(Some(Ok(token))) => {
match this.parser.push_and_parse(&token) {
Ok(json_value) => match serde_json::from_value::<T>(json_value) {
Ok(value) => {
this.last_value = Some(value.clone());
return std::task::Poll::Ready(Some(Ok(value)));
}
Err(_) => {
continue;
}
},
Err(PartialJsonError::Incomplete(_)) => {
continue;
}
Err(PartialJsonError::Invalid(_msg)) => {
continue;
}
}
}
std::task::Poll::Ready(Some(Err(e))) => {
return std::task::Poll::Ready(Some(Err(e)));
}
std::task::Poll::Ready(None) => {
this.done = true;
let parser = std::mem::take(&mut this.parser);
match parser.finalize() {
Ok(json_value) => match serde_json::from_value::<T>(json_value) {
Ok(value) => {
let is_new = this.last_value.as_ref() != Some(&value);
if is_new {
return std::task::Poll::Ready(Some(Ok(value)));
}
return std::task::Poll::Ready(None);
}
Err(e) => {
return std::task::Poll::Ready(Some(Err(
StructuredOutputError::ParseError(format!(
"Failed to deserialize final JSON into target type: {}",
e
)),
)));
}
},
Err(PartialJsonError::Invalid(msg)) => {
if this.last_value.is_some() {
return std::task::Poll::Ready(None);
}
return std::task::Poll::Ready(Some(Err(
StructuredOutputError::StreamIncomplete(msg),
)));
}
Err(PartialJsonError::Incomplete(msg)) => {
if this.last_value.is_some() {
return std::task::Poll::Ready(None);
}
return std::task::Poll::Ready(Some(Err(
StructuredOutputError::StreamIncomplete(msg),
)));
}
}
}
std::task::Poll::Pending => {
return std::task::Poll::Pending;
}
}
}
}
}
#[allow(dead_code)]
pub(crate) fn try_deserialize_partial<T: DeserializeOwned>(value: Value) -> Option<T> {
serde_json::from_value::<T>(value).ok()
}