use std::collections::{BTreeMap, BTreeSet, HashMap};
use rho_sdk::model::{ContentBlock, Message, ToolCall, ToolResult};
use serde::Serialize;
use crate::history_message::HistoryMessage;
pub(super) const PROBE_SET_VERSION: u32 = 4;
const MAX_REFERENCE_ITEMS: usize = 8;
const REFERENCE_ITEM_CHARS: usize = 600;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Serialize)]
#[serde(rename_all = "snake_case")]
pub(super) enum Probe {
FilesChanged,
TestResult,
UserRequests,
Errors,
}
impl Probe {
pub(super) fn question(self) -> &'static str {
match self {
Self::FilesChanged => {
"List every file path the agent created, edited, or deleted in this session, one per line."
}
Self::TestResult => {
"What was the most recent test, build, or lint command the agent ran (for example cargo test, cargo clippy, pytest, or make), and did it pass or fail? Quote the command."
}
Self::UserRequests => {
"What has the user asked for in this session, and what constraints or preferences did they state?"
}
Self::Errors => {
"What errors or failed tool calls came up in this session, and how were they handled?"
}
}
}
pub(super) fn scoring(self) -> Scoring {
match self {
Self::FilesChanged => Scoring::ExactMatch,
Self::TestResult | Self::UserRequests | Self::Errors => Scoring::Judge,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum Scoring {
ExactMatch,
Judge,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) enum Reference {
Paths(Vec<String>),
Facts(String),
}
impl Reference {
pub(super) fn text(&self) -> String {
match self {
Self::Paths(paths) => paths.join("\n"),
Self::Facts(facts) => facts.clone(),
}
}
}
pub(super) fn references(history: &[Message]) -> BTreeMap<Probe, Reference> {
let results = history
.iter()
.filter_map(|message| match message {
Message::ToolResult(result) => Some((result.id.as_str(), result)),
_ => None,
})
.collect::<HashMap<_, _>>();
let calls = history
.iter()
.filter_map(Message::completed_assistant_content)
.flatten()
.filter_map(|block| match block {
ContentBlock::ToolCall(call) => Some(call),
ContentBlock::Text(_) | ContentBlock::Image(_) => None,
})
.collect::<Vec<_>>();
let mut references = BTreeMap::new();
let paths = calls
.iter()
.flat_map(|call| changed_paths(call))
.collect::<BTreeSet<_>>();
if !paths.is_empty() {
references.insert(
Probe::FilesChanged,
Reference::Paths(paths.into_iter().collect()),
);
}
if let Some((call, result)) = calls.iter().rev().find_map(|call| {
let command = shell_command(call)?;
is_check_command(command).then_some((command, results.get(call.id.as_str())?))
}) {
references.insert(
Probe::TestResult,
Reference::Facts(format!(
"command: {}\nresult: {}",
excerpt(call),
status(result)
)),
);
}
let requests = history
.iter()
.filter_map(|message| match HistoryMessage::of(message) {
HistoryMessage::User(blocks) => Some(text_of(blocks)),
_ => None,
})
.filter(|text| !text.trim().is_empty() && !is_host_context(text))
.collect::<Vec<_>>();
if !requests.is_empty() {
references.insert(
Probe::UserRequests,
Reference::Facts(numbered(requests.iter().map(|text| excerpt(text)))),
);
}
let errors = calls
.iter()
.filter(|call| is_consequential(call))
.filter_map(|call| {
let result = results.get(call.id.as_str())?;
(!result.ok).then(|| {
format!(
"{} {}: {}",
call.name,
excerpt(&call.arguments.to_string()),
excerpt(&result.content)
)
})
})
.collect::<Vec<_>>();
if !errors.is_empty() {
references.insert(
Probe::Errors,
Reference::Facts(numbered(errors.into_iter())),
);
}
references
}
pub(super) fn path_recall(paths: &[String], answer: &str) -> f64 {
if paths.is_empty() {
return 0.0;
}
let found = paths
.iter()
.filter(|path| names_path(answer, &path_suffix(path)))
.count();
found as f64 / paths.len() as f64
}
fn names_path(answer: &str, suffix: &str) -> bool {
let is_path_char = |c: char| c.is_alphanumeric() || matches!(c, '.' | '_' | '-');
answer.match_indices(suffix).any(|(start, _)| {
let before = answer[..start].chars().next_back();
let after = answer[start + suffix.len()..].chars().next();
before.is_none_or(|c| !is_path_char(c)) && after.is_none_or(|c| !is_path_char(c))
})
}
fn path_suffix(path: &str) -> String {
let parts = path
.split(['/', '\\'])
.filter(|part| !part.is_empty())
.collect::<Vec<_>>();
parts[parts.len().saturating_sub(2)..].join("/")
}
fn changed_paths(call: &ToolCall) -> Vec<String> {
let string = |key: &str| call.arguments.get(key).and_then(|value| value.as_str());
match call.name.as_str() {
"write" | "str_replace" => string("path").map(str::to_owned).into_iter().collect(),
"edit" => string("input")
.into_iter()
.flat_map(str::lines)
.filter_map(|line| {
let header = line.trim().strip_prefix('[')?.strip_suffix(']')?;
Some(header.rsplit_once('#')?.0.to_owned())
})
.collect(),
"apply_patch" => string("input")
.into_iter()
.flat_map(str::lines)
.filter_map(|line| {
[
"*** Add File: ",
"*** Update File: ",
"*** Delete File: ",
"*** Move to: ",
]
.iter()
.find_map(|prefix| line.strip_prefix(prefix))
.map(|path| path.trim().to_owned())
})
.collect(),
_ => Vec::new(),
}
}
fn is_host_context(text: &str) -> bool {
const HEADERS: [&str; 11] = [
"[process notification]",
"[agent notification]",
"[workflow notification]",
"[workflow started]",
"[runtime notifications ",
"[computer use context]",
"[conversation model switched ",
"[advisor model switched ",
"[advisor mode on]",
"[advisor mode off]",
"[edit tool switched]",
];
let text = text.trim_start();
HEADERS.iter().any(|header| text.starts_with(header))
}
fn is_consequential(call: &ToolCall) -> bool {
!changed_paths(call).is_empty() || shell_command(call).is_some_and(is_check_command)
}
fn shell_command(call: &ToolCall) -> Option<&str> {
matches!(call.name.as_str(), "bash" | "powershell")
.then(|| call.arguments.get("command")?.as_str())
.flatten()
}
fn is_check_command(command: &str) -> bool {
const INVOCATIONS: [&[&str]; 18] = [
&["cargo", "test"],
&["cargo", "nextest"],
&["cargo", "clippy"],
&["cargo", "check"],
&["cargo", "build"],
&["pytest"],
&["python3", "-m", "pytest"],
&["python3", "-m", "unittest"],
&["python3", "scripts/validate.py"],
&["npm", "test"],
&["npm", "run", "test"],
&["npm", "run", "lint"],
&["npm", "run", "build"],
&["pnpm", "test"],
&["go", "test"],
&["make"],
&["tsc"],
&["vitest"],
];
command
.lines()
.next()
.unwrap_or_default()
.split([';', '|', '&'])
.map(|segment| {
segment
.split_whitespace()
.skip_while(|word| word.contains('=') || *word == "timeout")
.skip_while(|word| word.chars().all(|c| c.is_ascii_digit() || c == 's'))
.collect::<Vec<_>>()
})
.any(|words| {
INVOCATIONS
.iter()
.any(|invocation| words.starts_with(invocation))
})
}
fn status(result: &ToolResult) -> &'static str {
if result.ok {
"passed"
} else {
"failed"
}
}
fn text_of(blocks: &[ContentBlock]) -> String {
blocks
.iter()
.filter_map(|block| match block {
ContentBlock::Text(text) => Some(text.as_str()),
ContentBlock::Image(_) | ContentBlock::ToolCall(_) => None,
})
.collect::<Vec<_>>()
.join("\n")
}
fn numbered(items: impl DoubleEndedIterator<Item = String>) -> String {
let mut recent = items.rev().take(MAX_REFERENCE_ITEMS).collect::<Vec<_>>();
recent.reverse();
recent
.iter()
.enumerate()
.map(|(index, item)| format!("{}. {item}", index + 1))
.collect::<Vec<_>>()
.join("\n")
}
fn excerpt(text: &str) -> String {
let chars = text.chars().count();
if chars <= REFERENCE_ITEM_CHARS {
return text.to_owned();
}
let half = REFERENCE_ITEM_CHARS / 2;
let head = text.chars().take(half).collect::<String>();
let tail = text.chars().skip(chars - half).collect::<String>();
format!("{head}\n[...]\n{tail}")
}
#[cfg(test)]
#[path = "probes_tests.rs"]
mod tests;