Compare commits

...
Author SHA1 Message Date
Lewis 37234797b4 just: clippy over all targets, lint the bsky-off build
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:14:11 +03:00
Lewis 6297b1a451 cache: DID, SSO, & OAuth client metadata caches onto shared cache
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:05:16 +03:00
Lewis 107149f396 lexicon: schema docs & negative results via cluster cache
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:05:16 +03:00
Lewis bd47cbdaa4 plc: dedup fetch paths, cache TTL from config
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:05:16 +03:00
Lewis 9840ac77cf auth: EmailTokenPurpose from tranquil-types, shared cache key fns, MemoryCache in tests
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:05:16 +03:00
Lewis 420ce1e201 types: HttpUrl newtypes, shared cache key/JSON helpers
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:05:16 +03:00
Lewis c723bc2164 pds: compile bsky-specific proxy, CORS, & validation out under bsky features
Lewis: May this revision serve well! <did:plc:3fwecdnvtcscjnrx2p4n7alz>
2026-08-16 20:05:16 +03:00
52 changed files with 1911 additions and 1122 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))
}
+2 -7
View File
@@ -147,12 +147,7 @@ async fn try_reactivate_migration(
Json(CreateAccountOutput {
handle: handle.clone(),
did: did.clone(),
did_doc: state
.did_resolver
.fetch_did_document(did)
.await
.ok()
.map(|f| (*f).clone()),
did_doc: state.did_resolver.fetch_did_document(did).await.ok(),
access_jwt: access_meta.token,
refresh_jwt: refresh_meta.token,
verification_required,
@@ -568,7 +563,7 @@ pub async fn create_account(
Json(CreateAccountOutput {
handle: handle.clone(),
did,
did_doc: did_doc.map(|f| (*f).clone()),
did_doc,
access_jwt: session.access_jwt,
refresh_jwt: session.refresh_jwt,
verification_required: !is_migration,
@@ -164,6 +164,13 @@ pub async fn resolve_signing_key(
}
}
#[cfg_attr(
not(feature = "bsky"),
expect(
unused_variables,
reason = "only the bsky block writes display_name into the default profile record"
)
)]
pub async fn sequence_new_account(
state: &AppState,
did: &Did,
+3 -3
View File
@@ -351,7 +351,7 @@ pub async fn create_session(
refresh_jwt: refresh_meta.token,
handle,
did: row.did,
did_doc: did_doc.ok().map(|f| (*f).clone()),
did_doc: did_doc.ok(),
email: row.email,
email_confirmed: Some(row.channel_verification.email),
email_auth_factor: email_auth_factor_out,
@@ -444,7 +444,7 @@ pub async fn get_session(
status: account_state.status_for_session().map(String::from),
migrated_to_pds,
migrated_at,
did_doc: did_doc.ok().map(|f| (*f).clone()),
did_doc: did_doc.ok(),
}))
}
Ok(None) => Err(ApiError::AuthenticationFailed(None)),
@@ -800,7 +800,7 @@ async fn build_refresh_session_output(
preferred_locale: u.preferred_locale,
is_admin: u.is_admin,
active: account_state.is_active(),
did_doc: did_doc.ok().map(|f| (*f).clone()),
did_doc: did_doc.ok(),
status: account_state.status_for_session().map(String::from),
}))
}
+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}");
+1 -1
View File
@@ -835,7 +835,7 @@ pub struct PlcConfig {
#[config(env = "PLC_CONNECT_TIMEOUT_SECS", default = 5)]
pub connect_timeout_secs: u64,
/// Seconds to cache DID documents in memory.
/// Seconds to cache DID documents.
#[config(env = "DID_CACHE_TTL_SECS", default = 300)]
pub did_cache_ttl_secs: u64,
}
+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(())
}
}
+4 -3
View File
@@ -5,10 +5,11 @@ edition.workspace = true
license.workspace = true
[features]
resolve = ["dep:reqwest", "dep:hickory-resolver", "dep:tokio", "dep:parking_lot", "dep:tracing", "dep:urlencoding"]
resolve = ["dep:reqwest", "dep:hickory-resolver", "dep:tokio", "dep:parking_lot", "dep:tracing", "dep:tranquil-infra"]
[dependencies]
tranquil-types = { path = "../tranquil-types", default-features = false }
tranquil-types = { workspace = true }
tranquil-infra = { workspace = true, optional = true, features = ["cache-keys"] }
serde = { workspace = true }
serde_json = { workspace = true }
thiserror = { workspace = true }
@@ -19,9 +20,9 @@ hickory-resolver = { workspace = true, optional = true }
tokio = { workspace = true, optional = true }
parking_lot = { workspace = true, optional = true }
tracing = { workspace = true, optional = true }
urlencoding = { workspace = true, optional = true }
[dev-dependencies]
wiremock = { workspace = true }
tokio = { workspace = true }
futures = { workspace = true }
tranquil-infra = { workspace = true, features = ["testing", "cache-keys"] }
+217 -31
View File
@@ -6,9 +6,11 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::Notify;
use tranquil_infra::cache_keys::{lexicon_doc_key, lexicon_negative_key};
use tranquil_infra::{Cache, read_json, write_json};
use tranquil_types::Nsid;
const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(24 * 60 * 60);
const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(60 * 60);
const POSITIVE_CACHE_TTL: Duration = Duration::from_secs(24 * 60 * 60);
const REFRESH_FAILURE_BACKOFF: Duration = Duration::from_secs(60);
const MAX_DYNAMIC_SCHEMAS: usize = 1024;
@@ -17,6 +19,13 @@ struct NegativeEntry {
expires_at: Instant,
}
fn negative_ttl_for(error: &ResolveError) -> Duration {
match error.is_definitive() {
true => NEGATIVE_CACHE_TTL,
false => REFRESH_FAILURE_BACKOFF,
}
}
struct PositiveEntry {
doc: Arc<LexiconDoc>,
expires_at: Instant,
@@ -44,6 +53,7 @@ pub struct DynamicRegistry {
negative_cache: RwLock<HashMap<Nsid, NegativeEntry>>,
in_flight: RwLock<HashMap<Nsid, Arc<Notify>>>,
network_disabled: AtomicBool,
shared: RwLock<Option<Arc<dyn Cache>>>,
}
struct InFlightGuard<'a> {
@@ -70,9 +80,18 @@ impl DynamicRegistry {
negative_cache: RwLock::new(HashMap::new()),
in_flight: RwLock::new(HashMap::new()),
network_disabled: AtomicBool::new(false),
shared: RwLock::new(None),
}
}
pub fn set_shared_cache(&self, cache: Arc<dyn Cache>) {
*self.shared.write() = Some(cache);
}
fn shared_cache(&self) -> Option<Arc<dyn Cache>> {
self.shared.read().clone()
}
pub fn from_env() -> Self {
let registry = Self::new();
let disabled =
@@ -105,13 +124,17 @@ impl DynamicRegistry {
}
pub fn is_negative_cached(&self, nsid: &Nsid) -> bool {
let cache = self.negative_cache.read();
cache
.get(nsid)
.is_some_and(|entry| entry.expires_at > Instant::now())
self.negative_remaining(nsid).is_some()
}
fn insert_negative(&self, nsid: &Nsid) {
fn negative_remaining(&self, nsid: &Nsid) -> Option<Duration> {
self.negative_cache
.read()
.get(nsid)
.and_then(|entry| entry.expires_at.checked_duration_since(Instant::now()))
}
fn insert_negative(&self, nsid: &Nsid, ttl: Duration) {
let mut cache = self.negative_cache.write();
if cache.len() >= MAX_DYNAMIC_SCHEMAS {
let now = Instant::now();
@@ -120,7 +143,7 @@ impl DynamicRegistry {
cache.insert(
nsid.clone(),
NegativeEntry {
expires_at: Instant::now() + NEGATIVE_CACHE_TTL,
expires_at: Instant::now() + ttl,
},
);
}
@@ -159,6 +182,44 @@ impl DynamicRegistry {
arc
}
async fn shared_get(&self, nsid: &Nsid) -> Option<Arc<LexiconDoc>> {
let cache = self.shared_cache()?;
let doc = read_json::<LexiconDoc>(cache.as_ref(), &lexicon_doc_key(nsid)).await?;
Some(self.insert_schema(doc))
}
async fn shared_put(&self, doc: &LexiconDoc) {
let Some(cache) = self.shared_cache() else {
return;
};
write_json(
cache.as_ref(),
&lexicon_doc_key(&doc.id),
doc,
POSITIVE_CACHE_TTL,
)
.await;
let _ = cache.delete(&lexicon_negative_key(&doc.id)).await;
}
async fn shared_is_negative(&self, nsid: &Nsid) -> bool {
match self.shared_cache() {
Some(cache) => cache.get(&lexicon_negative_key(nsid)).await.is_some(),
None => false,
}
}
async fn shared_put_negative(&self, nsid: &Nsid, error: &ResolveError) {
if !error.is_definitive() {
return;
}
if let Some(cache) = self.shared_cache() {
let _ = cache
.set(&lexicon_negative_key(nsid), "1", NEGATIVE_CACHE_TTL)
.await;
}
}
fn bump_expiry(&self, nsid: &Nsid, duration: Duration) {
let mut store = self.store.write();
if let Some(entry) = store.schemas.get_mut(nsid) {
@@ -203,15 +264,23 @@ impl DynamicRegistry {
match self.acquire_leadership(nsid) {
Some(_guard) => match resolver(nsid.clone()).await {
Ok(doc) => Ok(self.insert_schema(doc)),
Ok(doc) => {
self.shared_put(&doc).await;
Ok(self.insert_schema(doc))
}
Err(e) => {
let (doc, source) = match self.shared_get(nsid).await {
Some(doc) => (doc, "shared"),
None => (stale, "local"),
};
self.bump_expiry(nsid, REFRESH_FAILURE_BACKOFF);
tracing::warn!(
nsid = %nsid,
error = %e,
"lexicon refresh failed, serving stale cached entry"
source,
"lexicon refresh failed, serving cached entry"
);
Ok(stale)
Ok(doc)
}
},
None => {
@@ -230,34 +299,59 @@ impl DynamicRegistry {
F: FnOnce(Nsid) -> Fut,
Fut: std::future::Future<Output = Result<LexiconDoc, ResolveError>>,
{
if self.network_disabled.load(Ordering::Relaxed) {
return Err(ResolveError::NetworkDisabled);
if let Some(doc) = self.shared_get(nsid).await {
return Ok(doc);
}
if self.is_negative_cached(nsid) {
if let Some(remaining) = self.negative_remaining(nsid) {
return Err(ResolveError::NegativelyCached {
nsid: nsid.clone(),
ttl_secs: NEGATIVE_CACHE_TTL.as_secs(),
ttl_secs: remaining.as_secs(),
});
}
if self.shared_is_negative(nsid).await {
// Cache reports 0 remaining TTL for shared negative hit,
// so we mirror for the backoff rather than a full `NEGATIVE_CACHE_TTL`.
self.insert_negative(nsid, REFRESH_FAILURE_BACKOFF);
return Err(ResolveError::NegativelyCached {
nsid: nsid.clone(),
ttl_secs: REFRESH_FAILURE_BACKOFF.as_secs(),
});
}
if self.network_disabled.load(Ordering::Relaxed) {
return Err(ResolveError::NetworkDisabled);
}
match self.acquire_leadership(nsid) {
Some(_guard) => match resolver(nsid.clone()).await {
Ok(doc) => Ok(self.insert_schema(doc)),
Ok(doc) => {
self.shared_put(&doc).await;
Ok(self.insert_schema(doc))
}
Err(e) => {
self.insert_negative(nsid);
tracing::debug!(nsid = %nsid, error = %e, "caching negative resolution result");
let ttl = negative_ttl_for(&e);
self.insert_negative(nsid, ttl);
self.shared_put_negative(nsid, &e).await;
tracing::debug!(
nsid = %nsid,
error = %e,
ttl_secs = ttl.as_secs(),
"caching negative resolution result"
);
Err(e)
}
},
None => {
self.wait_for_leader(nsid).await;
match self.get_cached(nsid) {
Some(doc) => Ok(doc),
None if self.is_negative_cached(nsid) => Err(ResolveError::NegativelyCached {
match (self.get_cached(nsid), self.negative_remaining(nsid)) {
(Some(doc), _) => Ok(doc),
(None, Some(remaining)) => Err(ResolveError::NegativelyCached {
nsid: nsid.clone(),
ttl_secs: NEGATIVE_CACHE_TTL.as_secs(),
ttl_secs: remaining.as_secs(),
}),
None => Err(ResolveError::LeaderAborted { nsid: nsid.clone() }),
(None, None) => Err(ResolveError::LeaderAborted { nsid: nsid.clone() }),
}
}
}
@@ -316,6 +410,7 @@ impl Default for DynamicRegistry {
#[cfg(test)]
mod tests {
use super::*;
use tranquil_infra::MemoryCache;
fn nsid(s: &str) -> Nsid {
s.parse().unwrap()
@@ -324,19 +419,19 @@ mod tests {
#[test]
fn test_negative_cache() {
let registry = DynamicRegistry::new();
assert!(!registry.is_negative_cached(&nsid("com.example.test")));
assert!(!registry.is_negative_cached(&nsid("pet.nel.negative")));
registry.insert_negative(&nsid("com.example.test"));
assert!(registry.is_negative_cached(&nsid("com.example.test")));
registry.insert_negative(&nsid("pet.nel.negative"), NEGATIVE_CACHE_TTL);
assert!(registry.is_negative_cached(&nsid("pet.nel.negative")));
}
#[tokio::test]
async fn test_negative_cache_returns_appropriate_error_variant() {
let registry = DynamicRegistry::new();
registry.insert_negative(&nsid("com.example.cached"));
registry.insert_negative(&nsid("pet.nel.cached"), NEGATIVE_CACHE_TTL);
let err = registry
.resolve_and_cache(&nsid("com.example.cached"))
.resolve_and_cache(&nsid("pet.nel.cached"))
.await
.unwrap_err();
@@ -383,17 +478,17 @@ mod tests {
fn test_negative_cache_cleared_on_insert() {
let registry = DynamicRegistry::new();
registry.insert_negative(&nsid("com.example.test"));
assert!(registry.is_negative_cached(&nsid("com.example.test")));
registry.insert_negative(&nsid("pet.nel.cleared"), NEGATIVE_CACHE_TTL);
assert!(registry.is_negative_cached(&nsid("pet.nel.cleared")));
let doc = LexiconDoc {
lexicon: 1,
id: nsid("com.example.test"),
id: nsid("pet.nel.cleared"),
defs: HashMap::new(),
};
registry.insert_schema(doc);
assert!(!registry.is_negative_cached(&nsid("com.example.test")));
assert!(!registry.is_negative_cached(&nsid("pet.nel.cleared")));
}
#[test]
@@ -692,4 +787,95 @@ mod tests {
"evicted Arc should be freed when no external references remain"
);
}
#[tokio::test]
async fn test_shared_positive_hit_skips_resolver() {
let registry = DynamicRegistry::new();
let cache = Arc::new(MemoryCache::new());
registry.set_shared_cache(cache.clone());
let doc = LexiconDoc {
lexicon: 1,
id: nsid("pet.nel.sharedDoc"),
defs: HashMap::new(),
};
cache
.set(
&lexicon_doc_key(&nsid("pet.nel.sharedDoc")),
&serde_json::to_string(&doc).unwrap(),
POSITIVE_CACHE_TTL,
)
.await
.unwrap();
let resolved = registry
.resolve_and_cache_with(&nsid("pet.nel.sharedDoc"), |_| async move {
panic!("resolver mustn't run on a shared positive hit")
})
.await
.unwrap();
assert_eq!(resolved.id, "pet.nel.sharedDoc");
assert!(registry.get_cached(&nsid("pet.nel.sharedDoc")).is_some());
}
#[tokio::test]
async fn test_definitive_failure_writes_shared_negative_and_peers_mirror_it() {
let cache = Arc::new(MemoryCache::new());
let registry = DynamicRegistry::new();
registry.set_shared_cache(cache.clone());
let _ = registry
.resolve_and_cache_with(&nsid("pet.nel.gone"), |n| async move {
Err::<LexiconDoc, _>(ResolveError::SchemaNotFound {
nsid: n,
url: "https://oyster.cafe".to_string(),
})
})
.await;
assert!(
cache
.get(&lexicon_negative_key(&nsid("pet.nel.gone")))
.await
.is_some(),
"definitive failure must write the shared negative key"
);
let _ = registry
.resolve_and_cache_with(&nsid("pet.nel.transient"), |n| async move {
Err::<LexiconDoc, _>(ResolveError::DnsLookup {
domain: n.into_inner(),
reason: "simulated".to_string(),
})
})
.await;
assert!(
cache
.get(&lexicon_negative_key(&nsid("pet.nel.transient")))
.await
.is_none(),
"transient failure must stay out of the shared negative key"
);
let peer = DynamicRegistry::new();
peer.set_shared_cache(cache);
let err = peer
.resolve_and_cache_with(&nsid("pet.nel.gone"), |_| async move {
panic!("resolver mustn't run on a shared negative hit")
})
.await
.unwrap_err();
match err {
ResolveError::NegativelyCached { ttl_secs, .. } => assert!(
ttl_secs <= REFRESH_FAILURE_BACKOFF.as_secs(),
"local mirror must use the backoff TTL, got {}s",
ttl_secs
),
other => panic!("expected NegativelyCached, got: {}", other),
}
assert!(
peer.negative_remaining(&nsid("pet.nel.gone"))
.expect("local mirror exists")
<= REFRESH_FAILURE_BACKOFF
);
}
}
+5
View File
@@ -125,6 +125,11 @@ impl LexiconRegistry {
pub fn is_negative_cached(&self, nsid: &Nsid) -> bool {
self.dynamic.is_negative_cached(nsid)
}
#[cfg(feature = "resolve")]
pub fn set_shared_cache(&self, cache: Arc<dyn tranquil_infra::Cache>) {
self.dynamic.set_shared_cache(cache);
}
}
pub struct ResolvedRef {
+95 -82
View File
@@ -4,7 +4,10 @@ use hickory_resolver::config::{ResolverConfig, ResolverOpts};
use reqwest::Client;
use std::sync::OnceLock;
use std::time::Duration;
use tranquil_types::{Did, Nsid};
use tranquil_types::did_doc::extract_pds_endpoint;
use tranquil_types::{
Did, Nsid, SchemaHostUrl, UrlKind, dns_guard, redirect_policy, url_kind, url_reach_permits,
};
static RESOLVER_CLIENT: OnceLock<Client> = OnceLock::new();
@@ -17,7 +20,8 @@ fn client() -> &'static Client {
.connect_timeout(Duration::from_secs(5))
.pool_max_idle_per_host(4)
.pool_idle_timeout(Duration::from_secs(60))
.redirect(reqwest::redirect::Policy::limited(3))
.redirect(redirect_policy(url_kind::SchemaHost::REACH_POLICY))
.dns_resolver(dns_guard(url_kind::SchemaHost::REACH_POLICY))
.build()
.expect("failed to build lexicon resolver HTTP client")
})
@@ -63,6 +67,8 @@ pub enum ResolveError {
NoPdsEndpoint { did: Did },
#[error("schema fetch failed from {url}: {reason}")]
SchemaFetch { url: String, reason: String },
#[error("no schema record for {nsid} at {url}")]
SchemaNotFound { nsid: Nsid, url: String },
#[error("schema deserialization failed: {0}")]
InvalidSchema(String),
#[error("schema resolution recently failed for {nsid}, cached for {ttl_secs}s")]
@@ -73,6 +79,23 @@ pub enum ResolveError {
LeaderAborted { nsid: Nsid },
}
impl ResolveError {
pub fn is_definitive(&self) -> bool {
match self {
Self::NoDid { .. }
| Self::NoPdsEndpoint { .. }
| Self::InvalidSchema(_)
| Self::SchemaNotFound { .. } => true,
Self::DnsLookup { .. }
| Self::DidResolution { .. }
| Self::SchemaFetch { .. }
| Self::NegativelyCached { .. }
| Self::NetworkDisabled
| Self::LeaderAborted { .. } => false,
}
}
}
pub fn nsid_to_authority(nsid: &Nsid) -> String {
let mut segments: Vec<&str> = nsid.split('.').collect();
segments.pop();
@@ -123,7 +146,7 @@ pub async fn resolve_did_from_dns(authority: &str) -> Result<Did, ResolveError>
pub async fn resolve_pds_endpoint(
did: &Did,
plc_directory_url: Option<&str>,
) -> Result<String, ResolveError> {
) -> Result<SchemaHostUrl, ResolveError> {
let plc_base = plc_directory_url.unwrap_or(DEFAULT_PLC_DIRECTORY);
let url = match did
@@ -131,7 +154,20 @@ pub async fn resolve_pds_endpoint(
.and_then(|(_, rest)| rest.split_once(':'))
{
Some(("plc", _)) => format!("{}/{}", plc_base.trim_end_matches('/'), did),
Some(("web", domain)) => format!("https://{}/.well-known/did.json", domain),
Some(("web", domain)) => {
let url = format!("https://{}/.well-known/did.json", domain);
let permitted = reqwest::Url::parse(&url)
.is_ok_and(|u| url_reach_permits(&u, url_kind::SchemaHost::REACH_POLICY));
match permitted {
true => url,
false => {
return Err(ResolveError::DidResolution {
did: did.clone(),
reason: "did:web host is outside the allowed host reach".to_string(),
});
}
}
}
_ => {
return Err(ResolveError::DidResolution {
did: did.clone(),
@@ -162,39 +198,29 @@ pub async fn resolve_pds_endpoint(
reason: e.to_string(),
})?;
extract_pds_endpoint(&doc).ok_or_else(|| ResolveError::NoPdsEndpoint { did: did.clone() })
extract_pds_endpoint(&doc).map_err(|_| ResolveError::NoPdsEndpoint { did: did.clone() })
}
fn extract_pds_endpoint(doc: &serde_json::Value) -> Option<String> {
doc.get("service")
.and_then(|s| s.as_array())
.and_then(|services| {
services.iter().find_map(|svc| {
let is_pds = svc
.get("type")
.and_then(|t| t.as_str())
.is_some_and(|t| t == "AtprotoPersonalDataServer");
is_pds
.then(|| svc.get("serviceEndpoint").and_then(|ep| ep.as_str()))?
.map(|s| s.to_string())
})
})
fn is_record_absent(xrpc_error: &str, xrpc_message: &str) -> bool {
xrpc_error == "RecordNotFound"
|| xrpc_error == "InvalidRequest" && xrpc_message.starts_with("Could not locate record")
}
pub async fn fetch_schema_from_pds(
pds_endpoint: &str,
pds_endpoint: &SchemaHostUrl,
did: &Did,
nsid: &Nsid,
) -> Result<LexiconDoc, ResolveError> {
let url = format!(
"{}/xrpc/com.atproto.repo.getRecord?repo={}&collection=com.atproto.lexicon.schema&rkey={}",
pds_endpoint.trim_end_matches('/'),
urlencoding::encode(did.as_str()),
urlencoding::encode(nsid.as_str())
);
let mut request_url = pds_endpoint.endpoint("xrpc/com.atproto.repo.getRecord");
request_url
.query_pairs_mut()
.append_pair("repo", did.as_str())
.append_pair("collection", "com.atproto.lexicon.schema")
.append_pair("rkey", nsid.as_str());
let url = request_url.to_string();
let resp = client()
.get(&url)
.get(request_url)
.send()
.await
.map_err(|e| ResolveError::SchemaFetch {
@@ -204,10 +230,27 @@ pub async fn fetch_schema_from_pds(
let status = resp.status();
if !status.is_success() {
return Err(ResolveError::SchemaFetch {
url,
reason: format!("HTTP {}", status),
});
let body = read_body_limited(resp, MAX_RESPONSE_BYTES)
.await
.ok()
.and_then(|bytes| serde_json::from_slice::<serde_json::Value>(&bytes).ok())
.unwrap_or(serde_json::Value::Null);
let field = |name: &str| {
body.get(name)
.and_then(|v| v.as_str())
.unwrap_or_default()
.to_string()
};
return match is_record_absent(&field("error"), &field("message")) {
true => Err(ResolveError::SchemaNotFound {
nsid: nsid.clone(),
url,
}),
false => Err(ResolveError::SchemaFetch {
url,
reason: format!("HTTP {}", status),
}),
};
}
let body = read_body_limited(resp, MAX_RESPONSE_BYTES)
@@ -292,6 +335,27 @@ mod tests {
s.parse().unwrap()
}
#[test]
fn is_record_absent_recognizes_only_the_reference_pds_absence_shapes() {
assert!(is_record_absent(
"RecordNotFound",
"Could not locate record: at://did:plc:nel/com.atproto.lexicon.schema/x"
));
assert!(is_record_absent("RecordNotFound", ""));
assert!(is_record_absent(
"InvalidRequest",
"Could not locate record"
));
assert!(!is_record_absent(
"InvalidRequest",
"Error: rkey must be a valid record key"
));
assert!(!is_record_absent("InvalidRequest", ""));
assert!(!is_record_absent("InternalServerError", ""));
assert!(!is_record_absent("RateLimitExceeded", ""));
assert!(!is_record_absent("", ""));
}
#[test]
fn test_nsid_to_authority() {
assert_eq!(
@@ -316,57 +380,6 @@ mod tests {
);
}
#[test]
fn test_extract_pds_endpoint_valid() {
let doc = serde_json::json!({
"service": [{
"type": "AtprotoPersonalDataServer",
"serviceEndpoint": "https://pds.example.com"
}]
});
assert_eq!(
extract_pds_endpoint(&doc),
Some("https://pds.example.com".to_string())
);
}
#[test]
fn test_extract_pds_endpoint_multiple_services() {
let doc = serde_json::json!({
"service": [
{
"type": "AtprotoLabeler",
"serviceEndpoint": "https://labeler.example.com"
},
{
"type": "AtprotoPersonalDataServer",
"serviceEndpoint": "https://pds.example.com"
}
]
});
assert_eq!(
extract_pds_endpoint(&doc),
Some("https://pds.example.com".to_string())
);
}
#[test]
fn test_extract_pds_endpoint_missing() {
let doc = serde_json::json!({
"service": [{
"type": "AtprotoLabeler",
"serviceEndpoint": "https://labeler.example.com"
}]
});
assert_eq!(extract_pds_endpoint(&doc), None);
}
#[test]
fn test_extract_pds_endpoint_no_services() {
let doc = serde_json::json!({});
assert_eq!(extract_pds_endpoint(&doc), None);
}
#[test]
fn test_validate_fetched_schema_ok() {
let doc = LexiconDoc {
+15 -15
View File
@@ -1,8 +1,8 @@
use serde::Deserialize;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use tranquil_types::Nsid;
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
pub struct LexiconDoc {
pub lexicon: u32,
pub id: Nsid,
@@ -10,7 +10,7 @@ pub struct LexiconDoc {
pub defs: HashMap<String, LexDef>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum LexDef {
#[serde(rename = "record")]
@@ -35,14 +35,14 @@ pub enum LexDef {
PermissionSet {},
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
pub struct LexRecord {
#[serde(default)]
pub key: Option<String>,
pub record: LexObject,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
pub struct LexObject {
#[serde(default)]
pub required: Vec<String>,
@@ -52,7 +52,7 @@ pub struct LexObject {
pub properties: HashMap<String, LexProperty>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum LexProperty {
#[serde(rename = "string")]
@@ -79,7 +79,7 @@ pub enum LexProperty {
Object(LexObject),
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LexString {
#[serde(default)]
@@ -102,7 +102,7 @@ pub struct LexString {
pub default: Option<String>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
pub struct LexInteger {
#[serde(default)]
pub minimum: Option<i64>,
@@ -116,7 +116,7 @@ pub struct LexInteger {
pub const_value: Option<i64>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LexBytes {
#[serde(default)]
@@ -125,7 +125,7 @@ pub struct LexBytes {
pub min_length: Option<u64>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LexBlob {
#[serde(default)]
@@ -134,7 +134,7 @@ pub struct LexBlob {
pub max_size: Option<u64>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LexArray {
pub items: Box<LexProperty>,
@@ -144,7 +144,7 @@ pub struct LexArray {
pub max_length: Option<u64>,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
pub struct LexUnion {
#[serde(default)]
pub refs: Vec<String>,
@@ -152,14 +152,14 @@ pub struct LexUnion {
pub closed: bool,
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LexRef {
#[serde(rename = "ref")]
pub reference: String,
}
#[derive(Debug, Clone, Deserialize)]
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum StringFormat {
#[serde(rename = "did")]
Did,
@@ -204,6 +204,6 @@ pub fn parse_ref(reference: &str) -> ParsedRef<'_> {
}
}
#[derive(Debug, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LexStringDef {}
@@ -74,7 +74,7 @@ async fn test_resolve_pds_endpoint_from_plc() {
let endpoint = resolve_pds_endpoint(&did.parse().unwrap(), Some(&plc_server.uri()))
.await
.unwrap();
assert_eq!(endpoint, "https://pds.example.com");
assert_eq!(endpoint.as_str(), "https://pds.example.com");
}
#[tokio::test]
@@ -130,14 +130,17 @@ async fn test_resolve_pds_endpoint_multiple_services_picks_pds() {
"id": did,
"service": [
{
"id": "#atproto_labeler",
"type": "AtprotoLabeler",
"serviceEndpoint": "https://labeler.example.com"
},
{
"id": "#bsky_notif",
"type": "BskyNotificationService",
"serviceEndpoint": "https://notify.example.com"
},
{
"id": "#atproto_pds",
"type": "AtprotoPersonalDataServer",
"serviceEndpoint": "https://pds.example.com"
}
@@ -149,7 +152,7 @@ async fn test_resolve_pds_endpoint_multiple_services_picks_pds() {
let endpoint = resolve_pds_endpoint(&did.parse().unwrap(), Some(&plc_server.uri()))
.await
.unwrap();
assert_eq!(endpoint, "https://pds.example.com");
assert_eq!(endpoint.as_str(), "https://pds.example.com");
}
#[tokio::test]
@@ -168,7 +171,7 @@ async fn test_fetch_schema_from_pds_success() {
.await;
let doc = fetch_schema_from_pds(
&pds_server.uri(),
&pds_server.uri().parse().unwrap(),
&did.parse().unwrap(),
&nsid.parse().unwrap(),
)
@@ -195,7 +198,7 @@ async fn test_fetch_schema_missing_value_field() {
.await;
let result = fetch_schema_from_pds(
&pds_server.uri(),
&pds_server.uri().parse().unwrap(),
&did.parse().unwrap(),
&nsid.parse().unwrap(),
)
@@ -222,7 +225,7 @@ async fn test_fetch_schema_invalid_lexicon_json() {
.await;
let result = fetch_schema_from_pds(
&pds_server.uri(),
&pds_server.uri().parse().unwrap(),
&did.parse().unwrap(),
&nsid.parse().unwrap(),
)
@@ -352,7 +355,7 @@ async fn test_pds_trailing_slash_handled() {
let pds_url_with_slash = format!("{}/", pds_server.uri());
let doc = fetch_schema_from_pds(
&pds_url_with_slash,
&pds_url_with_slash.parse().unwrap(),
&did.parse().unwrap(),
&nsid.parse().unwrap(),
)
@@ -377,7 +380,7 @@ async fn test_fetch_schema_error_status_gives_meaningful_error() {
.await;
let result = fetch_schema_from_pds(
&pds_server.uri(),
&pds_server.uri().parse().unwrap(),
&did.parse().unwrap(),
&nsid.parse().unwrap(),
)
+1
View File
@@ -37,6 +37,7 @@ webauthn-rs = { workspace = true }
[dev-dependencies]
async-trait = { workspace = true }
tranquil-infra = { workspace = true, features = ["testing"] }
[features]
bsky = []
@@ -120,7 +120,7 @@ pub async fn consent_get(
};
let did = flow_with_user.did().clone();
let client_cache = ClientMetadataCache::new(3600);
let client_cache = &state.client_metadata_cache;
let client_metadata = client_cache
.get(&request_data.parameters.client_id)
.await
@@ -80,7 +80,7 @@ pub async fn authorize_get(
"Authorization request has expired. Please start a new request.",
);
}
let client_cache = ClientMetadataCache::new(3600);
let client_cache = &state.client_metadata_cache;
let client_name = client_cache
.get(&request_data.parameters.client_id)
.await
@@ -14,8 +14,7 @@ use tranquil_db_traits::{ScopePreference, WebauthnChallengeType};
use tranquil_pds::auth::{BareLoginIdentifier, NormalizedLoginIdentifier};
use tranquil_pds::comms::comms_repo::enqueue_2fa_code;
use tranquil_pds::oauth::{
AuthFlow, ClientMetadataCache, DeviceData, DeviceId, OAuthError, Prompt, SessionId,
db::should_show_consent,
AuthFlow, DeviceData, DeviceId, OAuthError, Prompt, SessionId, db::should_show_consent,
};
use tranquil_pds::rate_limit::{
OAuthAuthorizeLimit, OAuthRateLimited, OAuthRegisterCompleteLimit, TotpVerifyLimit,
@@ -33,36 +33,12 @@ pub async fn resolve_effective_scopes(
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use std::sync::Mutex;
use std::time::Duration;
use tranquil_pds::cache::{Cache, CacheError};
use tranquil_infra::MemoryCache;
use tranquil_pds::cache::Cache;
#[derive(Default)]
struct MapCache(Mutex<HashMap<String, String>>);
#[async_trait::async_trait]
impl Cache for MapCache {
async fn get(&self, k: &str) -> Option<String> {
self.0.lock().unwrap().get(k).cloned()
}
async fn set(&self, k: &str, v: &str, _t: Duration) -> Result<(), CacheError> {
self.0.lock().unwrap().insert(k.into(), v.into());
Ok(())
}
async fn delete(&self, k: &str) -> Result<(), CacheError> {
self.0.lock().unwrap().remove(k);
Ok(())
}
async fn get_bytes(&self, _k: &str) -> Option<Vec<u8>> {
None
}
async fn set_bytes(&self, _k: &str, _v: &[u8], _t: Duration) -> Result<(), CacheError> {
Ok(())
}
}
fn cache_with(nsid: &str, scopes: &str) -> MapCache {
let c = MapCache::default();
async fn cache_with(nsid: &str, scopes: &str) -> MemoryCache {
let c = MemoryCache::new();
let key = tranquil_pds::cache_keys::permission_set_key(
&tranquil_types::Nsid::new(nsid).unwrap(),
None,
@@ -74,7 +50,7 @@ mod tests {
"refreshed_at": chrono::Utc::now().timestamp(),
})
.to_string();
c.0.lock().unwrap().insert(key, json);
let _ = c.set(&key, &json, Duration::from_secs(3600)).await;
c
}
@@ -83,7 +59,8 @@ mod tests {
let c = cache_with(
"io.atcr.authFullApp",
"repo:io.atcr.manifest?action=create identity:*",
);
)
.await;
let eff = resolve_effective_scopes(
&c,
"atproto include:io.atcr.authFullApp",
@@ -104,7 +81,8 @@ mod tests {
let c = cache_with(
"io.atcr.authFullApp",
"repo:io.atcr.manifest?action=create identity:*",
);
)
.await;
let granted = DbScope::new("atproto repo:* blob:*/* account:*?action=manage").unwrap();
let eff = resolve_effective_scopes(
&c,
@@ -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,
@@ -3,8 +3,8 @@ use axum::{Json, extract::State, http::HeaderMap};
use chrono::{Duration, Utc};
use serde::{Deserialize, Serialize};
use tranquil_pds::oauth::{
AuthorizationRequestParameters, ClientAuth, ClientMetadataCache, CodeChallengeMethod,
OAuthError, Prompt, RequestData, RequestId, ResponseMode, ResponseType,
AuthorizationRequestParameters, ClientAuth, CodeChallengeMethod, OAuthError, Prompt,
RequestData, RequestId, ResponseMode, ResponseType,
scopes::{ParsedScope, parse_scope},
};
use tranquil_pds::rate_limit::{OAuthParLimit, OAuthRateLimited};
@@ -80,7 +80,7 @@ pub async fn pushed_authorization_request(
.ok_or_else(|| OAuthError::InvalidRequest("code_challenge is required".to_string()))?;
let code_challenge_method =
parse_code_challenge_method(request.code_challenge_method.as_deref())?;
let client_cache = ClientMetadataCache::new(3600);
let client_cache = &state.client_metadata_cache;
let client_metadata = client_cache.get(&request.client_id).await?;
client_cache.validate_redirect_uri(&client_metadata, &request.redirect_uri)?;
let client_auth = determine_client_auth(&request)?;
@@ -8,8 +8,7 @@ use chrono::{Duration, Utc};
use tranquil_db_traits::RefreshTokenLookup;
use tranquil_pds::config::AuthConfig;
use tranquil_pds::oauth::{
AuthFlow, ClientAuth, ClientMetadataCache, DPoPVerifier, OAuthError, RefreshToken, TokenData,
TokenId,
AuthFlow, ClientAuth, DPoPVerifier, OAuthError, RefreshToken, TokenData, TokenId,
db::{enforce_token_limit_for_user, lookup_refresh_token},
verify_client_auth,
};
@@ -63,7 +62,7 @@ pub async fn handle_authorization_code_grant(
return Err(OAuthError::InvalidGrant("client_id mismatch".to_string()));
}
let did = authorized.did.clone();
let client_metadata_cache = ClientMetadataCache::new(3600);
let client_metadata_cache = &state.client_metadata_cache;
let client_metadata = client_metadata_cache.get(&authorized.client_id).await?;
let client_auth = match &request.client_auth {
RequestClientAuth::PrivateKeyJwt {
@@ -85,7 +84,7 @@ pub async fn handle_authorization_code_grant(
},
RequestClientAuth::None { .. } => ClientAuth::None,
};
verify_client_auth(&client_metadata_cache, &client_metadata, &client_auth).await?;
verify_client_auth(client_metadata_cache, &client_metadata, &client_auth).await?;
verify_pkce(&authorized.parameters.code_challenge, &code_verifier)?;
if let Some(req_redirect_uri) = &redirect_uri
&& req_redirect_uri != &authorized.parameters.redirect_uri
+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 }
+105 -89
View File
@@ -1,12 +1,19 @@
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use std::time::Duration;
use crate::OAuthError;
use crate::types::ClientAuth;
use tranquil_types::ClientId;
use tranquil_infra::cache_keys::{
oauth_client_jwks_cooldown_key, oauth_client_jwks_key, oauth_client_meta_key,
};
use tranquil_infra::{Cache, cached_json, write_json};
use tranquil_types::{
ClientId, JwksUri, ReachPolicy, dns_guard, redirect_policy, url_reach_permits,
};
const JWKS_REFRESH_COOLDOWN: Duration = Duration::from_secs(60);
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClientMetadata {
@@ -30,8 +37,12 @@ pub struct ClientMetadata {
pub dpop_bound_access_tokens: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub jwks: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub jwks_uri: Option<String>,
#[serde(
default,
skip_serializing_if = "Option::is_none",
deserialize_with = "tranquil_types::http_url::deserialize_optional"
)]
pub jwks_uri: Option<JwksUri>,
#[serde(skip_serializing_if = "Option::is_none")]
pub application_type: Option<String>,
}
@@ -58,33 +69,23 @@ impl Default for ClientMetadata {
#[derive(Clone)]
pub struct ClientMetadataCache {
cache: Arc<RwLock<HashMap<String, CachedMetadata>>>,
jwks_cache: Arc<RwLock<HashMap<String, CachedJwks>>>,
cache: Arc<dyn Cache>,
http_client: Client,
cache_ttl_secs: u64,
}
struct CachedMetadata {
metadata: ClientMetadata,
cached_at: std::time::Instant,
}
struct CachedJwks {
jwks: serde_json::Value,
cached_at: std::time::Instant,
cache_ttl: Duration,
}
impl ClientMetadataCache {
pub fn new(cache_ttl_secs: u64) -> Self {
pub fn new(cache: Arc<dyn Cache>, cache_ttl: Duration) -> Self {
Self {
cache: Arc::new(RwLock::new(HashMap::new())),
jwks_cache: Arc::new(RwLock::new(HashMap::new())),
cache,
http_client: {
let builder = Client::builder()
.timeout(std::time::Duration::from_secs(30))
.connect_timeout(std::time::Duration::from_secs(10))
.pool_max_idle_per_host(10)
.pool_idle_timeout(std::time::Duration::from_secs(90))
.redirect(redirect_policy(ReachPolicy::DEBUG_LOOPBACK))
.dns_resolver(dns_guard(ReachPolicy::DEBUG_LOOPBACK))
.user_agent(concat!(
"Tranquil-PDS/",
env!("CARGO_PKG_VERSION"),
@@ -92,9 +93,11 @@ impl ClientMetadataCache {
));
#[cfg(feature = "native-tls-roots")]
let builder = builder.danger_accept_invalid_certs(true);
builder.build().unwrap_or_else(|_| Client::new())
builder
.build()
.expect("failed to build client metadata HTTP client")
},
cache_ttl_secs,
cache_ttl,
}
}
@@ -150,26 +153,13 @@ impl ClientMetadataCache {
if Self::is_loopback_client(client_id) {
return Self::build_loopback_metadata(client_id);
}
{
let cache = self.cache.read().await;
if let Some(cached) = cache.get(client_id.as_str())
&& cached.cached_at.elapsed().as_secs() < self.cache_ttl_secs
{
return Ok(cached.metadata.clone());
}
}
let metadata = self.fetch_metadata(client_id).await?;
{
let mut cache = self.cache.write().await;
cache.insert(
client_id.to_string(),
CachedMetadata {
metadata: metadata.clone(),
cached_at: std::time::Instant::now(),
},
);
}
Ok(metadata)
cached_json(
self.cache.as_ref(),
&oauth_client_meta_key(client_id),
self.cache_ttl,
|| self.fetch_metadata(client_id),
)
.await
}
pub async fn get_jwks(
@@ -181,43 +171,57 @@ impl ClientMetadataCache {
}
let jwks_uri = metadata.jwks_uri.as_ref().ok_or_else(|| {
OAuthError::InvalidClient(
"Client using private_key_jwt must have jwks or jwks_uri".to_string(),
"Client using private_key_jwt must have jwks or a usable jwks_uri".to_string(),
)
})?;
{
let cache = self.jwks_cache.read().await;
if let Some(cached) = cache.get(jwks_uri)
&& cached.cached_at.elapsed().as_secs() < self.cache_ttl_secs
{
return Ok(cached.jwks.clone());
cached_json(
self.cache.as_ref(),
&oauth_client_jwks_key(jwks_uri),
self.cache_ttl,
|| self.fetch_jwks(jwks_uri),
)
.await
}
async fn refresh_jwks(
&self,
metadata: &ClientMetadata,
) -> Result<Option<serde_json::Value>, OAuthError> {
match (&metadata.jwks, &metadata.jwks_uri) {
(None, Some(jwks_uri)) => {
let cooldown_key = oauth_client_jwks_cooldown_key(jwks_uri);
if self.cache.get(&cooldown_key).await.is_some() {
return Ok(None);
}
let _ = self
.cache
.set(&cooldown_key, "1", JWKS_REFRESH_COOLDOWN)
.await;
self.fetch_and_store_jwks(jwks_uri).await.map(Some)
}
_ => Ok(None),
}
}
async fn fetch_and_store_jwks(
&self,
jwks_uri: &JwksUri,
) -> Result<serde_json::Value, OAuthError> {
let jwks = self.fetch_jwks(jwks_uri).await?;
{
let mut cache = self.jwks_cache.write().await;
cache.insert(
jwks_uri.clone(),
CachedJwks {
jwks: jwks.clone(),
cached_at: std::time::Instant::now(),
},
);
}
write_json(
self.cache.as_ref(),
&oauth_client_jwks_key(jwks_uri),
&jwks,
self.cache_ttl,
)
.await;
Ok(jwks)
}
async fn fetch_jwks(&self, jwks_uri: &str) -> Result<serde_json::Value, OAuthError> {
if !jwks_uri.starts_with("https://")
&& (!jwks_uri.starts_with("http://")
|| (!jwks_uri.contains("localhost") && !jwks_uri.contains("127.0.0.1")))
{
return Err(OAuthError::InvalidClient(
"jwks_uri must use https (except for localhost)".to_string(),
));
}
async fn fetch_jwks(&self, jwks_uri: &JwksUri) -> Result<serde_json::Value, OAuthError> {
let response = self
.http_client
.get(jwks_uri)
.get(jwks_uri.as_str())
.header("Accept", "application/json")
.send()
.await
@@ -243,22 +247,16 @@ impl ClientMetadataCache {
}
async fn fetch_metadata(&self, client_id: &ClientId) -> Result<ClientMetadata, OAuthError> {
if !client_id.starts_with("http://") && !client_id.starts_with("https://") {
let url = reqwest::Url::parse(client_id)
.map_err(|_| OAuthError::InvalidClient("client_id must be a URL".to_string()))?;
if !url_reach_permits(&url, ReachPolicy::DEBUG_LOOPBACK) {
return Err(OAuthError::InvalidClient(
"client_id must be a URL".to_string(),
));
}
if client_id.starts_with("http://")
&& !client_id.contains("localhost")
&& !client_id.contains("127.0.0.1")
{
return Err(OAuthError::InvalidClient(
"Non-localhost client_id must use https".to_string(),
"client_id must be an https URL inside the allowed host reach".to_string(),
));
}
let response = self
.http_client
.get(client_id.as_str())
.get(url)
.header("Accept", "application/json")
.send()
.await
@@ -514,7 +512,29 @@ async fn verify_private_key_jwt_async(
"client_assertion iat is in the future".to_string(),
));
}
let signing_input = format!("{}.{}", parts[0], parts[1]);
let signature_bytes = URL_SAFE_NO_PAD
.decode(parts[2])
.map_err(|_| OAuthError::InvalidClient("Invalid signature encoding".to_string()))?;
let jwks = cache.get_jwks(metadata).await?;
match verify_assertion_signature(&jwks, kid, alg, &signing_input, &signature_bytes) {
Ok(()) => Ok(()),
Err(cached_failure) => match cache.refresh_jwks(metadata).await {
Ok(Some(fresh)) => {
verify_assertion_signature(&fresh, kid, alg, &signing_input, &signature_bytes)
}
Ok(None) | Err(_) => Err(cached_failure),
},
}
}
fn verify_assertion_signature(
jwks: &serde_json::Value,
kid: Option<&str>,
alg: &str,
signing_input: &str,
signature: &[u8],
) -> Result<(), OAuthError> {
let keys = jwks
.get("keys")
.and_then(|k| k.as_array())
@@ -531,10 +551,6 @@ async fn verify_private_key_jwt_async(
"No matching key found in client JWKS".to_string(),
));
}
let signing_input = format!("{}.{}", parts[0], parts[1]);
let signature_bytes = URL_SAFE_NO_PAD
.decode(parts[2])
.map_err(|_| OAuthError::InvalidClient("Invalid signature encoding".to_string()))?;
matching_keys
.into_iter()
.filter(|key| {
@@ -544,12 +560,12 @@ async fn verify_private_key_jwt_async(
.find_map(|key| {
let kty = key.get("kty").and_then(|k| k.as_str()).unwrap_or("");
match (alg, kty) {
("ES256", "EC") => verify_es256(key, &signing_input, &signature_bytes).ok(),
("ES384", "EC") => verify_es384(key, &signing_input, &signature_bytes).ok(),
("ES256", "EC") => verify_es256(key, signing_input, signature).ok(),
("ES384", "EC") => verify_es384(key, signing_input, signature).ok(),
("RS256" | "RS384" | "RS512", "RSA") => {
verify_rsa(alg, key, &signing_input, &signature_bytes).ok()
verify_rsa(alg, key, signing_input, signature).ok()
}
("EdDSA", "OKP") => verify_eddsa(key, &signing_input, &signature_bytes).ok(),
("EdDSA", "OKP") => verify_eddsa(key, signing_input, signature).ok(),
_ => None,
}
})
+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 }
+22 -20
View File
@@ -335,27 +335,29 @@ async fn proxy_handler(
};
// BSKY: getFeed must be audienced to the feed generator, not the AppView.
let (token_aud, token_lxm) =
if cfg!(feature = "bsky-support") && method == "app.bsky.feed.getFeed" {
match resolve_feed_generator_did(&resolved.url, query.as_deref()).await {
Some(feed_did) => (
feed_did,
"app.bsky.feed.getFeedSkeleton"
.parse::<Nsid>()
.expect("getFeedSkeleton is a valid NSID"),
),
None => {
warn!(
"getFeed proxy: could not resolve feed generator DID; refusing \
to mint an AppView-audienced token"
);
return ApiError::InvalidRequest("Could not resolve feed".into())
.into_response();
}
#[cfg(feature = "bsky-support")]
let (token_aud, token_lxm) = if method == "app.bsky.feed.getFeed" {
match resolve_feed_generator_did(&resolved.url, query.as_deref()).await {
Some(feed_did) => (
feed_did,
"app.bsky.feed.getFeedSkeleton"
.parse::<Nsid>()
.expect("getFeedSkeleton is a valid NSID"),
),
None => {
warn!(
"getFeed proxy refuses to mint an AppView-audienced token \
because feed generator DID resolution failed"
);
return ApiError::InvalidRequest("Couldn't resolve feed".into())
.into_response();
}
} else {
(resolved.did.clone(), method_nsid.clone())
};
}
} else {
(resolved.did.clone(), method_nsid.clone())
};
#[cfg(not(feature = "bsky-support"))]
let (token_aud, token_lxm) = (resolved.did.clone(), method_nsid.clone());
match crate::auth::create_service_token(
&auth_user.did,
+14 -92
View File
@@ -2,32 +2,14 @@ use serde::{Deserialize, Serialize};
use std::time::Duration;
use crate::cache::Cache;
use crate::cache_keys::email_token_key;
use crate::types::Did;
use crate::util::{generate_token_code, normalize_token_code};
pub use tranquil_types::EmailTokenPurpose;
const TOKEN_TTL_SECS: u64 = 900;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EmailTokenPurpose {
UpdateEmail,
ConfirmEmail,
DeleteAccount,
ResetPassword,
PlcOperation,
}
impl EmailTokenPurpose {
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",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TokenData {
token: String,
@@ -42,10 +24,6 @@ pub enum TokenError {
ExpiredToken,
}
fn cache_key(did: &Did, purpose: EmailTokenPurpose) -> String {
format!("email_token:{}:{}", purpose.as_str(), did)
}
fn current_timestamp() -> u64 {
u64::try_from(chrono::Utc::now().timestamp()).unwrap_or(0)
}
@@ -69,7 +47,7 @@ pub async fn create_email_token(
cache
.set(
&cache_key(did, purpose),
&email_token_key(did, purpose),
&json,
Duration::from_secs(TOKEN_TTL_SECS),
)
@@ -89,7 +67,7 @@ pub async fn validate_email_token(
return Err(TokenError::CacheUnavailable);
}
let key = cache_key(did, purpose);
let key = email_token_key(did, purpose);
let json = cache.get(&key).await.ok_or(TokenError::InvalidToken)?;
let data: TokenData = serde_json::from_str(&json).map_err(|_| TokenError::InvalidToken)?;
@@ -112,7 +90,7 @@ pub async fn validate_email_token(
}
pub async fn delete_email_token(cache: &dyn Cache, did: &Did, purpose: EmailTokenPurpose) {
let _ = cache.delete(&cache_key(did, purpose)).await;
let _ = cache.delete(&email_token_key(did, purpose)).await;
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
@@ -128,67 +106,11 @@ fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
#[cfg(test)]
mod tests {
use super::*;
use crate::cache::CacheError;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Mutex;
struct MockCache {
data: Mutex<HashMap<String, (String, u64)>>,
}
impl MockCache {
fn new() -> Self {
Self {
data: Mutex::new(HashMap::new()),
}
}
}
#[async_trait]
impl Cache for MockCache {
async fn get(&self, key: &str) -> Option<String> {
let data = self.data.lock().unwrap();
let now = current_timestamp();
data.get(key)
.filter(|(_, exp)| *exp > now)
.map(|(v, _)| v.clone())
}
async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> {
let mut data = self.data.lock().unwrap();
let expires = current_timestamp() + ttl.as_secs();
data.insert(key.to_string(), (value.to_string(), expires));
Ok(())
}
async fn delete(&self, key: &str) -> Result<(), CacheError> {
let mut data = self.data.lock().unwrap();
data.remove(key);
Ok(())
}
async fn get_bytes(&self, _key: &str) -> Option<Vec<u8>> {
None
}
async fn set_bytes(
&self,
_key: &str,
_value: &[u8],
_ttl: Duration,
) -> Result<(), CacheError> {
Ok(())
}
fn is_available(&self) -> bool {
true
}
}
use tranquil_infra::MemoryCache;
#[tokio::test]
async fn test_create_and_validate_token() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail)
@@ -205,7 +127,7 @@ mod tests {
#[tokio::test]
async fn test_token_consumed_after_use() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail)
@@ -223,7 +145,7 @@ mod tests {
#[tokio::test]
async fn test_invalid_token_rejected() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
let _token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail)
@@ -237,7 +159,7 @@ mod tests {
#[tokio::test]
async fn test_wrong_purpose_rejected() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail)
@@ -252,7 +174,7 @@ mod tests {
#[tokio::test]
async fn test_token_format() {
// The emitted token is the display form: uppercase `XXXXX-XXXXX`.
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
(0..50).for_each(|_| {
let token = futures::executor::block_on(create_email_token(
@@ -269,7 +191,7 @@ mod tests {
#[tokio::test]
async fn test_case_insensitive_validation() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail)
@@ -284,7 +206,7 @@ mod tests {
#[tokio::test]
async fn test_hyphen_insensitive_validation() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:teq").expect("valid DID");
let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail)
+28 -91
View File
@@ -3,6 +3,7 @@ use serde::{Deserialize, Serialize};
use std::time::Duration;
use crate::cache::Cache;
use crate::cache_keys::{legacy_2fa_challenge_key, legacy_2fa_cooldown_key};
use crate::types::Did;
use crate::util::{generate_token_code, normalize_token_code};
@@ -58,8 +59,8 @@ pub async fn create_challenge(
}
pub async fn clear_challenge(cache: &dyn Cache, did: &Did) {
let _ = cache.delete(&challenge_key(did)).await;
let _ = cache.delete(&cooldown_key(did)).await;
let _ = cache.delete(&legacy_2fa_challenge_key(did)).await;
let _ = cache.delete(&legacy_2fa_cooldown_key(did)).await;
}
async fn validate_challenge_internal(
@@ -71,7 +72,7 @@ async fn validate_challenge_internal(
return Err(ValidationError::CacheUnavailable);
}
let challenge_k = challenge_key(did);
let challenge_k = legacy_2fa_challenge_key(did);
let json = cache
.get(&challenge_k)
@@ -114,19 +115,11 @@ async fn validate_challenge_internal(
}
let _ = cache.delete(&challenge_k).await;
let _ = cache.delete(&cooldown_key(did)).await;
let _ = cache.delete(&legacy_2fa_cooldown_key(did)).await;
Ok(())
}
fn challenge_key(did: &Did) -> String {
format!("legacy_2fa:{}", did)
}
fn cooldown_key(did: &Did) -> String {
format!("legacy_2fa_cooldown:{}", did)
}
fn current_timestamp() -> u64 {
u64::try_from(Utc::now().timestamp()).unwrap_or(0)
}
@@ -226,7 +219,7 @@ async fn create_challenge_code(
return Err(ChallengeError::CacheUnavailable);
}
let cooldown = cooldown_key(did);
let cooldown = legacy_2fa_cooldown_key(did);
if cache.get(&cooldown).await.is_some() {
return Err(ChallengeError::RateLimited);
}
@@ -244,7 +237,7 @@ async fn create_challenge_code(
cache
.set(
&challenge_key(did),
&legacy_2fa_challenge_key(did),
&json,
Duration::from_secs(CHALLENGE_TTL_SECS),
)
@@ -280,67 +273,11 @@ impl From<ValidationError> for Legacy2faFlowError {
#[cfg(test)]
mod tests {
use super::*;
use crate::cache::CacheError;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Mutex;
struct MockCache {
data: Mutex<HashMap<String, (String, u64)>>,
}
impl MockCache {
fn new() -> Self {
Self {
data: Mutex::new(HashMap::new()),
}
}
}
#[async_trait]
impl Cache for MockCache {
async fn get(&self, key: &str) -> Option<String> {
let data = self.data.lock().unwrap();
let now = current_timestamp();
data.get(key)
.filter(|(_, exp)| *exp > now)
.map(|(v, _)| v.clone())
}
async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> {
let mut data = self.data.lock().unwrap();
let expires = current_timestamp() + ttl.as_secs();
data.insert(key.to_string(), (value.to_string(), expires));
Ok(())
}
async fn delete(&self, key: &str) -> Result<(), CacheError> {
let mut data = self.data.lock().unwrap();
data.remove(key);
Ok(())
}
async fn get_bytes(&self, _key: &str) -> Option<Vec<u8>> {
None
}
async fn set_bytes(
&self,
_key: &str,
_value: &[u8],
_ttl: Duration,
) -> Result<(), CacheError> {
Ok(())
}
fn is_available(&self) -> bool {
true
}
}
use tranquil_infra::MemoryCache;
#[tokio::test]
async fn test_create_and_validate_challenge() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let code = create_challenge(&cache, &did).await.unwrap();
@@ -352,7 +289,7 @@ mod tests {
#[tokio::test]
async fn test_challenge_code_format() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let code = create_challenge(&cache, &did).await.unwrap();
@@ -364,7 +301,7 @@ mod tests {
#[tokio::test]
async fn test_case_insensitive_validation() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let code = create_challenge(&cache, &did).await.unwrap();
@@ -375,7 +312,7 @@ mod tests {
#[tokio::test]
async fn test_hyphen_insensitive_validation() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let code = create_challenge(&cache, &did).await.unwrap();
@@ -386,7 +323,7 @@ mod tests {
#[tokio::test]
async fn test_invalid_code_rejected() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let _code = create_challenge(&cache, &did).await.unwrap();
@@ -396,7 +333,7 @@ mod tests {
#[tokio::test]
async fn test_challenge_consumed_on_success() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let code = create_challenge(&cache, &did).await.unwrap();
@@ -410,7 +347,7 @@ mod tests {
#[tokio::test]
async fn test_max_attempts_exceeded() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let _code = create_challenge(&cache, &did).await.unwrap();
@@ -425,7 +362,7 @@ mod tests {
#[tokio::test]
async fn test_rate_limiting() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let _first = create_challenge(&cache, &did).await.unwrap();
@@ -453,7 +390,7 @@ mod tests {
#[tokio::test]
async fn test_process_flow_not_required() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -470,7 +407,7 @@ mod tests {
#[tokio::test]
async fn test_process_flow_not_required_because_app_password() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: true,
@@ -487,7 +424,7 @@ mod tests {
#[tokio::test]
async fn test_process_flow_blocked() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -504,7 +441,7 @@ mod tests {
#[tokio::test]
async fn test_process_flow_challenge_sent_totp() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -521,7 +458,7 @@ mod tests {
#[tokio::test]
async fn test_process_flow_challenge_sent_email_2fa_enabled() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test2".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -538,7 +475,7 @@ mod tests {
#[tokio::test]
async fn test_process_flow_verified() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -557,7 +494,7 @@ mod tests {
#[tokio::test]
async fn test_attempts_persist_across_failures() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:test123".to_string()).unwrap();
let code = create_challenge(&cache, &did).await.unwrap();
@@ -590,7 +527,7 @@ mod tests {
#[tokio::test]
async fn test_totp_shaped_token_accepted_via_verifier() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:totp1".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -607,7 +544,7 @@ mod tests {
#[tokio::test]
async fn test_totp_shaped_token_rejected_does_not_touch_email_challenge() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:totp2".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -641,7 +578,7 @@ mod tests {
#[tokio::test]
async fn test_email_shaped_token_routes_to_email_path_when_totp_present() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:totp3".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -662,7 +599,7 @@ mod tests {
#[tokio::test]
async fn test_backup_code_shaped_token_routes_to_verifier() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:totp4".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
@@ -681,7 +618,7 @@ mod tests {
#[tokio::test]
async fn test_totp_shaped_token_ignored_when_no_totp() {
let cache = MockCache::new();
let cache = MemoryCache::new();
let did = Did::new("did:plc:totp5".to_string()).unwrap();
let ctx = Legacy2faContext {
is_app_password: false,
+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,
})
+69 -197
View File
@@ -1,10 +1,9 @@
use crate::cache::Cache;
use crate::types::Did;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
use std::time::Duration;
use tracing::{debug, info, warn};
#[derive(Debug, thiserror::Error)]
@@ -13,6 +12,8 @@ pub enum DidResolutionError {
UnsupportedDidMethod(String),
#[error("Invalid did:web format")]
InvalidDidWeb,
#[error("did:web host {0} is outside the allowed host reach")]
DidWebHostRejected(String),
#[error("HTTP request failed: {0}")]
HttpFailed(String),
#[error("Invalid DID document: {0}")]
@@ -53,43 +54,50 @@ pub struct DidService {
pub struct ResolvedService {
pub url: String,
pub did: Did,
pub service_id: String,
}
type TimedCache<T> = RwLock<HashMap<Box<str>, (Instant, Arc<T>)>>;
pub struct DidResolver {
did_doc_cache: TimedCache<serde_json::Value>,
parsed_did_doc_cache: TimedCache<DidDocument>,
service_cache: TimedCache<ResolvedService>,
cache: Arc<dyn Cache>,
client: Client,
cache_ttl: Duration,
plc_directory_url: String,
}
impl DidResolver {
pub fn new() -> Self {
pub fn new(cache: Arc<dyn Cache>) -> Self {
let cfg = tranquil_config::get();
let cache_ttl_secs = cfg.plc.did_cache_ttl_secs;
let plc_directory_url = cfg.plc.directory_url.clone();
let client = Client::builder()
.timeout(Duration::from_secs(10))
.connect_timeout(Duration::from_secs(5))
.pool_max_idle_per_host(10)
.redirect(tranquil_types::redirect_policy(
tranquil_types::ReachPolicy::DEBUG_LOOPBACK,
))
.dns_resolver(tranquil_types::dns_guard(
tranquil_types::ReachPolicy::DEBUG_LOOPBACK,
))
.build()
.unwrap_or_else(|_| Client::new());
.expect("failed to build DID resolver HTTP client");
info!("DID resolver initialized");
Self {
did_doc_cache: RwLock::new(HashMap::new()),
parsed_did_doc_cache: RwLock::new(HashMap::new()),
service_cache: RwLock::new(HashMap::new()),
cache,
client,
cache_ttl: Duration::from_secs(cache_ttl_secs),
plc_directory_url,
cache_ttl: Duration::from_secs(cfg.plc.did_cache_ttl_secs),
plc_directory_url: cfg.plc.directory_url.clone(),
}
}
fn doc_cache_key(did: &Did) -> Result<String, DidResolutionError> {
match (did.is_plc(), did.is_web()) {
(true, _) => Ok(crate::cache_keys::plc_doc_key(did)),
(_, true) => Ok(crate::cache_keys::did_web_doc_key(did)),
_ => {
warn!("Unsupported DID method: {}", did);
Err(DidResolutionError::UnsupportedDidMethod(did.to_string()))
}
}
}
@@ -97,175 +105,50 @@ impl DidResolver {
&self,
did: &Did,
service_id: &str,
) -> Result<Arc<ResolvedService>, ServiceResolutionError> {
{
let cache = self.service_cache.read().await;
if let Some(cached) = cache.get(&*format!("{did}#{service_id}"))
&& cached.0.elapsed() < self.cache_ttl
{
return Ok(cached.1.clone());
}
}
) -> Result<ResolvedService, ServiceResolutionError> {
let did_doc = self.resolve_did(did).await?;
let Some(service) = did_doc
let suffix = format!("#{service_id}");
did_doc
.services
.iter()
.find(|s| s.id.ends_with(&format!("#{service_id}")))
else {
return Err(ServiceResolutionError::ServiceIdNotFound(service_id.into()));
};
let resolved = Arc::new(ResolvedService {
url: service.service_endpoint.clone(),
did: did.clone(),
service_id: service_id.into(),
});
{
let mut cache = self.service_cache.write().await;
cache.insert(
format!("{did}#{service_id}").into(),
(Instant::now(), resolved.clone()),
);
}
Ok(resolved)
.find(|s| s.id.ends_with(&suffix))
.map(|service| ResolvedService {
url: service.service_endpoint.clone(),
did: did.clone(),
})
.ok_or_else(|| ServiceResolutionError::ServiceIdNotFound(service_id.into()))
}
pub async fn resolve_did(&self, did: &Did) -> Result<Arc<DidDocument>, DidResolutionError> {
{
let cache = self.parsed_did_doc_cache.read().await;
if let Some(cached) = cache.get(did.as_str())
&& cached.0.elapsed() < self.cache_ttl
{
return Ok(cached.1.clone());
}
}
let resolved = Arc::new(self.resolve_did_uncached(did).await?);
{
let mut cache = self.parsed_did_doc_cache.write().await;
cache.insert(did.as_str().into(), (Instant::now(), resolved.clone()));
}
Ok(resolved)
pub async fn resolve_did(&self, did: &Did) -> Result<DidDocument, DidResolutionError> {
self.cached_did_document(did).await
}
pub async fn refresh_did(&self, did: &Did) -> Result<Arc<DidDocument>, DidResolutionError> {
{
let mut cache = self.parsed_did_doc_cache.write().await;
cache.remove(did.as_str());
let mut cache = self.service_cache.write().await;
cache.retain(|k, _| !k.starts_with(did.as_str()));
}
pub async fn refresh_did(&self, did: &Did) -> Result<DidDocument, DidResolutionError> {
let _ = self.cache.delete(&Self::doc_cache_key(did)?).await;
self.resolve_did(did).await
}
async fn resolve_did_uncached(&self, did: &Did) -> Result<DidDocument, DidResolutionError> {
if did.is_web() {
self.resolve_did_web(did).await
} else if did.is_plc() {
self.resolve_did_plc(did).await
} else {
warn!("Unsupported DID method: {}", did);
Err(DidResolutionError::UnsupportedDidMethod(did.to_string()))
}
}
async fn resolve_did_web(&self, did: &Did) -> Result<DidDocument, DidResolutionError> {
let url = build_did_web_url(did)?;
debug!("Resolving did:web {} via {}", did, url);
let resp = self
.client
.get(&url)
.send()
.await
.map_err(|e| DidResolutionError::HttpFailed(e.to_string()))?;
if !resp.status().is_success() {
return Err(DidResolutionError::HttpFailed(format!(
"HTTP {}",
resp.status()
)));
}
resp.json::<DidDocument>()
.await
.map_err(|e| DidResolutionError::InvalidDocument(e.to_string()))
}
async fn resolve_did_plc(&self, did: &Did) -> Result<DidDocument, DidResolutionError> {
let url = format!(
"{}/{}",
self.plc_directory_url,
urlencoding::encode(did.as_str())
);
debug!("Resolving did:plc {} via {}", did, url);
let resp = self
.client
.get(&url)
.send()
.await
.map_err(|e| DidResolutionError::HttpFailed(e.to_string()))?;
if resp.status() == reqwest::StatusCode::NOT_FOUND {
return Err(DidResolutionError::NotFound);
}
if !resp.status().is_success() {
return Err(DidResolutionError::HttpFailed(format!(
"HTTP {}",
resp.status()
)));
}
resp.json::<DidDocument>()
.await
.map_err(|e| DidResolutionError::InvalidDocument(e.to_string()))
}
pub async fn fetch_did_document(
&self,
did: &Did,
) -> Result<Arc<serde_json::Value>, DidResolutionError> {
{
let cache = self.did_doc_cache.read().await;
if let Some(cached) = cache.get(did.as_str())
&& cached.0.elapsed() < self.cache_ttl
{
return Ok(cached.1.clone());
}
}
let resolved = Arc::new(self.fetch_did_document_uncached(did).await?);
{
let mut cache = self.did_doc_cache.write().await;
cache.insert(did.as_str().into(), (Instant::now(), resolved.clone()));
}
Ok(resolved)
) -> Result<serde_json::Value, DidResolutionError> {
self.cached_did_document(did).await
}
// TODO: make cached version
async fn fetch_did_document_uncached(
async fn cached_did_document<T: serde::de::DeserializeOwned>(
&self,
did: &Did,
) -> Result<serde_json::Value, DidResolutionError> {
if did.is_web() {
self.fetch_did_document_web(did).await
} else if did.is_plc() {
self.fetch_did_document_plc(did).await
} else {
warn!("Unsupported DID method: {}", did);
Err(DidResolutionError::UnsupportedDidMethod(did.to_string()))
}
) -> Result<T, DidResolutionError> {
let cache_key = Self::doc_cache_key(did)?;
let doc =
crate::cache::cached_json(self.cache.as_ref(), &cache_key, self.cache_ttl, || async {
match did.is_plc() {
true => self.fetch_did_document_plc(did).await,
false => self.fetch_did_document_web(did).await,
}
})
.await?;
serde_json::from_value(doc).map_err(|e| DidResolutionError::InvalidDocument(e.to_string()))
}
async fn fetch_did_document_web(
@@ -274,6 +157,8 @@ impl DidResolver {
) -> Result<serde_json::Value, DidResolutionError> {
let url = build_did_web_url(did)?;
debug!("Resolving did:web {} via {}", did, url);
let resp = self
.client
.get(&url)
@@ -303,6 +188,8 @@ impl DidResolver {
urlencoding::encode(did.as_str())
);
debug!("Resolving did:plc {} via {}", did, url);
let resp = self
.client
.get(&url)
@@ -325,21 +212,6 @@ impl DidResolver {
.await
.map_err(|e| DidResolutionError::InvalidDocument(e.to_string()))
}
pub async fn invalidate_cache(&self, did: &Did) {
let mut doc_cache = self.parsed_did_doc_cache.write().await;
doc_cache.remove(did.as_str());
}
}
impl Default for DidResolver {
fn default() -> Self {
Self::new()
}
}
pub fn create_did_resolver() -> Arc<DidResolver> {
Arc::new(DidResolver::new())
}
fn build_did_web_url(did: &Did) -> Result<String, DidResolutionError> {
@@ -372,18 +244,18 @@ fn build_did_web_url(did: &Did) -> Result<String, DidResolutionError> {
}
};
let scheme =
if host.starts_with("localhost") || host.starts_with("127.0.0.1") || host.contains(':') {
"http"
} else {
"https"
};
let url = if path.is_empty() {
format!("{}://{}/.well-known/did.json", scheme, host)
let https = if path.is_empty() {
format!("https://{}/.well-known/did.json", host)
} else {
format!("{}://{}{}/did.json", scheme, host, path)
format!("https://{}{}/did.json", host, path)
};
Ok(url)
let mut url = reqwest::Url::parse(&https).map_err(|_| DidResolutionError::InvalidDidWeb)?;
if tranquil_types::url_reach(&url) == Some(tranquil_types::HostReach::Loopback) {
let _ = url.set_scheme("http");
}
match tranquil_types::url_reach_permits(&url, tranquil_types::ReachPolicy::DEBUG_LOOPBACK) {
true => Ok(url.to_string()),
false => Err(DidResolutionError::DidWebHostRejected(host)),
}
}
+15 -12
View File
@@ -35,7 +35,7 @@ use serde_json::json;
use state::AppState;
use tower::ServiceBuilder;
use tower_http::{
cors::{Any, CorsLayer},
cors::{AllowHeaders, Any, CorsLayer},
services::{ServeDir, ServeFile},
};
pub use tranquil_db_traits::AccountStatus;
@@ -106,17 +106,20 @@ pub fn app_with_routes(state: AppState, external: ExternalRoutes) -> Router {
CorsLayer::new()
.allow_origin(Any)
.allow_methods([Method::GET, Method::POST, Method::OPTIONS])
.allow_headers([
http::header::AUTHORIZATION,
http::header::CONTENT_TYPE,
http::header::CONTENT_ENCODING,
http::header::ACCEPT_ENCODING,
http::header::USER_AGENT,
util::HEADER_DPOP,
util::HEADER_ATPROTO_PROXY,
util::HEADER_ATPROTO_ACCEPT_LABELERS,
util::HEADER_X_BSKY_TOPICS,
])
.allow_headers(AllowHeaders::list(
[
http::header::AUTHORIZATION,
http::header::CONTENT_TYPE,
http::header::CONTENT_ENCODING,
http::header::ACCEPT_ENCODING,
http::header::USER_AGENT,
util::HEADER_DPOP,
util::HEADER_ATPROTO_PROXY,
util::HEADER_ATPROTO_ACCEPT_LABELERS,
]
.into_iter()
.chain(util::CORS_BSKY_ALLOW_HEADERS),
))
.expose_headers([
http::header::WWW_AUTHENTICATE,
util::HEADER_DPOP_NONCE,
+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)?;
@@ -137,39 +137,12 @@ fn map_err(e: &ScopeExpansionError) -> ResolveFailure {
#[cfg(test)]
mod tests {
use super::*;
use crate::cache::{Cache, CacheError};
use std::collections::HashMap;
use std::sync::Mutex;
use std::time::Duration;
use tranquil_infra::MemoryCache;
#[derive(Default)]
struct MapCache(Mutex<HashMap<String, String>>);
const SEED_TTL: Duration = Duration::from_secs(3600);
#[async_trait::async_trait]
impl Cache for MapCache {
async fn get(&self, key: &str) -> Option<String> {
self.0.lock().unwrap().get(key).cloned()
}
async fn set(&self, key: &str, value: &str, _ttl: Duration) -> Result<(), CacheError> {
self.0
.lock()
.unwrap()
.insert(key.to_string(), value.to_string());
Ok(())
}
async fn delete(&self, key: &str) -> Result<(), CacheError> {
self.0.lock().unwrap().remove(key);
Ok(())
}
async fn get_bytes(&self, _key: &str) -> Option<Vec<u8>> {
None
}
async fn set_bytes(&self, _k: &str, _v: &[u8], _t: Duration) -> Result<(), CacheError> {
Ok(())
}
}
fn seed_at(cache: &MapCache, nsid: &str, scope: &str, refreshed_at: i64) {
async fn seed_at(cache: &MemoryCache, nsid: &str, scope: &str, refreshed_at: i64) {
let key =
crate::cache_keys::permission_set_key(&tranquil_types::Nsid::new(nsid).unwrap(), None);
let val = serde_json::to_string(&CachedPermissionSet {
@@ -179,21 +152,22 @@ mod tests {
refreshed_at,
})
.unwrap();
cache.0.lock().unwrap().insert(key, val);
let _ = cache.set(&key, &val, SEED_TTL).await;
}
fn seed(cache: &MapCache, nsid: &str, scope: &str) {
seed_at(cache, nsid, scope, now_secs());
async fn seed(cache: &MemoryCache, nsid: &str, scope: &str) {
seed_at(cache, nsid, scope, now_secs()).await;
}
#[tokio::test]
async fn cache_hit_expands_without_network() {
let cache = MapCache::default();
let cache = MemoryCache::new();
seed(
&cache,
"io.atcr.authFullApp",
"repo:io.atcr.manifest?action=create identity:*",
);
)
.await;
let out = expand_scopes(&cache, "atproto include:io.atcr.authFullApp").await;
assert!(out.failures.is_empty());
assert_eq!(out.passthrough, vec!["atproto".to_string()]);
@@ -208,13 +182,14 @@ mod tests {
#[tokio::test]
async fn stale_entry_is_served_when_refresh_fails() {
let cache = MapCache::default();
let cache = MemoryCache::new();
seed_at(
&cache,
"nonexistent.fake.permissionSet",
"repo:nonexistent.fake.record?action=create",
now_secs() - STALE_AFTER_SECS - 1,
);
)
.await;
let out = expand_scopes(&cache, "include:nonexistent.fake.permissionSet").await;
assert!(
out.failures.is_empty(),
@@ -230,7 +205,7 @@ mod tests {
#[tokio::test]
async fn entry_without_refreshed_at_is_treated_as_stale_but_usable() {
let cache = MapCache::default();
let cache = MemoryCache::new();
let key = crate::cache_keys::permission_set_key(
&tranquil_types::Nsid::new("nonexistent.fake.permissionSet").unwrap(),
None,
@@ -238,7 +213,7 @@ mod tests {
// Shape written before `refreshed_at` existed.
let legacy =
r#"{"scope":"repo:nonexistent.fake.record?action=create","title":null,"detail":null}"#;
cache.0.lock().unwrap().insert(key, legacy.to_string());
let _ = cache.set(&key, legacy, SEED_TTL).await;
let out = expand_scopes(&cache, "include:nonexistent.fake.permissionSet").await;
assert!(out.failures.is_empty());
assert_eq!(out.sets.len(), 1);
@@ -246,7 +221,7 @@ mod tests {
#[tokio::test]
async fn passthrough_scopes_untouched() {
let cache = MapCache::default();
let cache = MemoryCache::new();
let out = expand_scopes(&cache, "atproto repo:app.bsky.feed.post?action=create").await;
assert!(out.failures.is_empty());
assert!(out.sets.is_empty());
@@ -255,7 +230,7 @@ mod tests {
#[tokio::test]
async fn cache_miss_unresolvable_is_a_failure() {
let cache = MapCache::default();
let cache = MemoryCache::new();
let out = expand_scopes(&cache, "include:nonexistent.fake.permissionSet").await;
assert_eq!(out.sets.len(), 0);
assert_eq!(out.failures.len(), 1);
+42 -93
View File
@@ -165,12 +165,11 @@ impl PlcOpOrTombstone {
}
}
const PLC_CACHE_TTL_SECS: u64 = 300;
pub struct PlcClient {
base_url: String,
client: Client,
cache: Option<Arc<dyn Cache>>,
cache_ttl: Duration,
}
impl PlcClient {
@@ -193,12 +192,19 @@ impl PlcClient {
.connect_timeout(Duration::from_secs(connect_timeout_secs))
.pool_max_idle_per_host(5)
.pool_idle_timeout(Duration::from_secs(90))
.redirect(tranquil_types::redirect_policy(
tranquil_types::ReachPolicy::DEBUG_LOOPBACK,
))
.dns_resolver(tranquil_types::dns_guard(
tranquil_types::ReachPolicy::DEBUG_LOOPBACK,
))
.build()
.unwrap_or_else(|_| Client::new());
.expect("failed to build PLC directory HTTP client");
Self {
base_url,
client,
cache,
cache_ttl: Duration::from_secs(cfg.map_or(300, |c| c.plc.did_cache_ttl_secs)),
}
}
@@ -206,15 +212,7 @@ impl PlcClient {
urlencoding::encode(did.as_str()).to_string()
}
pub async fn get_document(&self, did: &Did) -> Result<Value, PlcError> {
let cache_key = crate::cache_keys::plc_doc_key(did);
if let Some(ref cache) = self.cache
&& let Some(cached) = cache.get(&cache_key).await
&& let Ok(value) = serde_json::from_str(&cached)
{
return Ok(value);
}
let url = format!("{}/{}", self.base_url, Self::encode_did(did));
async fn fetch_json<T: serde::de::DeserializeOwned>(&self, url: String) -> Result<T, PlcError> {
let response = self.client.get(&url).send().await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Err(PlcError::NotFound);
@@ -227,101 +225,52 @@ impl PlcClient {
status, body
)));
}
let value: Value = response
response
.json()
.await
.map_err(|e| PlcError::InvalidResponse(e.to_string()))?;
if let Some(ref cache) = self.cache
&& let Ok(json_str) = serde_json::to_string(&value)
{
let _ = cache
.set(
&cache_key,
&json_str,
Duration::from_secs(PLC_CACHE_TTL_SECS),
)
.await;
.map_err(|e| PlcError::InvalidResponse(e.to_string()))
}
async fn cached_fetch(&self, cache_key: &str, url: String) -> Result<Value, PlcError> {
match &self.cache {
Some(cache) => {
crate::cache::cached_json(cache.as_ref(), cache_key, self.cache_ttl, || {
self.fetch_json(url)
})
.await
}
None => self.fetch_json(url).await,
}
Ok(value)
}
pub async fn get_document(&self, did: &Did) -> Result<Value, PlcError> {
let url = format!("{}/{}", self.base_url, Self::encode_did(did));
self.cached_fetch(&crate::cache_keys::plc_doc_key(did), url)
.await
}
pub async fn get_document_data(&self, did: &Did) -> Result<Value, PlcError> {
let cache_key = crate::cache_keys::plc_data_key(did);
if let Some(ref cache) = self.cache
&& let Some(cached) = cache.get(&cache_key).await
&& let Ok(value) = serde_json::from_str(&cached)
{
return Ok(value);
}
let url = format!("{}/{}/data", self.base_url, Self::encode_did(did));
let response = self.client.get(&url).send().await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Err(PlcError::NotFound);
}
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
return Err(PlcError::InvalidResponse(format!(
"HTTP {}: {}",
status, body
)));
}
let value: Value = response
.json()
self.cached_fetch(&crate::cache_keys::plc_data_key(did), url)
.await
.map_err(|e| PlcError::InvalidResponse(e.to_string()))?;
if let Some(ref cache) = self.cache
&& let Ok(json_str) = serde_json::to_string(&value)
{
let _ = cache
.set(
&cache_key,
&json_str,
Duration::from_secs(PLC_CACHE_TTL_SECS),
)
.await;
}
Ok(value)
}
pub async fn get_last_op(&self, did: &Did) -> Result<PlcOpOrTombstone, PlcError> {
let url = format!("{}/{}/log/last", self.base_url, Self::encode_did(did));
let response = self.client.get(&url).send().await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Err(PlcError::NotFound);
}
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
return Err(PlcError::InvalidResponse(format!(
"HTTP {}: {}",
status, body
)));
}
response
.json()
.await
.map_err(|e| PlcError::InvalidResponse(e.to_string()))
self.fetch_json(format!(
"{}/{}/log/last",
self.base_url,
Self::encode_did(did)
))
.await
}
pub async fn get_audit_log(&self, did: &Did) -> Result<Vec<Value>, PlcError> {
let url = format!("{}/{}/log/audit", self.base_url, Self::encode_did(did));
let response = self.client.get(&url).send().await?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Err(PlcError::NotFound);
}
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
return Err(PlcError::InvalidResponse(format!(
"HTTP {}: {}",
status, body
)));
}
response
.json()
.await
.map_err(|e| PlcError::InvalidResponse(e.to_string()))
self.fetch_json(format!(
"{}/{}/log/audit",
self.base_url,
Self::encode_did(did)
))
.await
}
pub async fn send_operation(&self, did: &Did, operation: &Value) -> Result<(), PlcError> {
+156 -113
View File
@@ -4,15 +4,23 @@ use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, jwk:
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::{Arc, LazyLock};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use thiserror::Error;
use tokio::sync::{OnceCell, RwLock};
use tokio::sync::RwLock;
use tranquil_db_traits::SsoProviderType;
use tranquil_types::{SsoIssuer, SsoJwksUri};
use super::config::{AppleProviderConfig, ProviderConfig, SsoConfig};
use crate::cache::{Cache, cached_json};
use crate::cache_keys::{oidc_discovery_key, sso_jwks_key};
const SSO_HTTP_TIMEOUT: Duration = Duration::from_secs(15);
const SSO_DISCOVERY_TTL: Duration = Duration::from_secs(3600);
static APPLE_JWKS_URI: LazyLock<SsoJwksUri> = LazyLock::new(|| {
SsoJwksUri::new("https://appleid.apple.com/auth/keys")
.expect("Apple JWKS URI is a valid https URL")
});
struct PkceChallenge {
code_verifier: String,
@@ -28,6 +36,12 @@ fn create_http_client() -> Client {
Client::builder()
.timeout(SSO_HTTP_TIMEOUT)
.connect_timeout(Duration::from_secs(5))
.redirect(tranquil_types::redirect_policy(
tranquil_types::ReachPolicy::AllowPrivate,
))
.dns_resolver(tranquil_types::dns_guard(
tranquil_types::ReachPolicy::AllowPrivate,
))
.build()
.expect("Failed to create HTTP client")
}
@@ -367,16 +381,21 @@ impl SsoProvider for DiscordProvider {
}
}
#[derive(Debug, Clone, Deserialize)]
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OidcDiscoveryConfig {
pub issuer: String,
pub issuer: SsoIssuer,
pub authorization_endpoint: String,
pub token_endpoint: String,
pub userinfo_endpoint: Option<String>,
pub jwks_uri: Option<String>,
#[serde(
default,
deserialize_with = "tranquil_types::http_url::deserialize_optional"
)]
pub jwks_uri: Option<SsoJwksUri>,
}
struct OidcDiscoveryCache {
#[derive(Serialize, Deserialize)]
struct OidcDiscovery {
config: OidcDiscoveryConfig,
jwks: Option<JwkSet>,
}
@@ -385,10 +404,10 @@ pub struct OidcProvider {
provider_type: SsoProviderType,
client_id: String,
client_secret: String,
issuer: String,
issuer: SsoIssuer,
display_name: String,
http_client: Client,
discovery_cache: OnceCell<OidcDiscoveryCache>,
cache: Arc<dyn Cache>,
}
impl OidcProvider {
@@ -397,11 +416,25 @@ impl OidcProvider {
config: &ProviderConfig,
default_issuer: Option<&str>,
default_name: &str,
cache: Arc<dyn Cache>,
) -> Option<Self> {
let issuer = config
let issuer = match config
.issuer
.clone()
.or_else(|| default_issuer.map(String::from))?;
.or_else(|| default_issuer.map(String::from))
.map(SsoIssuer::new)
{
Some(Ok(issuer)) => issuer,
Some(Err(e)) => {
tracing::error!(
provider = %provider_type.as_str(),
error = %e,
"SSO provider disabled because its issuer isn't a usable http or https URL"
);
return None;
}
None => return None,
};
Some(Self {
provider_type,
@@ -413,74 +446,80 @@ impl OidcProvider {
.clone()
.unwrap_or_else(|| default_name.to_string()),
http_client: create_http_client(),
discovery_cache: OnceCell::new(),
cache,
})
}
async fn get_discovery(&self) -> Result<&OidcDiscoveryCache, SsoError> {
self.discovery_cache
.get_or_try_init(|| async {
let discovery_url = format!(
"{}/.well-known/openid-configuration",
self.issuer.trim_end_matches('/')
);
async fn get_discovery(&self) -> Result<OidcDiscovery, SsoError> {
cached_json(
self.cache.as_ref(),
&oidc_discovery_key(&self.issuer),
SSO_DISCOVERY_TTL,
|| self.fetch_discovery(),
)
.await
}
tracing::debug!(
provider = %self.provider_type.as_str(),
url = %discovery_url,
"Fetching OIDC discovery document"
);
async fn fetch_discovery(&self) -> Result<OidcDiscovery, SsoError> {
let discovery_url = self.issuer.endpoint(".well-known/openid-configuration");
let resp = self
.http_client
.get(&discovery_url)
.send()
.await
.map_err(|e| SsoError::Discovery(e.to_string()))?;
tracing::debug!(
provider = %self.provider_type.as_str(),
url = %discovery_url,
"Fetching OIDC discovery document"
);
if !resp.status().is_success() {
return Err(SsoError::Discovery(format!(
"Discovery endpoint returned {}",
resp.status()
)));
}
let config: OidcDiscoveryConfig = resp
.json()
.await
.map_err(|e| SsoError::Discovery(e.to_string()))?;
let jwks = match &config.jwks_uri {
Some(jwks_uri) => {
tracing::debug!(
provider = %self.provider_type.as_str(),
url = %jwks_uri,
"Fetching JWKS"
);
let jwks_resp =
self.http_client.get(jwks_uri).send().await.map_err(|e| {
SsoError::Discovery(format!("JWKS fetch failed: {}", e))
})?;
if jwks_resp.status().is_success() {
Some(jwks_resp.json::<JwkSet>().await.map_err(|e| {
SsoError::Discovery(format!("JWKS parse failed: {}", e))
})?)
} else {
tracing::warn!(
provider = %self.provider_type.as_str(),
status = %jwks_resp.status(),
"JWKS fetch returned non-success status"
);
None
}
}
None => None,
};
Ok(OidcDiscoveryCache { config, jwks })
})
let resp = self
.http_client
.get(discovery_url)
.send()
.await
.map_err(|e| SsoError::Discovery(e.to_string()))?;
if !resp.status().is_success() {
return Err(SsoError::Discovery(format!(
"Discovery endpoint returned {}",
resp.status()
)));
}
let config: OidcDiscoveryConfig = resp
.json()
.await
.map_err(|e| SsoError::Discovery(e.to_string()))?;
let jwks =
match &config.jwks_uri {
Some(jwks_uri) => {
tracing::debug!(
provider = %self.provider_type.as_str(),
url = %jwks_uri,
"Fetching JWKS"
);
let jwks_resp = self
.http_client
.get(jwks_uri.as_str())
.send()
.await
.map_err(|e| SsoError::Discovery(format!("JWKS fetch failed: {}", e)))?;
if jwks_resp.status().is_success() {
Some(jwks_resp.json::<JwkSet>().await.map_err(|e| {
SsoError::Discovery(format!("JWKS parse failed: {}", e))
})?)
} else {
tracing::warn!(
provider = %self.provider_type.as_str(),
status = %jwks_resp.status(),
"JWKS fetch returned non-success status"
);
None
}
}
None => None,
};
Ok(OidcDiscovery { config, jwks })
}
fn generate_pkce() -> PkceChallenge {
@@ -602,9 +641,7 @@ impl SsoProvider for OidcProvider {
let auth_endpoint = match self.provider_type {
SsoProviderType::Google => "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
SsoProviderType::Gitlab => {
format!("{}/oauth/authorize", self.issuer.trim_end_matches('/'))
}
SsoProviderType::Gitlab => self.issuer.endpoint("oauth/authorize").to_string(),
_ => {
let discovery = self.get_discovery().await?;
discovery.config.authorization_endpoint.clone()
@@ -638,7 +675,7 @@ impl SsoProvider for OidcProvider {
) -> Result<SsoTokenResponse, SsoError> {
let token_endpoint = match self.provider_type {
SsoProviderType::Google => "https://oauth2.googleapis.com/token".to_string(),
SsoProviderType::Gitlab => format!("{}/oauth/token", self.issuer.trim_end_matches('/')),
SsoProviderType::Gitlab => self.issuer.endpoint("oauth/token").to_string(),
_ => {
let discovery = self.get_discovery().await?;
discovery.config.token_endpoint.clone()
@@ -721,9 +758,7 @@ impl SsoProvider for OidcProvider {
SsoProviderType::Google => {
"https://openidconnect.googleapis.com/v1/userinfo".to_string()
}
SsoProviderType::Gitlab => {
format!("{}/oauth/userinfo", self.issuer.trim_end_matches('/'))
}
SsoProviderType::Gitlab => self.issuer.endpoint("oauth/userinfo").to_string(),
_ => {
let discovery = self.get_discovery().await?;
discovery
@@ -777,11 +812,11 @@ pub struct AppleProvider {
private_key_pem: String,
http_client: Client,
client_secret_cache: RwLock<Option<CachedClientSecret>>,
jwks_cache: OnceCell<JwkSet>,
cache: Arc<dyn Cache>,
}
impl AppleProvider {
pub fn new(config: &AppleProviderConfig) -> Result<Self, SsoError> {
pub fn new(config: &AppleProviderConfig, cache: Arc<dyn Cache>) -> Result<Self, SsoError> {
let key_pem = config.private_key_pem.replace("\\n", "\n");
jsonwebtoken::EncodingKey::from_ec_pem(key_pem.as_bytes())
@@ -794,7 +829,7 @@ impl AppleProvider {
private_key_pem: key_pem,
http_client: create_http_client(),
client_secret_cache: RwLock::new(None),
jwks_cache: OnceCell::new(),
cache,
})
}
@@ -868,29 +903,35 @@ impl AppleProvider {
Ok(generated.secret)
}
async fn get_jwks(&self) -> Result<&JwkSet, SsoError> {
self.jwks_cache
.get_or_try_init(|| async {
tracing::debug!("Fetching Apple JWKS");
let resp = self
.http_client
.get("https://appleid.apple.com/auth/keys")
.send()
.await
.map_err(|e| SsoError::Discovery(format!("Apple JWKS fetch failed: {}", e)))?;
async fn get_jwks(&self) -> Result<JwkSet, SsoError> {
cached_json(
self.cache.as_ref(),
&sso_jwks_key(&APPLE_JWKS_URI),
SSO_DISCOVERY_TTL,
|| self.fetch_jwks(),
)
.await
}
if !resp.status().is_success() {
return Err(SsoError::Discovery(format!(
"Apple JWKS returned {}",
resp.status()
)));
}
resp.json::<JwkSet>()
.await
.map_err(|e| SsoError::Discovery(format!("Apple JWKS parse failed: {}", e)))
})
async fn fetch_jwks(&self) -> Result<JwkSet, SsoError> {
tracing::debug!("Fetching Apple JWKS");
let resp = self
.http_client
.get(APPLE_JWKS_URI.as_str())
.send()
.await
.map_err(|e| SsoError::Discovery(format!("Apple JWKS fetch failed: {}", e)))?;
if !resp.status().is_success() {
return Err(SsoError::Discovery(format!(
"Apple JWKS returned {}",
resp.status()
)));
}
resp.json()
.await
.map_err(|e| SsoError::Discovery(format!("Apple JWKS parse failed: {}", e)))
}
fn validate_id_token(
@@ -1043,7 +1084,7 @@ impl SsoProvider for AppleProvider {
})?;
let jwks = self.get_jwks().await?;
let claims = self.validate_id_token(id_token, jwks, expected_nonce)?;
let claims = self.validate_id_token(id_token, &jwks, expected_nonce)?;
tracing::debug!(
sub = %claims.sub,
@@ -1063,10 +1104,11 @@ impl SsoProvider for AppleProvider {
#[derive(Clone)]
pub struct SsoManager {
providers: HashMap<SsoProviderType, Arc<dyn SsoProvider>>,
config: &'static SsoConfig,
}
impl SsoManager {
pub fn from_config(config: &SsoConfig) -> Self {
pub fn from_config(config: &'static SsoConfig, cache: Arc<dyn Cache>) -> Self {
let mut providers: HashMap<SsoProviderType, Arc<dyn SsoProvider>> = HashMap::new();
if let Some(ref cfg) = config.github {
@@ -1086,13 +1128,15 @@ impl SsoManager {
cfg,
Some("https://accounts.google.com"),
"Google",
cache.clone(),
)
{
providers.insert(SsoProviderType::Google, Arc::new(provider));
}
if let Some(ref cfg) = config.gitlab
&& let Some(provider) = OidcProvider::new(SsoProviderType::Gitlab, cfg, None, "GitLab")
&& let Some(provider) =
OidcProvider::new(SsoProviderType::Gitlab, cfg, None, "GitLab", cache.clone())
{
providers.insert(SsoProviderType::Gitlab, Arc::new(provider));
}
@@ -1103,13 +1147,14 @@ impl SsoManager {
cfg,
None,
cfg.display_name.as_deref().unwrap_or("SSO"),
cache.clone(),
)
{
providers.insert(SsoProviderType::Oidc, Arc::new(provider));
}
if let Some(ref cfg) = config.apple {
match AppleProvider::new(cfg) {
match AppleProvider::new(cfg, cache.clone()) {
Ok(provider) => {
providers.insert(SsoProviderType::Apple, Arc::new(provider));
}
@@ -1119,7 +1164,11 @@ impl SsoManager {
}
}
Self { providers }
Self { providers, config }
}
pub fn config(&self) -> &'static SsoConfig {
self.config
}
pub fn get_provider(&self, provider_type: SsoProviderType) -> Option<Arc<dyn SsoProvider>> {
@@ -1137,9 +1186,3 @@ impl SsoManager {
!self.providers.is_empty()
}
}
impl Default for SsoManager {
fn default() -> Self {
Self::from_config(SsoConfig::get())
}
}
+34 -7
View File
@@ -15,10 +15,12 @@ use std::error::Error;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::sync::broadcast;
use tokio_util::sync::CancellationToken;
use tranquil_db::PostgresRepositories;
use tranquil_db_traits::SequencedEvent;
use tranquil_oauth::ClientMetadataCache;
static RATE_LIMITING_DISABLED: AtomicBool = AtomicBool::new(false);
@@ -49,6 +51,7 @@ pub struct AppState {
pub sso_manager: SsoManager,
pub webauthn_config: Arc<WebAuthnConfig>,
pub cross_pds_oauth: Arc<CrossPdsOAuthClient>,
pub client_metadata_cache: ClientMetadataCache,
pub shutdown: CancellationToken,
pub bootstrap_invite_code: Option<crate::types::InviteCode>,
pub signal_sender: Option<Arc<tranquil_signal::SignalSlot>>,
@@ -210,6 +213,27 @@ impl RateLimitKind {
}
}
const CLIENT_METADATA_TTL: Duration = Duration::from_secs(3600);
struct CacheBound {
did_resolver: Arc<DidResolver>,
cross_pds_oauth: Arc<CrossPdsOAuthClient>,
client_metadata_cache: ClientMetadataCache,
sso_manager: SsoManager,
}
impl CacheBound {
fn new(cache: &Arc<dyn Cache>, sso_config: &'static SsoConfig) -> Self {
tranquil_lexicon::LexiconRegistry::global().set_shared_cache(cache.clone());
Self {
did_resolver: Arc::new(DidResolver::new(cache.clone())),
cross_pds_oauth: Arc::new(CrossPdsOAuthClient::new(cache.clone())),
client_metadata_cache: ClientMetadataCache::new(cache.clone(), CLIENT_METADATA_TTL),
sso_manager: SsoManager::from_config(sso_config, cache.clone()),
}
}
}
impl AppState {
pub fn plc_client(&self) -> PlcClient {
PlcClient::with_cache(None, Some(self.cache.clone()))
@@ -366,10 +390,7 @@ impl AppState {
let (cache, distributed_rate_limiter) = create_cache(shutdown.clone())
.await
.expect("Failed to initialize cache and distributed rate limiter at startup");
let did_resolver = Arc::new(DidResolver::new());
let cross_pds_oauth = Arc::new(CrossPdsOAuthClient::new(cache.clone()));
let sso_config = SsoConfig::init();
let sso_manager = SsoManager::from_config(sso_config);
let bound = CacheBound::new(&cache, SsoConfig::init());
let webauthn_config = Arc::new(
WebAuthnConfig::new(&cfg.server.hostname)
.expect("Failed to create WebAuthn config at startup"),
@@ -385,9 +406,10 @@ impl AppState {
circuit_breakers,
cache,
distributed_rate_limiter,
did_resolver,
cross_pds_oauth,
sso_manager,
did_resolver: bound.did_resolver,
cross_pds_oauth: bound.cross_pds_oauth,
client_metadata_cache: bound.client_metadata_cache,
sso_manager: bound.sso_manager,
webauthn_config,
shutdown,
bootstrap_invite_code: None,
@@ -410,6 +432,11 @@ impl AppState {
cache: Arc<dyn Cache>,
distributed_rate_limiter: Arc<dyn DistributedRateLimiter>,
) -> Self {
let bound = CacheBound::new(&cache, self.sso_manager.config());
self.did_resolver = bound.did_resolver;
self.cross_pds_oauth = bound.cross_pds_oauth;
self.client_metadata_cache = bound.client_metadata_cache;
self.sso_manager = bound.sso_manager;
self.cache = cache;
self.distributed_rate_limiter = distributed_rate_limiter;
self
+1
View File
@@ -1,5 +1,6 @@
pub use tranquil_types::*;
#[cfg(feature = "bsky")]
use std::sync::LazyLock;
#[cfg(feature = "bsky")]
+9 -6
View File
@@ -89,6 +89,10 @@ pub const HEADER_ATPROTO_CONTENT_LABELERS: HeaderName =
HeaderName::from_static("atproto-content-labelers");
#[cfg(feature = "bsky-support")]
pub const HEADER_X_BSKY_TOPICS: HeaderName = HeaderName::from_static("x-bsky-topics");
#[cfg(feature = "bsky-support")]
pub const CORS_BSKY_ALLOW_HEADERS: [HeaderName; 1] = [HEADER_X_BSKY_TOPICS];
#[cfg(not(feature = "bsky-support"))]
pub const CORS_BSKY_ALLOW_HEADERS: [HeaderName; 0] = [];
pub fn get_header_str(
headers: &HeaderMap,
@@ -250,11 +254,7 @@ pub fn build_full_url(path: &str) -> String {
&& (path.starts_with("/com.atproto.")
// BSKY: Bluesky requires that the PDS implement some app.bsky.* endpoints so we need to deal with those here too.
// TODO: surely we can figure out a way to do this more generically?
|| (if cfg!(feature = "bsky-support") {
path.starts_with("/app.bsky.")
} else {
true
})
|| (cfg!(feature = "bsky-support") && path.starts_with("/app.bsky."))
|| path.starts_with("/_"))
{
format!("/xrpc{path}")
@@ -798,7 +798,10 @@ mod tests {
);
assert_eq!(
build_full_url("/app.bsky.feed.getTimeline"),
"https://example.com/xrpc/app.bsky.feed.getTimeline"
match cfg!(feature = "bsky-support") {
true => "https://example.com/xrpc/app.bsky.feed.getTimeline",
false => "https://example.com/app.bsky.feed.getTimeline",
}
);
assert_eq!(
build_full_url("/_health"),
@@ -132,6 +132,10 @@ fn validate_preamble<'a>(
Ok((record_type, obj))
}
#[cfg_attr(
not(feature = "bsky"),
expect(unused_variables, reason = "only bsky record checks read obj and rkey")
)]
fn check_banned_content(
record_type: &str,
obj: &serde_json::Map<String, Value>,
@@ -211,6 +215,7 @@ fn check_post_banned_content(obj: &serde_json::Map<String, Value>) -> Result<(),
Ok(())
}
#[cfg(feature = "bsky")]
fn check_string_field(
obj: &serde_json::Map<String, Value>,
field: &str,
+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::*;
+1 -1
View File
@@ -390,7 +390,7 @@
# Default value: 5
#connect_timeout_secs = 5
# Seconds to cache DID documents in memory.
# Seconds to cache DID documents.
#
# Can also be specified via environment variable `DID_CACHE_TTL_SECS`.
#
+5 -2
View File
@@ -16,12 +16,15 @@ build-release:
check:
cargo check
clippy:
cargo clippy -- -D warnings
cargo clippy --all-targets -- -D warnings
lint-no-bsky:
cargo clippy -p tranquil-server --no-default-features --features frontend,postgres,s3,valkey --all-targets -- -D warnings
cargo clippy -p tranquil-pds --no-default-features --all-targets -- -D warnings
fmt:
cargo fmt
fmt-check:
cargo fmt -- --check
lint: fmt-check clippy
lint: fmt-check clippy lint-no-bsky
test-store:
SQLX_OFFLINE=true cargo nextest run -p tranquil-store --features tranquil-store/test-harness