use std::collections::{BTreeMap, HashMap};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use edgecrab_tools::config_ref::AppConfigRef;
use edgecrab_tools::registry::{
ApprovalRequest, ApprovalResponse, ToolContext, ToolRegistry, to_llm_definitions,
};
use edgecrab_types::trajectory::{
TrajectoryMetadata, convert_scratchpad_to_think, save_trajectory,
};
use edgecrab_types::{
AgentError, Content, Cost, Message, Role, ToolError, ToolErrorResponse, Trajectory, Usage,
};
use edgequake_llm::traits::{StreamChunk, StreamUsage};
use edgequake_llm::{CachePromptConfig, LLMProvider, VsCodeCopilotProvider, apply_cache_control};
use futures::StreamExt;
use tokio_util::sync::CancellationToken;
use crate::agent::{Agent, ConversationResult, SessionState};
use crate::compression::{
CompressionParams, CompressionStatus, check_compression_status, compress_with_llm,
};
use crate::config::edgecrab_home;
use crate::context_references::expand_context_refs_with_policy;
use crate::model_router::{RoutingThresholds, SmartRoutingConfig, resolve_turn_route};
use crate::pricing::{CanonicalUsage, estimate_cost};
use crate::prompt_builder::{
PromptBuilder, load_global_soul, load_memory_sections, load_preloaded_skills,
load_skill_summary,
};
use crate::sub_agent_runner::CoreSubAgentRunner;
const MAX_RETRIES: u32 = 3;
const BASE_BACKOFF: Duration = Duration::from_millis(500);
const SKILL_REFLECTION_THRESHOLD: u32 = 5;
#[allow(clippy::too_many_arguments)]
fn build_tool_context(
cwd: &std::path::Path,
app_config_ref: AppConfigRef,
cancel: &CancellationToken,
state_db: &Option<Arc<edgecrab_state::SessionDb>>,
platform: edgecrab_types::Platform,
process_table: &Arc<edgecrab_tools::ProcessTable>,
provider: Option<Arc<dyn edgequake_llm::LLMProvider>>,
tool_registry: Option<Arc<ToolRegistry>>,
sub_agent_runner: Option<Arc<dyn edgecrab_tools::SubAgentRunner>>,
clarify_tx: Option<
tokio::sync::mpsc::UnboundedSender<edgecrab_tools::registry::ClarifyRequest>,
>,
approval_tx: Option<tokio::sync::mpsc::UnboundedSender<ApprovalRequest>>,
gateway_sender: Option<Arc<dyn edgecrab_tools::registry::GatewaySender>>,
origin_chat: Option<(String, String)>,
conversation_session_id: &str,
todo_store: Option<Arc<edgecrab_tools::TodoStore>>,
) -> ToolContext {
ToolContext {
task_id: uuid::Uuid::new_v4().to_string(),
cwd: cwd.to_path_buf(),
session_id: conversation_session_id.to_string(),
user_task: None,
cancel: cancel.clone(),
config: app_config_ref,
state_db: state_db.clone(),
platform,
process_table: Some(process_table.clone()),
provider,
tool_registry,
delegate_depth: 0,
sub_agent_runner,
clarify_tx,
approval_tx,
on_skills_changed: Some(std::sync::Arc::new(
crate::prompt_builder::invalidate_skills_cache,
)),
gateway_sender,
origin_chat: origin_chat.clone(),
session_key: Some(match &origin_chat {
Some((platform, chat_id)) => format!("{}:{}", platform, chat_id),
None => conversation_session_id.to_string(),
}),
todo_store,
}
}
enum LoopAction {
Continue,
Done(String),
}
struct DispatchContext<'a> {
cwd: std::path::PathBuf,
registry: Option<&'a Arc<ToolRegistry>>,
cancel: &'a CancellationToken,
state_db: &'a Option<Arc<edgecrab_state::SessionDb>>,
platform: edgecrab_types::Platform,
process_table: &'a Arc<edgecrab_tools::ProcessTable>,
provider: Option<Arc<dyn edgequake_llm::LLMProvider>>,
gateway_sender: Option<Arc<dyn edgecrab_tools::registry::GatewaySender>>,
sub_agent_runner: Option<Arc<dyn edgecrab_tools::SubAgentRunner>>,
event_tx: Option<&'a tokio::sync::mpsc::UnboundedSender<crate::StreamEvent>>,
clarify_tx:
Option<tokio::sync::mpsc::UnboundedSender<edgecrab_tools::registry::ClarifyRequest>>,
approval_tx: Option<tokio::sync::mpsc::UnboundedSender<ApprovalRequest>>,
origin_chat: Option<(String, String)>,
app_config_ref: AppConfigRef,
conversation_session_id: String,
todo_store: Option<Arc<edgecrab_tools::TodoStore>>,
capability_suppressions: Arc<Mutex<HashMap<String, ToolErrorResponse>>>,
}
impl Agent {
pub(crate) async fn execute_loop(
&self,
user_message: &str,
system_message: Option<&str>,
history: Option<Vec<Message>>,
event_tx: Option<&tokio::sync::mpsc::UnboundedSender<crate::StreamEvent>>,
cwd_override: Option<&std::path::Path>,
) -> Result<ConversationResult, AgentError> {
self.budget.reset();
let cancel = {
let mut guard = self.cancel.lock().expect("cancel mutex not poisoned");
if guard.is_cancelled() {
*guard = CancellationToken::new();
}
guard.clone()
};
let config = self.config.read().await.clone();
let provider = self.provider.read().await.clone();
let mut session = self.session.write().await;
if let Some(hist) = history {
session.messages = hist;
}
let cwd = cwd_override
.map(std::path::Path::to_path_buf)
.unwrap_or_else(|| {
std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."))
});
let expanded_enabled =
edgecrab_tools::toolsets::expand_toolset_names(&config.enabled_toolsets);
let expanded_disabled =
edgecrab_tools::toolsets::expand_toolset_names(&config.disabled_toolsets);
let parent_active_toolsets = if config.enabled_toolsets.is_empty()
|| edgecrab_tools::toolsets::contains_all_sentinel(&config.enabled_toolsets)
|| expanded_enabled.is_empty()
{
Vec::new()
} else {
expanded_enabled
.iter()
.filter(|toolset| !expanded_disabled.iter().any(|d| d == *toolset))
.cloned()
.collect()
};
let app_config_ref = AppConfigRef {
edgecrab_home: crate::config::edgecrab_home(),
file_allowed_roots: config.file_allowed_roots.clone(),
path_restrictions: config.path_restrictions.clone(),
delegation_enabled: config.delegation_enabled,
delegation_model: config.delegation_model.clone(),
delegation_provider: config.delegation_provider.clone(),
delegation_max_subagents: config.delegation_max_subagents,
delegation_max_iterations: config.delegation_max_iterations,
parent_active_toolsets,
disabled_toolsets: expanded_disabled.clone(),
external_skill_dirs: config.skills_config.external_dirs.clone(),
disabled_skills: {
let mut d = config.skills_config.disabled.clone();
let platform_str_cfg = config.platform.to_string();
if let Some(pd) = config
.skills_config
.platform_disabled
.get(&platform_str_cfg)
{
d.extend(pd.iter().cloned());
}
d
},
browser_record_sessions: config.browser.record_sessions,
browser_command_timeout: config.browser.command_timeout,
browser_recording_max_age_hours: config.browser.recording_max_age_hours,
checkpoints_enabled: config.checkpoints_enabled,
checkpoints_max_snapshots: config.checkpoints_max_snapshots,
preloaded_skills: config.skills_config.preloaded.clone(),
terminal_backend: config.terminal_backend.clone(),
terminal_docker: config.terminal_docker.clone(),
terminal_ssh: config.terminal_ssh.clone(),
terminal_modal: config.terminal_modal.clone(),
terminal_daytona: config.terminal_daytona.clone(),
terminal_singularity: config.terminal_singularity.clone(),
terminal_env_passthrough: config.terminal_env_passthrough.clone(),
auxiliary_provider: config.auxiliary.provider.clone(),
auxiliary_model: config.auxiliary.model.clone(),
auxiliary_base_url: config.auxiliary.base_url.clone(),
auxiliary_api_key_env: config.auxiliary.api_key_env.clone(),
..Default::default()
};
if !config.terminal_env_passthrough.is_empty() {
edgecrab_tools::tools::backends::local::register_env_passthrough(
&config.terminal_env_passthrough,
);
}
let (tool_defs, tool_names_for_prompt) = if let Some(ref registry) = self.tool_registry {
let ctx = build_tool_context(
&cwd,
app_config_ref.clone(),
&cancel,
&self.state_db,
config.platform,
&self.process_table,
Some(provider.clone()),
self.tool_registry.clone(),
None,
None, None, self.gateway_sender.read().await.clone(),
config.origin_chat.clone(),
"schema-resolution", Some(self.todo_store.clone()),
);
let enabled_filter = if config.enabled_toolsets.is_empty()
|| edgecrab_tools::toolsets::contains_all_sentinel(&config.enabled_toolsets)
|| expanded_enabled.is_empty()
{
None
} else {
Some(expanded_enabled.as_slice())
};
let disabled_filter = if expanded_disabled.is_empty() {
None
} else {
Some(expanded_disabled.as_slice())
};
let schemas = registry.get_definitions(enabled_filter, disabled_filter, &ctx);
let names: Vec<String> = schemas.iter().map(|s| s.name.clone()).collect();
(to_llm_definitions(&schemas), names)
} else {
(Vec::new(), Vec::new())
};
if session.cached_system_prompt.is_none() {
let prompt = if let Some(explicit) = system_message {
explicit.to_string()
} else {
let home = edgecrab_home();
let memory_sections = if config.skip_memory {
Vec::new()
} else {
load_memory_sections(&home)
};
let platform_str = config.platform.to_string();
let mut disabled_skills = config.skills_config.disabled.clone();
if let Some(platform_disabled) =
config.skills_config.platform_disabled.get(&platform_str)
{
disabled_skills.extend(platform_disabled.iter().cloned());
}
let skill_summary = load_skill_summary(&home, &disabled_skills, None, None);
let preloaded_content =
load_preloaded_skills(&home, &config.skills_config.preloaded);
let combined_skill_prompt: Option<String> =
match (preloaded_content.is_empty(), skill_summary) {
(false, Some(summary)) => Some(format!("{preloaded_content}\n\n{summary}")),
(false, None) => Some(preloaded_content),
(true, summary) => summary,
};
let global_soul = load_global_soul(&home);
PromptBuilder::new(config.platform)
.skip_context_files(config.skip_context_files)
.available_tools(tool_names_for_prompt)
.build(
global_soul.as_deref(), Some(&cwd),
&memory_sections,
combined_skill_prompt.as_deref(),
)
};
let prompt = if let Some(ref addon) = config.personality_addon {
format!("{prompt}\n\n## Personality\n\n{addon}")
} else {
prompt
};
session.cached_system_prompt = Some(prompt);
}
let context_path_policy = app_config_ref.file_path_policy(&cwd);
let mut expansion =
expand_context_refs_with_policy(user_message, &cwd, &context_path_policy);
if !expansion.refs_found.is_empty() {
tracing::debug!(
refs = expansion.refs_found.len(),
errors = expansion.errors.len(),
"expanded @context references"
);
}
for err in &expansion.errors {
tracing::warn!(error = %err, "context reference expansion error");
}
if !expansion.refs_found.is_empty() {
let context_window = CompressionParams::default().context_window;
let injected_chars = expansion.expanded.len().saturating_sub(user_message.len());
let injected_tokens = injected_chars / 4;
let hard_limit = context_window / 2; let soft_limit = context_window / 4;
if injected_tokens > hard_limit {
tracing::warn!(
injected_tokens,
hard_limit,
"@context injection exceeds 50% of context window — stripping injected content"
);
let notice = format!(
"{user_message}\n\n[Warning: @context injection (~{injected_tokens} tokens) \
exceeds the 50% context-window limit ({hard_limit} tokens). \
Injected content was removed to protect the context budget.]"
);
expansion.expanded = notice;
expansion.budget_blocked = true;
} else if injected_tokens > soft_limit {
tracing::warn!(
injected_tokens,
soft_limit,
"@context injection exceeds 25% of context window — approaching budget limit"
);
expansion.budget_warning = true;
}
}
let smart_routing = SmartRoutingConfig {
enabled: config.model_config.smart_routing.enabled,
cheap_model: config.model_config.smart_routing.cheap_model.clone(),
cheap_base_url: config.model_config.smart_routing.cheap_base_url.clone(),
cheap_api_key_env: config.model_config.smart_routing.cheap_api_key_env.clone(),
thresholds: RoutingThresholds::default(),
};
let route = resolve_turn_route(&expansion.expanded, &config.model_config, &smart_routing);
if let Some(ref label) = route.label {
tracing::info!(route = %label, "model routing decision");
}
let effective_provider = if !route.is_primary {
if let Some((prov_name, model_name)) = route.model.split_once('/') {
let canonical = match prov_name {
"copilot" => "vscode-copilot",
other => other,
};
let cheap_opt: Option<Arc<dyn LLMProvider>> = if canonical == "vscode-copilot" {
match VsCodeCopilotProvider::new()
.model(model_name)
.with_vision(true)
.build()
{
Ok(p) => Some(Arc::new(p) as Arc<dyn LLMProvider>),
Err(e) => {
tracing::warn!(error = %e, "failed to create copilot provider, using primary");
None
}
}
} else {
let is_gemini_canonical =
matches!(canonical, "google" | "gemini" | "vertex" | "vertexai");
let primary_is_vertex = provider.name() == "vertex-ai";
let (effective_canonical, effective_model) =
if is_gemini_canonical && primary_is_vertex {
tracing::info!(
cheap_model = %route.model,
"smart routing: using Vertex AI endpoint for cheap Gemini model \
(primary is vertex-ai)"
);
let bare = model_name.strip_prefix("vertexai:").unwrap_or(model_name);
("vertexai", bare)
} else {
(canonical, model_name)
};
match edgequake_llm::ProviderFactory::create_llm_provider(
effective_canonical,
effective_model,
) {
Ok(p) => Some(p),
Err(e) => {
tracing::warn!(error = %e, "failed to create cheap model provider, using primary");
None
}
}
};
match cheap_opt {
Some(cheap) => {
tracing::info!(model = %route.model, "using smart-routed cheap model");
cheap
}
None => provider.clone(),
}
} else {
provider.clone()
}
} else {
provider.clone()
};
let injection_threats = crate::prompt_builder::scan_for_injection(&expansion.expanded);
if !injection_threats.is_empty() {
tracing::warn!(
threats = injection_threats.len(),
"prompt injection patterns detected in user input"
);
for threat in &injection_threats {
tracing::warn!(
pattern = %threat.pattern_name,
severity = ?threat.severity,
"injection threat"
);
}
}
session.messages.push(Message::user(&expansion.expanded));
session.user_turn_count += 1;
let initial_turn_tool_call_count = session.session_tool_call_count;
let sub_agent_runner: Option<Arc<dyn edgecrab_tools::SubAgentRunner>> =
if let Some(ref registry) = self.tool_registry {
Some(Arc::new(CoreSubAgentRunner::new(
provider.clone(),
registry.clone(),
config.platform,
config.model.clone(),
)))
} else {
None
};
let (clarify_req_tx, mut clarify_req_rx) =
tokio::sync::mpsc::unbounded_channel::<edgecrab_tools::registry::ClarifyRequest>();
let (approval_req_tx, mut approval_req_rx) =
tokio::sync::mpsc::unbounded_channel::<ApprovalRequest>();
if let Some(ev_tx) = event_tx {
let clarify_ev_tx = ev_tx.clone();
tokio::spawn(async move {
while let Some(req) = clarify_req_rx.recv().await {
let _ = clarify_ev_tx.send(crate::StreamEvent::Clarify {
question: req.question,
choices: req.choices,
response_tx: req.response_tx,
});
}
});
let approval_ev_tx = ev_tx.clone();
tokio::spawn(async move {
while let Some(req) = approval_req_rx.recv().await {
let (decision_tx, decision_rx) =
tokio::sync::oneshot::channel::<crate::ApprovalChoice>();
let _ = approval_ev_tx.send(crate::StreamEvent::Approval {
command: req.command,
full_command: req.full_command,
reasons: req.reasons,
response_tx: decision_tx,
});
let mapped = match decision_rx.await {
Ok(crate::ApprovalChoice::Once) => ApprovalResponse::Once,
Ok(crate::ApprovalChoice::Session) => ApprovalResponse::Session,
Ok(crate::ApprovalChoice::Always) => ApprovalResponse::Always,
Ok(crate::ApprovalChoice::Deny) | Err(_) => ApprovalResponse::Deny,
};
let _ = req.response_tx.send(mapped);
}
});
}
let clarify_tx_for_dispatch = Some(clarify_req_tx);
let approval_tx_for_dispatch = Some(approval_req_tx);
let turn_started_at = std::time::Instant::now();
let conversation_session_id = session
.session_id
.clone()
.or_else(|| config.session_id.clone())
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
if session.session_id.is_none() {
session.session_id = Some(conversation_session_id.clone());
}
let mut final_response = String::new();
let mut interrupted = false;
let mut budget_exhausted = false;
let mut tool_errors_acc: Vec<edgecrab_types::ToolErrorRecord> = Vec::new();
let capability_suppressions: Arc<Mutex<HashMap<String, ToolErrorResponse>>> =
Arc::new(Mutex::new(HashMap::new()));
let mut pressure_warned = false;
loop {
if !self.budget.try_consume() {
tracing::warn!(
used = self.budget.used(),
max = self.budget.max(),
"iteration budget exhausted"
);
budget_exhausted = true;
break;
}
if cancel.is_cancelled() {
interrupted = true;
break;
}
let compression_params = CompressionParams::default();
match check_compression_status(&session.messages, &compression_params) {
CompressionStatus::NeedsCompression => {
tracing::info!(
messages = session.messages.len(),
"compressing context before API call"
);
session.messages =
compress_with_llm(&session.messages, &compression_params, &provider).await;
if let Some(snapshot) = self.todo_store.format_for_injection() {
session.messages.push(Message::user(&snapshot));
}
if check_compression_status(&session.messages, &compression_params)
== CompressionStatus::Ok
{
pressure_warned = false;
}
}
CompressionStatus::PressureWarning if !pressure_warned => {
let estimated = crate::compression::estimate_tokens(&session.messages);
let threshold_tokens = (compression_params.context_window as f32
* compression_params.threshold)
as usize;
tracing::warn!(
estimated_tokens = estimated,
threshold_tokens,
"context approaching compression threshold"
);
if let Some(tx) = event_tx {
let _ = tx.send(crate::StreamEvent::ContextPressure {
estimated_tokens: estimated,
threshold_tokens,
});
}
pressure_warned = true;
}
_ => {}
}
let cache_cfg = if config.model_config.prompt_caching {
Some(CachePromptConfig::default())
} else {
None
};
let chat_messages = build_chat_messages(
session.cached_system_prompt.as_deref(),
&session.messages,
cache_cfg.as_ref(),
);
sanitize_orphaned_tool_results(&mut session.messages);
let tool_defs = if let Some(ref registry) = self.tool_registry {
let schema_ctx = build_tool_context(
&cwd,
app_config_ref.clone(),
&cancel,
&self.state_db,
config.platform,
&self.process_table,
Some(effective_provider.clone()),
self.tool_registry.clone(),
None,
None,
None,
self.gateway_sender.read().await.clone(),
config.origin_chat.clone(),
&conversation_session_id,
Some(self.todo_store.clone()),
);
let enabled_filter = if config.enabled_toolsets.is_empty()
|| edgecrab_tools::toolsets::contains_all_sentinel(&config.enabled_toolsets)
|| expanded_enabled.is_empty()
{
None
} else {
Some(expanded_enabled.as_slice())
};
let disabled_filter = if expanded_disabled.is_empty() {
None
} else {
Some(expanded_disabled.as_slice())
};
let schemas =
registry.get_definitions(enabled_filter, disabled_filter, &schema_ctx);
to_llm_definitions(&schemas)
} else {
Vec::new()
};
{
let names: Vec<&str> = tool_defs.iter().map(|t| t.function.name.as_str()).collect();
tracing::debug!(
tools = ?names,
browser_present = names.iter().any(|&n| n.starts_with("browser_")),
"tool schema for next API call"
);
if names.iter().any(|&n| n.starts_with("browser_")) {
tracing::info!(
"✓ browser tools present in schema ({})",
names.iter().filter(|&&n| n.starts_with("browser_")).count()
);
} else {
tracing::warn!("✗ NO browser tools in schema — model will use mcp_call_tool");
}
}
let native_streaming_active = config.streaming
&& event_tx.is_some()
&& effective_provider.supports_tool_streaming();
let response = match api_call_with_retry(
&effective_provider,
&chat_messages,
&tool_defs,
MAX_RETRIES,
&cancel,
event_tx,
native_streaming_active,
)
.await
{
Ok(r) => r,
Err(AgentError::Interrupted) => {
interrupted = true;
break;
}
Err(primary_err) => {
if native_streaming_active {
return Err(primary_err);
}
if let Some(ref fb) = config.model_config.fallback {
let fb_route = crate::model_router::fallback_route(fb);
tracing::warn!(
primary_error = %primary_err,
fallback = %fb_route.model,
"primary API failed, trying fallback"
);
if let Some((fb_prov_name, fb_model_name)) = fb_route.model.split_once('/')
{
let fb_canonical = match fb_prov_name {
"copilot" => "vscode-copilot",
other => other,
};
let fb_prov_opt: Option<Arc<dyn LLMProvider>> =
if fb_canonical == "vscode-copilot" {
VsCodeCopilotProvider::new()
.model(fb_model_name)
.with_vision(true) .build()
.ok()
.map(|p| Arc::new(p) as Arc<dyn LLMProvider>)
} else {
edgequake_llm::ProviderFactory::create_llm_provider(
fb_canonical,
fb_model_name,
)
.ok()
};
if let Some(fb_prov) = fb_prov_opt {
let fallback_native_streaming = config.streaming
&& event_tx.is_some()
&& fb_prov.supports_tool_streaming();
match api_call_with_retry(
&fb_prov,
&chat_messages,
&tool_defs,
1,
&cancel,
event_tx,
fallback_native_streaming,
)
.await
{
Ok(r) => r,
Err(AgentError::Interrupted) => {
interrupted = true;
break;
}
Err(fb_err) => {
tracing::error!(fallback_error = %fb_err, "fallback also failed");
return Err(primary_err);
}
}
} else {
return Err(primary_err);
}
} else {
return Err(primary_err);
}
} else {
return Err(primary_err);
}
}
};
if cancel.is_cancelled() {
interrupted = true;
break;
}
session.api_call_count += 1;
session.session_input_tokens += response.prompt_tokens as u64;
session.session_output_tokens += response.completion_tokens as u64;
if let Some(cache_tokens) = response.cache_hit_tokens {
session.session_cache_read_tokens += cache_tokens as u64;
}
if let Some(reasoning_tokens) = response.thinking_tokens {
session.session_reasoning_tokens += reasoning_tokens as u64;
}
if response.content.trim().is_empty()
&& !response.has_tool_calls()
&& response.finish_reason.as_deref() != Some("length")
{
tracing::info!("empty response from LLM, nudging to continue");
session.messages.push(Message::user(
"[system: your response was empty — please provide a response]",
));
continue;
}
let dctx = DispatchContext {
cwd: cwd.clone(),
registry: self.tool_registry.as_ref(),
cancel: &cancel,
state_db: &self.state_db,
platform: config.platform,
process_table: &self.process_table,
provider: Some(provider.clone()),
gateway_sender: self.gateway_sender.read().await.clone(),
sub_agent_runner: sub_agent_runner.clone(),
event_tx,
clarify_tx: clarify_tx_for_dispatch.clone(),
approval_tx: approval_tx_for_dispatch.clone(),
origin_chat: config.origin_chat.clone(),
app_config_ref: app_config_ref.clone(),
conversation_session_id: conversation_session_id.clone(),
todo_store: Some(self.todo_store.clone()),
capability_suppressions: capability_suppressions.clone(),
};
let action =
process_response(&response, &mut session, &dctx, &mut tool_errors_acc).await?;
match action {
LoopAction::Done(text) => {
if response.finish_reason.as_deref() == Some("length") {
tracing::info!(
partial_len = text.len(),
"response truncated (finish_reason=length), auto-continuing"
);
session.messages.push(Message::user(
"[system: your response was truncated due to length — please continue exactly where you left off]",
));
continue;
}
final_response = text;
break;
}
LoopAction::Continue => {
if let Some(warning) =
get_budget_warning(session.api_call_count, config.max_iterations)
{
inject_budget_warning(&mut session.messages, &warning);
}
continue;
}
}
}
if budget_exhausted && final_response.is_empty() {
let msg = format!(
"[Agent reached the {} iteration limit before completing the task. \
Please try rephrasing your request or increase the iteration budget.]",
self.budget.max()
);
tracing::warn!(
max = self.budget.max(),
"emitting budget-exhausted fallback response"
);
session.messages.push(Message::assistant(&msg));
if let Some(tx) = event_tx {
let _ = tx.send(crate::StreamEvent::Token(msg.clone()));
}
final_response = msg;
}
let turn_tool_calls = session
.session_tool_call_count
.saturating_sub(initial_turn_tool_call_count);
if !interrupted
&& !config.skip_memory
&& self.tool_registry.is_some()
&& turn_tool_calls >= SKILL_REFLECTION_THRESHOLD
{
let bg_ctx = BackgroundReflectionCtx {
messages: session.messages.clone(),
system_prompt: session.cached_system_prompt.clone(),
tool_defs: tool_defs.clone(),
cwd: cwd.clone(),
registry: self.tool_registry.as_ref().map(Arc::clone),
cancel: cancel.clone(),
state_db: self.state_db.clone(),
platform: config.platform,
process_table: Arc::clone(&self.process_table),
provider: Arc::clone(&effective_provider),
gateway_sender: self.gateway_sender.read().await.clone(),
sub_agent_runner: sub_agent_runner.clone(),
app_config_ref: app_config_ref.clone(),
conversation_session_id: conversation_session_id.clone(),
origin_chat: config.origin_chat.clone(),
todo_store: Some(self.todo_store.clone()),
};
tokio::spawn(run_learning_reflection_bg(bg_ctx));
}
let session_id = session
.session_id
.clone()
.or_else(|| config.session_id.clone())
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
session.session_id = Some(session_id.clone());
let usage = Usage {
input_tokens: session.session_input_tokens,
output_tokens: session.session_output_tokens,
cache_read_tokens: session.session_cache_read_tokens,
cache_write_tokens: session.session_cache_write_tokens,
reasoning_tokens: session.session_reasoning_tokens,
..Default::default()
};
let canonical_usage = CanonicalUsage {
input_tokens: session.session_input_tokens,
output_tokens: session.session_output_tokens,
cache_read_tokens: session.session_cache_read_tokens,
cache_write_tokens: session.session_cache_write_tokens,
reasoning_tokens: session.session_reasoning_tokens,
};
let cost_result = estimate_cost(&canonical_usage, &config.model);
let cost = Cost {
input_cost: canonical_usage.input_tokens as f64 * cost_result.amount_usd.unwrap_or(0.0)
/ canonical_usage.total_tokens().max(1) as f64,
output_cost: canonical_usage.output_tokens as f64
* cost_result.amount_usd.unwrap_or(0.0)
/ canonical_usage.total_tokens().max(1) as f64,
total_cost: cost_result.amount_usd.unwrap_or(0.0),
..Default::default()
};
if let Some(ref db) = self.state_db {
let title = session
.messages
.iter()
.find(|m| m.role == Role::User)
.map(|m| {
let t = m.text_content();
if t.len() > 80 {
format!("{}…", crate::safe_truncate(&t, 80))
} else {
t
}
})
.unwrap_or_else(|| "Untitled session".to_string());
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs_f64();
let (source, routing_key) = match &config.origin_chat {
Some((platform, chat_id)) => (platform.clone(), Some(chat_id.clone())),
None => ("cli".to_string(), None),
};
let record = edgecrab_state::SessionRecord {
id: session_id.clone(),
source,
user_id: routing_key,
model: Some(config.model.clone()),
system_prompt: session.cached_system_prompt.clone(),
parent_session_id: None,
started_at: now,
ended_at: Some(now),
end_reason: if interrupted {
Some("interrupted".to_string())
} else {
None
},
message_count: session.messages.len() as i64,
tool_call_count: session.session_tool_call_count as i64,
input_tokens: session.session_input_tokens as i64,
output_tokens: session.session_output_tokens as i64,
cache_read_tokens: session.session_cache_read_tokens as i64,
cache_write_tokens: session.session_cache_write_tokens as i64,
reasoning_tokens: session.session_reasoning_tokens as i64,
estimated_cost_usd: cost_result.amount_usd,
title: Some(title),
};
if let Err(e) = db.save_session(&record) {
tracing::warn!(error = %e, "failed to save session to state DB");
}
if let Err(e) = db.replace_messages(&session_id, &session.messages, now) {
tracing::warn!(error = %e, "failed to save messages to state DB");
}
if session.user_turn_count == 1 && !final_response.is_empty() {
let user_snippet = user_message.chars().take(500).collect::<String>();
let asst_snippet = final_response.chars().take(500).collect::<String>();
let db_clone = db.clone();
let sid_clone = session_id.clone();
let prov_clone = effective_provider.clone();
tokio::spawn(async move {
auto_title_session(db_clone, sid_clone, user_snippet, asst_snippet, prov_clone)
.await;
});
}
}
let completed = !interrupted && !final_response.trim().is_empty();
if config.save_trajectories {
let trajectory_dir = edgecrab_home().join("trajectories");
let trajectory_path = trajectory_dir.join(if completed {
"trajectory_samples.jsonl"
} else {
"failed_trajectories.jsonl"
});
if let Err(e) = std::fs::create_dir_all(&trajectory_dir) {
tracing::warn!(error = %e, path = %trajectory_dir.display(), "failed to create trajectory directory");
} else {
let trajectory = build_trajectory(
&session_id,
&config.model,
&session.messages,
session.api_call_count,
cost.total_cost,
completed,
turn_started_at.elapsed().as_secs_f64(),
);
if let Err(e) = save_trajectory(&trajectory_path, &trajectory) {
tracing::warn!(error = %e, path = %trajectory_path.display(), "failed to save trajectory");
}
}
}
let messages = session.messages.clone();
let api_calls = session.api_call_count;
let model = config.model.clone();
Ok(ConversationResult {
final_response,
messages,
session_id,
api_calls,
interrupted,
budget_exhausted,
model,
usage,
cost,
tool_errors: tool_errors_acc,
})
}
}
pub fn build_chat_messages(
system_prompt: Option<&str>,
messages: &[Message],
cache_config: Option<&CachePromptConfig>,
) -> Vec<edgequake_llm::ChatMessage> {
let mut out = Vec::with_capacity(messages.len() + 1);
if let Some(sys) = system_prompt {
out.push(edgequake_llm::ChatMessage::system(sys));
}
for m in messages {
let text = m.text_content();
match m.role {
Role::System => out.push(edgequake_llm::ChatMessage::system(&text)),
Role::User => out.push(edgequake_llm::ChatMessage::user(&text)),
Role::Assistant => {
if let Some(ref tool_calls) = m.tool_calls {
if !tool_calls.is_empty() {
let llm_calls: Vec<edgequake_llm::ToolCall> =
tool_calls.iter().map(|tc| tc.to_llm()).collect();
out.push(edgequake_llm::ChatMessage::assistant_with_tools(
&text, llm_calls,
));
continue;
}
}
out.push(edgequake_llm::ChatMessage::assistant(&text));
}
Role::Tool => {
let tool_call_id = m.tool_call_id.as_deref().unwrap_or("unknown");
let mut chat_msg = edgequake_llm::ChatMessage::tool_result(tool_call_id, &text);
chat_msg.name = m.name.clone();
out.push(chat_msg);
}
}
}
if let Some(cfg) = cache_config {
apply_cache_control(&mut out, cfg);
}
out
}
#[inline]
fn parse_tool_error_response(result: &str) -> Option<ToolErrorResponse> {
let parsed = serde_json::from_str::<ToolErrorResponse>(result).ok()?;
(parsed.response_type == "tool_error").then_some(parsed)
}
fn tool_attempt_fingerprint(name: &str, args_json: &str) -> String {
let normalized_args = serde_json::from_str::<serde_json::Value>(args_json)
.ok()
.and_then(|value| serde_json::to_string(&value).ok())
.unwrap_or_else(|| args_json.trim().to_string());
format!("{name}:{normalized_args}")
}
fn suppressed_retry_response(
name: &str,
args_json: &str,
prior: &ToolErrorResponse,
) -> ToolErrorResponse {
let mut error = ToolError::capability_denied(
name,
"suppressed_capability_retry",
format!(
"EdgeCrab already blocked an equivalent `{name}` call earlier in this conversation because `{}` is unresolved. Repeating the same call would be flaky. Change the approach or complete the required user action first.",
prior.code
),
)
.with_suppression_key(tool_attempt_fingerprint(name, args_json));
if let Some(suggested_tool) = &prior.suggested_tool {
error = error.with_suggested_tool(suggested_tool.clone());
}
if let Some(suggested_action) = &prior.suggested_action {
error = error.with_suggested_action(suggested_action.clone());
}
error.to_llm_payload()
}
#[inline]
fn is_tool_error(result: &str) -> bool {
parse_tool_error_response(result).is_some() || result.starts_with("Tool error:")
}
fn emit_tool_done(
tx: Option<&tokio::sync::mpsc::UnboundedSender<crate::StreamEvent>>,
name: &str,
args_json: &str,
duration_ms: u64,
is_error: bool,
) {
if let Some(tx) = tx {
let _ = tx.send(crate::StreamEvent::ToolDone {
name: name.to_string(),
args_json: args_json.to_string(),
duration_ms,
is_error,
});
}
}
#[derive(Default)]
struct PartialToolCall {
id: Option<String>,
function_name: Option<String>,
arguments: String,
thought_signature: Option<String>,
}
async fn api_call_streaming(
provider: &Arc<dyn LLMProvider>,
messages: &[edgequake_llm::ChatMessage],
tool_defs: &[edgequake_llm::ToolDefinition],
event_tx: &tokio::sync::mpsc::UnboundedSender<crate::StreamEvent>,
any_tokens_sent: &std::sync::atomic::AtomicBool,
) -> edgequake_llm::Result<edgequake_llm::LLMResponse> {
let mut stream = provider
.chat_with_tools_stream(messages, tool_defs, None, None)
.await?;
let mut content = String::new();
let mut thinking = String::new();
let mut thinking_tokens = 0usize;
let mut final_usage: Option<StreamUsage> = None;
let mut finish_reason: Option<String> = None;
let mut tool_calls: BTreeMap<usize, PartialToolCall> = BTreeMap::new();
while let Some(chunk) = stream.next().await {
match chunk? {
StreamChunk::Content(delta) => {
if !delta.is_empty() {
any_tokens_sent.store(true, std::sync::atomic::Ordering::Relaxed);
content.push_str(&delta);
let _ = event_tx.send(crate::StreamEvent::Token(delta));
}
}
StreamChunk::ThinkingContent {
text, tokens_used, ..
} => {
if !text.is_empty() {
any_tokens_sent.store(true, std::sync::atomic::Ordering::Relaxed);
thinking.push_str(&text);
let _ = event_tx.send(crate::StreamEvent::Reasoning(text));
}
if let Some(tokens) = tokens_used {
thinking_tokens += tokens;
}
}
StreamChunk::ToolCallDelta {
index,
id,
function_name,
function_arguments,
thought_signature,
} => {
any_tokens_sent.store(true, std::sync::atomic::Ordering::Relaxed);
let entry = tool_calls.entry(index).or_default();
if let Some(id) = id {
entry.id = Some(id);
}
if let Some(name) = function_name {
entry.function_name = Some(name);
}
if let Some(args) = function_arguments {
entry.arguments.push_str(&args);
}
if thought_signature.is_some() {
entry.thought_signature = thought_signature;
}
}
StreamChunk::Finished { reason, usage, .. } => {
finish_reason = Some(reason);
if usage.is_some() {
final_usage = usage;
}
}
}
}
let mut response = edgequake_llm::LLMResponse::new(content, provider.model().to_string());
if let Some(reason) = finish_reason {
response.finish_reason = Some(reason);
}
if !thinking.is_empty() {
response.thinking_content = Some(thinking);
}
response.tool_calls = tool_calls
.into_iter()
.map(|(index, partial)| edgequake_llm::ToolCall {
id: partial.id.unwrap_or_else(|| format!("stream_call_{index}")),
call_type: "function".to_string(),
function: edgequake_llm::FunctionCall {
name: partial
.function_name
.unwrap_or_else(|| "unknown_tool".to_string()),
arguments: if partial.arguments.trim().is_empty() {
"{}".to_string()
} else {
partial.arguments
},
},
thought_signature: partial.thought_signature,
})
.collect();
if let Some(usage) = final_usage {
response = response.with_usage(usage.prompt_tokens, usage.completion_tokens);
if let Some(cache_hit_tokens) = usage.cache_hit_tokens {
response = response.with_cache_hit_tokens(cache_hit_tokens);
}
if let Some(authoritative_thinking_tokens) = usage.thinking_tokens {
response = response.with_thinking_tokens(authoritative_thinking_tokens);
} else if thinking_tokens > 0 {
response = response.with_thinking_tokens(thinking_tokens);
}
} else {
let estimated_prompt_tokens = estimate_stream_prompt_tokens(messages, tool_defs);
let estimated_completion_tokens = estimate_stream_completion_tokens(
&response.content,
response.thinking_content.as_deref(),
&response.tool_calls,
);
if estimated_prompt_tokens > 0 || estimated_completion_tokens > 0 {
response = response.with_usage(estimated_prompt_tokens, estimated_completion_tokens);
}
if thinking_tokens > 0 {
response = response.with_thinking_tokens(thinking_tokens);
}
}
Ok(response)
}
fn estimate_stream_prompt_tokens(
messages: &[edgequake_llm::ChatMessage],
tool_defs: &[edgequake_llm::ToolDefinition],
) -> usize {
estimate_tokens_from_json(&(messages, tool_defs))
}
fn estimate_stream_completion_tokens(
content: &str,
thinking: Option<&str>,
tool_calls: &[edgequake_llm::ToolCall],
) -> usize {
estimate_tokens_from_json(&(content, thinking, tool_calls))
}
fn estimate_tokens_from_json<T: serde::Serialize>(value: &T) -> usize {
let serialized = match serde_json::to_string(value) {
Ok(serialized) => serialized,
Err(_) => return 0,
};
estimate_tokens_from_text(&serialized)
}
fn estimate_tokens_from_text(text: &str) -> usize {
let trimmed = text.trim();
if trimmed.is_empty() {
return 0;
}
trimmed.chars().count().div_ceil(4)
}
async fn api_call_with_retry(
provider: &Arc<dyn LLMProvider>,
messages: &[edgequake_llm::ChatMessage],
tool_defs: &[edgequake_llm::ToolDefinition],
max_retries: u32,
cancel: &CancellationToken,
event_tx: Option<&tokio::sync::mpsc::UnboundedSender<crate::StreamEvent>>,
use_native_streaming: bool,
) -> Result<edgequake_llm::LLMResponse, AgentError> {
let mut last_err = None;
let retry_budget = max_retries;
for attempt in 0..=retry_budget {
if attempt > 0 {
let delay = BASE_BACKOFF * 2u32.saturating_pow(attempt - 1);
tokio::select! {
biased;
_ = cancel.cancelled() => {
tracing::debug!(attempt, "api_call_with_retry: cancelled during backoff sleep");
return Err(AgentError::Interrupted);
}
_ = tokio::time::sleep(delay) => {
tracing::debug!(attempt, "retrying API call after backoff");
}
}
}
let tokens_sent = std::sync::atomic::AtomicBool::new(false);
if let Some(tx) = event_tx {
let ctx_json = serde_json::json!({
"event": "llm:pre",
"model": provider.model(),
"attempt": attempt,
})
.to_string();
let _ = tx.send(crate::StreamEvent::HookEvent {
event: "llm:pre".to_string(),
context_json: ctx_json,
});
}
let call_fut = async {
if use_native_streaming {
let tx = event_tx.expect("native streaming requires event channel");
api_call_streaming(provider, messages, tool_defs, tx, &tokens_sent).await
} else if tool_defs.is_empty() {
provider.chat(messages, None).await
} else {
provider
.chat_with_tools(messages, tool_defs, None, None)
.await
}
};
let result = tokio::select! {
biased;
_ = cancel.cancelled() => {
tracing::debug!(attempt, "api_call_with_retry: cancelled during API call");
return Err(AgentError::Interrupted);
}
r = call_fut => r,
};
match result {
Ok(response) => {
if let Some(tx) = event_tx {
let ctx_json = serde_json::json!({
"event": "llm:post",
"model": provider.model(),
"prompt_tokens": response.prompt_tokens,
"completion_tokens": response.completion_tokens,
})
.to_string();
let _ = tx.send(crate::StreamEvent::HookEvent {
event: "llm:post".to_string(),
context_json: ctx_json,
});
}
return Ok(response);
}
Err(e) => {
tracing::warn!(attempt, error = %e, "API call failed");
if use_native_streaming {
let mid_stream = tokens_sent.load(std::sync::atomic::Ordering::Relaxed);
let is_prestream_retryable = !mid_stream
&& matches!(
e,
edgequake_llm::LlmError::RateLimited(_)
| edgequake_llm::LlmError::NetworkError(_)
| edgequake_llm::LlmError::Timeout
| edgequake_llm::LlmError::AuthError(_)
);
if !is_prestream_retryable {
return Err(AgentError::Llm(format!(
"API call failed after {} retries: {}",
attempt, e
)));
}
}
last_err = Some(e);
}
}
}
Err(AgentError::Llm(format!(
"API call failed after {} retries: {}",
retry_budget,
last_err.map_or_else(|| "unknown error".to_string(), |e| e.to_string())
)))
}
fn get_budget_warning(api_call_count: u32, max_iterations: u32) -> Option<String> {
if max_iterations == 0 {
return None;
}
let progress = api_call_count as f64 / max_iterations as f64;
if progress >= 0.9 {
Some(format!(
"[URGENT: {}% of iteration budget used ({}/{}). You MUST provide a final response NOW — do not make further tool calls.]",
(progress * 100.0) as u32,
api_call_count,
max_iterations
))
} else if progress >= 0.7 {
Some(format!(
"[BUDGET: {}% of iteration budget used ({}/{}). Start wrapping up — avoid multi-step tool chains.]",
(progress * 100.0) as u32,
api_call_count,
max_iterations
))
} else {
None
}
}
fn inject_budget_warning(messages: &mut Vec<Message>, warning: &str) {
if let Some(msg) = messages.iter_mut().rev().find(|m| m.role == Role::Tool) {
let current = msg.text_content();
let new_content = if let Ok(mut v) = serde_json::from_str::<serde_json::Value>(¤t) {
if let Some(obj) = v.as_object_mut() {
obj.insert(
"_budget_warning".to_string(),
serde_json::Value::String(warning.to_string()),
);
serde_json::to_string(&v).unwrap_or_else(|_| format!("{}\n\n{}", current, warning))
} else {
format!("{}\n\n{}", current, warning)
}
} else {
format!("{}\n\n{}", current, warning)
};
msg.content = Some(Content::Text(new_content));
} else {
tracing::debug!("no tool messages found, injecting budget warning as user message");
messages.push(Message::user(warning));
}
}
fn build_trajectory(
session_id: &str,
model: &str,
messages: &[Message],
api_calls: u32,
total_cost: f64,
completed: bool,
duration_seconds: f64,
) -> Trajectory {
let normalized_messages = normalize_messages_for_trajectory(messages);
let total_tokens = normalized_messages
.iter()
.map(|message| message.text_content().len() as u64 / 4)
.sum();
Trajectory {
session_id: session_id.to_string(),
model: model.to_string(),
timestamp: chrono::Utc::now().to_rfc3339(),
messages: normalized_messages,
metadata: TrajectoryMetadata {
task_id: None,
total_tokens,
total_cost,
api_calls,
tools_used: collect_used_tools(messages),
completed,
duration_seconds,
},
}
}
fn normalize_messages_for_trajectory(messages: &[Message]) -> Vec<Message> {
messages
.iter()
.cloned()
.map(|mut message| {
if let Some(Content::Text(text)) = &message.content {
message.content = Some(Content::Text(convert_scratchpad_to_think(text)));
}
if let Some(reasoning) = &message.reasoning {
message.reasoning = Some(convert_scratchpad_to_think(reasoning));
}
message
})
.collect()
}
fn collect_used_tools(messages: &[Message]) -> Vec<String> {
let mut tools = Vec::new();
for message in messages.iter().filter(|message| message.role == Role::Tool) {
if let Some(name) = &message.name {
if !tools.iter().any(|existing| existing == name) {
tools.push(name.clone());
}
}
}
tools
}
#[inline]
fn extract_tool_error_text(result: &str) -> String {
if let Some(payload) = parse_tool_error_response(result) {
return payload.error;
}
result.to_string()
}
fn remember_tool_suppression(
suppressions: &Arc<Mutex<HashMap<String, ToolErrorResponse>>>,
name: &str,
args_json: &str,
result: &str,
) {
let Some(payload) = parse_tool_error_response(result) else {
return;
};
if !payload.suppress_retry {
return;
}
let mut guard = suppressions
.lock()
.expect("capability suppression cache lock poisoned");
guard.insert(tool_attempt_fingerprint(name, args_json), payload.clone());
if let Some(extra_key) = payload.suppression_key.clone() {
guard.insert(extra_key, payload);
}
}
async fn process_response(
response: &edgequake_llm::LLMResponse,
session: &mut SessionState,
dctx: &DispatchContext<'_>,
tool_errors: &mut Vec<edgecrab_types::ToolErrorRecord>,
) -> Result<LoopAction, AgentError> {
if response.has_tool_calls() {
let our_tool_calls: Vec<edgecrab_types::ToolCall> = response
.tool_calls
.iter()
.map(edgecrab_types::ToolCall::from_llm)
.collect();
let assistant_text = response.content.clone();
let mut assistant_msg = Message::assistant_with_tool_calls(&assistant_text, our_tool_calls);
if let Some(ref thinking) = response.thinking_content {
assistant_msg.reasoning = Some(thinking.clone());
}
session.messages.push(assistant_msg);
session.session_tool_call_count += response.tool_calls.len() as u32;
let mut parallel_tasks = tokio::task::JoinSet::new();
let mut sequential_calls = Vec::new();
let mut parallel_submitted: Vec<(String, String)> = Vec::new();
for tc in &response.tool_calls {
let is_parallel = dctx
.registry
.as_ref()
.map(|r| r.is_parallel_safe(&tc.function.name))
.unwrap_or(false);
if let Some(tx) = dctx.event_tx {
let _ = tx.send(crate::StreamEvent::ToolExec {
name: tc.function.name.clone(),
args_json: tc.function.arguments.clone(),
});
}
if is_parallel {
parallel_submitted.push((tc.id.clone(), tc.function.name.clone()));
let tc_id = tc.id.clone();
let tc_name = tc.function.name.clone();
let tc_args = tc.function.arguments.clone();
let reg = dctx.registry.cloned();
let cancel_token = dctx.cancel.clone();
let state = dctx.state_db.clone();
let plat = dctx.platform;
let proc_table = dctx.process_table.clone();
let prov = dctx.provider.clone();
let gateway_sender = dctx.gateway_sender.clone();
let sar = dctx.sub_agent_runner.clone();
let clarify = dctx.clarify_tx.clone();
let approval = dctx.approval_tx.clone();
let args_for_done = tc.function.arguments.clone();
let origin = dctx.origin_chat.clone();
let app_cfg_ref = dctx.app_config_ref.clone();
let conv_sess_id = dctx.conversation_session_id.clone();
let todo_store_clone = dctx.todo_store.clone();
let capability_suppressions = dctx.capability_suppressions.clone();
let dispatch_cwd = dctx.cwd.clone();
parallel_tasks.spawn(async move {
let started = std::time::Instant::now();
let inner = DispatchContext {
cwd: dispatch_cwd,
registry: reg.as_ref(),
cancel: &cancel_token,
state_db: &state,
platform: plat,
process_table: &proc_table,
provider: prov,
gateway_sender,
sub_agent_runner: sar,
event_tx: None, clarify_tx: clarify,
approval_tx: approval,
origin_chat: origin,
app_config_ref: app_cfg_ref,
conversation_session_id: conv_sess_id,
todo_store: todo_store_clone,
capability_suppressions,
};
let result = dispatch_single_tool(&tc_name, &tc_args, &inner).await;
let duration_ms = started.elapsed().as_millis() as u64;
(tc_id, tc_name, args_for_done, result, duration_ms)
});
} else {
sequential_calls.push(tc);
}
}
let mut received_parallel_ids: std::collections::HashSet<String> =
std::collections::HashSet::new();
while let Some(join_result) = parallel_tasks.join_next().await {
match join_result {
Ok((tc_id, tc_name, args_json, tool_result, duration_ms)) => {
let is_error = is_tool_error(&tool_result);
emit_tool_done(dctx.event_tx, &tc_name, &args_json, duration_ms, is_error);
if is_error {
remember_tool_suppression(
&dctx.capability_suppressions,
&tc_name,
&args_json,
&tool_result,
);
tool_errors.push(edgecrab_types::ToolErrorRecord {
turn: session.api_call_count,
tool_name: tc_name.clone(),
arguments: args_json.clone(),
error: extract_tool_error_text(&tool_result),
tool_result: tool_result.clone(),
});
}
received_parallel_ids.insert(tc_id.clone());
session
.messages
.push(Message::tool_result(&tc_id, &tc_name, &tool_result));
}
Err(e) => {
tracing::error!(error = %e, "parallel tool task panicked");
}
}
}
for (tc_id, tc_name) in ¶llel_submitted {
if !received_parallel_ids.contains(tc_id) {
tracing::warn!(
tool_call_id = %tc_id,
tool_name = %tc_name,
"injecting error result for panicked parallel tool task"
);
session.messages.push(Message::tool_result(
tc_id,
tc_name,
&format!("Tool error: task panicked — internal error executing '{tc_name}'"),
));
}
}
for tc in sequential_calls {
let started = std::time::Instant::now();
let tool_result =
dispatch_single_tool(&tc.function.name, &tc.function.arguments, dctx).await;
let duration_ms = started.elapsed().as_millis() as u64;
let is_error = is_tool_error(&tool_result);
emit_tool_done(
dctx.event_tx,
&tc.function.name,
&tc.function.arguments,
duration_ms,
is_error,
);
if is_error {
remember_tool_suppression(
&dctx.capability_suppressions,
&tc.function.name,
&tc.function.arguments,
&tool_result,
);
tool_errors.push(edgecrab_types::ToolErrorRecord {
turn: session.api_call_count,
tool_name: tc.function.name.clone(),
arguments: tc.function.arguments.clone(),
error: extract_tool_error_text(&tool_result),
tool_result: tool_result.clone(),
});
}
session.messages.push(Message::tool_result(
&tc.id,
&tc.function.name,
&tool_result,
));
}
return Ok(LoopAction::Continue);
}
let text = response.content.clone();
let mut msg = Message::assistant(&text);
if let Some(ref thinking) = response.thinking_content {
msg.reasoning = Some(thinking.clone());
}
session.messages.push(msg);
Ok(LoopAction::Done(text))
}
async fn dispatch_single_tool(name: &str, args_json: &str, dctx: &DispatchContext<'_>) -> String {
let Some(reg) = dctx.registry else {
return format!(
"Tool '{}' execution is not yet wired (no ToolRegistry provided).",
name
);
};
let attempt_key = tool_attempt_fingerprint(name, args_json);
if let Some(prior) = dctx
.capability_suppressions
.lock()
.expect("capability suppression cache lock poisoned")
.get(&attempt_key)
.cloned()
{
return serde_json::to_string(&suppressed_retry_response(name, args_json, &prior))
.expect("suppressed retry payload serializes");
}
if let Some(tx) = dctx.event_tx {
let ctx_json = serde_json::json!({
"event": "tool:pre",
"tool_name": name,
"args_json": args_json,
"session_id": &dctx.conversation_session_id,
})
.to_string();
let _ = tx.send(crate::StreamEvent::HookEvent {
event: "tool:pre".to_string(),
context_json: ctx_json,
});
}
let ctx = build_tool_context(
&dctx.cwd,
dctx.app_config_ref.clone(),
dctx.cancel,
dctx.state_db,
dctx.platform,
dctx.process_table,
dctx.provider.clone(),
dctx.registry.cloned(), dctx.sub_agent_runner.clone(),
dctx.clarify_tx.clone(),
dctx.approval_tx.clone(),
dctx.gateway_sender.clone(),
dctx.origin_chat.clone(),
&dctx.conversation_session_id,
dctx.todo_store.clone(),
);
let args: serde_json::Value = match serde_json::from_str(args_json) {
Ok(v) => v,
Err(e) => {
tracing::warn!(tool_name = %name, error = %e, args_json = %args_json, "malformed tool arguments JSON");
return ToolError::InvalidArgs {
tool: name.to_string(),
message: format!("invalid JSON arguments: {e}"),
}
.to_llm_response();
}
};
let result = match reg.dispatch(name, args, &ctx).await {
Ok(output) => output,
Err(e) => e.to_llm_response(),
};
if let Some(tx) = dctx.event_tx {
let is_error = is_tool_error(&result);
let ctx_json = serde_json::json!({
"event": "tool:post",
"tool_name": name,
"session_id": &dctx.conversation_session_id,
"is_error": is_error,
})
.to_string();
let _ = tx.send(crate::StreamEvent::HookEvent {
event: "tool:post".to_string(),
context_json: ctx_json,
});
}
result
}
async fn auto_title_session(
db: Arc<edgecrab_state::SessionDb>,
session_id: String,
user_snippet: String,
assistant_snippet: String,
provider: Arc<dyn LLMProvider>,
) {
match db.get_session(&session_id) {
Ok(Some(rec)) => {
if let Some(ref existing) = rec.title {
if !existing.is_empty() && existing.len() < 80 && !existing.ends_with('…') {
tracing::debug!("session already has a title, skipping auto-title");
return;
}
}
}
_ => return,
}
let prompt = format!(
"Generate a short, descriptive title (3-7 words) for a conversation that starts with:\n\
User: {user_snippet}\n\nAssistant: {assistant_snippet}\n\n\
Return ONLY the title. No quotes, no punctuation at the end, no prefixes."
);
let messages = vec![
edgequake_llm::ChatMessage::system(
"You generate ultra-short session titles. Respond with ONLY the title, nothing else.",
),
edgequake_llm::ChatMessage::user(&prompt),
];
match provider.chat(&messages, None).await {
Ok(resp) => {
let mut title = resp.content.trim().to_string();
title = title.trim_matches(|c| c == '"' || c == '\'').to_string();
if title.to_lowercase().starts_with("title:") {
title = title[6..].trim().to_string();
}
if title.len() > 80 {
title = format!("{}…", crate::safe_truncate(&title, 77));
}
if !title.is_empty() {
if let Err(e) = db.update_session_title(&session_id, &title) {
tracing::debug!(error = %e, "auto-title DB update failed");
} else {
tracing::debug!(title, "auto-generated session title");
}
}
}
Err(e) => tracing::debug!(error = %e, "auto-title generation failed"),
}
}
struct BackgroundReflectionCtx {
messages: Vec<Message>,
system_prompt: Option<String>,
tool_defs: Vec<edgequake_llm::ToolDefinition>,
cwd: std::path::PathBuf,
registry: Option<Arc<ToolRegistry>>,
cancel: CancellationToken,
state_db: Option<Arc<edgecrab_state::SessionDb>>,
platform: edgecrab_types::Platform,
process_table: Arc<edgecrab_tools::ProcessTable>,
provider: Arc<dyn edgequake_llm::LLMProvider>,
gateway_sender: Option<Arc<dyn edgecrab_tools::registry::GatewaySender>>,
sub_agent_runner: Option<Arc<dyn edgecrab_tools::SubAgentRunner>>,
app_config_ref: AppConfigRef,
conversation_session_id: String,
origin_chat: Option<(String, String)>,
todo_store: Option<Arc<edgecrab_tools::TodoStore>>,
}
async fn run_learning_reflection_bg(ctx: BackgroundReflectionCtx) {
let registry_arc = ctx.registry.as_ref();
let state_db_ref = &ctx.state_db;
let process_table_ref = &ctx.process_table;
let dctx = DispatchContext {
cwd: ctx.cwd.clone(),
registry: registry_arc,
cancel: &ctx.cancel,
state_db: state_db_ref,
platform: ctx.platform,
process_table: process_table_ref,
provider: Some(Arc::clone(&ctx.provider)),
gateway_sender: ctx.gateway_sender.clone(),
sub_agent_runner: ctx.sub_agent_runner.clone(),
event_tx: None, clarify_tx: None, approval_tx: None, origin_chat: ctx.origin_chat.clone(),
app_config_ref: ctx.app_config_ref.clone(),
conversation_session_id: ctx.conversation_session_id.clone(),
todo_store: ctx.todo_store.clone(),
capability_suppressions: Arc::new(Mutex::new(HashMap::new())),
};
let mut session = SessionState {
messages: ctx.messages,
cached_system_prompt: ctx.system_prompt,
..Default::default()
};
run_learning_reflection(&mut session, &ctx.tool_defs, &ctx.provider, &dctx).await;
}
async fn run_learning_reflection(
session: &mut SessionState,
tool_defs: &[edgequake_llm::ToolDefinition],
provider: &Arc<dyn LLMProvider>,
dctx: &DispatchContext<'_>,
) {
const REFLECTION_PROMPT: &str = "\
[system: learning checkpoint] This session used multiple tool calls. \
Please reflect briefly (1-2 sentences of thinking, not shown to the user): \
Did you discover a reusable workflow, debugging technique, or non-trivial \
pattern worth saving? If yes, call skill_manage(action='create', name='...', \
content='---\\nname: ...\\ndescription: ...\\n---\\n# Steps\\n...') to save it. \
Did you learn something important about the user, their project, or environment \
that should persist? If yes, call memory_write to record it. \
If nothing is worth saving, respond with exactly 'reflection: nothing to save' \
and stop — do NOT call any tools.";
session.messages.push(Message::user(REFLECTION_PROMPT));
let chat_messages = build_chat_messages(
session.cached_system_prompt.as_deref(),
&session.messages,
None, );
let response = match provider
.chat_with_tools(&chat_messages, tool_defs, None, None)
.await
{
Ok(r) => r,
Err(e) => {
tracing::debug!(error = %e, "learning reflection API call failed (non-fatal)");
session.messages.pop();
return;
}
};
let mut _reflection_tool_errors: Vec<edgecrab_types::ToolErrorRecord> = Vec::new();
if let Err(e) = process_response(&response, session, dctx, &mut _reflection_tool_errors).await {
tracing::debug!(error = %e, "learning reflection tool dispatch failed (non-fatal)");
}
}
fn sanitize_orphaned_tool_results(messages: &mut Vec<Message>) {
use std::collections::HashSet;
let mut valid_ids: HashSet<String> = HashSet::new();
for msg in messages.iter() {
if msg.role == Role::Assistant {
if let Some(ref calls) = msg.tool_calls {
for tc in calls {
valid_ids.insert(tc.id.clone());
}
}
}
}
let before = messages.len();
messages.retain(|msg| {
if msg.role == Role::Tool {
msg.tool_call_id
.as_ref()
.is_some_and(|id| valid_ids.contains(id))
} else {
true
}
});
let removed = before - messages.len();
if removed > 0 {
tracing::info!(removed, "sanitized orphaned tool result messages");
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::AgentBuilder;
use async_trait::async_trait;
use edgecrab_tools::{ProcessTable, ToolRegistry};
use edgequake_llm::traits::StreamUsage;
use edgequake_llm::{ChatMessage, CompletionOptions, ToolChoice, ToolDefinition};
use serde_json::json;
use tempfile::TempDir;
#[derive(Clone)]
struct StreamingUsageProvider {
chunks: Vec<StreamChunk>,
}
#[async_trait]
impl LLMProvider for StreamingUsageProvider {
fn name(&self) -> &str {
"streaming-usage-test"
}
fn model(&self) -> &str {
"streaming-usage-test-model"
}
fn max_context_length(&self) -> usize {
8192
}
async fn complete(
&self,
prompt: &str,
) -> edgequake_llm::Result<edgequake_llm::LLMResponse> {
Ok(edgequake_llm::LLMResponse::new(prompt, self.model()))
}
async fn complete_with_options(
&self,
prompt: &str,
_options: &CompletionOptions,
) -> edgequake_llm::Result<edgequake_llm::LLMResponse> {
self.complete(prompt).await
}
async fn chat(
&self,
messages: &[ChatMessage],
options: Option<&CompletionOptions>,
) -> edgequake_llm::Result<edgequake_llm::LLMResponse> {
self.chat_with_tools(messages, &[], None, options).await
}
async fn chat_with_tools(
&self,
_messages: &[ChatMessage],
_tools: &[ToolDefinition],
_tool_choice: Option<ToolChoice>,
_options: Option<&CompletionOptions>,
) -> edgequake_llm::Result<edgequake_llm::LLMResponse> {
Ok(edgequake_llm::LLMResponse::new("non-stream", self.model()))
}
async fn chat_with_tools_stream(
&self,
_messages: &[ChatMessage],
_tools: &[ToolDefinition],
_tool_choice: Option<ToolChoice>,
_options: Option<&CompletionOptions>,
) -> edgequake_llm::Result<
futures::stream::BoxStream<'static, edgequake_llm::Result<StreamChunk>>,
> {
use futures::StreamExt;
Ok(futures::stream::iter(self.chunks.clone().into_iter().map(Ok)).boxed())
}
fn supports_tool_streaming(&self) -> bool {
true
}
}
#[tokio::test]
async fn api_call_streaming_preserves_authoritative_usage() {
let provider: Arc<dyn LLMProvider> = Arc::new(StreamingUsageProvider {
chunks: vec![
StreamChunk::Content("streamed answer".to_string()),
StreamChunk::Finished {
reason: "stop".to_string(),
ttft_ms: None,
usage: Some(
StreamUsage::new(11, 7)
.with_cache_hit_tokens(2)
.with_thinking_tokens(5),
),
},
],
});
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let tokens_sent = std::sync::atomic::AtomicBool::new(false);
let response = api_call_streaming(
&provider,
&[ChatMessage::user("hello")],
&[],
&tx,
&tokens_sent,
)
.await
.expect("stream call");
assert_eq!(response.prompt_tokens, 11);
assert_eq!(response.completion_tokens, 7);
assert_eq!(response.total_tokens, 18);
assert_eq!(response.cache_hit_tokens, Some(2));
assert_eq!(response.thinking_tokens, Some(5));
assert_eq!(response.finish_reason.as_deref(), Some("stop"));
}
#[tokio::test]
async fn api_call_streaming_estimates_usage_when_provider_omits_it() {
let provider: Arc<dyn LLMProvider> = Arc::new(StreamingUsageProvider {
chunks: vec![
StreamChunk::ThinkingContent {
text: "reasoning trace".to_string(),
tokens_used: Some(3),
budget_total: None,
},
StreamChunk::Content("streamed answer".to_string()),
StreamChunk::Finished {
reason: "stop".to_string(),
ttft_ms: None,
usage: None,
},
],
});
let tool_defs = vec![ToolDefinition::function(
"echo",
"Echo input",
json!({
"type": "object",
"properties": {
"text": {"type": "string"}
},
"required": ["text"]
}),
)];
let (tx, _rx) = tokio::sync::mpsc::unbounded_channel();
let tokens_sent = std::sync::atomic::AtomicBool::new(false);
let response = api_call_streaming(
&provider,
&[
ChatMessage::system("system"),
ChatMessage::user("hello world"),
],
&tool_defs,
&tx,
&tokens_sent,
)
.await
.expect("stream call");
assert!(
response.prompt_tokens > 0,
"prompt tokens should be estimated"
);
assert!(
response.completion_tokens > 0,
"completion tokens should be estimated"
);
assert_eq!(response.thinking_tokens, Some(3));
assert_eq!(response.finish_reason.as_deref(), Some("stop"));
}
#[tokio::test]
async fn execute_loop_basic() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.build()
.expect("build");
let result = agent
.execute_loop("hello", Some("Be helpful."), None, None, None)
.await
.expect("loop");
assert!(!result.final_response.is_empty());
assert_eq!(result.api_calls, 1);
assert!(!result.interrupted);
}
#[tokio::test]
async fn execute_loop_with_history() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.build()
.expect("build");
let history = vec![
Message::user("previous question"),
Message::assistant("previous answer"),
];
let result = agent
.execute_loop("follow-up", None, Some(history), None, None)
.await
.expect("loop");
assert_eq!(result.messages.len(), 4);
}
#[tokio::test]
async fn execute_loop_uses_cwd_override_for_context_discovery() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.build()
.expect("build");
let workspace = TempDir::new().expect("workspace");
std::fs::write(
workspace.path().join("AGENTS.md"),
"# Workspace Rules\n\nUse the override workspace.",
)
.expect("write AGENTS.md");
agent
.execute_loop("hello", None, None, None, Some(workspace.path()))
.await
.expect("loop");
let session = agent.session.read().await;
let prompt = session
.cached_system_prompt
.as_deref()
.expect("cached system prompt");
assert!(prompt.contains("Use the override workspace."));
}
#[test]
fn build_trajectory_normalizes_reasoning_and_collects_tools() {
let messages = vec![
Message::user("hello"),
Message::assistant("<REASONING_SCRATCHPAD>plan</REASONING_SCRATCHPAD>done"),
Message::tool_result("call_1", "read_file", "contents"),
Message::tool_result("call_2", "read_file", "more contents"),
];
let trajectory =
build_trajectory("session-1", "provider/model", &messages, 2, 0.25, true, 1.5);
assert_eq!(trajectory.session_id, "session-1");
assert_eq!(trajectory.metadata.api_calls, 2);
assert_eq!(
trajectory.metadata.tools_used,
vec!["read_file".to_string()]
);
assert!(
trajectory.messages[1]
.text_content()
.contains("<think>plan</think>")
);
}
#[tokio::test]
async fn execute_loop_resets_preexisting_interrupt() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.build()
.expect("build");
agent.interrupt();
let result = agent
.execute_loop("hello", None, None, None, None)
.await
.expect("loop");
assert!(
!result.interrupted,
"pre-loop interrupt must be reset, not permanently sticky"
);
assert!(!result.final_response.is_empty());
assert!(!agent.is_cancelled());
}
#[tokio::test]
async fn execute_loop_budget_exhaust() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(1)
.build()
.expect("build");
let result = agent
.execute_loop("hello", None, None, None, None)
.await
.expect("loop");
assert_eq!(result.api_calls, 1);
}
#[test]
fn build_chat_messages_prepends_system() {
let messages = vec![Message::user("hi")];
let chat_msgs = build_chat_messages(Some("system prompt"), &messages, None);
assert_eq!(chat_msgs.len(), 2);
}
#[test]
fn build_chat_messages_no_system() {
let messages = vec![Message::user("hi")];
let chat_msgs = build_chat_messages(None, &messages, None);
assert_eq!(chat_msgs.len(), 1);
}
#[test]
fn build_chat_messages_with_cache_config() {
let messages = vec![Message::user(
"a long user message that is at least one thousand chars. "
.repeat(20)
.as_str(),
)];
let cfg = CachePromptConfig::default();
let chat_msgs = build_chat_messages(Some("system prompt"), &messages, Some(&cfg));
assert_eq!(chat_msgs.len(), 2);
assert!(chat_msgs[0].cache_control.is_some());
}
#[test]
fn sanitize_removes_orphaned_tool_results() {
let mut messages = vec![
Message::user("hi"),
Message::tool_result("orphan-id", "read_file", "file content"),
Message::assistant("hello"),
];
sanitize_orphaned_tool_results(&mut messages);
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].role, Role::User);
assert_eq!(messages[1].role, Role::Assistant);
}
#[test]
fn sanitize_keeps_valid_tool_results() {
let tc = edgecrab_types::ToolCall {
id: "valid-id".into(),
r#type: "function".into(),
function: edgecrab_types::FunctionCall {
name: "read_file".into(),
arguments: "{}".into(),
},
thought_signature: None,
};
let mut messages = vec![
Message::user("hi"),
Message::assistant_with_tool_calls("calling tool", vec![tc]),
Message::tool_result("valid-id", "read_file", "file content"),
];
sanitize_orphaned_tool_results(&mut messages);
assert_eq!(messages.len(), 3, "valid tool result should be kept");
}
#[tokio::test]
async fn budget_exhaustion_at_gate_returns_synthetic_response() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(0)
.build()
.expect("build");
let result = agent
.execute_loop("do something", Some("Be helpful."), None, None, None)
.await
.expect("loop should not error on budget exhaustion");
assert!(
!result.final_response.is_empty(),
"budget-exhausted agent must not return empty response"
);
assert!(
result.final_response.contains("iteration limit"),
"synthetic response should mention 'iteration limit'; got: '{}'",
result.final_response
);
assert!(
result.budget_exhausted,
"budget_exhausted must be true when loop exits via budget gate"
);
assert!(!result.interrupted, "interrupted must be false");
assert_eq!(
result.api_calls, 0,
"no API calls should occur with budget=0"
);
}
#[tokio::test]
async fn chat_never_returns_empty_on_budget_exhaustion() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(0)
.build()
.expect("build");
let response = agent
.chat("do a lot of things")
.await
.expect("chat should not error");
assert!(
!response.is_empty(),
"chat() must not return empty string on budget exhaustion"
);
}
#[tokio::test]
async fn normal_completion_resets_budget_exhausted_flag() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(10)
.build()
.expect("build");
let result = agent
.execute_loop("hello", Some("Be helpful."), None, None, None)
.await
.expect("loop");
assert!(!result.final_response.is_empty());
assert!(
!result.budget_exhausted,
"normal completion must not set budget_exhausted"
);
assert!(!result.interrupted);
}
#[tokio::test]
async fn budget_exactly_one_produces_response() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(1)
.build()
.expect("build");
let result = agent
.execute_loop("hello", None, None, None, None)
.await
.expect("loop");
assert!(!result.final_response.is_empty());
assert!(
!result.budget_exhausted,
"text response was produced, not exhausted"
);
assert_eq!(result.api_calls, 1);
}
#[tokio::test]
async fn budget_exhausted_exactly_on_tool_turn_boundary() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(0)
.build()
.expect("build");
let result = agent
.execute_loop("run", None, None, None, None)
.await
.expect("loop");
assert!(result.budget_exhausted, "budget_exhausted must be true");
assert!(
!result.final_response.is_empty(),
"synthetic response must not be empty"
);
assert_eq!(result.api_calls, 0, "no API calls before budget gate");
}
#[tokio::test]
async fn multi_turn_tool_chain_completes_with_sufficient_budget() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(10)
.build()
.expect("build");
let result = agent
.execute_loop("do something and respond", None, None, None, None)
.await
.expect("loop");
assert!(
!result.final_response.is_empty(),
"response must be non-empty"
);
assert!(
!result.budget_exhausted,
"should complete normally without budget exhaustion"
);
assert!(!result.interrupted);
assert_eq!(result.api_calls, 1, "one API call for a text-only response");
}
#[test]
fn inject_budget_warning_appends_to_tool_message_json() {
let mut messages = vec![
Message::user("task"),
Message::tool_result("id1", "read_file", r#"{"output": "some content"}"#),
];
inject_budget_warning(&mut messages, "[URGENT: wrap up]");
let last = messages.last().expect("budget warning target exists");
let text = last.text_content();
assert!(
text.contains("_budget_warning"),
"budget warning should be injected into JSON tool message; got: {text}"
);
assert!(
text.contains("wrap up"),
"warning text should be present; got: {text}"
);
}
#[test]
fn inject_budget_warning_appends_to_tool_message_plain() {
let mut messages = vec![
Message::user("task"),
Message::tool_result("id1", "read_file", "plain text output"),
];
inject_budget_warning(&mut messages, "[URGENT: wrap up]");
let last = messages.last().expect("budget warning target exists");
let text = last.text_content();
assert!(
text.contains("wrap up"),
"plain-text warning should be appended; got: {text}"
);
}
#[test]
fn inject_budget_warning_falls_back_to_user_message_when_no_tools() {
let mut messages = vec![
Message::user("hello"),
Message::assistant("how can I help?"),
];
let before = messages.len();
inject_budget_warning(&mut messages, "[BUDGET: 70%]");
assert_eq!(
messages.len(),
before + 1,
"should inject a new user message as fallback"
);
let last = messages.last().expect("fallback warning message exists");
assert_eq!(last.role, Role::User);
assert!(last.text_content().contains("70%"));
}
#[test]
fn get_budget_warning_none_below_70_percent() {
assert!(
get_budget_warning(6, 10).is_none(),
"60% should produce no warning"
);
}
#[test]
fn get_budget_warning_at_70_percent() {
let w = get_budget_warning(7, 10);
assert!(w.is_some(), "70% should produce BUDGET warning");
assert!(
w.expect("70% warning should exist").contains("BUDGET"),
"should say BUDGET"
);
}
#[test]
fn get_budget_warning_at_90_percent() {
let w = get_budget_warning(9, 10);
assert!(w.is_some(), "90% should produce URGENT warning");
assert!(
w.expect("90% warning should exist").contains("URGENT"),
"should say URGENT"
);
}
#[test]
fn get_budget_warning_zero_max_iterations() {
assert!(
get_budget_warning(5, 0).is_none(),
"zero max_iterations should produce no warning (avoid div-by-zero)"
);
}
#[test]
fn sanitize_handles_empty_messages() {
let mut messages: Vec<Message> = Vec::new();
sanitize_orphaned_tool_results(&mut messages);
assert_eq!(messages.len(), 0);
}
#[test]
fn sanitize_removes_multiple_orphans() {
let mut messages = vec![
Message::user("hi"),
Message::tool_result("orphan-1", "read_file", "content-a"),
Message::tool_result("orphan-2", "write_file", "content-b"),
Message::assistant("done"),
];
sanitize_orphaned_tool_results(&mut messages);
assert_eq!(messages.len(), 2, "both orphans should be removed");
}
#[test]
fn sanitize_handles_tool_result_without_tool_call_id() {
let mut msg = Message::tool_result("some-id", "read_file", "data");
msg.tool_call_id = None; let mut messages = vec![Message::user("hi"), msg, Message::assistant("done")];
sanitize_orphaned_tool_results(&mut messages);
assert_eq!(messages.len(), 2, "None-id tool result should be removed");
}
#[test]
fn build_chat_messages_tool_role_uses_tool_call_id() {
let tc = edgecrab_types::ToolCall {
id: "tc-abc".into(),
r#type: "function".into(),
function: edgecrab_types::FunctionCall {
name: "read_file".into(),
arguments: "{}".into(),
},
thought_signature: None,
};
let messages = vec![
Message::user("read something"),
Message::assistant_with_tool_calls("sure", vec![tc]),
Message::tool_result("tc-abc", "read_file", "contents"),
];
let chat_msgs = build_chat_messages(None, &messages, None);
assert_eq!(chat_msgs.len(), 3);
}
#[test]
fn build_chat_messages_empty_input() {
let chat_msgs = build_chat_messages(None, &[], None);
assert_eq!(
chat_msgs.len(),
0,
"empty messages with no system → 0 chat messages"
);
}
fn make_dispatch_context_for_test<'a>(
registry: &'a Arc<ToolRegistry>,
cancel: &'a CancellationToken,
state_db: &'a Option<Arc<edgecrab_state::SessionDb>>,
process_table: &'a Arc<ProcessTable>,
capability_suppressions: Arc<Mutex<HashMap<String, ToolErrorResponse>>>,
) -> DispatchContext<'a> {
DispatchContext {
cwd: std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")),
registry: Some(registry),
cancel,
state_db,
platform: edgecrab_types::Platform::Cli,
process_table,
provider: None,
gateway_sender: None,
sub_agent_runner: None,
event_tx: None,
clarify_tx: None,
approval_tx: None,
origin_chat: None,
app_config_ref: AppConfigRef::default(),
conversation_session_id: "test-conversation".into(),
todo_store: None,
capability_suppressions,
}
}
#[tokio::test]
async fn dispatch_single_tool_uses_dispatch_context_cwd() {
let registry = Arc::new(ToolRegistry::new());
let cancel = CancellationToken::new();
let state_db = None;
let process_table = Arc::new(ProcessTable::new());
let capability_suppressions = Arc::new(Mutex::new(HashMap::new()));
let mut dctx = make_dispatch_context_for_test(
®istry,
&cancel,
&state_db,
&process_table,
capability_suppressions,
);
let workspace = TempDir::new().expect("workspace");
std::fs::write(workspace.path().join("proof.txt"), "dispatch cwd works").expect("write");
dctx.cwd = workspace.path().to_path_buf();
let result = dispatch_single_tool(
"read_file",
r#"{"path":"proof.txt","line_numbers":false}"#,
&dctx,
)
.await;
assert!(result.contains("dispatch cwd works"), "got: {result}");
}
#[tokio::test]
async fn dispatch_single_tool_returns_structured_json_error() {
let registry = Arc::new(ToolRegistry::new());
let cancel = CancellationToken::new();
let state_db = None;
let process_table = Arc::new(ProcessTable::new());
let capability_suppressions = Arc::new(Mutex::new(HashMap::new()));
let dctx = make_dispatch_context_for_test(
®istry,
&cancel,
&state_db,
&process_table,
capability_suppressions,
);
let result = dispatch_single_tool("read_file", "{}", &dctx).await;
let parsed = parse_tool_error_response(&result).expect("structured tool error");
assert_eq!(parsed.response_type, "tool_error");
assert_eq!(parsed.category, "arguments");
assert_eq!(parsed.code, "invalid_arguments");
assert_eq!(parsed.tool.as_deref(), Some("read_file"));
}
#[tokio::test]
async fn dispatch_single_tool_suppresses_repeated_capability_retry() {
let registry = Arc::new(ToolRegistry::new());
let cancel = CancellationToken::new();
let state_db = None;
let process_table = Arc::new(ProcessTable::new());
let capability_suppressions = Arc::new(Mutex::new(HashMap::new()));
let dctx = make_dispatch_context_for_test(
®istry,
&cancel,
&state_db,
&process_table,
capability_suppressions.clone(),
);
let args_json = r#"{"command":"top"}"#;
let first = dispatch_single_tool("terminal", args_json, &dctx).await;
let first_payload = parse_tool_error_response(&first).expect("structured error");
assert_eq!(first_payload.code, "non_interactive_terminal_required");
remember_tool_suppression(&capability_suppressions, "terminal", args_json, &first);
let second = dispatch_single_tool("terminal", args_json, &dctx).await;
let second_payload = parse_tool_error_response(&second).expect("structured error");
assert_eq!(second_payload.code, "suppressed_capability_retry");
assert!(second_payload.error.contains("already blocked"));
}
#[tokio::test]
async fn cancellation_sets_interrupted_not_budget_exhausted() {
let provider: Arc<dyn LLMProvider> = Arc::new(edgequake_llm::MockProvider::new());
let agent = AgentBuilder::new("mock")
.provider(provider)
.max_iterations(100)
.build()
.expect("build");
let result = agent
.execute_loop("hello", None, None, None, None)
.await
.expect("loop");
assert!(
!result.budget_exhausted,
"normal completion must not set budget_exhausted"
);
assert!(
!result.interrupted,
"normal completion must not set interrupted"
);
assert!(!result.final_response.is_empty());
}
}