Skip to main content

systemprompt_api/services/gateway/
policy.rs

1//! Resolution and caching of the effective gateway policy.
2//!
3//! [`PolicyResolver`] loads the global policy rows in ascending
4//! `(priority, name)` order and merges them into a single
5//! [`GatewayPolicySpec`] — each non-empty section overrides the previous, so
6//! the highest-priority row wins. The result is cached for a short TTL.
7//!
8//! A DB error is a fault governed by [`QuotaFaultMode`]: under `Open` the
9//! resolver degrades to a permissive policy, which drops quota windows *and*
10//! safety scanning for the request; under `Closed` it returns
11//! [`PolicyUnavailable`] and the request is denied. A malformed spec row is
12//! always skipped — the remaining rows still merge.
13//!
14//! Copyright (c) systemprompt.io — Business Source License 1.1.
15//! See <https://systemprompt.io> for licensing details.
16
17use std::sync::{Arc, RwLock};
18use std::time::{Duration, Instant};
19
20use systemprompt_ai::repository::AiGatewayPolicyRepository;
21use systemprompt_models::services::QuotaFaultMode;
22
23pub use systemprompt_ai::{GatewayPolicySpec, QuotaMode, QuotaWindow, SafetyConfig};
24
25const CACHE_TTL: Duration = Duration::from_secs(60);
26
27#[derive(Debug, thiserror::Error)]
28#[error("gateway policy unavailable: {reason}")]
29pub struct PolicyUnavailable {
30    pub reason: String,
31}
32
33#[derive(Clone)]
34pub struct PolicyResolver {
35    repo: Arc<AiGatewayPolicyRepository>,
36    cache: Arc<RwLock<Option<CachedEntry>>>,
37}
38
39impl std::fmt::Debug for PolicyResolver {
40    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41        f.debug_struct("PolicyResolver").finish()
42    }
43}
44
45#[derive(Clone)]
46struct CachedEntry {
47    spec: GatewayPolicySpec,
48    fetched_at: Instant,
49}
50
51impl PolicyResolver {
52    pub fn from_repository(repo: AiGatewayPolicyRepository) -> Self {
53        Self {
54            repo: Arc::new(repo),
55            cache: Arc::new(RwLock::new(None)),
56        }
57    }
58
59    pub async fn resolve(
60        &self,
61        fault_mode: QuotaFaultMode,
62    ) -> Result<GatewayPolicySpec, PolicyUnavailable> {
63        if let Ok(cache) = self.cache.read()
64            && let Some(entry) = cache.as_ref()
65            && entry.fetched_at.elapsed() < CACHE_TTL
66        {
67            return Ok(entry.spec.clone());
68        }
69
70        let rows = match self.repo.list_for_global().await {
71            Ok(r) => r,
72            Err(e) => {
73                if fault_mode.is_closed() {
74                    tracing::error!(
75                        error = %e,
76                        fault_mode = fault_mode.as_str(),
77                        "Gateway policy read failed; denying the request"
78                    );
79                    return Err(PolicyUnavailable {
80                        reason: e.to_string(),
81                    });
82                }
83                tracing::warn!(
84                    error = %e,
85                    fault_mode = fault_mode.as_str(),
86                    "Gateway policy read failed; falling back to a permissive policy \
87                     (quota windows and safety scanning are not applied)"
88                );
89                return Ok(GatewayPolicySpec::permissive());
90            },
91        };
92
93        let spec = merge(rows);
94        if let Ok(mut cache) = self.cache.write() {
95            *cache = Some(CachedEntry {
96                spec: spec.clone(),
97                fetched_at: Instant::now(),
98            });
99        }
100        Ok(spec)
101    }
102}
103
104fn merge(rows: Vec<systemprompt_ai::GatewayPolicyRow>) -> GatewayPolicySpec {
105    let mut merged = GatewayPolicySpec::permissive();
106    for row in rows {
107        let Ok(spec) = serde_json::from_value::<GatewayPolicySpec>(row.spec) else {
108            tracing::warn!(policy_id = %row.id, name = %row.name, "policy spec JSON malformed — skipped");
109            continue;
110        };
111        if !spec.quota_windows.is_empty() || spec.quota_mode.is_warn() {
112            merged.quota_mode = spec.quota_mode;
113        }
114        if !spec.quota_windows.is_empty() {
115            merged.quota_windows = spec.quota_windows;
116        }
117        if !spec.safety.scanners.is_empty()
118            || !spec.safety.block_categories.is_empty()
119            || !spec.safety.block_response_categories.is_empty()
120            || spec.safety.mode.is_warn()
121        {
122            merged.safety = spec.safety;
123        }
124    }
125    merged
126}