1use std::path::PathBuf;
2
3use serde::{Deserialize, Serialize};
4
5use crate::event::TurnId;
6
7#[derive(Debug, Clone, Serialize, PartialEq)]
8pub struct Message {
9 pub role: MessageRole,
10 pub parts: Vec<MessagePart>,
11 pub turn_id: TurnId,
12}
13
14impl<'de> Deserialize<'de> for Message {
15 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
16 where
17 D: serde::Deserializer<'de>,
18 {
19 #[derive(Deserialize)]
20 struct RawMessage {
21 role: MessageRole,
22 parts: Vec<MessagePart>,
23 turn_id: TurnId,
24 }
25
26 let raw = RawMessage::deserialize(deserializer)?;
27 let RawMessage {
28 role,
29 parts,
30 turn_id,
31 } = raw;
32 Ok(Self {
33 role,
34 parts: normalize_legacy_compact_summary(role, parts),
35 turn_id,
36 })
37 }
38}
39
40#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
41#[serde(rename_all = "snake_case")]
42pub enum MessageRole {
43 User,
44 Assistant,
45 System,
46 Tool,
47}
48
49#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
50#[serde(tag = "type", rename_all = "snake_case")]
51pub enum MessagePart {
52 CompactSummary {
53 summary: String,
54 seq_start: u64,
55 seq_end: u64,
56 count: usize,
57 },
58 Text {
59 text: String,
60 },
61 Thinking {
62 thinking: String,
63 #[serde(default, skip_serializing_if = "Option::is_none")]
64 signature: Option<String>,
65 },
66 Image {
67 source: ImageSource,
68 },
69 ToolUse {
70 id: String,
71 name: String,
72 input: serde_json::Value,
73 },
74 ToolResult {
75 tool_use_id: String,
76 content: String,
77 #[serde(default, skip_serializing_if = "core::ops::Not::not")]
78 is_error: bool,
79 },
80}
81
82#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
83pub struct ImageSource {
84 pub media_type: String,
85 pub data: ImageData,
86}
87
88#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
89#[serde(tag = "kind", rename_all = "snake_case")]
90pub enum ImageData {
91 Base64 { data: String },
92 Path { path: PathBuf },
93}
94
95impl Message {
96 pub fn user_text(turn_id: TurnId, text: impl Into<String>) -> Self {
97 Self {
98 role: MessageRole::User,
99 parts: vec![MessagePart::Text { text: text.into() }],
100 turn_id,
101 }
102 }
103
104 pub fn assistant_text(turn_id: TurnId, text: impl Into<String>) -> Self {
105 Self {
106 role: MessageRole::Assistant,
107 parts: vec![MessagePart::Text { text: text.into() }],
108 turn_id,
109 }
110 }
111
112 pub fn system_text(turn_id: TurnId, text: impl Into<String>) -> Self {
113 Self {
114 role: MessageRole::System,
115 parts: vec![MessagePart::Text { text: text.into() }],
116 turn_id,
117 }
118 }
119
120 pub fn system_compact_summary(
121 turn_id: TurnId,
122 summary: impl Into<String>,
123 seq_start: u64,
124 seq_end: u64,
125 count: usize,
126 ) -> Self {
127 Self {
128 role: MessageRole::System,
129 parts: vec![MessagePart::CompactSummary {
130 summary: summary.into(),
131 seq_start,
132 seq_end,
133 count,
134 }],
135 turn_id,
136 }
137 }
138
139 pub fn text_concat(&self) -> String {
140 let mut out = String::new();
141 for p in &self.parts {
142 match p {
143 MessagePart::Text { text } => out.push_str(text),
144 MessagePart::CompactSummary { summary, .. } => out.push_str(summary),
145 _ => {}
146 }
147 }
148 out
149 }
150
151 pub fn thinking_concat(&self) -> String {
152 let mut out = String::new();
153 for p in &self.parts {
154 if let MessagePart::Thinking { thinking, .. } = p {
155 out.push_str(thinking);
156 }
157 }
158 out
159 }
160
161 pub fn thinking_signature(&self) -> Option<String> {
162 self.parts.iter().rev().find_map(|p| {
163 if let MessagePart::Thinking { signature, .. } = p {
164 signature.clone()
165 } else {
166 None
167 }
168 })
169 }
170}
171
172impl MessageRole {
173 pub fn as_str(&self) -> &'static str {
174 match self {
175 MessageRole::User => "user",
176 MessageRole::Assistant => "assistant",
177 MessageRole::System => "system",
178 MessageRole::Tool => "tool",
179 }
180 }
181}
182
183fn normalize_legacy_compact_summary(
184 role: MessageRole,
185 parts: Vec<MessagePart>,
186) -> Vec<MessagePart> {
187 if role != MessageRole::System {
188 return parts;
189 }
190 if parts.len() != 1 {
191 return parts;
192 }
193 let MessagePart::Text { text } = &parts[0] else {
194 return parts;
195 };
196 let Some((summary, seq_start, seq_end, count)) = parse_legacy_compact_summary_text(text) else {
197 return parts;
198 };
199 vec![MessagePart::CompactSummary {
200 summary,
201 seq_start,
202 seq_end,
203 count,
204 }]
205}
206
207pub(crate) fn parse_legacy_compact_summary_text(text: &str) -> Option<(String, u64, u64, usize)> {
208 let start_marker = "[atman:compact ";
209 let start = text.rfind(start_marker)?;
210 let after = &text[start + start_marker.len()..];
211 let end = after.find(']')?;
212 let inner = &after[..end];
213 let mut seq_start = None;
214 let mut seq_end = None;
215 let mut count = None;
216 for token in inner.split_whitespace() {
217 let Some((k, v)) = token.split_once('=') else {
218 continue;
219 };
220 match k {
221 "seq_start" => seq_start = v.parse().ok(),
222 "seq_end" => seq_end = v.parse().ok(),
223 "count" => count = v.parse().ok(),
224 _ => {}
225 }
226 }
227 let summary = text[..start].trim_end().to_string();
228 Some((summary, seq_start?, seq_end?, count?))
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234
235 #[test]
236 fn user_text_roundtrips_via_serde_json() {
237 let msg = Message::user_text(TurnId::now(), "hello");
238 let s = serde_json::to_string(&msg).unwrap();
239 let back: Message = serde_json::from_str(&s).unwrap();
240 assert_eq!(msg, back);
241 }
242
243 #[test]
244 fn legacy_compact_summary_deserializes_to_structured_variant() {
245 let turn_id = TurnId::now();
246 let msg = Message {
247 role: MessageRole::System,
248 parts: vec![MessagePart::Text {
249 text: "handoff\n\n[atman:compact seq_start=2 seq_end=7 count=6]".into(),
250 }],
251 turn_id,
252 };
253 let s = serde_json::to_string(&msg).unwrap();
254 let back: Message = serde_json::from_str(&s).unwrap();
255 assert!(matches!(
256 back.parts.as_slice(),
257 [MessagePart::CompactSummary { .. }]
258 ));
259 assert_eq!(back.text_concat(), "handoff");
260 }
261
262 #[test]
263 fn text_concat_skips_non_text_parts() {
264 let msg = Message {
265 role: MessageRole::User,
266 parts: vec![
267 MessagePart::Text { text: "a ".into() },
268 MessagePart::Image {
269 source: ImageSource {
270 media_type: "image/png".into(),
271 data: ImageData::Path {
272 path: PathBuf::from("/tmp/x.png"),
273 },
274 },
275 },
276 MessagePart::Text { text: "b".into() },
277 ],
278 turn_id: TurnId::now(),
279 };
280 assert_eq!(msg.text_concat(), "a b");
281 }
282
283 #[test]
284 fn tool_result_is_error_defaults_to_false_and_skips_serialize_when_false() {
285 let msg = Message {
286 role: MessageRole::Tool,
287 parts: vec![MessagePart::ToolResult {
288 tool_use_id: "toolu_1".into(),
289 content: "ok".into(),
290 is_error: false,
291 }],
292 turn_id: TurnId::now(),
293 };
294 let s = serde_json::to_string(&msg).unwrap();
295 assert!(!s.contains("is_error"), "should skip when false: {s}");
296
297 let err_msg = Message {
298 role: MessageRole::Tool,
299 parts: vec![MessagePart::ToolResult {
300 tool_use_id: "toolu_1".into(),
301 content: "nope".into(),
302 is_error: true,
303 }],
304 turn_id: TurnId::now(),
305 };
306 let s = serde_json::to_string(&err_msg).unwrap();
307 assert!(s.contains("\"is_error\":true"), "{s}");
308 }
309
310 #[test]
311 fn role_as_str_matches_wire_format() {
312 assert_eq!(MessageRole::User.as_str(), "user");
313 assert_eq!(MessageRole::Assistant.as_str(), "assistant");
314 assert_eq!(MessageRole::System.as_str(), "system");
315 assert_eq!(MessageRole::Tool.as_str(), "tool");
316 }
317}