1use std::sync::{Arc, Mutex};
2use std::time::Duration;
3
4use serde::{Deserialize, Serialize};
5use tokio_util::sync::CancellationToken;
6
7use crate::{ErrorEnvelope, TurnOutcome};
8
9use super::{
10 AwaitEventKey, AwaitEventResolver, AwaitEventWaitIdentity, EffectHost, ExecutionScope,
11 Resolution, ResolveOutcome, RuntimeError,
12};
13
14#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
20pub struct TurnAddress {
21 pub session_id: String,
22 pub turn_id: String,
23}
24
25impl TurnAddress {
26 pub fn new(session_id: impl Into<String>, turn_id: impl Into<String>) -> Self {
27 Self {
28 session_id: session_id.into(),
29 turn_id: turn_id.into(),
30 }
31 }
32
33 fn scope(&self) -> ExecutionScope {
34 ExecutionScope::turn(&self.session_id, &self.turn_id)
35 }
36
37 fn validate(&self) -> Result<(), RuntimeError> {
38 self.scope().validate()
39 }
40}
41
42#[doc(hidden)]
49#[derive(Clone, Default)]
50pub struct TurnCancelOriginHint {
51 origin: Arc<Mutex<Option<Option<String>>>>,
52}
53
54impl TurnCancelOriginHint {
55 pub fn set(&self, origin: Option<String>) {
56 let mut hint = self.origin.lock().expect("turn cancel origin hint lock");
57 if hint.is_none() {
58 *hint = Some(origin);
59 }
60 }
61
62 pub(crate) fn get(&self) -> Option<String> {
63 self.origin
64 .lock()
65 .expect("turn cancel origin hint lock")
66 .clone()
67 .flatten()
68 }
69}
70
71#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
72pub struct TurnCancellationEvidence {
73 pub request_id: String,
74 #[serde(default, skip_serializing_if = "Option::is_none")]
76 pub origin: Option<String>,
77 #[serde(default, skip_serializing_if = "Option::is_none")]
78 pub reason: Option<String>,
79}
80
81#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
82pub struct TurnCancelRequest {
83 pub address: TurnAddress,
84 pub request_id: String,
85 #[serde(default, skip_serializing_if = "Option::is_none")]
87 pub origin: Option<String>,
88 #[serde(default, skip_serializing_if = "Option::is_none")]
89 pub reason: Option<String>,
90}
91
92impl TurnCancelRequest {
93 pub fn new(
94 address: TurnAddress,
95 request_id: impl Into<String>,
96 origin: Option<String>,
97 ) -> Self {
98 Self {
99 address,
100 request_id: request_id.into(),
101 origin,
102 reason: None,
103 }
104 }
105
106 pub fn with_reason(mut self, reason: impl Into<String>) -> Self {
107 self.reason = Some(reason.into());
108 self
109 }
110
111 fn validate(&self) -> Result<(), RuntimeError> {
112 self.address.validate()?;
113 if self.request_id.trim().is_empty() {
114 return Err(RuntimeError::new(
115 "invalid_turn_cancel_request",
116 "turn cancellation requires a non-empty request id",
117 ));
118 }
119 Ok(())
120 }
121
122 fn evidence(&self) -> TurnCancellationEvidence {
123 TurnCancellationEvidence {
124 request_id: self.request_id.clone(),
125 origin: self.origin.clone(),
126 reason: self.reason.clone(),
127 }
128 }
129}
130
131#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
132#[serde(tag = "outcome", content = "cancellation", rename_all = "snake_case")]
133pub enum TurnCancelOutcome {
134 Requested(TurnCancellationEvidence),
135 AlreadyRequested(TurnCancellationEvidence),
136 CompletionWonRace,
137 UnknownOrRevoked,
138}
139
140#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
148pub struct TurnCancelReceipt {
149 pub durability_tier: crate::DurabilityTier,
150 pub outcome: TurnCancelOutcome,
151}
152
153#[derive(Clone, Debug, Serialize, Deserialize)]
154#[serde(tag = "status", rename_all = "snake_case")]
155pub enum TurnTerminal {
156 Committed {
157 outcome: TurnOutcome,
158 #[serde(default, skip_serializing_if = "Option::is_none")]
159 cancellation: Option<TurnCancellationEvidence>,
160 #[serde(default, skip_serializing_if = "Option::is_none")]
161 session_revision: Option<u64>,
162 },
163 Failed {
164 error: ErrorEnvelope,
165 },
166}
167
168#[async_trait::async_trait]
170pub trait TurnAttach: Send + Sync {
171 async fn await_terminal(&self, address: &TurnAddress) -> Result<TurnTerminal, RuntimeError>;
172}
173
174#[derive(Clone)]
192pub struct TurnWorkDriver {
193 effect_host: Arc<dyn EffectHost>,
194 attach: Option<Arc<dyn TurnAttach>>,
195}
196
197impl TurnWorkDriver {
198 pub fn new(effect_host: Arc<dyn EffectHost>) -> Self {
199 Self {
200 effect_host,
201 attach: None,
202 }
203 }
204
205 pub fn with_attach(mut self, attach: Arc<dyn TurnAttach>) -> Self {
206 self.attach = Some(attach);
207 self
208 }
209
210 pub fn effect_host(&self) -> Arc<dyn EffectHost> {
211 Arc::clone(&self.effect_host)
212 }
213
214 pub async fn request_cancel(
215 &self,
216 request: TurnCancelRequest,
217 ) -> Result<TurnCancelReceipt, RuntimeError> {
218 request.validate()?;
219 let durability_tier = self.effect_host.durability_tier();
220 let key = cancel_gate_key(self.effect_host.as_ref(), &request.address).await?;
221 let evidence = request.evidence();
222 let resolution = gate_resolution(TurnGateTerminal::CancelRequested(evidence.clone()))?;
223 let outcome = match self
224 .effect_host
225 .resolve_await_event(&key, resolution)
226 .await?
227 {
228 ResolveOutcome::Accepted => Ok(TurnCancelOutcome::Requested(evidence)),
229 ResolveOutcome::AlreadyResolved { terminal } => match decode_gate(terminal)? {
230 TurnGateTerminal::CancelRequested(existing) => {
231 Ok(TurnCancelOutcome::AlreadyRequested(existing))
232 }
233 TurnGateTerminal::CompletionSealed => Ok(TurnCancelOutcome::CompletionWonRace),
234 },
235 ResolveOutcome::UnknownOrRevoked => Ok(TurnCancelOutcome::UnknownOrRevoked),
236 }?;
237 Ok(TurnCancelReceipt {
238 durability_tier,
239 outcome,
240 })
241 }
242
243 pub async fn await_terminal(
244 &self,
245 address: &TurnAddress,
246 ) -> Result<TurnTerminal, RuntimeError> {
247 address.validate()?;
248 if let Some(attach) = self.attach.as_ref() {
249 return attach.await_terminal(address).await;
250 }
251 let key = terminal_key(self.effect_host.as_ref(), address).await?;
252 let resolution = self
253 .effect_host
254 .await_await_event(&key, CancellationToken::new(), None)
255 .await?;
256 decode_terminal(address, resolution)
257 }
258
259 pub async fn await_terminal_with_timeout(
264 &self,
265 address: &TurnAddress,
266 timeout: Duration,
267 ) -> Result<TurnTerminal, RuntimeError> {
268 tokio::time::timeout(timeout, self.await_terminal(address))
269 .await
270 .map_err(|_| {
271 RuntimeError::new(
272 "turn_terminal_await_timeout",
273 format!(
274 "timed out awaiting terminal for turn `{}` in session `{}` after {} ms",
275 address.turn_id,
276 address.session_id,
277 timeout.as_millis()
278 ),
279 )
280 })?
281 }
282}
283
284#[derive(Clone, Debug, Serialize, Deserialize)]
285#[serde(tag = "state", content = "cancellation", rename_all = "snake_case")]
286enum TurnGateTerminal {
287 CancelRequested(TurnCancellationEvidence),
288 CompletionSealed,
289}
290
291fn gate_resolution(value: TurnGateTerminal) -> Result<Resolution, RuntimeError> {
292 serde_json::to_value(value)
293 .map(Resolution::Ok)
294 .map_err(|err| RuntimeError::new("turn_cancel_gate_encode", err.to_string()))
295}
296
297fn decode_gate(resolution: Resolution) -> Result<TurnGateTerminal, RuntimeError> {
298 match resolution {
299 Resolution::Ok(value) => serde_json::from_value(value)
300 .map_err(|err| RuntimeError::new("turn_cancel_gate_decode", err.to_string())),
301 other => Err(RuntimeError::new(
302 "turn_cancel_gate_invalid_terminal",
303 format!("turn cancellation gate resolved with {other:?}"),
304 )),
305 }
306}
307
308fn terminal_resolution(value: &TurnTerminal) -> Result<Resolution, RuntimeError> {
309 serde_json::to_value(value)
310 .map(Resolution::Ok)
311 .map_err(|err| RuntimeError::new("turn_terminal_encode", err.to_string()))
312}
313
314fn decode_terminal(
315 address: &TurnAddress,
316 resolution: Resolution,
317) -> Result<TurnTerminal, RuntimeError> {
318 match resolution {
319 Resolution::Ok(value) => serde_json::from_value(value).map_err(|err| {
320 RuntimeError::new(
321 "turn_terminal_decode",
322 format!(
323 "invalid terminal result for turn `{}` in session `{}`: {err}",
324 address.turn_id, address.session_id
325 ),
326 )
327 }),
328 other => Err(RuntimeError::new(
329 "turn_terminal_invalid_resolution",
330 format!(
331 "terminal result for turn `{}` in session `{}` resolved with {other:?}",
332 address.turn_id, address.session_id
333 ),
334 )),
335 }
336}
337
338async fn cancel_gate_key(
339 resolver: &dyn AwaitEventResolver,
340 address: &TurnAddress,
341) -> Result<AwaitEventKey, RuntimeError> {
342 resolver
343 .await_event_key(&address.scope(), AwaitEventWaitIdentity::TurnCancelGate)
344 .await
345}
346
347async fn terminal_key(
348 resolver: &dyn AwaitEventResolver,
349 address: &TurnAddress,
350) -> Result<AwaitEventKey, RuntimeError> {
351 resolver
352 .await_event_key(&address.scope(), AwaitEventWaitIdentity::TurnTerminal)
353 .await
354}
355
356pub(crate) struct ActiveTurnControl {
359 address: TurnAddress,
360 cancel_key: AwaitEventKey,
361 terminal_key: AwaitEventKey,
362 evidence: Mutex<Option<TurnCancellationEvidence>>,
363 local_cancel_origin: TurnCancelOriginHint,
364}
365
366impl ActiveTurnControl {
367 pub(crate) async fn new(
368 resolver: &dyn AwaitEventResolver,
369 address: TurnAddress,
370 ) -> Result<Self, RuntimeError> {
371 address.validate()?;
372 Ok(Self {
373 cancel_key: cancel_gate_key(resolver, &address).await?,
374 terminal_key: terminal_key(resolver, &address).await?,
375 address,
376 evidence: Mutex::new(None),
377 local_cancel_origin: TurnCancelOriginHint::default(),
378 })
379 }
380
381 pub(crate) fn with_local_cancel_origin(mut self, origin: TurnCancelOriginHint) -> Self {
382 self.local_cancel_origin = origin;
383 self
384 }
385
386 pub(crate) async fn await_cancel(
387 &self,
388 resolver: &dyn AwaitEventResolver,
389 stop_wait: CancellationToken,
390 ) -> Result<Option<TurnCancellationEvidence>, RuntimeError> {
391 let resolution = resolver
392 .await_await_event(&self.cancel_key, stop_wait, None)
393 .await?;
394 match decode_gate(resolution)? {
395 TurnGateTerminal::CancelRequested(evidence) => {
396 self.remember(evidence.clone());
397 Ok(Some(evidence))
398 }
399 TurnGateTerminal::CompletionSealed => Ok(None),
400 }
401 }
402
403 pub(crate) async fn settle_before_commit(
404 &self,
405 resolver: &dyn AwaitEventResolver,
406 locally_cancelled: bool,
407 ) -> Result<Option<TurnCancellationEvidence>, RuntimeError> {
408 if let Some(evidence) = self.evidence() {
409 return Ok(Some(evidence));
410 }
411 let proposed = if locally_cancelled {
412 TurnGateTerminal::CancelRequested(self.internal_evidence())
413 } else {
414 TurnGateTerminal::CompletionSealed
415 };
416 let outcome = resolver
417 .resolve_await_event(&self.cancel_key, gate_resolution(proposed.clone())?)
418 .await?;
419 let terminal = match outcome {
420 ResolveOutcome::Accepted => proposed,
421 ResolveOutcome::AlreadyResolved { terminal } => decode_gate(terminal)?,
422 ResolveOutcome::UnknownOrRevoked => {
423 return Err(RuntimeError::new(
424 "turn_control_unknown_or_revoked",
425 format!(
426 "turn `{}` in session `{}` was revoked before final commit",
427 self.address.turn_id, self.address.session_id
428 ),
429 ));
430 }
431 };
432 match terminal {
433 TurnGateTerminal::CancelRequested(evidence) => {
434 self.remember(evidence.clone());
435 Ok(Some(evidence))
436 }
437 TurnGateTerminal::CompletionSealed => Ok(None),
438 }
439 }
440
441 pub(crate) async fn publish_terminal(
442 &self,
443 resolver: &dyn AwaitEventResolver,
444 terminal: &TurnTerminal,
445 ) -> Result<(), RuntimeError> {
446 match resolver
447 .resolve_await_event(&self.terminal_key, terminal_resolution(terminal)?)
448 .await?
449 {
450 ResolveOutcome::Accepted | ResolveOutcome::AlreadyResolved { .. } => Ok(()),
451 ResolveOutcome::UnknownOrRevoked => Err(RuntimeError::new(
452 "turn_terminal_unknown_or_revoked",
453 format!(
454 "terminal promise for turn `{}` in session `{}` was revoked",
455 self.address.turn_id, self.address.session_id
456 ),
457 )),
458 }
459 }
460
461 pub(crate) fn evidence(&self) -> Option<TurnCancellationEvidence> {
462 self.evidence
463 .lock()
464 .expect("turn cancellation evidence lock")
465 .clone()
466 }
467
468 fn remember(&self, evidence: TurnCancellationEvidence) {
469 *self
470 .evidence
471 .lock()
472 .expect("turn cancellation evidence lock") = Some(evidence);
473 }
474
475 fn internal_evidence(&self) -> TurnCancellationEvidence {
476 TurnCancellationEvidence {
477 request_id: format!("internal:{}", self.address.turn_id),
478 origin: self.local_cancel_origin.get(),
479 reason: None,
480 }
481 }
482}
483
484#[cfg(test)]
485mod tests {
486 use super::*;
487 use crate::{InlineEffectHost, TurnFinish, TurnStop};
488
489 fn address(label: &str) -> TurnAddress {
490 TurnAddress::new(
491 format!("turn-control-{label}-{}", uuid::Uuid::new_v4()),
492 "turn-a",
493 )
494 }
495
496 fn request(address: TurnAddress, request_id: &str) -> TurnCancelRequest {
497 TurnCancelRequest::new(address, request_id, Some("user".to_string()))
498 .with_reason("stop button")
499 }
500
501 #[tokio::test]
502 async fn cancel_before_start_duplicate_and_terminal_attach() {
503 let host = Arc::new(InlineEffectHost::default());
504 let driver = TurnWorkDriver::new(host.clone());
505 let address = address("before-start");
506
507 let first = driver
508 .request_cancel(request(address.clone(), "request-1"))
509 .await
510 .expect("request cancellation");
511 assert_eq!(first.durability_tier, crate::DurabilityTier::Inline);
512 let evidence = match first.outcome {
513 TurnCancelOutcome::Requested(evidence) => evidence,
514 other => panic!("expected requested, got {other:?}"),
515 };
516 assert_eq!(evidence.request_id, "request-1");
517
518 let duplicate = driver
519 .request_cancel(request(address.clone(), "request-2"))
520 .await
521 .expect("duplicate cancellation");
522 assert!(matches!(
523 duplicate.outcome,
524 TurnCancelOutcome::AlreadyRequested(TurnCancellationEvidence { ref request_id, .. })
525 if request_id == "request-1"
526 ));
527
528 let active = ActiveTurnControl::new(host.as_ref(), address.clone())
529 .await
530 .expect("active control");
531 let observed = active
532 .settle_before_commit(host.as_ref(), false)
533 .await
534 .expect("settle")
535 .expect("cancellation won");
536 assert_eq!(observed, evidence);
537 let terminal = TurnTerminal::Committed {
538 outcome: TurnOutcome::Stopped(TurnStop::Cancelled),
539 cancellation: Some(observed),
540 session_revision: Some(7),
541 };
542 active
543 .publish_terminal(host.as_ref(), &terminal)
544 .await
545 .expect("publish terminal");
546 let attached = driver
547 .await_terminal(&address)
548 .await
549 .expect("attach terminal");
550 assert!(matches!(
551 attached,
552 TurnTerminal::Committed {
553 outcome: TurnOutcome::Stopped(TurnStop::Cancelled),
554 cancellation: Some(_),
555 session_revision: Some(7),
556 }
557 ));
558 }
559
560 #[tokio::test]
561 async fn concurrent_completion_seal_vs_cancel_is_first_writer_wins() {
562 let host = Arc::new(InlineEffectHost::default());
563 let driver = TurnWorkDriver::new(host.clone());
564 let address = address("race");
565 let active = ActiveTurnControl::new(host.as_ref(), address.clone())
566 .await
567 .expect("active control");
568
569 let (seal, cancel) = tokio::join!(
570 active.settle_before_commit(host.as_ref(), false),
571 driver.request_cancel(request(address, "race-request")),
572 );
573 match (seal.expect("seal"), cancel.expect("cancel").outcome) {
574 (None, TurnCancelOutcome::CompletionWonRace) => {}
575 (Some(evidence), TurnCancelOutcome::Requested(requested)) => {
576 assert_eq!(evidence, requested);
577 }
578 other => panic!("inconsistent gate race result: {other:?}"),
579 }
580 }
581
582 #[tokio::test]
583 async fn recovered_owner_observes_pending_cancel_after_control_recreation() {
584 let host = Arc::new(InlineEffectHost::default());
585 let driver = TurnWorkDriver::new(host.clone());
586 let address = address("replay");
587 let requested = driver
588 .request_cancel(request(address.clone(), "request-before-replay"))
589 .await
590 .expect("request cancellation");
591 let expected = match requested.outcome {
592 TurnCancelOutcome::Requested(evidence) => evidence,
593 other => panic!("expected requested, got {other:?}"),
594 };
595
596 let recovered = ActiveTurnControl::new(host.as_ref(), address)
597 .await
598 .expect("recreate active control under the recovered owner");
599 let observed = recovered
600 .settle_before_commit(host.as_ref(), false)
601 .await
602 .expect("settle recovered turn")
603 .expect("pending cancellation survives owner loss");
604 assert_eq!(observed, expected);
605 }
606
607 #[tokio::test]
608 async fn turn_control_is_exact_scope_and_excluded_from_wait_cancel_sweep() {
609 let host = Arc::new(InlineEffectHost::default());
610 let driver = TurnWorkDriver::new(host.clone());
611 let address_a = address("scope");
612 let address_b = TurnAddress::new(&address_a.session_id, "turn-b");
613 let address_future = TurnAddress::new(&address_a.session_id, "turn-future");
614
615 driver
616 .request_cancel(request(address_a.clone(), "request-a"))
617 .await
618 .expect("cancel a");
619
620 let tool_key = host
621 .await_event_key(
622 &ExecutionScope::turn(&address_a.session_id, "tool-turn"),
623 AwaitEventWaitIdentity::tool_completion("tool-call"),
624 )
625 .await
626 .expect("tool key");
627 let tool_host = host.clone();
628 let tool_wait = tokio::spawn(async move {
629 tool_host
630 .await_await_event(&tool_key, CancellationToken::new(), None)
631 .await
632 });
633 tokio::task::yield_now().await;
634 host.cancel_await_events_for_session(&address_a.session_id)
635 .await
636 .expect("cancel durable waits");
637 assert!(matches!(
638 tool_wait
639 .await
640 .expect("tool wait task")
641 .expect("tool resolution"),
642 Resolution::Cancelled
643 ));
644
645 assert!(matches!(
646 driver
647 .request_cancel(request(address_a.clone(), "request-a-duplicate"))
648 .await
649 .expect("duplicate a")
650 .outcome,
651 TurnCancelOutcome::AlreadyRequested(_)
652 ));
653 assert!(matches!(
654 driver
655 .request_cancel(request(address_b, "request-b"))
656 .await
657 .expect("cancel b")
658 .outcome,
659 TurnCancelOutcome::Requested(_)
660 ));
661 assert!(matches!(
662 driver
663 .request_cancel(request(address_future, "request-future"))
664 .await
665 .expect("cancel future")
666 .outcome,
667 TurnCancelOutcome::Requested(_)
668 ));
669 }
670
671 #[tokio::test]
672 async fn session_deletion_revokes_control_promises() {
673 let host = Arc::new(InlineEffectHost::default());
674 let driver = TurnWorkDriver::new(host.clone());
675 let address = address("revoke");
676 host.revoke_await_events_for_session(&address.session_id)
677 .await
678 .expect("revoke session");
679 assert!(matches!(
680 driver
681 .request_cancel(request(address, "request-after-delete"))
682 .await
683 .expect("revoked outcome")
684 .outcome,
685 TurnCancelOutcome::UnknownOrRevoked
686 ));
687 }
688
689 #[tokio::test]
690 async fn terminal_attachment_timeout_does_not_poison_later_publication() {
691 let host = Arc::new(InlineEffectHost::default());
692 let driver = TurnWorkDriver::new(host.clone());
693 let address = address("terminal-timeout");
694 let error = driver
695 .await_terminal_with_timeout(&address, Duration::from_millis(1))
696 .await
697 .expect_err("unpublished terminal must time out");
698 assert_eq!(error.code.as_str(), "turn_terminal_await_timeout");
699
700 let active = ActiveTurnControl::new(host.as_ref(), address.clone())
701 .await
702 .expect("active control after timed-out attach");
703 active
704 .settle_before_commit(host.as_ref(), false)
705 .await
706 .expect("seal after timed-out attach");
707 active
708 .publish_terminal(
709 host.as_ref(),
710 &TurnTerminal::Committed {
711 outcome: TurnOutcome::Finished(TurnFinish::AssistantMessage {
712 text: "done".to_string(),
713 }),
714 cancellation: None,
715 session_revision: None,
716 },
717 )
718 .await
719 .expect("publish after timed-out attach");
720 assert!(matches!(
721 driver.await_terminal(&address).await.expect("late attach"),
722 TurnTerminal::Committed {
723 outcome: TurnOutcome::Finished(_),
724 cancellation: None,
725 ..
726 }
727 ));
728 }
729
730 #[test]
731 fn local_cancel_origin_hint_preserves_first_origin() {
732 let hint = TurnCancelOriginHint::default();
733 hint.set(Some("shutdown".to_string()));
734 hint.set(Some("user".to_string()));
735
736 assert_eq!(hint.get().as_deref(), Some("shutdown"));
737 }
738
739 #[test]
740 fn local_cancel_origin_hint_preserves_explicit_absence() {
741 let hint = TurnCancelOriginHint::default();
742 hint.set(None);
743 hint.set(Some("user".to_string()));
744
745 assert_eq!(hint.get(), None);
746 }
747
748 #[test]
749 fn terminal_success_has_no_cancellation_evidence() {
750 let terminal = TurnTerminal::Committed {
751 outcome: TurnOutcome::Finished(TurnFinish::AssistantMessage {
752 text: "done".to_string(),
753 }),
754 cancellation: None,
755 session_revision: None,
756 };
757 let encoded = terminal_resolution(&terminal).expect("encode terminal");
758 assert!(matches!(encoded, Resolution::Ok(_)));
759 }
760}