feat: Redesign the tool system

This commit is contained in:
2026-01-10 00:31:14 +09:00
parent 5691b09fc8
commit 16fda38039
13 changed files with 897 additions and 396 deletions
+48 -49
View File
@@ -9,11 +9,11 @@ use std::time::{Duration, Instant};
use async_trait::async_trait;
use llm_worker::Worker;
use llm_worker::hook::{
AfterToolCall, AfterToolCallResult, BeforeToolCall, BeforeToolCallResult, Hook, HookError,
ToolCall, ToolResult,
Hook, HookError, PostToolCall, PostToolCallContext, PostToolCallResult, PreToolCall,
PreToolCallResult, ToolCallContext,
};
use llm_worker::llm_client::event::{Event, ResponseStatus, StatusEvent};
use llm_worker::tool::{Tool, ToolError};
use llm_worker::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
mod common;
use common::MockLlmClient;
@@ -42,25 +42,24 @@ impl SlowTool {
fn call_count(&self) -> usize {
self.call_count.load(Ordering::SeqCst)
}
/// ToolDefinition を作成
fn definition(&self) -> ToolDefinition {
let tool = self.clone();
Arc::new(move || {
let meta = ToolMeta::new(&tool.name)
.description("A tool that waits before responding")
.input_schema(serde_json::json!({
"type": "object",
"properties": {}
}));
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
})
}
}
#[async_trait]
impl Tool for SlowTool {
fn name(&self) -> &str {
&self.name
}
fn description(&self) -> &str {
"A tool that waits before responding"
}
fn input_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {}
})
}
async fn execute(&self, _input_json: &str) -> Result<String, ToolError> {
self.call_count.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(self.delay_ms)).await;
@@ -106,9 +105,9 @@ async fn test_parallel_tool_execution() {
let tool2_clone = tool2.clone();
let tool3_clone = tool3.clone();
worker.register_tool(tool1);
worker.register_tool(tool2);
worker.register_tool(tool3);
worker.register_tool(tool1.definition()).unwrap();
worker.register_tool(tool2.definition()).unwrap();
worker.register_tool(tool3.definition()).unwrap();
let start = Instant::now();
let _result = worker.run("Run all tools").await;
@@ -130,7 +129,7 @@ async fn test_parallel_tool_execution() {
println!("Parallel execution completed in {:?}", elapsed);
}
/// Hook: before_tool_call でスキップされたツールは実行されないことを確認
/// Hook: pre_tool_call でスキップされたツールは実行されないことを確認
#[tokio::test]
async fn test_before_tool_call_skip() {
let events = vec![
@@ -154,24 +153,24 @@ async fn test_before_tool_call_skip() {
let allowed_clone = allowed_tool.clone();
let blocked_clone = blocked_tool.clone();
worker.register_tool(allowed_tool);
worker.register_tool(blocked_tool);
worker.register_tool(allowed_tool.definition()).unwrap();
worker.register_tool(blocked_tool.definition()).unwrap();
// "blocked_tool" をスキップするHook
struct BlockingHook;
#[async_trait]
impl Hook<BeforeToolCall> for BlockingHook {
async fn call(&self, tool_call: &mut ToolCall) -> Result<BeforeToolCallResult, HookError> {
if tool_call.name == "blocked_tool" {
Ok(BeforeToolCallResult::Skip)
impl Hook<PreToolCall> for BlockingHook {
async fn call(&self, ctx: &mut ToolCallContext) -> Result<PreToolCallResult, HookError> {
if ctx.call.name == "blocked_tool" {
Ok(PreToolCallResult::Skip)
} else {
Ok(BeforeToolCallResult::Continue)
Ok(PreToolCallResult::Continue)
}
}
}
worker.add_before_tool_call_hook(BlockingHook);
worker.add_pre_tool_call_hook(BlockingHook);
let _result = worker.run("Test hook").await;
@@ -188,9 +187,9 @@ async fn test_before_tool_call_skip() {
);
}
/// Hook: after_tool_call で結果が改変されることを確認
/// Hook: post_tool_call で結果が改変されることを確認
#[tokio::test]
async fn test_after_tool_call_modification() {
async fn test_post_tool_call_modification() {
// 複数リクエストに対応するレスポンスを準備
let client = MockLlmClient::with_responses(vec![
// 1回目のリクエスト: ツール呼び出し
@@ -220,21 +219,21 @@ async fn test_after_tool_call_modification() {
#[async_trait]
impl Tool for SimpleTool {
fn name(&self) -> &str {
"test_tool"
}
fn description(&self) -> &str {
"Test"
}
fn input_schema(&self) -> serde_json::Value {
serde_json::json!({})
}
async fn execute(&self, _: &str) -> Result<String, ToolError> {
Ok("Original Result".to_string())
}
}
worker.register_tool(SimpleTool);
fn simple_tool_definition() -> ToolDefinition {
Arc::new(|| {
let meta = ToolMeta::new("test_tool")
.description("Test")
.input_schema(serde_json::json!({}));
(meta, Arc::new(SimpleTool) as Arc<dyn Tool>)
})
}
worker.register_tool(simple_tool_definition()).unwrap();
// 結果を改変するHook
struct ModifyingHook {
@@ -242,19 +241,19 @@ async fn test_after_tool_call_modification() {
}
#[async_trait]
impl Hook<AfterToolCall> for ModifyingHook {
impl Hook<PostToolCall> for ModifyingHook {
async fn call(
&self,
tool_result: &mut ToolResult,
) -> Result<AfterToolCallResult, HookError> {
tool_result.content = format!("[Modified] {}", tool_result.content);
*self.modified_content.lock().unwrap() = Some(tool_result.content.clone());
Ok(AfterToolCallResult::Continue)
ctx: &mut PostToolCallContext,
) -> Result<PostToolCallResult, HookError> {
ctx.result.content = format!("[Modified] {}", ctx.result.content);
*self.modified_content.lock().unwrap() = Some(ctx.result.content.clone());
Ok(PostToolCallResult::Continue)
}
}
let modified_content = Arc::new(std::sync::Mutex::new(None));
worker.add_after_tool_call_hook(ModifyingHook {
worker.add_post_tool_call_hook(ModifyingHook {
modified_content: modified_content.clone(),
});
+51 -31
View File
@@ -9,7 +9,7 @@ use std::sync::atomic::{AtomicUsize, Ordering};
use schemars;
use serde;
use llm_worker::tool::Tool;
use llm_worker::tool::{Tool, ToolMeta};
use llm_worker_macros::tool_registry;
// =============================================================================
@@ -51,30 +51,31 @@ async fn test_basic_tool_generation() {
prefix: "Hello".to_string(),
};
// ファクトリメソッドでツールを取得
let greet_tool = ctx.greet_tool();
// ファクトリメソッドでToolDefinitionを取得
let greet_definition = ctx.greet_definition();
// 名前の確認
assert_eq!(greet_tool.name(), "greet");
// ファクトリを呼び出してMetaとToolを取得
let (meta, tool) = greet_definition();
// 説明の確認(docコメントから取得)
let desc = greet_tool.description();
// メタ情報の確認
assert_eq!(meta.name, "greet");
assert!(
desc.contains("メッセージに挨拶を追加する"),
meta.description.contains("メッセージに挨拶を追加する"),
"Description should contain doc comment: {}",
desc
meta.description
);
// スキーマの確認
let schema = greet_tool.input_schema();
println!("Schema: {}", serde_json::to_string_pretty(&schema).unwrap());
assert!(
schema.get("properties").is_some(),
meta.input_schema.get("properties").is_some(),
"Schema should have properties"
);
println!(
"Schema: {}",
serde_json::to_string_pretty(&meta.input_schema).unwrap()
);
// 実行テスト
let result = greet_tool.execute(r#"{"message": "World"}"#).await;
let result = tool.execute(r#"{"message": "World"}"#).await;
assert!(result.is_ok(), "Should execute successfully");
let output = result.unwrap();
assert!(output.contains("Hello"), "Output should contain prefix");
@@ -87,11 +88,11 @@ async fn test_multiple_arguments() {
prefix: "".to_string(),
};
let add_tool = ctx.add_tool();
let (meta, tool) = ctx.add_definition()();
assert_eq!(add_tool.name(), "add");
assert_eq!(meta.name, "add");
let result = add_tool.execute(r#"{"a": 10, "b": 20}"#).await;
let result = tool.execute(r#"{"a": 10, "b": 20}"#).await;
assert!(result.is_ok());
let output = result.unwrap();
assert!(output.contains("30"), "Should contain sum: {}", output);
@@ -103,12 +104,12 @@ async fn test_no_arguments() {
prefix: "TestPrefix".to_string(),
};
let get_prefix_tool = ctx.get_prefix_tool();
let (meta, tool) = ctx.get_prefix_definition()();
assert_eq!(get_prefix_tool.name(), "get_prefix");
assert_eq!(meta.name, "get_prefix");
// 空のJSONオブジェクトで呼び出し
let result = get_prefix_tool.execute(r#"{}"#).await;
let result = tool.execute(r#"{}"#).await;
assert!(result.is_ok());
let output = result.unwrap();
assert!(
@@ -124,10 +125,10 @@ async fn test_invalid_arguments() {
prefix: "".to_string(),
};
let greet_tool = ctx.greet_tool();
let (_, tool) = ctx.greet_definition()();
// 不正なJSON
let result = greet_tool.execute(r#"{"wrong_field": "value"}"#).await;
let result = tool.execute(r#"{"wrong_field": "value"}"#).await;
assert!(result.is_err(), "Should fail with invalid arguments");
}
@@ -163,9 +164,9 @@ impl FallibleContext {
#[tokio::test]
async fn test_result_return_type_success() {
let ctx = FallibleContext;
let validate_tool = ctx.validate_tool();
let (_, tool) = ctx.validate_definition()();
let result = validate_tool.execute(r#"{"value": 42}"#).await;
let result = tool.execute(r#"{"value": 42}"#).await;
assert!(result.is_ok(), "Should succeed for positive value");
let output = result.unwrap();
assert!(output.contains("Valid"), "Should contain Valid: {}", output);
@@ -174,9 +175,9 @@ async fn test_result_return_type_success() {
#[tokio::test]
async fn test_result_return_type_error() {
let ctx = FallibleContext;
let validate_tool = ctx.validate_tool();
let (_, tool) = ctx.validate_definition()();
let result = validate_tool.execute(r#"{"value": -1}"#).await;
let result = tool.execute(r#"{"value": -1}"#).await;
assert!(result.is_err(), "Should fail for negative value");
let err = result.unwrap_err();
@@ -211,12 +212,12 @@ async fn test_sync_method() {
counter: Arc::new(AtomicUsize::new(0)),
};
let increment_tool = ctx.increment_tool();
let (_, tool) = ctx.increment_definition()();
// 3回実行
let result1 = increment_tool.execute(r#"{}"#).await;
let result2 = increment_tool.execute(r#"{}"#).await;
let result3 = increment_tool.execute(r#"{}"#).await;
let result1 = tool.execute(r#"{}"#).await;
let result2 = tool.execute(r#"{}"#).await;
let result3 = tool.execute(r#"{}"#).await;
assert!(result1.is_ok());
assert!(result2.is_ok());
@@ -225,3 +226,22 @@ async fn test_sync_method() {
// カウンターは3になっているはず
assert_eq!(ctx.counter.load(Ordering::SeqCst), 3);
}
// =============================================================================
// Test: ToolMeta Immutability
// =============================================================================
#[tokio::test]
async fn test_tool_meta_immutability() {
let ctx = SimpleContext {
prefix: "Test".to_string(),
};
// 2回取得しても同じメタ情報が得られることを確認
let (meta1, _) = ctx.greet_definition()();
let (meta2, _) = ctx.greet_definition()();
assert_eq!(meta1.name, meta2.name);
assert_eq!(meta1.description, meta2.description);
assert_eq!(meta1.input_schema, meta2.input_schema);
}
-1
View File
@@ -1,4 +1,3 @@
use llm_worker::llm_client::LlmClient;
use llm_worker::llm_client::providers::openai::OpenAIClient;
use llm_worker::{Worker, WorkerError};
+21 -23
View File
@@ -12,7 +12,7 @@ use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use common::MockLlmClient;
use llm_worker::Worker;
use llm_worker::tool::{Tool, ToolError};
use llm_worker::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
/// フィクスチャディレクトリのパス
fn fixtures_dir() -> std::path::PathBuf {
@@ -35,31 +35,29 @@ impl MockWeatherTool {
fn get_call_count(&self) -> usize {
self.call_count.load(Ordering::SeqCst)
}
fn definition(&self) -> ToolDefinition {
let tool = self.clone();
Arc::new(move || {
let meta = ToolMeta::new("get_weather")
.description("Get the current weather for a city")
.input_schema(serde_json::json!({
"type": "object",
"properties": {
"city": {
"type": "string",
"description": "The city name"
}
},
"required": ["city"]
}));
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
})
}
}
#[async_trait]
impl Tool for MockWeatherTool {
fn name(&self) -> &str {
"get_weather"
}
fn description(&self) -> &str {
"Get the current weather for a city"
}
fn input_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"city": {
"type": "string",
"description": "The city name"
}
},
"required": ["city"]
})
}
async fn execute(&self, input_json: &str) -> Result<String, ToolError> {
self.call_count.fetch_add(1, Ordering::SeqCst);
@@ -158,7 +156,7 @@ async fn test_worker_tool_call() {
// ツールを登録
let weather_tool = MockWeatherTool::new();
let tool_for_check = weather_tool.clone();
worker.register_tool(weather_tool);
worker.register_tool(weather_tool.definition()).unwrap();
// メッセージを送信
let _result = worker.run("What's the weather in Tokyo?").await;