use std::time::Duration;
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use hmac::{Hmac, KeyInit, Mac};
use serde::{Serialize, de::DeserializeOwned};
use sha2::Sha256;
use thiserror::Error;
type HmacSha256 = Hmac<Sha256>;
const VERSION: &str = "rs1";
const DOMAIN: &[u8] = b"rmcp/mrtr/request-state/v1";
const EXPIRY_LEN: usize = 8;
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum RequestStateError {
#[error("request state is malformed or uses an unsupported format")]
MalformedFormat,
#[error("request state is not valid base64url")]
InvalidEncoding,
#[error("request state failed integrity verification")]
IntegrityCheckFailed,
#[error("request state has expired")]
Expired,
#[error("failed to serialize request state payload: {0}")]
Serialization(#[source] serde_json::Error),
#[error("failed to deserialize request state payload: {0}")]
Deserialization(#[source] serde_json::Error),
}
#[derive(Clone, Copy, Debug, Default)]
pub struct SealOptions<'a> {
associated_data: &'a [u8],
ttl: Option<Duration>,
}
impl<'a> SealOptions<'a> {
pub fn new() -> Self {
Self::default()
}
pub fn associated_data(mut self, associated_data: &'a [u8]) -> Self {
self.associated_data = associated_data;
self
}
pub fn ttl(mut self, ttl: Duration) -> Self {
self.ttl = Some(ttl);
self
}
}
#[derive(Clone)]
pub struct RequestStateCodec {
key: Box<[u8]>,
}
impl std::fmt::Debug for RequestStateCodec {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RequestStateCodec")
.field("key", &"<redacted>")
.finish()
}
}
impl RequestStateCodec {
pub fn new(key: impl Into<Vec<u8>>) -> Self {
Self {
key: key.into().into_boxed_slice(),
}
}
pub fn seal(&self, payload: &[u8]) -> String {
self.seal_with(payload, &SealOptions::default())
}
pub fn seal_with(&self, payload: &[u8], options: &SealOptions<'_>) -> String {
self.seal_at(payload, options, Self::now_ms())
}
pub fn seal_json<T: Serialize>(&self, value: &T) -> Result<String, RequestStateError> {
self.seal_json_with(value, &SealOptions::default())
}
pub fn seal_json_with<T: Serialize>(
&self,
value: &T,
options: &SealOptions<'_>,
) -> Result<String, RequestStateError> {
let payload = serde_json::to_vec(value).map_err(RequestStateError::Serialization)?;
Ok(self.seal_with(&payload, options))
}
pub fn open(&self, sealed: &str) -> Result<Vec<u8>, RequestStateError> {
self.open_with(sealed, &[])
}
pub fn open_with(
&self,
sealed: &str,
associated_data: &[u8],
) -> Result<Vec<u8>, RequestStateError> {
self.open_at(sealed, associated_data, Self::now_ms())
}
pub fn open_json<T: DeserializeOwned>(&self, sealed: &str) -> Result<T, RequestStateError> {
self.open_json_with(sealed, &[])
}
pub fn open_json_with<T: DeserializeOwned>(
&self,
sealed: &str,
associated_data: &[u8],
) -> Result<T, RequestStateError> {
let payload = self.open_with(sealed, associated_data)?;
serde_json::from_slice(&payload).map_err(RequestStateError::Deserialization)
}
fn seal_at(&self, payload: &[u8], options: &SealOptions<'_>, now_ms: i64) -> String {
let expiry = match options.ttl {
Some(ttl) => now_ms.saturating_add(ttl.as_millis().min(i64::MAX as u128) as i64),
None => 0,
};
let mut body = Vec::with_capacity(EXPIRY_LEN + payload.len());
body.extend_from_slice(&expiry.to_be_bytes());
body.extend_from_slice(payload);
let tag = self
.mac_for(options.associated_data, &body)
.finalize()
.into_bytes();
let b64_len = |n: usize| n.div_ceil(3) * 4;
let mut out =
String::with_capacity(VERSION.len() + 2 + b64_len(body.len()) + b64_len(tag.len()));
out.push_str(VERSION);
out.push('.');
URL_SAFE_NO_PAD.encode_string(&body, &mut out);
out.push('.');
URL_SAFE_NO_PAD.encode_string(tag.as_slice(), &mut out);
out
}
fn open_at(
&self,
sealed: &str,
associated_data: &[u8],
now_ms: i64,
) -> Result<Vec<u8>, RequestStateError> {
let mut parts = sealed.split('.');
let version = parts.next().ok_or(RequestStateError::MalformedFormat)?;
let body_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?;
let tag_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?;
if parts.next().is_some() || version != VERSION {
return Err(RequestStateError::MalformedFormat);
}
let body = URL_SAFE_NO_PAD
.decode(body_b64)
.map_err(|_| RequestStateError::InvalidEncoding)?;
let tag = URL_SAFE_NO_PAD
.decode(tag_b64)
.map_err(|_| RequestStateError::InvalidEncoding)?;
self.mac_for(associated_data, &body)
.verify_slice(&tag)
.map_err(|_| RequestStateError::IntegrityCheckFailed)?;
if body.len() < EXPIRY_LEN {
return Err(RequestStateError::MalformedFormat);
}
let expiry = i64::from_be_bytes(body[..EXPIRY_LEN].try_into().expect("checked length"));
if expiry != 0 && now_ms > expiry {
return Err(RequestStateError::Expired);
}
Ok(body[EXPIRY_LEN..].to_vec())
}
fn mac_for(&self, associated_data: &[u8], body: &[u8]) -> HmacSha256 {
let mut mac =
HmacSha256::new_from_slice(&self.key).expect("HMAC accepts keys of any length");
mac.update(DOMAIN);
mac.update(&(associated_data.len() as u64).to_be_bytes());
mac.update(associated_data);
mac.update(body);
mac
}
fn now_ms() -> i64 {
chrono::Utc::now().timestamp_millis()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn seal_open_roundtrips_bytes() {
let codec = RequestStateCodec::new(b"test-key-test-key-test-key-32byte".to_vec());
let sealed = codec.seal(b"hello world");
assert!(sealed.starts_with("rs1."));
assert_eq!(codec.open(&sealed).unwrap(), b"hello world");
}
#[test]
fn seal_open_roundtrips_json() {
#[derive(serde::Serialize, serde::Deserialize, PartialEq, Debug)]
struct State {
tool: String,
round: u32,
}
let codec = RequestStateCodec::new(b"another-strong-signing-key-here!!".to_vec());
let state = State {
tool: "weather".into(),
round: 3,
};
let sealed = codec.seal_json(&state).unwrap();
let opened: State = codec.open_json(&sealed).unwrap();
assert_eq!(opened, state);
}
#[test]
fn empty_payload_roundtrips() {
let codec = RequestStateCodec::new(b"k".to_vec());
let sealed = codec.seal(b"");
assert_eq!(codec.open(&sealed).unwrap(), b"");
}
#[test]
fn tampered_payload_is_rejected() {
let codec = RequestStateCodec::new(b"signing-key-signing-key-signing!!".to_vec());
let sealed = codec.seal(b"amount=100");
let mut parts: Vec<&str> = sealed.split('.').collect();
let forged_body = URL_SAFE_NO_PAD.encode(b"amount=999");
parts[1] = &forged_body;
let forged = parts.join(".");
assert!(matches!(
codec.open(&forged),
Err(RequestStateError::IntegrityCheckFailed)
));
}
#[test]
fn different_key_is_rejected() {
let signer = RequestStateCodec::new(b"the-real-signing-key-value-here!!".to_vec());
let attacker = RequestStateCodec::new(b"a-totally-different-forged-key!!!".to_vec());
let sealed = signer.seal(b"trusted");
assert!(matches!(
attacker.open(&sealed),
Err(RequestStateError::IntegrityCheckFailed)
));
}
#[test]
fn appended_bytes_are_rejected() {
let codec = RequestStateCodec::new(b"key-key-key-key-key-key-key-key!!".to_vec());
let mut sealed = codec.seal(b"state");
sealed.push('x');
assert!(codec.open(&sealed).is_err());
}
#[test]
fn wrong_version_prefix_is_malformed() {
let codec = RequestStateCodec::new(b"key".to_vec());
let sealed = codec.seal(b"state");
let bumped = sealed.replacen("rs1.", "rs2.", 1);
assert!(matches!(
codec.open(&bumped),
Err(RequestStateError::MalformedFormat)
));
}
#[test]
fn missing_sections_are_malformed() {
let codec = RequestStateCodec::new(b"key".to_vec());
assert!(matches!(
codec.open("rs1"),
Err(RequestStateError::MalformedFormat)
));
assert!(matches!(
codec.open("rs1.onlybody"),
Err(RequestStateError::MalformedFormat)
));
assert!(matches!(
codec.open("rs1.a.b.c"),
Err(RequestStateError::MalformedFormat)
));
}
#[test]
fn non_base64_sections_are_invalid_encoding() {
let codec = RequestStateCodec::new(b"key".to_vec());
assert!(matches!(
codec.open("rs1.!!!!.!!!!"),
Err(RequestStateError::InvalidEncoding)
));
}
#[test]
fn debug_does_not_leak_key() {
let codec = RequestStateCodec::new(b"super-secret-key".to_vec());
let rendered = format!("{codec:?}");
assert!(!rendered.contains("super-secret-key"));
assert!(rendered.contains("redacted"));
}
mod associated_data {
use super::*;
#[test]
fn matching_context_opens() {
let codec = RequestStateCodec::new(b"key-key-key-key-key-key-key-key!!".to_vec());
let ctx = b"user:alice|tools/call:weather";
let sealed = codec.seal_with(b"state", &SealOptions::new().associated_data(ctx));
assert_eq!(codec.open_with(&sealed, ctx).unwrap(), b"state");
}
#[test]
fn different_context_is_rejected() {
let codec = RequestStateCodec::new(b"key-key-key-key-key-key-key-key!!".to_vec());
let sealed =
codec.seal_with(b"state", &SealOptions::new().associated_data(b"user:alice"));
assert!(matches!(
codec.open_with(&sealed, b"user:bob"),
Err(RequestStateError::IntegrityCheckFailed)
));
}
#[test]
fn missing_context_is_rejected() {
let codec = RequestStateCodec::new(b"key-key-key-key-key-key-key-key!!".to_vec());
let sealed =
codec.seal_with(b"state", &SealOptions::new().associated_data(b"user:alice"));
assert!(matches!(
codec.open(&sealed),
Err(RequestStateError::IntegrityCheckFailed)
));
}
}
mod ttl {
use super::*;
const KEY: &[u8] = b"ttl-signing-key-ttl-signing-key!!";
#[test]
fn within_ttl_opens() {
let codec = RequestStateCodec::new(KEY.to_vec());
let sealed = codec.seal_at(
b"state",
&SealOptions::new().ttl(Duration::from_secs(60)),
1_000,
);
assert_eq!(codec.open_at(&sealed, &[], 31_000).unwrap(), b"state");
}
#[test]
fn past_ttl_is_expired() {
let codec = RequestStateCodec::new(KEY.to_vec());
let sealed = codec.seal_at(
b"state",
&SealOptions::new().ttl(Duration::from_secs(60)),
1_000,
);
assert!(matches!(
codec.open_at(&sealed, &[], 62_000),
Err(RequestStateError::Expired)
));
}
#[test]
fn no_ttl_never_expires() {
let codec = RequestStateCodec::new(KEY.to_vec());
let sealed = codec.seal_at(b"state", &SealOptions::new(), 1_000);
assert_eq!(codec.open_at(&sealed, &[], i64::MAX).unwrap(), b"state");
}
#[test]
fn ttl_and_associated_data_combine() {
let codec = RequestStateCodec::new(KEY.to_vec());
let ctx = b"user:alice";
let sealed = codec.seal_at(
b"state",
&SealOptions::new()
.associated_data(ctx)
.ttl(Duration::from_secs(60)),
1_000,
);
assert_eq!(codec.open_at(&sealed, ctx, 10_000).unwrap(), b"state");
assert!(matches!(
codec.open_at(&sealed, b"user:bob", 10_000),
Err(RequestStateError::IntegrityCheckFailed)
));
assert!(matches!(
codec.open_at(&sealed, ctx, 99_000),
Err(RequestStateError::Expired)
));
}
}
}