rig_core/streaming/
mod.rs1mod event;
17
18use std::pin::Pin;
19use std::sync::{Arc, Mutex};
20use std::task::{Context, Poll};
21
22use futures::{Stream, StreamExt};
23use serde::{Deserialize, Serialize};
24
25use crate::completion::CompletionResponse;
26use crate::driver::{lock, record_request_id};
27use crate::error::{ErrorReport, ProviderError};
28pub use crate::json_utils::parse_partial_arguments;
29use crate::operation::{Completion, Turn};
30use crate::wasm_compat::WasmBoxedStream;
31use crate::wire::{Operation, Shared};
32pub use event::{Item, Part, PartKind, SequenceError, StreamEvent, Transcript};
33
34#[derive(Clone, PartialEq, Serialize, Deserialize)]
37#[serde(transparent)]
38pub struct UnknownPayload(serde_json::Value);
39
40impl UnknownPayload {
41 pub fn new(value: serde_json::Value) -> Self {
43 Self(value)
44 }
45
46 pub fn value(&self) -> &serde_json::Value {
48 &self.0
49 }
50}
51
52impl std::fmt::Debug for UnknownPayload {
53 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
55 let bytes = serde_json::to_vec(&self.0).map_or(0, |json| json.len());
56 write!(f, "UnknownPayload({bytes} bytes redacted)")
57 }
58}
59
60impl From<serde_json::Value> for UnknownPayload {
61 fn from(value: serde_json::Value) -> Self {
62 Self(value)
63 }
64}
65
66#[cfg(test)]
67mod unknown_payload_tests;
68
69#[derive(Debug, Clone, PartialEq)]
73pub enum Relayed {
74 Origin(crate::message::Origin),
77 Item(Item<StreamEvent>),
79 Done(Box<CompletionResponse>),
81}
82
83pub type StreamEvents = WasmBoxedStream<'static, Result<Relayed, ErrorReport>>;
85
86pub struct Streamed<Op: Operation> {
93 reading: Option<WasmBoxedStream<'static, ()>>,
95 shared: Arc<Mutex<Shared<Op>>>,
96 span: tracing::Span,
97 provider: String,
98 failed: Option<ProviderError>,
100}
101
102pub type CompletionStream = Streamed<Completion>;
104
105impl<Op: Operation> Streamed<Op> {
106 pub(crate) fn new(
108 reading: WasmBoxedStream<'static, ()>,
109 shared: Arc<Mutex<Shared<Op>>>,
110 span: tracing::Span,
111 provider: impl Into<String>,
112 ) -> Self {
113 Self {
114 reading: Some(reading),
115 shared,
116 span,
117 provider: provider.into(),
118 failed: None,
119 }
120 }
121
122 fn poll_item(
126 &mut self,
127 cx: &mut Context<'_>,
128 ) -> Poll<Option<Result<Item<Op::Event>, ProviderError>>> {
129 loop {
130 if self.failed.is_some() {
131 return Poll::Ready(None);
132 }
133 {
134 let mut shared = lock(&self.shared);
135 if let Some(item) = shared.take() {
136 if let Err(error) = &item {
137 record_request_id(&self.span, error.provider_request_id());
138 shared.items.clear();
139 self.failed = Some(error.clone());
140 self.reading = None;
141 }
142 return Poll::Ready(Some(item));
143 }
144 }
145 let Some(reading) = &mut self.reading else {
146 return Poll::Ready(None);
147 };
148 match reading.as_mut().poll_next(cx) {
149 Poll::Pending => return Poll::Pending,
150 Poll::Ready(Some(())) => {}
151 Poll::Ready(None) => self.reading = None,
152 }
153 }
154 }
155
156 pub async fn finish(self) -> Result<Op::Response, ProviderError> {
162 self.finish_routed().await.map_err(|(error, _)| error)
163 }
164
165 pub(crate) async fn finish_routed(mut self) -> Result<Op::Response, (ProviderError, String)> {
168 let route = |stream: &Self| lock(&stream.shared).route.clone();
169 while let Some(item) = futures::future::poll_fn(|cx| self.poll_item(cx)).await {
170 if let Err(error) = item {
171 return Err((error, route(&self)));
172 }
173 }
174 if let Some(error) = self.failed.take() {
175 return Err((error, route(&self)));
176 }
177 let path = route(&self);
178 let Ok(shared) = Arc::try_unwrap(self.shared) else {
179 return Err((
180 ProviderError::Response("the reply is still being read".to_owned()),
181 path,
182 ));
183 };
184 shared
185 .into_inner()
186 .unwrap_or_else(std::sync::PoisonError::into_inner)
187 .conclude(&self.provider)
188 .map_err(|error| (error, path))
189 }
190}
191
192pub fn delivered(items: &[Item<StreamEvent>]) -> Vec<crate::message::AssistantContent> {
197 use crate::wire::Fold;
198 let mut turn = Turn::relayed("delivered");
199 for item in items {
200 if let Item::Event(event) = item
201 && turn.absorb(event).is_err()
202 {
203 break;
204 }
205 }
206 let reply = crate::wire::Reply {
207 provider: String::new(),
208 raw: serde_json::Value::Null,
209 provider_request_id: None,
210 };
211 turn.partial(None, &reply, None).choice
212}
213
214impl Streamed<Completion> {
215 pub fn relay(label: impl Into<String>, mut events: StreamEvents) -> Self {
219 let label = label.into();
220 let shared = Arc::new(Mutex::new(Shared::new(Turn::relayed(label.clone()))));
221 let writer = Arc::clone(&shared);
222 let reading = async_stream::stream! {
223 while let Some(item) = events.next().await {
224 let ended = {
226 let mut shared = lock(&writer);
227 match item {
228 Ok(Relayed::Origin(origin)) => {
229 Turn::set_origin(&mut shared.fold, origin);
230 false
231 }
232 Ok(Relayed::Item(item)) => {
233 shared.items.push_back(Ok(item));
234 false
235 }
236 Ok(Relayed::Done(response)) => {
237 shared.response = Some(*response);
238 true
239 }
240 Err(report) => {
241 shared
242 .items
243 .push_back(Err(ProviderError::Relayed(Box::new(report))));
244 true
245 }
246 }
247 };
248 if ended {
249 return;
250 }
251 yield ();
252 }
253 lock(&writer).items.push_back(Err(ProviderError::Truncated));
254 };
255 Self::new(Box::pin(reading), shared, tracing::Span::none(), label)
256 }
257
258 pub fn into_relay(mut self) -> StreamEvents {
262 let origin = lock(&self.shared).fold.origin().clone();
263 Box::pin(async_stream::stream! {
264 yield Ok(Relayed::Origin(origin));
265 while let Some(item) = self.next().await {
266 match item {
267 Ok(item) => yield Ok(Relayed::Item(item)),
268 Err(ProviderError::Truncated) => return,
269 Err(error) => {
270 yield Err(ErrorReport::from(&error));
271 return;
272 }
273 }
274 }
275 match self.finish().await {
276 Ok(response) => yield Ok(Relayed::Done(Box::new(response))),
277 Err(ProviderError::Truncated) => {}
278 Err(error) => yield Err(ErrorReport::from(&error)),
279 }
280 })
281 }
282
283 pub fn partial(&self) -> CompletionResponse {
287 let shared = lock(&self.shared);
288 if let Some(response) = &shared.response {
289 return response.clone();
290 }
291 shared.fold.partial(
292 shared.end.as_ref(),
293 &shared.reply(&self.provider),
294 self.failed.as_ref(),
295 )
296 }
297}
298
299impl<Op: Operation> Stream for Streamed<Op> {
300 type Item = Result<Item<Op::Event>, ProviderError>;
301
302 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
303 self.get_mut().poll_item(cx)
304 }
305}
306
307#[cfg(test)]
308mod tests;