From a7de1974edc0ad1240d88652bae46894198edfb9 Mon Sep 17 00:00:00 2001 From: David Hagerty Date: Tue, 20 Jan 2026 21:00:11 -0500 Subject: [PATCH] feat: extend Message model with Role::Tool and tool_call_id for provider compatibility --- src/llm/anthropic.rs | 7 ++++-- src/llm/mod.rs | 37 ++++++++++++++++++++++++++++++ src/planning/mod.rs | 27 ++++++++-------------- src/ralph/mod.rs | 19 +++++----------- tests/llm_test.rs | 22 ++++-------------- tests/message_test.rs | 53 +++++++++++++++++++++++++++++++++++++++++++ 6 files changed, 116 insertions(+), 49 deletions(-) create mode 100644 tests/message_test.rs diff --git a/src/llm/anthropic.rs b/src/llm/anthropic.rs index 33cb2ef..5ab82af 100644 --- a/src/llm/anthropic.rs +++ b/src/llm/anthropic.rs @@ -78,15 +78,18 @@ impl AnthropicClient { .find(|m| m.role == Role::System) .map(|m| m.content.clone()); - // Filter out system messages from messages array + // Filter out system messages and tool messages from messages array + // (Tool messages need special handling in Anthropic - for now we skip them + // and rely on the calling code to format tool results as user messages) let anthropic_messages: Vec = messages .iter() - .filter(|m| m.role != Role::System) + .filter(|m| m.role != Role::System && m.role != Role::Tool) .map(|m| AnthropicMessage { role: match m.role { Role::User => "user".to_string(), Role::Assistant => "assistant".to_string(), Role::System => unreachable!("System messages filtered out"), + Role::Tool => unreachable!("Tool messages filtered out"), }, content: m.content.clone(), }) diff --git a/src/llm/mod.rs b/src/llm/mod.rs index da9e49c..ec56f7b 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -9,6 +9,7 @@ pub enum Role { User, Assistant, System, + Tool, } /// A message in the conversation @@ -16,6 +17,42 @@ pub enum Role { pub struct Message { pub role: Role, pub content: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, +} + +impl Message { + pub fn user(content: impl Into) -> Self { + Self { + role: Role::User, + content: content.into(), + tool_call_id: None, + } + } + + pub fn assistant(content: impl Into) -> Self { + Self { + role: Role::Assistant, + content: content.into(), + tool_call_id: None, + } + } + + pub fn system(content: impl Into) -> Self { + Self { + role: Role::System, + content: content.into(), + tool_call_id: None, + } + } + + pub fn tool_result(tool_call_id: impl Into, content: impl Into) -> Self { + Self { + role: Role::Tool, + content: content.into(), + tool_call_id: Some(tool_call_id.into()), + } + } } /// Definition of a tool that can be called by the LLM diff --git a/src/planning/mod.rs b/src/planning/mod.rs index 838595a..9ece9ad 100644 --- a/src/planning/mod.rs +++ b/src/planning/mod.rs @@ -1,6 +1,6 @@ use crate::config::{Config, LlmProvider}; use crate::llm::anthropic::AnthropicClient; -use crate::llm::{LlmClient, Message, ResponseContent, Role}; +use crate::llm::{LlmClient, Message, ResponseContent}; use crate::security::{SecurityValidator, permission::CliPermissionHandler}; use crate::tools::{ ToolRegistry, @@ -69,10 +69,7 @@ impl PlanningAgent { ))); // Initialize conversation with system message - let system_message = Message { - role: Role::System, - content: PLANNING_SYSTEM_PROMPT.to_string(), - }; + let system_message = Message::system(PLANNING_SYSTEM_PROMPT); Ok(Self { client, @@ -112,10 +109,7 @@ impl PlanningAgent { } // Add user message to conversation - self.conversation.push(Message { - role: Role::User, - content: input.to_string(), - }); + self.conversation.push(Message::user(input)); // Process the conversation turn match self.process_turn().await { @@ -151,10 +145,7 @@ impl PlanningAgent { match response.content { ResponseContent::Text(text) => { // Add assistant response to conversation - self.conversation.push(Message { - role: Role::Assistant, - content: text.clone(), - }); + self.conversation.push(Message::assistant(text.clone())); // Print assistant response println!("\nAssistant: {}\n", text); @@ -192,10 +183,12 @@ impl PlanningAgent { }; // Add tool result to conversation - self.conversation.push(Message { - role: Role::User, - content: format!("Tool result for {}:\n{}", tool_call.name, output), - }); + // Note: Using User role for now as Anthropic expects tool results + // in user messages. Future OpenAI provider will use Message::tool_result() + self.conversation.push(Message::user(format!( + "Tool result for {}:\n{}", + tool_call.name, output + ))); } // Continue the loop to get the next LLM response diff --git a/src/ralph/mod.rs b/src/ralph/mod.rs index cc8f310..c50b300 100644 --- a/src/ralph/mod.rs +++ b/src/ralph/mod.rs @@ -1,6 +1,6 @@ use crate::config::{Config, LlmProvider}; use crate::llm::anthropic::AnthropicClient; -use crate::llm::{LlmClient, Message, ResponseContent, Role}; +use crate::llm::{LlmClient, Message, ResponseContent}; use crate::security::{SecurityValidator, permission::CliPermissionHandler}; use crate::spec::{Spec, TaskStatus}; use crate::tools::ToolRegistry; @@ -178,10 +178,7 @@ impl RalphLoop { let context = self.build_context(task_id)?; let tool_definitions = self.tools.definitions(); - let mut messages = vec![Message { - role: Role::User, - content: context, - }]; + let mut messages = vec![Message::user(context)]; // Agentic loop - continue until we get a completion signal let max_turns = 50; @@ -207,10 +204,7 @@ impl RalphLoop { } // Add assistant response to conversation - messages.push(Message { - role: Role::Assistant, - content: text, - }); + messages.push(Message::assistant(text)); } ResponseContent::ToolCalls(tool_calls) => { println!(" Executing {} tool calls", tool_calls.len()); @@ -234,11 +228,10 @@ impl RalphLoop { } // Add tool results as user message + // Note: Using User role for now as Anthropic expects tool results + // in user messages. Future OpenAI provider will use Message::tool_result() let results_text = results.join("\n\n"); - messages.push(Message { - role: Role::User, - content: results_text, - }); + messages.push(Message::user(results_text)); } } } diff --git a/tests/llm_test.rs b/tests/llm_test.rs index 640191d..5d72901 100644 --- a/tests/llm_test.rs +++ b/tests/llm_test.rs @@ -1,12 +1,9 @@ -use rustagent::llm::{Message, Role, ToolDefinition}; +use rustagent::llm::{Message, ToolDefinition}; use serde_json::json; #[test] fn test_message_serialization() { - let msg = Message { - role: Role::User, - content: "Hello".to_string(), - }; + let msg = Message::user("Hello"); let json = serde_json::to_value(&msg).unwrap(); assert_eq!(json["role"], "user"); @@ -37,10 +34,7 @@ async fn test_anthropic_message_format() { // This test validates request structure, doesn't actually call API let client = AnthropicClient::new("test-key".to_string(), "claude-sonnet-4".to_string(), 4096); - let messages = vec![Message { - role: Role::User, - content: "Hello".to_string(), - }]; + let messages = vec![Message::user("Hello")]; // We'll test this by mocking in future, for now just construct assert!(client.format_request(&messages, &[]).is_ok()); @@ -51,14 +45,8 @@ fn test_format_request_with_system_message() { let client = AnthropicClient::new("test-key".to_string(), "claude-sonnet-4".to_string(), 4096); let messages = vec![ - Message { - role: Role::System, - content: "You are a helpful assistant".to_string(), - }, - Message { - role: Role::User, - content: "Hello".to_string(), - }, + Message::system("You are a helpful assistant"), + Message::user("Hello"), ]; let request = client.format_request(&messages, &[]).unwrap(); diff --git a/tests/message_test.rs b/tests/message_test.rs new file mode 100644 index 0000000..78958d2 --- /dev/null +++ b/tests/message_test.rs @@ -0,0 +1,53 @@ +use rustagent::llm::{Message, Role}; + +#[test] +fn test_tool_message_creation() { + let msg = Message::tool_result("call_123", "File contents here"); + assert_eq!(msg.role, Role::Tool); + assert_eq!(msg.tool_call_id, Some("call_123".to_string())); + assert_eq!(msg.content, "File contents here"); +} + +#[test] +fn test_user_message_has_no_tool_call_id() { + let msg = Message::user("Hello"); + assert_eq!(msg.role, Role::User); + assert!(msg.tool_call_id.is_none()); +} + +#[test] +fn test_assistant_message_constructor() { + let msg = Message::assistant("I can help with that"); + assert_eq!(msg.role, Role::Assistant); + assert_eq!(msg.content, "I can help with that"); + assert!(msg.tool_call_id.is_none()); +} + +#[test] +fn test_system_message_constructor() { + let msg = Message::system("You are a helpful assistant"); + assert_eq!(msg.role, Role::System); + assert_eq!(msg.content, "You are a helpful assistant"); + assert!(msg.tool_call_id.is_none()); +} + +#[test] +fn test_tool_message_serialization() { + let msg = Message::tool_result("call_abc", "result data"); + let json = serde_json::to_value(&msg).unwrap(); + + assert_eq!(json["role"], "tool"); + assert_eq!(json["tool_call_id"], "call_abc"); + assert_eq!(json["content"], "result data"); +} + +#[test] +fn test_user_message_serialization_skips_none_tool_call_id() { + let msg = Message::user("Hello"); + let json = serde_json::to_value(&msg).unwrap(); + + assert_eq!(json["role"], "user"); + assert_eq!(json["content"], "Hello"); + // tool_call_id should not be present (skip_serializing_if = "Option::is_none") + assert!(json.get("tool_call_id").is_none()); +} -- 2.51.2