Skip to main content

rig_core/serve/
writer.rs

1//! Co-polled streaming replies written part by part, closed when the reply
2//! ends.
3//!
4//! ```
5//! use rig_core::message::Origin;
6//! use rig_core::serve::Reply;
7//!
8//! let origin = Origin::new("example.api", "example", "example-1");
9//! let reply = Reply::written(origin, |mut writer| async move {
10//!     let _ = writer.text("hello").await;
11//! });
12//! # let _ = reply;
13//! ```
14
15use 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
28/// A streaming answer under construction, the bus's completion writer.
29/// Obtained by [`Reply::written`]; [`finish`](Self::finish) ends the reply
30/// with the response it folds into, while a writer dropped without it
31/// leaves a truncated stream.
32pub 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    /// Return a stream that owns and polls the writing future alongside its
43    /// private receiver. No task is spawned. The bridge has zero shared
44    /// capacity and one sender-reserved slot; it is not a rendezvous channel.
45    /// Dropping the returned stream drops the writing future and receiver.
46    /// `origin` names what the reply comes from; it is the stream's first
47    /// item, so a consumer that stops early still knows it.
48    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        // The sender's reserved slot holds the origin until the consumer
55        // polls.
56        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    /// A text fragment: extends the open text part, or opens one after
82    /// closing an open reasoning part.
83    pub async fn text(&mut self, text: impl Into<String>) -> Result<(), SinkClosed> {
84        self.write(Block::Text, &text.into()).await
85    }
86
87    /// A reasoning fragment: extends the open reasoning part, or opens one
88    /// after closing an open text part.
89    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    /// A whole tool call, under an id rig issues.
102    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    /// The reply's provider document, the response's `raw`.
129    pub fn raw(&mut self, raw: serde_json::Value) {
130        self.raw = raw;
131    }
132
133    /// The provider's transport request id, the response's
134    /// `provider_request_id`: what a writer relaying a provider's reply
135    /// reports in place of a transport. An empty id is no id.
136    pub fn request_id(&mut self, request_id: impl Into<String>) {
137        self.request_id = Some(request_id.into());
138    }
139
140    /// An in-band error: the consumer's last item.
141    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    /// End the reply: closes the parts still open, then sends the response
147    /// the reply folds into, under the writer's origin. The returned stream
148    /// ends when the writing future also finishes.
149    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    /// Whether the consumer has closed the receiving side.
166    pub fn is_closed(&self) -> bool {
167        self.events.is_closed()
168    }
169
170    /// Send what the writer emitted, folding each event as it leaves.
171    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}