systemprompt_api/services/gateway/
policy.rs1use 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}