mirror of
https://tangled.org/tranquil.farm/tranquil-pds
synced 2026-08-24 18:26:06 +00:00
Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
37234797b4 | ||
|
|
6297b1a451 | ||
|
|
107149f396 | ||
|
|
bd47cbdaa4 | ||
|
|
9840ac77cf | ||
|
|
420ce1e201 | ||
|
|
c723bc2164 | ||
|
|
18455f54f2 | ||
|
|
ce2f05b9d4 | ||
|
|
c88f69f31d | ||
|
|
b3c314ce66 | ||
|
|
434079a732 | ||
|
|
a5a2f30bbe | ||
|
|
dc2fbe6654 | ||
|
|
bc751b0ee2 | ||
|
|
9e78206cf4 | ||
|
|
779dc1b985 | ||
|
|
1dc0c40206 |
+1
-1
@@ -64,7 +64,7 @@ In order of importance the following rules describe what "correct" means for Tra
|
||||
and not something said application relies on for proper functioning.
|
||||
|
||||
There is bound to be edge cases that these rules don't fully cover.
|
||||
Here common sense, community sentiment, furthering the goals of atproto itself, and ultimately maintainer opinion take precedence over support for any individual applicaion.
|
||||
Here common sense, community sentiment, furthering the goals of atproto itself, and ultimately maintainer opinion take precedence over support for any individual application.
|
||||
Even Bluesky.
|
||||
|
||||
The rules above are meant to capture Tranquils goals of being correct while being community oriented and avoiding as much "Bluesky-defaultism" as possible.
|
||||
|
||||
Generated
+48
-1
@@ -105,6 +105,21 @@ version = "0.1.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "250f629c0161ad8107cf89319e990051fae62832fd343083bea452d93e2205fd"
|
||||
|
||||
[[package]]
|
||||
name = "alloc-no-stdlib"
|
||||
version = "2.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cc7bb162ec39d46ab1ca8c77bf72e890535becd1751bb45f64c597edb4c8c6b3"
|
||||
|
||||
[[package]]
|
||||
name = "alloc-stdlib"
|
||||
version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0e76a019e91224d279006ff972f1e984179a6e9feb050adba6ce8274aef23195"
|
||||
dependencies = [
|
||||
"alloc-no-stdlib",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "allocator-api2"
|
||||
version = "0.2.21"
|
||||
@@ -1250,6 +1265,27 @@ dependencies = [
|
||||
"cfg_aliases",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "brotli"
|
||||
version = "8.0.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5cc91aac060a7a1e25823bdccbfb6af1875b88f17c6daac97894eed8207166b3"
|
||||
dependencies = [
|
||||
"alloc-no-stdlib",
|
||||
"alloc-stdlib",
|
||||
"brotli-decompressor",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "brotli-decompressor"
|
||||
version = "5.0.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3a32acac15fe1967bc3986b2a6347dffc965602354ea6f450ad07e8bfd253583"
|
||||
dependencies = [
|
||||
"alloc-no-stdlib",
|
||||
"alloc-stdlib",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "bs58"
|
||||
version = "0.5.1"
|
||||
@@ -7682,6 +7718,7 @@ dependencies = [
|
||||
"base32",
|
||||
"base64 0.22.1",
|
||||
"bcrypt",
|
||||
"brotli",
|
||||
"chrono",
|
||||
"hmac",
|
||||
"k256",
|
||||
@@ -7801,7 +7838,10 @@ dependencies = [
|
||||
"async-trait",
|
||||
"bytes",
|
||||
"futures",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.18",
|
||||
"tranquil-types",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -7818,9 +7858,9 @@ dependencies = [
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tranquil-infra",
|
||||
"tranquil-types",
|
||||
"unicode-segmentation",
|
||||
"urlencoding",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
@@ -7843,6 +7883,7 @@ dependencies = [
|
||||
"sqlx",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tranquil-infra",
|
||||
"tranquil-types",
|
||||
"uuid",
|
||||
]
|
||||
@@ -7874,6 +7915,7 @@ dependencies = [
|
||||
"tranquil-config",
|
||||
"tranquil-crypto",
|
||||
"tranquil-db-traits",
|
||||
"tranquil-infra",
|
||||
"tranquil-pds",
|
||||
"tranquil-scopes",
|
||||
"tranquil-types",
|
||||
@@ -7954,6 +7996,7 @@ dependencies = [
|
||||
"tranquil-config",
|
||||
"tranquil-db",
|
||||
"tranquil-db-traits",
|
||||
"tranquil-infra",
|
||||
"tranquil-lexicon",
|
||||
"tranquil-oauth",
|
||||
"tranquil-oauth-server",
|
||||
@@ -8184,10 +8227,14 @@ dependencies = [
|
||||
"cid",
|
||||
"jacquard-common",
|
||||
"rand 0.8.5",
|
||||
"reqwest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sqlx",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"url",
|
||||
"uuid",
|
||||
]
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -489,9 +489,15 @@ pub fn webhook_routes() -> axum::Router<AppState> {
|
||||
pub fn misc_routes() -> axum::Router<AppState> {
|
||||
use axum::routing::get;
|
||||
|
||||
axum::Router::new()
|
||||
let router = axum::Router::new()
|
||||
.route("/health", get(server::health))
|
||||
.route("/robots.txt", get(server::robots_txt))
|
||||
.route("/favicon.ico", get(server::get_logo))
|
||||
.route("/u/{handle}/did.json", get(identity::user_did_doc))
|
||||
.route("/u/{handle}/did.json", get(identity::user_did_doc));
|
||||
|
||||
if tranquil_config::get().server.rfc_moo_compliance {
|
||||
router.route("/cow.txt", get(server::cow_txt))
|
||||
} else {
|
||||
router
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,7 +10,7 @@ use serde_json::Value;
|
||||
use std::str::FromStr;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use tracing::{error, info, warn};
|
||||
use tracing::{debug, error, info, warn};
|
||||
use tranquil_pds::api::EmptyResponse;
|
||||
use tranquil_pds::api::error::{ApiError, DbResultExt};
|
||||
use tranquil_pds::auth::{Auth, NotTakendown, Permissive, require_legacy_session_mfa};
|
||||
@@ -212,9 +212,10 @@ async fn assert_valid_did_document_for_service(
|
||||
if let Some(ref expected_rotation_key) = server_rotation_key
|
||||
&& !doc_rotation_keys.contains(&expected_rotation_key.as_str())
|
||||
{
|
||||
return Err(ApiError::InvalidRequest(
|
||||
"Server rotation key not included in PLC DID data".into(),
|
||||
));
|
||||
debug!(
|
||||
"DID {} rotation keys {:?} omit the PDS-managed server rotation key {}",
|
||||
did, doc_rotation_keys, expected_rotation_key
|
||||
);
|
||||
}
|
||||
|
||||
let doc_signing_key = doc_data
|
||||
@@ -251,13 +252,10 @@ async fn assert_valid_did_document_for_service(
|
||||
}
|
||||
|
||||
if !doc_rotation_keys.contains(&expected_did_key.as_str()) {
|
||||
warn!(
|
||||
debug!(
|
||||
"DID {} rotation keys {:?} omit the PDS-managed signing key {}",
|
||||
did, doc_rotation_keys, expected_did_key
|
||||
);
|
||||
return Err(ApiError::InvalidRequest(
|
||||
"PLC rotation keys omit the PDS-managed signing key required to sign operations for this identity".into(),
|
||||
));
|
||||
}
|
||||
}
|
||||
} else if let Some(host_and_path) = did.as_str().strip_prefix("did:web:") {
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
..........................
|
||||
....*o|||||||8#@@@@@@@@@@@@@@@@@@@@@@@###&|o:_..
|
||||
..*:o|||&8##@###8888888######@#@###########################|*...
|
||||
.:o|||8#####8888|:::**. *&########################@@################&o_
|
||||
.*o&8###@#8&o*_. :###@##############@########################@@##&o_
|
||||
.*o8########& :##@#@##############@############################@###|_
|
||||
.*o|8##########8o .#######################################################&o_
|
||||
*&##|_ ..*&##8&o*|88888|_ _#######################################@##################|.
|
||||
*#####& *&######&o_..*o|o:_ .&##o _###########################################################&_
|
||||
_##8*##8 .|88|:::|#######8###8|*:_ .&#@@8 _##@@@########################################################&_
|
||||
_#@8_##8_ *8#8|*_ _:|#####&&####8 .&##############################################################|
|
||||
_#@8.|##8_ _::o###8&##8 .|##@############################8###########################@@#|_
|
||||
*###o.|88o ..*&####|..##& _|##########################8|_ .|#############################8
|
||||
*|###|_ ._&####8|*_ _*_ _::&8888888888888888|::*_ .|##@####@@@##################|
|
||||
*&###|_ _:_ .&88###8|*_ ..... .|#####@@@##################8
|
||||
.##@#& _##& .|##o _#@@#@#| .|#######&:_ _|###@####################8
|
||||
.:8##8o _o:*&##| *##8_.&@@##@#| _::o8#8|::|#####|_ _|#################88###8
|
||||
.&##&*_ *###o_###| .|##8*&##|*###o _###8####8|_ _:|###|_ .*o|||o:_ _:::&8888888888|_ _##8
|
||||
.###|_. _###o *###|*&#######8 *##8 .##8_ _:|###|_ _|###|_ .&########o _##&
|
||||
_|####8|&##8:_ _|#########88o .##8 *##& _|###o .|###o .#########| .o##o
|
||||
o#8|*:#@@###o _:::*__*_ _##8_ _##8_ _&##|_ *##8_ *8#####8|_ .*oo:_ o##|
|
||||
*###o.&#####& _oo* .8##& .8##8_ .|##& o##& _::::_.*o|8######|_ .##8.
|
||||
_###&o&##8_:*_ .###& .###|_&#####| _##8 :###o *ooo&#########@#@#& ....:##&
|
||||
.|8||###&. _**_ .###88##|*&###|*._&##& *|##8_ *o&####@@####@@######& .*o||||||&#######8_
|
||||
*&###o _|88##8_ _:8######|*:###|_ _##################88|_ *&#################&
|
||||
*#####o *&8o *##& _:::*_.&##|_ _#@##############8_ :##################8*
|
||||
.###&##8_.|88o *&8o _@@& .###| .&####@#########|_ .####@###@@########8*
|
||||
_##&.|###|_.... .|88o _##8* *&###|_ *###o _:&88######8|_ .*o|||##################o
|
||||
_##8_ _|########|_ .*o8####&#@@#@##o *#@8 _*:*. .&###@###################|
|
||||
.|##& _:::::&##& .*&##############@#8_ .###o _#@####################|_
|
||||
.&##|_ .&##8_ *o&####################8_ *##& .&#####@@############8*
|
||||
.|##8_.&###&####8_ _########################8**##& _##################|_
|
||||
.&###&##888888##8_ .|88######@###########|*######o _|8###############|_
|
||||
.&#####o *###|_ _::::::*o##8**o##8 .|###8o .&##@#############&*
|
||||
.|####o _|###|_ _##8.*&##& _*_ .#################|
|
||||
_*_ _|###|_.. .|#####8|_ *&#@#########8###&
|
||||
*#######&|o:_... ..*:::*. ......._:o&8####888&o:#####8_
|
||||
_###|&888#####@#####&|||o:_........................._:o||||8##@@@@####8|:*_ _:::*_
|
||||
.|#@###o _:::o#@#888######@@@@@@@@@@@@@@@@@@@@@@@@#####888|::::::**_
|
||||
_::*_ :##& *&8|_:::::::::::::::::::::::::**_
|
||||
.###8||&##8o
|
||||
_|888888|_
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -29,6 +29,10 @@ pub async fn robots_txt() -> impl IntoResponse {
|
||||
"# Hello!\n\n# Crawling the public API is allowed\nUser-agent: *\nAllow: /\n",
|
||||
)
|
||||
}
|
||||
|
||||
pub async fn cow_txt() -> &'static str {
|
||||
include_str!("cow.txt")
|
||||
}
|
||||
pub fn is_self_hosted_did_web_enabled() -> bool {
|
||||
tranquil_config::get().server.enable_pds_hosted_did_web
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ pub use email::{
|
||||
};
|
||||
pub use invite::{create_invite_code, create_invite_codes, get_account_invite_codes};
|
||||
pub use logo::get_logo;
|
||||
pub use meta::{describe_server, health, robots_txt};
|
||||
pub use meta::{cow_txt, describe_server, health, robots_txt};
|
||||
pub use migration::{get_did_document, update_did_document};
|
||||
pub use passkey_account::{
|
||||
complete_passkey_setup, create_passkey_account, recover_passkey_account,
|
||||
|
||||
@@ -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),
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -24,3 +24,4 @@ subtle = { workspace = true }
|
||||
totp-rs = { workspace = true }
|
||||
urlencoding = { workspace = true }
|
||||
uuid = { workspace = true }
|
||||
brotli = "8.0.4"
|
||||
|
||||
@@ -0,0 +1,183 @@
|
||||
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use brotli::{CompressorWriter, Decompressor};
|
||||
use std::fmt;
|
||||
use std::io::{Read, Write};
|
||||
|
||||
const COMPRESSED_PREFIX: &str = "$br$";
|
||||
const QUALITY: u32 = 9;
|
||||
const WINDOW_BITS: u32 = 16;
|
||||
const BUFFER_SIZE: usize = 4096;
|
||||
const MAX_SCOPE_LEN: u64 = 64 * 1024;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ScopeDecodeError {
|
||||
Base64DecodeFailed,
|
||||
DecompressFailed,
|
||||
TooLarge,
|
||||
}
|
||||
|
||||
impl fmt::Display for ScopeDecodeError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Base64DecodeFailed => write!(f, "Base64 decode of compressed scope failed"),
|
||||
Self::DecompressFailed => write!(f, "Brotli decompression of scope failed"),
|
||||
Self::TooLarge => write!(f, "Decompressed scope exceeds maximum length"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for ScopeDecodeError {}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ScopeEncodeError {
|
||||
TooLarge,
|
||||
}
|
||||
|
||||
impl fmt::Display for ScopeEncodeError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::TooLarge => write!(f, "Scope exceeds maximum length"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for ScopeEncodeError {}
|
||||
|
||||
fn brotli_compress(input: &str) -> Vec<u8> {
|
||||
let mut writer = CompressorWriter::new(Vec::new(), BUFFER_SIZE, QUALITY, WINDOW_BITS);
|
||||
|
||||
writer
|
||||
.write_all(input.as_bytes())
|
||||
.expect("writing to a Vec cannot fail");
|
||||
|
||||
writer.into_inner()
|
||||
}
|
||||
|
||||
fn brotli_decompress(input: &[u8]) -> Result<String, ScopeDecodeError> {
|
||||
let mut output = String::new();
|
||||
|
||||
Decompressor::new(input, BUFFER_SIZE)
|
||||
.take(MAX_SCOPE_LEN + 1)
|
||||
.read_to_string(&mut output)
|
||||
.map_err(|_| ScopeDecodeError::DecompressFailed)?;
|
||||
|
||||
if output.len() as u64 > MAX_SCOPE_LEN {
|
||||
return Err(ScopeDecodeError::TooLarge);
|
||||
}
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
pub fn encode_scope(scope: &str) -> Result<String, ScopeEncodeError> {
|
||||
if scope.len() as u64 > MAX_SCOPE_LEN {
|
||||
return Err(ScopeEncodeError::TooLarge);
|
||||
}
|
||||
|
||||
let tagged = format!(
|
||||
"{COMPRESSED_PREFIX}{}",
|
||||
URL_SAFE_NO_PAD.encode(brotli_compress(scope))
|
||||
);
|
||||
|
||||
if tagged.len() < scope.len() || scope.starts_with(COMPRESSED_PREFIX) {
|
||||
Ok(tagged)
|
||||
} else {
|
||||
Ok(scope.to_owned())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decode_scope(scope: &str) -> Result<String, ScopeDecodeError> {
|
||||
let Some(encoded) = scope.strip_prefix(COMPRESSED_PREFIX) else {
|
||||
return Ok(scope.to_owned());
|
||||
};
|
||||
|
||||
let compressed = URL_SAFE_NO_PAD
|
||||
.decode(encoded)
|
||||
.map_err(|_| ScopeDecodeError::Base64DecodeFailed)?;
|
||||
|
||||
brotli_decompress(&compressed)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn long_scope() -> String {
|
||||
let mut scope = String::from("transition:generic transition:chat.bsky");
|
||||
for collection in [
|
||||
"social.colibri.message",
|
||||
"social.colibri.community",
|
||||
"social.colibri.reaction",
|
||||
"social.colibri.member",
|
||||
"social.colibri.channel.read",
|
||||
] {
|
||||
scope.push_str(&format!(" repo:{collection}?action=create&action=delete"));
|
||||
}
|
||||
scope
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn long_scope_roundtrips_through_compression() {
|
||||
let scope = long_scope();
|
||||
let encoded = encode_scope(&scope).unwrap();
|
||||
|
||||
assert!(encoded.starts_with(COMPRESSED_PREFIX));
|
||||
assert!(encoded.len() < scope.len());
|
||||
assert_eq!(decode_scope(&encoded).unwrap(), scope);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn short_scope_stays_plaintext() {
|
||||
let encoded = encode_scope("com.atproto.access").unwrap();
|
||||
|
||||
assert_eq!(encoded, "com.atproto.access");
|
||||
assert_eq!(decode_scope(&encoded).unwrap(), "com.atproto.access");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn untagged_scope_passes_through() {
|
||||
assert_eq!(
|
||||
decode_scope("com.atproto.refresh").unwrap(),
|
||||
"com.atproto.refresh"
|
||||
);
|
||||
assert_eq!(decode_scope("").unwrap(), "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_compressed_scope_errors_instead_of_panicking() {
|
||||
assert_eq!(
|
||||
decode_scope("$br$not valid base64!"),
|
||||
Err(ScopeDecodeError::Base64DecodeFailed)
|
||||
);
|
||||
assert_eq!(
|
||||
decode_scope("$br$AAAAAAAAAAAAAAAA"),
|
||||
Err(ScopeDecodeError::DecompressFailed)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compression_bomb_is_rejected() {
|
||||
let bomb = URL_SAFE_NO_PAD.encode(brotli_compress(&"a".repeat(MAX_SCOPE_LEN as usize * 2)));
|
||||
|
||||
assert_eq!(
|
||||
decode_scope(&format!("{COMPRESSED_PREFIX}{bomb}")),
|
||||
Err(ScopeDecodeError::TooLarge)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn plaintext_that_looks_compressed_roundtrips() {
|
||||
let scope = "$br$repo:*";
|
||||
let encoded = encode_scope(scope).unwrap();
|
||||
|
||||
assert!(encoded.starts_with(COMPRESSED_PREFIX));
|
||||
assert_eq!(decode_scope(&encoded).unwrap(), scope);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn encode_rejects_oversized_scope() {
|
||||
let oversized = "a".repeat(MAX_SCOPE_LEN as usize + 1);
|
||||
|
||||
assert_eq!(encode_scope(&oversized), Err(ScopeEncodeError::TooLarge));
|
||||
assert!(encode_scope(&"a".repeat(MAX_SCOPE_LEN as usize)).is_ok());
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
mod compress;
|
||||
mod token;
|
||||
mod totp;
|
||||
mod types;
|
||||
@@ -12,6 +13,8 @@ pub use token::{
|
||||
create_service_token_hs256,
|
||||
};
|
||||
|
||||
pub use compress::{ScopeDecodeError, ScopeEncodeError, decode_scope, encode_scope};
|
||||
|
||||
pub use totp::{
|
||||
TotpError, decrypt_totp_secret, encrypt_totp_secret, generate_backup_codes,
|
||||
generate_qr_png_base64, generate_totp_secret, generate_totp_uri, hash_backup_code,
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use crate::compress::encode_scope;
|
||||
|
||||
use super::types::{
|
||||
ActClaim, Claims, Header, SigningAlgorithm, TokenScope, TokenType, TokenWithMetadata,
|
||||
};
|
||||
use anyhow::Result;
|
||||
use anyhow::{Context, Result};
|
||||
use base64::Engine as _;
|
||||
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||
use chrono::{DateTime, Duration, Utc};
|
||||
@@ -205,7 +207,7 @@ fn create_signed_token_pinned(
|
||||
aud: format!("did:web:{}", aud_hostname),
|
||||
exp: expiration,
|
||||
iat: Utc::now().timestamp(),
|
||||
scope: Some(scope.to_string()),
|
||||
scope: Some(encode_scope(scope).context("Scope too large to encode")?),
|
||||
lxm: None,
|
||||
jti: jti.clone(),
|
||||
act,
|
||||
@@ -328,7 +330,7 @@ fn create_hs256_token_with_metadata(
|
||||
),
|
||||
exp: expiration,
|
||||
iat: Utc::now().timestamp(),
|
||||
scope: Some(scope.to_string()),
|
||||
scope: Some(encode_scope(scope).context("Scope too large to encode")?),
|
||||
lxm: None,
|
||||
jti: jti.clone(),
|
||||
act: None,
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use crate::compress::decode_scope;
|
||||
|
||||
use super::types::{
|
||||
Claims, Header, SigningAlgorithm, TokenData, TokenDecodeError, TokenScope, TokenType,
|
||||
TokenVerifyError, UnsafeClaims,
|
||||
@@ -164,9 +166,15 @@ pub fn verify_token_es256k(
|
||||
.decode(claims_b64)
|
||||
.map_err(|_| TokenVerifyError::Invalid("Base64 decode of claims failed"))?;
|
||||
|
||||
let claims: Claims = serde_json::from_slice(&claims_bytes)
|
||||
let mut claims: Claims = serde_json::from_slice(&claims_bytes)
|
||||
.map_err(|_| TokenVerifyError::Invalid("JSON decode of claims failed"))?;
|
||||
|
||||
if let Some(scope) = &claims.scope {
|
||||
claims.scope = Some(
|
||||
decode_scope(scope).map_err(|_| TokenVerifyError::Invalid("Invalid token scope"))?,
|
||||
);
|
||||
}
|
||||
|
||||
let now = Utc::now().timestamp();
|
||||
if claims.exp < now {
|
||||
return Err(TokenVerifyError::Expired);
|
||||
@@ -244,9 +252,13 @@ fn verify_token_hs256_internal(
|
||||
.decode(claims_b64)
|
||||
.context("Base64 decode of claims failed")?;
|
||||
|
||||
let claims: Claims =
|
||||
let mut claims: Claims =
|
||||
serde_json::from_slice(&claims_bytes).context("JSON decode of claims failed")?;
|
||||
|
||||
if let Some(scope) = &claims.scope {
|
||||
claims.scope = Some(decode_scope(scope).context("Invalid scope claim encoding")?);
|
||||
}
|
||||
|
||||
let now = Utc::now().timestamp();
|
||||
if claims.exp < now {
|
||||
return Err(anyhow!("Token expired"));
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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}");
|
||||
|
||||
@@ -445,6 +445,10 @@ pub struct ServerConfig {
|
||||
#[config(env = "ENABLE_PDS_HOSTED_DID_WEB", default = false)]
|
||||
pub enable_pds_hosted_did_web: bool,
|
||||
|
||||
/// iykyk!
|
||||
#[config(env = "RFC_MOO_COMPLIANCE", default = false)]
|
||||
pub rfc_moo_compliance: bool,
|
||||
|
||||
/// When set to true, skip age-assurance birthday prompt for all accounts.
|
||||
#[config(env = "PDS_AGE_ASSURANCE_OVERRIDE", default = false)]
|
||||
pub age_assurance_override: bool,
|
||||
@@ -831,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,
|
||||
}
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
@@ -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"] }
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -37,6 +37,7 @@ webauthn-rs = { workspace = true }
|
||||
|
||||
[dev-dependencies]
|
||||
async-trait = { workspace = true }
|
||||
tranquil-infra = { workspace = true, features = ["testing"] }
|
||||
|
||||
[features]
|
||||
bsky = []
|
||||
|
||||
@@ -10,6 +10,8 @@ pub struct ScopeInfo {
|
||||
pub display_name: String,
|
||||
pub granted: Option<bool>,
|
||||
pub restricted: bool,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub effective_scope: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -118,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
|
||||
@@ -186,20 +188,29 @@ pub async fn consent_get(
|
||||
|
||||
let grant_scope_str: Option<&str> =
|
||||
delegation_grant.as_ref().map(|g| g.granted_scopes.as_str());
|
||||
let is_restricted = |scope: &str| -> bool {
|
||||
grant_scope_str.is_some_and(|g| !tranquil_pds::delegation::grant_covers(g, scope))
|
||||
let coverage_of = |scope: &str| -> tranquil_pds::delegation::GrantCoverage {
|
||||
match grant_scope_str {
|
||||
Some(g) => tranquil_pds::delegation::grant_coverage(g, scope),
|
||||
None => tranquil_pds::delegation::GrantCoverage::Full,
|
||||
}
|
||||
};
|
||||
|
||||
let make_scope_info = |scope: &str| -> ScopeInfo {
|
||||
let (restricted, effective_scope) = match coverage_of(scope) {
|
||||
tranquil_pds::delegation::GrantCoverage::Full => (false, None),
|
||||
tranquil_pds::delegation::GrantCoverage::Narrowed(narrowed) => (false, Some(narrowed)),
|
||||
tranquil_pds::delegation::GrantCoverage::Withheld => (true, None),
|
||||
};
|
||||
let described = effective_scope.as_deref().unwrap_or(scope);
|
||||
let (category, required, description, display_name) =
|
||||
if let Some(def) = tranquil_pds::oauth::scopes::SCOPE_DEFINITIONS.get(scope) {
|
||||
let desc = if scope == "atproto" && has_granular_scopes {
|
||||
if let Some(def) = tranquil_pds::oauth::scopes::SCOPE_DEFINITIONS.get(described) {
|
||||
let desc = if described == "atproto" && has_granular_scopes {
|
||||
"AT Protocol baseline scope (permissions determined by selected options below)"
|
||||
.to_string()
|
||||
} else {
|
||||
def.description.to_string()
|
||||
};
|
||||
let name = if scope == "atproto" && has_granular_scopes {
|
||||
let name = if described == "atproto" && has_granular_scopes {
|
||||
"AT Protocol Access".to_string()
|
||||
} else {
|
||||
def.display_name.to_string()
|
||||
@@ -210,19 +221,19 @@ pub async fn consent_get(
|
||||
desc,
|
||||
name,
|
||||
)
|
||||
} else if scope.starts_with("ref:") {
|
||||
} else if described.starts_with("ref:") {
|
||||
(
|
||||
"Reference".to_string(),
|
||||
false,
|
||||
"Referenced scope".to_string(),
|
||||
scope.to_string(),
|
||||
described.to_string(),
|
||||
)
|
||||
} else {
|
||||
(
|
||||
"Other".to_string(),
|
||||
false,
|
||||
format!("Access to {}", scope),
|
||||
scope.to_string(),
|
||||
format!("Access to {}", described),
|
||||
described.to_string(),
|
||||
)
|
||||
};
|
||||
let granted = pref_map.get(scope).copied();
|
||||
@@ -233,7 +244,8 @@ pub async fn consent_get(
|
||||
description,
|
||||
display_name,
|
||||
granted,
|
||||
restricted: is_restricted(scope),
|
||||
restricted,
|
||||
effective_scope,
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -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 ¶ms.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
|
||||
|
||||
@@ -43,7 +43,8 @@ pub fn create_access_token_with_delegation(
|
||||
let issuer = format!("https://{}", pds_hostname);
|
||||
let now = Utc::now().timestamp();
|
||||
let exp = now + ACCESS_TOKEN_EXPIRY_SECONDS;
|
||||
let actual_scope = scope.unwrap_or("atproto");
|
||||
let actual_scope = tranquil_pds::auth::encode_scope(scope.unwrap_or("atproto"))
|
||||
.map_err(|_| OAuthError::InvalidScope("Scope too large".to_string()))?;
|
||||
let mut payload = json!({
|
||||
"iss": issuer,
|
||||
"sub": sub.as_str(),
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
})
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -43,14 +43,15 @@ pub use scope_verified::{
|
||||
pub use service::{ServiceTokenClaims, ServiceTokenError, ServiceTokenVerifier, is_service_token};
|
||||
|
||||
pub use tranquil_auth::{
|
||||
ActClaim, Claims, Header, SigningAlgorithm, TokenData, TokenDecodeError, TokenScope, TokenType,
|
||||
TokenVerifyError, TokenWithMetadata, TotpError, UnsafeClaims, create_access_token,
|
||||
create_access_token_hs256, create_access_token_hs256_with_metadata,
|
||||
create_access_token_with_delegation, create_access_token_with_jti,
|
||||
create_access_token_with_metadata, create_access_token_with_scope_metadata,
|
||||
create_refresh_token, create_refresh_token_hs256, create_refresh_token_hs256_with_metadata,
|
||||
create_refresh_token_with_jti, create_refresh_token_with_metadata, create_service_token,
|
||||
create_service_token_hs256, generate_backup_codes, generate_qr_png_base64,
|
||||
ActClaim, Claims, Header, ScopeDecodeError, ScopeEncodeError, SigningAlgorithm, TokenData,
|
||||
TokenDecodeError, TokenScope, TokenType, TokenVerifyError, TokenWithMetadata, TotpError,
|
||||
UnsafeClaims, create_access_token, create_access_token_hs256,
|
||||
create_access_token_hs256_with_metadata, create_access_token_with_delegation,
|
||||
create_access_token_with_jti, create_access_token_with_metadata,
|
||||
create_access_token_with_scope_metadata, create_refresh_token, create_refresh_token_hs256,
|
||||
create_refresh_token_hs256_with_metadata, create_refresh_token_with_jti,
|
||||
create_refresh_token_with_metadata, create_service_token, create_service_token_hs256,
|
||||
decode_scope, encode_scope, generate_backup_codes, generate_qr_png_base64,
|
||||
generate_totp_secret, generate_totp_uri, get_algorithm_from_token, get_did_from_token,
|
||||
get_jti_from_token, hash_backup_code, is_backup_code_format, verify_access_token,
|
||||
verify_access_token_hs256, verify_backup_code, verify_refresh_token,
|
||||
|
||||
@@ -67,6 +67,7 @@ impl WebAuthnConfig {
|
||||
.get_or_insert_with(AuthenticatorSelectionCriteria::default);
|
||||
sel.resident_key = Some(ResidentKeyRequirement::Required);
|
||||
sel.require_resident_key = true;
|
||||
ccr.public_key.hints = None;
|
||||
(ccr, state)
|
||||
})
|
||||
.map_err(|e| WebauthnError::RegistrationFailed(e.to_string()))
|
||||
@@ -88,6 +89,10 @@ impl WebAuthnConfig {
|
||||
) -> Result<(RequestChallengeResponse, SecurityKeyAuthentication), WebauthnError> {
|
||||
self.webauthn
|
||||
.start_securitykey_authentication(&credentials)
|
||||
.map(|(mut rcr, state)| {
|
||||
rcr.public_key.hints = None;
|
||||
(rcr, state)
|
||||
})
|
||||
.map_err(|e| WebauthnError::AuthenticationFailed(e.to_string()))
|
||||
}
|
||||
|
||||
|
||||
Vendored
+3
-1
@@ -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 +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::*;
|
||||
|
||||
@@ -5,14 +5,24 @@ pub use roles::{
|
||||
CanAddControllers, CanControlAccounts, verify_can_add_controllers, verify_can_control_accounts,
|
||||
};
|
||||
pub use scopes::{
|
||||
EDITOR_FULL_SCOPES, InvalidDelegationScopeError, OWNER_FULL_SCOPES, SCOPE_PRESETS, ScopePreset,
|
||||
ValidatedDelegationScope, grant_covers, intersect_scopes,
|
||||
EDITOR_FULL_SCOPES, GrantCoverage, InvalidDelegationScopeError, OWNER_FULL_SCOPES,
|
||||
SCOPE_PRESETS, ScopePreset, ValidatedDelegationScope, grant_coverage, intersect_scopes,
|
||||
};
|
||||
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,
|
||||
})
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use std::collections::HashSet;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use tranquil_scopes::{covers, parse_scope};
|
||||
use tranquil_scopes::{Coverage, ParsedScope, coverage, parse_scope};
|
||||
|
||||
pub use tranquil_db_traits::{
|
||||
DbScope as ValidatedDelegationScope, InvalidScopeError as InvalidDelegationScopeError,
|
||||
@@ -46,35 +46,51 @@ pub const SCOPE_PRESETS: &[ScopePreset] = &[
|
||||
},
|
||||
];
|
||||
|
||||
pub fn intersect_scopes(requested: &str, granted: &str) -> String {
|
||||
let requested_set: HashSet<&str> = requested.split_whitespace().collect();
|
||||
let granted_parsed: Vec<tranquil_scopes::ParsedScope> =
|
||||
granted.split_whitespace().map(parse_scope).collect();
|
||||
|
||||
let mut scopes: Vec<&str> = requested_set
|
||||
.iter()
|
||||
.filter(|requested_scope| {
|
||||
**requested_scope != "atproto" && any_granted_covers(requested_scope, &granted_parsed)
|
||||
})
|
||||
.copied()
|
||||
.chain(requested_set.contains("atproto").then_some("atproto"))
|
||||
.collect();
|
||||
scopes.sort();
|
||||
scopes.join(" ")
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GrantCoverage {
|
||||
Full,
|
||||
Narrowed(String),
|
||||
Withheld,
|
||||
}
|
||||
|
||||
pub fn grant_covers(granted: &str, scope: &str) -> bool {
|
||||
fn scope_coverage(granted: &[ParsedScope], scope: &str) -> GrantCoverage {
|
||||
if scope == "atproto" {
|
||||
return true;
|
||||
return GrantCoverage::Full;
|
||||
}
|
||||
|
||||
match coverage(granted, &parse_scope(scope)) {
|
||||
Coverage::Full => GrantCoverage::Full,
|
||||
Coverage::Narrowed(ParsedScope::Repo(repo)) => {
|
||||
GrantCoverage::Narrowed(repo.to_scope_string())
|
||||
}
|
||||
Coverage::Narrowed(_) => GrantCoverage::Full,
|
||||
Coverage::Withheld => GrantCoverage::Withheld,
|
||||
}
|
||||
let granted_parsed: Vec<tranquil_scopes::ParsedScope> =
|
||||
granted.split_whitespace().map(parse_scope).collect();
|
||||
any_granted_covers(scope, &granted_parsed)
|
||||
}
|
||||
|
||||
fn any_granted_covers(requested: &str, granted: &[tranquil_scopes::ParsedScope]) -> bool {
|
||||
let requested_parsed = parse_scope(requested);
|
||||
granted.iter().any(|g| covers(g, &requested_parsed))
|
||||
pub fn grant_coverage(granted: &str, scope: &str) -> GrantCoverage {
|
||||
scope_coverage(&parse_grant(granted), scope)
|
||||
}
|
||||
|
||||
fn parse_grant(granted: &str) -> Vec<ParsedScope> {
|
||||
granted.split_whitespace().map(parse_scope).collect()
|
||||
}
|
||||
|
||||
pub fn intersect_scopes(requested: &str, granted: &str) -> String {
|
||||
let granted_parsed = parse_grant(granted);
|
||||
|
||||
let scopes: BTreeSet<String> = requested
|
||||
.split_whitespace()
|
||||
.filter_map(
|
||||
|requested_scope| match scope_coverage(&granted_parsed, requested_scope) {
|
||||
GrantCoverage::Full => Some(requested_scope.to_string()),
|
||||
GrantCoverage::Narrowed(narrowed) => Some(narrowed),
|
||||
GrantCoverage::Withheld => None,
|
||||
},
|
||||
)
|
||||
.collect();
|
||||
|
||||
scopes.into_iter().collect::<Vec<String>>().join(" ")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -220,12 +236,33 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_intersect_partial_action_grant_drops_actionless_request() {
|
||||
fn test_intersect_partial_action_grant_narrows_actionless_request() {
|
||||
let result = intersect_scopes(
|
||||
"repo:app.bsky.feed.post",
|
||||
"repo:*?action=create&action=delete",
|
||||
);
|
||||
assert_eq!(result, "");
|
||||
assert_eq!(
|
||||
result,
|
||||
"repo:app.bsky.feed.post?action=create&action=delete"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_intersect_keeps_collapsed_request_under_split_action_grant() {
|
||||
assert_eq!(
|
||||
intersect_scopes(
|
||||
"repo:io.atcr.manifest?action=create&action=delete",
|
||||
EDITOR_FULL_SCOPES
|
||||
),
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
);
|
||||
assert_eq!(
|
||||
intersect_scopes(
|
||||
"repo:io.atcr.manifest?action=create&action=delete",
|
||||
"repo:*?action=create"
|
||||
),
|
||||
"repo:io.atcr.manifest?action=create"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -262,33 +299,35 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_grant_covers_matches_intersection() {
|
||||
fn test_grant_coverage_full_and_withheld() {
|
||||
let granted = "atproto repo:* blob:*/* account:*?action=manage";
|
||||
let intersected = intersect_scopes(
|
||||
"repo:app.bsky.feed.post?action=create identity:* account:*?action=manage",
|
||||
granted,
|
||||
);
|
||||
assert!(grant_covers(
|
||||
granted,
|
||||
"repo:app.bsky.feed.post?action=create"
|
||||
));
|
||||
assert!(grant_covers(granted, "account:*?action=manage"));
|
||||
assert!(!grant_covers(granted, "identity:*"));
|
||||
assert_eq!(grant_coverage(granted, "atproto"), GrantCoverage::Full);
|
||||
assert_eq!(
|
||||
grant_covers(granted, "identity:*"),
|
||||
intersected.contains("identity")
|
||||
grant_coverage(granted, "repo:app.bsky.feed.post?action=create"),
|
||||
GrantCoverage::Full
|
||||
);
|
||||
assert_eq!(
|
||||
grant_coverage(granted, "identity:*"),
|
||||
GrantCoverage::Withheld
|
||||
);
|
||||
assert_eq!(grant_coverage("", "identity:*"), GrantCoverage::Withheld);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_grant_covers_atproto_always_true() {
|
||||
assert!(grant_covers("", "atproto"));
|
||||
assert!(grant_covers("repo:*", "atproto"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_grant_covers_empty_grant_covers_nothing_else() {
|
||||
assert!(!grant_covers("", "repo:app.bsky.feed.post?action=create"));
|
||||
assert!(!grant_covers("", "identity:*"));
|
||||
fn test_grant_coverage_narrowed_when_grant_is_a_strict_action_subset() {
|
||||
assert_eq!(
|
||||
grant_coverage(
|
||||
EDITOR_FULL_SCOPES,
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
),
|
||||
GrantCoverage::Full
|
||||
);
|
||||
assert_eq!(
|
||||
grant_coverage(
|
||||
"atproto repo:*?action=create blob:*/*",
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
),
|
||||
GrantCoverage::Narrowed("repo:io.atcr.manifest?action=create".to_string())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+69
-197
@@ -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)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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, ¶ms, None)
|
||||
.send_with_dpop_retry(&signing_key, "POST", par_endpoint.as_str(), ¶ms, 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, ¶ms, None)
|
||||
.send_with_dpop_retry(
|
||||
&signing_key,
|
||||
"POST",
|
||||
meta.token_endpoint.as_str(),
|
||||
¶ms,
|
||||
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);
|
||||
|
||||
@@ -164,7 +164,9 @@ pub fn extract_oauth_token_info(token: &str) -> Result<OAuthTokenInfo, OAuthErro
|
||||
let scope = payload
|
||||
.get("scope")
|
||||
.and_then(|s| s.as_str())
|
||||
.map(|s| s.to_string());
|
||||
.map(crate::auth::decode_scope)
|
||||
.transpose()
|
||||
.map_err(|_| OAuthError::InvalidToken("Invalid scope claim encoding".to_string()))?;
|
||||
let controller_did = payload
|
||||
.get("act")
|
||||
.and_then(|a| a.get("sub"))
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,5 +1,6 @@
|
||||
pub use tranquil_types::*;
|
||||
|
||||
#[cfg(feature = "bsky")]
|
||||
use std::sync::LazyLock;
|
||||
|
||||
#[cfg(feature = "bsky")]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -15,6 +15,13 @@ use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||
const PERMISSION_SET_NSID: &str = "io.atcr.authFullApp";
|
||||
const PERMISSION_SET_GRANULAR_SCOPE: &str =
|
||||
"repo:io.atcr.manifest?action=create rpc:io.atcr.getManifest?aud=*";
|
||||
const PERMISSION_SET_MULTI_ACTION_SCOPE: &str =
|
||||
"repo:io.atcr.manifest?action=create&action=update&action=delete";
|
||||
const EDITOR_SET_NSID: &str = "io.atcr.authEditorApp";
|
||||
const SUBSET_SET_NSID: &str = "io.atcr.authSubsetApp";
|
||||
const PERMISSION_SET_CREATE_DELETE_SCOPE: &str =
|
||||
"repo:io.atcr.manifest?action=create&action=delete";
|
||||
const CREATE_ONLY_GRANT: &str = "atproto repo:*?action=create blob:*/*";
|
||||
|
||||
fn disable_rate_limiting_once() {
|
||||
static ONCE: std::sync::Once = std::sync::Once::new();
|
||||
@@ -105,6 +112,21 @@ async fn create_delegated_session_with_scope(
|
||||
handle_prefix: &str,
|
||||
redirect_uri: &str,
|
||||
scope: &str,
|
||||
) -> (DelegatedSession, Value, MockServer) {
|
||||
create_delegated_session_with_grant(
|
||||
handle_prefix,
|
||||
redirect_uri,
|
||||
scope,
|
||||
tranquil_pds::delegation::OWNER_FULL_SCOPES,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn create_delegated_session_with_grant(
|
||||
handle_prefix: &str,
|
||||
redirect_uri: &str,
|
||||
scope: &str,
|
||||
controller_scopes: &str,
|
||||
) -> (DelegatedSession, Value, MockServer) {
|
||||
let url = base_url().await;
|
||||
disable_rate_limiting_once();
|
||||
@@ -119,7 +141,7 @@ async fn create_delegated_session_with_scope(
|
||||
.bearer_auth(&controller_jwt)
|
||||
.json(&json!({
|
||||
"handle": delegated_handle,
|
||||
"controllerScopes": tranquil_pds::delegation::OWNER_FULL_SCOPES
|
||||
"controllerScopes": controller_scopes
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -375,6 +397,131 @@ async fn test_delegated_include_scope_shows_granular_on_consent() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_delegated_editor_grant_keeps_collapsed_permission_set() {
|
||||
seed_permission_set(EDITOR_SET_NSID, PERMISSION_SET_MULTI_ACTION_SCOPE).await;
|
||||
|
||||
let scope = format!("atproto include:{}", EDITOR_SET_NSID);
|
||||
let (session, consent_body, _mock) = create_delegated_session_with_grant(
|
||||
"pse",
|
||||
"https://example.com/permset-editor-callback",
|
||||
&scope,
|
||||
tranquil_pds::delegation::EDITOR_FULL_SCOPES,
|
||||
)
|
||||
.await;
|
||||
|
||||
let set_entry = consent_body["permission_sets"]
|
||||
.as_array()
|
||||
.expect("consent response should have a permission_sets array")
|
||||
.iter()
|
||||
.find(|s| s["nsid"].as_str() == Some(EDITOR_SET_NSID))
|
||||
.unwrap_or_else(|| {
|
||||
panic!(
|
||||
"permission_sets should contain an entry for nsid '{}'. Got: {:?}",
|
||||
EDITOR_SET_NSID, consent_body
|
||||
)
|
||||
});
|
||||
assert_eq!(
|
||||
set_entry["restricted"].as_bool(),
|
||||
Some(false),
|
||||
"an editor grant spells its actions as separate tokens, but it still permits every \
|
||||
action in the collapsed set, so the set must not be marked restricted. Got: {:?}",
|
||||
set_entry
|
||||
);
|
||||
|
||||
let payload = decode_jwt_payload(&session.access_token);
|
||||
let jwt_scope = tranquil_pds::auth::decode_scope(
|
||||
payload["scope"]
|
||||
.as_str()
|
||||
.expect("access token JWT should have a scope claim"),
|
||||
)
|
||||
.expect("JWT scope claim should decode");
|
||||
assert!(
|
||||
jwt_scope.contains(PERMISSION_SET_MULTI_ACTION_SCOPE),
|
||||
"delegated intersection must narrow the collapsed repo scope rather than discard it, \
|
||||
expected '{}' in decoded scope, got: {}",
|
||||
PERMISSION_SET_MULTI_ACTION_SCOPE,
|
||||
jwt_scope
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_delegated_consent_shows_the_scope_the_token_will_carry() {
|
||||
seed_permission_set(SUBSET_SET_NSID, PERMISSION_SET_CREATE_DELETE_SCOPE).await;
|
||||
|
||||
let scope = format!("atproto include:{}", SUBSET_SET_NSID);
|
||||
let (session, consent_body, _mock) = create_delegated_session_with_grant(
|
||||
"pss",
|
||||
"https://example.com/permset-subset-callback",
|
||||
&scope,
|
||||
CREATE_ONLY_GRANT,
|
||||
)
|
||||
.await;
|
||||
|
||||
let set_entry = consent_body["permission_sets"]
|
||||
.as_array()
|
||||
.expect("consent response should have a permission_sets array")
|
||||
.iter()
|
||||
.find(|s| s["nsid"].as_str() == Some(SUBSET_SET_NSID))
|
||||
.unwrap_or_else(|| {
|
||||
panic!(
|
||||
"permission_sets should contain an entry for nsid '{}'. Got: {:?}",
|
||||
SUBSET_SET_NSID, consent_body
|
||||
)
|
||||
});
|
||||
assert_eq!(
|
||||
set_entry["restricted"].as_bool(),
|
||||
Some(false),
|
||||
"the create action is still granted, so the set stays approvable. Got: {:?}",
|
||||
set_entry
|
||||
);
|
||||
|
||||
let repo = set_entry["expanded"]
|
||||
.as_array()
|
||||
.expect("permission_sets entry should have an expanded array")
|
||||
.iter()
|
||||
.find(|s| s["scope"].as_str() == Some(PERMISSION_SET_CREATE_DELETE_SCOPE))
|
||||
.unwrap_or_else(|| {
|
||||
panic!(
|
||||
"expanded[] should list the requested scope '{}'. Got: {:?}",
|
||||
PERMISSION_SET_CREATE_DELETE_SCOPE, set_entry
|
||||
)
|
||||
});
|
||||
assert_eq!(
|
||||
repo["restricted"].as_bool(),
|
||||
Some(false),
|
||||
"a partially-covered scope is neither fully granted nor withheld. Got: {:?}",
|
||||
repo
|
||||
);
|
||||
let effective_scope = repo["effective_scope"].as_str().unwrap_or_else(|| {
|
||||
panic!(
|
||||
"a scope the grant narrows must report the actions it actually confers. Got: {:?}",
|
||||
repo
|
||||
)
|
||||
});
|
||||
assert_eq!(effective_scope, "repo:io.atcr.manifest?action=create");
|
||||
|
||||
let payload = decode_jwt_payload(&session.access_token);
|
||||
let jwt_scope = tranquil_pds::auth::decode_scope(
|
||||
payload["scope"]
|
||||
.as_str()
|
||||
.expect("access token JWT should have a scope claim"),
|
||||
)
|
||||
.expect("JWT scope claim should decode");
|
||||
assert!(
|
||||
jwt_scope.split_whitespace().any(|s| s == effective_scope),
|
||||
"the consent screen must show the scope the token carries, expected '{}' in decoded \
|
||||
scope, got: {}",
|
||||
effective_scope,
|
||||
jwt_scope
|
||||
);
|
||||
assert!(
|
||||
!jwt_scope.contains("action=delete"),
|
||||
"the grant confers no delete action, so the token must not carry one, got: {}",
|
||||
jwt_scope
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_grant_row_keeps_include_jwt_carries_expanded() {
|
||||
seed_permission_set(PERMISSION_SET_NSID, PERMISSION_SET_GRANULAR_SCOPE).await;
|
||||
@@ -505,6 +652,129 @@ async fn test_enforcement_uses_expanded_jwt_scope() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_long_expanded_scope_is_compressed_in_jwt() {
|
||||
const BIG_SET_NSID: &str = "io.atcr.authBigApp";
|
||||
let collections = [
|
||||
"io.atcr.manifest",
|
||||
"io.atcr.sailor.star",
|
||||
"io.atcr.tag",
|
||||
"io.atcr.blueprint",
|
||||
"io.atcr.artifact",
|
||||
"io.atcr.channel.read",
|
||||
];
|
||||
let granular_scope = collections
|
||||
.iter()
|
||||
.map(|coll| format!("repo:{}?action=create&action=delete", coll))
|
||||
.collect::<Vec<_>>()
|
||||
.join(" ");
|
||||
seed_permission_set(BIG_SET_NSID, &granular_scope).await;
|
||||
|
||||
let scope = format!("atproto include:{}", BIG_SET_NSID);
|
||||
let (session, _consent_body, _mock) = create_delegated_session_with_scope(
|
||||
"psc",
|
||||
"https://example.com/permset-compress-callback",
|
||||
&scope,
|
||||
)
|
||||
.await;
|
||||
|
||||
let payload = decode_jwt_payload(&session.access_token);
|
||||
let jwt_scope = payload["scope"]
|
||||
.as_str()
|
||||
.expect("access token JWT should have a scope claim");
|
||||
assert!(
|
||||
jwt_scope.starts_with("$br$"),
|
||||
"an expanded scope this long should be compressed in the JWT claim, got: {}",
|
||||
jwt_scope
|
||||
);
|
||||
|
||||
let decoded = tranquil_pds::auth::decode_scope(jwt_scope).expect("scope claim should decode");
|
||||
for coll in collections {
|
||||
assert!(
|
||||
decoded.contains(&format!("repo:{}?action=create", coll)),
|
||||
"decoded scope should carry {}, got: {}",
|
||||
coll,
|
||||
decoded
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
!decoded.contains("include:"),
|
||||
"decoded scope should not contain the raw include: token, got: {}",
|
||||
decoded
|
||||
);
|
||||
|
||||
let url = base_url().await;
|
||||
let http_client = client();
|
||||
|
||||
let introspect_res = http_client
|
||||
.post(format!("{}/oauth/introspect", url))
|
||||
.form(&[("token", session.access_token.as_str())])
|
||||
.send()
|
||||
.await
|
||||
.expect("introspect request failed");
|
||||
assert_eq!(introspect_res.status(), StatusCode::OK);
|
||||
let introspect_body: Value = introspect_res.json().await.unwrap();
|
||||
let introspect_scope = introspect_body["scope"]
|
||||
.as_str()
|
||||
.expect("introspect response should have a scope string");
|
||||
assert_eq!(
|
||||
introspect_scope, decoded,
|
||||
"introspect should report the decoded scope"
|
||||
);
|
||||
|
||||
let collection = collections[0];
|
||||
let create_res = http_client
|
||||
.post(format!("{}/xrpc/com.atproto.repo.createRecord", url))
|
||||
.bearer_auth(&session.access_token)
|
||||
.json(&json!({
|
||||
"repo": session.delegated_did,
|
||||
"collection": collection,
|
||||
"validate": false,
|
||||
"record": {
|
||||
"$type": collection,
|
||||
"note": "compressed scope enforcement test",
|
||||
"createdAt": Utc::now().to_rfc3339()
|
||||
}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("createRecord request failed");
|
||||
assert_ne!(
|
||||
create_res.status(),
|
||||
StatusCode::FORBIDDEN,
|
||||
"a compressed scope claim must still authorize the collections it covers. Got body: {:?}",
|
||||
create_res.text().await
|
||||
);
|
||||
|
||||
let refresh_res = http_client
|
||||
.post(format!("{}/oauth/token", url))
|
||||
.form(&[
|
||||
("grant_type", "refresh_token"),
|
||||
("refresh_token", session.refresh_token.as_str()),
|
||||
("client_id", session.client_id.as_str()),
|
||||
])
|
||||
.send()
|
||||
.await
|
||||
.expect("Refresh request failed");
|
||||
assert_eq!(refresh_res.status(), StatusCode::OK);
|
||||
let refresh_body: Value = refresh_res.json().await.unwrap();
|
||||
let refreshed_token = refresh_body["access_token"].as_str().unwrap();
|
||||
let refreshed_claim = decode_jwt_payload(refreshed_token)["scope"]
|
||||
.as_str()
|
||||
.expect("refreshed JWT should have a scope claim")
|
||||
.to_string();
|
||||
assert!(
|
||||
refreshed_claim.starts_with("$br$"),
|
||||
"refreshed claim should also be compressed, got: {}",
|
||||
refreshed_claim
|
||||
);
|
||||
assert_eq!(
|
||||
tranquil_pds::auth::decode_scope(&refreshed_claim).expect("refreshed scope should decode"),
|
||||
decoded,
|
||||
"refresh must yield a byte-identical decoded scope"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_consent_post_errors_when_set_unresolvable() {
|
||||
const UNRESOLVABLE_NSID: &str = "io.atcr.authUnresolvableSet";
|
||||
|
||||
@@ -254,15 +254,18 @@ fn test_scope_with_multiple_params() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_scope_invalid_action_ignored() {
|
||||
let scope = parse_scope("repo:*?action=invalid");
|
||||
if let ParsedScope::Repo(repo) = scope {
|
||||
assert!(repo.actions.contains(&RepoAction::Create));
|
||||
assert!(repo.actions.contains(&RepoAction::Update));
|
||||
assert!(repo.actions.contains(&RepoAction::Delete));
|
||||
} else {
|
||||
panic!("Expected Repo scope");
|
||||
}
|
||||
fn test_scope_invalid_action_rejects_whole_scope() {
|
||||
assert!(
|
||||
matches!(
|
||||
parse_scope("repo:*?action=invalid"),
|
||||
ParsedScope::Unknown(_)
|
||||
),
|
||||
"an unrecognized action must not fall back to granting every action"
|
||||
);
|
||||
assert!(matches!(
|
||||
parse_scope("repo:*?action=create&action=invalid"),
|
||||
ParsedScope::Unknown(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
use crate::parser::{
|
||||
AccountAction, AccountAttr, AccountScope, BlobScope, IdentityAttr, IdentityScope, ParsedScope,
|
||||
RepoScope, RpcScope,
|
||||
RepoAction, RepoScope, RpcScope,
|
||||
};
|
||||
use std::collections::HashSet;
|
||||
|
||||
pub fn covers(granted: &ParsedScope, requested: &ParsedScope) -> bool {
|
||||
use ParsedScope::*;
|
||||
@@ -21,8 +22,8 @@ pub fn covers(granted: &ParsedScope, requested: &ParsedScope) -> bool {
|
||||
}
|
||||
}
|
||||
|
||||
fn repo_covers(g: &RepoScope, r: &RepoScope) -> bool {
|
||||
let collection_ok = match &g.collection {
|
||||
fn repo_collection_covers(g: &RepoScope, r: &RepoScope) -> bool {
|
||||
match &g.collection {
|
||||
None => true,
|
||||
Some(gc) => match &r.collection {
|
||||
None => false,
|
||||
@@ -33,8 +34,54 @@ fn repo_covers(g: &RepoScope, r: &RepoScope) -> bool {
|
||||
None => gc == rc,
|
||||
},
|
||||
},
|
||||
};
|
||||
collection_ok && r.actions.is_subset(&g.actions)
|
||||
}
|
||||
}
|
||||
|
||||
fn repo_covers(g: &RepoScope, r: &RepoScope) -> bool {
|
||||
repo_collection_covers(g, r) && r.actions.is_subset(&g.actions)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum Coverage {
|
||||
Full,
|
||||
Narrowed(ParsedScope),
|
||||
Withheld,
|
||||
}
|
||||
|
||||
pub fn coverage(granted: &[ParsedScope], requested: &ParsedScope) -> Coverage {
|
||||
if let ParsedScope::Repo(r) = requested {
|
||||
let actions: HashSet<RepoAction> = granted
|
||||
.iter()
|
||||
.filter_map(|g| match g {
|
||||
ParsedScope::Repo(g) if repo_collection_covers(g, r) => Some(&g.actions),
|
||||
_ => None,
|
||||
})
|
||||
.flat_map(|granted_actions| granted_actions.intersection(&r.actions).copied())
|
||||
.collect();
|
||||
|
||||
return match actions.len() {
|
||||
0 => Coverage::Withheld,
|
||||
_ if actions == r.actions => Coverage::Full,
|
||||
_ => Coverage::Narrowed(ParsedScope::Repo(RepoScope {
|
||||
collection: r.collection.clone(),
|
||||
actions,
|
||||
})),
|
||||
};
|
||||
}
|
||||
|
||||
if granted.iter().any(|g| covers(g, requested)) {
|
||||
Coverage::Full
|
||||
} else {
|
||||
Coverage::Withheld
|
||||
}
|
||||
}
|
||||
|
||||
pub fn narrow(granted: &[ParsedScope], requested: &ParsedScope) -> Option<ParsedScope> {
|
||||
match coverage(granted, requested) {
|
||||
Coverage::Full => Some(requested.clone()),
|
||||
Coverage::Narrowed(scope) => Some(scope),
|
||||
Coverage::Withheld => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn blob_covers(g: &BlobScope, r: &BlobScope) -> bool {
|
||||
@@ -74,13 +121,33 @@ fn identity_covers(g: &IdentityScope, r: &IdentityScope) -> bool {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::covers;
|
||||
use crate::parser::parse_scope;
|
||||
use super::{Coverage, coverage, covers, narrow};
|
||||
use crate::parser::{ParsedScope, parse_scope};
|
||||
|
||||
fn c(granted: &str, requested: &str) -> bool {
|
||||
covers(&parse_scope(granted), &parse_scope(requested))
|
||||
}
|
||||
|
||||
fn narrowed(granted: &str, requested: &str) -> Option<String> {
|
||||
let granted: Vec<ParsedScope> = granted.split_whitespace().map(parse_scope).collect();
|
||||
|
||||
match narrow(&granted, &parse_scope(requested)) {
|
||||
Some(ParsedScope::Repo(repo)) => Some(repo.to_scope_string()),
|
||||
Some(_) => Some(requested.to_string()),
|
||||
None => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn covered(granted: &str, requested: &str) -> Coverage {
|
||||
let granted: Vec<ParsedScope> = granted.split_whitespace().map(parse_scope).collect();
|
||||
|
||||
coverage(&granted, &parse_scope(requested))
|
||||
}
|
||||
|
||||
fn narrowed_to(scope: &str) -> Coverage {
|
||||
Coverage::Narrowed(parse_scope(scope))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repo_wildcard_covers_specific() {
|
||||
assert!(c("repo:*", "repo:app.bsky.feed.post"));
|
||||
@@ -193,4 +260,66 @@ mod tests {
|
||||
assert!(c("weird:token", "weird:token"));
|
||||
assert!(!c("weird:token", "other:token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn narrow_intersects_repo_actions() {
|
||||
assert_eq!(
|
||||
narrowed(
|
||||
"repo:*?action=create repo:*?action=update repo:*?action=delete",
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
),
|
||||
Some("repo:io.atcr.manifest?action=create&action=delete".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
narrowed(
|
||||
"repo:*?action=create",
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
),
|
||||
Some("repo:io.atcr.manifest?action=create".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
narrowed(
|
||||
"repo:*?action=create",
|
||||
"repo:io.atcr.manifest?action=delete"
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(narrowed("repo:app.bsky.*?action=create", "repo:*"), None);
|
||||
assert_eq!(
|
||||
narrowed("identity:*", "identity:handle"),
|
||||
Some("identity:handle".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn coverage_distinguishes_full_from_narrowed_repo_actions() {
|
||||
assert_eq!(
|
||||
covered(
|
||||
"repo:*?action=create repo:*?action=update repo:*?action=delete",
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
),
|
||||
Coverage::Full
|
||||
);
|
||||
assert_eq!(
|
||||
covered(
|
||||
"repo:*?action=create",
|
||||
"repo:io.atcr.manifest?action=create&action=delete"
|
||||
),
|
||||
narrowed_to("repo:io.atcr.manifest?action=create")
|
||||
);
|
||||
assert_eq!(
|
||||
covered(
|
||||
"repo:*?action=create&action=delete",
|
||||
"repo:io.atcr.manifest"
|
||||
),
|
||||
narrowed_to("repo:io.atcr.manifest?action=create&action=delete")
|
||||
);
|
||||
assert_eq!(
|
||||
covered(
|
||||
"repo:*?action=create",
|
||||
"repo:io.atcr.manifest?action=delete"
|
||||
),
|
||||
Coverage::Withheld
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ mod parser;
|
||||
mod permission_set;
|
||||
mod permissions;
|
||||
|
||||
pub use coverage::covers;
|
||||
pub use coverage::{Coverage, coverage, covers, narrow};
|
||||
pub use definitions::{
|
||||
SCOPE_DEFINITIONS, ScopeCategory, ScopeDefinition, format_scope_for_display,
|
||||
get_required_scopes, get_scope_definition, is_valid_scope,
|
||||
|
||||
@@ -28,7 +28,22 @@ pub struct RepoScope {
|
||||
pub actions: HashSet<RepoAction>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
impl RepoScope {
|
||||
pub fn to_scope_string(&self) -> String {
|
||||
let mut actions: Vec<RepoAction> = self.actions.iter().copied().collect();
|
||||
actions.sort();
|
||||
|
||||
let rendered: Vec<&str> = actions.iter().map(RepoAction::as_str).collect();
|
||||
|
||||
format!(
|
||||
"repo:{}?action={}",
|
||||
self.collection.as_deref().unwrap_or("*"),
|
||||
rendered.join("&action=")
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum RepoAction {
|
||||
Create,
|
||||
@@ -37,6 +52,8 @@ pub enum RepoAction {
|
||||
}
|
||||
|
||||
impl RepoAction {
|
||||
pub const ALL: [RepoAction; 3] = [Self::Create, Self::Update, Self::Delete];
|
||||
|
||||
pub fn parse_str(s: &str) -> Option<Self> {
|
||||
match s {
|
||||
"create" => Some(Self::Create),
|
||||
@@ -45,6 +62,14 @@ impl RepoAction {
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Create => "create",
|
||||
Self::Update => "update",
|
||||
Self::Delete => "delete",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -150,6 +175,14 @@ fn parse_query_params(query: &str) -> HashMap<String, Vec<String>> {
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_repo_actions(params: &HashMap<String, Vec<String>>) -> Option<HashSet<RepoAction>> {
|
||||
match params.get("action") {
|
||||
None => Some(RepoAction::ALL.into_iter().collect()),
|
||||
Some(values) if values.is_empty() => None,
|
||||
Some(values) => values.iter().map(|s| RepoAction::parse_str(s)).collect(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse_scope(scope: &str) -> ParsedScope {
|
||||
match scope {
|
||||
"atproto" => return ParsedScope::Atproto,
|
||||
@@ -169,20 +202,9 @@ pub fn parse_scope(scope: &str) -> ParsedScope {
|
||||
Some(rest.to_string())
|
||||
};
|
||||
|
||||
let actions: HashSet<RepoAction> = params
|
||||
.get("action")
|
||||
.map(|action_values| {
|
||||
action_values
|
||||
.iter()
|
||||
.filter_map(|s| RepoAction::parse_str(s))
|
||||
.collect()
|
||||
})
|
||||
.filter(|set: &HashSet<RepoAction>| !set.is_empty())
|
||||
.unwrap_or_else(|| {
|
||||
[RepoAction::Create, RepoAction::Update, RepoAction::Delete]
|
||||
.into_iter()
|
||||
.collect()
|
||||
});
|
||||
let Some(actions) = parse_repo_actions(¶ms) else {
|
||||
return ParsedScope::Unknown(scope.to_string());
|
||||
};
|
||||
|
||||
return ParsedScope::Repo(RepoScope {
|
||||
collection,
|
||||
@@ -191,20 +213,10 @@ pub fn parse_scope(scope: &str) -> ParsedScope {
|
||||
}
|
||||
|
||||
if base == "repo" {
|
||||
let actions: HashSet<RepoAction> = params
|
||||
.get("action")
|
||||
.map(|action_values| {
|
||||
action_values
|
||||
.iter()
|
||||
.filter_map(|s| RepoAction::parse_str(s))
|
||||
.collect()
|
||||
})
|
||||
.filter(|set: &HashSet<RepoAction>| !set.is_empty())
|
||||
.unwrap_or_else(|| {
|
||||
[RepoAction::Create, RepoAction::Update, RepoAction::Delete]
|
||||
.into_iter()
|
||||
.collect()
|
||||
});
|
||||
let Some(actions) = parse_repo_actions(¶ms) else {
|
||||
return ParsedScope::Unknown(scope.to_string());
|
||||
};
|
||||
|
||||
return ParsedScope::Repo(RepoScope {
|
||||
collection: None,
|
||||
actions,
|
||||
@@ -340,6 +352,22 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_repo_unrecognized_action_is_not_a_repo_scope() {
|
||||
assert!(matches!(
|
||||
parse_scope("repo:app.bsky.feed.post?action=read"),
|
||||
ParsedScope::Unknown(_)
|
||||
));
|
||||
assert!(matches!(
|
||||
parse_scope("repo:app.bsky.feed.post?action="),
|
||||
ParsedScope::Unknown(_)
|
||||
));
|
||||
assert!(matches!(
|
||||
parse_scope("repo?action=read"),
|
||||
ParsedScope::Unknown(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_blob_wildcard() {
|
||||
let scope = parse_scope("blob:*/*");
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
use crate::parser::RepoAction;
|
||||
use hickory_resolver::TokioAsyncResolver;
|
||||
use hickory_resolver::config::{ResolverConfig, ResolverOpts};
|
||||
use reqwest::Client;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use tracing::debug;
|
||||
use std::collections::{BTreeMap, BTreeSet, HashMap};
|
||||
use tracing::{debug, warn};
|
||||
use tranquil_types::{Did, Nsid};
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
@@ -332,33 +333,49 @@ fn is_under_authority(target_nsid: &str, authority: &str) -> bool {
|
||||
.is_some_and(|c| c == '.')
|
||||
}
|
||||
|
||||
const DEFAULT_ACTIONS: &[&str] = &["create", "update", "delete"];
|
||||
fn parse_permission_actions(actions: Option<&Vec<String>>) -> Option<BTreeSet<RepoAction>> {
|
||||
match actions {
|
||||
None => Some(RepoAction::ALL.into_iter().collect()),
|
||||
Some(values) => values
|
||||
.iter()
|
||||
.map(|value| {
|
||||
let parsed = RepoAction::parse_str(value);
|
||||
if parsed.is_none() {
|
||||
warn!(
|
||||
action = %value,
|
||||
"skipping permission entry with unrecognized repo action"
|
||||
);
|
||||
}
|
||||
parsed
|
||||
})
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
|
||||
fn build_expanded_scopes(
|
||||
permissions: &[PermissionEntry],
|
||||
default_aud: Option<&str>,
|
||||
namespace_authority: &str,
|
||||
) -> String {
|
||||
let mut scopes: Vec<String> = Vec::new();
|
||||
let mut ungrouped_repo_scopes: BTreeMap<String, BTreeSet<RepoAction>> = BTreeMap::new();
|
||||
let mut rpc_scopes: Vec<String> = Vec::new();
|
||||
|
||||
permissions
|
||||
.iter()
|
||||
.for_each(|perm| match perm.resource.as_str() {
|
||||
"repo" => {
|
||||
if let Some(collections) = &perm.collection {
|
||||
let actions: Vec<&str> = perm
|
||||
.action
|
||||
.as_ref()
|
||||
.map(|a| a.iter().map(String::as_str).collect())
|
||||
.unwrap_or_else(|| DEFAULT_ACTIONS.to_vec());
|
||||
|
||||
if let Some(collections) = &perm.collection
|
||||
&& let Some(actions) = parse_permission_actions(perm.action.as_ref())
|
||||
&& !actions.is_empty()
|
||||
{
|
||||
collections
|
||||
.iter()
|
||||
.filter(|coll| is_under_authority(coll, namespace_authority))
|
||||
.for_each(|coll| {
|
||||
actions.iter().for_each(|action| {
|
||||
scopes.push(format!("repo:{}?action={}", coll, action));
|
||||
});
|
||||
ungrouped_repo_scopes
|
||||
.entry(coll.to_string())
|
||||
.or_default()
|
||||
.extend(actions.iter().copied());
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -373,14 +390,31 @@ fn build_expanded_scopes(
|
||||
Some(aud) => format!("rpc:{}?aud={}", lxm, aud),
|
||||
None => format!("rpc:{}", lxm),
|
||||
};
|
||||
scopes.push(scope);
|
||||
|
||||
if !rpc_scopes.contains(&scope) {
|
||||
rpc_scopes.push(scope);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
});
|
||||
|
||||
scopes.join(" ")
|
||||
let grouped_repo_scopes: Vec<String> = ungrouped_repo_scopes
|
||||
.iter()
|
||||
.map(|(repo, actions)| {
|
||||
let rendered: Vec<&str> = actions.iter().map(RepoAction::as_str).collect();
|
||||
|
||||
format!("repo:{}?action={}", repo, rendered.join("&action="))
|
||||
})
|
||||
.collect();
|
||||
|
||||
let combined_repo_scopes = grouped_repo_scopes.join(" ");
|
||||
let combined_rpc_scopes = rpc_scopes.join(" ");
|
||||
|
||||
format!("{} {}", combined_repo_scopes, combined_rpc_scopes)
|
||||
.trim()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -459,10 +493,11 @@ mod tests {
|
||||
}];
|
||||
|
||||
let expanded = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert!(expanded.contains("repo:io.atcr.manifest?action=create"));
|
||||
assert!(expanded.contains("repo:io.atcr.manifest?action=delete"));
|
||||
assert!(expanded.contains("repo:io.atcr.sailor.star?action=create"));
|
||||
assert!(!expanded.contains("app.bsky.feed.post"));
|
||||
assert_eq!(
|
||||
expanded,
|
||||
"repo:io.atcr.manifest?action=create&action=delete \
|
||||
repo:io.atcr.sailor.star?action=create&action=delete"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -476,9 +511,138 @@ mod tests {
|
||||
}];
|
||||
|
||||
let expanded = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert!(expanded.contains("repo:io.atcr.manifest?action=create"));
|
||||
assert!(expanded.contains("repo:io.atcr.manifest?action=update"));
|
||||
assert!(expanded.contains("repo:io.atcr.manifest?action=delete"));
|
||||
assert_eq!(
|
||||
expanded,
|
||||
"repo:io.atcr.manifest?action=create&action=update&action=delete"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_expanded_scopes_repo_omitted_action_grants_all() {
|
||||
let permissions = vec![PermissionEntry {
|
||||
resource: "repo".to_string(),
|
||||
action: None,
|
||||
collection: Some(vec!["io.atcr.manifest".to_string()]),
|
||||
lxm: None,
|
||||
aud: None,
|
||||
}];
|
||||
|
||||
let expanded = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert_eq!(
|
||||
expanded, "repo:io.atcr.manifest?action=create&action=update&action=delete",
|
||||
"an omitted action list means all actions"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_expanded_scopes_repo_empty_action_list_skips_entry() {
|
||||
let permissions = vec![PermissionEntry {
|
||||
resource: "repo".to_string(),
|
||||
action: Some(vec![]),
|
||||
collection: Some(vec!["io.atcr.manifest".to_string()]),
|
||||
lxm: None,
|
||||
aud: None,
|
||||
}];
|
||||
|
||||
let expanded = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert!(
|
||||
expanded.is_empty(),
|
||||
"an explicitly empty action list is invalid, so the entry is skipped rather \
|
||||
than expanded to all actions or emitted as a bare `?action=`, got: {expanded}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_expanded_scopes_is_deterministic() {
|
||||
let permissions = vec![
|
||||
PermissionEntry {
|
||||
resource: "repo".to_string(),
|
||||
action: Some(vec!["create".to_string()]),
|
||||
collection: Some(vec![
|
||||
"io.atcr.sailor.star".to_string(),
|
||||
"io.atcr.manifest".to_string(),
|
||||
"io.atcr.blob".to_string(),
|
||||
]),
|
||||
lxm: None,
|
||||
aud: None,
|
||||
},
|
||||
PermissionEntry {
|
||||
resource: "rpc".to_string(),
|
||||
action: None,
|
||||
collection: None,
|
||||
lxm: Some(vec![
|
||||
"io.atcr.getManifest".to_string(),
|
||||
"io.atcr.listTags".to_string(),
|
||||
]),
|
||||
aud: Some("*".to_string()),
|
||||
},
|
||||
];
|
||||
|
||||
let first = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert_eq!(
|
||||
first,
|
||||
"repo:io.atcr.blob?action=create repo:io.atcr.manifest?action=create \
|
||||
repo:io.atcr.sailor.star?action=create \
|
||||
rpc:io.atcr.getManifest?aud=* rpc:io.atcr.listTags?aud=*"
|
||||
);
|
||||
|
||||
for _ in 0..16 {
|
||||
assert_eq!(build_expanded_scopes(&permissions, None, "io.atcr"), first);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_expanded_scopes_dedupes_and_canonicalizes_actions() {
|
||||
let permissions = vec![
|
||||
PermissionEntry {
|
||||
resource: "repo".to_string(),
|
||||
action: Some(vec!["delete".to_string(), "create".to_string()]),
|
||||
collection: Some(vec!["io.atcr.manifest".to_string()]),
|
||||
lxm: None,
|
||||
aud: None,
|
||||
},
|
||||
PermissionEntry {
|
||||
resource: "repo".to_string(),
|
||||
action: Some(vec!["create".to_string(), "update".to_string()]),
|
||||
collection: Some(vec!["io.atcr.manifest".to_string()]),
|
||||
lxm: None,
|
||||
aud: None,
|
||||
},
|
||||
PermissionEntry {
|
||||
resource: "rpc".to_string(),
|
||||
action: None,
|
||||
collection: None,
|
||||
lxm: Some(vec![
|
||||
"io.atcr.getManifest".to_string(),
|
||||
"io.atcr.getManifest".to_string(),
|
||||
]),
|
||||
aud: Some("*".to_string()),
|
||||
},
|
||||
];
|
||||
|
||||
let expanded = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert_eq!(
|
||||
expanded,
|
||||
"repo:io.atcr.manifest?action=create&action=update&action=delete \
|
||||
rpc:io.atcr.getManifest?aud=*"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_expanded_scopes_repo_unrecognized_action_skips_entry() {
|
||||
let permissions = vec![PermissionEntry {
|
||||
resource: "repo".to_string(),
|
||||
action: Some(vec!["read".to_string()]),
|
||||
collection: Some(vec!["io.atcr.manifest".to_string()]),
|
||||
lxm: None,
|
||||
aud: None,
|
||||
}];
|
||||
|
||||
let expanded = build_expanded_scopes(&permissions, None, "io.atcr");
|
||||
assert!(
|
||||
expanded.is_empty(),
|
||||
"an unrecognized repo action must not expand to all actions, got: {expanded}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -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 }
|
||||
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -10,7 +10,7 @@ Environment=SERVER_PORT=3000
|
||||
Volume=/srv/tranquil-pds/config/config.toml:/etc/tranquil-pds/config.toml:ro,Z
|
||||
Volume=/srv/tranquil-pds/blobs:/var/lib/tranquil-pds/blobs:Z
|
||||
Volume=/srv/tranquil-pds/store:/var/lib/tranquil-pds/store:Z
|
||||
HealthCmd=wget -q --spider http://localhost:3000/xrpc/_health
|
||||
HealthCmd=["/usr/local/bin/tranquil-pds", "healthcheck"]
|
||||
HealthInterval=30s
|
||||
HealthTimeout=10s
|
||||
HealthRetries=3
|
||||
|
||||
+11
-2
@@ -34,6 +34,13 @@
|
||||
# Default value: false
|
||||
#enable_pds_hosted_did_web = false
|
||||
|
||||
# iykyk!
|
||||
#
|
||||
# Can also be specified via environment variable `RFC_MOO_COMPLIANCE`.
|
||||
#
|
||||
# Default value: false
|
||||
#rfc_moo_compliance = false
|
||||
|
||||
# When set to true, skip age-assurance birthday prompt for all accounts.
|
||||
#
|
||||
# Can also be specified via environment variable `PDS_AGE_ASSURANCE_OVERRIDE`.
|
||||
@@ -383,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`.
|
||||
#
|
||||
@@ -416,7 +423,9 @@
|
||||
# List of relay / crawler notification URLs.
|
||||
#
|
||||
# Can also be specified via environment variable `CRAWLERS`.
|
||||
#crawlers =
|
||||
#
|
||||
# Default value: ["https://relay.fire.hose.cam", "https://relay3.fr.hose.cam", "https://bsky.network", "https://northamerica.firehose.network", "https://europe.firehose.network", "https://asia.firehose.network", "https://atproto.africa", "https://relay.upcloud.world"]
|
||||
#crawlers = ["https://relay.fire.hose.cam", "https://relay3.fr.hose.cam", "https://bsky.network", "https://northamerica.firehose.network", "https://europe.firehose.network", "https://asia.firehose.network", "https://atproto.africa", "https://relay.upcloud.world"]
|
||||
|
||||
[email]
|
||||
# Sender email address. When unset, email sending is disabled.
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
display_name: string
|
||||
granted: boolean | null
|
||||
restricted?: boolean
|
||||
effective_scope?: string
|
||||
}
|
||||
|
||||
const SCOPE_LOCALE_MAP: Record<string, string> = {
|
||||
@@ -320,7 +321,7 @@
|
||||
)
|
||||
|
||||
function getLocalizedScopeName(scope: ScopeInfo): string {
|
||||
const localeKey = SCOPE_LOCALE_MAP[scope.scope]
|
||||
const localeKey = SCOPE_LOCALE_MAP[scope.effective_scope ?? scope.scope]
|
||||
if (!localeKey) return scope.display_name
|
||||
|
||||
if (scope.scope === 'atproto' && hasGranularScopes) {
|
||||
@@ -333,7 +334,7 @@
|
||||
}
|
||||
|
||||
function getLocalizedScopeDescription(scope: ScopeInfo): string {
|
||||
const localeKey = SCOPE_LOCALE_MAP[scope.scope]
|
||||
const localeKey = SCOPE_LOCALE_MAP[scope.effective_scope ?? scope.scope]
|
||||
if (!localeKey) return scope.description
|
||||
|
||||
if (scope.scope === 'atproto' && hasGranularScopes) {
|
||||
@@ -355,7 +356,7 @@
|
||||
const rpc: string[] = []
|
||||
const other: ScopeInfo[] = []
|
||||
for (const s of expanded) {
|
||||
const [base, query = ''] = s.scope.split('?')
|
||||
const [base, query = ''] = (s.effective_scope ?? s.scope).split('?')
|
||||
const params = new URLSearchParams(query)
|
||||
if (base.startsWith('repo:')) {
|
||||
const collection = base.slice('repo:'.length) || '*'
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user