use std::collections::HashSet;
use std::fmt;
use std::path::Path;
use serde_json::{json, Value};
use super::client::{ChatClient, ChatRequest, ClientError, Message, Usage};
use super::config::ExploreConfig;
use super::trace::TraceWriter;
use super::{grounding, steering, toolset};
pub const MAX_TURNS: usize = 12;
const MAX_COMPLETION_TOKENS: u32 = 1024;
const TEMPERATURE: f32 = 0.0;
const THINK: bool = true;
const THRASH_LIMIT: f64 = 3.0;
const NUDGE_AT: usize = 2;
const TOKEN_BUDGET: u32 = 28_000;
const TIME_BUDGET_SECS: u64 = 90;
const RETRY_ON_LEAK: u32 = 1;
#[derive(Debug, Clone)]
pub struct ExploreAnswer {
pub text: String,
pub turns: usize,
pub truncated: bool,
}
#[derive(Debug)]
pub enum ExploreError {
ProviderDown {
url: String,
detail: String,
},
Client(String),
}
impl fmt::Display for ExploreError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ExploreError::ProviderDown { url, detail } => {
write!(f, "inference server unreachable at {url}: {detail}")
}
ExploreError::Client(msg) => write!(f, "chat client error: {msg}"),
}
}
}
impl std::error::Error for ExploreError {}
fn map_client_error(e: ClientError) -> ExploreError {
match e {
ClientError::Connection { url, detail } => ExploreError::ProviderDown { url, detail },
other => ExploreError::Client(other.to_string()),
}
}
pub trait ProgressReporter {
fn report(&self, progress: usize, total: usize, message: &str);
}
pub struct NoopReporter;
impl ProgressReporter for NoopReporter {
fn report(&self, _progress: usize, _total: usize, _message: &str) {}
}
pub fn run_explore(
question: &str,
root: &Path,
cfg: &ExploreConfig,
client: &dyn ChatClient,
) -> Result<ExploreAnswer, ExploreError> {
run_explore_reporting(question, root, cfg, client, &NoopReporter, None)
}
pub fn run_explore_reporting(
question: &str,
root: &Path,
cfg: &ExploreConfig,
client: &dyn ChatClient,
progress: &dyn ProgressReporter,
trace: Option<&TraceWriter>,
) -> Result<ExploreAnswer, ExploreError> {
let total = MAX_TURNS + 2;
let sys = steering::system_prompt(cfg.steering, root);
let mut messages: Vec<Message> = vec![Message::system(sys), Message::user(question)];
let call_id = trace.map(|tw| tw.call_start(question)).unwrap_or(0);
let call_t0 = std::time::Instant::now();
let mut agg = Usage::default();
let mut seen_sigs: HashSet<String> = HashSet::new();
let mut consec_unprod: f64 = 0.0;
let mut nudged = false;
let mut leak_retries_left = RETRY_ON_LEAK;
let mut ctx_tokens: u32 = 0;
let mut answer_raw: Option<String> = None;
let mut truncated = true; let mut turns_used = 0usize;
let mut activity = "exploring the codebase".to_string();
for turn in 0..MAX_TURNS {
turns_used = turn + 1;
if turn > 0 && call_t0.elapsed().as_secs() > TIME_BUDGET_SECS {
break;
}
progress.report(
turn + 1,
total,
&format!("turn {}/{MAX_TURNS} · {activity}", turn + 1),
);
let req = ChatRequest::new(messages.clone())
.with_tools(toolset::all_tools())
.with_tool_choice(json!("auto"))
.with_explore_sampling(TEMPERATURE, MAX_COMPLETION_TOKENS, THINK);
let req_trace = trace.map(|_| request_trace(&req, &cfg.model));
let t0 = std::time::Instant::now();
let resp = client.chat(req).map_err(map_client_error)?;
let wall = t0.elapsed().as_millis();
if let Some(u) = resp.usage {
agg.prompt_tokens = agg.prompt_tokens.saturating_add(u.prompt_tokens);
agg.completion_tokens = agg.completion_tokens.saturating_add(u.completion_tokens);
agg.total_tokens = agg.total_tokens.saturating_add(u.total_tokens);
if u.prompt_tokens > 0 {
ctx_tokens = u.prompt_tokens;
}
}
if let (Some(tw), Some(req_v)) = (trace, &req_trace) {
let resp_v = serde_json::to_value(&resp).unwrap_or(Value::Null);
tw.turn(call_id, turn + 1, req_v, &resp_v, resp.usage, wall);
}
let step = match resp.first_message() {
Some(m) => m.clone(),
None => break,
};
let content = step.content.clone().unwrap_or_default();
if step.tool_calls.is_empty() {
if leak_retries_left > 0
&& grounding::has_leak(&content)
&& grounding::extract_final(&content).is_none()
{
leak_retries_left -= 1;
messages.push(Message::assistant(grounding::neutralize_xml(&content)));
messages.push(Message::user(steering::LEAK_RETRY));
continue;
}
if grounding::extract_final(&content).is_some() {
truncated = false;
}
answer_raw = Some(content);
break;
}
let mut step = step.clone();
step.content = Some(grounding::neutralize_xml(&content));
messages.push(step.clone());
for c in &step.tool_calls {
let obs = grounding::neutralize_xml(&toolset::dispatch(&c.name, &c.arguments, root));
let sig = format!("{}:{}", c.name, compact_args(&c.arguments));
let is_dup = !seen_sigs.insert(sig);
let is_empty = toolset::is_empty_obs(&obs);
if is_dup {
consec_unprod += 1.0; } else if is_empty {
consec_unprod += 0.5; } else {
consec_unprod = 0.0;
}
messages.push(Message::tool(&c.id, obs));
}
activity = summarize_activity(&step.tool_calls);
if consec_unprod >= THRASH_LIMIT
|| ctx_tokens >= TOKEN_BUDGET
|| call_t0.elapsed().as_secs() > TIME_BUDGET_SECS
|| turn >= MAX_TURNS - 1
{
break;
}
if !nudged && (MAX_TURNS - 1 - turn) <= NUDGE_AT {
messages.push(Message::user(steering::NUDGE));
nudged = true;
}
}
if answer_raw.is_none() {
messages.push(Message::user(steering::FORCE_FINAL_ANSWER));
for att in 0..=RETRY_ON_LEAK {
progress.report(total - 1, total, "wrapping up: final answer");
let req = ChatRequest::new(messages.clone()).with_explore_sampling(
TEMPERATURE,
MAX_COMPLETION_TOKENS.max(1024),
THINK,
);
let req_trace = trace.map(|_| request_trace(&req, &cfg.model));
let t0 = std::time::Instant::now();
let resp = client.chat(req).map_err(map_client_error)?;
let wall = t0.elapsed().as_millis();
if let Some(u) = resp.usage {
agg.prompt_tokens = agg.prompt_tokens.saturating_add(u.prompt_tokens);
agg.completion_tokens = agg.completion_tokens.saturating_add(u.completion_tokens);
agg.total_tokens = agg.total_tokens.saturating_add(u.total_tokens);
}
if let (Some(tw), Some(req_v)) = (trace, &req_trace) {
let resp_v = serde_json::to_value(&resp).unwrap_or(Value::Null);
tw.turn(call_id, turns_used + 1, req_v, &resp_v, resp.usage, wall);
}
let fcontent = resp
.first_message()
.and_then(|m| m.content.clone())
.unwrap_or_default();
let salvaged = grounding::extract_final(&fcontent);
if grounding::has_leak(&fcontent) && salvaged.is_none() && att < RETRY_ON_LEAK {
messages.push(Message::assistant(grounding::neutralize_xml(&fcontent)));
messages.push(Message::user(steering::FORCE_STRICT));
continue;
}
answer_raw = salvaged.is_some().then_some(fcontent);
break;
}
}
progress.report(total, total, "grounding answer");
let text = answer_raw
.map(|c| grounding::get_final_answer(&c, root))
.unwrap_or_default();
if let Some(tw) = trace {
tw.call_end(call_id, &text, turns_used, agg, call_t0.elapsed().as_millis(), truncated);
}
Ok(ExploreAnswer { text, turns: turns_used, truncated })
}
fn request_trace(req: &ChatRequest, model: &str) -> Value {
let mut v = serde_json::to_value(req).unwrap_or(Value::Null);
if let Some(obj) = v.as_object_mut() {
obj.insert("model".to_string(), Value::String(model.to_string()));
}
v
}
fn compact_args(args: &Value) -> String {
if args.is_null() {
String::new()
} else {
serde_json::to_string(args).unwrap_or_default()
}
}
fn summarize_activity(calls: &[super::client::ToolCall]) -> String {
let mut parts: Vec<String> = Vec::new();
for c in calls {
let part = match c.name.as_str() {
toolset::READ => format!("Read {}", basename_arg(&c.arguments, "file_path")),
toolset::GLOB => format!("Glob {}", str_arg(&c.arguments, "pattern")),
toolset::GREP => format!("Grep {}", str_arg(&c.arguments, "pattern")),
other => format!("grove {}", other.trim_start_matches("mcp__grove__")),
};
parts.push(part);
}
let joined = parts.join(", ");
let s: String = if joined.chars().count() > 80 {
let mut t: String = joined.chars().take(77).collect();
t.push('…');
t
} else {
joined
};
if s.is_empty() {
"exploring the codebase".to_string()
} else {
s
}
}
fn str_arg(args: &Value, key: &str) -> String {
args.get(key)
.and_then(Value::as_str)
.unwrap_or("")
.chars()
.take(30)
.collect()
}
fn basename_arg(args: &Value, key: &str) -> String {
let p = args.get(key).and_then(Value::as_str).unwrap_or("");
Path::new(p)
.file_name()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| p.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::explore::client::{ChatResponse, Choice, Role, ToolCall};
use crate::explore::config::{Provider, Steering};
use std::cell::RefCell;
struct FakeClient {
scripted: RefCell<std::collections::VecDeque<ChatResponse>>,
seen_tool_names: RefCell<Vec<Vec<String>>>,
}
impl FakeClient {
fn new(responses: Vec<ChatResponse>) -> Self {
FakeClient {
scripted: RefCell::new(responses.into()),
seen_tool_names: RefCell::new(Vec::new()),
}
}
}
impl ChatClient for FakeClient {
fn chat(&self, req: ChatRequest) -> Result<ChatResponse, ClientError> {
self.seen_tool_names
.borrow_mut()
.push(req.tools.iter().map(|t| t.function.name.clone()).collect());
Ok(self
.scripted
.borrow_mut()
.pop_front()
.unwrap_or_else(|| text_response("(end)")))
}
}
fn text_response(s: &str) -> ChatResponse {
ChatResponse {
choices: vec![Choice {
message: Message {
role: Role::Assistant,
content: Some(s.to_string()),
tool_calls: vec![],
tool_call_id: None,
name: None,
},
finish_reason: None,
}],
usage: None,
}
}
fn tool_call_response(name: &str, args: Value) -> ChatResponse {
ChatResponse {
choices: vec![Choice {
message: Message {
role: Role::Assistant,
content: None,
tool_calls: vec![ToolCall {
id: "call_1".into(),
name: name.into(),
arguments: args,
}],
tool_call_id: None,
name: None,
},
finish_reason: None,
}],
usage: None,
}
}
fn cfg() -> ExploreConfig {
ExploreConfig {
provider: Provider::LlamaCpp,
base_url: "http://localhost:8080/v1".into(),
model: "qwen3.5-4b".into(),
steering: Steering::Standard,
allowed_tools: vec!["grove".into(), "rg".into()],
tap: false,
trace_retain: 50,
}
}
#[test]
fn voluntary_location_line_answer_is_not_truncated() {
let dir = std::env::temp_dir().join(format!("grove-agent-vol-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(dir.join("a.rs"), "fn a(){}\n").unwrap();
let client = FakeClient::new(vec![text_response("rust:a.rs#a@1")]);
let ans = run_explore("where is a", &dir, &cfg(), &client).unwrap();
assert!(!ans.truncated);
assert_eq!(ans.turns, 1);
assert_eq!(ans.text, "rust:a.rs#a@1");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn every_turn_offers_the_full_reference_toolset() {
let client = FakeClient::new(vec![text_response("done")]);
run_explore("q", Path::new("."), &cfg(), &client).unwrap();
let seen = &client.seen_tool_names.borrow()[0];
assert_eq!(
seen,
&vec![
"Glob",
"Grep",
"Read",
"mcp__grove__outline",
"mcp__grove__symbols",
"mcp__grove__source",
"mcp__grove__callers",
"mcp__grove__map",
"mcp__grove__definition",
]
);
}
#[test]
fn turn_cap_forces_an_answer_via_the_no_tools_turn() {
let mut responses = Vec::new();
for _ in 0..(MAX_TURNS + 2) {
responses.push(tool_call_response(
"mcp__grove__map",
json!({"dir": "."}),
));
}
responses.push(text_response("src/x.rs:1"));
let client = FakeClient::new(responses);
let ans = run_explore("q", Path::new("."), &cfg(), &client).unwrap();
assert!(ans.truncated, "answer came from the forced/backstop path");
let seen = client.seen_tool_names.borrow();
assert!(seen.last().unwrap().is_empty(), "forced turn has no tools");
}
#[test]
fn duplicate_calls_trip_the_thrash_backstop_early() {
let mut responses = Vec::new();
for _ in 0..5 {
responses.push(tool_call_response(
"mcp__grove__symbols",
json!({"dir": ".", "name": "zzz"}),
));
}
responses.push(text_response("(nothing)"));
let client = FakeClient::new(responses);
let ans = run_explore("q", Path::new("."), &cfg(), &client).unwrap();
assert!(ans.truncated);
assert!(ans.turns <= 4, "thrash broke early, turns={}", ans.turns);
}
#[test]
fn leaked_tool_call_is_retried_not_taken_as_answer() {
let dir = std::env::temp_dir().join(format!("grove-agent-leak-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(dir.join("a.rs"), "fn a(){}\n").unwrap();
let client = FakeClient::new(vec![
text_response("<tool_call>{\"name\":\"Grep\"}</tool_call>"),
text_response("rust:a.rs#a@1"),
]);
let ans = run_explore("q", &dir, &cfg(), &client).unwrap();
assert_eq!(ans.text, "rust:a.rs#a@1");
assert!(!ans.truncated);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn provider_down_maps_to_provider_down_error() {
struct DownClient;
impl ChatClient for DownClient {
fn chat(&self, _req: ChatRequest) -> Result<ChatResponse, ClientError> {
Err(ClientError::Connection {
url: "http://x".into(),
detail: "refused".into(),
})
}
}
let err = run_explore("q", Path::new("."), &cfg(), &DownClient).unwrap_err();
assert!(matches!(err, ExploreError::ProviderDown { .. }));
}
}