Skip to main content

pulse_system_types/
llm.rs

1//! LLM interaction types — the shared contract for model-agnostic design.
2//!
3//! These types define the conversation model, content blocks, and provider trait
4//! that pulse-null and its plugins use to interact with language models.
5
6use std::future::Future;
7use std::pin::Pin;
8
9/// Result type for LLM provider invocations
10pub type LlmResult<'a> = Pin<
11    Box<
12        dyn Future<Output = Result<LlmResponse, Box<dyn std::error::Error + Send + Sync>>>
13            + Send
14            + 'a,
15    >,
16>;
17
18/// A content block in a message or response.
19/// Claude API uses tagged unions — each block has a "type" field.
20#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
21#[serde(tag = "type")]
22pub enum ContentBlock {
23    #[serde(rename = "text")]
24    Text { text: String },
25
26    #[serde(rename = "tool_use")]
27    ToolUse {
28        id: String,
29        name: String,
30        input: serde_json::Value,
31    },
32
33    #[serde(rename = "tool_result")]
34    ToolResult {
35        tool_use_id: String,
36        content: String,
37        #[serde(skip_serializing_if = "Option::is_none")]
38        is_error: Option<bool>,
39    },
40}
41
42/// Message content can be a simple string or structured content blocks.
43/// The Claude API accepts both formats.
44#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
45#[serde(untagged)]
46pub enum MessageContent {
47    Text(String),
48    Blocks(Vec<ContentBlock>),
49}
50
51/// Why the model stopped generating.
52#[derive(Debug, Clone, PartialEq)]
53pub enum StopReason {
54    EndTurn,
55    ToolUse,
56    MaxTokens,
57    StopSequence,
58    Other(String),
59}
60
61/// Response from an LLM invocation
62#[derive(Debug, Clone)]
63pub struct LlmResponse {
64    pub content: Vec<ContentBlock>,
65    pub stop_reason: StopReason,
66    pub model: String,
67    pub input_tokens: Option<u32>,
68    pub output_tokens: Option<u32>,
69}
70
71impl LlmResponse {
72    /// Extract all text content from the response, concatenated.
73    pub fn text(&self) -> String {
74        self.content
75            .iter()
76            .filter_map(|block| match block {
77                ContentBlock::Text { text } => Some(text.as_str()),
78                _ => None,
79            })
80            .collect::<Vec<_>>()
81            .join("")
82    }
83
84    /// Check if the response contains any tool_use blocks.
85    pub fn has_tool_use(&self) -> bool {
86        self.content
87            .iter()
88            .any(|block| matches!(block, ContentBlock::ToolUse { .. }))
89    }
90}
91
92/// Origin of a message — distinguishes real human input from system-generated messages.
93///
94/// Used by the hallucination guard to detect self-conversation loops: if no message
95/// with `MessageSource::Human` has arrived in N rounds, the pulse is talking to itself.
96#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
97#[serde(tag = "type", rename_all = "snake_case")]
98pub enum MessageSource {
99    /// Message from a real human via an input channel (HTTP, Discord, voice, REPL).
100    Human { channel: String, sender: String },
101    /// Tool execution result injected by the tool loop.
102    ToolResult { tool_use_id: String },
103    /// Message generated by a scheduled/autonomous task.
104    ScheduledTask { task_name: String },
105    /// System-generated message (compaction summary, context injection, etc.).
106    System,
107    /// LLM-generated assistant response.
108    Assistant,
109}
110
111/// A message in a conversation
112#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
113pub struct Message {
114    pub role: Role,
115    pub content: MessageContent,
116    /// Origin of this message. None for legacy messages loaded from disk.
117    #[serde(default, skip_serializing_if = "Option::is_none")]
118    pub source: Option<MessageSource>,
119}
120
121/// Conversation role
122#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
123#[serde(rename_all = "lowercase")]
124pub enum Role {
125    User,
126    Assistant,
127}
128
129/// Trait for LLM providers — the core abstraction for model-agnostic design
130pub trait LmProvider: Send + Sync {
131    /// Send a message and get a response.
132    /// `tools` is an optional slice of tool definitions (JSON objects).
133    fn invoke(
134        &self,
135        system_prompt: &str,
136        messages: &[Message],
137        max_tokens: u32,
138        tools: Option<&[serde_json::Value]>,
139    ) -> LlmResult<'_>;
140
141    /// Provider name
142    fn name(&self) -> &str;
143
144    /// Whether this provider supports tool use
145    fn supports_tools(&self) -> bool {
146        false
147    }
148}
149
150#[cfg(test)]
151mod tests {
152    use super::*;
153
154    #[test]
155    fn content_block_text_serializes() {
156        let block = ContentBlock::Text {
157            text: "hello".into(),
158        };
159        let json = serde_json::to_string(&block).unwrap();
160        assert!(json.contains("\"type\":\"text\""));
161        assert!(json.contains("\"text\":\"hello\""));
162    }
163
164    #[test]
165    fn content_block_tool_use_serializes() {
166        let block = ContentBlock::ToolUse {
167            id: "t1".into(),
168            name: "file_read".into(),
169            input: serde_json::json!({"path": "/tmp/test"}),
170        };
171        let json = serde_json::to_string(&block).unwrap();
172        assert!(json.contains("\"type\":\"tool_use\""));
173        assert!(json.contains("\"name\":\"file_read\""));
174    }
175
176    #[test]
177    fn content_block_tool_result_serializes() {
178        let block = ContentBlock::ToolResult {
179            tool_use_id: "t1".into(),
180            content: "file contents".into(),
181            is_error: None,
182        };
183        let json = serde_json::to_string(&block).unwrap();
184        assert!(json.contains("\"type\":\"tool_result\""));
185        assert!(!json.contains("is_error")); // skipped when None
186    }
187
188    #[test]
189    fn message_content_text_roundtrip() {
190        let content = MessageContent::Text("hello".into());
191        let json = serde_json::to_string(&content).unwrap();
192        let back: MessageContent = serde_json::from_str(&json).unwrap();
193        matches!(back, MessageContent::Text(s) if s == "hello");
194    }
195
196    #[test]
197    fn message_content_blocks_roundtrip() {
198        let content = MessageContent::Blocks(vec![ContentBlock::Text {
199            text: "hello".into(),
200        }]);
201        let json = serde_json::to_string(&content).unwrap();
202        let back: MessageContent = serde_json::from_str(&json).unwrap();
203        matches!(back, MessageContent::Blocks(b) if b.len() == 1);
204    }
205
206    #[test]
207    fn llm_response_text_extraction() {
208        let response = LlmResponse {
209            content: vec![
210                ContentBlock::Text {
211                    text: "hello ".into(),
212                },
213                ContentBlock::ToolUse {
214                    id: "t1".into(),
215                    name: "test".into(),
216                    input: serde_json::Value::Null,
217                },
218                ContentBlock::Text {
219                    text: "world".into(),
220                },
221            ],
222            stop_reason: StopReason::EndTurn,
223            model: "test".into(),
224            input_tokens: None,
225            output_tokens: None,
226        };
227        assert_eq!(response.text(), "hello world");
228        assert!(response.has_tool_use());
229    }
230
231    #[test]
232    fn llm_response_no_tool_use() {
233        let response = LlmResponse {
234            content: vec![ContentBlock::Text {
235                text: "just text".into(),
236            }],
237            stop_reason: StopReason::EndTurn,
238            model: "test".into(),
239            input_tokens: None,
240            output_tokens: None,
241        };
242        assert!(!response.has_tool_use());
243    }
244
245    #[test]
246    fn stop_reason_equality() {
247        assert_eq!(StopReason::EndTurn, StopReason::EndTurn);
248        assert_ne!(StopReason::EndTurn, StopReason::ToolUse);
249        assert_eq!(
250            StopReason::Other("custom".into()),
251            StopReason::Other("custom".into())
252        );
253    }
254
255    #[test]
256    fn message_serializes() {
257        let msg = Message {
258            role: Role::User,
259            content: MessageContent::Text("hi".into()),
260            source: None,
261        };
262        let json = serde_json::to_string(&msg).unwrap();
263        assert!(json.contains("\"role\":\"user\""));
264        // source: None should be omitted from JSON
265        assert!(!json.contains("source"));
266    }
267
268    #[test]
269    fn message_source_human_roundtrip() {
270        let msg = Message {
271            role: Role::User,
272            content: MessageContent::Text("hello".into()),
273            source: Some(MessageSource::Human {
274                channel: "chat".into(),
275                sender: "dani".into(),
276            }),
277        };
278        let json = serde_json::to_string(&msg).unwrap();
279        assert!(json.contains("\"source\""));
280        assert!(json.contains("\"human\""));
281        let back: Message = serde_json::from_str(&json).unwrap();
282        assert_eq!(
283            back.source,
284            Some(MessageSource::Human {
285                channel: "chat".into(),
286                sender: "dani".into(),
287            })
288        );
289    }
290
291    #[test]
292    fn message_without_source_deserializes() {
293        // Legacy messages without a source field should deserialize with source: None
294        let json = r#"{"role":"user","content":"hello"}"#;
295        let msg: Message = serde_json::from_str(json).unwrap();
296        assert!(msg.source.is_none());
297    }
298
299    #[test]
300    fn message_source_tool_result_serializes() {
301        let source = MessageSource::ToolResult {
302            tool_use_id: "t1".into(),
303        };
304        let json = serde_json::to_string(&source).unwrap();
305        assert!(json.contains("\"tool_result\""));
306        assert!(json.contains("\"tool_use_id\""));
307    }
308
309    #[test]
310    fn message_source_scheduled_task_serializes() {
311        let source = MessageSource::ScheduledTask {
312            task_name: "morning_orientation".into(),
313        };
314        let json = serde_json::to_string(&source).unwrap();
315        assert!(json.contains("\"scheduled_task\""));
316        assert!(json.contains("\"task_name\""));
317    }
318
319    #[test]
320    fn message_source_system_serializes() {
321        let source = MessageSource::System;
322        let json = serde_json::to_string(&source).unwrap();
323        assert!(json.contains("\"system\""));
324    }
325}