use std::borrow::Cow;
use crate::{prompt, stream::MessageDelta, Model};
use serde::{Deserialize, Serialize};
#[derive(Debug, Serialize, Deserialize, derive_more::Display)]
#[cfg_attr(any(feature = "partial-eq", test), derive(PartialEq))]
#[display("{}", message)]
pub struct Message<'a> {
pub id: Cow<'a, str>,
#[serde(flatten)]
pub message: prompt::Message<'a>,
pub model: Model,
pub stop_reason: Option<StopReason>,
pub stop_sequence: Option<Cow<'a, str>>,
pub usage: Usage,
}
impl Message<'_> {
pub fn apply_delta(&mut self, delta: MessageDelta) {
self.stop_reason = delta.stop_reason;
self.stop_sequence = delta.stop_sequence;
if let Some(usage) = delta.usage {
self.usage = usage;
}
}
pub fn tool_use(&self) -> Option<&crate::tool::Use> {
if !matches!(self.stop_reason, Some(StopReason::ToolUse)) {
return None;
}
self.message.content.last()?.tool_use()
}
pub fn into_static(self) -> Message<'static> {
Message {
id: Cow::Owned(self.id.into_owned()),
message: self.message.into_static(),
model: self.model,
stop_reason: self.stop_reason,
stop_sequence: self
.stop_sequence
.map(|s| Cow::Owned(s.into_owned())),
usage: self.usage,
}
}
}
#[derive(Debug, Serialize, Deserialize)]
#[cfg_attr(any(feature = "partial-eq", test), derive(PartialEq))]
#[serde(rename_all = "snake_case")]
pub enum StopReason {
EndTurn,
MaxTokens,
StopSequence,
ToolUse,
}
#[derive(Debug, Serialize, Deserialize, Default)]
#[cfg_attr(any(feature = "partial-eq", test), derive(PartialEq))]
pub struct Usage {
pub input_tokens: u64,
#[cfg(feature = "prompt-caching")]
pub cache_creation_input_tokens: Option<u64>,
#[cfg(feature = "prompt-caching")]
pub cache_read_input_tokens: Option<u64>,
pub output_tokens: u64,
}
#[cfg(feature = "markdown")]
impl crate::markdown::ToMarkdown for Message<'_> {
fn markdown_events_custom<'a>(
&'a self,
options: crate::markdown::Options,
) -> Box<dyn Iterator<Item = pulldown_cmark::Event<'a>> + 'a> {
self.message.markdown_events_custom(options)
}
}
#[cfg(test)]
mod tests {
use super::*;
pub const RESPONSE_JSON: &str = r#"{
"content": [
{
"text": "Hi! My name is Claude.",
"type": "text"
}
],
"id": "msg_013Zva2CMHLNnXjNJJKqJ2EF",
"model": "claude-3-5-sonnet-20240620",
"role": "assistant",
"stop_reason": "end_turn",
"stop_sequence": null,
"type": "message",
"usage": {
"input_tokens": 2095,
"output_tokens": 503
}
}"#;
#[test]
fn deserialize_response_message() {
let message: Message = serde_json::from_str(RESPONSE_JSON).unwrap();
assert_eq!(message.message.content.len(), 22);
assert_eq!(message.id, "msg_013Zva2CMHLNnXjNJJKqJ2EF");
assert_eq!(message.model, crate::Model::Sonnet35_20240620);
assert!(matches!(message.stop_reason, Some(StopReason::EndTurn)));
assert_eq!(message.stop_sequence, None);
assert_eq!(message.usage.input_tokens, 2095);
assert_eq!(message.usage.output_tokens, 503);
}
#[test]
fn test_apply_delta() {
let mut message: Message = serde_json::from_str(RESPONSE_JSON).unwrap();
let delta = MessageDelta {
stop_reason: Some(StopReason::MaxTokens),
stop_sequence: Some("sequence".into()),
usage: Some(Usage {
input_tokens: 100,
output_tokens: 200,
..Default::default()
}),
};
message.apply_delta(delta);
assert_eq!(message.stop_reason, Some(StopReason::MaxTokens));
assert_eq!(message.stop_sequence, Some("sequence".into()));
assert_eq!(message.usage.input_tokens, 100);
assert_eq!(message.usage.output_tokens, 200);
}
#[test]
fn test_tool_use() {
let mut message: Message = serde_json::from_str(RESPONSE_JSON).unwrap();
assert!(message.tool_use().is_none());
message.stop_reason = Some(StopReason::ToolUse);
assert!(message.tool_use().is_none());
message.message.content.push(crate::tool::Use {
id: "id".into(),
name: "name".into(),
input: serde_json::json!({}),
#[cfg(feature = "prompt-caching")]
cache_control: None,
});
assert!(message.tool_use().is_some());
}
#[test]
fn test_into_static() {
let message: Message = serde_json::from_str(RESPONSE_JSON).unwrap();
let static_message = message.into_static();
assert_eq!(static_message.id, "msg_013Zva2CMHLNnXjNJJKqJ2EF");
assert_eq!(static_message.model, crate::Model::Sonnet35_20240620);
assert!(matches!(
static_message.stop_reason,
Some(StopReason::EndTurn)
));
assert_eq!(static_message.stop_sequence, None);
assert_eq!(static_message.usage.input_tokens, 2095);
assert_eq!(static_message.usage.output_tokens, 503);
}
#[test]
#[cfg(feature = "markdown")]
fn test_markdown() {
use crate::markdown::ToMarkdown;
let message = Message {
id: "id".into(),
message: prompt::Message {
role: prompt::message::Role::User,
content: prompt::message::Content::SinglePart(
"Hello, **world**!".into(),
),
},
model: crate::Model::Sonnet35,
stop_reason: None,
stop_sequence: None,
usage: Usage {
input_tokens: 1,
#[cfg(feature = "prompt-caching")]
cache_creation_input_tokens: Some(2),
#[cfg(feature = "prompt-caching")]
cache_read_input_tokens: Some(3),
output_tokens: 4,
},
};
let expected = "### User\n\nHello, **world**!";
let markdown = message.markdown();
assert_eq!(markdown.as_ref(), expected);
}
}