use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use crate::message::{ToolContent, ToolContentPart};
use crate::tool::ToolContext;
use super::{ToolDispatchContext, ToolDispatchResult, ToolMiddleware, ToolPipeline};
pub trait Verifier: Send + Sync {
fn verify<'a>(
&'a self,
ctx: &'a ToolContext,
tool_name: &'a str,
) -> Pin<Box<dyn Future<Output = VerifyResult> + Send + 'a>>;
}
#[derive(Debug, Clone)]
pub struct VerifyResult {
pub passed: bool,
pub diagnostics: String,
}
pub struct NoopVerifier;
impl Verifier for NoopVerifier {
fn verify<'a>(
&'a self,
_ctx: &'a ToolContext,
_tool_name: &'a str,
) -> Pin<Box<dyn Future<Output = VerifyResult> + Send + 'a>> {
Box::pin(async move {
VerifyResult {
passed: true,
diagnostics: String::new(),
}
})
}
}
pub struct VerifyMiddleware {
verifier: Arc<dyn Verifier>,
write_tools: Vec<String>,
}
impl VerifyMiddleware {
#[must_use]
pub fn new(verifier: Arc<dyn Verifier>, write_tools: Vec<String>) -> Self {
Self {
verifier,
write_tools,
}
}
}
impl std::fmt::Debug for VerifyMiddleware {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("VerifyMiddleware")
.field("verifier", &"<dyn Verifier>")
.field("write_tools", &self.write_tools)
.finish()
}
}
impl ToolMiddleware for VerifyMiddleware {
fn name(&self) -> &'static str {
"verify"
}
fn dispatch<'a>(
&'a self,
ctx: &'a mut ToolDispatchContext,
next: &'a ToolPipeline,
) -> Pin<Box<dyn Future<Output = ToolDispatchResult> + Send + 'a>> {
let verifier = &self.verifier;
let write_tools = &self.write_tools;
Box::pin(async move {
let mut result = next.dispatch(ctx).await;
let resolved = if result.resolved_tool_name.is_empty() {
&ctx.tool_name
} else {
&result.resolved_tool_name
};
let is_write = write_tools.iter().any(|t| t == resolved);
if !is_write || result.is_error {
return result;
}
let verify = verifier.verify(&ctx.tool_context, resolved).await;
append_verify_result(&mut result.output, &verify);
result
})
}
}
fn append_verify_result(output: &mut ToolContent, verify: &VerifyResult) {
let status = if verify.passed { "passed" } else { "failed" };
let block = format!("\n\n[verify] {status}: {}", verify.diagnostics);
match output {
ToolContent::Text(s) => s.push_str(&block),
ToolContent::Multipart(parts) => {
parts.push(ToolContentPart::Text {
text: block.trim_start().to_string(),
});
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::ToolContent;
use crate::middleware::{ToolDispatchContext, ToolPipeline};
use crate::tool::{PermissionCheck, ToolContext, ToolRegistry};
use std::sync::Arc;
struct CannedVerifier {
result: VerifyResult,
called: Arc<std::sync::Mutex<bool>>,
}
impl Verifier for CannedVerifier {
fn verify<'a>(
&'a self,
_ctx: &'a ToolContext,
_tool_name: &'a str,
) -> Pin<Box<dyn Future<Output = VerifyResult> + Send + 'a>> {
let result = self.result.clone();
let called = self.called.clone();
Box::pin(async move {
*called.lock().unwrap() = true;
result
})
}
}
fn make_verifier(
passed: bool,
diagnostics: &str,
) -> (Arc<CannedVerifier>, Arc<std::sync::Mutex<bool>>) {
let called = Arc::new(std::sync::Mutex::new(false));
let v = Arc::new(CannedVerifier {
result: VerifyResult {
passed,
diagnostics: diagnostics.to_string(),
},
called: called.clone(),
});
(v, called)
}
struct FixedOutputMiddleware {
output: ToolContent,
is_error: bool,
}
impl ToolMiddleware for FixedOutputMiddleware {
fn name(&self) -> &'static str {
"fixed_output"
}
fn dispatch<'a>(
&'a self,
_ctx: &'a mut ToolDispatchContext,
_next: &'a ToolPipeline,
) -> Pin<Box<dyn Future<Output = ToolDispatchResult> + Send + 'a>> {
let output = self.output.clone();
let is_error = self.is_error;
Box::pin(async move {
ToolDispatchResult {
output,
is_error,
resolved_tool_name: String::new(),
tool_call_id: String::new(),
duration: std::time::Duration::ZERO,
display_hint: None,
}
})
}
}
fn pipeline_with(
verify: VerifyMiddleware,
output: ToolContent,
is_error: bool,
) -> ToolPipeline {
let registry = Arc::new(ToolRegistry::new());
ToolPipeline::builder()
.with_middleware(verify)
.with_middleware(FixedOutputMiddleware { output, is_error })
.with_core(registry)
.build()
.expect("pipeline builds")
}
fn ctx_for(tool_name: &str) -> ToolDispatchContext {
ToolDispatchContext {
tool_name: tool_name.to_string(),
input: serde_json::json!({}),
call_id: "c1".to_string(),
turn_number: 0,
cancel: Arc::new(crate::cancel::CancelSignal::new()),
permission: PermissionCheck::Allow,
tool_context: ToolContext::default(),
}
}
fn write_tools() -> Vec<String> {
vec!["Write".to_string(), "Edit".to_string()]
}
#[tokio::test]
async fn verify_runs_after_write_tool() {
let (v, called) = make_verifier(true, "ok");
let mw = VerifyMiddleware::new(v, write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("wrote 42 bytes"), false);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
assert!(
*called.lock().unwrap(),
"verifier should be called for write tools"
);
assert!(
result.output.to_string().contains("[verify]"),
"output should contain verify block: {}",
result.output
);
}
#[tokio::test]
async fn verify_skipped_for_read_tool() {
let (v, called) = make_verifier(true, "ok");
let mw = VerifyMiddleware::new(v, write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("file contents"), false);
let mut ctx = ctx_for("Read");
let result = pipeline.dispatch(&mut ctx).await;
assert!(
!*called.lock().unwrap(),
"verifier must not be called for non-write tools"
);
assert!(
!result.output.to_string().contains("[verify]"),
"output should not contain verify block for read tools"
);
}
#[tokio::test]
async fn verify_result_appended_to_text_output() {
let (v, _) = make_verifier(true, "all checks pass");
let mw = VerifyMiddleware::new(v, write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("wrote 1 file"), false);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
let s = match result.output {
ToolContent::Text(s) => s,
ToolContent::Multipart(parts) => {
panic!("expected Text, got Multipart with {} parts", parts.len())
}
};
assert!(
s.contains("wrote 1 file"),
"original output should be preserved: {s}"
);
assert!(
s.contains("[verify] passed: all checks pass"),
"verify block should be appended: {s}"
);
}
#[tokio::test]
async fn verify_result_appended_to_multipart_output() {
let (v, _) = make_verifier(false, "1 error: expected `;`");
let mw = VerifyMiddleware::new(v, write_tools());
let existing = ToolContent::from_multipart(vec![ToolContentPart::Text {
text: "wrote 2 files".to_string(),
}]);
let pipeline = pipeline_with(mw, existing, false);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
let parts = match result.output {
ToolContent::Multipart(parts) => parts,
ToolContent::Text(t) => panic!("expected Multipart, got Text: {t}"),
};
assert_eq!(parts.len(), 2, "verify block should be a new part");
match &parts[0] {
ToolContentPart::Text { text } => {
assert!(
text.contains("wrote 2 files"),
"first part unchanged: {text}"
);
assert!(
!text.contains("[verify]"),
"verify block must not leak into the first part: {text}"
);
}
ToolContentPart::Image { .. } => panic!("expected first part to be Text, got Image"),
}
match &parts[1] {
ToolContentPart::Text { text } => {
assert!(
text.contains("[verify] failed: 1 error: expected `;`"),
"verify block in second part: {text}"
);
}
ToolContentPart::Image { .. } => panic!("expected second part to be Text, got Image"),
}
}
#[tokio::test]
async fn verify_pass_does_not_block() {
let (v, _) = make_verifier(true, "clean");
let mw = VerifyMiddleware::new(v, write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("ok"), false);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
assert!(
!result.is_error,
"verify pass must not mark the result as error"
);
assert!(
result.output.to_string().contains("[verify] passed"),
"diagnostics should be present: {}",
result.output
);
}
#[tokio::test]
async fn verify_fail_appended_not_raised() {
let (v, _) = make_verifier(false, "compile error in main.rs");
let mw = VerifyMiddleware::new(v, write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("wrote main.rs"), false);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
assert!(
!result.is_error,
"verify failure must not mark the result as error"
);
assert!(
result
.output
.to_string()
.contains("[verify] failed: compile error"),
"failed verify block should be appended: {}",
result.output
);
}
#[tokio::test]
async fn verify_skipped_when_tool_errored() {
let (v, called) = make_verifier(true, "ok");
let mw = VerifyMiddleware::new(v, write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("permission denied"), true);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
assert!(
!*called.lock().unwrap(),
"verifier must not be called when the tool errored"
);
assert!(
!result.output.to_string().contains("[verify]"),
"no verify block expected on a tool error: {}",
result.output
);
assert!(
result.is_error,
"is_error from the tool should be preserved"
);
}
#[tokio::test]
async fn verify_block_format() {
let (pass_v, _) = make_verifier(true, "diag-pass");
let mw_pass = VerifyMiddleware::new(pass_v, write_tools());
let pipeline = pipeline_with(mw_pass, ToolContent::from_string("x"), false);
let mut ctx = ctx_for("Write");
let pass_out = pipeline.dispatch(&mut ctx).await.output.to_string();
assert!(
pass_out.contains("[verify] passed: diag-pass"),
"pass format: {pass_out}"
);
let (fail_v, _) = make_verifier(false, "diag-fail");
let mw_fail = VerifyMiddleware::new(fail_v, write_tools());
let pipeline = pipeline_with(mw_fail, ToolContent::from_string("x"), false);
let mut ctx = ctx_for("Write");
let fail_out = pipeline.dispatch(&mut ctx).await.output.to_string();
assert!(
fail_out.contains("[verify] failed: diag-fail"),
"fail format: {fail_out}"
);
}
#[test]
fn verify_middleware_name() {
let (v, _) = make_verifier(true, "");
let mw = VerifyMiddleware::new(v, write_tools());
assert_eq!(mw.name(), "verify");
}
#[test]
fn verify_middleware_debug() {
let (v, _) = make_verifier(true, "");
let mw = VerifyMiddleware::new(v, write_tools());
let debug = format!("{mw:?}");
assert!(debug.contains("VerifyMiddleware"));
assert!(debug.contains("write_tools"));
}
#[tokio::test]
async fn noop_verifier_always_passes() {
let verifier = NoopVerifier;
let ctx = ToolContext::default();
let result = verifier.verify(&ctx, "Write").await;
assert!(result.passed, "NoopVerifier must always pass");
assert!(
result.diagnostics.is_empty(),
"NoopVerifier must produce empty diagnostics"
);
}
#[tokio::test]
async fn noop_verifier_in_verify_middleware_appends_passed_block() {
let mw = VerifyMiddleware::new(Arc::new(NoopVerifier), write_tools());
let pipeline = pipeline_with(mw, ToolContent::from_string("wrote 42 bytes"), false);
let mut ctx = ctx_for("Write");
let result = pipeline.dispatch(&mut ctx).await;
let rendered = result.output.to_string();
assert!(
rendered.contains("[verify]"),
"NoopVerifier still triggers the middleware's append: got {rendered:?}"
);
assert!(
rendered.contains("[verify] passed:"),
"appended block reports passed status: got {rendered:?}"
);
}
}