Skip to main content

runifold_agent/
builder.rs

1use std::sync::Arc;
2
3use runifold_core::CapabilitySet;
4use runifold_effect::{EffectExecutor, EffectRecoveryPolicy};
5use runifold_model::{
6    FeaturePolicy, GenerationOptions, Message, Model, ModelRef, OutputFormat, ProviderToolSpec,
7    ResponseMode,
8};
9use runifold_retrieval::{Document, RetrievalError, Retriever};
10use runifold_tool::{Tool, ToolRegistrationError};
11use schemars::JsonSchema;
12use thiserror::Error;
13
14use crate::agent::DynamicContext;
15use crate::{
16    Agent, AgentConfig, AgentDescriptor, AgentError, AgentFuture, AgentOutcome,
17    AgentRegistrationError, AgentRoute, GatewayMiddleware, StructuredAgent, ToolErrorPolicy,
18};
19
20/// Failure while assembling an [`Agent`].
21#[derive(Clone, Debug, Error, Eq, PartialEq)]
22#[non_exhaustive]
23pub enum AgentBuildError {
24    /// Tool registration failed.
25    #[error("agent Tool registration failed: {0}")]
26    Tool(#[from] ToolRegistrationError),
27    /// Child Agent route registration failed.
28    #[error("agent route registration failed: {0}")]
29    Route(#[from] AgentRegistrationError),
30    /// One model-facing name was used by both a Tool and an Agent.
31    #[error("callable name `{0}` is registered as both a Tool and an Agent")]
32    CallableNameCollision(String),
33    /// The Agent name is blank.
34    #[error("agent name cannot be empty")]
35    EmptyName,
36    /// The configured turn limit cannot execute any model turn.
37    #[error("max_turns must be greater than zero")]
38    ZeroMaxTurns,
39    /// A static or dynamic context registration was invalid.
40    #[error("agent retrieval configuration failed: {0}")]
41    Retrieval(#[from] RetrievalError),
42}
43
44/// Failure while building and immediately prompting an Agent.
45#[derive(Debug, Error)]
46#[non_exhaustive]
47pub enum AgentPromptError {
48    /// Fluent Agent assembly failed before model execution.
49    #[error("failed to build agent: {0}")]
50    Build(#[from] AgentBuildError),
51    /// Canonical Agent execution failed.
52    #[error("agent prompt failed: {0}")]
53    Run(#[from] AgentError),
54}
55
56/// Fluent assembly of one canonical [`Agent`].
57///
58/// Registration failures are retained and returned by [`Self::build`] so
59/// Tool and child Agent calls remain chainable without silently replacing an
60/// existing name.
61pub struct AgentBuilder {
62    agent: Agent,
63    error: Option<AgentBuildError>,
64}
65
66impl AgentBuilder {
67    /// Creates a builder around the same execution path as [`Agent::new`].
68    pub fn new(name: impl Into<String>, model: Arc<dyn Model>, model_ref: ModelRef) -> Self {
69        Self {
70            agent: Agent::new(name, model, model_ref),
71            error: None,
72        }
73    }
74
75    /// Appends a system instruction.
76    #[must_use]
77    pub fn system(mut self, instruction: impl Into<String>) -> Self {
78        self.agent
79            .instructions
80            .push(Message::system(instruction.into()));
81        self
82    }
83
84    /// Adds one static document as untrusted user-level context.
85    ///
86    /// The document is never promoted to a system instruction. Use
87    /// [`Self::system`] for trusted application policy.
88    #[must_use]
89    pub fn context(self, text: impl Into<String>) -> Self {
90        let id = format!("static-context-{}", self.agent.context.len() + 1);
91        match Document::new(id, text) {
92            Ok(document) => self.context_document(document),
93            Err(error) => self.with_error(error.into()),
94        }
95    }
96
97    /// Adds one validated static context document.
98    #[must_use]
99    pub fn context_document(mut self, document: Document) -> Self {
100        if self.error.is_none() {
101            self.agent.context.push(document);
102        }
103        self
104    }
105
106    /// Adds an owned dynamic context source.
107    #[must_use]
108    pub fn dynamic_context<R>(self, limit: usize, retriever: R) -> Self
109    where
110        R: Retriever + 'static,
111    {
112        self.shared_dynamic_context(limit, Arc::new(retriever))
113    }
114
115    /// Adds a shared, type-erased dynamic context source.
116    #[must_use]
117    pub fn shared_dynamic_context(mut self, limit: usize, retriever: Arc<dyn Retriever>) -> Self {
118        if self.error.is_none() {
119            if limit == 0 {
120                self.error = Some(RetrievalError::ZeroLimit.into());
121            } else {
122                self.agent
123                    .dynamic_context
124                    .push(DynamicContext { limit, retriever });
125            }
126        }
127        self
128    }
129
130    /// Registers an owned Tool.
131    #[must_use]
132    pub fn tool<T>(self, tool: T) -> Self
133    where
134        T: Tool + 'static,
135    {
136        self.shared_tool(Arc::new(tool))
137    }
138
139    /// Registers a shared, type-erased Tool.
140    #[must_use]
141    pub fn shared_tool(mut self, tool: Arc<dyn Tool>) -> Self {
142        if self.error.is_none()
143            && let Err(error) = self.agent.tools.register(tool)
144        {
145            self.error = Some(error.into());
146        }
147        self
148    }
149
150    fn with_error(mut self, error: AgentBuildError) -> Self {
151        if self.error.is_none() {
152            self.error = Some(error);
153        }
154        self
155    }
156
157    /// Registers a child Agent route with explicit delegated capabilities.
158    #[must_use]
159    pub fn child(
160        mut self,
161        descriptor: AgentDescriptor,
162        child: Arc<Agent>,
163        capabilities: CapabilitySet,
164    ) -> Self {
165        if self.error.is_none() {
166            let route = AgentRoute::new(descriptor, child).with_capabilities(capabilities);
167            if let Err(error) = self.agent.agents.register(route) {
168                self.error = Some(error.into());
169            }
170        }
171        self
172    }
173
174    /// Appends Gateway around-middleware.
175    #[must_use]
176    pub fn gateway_layer(mut self, middleware: Arc<dyn GatewayMiddleware>) -> Self {
177        self.agent.agents.push_middleware(middleware);
178        self
179    }
180
181    /// Sets the maximum nested Agent delegation depth.
182    #[must_use]
183    pub fn max_delegation_depth(mut self, max_depth: u32) -> Self {
184        self.agent.agents = self.agent.agents.with_max_depth(max_depth);
185        self
186    }
187
188    /// Sets the local model-turn limit.
189    #[must_use]
190    pub const fn max_turns(mut self, max_turns: u32) -> Self {
191        self.agent.config.max_turns = max_turns;
192        self
193    }
194
195    /// Sets Tool failure behavior.
196    #[must_use]
197    pub const fn tool_error_policy(mut self, policy: ToolErrorPolicy) -> Self {
198        self.agent.config.tool_error_policy = policy;
199        self
200    }
201
202    /// Sets provider feature-degradation behavior.
203    #[must_use]
204    pub const fn feature_policy(mut self, policy: FeaturePolicy) -> Self {
205        self.agent.config.feature_policy = policy;
206        self
207    }
208
209    /// Sets the desired terminal model-output format.
210    #[must_use]
211    pub fn output_format(mut self, output_format: OutputFormat) -> Self {
212        self.agent.output_format = output_format;
213        self
214    }
215
216    /// Requests strict structured output described by the Rust type `T`.
217    #[must_use]
218    pub fn structured_output<T>(self, name: impl Into<String>) -> Self
219    where
220        T: JsonSchema,
221    {
222        self.output_format(OutputFormat::typed::<T>(name))
223    }
224
225    /// Adds a provider-hosted tool such as Ark web search.
226    #[must_use]
227    pub fn provider_tool(mut self, tool: ProviderToolSpec) -> Self {
228        self.agent.provider_tools.push(tool);
229        self
230    }
231
232    /// Replaces common model generation controls for every Agent turn.
233    #[must_use]
234    pub fn generation(mut self, generation: GenerationOptions) -> Self {
235        self.agent.generation = generation;
236        self
237    }
238
239    /// Sets the sampling temperature for every Agent turn.
240    #[must_use]
241    pub fn temperature(mut self, temperature: f64) -> Self {
242        self.agent.generation.temperature = Some(temperature);
243        self
244    }
245
246    /// Sets nucleus sampling for every Agent turn.
247    #[must_use]
248    pub fn top_p(mut self, top_p: f64) -> Self {
249        self.agent.generation.top_p = Some(top_p);
250        self
251    }
252
253    /// Sets the maximum number of output tokens for every Agent turn.
254    #[must_use]
255    pub fn max_output_tokens(mut self, max_output_tokens: u64) -> Self {
256        self.agent.generation.max_output_tokens = Some(max_output_tokens);
257        self
258    }
259
260    /// Selects streaming or complete response delivery for every Agent turn.
261    #[must_use]
262    pub const fn response_mode(mut self, response_mode: ResponseMode) -> Self {
263        self.agent.response_mode = response_mode;
264        self
265    }
266
267    /// Adds namespaced provider options to every Agent turn.
268    #[must_use]
269    pub fn provider_options(
270        mut self,
271        provider: impl Into<String>,
272        options: serde_json::Value,
273    ) -> Self {
274        self.agent.provider_options.insert(provider.into(), options);
275        self
276    }
277
278    /// Replaces all local Agent configuration.
279    #[must_use]
280    pub const fn config(mut self, config: AgentConfig) -> Self {
281        self.agent.config = config;
282        self
283    }
284
285    /// Shares a write-ahead effect coordinator with this Agent.
286    #[must_use]
287    pub fn effect_executor(mut self, effects: EffectExecutor) -> Self {
288        self.agent.effects = effects;
289        self
290    }
291
292    /// Sets recovery behavior for ambiguous callable effects.
293    #[must_use]
294    pub const fn effect_recovery_policy(mut self, policy: EffectRecoveryPolicy) -> Self {
295        self.agent.effect_recovery = policy;
296        self
297    }
298
299    /// Validates registrations and returns the canonical Agent.
300    ///
301    /// # Errors
302    ///
303    /// Returns [`AgentBuildError`] for blank identity, invalid turn bounds,
304    /// duplicate registrations, or Tool/Agent name collisions.
305    pub fn build(self) -> Result<Agent, AgentBuildError> {
306        if let Some(error) = self.error {
307            return Err(error);
308        }
309        if self.agent.name.trim().is_empty() {
310            return Err(AgentBuildError::EmptyName);
311        }
312        if self.agent.config.max_turns == 0 {
313            return Err(AgentBuildError::ZeroMaxTurns);
314        }
315        if let Some(collision) = self
316            .agent
317            .agents
318            .model_specs()
319            .into_iter()
320            .find(|spec| self.agent.tools.contains(&spec.name))
321        {
322            return Err(AgentBuildError::CallableNameCollision(collision.name));
323        }
324        Ok(self.agent)
325    }
326
327    /// Builds the Agent and runs one ergonomic prompt.
328    ///
329    /// This removes the explicit build step for one-shot usage while retaining
330    /// the complete canonical outcome. Use [`Self::build`] when the Agent will
331    /// be reused or executed with an explicit runtime context.
332    pub fn prompt(
333        self,
334        input: impl Into<String> + Send + 'static,
335    ) -> AgentFuture<'static, Result<AgentOutcome, AgentPromptError>> {
336        let input = input.into();
337        Box::pin(async move {
338            let agent = self.build()?;
339            Ok(agent.prompt(input).await?)
340        })
341    }
342
343    /// Builds the Agent, runs one ergonomic prompt, and returns only
344    /// model-visible text.
345    ///
346    /// This is the shortest path from provider configuration to a text answer.
347    /// Use [`Self::prompt`] when transcript, usage, warnings, and provider
348    /// events must be preserved.
349    pub fn prompt_text(
350        self,
351        input: impl Into<String> + Send + 'static,
352    ) -> AgentFuture<'static, Result<String, AgentPromptError>> {
353        let input = input.into();
354        Box::pin(async move {
355            let agent = self.build()?;
356            Ok(agent.prompt_text(input).await?)
357        })
358    }
359
360    /// Builds an Agent whose schema and local decoder are bound to `T`.
361    ///
362    /// This is the preferred terminal builder operation for structured output
363    /// because a different decode type cannot be selected later by mistake.
364    ///
365    /// # Errors
366    ///
367    /// Returns [`AgentBuildError`] under the same validation rules as
368    /// [`Self::build`].
369    pub fn build_structured<T>(
370        self,
371        name: impl Into<String>,
372    ) -> Result<StructuredAgent<T>, AgentBuildError>
373    where
374        T: JsonSchema,
375    {
376        self.structured_output::<T>(name)
377            .build()
378            .map(StructuredAgent::new)
379    }
380}
381
382impl std::fmt::Debug for AgentBuilder {
383    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
384        formatter
385            .debug_struct("AgentBuilder")
386            .field("agent", &self.agent.name)
387            .field("error", &self.error)
388            .finish_non_exhaustive()
389    }
390}
391
392#[cfg(test)]
393mod tests {
394    use std::{collections::BTreeMap, sync::Arc};
395
396    use runifold_core::{CapabilityId, CapabilitySet, EffectClass, RiskLevel};
397    use runifold_model::{
398        ContentPart, FinishReason, ModelRef, ModelStreamEvent, OutputFormat, ProviderToolSpec,
399        ResponseMode,
400    };
401    use runifold_testkit::ScriptedModel;
402    use runifold_tool::{Tool, ToolContext, ToolDescriptor, ToolError, ToolFuture, ToolOutput};
403    use schemars::JsonSchema;
404    use serde::Deserialize;
405    use serde_json::json;
406
407    use crate::{Agent, AgentBuildError, AgentDescriptor, AgentPromptError};
408
409    struct TestTool {
410        descriptor: ToolDescriptor,
411    }
412
413    #[derive(Deserialize, JsonSchema)]
414    struct TypedAnswer {
415        value: u32,
416    }
417
418    impl TestTool {
419        fn named(name: &str) -> Self {
420            Self {
421                descriptor: ToolDescriptor {
422                    id: CapabilityId::new(),
423                    name: name.into(),
424                    version: "1".into(),
425                    description: "test".into(),
426                    input_schema: json!({"type": "object"}),
427                    output_schema: json!({"type": "object"}),
428                    effect: EffectClass::Pure,
429                    risk: RiskLevel::Low,
430                    metadata: BTreeMap::new(),
431                },
432            }
433        }
434    }
435
436    impl Tool for TestTool {
437        fn descriptor(&self) -> &ToolDescriptor {
438            &self.descriptor
439        }
440
441        fn invoke(
442            &self,
443            input: serde_json::Value,
444            _context: ToolContext,
445        ) -> ToolFuture<'_, Result<ToolOutput, ToolError>> {
446            Box::pin(async move { Ok(ToolOutput::model_visible(input)) })
447        }
448    }
449
450    #[test]
451    fn fluent_builder_assembles_the_canonical_agent() {
452        let model = Arc::new(ScriptedModel::new());
453        let agent = Agent::builder("worker", model, ModelRef::new("test", "scripted"))
454            .system("Be precise")
455            .tool(TestTool::named("lookup"))
456            .max_turns(4)
457            .build()
458            .unwrap();
459
460        assert_eq!(agent.name, "worker");
461        assert_eq!(agent.instructions.len(), 1);
462        assert!(agent.tools.contains("lookup"));
463        assert_eq!(agent.config.max_turns, 4);
464        assert_eq!(agent.callable_capabilities().len(), 1);
465    }
466
467    #[test]
468    fn builder_retains_generation_provider_and_delivery_controls() {
469        let provider_tool = ProviderToolSpec::new("ark", "web_search").unwrap();
470        let agent = Agent::builder(
471            "researcher",
472            Arc::new(ScriptedModel::new()),
473            ModelRef::new("ark", "doubao"),
474        )
475        .temperature(0.2)
476        .top_p(0.8)
477        .max_output_tokens(4_096)
478        .response_mode(ResponseMode::Complete)
479        .provider_tool(provider_tool)
480        .provider_options("ark", json!({"thinking": {"type": "enabled"}}))
481        .build()
482        .unwrap();
483
484        assert_eq!(agent.generation.temperature, Some(0.2));
485        assert_eq!(agent.generation.top_p, Some(0.8));
486        assert_eq!(agent.generation.max_output_tokens, Some(4_096));
487        assert_eq!(agent.response_mode, ResponseMode::Complete);
488        assert_eq!(agent.provider_tools[0].tool_type, "web_search");
489        assert_eq!(agent.provider_options["ark"]["thinking"]["type"], "enabled");
490    }
491
492    #[test]
493    fn build_rejects_tool_and_agent_name_collisions() {
494        let model = Arc::new(ScriptedModel::new());
495        let child = Arc::new(Agent::new(
496            "child",
497            model.clone(),
498            ModelRef::new("test", "child"),
499        ));
500        let error = Agent::builder("parent", model, ModelRef::new("test", "parent"))
501            .tool(TestTool::named("search"))
502            .child(
503                AgentDescriptor::new("search", "delegate search"),
504                child,
505                CapabilitySet::new(),
506            )
507            .build()
508            .unwrap_err();
509
510        assert_eq!(
511            error,
512            AgentBuildError::CallableNameCollision("search".into())
513        );
514    }
515
516    #[test]
517    fn builder_derives_a_strict_output_schema_from_a_rust_type() {
518        let example = TypedAnswer { value: 7 };
519        assert_eq!(example.value, 7);
520        let agent = Agent::builder(
521            "worker",
522            Arc::new(ScriptedModel::new()),
523            ModelRef::new("test", "scripted"),
524        )
525        .structured_output::<TypedAnswer>("typed_answer")
526        .build()
527        .unwrap();
528
529        let OutputFormat::JsonSchema {
530            name,
531            schema,
532            strict,
533        } = agent.output_format
534        else {
535            panic!("expected JSON-schema output");
536        };
537        assert_eq!(name, "typed_answer");
538        assert!(strict);
539        assert_eq!(schema["properties"]["value"]["type"], "integer");
540    }
541
542    #[test]
543    fn builder_prompt_text_is_a_single_use_golden_path() {
544        let model = ScriptedModel::new();
545        model.enqueue([
546            ModelStreamEvent::ResponseStarted {
547                id: Some("response-1".into()),
548                model: ModelRef::new("test", "scripted"),
549            },
550            ModelStreamEvent::ContentPartCompleted {
551                index: 0,
552                part: ContentPart::text("done"),
553            },
554            ModelStreamEvent::ResponseCompleted {
555                finish_reason: FinishReason::Stop,
556                provider_metadata: BTreeMap::new(),
557            },
558        ]);
559
560        let text = futures_executor::block_on(
561            Agent::builder("worker", Arc::new(model), ModelRef::new("test", "scripted"))
562                .system("Be precise")
563                .prompt_text("start"),
564        )
565        .unwrap();
566
567        assert_eq!(text, "done");
568    }
569
570    #[test]
571    fn builder_prompt_reports_build_failures_before_model_execution() {
572        let error = futures_executor::block_on(
573            Agent::builder(
574                "",
575                Arc::new(ScriptedModel::new()),
576                ModelRef::new("test", "scripted"),
577            )
578            .prompt("start"),
579        )
580        .unwrap_err();
581
582        assert!(matches!(
583            error,
584            AgentPromptError::Build(AgentBuildError::EmptyName)
585        ));
586    }
587}