use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
#[cfg(feature = "tools")]
use async_trait::async_trait;
use log;
use tokio::sync::mpsc;
use crate::backend::{GenerationResult, InferenceParams, LlmBackend};
use crate::context::{plan_prune, prepare_context, ContextConfig, PruneStrategy};
use crate::error::{CoreError, CoreResult};
use crate::events::AgentEvent;
use crate::messages::{Message, Role, ToolCall};
use crate::template::{ChatMLTemplate, ChatTemplate};
use crate::tools::{parse_tool_calls, ToolSchema};
#[cfg(feature = "tools")]
use crate::{
messages::ToolResult,
tools::{Tool, ToolOutput, ToolUpdateCallback},
};
const SUMMARY_MARKER: &str = "[Summary of earlier conversation]";
const SUMMARY_MAX_TOKENS: u32 = 320;
#[derive(Debug, Clone)]
pub struct AgentConfig {
pub system_prompt: String,
pub inference_params: InferenceParams,
pub context_config: ContextConfig,
pub max_tool_iterations: usize,
}
impl Default for AgentConfig {
fn default() -> Self {
Self {
system_prompt: "You are a helpful assistant.".to_string(),
inference_params: InferenceParams::default(),
context_config: ContextConfig::default(),
max_tool_iterations: 8,
}
}
}
#[cfg(feature = "tools")]
#[derive(Debug, Clone)]
pub enum ApprovalDecision {
Approve,
Deny {
reason: String,
},
}
#[cfg(feature = "tools")]
#[async_trait]
pub trait ApprovalHook: Send + Sync {
async fn review(&self, call: &ToolCall) -> ApprovalDecision;
}
pub struct Agent {
config: AgentConfig,
messages: Vec<Message>,
#[cfg(feature = "tools")]
tools: Vec<Box<dyn Tool>>,
#[cfg(feature = "tools")]
approval_hook: Option<Arc<dyn ApprovalHook>>,
template: Arc<dyn ChatTemplate>,
abort: Arc<AtomicBool>,
msg_counter: u64,
}
impl Agent {
pub fn new(config: AgentConfig) -> Self {
Self::with_template(config, Arc::new(ChatMLTemplate))
}
pub fn with_template(config: AgentConfig, template: Arc<dyn ChatTemplate>) -> Self {
log::debug!(
"Agent created: system_prompt_len={}, max_ctx={}, max_resp={}, template={}",
config.system_prompt.len(),
config.context_config.max_context_tokens,
config.context_config.max_response_tokens,
template.name(),
);
Self {
config,
messages: Vec::new(),
#[cfg(feature = "tools")]
tools: Vec::new(),
#[cfg(feature = "tools")]
approval_hook: None,
template,
abort: Arc::new(AtomicBool::new(false)),
msg_counter: 0,
}
}
pub fn messages(&self) -> &[Message] {
&self.messages
}
pub fn config(&self) -> &AgentConfig {
&self.config
}
pub fn template(&self) -> &dyn ChatTemplate {
self.template.as_ref()
}
pub fn set_system_prompt(&mut self, prompt: impl Into<String>) {
let prompt = prompt.into();
log::debug!("Agent system prompt updated: len={}", prompt.len());
self.config.system_prompt = prompt;
}
pub fn set_inference_params(&mut self, params: InferenceParams) {
log::debug!(
"Agent inference params: max_tokens={}, temp={}, ctx={}, threads={}",
params.max_tokens,
params.temperature,
params.context_size,
params.n_threads,
);
self.config.inference_params = params;
}
pub fn set_context_config(&mut self, config: ContextConfig) {
log::debug!(
"Agent context config: max_ctx={}, max_resp={}",
config.max_context_tokens,
config.max_response_tokens,
);
self.config.context_config = config;
}
pub fn set_prune_strategy(&mut self, strategy: PruneStrategy) {
log::debug!("Agent prune strategy: {strategy:?}");
self.config.context_config.prune_strategy = strategy;
}
pub fn set_pinned(&mut self, message_id: &str, pinned: bool) -> bool {
match self.messages.iter_mut().find(|m| m.id == message_id) {
Some(msg) => {
msg.pinned = pinned;
log::debug!("Agent message {message_id} pinned={pinned}");
true
}
None => {
log::warn!("set_pinned: message {message_id} not found");
false
}
}
}
pub fn set_template(&mut self, template: Arc<dyn ChatTemplate>) {
log::debug!("Agent template updated: {}", template.name());
self.template = template;
}
#[cfg(feature = "tools")]
pub fn set_tools(&mut self, tools: Vec<Box<dyn Tool>>) {
log::debug!("Agent tools set: count={}", tools.len());
self.tools = tools;
}
#[cfg(feature = "tools")]
pub fn set_approval_hook(&mut self, hook: Arc<dyn ApprovalHook>) {
log::debug!("Agent approval hook installed");
self.approval_hook = Some(hook);
}
pub fn clear(&mut self) {
let count = self.messages.len();
self.messages.clear();
log::debug!("Agent conversation cleared: {count} messages removed");
}
pub fn replace_messages(&mut self, messages: Vec<Message>) {
log::debug!("Agent messages replaced: count={}", messages.len());
let max_loaded = messages
.iter()
.filter_map(|m| m.id.strip_prefix("msg-"))
.filter_map(|n| n.parse::<u64>().ok())
.max()
.unwrap_or(0);
self.msg_counter = self.msg_counter.max(max_loaded);
self.messages = messages;
}
pub fn abort(&self) {
log::debug!("Agent abort requested");
self.abort.store(true, Ordering::Relaxed);
}
pub fn abort_flag(&self) -> Arc<AtomicBool> {
self.abort.clone()
}
fn next_id(&mut self) -> String {
self.msg_counter += 1;
format!("msg-{}", self.msg_counter)
}
#[cfg(feature = "tools")]
fn tool_schemas(&self) -> Vec<ToolSchema> {
self.tools.iter().map(|t| t.schema()).collect()
}
#[cfg(not(feature = "tools"))]
fn tool_schemas(&self) -> Vec<ToolSchema> {
Vec::new()
}
pub async fn prompt(
&mut self,
text: impl Into<String>,
backend: Arc<dyn LlmBackend>,
tx: mpsc::UnboundedSender<AgentEvent>,
) -> CoreResult<()> {
let text = text.into().trim().to_string();
if text.is_empty() {
return Err(CoreError::Agent("Empty message".into()));
}
self.abort.store(false, Ordering::Relaxed);
let user_msg = Message::user(self.next_id(), &text);
self.messages.push(user_msg.clone());
tx.send(AgentEvent::AgentStart).ok();
tx.send(AgentEvent::MessageStart {
message: user_msg.clone(),
})
.ok();
tx.send(AgentEvent::MessageEnd { message: user_msg }).ok();
self.compress_if_needed(&backend, &tx).await;
let mut new_messages: Vec<Message> = Vec::new();
#[cfg(feature = "tools")]
let has_tools = !self.tools.is_empty();
#[cfg(not(feature = "tools"))]
let has_tools = false;
for iteration in 0..self.config.max_tool_iterations {
tx.send(AgentEvent::TurnStart).ok();
let gen = match self.generate_once(backend.clone(), &tx).await {
Ok(gen) => gen,
Err(CoreError::Aborted) => {
log::info!("Agent::prompt: generation aborted by user");
let assistant_msg = Message::assistant(self.next_id(), "");
self.messages.push(assistant_msg.clone());
new_messages.push(assistant_msg.clone());
tx.send(AgentEvent::MessageEnd {
message: assistant_msg.clone(),
})
.ok();
tx.send(AgentEvent::TurnEnd {
message: assistant_msg,
tool_results: vec![],
})
.ok();
tx.send(AgentEvent::AgentEnd {
messages: new_messages,
})
.ok();
return Ok(());
}
Err(e) => {
log::error!("Agent::prompt: generation error: {e}");
if iteration == 0 {
self.messages.pop();
}
tx.send(AgentEvent::Error {
message: e.to_string(),
})
.ok();
tx.send(AgentEvent::AgentEnd { messages: vec![] }).ok();
return Ok(());
}
};
log::debug!(
"Agent::prompt: turn {} → {} tokens, {:.1} t/s, {:.1}ms ttft",
iteration,
gen.tokens_generated,
gen.tokens_per_sec,
gen.time_to_first_token_ms,
);
let mut assistant_msg = Message::assistant(self.next_id(), &gen.text);
let parsed = if has_tools {
parse_tool_calls(&gen.text)
} else {
Vec::new()
};
let tool_calls: Vec<ToolCall> = parsed
.iter()
.enumerate()
.map(|(i, p)| ToolCall {
id: format!("{}-call-{}", assistant_msg.id, i + 1),
name: p.name.clone(),
arguments: p.arguments.clone(),
})
.collect();
assistant_msg.tool_calls = tool_calls.clone();
self.messages.push(assistant_msg.clone());
new_messages.push(assistant_msg.clone());
tx.send(AgentEvent::GenerationStats {
tokens_generated: gen.tokens_generated,
prompt_tokens: gen.prompt_tokens,
tokens_per_sec: gen.tokens_per_sec,
time_to_first_token_ms: gen.time_to_first_token_ms,
generation_time_ms: gen.generation_time_ms,
})
.ok();
tx.send(AgentEvent::MessageEnd {
message: assistant_msg.clone(),
})
.ok();
if tool_calls.is_empty() {
tx.send(AgentEvent::TurnEnd {
message: assistant_msg,
tool_results: vec![],
})
.ok();
tx.send(AgentEvent::AgentEnd {
messages: new_messages,
})
.ok();
return Ok(());
}
#[cfg(feature = "tools")]
{
let aborted = self
.run_tool_calls(&tool_calls, assistant_msg, &mut new_messages, &tx)
.await;
if aborted {
tx.send(AgentEvent::AgentEnd {
messages: new_messages,
})
.ok();
return Ok(());
}
}
}
log::warn!(
"Agent::prompt: stopped after {} tool iterations",
self.config.max_tool_iterations
);
tx.send(AgentEvent::Warning {
message: format!(
"Stopped after {} tool iterations without a final answer",
self.config.max_tool_iterations
),
})
.ok();
tx.send(AgentEvent::AgentEnd {
messages: new_messages,
})
.ok();
Ok(())
}
#[cfg(feature = "tools")]
async fn run_tool_calls(
&mut self,
tool_calls: &[ToolCall],
assistant_msg: Message,
new_messages: &mut Vec<Message>,
tx: &mpsc::UnboundedSender<AgentEvent>,
) -> bool {
let mut tool_results: Vec<ToolResult> = Vec::new();
for call in tool_calls {
let decision = match &self.approval_hook {
Some(hook) => hook.review(call).await,
None => ApprovalDecision::Approve,
};
let (content, is_error) = match decision {
ApprovalDecision::Deny { reason } => {
log::info!("Agent::prompt: tool '{}' denied: {reason}", call.name);
tx.send(AgentEvent::ToolDenied {
tool_call_id: call.id.clone(),
tool_name: call.name.clone(),
reason: reason.clone(),
})
.ok();
(reason, true)
}
ApprovalDecision::Approve => {
tx.send(AgentEvent::ToolExecStart {
tool_call_id: call.id.clone(),
tool_name: call.name.clone(),
args: call.arguments.clone(),
})
.ok();
let (content, is_error) = match self.execute_tool(call, tx).await {
Ok(out) => (out.content, false),
Err(e) => {
log::warn!("Agent::prompt: tool '{}' failed: {e}", call.name);
(e.to_string(), true)
}
};
tx.send(AgentEvent::ToolExecEnd {
tool_call_id: call.id.clone(),
tool_name: call.name.clone(),
result: ToolResult {
tool_call_id: call.id.clone(),
tool_name: call.name.clone(),
content: content.clone(),
is_error,
},
})
.ok();
(content, is_error)
}
};
let result = ToolResult {
tool_call_id: call.id.clone(),
tool_name: call.name.clone(),
content: content.clone(),
is_error,
};
let result_msg =
Message::tool_result(self.next_id(), &call.id, &call.name, content, is_error);
self.messages.push(result_msg.clone());
new_messages.push(result_msg.clone());
tx.send(AgentEvent::MessageStart {
message: result_msg.clone(),
})
.ok();
tx.send(AgentEvent::MessageEnd {
message: result_msg,
})
.ok();
tool_results.push(result);
}
tx.send(AgentEvent::TurnEnd {
message: assistant_msg,
tool_results,
})
.ok();
self.abort.load(Ordering::Relaxed)
}
pub fn prompt_stream(
&mut self,
text: impl Into<String>,
backend: Arc<dyn LlmBackend>,
) -> (
mpsc::UnboundedReceiver<AgentEvent>,
impl std::future::Future<Output = CoreResult<()>> + '_,
) {
let (tx, rx) = mpsc::unbounded_channel();
let text = text.into();
let fut = async move { self.prompt(text, backend, tx).await };
(rx, fut)
}
async fn generate_once(
&self,
backend: Arc<dyn LlmBackend>,
tx: &mpsc::UnboundedSender<AgentEvent>,
) -> CoreResult<GenerationResult> {
let messages = self.messages.clone();
let system_prompt = self.config.system_prompt.clone();
let ctx_config = self.config.context_config.clone();
let tool_schemas = self.tool_schemas();
let params = self.config.inference_params.clone();
let abort = self.abort.clone();
let max_ctx = self.config.context_config.max_context_tokens;
let template = self.template.clone();
let token_tx = tx.clone();
let budget_tx = tx.clone();
log::debug!(
"Agent::generate_once: spawning blocking (max_tokens={}, temp={}, ctx={}, threads={})",
params.max_tokens,
params.temperature,
params.context_size,
params.n_threads,
);
let handle = tokio::task::spawn_blocking(move || {
if !backend.is_ready() {
return Err(CoreError::Backend("No model loaded".into()));
}
let prepared = prepare_context(
template.as_ref(),
&system_prompt,
&messages,
&tool_schemas,
&ctx_config,
&|text| backend.tokenize_count(text).unwrap_or(0),
)?;
log::debug!(
"Context prepared: tokens={}, kept={}, pruned={}",
prepared.token_count,
prepared.messages_included,
prepared.messages_pruned,
);
budget_tx
.send(AgentEvent::ContextBudget {
used_tokens: prepared.token_count,
max_tokens: max_ctx,
messages_in_context: prepared.messages_included,
messages_pruned: prepared.messages_pruned,
})
.ok();
backend.generate(
&prepared.prompt,
¶ms,
abort,
Box::new(move |token, count, tps| {
token_tx
.send(AgentEvent::MessageDelta {
delta: token.to_string(),
tokens_generated: count,
tokens_per_sec: tps,
})
.ok();
}),
)
});
handle.await.map_err(|e| {
log::error!("Agent::generate_once: blocking task panicked: {e}");
CoreError::Agent(format!("Inference task failed: {e}"))
})?
}
#[cfg(feature = "tools")]
async fn execute_tool(
&self,
call: &ToolCall,
tx: &mpsc::UnboundedSender<AgentEvent>,
) -> CoreResult<ToolOutput> {
let Some(tool) = self.tools.iter().find(|t| t.name() == call.name) else {
return Err(CoreError::Tool(format!("unknown tool: {}", call.name)));
};
let update_tx = tx.clone();
let tool_call_id = call.id.clone();
let tool_name = call.name.clone();
let on_update: ToolUpdateCallback = Box::new(move |partial: &str| {
update_tx
.send(AgentEvent::ToolExecUpdate {
tool_call_id: tool_call_id.clone(),
tool_name: tool_name.clone(),
partial: partial.to_string(),
})
.ok();
});
tool.execute(&call.id, call.arguments.clone(), Some(on_update))
.await
}
async fn compress_if_needed(
&mut self,
backend: &Arc<dyn LlmBackend>,
tx: &mpsc::UnboundedSender<AgentEvent>,
) {
if self.config.context_config.prune_strategy != PruneStrategy::Summarize {
return;
}
let messages = self.messages.clone();
let system_prompt = self.config.system_prompt.clone();
let tools = self.tool_schemas();
let ctx_config = self.config.context_config.clone();
let template = self.template.clone();
let abort = self.abort.clone();
let params = self.config.inference_params.clone();
let backend = backend.clone();
let outcome = tokio::task::spawn_blocking(move || -> Option<(Vec<usize>, String)> {
if !backend.is_ready() {
return None;
}
let counter = |t: &str| backend.tokenize_count(t).unwrap_or(0);
let plan = plan_prune(
template.as_ref(),
&system_prompt,
&messages,
&tools,
&ctx_config,
&counter,
)
.ok()?;
if plan.dropped.is_empty() {
return None; }
let mut remove: Vec<usize> = plan.dropped.iter().flat_map(|r| r.clone()).collect();
let prior_summary = messages
.iter()
.position(|m| m.pinned && m.content.starts_with(SUMMARY_MARKER));
let prior_body = prior_summary.map(|i| {
remove.push(i);
messages[i]
.content
.strip_prefix(SUMMARY_MARKER)
.unwrap_or(&messages[i].content)
.trim()
.to_string()
});
remove.sort_unstable();
remove.dedup();
let transcript = render_transcript(&messages, &remove);
let mut body = String::new();
if let Some(prev) = prior_body.filter(|s| !s.is_empty()) {
body.push_str("Earlier summary:\n");
body.push_str(&prev);
body.push_str("\n\n");
}
body.push_str("Conversation excerpt:\n");
body.push_str(&transcript);
let instruction = "You compress conversation history. Summarize the \
material below into a concise note that preserves key facts, \
decisions, names, and unresolved questions. Reply with only the \
summary.";
let req = Message::user("summary-req", format!("{instruction}\n\n{body}"));
let prompt = template.format(
"You summarize conversations faithfully and concisely.",
std::slice::from_ref(&req),
&[],
);
let sum_params = InferenceParams {
max_tokens: SUMMARY_MAX_TOKENS,
..params
};
let gen = backend
.generate(&prompt, &sum_params, abort, Box::new(|_, _, _| {}))
.ok()?;
let summary = gen.text.trim().to_string();
if summary.is_empty() {
return None;
}
Some((remove, summary))
})
.await;
let Some((remove, summary)) = outcome.ok().flatten() else {
return;
};
self.fold_into_summary(&remove, summary);
tx.send(AgentEvent::Warning {
message: format!(
"Summarized {} earlier message(s) to fit the context window",
remove.len()
),
})
.ok();
}
fn fold_into_summary(&mut self, remove: &[usize], summary: String) {
if remove.is_empty() {
return;
}
let insert_at = *remove.iter().min().unwrap();
let mut sorted = remove.to_vec();
sorted.sort_unstable();
for &i in sorted.iter().rev() {
if i < self.messages.len() {
self.messages.remove(i);
}
}
let summary_msg =
Message::user(self.next_id(), format!("{SUMMARY_MARKER}\n{summary}")).pinned();
let at = insert_at.min(self.messages.len());
self.messages.insert(at, summary_msg);
log::info!(
"Folded {} messages into a pinned summary at index {at}",
remove.len()
);
}
}
fn render_transcript(messages: &[Message], indices: &[usize]) -> String {
indices
.iter()
.filter_map(|&i| messages.get(i))
.map(|m| {
let role = match m.role {
Role::User => "User",
Role::Assistant | Role::ToolCall => "Assistant",
Role::ToolResult => "Tool",
Role::System => "System",
};
format!("{role}: {}", m.content)
})
.collect::<Vec<_>>()
.join("\n")
}