use serde::ser::SerializeStruct;
use serde::{Deserialize, Serialize, Serializer};
use super::tool::{CacheControl, ToolCall};
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct ChatMessage {
pub role: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(skip)]
pub cache_control: Option<CacheControl>,
}
impl Serialize for ChatMessage {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut state = serializer.serialize_struct("ChatMessage", 5)?;
state.serialize_field("role", &self.role)?;
match (&self.content, &self.cache_control) {
(Some(text), Some(cache_control)) => {
let block = [CachedTextBlock {
kind: "text",
text: text.as_str(),
cache_control,
}];
state.serialize_field("content", &block)?;
}
(Some(text), None) => state.serialize_field("content", text)?,
(None, _) => state.skip_field("content")?,
}
match &self.tool_calls {
Some(calls) => state.serialize_field("tool_calls", calls)?,
None => state.skip_field("tool_calls")?,
}
match &self.tool_call_id {
Some(id) => state.serialize_field("tool_call_id", id)?,
None => state.skip_field("tool_call_id")?,
}
match &self.name {
Some(name) => state.serialize_field("name", name)?,
None => state.skip_field("name")?,
}
state.end()
}
}
#[derive(Serialize)]
struct CachedTextBlock<'a> {
#[serde(rename = "type")]
kind: &'static str,
text: &'a str,
cache_control: &'a CacheControl,
}
impl ChatMessage {
pub fn system(content: impl Into<String>) -> Self {
Self::text("system", content)
}
pub fn user(content: impl Into<String>) -> Self {
Self::text("user", content)
}
pub fn assistant(content: impl Into<String>) -> Self {
Self::text("assistant", content)
}
pub fn tool_result(
tool_call_id: impl Into<String>,
name: impl Into<String>,
content: impl Into<String>,
) -> Self {
Self {
role: "tool".into(),
content: Some(content.into()),
tool_calls: None,
tool_call_id: Some(tool_call_id.into()),
name: Some(name.into()),
cache_control: None,
}
}
fn text(role: &str, content: impl Into<String>) -> Self {
Self {
role: role.into(),
content: Some(content.into()),
tool_calls: None,
tool_call_id: None,
name: None,
cache_control: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn constructors_set_role() {
assert_eq!(ChatMessage::system("s").role, "system");
assert_eq!(ChatMessage::user("u").role, "user");
assert_eq!(ChatMessage::assistant("a").role, "assistant");
}
#[test]
fn tool_result_fields() {
let m = ChatMessage::tool_result("call_abc", "get_weather", r#"{"t":72}"#);
assert_eq!(m.role, "tool");
assert_eq!(m.tool_call_id.as_deref(), Some("call_abc"));
assert_eq!(m.name.as_deref(), Some("get_weather"));
}
#[test]
fn serialises_all_roles() {
for (m, role) in [
(ChatMessage::system("x"), "system"),
(ChatMessage::user("x"), "user"),
(ChatMessage::assistant("x"), "assistant"),
(ChatMessage::tool_result("i", "f", "r"), "tool"),
] {
let v: serde_json::Value = serde_json::to_value(&m).expect("serialise");
assert_eq!(v["role"].as_str(), Some(role));
}
}
#[test]
fn cache_control_serialises_as_block() {
let mut m = ChatMessage::system("you are helpful");
m.cache_control = Some(CacheControl::ephemeral());
let v: serde_json::Value = serde_json::to_value(&m).expect("serialise");
assert_eq!(
v["content"],
serde_json::json!([{
"type": "text",
"text": "you are helpful",
"cache_control": {"type": "ephemeral"}
}])
);
}
#[test]
fn plain_content_when_no_cache_control() {
let m = ChatMessage::system("hi");
let v: serde_json::Value = serde_json::to_value(&m).expect("serialise");
assert_eq!(v["content"], serde_json::json!("hi"));
}
}