feat: Implement worker context management and cache protection mechanisms using type-state

This commit is contained in:
2026-01-08 17:57:03 +09:00
parent 45c8457b71
commit 2487d1ece7
9 changed files with 831 additions and 195 deletions
+4 -8
View File
@@ -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);
+6 -10
View File
@@ -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();
+3 -6
View File
@@ -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");
}
+372
View File
@@ -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"));
}