use async_trait::async_trait;
use dashmap::DashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::oneshot;
use tokio_util::sync::CancellationToken;
#[derive(Debug, Clone)]
pub struct InboundApprovalMessage {
pub channel: String,
pub account_id: String,
pub sender_id: String,
pub body: String,
pub received_at: i64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ApprovalCommand {
Approve {
patch_id: String,
},
Reject {
patch_id: String,
reason: Option<String>,
},
}
impl ApprovalCommand {
pub fn patch_id(&self) -> &str {
match self {
ApprovalCommand::Approve { patch_id } | ApprovalCommand::Reject { patch_id, .. } => {
patch_id.as_str()
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ApprovalDecision {
Approved,
Rejected { reason: Option<String> },
Expired,
}
pub struct PendingApproval {
pub patch_id: String,
pub binding_id: String,
pub agent_id: String,
pub channel: String,
pub account_id: String,
pub sender_id: String,
pub created_at: i64,
pub expires_at: i64,
}
struct PendingEntry {
patch: PendingApproval,
responder: oneshot::Sender<ApprovalDecision>,
}
#[derive(Debug, Clone)]
pub struct ApprovalCorrelatorConfig {
pub default_timeout: Duration,
pub reaper_interval: Duration,
}
impl Default for ApprovalCorrelatorConfig {
fn default() -> Self {
Self {
default_timeout: Duration::from_secs(86_400),
reaper_interval: Duration::from_secs(60),
}
}
}
pub struct ApprovalCorrelator {
pending: DashMap<String, PendingEntry>,
config: ApprovalCorrelatorConfig,
cancel: CancellationToken,
}
impl ApprovalCorrelator {
pub fn new(config: ApprovalCorrelatorConfig) -> Arc<Self> {
Arc::new(Self {
pending: DashMap::new(),
config,
cancel: CancellationToken::new(),
})
}
pub fn park(&self, patch: PendingApproval) -> oneshot::Receiver<ApprovalDecision> {
let (tx, rx) = oneshot::channel();
self.pending.insert(
patch.patch_id.clone(),
PendingEntry {
patch,
responder: tx,
},
);
rx
}
pub fn pending_count(&self) -> usize {
self.pending.len()
}
pub fn cancel_patch(&self, patch_id: &str) -> bool {
self.pending.remove(patch_id).is_some()
}
pub fn spawn_workers(self: &Arc<Self>, source: Arc<dyn ApprovalSource>) {
let me = Arc::clone(self);
let cancel = self.cancel.clone();
let src = Arc::clone(&source);
tokio::spawn(async move {
loop {
tokio::select! {
_ = cancel.cancelled() => break,
msg = src.next_message() => match msg {
Some(m) => me.on_inbound(m),
None => break,
}
}
}
});
let me = Arc::clone(self);
let cancel = self.cancel.clone();
let interval = self.config.reaper_interval;
tokio::spawn(async move {
loop {
tokio::select! {
_ = cancel.cancelled() => break,
_ = tokio::time::sleep(interval) => {
me.reap_expired();
}
}
}
});
}
pub fn on_inbound(&self, msg: InboundApprovalMessage) {
let cmd = match parse_approval_command(&msg.body) {
Some(c) => c,
None => return,
};
let patch_id = cmd.patch_id().to_string();
let Some((_, entry)) = self.pending.remove(&patch_id) else {
tracing::debug!(
target: "config::approval",
patch_id = %patch_id,
"[config] inbound matched no pending entry"
);
return;
};
if entry.patch.channel != msg.channel || entry.patch.account_id != msg.account_id {
tracing::warn!(
target: "config::approval_forgery_rejected",
patch_id = %patch_id,
expected_channel = %entry.patch.channel,
expected_account = %entry.patch.account_id,
got_channel = %msg.channel,
got_account = %msg.account_id,
"[config] approval came from wrong binding — discarded"
);
self.pending.insert(entry.patch.patch_id.clone(), entry);
return;
}
let decision = match cmd {
ApprovalCommand::Approve { .. } => ApprovalDecision::Approved,
ApprovalCommand::Reject { reason, .. } => ApprovalDecision::Rejected { reason },
};
let _ = entry.responder.send(decision);
}
fn reap_expired(&self) {
let now = chrono::Utc::now().timestamp();
let mut to_drop: Vec<String> = Vec::new();
for kv in self.pending.iter() {
if kv.value().patch.expires_at <= now {
to_drop.push(kv.key().clone());
}
}
for id in to_drop {
if let Some((_, entry)) = self.pending.remove(&id) {
tracing::info!(
target: "config::approval_expired",
patch_id = %id,
"[config] approval expired"
);
let _ = entry.responder.send(ApprovalDecision::Expired);
}
}
}
pub fn shutdown(&self) {
self.cancel.cancel();
}
}
#[async_trait]
pub trait ApprovalSource: Send + Sync {
async fn next_message(&self) -> Option<InboundApprovalMessage>;
}
pub struct MockApprovalSource {
queue: tokio::sync::Mutex<std::collections::VecDeque<InboundApprovalMessage>>,
notify: tokio::sync::Notify,
closed: std::sync::atomic::AtomicBool,
}
impl Default for MockApprovalSource {
fn default() -> Self {
Self::new()
}
}
impl MockApprovalSource {
pub fn new() -> Self {
Self {
queue: tokio::sync::Mutex::new(std::collections::VecDeque::new()),
notify: tokio::sync::Notify::new(),
closed: std::sync::atomic::AtomicBool::new(false),
}
}
pub async fn inject(&self, msg: InboundApprovalMessage) {
self.queue.lock().await.push_back(msg);
self.notify.notify_one();
}
pub fn close(&self) {
self.closed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.notify.notify_waiters();
}
}
#[async_trait]
impl ApprovalSource for MockApprovalSource {
async fn next_message(&self) -> Option<InboundApprovalMessage> {
loop {
if let Some(m) = self.queue.lock().await.pop_front() {
return Some(m);
}
if self.closed.load(std::sync::atomic::Ordering::Relaxed) {
return None;
}
self.notify.notified().await;
}
}
}
pub fn parse_approval_command(body: &str) -> Option<ApprovalCommand> {
use std::sync::OnceLock;
static RE: OnceLock<regex::Regex> = OnceLock::new();
let re = RE.get_or_init(|| {
regex::Regex::new(
r"(?x)
^\s*
\[
config-
(?P<verb>approve|reject)
\s+
patch_id=
(?P<id>[0-9A-HJKMNP-TV-Z]+)
(?:\s+reason=(?P<reason>.*?))?
\]
\s*$
",
)
.expect("approval-command regex must compile")
});
let caps = re.captures(body.trim())?;
let verb = caps.name("verb")?.as_str();
let id = caps.name("id")?.as_str().to_string();
match verb {
"approve" => Some(ApprovalCommand::Approve { patch_id: id }),
"reject" => {
let reason = caps.name("reason").map(|m| m.as_str().trim().to_string());
Some(ApprovalCommand::Reject {
patch_id: id,
reason,
})
}
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::Ordering;
fn fixture_pending(patch_id: &str) -> PendingApproval {
PendingApproval {
patch_id: patch_id.into(),
binding_id: "whatsapp:default".into(),
agent_id: "cody".into(),
channel: "whatsapp".into(),
account_id: "default".into(),
sender_id: "5511".into(),
created_at: 0,
expires_at: i64::MAX,
}
}
fn fixture_inbound(_patch_id: &str, body: &str) -> InboundApprovalMessage {
InboundApprovalMessage {
channel: "whatsapp".into(),
account_id: "default".into(),
sender_id: "5511".into(),
body: body.into(),
received_at: 0,
}
}
#[test]
fn parse_approve_minimal() {
let cmd = parse_approval_command("[config-approve patch_id=01J7HVK9MWXYZ]").unwrap();
assert_eq!(
cmd,
ApprovalCommand::Approve {
patch_id: "01J7HVK9MWXYZ".into()
}
);
}
#[test]
fn parse_reject_with_reason() {
let cmd = parse_approval_command(
"[config-reject patch_id=01J7HVK9MWXYZ reason=use sonnet por costo]",
)
.unwrap();
assert_eq!(
cmd,
ApprovalCommand::Reject {
patch_id: "01J7HVK9MWXYZ".into(),
reason: Some("use sonnet por costo".into())
}
);
}
#[test]
fn parse_reject_without_reason() {
let cmd = parse_approval_command("[config-reject patch_id=01J7HVK9MWXYZ]").unwrap();
assert_eq!(
cmd,
ApprovalCommand::Reject {
patch_id: "01J7HVK9MWXYZ".into(),
reason: None
}
);
}
#[test]
fn parse_ignores_garbage() {
assert!(parse_approval_command("hello world").is_none());
assert!(parse_approval_command("").is_none());
assert!(parse_approval_command("[config-approve]").is_none());
assert!(parse_approval_command("[config-approve patch_id=]").is_none());
assert!(parse_approval_command("[config-approve patch_id=foo]").is_none());
}
#[test]
fn parse_strict_anchors_full_message() {
assert!(
parse_approval_command("let's go [config-approve patch_id=01J7HVK9MWXYZ] thanks")
.is_none()
);
}
#[tokio::test]
async fn correlator_resolves_pending_on_match() {
let c = ApprovalCorrelator::new(ApprovalCorrelatorConfig::default());
let rx = c.park(fixture_pending("01J7HVK9MWXYZ"));
c.on_inbound(fixture_inbound(
"01J7HVK9MWXYZ",
"[config-approve patch_id=01J7HVK9MWXYZ]",
));
let decision = rx.await.unwrap();
assert_eq!(decision, ApprovalDecision::Approved);
assert_eq!(c.pending_count(), 0);
}
#[tokio::test]
async fn correlator_resolves_with_reason_on_reject() {
let c = ApprovalCorrelator::new(ApprovalCorrelatorConfig::default());
let rx = c.park(fixture_pending("01J7HVK9MWXYZ"));
c.on_inbound(fixture_inbound(
"01J7HVK9MWXYZ",
"[config-reject patch_id=01J7HVK9MWXYZ reason=cost]",
));
match rx.await.unwrap() {
ApprovalDecision::Rejected { reason } => assert_eq!(reason.as_deref(), Some("cost")),
other => panic!("expected Rejected, got {other:?}"),
}
}
#[tokio::test]
async fn correlator_rejects_cross_binding_message() {
let c = ApprovalCorrelator::new(ApprovalCorrelatorConfig::default());
let _rx = c.park(fixture_pending("01J7HVK9MWXYZ"));
let bad_binding = InboundApprovalMessage {
channel: "telegram".into(),
account_id: "other".into(),
sender_id: "abc".into(),
body: "[config-approve patch_id=01J7HVK9MWXYZ]".into(),
received_at: 0,
};
c.on_inbound(bad_binding);
assert_eq!(c.pending_count(), 1);
}
#[tokio::test]
async fn correlator_no_pending_entry_logs_and_returns() {
let c = ApprovalCorrelator::new(ApprovalCorrelatorConfig::default());
c.on_inbound(fixture_inbound(
"01J7UNKNOWNPATCHID",
"[config-approve patch_id=01J7UNKNOWNPATCHID]",
));
assert_eq!(c.pending_count(), 0);
}
#[tokio::test(start_paused = true)]
async fn correlator_expires_pending_after_timeout() {
let cfg = ApprovalCorrelatorConfig {
default_timeout: Duration::from_secs(1),
reaper_interval: Duration::from_millis(50),
};
let c = ApprovalCorrelator::new(cfg);
let mut p = fixture_pending("01J7HVK9MWXYZ");
p.expires_at = chrono::Utc::now().timestamp() - 1; let rx = c.park(p);
c.reap_expired();
let decision = rx.await.unwrap();
assert_eq!(decision, ApprovalDecision::Expired);
assert_eq!(c.pending_count(), 0);
}
#[tokio::test]
async fn mock_source_round_trips_messages() {
let src = Arc::new(MockApprovalSource::new());
let inbound = fixture_inbound("01J7HVK9MWXYZ", "[config-approve patch_id=01J7HVK9MWXYZ]");
src.inject(inbound.clone()).await;
let got = src.next_message().await.unwrap();
assert_eq!(got.body, inbound.body);
src.close();
assert!(src.next_message().await.is_none());
assert!(src.closed.load(Ordering::Relaxed));
}
}