use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use base64::Engine;
const TOKEN_VERSION_PREFIX: &str = "v1:";
pub(crate) const MAX_PAGE_TOKEN_CHARS: usize = 1024;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum PageTokenError {
TooLong,
NotBase64,
NotUtf8,
WrongVersion,
EmptyCursor,
}
pub(crate) fn encode_page_token(session_id: &str) -> String {
URL_SAFE_NO_PAD.encode(format!("{TOKEN_VERSION_PREFIX}{session_id}"))
}
pub(crate) fn decode_page_token(token: &str) -> Result<String, PageTokenError> {
if token.len() > MAX_PAGE_TOKEN_CHARS {
return Err(PageTokenError::TooLong);
}
let bytes = URL_SAFE_NO_PAD
.decode(token)
.map_err(|_| PageTokenError::NotBase64)?;
let plaintext = String::from_utf8(bytes).map_err(|_| PageTokenError::NotUtf8)?;
let cursor = plaintext
.strip_prefix(TOKEN_VERSION_PREFIX)
.ok_or(PageTokenError::WrongVersion)?;
if cursor.is_empty() {
return Err(PageTokenError::EmptyCursor);
}
Ok(cursor.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn token_round_trips_session_id() {
let ids: Vec<String> = vec![
"3f2504e0-4f89-41d3-9a0c-0305e82c3301".to_string(),
"AAECAwQFBgcICQoLDA0ODw".to_string(),
"legacy_session_1".to_string(),
"session-ünïcøde-☃-世界".to_string(),
"x".repeat(200),
];
for id in &ids {
let token = encode_page_token(id);
assert_eq!(decode_page_token(&token), Ok(id.clone()), "id = {id}");
}
}
#[test]
fn token_is_not_the_bare_session_id() {
for id in [
"3f2504e0-4f89-41d3-9a0c-0305e82c3301",
"AAECAwQFBgcICQoLDA0ODw",
"legacy_session_1",
] {
assert_ne!(encode_page_token(id), id);
}
}
#[test]
fn decode_rejects_garbage() {
let garbage = URL_SAFE_NO_PAD.encode("just some bytes");
assert_eq!(
decode_page_token(&garbage),
Err(PageTokenError::WrongVersion)
);
assert!(decode_page_token("").is_err());
}
#[test]
fn decode_rejects_non_base64() {
for token in ["not a token!", "***", "abc=", "a b c", "%%%%"] {
assert_eq!(
decode_page_token(token),
Err(PageTokenError::NotBase64),
"token literal rejected for the wrong reason"
);
}
}
#[test]
fn decode_rejects_truncated_token() {
let token = encode_page_token("3f2504e0-4f89-41d3-9a0c-0305e82c3301");
for cut in 1..token.len() {
let truncated = &token[..cut];
assert_ne!(
decode_page_token(truncated),
Ok("3f2504e0-4f89-41d3-9a0c-0305e82c3301".to_string()),
"truncation at {cut} produced the original cursor"
);
}
for plaintext in ["v", "v1"] {
let short = URL_SAFE_NO_PAD.encode(plaintext);
assert_eq!(
decode_page_token(&short),
Err(PageTokenError::WrongVersion),
"plaintext {plaintext:?}"
);
}
assert_ne!(
decode_page_token(&token[1..]),
Ok("3f2504e0-4f89-41d3-9a0c-0305e82c3301".to_string())
);
}
#[test]
fn decode_rejects_oversized_token_without_decoding() {
let huge = "A".repeat(2 * 1024 * 1024);
assert_eq!(decode_page_token(&huge), Err(PageTokenError::TooLong));
let at_limit = "A".repeat(MAX_PAGE_TOKEN_CHARS);
assert_ne!(decode_page_token(&at_limit), Err(PageTokenError::TooLong));
let over_limit = "A".repeat(MAX_PAGE_TOKEN_CHARS + 1);
assert_eq!(decode_page_token(&over_limit), Err(PageTokenError::TooLong));
}
#[test]
fn decode_rejects_wrong_version_prefix() {
for plaintext in ["v2:abc", "v:abc", "1:abc", "abc", "V1:abc", " v1:abc"] {
let token = URL_SAFE_NO_PAD.encode(plaintext);
assert_eq!(
decode_page_token(&token),
Err(PageTokenError::WrongVersion),
"plaintext {plaintext:?} was not rejected as a wrong version"
);
}
}
#[test]
fn decode_rejects_valid_base64_of_invalid_utf8() {
for bytes in [
vec![0x76, 0x31, 0x3a, 0xff], vec![0xc3, 0x28], vec![0xed, 0xa0, 0x80], vec![0x76, 0x31, 0x3a, 0xe2, 0x28], ] {
let token = URL_SAFE_NO_PAD.encode(&bytes);
assert_eq!(
decode_page_token(&token),
Err(PageTokenError::NotUtf8),
"bytes {bytes:?} were not rejected as invalid UTF-8"
);
}
}
#[test]
fn decode_rejects_empty_cursor_after_prefix() {
let token = URL_SAFE_NO_PAD.encode("v1:");
assert_eq!(decode_page_token(&token), Err(PageTokenError::EmptyCursor));
}
#[test]
fn decode_never_panics_on_adversarial_input() {
let seed = encode_page_token("3f2504e0-4f89-41d3-9a0c-0305e82c3301");
let seed_bytes = seed.as_bytes();
let mut cases: Vec<String> = Vec::new();
for i in 0..seed_bytes.len() {
for bit in 0..8u32 {
let mut b = seed_bytes.to_vec();
b[i] ^= 1 << bit;
cases.push(String::from_utf8_lossy(&b).into_owned());
}
}
for cut in 0..=seed_bytes.len() {
cases.push(String::from_utf8_lossy(&seed_bytes[..cut]).into_owned());
}
for junk in [
"", "=", "==", "!", "\n", "\r\n", "\0", "v1:", "AAAA", "../", "%00", "\u{feff}",
] {
cases.push(format!("{junk}{seed}"));
cases.push(format!("{seed}{junk}"));
cases.push(junk.to_string());
}
for bytes in [
vec![0x76, 0x31, 0x3a, 0x00, 0x41],
vec![0x00; 64],
vec![0xed, 0xa0, 0x80],
vec![0xed, 0xbf, 0xbf],
vec![0xff; 128],
vec![0x76, 0x31, 0x3a, 0xed, 0xa0, 0xbd, 0xed, 0xb8, 0x80],
] {
cases.push(URL_SAFE_NO_PAD.encode(&bytes));
cases.push(String::from_utf8_lossy(&bytes).into_owned());
}
for len in [0usize, 1, 2, 3, 4, 1023, 1024, 1025, 4096] {
cases.push("A".repeat(len));
cases.push("~".repeat(len.min(256)));
}
for n in 1..200usize {
cases.push(format!("v1:{}", "A".repeat(n)));
cases.push(URL_SAFE_NO_PAD.encode("v1:".repeat(n)));
cases.push(URL_SAFE_NO_PAD.encode("\u{0}".repeat(n)));
}
assert!(
cases.len() >= 1000,
"corpus too small: {} cases",
cases.len()
);
for case in &cases {
let _ = decode_page_token(case);
}
assert_eq!(
decode_page_token(&seed),
Ok("3f2504e0-4f89-41d3-9a0c-0305e82c3301".to_string())
);
}
}