use serde::{Deserialize, Serialize};
use serde_json::json;
use std::{any::Any, collections::BTreeMap, future::Future, marker::PhantomData, sync::Arc};
use crate::{
BoxError, BoxPinFut, Function, ToolGroup, ToolGroupInfo,
context::AgentContext,
model::{AgentOutput, FunctionDefinition, Resource},
registry::{collect_groups, select_by_names},
select_resources, validate_function_name,
};
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct AgentArgs {
pub prompt: String,
}
pub trait Agent<C>: Send + Sync
where
C: AgentContext + Send + Sync,
{
fn name(&self) -> String;
fn description(&self) -> String;
fn definition(&self) -> FunctionDefinition {
FunctionDefinition {
name: self.name().to_ascii_lowercase(),
description: self.description(),
parameters: json!({
"type": "object",
"description": "Run this agent on a focused task. Provide a self-contained prompt with the goal, relevant context, constraints, and expected output.",
"properties": {
"prompt": {
"type": "string",
"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.",
"minLength": 1
},
},
"required": ["prompt"],
"additionalProperties": false
}),
strict: Some(true),
}
}
fn group(&self) -> Option<ToolGroupInfo> {
None
}
fn supported_resource_tags(&self) -> Vec<String> {
Vec::new()
}
fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
let supported_tags = self.supported_resource_tags();
select_resources(resources, &supported_tags)
}
fn init(&self, _ctx: C) -> impl Future<Output = Result<(), BoxError>> + Send {
std::future::ready(Ok(()))
}
fn tool_dependencies(&self) -> Vec<String> {
Vec::new()
}
fn run(
&self,
ctx: C,
prompt: String,
resources: Vec<Resource>,
) -> impl Future<Output = Result<AgentOutput, BoxError>> + Send;
}
pub trait DynAgent<C>: Send + Sync
where
C: AgentContext + Send + Sync,
{
fn as_any(&self) -> &(dyn Any + Send + Sync);
fn into_any(self: Arc<Self>) -> Arc<dyn Any + Send + Sync>;
fn label(&self) -> &str;
fn name(&self) -> String;
fn definition(&self) -> FunctionDefinition;
fn tool_dependencies(&self) -> Vec<String>;
fn group(&self) -> Option<ToolGroupInfo>;
fn supported_resource_tags(&self) -> Vec<String>;
fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
select_resources(resources, &self.supported_resource_tags())
}
fn init(&self, ctx: C) -> BoxPinFut<Result<(), BoxError>>;
fn run(
&self,
ctx: C,
prompt: String,
resources: Vec<Resource>,
) -> BoxPinFut<Result<AgentOutput, BoxError>>;
}
impl<C> dyn DynAgent<C>
where
C: AgentContext + Send + Sync + 'static,
{
pub fn downcast_ref<T>(&self) -> Option<&T>
where
T: Agent<C> + 'static,
{
self.as_any().downcast_ref::<T>()
}
pub fn downcast<T>(self: Arc<Self>) -> Result<Arc<T>, Arc<Self>>
where
T: Agent<C> + 'static,
{
match self.clone().into_any().downcast::<T>() {
Ok(agent) => Ok(agent),
Err(_) => Err(self),
}
}
}
struct AgentWrapper<T, C>
where
T: Agent<C> + 'static,
C: AgentContext + Send + Sync + 'static,
{
inner: Arc<T>,
label: String,
_phantom: PhantomData<C>,
}
impl<T, C> DynAgent<C> for AgentWrapper<T, C>
where
T: Agent<C> + 'static,
C: AgentContext + Send + Sync + 'static,
{
fn as_any(&self) -> &(dyn Any + Send + Sync) {
self.inner.as_ref()
}
fn into_any(self: Arc<Self>) -> Arc<dyn Any + Send + Sync> {
self.inner.clone()
}
fn label(&self) -> &str {
&self.label
}
fn name(&self) -> String {
self.inner.name()
}
fn definition(&self) -> FunctionDefinition {
self.inner.definition()
}
fn tool_dependencies(&self) -> Vec<String> {
self.inner.tool_dependencies()
}
fn group(&self) -> Option<ToolGroupInfo> {
self.inner.group()
}
fn supported_resource_tags(&self) -> Vec<String> {
self.inner.supported_resource_tags()
}
fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
self.inner.select_resources(resources)
}
fn init(&self, ctx: C) -> BoxPinFut<Result<(), BoxError>> {
let agent = self.inner.clone();
Box::pin(async move { agent.init(ctx).await })
}
fn run(
&self,
ctx: C,
prompt: String,
resources: Vec<Resource>,
) -> BoxPinFut<Result<AgentOutput, BoxError>> {
let agent = self.inner.clone();
Box::pin(async move { agent.run(ctx, prompt, resources).await })
}
}
pub struct AgentSet<C: AgentContext> {
set: BTreeMap<String, Arc<dyn DynAgent<C>>>,
}
impl<C: AgentContext> Default for AgentSet<C> {
fn default() -> Self {
Self {
set: BTreeMap::new(),
}
}
}
impl<C> AgentSet<C>
where
C: AgentContext + Send + Sync + 'static,
{
pub fn new() -> Self {
Self::default()
}
pub fn contains(&self, name: &str) -> bool {
self.set.contains_key(&name.to_ascii_lowercase())
}
pub fn contains_lowercase(&self, lowercase_name: &str) -> bool {
self.set.contains_key(lowercase_name)
}
pub fn names(&self) -> Vec<String> {
self.set.keys().cloned().collect()
}
pub fn groups(&self) -> Vec<ToolGroup> {
collect_groups(self.set.iter().map(|(name, agent)| (name, agent.group())))
}
pub fn definition(&self, name: &str) -> Option<FunctionDefinition> {
self.set
.get(&name.to_ascii_lowercase())
.map(|agent| agent.definition())
}
pub fn definitions(&self, names: Option<&[String]>) -> Vec<FunctionDefinition> {
select_by_names(&self.set, names, |agent| agent.definition())
}
pub fn functions(&self, names: Option<&[String]>) -> Vec<Function> {
select_by_names(&self.set, names, |agent| Function {
definition: agent.definition(),
supported_resource_tags: agent.supported_resource_tags(),
})
}
pub fn select_resources(&self, name: &str, resources: &mut Vec<Resource>) -> Vec<Resource> {
if resources.is_empty() {
return Vec::new();
}
self.set
.get(&name.to_ascii_lowercase())
.map(|agent| agent.select_resources(resources))
.unwrap_or_default()
}
pub fn add<T>(&mut self, agent: Arc<T>, label: Option<String>) -> Result<(), BoxError>
where
T: Agent<C> + Send + Sync + 'static,
{
let label = label.unwrap_or_else(|| agent.name().to_ascii_lowercase());
self.add_dyn(Arc::new(AgentWrapper {
inner: agent,
label,
_phantom: PhantomData,
}))
}
pub fn add_dyn(&mut self, agent: Arc<dyn DynAgent<C>>) -> Result<(), BoxError> {
let name = agent.name().to_ascii_lowercase();
validate_function_name(&name)?;
if self.set.contains_key(&name) {
return Err(format!("agent {} already exists", name).into());
}
self.set.insert(name, agent);
Ok(())
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &Arc<dyn DynAgent<C>>)> {
self.set.iter().map(|(name, agent)| (name.as_str(), agent))
}
pub fn get(&self, name: &str) -> Option<Arc<dyn DynAgent<C>>> {
self.set.get(&name.to_ascii_lowercase()).cloned()
}
pub fn get_lowercase(&self, lowercase_name: &str) -> Option<Arc<dyn DynAgent<C>>> {
self.set.get(lowercase_name).cloned()
}
}
impl<C> IntoIterator for AgentSet<C>
where
C: AgentContext + Send + Sync + 'static,
{
type Item = Arc<dyn DynAgent<C>>;
type IntoIter = std::collections::btree_map::IntoValues<String, Arc<dyn DynAgent<C>>>;
fn into_iter(self) -> Self::IntoIter {
self.set.into_values()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::{MockContext, resource};
struct ExampleAgent {
id: usize,
}
struct OtherAgent;
struct TaggedAgent;
struct InvalidAgent;
impl Agent<MockContext> for ExampleAgent {
fn name(&self) -> String {
"example_agent".to_string()
}
fn description(&self) -> String {
"Example agent used for downcast tests".to_string()
}
fn group(&self) -> Option<ToolGroupInfo> {
Some(ToolGroupInfo {
id: "example_bundle".to_string(),
title: "Example bundle".to_string(),
description: "Agents used together in tests".to_string(),
instructions: Some("Combine these agents.".to_string()),
})
}
async fn run(
&self,
_ctx: MockContext,
_prompt: String,
_resources: Vec<Resource>,
) -> Result<AgentOutput, BoxError> {
Ok(AgentOutput {
content: self.id.to_string(),
..AgentOutput::default()
})
}
}
impl Agent<MockContext> for OtherAgent {
fn name(&self) -> String {
"other_agent".to_string()
}
fn description(&self) -> String {
"Other agent used for downcast tests".to_string()
}
fn group(&self) -> Option<ToolGroupInfo> {
Some(ToolGroupInfo {
id: "example_bundle".to_string(),
title: "Example bundle".to_string(),
description: "Agents used together in tests".to_string(),
instructions: Some("Combine these agents.".to_string()),
})
}
fn select_resources(&self, resources: &mut Vec<Resource>) -> Vec<Resource> {
resources
.extract_if(.., |resource| resource.name == "selected")
.collect()
}
async fn run(
&self,
_ctx: MockContext,
_prompt: String,
_resources: Vec<Resource>,
) -> Result<AgentOutput, BoxError> {
Ok(AgentOutput {
content: "other".to_string(),
..AgentOutput::default()
})
}
}
impl Agent<MockContext> for TaggedAgent {
fn name(&self) -> String {
"tagged_agent".to_string()
}
fn description(&self) -> String {
"Agent that consumes text and code resources".to_string()
}
fn supported_resource_tags(&self) -> Vec<String> {
vec!["text".to_string(), "code".to_string()]
}
fn tool_dependencies(&self) -> Vec<String> {
vec!["lookup".to_string(), "summarize".to_string()]
}
async fn run(
&self,
_ctx: MockContext,
prompt: String,
resources: Vec<Resource>,
) -> Result<AgentOutput, BoxError> {
Ok(AgentOutput {
content: format!("{prompt}:{}", resources.len()),
..AgentOutput::default()
})
}
}
impl Agent<MockContext> for InvalidAgent {
fn name(&self) -> String {
"bad.agent".to_string()
}
fn description(&self) -> String {
"Invalid function name".to_string()
}
async fn run(
&self,
_ctx: MockContext,
_prompt: String,
_resources: Vec<Resource>,
) -> Result<AgentOutput, BoxError> {
Ok(AgentOutput::default())
}
}
#[test]
fn dyn_agent_downcast_ref_returns_inner_agent() {
let agent = Arc::new(ExampleAgent { id: 7 });
let mut agent_set = AgentSet::<MockContext>::new();
agent_set
.add(agent, Some("test-label".to_string()))
.unwrap();
let dyn_agent = agent_set.get("example_agent").unwrap();
let concrete = dyn_agent.downcast_ref::<ExampleAgent>().unwrap();
assert_eq!(concrete.id, 7);
assert!(dyn_agent.downcast_ref::<OtherAgent>().is_none());
}
#[test]
fn agent_set_collects_declared_groups() {
let mut agent_set = AgentSet::<MockContext>::new();
agent_set
.add(Arc::new(ExampleAgent { id: 1 }), None)
.unwrap();
agent_set.add(Arc::new(OtherAgent), None).unwrap();
agent_set.add(Arc::new(TaggedAgent), None).unwrap();
let groups = agent_set.groups();
assert_eq!(groups.len(), 1);
assert_eq!(groups[0].id, "example_bundle");
assert_eq!(
groups[0].members,
vec!["example_agent".to_string(), "other_agent".to_string()]
);
assert_eq!(
groups[0].instructions.as_deref(),
Some("Combine these agents.")
);
}
#[test]
fn dyn_agent_downcast_returns_original_arc() {
let agent = Arc::new(ExampleAgent { id: 9 });
let mut agent_set = AgentSet::<MockContext>::new();
agent_set
.add(agent.clone(), Some("test-label".to_string()))
.unwrap();
let dyn_agent = agent_set.get("example_agent").unwrap();
let concrete = dyn_agent
.downcast::<ExampleAgent>()
.ok()
.expect("expected downcast to ExampleAgent to succeed");
assert_eq!(concrete.id, 9);
assert!(Arc::ptr_eq(&concrete, &agent));
}
#[test]
fn dyn_agent_downcast_mismatch_returns_original_arc() {
let agent = Arc::new(ExampleAgent { id: 11 });
let mut agent_set = AgentSet::<MockContext>::new();
agent_set
.add(agent, Some("test-label".to_string()))
.unwrap();
let dyn_agent = agent_set.get("example_agent").unwrap();
let original = dyn_agent.clone();
let err = dyn_agent
.downcast::<OtherAgent>()
.err()
.expect("expected downcast to OtherAgent to fail");
assert!(Arc::ptr_eq(&err, &original));
assert_eq!(err.name(), "example_agent");
assert_eq!(err.label(), "test-label");
}
#[test]
fn agent_default_methods_and_dyn_wrapper_forward_calls() {
futures::executor::block_on(async {
let agent = Arc::new(ExampleAgent { id: 42 });
let mut resources = vec![resource(1, &["text"])];
let definition = agent.definition();
assert_eq!(definition.name, "example_agent");
assert_eq!(definition.description, agent.description());
assert_eq!(definition.strict, Some(true));
assert_eq!(definition.parameters["type"], "object");
assert_eq!(
definition.parameters["required"].as_array().unwrap()[0],
"prompt"
);
assert!(agent.supported_resource_tags().is_empty());
assert!(agent.select_resources(&mut resources).is_empty());
assert_eq!(resources.len(), 1);
agent.init(MockContext::default()).await.unwrap();
assert!(agent.tool_dependencies().is_empty());
let mut agent_set = AgentSet::<MockContext>::new();
agent_set
.add(agent, Some("example label".to_string()))
.unwrap();
let dyn_agent = agent_set.get("EXAMPLE_AGENT").unwrap();
assert_eq!(dyn_agent.label(), "example label");
assert_eq!(dyn_agent.name(), "example_agent");
assert_eq!(dyn_agent.definition().name, "example_agent");
assert!(dyn_agent.tool_dependencies().is_empty());
assert!(dyn_agent.supported_resource_tags().is_empty());
dyn_agent.init(MockContext::default()).await.unwrap();
let output = dyn_agent
.run(MockContext::default(), "ignored".to_string(), Vec::new())
.await
.unwrap();
assert_eq!(output.content, "42");
});
}
#[test]
fn agent_set_registry_filters_resources_and_reports_errors() {
futures::executor::block_on(async {
let mut agent_set = AgentSet::<MockContext>::new();
agent_set
.add(Arc::new(ExampleAgent { id: 1 }), None)
.unwrap();
agent_set
.add(Arc::new(TaggedAgent), Some("tagged label".to_string()))
.unwrap();
assert!(agent_set.contains("EXAMPLE_AGENT"));
assert!(agent_set.contains_lowercase("tagged_agent"));
assert!(!agent_set.contains("missing_agent"));
assert_eq!(
agent_set.names(),
vec!["example_agent".to_string(), "tagged_agent".to_string()]
);
let definition = agent_set.definition("TAGGED_AGENT").unwrap();
assert_eq!(definition.name, "tagged_agent");
assert!(agent_set.definition("missing_agent").is_none());
let selected_names = vec!["TAGGED_AGENT".to_string(), "missing_agent".to_string()];
let selected_definitions = agent_set.definitions(Some(&selected_names));
assert_eq!(selected_definitions.len(), 1);
assert_eq!(selected_definitions[0].name, "tagged_agent");
assert_eq!(agent_set.definitions(None).len(), 2);
let selected_functions = agent_set.functions(Some(&selected_names));
assert_eq!(selected_functions.len(), 1);
assert_eq!(
selected_functions[0].supported_resource_tags,
vec!["text".to_string(), "code".to_string()]
);
assert_eq!(agent_set.functions(None).len(), 2);
let mut empty = Vec::new();
assert!(
agent_set
.select_resources("tagged_agent", &mut empty)
.is_empty()
);
let mut resources = vec![
resource(1, &["image"]),
resource(2, &["text"]),
resource(3, &["code", "text"]),
resource(4, &["audio"]),
];
let selected = agent_set.select_resources("TAGGED_AGENT", &mut resources);
assert_eq!(
selected
.iter()
.map(|resource| resource._id)
.collect::<Vec<_>>(),
vec![2, 3]
);
assert_eq!(
resources
.iter()
.map(|resource| resource._id)
.collect::<Vec<_>>(),
vec![1, 4]
);
assert!(
agent_set
.select_resources("missing_agent", &mut resources)
.is_empty()
);
let dyn_agent = agent_set.get_lowercase("tagged_agent").unwrap();
assert_eq!(dyn_agent.label(), "tagged label");
let output = dyn_agent
.run(
MockContext::default(),
"prompt".to_string(),
vec![resource(9, &["text"])],
)
.await
.unwrap();
assert_eq!(output.content, "prompt:1");
assert!(agent_set.get("missing_agent").is_none());
assert!(agent_set.get_lowercase("missing_agent").is_none());
let duplicate = agent_set
.add(Arc::new(ExampleAgent { id: 2 }), None)
.unwrap_err();
assert!(duplicate.to_string().contains("already exists"));
let invalid = agent_set.add(Arc::new(InvalidAgent), None).unwrap_err();
assert!(invalid.to_string().contains("invalid character"));
});
}
#[test]
fn agent_registry_preserves_custom_resource_selection() {
let agent = Arc::new(OtherAgent);
let mut direct = vec![
resource(1, &["text"]),
Resource {
name: "selected".into(),
..resource(2, &["image"])
},
];
let mut registered = direct.clone();
let expected = agent.select_resources(&mut direct);
assert_eq!(expected.iter().map(|r| r._id).collect::<Vec<_>>(), vec![2]);
let mut set = AgentSet::new();
set.add(agent, None).unwrap();
let actual = set.select_resources("OTHER_AGENT", &mut registered);
assert_eq!(
serde_json::to_value(actual).unwrap(),
serde_json::to_value(expected).unwrap()
);
assert_eq!(
registered.iter().map(|r| r._id).collect::<Vec<_>>(),
vec![1]
);
}
}