Skip to main content

agent_base/engine/
builder.rs

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    /// Create a new builder with an `LlmProvider`.
37    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    /// Set a callback to transform messages before they are sent to the LLM.
67    ///
68    /// The default behavior (when `None`) is to filter out
69    /// `ChatMessage::Custom` variants, which most providers don't understand.
70    /// Override this to inject custom serialization logic for application-specific
71    /// message types.
72    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    /// Set a guard from a pre-built `Arc<dyn ReactLoopGuard>`.
83    ///
84    /// Use this when you need to share the guard reference (e.g. for inspecting
85    /// recorded calls in tests).
86    pub fn guard_dyn(mut self, guard: Arc<dyn ReactLoopGuard>) -> Self {
87        self.guard = Some(guard);
88        self
89    }
90
91    /// Set an inline context compactor for the react loop.
92    ///
93    /// When set, the react loop will check token count after each tool turn
94    /// and compact the session if it exceeds the configured threshold.
95    pub fn context_compactor(mut self, compactor: Arc<dyn ContextCompaction>) -> Self {
96        self.context_compactor = Some(compactor);
97        self
98    }
99
100    /// Check if a guard has been set.
101    ///
102    /// Used by agent-works to inject a default guard if none was set.
103    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    /// Set whether to include the reasoning content in LLM responses.
113    ///
114    /// Controls whether the `reasoning_content` field from the LLM is forwarded
115    /// to consumers (i.e., "show the thinking process").
116    /// See [`AgentConfig::enable_thought`] for the distinction from `enable_thinking()`.
117    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    /// Set whether to enable the model's extended thinking / reasoning mode.
128    ///
129    /// Controls whether the model performs deep reasoning (i.e., "enable thinking mode").
130    /// See [`AgentConfig::enable_thought`] for the distinction from `enable_thought()`.
131    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    /// Cap the number of react-loop iterations allowed for a *single* run (one user
221    /// input). Distinct from [`Self::max_turns_per_session`], which caps turns across
222    /// the whole session. When unset, falls back to `DEFAULT_MAX_TURNS` (50).
223    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    /// Conditionally chain a builder call: apply `f` only when `value` is `Some`.
244    ///
245    /// # Example
246    /// ```ignore
247    /// let builder = AgentBuilder::new(client)
248    ///     .apply_if(config.timeout, |b, t| b.tool_timeout(t))
249    ///     .apply_if(config.max_chars, |b, c| b.max_tool_output_chars(c));
250    /// ```
251    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}