fix: preserve confirmed output on interceptor abort
This commit is contained in:
@@ -1375,6 +1375,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let call_info = call_info_map.get(&tool_result.tool_use_id);
|
let call_info = call_info_map.get(&tool_result.tool_use_id);
|
||||||
|
let mut abort_reason = None;
|
||||||
if let Some((tool_call, meta, tool, context)) = call_info {
|
if let Some((tool_call, meta, tool, context)) = call_info {
|
||||||
let mut info = ToolResultInfo {
|
let mut info = ToolResultInfo {
|
||||||
call: tool_call.clone(),
|
call: tool_call.clone(),
|
||||||
@@ -1387,7 +1388,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
match self.interceptor.post_tool_call(&mut info).await {
|
match self.interceptor.post_tool_call(&mut info).await {
|
||||||
PostToolAction::Continue => {}
|
PostToolAction::Continue => {}
|
||||||
PostToolAction::Abort(reason) => {
|
PostToolAction::Abort(reason) => {
|
||||||
return Err(EngineError::Aborted(reason));
|
abort_reason = Some(reason);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
tool_result = info.result;
|
tool_result = info.result;
|
||||||
@@ -1451,6 +1452,9 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
"Tool execution terminalized"
|
"Tool execution terminalized"
|
||||||
);
|
);
|
||||||
self.emit_tool_result(&tool_result);
|
self.emit_tool_result(&tool_result);
|
||||||
|
if let Some(reason) = abort_reason {
|
||||||
|
return Err(EngineError::Aborted(reason));
|
||||||
|
}
|
||||||
Ok(true)
|
Ok(true)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1021,3 +1021,76 @@ async fn test_before_tool_call_synthetic_result_committed() {
|
|||||||
} if call_id == "call_1" && summary == "permission denied"
|
} if call_id == "call_1" && summary == "permission denied"
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
|
||||||
|
let client = MockLlmClient::new(vec![
|
||||||
|
Event::tool_use_start(0, "call_confirmed", "confirmed"),
|
||||||
|
Event::tool_input_delta(0, r#"{}"#),
|
||||||
|
Event::tool_use_stop(0),
|
||||||
|
Event::Status(StatusEvent {
|
||||||
|
status: ResponseStatus::Completed,
|
||||||
|
}),
|
||||||
|
]);
|
||||||
|
let mut engine = Engine::new(client);
|
||||||
|
let tool = SlowTool::new("confirmed", 1);
|
||||||
|
engine.register_tool(tool.definition());
|
||||||
|
|
||||||
|
struct AbortAfterResult;
|
||||||
|
#[async_trait]
|
||||||
|
impl Interceptor for AbortAfterResult {
|
||||||
|
async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction {
|
||||||
|
PostToolAction::Abort("policy stopped the run".to_string())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
engine.set_interceptor(AbortAfterResult);
|
||||||
|
|
||||||
|
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
|
||||||
|
let published = observed.clone();
|
||||||
|
engine.on_tool_result(move |_| published.lock().unwrap().push("published"));
|
||||||
|
let committed = observed.clone();
|
||||||
|
let mut annotate = move |item: &Item| {
|
||||||
|
if matches!(item, Item::ToolResult { .. }) {
|
||||||
|
committed.lock().unwrap().push("committed");
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
};
|
||||||
|
|
||||||
|
let mut history = History::new();
|
||||||
|
let output = engine
|
||||||
|
.run_with_annotation(&mut history, "run confirmed tool", &mut annotate)
|
||||||
|
.await;
|
||||||
|
observed.lock().unwrap().push("run-returned");
|
||||||
|
|
||||||
|
assert_eq!(tool.call_count(), 1);
|
||||||
|
assert_eq!(
|
||||||
|
observed.lock().unwrap().as_slice(),
|
||||||
|
["committed", "published", "run-returned"]
|
||||||
|
);
|
||||||
|
assert!(matches!(
|
||||||
|
output.result,
|
||||||
|
agen::EngineRunExit::Interrupted(agen::StopReason::Unexpected(
|
||||||
|
agen::EngineError::Aborted(ref reason)
|
||||||
|
)) if reason == "policy stopped the run"
|
||||||
|
));
|
||||||
|
let terminal: Vec<_> = history
|
||||||
|
.iter()
|
||||||
|
.filter_map(|entry| match &entry.item {
|
||||||
|
Item::ToolResult {
|
||||||
|
call_id,
|
||||||
|
disposition,
|
||||||
|
..
|
||||||
|
} if call_id == "call_confirmed" => Some(*disposition),
|
||||||
|
_ => None,
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
assert_eq!(terminal, [ToolResultDisposition::Success]);
|
||||||
|
assert!(!history.iter().any(|entry| matches!(
|
||||||
|
&entry.item,
|
||||||
|
Item::ToolResult {
|
||||||
|
call_id,
|
||||||
|
disposition: ToolResultDisposition::OutcomeUnknown,
|
||||||
|
..
|
||||||
|
} if call_id == "call_confirmed"
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|||||||
@@ -8081,7 +8081,17 @@ mod build_summary_prompt_tests {
|
|||||||
|
|
||||||
worker.ensure_segment_head().unwrap();
|
worker.ensure_segment_head().unwrap();
|
||||||
worker.wire_history_persistence();
|
worker.wire_history_persistence();
|
||||||
worker.set_history_for_test(vec![Item::tool_call("call-1", "Bash", "{}")]);
|
worker.set_history_for_test(vec![
|
||||||
|
Item::tool_call("call-known", "Read", "{}"),
|
||||||
|
Item::tool_result_item_with_disposition_and_attachments(
|
||||||
|
"call-known",
|
||||||
|
"known result",
|
||||||
|
Some("confirmed output".to_string()),
|
||||||
|
agen::ToolResultDisposition::Success,
|
||||||
|
Vec::new(),
|
||||||
|
),
|
||||||
|
Item::tool_call("call-orphan", "Bash", "{}"),
|
||||||
|
]);
|
||||||
let _ = worker
|
let _ = worker
|
||||||
.handle_worker_result(
|
.handle_worker_result(
|
||||||
EngineRunExit::Interrupted(StopReason::Cancelled),
|
EngineRunExit::Interrupted(StopReason::Cancelled),
|
||||||
@@ -8090,6 +8100,44 @@ mod build_summary_prompt_tests {
|
|||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
|
let history = worker.history();
|
||||||
|
assert_eq!(
|
||||||
|
history
|
||||||
|
.iter()
|
||||||
|
.filter(|item| matches!(
|
||||||
|
item,
|
||||||
|
Item::ToolResult {
|
||||||
|
call_id,
|
||||||
|
disposition: agen::ToolResultDisposition::Success,
|
||||||
|
..
|
||||||
|
} if call_id == "call-known"
|
||||||
|
))
|
||||||
|
.count(),
|
||||||
|
1
|
||||||
|
);
|
||||||
|
assert!(!history.iter().any(|item| matches!(
|
||||||
|
item,
|
||||||
|
Item::ToolResult {
|
||||||
|
call_id,
|
||||||
|
disposition: agen::ToolResultDisposition::OutcomeUnknown,
|
||||||
|
..
|
||||||
|
} if call_id == "call-known"
|
||||||
|
)));
|
||||||
|
assert_eq!(
|
||||||
|
history
|
||||||
|
.iter()
|
||||||
|
.filter(|item| matches!(
|
||||||
|
item,
|
||||||
|
Item::ToolResult {
|
||||||
|
call_id,
|
||||||
|
disposition: agen::ToolResultDisposition::OutcomeUnknown,
|
||||||
|
..
|
||||||
|
} if call_id == "call-orphan"
|
||||||
|
))
|
||||||
|
.count(),
|
||||||
|
1
|
||||||
|
);
|
||||||
|
|
||||||
let entries = worker
|
let entries = worker
|
||||||
.store
|
.store
|
||||||
.read_all(
|
.read_all(
|
||||||
@@ -8105,13 +8153,14 @@ mod build_summary_prompt_tests {
|
|||||||
LogEntry::AnnotatedToolResult {
|
LogEntry::AnnotatedToolResult {
|
||||||
entry: session_store::LoggedHistoryEntry {
|
entry: session_store::LoggedHistoryEntry {
|
||||||
item: session_store::LoggedItem::ToolResult {
|
item: session_store::LoggedItem::ToolResult {
|
||||||
|
call_id,
|
||||||
disposition: agen::ToolResultDisposition::OutcomeUnknown,
|
disposition: agen::ToolResultDisposition::OutcomeUnknown,
|
||||||
..
|
..
|
||||||
},
|
},
|
||||||
..
|
..
|
||||||
},
|
},
|
||||||
..
|
..
|
||||||
}
|
} if call_id == "call-orphan"
|
||||||
)
|
)
|
||||||
})
|
})
|
||||||
.expect("durable OutcomeUnknown closure");
|
.expect("durable OutcomeUnknown closure");
|
||||||
|
|||||||
Reference in New Issue
Block a user