Skip to main content

ai_agents_runtime/spawner/
tools.rs

1//! Built-in tools for inter-agent messaging and dynamic agent management.
2
3use std::sync::Arc;
4
5use async_trait::async_trait;
6use schemars::JsonSchema;
7use serde::Deserialize;
8use serde_json::{Value, json};
9
10use ai_agents_core::{ChatMessage, LLMProvider, Tool, ToolResult};
11use ai_agents_llm::LLMRegistry;
12use ai_agents_observability::{ObservationPurpose, with_observation_purpose};
13use ai_agents_tools::generate_schema;
14
15use super::registry::AgentRegistry;
16use super::spawner::AgentSpawner;
17use crate::turn_context::current_turn_actor_context;
18
19//
20// GenerateAgentTool
21//
22
23/// Tool that lets a parent agent generate and spawn a new agent.
24///
25/// The tool description is built dynamically from template metadata so the LLM can discover available templates and their variables automatically.
26pub struct GenerateAgentTool {
27    spawner: Arc<AgentSpawner>,
28    registry: Arc<AgentRegistry>,
29    llm: Arc<LLMRegistry>,
30    /// Pre-built description including available templates and their variables.
31    enriched_description: String,
32}
33
34#[derive(Debug, Deserialize, JsonSchema)]
35#[allow(dead_code)]
36struct GenerateAgentInput {
37    /// Natural language description of the agent to create.
38    description: String,
39    /// Agent name / ID.
40    name: String,
41    /// Optional: named template to render instead of LLM generation.
42    #[serde(default)]
43    template: Option<String>,
44}
45
46impl GenerateAgentTool {
47    pub fn new(
48        spawner: Arc<AgentSpawner>,
49        registry: Arc<AgentRegistry>,
50        llm: Arc<LLMRegistry>,
51    ) -> Self {
52        let enriched_description = Self::build_description(&spawner);
53        Self {
54            spawner,
55            registry,
56            llm,
57            enriched_description,
58        }
59    }
60
61    /// Build tool description from template metadata (description + variables).
62    fn build_description(spawner: &AgentSpawner) -> String {
63        let mut desc = String::from(
64            "Generate and spawn a new AI agent from a description. \
65             Provide a natural language description of the agent's \
66             personality, capabilities, and purpose.",
67        );
68
69        let templates = spawner.templates();
70        if templates.is_empty() {
71            return desc;
72        }
73
74        desc.push_str("\n\nAvailable templates (pass name as \"template\" field):");
75        for (name, tpl) in templates {
76            desc.push_str("\n  ");
77            desc.push_str(name);
78            if let Some(ref d) = tpl.description {
79                desc.push_str(": ");
80                desc.push_str(d);
81            }
82            if let Some(ref vars) = tpl.variables {
83                for (var_name, var_desc) in vars {
84                    desc.push_str("\n    - ");
85                    desc.push_str(var_name);
86                    desc.push_str(": ");
87                    desc.push_str(var_desc);
88                }
89            }
90        }
91
92        desc.push_str(
93            "\n\nWhen using a template, pass its variables as additional fields \
94             alongside name and description.",
95        );
96
97        desc
98    }
99}
100
101#[async_trait]
102impl Tool for GenerateAgentTool {
103    fn id(&self) -> &str {
104        "spawn_agent"
105    }
106
107    fn name(&self) -> &str {
108        "Spawn Agent"
109    }
110
111    fn description(&self) -> &str {
112        &self.enriched_description
113    }
114
115    fn input_schema(&self) -> Value {
116        generate_schema::<GenerateAgentInput>()
117    }
118
119    // Keeps template bypass and one repair attempt while separating generation and repair providers.
120    async fn execute(&self, args: Value, _ctx: ai_agents_core::ToolExecutionContext) -> ToolResult {
121        let description = match args.get("description").and_then(|v| v.as_str()) {
122            Some(d) => d,
123            None => return ToolResult::error("missing required field: description"),
124        };
125        let name = match args.get("name").and_then(|v| v.as_str()) {
126            Some(n) => n,
127            None => return ToolResult::error("missing required field: name"),
128        };
129        let template = args.get("template").and_then(|v| v.as_str());
130
131        //
132        // Template path
133        if let Some(tpl_name) = template {
134            let mut vars = std::collections::HashMap::new();
135            vars.insert("name".to_string(), name.to_string());
136            vars.insert("description".to_string(), description.to_string());
137
138            // Forward any extra top-level string fields as template variables.
139            if let Some(obj) = args.as_object() {
140                for (k, v) in obj {
141                    if k == "description" || k == "name" || k == "template" {
142                        continue;
143                    }
144                    if let Some(s) = v.as_str() {
145                        vars.insert(k.clone(), s.to_string());
146                    }
147                }
148            }
149
150            return match self.spawner.spawn_from_template(tpl_name, vars).await {
151                Ok(agent) => {
152                    let id = agent.id.clone();
153                    match self.registry.register(agent).await {
154                        Ok(()) => ToolResult::ok(
155                            json!({"id": id, "source": "template", "template": tpl_name})
156                                .to_string(),
157                        ),
158                        Err(e) => ToolResult::error(format!("registry error: {}", e)),
159                    }
160                }
161                Err(e) => ToolResult::error(format!("template spawn failed: {}", e)),
162            };
163        }
164
165        //
166        // LLM generation
167        let llm: Arc<dyn LLMProvider> = match self
168            .llm
169            .resolve_role_override(ai_agents_llm::LLMRole::SpawnerGeneration, None)
170        {
171            Ok(Some(resolved)) => resolved.provider,
172            Ok(None) => match self.llm.router() {
173                Ok(l) => l,
174                Err(_) => match self.llm.default() {
175                    Ok(l) => l,
176                    Err(e) => return ToolResult::error(format!("no LLM available: {}", e)),
177                },
178            },
179            Err(error) => return ToolResult::error(error.to_string()),
180        };
181
182        let prompt = build_generation_prompt(name, description);
183        let messages = vec![ChatMessage::user(prompt)];
184
185        let yaml = match with_observation_purpose(
186            ObservationPurpose::OrchestrationRouting,
187            llm.complete(&messages, None),
188        )
189        .await
190        {
191            Ok(resp) => strip_code_fences(&resp.content),
192            Err(e) => return ToolResult::error(format!("LLM generation failed: {}", e)),
193        };
194
195        // First attempt: parse and spawn.
196        match self.spawner.spawn_from_yaml(&yaml).await {
197            Ok(agent) => {
198                let id = agent.id.clone();
199                return match self.registry.register(agent).await {
200                    Ok(()) => {
201                        ToolResult::ok(json!({"id": id, "source": "llm_generated"}).to_string())
202                    }
203                    Err(e) => ToolResult::error(format!("registry error: {}", e)),
204                };
205            }
206            Err(first_err) => {
207                // Retry once with an error-correction prompt.
208                let retry_prompt = format!(
209                    "The YAML you generated was invalid:\n{}\n\nError: {}\n\n\
210                     Please fix the YAML and return ONLY valid YAML with no markdown fences.",
211                    yaml, first_err
212                );
213                let retry_messages = vec![
214                    ChatMessage::user(build_generation_prompt(name, description)),
215                    ChatMessage::assistant(&yaml),
216                    ChatMessage::user(retry_prompt),
217                ];
218
219                let repair_llm = match self
220                    .llm
221                    .resolve_role_override(ai_agents_llm::LLMRole::SpawnerRepair, None)
222                {
223                    Ok(Some(resolved)) => resolved.provider,
224                    Ok(None) => llm.clone(),
225                    Err(error) => return ToolResult::error(error.to_string()),
226                };
227                let retry_yaml = match with_observation_purpose(
228                    ObservationPurpose::OrchestrationRouting,
229                    repair_llm.complete(&retry_messages, None),
230                )
231                .await
232                {
233                    Ok(resp) => strip_code_fences(&resp.content),
234                    Err(e) => {
235                        return ToolResult::error(format!(
236                            "LLM retry failed: {} (original error: {})",
237                            e, first_err
238                        ));
239                    }
240                };
241
242                match self.spawner.spawn_from_yaml(&retry_yaml).await {
243                    Ok(agent) => {
244                        let id = agent.id.clone();
245                        match self.registry.register(agent).await {
246                            Ok(()) => ToolResult::ok(
247                                json!({"id": id, "source": "llm_generated", "retried": true})
248                                    .to_string(),
249                            ),
250                            Err(e) => ToolResult::error(format!("registry error: {}", e)),
251                        }
252                    }
253                    Err(e) => ToolResult::error(format!(
254                        "spawn failed after retry: {} (original: {})",
255                        e, first_err
256                    )),
257                }
258            }
259        }
260    }
261}
262
263/// Build the YAML-generation prompt sent to the LLM.
264fn build_generation_prompt(name: &str, description: &str) -> String {
265    format!(
266        "Generate a valid YAML agent specification.\n\n\
267         Required fields:\n\
268         - name: string (the agent's name)\n\
269         - system_prompt: string (detailed behavioral instructions)\n\n\
270         Optional fields: memory (type, max_messages, compress_threshold), \
271         reasoning (mode: auto|cot|react), disambiguation (enabled: true/false).\n\n\
272         Example:\n\
273         ```yaml\n\
274         name: Helper\n\
275         system_prompt: |\n\
276           You are a helpful assistant who answers concisely.\n\
277         memory:\n\
278           type: compacting\n\
279           max_messages: 100\n\
280           compress_threshold: 20\n\
281         ```\n\n\
282         Now generate a spec for:\n\
283         Name: {}\n\
284         Description: {}\n\n\
285         Return ONLY the YAML content. No markdown fences, no commentary.",
286        name, description
287    )
288}
289
290/// Strip optional markdown code fences from LLM output.
291fn strip_code_fences(text: &str) -> String {
292    let trimmed = text.trim();
293    let trimmed = trimmed
294        .strip_prefix("```yaml")
295        .or_else(|| trimmed.strip_prefix("```"))
296        .unwrap_or(trimmed);
297    let trimmed = trimmed.strip_suffix("```").unwrap_or(trimmed);
298    trimmed.trim().to_string()
299}
300
301//
302// SendMessageTool
303//
304
305/// Tool that sends a message from the owning agent to another registered agent.
306pub struct SendMessageTool {
307    registry: Arc<AgentRegistry>,
308    /// ID of the agent that owns this tool (the sender).
309    sender_id: String,
310}
311
312#[derive(Debug, Deserialize, JsonSchema)]
313#[allow(dead_code)]
314struct SendMessageInput {
315    /// Target agent ID.
316    to: String,
317    /// Message to send.
318    message: String,
319}
320
321impl SendMessageTool {
322    pub fn new(registry: Arc<AgentRegistry>, sender_id: impl Into<String>) -> Self {
323        Self {
324            registry,
325            sender_id: sender_id.into(),
326        }
327    }
328}
329
330#[async_trait]
331impl Tool for SendMessageTool {
332    fn id(&self) -> &str {
333        "send_agent_message"
334    }
335
336    fn name(&self) -> &str {
337        "Send Agent Message"
338    }
339
340    fn description(&self) -> &str {
341        "Send a message to another registered agent and receive its response."
342    }
343
344    fn input_schema(&self) -> Value {
345        generate_schema::<SendMessageInput>()
346    }
347
348    async fn execute(&self, args: Value, _ctx: ai_agents_core::ToolExecutionContext) -> ToolResult {
349        let to = match args.get("to").and_then(|v| v.as_str()) {
350            Some(t) => t,
351            None => return ToolResult::error("missing required field: to"),
352        };
353        let message = match args.get("message").and_then(|v| v.as_str()) {
354            Some(m) => m,
355            None => return ToolResult::error("missing required field: message"),
356        };
357
358        let actor_context = current_turn_actor_context()
359            .unwrap_or_default()
360            .for_sender(self.sender_id.clone());
361        match self
362            .registry
363            .send_with_actor_context(&self.sender_id, to, message, actor_context)
364            .await
365        {
366            Ok(response) => {
367                ToolResult::ok(json!({"from": to, "response": response.content}).to_string())
368            }
369            Err(e) => ToolResult::error(format!("send failed: {}", e)),
370        }
371    }
372}
373
374//
375// ListAgentsTool
376//
377
378/// Tool that lists all agents currently registered in the registry.
379pub struct ListAgentsTool {
380    registry: Arc<AgentRegistry>,
381}
382
383#[derive(Debug, Deserialize, JsonSchema)]
384#[allow(dead_code)]
385struct ListAgentsInput {}
386
387impl ListAgentsTool {
388    pub fn new(registry: Arc<AgentRegistry>) -> Self {
389        Self { registry }
390    }
391}
392
393#[async_trait]
394impl Tool for ListAgentsTool {
395    fn id(&self) -> &str {
396        "list_agents"
397    }
398
399    fn name(&self) -> &str {
400        "List Agents"
401    }
402
403    fn description(&self) -> &str {
404        "List all currently registered agents with their IDs and names."
405    }
406
407    fn input_schema(&self) -> Value {
408        generate_schema::<ListAgentsInput>()
409    }
410
411    async fn execute(
412        &self,
413        _args: Value,
414        _ctx: ai_agents_core::ToolExecutionContext,
415    ) -> ToolResult {
416        let agents = self.registry.list();
417        match serde_json::to_string(&agents) {
418            Ok(json) => ToolResult::ok(json),
419            Err(e) => ToolResult::error(format!("serialization error: {}", e)),
420        }
421    }
422}
423
424//
425// RemoveAgentTool
426//
427
428/// Tool that removes an agent from the registry by ID.
429pub struct RemoveAgentTool {
430    registry: Arc<AgentRegistry>,
431}
432
433#[derive(Debug, Deserialize, JsonSchema)]
434#[allow(dead_code)]
435struct RemoveAgentInput {
436    /// Agent ID to remove.
437    id: String,
438}
439
440impl RemoveAgentTool {
441    pub fn new(registry: Arc<AgentRegistry>) -> Self {
442        Self { registry }
443    }
444}
445
446#[async_trait]
447impl Tool for RemoveAgentTool {
448    fn id(&self) -> &str {
449        "remove_agent"
450    }
451
452    fn name(&self) -> &str {
453        "Remove Agent"
454    }
455
456    fn description(&self) -> &str {
457        "Remove a registered agent by its ID."
458    }
459
460    fn input_schema(&self) -> Value {
461        generate_schema::<RemoveAgentInput>()
462    }
463
464    async fn execute(&self, args: Value, _ctx: ai_agents_core::ToolExecutionContext) -> ToolResult {
465        let id = match args.get("id").and_then(|v| v.as_str()) {
466            Some(i) => i,
467            None => return ToolResult::error("missing required field: id"),
468        };
469
470        match self.registry.remove(id).await {
471            Some(removed) => ToolResult::ok(json!({"removed": true, "id": removed.id}).to_string()),
472            None => ToolResult::error(format!("agent not found: {}", id)),
473        }
474    }
475}
476
477#[cfg(test)]
478mod tests {
479    use super::super::spawner::ResolvedTemplate;
480    use super::*;
481    use std::collections::HashMap;
482
483    #[test]
484    fn test_strip_code_fences_yaml() {
485        let input = "```yaml\nname: Test\nsystem_prompt: hi\n```";
486        assert_eq!(strip_code_fences(input), "name: Test\nsystem_prompt: hi");
487    }
488
489    #[test]
490    fn test_strip_code_fences_bare() {
491        let input = "```\nname: Test\n```";
492        assert_eq!(strip_code_fences(input), "name: Test");
493    }
494
495    #[test]
496    fn test_strip_code_fences_none() {
497        let input = "name: Test\nsystem_prompt: hi";
498        assert_eq!(strip_code_fences(input), input);
499    }
500
501    #[test]
502    fn test_build_generation_prompt_contains_name() {
503        let prompt = build_generation_prompt("Gormund", "A gruff blacksmith");
504        assert!(prompt.contains("Gormund"));
505        assert!(prompt.contains("gruff blacksmith"));
506    }
507
508    #[test]
509    fn test_tool_ids_are_unique() {
510        let ids = [
511            "spawn_agent",
512            "send_agent_message",
513            "list_agents",
514            "remove_agent",
515        ];
516        let unique: std::collections::HashSet<_> = ids.iter().collect();
517        assert_eq!(unique.len(), ids.len());
518    }
519
520    #[test]
521    fn test_build_description_no_templates() {
522        let spawner = AgentSpawner::new();
523        let desc = GenerateAgentTool::build_description(&spawner);
524        assert!(desc.contains("Generate and spawn"));
525        assert!(!desc.contains("Available templates"));
526    }
527
528    #[test]
529    fn test_build_description_with_templates() {
530        let mut templates = HashMap::new();
531        templates.insert(
532            "npc_base".to_string(),
533            ResolvedTemplate {
534                content: "name: test".to_string(),
535                description: Some("General-purpose NPC".to_string()),
536                variables: Some({
537                    let mut v = HashMap::new();
538                    v.insert("role".to_string(), "NPC occupation".to_string());
539                    v.insert(
540                        "personality".to_string(),
541                        "Personality description".to_string(),
542                    );
543                    v
544                }),
545            },
546        );
547        let spawner = AgentSpawner::new().with_templates(templates);
548        let desc = GenerateAgentTool::build_description(&spawner);
549        assert!(desc.contains("Available templates"));
550        assert!(desc.contains("npc_base"));
551        assert!(desc.contains("General-purpose NPC"));
552        assert!(desc.contains("role"));
553        assert!(desc.contains("NPC occupation"));
554        assert!(desc.contains("personality"));
555    }
556
557    #[test]
558    fn test_build_description_template_no_metadata() {
559        let mut templates = HashMap::new();
560        templates.insert(
561            "bare".to_string(),
562            ResolvedTemplate {
563                content: "name: test".to_string(),
564                description: None,
565                variables: None,
566            },
567        );
568        let spawner = AgentSpawner::new().with_templates(templates);
569        let desc = GenerateAgentTool::build_description(&spawner);
570        assert!(desc.contains("Available templates"));
571        assert!(desc.contains("bare"));
572        // No description or variables appended
573        assert!(!desc.contains("NPC"));
574    }
575
576    #[test]
577    fn test_spawn_agent_schema_has_required_fields() {
578        let schema = generate_schema::<GenerateAgentInput>();
579        let props = schema.get("properties").expect("should have properties");
580        assert!(props.get("description").is_some());
581        assert!(props.get("name").is_some());
582        assert!(props.get("template").is_some());
583        let required = schema.get("required").expect("should have required");
584        let req_arr: Vec<&str> = required
585            .as_array()
586            .unwrap()
587            .iter()
588            .map(|v| v.as_str().unwrap())
589            .collect();
590        assert!(req_arr.contains(&"description"));
591        assert!(req_arr.contains(&"name"));
592        // template is optional (Option<String>), should not be required
593        assert!(!req_arr.contains(&"template"));
594    }
595
596    #[test]
597    fn test_send_agent_message_schema_has_required_fields() {
598        let schema = generate_schema::<SendMessageInput>();
599        let props = schema.get("properties").expect("should have properties");
600        assert!(props.get("to").is_some());
601        assert!(props.get("message").is_some());
602    }
603
604    #[test]
605    fn test_remove_agent_schema_has_id() {
606        let schema = generate_schema::<RemoveAgentInput>();
607        let props = schema.get("properties").expect("should have properties");
608        assert!(props.get("id").is_some());
609    }
610
611    #[test]
612    fn test_list_agents_schema_is_object() {
613        let schema = generate_schema::<ListAgentsInput>();
614        assert_eq!(schema.get("type").and_then(|v| v.as_str()), Some("object"));
615    }
616}