Skip to main content

rig_core/streaming/
mod.rs

1//! Completion events, and [`Streamed`], the one stream every reply arrives
2//! as.
3//!
4//! ```
5//! use rig_core::streaming::{Item, StreamEvent};
6//!
7//! fn text(item: &Item<StreamEvent>) -> Option<&str> {
8//!     match item {
9//!         Item::Event(StreamEvent::Text { text, .. }) => Some(text),
10//!         _ => None,
11//!     }
12//! }
13//! # let _ = text;
14//! ```
15
16mod 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/// Unmodeled JSON payload with content-redacted `Debug` output.
35/// Serialization preserves the payload; [`Self::value`] explicitly exposes it.
36#[derive(Clone, PartialEq, Serialize, Deserialize)]
37#[serde(transparent)]
38pub struct UnknownPayload(serde_json::Value);
39
40impl UnknownPayload {
41    /// Wrap a raw unmodeled payload.
42    pub fn new(value: serde_json::Value) -> Self {
43        Self(value)
44    }
45
46    /// The raw payload, for consumers who opt in to the content.
47    pub fn value(&self) -> &serde_json::Value {
48        &self.0
49    }
50}
51
52impl std::fmt::Debug for UnknownPayload {
53    /// Reports serialized size without exposing payload content.
54    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/// One item of a completion stream relayed over the bus: who the reply is
70/// from, an item of the reply, or the response the origin folded when the
71/// provider ended it.
72#[derive(Debug, Clone, PartialEq)]
73pub enum Relayed {
74    /// The wire, provider and model the reply is from, sent before its
75    /// first item, so a reply cut short still knows its origin.
76    Origin(crate::message::Origin),
77    /// An event, or a payload the origin's decoder did not model.
78    Item(Item<StreamEvent>),
79    /// The provider ended the reply; the origin's response.
80    Done(Box<CompletionResponse>),
81}
82
83/// A completion stream as the bus carries it.
84pub type StreamEvents = WasmBoxedStream<'static, Result<Relayed, ErrorReport>>;
85
86/// One reply as it arrives: its events, and the fold that has seen each of
87/// them. It is what [`Model::stream`](crate::Model::stream) returns for
88/// every operation, and what a call finishes.
89///
90/// An error is the last item: the stream ends after it. Stop polling to
91/// pause; drop the stream to cancel.
92pub struct Streamed<Op: Operation> {
93    /// The reading of the reply; `None` once it ended.
94    reading: Option<WasmBoxedStream<'static, ()>>,
95    shared: Arc<Mutex<Shared<Op>>>,
96    span: tracing::Span,
97    provider: String,
98    /// The error the stream yielded, which ends it.
99    failed: Option<ProviderError>,
100}
101
102/// A streamed completion.
103pub type CompletionStream = Streamed<Completion>;
104
105impl<Op: Operation> Streamed<Op> {
106    /// The reply `reading` writes into `shared`, under `span`.
107    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    /// The one place an item is taken: an event is absorbed by the fold
123    /// before it leaves, and an error, stamped with the request id, ends the
124    /// stream.
125    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    /// Read the rest of the reply, then fold it with the provider's end into
157    /// the response: what [`Model::call`](crate::Model::call) returns for
158    /// the same reply. An error the stream yielded, now or before, is the
159    /// result; a reply the provider did not end is
160    /// [`ProviderError::Truncated`].
161    pub async fn finish(self) -> Result<Op::Response, ProviderError> {
162        self.finish_routed().await.map_err(|(error, _)| error)
163    }
164
165    /// [`Self::finish`], with a failure paired with the request path of the
166    /// reply that failed.
167    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
192/// The assistant content a stream's `items` delivered, as its partial reply
193/// holds it: every part that ended, in start order, and the text of a text
194/// or reasoning part still open. The items hold no provider end, so no
195/// block keeps its provider item and the content replays canonically.
196pub 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    /// A stream relayed over the bus under `label`: its events, then the
216    /// response the origin folded. A relay that ends without one was cut
217    /// short.
218    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                // The lock is released before the stream yields.
225                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    /// This stream as the bus carries it: its items, then the response it
259    /// folds into, or the error that ended it. A reply cut short closes the
260    /// relay without a response, as its transport closed.
261    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    /// What arrived so far: every part that ended, and the provider's end
284    /// once it arrived. Valid after an error, and after the caller stopped
285    /// polling.
286    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;