types: HttpUrl newtypes, shared cache key/JSON helpers

Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
This commit is contained in:
Lewis
2026-08-16 20:05:16 +03:00
parent c723bc2164
commit 420ce1e201
20 changed files with 1011 additions and 167 deletions
Generated
+11 -1
View File
@@ -7838,7 +7838,10 @@ dependencies = [
"async-trait",
"bytes",
"futures",
"serde",
"serde_json",
"thiserror 2.0.18",
"tranquil-types",
]
[[package]]
@@ -7855,9 +7858,9 @@ dependencies = [
"thiserror 2.0.18",
"tokio",
"tracing",
"tranquil-infra",
"tranquil-types",
"unicode-segmentation",
"urlencoding",
"wiremock",
]
@@ -7880,6 +7883,7 @@ dependencies = [
"sqlx",
"tokio",
"tracing",
"tranquil-infra",
"tranquil-types",
"uuid",
]
@@ -7911,6 +7915,7 @@ dependencies = [
"tranquil-config",
"tranquil-crypto",
"tranquil-db-traits",
"tranquil-infra",
"tranquil-pds",
"tranquil-scopes",
"tranquil-types",
@@ -7991,6 +7996,7 @@ dependencies = [
"tranquil-config",
"tranquil-db",
"tranquil-db-traits",
"tranquil-infra",
"tranquil-lexicon",
"tranquil-oauth",
"tranquil-oauth-server",
@@ -8221,10 +8227,14 @@ dependencies = [
"cid",
"jacquard-common",
"rand 0.8.5",
"reqwest",
"serde",
"serde_json",
"sqlx",
"thiserror 2.0.18",
"tokio",
"tracing",
"url",
"uuid",
]
+1
View File
@@ -137,6 +137,7 @@ tower-layer = "0.3"
tracing = "0.1"
tracing-subscriber = "0.3"
urlencoding = "2.1"
url = "2.5"
uuid = { version = "1.19", features = ["v4", "v5", "v7", "fast-rng", "serde"] }
webauthn-rs = { version = "0.5", features = ["danger-allow-state-serialisation", "danger-user-presence-only-security-keys", "conditional-ui"] }
webauthn-rs-proto = "0.5"
+14 -9
View File
@@ -12,8 +12,8 @@ use tranquil_pds::api::{
};
use tranquil_pds::auth::{Active, Auth};
use tranquil_pds::delegation::{
DelegationActionType, SCOPE_PRESETS, ValidatedDelegationScope, verify_can_add_controllers,
verify_can_control_accounts,
DelegationActionType, IdentityResolutionError, SCOPE_PRESETS, ValidatedDelegationScope,
verify_can_add_controllers, verify_can_control_accounts,
};
use tranquil_pds::rate_limit::{AccountCreationLimit, RateLimited};
use tranquil_pds::state::AppState;
@@ -65,16 +65,16 @@ pub async fn add_controller(
) -> Result<Json<SuccessResponse>, ApiError> {
let resolved = tranquil_pds::delegation::resolve_identity(&state, &input.controller_did)
.await
.map_err(|_| ApiError::ControllerNotFound)?;
.map_err(|e| match e {
IdentityResolutionError::PdsEndpoint(_) => ApiError::InvalidDelegation(
"Controller PDS endpoint isn't a usable https URL".into(),
),
IdentityResolutionError::DidResolution(_) => ApiError::ControllerNotFound,
})?;
if !resolved.is_local
&& let Some(ref pds_url) = resolved.pds_url
{
if !pds_url.starts_with("https://") {
return Err(ApiError::InvalidDelegation(
"Controller PDS must use HTTPS".into(),
));
}
match state
.cross_pds_oauth
.check_remote_is_delegated(pds_url, &input.controller_did)
@@ -477,7 +477,12 @@ pub async fn resolve_controller(
let resolved = tranquil_pds::delegation::resolve_identity(&state, &did)
.await
.map_err(|_| ApiError::ControllerNotFound)?;
.map_err(|e| match e {
IdentityResolutionError::PdsEndpoint(_) => ApiError::InvalidDelegation(
"Controller PDS endpoint isn't a usable https URL".into(),
),
IdentityResolutionError::DidResolution(_) => ApiError::ControllerNotFound,
})?;
Ok(Json(resolved))
}
+1 -1
View File
@@ -9,7 +9,7 @@ valkey = ["dep:redis"]
[dependencies]
tranquil-config = { workspace = true }
tranquil-infra = { workspace = true }
tranquil-infra = { workspace = true, features = ["cache-keys"] }
tranquil-ripple = { workspace = true }
async-trait = { workspace = true }
+4 -3
View File
@@ -1,4 +1,6 @@
pub use tranquil_infra::{Cache, CacheError, DistributedRateLimiter};
pub use tranquil_infra::{
Cache, CacheError, DistributedRateLimiter, cache_keys, cached_json, read_json, write_json,
};
use async_trait::async_trait;
use std::sync::Arc;
@@ -173,11 +175,10 @@ pub async fn create_cache(
) -> Result<(Arc<dyn Cache>, Arc<dyn DistributedRateLimiter>), CacheInitError> {
let cache_cfg = tranquil_config::try_get().map(|c| &c.cache);
let backend = cache_cfg.map(|c| c.backend.as_str()).unwrap_or("ripple");
let valkey_url = cache_cfg.and_then(|c| c.valkey_url.as_deref());
#[cfg(feature = "valkey")]
if backend == "valkey" {
if let Some(url) = valkey_url {
if let Some(url) = cache_cfg.and_then(|c| c.valkey_url.as_deref()) {
match ValkeyCache::new(url).await {
Ok(cache) => {
tracing::info!("using valkey cache at {url}");
+8
View File
@@ -4,8 +4,16 @@ version.workspace = true
edition.workspace = true
license.workspace = true
[features]
testing = []
cache-keys = ["dep:tranquil-types"]
[dependencies]
tranquil-types = { workspace = true, optional = true }
async-trait = { workspace = true }
bytes = { workspace = true }
futures = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
thiserror = { workspace = true }
+103
View File
@@ -0,0 +1,103 @@
use tranquil_types::{
CidLink, ClientId, CrossPdsState, Did, EmailTokenPurpose, Handle, Jti, JwksUri, Nsid, PdsUrl,
SsoIssuer, SsoJwksUri,
};
pub fn session_key(did: &Did, jti: &Jti) -> String {
format!("auth:session:{}:{}", did, jti)
}
pub fn signing_key_key(did: &Did) -> String {
format!("auth:key:{}", did)
}
pub fn user_status_key(did: &Did) -> String {
format!("auth:status:{}", did)
}
pub fn handle_key(handle: &Handle) -> String {
format!("handle:{}", handle)
}
pub fn reauth_key(did: &Did) -> String {
format!("reauth:{}", did)
}
pub fn plc_doc_key(did: &Did) -> String {
format!("plc:doc:{}", did)
}
pub fn plc_data_key(did: &Did) -> String {
format!("plc:data:{}", did)
}
pub fn did_web_doc_key(did: &Did) -> String {
format!("did:web:doc:{}", did)
}
pub fn email_update_key(did: &Did) -> String {
format!("email_update:{}", did)
}
pub fn email_token_key(did: &Did, purpose: EmailTokenPurpose) -> String {
format!("email_token:{}:{}", purpose, did)
}
pub fn legacy_2fa_challenge_key(did: &Did) -> String {
format!("legacy_2fa:{}", did)
}
pub fn legacy_2fa_cooldown_key(did: &Did) -> String {
format!("legacy_2fa_cooldown:{}", did)
}
pub fn scope_ref_key(cid: &CidLink) -> String {
format!("scope_ref:{}", cid)
}
pub fn auto_verify_sent_key(did: &Did) -> String {
format!("auto_verify_sent:{}", did)
}
pub fn permission_set_key(nsid: &Nsid, aud: Option<&str>) -> String {
match aud {
Some(a) => format!("permset:{}:{}", nsid, a),
None => format!("permset:{}", nsid),
}
}
pub fn oauth_client_meta_key(client_id: &ClientId) -> String {
format!("oauth:client_meta:{}", client_id)
}
pub fn oauth_client_jwks_key(jwks_uri: &JwksUri) -> String {
format!("oauth:jwks:{}", jwks_uri.canonical())
}
pub fn oauth_client_jwks_cooldown_key(jwks_uri: &JwksUri) -> String {
format!("oauth:jwks_cooldown:{}", jwks_uri.canonical())
}
pub fn sso_jwks_key(jwks_uri: &SsoJwksUri) -> String {
format!("sso:jwks:{}", jwks_uri.canonical())
}
pub fn oidc_discovery_key(issuer: &SsoIssuer) -> String {
format!("oidc:discovery:{}", issuer.canonical())
}
pub fn cross_pds_state_key(state: &CrossPdsState) -> String {
format!("cross_pds_state:{}", state)
}
pub fn cross_pds_oauth_meta_key(pds_url: &PdsUrl) -> String {
format!("cross_pds_oauth_meta:v2:{}", pds_url.canonical())
}
pub fn lexicon_doc_key(nsid: &Nsid) -> String {
format!("lexicon:doc:{}", nsid)
}
pub fn lexicon_negative_key(nsid: &Nsid) -> String {
format!("lexicon:neg:{}", nsid)
}
+45
View File
@@ -1,6 +1,15 @@
#[cfg(feature = "cache-keys")]
pub mod cache_keys;
#[cfg(feature = "testing")]
mod memory_cache;
#[cfg(feature = "testing")]
pub use memory_cache::MemoryCache;
use async_trait::async_trait;
use bytes::Bytes;
use futures::Stream;
use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
@@ -57,6 +66,42 @@ pub trait Cache: Send + Sync {
}
}
pub async fn read_json<T: serde::de::DeserializeOwned>(cache: &dyn Cache, key: &str) -> Option<T> {
let json = cache.get(key).await?;
serde_json::from_str(&json).ok()
}
pub async fn write_json<T: serde::Serialize>(
cache: &dyn Cache,
key: &str,
value: &T,
ttl: Duration,
) {
if let Ok(json) = serde_json::to_string(value) {
let _ = cache.set(key, &json, ttl).await;
}
}
pub async fn cached_json<T, E, Fut>(
cache: &dyn Cache,
key: &str,
ttl: Duration,
fetch: impl FnOnce() -> Fut,
) -> Result<T, E>
where
T: serde::Serialize + serde::de::DeserializeOwned,
Fut: Future<Output = Result<T, E>>,
{
match read_json(cache, key).await {
Some(value) => Ok(value),
None => {
let value = fetch().await?;
write_json(cache, key, &value, ttl).await;
Ok(value)
}
}
}
#[async_trait]
pub trait DistributedRateLimiter: Send + Sync {
async fn check_rate_limit(&self, key: &str, limit: u32, window_ms: u64) -> bool;
+74
View File
@@ -0,0 +1,74 @@
use crate::{Cache, CacheError};
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Mutex;
use std::time::{Duration, Instant};
struct Entry {
value: Vec<u8>,
expires_at: Instant,
}
#[derive(Default)]
pub struct MemoryCache {
entries: Mutex<HashMap<String, Entry>>,
}
impl MemoryCache {
pub fn new() -> Self {
Self::default()
}
fn read(&self, key: &str) -> Option<Vec<u8>> {
let now = Instant::now();
let mut entries = self.entries.lock().unwrap_or_else(|e| e.into_inner());
match entries.get(key) {
Some(entry) if entry.expires_at > now => Some(entry.value.clone()),
Some(_) => {
entries.remove(key);
None
}
None => None,
}
}
fn write(&self, key: &str, value: Vec<u8>, ttl: Duration) {
let entry = Entry {
value,
expires_at: Instant::now() + ttl,
};
self.entries
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(key.to_string(), entry);
}
}
#[async_trait]
impl Cache for MemoryCache {
async fn get(&self, key: &str) -> Option<String> {
self.read(key).and_then(|v| String::from_utf8(v).ok())
}
async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> {
self.write(key, value.as_bytes().to_vec(), ttl);
Ok(())
}
async fn delete(&self, key: &str) -> Result<(), CacheError> {
self.entries
.lock()
.unwrap_or_else(|e| e.into_inner())
.remove(key);
Ok(())
}
async fn get_bytes(&self, key: &str) -> Option<Vec<u8>> {
self.read(key)
}
async fn set_bytes(&self, key: &str, value: &[u8], ttl: Duration) -> Result<(), CacheError> {
self.write(key, value.to_vec(), ttl);
Ok(())
}
}
@@ -13,7 +13,8 @@ use tranquil_pds::rate_limit::{LoginLimit, OAuthRateLimited, TotpVerifyLimit};
use tranquil_pds::state::AppState;
use tranquil_pds::types::PlainPassword;
use tranquil_pds::util::ClientIp;
use tranquil_types::did_doc::{extract_handle, extract_pds_endpoint};
use tranquil_types::did_doc::{PdsEndpointError, extract_handle, extract_pds_endpoint};
use tranquil_types::url_kind;
use tranquil_types::{Did, RequestId};
#[allow(clippy::result_large_err)]
@@ -231,11 +232,17 @@ pub async fn delegation_auth(
}
};
let pds_url = match extract_pds_endpoint(&did_doc) {
Some(url) => url,
None => {
let pds_url = match extract_pds_endpoint::<url_kind::Pds>(&did_doc) {
Ok(url) => url,
Err(PdsEndpointError::Missing) => {
return DelegationAuthResponse::err("Controller has no PDS endpoint");
}
Err(PdsEndpointError::Invalid(e)) => {
tracing::warn!(controller = %controller_did, error = %e, "Controller PDS endpoint rejected");
return DelegationAuthResponse::err(
"Controller PDS endpoint isn't a usable https URL",
);
}
};
let hostname = &tranquil_config::get().server.hostname;
@@ -447,7 +454,7 @@ pub async fn delegation_auth_token(
#[derive(Debug, Deserialize)]
pub struct CrossPdsCallbackParams {
pub code: tranquil_types::AuthorizationCode,
pub state: String,
pub state: tranquil_types::CrossPdsState,
pub iss: Option<String>,
}
@@ -474,7 +481,7 @@ pub async fn delegation_callback(
if let Some(ref expected_issuer) = auth_state.expected_issuer {
match &params.iss {
Some(iss) if iss != expected_issuer => {
Some(iss) if iss.as_str() != expected_issuer.as_str() => {
tracing::error!(
"Cross-PDS issuer mismatch: expected {}, got {}",
expected_issuer,
+1
View File
@@ -6,6 +6,7 @@ license.workspace = true
[dependencies]
tranquil-types = { workspace = true }
tranquil-infra = { workspace = true, features = ["cache-keys"] }
anyhow = { workspace = true }
sqlx = { workspace = true }
+5 -5
View File
@@ -1,7 +1,7 @@
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use tranquil_types::{ClientId, Did};
use tranquil_types::{AuthServerEndpoint, ClientId, Did, Issuer};
pub use tranquil_types::{AuthorizationCode, DeviceId, RefreshToken, RequestId, TokenId};
@@ -195,9 +195,9 @@ pub struct ProtectedResourceMetadata {
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AuthorizationServerMetadata {
pub issuer: String,
pub authorization_endpoint: String,
pub token_endpoint: String,
pub issuer: Issuer,
pub authorization_endpoint: AuthServerEndpoint,
pub token_endpoint: AuthServerEndpoint,
pub jwks_uri: String,
pub registration_endpoint: Option<String>,
pub scopes_supported: Option<Vec<String>>,
@@ -206,7 +206,7 @@ pub struct AuthorizationServerMetadata {
pub grant_types_supported: Option<Vec<String>>,
pub token_endpoint_auth_methods_supported: Option<Vec<String>>,
pub code_challenge_methods_supported: Option<Vec<String>>,
pub pushed_authorization_request_endpoint: Option<String>,
pub pushed_authorization_request_endpoint: Option<AuthServerEndpoint>,
pub require_pushed_authorization_requests: Option<bool>,
pub dpop_signing_alg_values_supported: Option<Vec<String>>,
pub authorization_response_iss_parameter_supported: Option<bool>,
+2 -1
View File
@@ -15,7 +15,7 @@ tranquil-auth = { workspace = true }
tranquil-oauth = { workspace = true }
tranquil-comms = { workspace = true }
tranquil-signal = { workspace = true }
tranquil-db = { workspace = true }
tranquil-db = { workspace = true, features = ["postgres"] }
tranquil-db-traits = { workspace = true }
tranquil-store = { workspace = true }
tranquil-lexicon = { workspace = true, features = ["resolve"] }
@@ -86,6 +86,7 @@ frontend = []
native-tls-roots = ["tranquil-oauth/native-tls-roots"]
[dev-dependencies]
tranquil-infra = { workspace = true, features = ["testing"] }
tempfile = "3"
ciborium = { workspace = true }
ctor = { workspace = true }
+3 -1
View File
@@ -1,4 +1,6 @@
pub use tranquil_cache::{Cache, CacheError, DistributedRateLimiter, NoOpCache, create_cache};
pub use tranquil_cache::{
Cache, CacheError, DistributedRateLimiter, NoOpCache, cached_json, create_cache,
};
#[cfg(feature = "valkey")]
pub use tranquil_cache::{RedisRateLimiter, ValkeyCache};
+1 -48
View File
@@ -1,48 +1 @@
use crate::types::{CidLink, Did, Handle, Jti};
pub fn session_key(did: &Did, jti: &Jti) -> String {
format!("auth:session:{}:{}", did, jti)
}
pub fn signing_key_key(did: &Did) -> String {
format!("auth:key:{}", did)
}
pub fn user_status_key(did: &Did) -> String {
format!("auth:status:{}", did)
}
pub fn handle_key(handle: &Handle) -> String {
format!("handle:{}", handle)
}
pub fn reauth_key(did: &Did) -> String {
format!("reauth:{}", did)
}
pub fn plc_doc_key(did: &Did) -> String {
format!("plc:doc:{}", did)
}
pub fn plc_data_key(did: &Did) -> String {
format!("plc:data:{}", did)
}
pub fn email_update_key(did: &Did) -> String {
format!("email_update:{}", did)
}
pub fn scope_ref_key(cid: &CidLink) -> String {
format!("scope_ref:{}", cid)
}
pub fn auto_verify_sent_key(did: &Did) -> String {
format!("auto_verify_sent:{}", did)
}
pub fn permission_set_key(nsid: &tranquil_types::Nsid, aud: Option<&str>) -> String {
match aud {
Some(a) => format!("permset:{}:{}", nsid, a),
None => format!("permset:{}", nsid),
}
}
pub use tranquil_cache::cache_keys::*;
+23 -16
View File
@@ -13,6 +13,16 @@ pub use tranquil_db_traits::DelegationActionType;
use crate::did::DidResolutionError;
use crate::state::AppState;
use crate::types::{Did, Handle};
use tranquil_types::did_doc::{PdsEndpointError, extract_handle, extract_pds_endpoint};
use tranquil_types::{InvalidHttpUrl, PdsUrl};
#[derive(Debug, thiserror::Error)]
pub enum IdentityResolutionError {
#[error(transparent)]
DidResolution(#[from] DidResolutionError),
#[error("remote PDS endpoint is unusable: {0}")]
PdsEndpoint(InvalidHttpUrl),
}
#[derive(serde::Serialize)]
#[serde(rename_all = "camelCase")]
@@ -21,14 +31,14 @@ pub struct ResolvedIdentity {
#[serde(skip_serializing_if = "Option::is_none")]
pub handle: Option<Handle>,
#[serde(skip_serializing_if = "Option::is_none")]
pub pds_url: Option<String>,
pub pds_url: Option<PdsUrl>,
pub is_local: bool,
}
pub async fn resolve_identity(
state: &AppState,
did: &Did,
) -> Result<ResolvedIdentity, DidResolutionError> {
) -> Result<ResolvedIdentity, IdentityResolutionError> {
let is_local = state
.repos
.user
@@ -38,26 +48,23 @@ pub async fn resolve_identity(
.flatten()
.is_some();
let did_doc = state.did_resolver.resolve_did(did).await?;
let did_doc = state.did_resolver.fetch_did_document(did).await?;
let pds_url = did_doc.services.iter().find_map(|svc| {
if (svc.id == "#atproto_pds" || svc.id.ends_with("#atproto_pds"))
&& svc.service_type == "AtprotoPersonalDataServer"
{
Some(svc.service_endpoint.clone())
} else {
let pds_url = match (extract_pds_endpoint(&did_doc), is_local) {
(Ok(url), _) => Some(url),
(Err(PdsEndpointError::Missing), _) => None,
(Err(PdsEndpointError::Invalid(e)), true) => {
tracing::debug!(did = %did, error = %e, "local account has an unusable PDS endpoint");
None
}
});
let handle = did_doc
.also_known_as
.iter()
.find_map(|alias| alias.strip_prefix("at://"))
.and_then(|s| Handle::new(s).ok());
(Err(PdsEndpointError::Invalid(e)), false) => {
return Err(IdentityResolutionError::PdsEndpoint(e));
}
};
Ok(ResolvedIdentity {
did: did.clone(),
handle,
handle: extract_handle(&did_doc),
pds_url,
is_local,
})
+70 -65
View File
@@ -10,10 +10,12 @@ use tranquil_oauth::{
AuthorizationServerMetadata, ClientMetadata, compute_es256_jkt, compute_pkce_challenge,
create_dpop_proof,
};
use tranquil_types::{AuthorizationCode, ClientId, Did};
use tranquil_types::{AuthorizationCode, ClientId, CrossPdsState, Did, Issuer, PdsUrl};
use crate::cache::Cache;
const SERVER_METADATA_TTL: Duration = Duration::from_secs(300);
#[derive(Error, Debug)]
pub enum CrossPdsError {
#[error("failed to fetch OAuth metadata: {0}")]
@@ -32,11 +34,11 @@ pub enum CrossPdsError {
pub struct CrossPdsAuthState {
pub original_request_uri: String,
pub controller_did: Did,
pub controller_pds_url: String,
pub controller_pds_url: PdsUrl,
pub code_verifier: String,
pub dpop_private_key_der: String,
pub delegated_did: Did,
pub expected_issuer: Option<String>,
pub expected_issuer: Option<Issuer>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -70,17 +72,23 @@ impl CrossPdsOAuthClient {
let http = Client::builder()
.timeout(Duration::from_secs(15))
.connect_timeout(Duration::from_secs(5))
.redirect(tranquil_types::redirect_policy(
tranquil_types::ReachPolicy::GlobalOnly,
))
.dns_resolver(tranquil_types::dns_guard(
tranquil_types::ReachPolicy::GlobalOnly,
))
.build()
.unwrap_or_else(|_| Client::new());
.expect("failed to build cross-PDS OAuth HTTP client");
Self { http, cache }
}
pub async fn store_auth_state(
&self,
state_key: &str,
state_key: &CrossPdsState,
auth_state: &CrossPdsAuthState,
) -> Result<(), CrossPdsError> {
let cache_key = format!("cross_pds_state:{}", state_key);
let cache_key = crate::cache_keys::cross_pds_state_key(state_key);
let json_bytes = serde_json::to_vec(auth_state)
.map_err(|e| CrossPdsError::ParFailed(format!("serialize auth state: {}", e)))?;
let encrypted = crate::config::encrypt_key(&json_bytes)
@@ -93,9 +101,9 @@ impl CrossPdsOAuthClient {
pub async fn retrieve_auth_state(
&self,
state_key: &str,
state_key: &CrossPdsState,
) -> Result<CrossPdsAuthState, CrossPdsError> {
let cache_key = format!("cross_pds_state:{}", state_key);
let cache_key = crate::cache_keys::cross_pds_state_key(state_key);
let encrypted_bytes = self.cache.get_bytes(&cache_key).await.ok_or_else(|| {
CrossPdsError::TokenExchangeFailed("auth state expired or not found".into())
})?;
@@ -110,13 +118,11 @@ impl CrossPdsOAuthClient {
})
}
pub async fn check_remote_is_delegated(&self, pds_url: &str, did: &Did) -> Option<bool> {
let url = format!(
"{}/oauth/security-status?identifier={}",
pds_url.trim_end_matches('/'),
urlencoding::encode(did.as_str())
);
let resp = self.http.get(&url).send().await.ok()?;
pub async fn check_remote_is_delegated(&self, pds_url: &PdsUrl, did: &Did) -> Option<bool> {
let mut url = pds_url.endpoint("oauth/security-status");
url.query_pairs_mut()
.append_pair("identifier", did.as_str());
let resp = self.http.get(url).send().await.ok()?;
if !resp.status().is_success() {
return None;
}
@@ -176,24 +182,12 @@ impl CrossPdsOAuthClient {
Ok(resp)
}
fn require_https(url: &str, label: &str) -> Result<(), CrossPdsError> {
if !url.starts_with("https://") {
return Err(CrossPdsError::MetadataFetch(format!(
"{} must use HTTPS, got: {}",
label, url
)));
}
Ok(())
}
async fn resolve_authorization_server(&self, pds_url: &str) -> Result<String, CrossPdsError> {
Self::require_https(pds_url, "PDS URL")?;
let resource_url = format!(
"{}/.well-known/oauth-protected-resource",
pds_url.trim_end_matches('/')
);
if let Ok(resp) = self.http.get(&resource_url).send().await
async fn resolve_authorization_server(
&self,
pds_url: &PdsUrl,
) -> Result<Issuer, CrossPdsError> {
let resource_url = pds_url.endpoint(".well-known/oauth-protected-resource");
if let Ok(resp) = self.http.get(resource_url).send().await
&& resp.status().is_success()
{
#[derive(Deserialize)]
@@ -203,30 +197,36 @@ impl CrossPdsOAuthClient {
if let Ok(pr) = resp.json::<ProtectedResource>().await
&& let Some(server) = pr.authorization_servers.and_then(|s| s.into_iter().next())
{
Self::require_https(&server, "Authorization server")?;
return Ok(server);
return Issuer::new(server)
.map_err(|e| CrossPdsError::MetadataFetch(e.to_string()));
}
}
Ok(pds_url.trim_end_matches('/').to_string())
Issuer::new(pds_url.as_str()).map_err(|e| CrossPdsError::MetadataFetch(e.to_string()))
}
pub async fn fetch_server_metadata(
&self,
pds_url: &str,
pds_url: &PdsUrl,
) -> Result<AuthorizationServerMetadata, CrossPdsError> {
let cache_key = format!("cross_pds_oauth_meta:{}", pds_url);
if let Some(cached) = self.cache.get(&cache_key).await
&& let Ok(meta) = serde_json::from_str(&cached)
{
return Ok(meta);
}
crate::cache::cached_json(
self.cache.as_ref(),
&crate::cache_keys::cross_pds_oauth_meta_key(pds_url),
SERVER_METADATA_TTL,
|| self.fetch_verified_server_metadata(pds_url),
)
.await
}
async fn fetch_verified_server_metadata(
&self,
pds_url: &PdsUrl,
) -> Result<AuthorizationServerMetadata, CrossPdsError> {
let auth_server = self.resolve_authorization_server(pds_url).await?;
let url = format!("{}/.well-known/oauth-authorization-server", auth_server);
let url = auth_server.endpoint(".well-known/oauth-authorization-server");
let resp = self
.http
.get(&url)
.get(url.clone())
.send()
.await
.map_err(|e| CrossPdsError::MetadataFetch(e.to_string()))?;
@@ -244,11 +244,11 @@ impl CrossPdsOAuthClient {
.await
.map_err(|e| CrossPdsError::MetadataFetch(e.to_string()))?;
if let Ok(json_str) = serde_json::to_string(&meta) {
let _ = self
.cache
.set(&cache_key, &json_str, Duration::from_secs(300))
.await;
if meta.issuer != auth_server {
return Err(CrossPdsError::MetadataFetch(format!(
"issuer mismatch: {} serves metadata for {}",
auth_server, meta.issuer
)));
}
Ok(meta)
@@ -256,22 +256,22 @@ impl CrossPdsOAuthClient {
pub async fn initiate_par(
&self,
pds_url: &str,
pds_url: &PdsUrl,
urls: &DelegationOAuthUrls,
login_hint: Option<&str>,
original_request_uri: &str,
controller_did: &Did,
delegated_did: &Did,
) -> Result<(ParResult, CrossPdsAuthState, String), CrossPdsError> {
) -> Result<(ParResult, CrossPdsAuthState, CrossPdsState), CrossPdsError> {
let meta = self.fetch_server_metadata(pds_url).await?;
let par_endpoint = meta
.pushed_authorization_request_endpoint
.as_deref()
.as_ref()
.ok_or(CrossPdsError::NoParEndpoint)?;
let code_verifier = crate::util::generate_random_token();
let code_challenge = compute_pkce_challenge(&code_verifier);
let state = crate::util::generate_random_token();
let state = CrossPdsState::new(crate::util::generate_random_token());
let signing_key = SigningKey::random(&mut OsRng);
let dpop_key_der = URL_SAFE_NO_PAD.encode(signing_key.to_bytes());
@@ -284,7 +284,7 @@ impl CrossPdsOAuthClient {
("client_id", urls.client_id.to_string()),
("redirect_uri", urls.redirect_uri.clone()),
("scope", "atproto".to_string()),
("state", state.clone()),
("state", state.to_string()),
("code_challenge", code_challenge),
("code_challenge_method", "S256".to_string()),
("dpop_jkt", dpop_jkt),
@@ -294,7 +294,7 @@ impl CrossPdsOAuthClient {
}
let resp = self
.send_with_dpop_retry(&signing_key, "POST", par_endpoint, &params, None)
.send_with_dpop_retry(&signing_key, "POST", par_endpoint.as_str(), &params, None)
.await
.map_err(|e| CrossPdsError::ParFailed(e.to_string()))?;
@@ -313,17 +313,16 @@ impl CrossPdsOAuthClient {
.await
.map_err(|e| CrossPdsError::ParFailed(e.to_string()))?;
let authorize_url = format!(
"{}?request_uri={}&client_id={}",
meta.authorization_endpoint,
urlencoding::encode(&par_resp.request_uri),
urlencoding::encode(&urls.client_id)
);
let mut authorize_url = meta.authorization_endpoint.url().clone();
authorize_url
.query_pairs_mut()
.append_pair("request_uri", &par_resp.request_uri)
.append_pair("client_id", &urls.client_id);
let auth_state = CrossPdsAuthState {
original_request_uri: original_request_uri.to_string(),
controller_did: controller_did.clone(),
controller_pds_url: pds_url.to_string(),
controller_pds_url: pds_url.clone(),
code_verifier,
dpop_private_key_der: dpop_key_der,
delegated_did: delegated_did.clone(),
@@ -333,7 +332,7 @@ impl CrossPdsOAuthClient {
Ok((
ParResult {
request_uri: par_resp.request_uri,
authorize_url,
authorize_url: authorize_url.into(),
},
auth_state,
state,
@@ -366,7 +365,13 @@ impl CrossPdsOAuthClient {
];
let resp = self
.send_with_dpop_retry(&signing_key, "POST", &meta.token_endpoint, &params, None)
.send_with_dpop_retry(
&signing_key,
"POST",
meta.token_endpoint.as_str(),
&params,
None,
)
.await
.map_err(CrossPdsError::TokenExchangeFailed)?;
+1 -1
View File
@@ -18,7 +18,7 @@ tokio = { workspace = true }
tokio-util = { workspace = true }
futures = { workspace = true }
serde_json = { workspace = true }
url = "2.5"
url = { workspace = true }
uuid = { workspace = true }
thiserror = { workspace = true }
+7
View File
@@ -10,8 +10,15 @@ chrono = { workspace = true }
cid = { workspace = true }
jacquard-common = { workspace = true }
rand = { workspace = true }
reqwest = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
sqlx = { workspace = true }
thiserror = { workspace = true }
tokio = { workspace = true, features = ["net", "rt"] }
tracing = { workspace = true }
url = { workspace = true }
uuid = { workspace = true }
[dev-dependencies]
tokio = { workspace = true }
+624 -10
View File
@@ -1,6 +1,8 @@
use serde::{Deserialize, Serialize};
use std::borrow::Cow;
use std::fmt;
use std::hash::Hash;
use std::marker::PhantomData;
use std::ops::Deref;
use std::str::FromStr;
@@ -813,6 +815,10 @@ simple_string_newtype! {
pub struct Jti;
}
simple_string_newtype_no_sqlx! {
pub struct CrossPdsState;
}
simple_string_newtype! {
pub struct AuthorizationCode;
}
@@ -881,6 +887,425 @@ impl fmt::Display for CommsChannel {
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HostReach {
Global,
Loopback,
Private,
}
fn ipv4_reach(ip: std::net::Ipv4Addr) -> HostReach {
let [a, b, c, _] = ip.octets();
match ip {
_ if ip.is_loopback() => HostReach::Loopback,
_ if ip.is_private()
|| ip.is_link_local()
|| ip.is_multicast()
|| ip.is_documentation()
|| a == 0
|| a == 100 && (64..128).contains(&b)
|| a == 192 && b == 0 && c == 0
|| a == 192 && b == 88 && c == 99
|| a == 198 && (18..20).contains(&b)
|| a & 0xf0 == 240 =>
{
HostReach::Private
}
_ => HostReach::Global,
}
}
fn ipv6_reach(ip: std::net::Ipv6Addr) -> HostReach {
let seg = ip.segments();
let embedded_ipv4 =
|hi: u16, lo: u16| std::net::Ipv4Addr::from((u32::from(hi) << 16) | u32::from(lo));
match ip.to_ipv4_mapped() {
Some(mapped) => ipv4_reach(mapped),
None => match ip {
_ if ip.is_loopback() => HostReach::Loopback,
_ if seg[..6] == [0, 0, 0, 0, 0, 0] => ipv4_reach(embedded_ipv4(seg[6], seg[7])),
_ if seg[..2] == [0x2001, 0] => ipv4_reach(embedded_ipv4(!seg[6], !seg[7])),
_ if seg[0] == 0x2002 => ipv4_reach(embedded_ipv4(seg[1], seg[2])),
_ if seg[..6] == [0x64, 0xff9b, 0, 0, 0, 0] => {
ipv4_reach(embedded_ipv4(seg[6], seg[7]))
}
_ if seg[..3] == [0x64, 0xff9b, 1] => HostReach::Private,
_ if ip.is_unspecified()
|| ip.is_multicast()
|| seg[0] & 0xfe00 == 0xfc00
|| seg[0] & 0xffc0 == 0xfe80
|| seg[..2] == [0x2001, 0x0db8] =>
{
HostReach::Private
}
_ => HostReach::Global,
},
}
}
fn host_reach(host: url::Host<&str>) -> HostReach {
match host {
url::Host::Ipv4(ip) => ipv4_reach(ip),
url::Host::Ipv6(ip) => ipv6_reach(ip),
url::Host::Domain(name) => {
let name = name.trim_end_matches('.').to_ascii_lowercase();
match name.as_str() {
"localhost" => HostReach::Loopback,
_ if name.ends_with(".localhost") => HostReach::Loopback,
_ if name.ends_with(".local")
|| name.ends_with(".internal")
|| name.ends_with(".home.arpa")
|| name == "home.arpa" =>
{
HostReach::Private
}
_ => HostReach::Global,
}
}
}
}
pub fn url_reach(url: &url::Url) -> Option<HostReach> {
url.host().map(host_reach)
}
pub fn ip_reach(ip: std::net::IpAddr) -> HostReach {
match ip {
std::net::IpAddr::V4(v4) => ipv4_reach(v4),
std::net::IpAddr::V6(v6) => ipv6_reach(v6),
}
}
pub fn reach_permits(reach: HostReach, policy: ReachPolicy) -> bool {
matches!(
(reach, policy),
(HostReach::Global, _)
| (
HostReach::Loopback,
ReachPolicy::AllowLoopback | ReachPolicy::AllowPrivate,
)
| (HostReach::Private, ReachPolicy::AllowPrivate)
)
}
pub fn url_reach_permits(url: &url::Url, policy: ReachPolicy) -> bool {
let Some(reach) = url_reach(url) else {
return false;
};
let scheme_permits = matches!(
(url.scheme(), reach),
("https", _) | ("http", HostReach::Loopback | HostReach::Private)
);
scheme_permits && reach_permits(reach, policy)
}
fn parse_http_url(s: &str, policy: ReachPolicy, allow_query: bool) -> Option<url::Url> {
let parsed = url::Url::parse(s).ok()?;
let rejected = (parsed.query().is_some() && !allow_query)
|| parsed.fragment().is_some()
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| !url_reach_permits(&parsed, policy);
match rejected {
true => None,
false => Some(parsed),
}
}
const REDIRECT_HOP_LIMIT: usize = 5;
pub fn redirect_policy(policy: ReachPolicy) -> reqwest::redirect::Policy {
reqwest::redirect::Policy::custom(move |attempt| {
let over_limit = attempt.previous().len() > REDIRECT_HOP_LIMIT;
let permitted = url_reach_permits(attempt.url(), policy);
let target = attempt.url().clone();
match (over_limit, permitted) {
(true, _) => attempt.error(format!("more than {} redirect hops", REDIRECT_HOP_LIMIT)),
(false, false) => attempt.error(format!(
"redirect target {} is outside the allowed host reach",
target
)),
(false, true) => attempt.follow(),
}
})
}
pub struct ReachGuardedDns(ReachPolicy);
impl reqwest::dns::Resolve for ReachGuardedDns {
fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving {
let policy = self.0;
Box::pin(async move {
let host = name.as_str().to_owned();
let permitted: Vec<std::net::SocketAddr> = tokio::net::lookup_host((host.as_str(), 0))
.await?
.filter(|addr| reach_permits(ip_reach(addr.ip()), policy))
.collect();
match permitted.is_empty() {
true => Err(format!(
"no resolved address for {} is inside the allowed host reach",
host
)
.into()),
false => Ok(Box::new(permitted.into_iter()) as reqwest::dns::Addrs),
}
})
}
}
pub fn dns_guard(policy: ReachPolicy) -> std::sync::Arc<ReachGuardedDns> {
std::sync::Arc::new(ReachGuardedDns(policy))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReachPolicy {
AllowLoopback,
AllowPrivate,
GlobalOnly,
}
impl ReachPolicy {
#[cfg(debug_assertions)]
pub const DEBUG_LOOPBACK: ReachPolicy = ReachPolicy::AllowLoopback;
#[cfg(not(debug_assertions))]
pub const DEBUG_LOOPBACK: ReachPolicy = ReachPolicy::GlobalOnly;
}
pub trait UrlKind {
const LABEL: &'static str;
const REACH_POLICY: ReachPolicy;
const ALLOW_QUERY: bool;
}
pub mod url_kind {
use super::{ReachPolicy, UrlKind};
pub struct AuthServerEndpoint;
impl UrlKind for AuthServerEndpoint {
const LABEL: &'static str = "authorization server endpoint";
const REACH_POLICY: ReachPolicy = ReachPolicy::GlobalOnly;
const ALLOW_QUERY: bool = true;
}
pub struct Issuer;
impl UrlKind for Issuer {
const LABEL: &'static str = "issuer";
const REACH_POLICY: ReachPolicy = ReachPolicy::GlobalOnly;
const ALLOW_QUERY: bool = false;
}
pub struct Jwks;
impl UrlKind for Jwks {
const LABEL: &'static str = "JWKS URI";
const REACH_POLICY: ReachPolicy = ReachPolicy::DEBUG_LOOPBACK;
const ALLOW_QUERY: bool = true;
}
pub struct Pds;
impl UrlKind for Pds {
const LABEL: &'static str = "PDS URL";
const REACH_POLICY: ReachPolicy = ReachPolicy::GlobalOnly;
const ALLOW_QUERY: bool = false;
}
pub struct SchemaHost;
impl UrlKind for SchemaHost {
const LABEL: &'static str = "schema host URL";
const REACH_POLICY: ReachPolicy = ReachPolicy::DEBUG_LOOPBACK;
const ALLOW_QUERY: bool = false;
}
pub struct SsoIssuer;
impl UrlKind for SsoIssuer {
const LABEL: &'static str = "SSO issuer";
const REACH_POLICY: ReachPolicy = ReachPolicy::AllowPrivate;
const ALLOW_QUERY: bool = false;
}
pub struct SsoJwks;
impl UrlKind for SsoJwks {
const LABEL: &'static str = "SSO JWKS URI";
const REACH_POLICY: ReachPolicy = ReachPolicy::AllowPrivate;
const ALLOW_QUERY: bool = true;
}
}
pub struct HttpUrl<K: UrlKind> {
raw: String,
parsed: url::Url,
kind: PhantomData<fn() -> K>,
}
pub type AuthServerEndpoint = HttpUrl<url_kind::AuthServerEndpoint>;
pub type Issuer = HttpUrl<url_kind::Issuer>;
pub type JwksUri = HttpUrl<url_kind::Jwks>;
pub type PdsUrl = HttpUrl<url_kind::Pds>;
pub type SchemaHostUrl = HttpUrl<url_kind::SchemaHost>;
pub type SsoIssuer = HttpUrl<url_kind::SsoIssuer>;
pub type SsoJwksUri = HttpUrl<url_kind::SsoJwks>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InvalidHttpUrl {
pub kind: &'static str,
pub value: String,
}
impl fmt::Display for InvalidHttpUrl {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "invalid {}: {}", self.kind, self.value)
}
}
impl std::error::Error for InvalidHttpUrl {}
impl<K: UrlKind> HttpUrl<K> {
pub fn new(s: impl Into<String>) -> Result<Self, InvalidHttpUrl> {
let raw = s.into();
match parse_http_url(&raw, K::REACH_POLICY, K::ALLOW_QUERY) {
Some(parsed) => Ok(Self {
raw,
parsed,
kind: PhantomData,
}),
None => Err(InvalidHttpUrl {
kind: K::LABEL,
value: raw,
}),
}
}
/// The URL as given.
/// OIDC & OAuth define issuer comparison as an
/// exact string match,
/// so anything sent to or compared against a peer uses this.
/// Give it to us raw & wriggling!!
pub fn as_str(&self) -> &str {
&self.raw
}
/// The parsed form: lowercased scheme and host, with `/` for a bare authority.
/// Cache keys use this so `https://oyster.cafe` and `https://oyster.cafe/` share one entry.
pub fn canonical(&self) -> &str {
self.parsed.as_str()
}
pub fn url(&self) -> &url::Url {
&self.parsed
}
pub fn endpoint(&self, path: &str) -> url::Url {
let mut url = self.parsed.clone();
let base = url.path().trim_end_matches('/').to_owned();
url.set_path(&format!("{}/{}", base, path.trim_start_matches('/')));
url
}
}
pub mod http_url {
use super::{HttpUrl, UrlKind};
use serde::Deserialize;
pub fn deserialize_optional<'de, D, K>(deserializer: D) -> Result<Option<HttpUrl<K>>, D::Error>
where
D: serde::Deserializer<'de>,
K: UrlKind,
{
Ok(Option::<String>::deserialize(deserializer)?.and_then(|s| {
HttpUrl::new(s)
.inspect_err(|e| tracing::warn!(error = %e, "discarding unusable URL field"))
.ok()
}))
}
}
impl<K: UrlKind> FromStr for HttpUrl<K> {
type Err = InvalidHttpUrl;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Self::new(s)
}
}
impl<K: UrlKind> fmt::Debug for HttpUrl<K> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}({})", K::LABEL, self.raw)
}
}
impl<K: UrlKind> fmt::Display for HttpUrl<K> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.raw)
}
}
impl<K: UrlKind> Clone for HttpUrl<K> {
fn clone(&self) -> Self {
Self {
raw: self.raw.clone(),
parsed: self.parsed.clone(),
kind: PhantomData,
}
}
}
impl<K: UrlKind> PartialEq for HttpUrl<K> {
fn eq(&self, other: &Self) -> bool {
self.parsed == other.parsed
}
}
impl<K: UrlKind> Eq for HttpUrl<K> {}
impl<K: UrlKind> Hash for HttpUrl<K> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.parsed.hash(state);
}
}
impl<K: UrlKind> Serialize for HttpUrl<K> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(&self.raw)
}
}
impl<'de, K: UrlKind> Deserialize<'de> for HttpUrl<K> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
Self::new(s).map_err(|e| serde::de::Error::custom(e.to_string()))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum EmailTokenPurpose {
UpdateEmail,
ConfirmEmail,
DeleteAccount,
ResetPassword,
PlcOperation,
}
impl EmailTokenPurpose {
pub fn as_str(&self) -> &'static str {
match self {
Self::UpdateEmail => "update_email",
Self::ConfirmEmail => "confirm_email",
Self::DeleteAccount => "delete_account",
Self::ResetPassword => "reset_password",
Self::PlcOperation => "plc_operation",
}
}
}
impl fmt::Display for EmailTokenPurpose {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)]
#[serde(rename_all = "snake_case")]
#[sqlx(type_name = "comms_type", rename_all = "snake_case")]
@@ -894,24 +1319,32 @@ pub enum CommsType {
}
pub mod did_doc {
pub fn extract_pds_endpoint(doc: &serde_json::Value) -> Option<String> {
use crate::{HttpUrl, InvalidHttpUrl, UrlKind};
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum PdsEndpointError {
#[error("DID document has no atproto PDS service entry")]
Missing,
#[error(transparent)]
Invalid(#[from] InvalidHttpUrl),
}
pub fn extract_pds_endpoint<K: UrlKind>(
doc: &serde_json::Value,
) -> Result<HttpUrl<K>, PdsEndpointError> {
doc.get("service")
.and_then(|s| s.as_array())
.and_then(|services| {
services.iter().find_map(|svc| {
let id = svc.get("id").and_then(|v| v.as_str()).unwrap_or_default();
let svc_type = svc.get("type").and_then(|v| v.as_str()).unwrap_or_default();
if (id == "#atproto_pds" || id.ends_with("#atproto_pds"))
&& svc_type == "AtprotoPersonalDataServer"
{
svc.get("serviceEndpoint")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
} else {
None
}
((id == "#atproto_pds" || id.ends_with("#atproto_pds"))
&& svc_type == "AtprotoPersonalDataServer")
.then(|| svc.get("serviceEndpoint").and_then(|v| v.as_str()))?
})
})
.ok_or(PdsEndpointError::Missing)
.and_then(|endpoint| HttpUrl::new(endpoint).map_err(PdsEndpointError::Invalid))
}
pub fn extract_handle(doc: &serde_json::Value) -> Option<crate::Handle> {
@@ -928,6 +1361,187 @@ pub mod did_doc {
}
}
#[cfg(test)]
mod http_url_tests {
use super::did_doc::{PdsEndpointError, extract_pds_endpoint};
use super::{
AuthServerEndpoint, Issuer, JwksUri, PdsUrl, SchemaHostUrl, SsoIssuer, SsoJwksUri,
};
#[test]
fn extract_pds_endpoint_selects_the_pds_service_and_reports_missing_or_invalid() {
let labeler = serde_json::json!({
"id": "#atproto_labeler",
"type": "AtprotoLabeler",
"serviceEndpoint": "https://labeler.nel.pet"
});
let pds = |endpoint: &str| {
serde_json::json!({
"id": "#atproto_pds",
"type": "AtprotoPersonalDataServer",
"serviceEndpoint": endpoint
})
};
let both = serde_json::json!({ "service": [labeler.clone(), pds("https://oyster.cafe")] });
assert_eq!(
extract_pds_endpoint::<super::url_kind::Pds>(&both)
.unwrap()
.as_str(),
"https://oyster.cafe"
);
[
serde_json::json!({ "service": [labeler] }),
serde_json::json!({}),
]
.iter()
.for_each(|doc| {
assert_eq!(
extract_pds_endpoint::<super::url_kind::Pds>(doc).unwrap_err(),
PdsEndpointError::Missing
);
});
let plain_http = serde_json::json!({ "service": [pds("http://oyster.cafe")] });
assert!(matches!(
extract_pds_endpoint::<super::url_kind::Pds>(&plain_http),
Err(PdsEndpointError::Invalid(_))
));
}
#[test]
fn pds_and_jwks_kinds_reject_private_and_reserved_addresses() {
[
"https://10.0.0.1",
"https://192.168.1.1",
"https://172.16.0.1",
"https://169.254.169.254/latest/meta-data",
"https://100.64.0.1",
"https://0.1.2.3",
"https://192.0.0.8",
"https://192.88.99.1",
"https://[fd00::1]",
"https://[fe80::1]",
"https://[ff02::1]",
"https://[::ffff:10.0.0.1]",
"https://[64:ff9b::a00:1]",
"https://[64:ff9b:1::1]",
"https://[2002:a00:1::]",
"https://[::10.0.0.1]",
"https://[2001:0:0:0:0:0:f5ff:fffe]",
"https://kelp.internal",
"https://whelk.local",
"https://limpet.home.arpa",
]
.iter()
.for_each(|url| {
assert!(PdsUrl::new(*url).is_err(), "PdsUrl must reject {url}");
assert!(JwksUri::new(*url).is_err(), "JwksUri must reject {url}");
});
[
"https://oyster.cafe",
"https://[64:ff9b::808:808]",
"https://[::8.8.8.8]",
"https://[2001::f7f7:f7f7]",
]
.iter()
.for_each(|url| assert!(PdsUrl::new(*url).is_ok(), "PdsUrl must accept {url}"));
}
#[test]
fn each_kind_applies_its_own_local_host_policy() {
assert!(PdsUrl::new("http://127.0.0.1:2583").is_err());
assert!(PdsUrl::new("https://localhost").is_err());
assert!(Issuer::new("http://localhost:8080").is_err());
assert_eq!(
JwksUri::new("http://localhost:8080/keys").is_ok(),
cfg!(debug_assertions)
);
assert_eq!(
SchemaHostUrl::new("http://127.0.0.1:2583").is_ok(),
cfg!(debug_assertions)
);
assert!(SsoJwksUri::new("http://127.0.0.1:8080/keys").is_ok());
assert!(SsoJwksUri::new("http://[::1]:8080/keys").is_ok());
assert!(SsoJwksUri::new("http://squid.localhost:8080/keys").is_ok());
assert!(SsoJwksUri::new("https://keycloak.internal/keys?client=squid").is_ok());
assert!(SsoJwksUri::new("http://oyster.cafe/keys").is_err());
assert!(SsoIssuer::new("https://keycloak.internal/realms/uni").is_ok());
assert!(SsoIssuer::new("http://10.0.0.5:8080").is_ok());
assert!(SsoIssuer::new("http://localhost:8080").is_ok());
assert!(SsoIssuer::new("http://oyster.cafe").is_err());
assert!(AuthServerEndpoint::new("https://oyster.cafe/oauth/par?tenant=uni").is_ok());
[
"https://169.254.169.254/oauth/par",
"https://[fd00::1]/oauth/par",
"http://127.0.0.1:2583/oauth/par",
"http://oyster.cafe/oauth/par",
]
.iter()
.for_each(|url| {
assert!(
AuthServerEndpoint::new(*url).is_err(),
"AuthServerEndpoint must reject {url}"
);
});
}
#[test]
fn canonicalization_keeps_identity_and_rejects_query_fragment_and_userinfo() {
assert_eq!(
PdsUrl::new("HTTPS://oyster.cafe").unwrap().canonical(),
"https://oyster.cafe/"
);
let bare = PdsUrl::new("https://oyster.cafe").unwrap();
let slashed = PdsUrl::new("https://oyster.cafe/").unwrap();
assert_eq!(bare, slashed);
assert_eq!(bare.canonical(), slashed.canonical());
let issuer = Issuer::new("https://accounts.google.com").unwrap();
assert_eq!(issuer.as_str(), "https://accounts.google.com");
assert_eq!(issuer.canonical(), "https://accounts.google.com/");
assert_eq!(
PdsUrl::new("https://oyster.cafe/pds/")
.unwrap()
.endpoint(".well-known/oauth-protected-resource")
.as_str(),
"https://oyster.cafe/pds/.well-known/oauth-protected-resource"
);
assert_eq!(
JwksUri::new("https://oyster.cafe/keys?appid=abc")
.expect("JwksUri keeps the query")
.canonical(),
"https://oyster.cafe/keys?appid=abc"
);
assert!(PdsUrl::new("https://oyster.cafe/?x=1").is_err());
assert!(PdsUrl::new("https://oyster.cafe/#frag").is_err());
assert!(PdsUrl::new("https://nel:pw@oyster.cafe").is_err());
assert!(Issuer::new("https://oyster.cafe/?x=1").is_err());
assert!(JwksUri::new("https://oyster.cafe/keys#frag").is_err());
}
}
#[cfg(test)]
mod dns_guard_tests {
use super::{ReachPolicy, dns_guard};
use reqwest::dns::Resolve;
#[tokio::test]
async fn the_policy_gates_loopback_resolution() {
let name = |host: &str| host.parse::<reqwest::dns::Name>().expect("valid hostname");
assert!(
dns_guard(ReachPolicy::GlobalOnly)
.resolve(name("localhost"))
.await
.is_err()
);
let addrs: Vec<_> = dns_guard(ReachPolicy::AllowLoopback)
.resolve(name("localhost"))
.await
.expect("localhost resolves")
.collect();
assert!(!addrs.is_empty());
assert!(addrs.iter().all(|a| a.ip().is_loopback()));
}
}
#[cfg(test)]
mod validated_newtype_tests {
use super::*;