worker: treat rolled-back extraction as cancelled
This commit is contained in:
@@ -72,6 +72,16 @@ pub(crate) struct InternalWorkerError {
|
|||||||
pub(crate) async fn run_internal_worker(
|
pub(crate) async fn run_internal_worker(
|
||||||
spec: InternalWorkerSpec,
|
spec: InternalWorkerSpec,
|
||||||
) -> Result<InternalWorkerResult, InternalWorkerError> {
|
) -> Result<InternalWorkerResult, InternalWorkerError> {
|
||||||
|
run_internal_worker_with_prepare(spec, |_| {}).await
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn run_internal_worker_with_prepare<F>(
|
||||||
|
spec: InternalWorkerSpec,
|
||||||
|
prepare: F,
|
||||||
|
) -> Result<InternalWorkerResult, InternalWorkerError>
|
||||||
|
where
|
||||||
|
F: FnOnce(&mut Worker<Box<dyn LlmClient>, EphemeralSessionStore>),
|
||||||
|
{
|
||||||
let InternalWorkerSpec {
|
let InternalWorkerSpec {
|
||||||
identity,
|
identity,
|
||||||
mut manifest,
|
mut manifest,
|
||||||
@@ -163,6 +173,7 @@ pub(crate) async fn run_internal_worker(
|
|||||||
}
|
}
|
||||||
let session_id = worker.session_id();
|
let session_id = worker.session_id();
|
||||||
let segment_id = worker.segment_id();
|
let segment_id = worker.segment_id();
|
||||||
|
prepare(&mut worker);
|
||||||
|
|
||||||
match worker.run_text(&input).await {
|
match worker.run_text(&input).await {
|
||||||
Ok(lifecycle) => Ok(InternalWorkerResult {
|
Ok(lifecycle) => Ok(InternalWorkerResult {
|
||||||
@@ -355,6 +366,35 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
struct CancelBeforeAiClient {
|
||||||
|
calls: Arc<AtomicUsize>,
|
||||||
|
cancel_sender: Arc<Mutex<Option<tokio::sync::mpsc::Sender<()>>>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl LlmClient for CancelBeforeAiClient {
|
||||||
|
fn clone_boxed(&self) -> Box<dyn LlmClient> {
|
||||||
|
Box::new(self.clone())
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn stream(
|
||||||
|
&self,
|
||||||
|
_request: Request,
|
||||||
|
) -> Result<Pin<Box<dyn Stream<Item = Result<LlmEvent, ClientError>> + Send>>, ClientError>
|
||||||
|
{
|
||||||
|
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||||
|
let sender = self
|
||||||
|
.cancel_sender
|
||||||
|
.lock()
|
||||||
|
.expect("cancel sender lock")
|
||||||
|
.clone()
|
||||||
|
.expect("cancel sender installed before run");
|
||||||
|
sender.send(()).await.expect("internal engine is live");
|
||||||
|
Ok(Box::pin(futures::stream::pending()))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn manifest() -> WorkerManifest {
|
fn manifest() -> WorkerManifest {
|
||||||
WorkerManifest::from_toml(
|
WorkerManifest::from_toml(
|
||||||
r#"
|
r#"
|
||||||
@@ -414,6 +454,34 @@ permission = "write"
|
|||||||
assert_eq!(result.identity.kind, "test");
|
assert_eq!(result.identity.kind, "test");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cancellation_before_ai_item_returns_rolled_back_lifecycle() {
|
||||||
|
let calls = Arc::new(AtomicUsize::new(0));
|
||||||
|
let cancel_sender = Arc::new(Mutex::new(None));
|
||||||
|
let mut internal_spec = spec(calls.clone(), &[]);
|
||||||
|
internal_spec.client = Box::new(CancelBeforeAiClient {
|
||||||
|
calls: calls.clone(),
|
||||||
|
cancel_sender: cancel_sender.clone(),
|
||||||
|
});
|
||||||
|
let prepare_sender = cancel_sender.clone();
|
||||||
|
|
||||||
|
let result = match run_internal_worker_with_prepare(internal_spec, move |worker| {
|
||||||
|
*prepare_sender.lock().expect("cancel sender lock") =
|
||||||
|
Some(worker.engine_mut().cancel_sender());
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
Ok(result) => result,
|
||||||
|
Err(error) => panic!(
|
||||||
|
"Worker rollback should remain a lifecycle result: {:?}",
|
||||||
|
error.source
|
||||||
|
),
|
||||||
|
};
|
||||||
|
|
||||||
|
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||||
|
assert!(matches!(result.lifecycle, WorkerRunResult::RolledBack));
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn rejects_missing_explicit_tools_before_model_execution() {
|
async fn rejects_missing_explicit_tools_before_model_execution() {
|
||||||
let calls = Arc::new(AtomicUsize::new(0));
|
let calls = Arc::new(AtomicUsize::new(0));
|
||||||
|
|||||||
@@ -3475,7 +3475,22 @@ impl<C: LlmClient, St: Store> Worker<C, St> {
|
|||||||
lifecycle = ?result.lifecycle,
|
lifecycle = ?result.lifecycle,
|
||||||
"internal Worker execution completed"
|
"internal Worker execution completed"
|
||||||
);
|
);
|
||||||
result.usage.as_ref().map(usage_audit_from_event)
|
let usage = result.usage.as_ref().map(usage_audit_from_event);
|
||||||
|
if let Some(error) = extract_internal_worker_lifecycle_error(&result.lifecycle) {
|
||||||
|
audit
|
||||||
|
.emit(
|
||||||
|
self.workspace_client(),
|
||||||
|
event_tx,
|
||||||
|
memory::audit::WorkerLifecycleStatus::Cancelled,
|
||||||
|
"worker_cancelled: internal Worker run rolled back before AI output",
|
||||||
|
usage,
|
||||||
|
Some(extract_audit_base),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
return Err(error);
|
||||||
|
}
|
||||||
|
usage
|
||||||
}
|
}
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
tracing::debug!(
|
tracing::debug!(
|
||||||
@@ -3629,6 +3644,13 @@ impl<C: LlmClient, St: Store> Worker<C, St> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn extract_internal_worker_lifecycle_error(lifecycle: &WorkerRunResult) -> Option<WorkerError> {
|
||||||
|
match lifecycle {
|
||||||
|
WorkerRunResult::RolledBack => Some(WorkerError::Engine(EngineError::Cancelled)),
|
||||||
|
WorkerRunResult::Finished | WorkerRunResult::Paused | WorkerRunResult::LimitReached => None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn lifecycle_status_for_worker_error(err: &WorkerError) -> memory::audit::WorkerLifecycleStatus {
|
fn lifecycle_status_for_worker_error(err: &WorkerError) -> memory::audit::WorkerLifecycleStatus {
|
||||||
if matches!(err, WorkerError::Engine(EngineError::Cancelled)) {
|
if matches!(err, WorkerError::Engine(EngineError::Cancelled)) {
|
||||||
memory::audit::WorkerLifecycleStatus::Cancelled
|
memory::audit::WorkerLifecycleStatus::Cancelled
|
||||||
@@ -6323,6 +6345,19 @@ mod build_summary_prompt_tests {
|
|||||||
server.join().unwrap();
|
server.join().unwrap();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn rolled_back_internal_extract_is_cancelled_before_pointer_commit() {
|
||||||
|
let error = extract_internal_worker_lifecycle_error(&WorkerRunResult::RolledBack)
|
||||||
|
.expect("rolled-back extract must not enter the success path");
|
||||||
|
|
||||||
|
assert!(matches!(error, WorkerError::Engine(EngineError::Cancelled)));
|
||||||
|
assert!(matches!(
|
||||||
|
lifecycle_status_for_worker_error(&error),
|
||||||
|
memory::audit::WorkerLifecycleStatus::Cancelled
|
||||||
|
));
|
||||||
|
assert!(extract_internal_worker_lifecycle_error(&WorkerRunResult::Finished).is_none());
|
||||||
|
}
|
||||||
|
|
||||||
fn minimal_manifest() -> WorkerManifest {
|
fn minimal_manifest() -> WorkerManifest {
|
||||||
let toml_str = r#"
|
let toml_str = r#"
|
||||||
[worker]
|
[worker]
|
||||||
|
|||||||
Reference in New Issue
Block a user