1use std::future::Future;
7use std::pin::Pin;
8
9pub 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#[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#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
45#[serde(untagged)]
46pub enum MessageContent {
47 Text(String),
48 Blocks(Vec<ContentBlock>),
49}
50
51#[derive(Debug, Clone, PartialEq)]
53pub enum StopReason {
54 EndTurn,
55 ToolUse,
56 MaxTokens,
57 StopSequence,
58 Other(String),
59}
60
61#[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 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 pub fn has_tool_use(&self) -> bool {
86 self.content
87 .iter()
88 .any(|block| matches!(block, ContentBlock::ToolUse { .. }))
89 }
90}
91
92#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
97#[serde(tag = "type", rename_all = "snake_case")]
98pub enum MessageSource {
99 Human { channel: String, sender: String },
101 ToolResult { tool_use_id: String },
103 ScheduledTask { task_name: String },
105 System,
107 Assistant,
109}
110
111#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
113pub struct Message {
114 pub role: Role,
115 pub content: MessageContent,
116 #[serde(default, skip_serializing_if = "Option::is_none")]
118 pub source: Option<MessageSource>,
119}
120
121#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
123#[serde(rename_all = "lowercase")]
124pub enum Role {
125 User,
126 Assistant,
127}
128
129pub trait LmProvider: Send + Sync {
131 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 fn name(&self) -> &str;
143
144 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")); }
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 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 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}