Skip to main content

atman_runtime/
message.rs

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}