1use std::sync::Arc;
2
3use crate::llm::ReasoningConfig;
4use crate::tool::{Tool, ToolPolicy, ToolRegistry};
5use crate::types::{
6 AgentConfig, AtomicU64SessionIdGenerator, ConvertToLlmFn, ResponseFormat, RetryConfig,
7 SessionIdGenerator,
8};
9
10use super::AgentRuntime;
11use super::approval::ApprovalHandler;
12use super::context::{ContextCompaction, ContextWindowManager};
13use super::middleware::{Middleware, MiddlewareRef};
14use super::react_loop_guard::ReactLoopGuard;
15use super::recovery::{StopOnError, ToolErrorRecovery};
16use super::session_store::{InMemorySessionStore, SessionStore};
17
18pub struct AgentBuilder {
19 provider: Arc<dyn llm_trait::LlmProvider>,
20 config: AgentConfig,
21 tools: ToolRegistry,
22 approval_handler: Option<Arc<dyn ApprovalHandler>>,
23 tool_policy: Option<Arc<dyn ToolPolicy>>,
24 middlewares: Vec<MiddlewareRef>,
25 context_manager: Option<ContextWindowManager>,
26 session_store: Option<Arc<dyn SessionStore>>,
27 error_recovery: Option<Arc<dyn ToolErrorRecovery>>,
28 event_bus_capacity: usize,
29 session_id_generator: Option<Arc<dyn SessionIdGenerator>>,
30 convert_to_llm: Option<ConvertToLlmFn>,
31 guard: Option<Arc<dyn ReactLoopGuard>>,
32 context_compactor: Option<Arc<dyn ContextCompaction>>,
33}
34
35impl AgentBuilder {
36 pub fn new(provider: Arc<dyn llm_trait::LlmProvider>) -> Self {
38 Self {
39 provider,
40 config: AgentConfig::default(),
41 tools: ToolRegistry::default(),
42 approval_handler: None,
43 tool_policy: None,
44 middlewares: Vec::new(),
45 context_manager: None,
46 session_store: None,
47 error_recovery: None,
48 event_bus_capacity: 2048,
49 session_id_generator: None,
50 convert_to_llm: None,
51 guard: None,
52 context_compactor: None,
53 }
54 }
55
56 pub fn event_bus_capacity(mut self, capacity: usize) -> Self {
57 self.event_bus_capacity = capacity;
58 self
59 }
60
61 pub fn session_id_generator(mut self, generator: Arc<dyn SessionIdGenerator>) -> Self {
62 self.session_id_generator = Some(generator);
63 self
64 }
65
66 pub fn convert_to_llm(mut self, cb: ConvertToLlmFn) -> Self {
73 self.convert_to_llm = Some(cb);
74 self
75 }
76
77 pub fn guard(mut self, guard: impl ReactLoopGuard + 'static) -> Self {
78 self.guard = Some(Arc::new(guard));
79 self
80 }
81
82 pub fn guard_dyn(mut self, guard: Arc<dyn ReactLoopGuard>) -> Self {
87 self.guard = Some(guard);
88 self
89 }
90
91 pub fn context_compactor(mut self, compactor: Arc<dyn ContextCompaction>) -> Self {
96 self.context_compactor = Some(compactor);
97 self
98 }
99
100 pub fn get_guard(&self) -> Option<&Arc<dyn ReactLoopGuard>> {
104 self.guard.as_ref()
105 }
106
107 pub fn system_prompt(mut self, system_prompt: impl Into<String>) -> Self {
108 self.config.system_prompt = Some(system_prompt.into());
109 self
110 }
111
112 pub fn enable_thought(mut self, enable: bool) -> Self {
118 self.config.enable_thought = enable;
119 self
120 }
121
122 pub fn reasoning(mut self, config: ReasoningConfig) -> Self {
123 self.config.reasoning = Some(config);
124 self
125 }
126
127 pub fn enable_thinking(mut self, enable: bool) -> Self {
132 let mut config = self.config.reasoning.take().unwrap_or_default();
133 config.enabled = Some(enable);
134 self.config.reasoning = Some(config);
135 self
136 }
137
138 pub fn thinking_budget(mut self, budget: u64) -> Self {
139 let mut config = self.config.reasoning.take().unwrap_or_default();
140 config.budget_tokens = Some(budget);
141 self.config.reasoning = Some(config);
142 self
143 }
144
145 pub fn tool_timeout(mut self, timeout_ms: u64) -> Self {
146 self.config.tool.tool_timeout_ms = Some(timeout_ms);
147 self
148 }
149
150 pub fn max_tool_output_chars(mut self, max_chars: usize) -> Self {
151 self.config.tool.max_tool_output_chars = Some(max_chars);
152 self
153 }
154
155 pub fn register_tool(mut self, tool: impl Tool + 'static) -> Self {
156 self.tools.register(tool);
157 self
158 }
159
160 pub fn register_tool_arc(mut self, tool: Arc<dyn Tool>) -> Self {
161 self.tools.register_arc(tool);
162 self
163 }
164
165 pub fn approval_handler(mut self, handler: Arc<dyn ApprovalHandler>) -> Self {
166 self.approval_handler = Some(handler);
167 self
168 }
169
170 pub fn tool_policy(mut self, policy: Arc<dyn ToolPolicy>) -> Self {
171 self.tool_policy = Some(policy);
172 self
173 }
174
175 pub fn middleware(mut self, mw: impl Middleware + 'static) -> Self {
176 self.middlewares.push(Arc::new(mw));
177 self
178 }
179
180 pub fn context_window(mut self, max_tokens: usize) -> Self {
181 self.context_manager = Some(ContextWindowManager::new(max_tokens));
182 self
183 }
184
185 pub fn context_window_manager(mut self, manager: ContextWindowManager) -> Self {
186 self.context_manager = Some(manager);
187 self
188 }
189
190 pub fn response_format(mut self, format: ResponseFormat) -> Self {
191 self.config.llm.response_format = Some(format);
192 self
193 }
194
195 pub fn llm_retry(mut self, retry: RetryConfig) -> Self {
196 self.config.llm.llm_retry = Some(retry);
197 self
198 }
199
200 pub fn session_store(mut self, store: Arc<dyn SessionStore>) -> Self {
201 self.session_store = Some(store);
202 self
203 }
204
205 pub fn error_recovery(mut self, recovery: Arc<dyn ToolErrorRecovery>) -> Self {
206 self.error_recovery = Some(recovery);
207 self
208 }
209
210 pub fn max_sessions(mut self, max: usize) -> Self {
211 self.config.session.max_sessions = Some(max);
212 self
213 }
214
215 pub fn max_turns_per_session(mut self, max: usize) -> Self {
216 self.config.session.max_turns_per_session = Some(max);
217 self
218 }
219
220 pub fn execution_max_turns(mut self, max: u32) -> Self {
224 self.config.execution.max_turns = Some(max);
225 self
226 }
227
228 pub fn max_message_tokens(mut self, max: usize) -> Self {
229 self.config.session.max_message_tokens = Some(max);
230 self
231 }
232
233 pub fn tool_error_retry_prompt(mut self, prompt: impl Into<String>) -> Self {
234 self.config.tool.tool_error_retry_prompt = Some(prompt.into());
235 self
236 }
237
238 pub fn language(mut self, language: crate::types::Language) -> Self {
239 self.config.language = language;
240 self
241 }
242
243 pub fn apply_if<T>(self, value: Option<T>, f: impl FnOnce(Self, T) -> Self) -> Self {
252 match value {
253 Some(v) => f(self, v),
254 None => self,
255 }
256 }
257
258 pub fn build(self) -> crate::types::AgentResult<AgentRuntime> {
259 self.config.validate()?;
260
261 tracing::info!(
262 tool_count = self.tools.len(),
263 middleware_count = self.middlewares.len(),
264 has_approval = self.approval_handler.is_some(),
265 has_context_window = self.context_manager.is_some(),
266 "building agent runtime"
267 );
268
269 let event_bus = super::runtime::EventBus::new(self.event_bus_capacity);
270
271 let session_store = self
272 .session_store
273 .unwrap_or_else(|| Arc::new(InMemorySessionStore::new()));
274 let error_recovery = self.error_recovery.unwrap_or_else(|| Arc::new(StopOnError));
275 let session_id_generator = self
276 .session_id_generator
277 .unwrap_or_else(|| Arc::new(AtomicU64SessionIdGenerator::default()));
278
279 let session_manager = super::runtime::SessionManager::new(
280 session_id_generator,
281 session_store,
282 self.config.session.clone(),
283 );
284
285 let llm_engine = super::runtime::LlmEngine::new(self.provider.clone(), event_bus.clone());
286
287 let tool_engine = super::runtime::ToolEngine::new(
288 self.tools,
289 self.approval_handler,
290 self.tool_policy,
291 error_recovery,
292 event_bus.clone(),
293 );
294
295 let runner = Arc::new(super::runtime::RuntimeCore::new(
296 self.config,
297 llm_engine,
298 tool_engine,
299 session_manager,
300 event_bus,
301 self.context_manager,
302 self.middlewares,
303 self.convert_to_llm,
304 self.guard,
305 self.context_compactor,
306 ));
307
308 Ok(AgentRuntime { runner })
309 }
310}
311
312#[cfg(test)]
313mod tests {
314 use super::*;
315 use crate::engine::DenyAllApprovalHandler;
316 use crate::llm::ReasoningEffort;
317 use crate::tool::{Content, ToolContext};
318 use crate::types::{AgentError, AgentResult, ApprovalRequest, Language, ResponseFormat};
319 use async_trait::async_trait;
320 use llm_trait::{Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, ProviderInfo};
321 use serde_json::Value;
322
323 struct DummyProvider;
324
325 #[async_trait]
326 impl llm_trait::LlmProvider for DummyProvider {
327 async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
328 Ok(ChatStream::new(Box::pin(futures_util::stream::empty())))
329 }
330
331 async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
332 Ok(ChatResponse {
333 content: String::new(),
334 reasoning_content: None,
335 tool_calls: vec![],
336 usage: Default::default(),
337 finish_reason: llm_trait::FinishReason::Stop,
338 raw: None,
339 thinking_signature: None,
340 })
341 }
342
343 fn capabilities(&self) -> Capabilities {
344 Capabilities::default()
345 }
346
347 fn info(&self) -> ProviderInfo {
348 ProviderInfo {
349 name: "dummy".to_string(),
350 model: "dummy".to_string(),
351 version: None,
352 }
353 }
354 }
355
356 #[test]
357 fn execution_max_turns_writes_per_run_config() {
358 let builder = AgentBuilder::new(Arc::new(DummyProvider)).execution_max_turns(200);
359 assert_eq!(builder.config.execution.max_turns, Some(200));
360 }
361
362 #[test]
363 fn execution_max_turns_defaults_to_none() {
364 let builder = AgentBuilder::new(Arc::new(DummyProvider));
365 assert_eq!(builder.config.execution.max_turns, None);
366 }
367
368 fn b() -> AgentBuilder {
369 AgentBuilder::new(Arc::new(DummyProvider))
370 }
371
372 struct NoopTool;
373
374 #[async_trait]
375 impl Tool for NoopTool {
376 fn name(&self) -> &'static str {
377 "noop"
378 }
379 fn description(&self) -> &'static str {
380 "noop tool"
381 }
382 fn schema(&self) -> Value {
383 serde_json::json!({"type": "object"})
384 }
385 async fn call(&self, _args: &Value, _ctx: &ToolContext) -> AgentResult<Vec<Content>> {
386 Ok(vec![Content::text("ok")])
387 }
388 }
389
390 struct AutoApprovePolicy;
391
392 #[async_trait]
393 impl ToolPolicy for AutoApprovePolicy {
394 async fn evaluate_approval(
395 &self,
396 _tool_name: &str,
397 _args: &Value,
398 ) -> Option<ApprovalRequest> {
399 None
400 }
401 }
402
403 struct NoopMiddleware;
404
405 impl Middleware for NoopMiddleware {}
406
407 #[test]
408 fn system_prompt_sets_config() {
409 assert_eq!(
410 b().system_prompt("be helpful")
411 .config
412 .system_prompt
413 .as_deref(),
414 Some("be helpful")
415 );
416 }
417
418 #[test]
419 fn enable_thought_sets_config() {
420 assert!(b().enable_thought(true).config.enable_thought);
421 }
422
423 #[test]
424 fn reasoning_sets_config() {
425 let rc = ReasoningConfig {
426 enabled: Some(true),
427 budget_tokens: Some(64),
428 effort: Some(ReasoningEffort::Medium),
429 };
430 let builder = b().reasoning(rc);
431 let got = builder.config.reasoning.as_ref().unwrap();
432 assert_eq!(got.enabled, Some(true));
433 assert_eq!(got.budget_tokens, Some(64));
434 assert!(matches!(got.effort.as_ref(), Some(ReasoningEffort::Medium)));
435 }
436
437 #[test]
438 fn enable_thinking_and_budget_set_reasoning() {
439 let builder = b().enable_thinking(true).thinking_budget(128);
440 let got = builder.config.reasoning.as_ref().unwrap();
441 assert_eq!(got.enabled, Some(true));
442 assert_eq!(got.budget_tokens, Some(128));
443 }
444
445 #[test]
446 fn tool_limits_set_config() {
447 let builder = b().tool_timeout(5_000).max_tool_output_chars(1_024);
448 assert_eq!(builder.config.tool.tool_timeout_ms, Some(5_000));
449 assert_eq!(builder.config.tool.max_tool_output_chars, Some(1_024));
450 }
451
452 #[test]
453 fn register_tool_adds_to_registry() {
454 assert_eq!(b().register_tool(NoopTool).tools.len(), 1);
455 }
456
457 #[test]
458 fn approval_handler_and_tool_policy_are_set() {
459 let builder = b()
460 .approval_handler(Arc::new(DenyAllApprovalHandler))
461 .tool_policy(Arc::new(AutoApprovePolicy));
462 assert!(builder.approval_handler.is_some());
463 assert!(builder.tool_policy.is_some());
464 }
465
466 #[test]
467 fn middleware_and_context_window_are_set() {
468 let builder = b().middleware(NoopMiddleware).context_window(8_000);
469 assert_eq!(builder.middlewares.len(), 1);
470 assert!(builder.context_manager.is_some());
471 }
472
473 #[test]
474 fn response_format_and_retry_set_config() {
475 let builder = b()
476 .response_format(ResponseFormat::JsonObject)
477 .llm_retry(RetryConfig::default().max_retries(5));
478 assert!(builder.config.llm.response_format.is_some());
479 assert_eq!(
480 builder.config.llm.llm_retry.as_ref().unwrap().max_retries,
481 5
482 );
483 }
484
485 #[test]
486 fn session_store_and_error_recovery_are_set() {
487 let builder = b()
488 .session_store(Arc::new(InMemorySessionStore::new()))
489 .error_recovery(Arc::new(StopOnError));
490 assert!(builder.session_store.is_some());
491 assert!(builder.error_recovery.is_some());
492 }
493
494 #[test]
495 fn session_limits_set_config() {
496 let builder = b()
497 .max_sessions(10)
498 .max_turns_per_session(20)
499 .max_message_tokens(30);
500 assert_eq!(builder.config.session.max_sessions, Some(10));
501 assert_eq!(builder.config.session.max_turns_per_session, Some(20));
502 assert_eq!(builder.config.session.max_message_tokens, Some(30));
503 }
504
505 #[test]
506 fn tool_error_retry_prompt_and_language_set_config() {
507 let builder = b()
508 .tool_error_retry_prompt("try again")
509 .language(Language::Zh);
510 assert_eq!(
511 builder.config.tool.tool_error_retry_prompt.as_deref(),
512 Some("try again")
513 );
514 assert_eq!(builder.config.language, Language::Zh);
515 }
516
517 #[test]
518 fn apply_if_applies_when_some_and_skips_when_none() {
519 let applied = b().apply_if(Some(3_000_u64), |b, t| b.tool_timeout(t));
520 assert_eq!(applied.config.tool.tool_timeout_ms, Some(3_000));
521
522 let skipped = b().apply_if(None, |b, t| b.tool_timeout(t));
523 assert_eq!(skipped.config.tool.tool_timeout_ms, None);
524 }
525
526 #[test]
527 fn build_ok_with_defaults() {
528 assert!(b().build().is_ok());
529 }
530
531 #[test]
532 fn build_err_on_invalid_config() {
533 assert!(matches!(
534 b().execution_max_turns(0).build(),
535 Err(AgentError::ConfigError(_))
536 ));
537 }
538}