pub mod asynchronous;
use std::fmt;
use std::io::Write;
use std::time::Instant;
use asupersync::Cx;
use fastmcp_core::crypto::sha256_bounded;
use fastmcp_core::partition::{CredentialStoreKey, PartitionAuthorization};
use super::{
OAuthClient, OAuthClientConfiguration, OAuthCredentials, OAuthError, TOKEN_TIMEOUT,
admit_token_response, encode_form, operation_deadline, valid_opaque, validate_scopes,
};
use crate::http_auth::secure_file::SecureAtomicFile;
use crate::http_auth::secure_file::slot::coordinator::{
CoordinatedCredentialSlot, CoordinatedSlotError, CredentialAnchorBinding,
CredentialCommitAnchor,
};
use crate::http_auth::secure_file::slot::{SlotRecoveryOutcome, SlotRevision};
const MAGIC: &[u8; 8] = b"FCPORF01";
const MAX_CONFIGURATION_BYTES: usize = 256 * 1024;
pub const MAX_ENCODED_REFRESH_GRANT_BYTES: usize = 16 * 1024;
pub const MAX_PROTECTED_REFRESH_GRANT_BYTES: usize = 64 * 1024;
#[derive(Clone, Copy, Eq, PartialEq)]
pub struct OAuthGrantBinding([u8; 32]);
impl OAuthGrantBinding {
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
impl fmt::Debug for OAuthGrantBinding {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OAuthGrantBinding").finish_non_exhaustive()
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum OAuthGrantProtectionError {
Unavailable,
InvalidEnvelope,
Cancelled,
TooLarge,
EncodingFailed,
}
pub struct OAuthGrantEncoding<'a> {
binding: OAuthGrantBinding,
refresh_token: &'a str,
scopes: &'a [String],
}
impl OAuthGrantEncoding<'_> {
pub fn encoded_len(&self) -> usize {
8 + 32
+ 4
+ self.refresh_token.len()
+ 1
+ self
.scopes
.iter()
.map(|scope| 2 + scope.len())
.sum::<usize>()
}
pub fn write_to(&self, writer: &mut dyn Write) -> Result<(), OAuthGrantProtectionError> {
let write = |writer: &mut dyn Write, bytes: &[u8]| {
writer
.write_all(bytes)
.map_err(|_| OAuthGrantProtectionError::EncodingFailed)
};
write(writer, MAGIC)?;
write(writer, self.binding.as_bytes())?;
write(writer, &(self.refresh_token.len() as u32).to_be_bytes())?;
write(writer, self.refresh_token.as_bytes())?;
write(writer, &[self.scopes.len() as u8])?;
for scope in self.scopes {
write(writer, &(scope.len() as u16).to_be_bytes())?;
write(writer, scope.as_bytes())?;
}
Ok(())
}
}
pub trait OAuthGrantProtector: Send {
type Plaintext: AsRef<[u8]>;
fn seal(
&mut self,
cx: &Cx,
binding: &OAuthGrantBinding,
grant: &OAuthGrantEncoding<'_>,
) -> Result<Vec<u8>, OAuthGrantProtectionError>;
fn open(
&mut self,
cx: &Cx,
binding: &OAuthGrantBinding,
protected: &[u8],
) -> Result<Self::Plaintext, OAuthGrantProtectionError>;
}
pub struct OAuthRefreshGrant {
configuration: OAuthClientConfiguration,
refresh_token: String,
scopes: Vec<String>,
}
impl OAuthRefreshGrant {
pub fn scopes(&self) -> &[String] {
&self.scopes
}
}
impl OAuthCredentials {
pub fn take_refresh_grant(&mut self) -> Result<OAuthRefreshGrant, OAuthError> {
let refresh = self
.refresh_token
.as_deref()
.ok_or(OAuthError::RefreshUnavailable)?;
validate_refresh(refresh, &self.scopes, &self.configuration)
.map_err(|_| OAuthError::InvalidTokenResponse)?;
let configuration = self.configuration.clone();
let scopes = self.scopes.clone();
let refresh_token = self
.refresh_token
.take()
.ok_or(OAuthError::RefreshUnavailable)?;
Ok(OAuthRefreshGrant {
configuration,
refresh_token,
scopes,
})
}
}
impl OAuthClient {
pub async fn refresh_grant(
&self,
cx: &Cx,
grant: OAuthRefreshGrant,
) -> Result<OAuthCredentials, OAuthError> {
let deadline = operation_deadline(cx, TOKEN_TIMEOUT)?;
if grant.configuration != self.configuration {
return Err(OAuthError::CredentialBindingMismatch);
}
let scope = grant.scopes.join(" ");
let mut fields = vec![
("grant_type", "refresh_token"),
("client_id", self.configuration.client_id.as_str()),
("refresh_token", grant.refresh_token.as_str()),
("resource", self.configuration.resource.as_str()),
];
if !scope.is_empty() {
fields.push(("scope", scope.as_str()));
}
let body = encode_form(&fields)?;
let started = Instant::now();
let response = self.exchange(cx, deadline, body).await?;
if cx.checkpoint().is_err() {
return Err(OAuthError::Cancelled);
}
if cx.now() >= deadline {
return Err(OAuthError::TimedOut);
}
let mut credentials =
admit_token_response(&self.configuration, &grant.scopes, &response, started)?;
if credentials.refresh_token.is_none() {
credentials.refresh_token = Some(grant.refresh_token);
}
if cx.checkpoint().is_err() {
return Err(OAuthError::Cancelled);
}
if cx.now() >= deadline {
return Err(OAuthError::TimedOut);
}
Ok(credentials)
}
pub fn accepts_credentials(&self, credentials: &OAuthCredentials) -> bool {
credentials.configuration == self.configuration
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum OAuthRefreshStoreError {
ContextStopped,
ConfigurationMismatch,
RefreshUnavailable,
InvalidGrant,
TooLarge,
GenerationExhausted,
RevisionMismatch,
Protection(OAuthGrantProtectionError),
Storage(CoordinatedSlotError),
}
impl fmt::Display for OAuthRefreshStoreError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ContextStopped => f.write_str("OAuth refresh custody context stopped"),
Self::ConfigurationMismatch => {
f.write_str("OAuth refresh custody configuration mismatch")
}
Self::RefreshUnavailable => {
f.write_str("OAuth grant has no transferable refresh token")
}
Self::InvalidGrant => f.write_str("stored OAuth refresh grant is invalid"),
Self::TooLarge => f.write_str("stored OAuth refresh grant exceeds its bound"),
Self::GenerationExhausted => f.write_str("OAuth refresh custody generation exhausted"),
Self::RevisionMismatch => f.write_str("OAuth refresh custody revision mismatch"),
Self::Protection(_) => f.write_str("OAuth refresh grant protection failed"),
Self::Storage(error) => error.fmt(f),
}
}
}
impl std::error::Error for OAuthRefreshStoreError {}
impl From<CoordinatedSlotError> for OAuthRefreshStoreError {
fn from(error: CoordinatedSlotError) -> Self {
Self::Storage(error)
}
}
impl From<OAuthGrantProtectionError> for OAuthRefreshStoreError {
fn from(error: OAuthGrantProtectionError) -> Self {
Self::Protection(error)
}
}
pub struct OAuthRefreshStore<A, P> {
slot: CoordinatedCredentialSlot<A>,
protector: P,
configuration: OAuthClientConfiguration,
configuration_digest: [u8; 32],
anchor_binding: CredentialAnchorBinding,
}
impl<A: CredentialCommitAnchor, P: OAuthGrantProtector> OAuthRefreshStore<A, P> {
pub fn open(
cx: &Cx,
file: SecureAtomicFile,
key: &CredentialStoreKey,
authorization: &PartitionAuthorization,
namespace: &str,
anchor: A,
protector: P,
client: &OAuthClient,
) -> Result<(Self, Option<SlotRecoveryOutcome>), OAuthRefreshStoreError> {
checkpoint(cx)?;
let configuration_digest = configuration_digest(&client.configuration)?;
let anchor_binding = CredentialAnchorBinding::for_store(namespace, key, authorization)?;
let (slot, recovery) =
CoordinatedCredentialSlot::open(cx, file, key, authorization, namespace, anchor)?;
Ok((
Self {
slot,
protector,
configuration: client.configuration.clone(),
configuration_digest,
anchor_binding,
},
recovery,
))
}
pub fn revision(&self) -> Option<SlotRevision> {
self.slot.revision()
}
pub fn requires_recovery(&self) -> bool {
self.slot.requires_recovery()
}
pub fn store_refresh(
&mut self,
cx: &Cx,
authorization: &PartitionAuthorization,
expected: Option<SlotRevision>,
credentials: &mut OAuthCredentials,
) -> Result<SlotRevision, OAuthRefreshStoreError> {
checkpoint(cx)?;
if credentials.configuration != self.configuration {
return Err(OAuthRefreshStoreError::ConfigurationMismatch);
}
if expected != self.slot.revision() {
return Err(OAuthRefreshStoreError::RevisionMismatch);
}
let refresh_token = credentials
.refresh_token
.as_deref()
.ok_or(OAuthRefreshStoreError::RefreshUnavailable)?;
validate_refresh(refresh_token, &credentials.scopes, &self.configuration)?;
let generation = expected
.map_or(0, SlotRevision::generation)
.checked_add(1)
.ok_or(OAuthRefreshStoreError::GenerationExhausted)?;
drop(self.slot.load(cx, authorization)?);
let binding = self.binding(generation)?;
let encoding = OAuthGrantEncoding {
binding,
refresh_token,
scopes: &credentials.scopes,
};
let protected = self.protector.seal(cx, &binding, &encoding)?;
self.admit_protected(&protected)?;
checkpoint(cx)?;
let _transferred = credentials
.refresh_token
.take()
.ok_or(OAuthRefreshStoreError::RefreshUnavailable)?;
Ok(self.slot.replace(cx, authorization, expected, &protected)?)
}
pub fn take_refresh(
&mut self,
cx: &Cx,
authorization: &PartitionAuthorization,
) -> Result<Option<OAuthRefreshGrant>, OAuthRefreshStoreError> {
checkpoint(cx)?;
let Some(protected) = self.slot.load(cx, authorization)? else {
return Ok(None);
};
self.admit_protected(&protected)?;
let revision = self
.slot
.revision()
.ok_or(OAuthRefreshStoreError::InvalidGrant)?;
let binding = self.binding(revision.generation())?;
let plaintext = self.protector.open(cx, &binding, &protected)?;
let grant = decode_grant(&self.configuration, binding, plaintext.as_ref())?;
drop(plaintext);
checkpoint(cx)?;
let committed = self.slot.take(cx, authorization, revision)?;
let consumed = committed
.into_consumed()
.ok_or(OAuthRefreshStoreError::InvalidGrant)?;
if consumed != protected {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
Ok(Some(grant))
}
pub fn invalidate(
&mut self,
cx: &Cx,
authorization: &PartitionAuthorization,
) -> Result<SlotRevision, OAuthRefreshStoreError> {
Ok(self.slot.invalidate(cx, authorization)?)
}
fn binding(&self, generation: u64) -> Result<OAuthGrantBinding, OAuthRefreshStoreError> {
grant_binding(
self.configuration_digest,
self.anchor_binding.as_bytes(),
generation,
)
}
fn admit_protected(&self, protected: &[u8]) -> Result<(), OAuthRefreshStoreError> {
if protected.is_empty()
|| protected.len() > MAX_PROTECTED_REFRESH_GRANT_BYTES
|| protected.len() > self.slot.maximum_payload_bytes()
{
return Err(OAuthRefreshStoreError::TooLarge);
}
Ok(())
}
}
fn checkpoint(cx: &Cx) -> Result<(), OAuthRefreshStoreError> {
cx.checkpoint()
.map_err(|_| OAuthRefreshStoreError::ContextStopped)
}
fn validate_refresh(
token: &str,
scopes: &[String],
configuration: &OAuthClientConfiguration,
) -> Result<(), OAuthRefreshStoreError> {
if !valid_opaque(token, super::MAX_CODE_BYTES) {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
validate_scopes(scopes).map_err(|_| OAuthRefreshStoreError::InvalidGrant)?;
if scopes
.iter()
.any(|scope| !configuration.scopes.contains(scope))
{
return Err(OAuthRefreshStoreError::InvalidGrant);
}
Ok(())
}
fn field(output: &mut Vec<u8>, bytes: &[u8]) -> Result<(), OAuthRefreshStoreError> {
let length = u32::try_from(bytes.len()).map_err(|_| OAuthRefreshStoreError::TooLarge)?;
let total = output
.len()
.checked_add(4)
.and_then(|n| n.checked_add(bytes.len()))
.ok_or(OAuthRefreshStoreError::TooLarge)?;
if total > MAX_CONFIGURATION_BYTES {
return Err(OAuthRefreshStoreError::TooLarge);
}
output.extend_from_slice(&length.to_be_bytes());
output.extend_from_slice(bytes);
Ok(())
}
fn configuration_digest(
configuration: &OAuthClientConfiguration,
) -> Result<[u8; 32], OAuthRefreshStoreError> {
if configuration.resource_tls.is_some() != configuration.resource_tls_fingerprint.is_some() {
return Err(OAuthRefreshStoreError::ConfigurationMismatch);
}
let mut bytes = b"fastmcp/oauth-refresh-configuration/v1\0".to_vec();
for value in [
configuration.issuer.as_str(),
configuration.authorization_endpoint.as_str(),
configuration.token_endpoint.as_str(),
configuration.resource.as_str(),
configuration.client_id.as_str(),
] {
field(&mut bytes, value.as_bytes())?;
}
field(
&mut bytes,
&[u8::from(configuration.revocation_endpoint.is_some())],
)?;
if let Some(endpoint) = &configuration.revocation_endpoint {
field(&mut bytes, endpoint.as_str().as_bytes())?;
}
field(
&mut bytes,
&configuration.authorization_timeout.as_nanos().to_be_bytes(),
)?;
field(
&mut bytes,
&configuration
.max_access_token_lifetime
.as_nanos()
.to_be_bytes(),
)?;
field(
&mut bytes,
&(configuration.scopes.len() as u32).to_be_bytes(),
)?;
for scope in &configuration.scopes {
field(&mut bytes, scope.as_bytes())?;
}
field(
&mut bytes,
&(configuration.extra_root_certificates.len() as u32).to_be_bytes(),
)?;
for certificate in &configuration.extra_root_certificates {
field(&mut bytes, certificate)?;
}
field(
&mut bytes,
&[u8::from(configuration.resource_tls_fingerprint.is_some())],
)?;
if let Some(fingerprint) = configuration.resource_tls_fingerprint {
field(&mut bytes, &fingerprint)?;
}
sha256_bounded(&bytes, MAX_CONFIGURATION_BYTES)
.map(|digest| digest.into_bytes())
.map_err(|_| OAuthRefreshStoreError::TooLarge)
}
fn grant_binding(
configuration: [u8; 32],
anchor: &[u8; 32],
generation: u64,
) -> Result<OAuthGrantBinding, OAuthRefreshStoreError> {
if generation == 0 {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
let mut bytes = b"fastmcp/oauth-refresh-custody/v1\0".to_vec();
bytes.extend_from_slice(&configuration);
bytes.extend_from_slice(anchor);
bytes.extend_from_slice(&generation.to_be_bytes());
sha256_bounded(&bytes, 128)
.map(|digest| OAuthGrantBinding(digest.into_bytes()))
.map_err(|_| OAuthRefreshStoreError::InvalidGrant)
}
fn take<'a>(input: &mut &'a [u8], count: usize) -> Result<&'a [u8], OAuthRefreshStoreError> {
if count > input.len() {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
let (head, tail) = input.split_at(count);
*input = tail;
Ok(head)
}
fn decode_grant(
configuration: &OAuthClientConfiguration,
binding: OAuthGrantBinding,
mut input: &[u8],
) -> Result<OAuthRefreshGrant, OAuthRefreshStoreError> {
if input.len() > MAX_ENCODED_REFRESH_GRANT_BYTES {
return Err(OAuthRefreshStoreError::TooLarge);
}
if take(&mut input, 8)? != MAGIC || take(&mut input, 32)? != binding.as_bytes() {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
let length = u32::from_be_bytes(
take(&mut input, 4)?
.try_into()
.map_err(|_| OAuthRefreshStoreError::InvalidGrant)?,
) as usize;
if length > super::MAX_CODE_BYTES {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
let token = std::str::from_utf8(take(&mut input, length)?)
.map_err(|_| OAuthRefreshStoreError::InvalidGrant)?;
let count = usize::from(take(&mut input, 1)?[0]);
if count > 32 {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
let mut scopes = Vec::with_capacity(count);
for _ in 0..count {
let length = usize::from(u16::from_be_bytes(
take(&mut input, 2)?
.try_into()
.map_err(|_| OAuthRefreshStoreError::InvalidGrant)?,
));
if length > 256 {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
let scope = std::str::from_utf8(take(&mut input, length)?)
.map_err(|_| OAuthRefreshStoreError::InvalidGrant)?;
scopes.push(scope.to_owned());
}
if !input.is_empty() {
return Err(OAuthRefreshStoreError::InvalidGrant);
}
validate_refresh(token, &scopes, configuration)?;
Ok(OAuthRefreshGrant {
configuration: configuration.clone(),
refresh_token: token.to_owned(),
scopes,
})
}
#[cfg(test)]
mod tests;
#[cfg(test)]
mod resource_trust_tests {
use super::super::tests as native;
use super::*;
#[test]
fn persistent_binding_distinguishes_resource_ca_from_issuer_only_trust() {
let plain = native::config();
let resource = plain
.clone()
.with_resource_root_certificate(native::test_root())
.unwrap();
let same = plain
.clone()
.with_resource_root_certificate(native::test_root())
.unwrap();
let issuer = plain
.clone()
.with_extra_root_certificate(native::test_root())
.unwrap();
assert_eq!(
configuration_digest(&resource).unwrap(),
configuration_digest(&same).unwrap()
);
assert_ne!(
configuration_digest(&resource).unwrap(),
configuration_digest(&plain).unwrap()
);
assert_ne!(
configuration_digest(&resource).unwrap(),
configuration_digest(&issuer).unwrap()
);
let scopes = vec!["tools:read".to_owned()];
let binding = grant_binding(configuration_digest(&resource).unwrap(), &[8; 32], 1).unwrap();
let encoding = OAuthGrantEncoding {
binding,
refresh_token: "private-resource-refresh",
scopes: &scopes,
};
let mut bytes = Vec::new();
encoding.write_to(&mut bytes).unwrap();
let changed = grant_binding(configuration_digest(&issuer).unwrap(), &[8; 32], 1).unwrap();
assert!(decode_grant(&issuer, changed, &bytes).is_err());
assert!(decode_grant(&resource, binding, &bytes).is_ok());
}
#[test]
fn inconsistent_resource_trust_fingerprint_refuses_custody_binding() {
let mut configuration = native::config()
.with_resource_root_certificate(native::test_root())
.unwrap();
assert!(configuration_digest(&configuration).is_ok());
configuration.resource_tls_fingerprint = None;
assert_eq!(
configuration_digest(&configuration),
Err(OAuthRefreshStoreError::ConfigurationMismatch)
);
}
}