feat: Implement worker context management and cache protection mechanisms using type-state
This commit is contained in:
@@ -9,7 +9,7 @@ use std::time::{Duration, Instant};
|
||||
use async_trait::async_trait;
|
||||
use worker::Worker;
|
||||
use worker_types::{
|
||||
ControlFlow, Event, HookError, Message, ResponseStatus, StatusEvent, Tool, ToolCall, ToolError,
|
||||
ControlFlow, Event, HookError, ResponseStatus, StatusEvent, Tool, ToolCall, ToolError,
|
||||
ToolResult, WorkerHook,
|
||||
};
|
||||
|
||||
@@ -108,10 +108,8 @@ async fn test_parallel_tool_execution() {
|
||||
worker.register_tool(tool2);
|
||||
worker.register_tool(tool3);
|
||||
|
||||
let messages = vec![Message::user("Run all tools")];
|
||||
|
||||
let start = Instant::now();
|
||||
let _result = worker.run(messages).await;
|
||||
let _result = worker.run("Run all tools").await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
// 全ツールが呼び出されたことを確認
|
||||
@@ -176,8 +174,7 @@ async fn test_before_tool_call_skip() {
|
||||
|
||||
worker.add_hook(BlockingHook);
|
||||
|
||||
let messages = vec![Message::user("Test hook")];
|
||||
let _result = worker.run(messages).await;
|
||||
let _result = worker.run("Test hook").await;
|
||||
|
||||
// allowed_tool は呼び出されるが、blocked_tool は呼び出されない
|
||||
assert_eq!(
|
||||
@@ -262,8 +259,7 @@ async fn test_after_tool_call_modification() {
|
||||
modified_content: modified_content.clone(),
|
||||
});
|
||||
|
||||
let messages = vec![Message::user("Test modification")];
|
||||
let result = worker.run(messages).await;
|
||||
let result = worker.run("Test modification").await;
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete: {:?}", result);
|
||||
|
||||
|
||||
@@ -9,8 +9,8 @@ use std::sync::{Arc, Mutex};
|
||||
use common::MockLlmClient;
|
||||
use worker::{Worker, WorkerSubscriber};
|
||||
use worker_types::{
|
||||
ErrorEvent, Event, Message, ResponseStatus, StatusEvent, TextBlockEvent, ToolCall,
|
||||
ToolUseBlockEvent, UsageEvent,
|
||||
ErrorEvent, Event, ResponseStatus, StatusEvent, TextBlockEvent, ToolCall, ToolUseBlockEvent,
|
||||
UsageEvent,
|
||||
};
|
||||
|
||||
// =============================================================================
|
||||
@@ -115,8 +115,7 @@ async fn test_subscriber_text_block_events() {
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
let messages = vec![Message::user("Greet me")];
|
||||
let result = worker.run(messages).await;
|
||||
let result = worker.run("Greet me").await;
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete: {:?}", result);
|
||||
|
||||
@@ -155,8 +154,7 @@ async fn test_subscriber_tool_call_complete() {
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
let messages = vec![Message::user("Weather please")];
|
||||
let _ = worker.run(messages).await;
|
||||
let _ = worker.run("Weather please").await;
|
||||
|
||||
// ツール呼び出し完了が収集されていることを確認
|
||||
let completes = tool_call_completes.lock().unwrap();
|
||||
@@ -188,8 +186,7 @@ async fn test_subscriber_turn_events() {
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
let messages = vec![Message::user("Do something")];
|
||||
let result = worker.run(messages).await;
|
||||
let result = worker.run("Do something").await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
|
||||
@@ -226,8 +223,7 @@ async fn test_subscriber_usage_events() {
|
||||
worker.subscribe(subscriber);
|
||||
|
||||
// 実行
|
||||
let messages = vec![Message::user("Hello")];
|
||||
let _ = worker.run(messages).await;
|
||||
let _ = worker.run("Hello").await;
|
||||
|
||||
// Usageイベントが収集されていることを確認
|
||||
let usages = usage_events.lock().unwrap();
|
||||
|
||||
@@ -134,8 +134,7 @@ async fn test_worker_simple_text_response() {
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// シンプルなメッセージを送信
|
||||
let messages = vec![worker_types::Message::user("Hello")];
|
||||
let result = worker.run(messages).await;
|
||||
let result = worker.run("Hello").await;
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete successfully");
|
||||
}
|
||||
@@ -162,8 +161,7 @@ async fn test_worker_tool_call() {
|
||||
worker.register_tool(weather_tool);
|
||||
|
||||
// メッセージを送信
|
||||
let messages = vec![worker_types::Message::user("What's the weather in Tokyo?")];
|
||||
let _result = worker.run(messages).await;
|
||||
let _result = worker.run("What's the weather in Tokyo?").await;
|
||||
|
||||
// ツールが呼び出されたことを確認
|
||||
// Note: max_turns=1なのでツール結果後のリクエストは送信されない
|
||||
@@ -196,8 +194,7 @@ async fn test_worker_with_programmatic_events() {
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
let messages = vec![worker_types::Message::user("Greet me")];
|
||||
let result = worker.run(messages).await;
|
||||
let result = worker.run("Greet me").await;
|
||||
|
||||
assert!(result.is_ok(), "Worker should complete successfully");
|
||||
}
|
||||
|
||||
@@ -0,0 +1,372 @@
|
||||
//! Worker状態管理のテスト
|
||||
//!
|
||||
//! Type-stateパターン(Mutable/Locked)による状態遷移と
|
||||
//! ターン間の状態保持をテストする。
|
||||
|
||||
mod common;
|
||||
|
||||
use common::MockLlmClient;
|
||||
use worker::Worker;
|
||||
use worker_types::{Event, Message, MessageContent, ResponseStatus, StatusEvent};
|
||||
|
||||
// =============================================================================
|
||||
// Mutable状態のテスト
|
||||
// =============================================================================
|
||||
|
||||
/// Mutable状態でシステムプロンプトを設定できることを確認
|
||||
#[test]
|
||||
fn test_mutable_set_system_prompt() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
assert!(worker.get_system_prompt().is_none());
|
||||
|
||||
worker.set_system_prompt("You are a helpful assistant.");
|
||||
assert_eq!(
|
||||
worker.get_system_prompt(),
|
||||
Some("You are a helpful assistant.")
|
||||
);
|
||||
}
|
||||
|
||||
/// Mutable状態で履歴を自由に編集できることを確認
|
||||
#[test]
|
||||
fn test_mutable_history_manipulation() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 初期状態は空
|
||||
assert!(worker.history().is_empty());
|
||||
|
||||
// 履歴を追加
|
||||
worker.push_message(Message::user("Hello"));
|
||||
worker.push_message(Message::assistant("Hi there!"));
|
||||
assert_eq!(worker.history().len(), 2);
|
||||
|
||||
// 履歴への可変アクセス
|
||||
worker.history_mut().push(Message::user("How are you?"));
|
||||
assert_eq!(worker.history().len(), 3);
|
||||
|
||||
// 履歴をクリア
|
||||
worker.clear_history();
|
||||
assert!(worker.history().is_empty());
|
||||
|
||||
// 履歴を設定
|
||||
let messages = vec![Message::user("Test"), Message::assistant("Response")];
|
||||
worker.set_history(messages);
|
||||
assert_eq!(worker.history().len(), 2);
|
||||
}
|
||||
|
||||
/// ビルダーパターンでWorkerを構築できることを確認
|
||||
#[test]
|
||||
fn test_mutable_builder_pattern() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let worker = Worker::new(client)
|
||||
.system_prompt("System prompt")
|
||||
.with_message(Message::user("Hello"))
|
||||
.with_message(Message::assistant("Hi!"))
|
||||
.with_messages(vec![
|
||||
Message::user("How are you?"),
|
||||
Message::assistant("I'm fine!"),
|
||||
]);
|
||||
|
||||
assert_eq!(worker.get_system_prompt(), Some("System prompt"));
|
||||
assert_eq!(worker.history().len(), 4);
|
||||
}
|
||||
|
||||
/// extend_historyで複数メッセージを追加できることを確認
|
||||
#[test]
|
||||
fn test_mutable_extend_history() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
worker.push_message(Message::user("First"));
|
||||
|
||||
worker.extend_history(vec![
|
||||
Message::assistant("Response 1"),
|
||||
Message::user("Second"),
|
||||
Message::assistant("Response 2"),
|
||||
]);
|
||||
|
||||
assert_eq!(worker.history().len(), 4);
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// 状態遷移テスト
|
||||
// =============================================================================
|
||||
|
||||
/// lock()でMutable -> Locked状態に遷移することを確認
|
||||
#[test]
|
||||
fn test_lock_transition() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
worker.set_system_prompt("System");
|
||||
worker.push_message(Message::user("Hello"));
|
||||
worker.push_message(Message::assistant("Hi"));
|
||||
|
||||
// ロック
|
||||
let locked_worker = worker.lock();
|
||||
|
||||
// Locked状態でも履歴とシステムプロンプトにアクセス可能
|
||||
assert_eq!(locked_worker.get_system_prompt(), Some("System"));
|
||||
assert_eq!(locked_worker.history().len(), 2);
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2);
|
||||
}
|
||||
|
||||
/// unlock()でLocked -> Mutable状態に遷移することを確認
|
||||
#[test]
|
||||
fn test_unlock_transition() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
worker.push_message(Message::user("Hello"));
|
||||
let locked_worker = worker.lock();
|
||||
|
||||
// アンロック
|
||||
let mut worker = locked_worker.unlock();
|
||||
|
||||
// Mutable状態に戻ったので履歴操作が可能
|
||||
worker.push_message(Message::assistant("Hi"));
|
||||
worker.clear_history();
|
||||
assert!(worker.history().is_empty());
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ターン実行と状態保持のテスト
|
||||
// =============================================================================
|
||||
|
||||
/// Mutable状態でターンを実行し、履歴が正しく更新されることを確認
|
||||
#[tokio::test]
|
||||
async fn test_mutable_run_updates_history() {
|
||||
let events = vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Hello, I'm an assistant!"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
];
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 実行
|
||||
let result = worker.run("Hi there").await;
|
||||
assert!(result.is_ok());
|
||||
|
||||
// 履歴が更新されている
|
||||
let history = worker.history();
|
||||
assert_eq!(history.len(), 2); // user + assistant
|
||||
|
||||
// ユーザーメッセージ
|
||||
assert!(matches!(
|
||||
&history[0].content,
|
||||
MessageContent::Text(t) if t == "Hi there"
|
||||
));
|
||||
|
||||
// アシスタントメッセージ
|
||||
assert!(matches!(
|
||||
&history[1].content,
|
||||
MessageContent::Text(t) if t == "Hello, I'm an assistant!"
|
||||
));
|
||||
}
|
||||
|
||||
/// Locked状態で複数ターンを実行し、履歴が正しく累積することを確認
|
||||
#[tokio::test]
|
||||
async fn test_locked_multi_turn_history_accumulation() {
|
||||
// 2回のリクエストに対応するレスポンスを準備
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
// 1回目のレスポンス
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Nice to meet you!"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
// 2回目のレスポンス
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "I can help with that."),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
]);
|
||||
|
||||
let worker = Worker::new(client).system_prompt("You are helpful.");
|
||||
|
||||
// ロック(システムプロンプト設定後)
|
||||
let mut locked_worker = worker.lock();
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 0); // メッセージはまだない
|
||||
|
||||
// 1ターン目
|
||||
let result1 = locked_worker.run("Hello!").await;
|
||||
assert!(result1.is_ok());
|
||||
assert_eq!(locked_worker.history().len(), 2); // user + assistant
|
||||
|
||||
// 2ターン目
|
||||
let result2 = locked_worker.run("Can you help me?").await;
|
||||
assert!(result2.is_ok());
|
||||
assert_eq!(locked_worker.history().len(), 4); // 2 * (user + assistant)
|
||||
|
||||
// 履歴の内容を確認
|
||||
let history = locked_worker.history();
|
||||
|
||||
// 1ターン目のユーザーメッセージ
|
||||
assert!(matches!(&history[0].content, MessageContent::Text(t) if t == "Hello!"));
|
||||
|
||||
// 1ターン目のアシスタントメッセージ
|
||||
assert!(matches!(&history[1].content, MessageContent::Text(t) if t == "Nice to meet you!"));
|
||||
|
||||
// 2ターン目のユーザーメッセージ
|
||||
assert!(matches!(&history[2].content, MessageContent::Text(t) if t == "Can you help me?"));
|
||||
|
||||
// 2ターン目のアシスタントメッセージ
|
||||
assert!(matches!(&history[3].content, MessageContent::Text(t) if t == "I can help with that."));
|
||||
}
|
||||
|
||||
/// locked_prefix_lenがロック時点の履歴長を正しく記録することを確認
|
||||
#[tokio::test]
|
||||
async fn test_locked_prefix_len_tracking() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Response 1"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Response 2"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
]);
|
||||
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
// 事前にメッセージを追加
|
||||
worker.push_message(Message::user("Pre-existing message 1"));
|
||||
worker.push_message(Message::assistant("Pre-existing response 1"));
|
||||
|
||||
assert_eq!(worker.history().len(), 2);
|
||||
|
||||
// ロック
|
||||
let mut locked_worker = worker.lock();
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2); // ロック時点で2メッセージ
|
||||
|
||||
// ターン実行
|
||||
locked_worker.run("New message").await.unwrap();
|
||||
|
||||
// 履歴は増えるが、locked_prefix_lenは変わらない
|
||||
assert_eq!(locked_worker.history().len(), 4); // 2 + 2
|
||||
assert_eq!(locked_worker.locked_prefix_len(), 2); // 変わらない
|
||||
}
|
||||
|
||||
/// ターンカウントが正しくインクリメントされることを確認
|
||||
#[tokio::test]
|
||||
async fn test_turn_count_increment() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Turn 1"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Turn 2"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
]);
|
||||
|
||||
let mut worker = Worker::new(client);
|
||||
|
||||
assert_eq!(worker.turn_count(), 0);
|
||||
|
||||
worker.run("First").await.unwrap();
|
||||
assert_eq!(worker.turn_count(), 1);
|
||||
|
||||
worker.run("Second").await.unwrap();
|
||||
assert_eq!(worker.turn_count(), 2);
|
||||
}
|
||||
|
||||
/// unlock後に履歴を編集し、再度lockできることを確認
|
||||
#[tokio::test]
|
||||
async fn test_unlock_edit_relock() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Response"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]]);
|
||||
|
||||
let worker = Worker::new(client)
|
||||
.with_message(Message::user("Hello"))
|
||||
.with_message(Message::assistant("Hi"));
|
||||
|
||||
// ロック -> アンロック
|
||||
let locked = worker.lock();
|
||||
assert_eq!(locked.locked_prefix_len(), 2);
|
||||
|
||||
let mut unlocked = locked.unlock();
|
||||
|
||||
// 履歴を編集
|
||||
unlocked.clear_history();
|
||||
unlocked.push_message(Message::user("Fresh start"));
|
||||
|
||||
// 再ロック
|
||||
let relocked = unlocked.lock();
|
||||
assert_eq!(relocked.history().len(), 1);
|
||||
assert_eq!(relocked.locked_prefix_len(), 1);
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// システムプロンプト保持のテスト
|
||||
// =============================================================================
|
||||
|
||||
/// Locked状態でもシステムプロンプトが保持されることを確認
|
||||
#[test]
|
||||
fn test_system_prompt_preserved_in_locked_state() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let worker = Worker::new(client).system_prompt("Important system prompt");
|
||||
|
||||
let locked = worker.lock();
|
||||
assert_eq!(locked.get_system_prompt(), Some("Important system prompt"));
|
||||
|
||||
let unlocked = locked.unlock();
|
||||
assert_eq!(
|
||||
unlocked.get_system_prompt(),
|
||||
Some("Important system prompt")
|
||||
);
|
||||
}
|
||||
|
||||
/// unlock -> 再lock でシステムプロンプトを変更できることを確認
|
||||
#[test]
|
||||
fn test_system_prompt_change_after_unlock() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let worker = Worker::new(client).system_prompt("Original prompt");
|
||||
|
||||
let locked = worker.lock();
|
||||
let mut unlocked = locked.unlock();
|
||||
|
||||
unlocked.set_system_prompt("New prompt");
|
||||
assert_eq!(unlocked.get_system_prompt(), Some("New prompt"));
|
||||
|
||||
let relocked = unlocked.lock();
|
||||
assert_eq!(relocked.get_system_prompt(), Some("New prompt"));
|
||||
}
|
||||
Reference in New Issue
Block a user