use futures::{SinkExt, StreamExt, channel::mpsc};
use crate::{
error::{ErrorReport, ProviderError},
message::{AssistantContent, CallId, LocalCallId, Origin, ToolCall, ToolFunction, ToolName},
operation::{Block, Finish, Turn},
streaming::{Item, Relayed, StreamEvent},
wire::{Fold, Reply as WireReply},
};
use super::{Reply, SinkClosed};
use crate::wasm_compat::WasmCompatSend;
pub struct StreamWriter {
events: mpsc::Sender<Result<Relayed, ErrorReport>>,
origin: Origin,
turn: Turn,
items: std::collections::VecDeque<Result<Item<StreamEvent>, ProviderError>>,
raw: serde_json::Value,
request_id: Option<String>,
}
impl Reply {
pub fn written<F, Fut>(origin: Origin, write: F) -> Self
where
F: FnOnce(StreamWriter) -> Fut,
Fut: Future<Output = ()> + WasmCompatSend + 'static,
{
let (mut events, mut receiver) = mpsc::channel(0);
let _ = events.try_send(Ok(Relayed::Origin(origin.clone())));
let writer = StreamWriter {
events,
turn: Turn::new(origin.clone()),
origin,
items: std::collections::VecDeque::new(),
raw: serde_json::Value::Null,
request_id: None,
};
let mut writing = Some(Box::pin(write(writer)));
Self::Stream(Box::pin(futures::stream::poll_fn(move |cx| {
if let Some(future) = &mut writing
&& future.as_mut().poll(cx).is_ready()
{
writing = None;
}
match receiver.poll_next_unpin(cx) {
std::task::Poll::Ready(None) if writing.is_some() => std::task::Poll::Pending,
next => next,
}
})))
}
}
impl StreamWriter {
pub async fn text(&mut self, text: impl Into<String>) -> Result<(), SinkClosed> {
self.write(Block::Text, &text.into()).await
}
pub async fn reasoning(&mut self, text: impl Into<String>) -> Result<(), SinkClosed> {
self.write(Block::Reasoning { redacted: false }, &text.into())
.await
}
async fn write(&mut self, block: Block, fragment: &str) -> Result<(), SinkClosed> {
if let Err(error) = self.turn.run_item(&mut self.items, block, fragment) {
return self.error(ErrorReport::from(&error)).await;
}
self.flush().await
}
pub async fn tool_call(
&mut self,
name: impl Into<String>,
arguments: serde_json::Value,
) -> Result<(), SinkClosed> {
let Ok(name) = ToolName::new(name) else {
return self
.error(ErrorReport::from(&ProviderError::Response(
"a tool call needs a name".to_owned(),
)))
.await;
};
let call = ToolCall::new(
CallId::Local(LocalCallId::new()),
ToolFunction::new(name, arguments),
);
let written = self.turn.end_run(&mut self.items).and_then(|()| {
self.turn
.write_content(&mut self.items, AssistantContent::ToolCall(call))
});
if let Err(error) = written {
return self.error(ErrorReport::from(&error)).await;
}
self.flush().await
}
pub fn raw(&mut self, raw: serde_json::Value) {
self.raw = raw;
}
pub fn request_id(&mut self, request_id: impl Into<String>) {
self.request_id = Some(request_id.into());
}
pub async fn error(&mut self, report: ErrorReport) -> Result<(), SinkClosed> {
self.flush().await.map_err(|_| SinkClosed)?;
self.events.send(Err(report)).await.map_err(|_| SinkClosed)
}
pub async fn finish(mut self, finish: Finish) -> Result<(), SinkClosed> {
self.turn.close_open(&mut self.items);
self.flush().await?;
let reply = WireReply {
provider: self.origin.provider.clone(),
raw: std::mem::take(&mut self.raw),
provider_request_id: self.request_id.take(),
};
let turn = std::mem::replace(&mut self.turn, Turn::new(self.origin.clone()));
let item = match turn.finish(finish, reply) {
Ok(response) => Ok(Relayed::Done(Box::new(response))),
Err(error) => Err(ErrorReport::from(&error)),
};
self.events.send(item).await.map_err(|_| SinkClosed)
}
pub fn is_closed(&self) -> bool {
self.events.is_closed()
}
async fn flush(&mut self) -> Result<(), SinkClosed> {
while let Some(item) = self.items.pop_front() {
let item = match item {
Ok(item) => {
if let Item::Event(event) = &item {
let _ = self.turn.absorb(event);
}
Ok(Relayed::Item(item))
}
Err(error) => Err(ErrorReport::from(&error)),
};
self.events.send(item).await.map_err(|_| SinkClosed)?;
}
Ok(())
}
}