1use 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
19pub struct GenerateAgentTool {
27 spawner: Arc<AgentSpawner>,
28 registry: Arc<AgentRegistry>,
29 llm: Arc<LLMRegistry>,
30 enriched_description: String,
32}
33
34#[derive(Debug, Deserialize, JsonSchema)]
35#[allow(dead_code)]
36struct GenerateAgentInput {
37 description: String,
39 name: String,
41 #[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 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 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 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 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 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 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 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
263fn 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
290fn 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
301pub struct SendMessageTool {
307 registry: Arc<AgentRegistry>,
308 sender_id: String,
310}
311
312#[derive(Debug, Deserialize, JsonSchema)]
313#[allow(dead_code)]
314struct SendMessageInput {
315 to: String,
317 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
374pub 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
424pub struct RemoveAgentTool {
430 registry: Arc<AgentRegistry>,
431}
432
433#[derive(Debug, Deserialize, JsonSchema)]
434#[allow(dead_code)]
435struct RemoveAgentInput {
436 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 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 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}