use std::time::Duration;
use serde::{Deserialize, Serialize};
use crate::ToolCall;
const MIN_RATE_WINDOW: Duration = Duration::from_millis(50);
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Event {
Text(String),
Reasoning(Reasoning),
ToolCallDelta(ToolCallDelta),
Completed(Completion),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ToolCallDelta {
pub index: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "String::is_empty")]
pub arguments: String,
}
impl ToolCallDelta {
pub fn new(index: usize, arguments: impl Into<String>) -> Self {
Self {
index,
id: None,
name: None,
arguments: arguments.into(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Reasoning {
pub source: ReasoningSource,
pub text: String,
}
impl Reasoning {
pub fn new(source: ReasoningSource, text: impl Into<String>) -> Self {
Self {
source,
text: text.into(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ReasoningSource {
ReasoningContent,
Reasoning,
Think,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FinishReason {
Stop,
ToolCalls,
Length,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Usage {
pub input: u64,
pub output: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub total: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning: Option<u64>,
}
impl Usage {
pub fn new(input: u64, output: u64) -> Self {
Self {
input,
output,
total: None,
reasoning: None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Timing {
#[serde(rename = "first_token_ms", with = "millis")]
pub first_token: Duration,
#[serde(rename = "last_token_ms", with = "millis")]
pub last_token: Duration,
}
impl Timing {
pub fn new(first_token: Duration, last_token: Duration) -> Self {
Self {
first_token,
last_token,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Completion {
pub finish: FinishReason,
#[serde(default)]
pub text: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub reasoning: Vec<Reasoning>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub calls: Vec<ToolCall>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub usage: Option<Usage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timing: Option<Timing>,
}
impl Completion {
pub fn new(finish: FinishReason) -> Self {
Self {
finish,
text: String::new(),
reasoning: Vec::new(),
calls: Vec::new(),
usage: None,
timing: None,
}
}
pub fn tokens_per_second(&self) -> Option<f64> {
let usage = self.usage.as_ref()?;
let timing = self.timing.as_ref()?;
let window = timing.last_token.checked_sub(timing.first_token)?;
(usage.output > 1 && window >= MIN_RATE_WINDOW)
.then(|| usage.output as f64 / window.as_secs_f64())
}
}
mod millis {
use std::time::Duration;
use serde::{Deserialize, Deserializer, Serializer};
pub(super) fn serialize<S: Serializer>(
value: &Duration,
serializer: S,
) -> Result<S::Ok, S::Error> {
let millis = u64::try_from(value.as_millis()).unwrap_or(u64::MAX);
serializer.serialize_u64(millis)
}
pub(super) fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Duration, D::Error> {
u64::deserialize(deserializer).map(Duration::from_millis)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn timed(output: u64, first_ms: u64, last_ms: u64) -> Completion {
let mut done = Completion::new(FinishReason::Stop);
done.usage = Some(Usage::new(10, output));
done.timing = Some(Timing::new(
Duration::from_millis(first_ms),
Duration::from_millis(last_ms),
));
done
}
#[test]
fn the_rate_runs_from_the_first_visible_token_to_the_last() {
let rate = timed(100, 2_000, 4_000).tokens_per_second().unwrap();
assert!((rate - 50.0).abs() < 1e-9, "{rate}");
}
#[test]
fn there_is_no_honest_rate_for_one_token_or_a_short_window() {
assert_eq!(timed(1, 0, 1_000).tokens_per_second(), None);
assert_eq!(timed(12, 100, 149).tokens_per_second(), None);
assert!(timed(12, 100, 150).tokens_per_second().is_some());
assert_eq!(
Completion::new(FinishReason::Stop).tokens_per_second(),
None
);
}
}