1use std::fmt;
7use std::future::Future;
8use std::panic::AssertUnwindSafe;
9use std::panic::catch_unwind;
10use std::sync::Arc;
11use std::sync::Mutex;
12
13use chrono::DateTime;
14use chrono::Utc;
15use ferrin_core::generate_text::ApprovalContext;
16use ferrin_core::generate_text::ApprovalPolicy;
17use ferrin_core::generate_text::ApprovalStatus;
18use ferrin_core::generate_text::ParsedToolCall;
19use ferrin_spec::BoxFuture;
20use ferrin_spec::JsonValue;
21use ferrin_spec::ToolCallId;
22use ferrin_spec::ToolName;
23use serde::Deserialize;
24use serde::Serialize;
25use tokio::task::JoinSet;
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
29#[non_exhaustive]
30pub enum Enforcement {
31 #[default]
33 Observe,
34 Enforce,
36}
37
38#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
40pub struct PolicyDecisionToolCall {
41 pub tool_name: ToolName,
43 pub tool_call_id: ToolCallId,
45 pub input: JsonValue,
47}
48
49#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
51pub struct PolicyDecisionEvent {
52 pub tool_call: PolicyDecisionToolCall,
54 pub decision: ApprovalStatus,
56 pub enforced: bool,
58 pub effective: ApprovalStatus,
60 pub timestamp: DateTime<Utc>,
62}
63
64pub type OnDecisionFn = Arc<dyn Fn(PolicyDecisionEvent) -> BoxFuture<'static, ()> + Send + Sync>;
66
67pub type OnDecisionSyncFn = Arc<dyn Fn(&ParsedToolCall, Option<&ApprovalStatus>) + Send + Sync>;
69
70enum Observer {
71 Async(OnDecisionFn),
72 Sync(OnDecisionSyncFn),
73}
74
75pub struct Shadow<P> {
77 inner: P,
78 enforcement: Enforcement,
79 on_decision: Option<Observer>,
80 audit_tasks: Mutex<JoinSet<()>>,
81}
82
83pub fn shadow<P: ApprovalPolicy>(policy: P) -> Shadow<P> {
88 Shadow {
89 inner: policy,
90 enforcement: Enforcement::Observe,
91 on_decision: None,
92 audit_tasks: Mutex::new(JoinSet::new()),
93 }
94}
95
96impl<P> Shadow<P> {
97 #[must_use]
99 pub fn enforcement(mut self, enforcement: Enforcement) -> Self {
100 self.enforcement = enforcement;
101 self
102 }
103
104 #[must_use]
123 pub fn on_decision<F, Fut>(mut self, observer: F) -> Self
124 where
125 F: Fn(PolicyDecisionEvent) -> Fut + Send + Sync + 'static,
126 Fut: Future + Send + 'static,
127 {
128 let observer = Arc::new(observer);
129 self.on_decision = Some(Observer::Async(Arc::new(move |event| {
130 let observer = Arc::clone(&observer);
131 Box::pin(async move {
132 let _ = observer(event).await;
133 })
134 })));
135 self
136 }
137
138 #[must_use]
144 pub fn on_decision_sync(
145 mut self,
146 observer: impl Fn(&ParsedToolCall, Option<&ApprovalStatus>) + Send + Sync + 'static,
147 ) -> Self {
148 self.on_decision = Some(Observer::Sync(Arc::new(observer)));
149 self
150 }
151
152 pub async fn flush_decisions(&self) {
157 let mut pending = {
158 let mut tasks = self
159 .audit_tasks
160 .lock()
161 .unwrap_or_else(std::sync::PoisonError::into_inner);
162 std::mem::take(&mut *tasks)
163 };
164 while pending.join_next().await.is_some() {}
165 }
166
167 fn report(
168 &self,
169 call: &ParsedToolCall,
170 status: Option<&ApprovalStatus>,
171 effective: &ApprovalStatus,
172 ) {
173 match &self.on_decision {
174 Some(Observer::Async(observer)) => {
175 let Ok(runtime) = tokio::runtime::Handle::try_current() else {
176 return;
177 };
178 let event = PolicyDecisionEvent {
179 tool_call: PolicyDecisionToolCall {
180 tool_name: call.tool_name.clone(),
181 tool_call_id: call.tool_call_id.clone(),
182 input: call.input.clone(),
183 },
184 decision: status.cloned().unwrap_or(ApprovalStatus::NotApplicable),
185 enforced: self.enforcement == Enforcement::Enforce,
186 effective: effective.clone(),
187 timestamp: Utc::now(),
188 };
189 let observer = Arc::clone(observer);
190 let mut tasks = self
191 .audit_tasks
192 .lock()
193 .unwrap_or_else(std::sync::PoisonError::into_inner);
194 while tasks.try_join_next().is_some() {}
195 tasks.spawn_on(async move { observer(event).await }, &runtime);
196 }
197 Some(Observer::Sync(observer)) => {
198 let _ = catch_unwind(AssertUnwindSafe(|| observer(call, status)));
199 }
200 None => {}
201 }
202 }
203}
204
205impl<P> fmt::Debug for Shadow<P> {
206 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
207 f.debug_struct("Shadow")
208 .field("enforcement", &self.enforcement)
209 .field("on_decision", &self.on_decision.is_some())
210 .finish_non_exhaustive()
211 }
212}
213
214impl<P: ApprovalPolicy> ApprovalPolicy for Shadow<P> {
215 fn resolve<'a>(
216 &'a self,
217 call: &'a ParsedToolCall,
218 ctx: ApprovalContext<'a>,
219 ) -> BoxFuture<'a, Option<ApprovalStatus>> {
220 Box::pin(async move {
221 let status = self.inner.resolve(call, ctx).await;
222 tracing::debug!(
223 tool = %call.tool_name,
224 status = crate::diagnostics::status_kind(status.as_ref()),
225 enforcement = ?self.enforcement,
226 "shadow policy decision"
227 );
228 let effective = match self.enforcement {
229 Enforcement::Observe => ApprovalStatus::approved(),
230 Enforcement::Enforce => status.clone().unwrap_or(ApprovalStatus::NotApplicable),
231 };
232 self.report(call, status.as_ref(), &effective);
233 Some(effective)
234 })
235 }
236}