| use derive_more::derive::From; |
| use derive_setters::Setters; |
| use serde::{Deserialize, Serialize}; |
| use strum_macros::{EnumString, IntoStaticStr}; |
|
|
| use super::{ToolCall, ToolCallFull}; |
| use crate::TokenCount; |
| use crate::reasoning::{Reasoning, ReasoningFull}; |
|
|
| |
| |
| |
| |
| |
| #[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq)] |
| #[serde(rename_all = "snake_case")] |
| pub enum MessagePhase { |
| |
| Commentary, |
| |
| FinalAnswer, |
| } |
|
|
| #[derive(Default, Clone, Copy, Debug, Serialize, Deserialize, PartialEq)] |
| pub struct Usage { |
| pub prompt_tokens: TokenCount, |
| pub completion_tokens: TokenCount, |
| pub total_tokens: TokenCount, |
| pub cached_tokens: TokenCount, |
| pub cost: Option<f64>, |
| } |
|
|
| impl Usage { |
| |
| |
| |
| |
| pub fn accumulate(mut self, other: &Usage) -> Self { |
| self.prompt_tokens = self.prompt_tokens + other.prompt_tokens; |
| self.completion_tokens = self.completion_tokens + other.completion_tokens; |
| self.total_tokens = self.total_tokens + other.total_tokens; |
| self.cached_tokens = self.cached_tokens + other.cached_tokens; |
| self.cost = match (self.cost, other.cost) { |
| (Some(a), Some(b)) => Some(a + b), |
| (Some(a), None) => Some(a), |
| (None, Some(b)) => Some(b), |
| (None, None) => None, |
| }; |
| self |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| pub fn merge(mut self, other: &Usage) -> Self { |
| self.prompt_tokens = self.prompt_tokens.max(other.prompt_tokens); |
| self.completion_tokens = self.completion_tokens.max(other.completion_tokens); |
| self.total_tokens = self.total_tokens.max(other.total_tokens); |
| self.cached_tokens = self.cached_tokens.max(other.cached_tokens); |
| self.cost = match (self.cost, other.cost) { |
| (Some(a), Some(b)) => Some(a + b), |
| (Some(a), None) => Some(a), |
| (None, Some(b)) => Some(b), |
| (None, None) => None, |
| }; |
| self |
| } |
| } |
|
|
| |
| |
| |
| #[derive(Default, Clone, Debug, Setters, PartialEq)] |
| #[setters(into, strip_option)] |
| pub struct ChatCompletionMessage { |
| pub content: Option<Content>, |
| pub thought_signature: Option<String>, |
| pub reasoning: Option<Content>, |
| pub reasoning_details: Option<Vec<Reasoning>>, |
| pub tool_calls: Vec<ToolCall>, |
| pub finish_reason: Option<FinishReason>, |
| pub usage: Option<Usage>, |
| |
| |
| pub phase: Option<MessagePhase>, |
| } |
|
|
| impl From<FinishReason> for ChatCompletionMessage { |
| fn from(value: FinishReason) -> Self { |
| ChatCompletionMessage::default().finish_reason(value) |
| } |
| } |
|
|
| |
| #[derive(Clone, Debug, PartialEq, Eq, From)] |
| pub enum Content { |
| Part(ContentPart), |
| Full(ContentFull), |
| } |
|
|
| impl Content { |
| pub fn as_str(&self) -> &str { |
| match self { |
| Content::Part(part) => &part.0, |
| Content::Full(full) => &full.0, |
| } |
| } |
|
|
| pub fn part(content: impl ToString) -> Self { |
| Content::Part(ContentPart(content.to_string())) |
| } |
|
|
| pub fn full(content: impl ToString) -> Self { |
| Content::Full(ContentFull(content.to_string())) |
| } |
|
|
| pub fn is_empty(&self) -> bool { |
| self.as_str().is_empty() |
| } |
|
|
| pub fn is_part(&self) -> bool { |
| matches!(self, Content::Part(_)) |
| } |
| } |
|
|
| |
| #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] |
| #[serde(transparent)] |
| pub struct ContentPart(String); |
|
|
| |
| #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] |
| #[serde(transparent)] |
| pub struct ContentFull(String); |
|
|
| impl<T: AsRef<str>> From<T> for Content { |
| fn from(value: T) -> Self { |
| Content::Full(ContentFull(value.as_ref().to_string())) |
| } |
| } |
|
|
| |
| |
| #[derive(Clone, Debug, Deserialize, Serialize, EnumString, IntoStaticStr, PartialEq, Eq)] |
| pub enum FinishReason { |
| |
| |
| #[strum(serialize = "length")] |
| Length, |
| |
| |
| #[strum(serialize = "content_filter")] |
| ContentFilter, |
| |
| #[strum(serialize = "tool_calls")] |
| ToolCalls, |
| |
| #[strum(serialize = "stop", serialize = "end_turn")] |
| Stop, |
| } |
|
|
| impl ChatCompletionMessage { |
| pub fn assistant(content: impl Into<Content>) -> ChatCompletionMessage { |
| ChatCompletionMessage::default().content(content.into()) |
| } |
|
|
| pub fn add_reasoning_detail(mut self, detail: impl Into<Reasoning>) -> Self { |
| let detail = detail.into(); |
| if let Some(ref mut details) = self.reasoning_details { |
| details.push(detail); |
| } else { |
| self.reasoning_details = Some(vec![detail]); |
| } |
| self |
| } |
|
|
| pub fn add_tool_call(mut self, call_tool: impl Into<ToolCall>) -> Self { |
| self.tool_calls.push(call_tool.into()); |
| self |
| } |
|
|
| pub fn extend_calls(mut self, calls: Vec<impl Into<ToolCall>>) -> Self { |
| self.tool_calls.extend(calls.into_iter().map(Into::into)); |
| self |
| } |
|
|
| pub fn finish_reason_opt(mut self, reason: Option<FinishReason>) -> Self { |
| self.finish_reason = reason; |
| self |
| } |
|
|
| pub fn content_part(mut self, content: impl ToString) -> Self { |
| self.content = Some(Content::Part(ContentPart(content.to_string()))); |
| self |
| } |
|
|
| pub fn content_full(mut self, content: impl ToString) -> Self { |
| self.content = Some(Content::Full(ContentFull(content.to_string()))); |
| self |
| } |
| } |
|
|
| |
| |
| |
| #[derive(Clone, Debug, PartialEq)] |
| pub struct ChatCompletionMessageFull { |
| pub content: String, |
| pub thought_signature: Option<String>, |
| pub reasoning: Option<String>, |
| pub tool_calls: Vec<ToolCallFull>, |
| pub reasoning_details: Option<Vec<ReasoningFull>>, |
| pub usage: Usage, |
| pub finish_reason: Option<FinishReason>, |
| |
| |
| pub phase: Option<MessagePhase>, |
| } |
|
|
| #[cfg(test)] |
| mod tests { |
| use std::str::FromStr; |
|
|
| use pretty_assertions::assert_eq; |
|
|
| use super::*; |
| #[test] |
| fn test_usage_accumulate_with_both_costs() { |
| let fixture_usage_1 = Usage { |
| prompt_tokens: TokenCount::Actual(100), |
| completion_tokens: TokenCount::Actual(50), |
| total_tokens: TokenCount::Actual(150), |
| cached_tokens: TokenCount::Actual(20), |
| cost: Some(0.01), |
| }; |
|
|
| let fixture_usage_2 = Usage { |
| prompt_tokens: TokenCount::Actual(200), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(275), |
| cached_tokens: TokenCount::Actual(30), |
| cost: Some(0.02), |
| }; |
|
|
| let actual = fixture_usage_1.accumulate(&fixture_usage_2); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Actual(300), |
| completion_tokens: TokenCount::Actual(125), |
| total_tokens: TokenCount::Actual(425), |
| cached_tokens: TokenCount::Actual(50), |
| cost: Some(0.03), |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_usage_accumulate_mixed_token_types() { |
| let fixture_usage_1 = Usage { |
| prompt_tokens: TokenCount::Actual(100), |
| completion_tokens: TokenCount::Approx(50), |
| total_tokens: TokenCount::Actual(150), |
| cached_tokens: TokenCount::Actual(20), |
| cost: Some(0.01), |
| }; |
|
|
| let fixture_usage_2 = Usage { |
| prompt_tokens: TokenCount::Approx(200), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Approx(275), |
| cached_tokens: TokenCount::Approx(30), |
| cost: Some(0.02), |
| }; |
|
|
| let actual = fixture_usage_1.accumulate(&fixture_usage_2); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Approx(300), |
| completion_tokens: TokenCount::Approx(125), |
| total_tokens: TokenCount::Approx(425), |
| cached_tokens: TokenCount::Approx(50), |
| cost: Some(0.03), |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_usage_accumulate_partial_costs() { |
| let fixture_usage_1 = Usage { |
| prompt_tokens: TokenCount::Actual(100), |
| completion_tokens: TokenCount::Actual(50), |
| total_tokens: TokenCount::Actual(150), |
| cached_tokens: TokenCount::Actual(20), |
| cost: Some(0.01), |
| }; |
|
|
| let fixture_usage_2 = Usage { |
| prompt_tokens: TokenCount::Actual(200), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(275), |
| cached_tokens: TokenCount::Actual(30), |
| cost: None, |
| }; |
|
|
| let actual = fixture_usage_1.accumulate(&fixture_usage_2); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Actual(300), |
| completion_tokens: TokenCount::Actual(125), |
| total_tokens: TokenCount::Actual(425), |
| cached_tokens: TokenCount::Actual(50), |
| cost: Some(0.01), |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_usage_accumulate_no_costs() { |
| let fixture_usage_1 = Usage { |
| prompt_tokens: TokenCount::Actual(100), |
| completion_tokens: TokenCount::Actual(50), |
| total_tokens: TokenCount::Actual(150), |
| cached_tokens: TokenCount::Actual(20), |
| cost: None, |
| }; |
|
|
| let fixture_usage_2 = Usage { |
| prompt_tokens: TokenCount::Actual(200), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(275), |
| cached_tokens: TokenCount::Actual(30), |
| cost: None, |
| }; |
|
|
| let actual = fixture_usage_1.accumulate(&fixture_usage_2); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Actual(300), |
| completion_tokens: TokenCount::Actual(125), |
| total_tokens: TokenCount::Actual(425), |
| cached_tokens: TokenCount::Actual(50), |
| cost: None, |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_usage_accumulate_with_defaults() { |
| let fixture_usage_1 = Usage::default(); |
|
|
| let fixture_usage_2 = Usage { |
| prompt_tokens: TokenCount::Actual(200), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(275), |
| cached_tokens: TokenCount::Actual(30), |
| cost: Some(0.05), |
| }; |
|
|
| let actual = fixture_usage_1.accumulate(&fixture_usage_2); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Actual(200), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(275), |
| cached_tokens: TokenCount::Actual(30), |
| cost: Some(0.05), |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_finish_reason_from_str() { |
| assert_eq!( |
| FinishReason::from_str("length").unwrap(), |
| FinishReason::Length |
| ); |
| assert_eq!( |
| FinishReason::from_str("content_filter").unwrap(), |
| FinishReason::ContentFilter |
| ); |
| assert_eq!( |
| FinishReason::from_str("tool_calls").unwrap(), |
| FinishReason::ToolCalls |
| ); |
| assert_eq!(FinishReason::from_str("stop").unwrap(), FinishReason::Stop); |
| assert_eq!( |
| FinishReason::from_str("end_turn").unwrap(), |
| FinishReason::Stop |
| ); |
| } |
|
|
| #[test] |
| fn test_usage_merge_anthropic_cumulative() { |
| |
| |
| let fixture_message_start = Usage { |
| prompt_tokens: TokenCount::Actual(1000), |
| completion_tokens: TokenCount::Actual(1), |
| total_tokens: TokenCount::Actual(1001), |
| cached_tokens: TokenCount::Actual(300), |
| cost: None, |
| }; |
|
|
| let fixture_message_delta = Usage { |
| prompt_tokens: TokenCount::Actual(0), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(75), |
| cached_tokens: TokenCount::Actual(0), |
| cost: None, |
| }; |
|
|
| let actual = fixture_message_start.merge(&fixture_message_delta); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Actual(1000), |
| completion_tokens: TokenCount::Actual(75), |
| total_tokens: TokenCount::Actual(1001), |
| cached_tokens: TokenCount::Actual(300), |
| cost: None, |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_usage_merge_preserves_costs() { |
| let fixture_usage_1 = Usage { |
| prompt_tokens: TokenCount::Actual(100), |
| completion_tokens: TokenCount::Actual(0), |
| total_tokens: TokenCount::Actual(100), |
| cached_tokens: TokenCount::Actual(0), |
| cost: Some(0.01), |
| }; |
|
|
| let fixture_usage_2 = Usage { |
| prompt_tokens: TokenCount::Actual(0), |
| completion_tokens: TokenCount::Actual(50), |
| total_tokens: TokenCount::Actual(50), |
| cached_tokens: TokenCount::Actual(0), |
| cost: Some(0.02), |
| }; |
|
|
| let actual = fixture_usage_1.merge(&fixture_usage_2); |
|
|
| let expected = Usage { |
| prompt_tokens: TokenCount::Actual(100), |
| completion_tokens: TokenCount::Actual(50), |
| total_tokens: TokenCount::Actual(100), |
| cached_tokens: TokenCount::Actual(0), |
| cost: Some(0.03), |
| }; |
|
|
| assert_eq!(actual, expected); |
| } |
| } |
|
|