Skip to main content

rig_core/providers/cohere/
streaming.rs

1use crate::completion::{CompletionError, CompletionRequest};
2use crate::http_client::HttpClientExt;
3use crate::http_client::sse::GenericEventSource;
4use crate::providers::cohere::CompletionModel;
5use crate::providers::cohere::completion::{
6    CohereCompletionRequest, FinishReason, PROVIDER_NAME, Usage, map_finish_reason,
7};
8use crate::providers::internal::adapter::{AdapterOutput, WireAdapter, WireFrame};
9use crate::providers::internal::sse_transport::{
10    OpenLog, SseTransportOptions, open_wire_stream, skip_blank_and_done,
11};
12use crate::providers::internal::wire;
13use crate::streaming::{
14    MintKind, RawStreamingChoice, RawStreamingResult, StreamFinal, StreamPartId,
15    ToolCallDeltaContent, ToolInputEnd, UnparseableToolInput,
16};
17use crate::telemetry::{CompletionOperation, CompletionSpanBuilder, SpanCombinator};
18
19/// Cohere thinking deltas carry no id; a per-stream constant minted identity
20/// keys their accumulation and can never reach a request.
21const REASONING_ID: StreamPartId = StreamPartId::minted(MintKind::Reasoning, 0);
22use crate::{json_utils, streaming};
23use serde::{Deserialize, Serialize};
24
25#[derive(Debug, Deserialize)]
26#[serde(rename_all = "kebab-case", tag = "type")]
27enum StreamingEvent {
28    MessageStart {
29        #[serde(default)]
30        id: Option<String>,
31    },
32    ContentStart,
33    ContentDelta {
34        delta: Option<Delta>,
35    },
36    ContentEnd,
37    ToolPlan,
38    ToolCallStart {
39        delta: Option<Delta>,
40    },
41    ToolCallDelta {
42        delta: Option<Delta>,
43    },
44    ToolCallEnd,
45    MessageEnd {
46        delta: Option<MessageEndDelta>,
47    },
48}
49
50/// The kebab-case `type` values [`StreamingEvent`] can deserialize. A frame
51/// whose `type` is in this set but fails the full parse has a data-level
52/// defect and is surfaced as an `Err` item; a `type` outside this set is an
53/// event this client doesn't know yet and is skipped.
54const KNOWN_EVENT_TYPES: [&str; 9] = [
55    "message-start",
56    "content-start",
57    "content-delta",
58    "content-end",
59    "tool-plan",
60    "tool-call-start",
61    "tool-call-delta",
62    "tool-call-end",
63    "message-end",
64];
65
66#[derive(Debug, Deserialize)]
67struct MessageContentDelta {
68    text: Option<String>,
69    /// Cohere v2 reasoning models stream thought text as `content-delta`
70    /// frames whose content carries `thinking` instead of `text`.
71    thinking: Option<String>,
72}
73
74#[derive(Debug, Deserialize)]
75struct MessageToolFunctionDelta {
76    name: Option<String>,
77    arguments: Option<String>,
78}
79
80#[derive(Debug, Deserialize)]
81struct MessageToolCallDelta {
82    id: Option<String>,
83    function: Option<MessageToolFunctionDelta>,
84}
85
86#[derive(Debug, Deserialize)]
87struct MessageDelta {
88    content: Option<MessageContentDelta>,
89    tool_calls: Option<MessageToolCallDelta>,
90}
91
92#[derive(Debug, Deserialize)]
93struct Delta {
94    message: Option<MessageDelta>,
95}
96
97#[derive(Debug, Deserialize)]
98struct MessageEndDelta {
99    usage: Option<Usage>,
100    #[serde(default)]
101    finish_reason: Option<FinishReason>,
102}
103
104/// Cohere's terminal stream record, kept provider-native for
105/// [`CompletionModel::raw_stream`].
106#[derive(Clone, Debug, Serialize, Deserialize)]
107pub struct StreamingCompletionResponse {
108    pub usage: Option<Usage>,
109    /// Cohere's own `finish_reason` from the `message-end` event, when reported.
110    #[serde(default)]
111    pub finish_reason: Option<FinishReason>,
112    /// The `message-start` event's message identifier, when reported.
113    #[serde(default)]
114    pub message_id: Option<String>,
115}
116
117impl From<&StreamingCompletionResponse> for crate::completion::Usage {
118    fn from(response: &StreamingCompletionResponse) -> crate::completion::Usage {
119        response
120            .usage
121            .as_ref()
122            .map(crate::completion::Usage::from)
123            .unwrap_or_default()
124    }
125}
126
127impl From<StreamingCompletionResponse> for StreamFinal {
128    fn from(response: StreamingCompletionResponse) -> StreamFinal {
129        // Cohere's streaming events carry no model identifier, so the
130        // normalized `model` stays unset.
131        StreamFinal::new(PROVIDER_NAME, crate::completion::Usage::from(&response))
132            .with_optional_finish_reason(response.finish_reason.as_ref().map(map_finish_reason))
133            .with_optional_response_id(response.message_id)
134    }
135}
136
137/// The Cohere v2 chat SSE wire as a [`WireAdapter`].
138///
139/// Holds the per-stream state (open tool call, message id); frame-triage
140/// policy (warn-skip `Unknown` for forward compatibility, in-band `Err` on
141/// `Corrupt` so a later genuine `message-end` can still complete the stream)
142/// lives in [`run_wire_stream`], not here.
143struct CohereAdapter {
144    /// Wire id of the open tool call, when one is streaming. Only the wire
145    /// identity is tracked here; fragment assembly, internal-id minting, and
146    /// finalize policy live in the shared accumulator.
147    current_tool_call: Option<String>,
148    message_id: Option<String>,
149    /// Owns the constant-key reasoning lifecycle — the boundary end this
150    /// wire never announces is derived, not hand-rolled here.
151    reasoning: crate::providers::internal::chunk_lifecycle::MintedReasoningLifecycle,
152}
153
154impl Default for CohereAdapter {
155    fn default() -> Self {
156        Self {
157            current_tool_call: None,
158            message_id: None,
159            reasoning: crate::providers::internal::chunk_lifecycle::MintedReasoningLifecycle::new(
160                REASONING_ID,
161            ),
162        }
163    }
164}
165
166impl WireAdapter for CohereAdapter {
167    type Frame = WireFrame;
168    type Event = StreamingEvent;
169    type Response = StreamingCompletionResponse;
170
171    fn classify(&self, frame: WireFrame) -> wire::WireEvent<StreamingEvent> {
172        wire::classify_tagged_frame(&frame.as_str(), "type", |event_type| {
173            KNOWN_EVENT_TYPES.contains(&event_type)
174        })
175    }
176
177    fn interpret(&mut self, event: StreamingEvent, out: &mut AdapterOutput<Self::Response>) {
178        match event {
179            StreamingEvent::MessageStart { id: Some(id) } => {
180                self.message_id = Some(id);
181            }
182
183            StreamingEvent::ContentDelta { delta: Some(delta) } => {
184                let Some(message) = &delta.message else {
185                    return;
186                };
187                let Some(content) = &message.content else {
188                    return;
189                };
190
191                // Declare what the delta carried (thinking merges under the
192                // per-stream constant minted key); the shared lifecycle
193                // derives the canonical sequence, boundary end included.
194                self.reasoning.emit_chunk(
195                    crate::providers::internal::chunk_lifecycle::ChunkParts {
196                        reasoning: content.thinking.clone(),
197                        reasoning_signature: None,
198                        text: content.text.clone(),
199                        tool_events: Vec::new(),
200                    },
201                    out,
202                );
203            }
204
205            StreamingEvent::MessageEnd { delta } => {
206                // `message-end` is the genuine terminal even when its optional
207                // payload is absent; usage and finish reason then default. The
208                // driver stops consuming after the terminal record.
209                let span = tracing::Span::current();
210                let (usage, finish_reason) = match delta {
211                    Some(delta) => (delta.usage, delta.finish_reason),
212                    None => (None, None),
213                };
214                let recorded_usage = usage
215                    .as_ref()
216                    .map(crate::completion::Usage::from)
217                    .unwrap_or_default();
218                span.record_token_usage(&recorded_usage);
219                out.push(Ok(RawStreamingChoice::FinalResponse(
220                    StreamingCompletionResponse {
221                        usage,
222                        finish_reason,
223                        message_id: self.message_id.take(),
224                    },
225                )));
226            }
227
228            StreamingEvent::ToolCallStart { delta: Some(delta) } => {
229                let Some(message) = &delta.message else {
230                    return;
231                };
232                let Some(tool_calls) = &message.tool_calls else {
233                    return;
234                };
235                let Some(id) = tool_calls.id.clone() else {
236                    return;
237                };
238                let Some(function) = &tool_calls.function else {
239                    return;
240                };
241                let Some(name) = function.name.clone() else {
242                    return;
243                };
244                let Some(arguments) = function.arguments.clone() else {
245                    return;
246                };
247
248                self.current_tool_call = Some(id.clone());
249
250                let mut tool_events = vec![RawStreamingChoice::ToolCallDelta {
251                    id: StreamPartId::wire(id.clone()),
252                    content: ToolCallDeltaContent::Name(name),
253                }];
254                // `tool-call-start` may carry initial argument text; on the
255                // wire it is empty, but any payload must still enter assembly.
256                if !arguments.is_empty() {
257                    tool_events.push(RawStreamingChoice::ToolCallDelta {
258                        id: StreamPartId::wire(id),
259                        content: ToolCallDeltaContent::Delta(arguments),
260                    });
261                }
262                // Tool content interleaving an open thinking block: the
263                // shared lifecycle synthesizes the boundary end.
264                self.reasoning.emit_chunk(
265                    crate::providers::internal::chunk_lifecycle::ChunkParts {
266                        reasoning: None,
267                        reasoning_signature: None,
268                        text: None,
269                        tool_events,
270                    },
271                    out,
272                );
273            }
274
275            StreamingEvent::ToolCallDelta { delta: Some(delta) } => {
276                let Some(message) = &delta.message else {
277                    return;
278                };
279                let Some(tool_calls) = &message.tool_calls else {
280                    return;
281                };
282                let Some(function) = &tool_calls.function else {
283                    return;
284                };
285                let Some(arguments) = function.arguments.clone() else {
286                    return;
287                };
288
289                // A delta with no open call has nothing to extend; skip it, as
290                // the wire never starts a call mid-delta.
291                let Some(id) = self.current_tool_call.clone() else {
292                    return;
293                };
294
295                // Emit the delta so UI can show progress
296                out.push(Ok(RawStreamingChoice::ToolCallDelta {
297                    id: StreamPartId::wire(id),
298                    content: ToolCallDeltaContent::Delta(arguments),
299                }));
300            }
301
302            StreamingEvent::ToolCallEnd => {
303                let Some(id) = self.current_tool_call.take() else {
304                    return;
305                };
306                // Unparseable assembled input drops in the accumulator,
307                // matching the old skip.
308                out.push(Ok(RawStreamingChoice::ToolInputEnd(ToolInputEnd::new(
309                    id,
310                    UnparseableToolInput::Drop,
311                ))));
312            }
313
314            _ => {}
315        }
316    }
317
318    fn finish(&mut self, _out: &mut AdapterOutput<Self::Response>) {
319        // Only Cohere's `message-end` event counts as the provider completing
320        // the turn. A stream that reached EOF without it (truncation) has no
321        // terminal record to report; synthesizing one would present a partial
322        // turn as a successful, zero-usage completion.
323    }
324}
325
326impl<T> CompletionModel<T>
327where
328    T: HttpClientExt + Clone + 'static,
329{
330    /// Open a stream whose terminal record stays Cohere-native.
331    ///
332    /// This is the escape hatch for Cohere's own terminal payload; it shares the
333    /// request builder, transport, telemetry, and error handling with
334    /// [`CompletionModel::stream`](crate::completion::CompletionModel::stream),
335    /// which calls it and normalizes the terminal record once through
336    /// [`streaming::normalize_stream`] — one network request either way.
337    pub async fn raw_stream(
338        &self,
339        request: CompletionRequest,
340    ) -> Result<RawStreamingResult<StreamingCompletionResponse>, CompletionError> {
341        let system_instructions = request.preamble.clone();
342        let record_telemetry_content = request.record_telemetry_content;
343        let mut request = CohereCompletionRequest::try_from((self.model.as_ref(), request))?;
344        let span = CompletionSpanBuilder::new(
345            PROVIDER_NAME,
346            &request.model,
347            CompletionOperation::ChatStreaming,
348        )
349        .system_instructions(system_instructions.as_deref(), record_telemetry_content)
350        .build();
351
352        let params = json_utils::merge(
353            request.additional_params.unwrap_or(serde_json::json!({})),
354            serde_json::json!({"stream": true}),
355        );
356
357        request.additional_params = Some(params);
358
359        crate::providers::internal::trace_json(
360            crate::providers::internal::LogTarget::Streaming,
361            "Cohere streaming completion input",
362            &request,
363        );
364
365        let body = serde_json::to_vec(&request)?;
366
367        let req = self
368            .client
369            .post("/v2/chat")?
370            .body(body)
371            .map_err(|e| CompletionError::HttpError(e.into()))?;
372
373        Ok(open_wire_stream(
374            GenericEventSource::new(self.client.clone(), req),
375            SseTransportOptions {
376                open_log: OpenLog::Trace,
377                stream_ended_is_error: false,
378                log_transport_errors: true,
379            },
380            skip_blank_and_done,
381            CohereAdapter::default(),
382            span,
383        ))
384    }
385
386    pub(crate) async fn stream(
387        &self,
388        request: CompletionRequest,
389    ) -> Result<streaming::StreamingCompletionResponse, CompletionError> {
390        let stream = self.raw_stream(request).await?;
391        let normalized =
392            streaming::normalize_stream(stream, |response: StreamingCompletionResponse| {
393                Ok(response.into())
394            });
395
396        Ok(streaming::StreamingCompletionResponse::stream(
397            PROVIDER_NAME,
398            normalized,
399        ))
400    }
401}
402
403#[cfg(test)]
404mod tests {
405    use super::*;
406    use serde_json::json;
407
408    fn cohere_client<H>(http_client: H) -> crate::providers::cohere::Client<H>
409    where
410        H: HttpClientExt,
411    {
412        crate::providers::cohere::Client::builder()
413            .api_key("test-key")
414            .http_client(http_client)
415            .build()
416            .expect("client should build")
417    }
418
419    fn classify(data: &str) -> wire::WireEvent<StreamingEvent> {
420        wire::classify_tagged_frame(data, "type", |event_type| {
421            KNOWN_EVENT_TYPES.contains(&event_type)
422        })
423    }
424
425    #[test]
426    fn classify_known_event_decodes() {
427        let frame = json!({
428            "type": "content-delta",
429            "delta": {"message": {"content": {"text": "hi"}}},
430        })
431        .to_string();
432        assert!(matches!(
433            classify(&frame),
434            wire::WireEvent::Known(StreamingEvent::ContentDelta { .. })
435        ));
436    }
437
438    #[test]
439    fn classify_unknown_event_type_is_unknown() {
440        let frame = json!({"type": "citation-start"}).to_string();
441        assert!(matches!(
442            classify(&frame),
443            wire::WireEvent::Unknown { event_type, .. } if event_type == "citation-start"
444        ));
445    }
446
447    #[test]
448    fn classify_invalid_json_is_corrupt() {
449        assert!(matches!(classify("{not json"), wire::WireEvent::Corrupt(_)));
450    }
451
452    #[test]
453    fn classify_known_event_with_defective_payload_is_corrupt() {
454        let frame = json!({"type": "content-delta", "delta": 42}).to_string();
455        assert!(matches!(classify(&frame), wire::WireEvent::Corrupt(_)));
456    }
457
458    #[tokio::test]
459    async fn stream_terminal_record_is_normalized() {
460        use crate::client::CompletionClient;
461        use crate::completion::CompletionModel as _;
462        use crate::streaming::StreamedAssistantContent;
463        use crate::test_utils::MockStreamingClient;
464        use futures::StreamExt;
465
466        let sse_bytes = bytes::Bytes::from(
467            [
468                r#"{"type":"message-start","id":"msg_1"}"#,
469                r#"{"type":"content-delta","delta":{"message":{"content":{"text":"hi"}}}}"#,
470                r#"{"type":"message-end","delta":{"finish_reason":"MAX_TOKENS","usage":{"tokens":{"input_tokens":10,"output_tokens":4}}}}"#,
471            ]
472            .iter()
473            .map(|event| format!("data: {event}\n\n"))
474            .collect::<String>(),
475        );
476
477        let client = cohere_client(MockStreamingClient { sse_bytes });
478        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
479        let request = model.completion_request("hello").build();
480
481        let mut stream = crate::completion::CompletionModel::stream(&model, request)
482            .await
483            .expect("stream should open");
484
485        let mut terminal = None;
486        while let Some(item) = stream.next().await {
487            if let StreamedAssistantContent::Final(final_response) =
488                item.expect("stream item should be Ok")
489            {
490                terminal = Some(final_response);
491            }
492        }
493
494        let terminal = terminal.expect("stream should yield a terminal record");
495        assert_eq!(terminal.provider, PROVIDER_NAME);
496        assert_eq!(terminal.response_id.as_deref(), Some("msg_1"));
497        assert_eq!(terminal.message_id, None);
498        assert_eq!(
499            terminal.finish_reason,
500            Some(crate::completion::FinishReason::Length)
501        );
502        assert_eq!(terminal.usage.input_tokens, 10);
503        assert_eq!(terminal.usage.output_tokens, 4);
504        assert_eq!(terminal.usage.total_tokens, 14);
505        // Cohere's stream never names the model.
506        assert_eq!(terminal.model, None);
507    }
508
509    #[tokio::test]
510    async fn truncated_stream_does_not_synthesize_a_terminal_record() {
511        use crate::client::CompletionClient;
512        use crate::completion::CompletionModel as _;
513        use crate::streaming::StreamedAssistantContent;
514        use crate::test_utils::MockStreamingClient;
515        use futures::StreamExt;
516
517        // No `message-end`: the stream was cut off mid-response.
518        let sse_bytes = bytes::Bytes::from(
519            [
520                r#"{"type":"message-start","id":"msg_1"}"#,
521                r#"{"type":"content-delta","delta":{"message":{"content":{"text":"hi"}}}}"#,
522            ]
523            .iter()
524            .map(|event| format!("data: {event}\n\n"))
525            .collect::<String>(),
526        );
527
528        let client = cohere_client(MockStreamingClient { sse_bytes });
529        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
530        let request = model.completion_request("hello").build();
531
532        let mut stream = crate::completion::CompletionModel::stream(&model, request)
533            .await
534            .expect("stream should open");
535
536        let mut texts = Vec::new();
537        let mut saw_terminal = false;
538        while let Some(item) = stream.next().await {
539            match item.expect("stream item should be Ok") {
540                StreamedAssistantContent::Text(text) => texts.push(text.text),
541                StreamedAssistantContent::Final(_) => saw_terminal = true,
542                _ => {}
543            }
544        }
545
546        assert_eq!(texts, ["hi"]);
547        assert!(
548            !saw_terminal,
549            "EOF without message-end must not synthesize a terminal record"
550        );
551        assert!(stream.response.is_none());
552    }
553
554    #[tokio::test]
555    async fn malformed_frame_is_surfaced_and_the_terminal_still_arrives() {
556        use crate::client::CompletionClient;
557        use crate::completion::CompletionModel as _;
558        use crate::streaming::StreamedAssistantContent;
559        use crate::test_utils::MockStreamingClient;
560        use futures::StreamExt;
561
562        // A malformed frame between valid content and the genuine terminal
563        // must surface as an `Err` item without derailing the rest of the
564        // stream.
565        let sse_bytes = bytes::Bytes::from(
566            [
567                r#"{"type":"message-start","id":"msg_1"}"#,
568                r#"{"type":"content-delta","delta":{"message":{"content":{"text":"hi"}}}}"#,
569                "{not json",
570                r#"{"type":"message-end","delta":{"finish_reason":"COMPLETE","usage":{"tokens":{"input_tokens":10,"output_tokens":4}}}}"#,
571            ]
572            .iter()
573            .map(|event| format!("data: {event}\n\n"))
574            .collect::<String>(),
575        );
576
577        let client = cohere_client(MockStreamingClient { sse_bytes });
578        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
579        let request = model.completion_request("hello").build();
580
581        let mut stream = crate::completion::CompletionModel::stream(&model, request)
582            .await
583            .expect("stream should open");
584
585        let mut texts = Vec::new();
586        let mut saw_error = false;
587        let mut terminal = None;
588        while let Some(item) = stream.next().await {
589            match item {
590                Ok(StreamedAssistantContent::Text(text)) => texts.push(text.text),
591                Ok(StreamedAssistantContent::Final(final_response)) => {
592                    terminal = Some(final_response)
593                }
594                Ok(_) => {}
595                Err(_) => saw_error = true,
596            }
597        }
598
599        assert_eq!(texts, ["hi"]);
600        assert!(saw_error, "the malformed frame must reach the consumer");
601        let terminal = terminal.expect("the genuine terminal record must still arrive");
602        assert_eq!(terminal.usage.input_tokens, 10);
603        assert_eq!(terminal.usage.output_tokens, 4);
604    }
605
606    #[tokio::test]
607    async fn known_event_with_malformed_field_is_surfaced_as_an_error() {
608        use crate::client::CompletionClient;
609        use crate::completion::CompletionModel as _;
610        use crate::streaming::StreamedAssistantContent;
611        use crate::test_utils::MockStreamingClient;
612        use futures::StreamExt;
613
614        // A known `type` whose payload fails the full parse (text should be a
615        // string) is a data-level defect, not a forward-compatibility event.
616        let sse_bytes = bytes::Bytes::from(
617            [
618                r#"{"type":"message-start","id":"msg_1"}"#,
619                r#"{"type":"content-delta","delta":{"message":{"content":{"text":42}}}}"#,
620                r#"{"type":"message-end","delta":{"finish_reason":"COMPLETE","usage":{"tokens":{"input_tokens":10,"output_tokens":4}}}}"#,
621            ]
622            .iter()
623            .map(|event| format!("data: {event}\n\n"))
624            .collect::<String>(),
625        );
626
627        let client = cohere_client(MockStreamingClient { sse_bytes });
628        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
629        let request = model.completion_request("hello").build();
630
631        let mut stream = crate::completion::CompletionModel::stream(&model, request)
632            .await
633            .expect("stream should open");
634
635        let mut saw_error = false;
636        let mut terminal = None;
637        while let Some(item) = stream.next().await {
638            match item {
639                Ok(StreamedAssistantContent::Final(final_response)) => {
640                    terminal = Some(final_response)
641                }
642                Ok(_) => {}
643                Err(err) => {
644                    assert!(
645                        matches!(err, CompletionError::JsonError(_)),
646                        "expected a JSON parse error item, got {err:?}"
647                    );
648                    saw_error = true;
649                }
650            }
651        }
652
653        assert!(
654            saw_error,
655            "a known event with a malformed field must surface an error item"
656        );
657        let terminal = terminal.expect("the genuine terminal record must still arrive");
658        assert_eq!(terminal.usage.input_tokens, 10);
659    }
660
661    #[tokio::test]
662    async fn unknown_event_type_is_skipped_and_the_terminal_still_arrives() {
663        use crate::client::CompletionClient;
664        use crate::completion::CompletionModel as _;
665        use crate::streaming::StreamedAssistantContent;
666        use crate::test_utils::MockStreamingClient;
667        use futures::StreamExt;
668
669        // An invented `type` is an event this client doesn't know yet: it is
670        // skipped for forward compatibility, not surfaced as an error.
671        let sse_bytes = bytes::Bytes::from(
672            [
673                r#"{"type":"message-start","id":"msg_1"}"#,
674                r#"{"type":"citation-start","delta":{"whatever":true}}"#,
675                r#"{"type":"content-delta","delta":{"message":{"content":{"text":"hi"}}}}"#,
676                r#"{"type":"message-end","delta":{"finish_reason":"COMPLETE","usage":{"tokens":{"input_tokens":10,"output_tokens":4}}}}"#,
677            ]
678            .iter()
679            .map(|event| format!("data: {event}\n\n"))
680            .collect::<String>(),
681        );
682
683        let client = cohere_client(MockStreamingClient { sse_bytes });
684        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
685        let request = model.completion_request("hello").build();
686
687        let mut stream = crate::completion::CompletionModel::stream(&model, request)
688            .await
689            .expect("stream should open");
690
691        let mut texts = Vec::new();
692        let mut terminal = None;
693        while let Some(item) = stream.next().await {
694            match item.expect("unknown event types must not surface errors") {
695                StreamedAssistantContent::Text(text) => texts.push(text.text),
696                StreamedAssistantContent::Final(final_response) => terminal = Some(final_response),
697                _ => {}
698            }
699        }
700
701        assert_eq!(texts, ["hi"]);
702        let terminal = terminal.expect("the genuine terminal record must still arrive");
703        assert_eq!(terminal.usage.output_tokens, 4);
704    }
705
706    #[tokio::test]
707    async fn message_end_without_delta_still_emits_the_terminal_record() {
708        use crate::client::CompletionClient;
709        use crate::completion::CompletionModel as _;
710        use crate::streaming::StreamedAssistantContent;
711        use crate::test_utils::MockStreamingClient;
712        use futures::StreamExt;
713
714        // `message-end` with no payload is still the provider completing the
715        // turn; the terminal record arrives with default usage.
716        let sse_bytes = bytes::Bytes::from(
717            [
718                r#"{"type":"message-start","id":"msg_1"}"#,
719                r#"{"type":"content-delta","delta":{"message":{"content":{"text":"hi"}}}}"#,
720                r#"{"type":"message-end"}"#,
721            ]
722            .iter()
723            .map(|event| format!("data: {event}\n\n"))
724            .collect::<String>(),
725        );
726
727        let client = cohere_client(MockStreamingClient { sse_bytes });
728        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
729        let request = model.completion_request("hello").build();
730
731        let mut stream = crate::completion::CompletionModel::stream(&model, request)
732            .await
733            .expect("stream should open");
734
735        let mut texts = Vec::new();
736        let mut terminal = None;
737        while let Some(item) = stream.next().await {
738            match item.expect("stream item should be Ok") {
739                StreamedAssistantContent::Text(text) => texts.push(text.text),
740                StreamedAssistantContent::Final(final_response) => terminal = Some(final_response),
741                _ => {}
742            }
743        }
744
745        assert_eq!(texts, ["hi"]);
746        let terminal = terminal.expect("message-end without a delta is still the terminal");
747        assert_eq!(terminal.usage, crate::completion::Usage::default());
748        assert_eq!(terminal.finish_reason, None);
749        assert_eq!(terminal.response_id.as_deref(), Some("msg_1"));
750    }
751
752    #[tokio::test]
753    async fn thinking_deltas_aggregate_into_one_reasoning_part_before_the_text() {
754        use crate::client::CompletionClient;
755        use crate::completion::CompletionModel as _;
756        use crate::message::AssistantContent;
757        use crate::streaming::StreamedAssistantContent;
758        use crate::test_utils::MockStreamingClient;
759        use futures::StreamExt;
760
761        // Cohere v2 reasoning models stream `content-delta` frames carrying
762        // `thinking` before the answer's `text` frames (documented `thinking`
763        // deltas; #2258 F8 — previously these fell through the `text` guard
764        // and the thought text was lost).
765        let sse_bytes = bytes::Bytes::from(
766            [
767                r#"{"type":"message-start","id":"msg_1"}"#,
768                r#"{"type":"content-delta","delta":{"message":{"content":{"thinking":"step one, "}}}}"#,
769                r#"{"type":"content-delta","delta":{"message":{"content":{"thinking":"step two"}}}}"#,
770                r#"{"type":"content-delta","delta":{"message":{"content":{"text":"answer"}}}}"#,
771                r#"{"type":"message-end","delta":{"finish_reason":"COMPLETE","usage":{"tokens":{"input_tokens":10,"output_tokens":4}}}}"#,
772            ]
773            .iter()
774            .map(|event| format!("data: {event}\n\n"))
775            .collect::<String>(),
776        );
777
778        let client = cohere_client(MockStreamingClient { sse_bytes });
779        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
780        let request = model.completion_request("hello").build();
781
782        let mut stream = crate::completion::CompletionModel::stream(&model, request)
783            .await
784            .expect("stream should open");
785
786        let mut reasoning_deltas = Vec::new();
787        while let Some(item) = stream.next().await {
788            if let StreamedAssistantContent::ReasoningDelta { reasoning, .. } =
789                item.expect("stream item should be Ok")
790            {
791                reasoning_deltas.push(reasoning);
792            }
793        }
794        assert_eq!(reasoning_deltas, ["step one, ", "step two"]);
795
796        let parts: Vec<_> = stream.choice.clone();
797        assert_eq!(parts.len(), 2, "one reasoning part, one text part");
798        assert!(matches!(
799            parts.first(),
800            Some(AssistantContent::Reasoning(reasoning))
801                if reasoning.content.iter().any(|content| matches!(
802                    content,
803                    crate::message::ReasoningContent::Text { text, .. }
804                        if text == "step one, step two"
805                ))
806        ));
807        assert!(matches!(
808            parts.get(1),
809            Some(AssistantContent::Text(text)) if text.text == "answer"
810        ));
811    }
812
813    #[tokio::test]
814    async fn errored_stream_does_not_synthesize_a_terminal_record() {
815        use crate::client::CompletionClient;
816        use crate::completion::CompletionModel as _;
817        use crate::streaming::StreamedAssistantContent;
818        use crate::test_utils::HttpErrorStreamingClient;
819        use futures::StreamExt;
820
821        let client = cohere_client(HttpErrorStreamingClient::new(
822            http::StatusCode::TOO_MANY_REQUESTS,
823            r#"{"message":"slow down"}"#,
824        ));
825        let model = client.completion_model(crate::providers::cohere::COMMAND_R_08_2024);
826        let request = model.completion_request("hello").build();
827
828        let mut stream = crate::completion::CompletionModel::stream(&model, request)
829            .await
830            .expect("stream should open");
831
832        let mut saw_error = false;
833        let mut saw_terminal = false;
834        while let Some(item) = stream.next().await {
835            match item {
836                Ok(StreamedAssistantContent::Final(_)) => saw_terminal = true,
837                Ok(_) => {}
838                Err(_) => saw_error = true,
839            }
840        }
841
842        assert!(saw_error, "the transport failure must reach the consumer");
843        assert!(
844            !saw_terminal,
845            "a failed stream must not be reported as a successful, zero-usage completion"
846        );
847        assert!(stream.response.is_none());
848    }
849
850    #[test]
851    fn test_message_content_delta_deserialization() {
852        let json = json!({
853            "type": "content-delta",
854            "delta": {
855                "message": {
856                    "content": {
857                        "text": "Hello world"
858                    }
859                }
860            }
861        });
862
863        let event: StreamingEvent = serde_json::from_value(json).unwrap();
864        match event {
865            StreamingEvent::ContentDelta { delta } => {
866                assert!(delta.is_some());
867                let message = delta.unwrap().message.unwrap();
868                let content = message.content.unwrap();
869                assert_eq!(content.text, Some("Hello world".to_string()));
870            }
871            _ => panic!("Expected ContentDelta"),
872        }
873    }
874
875    #[test]
876    fn test_tool_call_start_deserialization() {
877        let json = json!({
878            "type": "tool-call-start",
879            "delta": {
880                "message": {
881                    "tool_calls": {
882                        "id": "call_123",
883                        "function": {
884                            "name": "get_weather",
885                            "arguments": "{"
886                        }
887                    }
888                }
889            }
890        });
891
892        let event: StreamingEvent = serde_json::from_value(json).unwrap();
893        match event {
894            StreamingEvent::ToolCallStart { delta } => {
895                assert!(delta.is_some());
896                let tool_call = delta.unwrap().message.unwrap().tool_calls.unwrap();
897                assert_eq!(tool_call.id, Some("call_123".to_string()));
898                assert_eq!(
899                    tool_call.function.unwrap().name,
900                    Some("get_weather".to_string())
901                );
902            }
903            _ => panic!("Expected ToolCallStart"),
904        }
905    }
906
907    #[test]
908    fn test_tool_call_delta_deserialization() {
909        let json = json!({
910            "type": "tool-call-delta",
911            "delta": {
912                "message": {
913                    "tool_calls": {
914                        "function": {
915                            "arguments": "\"location\""
916                        }
917                    }
918                }
919            }
920        });
921
922        let event: StreamingEvent = serde_json::from_value(json).unwrap();
923        match event {
924            StreamingEvent::ToolCallDelta { delta } => {
925                assert!(delta.is_some());
926                let tool_call = delta.unwrap().message.unwrap().tool_calls.unwrap();
927                let function = tool_call.function.unwrap();
928                assert_eq!(function.arguments, Some("\"location\"".to_string()));
929            }
930            _ => panic!("Expected ToolCallDelta"),
931        }
932    }
933
934    #[test]
935    fn test_tool_call_end_deserialization() {
936        let json = json!({
937            "type": "tool-call-end"
938        });
939
940        let event: StreamingEvent = serde_json::from_value(json).unwrap();
941        match event {
942            StreamingEvent::ToolCallEnd => {
943                // Success
944            }
945            _ => panic!("Expected ToolCallEnd"),
946        }
947    }
948
949    #[test]
950    fn test_message_end_with_usage_deserialization() {
951        let json = json!({
952            "type": "message-end",
953            "delta": {
954                "usage": {
955                    "tokens": {
956                        "input_tokens": 100,
957                        "output_tokens": 50
958                    }
959                }
960            }
961        });
962
963        let event: StreamingEvent = serde_json::from_value(json).unwrap();
964        match event {
965            StreamingEvent::MessageEnd { delta } => {
966                assert!(delta.is_some());
967                let usage = delta.unwrap().usage.unwrap();
968                let tokens = usage.tokens.unwrap();
969                assert_eq!(tokens.input_tokens, Some(100.0));
970                assert_eq!(tokens.output_tokens, Some(50.0));
971            }
972            _ => panic!("Expected MessageEnd"),
973        }
974    }
975
976    #[test]
977    fn test_streaming_event_order() {
978        // Test that a typical sequence of events deserializes correctly
979        let events = vec![
980            json!({"type": "message-start"}),
981            json!({"type": "content-start"}),
982            json!({
983                "type": "content-delta",
984                "delta": {
985                    "message": {
986                        "content": {
987                            "text": "Sure, "
988                        }
989                    }
990                }
991            }),
992            json!({
993                "type": "content-delta",
994                "delta": {
995                    "message": {
996                        "content": {
997                            "text": "I can help with that."
998                        }
999                    }
1000                }
1001            }),
1002            json!({"type": "content-end"}),
1003            json!({"type": "tool-plan"}),
1004            json!({
1005                "type": "tool-call-start",
1006                "delta": {
1007                    "message": {
1008                        "tool_calls": {
1009                            "id": "call_abc",
1010                            "function": {
1011                                "name": "search",
1012                                "arguments": ""
1013                            }
1014                        }
1015                    }
1016                }
1017            }),
1018            json!({
1019                "type": "tool-call-delta",
1020                "delta": {
1021                    "message": {
1022                        "tool_calls": {
1023                            "function": {
1024                                "arguments": "{\"query\":"
1025                            }
1026                        }
1027                    }
1028                }
1029            }),
1030            json!({
1031                "type": "tool-call-delta",
1032                "delta": {
1033                    "message": {
1034                        "tool_calls": {
1035                            "function": {
1036                                "arguments": "\"Rust\"}"
1037                            }
1038                        }
1039                    }
1040                }
1041            }),
1042            json!({"type": "tool-call-end"}),
1043            json!({
1044                "type": "message-end",
1045                "delta": {
1046                    "usage": {
1047                        "tokens": {
1048                            "input_tokens": 50,
1049                            "output_tokens": 25
1050                        }
1051                    }
1052                }
1053            }),
1054        ];
1055
1056        for (i, event_json) in events.iter().enumerate() {
1057            let result = serde_json::from_value::<StreamingEvent>(event_json.clone());
1058            assert!(
1059                result.is_ok(),
1060                "Failed to deserialize event at index {}: {:?}",
1061                i,
1062                result.err()
1063            );
1064        }
1065    }
1066}