#![cfg(feature = "llm")]
use std::{sync::Arc, vec};
use im::Vector;
use serde::{Deserialize, Serialize};
use crate::error::AgentError;
use crate::value::AgentValue;
#[cfg(feature = "image")]
use photon_rs::PhotonImage;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ContentBlock {
Text {
text: String,
},
Thinking {
thinking: String,
#[serde(skip_serializing_if = "Option::is_none", default)]
signature: Option<String>,
#[serde(skip_serializing_if = "std::ops::Not::not", default)]
redacted: bool,
},
#[cfg(feature = "image")]
Image {
data: String,
mime_type: String,
},
}
#[derive(Debug, Clone, PartialEq)]
pub enum MessageContent {
Text(String),
Blocks(Vec<ContentBlock>),
}
impl Default for MessageContent {
fn default() -> Self {
MessageContent::Text(String::new())
}
}
impl From<String> for MessageContent {
fn from(s: String) -> Self {
MessageContent::Text(s)
}
}
impl From<&str> for MessageContent {
fn from(s: &str) -> Self {
MessageContent::Text(s.to_string())
}
}
impl MessageContent {
pub fn text(&self) -> String {
match self {
MessageContent::Text(s) => s.clone(),
MessageContent::Blocks(blocks) => blocks
.iter()
.filter_map(|b| match b {
ContentBlock::Text { text } => Some(text.as_str()),
_ => None,
})
.collect(),
}
}
}
fn absorb_legacy_thinking(content: MessageContent, thinking: String) -> MessageContent {
let mut blocks = match content {
MessageContent::Text(s) if s.is_empty() => vec![],
MessageContent::Text(s) => vec![ContentBlock::Text { text: s }],
MessageContent::Blocks(blocks) => blocks,
};
blocks.insert(
0,
ContentBlock::Thinking {
thinking,
signature: None,
redacted: false,
},
);
MessageContent::Blocks(blocks)
}
#[derive(Debug, Default, Clone)]
pub struct Message {
pub id: Option<String>,
pub role: String,
pub content: MessageContent,
pub tokens: Option<usize>,
pub streaming: bool,
pub tool_calls: Option<Vector<ToolCall>>,
pub tool_name: Option<String>,
pub is_error: Option<bool>,
pub stop_reason: Option<String>,
pub usage: Option<Usage>,
#[cfg(feature = "image")]
pub image: Option<Arc<PhotonImage>>,
}
impl Message {
pub fn new(role: String, content: String) -> Self {
Self {
id: None,
role,
content: MessageContent::Text(content),
tokens: None,
streaming: false,
tool_calls: None,
tool_name: None,
is_error: None,
stop_reason: None,
usage: None,
#[cfg(feature = "image")]
image: None,
}
}
pub fn assistant(content: String) -> Self {
Message::new("assistant".to_string(), content)
}
pub fn system(content: String) -> Self {
Message::new("system".to_string(), content)
}
pub fn user(content: String) -> Self {
Message::new("user".to_string(), content)
}
pub fn tool(tool_name: String, content: String) -> Self {
let mut message = Message::new("tool".to_string(), content);
message.tool_name = Some(tool_name);
message
}
pub fn tool_with_content(tool_name: String, content: MessageContent) -> Self {
let mut message = Message::new("tool".to_string(), String::new());
message.content = content;
message.tool_name = Some(tool_name);
message
}
#[cfg(feature = "image")]
pub fn with_image(mut self, image: Arc<PhotonImage>) -> Self {
self.image = Some(image);
self
}
pub fn text(&self) -> String {
self.content.text()
}
pub fn thinking(&self) -> Option<String> {
let MessageContent::Blocks(blocks) = &self.content else {
return None;
};
let parts: Vec<&str> = blocks
.iter()
.filter_map(|b| match b {
ContentBlock::Thinking { redacted: true, .. } => Some("[redacted]"),
ContentBlock::Thinking { thinking, .. } => Some(thinking.as_str()),
_ => None,
})
.collect();
if parts.is_empty() {
None
} else {
Some(parts.join("\n"))
}
}
}
impl PartialEq for Message {
fn eq(&self, other: &Self) -> bool {
self.id == other.id && self.role == other.role && self.content == other.content
}
}
impl Serialize for Message {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
let mut map = serde_json::Map::new();
if let Some(id) = &self.id {
map.insert("id".to_string(), serde_json::Value::String(id.clone()));
}
map.insert(
"role".to_string(),
serde_json::Value::String(self.role.clone()),
);
let content_value = match &self.content {
MessageContent::Text(s) => serde_json::Value::String(s.clone()),
MessageContent::Blocks(blocks)
if blocks
.iter()
.all(|b| matches!(b, ContentBlock::Text { .. })) =>
{
serde_json::Value::String(self.content.text())
}
MessageContent::Blocks(blocks) => {
serde_json::to_value(blocks).map_err(serde::ser::Error::custom)?
}
};
map.insert("content".to_string(), content_value);
if let Some(tokens) = &self.tokens {
map.insert(
"tokens".to_string(),
serde_json::Value::Number((*tokens).into()),
);
}
if self.streaming {
map.insert("streaming".to_string(), serde_json::Value::Bool(true));
}
if let Some(tool_calls) = &self.tool_calls {
let mut tool_calls_vec = vec![];
for call in tool_calls {
tool_calls_vec.push(serde_json::to_value(call).map_err(serde::ser::Error::custom)?);
}
map.insert(
"tool_calls".to_string(),
serde_json::Value::Array(tool_calls_vec),
);
}
if let Some(tool_name) = &self.tool_name {
map.insert(
"tool_name".to_string(),
serde_json::Value::String(tool_name.clone()),
);
}
if let Some(is_error) = &self.is_error {
map.insert("is_error".to_string(), serde_json::Value::Bool(*is_error));
}
if let Some(stop_reason) = &self.stop_reason {
map.insert(
"stop_reason".to_string(),
serde_json::Value::String(stop_reason.clone()),
);
}
if let Some(usage) = &self.usage {
map.insert(
"usage".to_string(),
serde_json::to_value(usage).map_err(serde::ser::Error::custom)?,
);
}
#[cfg(feature = "image")]
{
if let Some(image) = &self.image {
map.insert(
"image".to_string(),
serde_json::Value::String(image.get_base64()),
);
}
}
map.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for Message {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let mut message = Message::user(String::default());
let map = serde_json::Map::deserialize(deserializer)?;
if let Some(id) = map.get("id") {
message.id = id.as_str().map(|s| s.to_string());
}
if let Some(role) = map.get("role") {
message.role = role
.as_str()
.ok_or_else(|| serde::de::Error::custom("role must be a string"))?
.to_string();
}
if let Some(content) = map.get("content") {
message.content = match content {
serde_json::Value::String(s) => MessageContent::Text(s.clone()),
serde_json::Value::Array(_) => {
let blocks: Vec<ContentBlock> = serde_json::from_value(content.clone())
.map_err(|e| {
serde::de::Error::custom(format!("invalid content blocks: {e}"))
})?;
MessageContent::Blocks(blocks)
}
_ => {
return Err(serde::de::Error::custom(
"content must be a string or an array of content blocks",
));
}
};
}
if let Some(tokens) = map.get("tokens") {
message.tokens = tokens.as_u64().map(|u| u as usize);
}
if let Some(thinking) = map.get("thinking").and_then(|v| v.as_str()) {
message.content =
absorb_legacy_thinking(std::mem::take(&mut message.content), thinking.to_string());
}
if let Some(streaming) = map.get("streaming") {
message.streaming = streaming.as_bool().unwrap_or(false);
}
if let Some(tool_calls) = map.get("tool_calls") {
let tool_calls = serde_json::from_value::<Vec<ToolCall>>(tool_calls.clone())
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
message.tool_calls = Some(tool_calls.into());
}
if let Some(tool_name) = map.get("tool_name") {
message.tool_name = tool_name.as_str().map(|s| s.to_string());
}
message.is_error = map.get("is_error").and_then(|v| v.as_bool());
message.stop_reason = map
.get("stop_reason")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
message.usage = map
.get("usage")
.and_then(|v| serde_json::from_value(v.clone()).ok());
#[cfg(feature = "image")]
if let Some(image) = map.get("image") {
let image_str = image
.as_str()
.ok_or_else(|| serde::de::Error::custom("image must be a string"))?;
let image = Arc::new(PhotonImage::new_from_base64(image_str));
message.image = Some(image);
}
Ok(message)
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct Usage {
#[serde(default)]
pub input_tokens: u64,
#[serde(default)]
pub output_tokens: u64,
#[serde(default)]
pub cache_read_tokens: u64,
#[serde(default)]
pub cache_write_tokens: u64,
}
#[cfg(feature = "image")]
const IMAGE_TOKENS: u64 = 1200;
pub fn estimate_message_tokens(m: &Message) -> u64 {
let mut chars: usize = 0;
#[cfg_attr(not(feature = "image"), allow(unused_mut))]
let mut image_tokens: u64 = 0;
match &m.content {
MessageContent::Text(s) => chars += s.len(),
MessageContent::Blocks(blocks) => {
for block in blocks {
match block {
ContentBlock::Text { text } => chars += text.len(),
ContentBlock::Thinking { thinking, .. } => chars += thinking.len(),
#[cfg(feature = "image")]
ContentBlock::Image { .. } => image_tokens += IMAGE_TOKENS,
}
}
}
}
if let Some(tool_calls) = &m.tool_calls {
for call in tool_calls {
chars += call.function.name.len();
chars += serde_json::to_string(&call.function.parameters).map_or(0, |s| s.len());
}
}
#[cfg(feature = "image")]
if m.image.is_some() {
image_tokens += IMAGE_TOKENS;
}
(chars as u64).div_ceil(4) + image_tokens
}
pub fn estimate_context_tokens(messages: &[Message]) -> u64 {
for (i, m) in messages.iter().enumerate().rev() {
if m.role == "assistant"
&& let Some(usage) = &m.usage
{
let anchor = usage.input_tokens
+ usage.output_tokens
+ usage.cache_read_tokens
+ usage.cache_write_tokens;
return anchor
+ messages[i + 1..]
.iter()
.map(estimate_message_tokens)
.sum::<u64>();
}
}
messages.iter().map(estimate_message_tokens).sum()
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolCall {
pub function: ToolCallFunction,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolCallFunction {
pub name: String,
pub parameters: serde_json::Value,
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub parse_error: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum MessageEvent {
Start {
partial: Message,
},
TextDelta {
delta: String,
partial: Message,
},
ThinkingDelta {
delta: String,
partial: Message,
},
ToolCallStart {
index: usize,
partial: Message,
},
ToolCallDelta {
index: usize,
delta: String,
partial: Message,
},
ToolCallEnd {
index: usize,
tool_call: ToolCall,
partial: Message,
},
Done {
message: Message,
},
Error {
message: Message,
error: String,
},
}
impl TryFrom<MessageEvent> for AgentValue {
type Error = AgentError;
fn try_from(event: MessageEvent) -> Result<Self, AgentError> {
let json = serde_json::to_value(&event).map_err(|e| {
AgentError::InvalidValue(format!("Failed to serialize MessageEvent: {e}"))
})?;
AgentValue::from_json(json)
}
}
impl TryFrom<AgentValue> for Message {
type Error = AgentError;
fn try_from(value: AgentValue) -> Result<Self, Self::Error> {
match value {
AgentValue::Message(msg) => Ok((*msg).clone()),
AgentValue::String(s) => Ok(Message::user(s.to_string())),
#[cfg(feature = "image")]
AgentValue::Image(img) => {
let mut message = Message::user("".to_string());
message.image = Some(img.clone());
Ok(message)
}
AgentValue::Object(obj) => {
let role = obj
.get("role")
.and_then(|r| r.as_str())
.unwrap_or("user")
.to_string();
let content_value = obj.get("content").ok_or_else(|| {
AgentError::InvalidValue("Message object missing 'content' field".to_string())
})?;
let content = match content_value {
AgentValue::String(s) => MessageContent::Text(s.to_string()),
AgentValue::Array(_) => {
let blocks: Vec<ContentBlock> =
serde_json::from_value(content_value.to_json()).map_err(|e| {
AgentError::InvalidValue(format!("Invalid content blocks: {e}"))
})?;
MessageContent::Blocks(blocks)
}
_ => {
return Err(AgentError::InvalidValue(
"'content' field must be a string or an array of content blocks"
.to_string(),
));
}
};
let mut message = Message::new(role, String::new());
message.content = content;
let id = obj
.get("id")
.and_then(|i| i.as_str())
.map(|s| s.to_string());
message.id = id;
if let Some(thinking) = obj.get("thinking").and_then(|t| t.as_str()) {
message.content = absorb_legacy_thinking(
std::mem::take(&mut message.content),
thinking.to_string(),
);
}
message.streaming = obj
.get("streaming")
.and_then(|st| st.as_bool())
.unwrap_or_default();
message.is_error = obj.get("is_error").and_then(|v| v.as_bool());
message.stop_reason = obj
.get("stop_reason")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
message.usage = obj
.get("usage")
.and_then(|v| serde_json::from_value(v.to_json()).ok());
if let Some(tool_name) = obj.get("tool_name") {
message.tool_name = Some(
tool_name
.as_str()
.ok_or_else(|| {
AgentError::InvalidValue(
"'tool_name' field must be a string".to_string(),
)
})?
.to_string(),
);
}
if let Some(tool_calls) = obj.get("tool_calls") {
let mut calls = vec![];
for call_value in tool_calls.as_array().ok_or_else(|| {
AgentError::InvalidValue("'tool_calls' field must be an array".to_string())
})? {
let id = call_value
.get("id")
.and_then(|i| i.as_str())
.map(|s| s.to_string());
let function = call_value.get("function").ok_or_else(|| {
AgentError::InvalidValue(
"Tool call missing 'function' field".to_string(),
)
})?;
let tool_name = function.get_str("name").ok_or_else(|| {
AgentError::InvalidValue(
"Tool call function missing 'name' field".to_string(),
)
})?;
let parameters = function.get("parameters").ok_or_else(|| {
AgentError::InvalidValue(
"Tool call function missing 'parameters' field".to_string(),
)
})?;
let call = ToolCall {
function: ToolCallFunction {
id,
name: tool_name.to_string(),
parameters: parameters.to_json(),
parse_error: None,
},
};
calls.push(call);
}
message.tool_calls = Some(calls.into());
}
#[cfg(feature = "image")]
{
if let Some(image_value) = obj.get("image") {
match image_value {
AgentValue::String(s) => {
message.image = Some(Arc::new(PhotonImage::new_from_base64(
s.trim_start_matches("data:image/png;base64,"),
)));
}
AgentValue::Image(img) => {
message.image = Some(img.clone());
}
_ => {}
}
}
}
Ok(message)
}
_ => Err(AgentError::InvalidValue(
"Cannot convert AgentValue to Message".to_string(),
)),
}
}
}
impl From<Message> for AgentValue {
fn from(msg: Message) -> Self {
AgentValue::Message(Arc::new(msg))
}
}
impl From<Vec<Message>> for AgentValue {
fn from(msgs: Vec<Message>) -> Self {
let agent_msgs: Vector<AgentValue> = msgs.into_iter().map(|m| m.into()).collect();
AgentValue::Array(agent_msgs)
}
}
#[cfg(test)]
mod tests {
use im::{hashmap, vector};
use super::*;
#[test]
fn test_tool_call_function_parse_error_serde() {
let func = ToolCallFunction {
name: "t".to_string(),
parameters: serde_json::json!({}),
id: Some("call1".to_string()),
parse_error: None,
};
let json = serde_json::to_value(&func).unwrap();
assert!(json.get("parse_error").is_none());
let restored: ToolCallFunction = serde_json::from_value(json).unwrap();
assert_eq!(restored.parse_error, None);
let func = ToolCallFunction {
name: "t".to_string(),
parameters: serde_json::json!({}),
id: Some("call1".to_string()),
parse_error: Some("bad json".to_string()),
};
let json = serde_json::to_value(&func).unwrap();
assert_eq!(
json.get("parse_error").and_then(|v| v.as_str()),
Some("bad json")
);
let restored: ToolCallFunction = serde_json::from_value(json).unwrap();
assert_eq!(restored.parse_error.as_deref(), Some("bad json"));
}
#[test]
fn test_message_to_from_agent_value() {
let msg = Message::user("What is the weather today?".to_string());
let value: AgentValue = msg.into();
assert!(value.is_message());
let msg_ref = value.as_message().unwrap();
assert_eq!(msg_ref.role, "user");
assert_eq!(msg_ref.text(), "What is the weather today?");
let msg_converted: Message = value.try_into().unwrap();
assert_eq!(msg_converted.role, "user");
assert_eq!(msg_converted.text(), "What is the weather today?");
}
#[test]
fn test_message_with_tool_calls_to_from_agent_value() {
let mut msg = Message::assistant("".to_string());
msg.tool_calls = Some(vector![ToolCall {
function: ToolCallFunction {
id: Some("call1".to_string()),
name: "get_weather".to_string(),
parameters: serde_json::json!({"location": "San Francisco"}),
parse_error: None,
},
}]);
let value: AgentValue = msg.into();
assert!(value.is_message());
let msg_ref = value.as_message().unwrap();
assert_eq!(msg_ref.role, "assistant");
assert_eq!(msg_ref.text(), "");
let tool_calls = msg_ref.tool_calls.as_ref().unwrap();
assert_eq!(tool_calls.len(), 1);
let first_call = &tool_calls[0];
assert_eq!(first_call.function.name, "get_weather");
assert_eq!(first_call.function.parameters["location"], "San Francisco");
let msg_converted: Message = value.try_into().unwrap();
dbg!(&msg_converted);
assert_eq!(msg_converted.role, "assistant");
assert_eq!(msg_converted.text(), "");
let tool_calls = msg_converted.tool_calls.unwrap();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].function.name, "get_weather");
assert_eq!(
tool_calls[0].function.parameters,
serde_json::json!({"location": "San Francisco"})
);
}
#[test]
fn test_tool_message_to_from_agent_value() {
let msg = Message::tool("get_time".to_string(), "2025-01-02 03:04:05".to_string());
let value: AgentValue = msg.clone().into();
let msg_ref = value.as_message().unwrap();
assert_eq!(msg_ref.role, "tool");
assert_eq!(msg_ref.tool_name.as_deref().unwrap(), "get_time");
assert_eq!(msg_ref.text(), "2025-01-02 03:04:05");
let msg_converted: Message = value.try_into().unwrap();
assert_eq!(msg_converted.role, "tool");
assert_eq!(msg_converted.tool_name.as_deref(), Some("get_time"));
assert_eq!(msg_converted.text(), "2025-01-02 03:04:05");
}
#[test]
fn test_message_from_string_value() {
let value = AgentValue::string("Just a simple message");
let msg: Message = value.try_into().unwrap();
assert_eq!(msg.role, "user");
assert_eq!(msg.text(), "Just a simple message");
}
#[test]
fn test_message_from_object_value() {
let value = AgentValue::object(hashmap! {
"role".into() => AgentValue::string("assistant"),
"content".into() =>
AgentValue::string("Here is some information."),
});
let msg: Message = value.try_into().unwrap();
assert_eq!(msg.role, "assistant");
assert_eq!(msg.text(), "Here is some information.");
}
#[test]
fn test_message_from_object_value_reads_is_error() {
let value = AgentValue::object(hashmap! {
"role".into() => AgentValue::string("tool"),
"content".into() => AgentValue::string("boom"),
"tool_name".into() => AgentValue::string("failing_tool"),
"is_error".into() => AgentValue::boolean(true),
});
let msg: Message = value.try_into().unwrap();
assert_eq!(msg.is_error, Some(true));
}
#[test]
fn test_message_from_invalid_value() {
let value = AgentValue::integer(42);
let result: Result<Message, AgentError> = value.try_into();
assert!(result.is_err());
}
#[test]
fn test_message_invalid_object() {
let value =
AgentValue::object(hashmap! {"some_key".into() => AgentValue::string("some_value")});
let result: Result<Message, AgentError> = value.try_into();
assert!(result.is_err());
}
#[test]
fn test_message_to_agent_value_with_tool_calls() {
let message = Message {
role: "assistant".to_string(),
content: MessageContent::default(),
tokens: None,
streaming: false,
tool_calls: Some(vector![ToolCall {
function: ToolCallFunction {
id: Some("call1".to_string()),
name: "active_applications".to_string(),
parameters: serde_json::json!({}),
parse_error: None,
},
}]),
id: None,
tool_name: None,
is_error: None,
stop_reason: None,
usage: None,
#[cfg(feature = "image")]
image: None,
};
let value: AgentValue = message.into();
let msg_ref = value.as_message().unwrap();
assert_eq!(msg_ref.role, "assistant");
assert_eq!(msg_ref.text(), "");
let tool_calls = msg_ref.tool_calls.as_ref().unwrap();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].function.name, "active_applications");
assert!(
tool_calls[0]
.function
.parameters
.as_object()
.unwrap()
.is_empty()
);
}
#[test]
fn test_message_is_error_serde_round_trip() {
let mut msg = Message::tool("failing_tool".to_string(), "boom".to_string());
msg.id = Some("call1".to_string());
msg.is_error = Some(true);
let json = serde_json::to_value(&msg).unwrap();
assert_eq!(json["is_error"], serde_json::json!(true));
let restored: Message = serde_json::from_value(json).unwrap();
assert_eq!(restored.is_error, Some(true));
assert_eq!(restored.id.as_deref(), Some("call1"));
assert_eq!(restored.tool_name.as_deref(), Some("failing_tool"));
}
#[test]
fn test_message_without_is_error_deserializes_to_none() {
let json = serde_json::json!({
"role": "tool",
"content": "ok",
"tool_name": "some_tool",
});
let msg: Message = serde_json::from_value(json).unwrap();
assert_eq!(msg.is_error, None);
}
#[test]
fn test_message_is_error_none_serializes_without_key() {
let msg = Message::tool("some_tool".to_string(), "ok".to_string());
assert_eq!(msg.is_error, None);
let json = serde_json::to_value(&msg).unwrap();
assert!(json.as_object().unwrap().get("is_error").is_none());
}
#[test]
fn test_message_stop_reason_serde_round_trip() {
let mut msg = Message::assistant("partial answer".to_string());
msg.stop_reason = Some("length".to_string());
let json = serde_json::to_value(&msg).unwrap();
assert_eq!(json["stop_reason"], serde_json::json!("length"));
let restored: Message = serde_json::from_value(json).unwrap();
assert_eq!(restored.stop_reason.as_deref(), Some("length"));
}
#[test]
fn test_message_without_stop_reason_deserializes_to_none() {
let json = serde_json::json!({
"role": "assistant",
"content": "ok",
});
let msg: Message = serde_json::from_value(json).unwrap();
assert_eq!(msg.stop_reason, None);
}
#[test]
fn test_message_stop_reason_none_serializes_without_key() {
let msg = Message::assistant("ok".to_string());
assert_eq!(msg.stop_reason, None);
let json = serde_json::to_value(&msg).unwrap();
assert!(json.as_object().unwrap().get("stop_reason").is_none());
}
#[test]
fn test_message_from_object_value_reads_stop_reason() {
let value = AgentValue::object(hashmap! {
"role".into() => AgentValue::string("assistant"),
"content".into() => AgentValue::string("truncated"),
"stop_reason".into() => AgentValue::string("length"),
});
let msg: Message = value.try_into().unwrap();
assert_eq!(msg.stop_reason.as_deref(), Some("length"));
}
#[test]
fn test_message_usage_serde_round_trip() {
let mut msg = Message::assistant("ok".to_string());
msg.usage = Some(Usage {
input_tokens: 100,
output_tokens: 20,
cache_read_tokens: 50,
cache_write_tokens: 10,
});
let json = serde_json::to_value(&msg).unwrap();
assert_eq!(
json["usage"],
serde_json::json!({
"input_tokens": 100,
"output_tokens": 20,
"cache_read_tokens": 50,
"cache_write_tokens": 10,
})
);
let restored: Message = serde_json::from_value(json).unwrap();
assert_eq!(restored.usage, msg.usage);
}
#[test]
fn test_message_without_usage_deserializes_to_none() {
let json = serde_json::json!({
"role": "assistant",
"content": "ok",
});
let msg: Message = serde_json::from_value(json).unwrap();
assert_eq!(msg.usage, None);
}
#[test]
fn test_message_usage_none_serializes_without_key() {
let msg = Message::assistant("ok".to_string());
assert_eq!(msg.usage, None);
let json = serde_json::to_value(&msg).unwrap();
assert!(json.as_object().unwrap().get("usage").is_none());
}
#[test]
fn test_message_from_object_value_reads_usage() {
let value = AgentValue::object(hashmap! {
"role".into() => AgentValue::string("assistant"),
"content".into() => AgentValue::string("ok"),
"usage".into() => AgentValue::object(hashmap! {
"input_tokens".into() => AgentValue::integer(7),
"output_tokens".into() => AgentValue::integer(3),
}),
});
let msg: Message = value.try_into().unwrap();
assert_eq!(
msg.usage,
Some(Usage {
input_tokens: 7,
output_tokens: 3,
cache_read_tokens: 0,
cache_write_tokens: 0,
})
);
}
#[test]
fn test_message_partial_usage_object_deserializes_with_defaults() {
let json = serde_json::json!({
"role": "assistant",
"content": "ok",
"usage": { "input_tokens": 42 },
});
let msg: Message = serde_json::from_value(json).unwrap();
assert_eq!(
msg.usage,
Some(Usage {
input_tokens: 42,
output_tokens: 0,
cache_read_tokens: 0,
cache_write_tokens: 0,
})
);
}
#[test]
fn test_message_unparseable_usage_deserializes_to_none() {
let json = serde_json::json!({
"role": "assistant",
"content": "ok",
"usage": "not an object",
});
let msg: Message = serde_json::from_value(json).unwrap();
assert_eq!(msg.usage, None);
}
#[test]
fn test_message_event_text_delta_serde_round_trip() {
let mut partial = Message::assistant("Hel".to_string());
partial.streaming = true;
let event = MessageEvent::TextDelta {
delta: "l".to_string(),
partial,
};
let json = serde_json::to_value(&event).unwrap();
assert_eq!(json["type"], serde_json::json!("text_delta"));
assert_eq!(json["delta"], serde_json::json!("l"));
assert_eq!(json["partial"]["content"], serde_json::json!("Hel"));
let restored: MessageEvent = serde_json::from_value(json).unwrap();
assert_eq!(restored, event);
let MessageEvent::TextDelta { delta, partial } = restored else {
panic!("wrong variant");
};
assert_eq!(delta, "l");
assert!(partial.streaming);
}
#[test]
fn test_message_event_done_serde_round_trip() {
let mut msg = Message::assistant("Hello".to_string());
msg.id = Some("msg1".to_string());
msg.stop_reason = Some("stop".to_string());
msg.usage = Some(Usage {
input_tokens: 10,
output_tokens: 5,
cache_read_tokens: 2,
cache_write_tokens: 1,
});
let event = MessageEvent::Done {
message: msg.clone(),
};
let json = serde_json::to_value(&event).unwrap();
assert_eq!(json["type"], serde_json::json!("done"));
assert_eq!(json["message"]["role"], serde_json::json!("assistant"));
assert_eq!(json["message"]["content"], serde_json::json!("Hello"));
let restored: MessageEvent = serde_json::from_value(json).unwrap();
assert_eq!(restored, event);
let MessageEvent::Done { message } = restored else {
panic!("wrong variant");
};
assert!(!message.streaming);
assert_eq!(message.stop_reason, msg.stop_reason);
assert_eq!(message.usage, msg.usage);
}
#[test]
fn test_message_event_tool_call_end_serde_round_trip() {
let tool_call = ToolCall {
function: ToolCallFunction {
id: Some("call1".to_string()),
name: "get_weather".to_string(),
parameters: serde_json::json!({"location": "Tokyo"}),
parse_error: None,
},
};
let mut partial = Message::assistant("".to_string());
partial.streaming = true;
partial.tool_calls = Some(vector![tool_call.clone()]);
let event = MessageEvent::ToolCallEnd {
index: 0,
tool_call,
partial,
};
let json = serde_json::to_value(&event).unwrap();
assert_eq!(json["type"], serde_json::json!("tool_call_end"));
assert_eq!(json["index"], serde_json::json!(0));
assert_eq!(
json["tool_call"]["function"]["name"],
serde_json::json!("get_weather")
);
let restored: MessageEvent = serde_json::from_value(json).unwrap();
assert_eq!(restored, event);
let MessageEvent::ToolCallEnd {
index,
tool_call,
partial,
} = restored
else {
panic!("wrong variant");
};
assert_eq!(index, 0);
assert_eq!(tool_call.function.name, "get_weather");
assert_eq!(
tool_call.function.parameters,
serde_json::json!({"location": "Tokyo"})
);
assert!(partial.streaming);
let restored_calls = partial.tool_calls.unwrap();
assert_eq!(restored_calls.len(), 1);
assert_eq!(restored_calls[0].function.id, Some("call1".to_string()));
}
#[test]
fn test_message_event_to_agent_value() {
let event = MessageEvent::Done {
message: Message::assistant("Hello".to_string()),
};
let value = AgentValue::try_from(event).unwrap();
assert!(value.is_object());
assert_eq!(value.get_str("type"), Some("done"));
let message = value.get("message").unwrap();
assert_eq!(message.get_str("role"), Some("assistant"));
assert_eq!(message.get_str("content"), Some("Hello"));
}
#[test]
fn test_message_event_error_to_agent_value() {
let event = MessageEvent::Error {
message: Message::assistant("partial".to_string()),
error: "connection reset".to_string(),
};
let value = AgentValue::try_from(event).unwrap();
assert_eq!(value.get_str("type"), Some("error"));
assert_eq!(value.get_str("error"), Some("connection reset"));
}
#[test]
fn test_message_partial_eq() {
let msg1 = Message::user("hello".to_string());
let msg2 = Message::user("hello".to_string());
let msg3 = Message::user("world".to_string());
assert_eq!(msg1, msg2);
assert_ne!(msg1, msg3);
let mut msg4 = Message::user("hello".to_string());
msg4.id = Some("123".to_string());
assert_ne!(msg1, msg4);
}
#[test]
fn test_message_legacy_thinking_field_absorbed_on_deserialize() {
let json = serde_json::json!({
"role": "assistant",
"content": "hi",
"thinking": "t",
});
let msg: Message = serde_json::from_value(json).unwrap();
assert_eq!(
msg.content,
MessageContent::Blocks(vec![
ContentBlock::Thinking {
thinking: "t".to_string(),
signature: None,
redacted: false,
},
ContentBlock::Text {
text: "hi".to_string()
},
])
);
assert_eq!(msg.text(), "hi");
assert_eq!(msg.thinking().as_deref(), Some("t"));
}
#[test]
fn test_message_pure_text_serializes_as_plain_string() {
let msg = Message::assistant("hello".to_string());
let json = serde_json::to_value(&msg).unwrap();
assert_eq!(json["content"], serde_json::json!("hello"));
let mut msg = Message::assistant(String::new());
msg.content = MessageContent::Blocks(vec![
ContentBlock::Text {
text: "hel".to_string(),
},
ContentBlock::Text {
text: "lo".to_string(),
},
]);
let json = serde_json::to_value(&msg).unwrap();
assert_eq!(json["content"], serde_json::json!("hello"));
}
#[test]
fn test_message_thinking_blocks_serde_round_trip() {
let mut msg = Message::assistant(String::new());
msg.content = MessageContent::Blocks(vec![
ContentBlock::Thinking {
thinking: "reasoning".to_string(),
signature: Some("sig123".to_string()),
redacted: false,
},
ContentBlock::Thinking {
thinking: "opaque-payload".to_string(),
signature: None,
redacted: true,
},
ContentBlock::Text {
text: "answer".to_string(),
},
]);
let json = serde_json::to_value(&msg).unwrap();
assert!(json["content"].is_array());
assert!(json.get("thinking").is_none());
let restored: Message = serde_json::from_value(json).unwrap();
assert_eq!(restored.content, msg.content);
}
#[test]
fn test_message_thinking_redacts_and_joins_with_newline() {
let mut msg = Message::assistant(String::new());
msg.content = MessageContent::Blocks(vec![
ContentBlock::Thinking {
thinking: "Let me think...".to_string(),
signature: Some("sig".to_string()),
redacted: false,
},
ContentBlock::Thinking {
thinking: "EqQBCgIYAg-ciphertext".to_string(),
signature: None,
redacted: true,
},
ContentBlock::Text {
text: "answer".to_string(),
},
]);
assert_eq!(
msg.thinking().as_deref(),
Some("Let me think...\n[redacted]")
);
}
#[test]
fn test_message_mixed_block_order_preserved() {
let blocks = vec![
ContentBlock::Text {
text: "before".to_string(),
},
ContentBlock::Thinking {
thinking: "mid".to_string(),
signature: Some("s".to_string()),
redacted: false,
},
ContentBlock::Text {
text: "after".to_string(),
},
];
let mut msg = Message::assistant(String::new());
msg.content = MessageContent::Blocks(blocks.clone());
let json = serde_json::to_value(&msg).unwrap();
let restored: Message = serde_json::from_value(json).unwrap();
assert_eq!(restored.content, MessageContent::Blocks(blocks));
assert_eq!(restored.text(), "beforeafter");
assert_eq!(restored.thinking().as_deref(), Some("mid"));
}
#[test]
fn test_estimate_message_tokens_rounds_up() {
let msg = Message::user("hello".to_string());
assert_eq!(estimate_message_tokens(&msg), 2);
let msg = Message::user("abcd".to_string());
assert_eq!(estimate_message_tokens(&msg), 1);
}
#[test]
fn test_estimate_message_tokens_counts_tool_calls() {
let parameters = serde_json::json!({"location": "Tokyo"});
let mut msg = Message::assistant(String::new());
msg.tool_calls = Some(vector![ToolCall {
function: ToolCallFunction {
id: Some("call1".to_string()),
name: "get_weather".to_string(),
parameters: parameters.clone(),
parse_error: None,
},
}]);
let chars = "get_weather".len() + serde_json::to_string(¶meters).unwrap().len();
assert_eq!(estimate_message_tokens(&msg), (chars as u64).div_ceil(4));
}
#[test]
fn test_estimate_message_tokens_counts_thinking_blocks() {
let mut msg = Message::assistant(String::new());
msg.content = MessageContent::Blocks(vec![
ContentBlock::Thinking {
thinking: "abcd".to_string(),
signature: None,
redacted: false,
},
ContentBlock::Thinking {
thinking: "wxyz".to_string(),
signature: None,
redacted: true,
},
ContentBlock::Text {
text: "efgh".to_string(),
},
]);
assert_eq!(estimate_message_tokens(&msg), 3);
}
#[cfg(feature = "image")]
#[test]
fn test_estimate_message_tokens_image_block_adds_flat_cost() {
let mut msg = Message::user(String::new());
msg.content = MessageContent::Blocks(vec![
ContentBlock::Text {
text: "abcd".to_string(),
},
ContentBlock::Image {
data: "base64-payload-not-counted-as-chars".to_string(),
mime_type: "image/png".to_string(),
},
]);
assert_eq!(estimate_message_tokens(&msg), 1 + 1200);
}
#[test]
fn test_estimate_context_tokens_anchors_on_latest_usage() {
let mut anchored = Message::assistant("answer".to_string());
anchored.usage = Some(Usage {
input_tokens: 100,
output_tokens: 20,
cache_read_tokens: 50,
cache_write_tokens: 10,
});
let earlier = Message::user("long history covered by the anchor".to_string());
let trailing = Message::user("12345678".to_string());
let messages = vec![earlier, anchored, trailing];
assert_eq!(estimate_context_tokens(&messages), 180 + 2);
}
#[test]
fn test_estimate_context_tokens_sums_all_without_usage() {
let messages = vec![
Message::user("abcd".to_string()), Message::assistant("efghijkl".to_string()), ];
assert_eq!(estimate_context_tokens(&messages), 3);
}
#[test]
fn test_estimate_context_tokens_usage_on_last_message() {
let mut msg = Message::assistant("whatever".to_string());
msg.usage = Some(Usage {
input_tokens: 7,
output_tokens: 3,
cache_read_tokens: 0,
cache_write_tokens: 0,
});
let messages = vec![Message::user("earlier".to_string()), msg];
assert_eq!(estimate_context_tokens(&messages), 10);
}
}