diff --git a/Cargo.lock b/Cargo.lock index 39254f358..7fcfa52d5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -23,12 +23,6 @@ version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" -[[package]] -name = "adler32" -version = "1.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "aae1277d39aeec15cb388266ecc24b11c80469deae6067e17a1a7aa9e5c1f234" - [[package]] name = "aead" version = "0.5.2" @@ -91,21 +85,6 @@ dependencies = [ "equator", ] -[[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.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94fb8275041c72129eb51b7d0322c29b8387a0386127718b096429201a5d6ece" -dependencies = [ - "alloc-no-stdlib", -] - [[package]] name = "allocator-api2" version = "0.2.21" @@ -224,12 +203,6 @@ dependencies = [ "stable_deref_trait", ] -[[package]] -name = "ascii" -version = "1.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d92bec98840b8f03a5ff5413de5293bfcd8bf96467cf5452609f939ec6f5de16" - [[package]] name = "async-compression" version = "0.4.41" @@ -478,12 +451,6 @@ dependencies = [ "match-lookup", ] -[[package]] -name = "base64" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9e1b586273c5702936fe7b7d6896644d8be71e6314cfe09d3167c95f712589e8" - [[package]] name = "base64" version = "0.22.1" @@ -582,37 +549,6 @@ dependencies = [ "cfg_aliases", ] -[[package]] -name = "brotli" -version = "3.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d640d25bc63c50fb1f0b545ffd80207d2e10a4c965530809b40ba3386825c391" -dependencies = [ - "alloc-no-stdlib", - "alloc-stdlib", - "brotli-decompressor", -] - -[[package]] -name = "brotli-decompressor" -version = "2.5.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4e2e4afe60d7dd600fdd3de8d0f08c2b7ec039712e3b6137ff98b7004e82de4f" -dependencies = [ - "alloc-no-stdlib", - "alloc-stdlib", -] - -[[package]] -name = "buf_redux" -version = "0.8.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b953a6887648bb07a535631f2bc00fbdb2a2216f135552cb3f534ed136b9c07f" -dependencies = [ - "memchr", - "safemem", -] - [[package]] name = "buffer" version = "0.1.9" @@ -721,12 +657,6 @@ dependencies = [ "windows-link", ] -[[package]] -name = "chunked_transfer" -version = "1.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e4de3bc4ea267985becf712dc6d9eed8b04c953b3fcfb339ebc87acd9804901" - [[package]] name = "ciborium" version = "0.2.2" @@ -916,7 +846,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4ddef33a339a91ea89fb53151bd0a4689cfce27055c291dfa69945475d22c747" dependencies = [ "aes-gcm", - "base64 0.22.1", + "base64", "percent-encoding", "rand 0.8.5", "subtle", @@ -1188,16 +1118,6 @@ dependencies = [ "syn", ] -[[package]] -name = "deflate" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c86f7e25f518f4b81808a2cf1c50996a61f5c2eb394b2393bd87f2a4780a432f" -dependencies = [ - "adler32", - "gzip-header", -] - [[package]] name = "der" version = "0.7.10" @@ -1550,17 +1470,6 @@ version = "0.2.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" -[[package]] -name = "filetime" -version = "0.2.27" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f98844151eee8917efc50bd9e8318cb963ae8b297431495d3f758616ea5c57db" -dependencies = [ - "cfg-if", - "libc", - "libredox", -] - [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -1870,15 +1779,6 @@ dependencies = [ "subtle", ] -[[package]] -name = "gzip-header" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "95cc527b92e6029a62960ad99aa8a6660faa4555fe5f731aab13aa6a921795a2" -dependencies = [ - "crc32fast", -] - [[package]] name = "h2" version = "0.4.13" @@ -1980,12 +1880,6 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" -[[package]] -name = "hermit-abi" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" - [[package]] name = "hex" version = "0.4.3" @@ -2169,7 +2063,7 @@ version = "0.1.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" dependencies = [ - "base64 0.22.1", + "base64", "bytes", "futures-channel", "futures-util", @@ -2544,7 +2438,7 @@ dependencies = [ "axum-extra", "axum-macros", "axum-test", - "base64 0.22.1", + "base64", "bytes", "chrono", "clap", @@ -2590,7 +2484,7 @@ dependencies = [ name = "jacquard-common" version = "0.12.0-beta.2" dependencies = [ - "base64 0.22.1", + "base64", "bon", "bytes", "chrono", @@ -2737,7 +2631,7 @@ dependencies = [ name = "jacquard-oauth" version = "0.12.0-beta.2" dependencies = [ - "base64 0.22.1", + "base64", "bytes", "chrono", "dashmap", @@ -2755,7 +2649,6 @@ dependencies = [ "p256", "p384", "rand 0.8.5", - "rouille", "serde", "serde_html_form", "serde_json", @@ -3009,18 +2902,6 @@ version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" -[[package]] -name = "libredox" -version = "0.1.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7ddbf48fd451246b1f8c2610bd3b4ac0cc6e149d89832867093ab69a17194f08" -dependencies = [ - "bitflags", - "libc", - "plain", - "redox_syscall 0.7.3", -] - [[package]] name = "linked-hash-map" version = "0.5.6" @@ -3247,16 +3128,6 @@ version = "0.3.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" -[[package]] -name = "mime_guess" -version = "2.0.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f7c44f8e672c00fe5308fa235f821cb4198414e1c77935c1ab6948d3fd78550e" -dependencies = [ - "mime", - "unicase", -] - [[package]] name = "mini-moka-wasm" version = "0.10.99" @@ -3349,24 +3220,6 @@ dependencies = [ "unsigned-varint 0.8.0", ] -[[package]] -name = "multipart" -version = "0.18.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "00dec633863867f29cb39df64a397cdf4a6354708ddd7759f70c7fb51c5f9182" -dependencies = [ - "buf_redux", - "httparse", - "log", - "mime", - "mime_guess", - "quick-error 1.2.3", - "rand 0.8.5", - "safemem", - "tempfile", - "twoway", -] - [[package]] name = "mutex-traits" version = "1.0.1" @@ -3543,25 +3396,6 @@ dependencies = [ "libm", ] -[[package]] -name = "num_cpus" -version = "1.17.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" -dependencies = [ - "hermit-abi", - "libc", -] - -[[package]] -name = "num_threads" -version = "0.1.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c7398b9c8b70908f6371f47ed36737907c87c52af34c268fed0bf0ceb92ead9" -dependencies = [ - "libc", -] - [[package]] name = "objc2" version = "0.6.4" @@ -3689,7 +3523,7 @@ checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" dependencies = [ "cfg-if", "libc", - "redox_syscall 0.5.18", + "redox_syscall", "smallvec", "windows-link", ] @@ -3832,12 +3666,6 @@ version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7edddbd0b52d732b21ad9a5fab5c704c14cd949e5e9a1ec5929a24fded1b904c" -[[package]] -name = "plain" -version = "0.2.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6" - [[package]] name = "png" version = "0.18.1" @@ -4238,15 +4066,6 @@ dependencies = [ "bitflags", ] -[[package]] -name = "redox_syscall" -version = "0.7.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6ce70a74e890531977d37e532c34d45e9055d2409ed08ddba14529471ed0be16" -dependencies = [ - "bitflags", -] - [[package]] name = "ref-cast" version = "1.0.25" @@ -4308,7 +4127,7 @@ version = "0.12.28" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" dependencies = [ - "base64 0.22.1", + "base64", "bytes", "encoding_rs", "futures-core", @@ -4400,30 +4219,6 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dbf2048e0e979efb2ca7b91c4f1a8d77c91853e9b987c94c555668a8994915ad" -[[package]] -name = "rouille" -version = "3.6.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3716fbf57fc1084d7a706adf4e445298d123e4a44294c4e8213caf1b85fcc921" -dependencies = [ - "base64 0.13.1", - "brotli", - "chrono", - "deflate", - "filetime", - "multipart", - "percent-encoding", - "rand 0.8.5", - "serde", - "serde_derive", - "serde_json", - "sha1_smol", - "threadpool", - "time", - "tiny_http", - "url", -] - [[package]] name = "rsa" version = "0.9.10" @@ -4577,12 +4372,6 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" -[[package]] -name = "safemem" -version = "0.3.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ef703b7cb59335eae2eb93ceb664c0eb7ea6bf567079d843e09420219668e072" - [[package]] name = "same-file" version = "1.0.6" @@ -4809,7 +4598,7 @@ version = "3.18.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dd5414fad8e6907dbdd5bc441a50ae8d6e26151a03b1de04d89a5576de61d01f" dependencies = [ - "base64 0.22.1", + "base64", "chrono", "hex", "serde_core", @@ -4841,12 +4630,6 @@ dependencies = [ "digest", ] -[[package]] -name = "sha1_smol" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bbfa15b3dddfee50a0fff136974b3e1bde555604ba463834a7eb7deb6417705d" - [[package]] name = "sha2" version = "0.10.9" @@ -5217,15 +5000,6 @@ dependencies = [ "cfg-if", ] -[[package]] -name = "threadpool" -version = "1.8.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d050e60b33d41c19108b32cea32164033a9013fe3b46cbd4457559bfbf77afaa" -dependencies = [ - "num_cpus", -] - [[package]] name = "tiff" version = "0.6.1" @@ -5259,9 +5033,7 @@ checksum = "743bd48c283afc0388f9b8827b976905fb217ad9e647fae3a379a9283c4def2c" dependencies = [ "deranged", "itoa", - "libc", "num-conv", - "num_threads", "powerfmt", "serde_core", "time-core", @@ -5284,18 +5056,6 @@ dependencies = [ "time-core", ] -[[package]] -name = "tiny_http" -version = "0.12.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "389915df6413a2e74fb181895f933386023c71110878cd0825588928e64cdc82" -dependencies = [ - "ascii", - "chunked_transfer", - "httpdate", - "log", -] - [[package]] name = "tinystr" version = "0.8.2" @@ -5716,15 +5476,6 @@ dependencies = [ "utf-8", ] -[[package]] -name = "twoway" -version = "0.1.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "59b11b2b5241ba34be09c3cc85a36e56e48f9888862e19cedf23336d35316ed1" -dependencies = [ - "memchr", -] - [[package]] name = "typeid" version = "1.0.3" @@ -5767,12 +5518,6 @@ version = "0.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "eaea85b334db583fe3274d12b4cd1880032beab409c0d774be044d4480ab9a94" -[[package]] -name = "unicase" -version = "2.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" - [[package]] name = "unicode-ident" version = "1.0.24" @@ -5913,7 +5658,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0ae7c6870b98c838123f22cac9a594cbe2d74ea48d79271c08f8c9e680b40fac" dependencies = [ "ansi_colours", - "base64 0.22.1", + "base64", "console", "crossterm", "image", diff --git a/crates/jacquard-oauth/Cargo.toml b/crates/jacquard-oauth/Cargo.toml index 1fcbf6571..9a588acac 100644 --- a/crates/jacquard-oauth/Cargo.toml +++ b/crates/jacquard-oauth/Cargo.toml @@ -14,7 +14,7 @@ license.workspace = true [features] default = [] -loopback = ["dep:rouille"] +loopback = [] browser-open = ["dep:webbrowser"] tracing = ["dep:tracing"] websocket = ["jacquard-common/websocket"] @@ -52,9 +52,8 @@ webbrowser = { version = "1", optional = true } tracing = { workspace = true, optional = true } smallvec.workspace = true -[target.'cfg(not(target_arch = "wasm32"))'.dependencies] -tokio = { workspace = true, features = ["rt", "net", "time"] } -rouille = { version = "3.6.2", optional = true } +[target.'cfg(not(all(target_arch = "wasm32", target_os = "unknown")))'.dependencies] +tokio = { workspace = true, features = ["rt", "net", "time", "io-util"] } [target.'cfg(target_arch = "wasm32")'.dependencies] diff --git a/crates/jacquard-oauth/src/error.rs b/crates/jacquard-oauth/src/error.rs index 9fc0bc81f..f7e95edb5 100644 --- a/crates/jacquard-oauth/src/error.rs +++ b/crates/jacquard-oauth/src/error.rs @@ -107,6 +107,10 @@ pub enum CallbackError { #[error("timeout")] #[diagnostic(code(jacquard_oauth::callback::timeout))] Timeout, + /// The local loopback callback server could not start or accept callbacks. + #[error("loopback callback server error: {0}")] + #[diagnostic(code(jacquard_oauth::callback::loopback_server))] + LoopbackServer(String), /// An error occurred resolving permission sets during session creation. #[cfg(feature = "scope-check")] #[error("scope resolution failed: {detail}")] diff --git a/crates/jacquard-oauth/src/lib.rs b/crates/jacquard-oauth/src/lib.rs index 36fcc3b12..4bf080516 100644 --- a/crates/jacquard-oauth/src/lib.rs +++ b/crates/jacquard-oauth/src/lib.rs @@ -78,5 +78,8 @@ pub mod utils; pub const FALLBACK_ALG: &str = "ES256"; /// Loopback server helpers for the local redirect-based OAuth flow. -#[cfg(feature = "loopback")] +#[cfg(all( + feature = "loopback", + not(all(target_arch = "wasm32", target_os = "unknown")) +))] pub mod loopback; diff --git a/crates/jacquard-oauth/src/loopback.rs b/crates/jacquard-oauth/src/loopback.rs index f7e31ab81..457501074 100644 --- a/crates/jacquard-oauth/src/loopback.rs +++ b/crates/jacquard-oauth/src/loopback.rs @@ -43,7 +43,10 @@ //! ``` //! //! -#![cfg(feature = "loopback")] +#![cfg(all( + feature = "loopback", + not(all(target_arch = "wasm32", target_os = "unknown")) +))] use crate::{ atproto::AtprotoClientMetadata, authstore::{ClientAuthStore, OAuthSessionMatch}, @@ -57,10 +60,13 @@ use jacquard_common::IntoStatic; use jacquard_common::deps::fluent_uri::Uri; use jacquard_common::session::{SessionHint, SessionSelector, SessionStoreError}; use jacquard_common::types::{did::Did, string::Handle}; -use rouille::Server; -use smol_str::{SmolStr, ToSmolStr}; +use smol_str::SmolStr; use std::net::SocketAddr; -use tokio::sync::mpsc; +use tokio::{ + io::{AsyncBufReadExt, AsyncWriteExt, BufReader}, + net::{TcpListener, TcpStream, ToSocketAddrs}, + sync::{mpsc, oneshot}, +}; fn oauth_hint_from_input(input: &str) -> SessionHint { if let Ok(did) = Did::new(input) { @@ -118,32 +124,62 @@ pub fn try_open_in_browser(_url: &str) -> bool { false } -fn create_callback_router( - request: &rouille::Request, - tx: mpsc::Sender, -) -> rouille::Response { - rouille::router!(request, - (GET) (/oauth/callback) => { - let state = request.get_param("state").unwrap(); - let code = request.get_param("code").unwrap(); - let iss = request.get_param("iss").unwrap(); - let callback_params = CallbackParams { - state: Some(state.to_smolstr()), - code: code.to_smolstr(), - iss: Some(iss.to_smolstr()), - }; - tx.try_send(callback_params).unwrap(); - rouille::Response::text("Logged in!") - }, - _ => rouille::Response::empty_404() - ) +async fn handle_callback_connection(mut stream: TcpStream, tx: mpsc::Sender) { + let Some(Some(params)) = read_callback_params(&mut stream).await else { + let _ = write_http_response(&mut stream, 404, "Not found").await; + return; + }; + + match tx.try_send(params) { + Ok(()) => { + let _ = write_http_response(&mut stream, 200, "Logged in!").await; + } + Err(_) => { + let _ = write_http_response(&mut stream, 500, "Could not deliver OAuth callback").await; + } + } +} + +async fn read_callback_params(stream: &mut TcpStream) -> Option> { + let mut reader = BufReader::new(stream); + let mut request_line = String::new(); + reader.read_line(&mut request_line).await.ok()?; + let mut parts = request_line.split_whitespace(); + let method = parts.next()?; + let target = parts.next()?; + if method != "GET" { + return Some(None); + } + let (path, query) = target.split_once('?').unwrap_or((target, "")); + if path != "/oauth/callback" { + return Some(None); + } + serde_html_form::from_str(query).ok().map(Some) +} + +async fn write_http_response( + stream: &mut TcpStream, + status: u16, + body: &str, +) -> std::io::Result<()> { + let reason = match status { + 200 => "OK", + 404 => "Not Found", + 500 => "Internal Server Error", + _ => "OK", + }; + let response = format!( + "HTTP/1.1 {status} {reason}\r\ncontent-type: text/plain; charset=utf-8\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + body + ); + stream.write_all(response.as_bytes()).await } /// Handle to a running loopback callback server, used to await the OAuth redirect. pub struct CallbackHandle { - #[allow(dead_code)] - server_handle: std::thread::JoinHandle<()>, - server_stop: std::sync::mpsc::Sender<()>, + server_handle: tokio::task::JoinHandle<()>, + server_stop: oneshot::Sender<()>, callback_rx: mpsc::Receiver, } @@ -155,19 +191,57 @@ pub struct CallbackHandle { /// /// Use in combination with [`handle_localhost_callback`] to handle the /// callback for the localhost loopback server. -pub fn one_shot_server(addr: SocketAddr) -> (SocketAddr, CallbackHandle) { +pub async fn one_shot_server( + addr: impl ToSocketAddrs, +) -> std::io::Result<(SocketAddr, CallbackHandle)> { let (tx, callback_rx) = mpsc::channel(5); - let server = Server::new(addr, move |request| { - create_callback_router(request, tx.clone()) - }) - .expect("Could not start server"); - let (server_handle, server_stop) = server.stoppable(); + let listener = TcpListener::bind(addr).await?; + let local_addr = listener.local_addr()?; + let (server_stop, mut stop_rx) = oneshot::channel(); + let server_handle = tokio::spawn(async move { + loop { + tokio::select! { + _ = &mut stop_rx => break, + accepted = listener.accept() => { + match accepted { + Ok((stream, _)) => { + tokio::spawn(handle_callback_connection(stream, tx.clone())); + } + Err(_) => break, + } + } + } + } + }); let handle = CallbackHandle { server_handle, server_stop, callback_rx, }; - (addr, handle) + Ok((local_addr, handle)) +} + +async fn wait_for_callback( + handle: CallbackHandle, + timeout_ms: u64, +) -> Result { + let CallbackHandle { + server_handle, + server_stop, + mut callback_rx, + } = handle; + let cb = tokio::time::timeout( + std::time::Duration::from_millis(timeout_ms), + callback_rx.recv(), + ) + .await; + let _ = server_stop.send(()); + let _ = server_handle.await; + if let Ok(Some(cb)) = cb { + Ok(cb) + } else { + Err(OAuthError::Callback(CallbackError::Timeout)) + } } /// Handles the OAuth callback for the localhost loopback server. @@ -187,21 +261,9 @@ where T: OAuthResolver + DpopExt + Send + Sync + 'static, S: ClientAuthStore + Send + Sync + 'static, { - // Await callback or timeout - let mut callback_rx = handle.callback_rx; - let cb = tokio::time::timeout( - std::time::Duration::from_millis(cfg.timeout_ms), - callback_rx.recv(), - ) - .await; - // trigger shutdown - let _ = handle.server_stop.send(()); - if let Ok(Some(cb)) = cb { - // Handle callback and create a session - Ok(flow_client.callback(cb).await?) - } else { - Err(OAuthError::Callback(CallbackError::Timeout)) - } + Ok(flow_client + .callback(wait_for_callback(handle, cfg.timeout_ms).await?) + .await?) } /// Handles the OAuth callback for the localhost loopback server. @@ -226,83 +288,56 @@ where + 'static, S: ClientAuthStore + Send + Sync + 'static, { - // Await callback or timeout - let mut callback_rx = handle.callback_rx; - let cb = tokio::time::timeout( - std::time::Duration::from_millis(cfg.timeout_ms), - callback_rx.recv(), - ) - .await; - // trigger shutdown - let _ = handle.server_stop.send(()); - if let Ok(Some(cb)) = cb { - // Handle callback and create a session - Ok(flow_client.callback(cb).await?) + Ok(flow_client + .callback(wait_for_callback(handle, cfg.timeout_ms).await?) + .await?) +} + +fn loopback_port(cfg: &LoopbackConfig) -> u16 { + match cfg.port { + LoopbackPort::Fixed(port) => port, + LoopbackPort::Ephemeral => 0, + } +} + +fn redirect_host(host: &str) -> String { + if host.contains(':') && !host.starts_with('[') { + format!("[{host}]") } else { - Err(OAuthError::Callback(CallbackError::Timeout)) + host.to_owned() } } -#[cfg(not(feature = "scope-check"))] impl OAuthClient where T: OAuthResolver + DpopExt + Send + Sync + 'static, S: ClientAuthStore + Send + Sync + 'static, { - /// Drive the full OAuth flow using a local loopback server. - /// - /// This uses localhost OAuth and an ephemeral in-process web server to - /// handle the OAuth callback redirect. It has a bunch of nice friendly - /// defaults to help you get started and will basically drive the *entire* - /// callback flow itself. - /// - /// Best used for development and for small CLI applications that don't - /// require long session lengths. For long-running unattended sessions, - /// app passwords (via CredentialSession in the jacquard crate) remain - /// the best option. For more complex OAuth, or if you want more control - /// over the process, use the other methods on OAuthClient. - /// - /// 'input' parameter is what you type in the login box (usually, your handle) - /// for it to look up your PDS and redirect to its authentication interface. - /// - /// If the `browser-open` feature is enabled, this will open a web browser - /// for you to authenticate with your PDS. It will also print the - /// callback url to the console for you to copy. - pub async fn login_with_local_server( + async fn start_loopback_flow( &self, - input: impl AsRef, + input: &str, opts: AuthorizeOptions, - cfg: LoopbackConfig, - ) -> crate::error::Result> { - let port = match cfg.port { - LoopbackPort::Fixed(p) => p, - LoopbackPort::Ephemeral => 0, - }; - // TODO: fix this to it also accepts ipv6 and properly finds a free port - let bind_addr: SocketAddr = format!("0.0.0.0:{}", port) - .parse() - .expect("invalid loopback host/port"); - let (local_addr, handle) = one_shot_server(bind_addr); + cfg: &LoopbackConfig, + ) -> crate::error::Result<(OAuthClient, CallbackHandle)> { + let (local_addr, handle) = one_shot_server((cfg.host.as_str(), loopback_port(cfg))) + .await + .map_err(|err| OAuthError::Callback(CallbackError::LoopbackServer(err.to_string())))?; println!("Listening on {}", local_addr); - let client_data = self.build_localhost_client_data(&cfg, &opts, local_addr); - // Build client using store and resolver + let client_data = self.build_localhost_client_data(cfg, &opts, local_addr); let flow_client = OAuthClient::new_with_shared( self.registry.store.clone(), self.client.clone(), client_data, ); - // Start auth and get authorization URL - let auth_url = flow_client.start_auth(input.as_ref(), opts).await?; - // Print URL for copy/paste + let auth_url = flow_client.start_auth(input, opts).await?; println!("To authenticate with your PDS, visit:\n{}\n", auth_url); - // Optionally open browser if cfg.open_browser { let _ = try_open_in_browser(&auth_url); } - handle_localhost_callback(handle, &flow_client, &cfg).await + Ok((flow_client, handle)) } /// Builds a [`crate::session::ClientData`] for use with the local loopback server method of OAuth. @@ -312,7 +347,11 @@ where opts: &AuthorizeOptions, local_addr: SocketAddr, ) -> crate::session::ClientData { - let redirect_uri = format!("http://{}:{}/oauth/callback", cfg.host, local_addr.port(),); + let redirect_uri = format!( + "http://{}:{}/oauth/callback", + redirect_host(&cfg.host), + local_addr.port(), + ); let redirect = Uri::parse(redirect_uri).unwrap(); let scopes = if opts.scopes.is_empty() { @@ -328,6 +367,46 @@ where .into_static() } + async fn restore_matching_session( + &self, + input: &str, + ) -> crate::error::Result>> + where + S: SessionSelector, + { + let hint = oauth_hint_from_input(input); + if let Some(matched) = self.registry.store.select_session(&hint).await? { + Ok(Some( + self.restore(&matched.key.did, matched.key.session_id.as_str()) + .await?, + )) + } else { + Ok(None) + } + } +} + +#[cfg(not(feature = "scope-check"))] +impl OAuthClient +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, +{ + /// Drive the full OAuth flow using a local loopback server. + /// + /// This uses localhost OAuth and an ephemeral in-process web server to + /// handle the OAuth callback redirect. It has friendly defaults to drive + /// the entire callback flow for development and small CLI applications. + pub async fn login_with_local_server( + &self, + input: impl AsRef, + opts: AuthorizeOptions, + cfg: LoopbackConfig, + ) -> crate::error::Result> { + let (flow_client, handle) = self.start_loopback_flow(input.as_ref(), opts, &cfg).await?; + handle_localhost_callback(handle, &flow_client, &cfg).await + } + /// Resume a stored session for the input identity, or drive the full OAuth flow using a local loopback server. pub async fn resume_or_login_with_local_server( &self, @@ -339,11 +418,8 @@ where S: SessionSelector, { let input_ref = input.as_ref(); - let hint = oauth_hint_from_input(input_ref); - if let Some(matched) = self.registry.store.select_session(&hint).await? { - return self - .restore(&matched.key.did, matched.key.session_id.as_str()) - .await; + if let Some(session) = self.restore_matching_session(input_ref).await? { + return Ok(session); } self.login_with_local_server(input_ref, opts, cfg).await } @@ -363,56 +439,15 @@ where /// Drive the full OAuth flow using a local loopback server. /// /// This uses localhost OAuth and an ephemeral in-process web server to - /// handle the OAuth callback redirect. It has a bunch of nice friendly - /// defaults to help you get started and will basically drive the *entire* - /// callback flow itself. - /// - /// Best used for development and for small CLI applications that don't - /// require long session lengths. For long-running unattended sessions, - /// app passwords (via CredentialSession in the jacquard crate) remain - /// the best option. For more complex OAuth, or if you want more control - /// over the process, use the other methods on OAuthClient. - /// - /// 'input' parameter is what you type in the login box (usually, your handle) - /// for it to look up your PDS and redirect to its authentication interface. - /// - /// If the `browser-open` feature is enabled, this will open a web browser - /// for you to authenticate with your PDS. It will also print the - /// callback url to the console for you to copy. + /// handle the OAuth callback redirect. It has friendly defaults to drive + /// the entire callback flow for development and small CLI applications. pub async fn login_with_local_server( &self, input: impl AsRef, opts: AuthorizeOptions, cfg: LoopbackConfig, ) -> crate::error::Result> { - let port = match cfg.port { - LoopbackPort::Fixed(p) => p, - LoopbackPort::Ephemeral => 0, - }; - // TODO: fix this to it also accepts ipv6 and properly finds a free port - let bind_addr: SocketAddr = format!("0.0.0.0:{}", port) - .parse() - .expect("invalid loopback host/port"); - let (local_addr, handle) = one_shot_server(bind_addr); - println!("Listening on {}", local_addr); - - let client_data = self.build_localhost_client_data(&cfg, &opts, local_addr); - // Build client using store and resolver - let flow_client = OAuthClient::new_with_shared( - self.registry.store.clone(), - self.client.clone(), - client_data, - ); - - // Start auth and get authorization URL - let auth_url = flow_client.start_auth(input.as_ref(), opts).await?; - // Print URL for copy/paste - println!("To authenticate with your PDS, visit:\n{}\n", auth_url); - // Optionally open browser - if cfg.open_browser { - let _ = try_open_in_browser(&auth_url); - } - + let (flow_client, handle) = self.start_loopback_flow(input.as_ref(), opts, &cfg).await?; handle_localhost_callback(handle, &flow_client, &cfg).await } @@ -427,35 +462,9 @@ where S: SessionSelector, { let input_ref = input.as_ref(); - let hint = oauth_hint_from_input(input_ref); - if let Some(matched) = self.registry.store.select_session(&hint).await? { - return self - .restore(&matched.key.did, matched.key.session_id.as_str()) - .await; + if let Some(session) = self.restore_matching_session(input_ref).await? { + return Ok(session); } self.login_with_local_server(input_ref, opts, cfg).await } - - /// Builds a [`crate::session::ClientData`] for use with the local loopback server method of OAuth. - pub fn build_localhost_client_data( - &self, - cfg: &LoopbackConfig, - opts: &AuthorizeOptions, - local_addr: SocketAddr, - ) -> crate::session::ClientData { - let redirect_uri = format!("http://{}:{}/oauth/callback", cfg.host, local_addr.port(),); - let redirect = Uri::parse(redirect_uri).unwrap(); - - let scopes = if opts.scopes.is_empty() { - Some(self.registry.client_data.config.scopes.clone()) - } else { - Some(opts.scopes.clone()) - }; - - crate::session::ClientData { - keyset: self.registry.client_data.keyset.clone(), - config: AtprotoClientMetadata::new_localhost(Some(vec![redirect]), scopes), - } - .into_static() - } }