diff --git a/examples/Caddyfile b/examples/Caddyfile index 9832246..26b8fa3 100644 --- a/examples/Caddyfile +++ b/examples/Caddyfile @@ -14,6 +14,7 @@ path /xrpc/com.atproto.server.getSession path /xrpc/com.atproto.server.updateEmail path /xrpc/com.atproto.server.createSession + path /xrpc/com.atproto.server.createAccount path /@atproto/oauth-provider/~api/sign-in } diff --git a/src/main.rs b/src/main.rs index 42fcb1a..2381dc5 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,6 @@ #![warn(clippy::unwrap_used)] use crate::oauth_provider::sign_in; -use crate::xrpc::com_atproto_server::{create_session, get_session, update_email}; +use crate::xrpc::com_atproto_server::{create_account, create_session, get_session, update_email}; use axum::body::Body; use axum::handler::Handler; use axum::http::{Method, header}; @@ -20,7 +20,6 @@ use std::time::Duration; use std::{env, net::SocketAddr}; use tower_governor::GovernorLayer; use tower_governor::governor::{GovernorConfig, GovernorConfigBuilder}; -use tower_governor::key_extractor::PeerIpKeyExtractor; use tower_http::compression::CompressionLayer; use tower_http::cors::{Any, CorsLayer}; use tracing::log; @@ -92,7 +91,12 @@ async fn main() -> Result<(), Box> { let pds_env_location = env::var("PDS_ENV_LOCATION").unwrap_or_else(|_| "/pds/pds.env".to_string()); - dotenvy::from_path(Path::new(&pds_env_location))?; + let result_of_finding_pds_env = dotenvy::from_path(Path::new(&pds_env_location)); + if let Err(e) = result_of_finding_pds_env { + log::error!( + "Error loading pds.env file (ignore if you loaded your variables in the environment somehow else): {e}" + ); + } let pds_root = env::var("PDS_DATA_DIRECTORY")?; let account_db_url = format!("{pds_root}/account.sqlite"); @@ -182,33 +186,32 @@ async fn main() -> Result<(), Box> { env::var("GATEKEEPER_CREATE_ACCOUNT_PER_SECOND").ok(); let create_account_limiter_burst: Option = env::var("GATEKEEPER_CREATE_ACCOUNT_BURST").ok(); - let mut create_account_governor_conf = None; - if create_account_governor_conf.is_some() && create_account_limiter_time.is_some() { + //Default should be 608 requests per 5 minutes, PDS is 300 per 500 so will never hit it ideally + let mut create_account_governor_conf = GovernorConfigBuilder::default(); + if create_account_limiter_time.is_some() { let time = create_account_limiter_time .expect("GATEKEEPER_CREATE_ACCOUNT_PER_SECOND not set") .parse::() .expect("GATEKEEPER_CREATE_ACCOUNT_PER_SECOND must be a valid integer"); + create_account_governor_conf.per_second(time); + } + + if create_account_limiter_burst.is_some() { let burst = create_account_limiter_burst .expect("GATEKEEPER_CREATE_ACCOUNT_BURST not set") .parse::() .expect("GATEKEEPER_CREATE_ACCOUNT_BURST must be a valid integer"); - - create_account_governor_conf = Some( - GovernorConfigBuilder::default() - .per_second(time) - .burst_size(burst) - .finish() - .expect("failed to create governor config for create account. this should not happen and is a bug"), - ) + create_account_governor_conf.burst_size(burst); } + let create_account_governor_conf = create_account_governor_conf.finish().expect( + "failed to create governor config for create account. this should not happen and is a bug", + ); + let create_session_governor_limiter = create_session_governor_conf.limiter().clone(); let sign_in_governor_limiter = sign_in_governor_conf.limiter().clone(); - let create_account_governor_limiter = match create_account_governor_conf { - None => None, - Some(conf) => Some(conf.limiter().clone()), - }; + let create_account_governor_limiter = create_account_governor_conf.limiter().clone(); let interval = Duration::from_secs(60); // a separate background task to clean up @@ -217,9 +220,7 @@ async fn main() -> Result<(), Box> { std::thread::sleep(interval); create_session_governor_limiter.retain_recent(); sign_in_governor_limiter.retain_recent(); - if let Some(ref limiter) = create_account_governor_limiter { - limiter.retain_recent(); - } + create_account_governor_limiter.retain_recent(); } }); @@ -243,7 +244,10 @@ async fn main() -> Result<(), Box> { "/xrpc/com.atproto.server.createSession", post(create_session.layer(GovernorLayer::new(create_session_governor_conf))), ) - .route("/xrpc/com.atproto.server.createAccount") + .route( + "/xrpc/com.atproto.server.createAccount", + post(create_account).layer(GovernorLayer::new(create_account_governor_conf)), + ) .layer(CompressionLayer::new()) .layer(cors) .with_state(state); diff --git a/src/xrpc/com_atproto_server.rs b/src/xrpc/com_atproto_server.rs index b05676f..552eddf 100644 --- a/src/xrpc/com_atproto_server.rs +++ b/src/xrpc/com_atproto_server.rs @@ -264,3 +264,27 @@ pub async fn get_session( ProxiedResult::Passthrough(resp) => Ok(resp), } } + +pub async fn create_account( + State(state): State, + mut req: Request, +) -> Result, StatusCode> { + let uri = format!( + "{}{}", + state.pds_base_url, "/xrpc/com.atproto.server.createAccount" + ); + + // Rewrite the URI to point at the upstream PDS; keep headers, method, and body intact + *req.uri_mut() = uri + .parse() + .map_err(|_| StatusCode::BAD_REQUEST)?; + + let proxied = state + .reverse_proxy_client + .request(req) + .await + .map_err(|_| StatusCode::BAD_REQUEST)? + .into_response(); + + Ok(proxied) +}