use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use tokio::sync::Notify;
use super::{HostBridge, HostBridgeInjectionState};
use crate::tool_call_cancellations::{fresh_registry, CancellationRegistry};
pub(super) struct DaemonIdleState {
idle: AtomicBool,
changed: Arc<Notify>,
}
impl Default for DaemonIdleState {
fn default() -> Self {
Self {
idle: AtomicBool::new(false),
changed: Arc::new(Notify::new()),
}
}
}
impl HostBridge {
pub fn set_daemon_idle(&self, idle: bool) {
if self.daemon_idle.idle.swap(idle, Ordering::SeqCst) != idle {
self.daemon_idle.changed.notify_waiters();
}
}
pub fn is_daemon_idle(&self) -> bool {
self.daemon_idle.idle.load(Ordering::SeqCst)
}
pub(crate) fn daemon_idle_notifier(&self) -> Arc<Notify> {
self.daemon_idle.changed.clone()
}
}
#[derive(Clone)]
pub struct HostBridgeControlState {
pub(super) cancelled: Arc<AtomicBool>,
pub(super) cancel_notify: Arc<Notify>,
pub(super) queued_transcript_injections: HostBridgeInjectionState,
pub(super) tool_call_cancellations: Arc<CancellationRegistry>,
}
impl HostBridgeControlState {
pub fn new(
cancelled: Arc<AtomicBool>,
cancel_notify: Arc<Notify>,
queued_transcript_injections: HostBridgeInjectionState,
tool_call_cancellations: Arc<CancellationRegistry>,
) -> Self {
Self {
cancelled,
cancel_notify,
queued_transcript_injections,
tool_call_cancellations,
}
}
pub(super) fn isolated(cancelled: Arc<AtomicBool>) -> Self {
Self::new(
cancelled,
Arc::new(Notify::new()),
HostBridgeInjectionState::default(),
fresh_registry(),
)
}
}
pub(super) fn handle_cancel_tool_call_notification(
registry: &CancellationRegistry,
params: &serde_json::Value,
) {
let session_id = params
.get("sessionId")
.or_else(|| params.get("session_id"))
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
let call_id = params
.get("toolCallId")
.or_else(|| params.get("tool_call_id"))
.or_else(|| params.get("callId"))
.or_else(|| params.get("call_id"))
.and_then(serde_json::Value::as_str)
.unwrap_or_default();
if call_id.is_empty() {
return;
}
let reason = params
.get("reason")
.and_then(serde_json::Value::as_str)
.unwrap_or("host cancelled in-flight tool call")
.to_string();
let inject_reminder = params
.get("injectReminder")
.or_else(|| params.get("inject_reminder"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(true);
registry.cancel(session_id, call_id, reason, inject_reminder);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn targeted_cancel_notification_uses_the_supplied_registry() {
let registry = Arc::new(CancellationRegistry::default());
let (handle, _guard) = registry.register("session", "call", "shell");
handle_cancel_tool_call_notification(
®istry,
&serde_json::json!({
"sessionId": "session",
"toolCallId": "call",
"reason": "host stop",
"injectReminder": false,
}),
);
assert!(handle.is_cancelled());
assert_eq!(handle.reason().as_deref(), Some("host stop"));
}
}