feat: Redesign the tool system
This commit is contained in:
@@ -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(),
|
||||
});
|
||||
|
||||
|
||||
@@ -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,4 +1,3 @@
|
||||
use llm_worker::llm_client::LlmClient;
|
||||
use llm_worker::llm_client::providers::openai::OpenAIClient;
|
||||
use llm_worker::{Worker, WorkerError};
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user