use agent_base::llm_trait::{ChatRequest, LlmProvider};
use agent_base::types::ChatMessage;
use std::sync::Arc;
use std::time::Duration;
#[derive(serde::Deserialize, Debug)]
pub(crate) struct JudgeResult {
pub done: bool,
pub reason: String,
}
pub(crate) async fn call_completion_judge(
client: Option<&Arc<dyn LlmProvider>>,
user_input: &str,
model_response: &str,
all_user_inputs: &[String],
judge_fail_open: bool,
judge_timeout_secs: u64,
recent_user_count: usize,
) -> Result<JudgeResult, String> {
let Some(client) = client else {
if judge_fail_open {
return Ok(JudgeResult {
done: true,
reason: "no LLM client available for judge".to_string(),
});
} else {
return Err("no LLM client available for judge".to_string());
}
};
let system_prompt = "You are a task completion judge. \
Given the user's conversation history and the agent's response, \
determine if the agent has sufficiently answered the task. \
Reply with JSON: {\"done\": true/false, \"reason\": \"brief explanation\"}";
let user_context = if all_user_inputs.is_empty() {
user_input.to_string()
} else {
let n = recent_user_count;
let start = all_user_inputs.len().saturating_sub(n);
let recent = &all_user_inputs[start..];
if recent.len() <= 1 {
user_input.to_string()
} else {
recent
.iter()
.enumerate()
.map(|(i, msg)| format!("{}. {}", start + i + 1, msg))
.collect::<Vec<_>>()
.join("\n")
}
};
let user_prompt = format!(
"【User Messages】\n{}\n\n【Agent Response】\n{}",
user_context, model_response
);
let messages = vec![
ChatMessage::system(system_prompt.to_string()),
ChatMessage::user(user_prompt),
];
let timeout_duration = Duration::from_secs(judge_timeout_secs);
let result = tokio::time::timeout(timeout_duration, async {
let request = ChatRequest::new(messages)
.with_response_format(agent_base::llm_trait::request::ResponseFormat::JsonObject);
let response = client
.chat(request)
.await
.map_err(|e| format!("LLM judge call failed: {}", e))?;
let result: JudgeResult = serde_json::from_str(&response.content)
.map_err(|e| format!("Failed to parse judge response: {}", e))?;
Ok(result)
})
.await;
let result: Result<JudgeResult, String> = match result {
Ok(inner) => inner, Err(_elapsed) => Err(format!("judge timeout after {}s", judge_timeout_secs)),
};
match result {
Ok(judge_result) => Ok(judge_result),
Err(e) => {
tracing::warn!(
error = %e,
fail_open = judge_fail_open,
"completion judge failed"
);
if judge_fail_open {
Ok(JudgeResult {
done: true,
reason: format!("judge failed ({}), trusting model", e),
})
} else {
Err(e)
}
}
}
}