1use std::collections::HashSet;
4use std::num::NonZeroUsize;
5
6use async_trait::async_trait;
7use jiff::Timestamp;
8use tollgate_core::{KeyId, Principal};
9
10use crate::{KeyRecord, StoreError};
11
12pub const MAX_KEY_PAGE_LIMIT: usize = 4096;
15pub const DEFAULT_KEY_PAGE_LIMIT: NonZeroUsize = NonZeroUsize::new(256).unwrap();
18pub const MAX_KEY_REVISION: u64 = i64::MAX as u64;
20
21pub fn validate_key_page_limit(limit: NonZeroUsize) -> Result<(), StoreError> {
27 if limit.get() > MAX_KEY_PAGE_LIMIT {
28 return Err(StoreError("credential page limit exceeds 4096".into()));
29 }
30 Ok(())
31}
32
33#[derive(Clone, Copy, PartialEq, Eq)]
36#[cfg_attr(feature = "wire", derive(serde::Serialize, serde::Deserialize))]
37pub struct CredentialRecord {
38 pub key_id: KeyId,
40 pub principal: Principal,
43 #[cfg_attr(feature = "wire", serde(with = "digest_hex"))]
46 pub digest: [u8; 32],
47 #[cfg_attr(feature = "wire", serde(deserialize_with = "required_option"))]
49 pub not_after: Option<Timestamp>,
50}
51
52impl std::fmt::Debug for CredentialRecord {
53 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54 f.debug_struct("CredentialRecord")
55 .field("key_id", &self.key_id)
56 .field("principal", &self.principal)
57 .field("digest", &"[redacted]")
58 .field("not_after", &self.not_after)
59 .finish()
60 }
61}
62
63impl From<KeyRecord> for CredentialRecord {
64 fn from(record: KeyRecord) -> Self {
65 Self {
66 key_id: record.key_id,
67 principal: record.principal,
68 digest: record.digest,
69 not_after: record.not_after,
70 }
71 }
72}
73
74#[derive(Debug, Clone, PartialEq, Eq)]
77pub struct CredentialSet(Vec<CredentialRecord>);
78
79impl CredentialSet {
80 pub fn try_new(records: Vec<CredentialRecord>) -> Result<Self, StoreError> {
87 let mut principals = HashSet::with_capacity(records.len());
88 let mut ids = HashSet::with_capacity(records.len());
89 for record in &records {
90 let mut prefix = [0; 16];
91 prefix.copy_from_slice(&record.digest[..16]);
92 if record.principal.0 != u128::from_be_bytes(prefix) {
93 return Err(StoreError(
94 "credential digest does not identify its principal".into(),
95 ));
96 }
97 if !principals.insert(record.principal) || !ids.insert(record.key_id) {
98 return Err(StoreError(
99 "credential projection contains duplicate identities".into(),
100 ));
101 }
102 }
103 Ok(Self(records))
104 }
105
106 pub fn records(&self) -> &[CredentialRecord] {
108 &self.0
109 }
110 pub fn into_records(self) -> Vec<CredentialRecord> {
112 self.0
113 }
114}
115
116#[derive(Debug, Clone, PartialEq, Eq)]
121pub struct KeyPage {
122 revision: u64,
123 as_of: Timestamp,
124 keys: CredentialSet,
125 next_after: Option<KeyId>,
126 after: Option<KeyId>,
127 limit: NonZeroUsize,
128}
129
130impl KeyPage {
131 pub fn try_new(
142 revision: u64,
143 as_of: Timestamp,
144 after: Option<KeyId>,
145 limit: NonZeroUsize,
146 keys: Vec<CredentialRecord>,
147 next_after: Option<KeyId>,
148 ) -> Result<Self, StoreError> {
149 validate_key_page_limit(limit)?;
150 if revision > MAX_KEY_REVISION || keys.len() > limit.get() {
151 return Err(StoreError(
152 "credential page exceeds its revision or record domain".into(),
153 ));
154 }
155 let mut previous = after;
156 for key in &keys {
157 if previous.is_some_and(|id| key.key_id <= id)
158 || key.not_after.is_some_and(|end| as_of >= end)
159 {
160 return Err(StoreError(
161 "credential page is unordered, repeated, or expired at its source".into(),
162 ));
163 }
164 previous = Some(key.key_id);
165 }
166 if let Some(next) = next_after
167 && (keys.len() != limit.get() || Some(next) != previous)
168 {
169 return Err(StoreError(
170 "credential page continuation does not advance its full page".into(),
171 ));
172 }
173 Ok(Self {
174 revision,
175 as_of,
176 keys: CredentialSet::try_new(keys)?,
177 next_after,
178 after,
179 limit,
180 })
181 }
182
183 pub fn validate_request(
189 &self,
190 after: Option<KeyId>,
191 limit: NonZeroUsize,
192 ) -> Result<(), StoreError> {
193 if self.after != after || self.limit != limit {
194 return Err(StoreError(
195 "credential page belongs to a different request".into(),
196 ));
197 }
198 Ok(())
199 }
200
201 pub fn revision(&self) -> u64 {
203 self.revision
204 }
205 pub fn as_of(&self) -> Timestamp {
207 self.as_of
208 }
209 pub fn records(&self) -> &[CredentialRecord] {
211 self.keys.records()
212 }
213 pub fn next_after(&self) -> Option<KeyId> {
216 self.next_after
217 }
218 pub fn into_records(self) -> Vec<CredentialRecord> {
220 self.keys.into_records()
221 }
222}
223
224#[async_trait]
229pub trait KeySource: Send + Sync {
230 async fn active_keys_page(
235 &self,
236 now: Timestamp,
237 after: Option<KeyId>,
238 limit: NonZeroUsize,
239 ) -> Result<KeyPage, StoreError>;
240}
241
242#[cfg(feature = "wire")]
243pub(crate) fn required_option<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
244where
245 D: serde::Deserializer<'de>,
246 T: serde::Deserialize<'de>,
247{
248 serde::Deserialize::deserialize(deserializer)
249}
250
251#[cfg(feature = "wire")]
252mod digest_hex {
253 use serde::{Deserialize, Deserializer, Serializer};
254
255 pub fn serialize<S: Serializer>(digest: &[u8; 32], serializer: S) -> Result<S::Ok, S::Error> {
256 use std::fmt::Write;
257 let mut text = String::with_capacity(64);
258 for byte in digest {
259 write!(text, "{byte:02x}").expect("writing into a String cannot fail");
260 }
261 serializer.serialize_str(&text)
262 }
263
264 pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<[u8; 32], D::Error> {
265 let text = String::deserialize(deserializer)?;
266 if text.len() != 64
267 || !text
268 .bytes()
269 .all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(&b))
270 {
271 return Err(serde::de::Error::custom(
272 "credential digest must be 64 lowercase hexadecimal characters",
273 ));
274 }
275 let mut digest = [0; 32];
276 for (byte, pair) in digest.iter_mut().zip(text.as_bytes().chunks_exact(2)) {
277 let pair = std::str::from_utf8(pair).map_err(serde::de::Error::custom)?;
278 *byte = u8::from_str_radix(pair, 16).map_err(serde::de::Error::custom)?;
279 }
280 Ok(digest)
281 }
282}
283
284#[cfg(test)]
285mod tests {
286 use super::*;
287 fn key(id: u128) -> CredentialRecord {
288 let mut digest = [0; 32];
289 digest[..16].copy_from_slice(&id.to_be_bytes());
290 CredentialRecord {
291 key_id: KeyId(id),
292 principal: Principal(id),
293 digest,
294 not_after: None,
295 }
296 }
297 #[test]
298 fn page_owns_identity_order_cursor_expiry_and_limit_validation() {
299 let now = Timestamp::UNIX_EPOCH;
300 let limit = NonZeroUsize::new(2).unwrap();
301 let page = |records, after, next| KeyPage::try_new(0, now, after, limit, records, next);
302 let valid = page(vec![key(1), key(2)], None, Some(KeyId(2))).unwrap();
303 assert!(valid.validate_request(None, limit).is_ok());
304 assert!(valid.validate_request(Some(KeyId(1)), limit).is_err());
305 assert!(
306 valid
307 .validate_request(None, NonZeroUsize::new(1).unwrap())
308 .is_err()
309 );
310 assert!(page(vec![key(1), key(2)], None, None).is_ok());
311 assert!(page(vec![], None, None).is_ok());
312 assert!(page(vec![], None, Some(KeyId(1))).is_err());
313 assert!(page(vec![key(1)], None, Some(KeyId(1))).is_err());
314 assert!(page(vec![key(1), key(2)], None, Some(KeyId(1))).is_err());
315 assert!(page(vec![key(2), key(1)], None, None).is_err());
316 assert!(page(vec![key(1), key(1)], None, None).is_err());
317 assert!(page(vec![key(1)], Some(KeyId(1)), None).is_err());
318 assert!(page(vec![key(1), key(2), key(3)], None, None).is_err());
319 let mut expired = key(1);
320 expired.not_after = Some(now);
321 assert!(page(vec![expired], None, None).is_err());
322 expired.not_after = now.checked_add(jiff::SignedDuration::from_nanos(1)).ok();
323 assert!(page(vec![expired], None, None).is_ok());
324 let mut corrupt = key(1);
325 corrupt.digest[0] ^= 1;
326 assert!(page(vec![corrupt], None, None).is_err());
327 let mut repeated_principal = key(1);
328 repeated_principal.key_id = KeyId(2);
329 assert!(page(vec![key(1), repeated_principal], None, None).is_err());
330 assert!(KeyPage::try_new(MAX_KEY_REVISION + 1, now, None, limit, vec![], None).is_err());
331 assert!(KeyPage::try_new(MAX_KEY_REVISION, now, None, limit, vec![], None).is_ok());
332 assert!(validate_key_page_limit(NonZeroUsize::new(MAX_KEY_PAGE_LIMIT).unwrap()).is_ok());
333 assert!(
334 validate_key_page_limit(NonZeroUsize::new(MAX_KEY_PAGE_LIMIT + 1).unwrap()).is_err()
335 );
336 }
337 #[test]
338 fn record_diagnostics_redact_the_digest() {
339 let record = key(0xabc);
340 assert!(!format!("{record:?}").contains(&format!("{:?}", record.digest)));
341 }
342}