Skip to main content

toolkit_contract/
policy.rs

1//! Per-call policy stack for cross-cutting concerns.
2//!
3//! Policies run hooks before and after the transport call. The
4//! [`PolicyStack`] composes an ordered list of [`Policy`] implementations
5//! and drives execution through them.
6
7use async_trait::async_trait;
8use std::future::Future;
9use std::sync::Arc;
10
11use crate::error::ContractError;
12use crate::ir::contract::{Idempotency, MethodKind};
13
14/// Context passed to policy hooks for each contract call.
15pub struct PolicyContext {
16    /// Contract name being invoked.
17    pub service: &'static str,
18    /// Method name being invoked.
19    pub method: &'static str,
20    /// Idempotency classification (used for retry decisions).
21    pub idempotency: Idempotency,
22    /// Whether the method is unary or streaming.
23    pub kind: MethodKind,
24}
25
26/// A policy that can intercept contract calls before and after transport.
27///
28/// Implement this trait to add cross-cutting concerns such as tracing,
29/// metrics, or authorization checks.
30#[async_trait]
31pub trait Policy: Send + Sync {
32    /// Called before the transport call is made.
33    ///
34    /// # Errors
35    ///
36    /// Return an error to short-circuit the call (subsequent policies
37    /// and the transport call will be skipped).
38    async fn on_request(&self, ctx: &PolicyContext) -> Result<(), ContractError>;
39
40    /// Called after the transport call completes.
41    ///
42    /// # Errors
43    ///
44    /// Returning an error replaces the original transport result.
45    async fn on_response(&self, ctx: &PolicyContext, success: bool) -> Result<(), ContractError>;
46}
47
48/// Ordered list of policies applied to every contract call.
49///
50/// Policies run `on_request` in insertion order and `on_response` in
51/// reverse order (like middleware stacks).
52pub struct PolicyStack {
53    policies: Vec<Arc<dyn Policy>>,
54}
55
56impl PolicyStack {
57    /// Create an empty policy stack.
58    #[must_use]
59    pub fn new() -> Self {
60        Self {
61            policies: Vec::new(),
62        }
63    }
64
65    /// Append a policy to the end of the stack.
66    pub fn push(&mut self, policy: Arc<dyn Policy>) {
67        self.policies.push(policy);
68    }
69
70    /// Execute a contract call through the policy stack.
71    ///
72    /// 1. Runs `on_request` for each policy in order.
73    /// 2. Invokes the transport closure `f`.
74    /// 3. Runs `on_response` for each policy in reverse order.
75    ///
76    /// # Errors
77    ///
78    /// Returns the first error from any policy hook, or the transport
79    /// error if the call itself fails.
80    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        // Track the highest index for which `on_request` succeeded so that
91        // we can symmetrically unwind those policies on the error path.
92        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            // Unwind already-succeeded policies in reverse with success=false.
106            if let Some(top) = last_ok {
107                for policy in self.policies[..=top].iter().rev() {
108                    // Best-effort cleanup unwind: original request error takes precedence,
109                    // so any error from on_response here is intentionally dropped.
110                    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
135/// Policy that emits `tracing` spans and log events for each contract call.
136pub 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}