use std::collections::BTreeMap;
use serde_json::{Map, Value, json};
use super::chat::PLAN;
use crate::completion::{FinishReason, Usage};
use crate::error::ProviderError;
use crate::json_utils::Lenient;
use crate::message::{CallId, Source, SourceLocation, ToolName};
use crate::operation::{Block, CallFragment, Completion, Finish};
use crate::providers::internal::wire;
use crate::wire::{
AdapterEvent, AdapterUsage, AdapterVerdict, Decoder, Flow, ObservationSink, Out, SpanUnit,
WireCitation, WireEvent, WireFrame, WireSpan,
};
const KNOWN_EVENT_TYPES: &[&str] = &[
"message-start",
"content-start",
"content-delta",
"content-end",
"tool-plan-delta",
"tool-call-start",
"tool-call-delta",
"tool-call-end",
"citation-start",
"citation-end",
"message-end",
"debug",
];
const PLAN_INDEX: usize = 1 << 21;
const CALLS: usize = 1 << 20;
fn checked(index: usize) -> Result<usize, ProviderError> {
if index < CALLS {
Ok(index)
} else {
Err(ProviderError::Response(format!(
"Cohere stated index {index}, past the {CALLS} parts or calls a reply may hold"
)))
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ChatEvent {
pub fields: Value,
}
impl ChatEvent {
fn kind(&self) -> &str {
self.fields.str("type").unwrap_or_default()
}
fn index(&self) -> Result<usize, ProviderError> {
self.fields
.u64("index")
.and_then(|index| usize::try_from(index).ok())
.ok_or_else(|| {
ProviderError::Response(format!("Cohere `{}` names no index", self.kind()))
})
.and_then(checked)
}
fn message(&self, key: &str) -> Option<&Map<String, Value>> {
self.fields
.at(&format!("/delta/message/{key}"))
.and_then(Value::as_object)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Kind {
Text,
Thinking,
Plan,
Call,
Opaque,
}
#[derive(Debug, Default)]
pub struct ChatDecoder {
open: BTreeMap<usize, (Kind, String)>,
started: BTreeMap<usize, Kind>,
message_id: Option<String>,
}
impl ChatDecoder {
fn content(
&mut self,
index: usize,
part: &Map<String, Value>,
out: &mut Out<'_, Completion>,
) -> Result<(), ProviderError> {
let index = checked(index)?;
let part = Value::Object(part.clone());
let (kind, block, key) = match part.str("type") {
Some("text") | None => (Kind::Text, Block::Text, "text"),
Some("thinking") => (
Kind::Thinking,
Block::Reasoning { redacted: false },
"thinking",
),
Some(_) => (Kind::Opaque, Block::Opaque { replay: true }, ""),
};
let text = part.str(key).unwrap_or_default().to_owned();
self.open.insert(index, (kind, String::new()));
self.started.insert(index, kind);
out.open(index, block, part)?;
out.push(index, &text)
}
fn grow(
&mut self,
index: usize,
delta: &Map<String, Value>,
out: &mut Out<'_, Completion>,
) -> Result<(), ProviderError> {
let Some((kind, _)) = self.open.get(&index) else {
return Err(ProviderError::Response(format!(
"Cohere streamed content to part {index}, which is not open"
)));
};
let key = match kind {
Kind::Thinking => "thinking",
Kind::Plan => PLAN,
Kind::Text | Kind::Call => "text",
Kind::Opaque => {
return out.edit(index, |item| {
crate::operation::completion::merge(item, delta)
});
}
};
let Some(text) = delta.get(key).and_then(Value::as_str) else {
return Ok(());
};
out.push(index, text)?;
let delta = Map::from_iter([(key.to_owned(), Value::from(text))]);
out.edit(index, |item| {
crate::operation::completion::merge(item, &delta)
})
}
fn plan(&mut self, text: &str, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
if !self.started.contains_key(&PLAN_INDEX) {
self.open.insert(PLAN_INDEX, (Kind::Plan, String::new()));
self.started.insert(PLAN_INDEX, Kind::Plan);
let item = json!({"type": PLAN, PLAN: ""});
out.open(PLAN_INDEX, Block::Reasoning { redacted: false }, item)?;
}
let delta = Map::from_iter([(PLAN.to_owned(), Value::from(text))]);
self.grow(PLAN_INDEX, &delta, out)
}
fn call(
&mut self,
index: usize,
call: &Map<String, Value>,
out: &mut Out<'_, Completion>,
) -> Result<(), ProviderError> {
if self.open.contains_key(&PLAN_INDEX) {
self.stop(PLAN_INDEX, out)?;
}
let index = CALLS + checked(index)?;
let item = Value::Object(call.clone());
let id = item.str("id").unwrap_or_default().to_owned();
let arguments = item
.at("/function/arguments")
.and_then(Value::as_str)
.unwrap_or_default()
.to_owned();
self.open.insert(index, (Kind::Call, String::new()));
self.started.insert(index, Kind::Call);
match ToolName::new(
item.at("/function/name")
.and_then(Value::as_str)
.unwrap_or_default(),
) {
Ok(name) => {
let id = CallId::from_wire(&id);
out.open(index, Block::Call { id, name }, item)?;
}
Err(_) => out.fragment(
Some(index),
CallFragment {
id: Some(&id),
..CallFragment::default()
},
)?,
}
self.arguments(index, &arguments, out)
}
fn arguments(
&mut self,
index: usize,
fragment: &str,
out: &mut Out<'_, Completion>,
) -> Result<(), ProviderError> {
let Some((Kind::Call, json)) = self.open.get_mut(&index) else {
return Err(ProviderError::Response(format!(
"Cohere streamed arguments to call {}, which is not open",
index.saturating_sub(CALLS)
)));
};
json.push_str(fragment);
out.push(index, fragment)
}
fn cite(&self, citation: Value, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
let index = if citation.str("type") == Some("PLAN") {
PLAN_INDEX
} else {
checked(
citation
.u64("content_index")
.map_or(Ok(0), usize::try_from)
.unwrap_or(usize::MAX),
)?
};
let Some(kind) = self.started.get(&index) else {
tracing::warn!(
index,
"Cohere cited a block the reply never opened; dropping it"
);
return Ok(());
};
if *kind == Kind::Text
&& let Some(cited) = citation_of(&citation)
{
out.cite(index, cited);
}
out.edit(index, |item| {
if let Some(item) = item.as_object_mut() {
match item.get_mut("citations") {
Some(Value::Array(citations)) => citations.push(citation),
_ => {
item.insert("citations".to_owned(), Value::Array(vec![citation]));
}
}
}
})
}
fn stop(&mut self, index: usize, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
let Some((kind, json)) = self.open.remove(&index) else {
return Err(ProviderError::Response(format!(
"Cohere ended block {index}, which is not open"
)));
};
out.edit(index, |item| {
let kept = match kind {
Kind::Text => !item.str("text").unwrap_or_default().trim().is_empty(),
Kind::Thinking | Kind::Plan | Kind::Opaque => true,
Kind::Call => {
let arguments = if json.trim().is_empty() {
"{}"
} else {
json.as_str()
};
let object = crate::json_utils::parse_tool_arguments(arguments)
.is_ok_and(|parsed| parsed.is_object());
if let Some(function) = item.get_mut("function").and_then(Value::as_object_mut)
{
function.insert("arguments".to_owned(), Value::from(arguments));
}
object
}
};
if !kept {
*item = Value::Null;
}
})?;
out.finish(index)
}
fn end(
&mut self,
usage_value: Option<&Value>,
reason: Option<&str>,
error: Option<&str>,
mut out: Out<'_, Completion>,
) -> Result<Flow, ProviderError> {
let open: Vec<usize> = self
.open
.iter()
.filter(|(_, (kind, json))| {
*kind != Kind::Call
|| json.trim().is_empty()
|| crate::json_utils::parse_tool_arguments(json).is_ok_and(|v| v.is_object())
})
.map(|(index, _)| *index)
.collect();
for index in open {
self.stop(index, &mut out)?;
}
let error = error
.filter(|error| !error.is_empty())
.map(str::to_owned)
.or_else(|| {
(reason == Some("ERROR")).then(|| "Cohere ended the reply with an error".to_owned())
});
let mut usage = usage_of(usage_value);
if let Some(billed) = billed_of(usage_value) {
usage.cost = out.catalog_cost(&billed);
}
Ok(out.end(Finish {
usage,
reason: reason.map(finish_of),
response_id: self.message_id.clone(),
model: None,
error,
}))
}
fn whole(
&mut self,
reply: &Value,
mut out: Out<'_, Completion>,
) -> Result<Flow, ProviderError> {
self.message_id = reply.str("id").map(str::to_owned);
let message = match reply.get("message") {
Some(Value::String(_)) => {
return Err(ProviderError::from_provider_body(reply.to_string()));
}
message => message.unwrap_or(&Value::Null),
};
if let Some(plan) = message.str(PLAN).filter(|plan| !plan.is_empty()) {
self.plan(plan, &mut out)?;
self.stop(PLAN_INDEX, &mut out)?;
}
for (index, part) in message.arr("content").iter().enumerate() {
if let Some(part) = part.as_object() {
self.content(index, part, &mut out)?;
self.stop(index, &mut out)?;
}
}
for (index, call) in message.arr("tool_calls").iter().enumerate() {
if let Some(call) = call.as_object() {
self.call(index, call, &mut out)?;
self.stop(CALLS + index, &mut out)?;
}
}
for citation in message.arr("citations") {
self.cite(citation.clone(), &mut out)?;
}
self.end(reply.get("usage"), reply.str("finish_reason"), None, out)
}
}
fn finish_of(reason: &str) -> FinishReason {
match reason {
"COMPLETE" | "STOP_SEQUENCE" => FinishReason::Stop,
"MAX_TOKENS" => FinishReason::Length,
"TOOL_CALL" => FinishReason::ToolCalls,
other => FinishReason::Other(other.to_owned()),
}
}
fn citation_of(citation: &Value) -> Option<WireCitation> {
let mut span = WireSpan::new(
citation.u64("start")?,
citation.u64("end")?,
SpanUnit::Chars,
);
if let Some(text) = citation.str("text") {
span = span.quoted(text);
}
let sources = citation
.arr("sources")
.iter()
.filter_map(|source| {
let id = source.str("id")?.to_owned();
match source.str("type") {
Some("document") => {
let cited = Source::new(SourceLocation::Document {
index: None,
id: Some(id),
within: None,
});
Some(match source.at("/document/title").and_then(Value::as_str) {
Some(title) => cited.title(title),
None => cited,
})
}
Some("tool") => Some(Source::new(SourceLocation::ToolOutput { id })),
_ => None,
}
})
.collect();
Some(WireCitation::new(Some(span), sources))
}
fn billed_of(usage: Option<&Value>) -> Option<Usage> {
let count = |pointer: &str| usage?.at(pointer)?.as_u64_lenient();
let input = count("/billed_units/input_tokens")?;
let output = count("/billed_units/output_tokens")?;
Some(Usage::new().input_tokens(input).output_tokens(output))
}
fn usage_of(usage: Option<&Value>) -> Usage {
let count = |pointer: &str| usage?.at(pointer)?.as_u64_lenient();
let input = count("/tokens/input_tokens").or_else(|| count("/billed_units/input_tokens"));
let output = count("/tokens/output_tokens").or_else(|| count("/billed_units/output_tokens"));
Usage {
input_tokens: input,
output_tokens: output,
cached_input_tokens: count("/cached_tokens"),
cache_creation_input_tokens: None,
reasoning_tokens: count("/tokens/reasoning_tokens"),
total_tokens: input.zip(output).map(|(input, output)| input + output),
tool_use_prompt_tokens: None,
cost: None,
}
}
impl<'id> Decoder<'id, Completion> for ChatDecoder {
type Event = ChatEvent;
fn classify(&self, frame: WireFrame) -> WireEvent<ChatEvent> {
let data = frame.as_str();
wire::classify_or_untagged(
&data,
"type",
|data| {
wire::classify_tagged_frame::<Value>(data, "type", |tag| {
KNOWN_EVENT_TYPES.contains(&tag)
})
},
|data| wire::classify_marker_keyed_frame::<Value>(data, &["message"]),
)
.map(|fields| ChatEvent { fields })
}
fn decode(
&mut self,
event: ChatEvent,
mut out: Out<'id, Completion>,
) -> Result<Flow, ProviderError> {
match event.kind() {
"" => return self.whole(&event.fields, out),
"message-start" => self.message_id = event.fields.str("id").map(str::to_owned),
"content-start" => {
let part = event.message("content").cloned().unwrap_or_default();
self.content(event.index()?, &part, &mut out)?;
}
"content-delta" => {
let delta = event.message("content").cloned().unwrap_or_default();
self.grow(event.index()?, &delta, &mut out)?;
}
"content-end" => self.stop(event.index()?, &mut out)?,
"tool-plan-delta" => {
if let Some(text) = event
.fields
.at("/delta/message/tool_plan")
.and_then(Value::as_str)
{
self.plan(text, &mut out)?;
}
}
"tool-call-start" => {
let call = event.message("tool_calls").cloned().unwrap_or_default();
self.call(event.index()?, &call, &mut out)?;
}
"tool-call-delta" => {
let fragment = event
.fields
.at("/delta/message/tool_calls/function/arguments")
.and_then(Value::as_str)
.unwrap_or_default();
self.arguments(CALLS + event.index()?, fragment, &mut out)?;
}
"tool-call-end" => self.stop(CALLS + event.index()?, &mut out)?,
"citation-start" => {
if let Some(citation) = event.message("citations") {
self.cite(Value::Object(citation.clone()), &mut out)?;
}
}
"message-end" => {
let delta = event.fields.get("delta");
return self.end(
delta.and_then(|delta| delta.get("usage")),
delta.and_then(|delta| delta.str("finish_reason")),
delta.and_then(|delta| delta.str("error")),
out,
);
}
_ => {}
}
Ok(Flow::More)
}
}
impl ChatDecoder {
pub(crate) fn project(payload: &[u8], sink: &mut ObservationSink<'_>) {
let Ok(payload) = serde_json::from_slice::<Value>(payload) else {
return;
};
let end = payload.get("delta").unwrap_or(&payload);
if let Some(usage) = end.get("usage") {
let usage = usage_of(Some(usage));
sink.emit(AdapterEvent::Usage {
usage: AdapterUsage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
total_tokens: usage.total_tokens,
cached_input_tokens: usage.cached_input_tokens,
reasoning_tokens: usage.reasoning_tokens,
tool_input_tokens: None,
},
});
}
let verdict = AdapterVerdict {
finish_reason: end.str("finish_reason").map(|reason| sink.scrub(reason)),
block_reason: None,
detail: None,
model: None,
};
let response_id = payload.str("id").map(|id| sink.scrub(id));
sink.provider(verdict, response_id);
}
}
pub(crate) mod document;
#[cfg(test)]
mod tests;