use std::collections::HashMap;
use std::sync::Arc;
use async_stream::try_stream;
use futures_util::stream::BoxStream;
use futures_util::StreamExt;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use tokio::sync::RwLock;
use crate::a2a::RemoteAgent;
use crate::agent::{AgentBuilder, AgentManifest, Capability};
use crate::error::{Error, Result};
use crate::identity::AgentIdentity;
use crate::llm::{
ChatMessage, ChatRequest, ChatResponse, EmbeddingProvider, LlmProvider, StreamChunk, ToolSpec,
UsageEvent, UsageObserver, UsageSnapshot, UsageTotals,
};
use crate::mcp::{McpClient, McpTool};
use crate::session::SessionStore;
use crate::skills::{SkillPolicy, SkillRegistry};
use crate::vector::{chunk_text, Document, MetadataFilter, SearchResult, VectorStore};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
pub struct AgentInfo {
pub name: String,
pub version: String,
pub description: String,
pub capabilities: Vec<Capability>,
pub skills: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub public_key: Option<String>,
pub signed: bool,
}
pub struct Agent {
pub(crate) manifest: AgentManifest,
pub(crate) identity: Option<AgentIdentity>,
pub(crate) llm: Option<Arc<dyn LlmProvider>>,
pub(crate) embeddings: Option<Arc<dyn EmbeddingProvider>>,
pub(crate) vector_store: Option<Arc<dyn VectorStore>>,
pub(crate) skills: SkillRegistry,
pub(crate) skill_policy: SkillPolicy,
pub(crate) mcp_clients: RwLock<HashMap<String, Arc<McpClient>>>,
pub(crate) sessions: Arc<dyn SessionStore>,
pub(crate) peers: RwLock<HashMap<String, Arc<RemoteAgent>>>,
pub(crate) usage_totals: Arc<UsageTotals>,
pub(crate) usage_observer: Option<Arc<dyn UsageObserver>>,
pub(crate) ready: std::sync::atomic::AtomicBool,
}
impl std::fmt::Debug for Agent {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Agent")
.field("name", &self.manifest.name)
.field("version", &self.manifest.version)
.field("skills", &self.skills)
.finish_non_exhaustive()
}
}
impl Agent {
pub fn builder() -> AgentBuilder {
AgentBuilder::new()
}
pub fn name(&self) -> &str {
&self.manifest.name
}
pub fn version(&self) -> &str {
&self.manifest.version
}
pub fn manifest(&self) -> &AgentManifest {
&self.manifest
}
pub fn identity(&self) -> Option<&AgentIdentity> {
self.identity.as_ref()
}
pub fn public_key(&self) -> Option<String> {
self.identity.as_ref().map(AgentIdentity::public_key_base64)
}
pub fn verify(&self) -> Result<()> {
self.manifest.verify()
}
pub fn info(&self) -> AgentInfo {
let mut skills: Vec<String> = self
.skills
.list()
.iter()
.map(|s| s.name().to_string())
.collect();
skills.sort();
AgentInfo {
name: self.manifest.name.clone(),
version: self.manifest.version.clone(),
description: self.manifest.description.clone(),
capabilities: self.manifest.capabilities.clone(),
skills,
model: self.manifest.model.clone(),
public_key: self.manifest.public_key.clone(),
signed: self.manifest.signature.is_some(),
}
}
pub fn active_capabilities(&self) -> Vec<&Capability> {
self.manifest
.capabilities
.iter()
.filter(|c| c.enabled)
.collect()
}
pub fn skills(&self) -> &SkillRegistry {
&self.skills
}
pub async fn execute_skill(&self, name: &str, input: Value) -> Result<Value> {
let skill = self
.skills
.get(name)
.ok_or_else(|| Error::SkillNotFound(name.to_string()))?;
self.skill_policy.check(skill.as_ref())?;
let handle = tokio::spawn(async move { skill.execute(input).await });
let abort = handle.abort_handle();
let outcome = match self.skill_policy.timeout() {
Some(limit) => tokio::time::timeout(limit, handle).await.map_err(|_| {
abort.abort(); Error::Skill(format!("skill '{name}' timed out after {limit:?}"))
})?,
None => handle.await,
};
outcome.map_err(|e| Error::Skill(format!("skill '{name}' panicked: {e}")))?
}
fn record_usage(&self, session_id: &str, response: &ChatResponse) {
if let Some(usage) = &response.usage {
let event = UsageEvent {
session_id: session_id.to_string(),
model: response.model.clone(),
usage: usage.clone(),
};
self.usage_totals.record(&event);
if let Some(observer) = &self.usage_observer {
observer.on_usage(&event);
}
}
}
pub fn usage(&self) -> UsageSnapshot {
self.usage_totals.snapshot()
}
fn require_llm(&self) -> Result<Arc<dyn LlmProvider>> {
self.llm.clone().ok_or(Error::NotConfigured("LLM provider"))
}
async fn conversation(&self, session_id: &str, message: &str) -> Result<Vec<ChatMessage>> {
let mut messages = Vec::new();
if let Some(system_prompt) = &self.manifest.system_prompt {
messages.push(ChatMessage::system(system_prompt));
}
messages.extend(self.sessions.load(session_id).await?);
messages.push(ChatMessage::user(message));
Ok(messages)
}
fn request_for(&self, messages: Vec<ChatMessage>) -> ChatRequest {
let mut request = ChatRequest::new(messages);
request.model = self.manifest.model.clone();
request
}
pub async fn chat(&self, session_id: &str, message: impl AsRef<str>) -> Result<String> {
let message = message.as_ref();
let llm = self.require_llm()?;
let request = self.request_for(self.conversation(session_id, message).await?);
let response = llm.chat(request).await?;
self.record_usage(session_id, &response);
self.sessions
.append(
session_id,
&[
ChatMessage::user(message),
ChatMessage::assistant(&response.content),
],
)
.await?;
Ok(response.content)
}
pub async fn chat_with_tools(
&self,
session_id: &str,
message: impl AsRef<str>,
max_rounds: usize,
) -> Result<String> {
let message = message.as_ref();
let llm = self.require_llm()?;
let tools: Vec<ToolSpec> = self
.skills
.list()
.iter()
.map(|s| ToolSpec::from_skill(s.as_ref()))
.collect();
let mut messages = self.conversation(session_id, message).await?;
let mut transcript: Vec<ChatMessage> = vec![ChatMessage::user(message)];
for _ in 0..max_rounds.max(1) {
let mut request = self.request_for(messages.clone());
if !tools.is_empty() {
request.tools = Some(tools.clone());
}
let response = llm.chat(request).await?;
self.record_usage(session_id, &response);
if response.tool_calls.is_empty() {
transcript.push(ChatMessage::assistant(&response.content));
self.sessions.append(session_id, &transcript).await?;
return Ok(response.content);
}
let assistant = ChatMessage::assistant_tool_calls(
response.content.clone(),
response.tool_calls.clone(),
);
messages.push(assistant.clone());
transcript.push(assistant);
for call in response.tool_calls {
let output = match self.execute_skill(&call.name, call.arguments.clone()).await {
Ok(value) => value.to_string(),
Err(e) => serde_json::json!({ "error": e.to_string() }).to_string(),
};
let result = ChatMessage::tool_result(&call.id, output);
messages.push(result.clone());
transcript.push(result);
}
}
Err(Error::Llm(format!(
"tool-calling loop did not converge within {max_rounds} rounds"
)))
}
pub async fn chat_stream(
&self,
session_id: &str,
message: impl AsRef<str>,
) -> Result<BoxStream<'static, Result<StreamChunk>>> {
let message = message.as_ref().to_string();
let llm = self.require_llm()?;
let request = self.request_for(self.conversation(session_id, &message).await?);
let mut inner = llm.chat_stream(request).await?;
let sessions = Arc::clone(&self.sessions);
let session_id = session_id.to_string();
let stream = try_stream! {
let mut reply = String::new();
while let Some(chunk) = inner.next().await {
let chunk = chunk?;
if chunk.done {
sessions
.append(
&session_id,
&[
ChatMessage::user(message.clone()),
ChatMessage::assistant(reply.clone()),
],
)
.await?;
yield chunk;
break;
}
reply.push_str(&chunk.delta);
yield chunk;
}
};
Ok(Box::pin(stream))
}
pub async fn session_history(&self, session_id: &str) -> Result<Vec<ChatMessage>> {
self.sessions.load(session_id).await
}
pub async fn list_sessions(&self) -> Result<Vec<String>> {
self.sessions.list_sessions().await
}
pub async fn clear_session(&self, session_id: &str) -> Result<()> {
self.sessions.clear(session_id).await
}
pub fn session_store(&self) -> Arc<dyn SessionStore> {
Arc::clone(&self.sessions)
}
pub async fn add_peer(&self, name: impl Into<String>, peer: RemoteAgent) {
self.peers.write().await.insert(name.into(), Arc::new(peer));
}
pub async fn peer(&self, name: &str) -> Option<Arc<RemoteAgent>> {
self.peers.read().await.get(name).cloned()
}
pub async fn list_peers(&self) -> Vec<String> {
self.peers.read().await.keys().cloned().collect()
}
async fn require_peer(&self, name: &str) -> Result<Arc<RemoteAgent>> {
self.peer(name)
.await
.ok_or_else(|| Error::A2a(format!("peer '{name}' is not registered")))
}
pub async fn delegate_chat(
&self,
peer: &str,
session_id: &str,
message: &str,
) -> Result<String> {
self.require_peer(peer)
.await?
.chat(session_id, message)
.await
}
pub async fn delegate_skill(&self, peer: &str, skill: &str, input: Value) -> Result<Value> {
self.require_peer(peer)
.await?
.execute_skill(skill, input)
.await
}
pub async fn connect_mcp_servers(&self) -> Result<Vec<String>> {
let mut connected = Vec::new();
for config in &self.manifest.mcp_servers {
let client = McpClient::connect(config).await?;
self.mcp_clients
.write()
.await
.insert(config.name.clone(), Arc::new(client));
connected.push(config.name.clone());
}
Ok(connected)
}
pub async fn mcp_client(&self, name: &str) -> Option<Arc<McpClient>> {
self.mcp_clients.read().await.get(name).cloned()
}
pub async fn mcp_tools(&self, server: &str) -> Result<Vec<McpTool>> {
let client = self
.mcp_client(server)
.await
.ok_or_else(|| Error::Mcp(format!("MCP server '{server}' is not connected")))?;
client.list_tools().await
}
pub async fn call_mcp_tool(&self, server: &str, tool: &str, arguments: Value) -> Result<Value> {
let client = self
.mcp_client(server)
.await
.ok_or_else(|| Error::Mcp(format!("MCP server '{server}' is not connected")))?;
client.call_tool(tool, arguments).await
}
pub async fn remember(&self, text: &str, metadata: Value) -> Result<String> {
let embeddings = self
.embeddings
.clone()
.ok_or(Error::NotConfigured("embedding provider"))?;
let store = self
.vector_store
.clone()
.ok_or(Error::NotConfigured("vector store"))?;
let mut vectors = embeddings.embed_documents(&[text.to_string()]).await?;
let vector = vectors
.pop()
.ok_or_else(|| Error::Llm("embedding provider returned no vectors".into()))?;
let id = uuid::Uuid::new_v4().to_string();
store
.upsert(vec![Document::new(&id, vector)
.with_text(text)
.with_metadata(metadata)])
.await?;
Ok(id)
}
pub async fn remember_batch(&self, texts: &[String], metadata: Value) -> Result<Vec<String>> {
let embeddings = self
.embeddings
.clone()
.ok_or(Error::NotConfigured("embedding provider"))?;
let store = self
.vector_store
.clone()
.ok_or(Error::NotConfigured("vector store"))?;
if texts.is_empty() {
return Ok(Vec::new());
}
let vectors = embeddings.embed_documents(texts).await?;
if vectors.len() != texts.len() {
return Err(Error::Llm(format!(
"embedding provider returned {} vectors for {} texts",
vectors.len(),
texts.len()
)));
}
let mut ids = Vec::with_capacity(texts.len());
let documents: Vec<Document> = texts
.iter()
.zip(vectors)
.map(|(text, vector)| {
let id = uuid::Uuid::new_v4().to_string();
ids.push(id.clone());
Document::new(id, vector)
.with_text(text)
.with_metadata(metadata.clone())
})
.collect();
store.upsert_batched(documents, 64).await?;
Ok(ids)
}
pub async fn remember_document(
&self,
text: &str,
metadata: Value,
max_chars: usize,
overlap: usize,
) -> Result<Vec<String>> {
let chunks = chunk_text(text, max_chars, overlap);
let mut ids = Vec::with_capacity(chunks.len());
for (index, chunk) in chunks.iter().enumerate() {
let mut chunk_metadata = metadata.clone();
if let Value::Object(map) = &mut chunk_metadata {
map.insert("_chunk".into(), Value::from(index));
}
ids.extend(
self.remember_batch(std::slice::from_ref(chunk), chunk_metadata)
.await?,
);
}
Ok(ids)
}
pub async fn recall(&self, query: &str, top_k: usize) -> Result<Vec<SearchResult>> {
self.recall_filtered(query, top_k, &MetadataFilter::new())
.await
}
pub async fn recall_filtered(
&self,
query: &str,
top_k: usize,
filter: &MetadataFilter,
) -> Result<Vec<SearchResult>> {
let embeddings = self
.embeddings
.clone()
.ok_or(Error::NotConfigured("embedding provider"))?;
let store = self
.vector_store
.clone()
.ok_or(Error::NotConfigured("vector store"))?;
let vector = embeddings.embed_query(query).await?;
store.search_filtered(vector, top_k, filter).await
}
pub fn vector_store(&self) -> Option<Arc<dyn VectorStore>> {
self.vector_store.clone()
}
pub fn set_ready(&self, ready: bool) {
self.ready
.store(ready, std::sync::atomic::Ordering::Relaxed);
}
pub fn is_ready(&self) -> bool {
self.ready.load(std::sync::atomic::Ordering::Relaxed)
}
}