1use futures::{SinkExt, StreamExt, channel::mpsc};
16
17use crate::{
18 error::{ErrorReport, ProviderError},
19 message::{AssistantContent, CallId, LocalCallId, Origin, ToolCall, ToolFunction, ToolName},
20 operation::{Block, Finish, Turn},
21 streaming::{Item, Relayed, StreamEvent},
22 wire::{Fold, Reply as WireReply},
23};
24
25use super::{Reply, SinkClosed};
26use crate::wasm_compat::WasmCompatSend;
27
28pub struct StreamWriter {
33 events: mpsc::Sender<Result<Relayed, ErrorReport>>,
34 origin: Origin,
35 turn: Turn,
36 items: std::collections::VecDeque<Result<Item<StreamEvent>, ProviderError>>,
37 raw: serde_json::Value,
38 request_id: Option<String>,
39}
40
41impl Reply {
42 pub fn written<F, Fut>(origin: Origin, write: F) -> Self
49 where
50 F: FnOnce(StreamWriter) -> Fut,
51 Fut: Future<Output = ()> + WasmCompatSend + 'static,
52 {
53 let (mut events, mut receiver) = mpsc::channel(0);
54 let _ = events.try_send(Ok(Relayed::Origin(origin.clone())));
57 let writer = StreamWriter {
58 events,
59 turn: Turn::new(origin.clone()),
60 origin,
61 items: std::collections::VecDeque::new(),
62 raw: serde_json::Value::Null,
63 request_id: None,
64 };
65 let mut writing = Some(Box::pin(write(writer)));
66 Self::Stream(Box::pin(futures::stream::poll_fn(move |cx| {
67 if let Some(future) = &mut writing
68 && future.as_mut().poll(cx).is_ready()
69 {
70 writing = None;
71 }
72 match receiver.poll_next_unpin(cx) {
73 std::task::Poll::Ready(None) if writing.is_some() => std::task::Poll::Pending,
74 next => next,
75 }
76 })))
77 }
78}
79
80impl StreamWriter {
81 pub async fn text(&mut self, text: impl Into<String>) -> Result<(), SinkClosed> {
84 self.write(Block::Text, &text.into()).await
85 }
86
87 pub async fn reasoning(&mut self, text: impl Into<String>) -> Result<(), SinkClosed> {
90 self.write(Block::Reasoning { redacted: false }, &text.into())
91 .await
92 }
93
94 async fn write(&mut self, block: Block, fragment: &str) -> Result<(), SinkClosed> {
95 if let Err(error) = self.turn.run_item(&mut self.items, block, fragment) {
96 return self.error(ErrorReport::from(&error)).await;
97 }
98 self.flush().await
99 }
100
101 pub async fn tool_call(
103 &mut self,
104 name: impl Into<String>,
105 arguments: serde_json::Value,
106 ) -> Result<(), SinkClosed> {
107 let Ok(name) = ToolName::new(name) else {
108 return self
109 .error(ErrorReport::from(&ProviderError::Response(
110 "a tool call needs a name".to_owned(),
111 )))
112 .await;
113 };
114 let call = ToolCall::new(
115 CallId::Local(LocalCallId::new()),
116 ToolFunction::new(name, arguments),
117 );
118 let written = self.turn.end_run(&mut self.items).and_then(|()| {
119 self.turn
120 .write_content(&mut self.items, AssistantContent::ToolCall(call))
121 });
122 if let Err(error) = written {
123 return self.error(ErrorReport::from(&error)).await;
124 }
125 self.flush().await
126 }
127
128 pub fn raw(&mut self, raw: serde_json::Value) {
130 self.raw = raw;
131 }
132
133 pub fn request_id(&mut self, request_id: impl Into<String>) {
137 self.request_id = Some(request_id.into());
138 }
139
140 pub async fn error(&mut self, report: ErrorReport) -> Result<(), SinkClosed> {
142 self.flush().await.map_err(|_| SinkClosed)?;
143 self.events.send(Err(report)).await.map_err(|_| SinkClosed)
144 }
145
146 pub async fn finish(mut self, finish: Finish) -> Result<(), SinkClosed> {
150 self.turn.close_open(&mut self.items);
151 self.flush().await?;
152 let reply = WireReply {
153 provider: self.origin.provider.clone(),
154 raw: std::mem::take(&mut self.raw),
155 provider_request_id: self.request_id.take(),
156 };
157 let turn = std::mem::replace(&mut self.turn, Turn::new(self.origin.clone()));
158 let item = match turn.finish(finish, reply) {
159 Ok(response) => Ok(Relayed::Done(Box::new(response))),
160 Err(error) => Err(ErrorReport::from(&error)),
161 };
162 self.events.send(item).await.map_err(|_| SinkClosed)
163 }
164
165 pub fn is_closed(&self) -> bool {
167 self.events.is_closed()
168 }
169
170 async fn flush(&mut self) -> Result<(), SinkClosed> {
172 while let Some(item) = self.items.pop_front() {
173 let item = match item {
174 Ok(item) => {
175 if let Item::Event(event) = &item {
176 let _ = self.turn.absorb(event);
177 }
178 Ok(Relayed::Item(item))
179 }
180 Err(error) => Err(ErrorReport::from(&error)),
181 };
182 self.events.send(item).await.map_err(|_| SinkClosed)?;
183 }
184 Ok(())
185 }
186}