use std::io::{Cursor as IoCursor, Read};
use base64::Engine;
use hmac::{Hmac, KeyInit, Mac};
use serde::{Deserialize, Serialize};
use sha2::Sha256;
use crate::error::AppError;
use crate::store::RawKvPair;
const B64: base64::engine::general_purpose::GeneralPurpose =
base64::engine::general_purpose::URL_SAFE_NO_PAD;
const HMAC_TAG_LEN: usize = 32;
const MAX_LAST_KEY_LEN: u32 = 1024;
pub const MIN_LIMIT: usize = 1;
pub const MAX_LIMIT: usize = 200;
pub const DEFAULT_LIMIT: usize = 50;
type HmacSha256 = Hmac<Sha256>;
#[derive(Debug, Clone, Deserialize, Default)]
pub struct PaginationParams {
pub cursor: Option<String>,
pub limit: Option<usize>,
}
impl PaginationParams {
pub fn effective_limit(&self) -> usize {
self.limit
.unwrap_or(DEFAULT_LIMIT)
.clamp(MIN_LIMIT, MAX_LIMIT)
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))]
pub struct Paginated<T> {
pub items: Vec<T>,
pub next_cursor: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub total_estimate: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Cursor {
pub last_key: Vec<u8>,
pub snapshot_id: u64,
}
impl Cursor {
pub fn new(last_key: Vec<u8>, snapshot_id: u64) -> Self {
Self {
last_key,
snapshot_id,
}
}
pub fn encode(&self, audit_key: &[u8; 32]) -> String {
self.encode_bound(audit_key, &[])
}
pub fn encode_bound(&self, audit_key: &[u8; 32], binding: &[u8]) -> String {
let mut buf = Vec::with_capacity(4 + self.last_key.len() + 8 + HMAC_TAG_LEN);
buf.extend_from_slice(&(self.last_key.len() as u32).to_be_bytes());
buf.extend_from_slice(&self.last_key);
buf.extend_from_slice(&self.snapshot_id.to_be_bytes());
let mut mac = HmacSha256::new_from_slice(audit_key).expect("32-byte HMAC key");
mac.update(&buf);
mac.update(binding);
let tag = mac.finalize().into_bytes();
buf.extend_from_slice(&tag);
B64.encode(&buf)
}
pub fn decode(wire: &str, audit_key: &[u8; 32]) -> Result<Self, AppError> {
Self::decode_bound(wire, audit_key, &[])
}
pub fn decode_bound(
wire: &str,
audit_key: &[u8; 32],
binding: &[u8],
) -> Result<Self, AppError> {
let raw = B64.decode(wire).map_err(|_| AppError::InvalidCursor)?;
if raw.len() <= HMAC_TAG_LEN + 4 + 8 {
return Err(AppError::InvalidCursor);
}
let payload_len = raw.len() - HMAC_TAG_LEN;
let payload = &raw[..payload_len];
let received_tag = &raw[payload_len..];
let mut mac = HmacSha256::new_from_slice(audit_key).expect("32-byte HMAC key");
mac.update(payload);
mac.update(binding);
mac.verify_slice(received_tag)
.map_err(|_| AppError::InvalidCursor)?;
let mut reader = IoCursor::new(payload);
let mut key_len_buf = [0u8; 4];
reader
.read_exact(&mut key_len_buf)
.map_err(|_| AppError::InvalidCursor)?;
let key_len = u32::from_be_bytes(key_len_buf);
if key_len > MAX_LAST_KEY_LEN {
return Err(AppError::InvalidCursor);
}
let mut last_key = vec![0u8; key_len as usize];
reader
.read_exact(&mut last_key)
.map_err(|_| AppError::InvalidCursor)?;
let mut snapshot_buf = [0u8; 8];
reader
.read_exact(&mut snapshot_buf)
.map_err(|_| AppError::InvalidCursor)?;
let snapshot_id = u64::from_be_bytes(snapshot_buf);
if reader.position() as usize != payload_len {
return Err(AppError::InvalidCursor);
}
Ok(Self {
last_key,
snapshot_id,
})
}
}
pub fn paginate<T, F>(
pairs: Vec<RawKvPair>,
cursor: Option<&Cursor>,
limit: usize,
audit_key: &[u8; 32],
snapshot_id: u64,
mut map_value: F,
) -> Result<Paginated<T>, AppError>
where
F: FnMut(&[u8]) -> Result<T, AppError>,
{
let limit = limit.clamp(MIN_LIMIT, MAX_LIMIT);
let start = match cursor {
Some(c) => {
pairs
.iter()
.position(|(k, _)| k.as_slice() > c.last_key.as_slice())
.unwrap_or(pairs.len())
}
None => 0,
};
let mut items = Vec::with_capacity(limit.min(pairs.len().saturating_sub(start)));
let mut idx = start;
let mut last_seen_key: Option<Vec<u8>> = None;
while items.len() < limit && idx < pairs.len() {
let (key, value) = &pairs[idx];
items.push(map_value(value)?);
last_seen_key = Some(key.clone());
idx += 1;
}
let next_cursor = if idx < pairs.len() {
last_seen_key.map(|k| Cursor::new(k, snapshot_id).encode(audit_key))
} else {
None
};
Ok(Paginated {
items,
next_cursor,
total_estimate: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
const KEY_A: [u8; 32] = [0xAA; 32];
const KEY_B: [u8; 32] = [0xBB; 32];
#[test]
fn a_bound_cursor_rejects_a_different_binding() {
let c = Cursor::new(b"audit:2026-01-01:abc".to_vec(), 7);
let wire = c.encode_bound(&KEY_A, b"action=MemberAdded");
assert_eq!(
Cursor::decode_bound(&wire, &KEY_A, b"action=MemberAdded").unwrap(),
c
);
assert!(matches!(
Cursor::decode_bound(&wire, &KEY_A, b"action=MemberRemoved"),
Err(AppError::InvalidCursor)
));
assert!(matches!(
Cursor::decode(&wire, &KEY_A),
Err(AppError::InvalidCursor)
));
assert!(matches!(
Cursor::decode_bound(&wire, &KEY_B, b"action=MemberAdded"),
Err(AppError::InvalidCursor)
));
}
#[test]
fn an_unbound_cursor_round_trips_as_before() {
let c = Cursor::new(b"policy:001".to_vec(), 3);
let wire = c.encode(&KEY_A);
assert_eq!(Cursor::decode(&wire, &KEY_A).unwrap(), c);
assert_eq!(wire, c.encode_bound(&KEY_A, &[]));
}
fn make_pairs(count: usize) -> Vec<RawKvPair> {
(0..count)
.map(|i| {
(
format!("item:{i:03}").into_bytes(),
format!(r#"{{"n":{i}}}"#).into_bytes(),
)
})
.collect()
}
fn deserialize_value(bytes: &[u8]) -> Result<serde_json::Value, AppError> {
serde_json::from_slice(bytes)
.map_err(|e| AppError::Internal(format!("test value deserialize failed: {e}")))
}
#[test]
fn cursor_round_trips_through_encode_decode() {
let c = Cursor::new(b"member:did:key:z6Mk".to_vec(), 1_700_000_000);
let wire = c.encode(&KEY_A);
let back = Cursor::decode(&wire, &KEY_A).unwrap();
assert_eq!(back, c);
}
#[test]
fn cursor_decoded_with_different_key_is_rejected() {
let c = Cursor::new(b"x".to_vec(), 42);
let wire = c.encode(&KEY_A);
let err = Cursor::decode(&wire, &KEY_B).expect_err("must reject");
assert!(matches!(err, AppError::InvalidCursor));
}
#[test]
fn foreign_audit_key_cursor_maps_to_http_400() {
use axum::response::IntoResponse;
let c = Cursor::new(b"member:did:key:zAttacker".to_vec(), 17);
let wire = c.encode(&KEY_A);
let err = Cursor::decode(&wire, &KEY_B).expect_err("must reject");
let response = err.into_response();
assert_eq!(response.status(), 400);
}
#[test]
fn cursor_with_tampered_payload_is_rejected() {
let c = Cursor::new(b"safe-key".to_vec(), 1);
let wire = c.encode(&KEY_A);
let mut bytes = B64.decode(&wire).unwrap();
bytes[5] ^= 0xFF;
let tampered = B64.encode(&bytes);
let err = Cursor::decode(&tampered, &KEY_A).expect_err("tampered");
assert!(matches!(err, AppError::InvalidCursor));
}
#[test]
fn cursor_malformed_base64_is_rejected() {
for bad in ["", "not!base64", "AAAA"] {
let err = Cursor::decode(bad, &KEY_A).expect_err("malformed");
assert!(matches!(err, AppError::InvalidCursor), "input {bad}");
}
}
#[test]
fn cursor_with_oversized_key_length_is_rejected() {
let mut payload = Vec::new();
payload.extend_from_slice(&u32::MAX.to_be_bytes());
payload.extend_from_slice(&0u64.to_be_bytes());
let mut mac = HmacSha256::new_from_slice(&KEY_A).unwrap();
mac.update(&payload);
let tag = mac.finalize().into_bytes();
payload.extend_from_slice(&tag);
let wire = B64.encode(&payload);
let err = Cursor::decode(&wire, &KEY_A).expect_err("oversized");
assert!(matches!(err, AppError::InvalidCursor));
}
#[test]
fn paginate_first_page_no_cursor() {
let pairs = make_pairs(7);
let out: Paginated<serde_json::Value> =
paginate(pairs, None, 3, &KEY_A, 100, deserialize_value).unwrap();
assert_eq!(out.items.len(), 3);
assert_eq!(out.items[0]["n"], 0);
assert_eq!(out.items[2]["n"], 2);
assert!(out.next_cursor.is_some());
}
#[test]
fn paginate_walks_entire_collection_without_duplicates() {
let pairs = make_pairs(7);
let mut cursor: Option<Cursor> = None;
let mut seen = Vec::new();
for _ in 0..5 {
let out: Paginated<serde_json::Value> = paginate(
pairs.clone(),
cursor.as_ref(),
3,
&KEY_A,
100,
deserialize_value,
)
.unwrap();
for item in &out.items {
seen.push(item["n"].as_u64().unwrap());
}
match out.next_cursor {
Some(wire) => cursor = Some(Cursor::decode(&wire, &KEY_A).unwrap()),
None => break,
}
}
assert_eq!(seen, (0..7).collect::<Vec<_>>());
}
#[test]
fn paginate_last_page_returns_no_next_cursor() {
let pairs = make_pairs(3);
let out: Paginated<serde_json::Value> =
paginate(pairs, None, 10, &KEY_A, 100, deserialize_value).unwrap();
assert_eq!(out.items.len(), 3);
assert!(out.next_cursor.is_none(), "last page must not link onward");
}
#[test]
fn paginate_clamps_limit() {
let pairs = make_pairs(500);
let out: Paginated<serde_json::Value> =
paginate(pairs, None, 9_999, &KEY_A, 100, deserialize_value).unwrap();
assert_eq!(out.items.len(), MAX_LIMIT);
}
#[test]
fn paginate_skips_past_cursor_key_exclusive() {
let pairs = make_pairs(5);
let cursor = Cursor::new(b"item:001".to_vec(), 1);
let out: Paginated<serde_json::Value> =
paginate(pairs, Some(&cursor), 10, &KEY_A, 100, deserialize_value).unwrap();
let ns: Vec<_> = out.items.iter().map(|v| v["n"].as_u64().unwrap()).collect();
assert_eq!(ns, vec![2, 3, 4]);
}
#[test]
fn paginate_returns_empty_when_cursor_already_past_end() {
let pairs = make_pairs(3);
let cursor = Cursor::new(b"zzz".to_vec(), 1);
let out: Paginated<serde_json::Value> =
paginate(pairs, Some(&cursor), 10, &KEY_A, 100, deserialize_value).unwrap();
assert!(out.items.is_empty());
assert!(out.next_cursor.is_none());
}
#[test]
fn pagination_params_effective_limit_clamps_and_defaults() {
assert_eq!(
PaginationParams {
cursor: None,
limit: None,
}
.effective_limit(),
DEFAULT_LIMIT
);
assert_eq!(
PaginationParams {
cursor: None,
limit: Some(0),
}
.effective_limit(),
MIN_LIMIT
);
assert_eq!(
PaginationParams {
cursor: None,
limit: Some(9999),
}
.effective_limit(),
MAX_LIMIT
);
assert_eq!(
PaginationParams {
cursor: None,
limit: Some(50),
}
.effective_limit(),
50
);
}
}