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, 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}