1use serde::{Deserialize, Serialize};
16use serde_json::json;
17use std::{any::Any, collections::BTreeMap, future::Future, marker::PhantomData, sync::Arc};
18
19use crate::{
20 BoxError, BoxPinFut, Function, ToolGroup, ToolGroupInfo,
21 context::AgentContext,
22 model::{AgentOutput, FunctionDefinition, Resource},
23 registry::{collect_groups, select_by_names},
24 select_resources, validate_function_name,
25};
26
27#[derive(Debug, Clone, Deserialize, Serialize)]
29pub struct AgentArgs {
30 pub prompt: String,
32}
33
34pub trait Agent<C>: Send + Sync
39where
40 C: AgentContext + Send + Sync,
41{
42 fn name(&self) -> String;
53
54 fn description(&self) -> String;
56
57 fn definition(&self) -> FunctionDefinition {
62 FunctionDefinition {
63 name: self.name().to_ascii_lowercase(),
64 description: self.description(),
65 parameters: json!({
66 "type": "object",
67 "description": "Run this agent on a focused task. Provide a self-contained prompt with the goal, relevant context, constraints, and expected output.",
68 "properties": {
69 "prompt": {
70 "type": "string",
71 "description": "The task for this agent. Include the objective, relevant context, constraints, preferred workflow or deliverable, and any success criteria needed to complete the work.",
72 "minLength": 1
73 },
74 },
75 "required": ["prompt"],
76 "additionalProperties": false
77 }),
78 strict: Some(true),
79 }
80 }
81
82 fn group(&self) -> Option<ToolGroupInfo> {
89 None
90 }
91
92 fn supported_resource_tags(&self) -> Vec<String> {
101 Vec::new()
102 }
103
104 fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
106 let supported_tags = self.supported_resource_tags();
107 select_resources(resources, &supported_tags)
108 }
109
110 fn init(&self, _ctx: C) -> impl Future<Output = Result<(), BoxError>> + Send {
114 std::future::ready(Ok(()))
115 }
116
117 fn tool_dependencies(&self) -> Vec<String> {
121 Vec::new()
122 }
123
124 fn run(
134 &self,
135 ctx: C,
136 prompt: String,
137 resources: Vec<Resource>,
138 ) -> impl Future<Output = Result<AgentOutput, BoxError>> + Send;
139}
140
141pub trait DynAgent<C>: Send + Sync
146where
147 C: AgentContext + Send + Sync,
148{
149 fn as_any(&self) -> &(dyn Any + Send + Sync);
151
152 fn into_any(self: Arc<Self>) -> Arc<dyn Any + Send + Sync>;
154
155 fn label(&self) -> &str;
162
163 fn name(&self) -> String;
165
166 fn definition(&self) -> FunctionDefinition;
168
169 fn tool_dependencies(&self) -> Vec<String>;
171
172 fn group(&self) -> Option<ToolGroupInfo>;
174
175 fn supported_resource_tags(&self) -> Vec<String>;
177
178 fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
180 select_resources(resources, &self.supported_resource_tags())
181 }
182
183 fn init(&self, ctx: C) -> BoxPinFut<Result<(), BoxError>>;
185
186 fn run(
188 &self,
189 ctx: C,
190 prompt: String,
191 resources: Vec<Resource>,
192 ) -> BoxPinFut<Result<AgentOutput, BoxError>>;
193}
194
195impl<C> dyn DynAgent<C>
196where
197 C: AgentContext + Send + Sync + 'static,
198{
199 pub fn downcast_ref<T>(&self) -> Option<&T>
201 where
202 T: Agent<C> + 'static,
203 {
204 self.as_any().downcast_ref::<T>()
205 }
206
207 pub fn downcast<T>(self: Arc<Self>) -> Result<Arc<T>, Arc<Self>>
209 where
210 T: Agent<C> + 'static,
211 {
212 match self.clone().into_any().downcast::<T>() {
213 Ok(agent) => Ok(agent),
214 Err(_) => Err(self),
215 }
216 }
217}
218
219struct AgentWrapper<T, C>
221where
222 T: Agent<C> + 'static,
223 C: AgentContext + Send + Sync + 'static,
224{
225 inner: Arc<T>,
226 label: String,
227 _phantom: PhantomData<C>,
228}
229
230impl<T, C> DynAgent<C> for AgentWrapper<T, C>
231where
232 T: Agent<C> + 'static,
233 C: AgentContext + Send + Sync + 'static,
234{
235 fn as_any(&self) -> &(dyn Any + Send + Sync) {
236 self.inner.as_ref()
237 }
238
239 fn into_any(self: Arc<Self>) -> Arc<dyn Any + Send + Sync> {
240 self.inner.clone()
241 }
242
243 fn label(&self) -> &str {
244 &self.label
245 }
246
247 fn name(&self) -> String {
248 self.inner.name()
249 }
250
251 fn definition(&self) -> FunctionDefinition {
252 self.inner.definition()
253 }
254
255 fn tool_dependencies(&self) -> Vec<String> {
256 self.inner.tool_dependencies()
257 }
258
259 fn group(&self) -> Option<ToolGroupInfo> {
260 self.inner.group()
261 }
262
263 fn supported_resource_tags(&self) -> Vec<String> {
264 self.inner.supported_resource_tags()
265 }
266
267 fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
268 self.inner.select_resources(resources)
269 }
270
271 fn init(&self, ctx: C) -> BoxPinFut<Result<(), BoxError>> {
272 let agent = self.inner.clone();
273 Box::pin(async move { agent.init(ctx).await })
274 }
275
276 fn run(
277 &self,
278 ctx: C,
279 prompt: String,
280 resources: Vec<Resource>,
281 ) -> BoxPinFut<Result<AgentOutput, BoxError>> {
282 let agent = self.inner.clone();
283 Box::pin(async move { agent.run(ctx, prompt, resources).await })
284 }
285}
286
287pub struct AgentSet<C: AgentContext> {
292 set: BTreeMap<String, Arc<dyn DynAgent<C>>>,
298}
299
300impl<C: AgentContext> Default for AgentSet<C> {
301 fn default() -> Self {
302 Self {
303 set: BTreeMap::new(),
304 }
305 }
306}
307
308impl<C> AgentSet<C>
309where
310 C: AgentContext + Send + Sync + 'static,
311{
312 pub fn new() -> Self {
314 Self::default()
315 }
316
317 pub fn contains(&self, name: &str) -> bool {
319 self.set.contains_key(&name.to_ascii_lowercase())
320 }
321
322 pub fn contains_lowercase(&self, lowercase_name: &str) -> bool {
324 self.set.contains_key(lowercase_name)
325 }
326
327 pub fn names(&self) -> Vec<String> {
329 self.set.keys().cloned().collect()
330 }
331
332 pub fn groups(&self) -> Vec<ToolGroup> {
339 collect_groups(self.set.iter().map(|(name, agent)| (name, agent.group())))
340 }
341
342 pub fn definition(&self, name: &str) -> Option<FunctionDefinition> {
344 self.set
345 .get(&name.to_ascii_lowercase())
346 .map(|agent| agent.definition())
347 }
348
349 pub fn definitions(&self, names: Option<&[String]>) -> Vec<FunctionDefinition> {
359 select_by_names(&self.set, names, |agent| agent.definition())
360 }
361
362 pub fn functions(&self, names: Option<&[String]>) -> Vec<Function> {
372 select_by_names(&self.set, names, |agent| Function {
373 definition: agent.definition(),
374 supported_resource_tags: agent.supported_resource_tags(),
375 })
376 }
377
378 pub fn select_resources(&self, name: &str, resources: &mut Vec<Resource>) -> Vec<Resource> {
380 if resources.is_empty() {
381 return Vec::new();
382 }
383
384 self.set
385 .get(&name.to_ascii_lowercase())
386 .map(|agent| agent.select_resources(resources))
387 .unwrap_or_default()
388 }
389
390 pub fn add<T>(&mut self, agent: Arc<T>, label: Option<String>) -> Result<(), BoxError>
396 where
397 T: Agent<C> + Send + Sync + 'static,
398 {
399 let label = label.unwrap_or_else(|| agent.name().to_ascii_lowercase());
400 self.add_dyn(Arc::new(AgentWrapper {
401 inner: agent,
402 label,
403 _phantom: PhantomData,
404 }))
405 }
406
407 pub fn add_dyn(&mut self, agent: Arc<dyn DynAgent<C>>) -> Result<(), BoxError> {
412 let name = agent.name().to_ascii_lowercase();
413 validate_function_name(&name)?;
414 if self.set.contains_key(&name) {
415 return Err(format!("agent {} already exists", name).into());
416 }
417
418 self.set.insert(name, agent);
419 Ok(())
420 }
421
422 pub fn iter(&self) -> impl Iterator<Item = (&str, &Arc<dyn DynAgent<C>>)> {
424 self.set.iter().map(|(name, agent)| (name.as_str(), agent))
425 }
426
427 pub fn get(&self, name: &str) -> Option<Arc<dyn DynAgent<C>>> {
429 self.set.get(&name.to_ascii_lowercase()).cloned()
430 }
431
432 pub fn get_lowercase(&self, lowercase_name: &str) -> Option<Arc<dyn DynAgent<C>>> {
434 self.set.get(lowercase_name).cloned()
435 }
436}
437
438impl<C> IntoIterator for AgentSet<C>
439where
440 C: AgentContext + Send + Sync + 'static,
441{
442 type Item = Arc<dyn DynAgent<C>>;
443 type IntoIter = std::collections::btree_map::IntoValues<String, Arc<dyn DynAgent<C>>>;
444
445 fn into_iter(self) -> Self::IntoIter {
447 self.set.into_values()
448 }
449}
450
451#[cfg(test)]
452mod tests {
453 use super::*;
454 use crate::test_support::{MockContext, resource};
455
456 struct ExampleAgent {
457 id: usize,
458 }
459
460 struct OtherAgent;
461
462 struct TaggedAgent;
463
464 struct InvalidAgent;
465
466 impl Agent<MockContext> for ExampleAgent {
467 fn name(&self) -> String {
468 "example_agent".to_string()
469 }
470
471 fn description(&self) -> String {
472 "Example agent used for downcast tests".to_string()
473 }
474
475 fn group(&self) -> Option<ToolGroupInfo> {
476 Some(ToolGroupInfo {
477 id: "example_bundle".to_string(),
478 title: "Example bundle".to_string(),
479 description: "Agents used together in tests".to_string(),
480 instructions: Some("Combine these agents.".to_string()),
481 })
482 }
483
484 async fn run(
485 &self,
486 _ctx: MockContext,
487 _prompt: String,
488 _resources: Vec<Resource>,
489 ) -> Result<AgentOutput, BoxError> {
490 Ok(AgentOutput {
491 content: self.id.to_string(),
492 ..AgentOutput::default()
493 })
494 }
495 }
496
497 impl Agent<MockContext> for OtherAgent {
498 fn name(&self) -> String {
499 "other_agent".to_string()
500 }
501
502 fn description(&self) -> String {
503 "Other agent used for downcast tests".to_string()
504 }
505
506 fn group(&self) -> Option<ToolGroupInfo> {
507 Some(ToolGroupInfo {
508 id: "example_bundle".to_string(),
509 title: "Example bundle".to_string(),
510 description: "Agents used together in tests".to_string(),
511 instructions: Some("Combine these agents.".to_string()),
512 })
513 }
514
515 fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
516 resources
517 .extract_if(.., |resource| resource.name == "selected")
518 .collect()
519 }
520
521 async fn run(
522 &self,
523 _ctx: MockContext,
524 _prompt: String,
525 _resources: Vec<Resource>,
526 ) -> Result<AgentOutput, BoxError> {
527 Ok(AgentOutput {
528 content: "other".to_string(),
529 ..AgentOutput::default()
530 })
531 }
532 }
533
534 impl Agent<MockContext> for TaggedAgent {
535 fn name(&self) -> String {
536 "tagged_agent".to_string()
537 }
538
539 fn description(&self) -> String {
540 "Agent that consumes text and code resources".to_string()
541 }
542
543 fn supported_resource_tags(&self) -> Vec<String> {
544 vec!["text".to_string(), "code".to_string()]
545 }
546
547 fn tool_dependencies(&self) -> Vec<String> {
548 vec!["lookup".to_string(), "summarize".to_string()]
549 }
550
551 async fn run(
552 &self,
553 _ctx: MockContext,
554 prompt: String,
555 resources: Vec<Resource>,
556 ) -> Result<AgentOutput, BoxError> {
557 Ok(AgentOutput {
558 content: format!("{prompt}:{}", resources.len()),
559 ..AgentOutput::default()
560 })
561 }
562 }
563
564 impl Agent<MockContext> for InvalidAgent {
565 fn name(&self) -> String {
566 "bad.agent".to_string()
567 }
568
569 fn description(&self) -> String {
570 "Invalid function name".to_string()
571 }
572
573 async fn run(
574 &self,
575 _ctx: MockContext,
576 _prompt: String,
577 _resources: Vec<Resource>,
578 ) -> Result<AgentOutput, BoxError> {
579 Ok(AgentOutput::default())
580 }
581 }
582
583 #[test]
584 fn dyn_agent_downcast_ref_returns_inner_agent() {
585 let agent = Arc::new(ExampleAgent { id: 7 });
586 let mut agent_set = AgentSet::<MockContext>::new();
587 agent_set
588 .add(agent, Some("test-label".to_string()))
589 .unwrap();
590
591 let dyn_agent = agent_set.get("example_agent").unwrap();
592 let concrete = dyn_agent.downcast_ref::<ExampleAgent>().unwrap();
593
594 assert_eq!(concrete.id, 7);
595 assert!(dyn_agent.downcast_ref::<OtherAgent>().is_none());
596 }
597
598 #[test]
599 fn agent_set_collects_declared_groups() {
600 let mut agent_set = AgentSet::<MockContext>::new();
601 agent_set
602 .add(Arc::new(ExampleAgent { id: 1 }), None)
603 .unwrap();
604 agent_set.add(Arc::new(OtherAgent), None).unwrap();
605 agent_set.add(Arc::new(TaggedAgent), None).unwrap();
607
608 let groups = agent_set.groups();
609 assert_eq!(groups.len(), 1);
610 assert_eq!(groups[0].id, "example_bundle");
611 assert_eq!(
613 groups[0].members,
614 vec!["example_agent".to_string(), "other_agent".to_string()]
615 );
616 assert_eq!(
617 groups[0].instructions.as_deref(),
618 Some("Combine these agents.")
619 );
620 }
621
622 #[test]
623 fn dyn_agent_downcast_returns_original_arc() {
624 let agent = Arc::new(ExampleAgent { id: 9 });
625 let mut agent_set = AgentSet::<MockContext>::new();
626 agent_set
627 .add(agent.clone(), Some("test-label".to_string()))
628 .unwrap();
629
630 let dyn_agent = agent_set.get("example_agent").unwrap();
631 let concrete = dyn_agent
632 .downcast::<ExampleAgent>()
633 .ok()
634 .expect("expected downcast to ExampleAgent to succeed");
635
636 assert_eq!(concrete.id, 9);
637 assert!(Arc::ptr_eq(&concrete, &agent));
638 }
639
640 #[test]
641 fn dyn_agent_downcast_mismatch_returns_original_arc() {
642 let agent = Arc::new(ExampleAgent { id: 11 });
643 let mut agent_set = AgentSet::<MockContext>::new();
644 agent_set
645 .add(agent, Some("test-label".to_string()))
646 .unwrap();
647
648 let dyn_agent = agent_set.get("example_agent").unwrap();
649 let original = dyn_agent.clone();
650 let err = dyn_agent
651 .downcast::<OtherAgent>()
652 .err()
653 .expect("expected downcast to OtherAgent to fail");
654
655 assert!(Arc::ptr_eq(&err, &original));
656 assert_eq!(err.name(), "example_agent");
657 assert_eq!(err.label(), "test-label");
658 }
659
660 #[test]
661 fn agent_default_methods_and_dyn_wrapper_forward_calls() {
662 futures::executor::block_on(async {
663 let agent = Arc::new(ExampleAgent { id: 42 });
664 let mut resources = vec![resource(1, &["text"])];
665
666 let definition = agent.definition();
667 assert_eq!(definition.name, "example_agent");
668 assert_eq!(definition.description, agent.description());
669 assert_eq!(definition.strict, Some(true));
670 assert_eq!(definition.parameters["type"], "object");
671 assert_eq!(
672 definition.parameters["required"].as_array().unwrap()[0],
673 "prompt"
674 );
675 assert!(agent.supported_resource_tags().is_empty());
676 assert!(agent.select_resources(&mut resources).is_empty());
677 assert_eq!(resources.len(), 1);
678 agent.init(MockContext::default()).await.unwrap();
679 assert!(agent.tool_dependencies().is_empty());
680
681 let mut agent_set = AgentSet::<MockContext>::new();
682 agent_set
683 .add(agent, Some("example label".to_string()))
684 .unwrap();
685 let dyn_agent = agent_set.get("EXAMPLE_AGENT").unwrap();
686
687 assert_eq!(dyn_agent.label(), "example label");
688 assert_eq!(dyn_agent.name(), "example_agent");
689 assert_eq!(dyn_agent.definition().name, "example_agent");
690 assert!(dyn_agent.tool_dependencies().is_empty());
691 assert!(dyn_agent.supported_resource_tags().is_empty());
692 dyn_agent.init(MockContext::default()).await.unwrap();
693
694 let output = dyn_agent
695 .run(MockContext::default(), "ignored".to_string(), Vec::new())
696 .await
697 .unwrap();
698 assert_eq!(output.content, "42");
699 });
700 }
701
702 #[test]
703 fn agent_set_registry_filters_resources_and_reports_errors() {
704 futures::executor::block_on(async {
705 let mut agent_set = AgentSet::<MockContext>::new();
706 agent_set
707 .add(Arc::new(ExampleAgent { id: 1 }), None)
708 .unwrap();
709 agent_set
710 .add(Arc::new(TaggedAgent), Some("tagged label".to_string()))
711 .unwrap();
712
713 assert!(agent_set.contains("EXAMPLE_AGENT"));
714 assert!(agent_set.contains_lowercase("tagged_agent"));
715 assert!(!agent_set.contains("missing_agent"));
716 assert_eq!(
717 agent_set.names(),
718 vec!["example_agent".to_string(), "tagged_agent".to_string()]
719 );
720
721 let definition = agent_set.definition("TAGGED_AGENT").unwrap();
722 assert_eq!(definition.name, "tagged_agent");
723 assert!(agent_set.definition("missing_agent").is_none());
724
725 let selected_names = vec!["TAGGED_AGENT".to_string(), "missing_agent".to_string()];
726 let selected_definitions = agent_set.definitions(Some(&selected_names));
727 assert_eq!(selected_definitions.len(), 1);
728 assert_eq!(selected_definitions[0].name, "tagged_agent");
729 assert_eq!(agent_set.definitions(None).len(), 2);
730
731 let selected_functions = agent_set.functions(Some(&selected_names));
732 assert_eq!(selected_functions.len(), 1);
733 assert_eq!(
734 selected_functions[0].supported_resource_tags,
735 vec!["text".to_string(), "code".to_string()]
736 );
737 assert_eq!(agent_set.functions(None).len(), 2);
738
739 let mut empty = Vec::new();
740 assert!(
741 agent_set
742 .select_resources("tagged_agent", &mut empty)
743 .is_empty()
744 );
745 let mut resources = vec![
746 resource(1, &["image"]),
747 resource(2, &["text"]),
748 resource(3, &["code", "text"]),
749 resource(4, &["audio"]),
750 ];
751 let selected = agent_set.select_resources("TAGGED_AGENT", &mut resources);
752 assert_eq!(
753 selected
754 .iter()
755 .map(|resource| resource._id)
756 .collect::<Vec<_>>(),
757 vec![2, 3]
758 );
759 assert_eq!(
760 resources
761 .iter()
762 .map(|resource| resource._id)
763 .collect::<Vec<_>>(),
764 vec![1, 4]
765 );
766 assert!(
767 agent_set
768 .select_resources("missing_agent", &mut resources)
769 .is_empty()
770 );
771
772 let dyn_agent = agent_set.get_lowercase("tagged_agent").unwrap();
773 assert_eq!(dyn_agent.label(), "tagged label");
774 let output = dyn_agent
775 .run(
776 MockContext::default(),
777 "prompt".to_string(),
778 vec![resource(9, &["text"])],
779 )
780 .await
781 .unwrap();
782 assert_eq!(output.content, "prompt:1");
783 assert!(agent_set.get("missing_agent").is_none());
784 assert!(agent_set.get_lowercase("missing_agent").is_none());
785
786 let duplicate = agent_set
787 .add(Arc::new(ExampleAgent { id: 2 }), None)
788 .unwrap_err();
789 assert!(duplicate.to_string().contains("already exists"));
790
791 let invalid = agent_set.add(Arc::new(InvalidAgent), None).unwrap_err();
792 assert!(invalid.to_string().contains("invalid character"));
793 });
794 }
795
796 #[test]
797 fn agent_registry_preserves_custom_resource_selection() {
798 let agent = Arc::new(OtherAgent);
799 let mut direct = vec![
800 resource(1, &["text"]),
801 Resource {
802 name: "selected".into(),
803 ..resource(2, &["image"])
804 },
805 ];
806 let mut registered = direct.clone();
807 let expected = agent.select_resources(&mut direct);
808 assert_eq!(expected.iter().map(|r| r._id).collect::<Vec<_>>(), vec![2]);
809 let mut set = AgentSet::new();
810 set.add(agent, None).unwrap();
811 let actual = set.select_resources("OTHER_AGENT", &mut registered);
812 assert_eq!(
813 serde_json::to_value(actual).unwrap(),
814 serde_json::to_value(expected).unwrap()
815 );
816 assert_eq!(
817 registered.iter().map(|r| r._id).collect::<Vec<_>>(),
818 vec![1]
819 );
820 }
821}