Skip to main content

rig_core/providers/cohere/
streaming.rs

1//! The native chat reply decoder, for a whole reply and for the event
2//! stream alike. Each content part, the tool plan and each tool call
3//! becomes one block in the order it arrives. A block's provider item is
4//! the part as Cohere stated it, with the citations that point at it, so
5//! citations and the tool plan round-trip through history.
6//!
7//! ```
8//! use rig_core::providers::cohere::streaming::ChatDecoder;
9//!
10//! let decoder = ChatDecoder::default();
11//! # let _ = decoder;
12//! ```
13
14use std::collections::BTreeMap;
15
16use serde_json::{Map, Value, json};
17
18use super::chat::PLAN;
19use crate::completion::{FinishReason, Usage};
20use crate::error::ProviderError;
21use crate::json_utils::Lenient;
22use crate::message::{CallId, Source, SourceLocation, ToolName};
23use crate::operation::{Block, CallFragment, Completion, Finish};
24use crate::providers::internal::wire;
25use crate::wire::{
26    AdapterEvent, AdapterUsage, AdapterVerdict, Decoder, Flow, ObservationSink, Out, SpanUnit,
27    WireCitation, WireEvent, WireFrame, WireSpan,
28};
29
30/// The stream's event tags; any other tag classifies as unknown.
31const KNOWN_EVENT_TYPES: &[&str] = &[
32    "message-start",
33    "content-start",
34    "content-delta",
35    "content-end",
36    "tool-plan-delta",
37    "tool-call-start",
38    "tool-call-delta",
39    "tool-call-end",
40    "citation-start",
41    "citation-end",
42    "message-end",
43    "debug",
44];
45
46/// The wire index of the tool plan. Content parts keep their own indices,
47/// and calls theirs from [`CALLS`] on.
48const PLAN_INDEX: usize = 1 << 21;
49
50/// The wire index of the first tool call. A content or call index Cohere
51/// states must stay below it.
52const CALLS: usize = 1 << 20;
53
54/// `index` as Cohere stated it for a content part or call, failing when it
55/// would reach the indices the decoder keeps for calls and the plan.
56fn checked(index: usize) -> Result<usize, ProviderError> {
57    if index < CALLS {
58        Ok(index)
59    } else {
60        Err(ProviderError::Response(format!(
61            "Cohere stated index {index}, past the {CALLS} parts or calls a reply may hold"
62        )))
63    }
64}
65
66/// One stream event, or the whole reply a unary call answers (`kind` is
67/// then empty), as Cohere sent it.
68#[derive(Debug, Clone, PartialEq)]
69pub struct ChatEvent {
70    /// The event as sent.
71    pub fields: Value,
72}
73
74impl ChatEvent {
75    fn kind(&self) -> &str {
76        self.fields.str("type").unwrap_or_default()
77    }
78
79    /// The part or call the event addresses.
80    fn index(&self) -> Result<usize, ProviderError> {
81        self.fields
82            .u64("index")
83            .and_then(|index| usize::try_from(index).ok())
84            .ok_or_else(|| {
85                ProviderError::Response(format!("Cohere `{}` names no index", self.kind()))
86            })
87            .and_then(checked)
88    }
89
90    /// The object at `pointer` under the event's `delta.message`.
91    fn message(&self, key: &str) -> Option<&Map<String, Value>> {
92        self.fields
93            .at(&format!("/delta/message/{key}"))
94            .and_then(Value::as_object)
95    }
96}
97
98/// What an open block is.
99#[derive(Debug, Clone, Copy, PartialEq, Eq)]
100enum Kind {
101    Text,
102    Thinking,
103    Plan,
104    Call,
105    Opaque,
106}
107
108/// Decodes native chat replies, a whole reply or a stream of events. An
109/// `*-end` event states its block complete, and `message-end` states every
110/// block still open complete, but a call whose arguments are not JSON.
111#[derive(Debug, Default)]
112pub struct ChatDecoder {
113    /// Each open block's kind and, for a call, the argument text streamed.
114    open: BTreeMap<usize, (Kind, String)>,
115    /// Every index a block opened at and its kind, so a citation finds its
116    /// block.
117    started: BTreeMap<usize, Kind>,
118    message_id: Option<String>,
119}
120
121impl ChatDecoder {
122    /// Open the content part at `index` as Cohere states it.
123    fn content(
124        &mut self,
125        index: usize,
126        part: &Map<String, Value>,
127        out: &mut Out<'_, Completion>,
128    ) -> Result<(), ProviderError> {
129        let index = checked(index)?;
130        let part = Value::Object(part.clone());
131        let (kind, block, key) = match part.str("type") {
132            Some("text") | None => (Kind::Text, Block::Text, "text"),
133            Some("thinking") => (
134                Kind::Thinking,
135                Block::Reasoning { redacted: false },
136                "thinking",
137            ),
138            // A part rig does not know goes back as it came.
139            Some(_) => (Kind::Opaque, Block::Opaque { replay: true }, ""),
140        };
141        let text = part.str(key).unwrap_or_default().to_owned();
142        self.open.insert(index, (kind, String::new()));
143        self.started.insert(index, kind);
144        out.open(index, block, part)?;
145        out.push(index, &text)
146    }
147
148    /// Grow the open part at `index`, and its item, by a delta's text.
149    fn grow(
150        &mut self,
151        index: usize,
152        delta: &Map<String, Value>,
153        out: &mut Out<'_, Completion>,
154    ) -> Result<(), ProviderError> {
155        let Some((kind, _)) = self.open.get(&index) else {
156            return Err(ProviderError::Response(format!(
157                "Cohere streamed content to part {index}, which is not open"
158            )));
159        };
160        let key = match kind {
161            Kind::Thinking => "thinking",
162            Kind::Plan => PLAN,
163            Kind::Text | Kind::Call => "text",
164            // An unknown part's deltas merge into its item.
165            Kind::Opaque => {
166                return out.edit(index, |item| {
167                    crate::operation::completion::merge(item, delta)
168                });
169            }
170        };
171        let Some(text) = delta.get(key).and_then(Value::as_str) else {
172            return Ok(());
173        };
174        out.push(index, text)?;
175        let delta = Map::from_iter([(key.to_owned(), Value::from(text))]);
176        out.edit(index, |item| {
177            crate::operation::completion::merge(item, &delta)
178        })
179    }
180
181    /// Append a fragment of the tool plan, opening it at its first.
182    fn plan(&mut self, text: &str, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
183        if !self.started.contains_key(&PLAN_INDEX) {
184            self.open.insert(PLAN_INDEX, (Kind::Plan, String::new()));
185            self.started.insert(PLAN_INDEX, Kind::Plan);
186            let item = json!({"type": PLAN, PLAN: ""});
187            out.open(PLAN_INDEX, Block::Reasoning { redacted: false }, item)?;
188        }
189        let delta = Map::from_iter([(PLAN.to_owned(), Value::from(text))]);
190        self.grow(PLAN_INDEX, &delta, out)
191    }
192
193    /// Open the tool call at `index` from its stated id, name and the
194    /// arguments it opens with.
195    fn call(
196        &mut self,
197        index: usize,
198        call: &Map<String, Value>,
199        out: &mut Out<'_, Completion>,
200    ) -> Result<(), ProviderError> {
201        // The plan is whole once the calls it plans begin.
202        if self.open.contains_key(&PLAN_INDEX) {
203            self.stop(PLAN_INDEX, out)?;
204        }
205        let index = CALLS + checked(index)?;
206        let item = Value::Object(call.clone());
207        let id = item.str("id").unwrap_or_default().to_owned();
208        let arguments = item
209            .at("/function/arguments")
210            .and_then(Value::as_str)
211            .unwrap_or_default()
212            .to_owned();
213        self.open.insert(index, (Kind::Call, String::new()));
214        self.started.insert(index, Kind::Call);
215        match ToolName::new(
216            item.at("/function/name")
217                .and_then(Value::as_str)
218                .unwrap_or_default(),
219        ) {
220            Ok(name) => {
221                let id = CallId::from_wire(&id);
222                out.open(index, Block::Call { id, name }, item)?;
223            }
224            // A nameless call: the writer drops it with a warning.
225            Err(_) => out.fragment(
226                Some(index),
227                CallFragment {
228                    id: Some(&id),
229                    ..CallFragment::default()
230                },
231            )?,
232        }
233        self.arguments(index, &arguments, out)
234    }
235
236    /// Append argument text to the open call at `index`.
237    fn arguments(
238        &mut self,
239        index: usize,
240        fragment: &str,
241        out: &mut Out<'_, Completion>,
242    ) -> Result<(), ProviderError> {
243        let Some((Kind::Call, json)) = self.open.get_mut(&index) else {
244            return Err(ProviderError::Response(format!(
245                "Cohere streamed arguments to call {}, which is not open",
246                index.saturating_sub(CALLS)
247            )));
248        };
249        json.push_str(fragment);
250        out.push(index, fragment)
251    }
252
253    /// Attach `citation` to the item of the block it cites: the tool plan,
254    /// or the content part at its `content_index` (the first by default).
255    /// A citation of a text part is also cited on its block. A citation of
256    /// a block that never opened has nowhere to go, and one whose
257    /// `content_index` reaches [`CALLS`] fails the reply.
258    fn cite(&self, citation: Value, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
259        let index = if citation.str("type") == Some("PLAN") {
260            PLAN_INDEX
261        } else {
262            checked(
263                citation
264                    .u64("content_index")
265                    .map_or(Ok(0), usize::try_from)
266                    .unwrap_or(usize::MAX),
267            )?
268        };
269        let Some(kind) = self.started.get(&index) else {
270            tracing::warn!(
271                index,
272                "Cohere cited a block the reply never opened; dropping it"
273            );
274            return Ok(());
275        };
276        if *kind == Kind::Text
277            && let Some(cited) = citation_of(&citation)
278        {
279            out.cite(index, cited);
280        }
281        out.edit(index, |item| {
282            if let Some(item) = item.as_object_mut() {
283                match item.get_mut("citations") {
284                    Some(Value::Array(citations)) => citations.push(citation),
285                    _ => {
286                        item.insert("citations".to_owned(), Value::Array(vec![citation]));
287                    }
288                }
289            }
290        })
291    }
292
293    /// End the block at `index` as stated complete. Its item becomes the
294    /// block's native unless it would replay nothing: blank text, or a call
295    /// whose arguments are not a JSON object.
296    fn stop(&mut self, index: usize, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
297        let Some((kind, json)) = self.open.remove(&index) else {
298            return Err(ProviderError::Response(format!(
299                "Cohere ended block {index}, which is not open"
300            )));
301        };
302        out.edit(index, |item| {
303            let kept = match kind {
304                Kind::Text => !item.str("text").unwrap_or_default().trim().is_empty(),
305                Kind::Thinking | Kind::Plan | Kind::Opaque => true,
306                Kind::Call => {
307                    let arguments = if json.trim().is_empty() {
308                        "{}"
309                    } else {
310                        json.as_str()
311                    };
312                    let object = crate::json_utils::parse_tool_arguments(arguments)
313                        .is_ok_and(|parsed| parsed.is_object());
314                    if let Some(function) = item.get_mut("function").and_then(Value::as_object_mut)
315                    {
316                        function.insert("arguments".to_owned(), Value::from(arguments));
317                    }
318                    object
319                }
320            };
321            if !kept {
322                *item = Value::Null;
323            }
324        })?;
325        out.finish(index)
326    }
327
328    /// End the reply: every block still open is complete, but a call whose
329    /// arguments are not JSON yet, which the end closes unfinished. Its
330    /// cost is the `billed_units` priced at the catalog, since the `tokens`
331    /// it reports as usage include tokens Cohere does not bill.
332    fn end(
333        &mut self,
334        usage_value: Option<&Value>,
335        reason: Option<&str>,
336        error: Option<&str>,
337        mut out: Out<'_, Completion>,
338    ) -> Result<Flow, ProviderError> {
339        let open: Vec<usize> = self
340            .open
341            .iter()
342            .filter(|(_, (kind, json))| {
343                *kind != Kind::Call
344                    || json.trim().is_empty()
345                    || crate::json_utils::parse_tool_arguments(json).is_ok_and(|v| v.is_object())
346            })
347            .map(|(index, _)| *index)
348            .collect();
349        for index in open {
350            self.stop(index, &mut out)?;
351        }
352        let error = error
353            .filter(|error| !error.is_empty())
354            .map(str::to_owned)
355            .or_else(|| {
356                (reason == Some("ERROR")).then(|| "Cohere ended the reply with an error".to_owned())
357            });
358        let mut usage = usage_of(usage_value);
359        if let Some(billed) = billed_of(usage_value) {
360            usage.cost = out.catalog_cost(&billed);
361        }
362        Ok(out.end(Finish {
363            usage,
364            reason: reason.map(finish_of),
365            response_id: self.message_id.clone(),
366            model: None,
367            error,
368        }))
369    }
370
371    /// A whole reply, written block by block through the calls a stream
372    /// makes, then ended.
373    fn whole(
374        &mut self,
375        reply: &Value,
376        mut out: Out<'_, Completion>,
377    ) -> Result<Flow, ProviderError> {
378        self.message_id = reply.str("id").map(str::to_owned);
379        // A `message` string is Cohere's error body, sent with a success
380        // status.
381        let message = match reply.get("message") {
382            Some(Value::String(_)) => {
383                return Err(ProviderError::from_provider_body(reply.to_string()));
384            }
385            message => message.unwrap_or(&Value::Null),
386        };
387        // The plan leads, as it streams first.
388        if let Some(plan) = message.str(PLAN).filter(|plan| !plan.is_empty()) {
389            self.plan(plan, &mut out)?;
390            self.stop(PLAN_INDEX, &mut out)?;
391        }
392        // A part or call that is not an object states nothing to keep.
393        for (index, part) in message.arr("content").iter().enumerate() {
394            if let Some(part) = part.as_object() {
395                self.content(index, part, &mut out)?;
396                self.stop(index, &mut out)?;
397            }
398        }
399        for (index, call) in message.arr("tool_calls").iter().enumerate() {
400            if let Some(call) = call.as_object() {
401                self.call(index, call, &mut out)?;
402                self.stop(CALLS + index, &mut out)?;
403            }
404        }
405        for citation in message.arr("citations") {
406            self.cite(citation.clone(), &mut out)?;
407        }
408        self.end(reply.get("usage"), reply.str("finish_reason"), None, out)
409    }
410}
411
412/// How a native `finish_reason` ends the turn. `ERROR`, `TIMEOUT` and any
413/// reason Cohere does not document are [`FinishReason::Other`], which fails
414/// the turn.
415fn finish_of(reason: &str) -> FinishReason {
416    match reason {
417        "COMPLETE" | "STOP_SEQUENCE" => FinishReason::Stop,
418        "MAX_TOKENS" => FinishReason::Length,
419        "TOOL_CALL" => FinishReason::ToolCalls,
420        other => FinishReason::Other(other.to_owned()),
421    }
422}
423
424/// One native citation as a [`WireCitation`]: a span of its content part
425/// in characters, checked against the `text` Cohere quotes, cited by each
426/// document or tool output it names. A citation without both offsets is
427/// `None` and stays only in the item.
428fn citation_of(citation: &Value) -> Option<WireCitation> {
429    let mut span = WireSpan::new(
430        citation.u64("start")?,
431        citation.u64("end")?,
432        SpanUnit::Chars,
433    );
434    if let Some(text) = citation.str("text") {
435        span = span.quoted(text);
436    }
437    let sources = citation
438        .arr("sources")
439        .iter()
440        .filter_map(|source| {
441            let id = source.str("id")?.to_owned();
442            match source.str("type") {
443                Some("document") => {
444                    let cited = Source::new(SourceLocation::Document {
445                        index: None,
446                        id: Some(id),
447                        within: None,
448                    });
449                    Some(match source.at("/document/title").and_then(Value::as_str) {
450                        Some(title) => cited.title(title),
451                        None => cited,
452                    })
453                }
454                Some("tool") => Some(Source::new(SourceLocation::ToolOutput { id })),
455                _ => None,
456            }
457        })
458        .collect();
459    Some(WireCitation::new(Some(span), sources))
460}
461
462/// The `billed_units` Cohere charges for, as usage to price: input and
463/// output, which leave out the tokens Cohere adds and does not bill.
464/// `None` unless both are reported.
465fn billed_of(usage: Option<&Value>) -> Option<Usage> {
466    let count = |pointer: &str| usage?.at(pointer)?.as_u64_lenient();
467    let input = count("/billed_units/input_tokens")?;
468    let output = count("/billed_units/output_tokens")?;
469    Some(Usage::new().input_tokens(input).output_tokens(output))
470}
471
472/// Rig's usage from Cohere's: the `tokens` the model read and wrote,
473/// else the `billed_units`, with `cached_tokens` among the input.
474fn usage_of(usage: Option<&Value>) -> Usage {
475    let count = |pointer: &str| usage?.at(pointer)?.as_u64_lenient();
476    let input = count("/tokens/input_tokens").or_else(|| count("/billed_units/input_tokens"));
477    let output = count("/tokens/output_tokens").or_else(|| count("/billed_units/output_tokens"));
478    Usage {
479        input_tokens: input,
480        output_tokens: output,
481        cached_input_tokens: count("/cached_tokens"),
482        cache_creation_input_tokens: None,
483        reasoning_tokens: count("/tokens/reasoning_tokens"),
484        total_tokens: input.zip(output).map(|(input, output)| input + output),
485        tool_use_prompt_tokens: None,
486        cost: None,
487    }
488}
489
490impl<'id> Decoder<'id, Completion> for ChatDecoder {
491    type Event = ChatEvent;
492
493    /// A stream event by its `type`; a frame without one is the whole
494    /// reply, recognized by its `message`.
495    fn classify(&self, frame: WireFrame) -> WireEvent<ChatEvent> {
496        let data = frame.as_str();
497        wire::classify_or_untagged(
498            &data,
499            "type",
500            |data| {
501                wire::classify_tagged_frame::<Value>(data, "type", |tag| {
502                    KNOWN_EVENT_TYPES.contains(&tag)
503                })
504            },
505            |data| wire::classify_marker_keyed_frame::<Value>(data, &["message"]),
506        )
507        .map(|fields| ChatEvent { fields })
508    }
509
510    fn decode(
511        &mut self,
512        event: ChatEvent,
513        mut out: Out<'id, Completion>,
514    ) -> Result<Flow, ProviderError> {
515        match event.kind() {
516            "" => return self.whole(&event.fields, out),
517            "message-start" => self.message_id = event.fields.str("id").map(str::to_owned),
518            "content-start" => {
519                let part = event.message("content").cloned().unwrap_or_default();
520                self.content(event.index()?, &part, &mut out)?;
521            }
522            "content-delta" => {
523                let delta = event.message("content").cloned().unwrap_or_default();
524                self.grow(event.index()?, &delta, &mut out)?;
525            }
526            "content-end" => self.stop(event.index()?, &mut out)?,
527            "tool-plan-delta" => {
528                if let Some(text) = event
529                    .fields
530                    .at("/delta/message/tool_plan")
531                    .and_then(Value::as_str)
532                {
533                    self.plan(text, &mut out)?;
534                }
535            }
536            "tool-call-start" => {
537                let call = event.message("tool_calls").cloned().unwrap_or_default();
538                self.call(event.index()?, &call, &mut out)?;
539            }
540            "tool-call-delta" => {
541                let fragment = event
542                    .fields
543                    .at("/delta/message/tool_calls/function/arguments")
544                    .and_then(Value::as_str)
545                    .unwrap_or_default();
546                self.arguments(CALLS + event.index()?, fragment, &mut out)?;
547            }
548            "tool-call-end" => self.stop(CALLS + event.index()?, &mut out)?,
549            "citation-start" => {
550                if let Some(citation) = event.message("citations") {
551                    self.cite(Value::Object(citation.clone()), &mut out)?;
552                }
553            }
554            "message-end" => {
555                let delta = event.fields.get("delta");
556                return self.end(
557                    delta.and_then(|delta| delta.get("usage")),
558                    delta.and_then(|delta| delta.str("finish_reason")),
559                    delta.and_then(|delta| delta.str("error")),
560                    out,
561                );
562            }
563            // `citation-end` and `debug`: the classifier passes only the
564            // listed tags.
565            _ => {}
566        }
567        Ok(Flow::More)
568    }
569}
570
571impl ChatDecoder {
572    /// Native chat metadata projected before normalization can discard it:
573    /// the finish reason, the reply id and the usage, on the unary reply
574    /// and on the stream's `message-start` and `message-end` events.
575    pub(crate) fn project(payload: &[u8], sink: &mut ObservationSink<'_>) {
576        let Ok(payload) = serde_json::from_slice::<Value>(payload) else {
577            return;
578        };
579        let end = payload.get("delta").unwrap_or(&payload);
580        if let Some(usage) = end.get("usage") {
581            let usage = usage_of(Some(usage));
582            sink.emit(AdapterEvent::Usage {
583                usage: AdapterUsage {
584                    input_tokens: usage.input_tokens,
585                    output_tokens: usage.output_tokens,
586                    total_tokens: usage.total_tokens,
587                    cached_input_tokens: usage.cached_input_tokens,
588                    reasoning_tokens: usage.reasoning_tokens,
589                    tool_input_tokens: None,
590                },
591            });
592        }
593        let verdict = AdapterVerdict {
594            finish_reason: end.str("finish_reason").map(|reason| sink.scrub(reason)),
595            block_reason: None,
596            detail: None,
597            model: None,
598        };
599        let response_id = payload.str("id").map(|id| sink.scrub(id));
600        sink.provider(verdict, response_id);
601    }
602}
603
604pub(crate) mod document;
605
606#[cfg(test)]
607mod tests;