1use async_trait::async_trait;
8use std::future::Future;
9use std::sync::Arc;
10
11use crate::error::ContractError;
12use crate::ir::contract::{Idempotency, MethodKind};
13
14pub struct PolicyContext {
16 pub service: &'static str,
18 pub method: &'static str,
20 pub idempotency: Idempotency,
22 pub kind: MethodKind,
24}
25
26#[async_trait]
31pub trait Policy: Send + Sync {
32 async fn on_request(&self, ctx: &PolicyContext) -> Result<(), ContractError>;
39
40 async fn on_response(&self, ctx: &PolicyContext, success: bool) -> Result<(), ContractError>;
46}
47
48pub struct PolicyStack {
53 policies: Vec<Arc<dyn Policy>>,
54}
55
56impl PolicyStack {
57 #[must_use]
59 pub fn new() -> Self {
60 Self {
61 policies: Vec::new(),
62 }
63 }
64
65 pub fn push(&mut self, policy: Arc<dyn Policy>) {
67 self.policies.push(policy);
68 }
69
70 pub async fn execute<F, Fut, T, E>(
81 &self,
82 ctx: &PolicyContext,
83 f: F,
84 map_policy_err: fn(ContractError) -> E,
85 ) -> Result<T, E>
86 where
87 F: FnOnce() -> Fut,
88 Fut: Future<Output = Result<T, E>>,
89 {
90 let mut last_ok: Option<usize> = None;
93 let mut request_err: Option<ContractError> = None;
94 for (idx, policy) in self.policies.iter().enumerate() {
95 match policy.on_request(ctx).await {
96 Ok(()) => last_ok = Some(idx),
97 Err(e) => {
98 request_err = Some(e);
99 break;
100 }
101 }
102 }
103
104 if let Some(e) = request_err {
105 if let Some(top) = last_ok {
107 for policy in self.policies[..=top].iter().rev() {
108 drop(policy.on_response(ctx, false).await);
111 }
112 }
113 return Err(map_policy_err(e));
114 }
115
116 let result = f().await;
117 let success = result.is_ok();
118
119 for policy in self.policies.iter().rev() {
120 if let Err(e) = policy.on_response(ctx, success).await {
121 return Err(map_policy_err(e));
122 }
123 }
124
125 result
126 }
127}
128
129impl Default for PolicyStack {
130 fn default() -> Self {
131 Self::new()
132 }
133}
134
135pub struct TracingPolicy;
137
138#[async_trait]
139impl Policy for TracingPolicy {
140 async fn on_request(&self, ctx: &PolicyContext) -> Result<(), ContractError> {
141 tracing::info!(
142 service = ctx.service,
143 method = ctx.method,
144 idempotency = ?ctx.idempotency,
145 kind = ?ctx.kind,
146 "contract call started"
147 );
148 Ok(())
149 }
150
151 async fn on_response(&self, ctx: &PolicyContext, success: bool) -> Result<(), ContractError> {
152 if success {
153 tracing::info!(
154 service = ctx.service,
155 method = ctx.method,
156 "contract call succeeded"
157 );
158 } else {
159 tracing::warn!(
160 service = ctx.service,
161 method = ctx.method,
162 "contract call failed"
163 );
164 }
165 Ok(())
166 }
167}
168
169#[cfg(test)]
170#[cfg_attr(coverage_nightly, coverage(off))]
171#[allow(clippy::unwrap_used)]
172mod tests {
173 use super::*;
174 use std::sync::atomic::{AtomicUsize, Ordering};
175
176 struct OrderRecorder {
177 id: usize,
178 log: Arc<parking_lot::Mutex<Vec<String>>>,
179 }
180
181 #[async_trait]
182 impl Policy for OrderRecorder {
183 async fn on_request(&self, _ctx: &PolicyContext) -> Result<(), ContractError> {
184 self.log.lock().push(format!("on_request:{}", self.id));
185 Ok(())
186 }
187
188 async fn on_response(
189 &self,
190 _ctx: &PolicyContext,
191 success: bool,
192 ) -> Result<(), ContractError> {
193 self.log
194 .lock()
195 .push(format!("on_response:{}:{success}", self.id));
196 Ok(())
197 }
198 }
199
200 fn test_ctx() -> PolicyContext {
201 PolicyContext {
202 service: "TestService",
203 method: "test_method",
204 idempotency: Idempotency::SafeRead,
205 kind: MethodKind::Unary,
206 }
207 }
208
209 #[tokio::test]
210 async fn policy_stack_calls_in_order() {
211 let log: Arc<parking_lot::Mutex<Vec<String>>> =
212 Arc::new(parking_lot::Mutex::new(Vec::new()));
213
214 let mut stack = PolicyStack::new();
215 stack.push(Arc::new(OrderRecorder {
216 id: 1,
217 log: Arc::clone(&log),
218 }));
219 stack.push(Arc::new(OrderRecorder {
220 id: 2,
221 log: Arc::clone(&log),
222 }));
223
224 let ctx = test_ctx();
225 let call_count = Arc::new(AtomicUsize::new(0));
226 let call_count_inner = Arc::clone(&call_count);
227
228 let result: Result<&str, ContractError> = stack
229 .execute(
230 &ctx,
231 || async move {
232 call_count_inner.fetch_add(1, Ordering::Relaxed);
233 Ok("done")
234 },
235 std::convert::identity,
236 )
237 .await;
238
239 assert_eq!(result.unwrap(), "done");
240 assert_eq!(call_count.load(Ordering::Relaxed), 1);
241
242 let entries = log.lock().clone();
243 assert_eq!(
244 entries,
245 vec![
246 "on_request:1",
247 "on_request:2",
248 "on_response:2:true",
249 "on_response:1:true",
250 ]
251 );
252 }
253
254 struct FailPolicy;
255
256 #[async_trait]
257 impl Policy for FailPolicy {
258 async fn on_request(&self, _ctx: &PolicyContext) -> Result<(), ContractError> {
259 Err(ContractError::Validation("blocked by policy".to_owned()))
260 }
261
262 async fn on_response(
263 &self,
264 _ctx: &PolicyContext,
265 _success: bool,
266 ) -> Result<(), ContractError> {
267 Ok(())
268 }
269 }
270
271 #[tokio::test]
272 async fn policy_stack_short_circuits_on_request_error() {
273 let log: Arc<parking_lot::Mutex<Vec<String>>> =
274 Arc::new(parking_lot::Mutex::new(Vec::new()));
275
276 let mut stack = PolicyStack::new();
277 stack.push(Arc::new(FailPolicy));
278 stack.push(Arc::new(OrderRecorder {
279 id: 2,
280 log: Arc::clone(&log),
281 }));
282
283 let ctx = test_ctx();
284 let result: Result<&str, ContractError> = stack
285 .execute(
286 &ctx,
287 || async { Ok("should not run") },
288 std::convert::identity,
289 )
290 .await;
291
292 assert!(result.is_err());
293 let entries = log.lock().clone();
294 assert!(entries.is_empty());
295 }
296
297 struct RecordCleanupPolicy {
298 cleaned: Arc<std::sync::atomic::AtomicBool>,
299 }
300
301 #[async_trait]
302 impl Policy for RecordCleanupPolicy {
303 async fn on_request(&self, _ctx: &PolicyContext) -> Result<(), ContractError> {
304 Ok(())
305 }
306
307 async fn on_response(
308 &self,
309 _ctx: &PolicyContext,
310 _success: bool,
311 ) -> Result<(), ContractError> {
312 self.cleaned
313 .store(true, std::sync::atomic::Ordering::SeqCst);
314 Ok(())
315 }
316 }
317
318 #[tokio::test]
319 async fn on_request_error_invokes_on_response_for_succeeded_policies() {
320 let cleaned = Arc::new(std::sync::atomic::AtomicBool::new(false));
321 let mut stack = PolicyStack::new();
322 stack.push(Arc::new(RecordCleanupPolicy {
323 cleaned: Arc::clone(&cleaned),
324 }));
325 stack.push(Arc::new(FailPolicy));
326
327 let ctx = test_ctx();
328 let result: Result<&str, ContractError> = stack
329 .execute(
330 &ctx,
331 || async { Ok("should not run") },
332 std::convert::identity,
333 )
334 .await;
335
336 assert!(result.is_err());
337 assert!(
338 cleaned.load(std::sync::atomic::Ordering::SeqCst),
339 "expected first policy's on_response to fire after second policy's on_request failed"
340 );
341 }
342
343 #[tokio::test]
344 async fn tracing_policy_does_not_error() {
345 let policy = TracingPolicy;
346 let ctx = test_ctx();
347
348 policy.on_request(&ctx).await.unwrap();
349 policy.on_response(&ctx, true).await.unwrap();
350 policy.on_response(&ctx, false).await.unwrap();
351 }
352}