use std::collections::BTreeMap;
use std::fmt;
use std::time::Duration;
use super::{
EnvelopeBinding, EnvelopeError, EnvelopePolicy, EnvelopePurpose, EphemeralEnvelopeProtector,
HEADER_BYTES, OpenedState, ProcessGenerationGuard, SnapshotCloneStance, TAG_BYTES,
};
use crate::McpRequestCancellation;
use crate::crypto::{draw_security_identifier, sha256_bounded};
use crate::partition::{ContinuationPartitionKey, PartitionAuthorization};
use asupersync::Cx;
const HANDLE_BYTES: usize = 40;
const ENTRY_IDENTITY_BYTES: usize = HANDLE_BYTES + 32 + 8;
pub struct ContinuationHandle([u8; HANDLE_BYTES]);
impl ContinuationHandle {
pub fn from_wire(wire: &str) -> Result<Self, ContinuationStoreError> {
if wire.len() != HANDLE_BYTES * 2 {
return Err(ContinuationStoreError::Unavailable);
}
fn digit(byte: u8) -> Result<u8, ContinuationStoreError> {
match byte {
b'0'..=b'9' => Ok(byte - b'0'),
b'a'..=b'f' => Ok(byte - b'a' + 10),
_ => Err(ContinuationStoreError::Unavailable),
}
}
let mut bytes = [0; HANDLE_BYTES];
for (output, pair) in bytes.iter_mut().zip(wire.as_bytes().as_chunks::<2>().0) {
*output = digit(pair[0])? * 16 + digit(pair[1])?;
}
Ok(Self(bytes))
}
pub fn to_wire(&self) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut wire = String::with_capacity(HANDLE_BYTES * 2);
for byte in self.0 {
wire.push(char::from(HEX[usize::from(byte >> 4)]));
wire.push(char::from(HEX[usize::from(byte & 15)]));
}
wire
}
}
impl fmt::Debug for ContinuationHandle {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("ContinuationHandle(<redacted>)")
}
}
#[derive(Clone, Copy, Debug)]
pub struct ContinuationStorePolicy {
maximum_entries: usize,
maximum_bytes: usize,
}
impl Default for ContinuationStorePolicy {
fn default() -> Self {
Self {
maximum_entries: 1024,
maximum_bytes: 8 * 1024 * 1024,
}
}
}
impl ContinuationStorePolicy {
pub fn new(
maximum_entries: usize,
maximum_bytes: usize,
) -> Result<Self, ContinuationStoreError> {
if !(1..=4096).contains(&maximum_entries)
|| !(1..=64 * 1024 * 1024).contains(&maximum_bytes)
{
return Err(ContinuationStoreError::InvalidPolicy);
}
Ok(Self {
maximum_entries,
maximum_bytes,
})
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ContinuationStoreError {
InvalidPolicy,
Capacity,
HandleExhausted,
Unavailable,
Protection(EnvelopeError),
}
impl From<EnvelopeError> for ContinuationStoreError {
fn from(error: EnvelopeError) -> Self {
match error {
EnvelopeError::InvalidEnvelope => Self::Unavailable,
error => Self::Protection(error),
}
}
}
impl fmt::Display for ContinuationStoreError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InvalidPolicy => f.write_str("invalid continuation store policy"),
Self::Capacity => f.write_str("continuation store capacity exhausted"),
Self::HandleExhausted => f.write_str("continuation handle sequence exhausted"),
Self::Unavailable => f.write_str("continuation is unavailable"),
Self::Protection(error) => error.fmt(f),
}
}
}
impl std::error::Error for ContinuationStoreError {}
struct Entry {
binding: [u8; 32],
envelope: Vec<u8>,
expires_at: u64,
owner: McpRequestCancellation,
}
impl Entry {
fn charge(&self) -> usize {
self.envelope.len() + ENTRY_IDENTITY_BYTES
}
}
pub struct EphemeralContinuationStore {
protector: EphemeralEnvelopeProtector,
namespace: String,
policy: ContinuationStorePolicy,
entries: BTreeMap<[u8; HANDLE_BYTES], Entry>,
retained_bytes: usize,
sequence: u64,
}
impl EphemeralContinuationStore {
pub fn new(
cx: &Cx,
guard: &ProcessGenerationGuard,
stance: SnapshotCloneStance,
namespace: &str,
protection: EnvelopePolicy,
policy: ContinuationStorePolicy,
) -> Result<Self, ContinuationStoreError> {
if namespace.is_empty()
|| namespace.len() > 128
|| !namespace
.bytes()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'.' | b'_' | b'-' | b':'))
{
return Err(EnvelopeError::InvalidBinding.into());
}
Ok(Self {
protector: EphemeralEnvelopeProtector::new(
cx,
guard,
stance,
EnvelopePurpose::Continuation,
protection,
)?,
namespace: namespace.to_owned(),
policy,
entries: BTreeMap::new(),
retained_bytes: 0,
sequence: 0,
})
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn retained_bytes(&self) -> usize {
self.retained_bytes
}
pub fn put(
&mut self,
cx: &Cx,
key: &ContinuationPartitionKey,
authorization: &PartitionAuthorization,
owner: &McpRequestCancellation,
plaintext: &[u8],
lifetime: Duration,
) -> Result<ContinuationHandle, ContinuationStoreError> {
self.protector.check(cx)?;
if owner.is_cancel_requested() {
return Err(ContinuationStoreError::Unavailable);
}
let binding = EnvelopeBinding::continuation(key, authorization, &self.namespace)?;
if plaintext.len() > self.protector.policy.plaintext_limit {
return Err(EnvelopeError::TooLarge.into());
}
if lifetime.is_zero() || lifetime > self.protector.policy.lifetime_bound {
return Err(EnvelopeError::InvalidLifetime.into());
}
let charge = HEADER_BYTES + TAG_BYTES + plaintext.len() + ENTRY_IDENTITY_BYTES;
if charge > self.policy.maximum_bytes {
return Err(ContinuationStoreError::Capacity);
}
if self.entries.len() >= self.policy.maximum_entries
|| self
.retained_bytes
.checked_add(charge)
.is_none_or(|total| total > self.policy.maximum_bytes)
{
self.prune(cx)?;
}
let retained = self
.retained_bytes
.checked_add(charge)
.filter(|total| *total <= self.policy.maximum_bytes)
.ok_or(ContinuationStoreError::Capacity)?;
if self.entries.len() >= self.policy.maximum_entries {
return Err(ContinuationStoreError::Capacity);
}
let sequence = self
.sequence
.checked_add(1)
.ok_or(ContinuationStoreError::HandleExhausted)?;
let expires_at = self
.protector
.elapsed()?
.checked_add(
u64::try_from(lifetime.as_nanos()).map_err(|_| EnvelopeError::InvalidLifetime)?,
)
.ok_or(EnvelopeError::InvalidLifetime)?;
let random = draw_security_identifier().map_err(|_| EnvelopeError::EntropyUnavailable)?;
let mut handle = ContinuationHandle([0; HANDLE_BYTES]);
handle.0[..32].copy_from_slice(random.as_bytes());
handle.0[32..].copy_from_slice(&sequence.to_be_bytes());
if self.entries.contains_key(&handle.0) {
return Err(ContinuationStoreError::HandleExhausted);
}
self.sequence = sequence;
let sealed_binding = bind_handle(binding, &handle)?;
let envelope = self
.protector
.seal(cx, &sealed_binding, plaintext, lifetime)?;
self.protector.check(cx)?;
if owner.is_cancel_requested() || self.protector.elapsed()? >= expires_at {
return Err(ContinuationStoreError::Unavailable);
}
debug_assert_eq!(envelope.len() + ENTRY_IDENTITY_BYTES, charge);
self.entries.insert(
handle.0,
Entry {
binding: binding.digest,
envelope,
expires_at,
owner: owner.clone(),
},
);
self.retained_bytes = retained;
Ok(handle)
}
pub fn take(
&mut self,
cx: &Cx,
key: &ContinuationPartitionKey,
authorization: &PartitionAuthorization,
handle: &ContinuationHandle,
) -> Result<OpenedState, ContinuationStoreError> {
self.protector.check(cx)?;
let binding = EnvelopeBinding::continuation(key, authorization, &self.namespace)?;
let entry = self
.entries
.get(&handle.0)
.ok_or(ContinuationStoreError::Unavailable)?;
if entry.binding != binding.digest
|| entry.owner.is_cancel_requested()
|| self.protector.elapsed()? >= entry.expires_at
{
return Err(ContinuationStoreError::Unavailable);
}
let opened = self
.protector
.open(cx, &bind_handle(binding, handle)?, &entry.envelope)?;
self.protector.check(cx)?;
if entry.owner.is_cancel_requested() || self.protector.elapsed()? >= entry.expires_at {
return Err(ContinuationStoreError::Unavailable);
}
let retired = self
.entries
.remove(&handle.0)
.ok_or(ContinuationStoreError::Unavailable)?;
self.retained_bytes -= retired.charge();
Ok(opened)
}
pub fn prune(&mut self, cx: &Cx) -> Result<usize, ContinuationStoreError> {
self.protector.check(cx)?;
let now = self.protector.elapsed()?;
let before = self.entries.len();
self.entries.retain(|_, entry| {
let keep = entry.expires_at > now && !entry.owner.is_cancel_requested();
if !keep {
self.retained_bytes -= entry.charge();
}
keep
});
Ok(before - self.entries.len())
}
pub fn rotate(&mut self, cx: &Cx) -> Result<u64, ContinuationStoreError> {
Ok(self.protector.rotate(cx)?)
}
pub fn close(&mut self) {
self.protector.close();
self.entries.clear();
self.retained_bytes = 0;
}
}
impl fmt::Debug for EphemeralContinuationStore {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EphemeralContinuationStore")
.field("entries", &self.entries.len())
.field("retained_bytes", &self.retained_bytes)
.finish_non_exhaustive()
}
}
fn bind_handle(
binding: EnvelopeBinding,
handle: &ContinuationHandle,
) -> Result<EnvelopeBinding, EnvelopeError> {
let mut input = Vec::with_capacity(128);
input.extend_from_slice(b"fastmcp/continuation-handle/v1\0");
input.extend_from_slice(&binding.digest);
input.extend_from_slice(&handle.0);
Ok(EnvelopeBinding {
purpose: EnvelopePurpose::Continuation,
digest: sha256_bounded(&input, 128)
.map_err(|_| EnvelopeError::InvalidBinding)?
.into_bytes(),
})
}
#[cfg(test)]
mod tests;