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