use serde_json::Value;
use crate::canonical::{CanonicalError, ContentKind, Event, FinishReason, Role, Usage};
use crate::protocol::json::{http_error, parse};
use crate::protocol::synth::drain;
use crate::protocol::{DecodeState, Frame};
mod blocks;
mod errors;
pub(super) fn decode(frame: Frame, state: &mut DecodeState) -> Result<Vec<Event>, CanonicalError> {
if let Some(status) = frame.status {
return Ok(vec![Event::Error(http_error(&frame.data, status))]); }
let v = parse(&frame.data)?;
if v["error"].is_object() {
return Ok(vec![Event::Error(errors::stream_error(&v["error"]))]); }
Ok(chunk(&v, state))
}
pub(super) fn decode_full(
body: &[u8],
state: &mut DecodeState,
) -> Result<Vec<Event>, CanonicalError> {
Ok(chunk(&parse(body)?, state))
}
fn chunk(v: &Value, state: &mut DecodeState) -> Vec<Event> {
let mut out = Vec::new();
if !state.started {
state.started = true;
out.push(Event::message_start(
None, v["modelVersion"].as_str().map(str::to_owned),
Role::Assistant,
));
}
let cand = &v["candidates"][0];
for part in cand["content"]["parts"].as_array().into_iter().flatten() {
blocks::part_events(part, state, &mut out);
}
match cand["finishReason"].as_str() {
Some(reason) => finish(reason, v, state, &mut out), None => match v["promptFeedback"]["blockReason"].as_str() {
Some(reason) => prompt_block(reason, v, state, &mut out),
None => {
if let Some(u) = usage(v) {
out.push(Event::Usage(u)); }
}
},
}
out
}
fn finish(reason: &str, v: &Value, state: &mut DecodeState, out: &mut Vec<Event>) {
terminate(finish_reason(reason, state), v, state, out);
}
fn prompt_block(reason: &str, v: &Value, state: &mut DecodeState, out: &mut Vec<Event>) {
terminate(
FinishReason::Refusal {
category: reason.to_lowercase(),
explanation: None,
},
v,
state,
out,
);
}
fn terminate(reason: FinishReason, v: &Value, state: &mut DecodeState, out: &mut Vec<Event>) {
drain(state, out);
if let Some(u) = usage(v) {
out.push(Event::Usage(u));
}
out.push(Event::Finish { reason });
state.terminated = true; }
fn finish_reason(reason: &str, state: &DecodeState) -> FinishReason {
if state
.open
.values()
.any(|b| matches!(b.kind, ContentKind::ToolUse { .. }))
{
return FinishReason::ToolUse;
}
match reason {
"STOP" => FinishReason::Stop,
"MAX_TOKENS" => FinishReason::Length,
"SAFETY" | "PROHIBITED_CONTENT" | "BLOCKLIST" => FinishReason::Refusal {
category: reason.to_lowercase(),
explanation: None,
},
other => FinishReason::Other(other.to_owned()),
}
}
fn usage(v: &Value) -> Option<Usage> {
let u = v.get("usageMetadata").filter(|u| u.is_object())?;
Some(Usage {
input_tokens: u["promptTokenCount"].as_u64().map(|x| x as u32),
output_tokens: u["candidatesTokenCount"].as_u64().map(|x| x as u32),
cache_read_tokens: u["cachedContentTokenCount"].as_u64().map(|x| x as u32),
cache_write_tokens: None,
})
}