1use moq_pattern::Patterns;
2use serde::{Deserialize, Serialize};
3use serde_with::{DurationSeconds, TimestampSeconds, serde_as};
4use std::collections::BTreeMap;
5use std::time::{Duration, SystemTime};
6
7#[serde_as]
14#[serde_with::skip_serializing_none]
15#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
16#[serde(default)]
17#[non_exhaustive]
18pub struct Grant {
19 #[serde(skip_serializing_if = "Patterns::is_empty")]
21 pub publish: Patterns,
22
23 #[serde(skip_serializing_if = "Patterns::is_empty")]
25 pub subscribe: Patterns,
26
27 pub root: Option<String>,
30
31 #[serde(skip_serializing_if = "BTreeMap::is_empty")]
35 pub mounts: BTreeMap<String, String>,
36
37 #[serde_as(as = "Option<TimestampSeconds<i64>>")]
39 pub expires: Option<SystemTime>,
40
41 #[serde_as(as = "Option<DurationSeconds<u64>>")]
43 pub revalidate: Option<Duration>,
44
45 pub tier: Option<String>,
47
48 #[serde(skip_serializing_if = "std::ops::Not::not")]
51 pub peer: bool,
52}
53
54impl Grant {
55 pub fn new(publish: Patterns, subscribe: Patterns) -> Self {
57 Self {
58 publish,
59 subscribe,
60 ..Default::default()
61 }
62 }
63
64 #[cfg(feature = "tokio")]
67 pub fn deadline(&self) -> Option<tokio::time::Instant> {
68 let remaining = self.expires?.duration_since(SystemTime::now()).unwrap_or_default();
69 tokio::time::Instant::now().checked_add(remaining)
70 }
71
72 pub fn validate(&self) -> crate::Result<()> {
75 if self.publish.is_empty() && self.subscribe.is_empty() {
76 return Err(crate::Error::UselessGrant);
77 }
78 if self.revalidate.is_some() && self.expires.is_none() {
79 return Err(crate::Error::UnboundedRevalidate);
80 }
81 if self.revalidate.is_some_and(|cadence| cadence.is_zero()) {
83 return Err(crate::Error::ZeroRevalidate);
84 }
85 if self.expires.is_some_and(|expires| expires <= SystemTime::now()) {
86 return Err(crate::Error::GrantExpired);
87 }
88 Ok(())
89 }
90}
91
92#[cfg(test)]
93mod tests {
94 use super::*;
95
96 fn patterns(texts: &[&str]) -> Patterns {
97 texts.iter().map(|text| text.parse().unwrap()).collect()
98 }
99
100 #[test]
101 fn round_trips_in_seconds() {
102 let grant = Grant {
103 publish: patterns(&["alice/**"]),
104 subscribe: patterns(&["**"]),
105 root: Some("pid/room".into()),
106 mounts: BTreeMap::new(),
107 expires: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(4_102_444_800)),
108 revalidate: Some(Duration::from_secs(60)),
109 tier: Some("websocket".into()),
110 peer: true,
111 };
112 let json = serde_json::to_value(&grant).unwrap();
113 assert_eq!(json["expires"], 4_102_444_800_i64);
114 assert_eq!(json["revalidate"], 60);
115 assert_eq!(json["publish"], serde_json::json!(["alice/**"]));
116 assert_eq!(json["peer"], true);
117 assert_eq!(serde_json::from_value::<Grant>(json).unwrap(), grant);
118 }
119
120 #[test]
123 fn serializes_to_the_cross_language_vector() {
124 let grant = Grant {
125 publish: patterns(&["alice/**"]),
126 subscribe: patterns(&["**"]),
127 root: Some("pid/room".into()),
128 mounts: BTreeMap::new(),
129 expires: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(4_102_444_800)),
130 revalidate: Some(Duration::from_secs(60)),
131 tier: Some("websocket".into()),
132 peer: true,
133 };
134 assert_eq!(
135 serde_json::to_string(&grant).unwrap(),
136 r#"{"publish":["alice/**"],"subscribe":["**"],"root":"pid/room","expires":4102444800,"revalidate":60,"tier":"websocket","peer":true}"#
137 );
138 }
139
140 #[test]
141 fn mounts_round_trip_as_an_object() {
142 let mut grant = Grant::new(Patterns::new(), patterns(&["**"]));
143 grant.mounts.insert(".svc".into(), ".svc/pid".into());
144 let json = serde_json::to_string(&grant).unwrap();
145 assert_eq!(json, r#"{"subscribe":["**"],"mounts":{".svc":".svc/pid"}}"#);
146 assert_eq!(serde_json::from_str::<Grant>(&json).unwrap(), grant);
147 }
148
149 #[test]
150 fn empty_fields_are_omitted_and_defaulted() {
151 let grant = Grant::new(patterns(&["**"]), Patterns::new());
152 assert_eq!(serde_json::to_string(&grant).unwrap(), r#"{"publish":["**"]}"#);
153 assert_eq!(serde_json::from_str::<Grant>(r#"{"publish":["**"]}"#).unwrap(), grant);
154 }
155
156 #[test]
157 fn validate_refuses_nothing_unbounded_and_expired() {
158 assert!(matches!(Grant::default().validate(), Err(crate::Error::UselessGrant)));
159
160 let mut grant = Grant::new(patterns(&["**"]), Patterns::new());
161 grant.validate().unwrap();
162
163 grant.revalidate = Some(Duration::from_secs(1));
164 assert!(matches!(grant.validate(), Err(crate::Error::UnboundedRevalidate)));
165
166 grant.expires = Some(SystemTime::now());
168 assert!(matches!(grant.validate(), Err(crate::Error::GrantExpired)));
169
170 grant.expires = Some(SystemTime::now() - Duration::from_secs(1));
171 assert!(matches!(grant.validate(), Err(crate::Error::GrantExpired)));
172
173 grant.expires = Some(SystemTime::now() + Duration::from_secs(60));
174 grant.validate().unwrap();
175
176 grant.revalidate = Some(Duration::ZERO);
177 assert!(matches!(grant.validate(), Err(crate::Error::ZeroRevalidate)));
178 }
179
180 #[cfg(feature = "tokio")]
181 #[tokio::test(start_paused = true)]
182 async fn deadline_is_the_exact_expiry() {
183 let start = tokio::time::Instant::now();
184 let mut grant = Grant::new(patterns(&["**"]), Patterns::new());
185 assert_eq!(grant.deadline(), None);
186
187 grant.expires = Some(SystemTime::now() + Duration::from_secs(10));
188 let deadline = grant.deadline().unwrap();
189 assert!(deadline <= start + Duration::from_secs(10), "not later than the expiry");
190 assert!(deadline > start + Duration::from_secs(9));
191
192 grant.expires = Some(SystemTime::now() - Duration::from_secs(1));
193 assert_eq!(
194 grant.deadline(),
195 Some(start),
196 "a past expiry is now, not a grace window"
197 );
198 }
199
200 #[cfg(all(feature = "tokio", unix))]
204 #[tokio::test(start_paused = true)]
205 async fn deadline_past_the_clock_is_never() {
206 let start = tokio::time::Instant::now();
207 let grant: Grant = serde_json::from_str(r#"{"publish":["**"],"expires":9223372036854775807}"#).unwrap();
208 let deadline = grant.deadline();
210 assert!(deadline.is_none_or(|at| at > start + Duration::from_secs(1 << 40)));
211 }
212}