use std::fmt;
use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::panic::catch_unwind;
use std::sync::Arc;
use std::sync::Mutex;
use chrono::DateTime;
use chrono::Utc;
use ferrin_core::generate_text::ApprovalContext;
use ferrin_core::generate_text::ApprovalPolicy;
use ferrin_core::generate_text::ApprovalStatus;
use ferrin_core::generate_text::ParsedToolCall;
use ferrin_spec::BoxFuture;
use ferrin_spec::JsonValue;
use ferrin_spec::ToolCallId;
use ferrin_spec::ToolName;
use serde::Deserialize;
use serde::Serialize;
use tokio::task::JoinSet;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum Enforcement {
#[default]
Observe,
Enforce,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PolicyDecisionToolCall {
pub tool_name: ToolName,
pub tool_call_id: ToolCallId,
pub input: JsonValue,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PolicyDecisionEvent {
pub tool_call: PolicyDecisionToolCall,
pub decision: ApprovalStatus,
pub enforced: bool,
pub effective: ApprovalStatus,
pub timestamp: DateTime<Utc>,
}
pub type OnDecisionFn = Arc<dyn Fn(PolicyDecisionEvent) -> BoxFuture<'static, ()> + Send + Sync>;
pub type OnDecisionSyncFn = Arc<dyn Fn(&ParsedToolCall, Option<&ApprovalStatus>) + Send + Sync>;
enum Observer {
Async(OnDecisionFn),
Sync(OnDecisionSyncFn),
}
pub struct Shadow<P> {
inner: P,
enforcement: Enforcement,
on_decision: Option<Observer>,
audit_tasks: Mutex<JoinSet<()>>,
}
pub fn shadow<P: ApprovalPolicy>(policy: P) -> Shadow<P> {
Shadow {
inner: policy,
enforcement: Enforcement::Observe,
on_decision: None,
audit_tasks: Mutex::new(JoinSet::new()),
}
}
impl<P> Shadow<P> {
#[must_use]
pub fn enforcement(mut self, enforcement: Enforcement) -> Self {
self.enforcement = enforcement;
self
}
#[must_use]
pub fn on_decision<F, Fut>(mut self, observer: F) -> Self
where
F: Fn(PolicyDecisionEvent) -> Fut + Send + Sync + 'static,
Fut: Future + Send + 'static,
{
let observer = Arc::new(observer);
self.on_decision = Some(Observer::Async(Arc::new(move |event| {
let observer = Arc::clone(&observer);
Box::pin(async move {
let _ = observer(event).await;
})
})));
self
}
#[must_use]
pub fn on_decision_sync(
mut self,
observer: impl Fn(&ParsedToolCall, Option<&ApprovalStatus>) + Send + Sync + 'static,
) -> Self {
self.on_decision = Some(Observer::Sync(Arc::new(observer)));
self
}
pub async fn flush_decisions(&self) {
let mut pending = {
let mut tasks = self
.audit_tasks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
std::mem::take(&mut *tasks)
};
while pending.join_next().await.is_some() {}
}
fn report(
&self,
call: &ParsedToolCall,
status: Option<&ApprovalStatus>,
effective: &ApprovalStatus,
) {
match &self.on_decision {
Some(Observer::Async(observer)) => {
let Ok(runtime) = tokio::runtime::Handle::try_current() else {
return;
};
let event = PolicyDecisionEvent {
tool_call: PolicyDecisionToolCall {
tool_name: call.tool_name.clone(),
tool_call_id: call.tool_call_id.clone(),
input: call.input.clone(),
},
decision: status.cloned().unwrap_or(ApprovalStatus::NotApplicable),
enforced: self.enforcement == Enforcement::Enforce,
effective: effective.clone(),
timestamp: Utc::now(),
};
let observer = Arc::clone(observer);
let mut tasks = self
.audit_tasks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
while tasks.try_join_next().is_some() {}
tasks.spawn_on(async move { observer(event).await }, &runtime);
}
Some(Observer::Sync(observer)) => {
let _ = catch_unwind(AssertUnwindSafe(|| observer(call, status)));
}
None => {}
}
}
}
impl<P> fmt::Debug for Shadow<P> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Shadow")
.field("enforcement", &self.enforcement)
.field("on_decision", &self.on_decision.is_some())
.finish_non_exhaustive()
}
}
impl<P: ApprovalPolicy> ApprovalPolicy for Shadow<P> {
fn resolve<'a>(
&'a self,
call: &'a ParsedToolCall,
ctx: ApprovalContext<'a>,
) -> BoxFuture<'a, Option<ApprovalStatus>> {
Box::pin(async move {
let status = self.inner.resolve(call, ctx).await;
tracing::debug!(
tool = %call.tool_name,
status = crate::diagnostics::status_kind(status.as_ref()),
enforcement = ?self.enforcement,
"shadow policy decision"
);
let effective = match self.enforcement {
Enforcement::Observe => ApprovalStatus::approved(),
Enforcement::Enforce => status.clone().unwrap_or(ApprovalStatus::NotApplicable),
};
self.report(call, status.as_ref(), &effective);
Some(effective)
})
}
}