//! Interactive OAuth authorization for explicitly configured native public clients.
//!
//! The caller supplies trusted issuer/endpoint configuration and the browser
//! launcher. The driver binds an IP-literal loopback listener *before* invoking
//! that launcher, admits one issuer- and state-bound callback, and redeems its
//! code with S256 PKCE over HTTPS. It creates no runtime or background task.
//!
//! This is the preregistered public-client slice of AUTH-07, not discovery,
//! dynamic registration or OIDC authentication. The Linux `persistence` adapter
//! connects refresh grants to caller-supplied durable protection and independently
//! anchored storage; it does not qualify those deployment providers. Credential
//! and parsed token-response types omit serialization and diagnostic formatting.
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use std::future::{Future, poll_fn};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::task::Poll;
use std::time::{Duration, Instant};
use asupersync::Cx;
use asupersync::channel::oneshot;
use asupersync::http::h1::{HttpClient, Method, RedirectPolicy, RetryPolicy};
use asupersync::io::{AsyncReadExt, AsyncWriteExt};
use asupersync::net::{TcpListener, TcpStream};
use asupersync::time::Sleep;
use asupersync::types::Time;
use fastmcp_core::crypto::{
HmacSha256Key, HmacSha256Tag, SecurityIdentifier, draw_hmac_sha256_key,
draw_security_identifier, sha256_bounded,
};
use fastmcp_core::{AccessToken, CanonicalResourceId, CanonicalResourceIdPolicy};
use serde::{Deserialize, Deserializer};
use super::{BoundBearerCredential, CanonicalHttpUrl};
/// Anchored refresh-grant custody and native renewal after a Linux file reopen.
#[cfg(target_os = "linux")]
pub mod persistence;
/// RFC 7009 remote token revocation for explicitly trusted native clients.
pub mod revocation;
const CALLBACK_PATH: &str = "/oauth/callback";
const STATE_DOMAIN: &[u8] = b"fastmcp/oauth-loopback-state/v1\0";
const MAX_CALLBACK_BYTES: usize = 16 * 1024;
const MAX_CALLBACK_CONNECTIONS: usize = 32;
const MAX_FORM_FIELDS: usize = 32;
const MAX_CODE_BYTES: usize = 4096;
const MAX_TOKEN_RESPONSE_BYTES: usize = 64 * 1024;
const MAX_FORM_BYTES: usize = 64 * 1024;
const CALLBACK_READ_TIMEOUT: Duration = Duration::from_secs(5);
const TOKEN_TIMEOUT: Duration = Duration::from_secs(30);
/// Sanitized errors. No variant retains a code, token, callback URL, response
/// body, or a third-party transport error that could reflect credentials.
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum OAuthError {
InvalidConfiguration,
RuntimeTimerUnavailable,
RuntimeCapabilityUnavailable,
Cancelled,
TimedOut,
RandomSourceUnavailable,
CallbackBindFailed,
BrowserLaunchFailed,
CallbackRejected,
IssuerMismatch,
AuthorizationDenied,
CallbackLimitExceeded,
TransportFailed,
TokenEndpointRejected,
InvalidTokenResponse,
ScopeExpansion,
ExpiredCredential,
CredentialBindingMismatch,
RefreshUnavailable,
}
impl fmt::Display for OAuthError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
Self::InvalidConfiguration => "invalid native OAuth configuration",
Self::RuntimeTimerUnavailable => "OAuth requires the caller's timer capability",
Self::RuntimeCapabilityUnavailable => {
"OAuth requires the caller's I/O and entropy authority"
}
Self::Cancelled => "OAuth operation cancelled",
Self::TimedOut => "OAuth operation deadline exceeded",
Self::RandomSourceUnavailable => "OAuth security randomness unavailable",
Self::CallbackBindFailed => "OAuth loopback listener could not bind",
Self::BrowserLaunchFailed => "OAuth browser launcher failed",
Self::CallbackRejected => "OAuth callback rejected",
Self::IssuerMismatch => "OAuth callback issuer does not match the configured issuer",
Self::AuthorizationDenied => "OAuth authorization was denied",
Self::CallbackLimitExceeded => "OAuth callback admission limit exceeded",
Self::TransportFailed => "OAuth token transport failed",
Self::TokenEndpointRejected => "OAuth token endpoint rejected the exchange",
Self::InvalidTokenResponse => "OAuth token response rejected",
Self::ScopeExpansion => "OAuth response expanded the requested scopes",
Self::ExpiredCredential => "OAuth credential already expired",
Self::CredentialBindingMismatch => {
"OAuth credential belongs to a different client binding"
}
Self::RefreshUnavailable => "OAuth credential has no reusable refresh token",
})
}
}
impl std::error::Error for OAuthError {}
/// The resource-identifier policy for a configured MCP endpoint. An endpoint
/// at the origin root (`https://mcp.example.com`, a canonical server URI MCP
/// lists as valid) needs the root policy; every other endpoint keeps the safe
/// non-root, no-trailing-slash default.
pub(crate) fn endpoint_resource_policy(resource: &CanonicalHttpUrl) -> CanonicalResourceIdPolicy {
if resource.path() == "/" {
CanonicalResourceIdPolicy::root_endpoint()
} else {
CanonicalResourceIdPolicy::DEFAULT
}
}
/// Immutable, administrator-supplied configuration for an RFC 8252 native
/// public client. The registration must allow the `/oauth/callback` path on
/// an IP-literal loopback URI with an ephemeral port. The authorization server
/// must return RFC 9207 `iss` on both successful and unsuccessful callbacks.
///
/// This constructor does not establish trust in URLs obtained from a peer.
/// Discovery and registration must be validated before constructing this value.
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct OAuthClientConfiguration {
issuer: String,
authorization_endpoint: CanonicalHttpUrl,
token_endpoint: CanonicalHttpUrl,
revocation_endpoint: Option<CanonicalHttpUrl>,
resource: CanonicalHttpUrl,
client_id: String,
scopes: Vec<String>,
authorization_timeout: Duration,
max_access_token_lifetime: Duration,
extra_root_certificates: Vec<Vec<u8>>,
resource_tls: Option<crate::http_executor::ResourceTlsTrust>,
// Stable, domain-separated digest of the admitted ordered resource roots.
// Updated only after ResourceTlsTrust admission succeeds; persistent OAuth
// bindings must not depend on Debug output or duplicate all DER buffers.
resource_tls_fingerprint: Option<[u8; 32]>,
}
impl OAuthClientConfiguration {
pub fn from_trusted_endpoints(
issuer: impl Into<String>,
authorization_endpoint: CanonicalHttpUrl,
token_endpoint: CanonicalHttpUrl,
resource: CanonicalHttpUrl,
client_id: impl Into<String>,
scopes: Vec<String>,
) -> Result<Self, OAuthError> {
let issuer = issuer.into();
let issuer_url =
CanonicalHttpUrl::parse(&issuer).map_err(|_| OAuthError::InvalidConfiguration)?;
for endpoint in [&issuer_url, &authorization_endpoint, &token_endpoint] {
if endpoint.scheme() != "https"
|| endpoint.has_userinfo()
|| endpoint.query().is_some()
|| endpoint.fragment().is_some()
{
return Err(OAuthError::InvalidConfiguration);
}
}
// Keep the original issuer spelling for the exact RFC 9207 comparison;
// canonical network identity must not replace that comparison.
if issuer.len() > 4096 || issuer.chars().any(|c| c.is_whitespace() || c.is_control()) {
return Err(OAuthError::InvalidConfiguration);
}
if resource.scheme() != "https" {
return Err(OAuthError::InvalidConfiguration);
}
CanonicalResourceId::parse_for_endpoint(
resource.as_str(),
&resource,
endpoint_resource_policy(&resource),
)
.map_err(|_| OAuthError::InvalidConfiguration)?;
let client_id = client_id.into();
if client_id.is_empty() || client_id.len() > 1024 || client_id.chars().any(char::is_control)
{
return Err(OAuthError::InvalidConfiguration);
}
validate_scopes(&scopes).map_err(|_| OAuthError::InvalidConfiguration)?;
Ok(Self {
issuer,
authorization_endpoint,
token_endpoint,
revocation_endpoint: None,
resource,
client_id,
scopes,
authorization_timeout: Duration::from_secs(300),
max_access_token_lifetime: Duration::from_secs(3600),
extra_root_certificates: Vec::new(),
resource_tls: None,
resource_tls_fingerprint: None,
})
}
/// Sets the single deadline covering bind, browser launch, callbacks and
/// redemption. Callback traffic cannot reset it. The caller's budget may
/// further shorten it.
pub fn with_authorization_timeout(mut self, timeout: Duration) -> Result<Self, OAuthError> {
if timeout.is_zero() || timeout > Duration::from_mins(15) {
return Err(OAuthError::InvalidConfiguration);
}
self.authorization_timeout = timeout;
Ok(self)
}
/// Sets a local upper bound on access-token reuse, including responses
/// omitting `expires_in`. This is a client safety limit, not an assertion
/// about the issuer's actual expiration policy.
pub fn with_max_access_token_lifetime(
mut self,
lifetime: Duration,
) -> Result<Self, OAuthError> {
if lifetime.is_zero() || lifetime > Duration::from_hours(24) {
return Err(OAuthError::InvalidConfiguration);
}
self.max_access_token_lifetime = lifetime;
Ok(self)
}
/// Adds an explicitly trusted private CA for this client's token endpoint.
/// This does not disable hostname/certificate verification or alter the
/// host's browser trust store. The exact root bytes become part of the
/// immutable credential binding, so refresh cannot switch trust policies.
pub fn with_extra_root_certificate(
mut self,
certificate: asupersync::tls::Certificate,
) -> Result<Self, OAuthError> {
let der = certificate.as_der();
if der.is_empty()
|| der.len() > 16 * 1024
|| self.extra_root_certificates.len() >= 8
|| self
.extra_root_certificates
.iter()
.any(|existing| existing.as_slice() == der)
{
return Err(OAuthError::InvalidConfiguration);
}
asupersync::tls::RootCertStore::empty()
.add(&certificate)
.map_err(|_| OAuthError::InvalidConfiguration)?;
self.extra_root_certificates.push(der.to_vec());
Ok(self)
}
/// Adds a private CA for this configuration's exact MCP resource. Managed
/// sessions retain this trust through login, refresh and protected POSTs.
/// It does not grant trust to token endpoints or the host's browser.
pub fn with_resource_root_certificate(
mut self,
certificate: asupersync::tls::Certificate,
) -> Result<Self, OAuthError> {
let der = certificate.as_der();
if der.is_empty() || der.len() > 16 * 1024 {
return Err(OAuthError::InvalidConfiguration);
}
let mut binding = b"fastmcp/oauth-resource-tls/v1\0".to_vec();
binding.extend_from_slice(&self.resource_tls_fingerprint.unwrap_or([0; 32]));
binding.extend_from_slice(&(der.len() as u32).to_be_bytes());
binding.extend_from_slice(der);
let fingerprint = sha256_bounded(&binding, 16 * 1024 + 128)
.map_err(|_| OAuthError::InvalidConfiguration)?
.into_bytes();
crate::http_executor::ResourceTlsTrust::add_root(
&mut self.resource_tls,
self.resource.clone(),
certificate,
)
.map_err(|_| OAuthError::InvalidConfiguration)?;
self.resource_tls_fingerprint = Some(fingerprint);
Ok(self)
}
}
/// An admitted grant bound to one issuer, registration and MCP resource.
/// There is deliberately no `Clone`, `Debug`, `Display`, or serde implementation.
/// Use `bearer_credential()` to supply the resource-bound access token to the
/// existing HTTP client. Refresh-token bytes are not publicly exposed.
pub struct OAuthCredentials {
configuration: OAuthClientConfiguration,
access: BoundBearerCredential,
refresh_token: Option<String>,
scopes: Vec<String>,
expires_at: Instant,
}
impl OAuthCredentials {
pub fn bearer_credential(&self) -> &BoundBearerCredential {
&self.access
}
pub fn scopes(&self) -> &[String] {
&self.scopes
}
pub fn expires_at(&self) -> Instant {
self.expires_at
}
pub fn has_refresh_token(&self) -> bool {
self.refresh_token.is_some()
}
}
/// Native OAuth driver. All I/O is polled under the supplied `Cx`, without
/// creating another runtime, detached task, browser subprocess, or global client.
#[derive(Clone, Debug)]
pub struct OAuthClient {
configuration: OAuthClientConfiguration,
}
impl OAuthClient {
pub fn new(configuration: OAuthClientConfiguration) -> Self {
Self { configuration }
}
pub(crate) fn resource_http_executor(&self) -> crate::http_executor::ModernHttpExecutor {
crate::http_executor::ModernHttpExecutor::new()
.with_resource_tls(self.configuration.resource_tls.clone())
}
/// Runs a preregistered public-client authorization-code flow.
///
/// `launch_browser` is invoked exactly once, after binding the callback
/// listener. The host chooses how to present/open the URL; its future must
/// return after launching, not wait for the OAuth callback. The launcher is
/// part of the host's trusted computing base and must not log the URL.
/// Dropping this future drops its listener, connection and pending exchange.
pub async fn authorize<L, F>(
&self,
cx: &Cx,
launch_browser: L,
) -> Result<OAuthCredentials, OAuthError>
where
L: FnOnce(CanonicalHttpUrl) -> F,
F: Future<Output = Result<(), OAuthError>>,
{
let deadline = operation_deadline(cx, self.configuration.authorization_timeout)?;
if !cx.capabilities().entropy {
return Err(OAuthError::RuntimeCapabilityUnavailable);
}
let listener = within(cx, deadline, bind_loopback()).await?;
let address = listener
.local_addr()
.map_err(|_| OAuthError::CallbackBindFailed)?;
if !address.ip().is_loopback() || address.port() == 0 {
return Err(OAuthError::CallbackBindFailed);
}
let redirect_uri = format!("http://{address}{CALLBACK_PATH}");
let attempt = AuthorizationAttempt::new()?;
let authorization_url = attempt.authorization_url(&self.configuration, &redirect_uri)?;
within(cx, deadline, async {
launch_browser(authorization_url)
.await
.map_err(|_| OAuthError::BrowserLaunchFailed)
})
.await?;
let code = wait_for_code(
cx,
deadline,
&listener,
address,
&attempt,
&self.configuration,
)
.await?;
// A callback can authorize only one POST. Close the listener before
// redemption; neither a duplicate callback nor a network error retries it.
drop(listener);
let verifier = attempt.verifier();
let body = encode_form(&[
("grant_type", "authorization_code"),
("client_id", &self.configuration.client_id),
("code", &code),
("redirect_uri", &redirect_uri),
("code_verifier", &verifier),
("resource", self.configuration.resource.as_str()),
])?;
let started = Instant::now();
let response = self.exchange(cx, deadline, body).await?;
let credentials = admit_token_response(
&self.configuration,
&self.configuration.scopes,
&response,
started,
)?;
if cx.checkpoint().is_err() {
return Err(OAuthError::Cancelled);
}
if cx.now() >= deadline {
return Err(OAuthError::TimedOut);
}
Ok(credentials)
}
/// Renews an access token through the same trusted issuer, registration,
/// resource and client policy that admitted the original grant.
///
/// Exclusive access to the credentials serializes refresh operations. A
/// successful response replaces the complete token pair at once. Narrowed
/// scopes become the next refresh ceiling; they cannot silently re-expand.
/// When the issuer omits a new refresh token, the previous one is retained.
///
/// After dispatch becomes possible, failure or cancellation discards the
/// old refresh token rather than retrying a possibly consumed/rotated token.
/// The previous access token and its original expiry remain unchanged;
/// `has_refresh_token()` becomes false and another login is required for
/// future renewal. Preflight failures leave the entire credential untouched.
pub async fn refresh(
&self,
cx: &Cx,
credentials: &mut OAuthCredentials,
) -> Result<(), OAuthError> {
let deadline = operation_deadline(cx, TOKEN_TIMEOUT)?;
let (body, previous_refresh) = self.prepare_refresh(credentials)?;
let started = Instant::now();
let response = self.exchange(cx, deadline, body).await?;
// Do not publish a new token pair after the caller has cancelled.
if cx.checkpoint().is_err() {
return Err(OAuthError::Cancelled);
}
if cx.now() >= deadline {
return Err(OAuthError::TimedOut);
}
let mut replacement =
self.admit_refresh(credentials, previous_refresh, &response, started)?;
if cx.checkpoint().is_err() {
return Err(OAuthError::Cancelled);
}
if cx.now() >= deadline {
return Err(OAuthError::TimedOut);
}
std::mem::swap(credentials, &mut replacement);
Ok(())
}
fn prepare_refresh(
&self,
credentials: &mut OAuthCredentials,
) -> Result<(String, String), OAuthError> {
if credentials.configuration != self.configuration {
return Err(OAuthError::CredentialBindingMismatch);
}
let previous = credentials
.refresh_token
.as_deref()
.ok_or(OAuthError::RefreshUnavailable)?;
let scope = credentials.scopes.join(" ");
let mut fields = vec![
("grant_type", "refresh_token"),
("client_id", self.configuration.client_id.as_str()),
("refresh_token", previous),
("resource", self.configuration.resource.as_str()),
];
if !scope.is_empty() {
fields.push(("scope", scope.as_str()));
}
let body = encode_form(&fields)?;
// All fallible local validation precedes this ownership transfer. Once
// an exchange is possible, cancellation cannot put this secret back.
let previous = credentials
.refresh_token
.take()
.ok_or(OAuthError::RefreshUnavailable)?;
Ok((body, previous))
}
fn admit_refresh(
&self,
previous: &OAuthCredentials,
previous_refresh: String,
response: &[u8],
started: Instant,
) -> Result<OAuthCredentials, OAuthError> {
let mut replacement =
admit_token_response(&self.configuration, &previous.scopes, response, started)?;
if replacement.refresh_token.is_none() {
replacement.refresh_token = Some(previous_refresh);
}
Ok(replacement)
}
async fn exchange(&self, cx: &Cx, deadline: Time, body: String) -> Result<Vec<u8>, OAuthError> {
let deadline = deadline.min(operation_deadline(cx, TOKEN_TIMEOUT)?);
let mut builder = HttpClient::builder()
.redirect_policy(RedirectPolicy::None)
.retry_policy(RetryPolicy::None)
.no_proxy()
.no_cookie_store()
.max_body_size(MAX_TOKEN_RESPONSE_BYTES)
.max_total_connections(1);
for der in &self.configuration.extra_root_certificates {
builder =
builder.add_root_certificate(asupersync::tls::Certificate::from_der(der.clone()));
}
let client = builder.build();
let response = within(cx, deadline, async {
client
.request(
cx,
Method::Post,
self.configuration.token_endpoint.as_str(),
vec![
(
"Content-Type".to_owned(),
"application/x-www-form-urlencoded".to_owned(),
),
("Accept".to_owned(), "application/json".to_owned()),
("Accept-Encoding".to_owned(), "identity".to_owned()),
("Connection".to_owned(), "close".to_owned()),
],
body.into_bytes(),
)
.await
.map_err(|_| OAuthError::TransportFailed)
})
.await?;
if response.status != 200 {
return Err(OAuthError::TokenEndpointRejected);
}
validate_token_headers(&response.headers)?;
if response.body.len() > MAX_TOKEN_RESPONSE_BYTES {
return Err(OAuthError::InvalidTokenResponse);
}
Ok(response.body)
}
}
struct AuthorizationAttempt {
verifier_material: SecurityIdentifier,
state_key: HmacSha256Key,
}
impl AuthorizationAttempt {
fn new() -> Result<Self, OAuthError> {
Ok(Self {
verifier_material: draw_security_identifier()
.map_err(|_| OAuthError::RandomSourceUnavailable)?,
state_key: draw_hmac_sha256_key().map_err(|_| OAuthError::RandomSourceUnavailable)?,
})
}
fn verifier(&self) -> String {
// 64 unreserved ASCII characters carrying 256 independent random bits.
hex(self.verifier_material.as_bytes())
}
fn state(&self) -> Result<String, OAuthError> {
self.state_key
.authenticate_bounded(STATE_DOMAIN, STATE_DOMAIN.len())
.map(|tag| hex(tag.as_bytes()))
.map_err(|_| OAuthError::InvalidConfiguration)
}
fn accepts_state(&self, state: &str) -> bool {
let Some(bytes) = decode_state(state) else {
return false;
};
self.state_key
.verify_bounded(
STATE_DOMAIN,
STATE_DOMAIN.len(),
&HmacSha256Tag::from_bytes(bytes),
)
.is_ok()
}
fn authorization_url(
&self,
config: &OAuthClientConfiguration,
redirect_uri: &str,
) -> Result<CanonicalHttpUrl, OAuthError> {
let state = self.state()?;
let challenge = pkce_challenge(&self.verifier())?;
let scopes = config.scopes.join(" ");
let mut fields = vec![
("response_type", "code"),
("client_id", config.client_id.as_str()),
("redirect_uri", redirect_uri),
("resource", config.resource.as_str()),
("state", state.as_str()),
("code_challenge", challenge.as_str()),
("code_challenge_method", "S256"),
];
if !scopes.is_empty() {
fields.push(("scope", scopes.as_str()));
}
let query = encode_form(&fields)?;
CanonicalHttpUrl::parse(&format!(
"{}?{query}",
config.authorization_endpoint.as_str()
))
.map_err(|_| OAuthError::InvalidConfiguration)
}
}
async fn bind_loopback() -> Result<TcpListener, OAuthError> {
let ipv4 = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
if let Ok(listener) = TcpListener::bind(ipv4).await {
return Ok(listener);
}
let ipv6 = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 0);
TcpListener::bind(ipv6)
.await
.map_err(|_| OAuthError::CallbackBindFailed)
}
fn operation_deadline(cx: &Cx, timeout: Duration) -> Result<Time, OAuthError> {
if cx.checkpoint().is_err() {
return Err(OAuthError::Cancelled);
}
if !cx.capabilities().io {
return Err(OAuthError::RuntimeCapabilityUnavailable);
}
if cx.timer_driver().is_none() {
return Err(OAuthError::RuntimeTimerUnavailable);
}
let nanos = u64::try_from(timeout.as_nanos()).map_err(|_| OAuthError::InvalidConfiguration)?;
let end = cx
.now()
.as_nanos()
.checked_add(nanos)
.ok_or(OAuthError::InvalidConfiguration)?;
let end = cx
.budget()
.deadline
.map_or(Time::from_nanos(end), |parent| {
parent.min(Time::from_nanos(end))
});
if cx.now() >= end {
return Err(OAuthError::TimedOut);
}
Ok(end)
}
async fn within<T>(
cx: &Cx,
deadline: Time,
future: impl Future<Output = Result<T, OAuthError>>,
) -> Result<T, OAuthError> {
let deadline = cx
.budget()
.deadline
.map_or(deadline, |parent| parent.min(deadline));
let mut future = std::pin::pin!(future);
let sleep = {
let _caller = Cx::set_current(Some(cx.clone()));
Sleep::new(deadline)
};
let mut sleep = std::pin::pin!(sleep);
// A pending cancel-correct receive registers a cancellation wake even when
// the socket or the host's browser future has no traffic of its own.
let (_sender, mut receiver) = oneshot::channel::<()>();
let mut cancelled = std::pin::pin!(receiver.recv(cx));
poll_fn(|task| {
if cx.checkpoint().is_err() {
return Poll::Ready(Err(OAuthError::Cancelled));
}
if cx.now() >= deadline {
return Poll::Ready(Err(OAuthError::TimedOut));
}
// Install the caller only for this poll, never across an await/yield.
let _caller = Cx::set_current(Some(cx.clone()));
if cancelled.as_mut().poll(task).is_ready() {
return Poll::Ready(Err(OAuthError::Cancelled));
}
if sleep.as_mut().poll(task).is_ready() {
return Poll::Ready(Err(OAuthError::TimedOut));
}
let result = future.as_mut().poll(task);
if cx.checkpoint().is_err() {
return Poll::Ready(Err(OAuthError::Cancelled));
}
if cx.now() >= deadline {
return Poll::Ready(Err(OAuthError::TimedOut));
}
result
})
.await
}
async fn wait_for_code(
cx: &Cx,
deadline: Time,
listener: &TcpListener,
address: SocketAddr,
attempt: &AuthorizationAttempt,
config: &OAuthClientConfiguration,
) -> Result<String, OAuthError> {
for _ in 0..MAX_CALLBACK_CONNECTIONS {
let (mut stream, peer) = within(cx, deadline, async {
listener
.accept()
.await
.map_err(|_| OAuthError::CallbackRejected)
})
.await?;
if !peer.ip().is_loopback() {
return Err(OAuthError::CallbackRejected);
}
let read_deadline = deadline.min(operation_deadline(cx, CALLBACK_READ_TIMEOUT)?);
let head = within(cx, read_deadline, read_callback_head(&mut stream)).await;
let outcome = head.and_then(|head| admit_callback(&head, address, attempt, &config.issuer));
let accepted = outcome.is_ok();
let response = callback_response(accepted);
// The browser sees only receipt, not a false claim that token redemption
// succeeded. A failed write does not erase an already-admitted callback.
let write_deadline = deadline.min(operation_deadline(cx, Duration::from_secs(1))?);
let _ = within(cx, write_deadline, async {
stream
.write_all(response.as_bytes())
.await
.map_err(|_| OAuthError::CallbackRejected)
})
.await;
match outcome {
Ok(code) => return Ok(code),
Err(OAuthError::CallbackRejected | OAuthError::TimedOut) => {}
Err(error) => return Err(error),
}
}
Err(OAuthError::CallbackLimitExceeded)
}
async fn read_callback_head(stream: &mut TcpStream) -> Result<Vec<u8>, OAuthError> {
let mut head = Vec::new();
let mut chunk = [0_u8; 1024];
loop {
let count = stream
.read(&mut chunk)
.await
.map_err(|_| OAuthError::CallbackRejected)?;
if count == 0 || count > MAX_CALLBACK_BYTES.saturating_sub(head.len()) {
return Err(OAuthError::CallbackRejected);
}
head.extend_from_slice(&chunk[..count]);
if head.windows(4).any(|window| window == b"\r\n\r\n") {
return Ok(head);
}
}
}
fn admit_callback(
head: &[u8],
address: SocketAddr,
attempt: &AuthorizationAttempt,
issuer: &str,
) -> Result<String, OAuthError> {
if head.len() > MAX_CALLBACK_BYTES {
return Err(OAuthError::CallbackRejected);
}
let head = std::str::from_utf8(head).map_err(|_| OAuthError::CallbackRejected)?;
let Some(headers) = head.strip_suffix("\r\n\r\n") else {
return Err(OAuthError::CallbackRejected);
};
let mut lines = headers.split("\r\n");
let mut request = lines.next().ok_or(OAuthError::CallbackRejected)?.split(' ');
if request.next() != Some("GET") {
return Err(OAuthError::CallbackRejected);
}
let target = request.next().ok_or(OAuthError::CallbackRejected)?;
if request.next() != Some("HTTP/1.1")
|| request.next().is_some()
|| target.bytes().any(|b| !(0x21..=0x7e).contains(&b))
|| target.contains('#')
{
return Err(OAuthError::CallbackRejected);
}
let (path, query) = target.split_once('?').ok_or(OAuthError::CallbackRejected)?;
if path != CALLBACK_PATH {
return Err(OAuthError::CallbackRejected);
}
let mut host = None;
let mut content_length = false;
for (index, line) in lines.enumerate() {
if index >= 64 {
return Err(OAuthError::CallbackRejected);
}
let (name, value) = line.split_once(':').ok_or(OAuthError::CallbackRejected)?;
if !AccessToken::is_valid_http_scheme(name)
|| value.bytes().any(|b| b == 0x7f || (b < 0x20 && b != b'\t'))
{
return Err(OAuthError::CallbackRejected);
}
let value = value.trim_matches([' ', '\t']);
if name.eq_ignore_ascii_case("host") {
if host.replace(value).is_some() {
return Err(OAuthError::CallbackRejected);
}
} else if name.eq_ignore_ascii_case("transfer-encoding") {
return Err(OAuthError::CallbackRejected);
} else if name.eq_ignore_ascii_case("content-length") {
if content_length || value != "0" {
return Err(OAuthError::CallbackRejected);
}
content_length = true;
}
}
if host != Some(address.to_string().as_str()) {
return Err(OAuthError::CallbackRejected);
}
let fields = decode_form(query)?;
if !fields
.get("state")
.is_some_and(|state| attempt.accepts_state(state))
{
return Err(OAuthError::CallbackRejected);
}
if fields.get("iss").map(String::as_str) != Some(issuer) {
return Err(OAuthError::IssuerMismatch);
}
match (fields.get("code"), fields.get("error")) {
(None, Some(error)) if valid_opaque(error, 256) => Err(OAuthError::AuthorizationDenied),
(Some(code), None) if valid_opaque(code, MAX_CODE_BYTES) => Ok(code.clone()),
_ => Err(OAuthError::CallbackRejected),
}
}
fn callback_response(accepted: bool) -> String {
let (status, body) = if accepted {
(
"200 OK",
"Authorization response received. Return to the application.",
)
} else {
("400 Bad Request", "Authorization response rejected.")
};
format!(
"HTTP/1.1 {status}\r\nContent-Type: text/plain; charset=utf-8\r\nContent-Length: {}\r\nCache-Control: no-store\r\nPragma: no-cache\r\nReferrer-Policy: no-referrer\r\nContent-Security-Policy: default-src 'none'\r\nConnection: close\r\n\r\n{body}",
body.len()
)
}
fn valid_opaque(value: &str, maximum: usize) -> bool {
!value.is_empty() && value.len() <= maximum && value.bytes().all(|b| (0x20..=0x7e).contains(&b))
}
fn validate_scopes(scopes: &[String]) -> Result<(), OAuthError> {
if scopes.len() > 32 || scopes.iter().map(String::len).sum::<usize>() > 4096 {
return Err(OAuthError::InvalidTokenResponse);
}
let mut seen = BTreeSet::new();
for scope in scopes {
if scope.is_empty()
|| scope.len() > 256
|| !seen.insert(scope.as_str())
|| !scope
.bytes()
.all(|b| b == 0x21 || (0x23..=0x5b).contains(&b) || (0x5d..=0x7e).contains(&b))
{
return Err(OAuthError::InvalidTokenResponse);
}
}
Ok(())
}
fn encode_form(fields: &[(&str, &str)]) -> Result<String, OAuthError> {
let mut output = String::new();
for (index, (name, value)) in fields.iter().enumerate() {
if index > 0 {
output.push('&');
}
for (part_index, part) in [*name, *value].into_iter().enumerate() {
if part_index == 1 {
output.push('=');
}
for byte in part.bytes() {
if output.len() > MAX_FORM_BYTES - 3 {
return Err(OAuthError::InvalidConfiguration);
}
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' => {
output.push(char::from(byte));
}
b' ' => output.push('+'),
_ => {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
output.push('%');
output.push(char::from(HEX[usize::from(byte >> 4)]));
output.push(char::from(HEX[usize::from(byte & 15)]));
}
}
}
}
}
Ok(output)
}
fn decode_form(query: &str) -> Result<BTreeMap<String, String>, OAuthError> {
if query.len() > MAX_CALLBACK_BYTES {
return Err(OAuthError::CallbackRejected);
}
let mut fields = BTreeMap::new();
for (index, field) in query.split('&').enumerate() {
if index >= MAX_FORM_FIELDS {
return Err(OAuthError::CallbackRejected);
}
let (key, value) = field.split_once('=').ok_or(OAuthError::CallbackRejected)?;
let key = decode_component(key)?;
let value = decode_component(value)?;
if key.is_empty()
|| key.len() > 128
|| value.len() > MAX_CODE_BYTES
|| fields.insert(key, value).is_some()
{
return Err(OAuthError::CallbackRejected);
}
}
Ok(fields)
}
fn decode_component(input: &str) -> Result<String, OAuthError> {
let mut decoded = Vec::with_capacity(input.len());
let mut bytes = input.bytes();
while let Some(byte) = bytes.next() {
let byte = match byte {
b'+' => b' ',
b'%' => {
let high = bytes
.next()
.and_then(hex_digit)
.ok_or(OAuthError::CallbackRejected)?;
let low = bytes
.next()
.and_then(hex_digit)
.ok_or(OAuthError::CallbackRejected)?;
(high << 4) | low
}
byte => byte,
};
if byte.is_ascii_control() {
return Err(OAuthError::CallbackRejected);
}
decoded.push(byte);
}
String::from_utf8(decoded).map_err(|_| OAuthError::CallbackRejected)
}
fn hex(bytes: &[u8]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut encoded = String::with_capacity(bytes.len() * 2);
for byte in bytes {
encoded.push(char::from(HEX[usize::from(byte >> 4)]));
encoded.push(char::from(HEX[usize::from(byte & 15)]));
}
encoded
}
fn hex_digit(byte: u8) -> Option<u8> {
match byte {
b'0'..=b'9' => Some(byte - b'0'),
b'a'..=b'f' => Some(byte - b'a' + 10),
b'A'..=b'F' => Some(byte - b'A' + 10),
_ => None,
}
}
fn decode_state(state: &str) -> Option<[u8; 32]> {
if state.len() != 64
|| !state
.bytes()
.all(|b| b.is_ascii_digit() || matches!(b, b'a'..=b'f'))
{
return None;
}
let mut decoded = [0; 32];
// `as_chunks` yields `[u8; 2]`, so the pair indexing below is checked at
// compile time. The length test above makes the remainder provably empty.
for (index, pair) in state.as_bytes().as_chunks::<2>().0.iter().enumerate() {
decoded[index] = (hex_digit(pair[0])? << 4) | hex_digit(pair[1])?;
}
Some(decoded)
}
fn pkce_challenge(verifier: &str) -> Result<String, OAuthError> {
if !(43..=128).contains(&verifier.len())
|| !verifier
.bytes()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-' | b'.' | b'_' | b'~'))
{
return Err(OAuthError::InvalidConfiguration);
}
let digest =
sha256_bounded(verifier.as_bytes(), 128).map_err(|_| OAuthError::InvalidConfiguration)?;
// The only Base64 input here is the fixed-width SHA-256 digest, not an
// extensible codec. Emit the RFC 7636 URL-safe alphabet without padding.
const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
let mut output = String::with_capacity(43);
let mut accumulator = 0_u32;
let mut bits = 0;
for byte in digest.as_bytes() {
accumulator = (accumulator << 8) | u32::from(*byte);
bits += 8;
while bits >= 6 {
bits -= 6;
output.push(char::from(ALPHABET[((accumulator >> bits) & 63) as usize]));
}
accumulator &= (1 << bits) - 1;
}
if bits != 0 {
output.push(char::from(
ALPHABET[((accumulator << (6 - bits)) & 63) as usize],
));
}
Ok(output)
}
fn validate_token_headers(headers: &[(String, String)]) -> Result<(), OAuthError> {
let mut content_type = None;
let mut encoding = false;
for (name, value) in headers {
if name.eq_ignore_ascii_case("content-type") {
if content_type.replace(value.as_str()).is_some() {
return Err(OAuthError::InvalidTokenResponse);
}
}
if name.eq_ignore_ascii_case("content-encoding") {
if encoding || !value.trim().eq_ignore_ascii_case("identity") {
return Err(OAuthError::InvalidTokenResponse);
}
encoding = true;
}
}
let value = content_type.ok_or(OAuthError::InvalidTokenResponse)?;
let mut parts = value.split(';');
if !parts
.next()
.is_some_and(|mime| mime.trim().eq_ignore_ascii_case("application/json"))
{
return Err(OAuthError::InvalidTokenResponse);
}
if let Some(parameter) = parts.next() {
let (name, value) = parameter
.trim()
.split_once('=')
.ok_or(OAuthError::InvalidTokenResponse)?;
if !name.trim().eq_ignore_ascii_case("charset")
|| !value.trim().trim_matches('"').eq_ignore_ascii_case("utf-8")
|| parts.next().is_some()
{
return Err(OAuthError::InvalidTokenResponse);
}
}
Ok(())
}
fn present<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
where
D: Deserializer<'de>,
T: Deserialize<'de>,
{
// `default` is absence; a present JSON null must not become absence.
T::deserialize(deserializer).map(Some)
}
#[derive(Deserialize)]
struct TokenResponse {
access_token: String,
token_type: String,
#[serde(default, deserialize_with = "present")]
expires_in: Option<u64>,
#[serde(default, deserialize_with = "present")]
scope: Option<String>,
#[serde(default, deserialize_with = "present")]
refresh_token: Option<String>,
#[serde(default, deserialize_with = "present")]
resource: Option<String>,
#[serde(default, deserialize_with = "present")]
error: Option<String>,
}
fn admit_token_response(
configuration: &OAuthClientConfiguration,
scope_ceiling: &[String],
bytes: &[u8],
started: Instant,
) -> Result<OAuthCredentials, OAuthError> {
if bytes.len() > MAX_TOKEN_RESPONSE_BYTES {
return Err(OAuthError::InvalidTokenResponse);
}
let response: TokenResponse =
serde_json::from_slice(bytes).map_err(|_| OAuthError::InvalidTokenResponse)?;
if !response.token_type.eq_ignore_ascii_case("Bearer")
|| !AccessToken::is_valid_token68(&response.access_token)
|| response
.refresh_token
.as_ref()
.is_some_and(|token| !valid_opaque(token, MAX_CODE_BYTES))
|| response
.resource
.as_ref()
.is_some_and(|resource| resource != configuration.resource.as_str())
|| response.error.is_some()
{
return Err(OAuthError::InvalidTokenResponse);
}
let scopes = response.scope.map_or_else(
|| scope_ceiling.to_vec(),
|scope| scope.split(' ').map(str::to_owned).collect(),
);
validate_scopes(&scopes)?;
if scopes.iter().any(|scope| !scope_ceiling.contains(scope)) {
return Err(OAuthError::ScopeExpansion);
}
let lifetime = response
.expires_in
.map_or(configuration.max_access_token_lifetime, |seconds| {
Duration::from_secs(seconds).min(configuration.max_access_token_lifetime)
});
let expires_at = started
.checked_add(lifetime)
.ok_or(OAuthError::InvalidTokenResponse)?;
if Instant::now() >= expires_at {
return Err(OAuthError::ExpiredCredential);
}
let access = BoundBearerCredential::bind_with_expiry(
configuration.resource.clone(),
response.access_token,
expires_at,
)
.map_err(|_| OAuthError::InvalidTokenResponse)?;
Ok(OAuthCredentials {
configuration: configuration.clone(),
access,
refresh_token: response.refresh_token,
scopes,
expires_at,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn url(value: &str) -> CanonicalHttpUrl {
CanonicalHttpUrl::parse(value).unwrap()
}
pub(super) fn config() -> OAuthClientConfiguration {
OAuthClientConfiguration::from_trusted_endpoints(
"https://issuer.example",
url("https://issuer.example/authorize"),
url("https://issuer.example/token"),
url("https://mcp.example/mcp"),
"native-client",
vec!["tools:read".to_owned(), "tools:write".to_owned()],
)
.unwrap()
}
fn wire(attempt: &AuthorizationAttempt, issuer: &str, extra: &str) -> Vec<u8> {
let query = encode_form(&[
("state", &attempt.state().unwrap()),
("iss", issuer),
("code", "code+/%"),
])
.unwrap();
format!("GET {CALLBACK_PATH}?{query}{extra} HTTP/1.1\r\nHost: 127.0.0.1:43210\r\n\r\n")
.into_bytes()
}
#[test]
fn pkce_matches_rfc_7636_appendix_b_without_plain_fallback() {
assert_eq!(
pkce_challenge("dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk").unwrap(),
"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
);
assert!(pkce_challenge("too-short").is_err());
let first = AuthorizationAttempt::new().unwrap();
let second = AuthorizationAttempt::new().unwrap();
assert_eq!(first.verifier().len(), 64);
assert_ne!(first.verifier(), second.verifier());
assert!(first.accepts_state(&first.state().unwrap()));
assert!(!second.accepts_state(&first.state().unwrap()));
let target = first
.authorization_url(&config(), "http://127.0.0.1:43210/oauth/callback")
.unwrap();
assert!(target.as_str().contains("code_challenge_method=S256"));
assert!(!target.as_str().contains(&first.verifier()));
let fields = decode_form(target.as_str().split_once('?').unwrap().1).unwrap();
assert_eq!(fields["resource"], "https://mcp.example/mcp");
assert_eq!(
fields["redirect_uri"],
"http://127.0.0.1:43210/oauth/callback"
);
}
#[test]
fn callback_requires_matching_state_issuer_host_and_one_code() {
let attempt = AuthorizationAttempt::new().unwrap();
let address = "127.0.0.1:43210".parse().unwrap();
let request = wire(&attempt, &config().issuer, "");
assert_eq!(
admit_callback(&request, address, &attempt, &config().issuer).unwrap(),
"code+/%"
);
let other = AuthorizationAttempt::new().unwrap();
assert_eq!(
admit_callback(&request, address, &other, &config().issuer),
Err(OAuthError::CallbackRejected)
);
assert_eq!(
admit_callback(&request, address, &attempt, "https://different.example"),
Err(OAuthError::IssuerMismatch)
);
for suffix in [
"&code=second",
"&co%64e=second",
"&error=access_denied",
"&state=other",
"&unknown=%ff",
] {
assert!(
admit_callback(
&wire(&attempt, &config().issuer, suffix),
address,
&attempt,
&config().issuer
)
.is_err()
);
}
for (from, to) in [
("Host: 127.0.0.1:43210", "Host: evil.example"),
("GET /oauth/callback?", "GET /other?"),
("\r\n\r\n", "\r\nContent-Length: 1\r\n\r\nx"),
("\r\n\r\n", "\r\nHost: 127.0.0.1:43210\r\n\r\n"),
] {
let changed = String::from_utf8(request.clone())
.unwrap()
.replace(from, to);
assert!(
admit_callback(changed.as_bytes(), address, &attempt, &config().issuer).is_err()
);
}
let response = callback_response(true);
assert!(!response.contains("code+/%"));
assert!(!response.contains(&attempt.state().unwrap()));
assert!(response.contains("Cache-Control: no-store"));
}
#[test]
fn token_admission_binds_resource_scopes_and_expiry() {
let config = config();
let now = Instant::now();
let grant = admit_token_response(&config, &config.scopes,
br#"{"access_token":"access-secret","token_type":"Bearer","expires_in":60,"refresh_token":"refresh-secret","scope":"tools:read"}"#, now).unwrap();
assert_eq!(grant.scopes(), &["tools:read".to_owned()]);
assert!(grant.has_refresh_token());
assert_eq!(grant.expires_at(), now + Duration::from_secs(60));
assert!(
grant
.bearer_credential()
.authorization_for_target(&config.resource)
.is_some()
);
assert!(
grant
.bearer_credential()
.authorization_for_target(&url("https://issuer.example/token"))
.is_none()
);
for invalid in [
r#"{"access_token":"access-secret","token_type":"Basic"}"#,
r#"{"access_token":"access-secret","token_type":"Bearer","expires_in":0}"#,
r#"{"access_token":"access-secret","token_type":"Bearer","expires_in":null}"#,
r#"{"access_token":"access-secret","token_type":"Bearer","scope":"admin"}"#,
r#"{"access_token":"access-secret","token_type":"Bearer","resource":"https://wrong.example/mcp"}"#,
r#"{"access_token":"first","access_token":"second","token_type":"Bearer"}"#,
r#"{"access_token":"access-secret","token_type":"Bearer","refresh_token":null}"#,
] {
let error = admit_token_response(&config, &config.scopes, invalid.as_bytes(), now)
.err()
.unwrap();
assert!(!format!("{error:?} {error}").contains("access-secret"));
}
}
#[test]
fn native_configuration_rejects_cleartext_endpoints_and_scope_ambiguity() {
let template = config();
for endpoint in [
"http://127.0.0.1/token",
"https://issuer.example/token?key=value",
] {
assert!(
OAuthClientConfiguration::from_trusted_endpoints(
template.issuer.clone(),
template.authorization_endpoint.clone(),
url(endpoint),
template.resource.clone(),
template.client_id.clone(),
template.scopes.clone(),
)
.is_err()
);
}
for scopes in [
vec!["duplicate".to_owned(); 2],
vec!["two scopes".to_owned()],
vec!["".to_owned()],
] {
assert!(validate_scopes(&scopes).is_err());
}
assert!(
template
.clone()
.with_authorization_timeout(Duration::ZERO)
.is_err()
);
assert!(
template
.with_max_access_token_lifetime(Duration::from_secs(86_401))
.is_err()
);
}
#[test]
fn token_media_admission_does_not_follow_or_decode_other_representations() {
let valid = vec![(
"Content-Type".to_owned(),
"application/json; charset=utf-8".to_owned(),
)];
assert!(validate_token_headers(&valid).is_ok());
let mut duplicate = valid.clone();
duplicate.push(("content-type".to_owned(), "application/json".to_owned()));
assert!(validate_token_headers(&duplicate).is_err());
let mut coded = valid;
coded.push(("Content-Encoding".to_owned(), "gzip".to_owned()));
assert!(validate_token_headers(&coded).is_err());
assert!(validate_token_headers(&[]).is_err());
assert!(decode_form("state=one&st%61te=two").is_err());
assert!(decode_form("code=%00").is_err());
assert!(decode_form("code=%").is_err());
}
#[test]
fn live_loopback_accepts_only_the_matching_attempt_and_releases_the_listener() {
asupersync::runtime::RuntimeBuilder::current_thread()
.build()
.unwrap()
.block_on(async {
let cx = Cx::current().expect("caller-owned runtime context");
let deadline = operation_deadline(&cx, Duration::from_secs(10)).unwrap();
let listener = within(&cx, deadline, bind_loopback()).await.unwrap();
let address = listener.local_addr().unwrap();
assert!(address.ip().is_loopback());
let attempt = AuthorizationAttempt::new().unwrap();
let impostor = AuthorizationAttempt::new().unwrap();
let config = config();
// Both peers use real sockets and the exact production parser.
// Only the state changes; the forged attempt cannot supply a code.
let mut peers = Vec::new();
for current in [&impostor, &attempt] {
let query = encode_form(&[
("state", ¤t.state().unwrap()),
("iss", &config.issuer),
("code", "live-code"),
])
.unwrap();
let request =
format!("GET {CALLBACK_PATH}?{query} HTTP/1.1\r\nHost: {address}\r\n\r\n");
let mut peer = within(&cx, deadline, async {
TcpStream::connect(address)
.await
.map_err(|_| OAuthError::CallbackRejected)
})
.await
.unwrap();
within(&cx, deadline, async {
peer.write_all(request.as_bytes())
.await
.map_err(|_| OAuthError::CallbackRejected)
})
.await
.unwrap();
peers.push(peer);
}
let code = wait_for_code(&cx, deadline, &listener, address, &attempt, &config)
.await
.unwrap();
assert_eq!(code, "live-code");
for (index, mut peer) in peers.into_iter().enumerate() {
let mut reply = Vec::new();
within(&cx, deadline, async {
peer.read_to_end(&mut reply)
.await
.map_err(|_| OAuthError::CallbackRejected)
})
.await
.unwrap();
let reply = String::from_utf8(reply).unwrap();
assert!(reply.starts_with(if index == 0 {
"HTTP/1.1 400"
} else {
"HTTP/1.1 200"
}));
assert!(!reply.contains("live-code"));
}
// Observe OUR listener directly rather than the port namespace.
// See `assert_listener_drained` for why the port probe was wrong
// here specifically.
assert_listener_drained(&listener);
drop(listener);
});
}
#[test]
fn callback_deadline_releases_idle_work_without_needing_peer_traffic() {
asupersync::runtime::RuntimeBuilder::current_thread()
.build()
.unwrap()
.block_on(async {
let cx = Cx::current().expect("caller-owned runtime context");
let setup_deadline = operation_deadline(&cx, Duration::from_secs(10)).unwrap();
let listener = within(&cx, setup_deadline, bind_loopback()).await.unwrap();
let address = listener.local_addr().unwrap();
let attempt = AuthorizationAttempt::new().unwrap();
let deadline = operation_deadline(&cx, Duration::from_millis(20)).unwrap();
let outcome =
wait_for_code(&cx, deadline, &listener, address, &attempt, &config()).await;
assert_eq!(outcome, Err(OAuthError::TimedOut));
// No port probe here. This test binds and drops its own listener, so
// its release is guaranteed by ownership, not by anything
// `wait_for_code` does - see `assert_listener_drained`. The subject of
// this test is the deadline expiring without peer traffic, which the
// `TimedOut` assertion above proves completely. A port probe could
// only ever add a false positive.
drop(listener);
});
}
/// Observes THIS listener directly: no further connection is pending on it.
///
/// # Why the port probe was the wrong instrument in the tests that use this
///
/// A bind probe asks a question about the **port namespace**, which is shared
/// with every other test in this binary. That is the right question when the
/// listener is owned by something we cannot inspect - the `authorize` future
/// in `dropping_public_login_future_closes_its_bound_callback_listener` - and
/// it stays a bind probe there.
///
/// It is the wrong question when the TEST itself binds and drops the
/// listener, as the two callers of this helper do. There, release is
/// guaranteed by ownership: `wait_for_code` takes `&TcpListener`, so the
/// borrow checker forbids it retaining one, and `asupersync`'s `TcpListener`
/// has no `Drop` of its own and closes its `std::net::TcpListener` descriptor
/// synchronously. A port probe in that position cannot fail for a real
/// reason - only when a concurrent test wins the freed ephemeral port, which
/// is what happened in wave-12b.
///
/// A probe that can only produce false positives is worse than no probe: it
/// fails on schedule, gets dismissed as flake, and trains everyone to ignore
/// the one time it means something.
///
/// So this asks a question about **our own listener** instead, which no other
/// test can influence: after `wait_for_code` has taken what it needed, is
/// anything still queued on it? That is race-free, and it is an OAuth
/// property rather than a restatement of Rust ownership.
///
/// The probe mints its own context. `poll_accept` returns
/// `Ready(Err(Interrupted))` whenever the ambient `Cx` is cancelled, before
/// it looks at the accept queue, so a probe inheriting a cancelled caller
/// would report a pending connection that does not exist. `Cx::clone` would
/// not do - it is an alias sharing the cancelled domain.
fn assert_listener_drained(listener: &TcpListener) {
let _frame = Cx::set_current(Some(Cx::for_request()));
let mut task = std::task::Context::from_waker(std::task::Waker::noop());
match listener.poll_accept(&mut task) {
std::task::Poll::Pending => {}
std::task::Poll::Ready(Ok((_stream, peer))) => panic!(
"a connection from {peer} is still queued after the callback was resolved; \
wait_for_code must consume exactly the connections it admits"
),
std::task::Poll::Ready(Err(error)) => panic!(
"the drained-listener probe is INCONCLUSIVE ({error}); it proves neither that \
the listener is quiet nor that a connection is pending"
),
}
}
/// Asserts the loopback callback listener at `address` is CLOSED, not merely
/// that its port is unreachable.
///
/// These are different claims. A `TcpStream::connect` that fails proves only
/// that nobody answered; a connect that SUCCEEDS proves only that *somebody*
/// is listening, never that *our* listener is. This is a lib unit test
/// sharing one binary with ~800 others, ~100 of which bind ephemeral loopback
/// ports in parallel, so a port freed a microsecond earlier can be handed
/// straight to a concurrent test and satisfy a connect against a listener
/// that has nothing to do with OAuth.
///
/// Binding the same address inverts the evidence: a successful bind proves
/// NOBODY is listening, which is direct proof the descriptor was closed, and
/// it reclaims the port so no concurrent test can occupy it mid-observation.
/// A genuine leak still fails here and fails harder - deterministic
/// `AddrInUse` instead of a probabilistic connect.
///
/// DO NOT add `SO_REUSEADDR` or `SO_REUSEPORT` to this bind. `SO_REUSEADDR`
/// only permits rebinding a port left in `TIME_WAIT` by an already-closed
/// socket; it does NOT permit two live listeners on one port, and that
/// refusal is the entire mechanism of this probe. `SO_REUSEPORT` does permit
/// exactly that, so setting it would let a leaked listener coexist with the
/// probe bind and silently turn this proof into a no-op.
///
/// No sleep and no timing tolerance: the close is synchronous in `Drop`, so
/// there is nothing to wait for and a tolerance would hide the very leak
/// being tested. The release case does run the whole experiment more than
/// once, but that is independent sampling on a fresh port each time, not a
/// re-observation of one port after a delay - see `RELEASE_TRIALS`.
/// Positive control for [`probe_listener_released`], run while the login
/// future is still alive, in every trial.
///
/// Without it the release probe is vacuous in one direction: a bind that
/// succeeds after the drop proves the port is free, but not that this
/// listener ever held it. A callback that was never bound, or an address
/// parsed out of the wrong field, would make the release probe pass while
/// proving nothing. This establishes that the probe can detect *this*
/// listener on *this* port, so the later success is a state change rather
/// than a constant.
///
/// This direction cannot be stolen by a concurrent test: we hold the port.
fn assert_listener_bound(address: SocketAddr) {
match std::net::TcpListener::bind(address) {
Err(error) if error.kind() == std::io::ErrorKind::AddrInUse => {}
Ok(listener) => {
drop(listener);
panic!(
"{address} was bindable while the login future was still alive, so the \
release probe cannot prove anything: either the callback listener was \
never bound or this is not the address it bound"
);
}
Err(error) => panic!(
"the bind control for {address} is INCONCLUSIVE ({error}); it establishes \
neither that the callback listener holds the port nor that it does not"
),
}
}
/// Asserts the callback listener at `address` is CLOSED.
///
/// SYNCHRONOUS ON PURPOSE, and this is the whole robustness argument.
///
/// The previous form awaited `asupersync::net::TcpListener::bind`. That
/// `.await` is a scheduling point: between `drop(login)` closing the
/// descriptor and the bind syscall reaching the kernel, the runtime polls
/// other tasks and the harness's other threads get a full async round trip
/// in which to claim the just-freed ephemeral port. This binary runs ~655
/// tests, ~100 of which bind ephemeral loopback ports, so that window was
/// routinely lost: the probe reported `AddrInUse` and that failure was
/// indistinguishable from the leak it exists to catch. A test that only
/// passes in isolation lies in every wave.
///
/// `std::net::TcpListener::bind` removes the await entirely - the drop and
/// the bind become straight-line code in one task with no yield between
/// them. This does not make theft impossible, since other OS threads still
/// run, but it removes the structural cause instead of tolerating the
/// symptom, and it costs no sleep and no timing tolerance. What remains of
/// the theft window is closed by independent repetition, not by waiting -
/// see `RELEASE_TRIALS`, and note that repeating the whole experiment on a
/// new port is a different thing from re-observing this one. The
/// close is synchronous: the listener is a local of the dropped `authorize`
/// future, asupersync's `TcpListener` has no `Drop` of its own, and its
/// `std::net::TcpListener` closes the descriptor in `Drop`. There is
/// nothing to wait for.
///
/// DO NOT reintroduce an `.await` between the drop and this bind, and DO
/// NOT add `SO_REUSEADDR` or `SO_REUSEPORT`. `SO_REUSEADDR` only permits
/// rebinding a port left in `TIME_WAIT` by an already-closed socket and
/// that refusal is the entire mechanism; `SO_REUSEPORT` permits two live
/// listeners on one port, which would let a leaked listener coexist with
/// this bind and silently turn the proof into a no-op.
///
/// The ambiguous arm still fails closed. It is narrower than it was, not
/// gone, so it names both hypotheses rather than asserting either.
/// What one release trial observed. `StillBound` is deliberately NOT a
/// panic: a single trial cannot distinguish a leak from a stolen port, and
/// the caller resolves that by repeating the whole experiment.
#[derive(Debug, PartialEq, Eq)]
enum ReleaseTrial {
/// The bind succeeded, so nothing holds the port. The descriptor was
/// closed. One such observation is conclusive on its own.
Released,
/// The port was still held. Either the listener leaked or a concurrent
/// ephemeral bind won it inside the straight-line window.
StillBound,
}
/// Number of independent release trials before a leak is declared.
///
/// Sized against the two hypotheses, not against a clock. A leaked listener
/// fails EVERY trial - the close is unconditional, so a leak is not a
/// probabilistic event. A thief must win a fresh, kernel-assigned ephemeral
/// port on each trial; the probability of that happening five times running
/// is the per-trial probability raised to the fifth power. Five is where
/// theft stops being a plausible explanation for a total failure while the
/// cost stays at a few binds and no sleep at all.
const RELEASE_TRIALS: usize = 5;
/// Runs ONE release trial and reports what it saw.
///
/// Everything the old assertion did, it still does: `std::net`'s blocking
/// bind with no `.await` between the drop and the syscall, no sleep, no
/// timing tolerance, no `SO_REUSEADDR`, no `SO_REUSEPORT`. The observation
/// is byte-for-byte the same question. What changed is only who answers it:
/// this reports `StillBound` instead of panicking, because that single
/// observation is genuinely ambiguous and the ambiguity is resolvable by
/// repetition rather than by tolerance.
///
/// The distinction matters and is not a retry in the forbidden sense. A
/// retry would re-observe THE SAME port after a delay, which is exactly the
/// tolerance that would hide a slow close. This instead re-runs the entire
/// experiment from a fresh `authorize` on a fresh ephemeral port, so each
/// trial is an independent sample of the same property.
///
/// The unexpected-errno arm still fails immediately and is never retried:
/// an unrecognised error is not a race, and repeating it would only convert
/// a clear diagnostic into a confusing one.
fn probe_listener_released(address: SocketAddr) -> ReleaseTrial {
match std::net::TcpListener::bind(address) {
Ok(listener) => {
drop(listener);
ReleaseTrial::Released
}
Err(error) if error.kind() == std::io::ErrorKind::AddrInUse => ReleaseTrial::StillBound,
Err(error) => panic!(
"the bind probe for {address} is INCONCLUSIVE ({error}); it proves neither \
closure nor a leak and must not be read as either"
),
}
}
pub(super) fn renewable_grant(config: &OAuthClientConfiguration) -> OAuthCredentials {
admit_token_response(config, &config.scopes,
br#"{"access_token":"access-one","token_type":"Bearer","expires_in":3600,"refresh_token":"refresh-one"}"#,
Instant::now(),
).unwrap()
}
#[test]
fn refresh_rotation_replaces_the_pair_and_preserves_a_narrowed_scope_ceiling() {
let config = config();
let client = OAuthClient::new(config.clone());
let mut grant = renewable_grant(&config);
let (body, previous) = client.prepare_refresh(&mut grant).unwrap();
let fields = decode_form(&body).unwrap();
assert_eq!(fields["grant_type"], "refresh_token");
assert_eq!(fields["refresh_token"], "refresh-one");
assert_eq!(fields["resource"], "https://mcp.example/mcp");
assert!(!fields.contains_key("code_verifier"));
assert!(!grant.has_refresh_token());
let replacement = client.admit_refresh(&grant, previous,
br#"{"access_token":"access-two","token_type":"Bearer","expires_in":120,"refresh_token":"refresh-two","scope":"tools:read"}"#,
Instant::now(),
).unwrap();
grant = replacement;
assert_eq!(
grant.access.authorization_for_target(&config.resource),
Some("Bearer access-two".to_owned())
);
let (body, previous) = client.prepare_refresh(&mut grant).unwrap();
let fields = decode_form(&body).unwrap();
assert_eq!(fields["refresh_token"], "refresh-two");
assert_eq!(fields["scope"], "tools:read");
let invalid = client.admit_refresh(
&grant,
previous,
br#"{"access_token":"access-three","token_type":"Bearer","scope":"tools:write"}"#,
Instant::now(),
);
assert_eq!(invalid.err(), Some(OAuthError::ScopeExpansion));
assert!(!grant.has_refresh_token());
assert_eq!(
grant.access.authorization_for_target(&config.resource),
Some("Bearer access-two".to_owned())
);
}
#[test]
fn refresh_without_rotation_retains_the_previous_token_only_after_valid_admission() {
let config = config();
let client = OAuthClient::new(config.clone());
let mut grant = renewable_grant(&config);
let (_, previous) = client.prepare_refresh(&mut grant).unwrap();
let replacement = client
.admit_refresh(
&grant,
previous,
br#"{"access_token":"access-two","token_type":"Bearer","expires_in":120}"#,
Instant::now(),
)
.unwrap();
assert_eq!(replacement.refresh_token.as_deref(), Some("refresh-one"));
assert_eq!(replacement.scopes, config.scopes);
}
#[test]
fn refresh_preflight_refuses_cross_binding_without_consuming_any_credential() {
let config = config();
let mut grant = renewable_grant(&config);
let access_before = grant.access.authorization_for_target(&config.resource);
let expiry_before = grant.expires_at;
for dimension in 0..5 {
let mut other = config.clone();
match dimension {
0 => other.issuer = "https://other.example".to_owned(),
1 => other.token_endpoint = url("https://other.example/token"),
2 => other.resource = url("https://other.example/mcp"),
3 => other.client_id = "another-client".to_owned(),
_ => other.max_access_token_lifetime = Duration::from_secs(1),
}
assert_eq!(
OAuthClient::new(other).prepare_refresh(&mut grant).err(),
Some(OAuthError::CredentialBindingMismatch)
);
assert_eq!(grant.refresh_token.as_deref(), Some("refresh-one"));
assert_eq!(
grant.access.authorization_for_target(&config.resource),
access_before
);
assert_eq!(grant.expires_at, expiry_before);
}
}
#[test]
fn abandoned_refresh_custody_cannot_replay_the_previous_refresh_token() {
let config = config();
let client = OAuthClient::new(config.clone());
let mut grant = renewable_grant(&config);
let access_before = grant.access.authorization_for_target(&config.resource);
let expiry_before = grant.expires_at;
// The transport owns the body and previous token after this point.
// Losing that future has the same ownership outcome as a lost reply.
drop(client.prepare_refresh(&mut grant).unwrap());
assert_eq!(
client.prepare_refresh(&mut grant).err(),
Some(OAuthError::RefreshUnavailable)
);
assert_eq!(
grant.access.authorization_for_target(&config.resource),
access_before
);
assert_eq!(grant.expires_at, expiry_before);
assert!(!grant.has_refresh_token());
}
// TEST ONLY credentials. The CA and localhost leaf are valid 2020-2049;
// none of these keys are used outside this in-process TLS fixture.
const TEST_ROOT: &[u8] = b"-----BEGIN CERTIFICATE-----\nMIIBgzCCASmgAwIBAgICA+kwCgYIKoZIzj0EAwIwJzElMCMGA1UEAwwcRmFzdE1D\nUCBPQXV0aCBURVNUIE9OTFkgUm9vdDAeFw0yMDAxMDEwMDAwMDBaFw00OTEyMzEw\nMDAwMDBaMCcxJTAjBgNVBAMMHEZhc3RNQ1AgT0F1dGggVEVTVCBPTkxZIFJvb3Qw\nWTATBgcqhkjOPQIBBggqhkjOPQMBBwNCAAS5t2O8JZ0hNjgI38E9Ov6i6mKoDRGo\nApMsykFkvgb6Zm9/5gCZ90eIKw7aWgK6iNs7lbtVY9mysZBIqm6pKQO2o0UwQzAS\nBgNVHRMBAf8ECDAGAQH/AgEAMA4GA1UdDwEB/wQEAwIBhjAdBgNVHQ4EFgQU6QNI\nrmvMiLoV3jIoCyohXARwI8gwCgYIKoZIzj0EAwIDSAAwRQIgCKOrW3vhzUJ2EyuY\nvQUTdqGFhy0zEHj4ITFLvXPz1X8CIQCLKD4EKCvS/zkBSu/6uee1WV9d97UpK3yW\nX/aCEJ5+hA==\n-----END CERTIFICATE-----\n";
const TEST_LEAF: &[u8] = b"-----BEGIN CERTIFICATE-----\nMIIBjjCCATSgAwIBAgICA+owCgYIKoZIzj0EAwIwJzElMCMGA1UEAwwcRmFzdE1D\nUCBPQXV0aCBURVNUIE9OTFkgUm9vdDAeFw0yMDAxMDEwMDAwMDBaFw00OTEyMzEw\nMDAwMDBaMBQxEjAQBgNVBAMMCWxvY2FsaG9zdDBZMBMGByqGSM49AgEGCCqGSM49\nAwEHA0IABPPKylLna9VpWAlpshHBhSsQHNOv3BaEGX4HSBhHiBVel0ce+qfHF15O\n0T63Zlp7TtxlMdEY+rPpgioSFDQVadijYzBhMAwGA1UdEwEB/wQCMAAwLAYDVR0R\nBCUwI4IJbG9jYWxob3N0hwR/AAABhxAAAAAAAAAAAAAAAAAAAAABMBMGA1UdJQQM\nMAoGCCsGAQUFBwMBMA4GA1UdDwEB/wQEAwIHgDAKBggqhkjOPQQDAgNIADBFAiEA\n6qrAr2qp/t6K62T9Et2mUU/zfd4kJb+ekyoAim1yTFcCICb6SdVY2fg15/SXf0vE\nIvYelqtTk8FQInCEcIxvfF3m\n-----END CERTIFICATE-----\n";
const TEST_KEY: &[u8] = b"-----BEGIN PRIVATE KEY-----\nMIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgcCe44IBKhbw+D/s7\nBjDHOOV0g+EoxFno7VJGKhJeer2hRANCAATzyspS52vVaVgJabIRwYUrEBzTr9wW\nhBl+B0gYR4gVXpdHHvqnxxdeTtE+t2Zae07cZTHRGPqz6YIqEhQ0FWnY\n-----END PRIVATE KEY-----\n";
pub(super) fn test_root() -> asupersync::tls::Certificate {
asupersync::tls::Certificate::from_pem(TEST_ROOT)
.unwrap()
.remove(0)
}
pub(super) fn test_acceptor() -> asupersync::tls::TlsAcceptor {
asupersync::tls::TlsAcceptorBuilder::new(
asupersync::tls::CertificateChain::from_pem(TEST_LEAF).unwrap(),
asupersync::tls::PrivateKey::from_pem(TEST_KEY).unwrap(),
)
.alpn_protocols(vec![b"http/1.1".to_vec()])
.build()
.unwrap()
}
pub(super) async fn pair<L: Future, R: Future>(left: L, right: R) -> (L::Output, R::Output) {
let mut left = std::pin::pin!(left);
let mut right = std::pin::pin!(right);
let mut left_output = None;
let mut right_output = None;
poll_fn(|task| {
if left_output.is_none() {
if let Poll::Ready(value) = left.as_mut().poll(task) {
left_output = Some(value);
}
}
if right_output.is_none() {
if let Poll::Ready(value) = right.as_mut().poll(task) {
right_output = Some(value);
}
}
if left_output.is_some() && right_output.is_some() {
Poll::Ready((left_output.take().unwrap(), right_output.take().unwrap()))
} else {
Poll::Pending
}
})
.await
}
pub(super) async fn read_token_request<IO: asupersync::io::AsyncRead + Unpin>(
stream: &mut IO,
) -> Result<(String, BTreeMap<String, String>), OAuthError> {
let mut wire = Vec::new();
let mut buffer = [0; 1024];
let head_end = loop {
let count = stream
.read(&mut buffer)
.await
.map_err(|_| OAuthError::TransportFailed)?;
if count == 0 || wire.len() + count > MAX_FORM_BYTES {
return Err(OAuthError::TransportFailed);
}
wire.extend_from_slice(&buffer[..count]);
if let Some(index) = wire.windows(4).position(|part| part == b"\r\n\r\n") {
break index + 4;
}
};
let head = std::str::from_utf8(&wire[..head_end])
.map_err(|_| OAuthError::TransportFailed)?
.to_owned();
let count = head
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.ok_or(OAuthError::TransportFailed)?;
if count > MAX_FORM_BYTES - head_end {
return Err(OAuthError::TransportFailed);
}
while wire.len() < head_end + count {
let n = stream
.read(&mut buffer)
.await
.map_err(|_| OAuthError::TransportFailed)?;
if n == 0 || wire.len() + n > MAX_FORM_BYTES {
return Err(OAuthError::TransportFailed);
}
wire.extend_from_slice(&buffer[..n]);
}
if wire.len() != head_end + count {
return Err(OAuthError::TransportFailed);
}
let form =
std::str::from_utf8(&wire[head_end..]).map_err(|_| OAuthError::TransportFailed)?;
Ok((head, decode_form(form)?))
}
async fn send_callback(
cx: &Cx,
authorization: CanonicalHttpUrl,
) -> Result<BTreeMap<String, String>, OAuthError> {
let fields = decode_form(authorization.query().ok_or(OAuthError::CallbackRejected)?)?;
let callback = CanonicalHttpUrl::parse(&fields["redirect_uri"])
.map_err(|_| OAuthError::CallbackRejected)?;
let address: SocketAddr = callback
.as_str()
.strip_prefix("http://")
.unwrap()
.split('/')
.next()
.unwrap()
.parse()
.unwrap();
let query = encode_form(&[
("code", "issued-code"),
("iss", "https://issuer.example"),
("state", &fields["state"]),
])?;
let request = format!("GET {CALLBACK_PATH}?{query} HTTP/1.1\r\nHost: {address}\r\n\r\n");
let deadline = operation_deadline(cx, Duration::from_secs(10))?;
within(cx, deadline, async {
let mut stream = TcpStream::connect(address)
.await
.map_err(|_| OAuthError::CallbackRejected)?;
stream
.write_all(request.as_bytes())
.await
.map_err(|_| OAuthError::CallbackRejected)?;
Ok(())
})
.await?;
Ok(fields)
}
#[test]
fn native_https_login_and_refresh_verify_pkce_and_rotate_the_live_grant() {
asupersync::runtime::RuntimeBuilder::current_thread().build().unwrap().block_on(async {
let cx = Cx::current().unwrap();
let deadline = operation_deadline(&cx, Duration::from_secs(20)).unwrap();
let listener = within(&cx, deadline, bind_loopback()).await.unwrap();
let address = listener.local_addr().unwrap();
let mut configuration = config().with_extra_root_certificate(test_root()).unwrap();
configuration.token_endpoint = url(&format!("https://{address}/token"));
let client = OAuthClient::new(configuration.clone());
let browser = std::sync::Mutex::new(None::<BTreeMap<String, String>>);
let acceptor = test_acceptor();
let server = within(&cx, deadline, async {
for index in 0..2 {
let (socket, _) = listener.accept().await.map_err(|_| OAuthError::TransportFailed)?;
let mut tls = acceptor.accept(socket).await.map_err(|_| OAuthError::TransportFailed)?;
let (head, form) = read_token_request(&mut tls).await?;
assert!(head.starts_with("POST /token HTTP/1.1\r\n"));
assert!(!head.to_ascii_lowercase().contains("authorization:"));
assert!(!head.to_ascii_lowercase().contains("cookie:"));
assert_eq!(form["client_id"], "native-client");
assert_eq!(form["resource"], "https://mcp.example/mcp");
let body = if index == 0 {
let observed = browser.lock().unwrap();
let observed = observed.as_ref().unwrap();
assert_eq!(form["grant_type"], "authorization_code");
assert_eq!(form["code"], "issued-code");
assert_eq!(pkce_challenge(&form["code_verifier"]).unwrap(), observed["code_challenge"]);
assert_eq!(form["redirect_uri"], observed["redirect_uri"]);
r#"{"access_token":"live-access-one","token_type":"Bearer","expires_in":600,"refresh_token":"live-refresh-one"}"#
} else {
assert_eq!(form["grant_type"], "refresh_token");
assert_eq!(form["refresh_token"], "live-refresh-one");
assert!(!form.contains_key("code_verifier"));
r#"{"access_token":"live-access-two","token_type":"Bearer","expires_in":300,"refresh_token":"live-refresh-two","scope":"tools:read"}"#
};
let response = format!("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nCache-Control: no-store\r\nConnection: close\r\n\r\n{body}", body.len());
tls.write_all(response.as_bytes()).await.map_err(|_| OAuthError::TransportFailed)?;
tls.shutdown().await.map_err(|_| OAuthError::TransportFailed)?;
}
Ok(())
});
let application = async {
let caller = &cx;
let observed_browser = &browser;
let mut credentials = client.authorize(&cx, |authorization| async move {
let fields = send_callback(caller, authorization).await?;
*observed_browser.lock().unwrap() = Some(fields);
Ok(())
}).await?;
assert_eq!(credentials.access.authorization_for_target(&configuration.resource), Some("Bearer live-access-one".to_owned()));
client.refresh(&cx, &mut credentials).await?;
assert_eq!(credentials.access.authorization_for_target(&configuration.resource), Some("Bearer live-access-two".to_owned()));
assert_eq!(credentials.refresh_token.as_deref(), Some("live-refresh-two"));
assert_eq!(credentials.scopes, ["tools:read".to_owned()]);
Ok::<(), OAuthError>(())
};
let (server, application) = Box::pin(pair(server, application)).await;
assert_eq!(server, Ok(()));
assert_eq!(application, Ok(()));
});
}
#[test]
fn native_https_rejects_untrusted_certificate_before_sending_the_code() {
asupersync::runtime::RuntimeBuilder::current_thread()
.build()
.unwrap()
.block_on(async {
let cx = Cx::current().unwrap();
let deadline = operation_deadline(&cx, Duration::from_secs(20)).unwrap();
let listener = within(&cx, deadline, bind_loopback()).await.unwrap();
let mut configuration = config();
configuration.token_endpoint =
url(&format!("https://{}/token", listener.local_addr().unwrap()));
// The only changed trust dimension: do not install the test root.
let client = OAuthClient::new(configuration);
let acceptor = test_acceptor();
let server = within(&cx, deadline, async {
let (socket, _) = listener
.accept()
.await
.map_err(|_| OAuthError::TransportFailed)?;
assert!(
acceptor.accept(socket).await.is_err(),
"untrusted TLS cannot become an HTTP request stream"
);
Ok(())
});
let caller = &cx;
let application = client.authorize(&cx, |authorization| async move {
send_callback(caller, authorization).await.map(|_| ())
});
let (server, application) = Box::pin(pair(server, application)).await;
assert_eq!(server, Ok(()));
assert_eq!(application.err(), Some(OAuthError::TransportFailed));
});
}
#[test]
fn lost_https_refresh_reply_and_redirect_never_replay_a_consumed_token() {
for redirect in [false, true] {
asupersync::runtime::RuntimeBuilder::current_thread().build().unwrap().block_on(async {
let cx = Cx::current().unwrap();
let deadline = operation_deadline(&cx, Duration::from_secs(20)).unwrap();
let listener = within(&cx, deadline, bind_loopback()).await.unwrap();
let mut configuration = config().with_extra_root_certificate(test_root()).unwrap();
configuration.token_endpoint = url(&format!("https://{}/token", listener.local_addr().unwrap()));
let client = OAuthClient::new(configuration.clone());
let mut credentials = renewable_grant(&configuration);
let before = credentials.access.authorization_for_target(&configuration.resource);
let expires_before = credentials.expires_at;
let acceptor = test_acceptor();
let server = within(&cx, deadline, async {
let (socket, _) = listener.accept().await.map_err(|_| OAuthError::TransportFailed)?;
let mut tls = acceptor.accept(socket).await.map_err(|_| OAuthError::TransportFailed)?;
let (_, form) = read_token_request(&mut tls).await?;
assert_eq!(form["refresh_token"], "refresh-one");
if redirect {
tls.write_all(b"HTTP/1.1 307 Temporary Redirect\r\nLocation: https://127.0.0.1:9/forbidden\r\nContent-Length: 0\r\nConnection: close\r\n\r\n").await.map_err(|_| OAuthError::TransportFailed)?;
tls.shutdown().await.map_err(|_| OAuthError::TransportFailed)?;
}
// Without a response, the client cannot know whether the
// server rotated the refresh token. The connection is lost.
Ok(())
});
let (server, result) = pair(server, client.refresh(&cx, &mut credentials)).await;
assert_eq!(server, Ok(()));
assert_eq!(result, Err(if redirect { OAuthError::TokenEndpointRejected } else { OAuthError::TransportFailed }));
assert_eq!(credentials.access.authorization_for_target(&configuration.resource), before);
assert_eq!(credentials.expires_at, expires_before);
assert!(!credentials.has_refresh_token());
assert_eq!(client.refresh(&cx, &mut credentials).await, Err(OAuthError::RefreshUnavailable));
let waker = std::task::Waker::noop();
let mut task = std::task::Context::from_waker(waker);
assert!(listener.poll_accept(&mut task).is_pending(), "no retry opened another connection");
});
}
}
#[test]
fn dropping_public_login_future_closes_its_bound_callback_listener() {
asupersync::runtime::RuntimeBuilder::current_thread()
.build()
.unwrap()
.block_on(async {
let cx = Cx::current().unwrap();
// Each iteration is a COMPLETE, independent experiment: its own
// login, its own listener, its own kernel-assigned ephemeral port.
// A leaked listener fails all of them; a port thief would have to
// win a different port every time.
let mut still_bound = Vec::new();
let mut released = false;
for _ in 0..RELEASE_TRIALS {
let deadline = operation_deadline(&cx, Duration::from_secs(10)).unwrap();
let client = OAuthClient::new(config());
let bound = std::cell::Cell::new(None::<SocketAddr>);
let observed_bound = &bound;
let mut login = Box::pin(client.authorize(&cx, |authorization| async move {
let fields = decode_form(authorization.query().unwrap())?;
let address = fields["redirect_uri"]
.strip_prefix("http://")
.unwrap()
.split('/')
.next()
.unwrap()
.parse()
.unwrap();
observed_bound.set(Some(address));
Ok(())
}));
let address = within(
&cx,
deadline,
poll_fn(|task| {
if let Poll::Ready(result) = login.as_mut().poll(task) {
return Poll::Ready(Err(result
.err()
.unwrap_or(OAuthError::CallbackRejected)));
}
match bound.get() {
Some(address) => Poll::Ready(Ok(address)),
None => Poll::Pending,
}
}),
)
.await
.unwrap();
// Positive control BEFORE the drop: while the login future is
// alive its callback listener holds the port, so this bind must be
// refused. It makes the release probe a state change rather than a
// constant, and it cannot be stolen because we hold the port.
// It runs in EVERY trial, so no trial is vacuous.
assert_listener_bound(address);
drop(login);
// No await between the drop and the probe. The probe is synchronous
// precisely so the runtime cannot schedule another task into the
// window where the freed port is unclaimed - see
// `probe_listener_released` for the full argument.
match probe_listener_released(address) {
ReleaseTrial::Released => {
released = true;
break;
}
ReleaseTrial::StillBound => still_bound.push(address),
}
}
assert!(
released,
"the callback listener was still bound after the login future was dropped in \
all {RELEASE_TRIALS} independent trials, on these separately assigned \
ephemeral ports: {still_bound:?}. Each trial issued its bind with no await \
between it and the drop, and each was preceded by a positive control proving \
the probe can see that listener. A concurrent test can steal one freed port; \
it cannot steal {RELEASE_TRIALS} different ones in a row. The listener leaked."
);
});
}
#[test]
fn private_ca_policy_is_validated_and_bound_to_refresh_credentials() {
assert!(
config()
.with_extra_root_certificate(asupersync::tls::Certificate::from_der(vec![0; 32]))
.is_err()
);
let configuration = config().with_extra_root_certificate(test_root()).unwrap();
assert!(
configuration
.clone()
.with_extra_root_certificate(test_root())
.is_err()
);
let mut credentials = renewable_grant(&configuration);
let before = credentials.refresh_token.clone();
assert_eq!(
OAuthClient::new(config())
.prepare_refresh(&mut credentials)
.err(),
Some(OAuthError::CredentialBindingMismatch)
);
assert_eq!(credentials.refresh_token, before);
}
}