mod event;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use futures::{Stream, StreamExt};
use serde::{Deserialize, Serialize};
use crate::completion::CompletionResponse;
use crate::driver::{lock, record_request_id};
use crate::error::{ErrorReport, ProviderError};
pub use crate::json_utils::parse_partial_arguments;
use crate::operation::{Completion, Turn};
use crate::wasm_compat::WasmBoxedStream;
use crate::wire::{Operation, Shared};
pub use event::{Item, Part, PartKind, SequenceError, StreamEvent, Transcript};
#[derive(Clone, PartialEq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct UnknownPayload(serde_json::Value);
impl UnknownPayload {
pub fn new(value: serde_json::Value) -> Self {
Self(value)
}
pub fn value(&self) -> &serde_json::Value {
&self.0
}
}
impl std::fmt::Debug for UnknownPayload {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let bytes = serde_json::to_vec(&self.0).map_or(0, |json| json.len());
write!(f, "UnknownPayload({bytes} bytes redacted)")
}
}
impl From<serde_json::Value> for UnknownPayload {
fn from(value: serde_json::Value) -> Self {
Self(value)
}
}
#[cfg(test)]
mod unknown_payload_tests;
#[derive(Debug, Clone, PartialEq)]
pub enum Relayed {
Origin(crate::message::Origin),
Item(Item<StreamEvent>),
Done(Box<CompletionResponse>),
}
pub type StreamEvents = WasmBoxedStream<'static, Result<Relayed, ErrorReport>>;
pub struct Streamed<Op: Operation> {
reading: Option<WasmBoxedStream<'static, ()>>,
shared: Arc<Mutex<Shared<Op>>>,
span: tracing::Span,
provider: String,
failed: Option<ProviderError>,
}
pub type CompletionStream = Streamed<Completion>;
impl<Op: Operation> Streamed<Op> {
pub(crate) fn new(
reading: WasmBoxedStream<'static, ()>,
shared: Arc<Mutex<Shared<Op>>>,
span: tracing::Span,
provider: impl Into<String>,
) -> Self {
Self {
reading: Some(reading),
shared,
span,
provider: provider.into(),
failed: None,
}
}
fn poll_item(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Item<Op::Event>, ProviderError>>> {
loop {
if self.failed.is_some() {
return Poll::Ready(None);
}
{
let mut shared = lock(&self.shared);
if let Some(item) = shared.take() {
if let Err(error) = &item {
record_request_id(&self.span, error.provider_request_id());
shared.items.clear();
self.failed = Some(error.clone());
self.reading = None;
}
return Poll::Ready(Some(item));
}
}
let Some(reading) = &mut self.reading else {
return Poll::Ready(None);
};
match reading.as_mut().poll_next(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Some(())) => {}
Poll::Ready(None) => self.reading = None,
}
}
}
pub async fn finish(self) -> Result<Op::Response, ProviderError> {
self.finish_routed().await.map_err(|(error, _)| error)
}
pub(crate) async fn finish_routed(mut self) -> Result<Op::Response, (ProviderError, String)> {
let route = |stream: &Self| lock(&stream.shared).route.clone();
while let Some(item) = futures::future::poll_fn(|cx| self.poll_item(cx)).await {
if let Err(error) = item {
return Err((error, route(&self)));
}
}
if let Some(error) = self.failed.take() {
return Err((error, route(&self)));
}
let path = route(&self);
let Ok(shared) = Arc::try_unwrap(self.shared) else {
return Err((
ProviderError::Response("the reply is still being read".to_owned()),
path,
));
};
shared
.into_inner()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.conclude(&self.provider)
.map_err(|error| (error, path))
}
}
pub fn delivered(items: &[Item<StreamEvent>]) -> Vec<crate::message::AssistantContent> {
use crate::wire::Fold;
let mut turn = Turn::relayed("delivered");
for item in items {
if let Item::Event(event) = item
&& turn.absorb(event).is_err()
{
break;
}
}
let reply = crate::wire::Reply {
provider: String::new(),
raw: serde_json::Value::Null,
provider_request_id: None,
};
turn.partial(None, &reply, None).choice
}
impl Streamed<Completion> {
pub fn relay(label: impl Into<String>, mut events: StreamEvents) -> Self {
let label = label.into();
let shared = Arc::new(Mutex::new(Shared::new(Turn::relayed(label.clone()))));
let writer = Arc::clone(&shared);
let reading = async_stream::stream! {
while let Some(item) = events.next().await {
let ended = {
let mut shared = lock(&writer);
match item {
Ok(Relayed::Origin(origin)) => {
Turn::set_origin(&mut shared.fold, origin);
false
}
Ok(Relayed::Item(item)) => {
shared.items.push_back(Ok(item));
false
}
Ok(Relayed::Done(response)) => {
shared.response = Some(*response);
true
}
Err(report) => {
shared
.items
.push_back(Err(ProviderError::Relayed(Box::new(report))));
true
}
}
};
if ended {
return;
}
yield ();
}
lock(&writer).items.push_back(Err(ProviderError::Truncated));
};
Self::new(Box::pin(reading), shared, tracing::Span::none(), label)
}
pub fn into_relay(mut self) -> StreamEvents {
let origin = lock(&self.shared).fold.origin().clone();
Box::pin(async_stream::stream! {
yield Ok(Relayed::Origin(origin));
while let Some(item) = self.next().await {
match item {
Ok(item) => yield Ok(Relayed::Item(item)),
Err(ProviderError::Truncated) => return,
Err(error) => {
yield Err(ErrorReport::from(&error));
return;
}
}
}
match self.finish().await {
Ok(response) => yield Ok(Relayed::Done(Box::new(response))),
Err(ProviderError::Truncated) => {}
Err(error) => yield Err(ErrorReport::from(&error)),
}
})
}
pub fn partial(&self) -> CompletionResponse {
let shared = lock(&self.shared);
if let Some(response) = &shared.response {
return response.clone();
}
shared.fold.partial(
shared.end.as_ref(),
&shared.reply(&self.provider),
self.failed.as_ref(),
)
}
}
impl<Op: Operation> Stream for Streamed<Op> {
type Item = Result<Item<Op::Event>, ProviderError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.get_mut().poll_item(cx)
}
}
#[cfg(test)]
mod tests;