use std::time::Duration;
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use tokio::sync::oneshot;
use tracing::{debug, info, warn};
use crate::core::task::{TaskId, TaskRisk};
const DEFAULT_APPROVAL_TIMEOUT: Duration = Duration::from_secs(300);
const MAX_PENDING_APPROVALS: usize = 1_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ApprovalDecision {
Approved,
Rejected,
}
#[derive(Debug, Clone, Serialize)]
#[non_exhaustive]
pub struct ApprovalRequest {
pub task_id: TaskId,
pub output: String,
pub risk: TaskRisk,
pub agent: Option<String>,
}
pub struct ApprovalGate {
pending: DashMap<TaskId, oneshot::Sender<ApprovalDecision>>,
timeout: Duration,
}
impl ApprovalGate {
pub fn new() -> Self {
Self {
pending: DashMap::new(),
timeout: DEFAULT_APPROVAL_TIMEOUT,
}
}
pub fn with_timeout(timeout: Duration) -> Self {
Self {
pending: DashMap::new(),
timeout,
}
}
#[must_use]
pub fn requires_approval(risk: TaskRisk, require_medium: bool) -> bool {
match risk {
TaskRisk::High => true,
TaskRisk::Medium => require_medium,
TaskRisk::Low => false,
}
}
pub fn request_approval(&self, task_id: TaskId) -> Option<oneshot::Receiver<ApprovalDecision>> {
if self.pending.len() >= MAX_PENDING_APPROVALS {
warn!(
task_id = %task_id,
"approval gate at capacity, auto-rejecting"
);
return None;
}
let (tx, rx) = oneshot::channel();
self.pending.insert(task_id, tx);
debug!(task_id = %task_id, "approval requested");
Some(rx)
}
pub fn submit_decision(&self, task_id: TaskId, decision: ApprovalDecision) -> bool {
if let Some((_, tx)) = self.pending.remove(&task_id) {
info!(task_id = %task_id, decision = ?decision, "approval decision submitted");
tx.send(decision).is_ok()
} else {
warn!(task_id = %task_id, "no pending approval found");
false
}
}
pub async fn wait_for_decision(
&self,
rx: oneshot::Receiver<ApprovalDecision>,
task_id: TaskId,
) -> ApprovalDecision {
match tokio::time::timeout(self.timeout, rx).await {
Ok(Ok(decision)) => decision,
Ok(Err(_)) => {
warn!(task_id = %task_id, "approval channel dropped, rejecting");
ApprovalDecision::Rejected
}
Err(_) => {
warn!(
task_id = %task_id,
timeout_secs = self.timeout.as_secs(),
"approval timed out, rejecting"
);
self.pending.remove(&task_id);
ApprovalDecision::Rejected
}
}
}
#[inline]
#[must_use]
pub fn pending_count(&self) -> usize {
self.pending.len()
}
#[must_use]
pub fn pending_tasks(&self) -> Vec<TaskId> {
self.pending.iter().map(|e| *e.key()).collect()
}
pub fn cancel(&self, task_id: TaskId) -> bool {
self.pending.remove(&task_id).is_some()
}
}
impl Default for ApprovalGate {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use uuid::Uuid;
#[test]
fn requires_approval_high_always() {
assert!(ApprovalGate::requires_approval(TaskRisk::High, false));
assert!(ApprovalGate::requires_approval(TaskRisk::High, true));
}
#[test]
fn requires_approval_medium_configurable() {
assert!(!ApprovalGate::requires_approval(TaskRisk::Medium, false));
assert!(ApprovalGate::requires_approval(TaskRisk::Medium, true));
}
#[test]
fn requires_approval_low_never() {
assert!(!ApprovalGate::requires_approval(TaskRisk::Low, false));
assert!(!ApprovalGate::requires_approval(TaskRisk::Low, true));
}
#[tokio::test]
async fn approval_flow_approved() {
let gate = ApprovalGate::new();
let task_id = Uuid::new_v4();
let rx = gate.request_approval(task_id).unwrap();
assert_eq!(gate.pending_count(), 1);
gate.submit_decision(task_id, ApprovalDecision::Approved);
let decision = gate.wait_for_decision(rx, task_id).await;
assert_eq!(decision, ApprovalDecision::Approved);
assert_eq!(gate.pending_count(), 0);
}
#[tokio::test]
async fn approval_flow_rejected() {
let gate = ApprovalGate::new();
let task_id = Uuid::new_v4();
let rx = gate.request_approval(task_id).unwrap();
gate.submit_decision(task_id, ApprovalDecision::Rejected);
let decision = gate.wait_for_decision(rx, task_id).await;
assert_eq!(decision, ApprovalDecision::Rejected);
}
#[tokio::test]
async fn approval_timeout_rejects() {
let gate = ApprovalGate::with_timeout(Duration::from_millis(10));
let task_id = Uuid::new_v4();
let rx = gate.request_approval(task_id).unwrap();
let decision = gate.wait_for_decision(rx, task_id).await;
assert_eq!(decision, ApprovalDecision::Rejected);
}
#[test]
fn submit_decision_for_unknown_task_returns_false() {
let gate = ApprovalGate::new();
assert!(!gate.submit_decision(Uuid::new_v4(), ApprovalDecision::Approved));
}
#[test]
fn cancel_removes_pending() {
let gate = ApprovalGate::new();
let task_id = Uuid::new_v4();
let _rx = gate.request_approval(task_id).unwrap();
assert_eq!(gate.pending_count(), 1);
assert!(gate.cancel(task_id));
assert_eq!(gate.pending_count(), 0);
}
#[test]
fn pending_tasks_lists_all() {
let gate = ApprovalGate::new();
let id1 = Uuid::new_v4();
let id2 = Uuid::new_v4();
let _rx1 = gate.request_approval(id1).unwrap();
let _rx2 = gate.request_approval(id2).unwrap();
let tasks = gate.pending_tasks();
assert_eq!(tasks.len(), 2);
assert!(tasks.contains(&id1));
assert!(tasks.contains(&id2));
}
}