use agent_base::engine::react_loop_guard::{GuardAction, GuardCtx, ReactLoopGuard};
use agent_base::llm::StreamClient;
use agent_base::types::{ChatMessage, ResponseFormat};
use async_trait::async_trait;
use std::sync::Arc;
use std::time::Duration;
pub struct DefaultGuardConfig {
pub reasoning_only_max_strikes: usize,
pub empty_response_max_strikes: usize,
pub reasoning_only_nudge: String,
pub empty_response_nudge: String,
pub use_llm_judge: bool,
pub judge_timeout_secs: u64,
pub judge_skip_threshold: usize,
pub judge_fail_open: bool,
pub detect_short_response: bool,
pub short_response_min_input: usize,
pub short_response_max_output: usize,
pub short_response_nudge: String,
}
impl Default for DefaultGuardConfig {
fn default() -> Self {
Self {
reasoning_only_max_strikes: 3,
empty_response_max_strikes: 3,
reasoning_only_nudge: "You produced internal reasoning but no tool call \
and no final answer. Make a decision now: call a tool to make progress, \
or write your final answer as plain text."
.to_string(),
empty_response_nudge: "Your response was empty. Please provide a response \
with either a tool call or your final answer."
.to_string(),
use_llm_judge: true,
judge_timeout_secs: 10,
judge_skip_threshold: 256,
judge_fail_open: false, detect_short_response: true,
short_response_min_input: 128,
short_response_max_output: 64,
short_response_nudge: "Your response may be incomplete — \
you may need to continue."
.to_string(),
}
}
}
pub struct DefaultGuard {
config: DefaultGuardConfig,
llm_client: Option<Arc<dyn StreamClient>>,
}
impl DefaultGuard {
pub fn new(config: DefaultGuardConfig) -> Self {
Self {
config,
llm_client: None,
}
}
pub fn with_llm_client(config: DefaultGuardConfig, llm_client: Arc<dyn StreamClient>) -> Self {
Self {
config,
llm_client: Some(llm_client),
}
}
async fn call_completion_judge(
&self,
user_input: &str,
model_response: &str,
) -> Result<JudgeResult, String> {
let Some(client) = &self.llm_client else {
if self.config.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 original question 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_prompt = format!(
"【User Question】\n{}\n\n【Agent Response】\n{}",
user_input, model_response
);
let messages = vec![
ChatMessage::system(system_prompt.to_string()),
ChatMessage::user(user_prompt),
];
let timeout_duration = Duration::from_secs(self.config.judge_timeout_secs);
let result = tokio::time::timeout(timeout_duration, async {
let raw_response = client
.chat(&messages, &[], None, Some(&ResponseFormat::JsonObject))
.await
.map_err(|e| format!("LLM judge call failed: {}", e))?;
let result: JudgeResult = serde_json::from_str(&raw_response)
.map_err(|e| format!("Failed to parse judge response: {}", e))?;
Ok(result)
})
.await;
match result {
Ok(judge_result) => judge_result,
Err(_) => {
tracing::warn!(
timeout_secs = self.config.judge_timeout_secs,
fail_open = self.config.judge_fail_open,
"completion judge timeout or failure"
);
if self.config.judge_fail_open {
Ok(JudgeResult {
done: true,
reason: format!(
"judge timeout after {}s, trusting model",
self.config.judge_timeout_secs
),
})
} else {
Err(format!(
"judge timeout after {}s, not trusting model",
self.config.judge_timeout_secs
))
}
}
}
}
}
#[derive(serde::Deserialize, Debug)]
struct JudgeResult {
done: bool,
reason: String,
}
#[async_trait]
impl ReactLoopGuard for DefaultGuard {
async fn on_reasoning_only(&self, ctx: &GuardCtx) -> GuardAction {
let strikes = ctx.reasoning_only_strikes;
if strikes >= self.config.reasoning_only_max_strikes {
return GuardAction::Fail(
"model produced only reasoning across multiple turns".to_string(),
);
}
GuardAction::Continue(self.config.reasoning_only_nudge.clone())
}
async fn on_empty_response(&self, ctx: &GuardCtx) -> GuardAction {
let strikes = ctx.empty_response_strikes;
if strikes >= self.config.empty_response_max_strikes {
return GuardAction::Fail("model returned empty responses repeatedly".to_string());
}
GuardAction::Continue(self.config.empty_response_nudge.clone())
}
async fn on_text_only(&self, ctx: &GuardCtx) -> GuardAction {
let input_len = ctx.user_input.chars().count();
let output_len = ctx.model_response.chars().count();
let is_short_response = self.config.detect_short_response
&& input_len > self.config.short_response_min_input
&& output_len < self.config.short_response_max_output
&& input_len > output_len;
if is_short_response {
tracing::info!(
input_chars = input_len,
output_chars = output_len,
min_input = self.config.short_response_min_input,
max_output = self.config.short_response_max_output,
run_has_tool_calls = ctx.run_has_tool_calls,
"short response detected in text-only branch"
);
if ctx.run_has_tool_calls && self.config.use_llm_judge {
const INPUT_LEN_LIMIT: usize = 10_000;
if input_len > INPUT_LEN_LIMIT {
tracing::info!(
input_chars = input_len,
input_limit = INPUT_LEN_LIMIT,
"skipping LLM judge — user input too large, trusting model"
);
return GuardAction::Done;
}
match self
.call_completion_judge(&ctx.user_input, &ctx.model_response)
.await
{
Ok(judge) => {
if judge.done {
GuardAction::Done
} else {
GuardAction::Continue(format!(
"Your answer is incomplete: {}. Continue working on the task.",
judge.reason
))
}
}
Err(e) => {
tracing::warn!("completion judge failed: {}", e);
if self.config.judge_fail_open {
GuardAction::Done
} else {
GuardAction::Continue(
"Cannot verify task completion, please continue working."
.to_string(),
)
}
}
}
} else {
GuardAction::Continue(self.config.short_response_nudge.clone())
}
} else if ctx.run_has_tool_calls && self.config.use_llm_judge {
if output_len >= self.config.judge_skip_threshold {
tracing::debug!(
response_chars = output_len,
threshold = self.config.judge_skip_threshold,
"text-only response long enough, skipping judge"
);
return GuardAction::Done;
}
const INPUT_LEN_LIMIT: usize = 10_000;
if input_len > INPUT_LEN_LIMIT {
tracing::info!(
input_chars = input_len,
input_limit = INPUT_LEN_LIMIT,
"skipping LLM judge — user input too large, trusting model"
);
return GuardAction::Done;
}
tracing::info!(
response_chars = output_len,
threshold = self.config.judge_skip_threshold,
"text-only response short, calling judge"
);
match self
.call_completion_judge(&ctx.user_input, &ctx.model_response)
.await
{
Ok(judge) => {
if judge.done {
GuardAction::Done
} else {
GuardAction::Continue(format!(
"Your answer is incomplete: {}. Continue working on the task.",
judge.reason
))
}
}
Err(e) => {
tracing::warn!("completion judge failed: {}", e);
if self.config.judge_fail_open {
GuardAction::Done
} else {
GuardAction::Continue(
"Cannot verify task completion, please continue working."
.to_string(),
)
}
}
}
} else {
GuardAction::Done
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use agent_base::llm::{LlmCapabilities, StreamChunk};
use agent_base::types::{AgentResult, FinishReason, SessionId};
use futures_core::Stream;
use serde_json::Value;
use std::pin::Pin;
use std::task::{Context, Poll};
fn make_ctx(
reasoning_only_strikes: usize,
empty_response_strikes: usize,
run_has_tool_calls: bool,
) -> GuardCtx {
GuardCtx {
session_id: SessionId {
id: 1,
external_id: None,
},
turn_count: 1,
user_input: "test".to_string(),
model_response: "response".to_string(),
finish_reason: FinishReason::Stop,
available_tools: vec![],
reasoning_only_strikes,
empty_response_strikes,
run_has_tool_calls,
}
}
fn make_ctx_with_text(
user_input: &str,
model_response: &str,
run_has_tool_calls: bool,
) -> GuardCtx {
GuardCtx {
session_id: SessionId {
id: 1,
external_id: None,
},
turn_count: 2,
user_input: user_input.to_string(),
model_response: model_response.to_string(),
finish_reason: FinishReason::Stop,
available_tools: vec!["echo".to_string()],
reasoning_only_strikes: 0,
empty_response_strikes: 0,
run_has_tool_calls,
}
}
struct MockJudgeClient {
response: String,
}
impl MockJudgeClient {
fn new(response: Value) -> Self {
Self {
response: response.to_string(),
}
}
}
#[async_trait]
impl StreamClient for MockJudgeClient {
async fn stream(
&self,
_messages: &[ChatMessage],
_tools: &[Value],
_reasoning: Option<&agent_base::ReasoningConfig>,
_response_format: Option<&ResponseFormat>,
) -> AgentResult<Pin<Box<dyn Stream<Item = AgentResult<StreamChunk>> + Send>>> {
let chunks = vec![
Ok(StreamChunk::Text(self.response.clone())),
Ok(StreamChunk::Stop {
finish_reason: Some("stop".to_string()),
}),
];
Ok(Box::pin(futures_util::stream::iter(chunks)))
}
fn capabilities(&self) -> LlmCapabilities {
LlmCapabilities::default()
}
}
struct MockTimeoutClient;
#[async_trait]
impl StreamClient for MockTimeoutClient {
async fn stream(
&self,
_messages: &[ChatMessage],
_tools: &[Value],
_reasoning: Option<&agent_base::ReasoningConfig>,
_response_format: Option<&ResponseFormat>,
) -> AgentResult<Pin<Box<dyn Stream<Item = AgentResult<StreamChunk>> + Send>>> {
struct HangingStream;
impl Stream for HangingStream {
type Item = AgentResult<StreamChunk>;
fn poll_next(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Option<Self::Item>> {
Poll::Pending }
}
Ok(Box::pin(HangingStream))
}
fn capabilities(&self) -> LlmCapabilities {
LlmCapabilities::default()
}
}
#[tokio::test]
async fn test_reasoning_only_below_threshold() {
let guard = DefaultGuard::new(DefaultGuardConfig::default());
let ctx = make_ctx(1, 0, false);
let action = guard.on_reasoning_only(&ctx).await;
assert!(matches!(action, GuardAction::Continue(_)));
}
#[tokio::test]
async fn test_reasoning_only_at_threshold() {
let guard = DefaultGuard::new(DefaultGuardConfig::default());
let ctx = make_ctx(3, 0, false);
let action = guard.on_reasoning_only(&ctx).await;
assert!(matches!(action, GuardAction::Fail(_)));
}
#[tokio::test]
async fn test_empty_response_below_threshold() {
let guard = DefaultGuard::new(DefaultGuardConfig::default());
let ctx = make_ctx(0, 1, false);
let action = guard.on_empty_response(&ctx).await;
assert!(matches!(action, GuardAction::Continue(_)));
}
#[tokio::test]
async fn test_empty_response_at_threshold() {
let guard = DefaultGuard::new(DefaultGuardConfig::default());
let ctx = make_ctx(0, 3, false);
let action = guard.on_empty_response(&ctx).await;
assert!(matches!(action, GuardAction::Fail(_)));
}
#[tokio::test]
async fn test_text_only_without_tool_calls() {
let guard = DefaultGuard::new(DefaultGuardConfig::default());
let ctx = make_ctx(0, 0, false);
let action = guard.on_text_only(&ctx).await;
assert!(matches!(action, GuardAction::Done));
}
#[tokio::test]
async fn test_text_only_with_tool_calls_no_llm_client() {
let config = DefaultGuardConfig {
use_llm_judge: true,
..DefaultGuardConfig::default()
};
let guard = DefaultGuard::new(config);
let ctx = make_ctx(0, 0, true);
let action = guard.on_text_only(&ctx).await;
assert!(matches!(action, GuardAction::Continue(_)));
}
#[tokio::test]
async fn test_text_only_with_tool_calls_no_llm_client_fail_open() {
let config = DefaultGuardConfig {
use_llm_judge: true,
judge_fail_open: true,
..DefaultGuardConfig::default()
};
let guard = DefaultGuard::new(config);
let ctx = make_ctx(0, 0, true);
let action = guard.on_text_only(&ctx).await;
assert!(matches!(action, GuardAction::Done));
}
#[tokio::test]
async fn test_text_only_with_tool_calls_judge_disabled() {
let config = DefaultGuardConfig {
use_llm_judge: false,
..DefaultGuardConfig::default()
};
let guard = DefaultGuard::new(config);
let ctx = make_ctx(0, 0, true);
let action = guard.on_text_only(&ctx).await;
assert!(matches!(action, GuardAction::Done));
}
#[tokio::test]
async fn test_custom_config() {
let config = DefaultGuardConfig {
reasoning_only_max_strikes: 5,
empty_response_max_strikes: 2,
..DefaultGuardConfig::default()
};
let guard = DefaultGuard::new(config);
let ctx = make_ctx(4, 0, false);
let action = guard.on_reasoning_only(&ctx).await;
assert!(matches!(action, GuardAction::Continue(_)));
let ctx = make_ctx(0, 2, false);
let action = guard.on_empty_response(&ctx).await;
assert!(matches!(action, GuardAction::Fail(_)));
}
#[tokio::test]
async fn test_text_only_judge_says_done() {
let judge_client = Arc::new(MockJudgeClient::new(serde_json::json!({
"done": true,
"reason": "task is complete"
})));
let guard = DefaultGuard::with_llm_client(
DefaultGuardConfig::default(),
judge_client,
);
let ctx = make_ctx_with_text(
"what is 2+2?",
"The answer is 4.",
true, );
let action = guard.on_text_only(&ctx).await;
assert!(
matches!(action, GuardAction::Done),
"judge says done → Done, got: {:?}",
action
);
}
#[tokio::test]
async fn test_text_only_judge_says_not_done() {
let judge_client = Arc::new(MockJudgeClient::new(serde_json::json!({
"done": false,
"reason": "only answered part of the question"
})));
let guard = DefaultGuard::with_llm_client(
DefaultGuardConfig::default(),
judge_client,
);
let ctx = make_ctx_with_text(
"list all files and explain each",
"Here are the files:", true,
);
let action = guard.on_text_only(&ctx).await;
match &action {
GuardAction::Continue(msg) => {
assert!(
msg.contains("incomplete"),
"nudge should mention incomplete: {}",
msg
);
assert!(
msg.contains("only answered part"),
"nudge should include judge reason: {}",
msg
);
}
other => panic!("expected Continue, got: {:?}", other),
}
}
#[tokio::test]
async fn test_text_only_short_response_detected_no_tools() {
let long_input = "a".repeat(200);
let short_output = "done";
let guard = DefaultGuard::new(DefaultGuardConfig::default());
let ctx = make_ctx_with_text(&long_input, short_output, false);
let action = guard.on_text_only(&ctx).await;
match &action {
GuardAction::Continue(msg) => {
assert!(
msg.contains("incomplete"),
"short response nudge: {}",
msg
);
}
other => panic!("expected Continue for short response, got: {:?}", other),
}
}
#[tokio::test]
async fn test_text_only_short_response_with_judge_done() {
let long_input = "a".repeat(200);
let short_output = "42";
let judge_client = Arc::new(MockJudgeClient::new(serde_json::json!({
"done": true,
"reason": "answer is correct"
})));
let guard = DefaultGuard::with_llm_client(
DefaultGuardConfig::default(),
judge_client,
);
let ctx = make_ctx_with_text(&long_input, short_output, true);
let action = guard.on_text_only(&ctx).await;
assert!(
matches!(action, GuardAction::Done),
"short response + judge done → Done, got: {:?}",
action
);
}
#[tokio::test]
async fn test_text_only_skip_threshold() {
let long_output = "x".repeat(300);
let judge_client = Arc::new(MockJudgeClient::new(serde_json::json!({
"done": false,
"reason": "would be incomplete but skipped"
})));
let guard = DefaultGuard::with_llm_client(
DefaultGuardConfig::default(),
judge_client,
);
let ctx = make_ctx_with_text("query", &long_output, true);
let action = guard.on_text_only(&ctx).await;
assert!(
matches!(action, GuardAction::Done),
"response >= skip_threshold → Done without calling judge, got: {:?}",
action
);
}
#[tokio::test]
async fn test_text_only_judge_timeout_fail_open() {
let judge_client = Arc::new(MockTimeoutClient);
let config = DefaultGuardConfig {
judge_fail_open: true,
judge_timeout_secs: 1, ..DefaultGuardConfig::default()
};
let guard = DefaultGuard::with_llm_client(config, judge_client);
let ctx = make_ctx_with_text("query", "short answer", true);
let action = guard.on_text_only(&ctx).await;
assert!(
matches!(action, GuardAction::Done),
"timeout + fail_open → Done (trust model), got: {:?}",
action
);
}
#[tokio::test]
async fn test_text_only_judge_timeout_fail_closed() {
let judge_client = Arc::new(MockTimeoutClient);
let config = DefaultGuardConfig {
judge_fail_open: false,
judge_timeout_secs: 1, ..DefaultGuardConfig::default()
};
let guard = DefaultGuard::with_llm_client(config, judge_client);
let ctx = make_ctx_with_text("query", "short answer", true);
let action = guard.on_text_only(&ctx).await;
match &action {
GuardAction::Continue(msg) => {
assert!(
msg.contains("Cannot verify"),
"fail_closed timeout → Continue with verify message: {}",
msg
);
}
other => panic!(
"expected Continue for timeout + fail_closed, got: {:?}",
other
),
}
}
#[tokio::test]
async fn test_text_only_skip_judge_when_input_too_large() {
let long_input = "a".repeat(10_001);
let judge_client = Arc::new(MockJudgeClient::new(serde_json::json!({
"done": false,
"reason": "would be incomplete but skipped"
})));
let guard = DefaultGuard::with_llm_client(
DefaultGuardConfig::default(),
judge_client,
);
let ctx = make_ctx_with_text(&long_input, "short answer", true);
let action = guard.on_text_only(&ctx).await;
assert!(
matches!(action, GuardAction::Done),
"input > 10k → skip judge → Done, got: {:?}",
action
);
}
}