Migrate Rig tool flow to domain runtime
This commit is contained in:
@@ -180,12 +180,14 @@ where
|
||||
}
|
||||
}
|
||||
Ok(StreamedAssistantContent::ToolCall { tool_call, .. }) => {
|
||||
yield Ok(AgentEvent::ToolProposed {
|
||||
call: ToolCall {
|
||||
yield Ok(AgentEvent::Tool {
|
||||
event: galaxy_agent_core::ToolEvent::Proposed {
|
||||
call: ToolCall {
|
||||
id: tool_call.id,
|
||||
name: tool_call.function.name,
|
||||
arguments: tool_call.function.arguments,
|
||||
},
|
||||
},
|
||||
});
|
||||
}
|
||||
Ok(StreamedAssistantContent::ToolCallDelta { .. }) => {
|
||||
@@ -298,10 +300,10 @@ fn user_content(content: MessageContent) -> Result<OneOrMany<UserContent>, Agent
|
||||
MessageContent::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
..
|
||||
is_error,
|
||||
} => vec![UserContent::tool_result(
|
||||
tool_use_id,
|
||||
OneOrMany::one(ToolResultContent::text(content)),
|
||||
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
|
||||
)],
|
||||
MessageContent::MultiPart(parts) => parts
|
||||
.into_iter()
|
||||
@@ -344,10 +346,10 @@ fn convert_user_part(part: ContentPart) -> Result<UserContent, AgentError> {
|
||||
ContentPart::ToolResult {
|
||||
tool_use_id,
|
||||
content,
|
||||
..
|
||||
is_error,
|
||||
} => Ok(UserContent::tool_result(
|
||||
tool_use_id,
|
||||
OneOrMany::one(ToolResultContent::text(content)),
|
||||
OneOrMany::one(ToolResultContent::text(tool_result_text(content, is_error))),
|
||||
)),
|
||||
ContentPart::ToolUse { .. } => Err(invalid_role("tool use", "user")),
|
||||
}
|
||||
@@ -371,6 +373,14 @@ fn convert_assistant_part(part: ContentPart) -> Result<AssistantContent, AgentEr
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_result_text(content: String, is_error: bool) -> String {
|
||||
if is_error {
|
||||
format!("[ERROR] {content}")
|
||||
} else {
|
||||
content
|
||||
}
|
||||
}
|
||||
|
||||
fn one_or_many<T: Clone>(parts: Vec<T>, role: &str) -> Result<OneOrMany<T>, AgentError> {
|
||||
OneOrMany::many(parts).map_err(|_| {
|
||||
AgentError::new(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use futures::StreamExt;
|
||||
use galaxy_agent_core::{
|
||||
AgentEvent, AgentRuntime, ConversationMessage, MessageContent, MessageRole,
|
||||
AgentEvent, AgentRuntime, ConversationMessage, MessageContent, MessageRole, ToolEvent,
|
||||
};
|
||||
use rig_core::client::CompletionClient;
|
||||
use rig_core::providers::openai;
|
||||
@@ -152,6 +152,71 @@ async fn usage_at_the_requested_limit_maps_to_max_tokens() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rig_stream_maps_complete_tool_call_without_executing_it() {
|
||||
let http_client = MockStreamingClient {
|
||||
sse_bytes: sse(&[
|
||||
r#"{"id":"cmpl-1","model":"test-model","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call-1","type":"function","function":{"name":"read_files","arguments":"{\"files\":[\"Cargo.toml\"]}"}}]},"finish_reason":null}],"usage":null}"#,
|
||||
r#"{"id":"cmpl-1","model":"test-model","choices":[{"delta":{"tool_calls":[]},"finish_reason":"tool_calls"}],"usage":null}"#,
|
||||
r#"{"choices":[],"usage":{"prompt_tokens":8,"completion_tokens":4,"total_tokens":12}}"#,
|
||||
"[DONE]",
|
||||
]),
|
||||
};
|
||||
let client = openai::CompletionsClient::builder()
|
||||
.api_key("test-key")
|
||||
.base_url("http://localhost/v1")
|
||||
.http_client(http_client)
|
||||
.build()
|
||||
.unwrap();
|
||||
let model = client.completion_model("test-model");
|
||||
let (_, control) = galaxy_agent_core::turn_control();
|
||||
let mut request = text_request();
|
||||
request.tools.push(galaxy_agent_core::ToolDefinition {
|
||||
name: "read_files".to_string(),
|
||||
description: "Read files".to_string(),
|
||||
input_schema: serde_json::json!({"type": "object"}),
|
||||
});
|
||||
|
||||
let events = start_model_turn(model, request, control, None, true)
|
||||
.await
|
||||
.unwrap()
|
||||
.collect::<Vec<_>>()
|
||||
.await
|
||||
.into_iter()
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.unwrap();
|
||||
|
||||
assert!(events.iter().any(|event| {
|
||||
matches!(
|
||||
event,
|
||||
AgentEvent::Tool {
|
||||
event: ToolEvent::Proposed { call },
|
||||
}
|
||||
if call.id == "call-1"
|
||||
&& call.name == "read_files"
|
||||
&& call.arguments == serde_json::json!({"files": ["Cargo.toml"]})
|
||||
)
|
||||
}));
|
||||
assert_eq!(
|
||||
events.last(),
|
||||
Some(&AgentEvent::TurnStopped {
|
||||
reason: StopReason::Completed,
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
events
|
||||
.iter()
|
||||
.filter(|event| matches!(
|
||||
event,
|
||||
AgentEvent::Tool {
|
||||
event: ToolEvent::Started { .. } | ToolEvent::Completed { .. },
|
||||
}
|
||||
))
|
||||
.count(),
|
||||
0
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_conversion_preserves_history_tools_and_limits() {
|
||||
let mut request = text_request();
|
||||
@@ -175,6 +240,57 @@ fn request_conversion_preserves_history_tools_and_limits() {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_conversion_preserves_tool_call_and_denied_result_for_the_next_turn() {
|
||||
let request = TurnRequest::new(
|
||||
"test-model",
|
||||
vec![
|
||||
ConversationMessage {
|
||||
role: MessageRole::Assistant,
|
||||
content: MessageContent::ToolUse {
|
||||
tool_use_id: "call-1".to_string(),
|
||||
name: "run_shell_command".to_string(),
|
||||
input: serde_json::json!({"command": "cargo test"}),
|
||||
},
|
||||
},
|
||||
ConversationMessage {
|
||||
role: MessageRole::User,
|
||||
content: MessageContent::ToolResult {
|
||||
tool_use_id: "call-1".to_string(),
|
||||
content: "Command not executed — permission denied.".to_string(),
|
||||
is_error: true,
|
||||
},
|
||||
},
|
||||
],
|
||||
);
|
||||
|
||||
let converted = build_completion_request(request, None, true).unwrap();
|
||||
let messages = converted.chat_history.iter().collect::<Vec<_>>();
|
||||
let Message::Assistant { content, .. } = messages[0] else {
|
||||
panic!("expected assistant tool call");
|
||||
};
|
||||
let Some(AssistantContent::ToolCall(call)) = content.iter().next() else {
|
||||
panic!("expected assistant tool call content");
|
||||
};
|
||||
assert_eq!(call.id, "call-1");
|
||||
assert_eq!(call.function.name, "run_shell_command");
|
||||
|
||||
let Message::User { content } = messages[1] else {
|
||||
panic!("expected user tool result");
|
||||
};
|
||||
let Some(UserContent::ToolResult(result)) = content.iter().next() else {
|
||||
panic!("expected user tool result content");
|
||||
};
|
||||
assert_eq!(result.id, "call-1");
|
||||
let Some(ToolResultContent::Text(text)) = result.content.iter().next() else {
|
||||
panic!("expected text tool result");
|
||||
};
|
||||
assert_eq!(
|
||||
text.text,
|
||||
"[ERROR] Command not executed — permission denied."
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_conversion_places_system_prompt_in_user_message_when_system_role_is_unsupported() {
|
||||
let mut request = text_request();
|
||||
|
||||
Reference in New Issue
Block a user