Skip to main content

rig_core/providers/cohere/
streaming.rs

1//! Cohere chat event decoding and terminal metadata for unary and streamed replies.
2//!
3//! ```
4//! use rig_core::providers::cohere::streaming::StreamingEvent;
5//! let event: StreamingEvent = serde_json::from_str(r#"{"type":"message-end"}"#)?;
6//! assert!(matches!(event, StreamingEvent::MessageEnd { delta: None }));
7//! # Ok::<(), serde_json::Error>(())
8//! ```
9
10use crate::error::ProviderError;
11use crate::operation::{CallFragment, Completion, Finish, IfMalformed, TextPart};
12use crate::providers::cohere::completion::{
13    AssistantContent, CompletionResponse, FinishReason, Usage, map_finish_reason,
14};
15use crate::providers::internal::thoughts::Thoughts;
16use crate::providers::internal::wire;
17use crate::wire::{Flow, Out, WireFrame};
18use serde::{Deserialize, Serialize};
19
20/// One streamed frame of Cohere's `/v2/chat`, named by its `type`.
21#[derive(Debug, Deserialize)]
22#[serde(rename_all = "kebab-case", tag = "type")]
23pub enum StreamingEvent {
24    MessageStart {
25        /// The message identifier, when the wire named one.
26        #[serde(default)]
27        id: Option<String>,
28    },
29    /// A content block opens.
30    ContentStart,
31    /// One content fragment: text, or a reasoning model's thought text.
32    ContentDelta {
33        /// The fragment, absent on a frame that carries none.
34        delta: Option<Delta>,
35    },
36    /// A content block closes.
37    ContentEnd,
38    /// The model's plan for the tool calls that follow.
39    ToolPlan,
40    /// A tool call opens, naming the function it calls.
41    ToolCallStart {
42        /// The call's identity and name.
43        delta: Option<Delta>,
44    },
45    /// One argument fragment of the open tool call.
46    ToolCallDelta {
47        /// The fragment.
48        delta: Option<Delta>,
49    },
50    /// The open tool call closes.
51    ToolCallEnd,
52    /// The turn ends: the wire's genuine terminal.
53    MessageEnd {
54        /// Usage and finish reason, absent on a bare terminal.
55        delta: Option<MessageEndDelta>,
56    },
57}
58
59/// The kebab-case `type` values [`StreamingEvent`] can deserialize. A frame
60/// whose `type` is in this set but fails the full parse has a data-level
61/// defect and is surfaced as an `Err` item; a `type` outside this set is an
62/// event this client doesn't know yet and is skipped.
63const KNOWN_EVENT_TYPES: [&str; 9] = [
64    "message-start",
65    "content-start",
66    "content-delta",
67    "content-end",
68    "tool-plan",
69    "tool-call-start",
70    "tool-call-delta",
71    "tool-call-end",
72    "message-end",
73];
74
75/// One content fragment of a `content-delta` frame.
76#[derive(Debug, Deserialize)]
77pub struct MessageContentDelta {
78    /// Assistant text.
79    pub text: Option<String>,
80    /// Cohere v2 reasoning models stream thought text as `content-delta`
81    /// frames whose content carries `thinking` instead of `text`.
82    pub thinking: Option<String>,
83}
84
85/// The function half of a tool-call frame.
86#[derive(Debug, Deserialize)]
87pub struct MessageToolFunctionDelta {
88    /// The tool's name, on the frame that opens the call.
89    pub name: Option<String>,
90    /// One fragment of the call's JSON arguments.
91    pub arguments: Option<String>,
92}
93
94/// The tool-call half of a message delta.
95#[derive(Debug, Deserialize)]
96pub struct MessageToolCallDelta {
97    /// The call's wire id, on the frame that opens it.
98    pub id: Option<String>,
99    /// The function the call names.
100    pub function: Option<MessageToolFunctionDelta>,
101}
102
103/// What one frame's message delta carried.
104#[derive(Debug, Deserialize)]
105pub struct MessageDelta {
106    /// A content fragment.
107    pub content: Option<MessageContentDelta>,
108    /// A tool-call fragment.
109    pub tool_calls: Option<MessageToolCallDelta>,
110}
111
112/// One frame's delta envelope.
113#[derive(Debug, Deserialize)]
114pub struct Delta {
115    /// The message the delta applies to.
116    pub message: Option<MessageDelta>,
117}
118
119/// The `message-end` payload: what the turn cost and why it stopped.
120#[derive(Debug, Deserialize)]
121pub struct MessageEndDelta {
122    /// Token counters, when Cohere reported them.
123    pub usage: Option<Usage>,
124    /// Cohere's own finish reason.
125    #[serde(default)]
126    pub finish_reason: Option<FinishReason>,
127}
128
129/// Cohere's terminal stream record: the `message-end` payload as rig parsed
130/// it: a streamed response's `raw`.
131#[derive(Clone, Debug, Serialize, Deserialize)]
132pub struct StreamingCompletionResponse {
133    pub usage: Option<Usage>,
134    /// Cohere's own `finish_reason` from the `message-end` event, when reported.
135    #[serde(default)]
136    pub finish_reason: Option<FinishReason>,
137    /// The `message-start` event's message identifier, when reported.
138    #[serde(default)]
139    pub message_id: Option<String>,
140}
141
142/// The `/v2/chat` decoder: one state machine for the whole reply and its
143/// stream of events.
144#[derive(Default)]
145pub struct ChatDecoder<'id> {
146    /// The wire index the open tool call's fragments are buffered under.
147    current_tool_call: Option<usize>,
148    /// Tool calls opened so far.
149    calls: usize,
150    message_id: Option<String>,
151    /// Reasoning closes when subsequent content changes block type.
152    thoughts: Thoughts<'id>,
153    text: Option<TextPart<'id>>,
154}
155
156/// Tagged streaming event or untagged unary reply from `/v2/chat`.
157#[derive(Debug, Deserialize)]
158#[serde(untagged)]
159pub enum ChatEvent {
160    /// One SSE frame of `POST /v2/chat` with `stream: true`.
161    Stream(StreamingEvent),
162    /// The whole reply of `POST /v2/chat` without `stream`.
163    Reply(CompletionResponse),
164}
165
166impl<'id> ChatDecoder<'id> {
167    fn close_text(&mut self, out: &mut Out<'id, Completion>) {
168        if let Some(part) = self.text.take() {
169            out.close_text(part);
170        }
171    }
172
173    /// Reasoning then text, as one content fragment carries them.
174    fn content(
175        &mut self,
176        out: &mut Out<'id, Completion>,
177        thinking: Option<&str>,
178        text: Option<&str>,
179    ) {
180        if let Some(thinking) = thinking.filter(|thinking| !thinking.is_empty()) {
181            self.close_text(out);
182            self.thoughts.fragment(out, thinking);
183        }
184        if let Some(text) = text.filter(|text| !text.is_empty()) {
185            self.thoughts.boundary();
186            let part = self.text.get_or_insert_with(|| out.text());
187            out.push_text(part, text);
188        }
189    }
190
191    /// Interpret one streamed `/v2/chat` frame.
192    fn interpret_stream(
193        &mut self,
194        event: StreamingEvent,
195        mut out: Out<'id, Completion>,
196    ) -> Result<Flow, ProviderError> {
197        match event {
198            StreamingEvent::MessageStart { id: Some(id) } => {
199                self.message_id = Some(id);
200            }
201
202            StreamingEvent::ContentDelta { delta: Some(delta) } => {
203                if let Some(content) = delta
204                    .message
205                    .as_ref()
206                    .and_then(|message| message.content.as_ref())
207                {
208                    self.content(
209                        &mut out,
210                        content.thinking.as_deref(),
211                        content.text.as_deref(),
212                    );
213                }
214            }
215
216            StreamingEvent::MessageEnd { delta } => {
217                // A bare message-end still completes the turn with unknown usage and reason.
218                let (usage, finish_reason) = match delta {
219                    Some(delta) => (delta.usage, delta.finish_reason),
220                    None => (None, None),
221                };
222                let message_id = self.message_id.take();
223                return self.end(usage, finish_reason, message_id, out, true);
224            }
225
226            StreamingEvent::ToolCallStart { delta: Some(delta) } => {
227                let Some(tool_calls) = delta
228                    .message
229                    .as_ref()
230                    .and_then(|message| message.tool_calls.as_ref())
231                else {
232                    return Ok(Flow::More);
233                };
234                let (Some(id), Some(function)) = (&tool_calls.id, &tool_calls.function) else {
235                    return Ok(Flow::More);
236                };
237                let (Some(name), Some(arguments)) = (&function.name, &function.arguments) else {
238                    return Ok(Flow::More);
239                };
240                // Tool content interleaving an open thinking part stops it.
241                self.thoughts.boundary();
242                self.close_text(&mut out);
243                let index = self.calls;
244                self.calls += 1;
245                self.current_tool_call = Some(index);
246                // `tool-call-start` may carry initial argument text; on the
247                // wire it is empty, but any payload is part of the call.
248                out.call_fragment(
249                    index,
250                    CallFragment {
251                        id: Some(id.as_str()),
252                        name: Some(name.as_str()),
253                        arguments: Some(arguments.as_str()),
254                        ..CallFragment::default()
255                    },
256                )?;
257            }
258
259            StreamingEvent::ToolCallDelta { delta: Some(delta) } => {
260                let Some(arguments) = delta
261                    .message
262                    .as_ref()
263                    .and_then(|message| message.tool_calls.as_ref())
264                    .and_then(|tool_calls| tool_calls.function.as_ref())
265                    .and_then(|function| function.arguments.as_deref())
266                else {
267                    return Ok(Flow::More);
268                };
269                // A delta with no open call has nothing to extend; the wire
270                // never starts a call mid-delta.
271                if let Some(index) = self.current_tool_call {
272                    out.call_fragment(
273                        index,
274                        CallFragment {
275                            arguments: Some(arguments),
276                            ..CallFragment::default()
277                        },
278                    )?;
279                }
280            }
281
282            StreamingEvent::ToolCallEnd => {
283                // This endpoint drops calls whose assembled arguments are unparseable.
284                if let Some(index) = self.current_tool_call.take() {
285                    out.close_pending(index, IfMalformed::Drop)?;
286                }
287            }
288
289            _ => {}
290        }
291        Ok(Flow::More)
292    }
293
294    /// The unary reply, written as the stream it would have been: its
295    /// content parts, each tool call whole, then the end the `message-end`
296    /// event carries.
297    fn interpret_reply(
298        &mut self,
299        reply: CompletionResponse,
300        mut out: Out<'id, Completion>,
301    ) -> Result<Flow, ProviderError> {
302        let response_id = Some(reply.id.clone()).filter(|id| !id.is_empty());
303        let finish_reason = Some(reply.finish_reason.clone());
304        let usage = reply.usage;
305        let (content, _citations, tool_calls) = reply.message()?;
306
307        for part in content {
308            match part {
309                AssistantContent::Text { text } => self.content(&mut out, None, Some(&text)),
310                AssistantContent::Thinking { thinking } => {
311                    self.content(&mut out, Some(&thinking), None);
312                }
313            }
314        }
315        self.thoughts.boundary();
316        self.close_text(&mut out);
317        for call in tool_calls {
318            let Some(function) = call.function else {
319                continue;
320            };
321            // An absent id is issued by rig, never taken from the tool name,
322            // which cannot tell repeated calls apart.
323            let index = self.calls;
324            self.calls += 1;
325            out.call_fragment(
326                index,
327                CallFragment {
328                    id: call.id.as_deref(),
329                    name: Some(function.name.as_str()),
330                    ..CallFragment::default()
331                },
332            )?;
333            out.announce_pending(index, function.arguments);
334            out.close_pending(index, IfMalformed::Fail)?;
335        }
336        self.end(usage, finish_reason, response_id, out, false)
337    }
338
339    /// The end both replies finish with: Cohere's usage, its finish reason,
340    /// and the message id it named. A stream's `raw` is this native record;
341    /// a whole reply's is the reply itself.
342    fn end(
343        &mut self,
344        usage: Option<Usage>,
345        finish_reason: Option<FinishReason>,
346        message_id: Option<String>,
347        mut out: Out<'id, Completion>,
348        streamed: bool,
349    ) -> Result<Flow, ProviderError> {
350        self.close_text(&mut out);
351        self.thoughts.close(&mut out, None);
352        let recorded_usage = usage
353            .as_ref()
354            .map(crate::completion::Usage::from)
355            .unwrap_or_default();
356        let native = StreamingCompletionResponse {
357            usage,
358            finish_reason,
359            message_id,
360        };
361        if streamed {
362            out.raw(serde_json::to_value(&native)?);
363        }
364        // Cohere's `/v2/chat` reports no model identifier in either mode, so
365        // the normalized `model` stays unset.
366        Ok(out.end(Finish {
367            usage: recorded_usage,
368            reason: native.finish_reason.as_ref().map(map_finish_reason),
369            response_id: native.message_id,
370            ..Finish::default()
371        }))
372    }
373}
374
375impl<'id> crate::wire::Decoder<'id, Completion> for ChatDecoder<'id> {
376    type Event = ChatEvent;
377
378    fn classify(&self, frame: WireFrame) -> crate::wire::WireEvent<ChatEvent> {
379        // One classifier for both shapes: a modeled `type` decodes as a
380        // streamed frame, an unmodeled one stays skippable, and a body with
381        // no `type` at all can only be the unary reply.
382        wire::classify_tagged_frame(&frame.as_str(), "type", |event_type| {
383            KNOWN_EVENT_TYPES.contains(&event_type)
384        })
385    }
386
387    /// EOF without message-end is truncation, not successful completion.
388    fn decode(
389        &mut self,
390        event: ChatEvent,
391        out: Out<'id, Completion>,
392    ) -> Result<Flow, ProviderError> {
393        match event {
394            ChatEvent::Stream(event) => self.interpret_stream(event, out),
395            ChatEvent::Reply(reply) => self.interpret_reply(reply, out),
396        }
397    }
398}
399
400#[cfg(test)]
401mod tests;