mirror of
https://github.com/mountain-loop/yaak.git
synced 2026-08-25 04:44:12 +02:00
Handle OAuth callback connections concurrently so browser preconnects can't block login
This commit is contained in:
@@ -16,6 +16,7 @@ use tokio::net::{TcpListener, TcpStream};
|
|||||||
const OAUTH_CLIENT_ID: &str = "a1fe44800c2d7e803cad1b4bf07a291c";
|
const OAUTH_CLIENT_ID: &str = "a1fe44800c2d7e803cad1b4bf07a291c";
|
||||||
const KEYRING_USER: &str = "yaak";
|
const KEYRING_USER: &str = "yaak";
|
||||||
const AUTH_TIMEOUT: Duration = Duration::from_secs(300);
|
const AUTH_TIMEOUT: Duration = Duration::from_secs(300);
|
||||||
|
const CALLBACK_READ_TIMEOUT: Duration = Duration::from_secs(10);
|
||||||
const MAX_REQUEST_BYTES: usize = 16 * 1024;
|
const MAX_REQUEST_BYTES: usize = 16 * 1024;
|
||||||
|
|
||||||
type CommandResult<T = ()> = std::result::Result<T, String>;
|
type CommandResult<T = ()> = std::result::Result<T, String>;
|
||||||
@@ -209,35 +210,71 @@ async fn receive_oauth_code(
|
|||||||
expected_state: &str,
|
expected_state: &str,
|
||||||
app_base_url: &str,
|
app_base_url: &str,
|
||||||
) -> CommandResult<String> {
|
) -> CommandResult<String> {
|
||||||
|
// Browsers speculatively open extra connections that may never carry a
|
||||||
|
// request. Handle each connection concurrently so an idle socket can't
|
||||||
|
// block the one carrying the real callback.
|
||||||
|
let (tx, mut rx) = tokio::sync::mpsc::channel::<CommandResult<String>>(1);
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
let (mut stream, _) = listener
|
tokio::select! {
|
||||||
.accept()
|
accepted = listener.accept() => {
|
||||||
.await
|
let (stream, _) = accepted
|
||||||
.map_err(|e| format!("OAuth callback server accept error: {e}"))?;
|
.map_err(|e| format!("OAuth callback server accept error: {e}"))?;
|
||||||
|
tokio::spawn(handle_callback_connection(
|
||||||
match parse_callback_request(&mut stream).await {
|
stream,
|
||||||
Ok((state, code)) => {
|
expected_state.to_string(),
|
||||||
if state != expected_state {
|
app_base_url.to_string(),
|
||||||
let _ = write_bad_request(&mut stream, "Invalid OAuth state").await;
|
tx.clone(),
|
||||||
continue;
|
));
|
||||||
}
|
|
||||||
|
|
||||||
let success_redirect = format!("{app_base_url}/login/oauth/success");
|
|
||||||
write_redirect(&mut stream, &success_redirect)
|
|
||||||
.await
|
|
||||||
.map_err(|e| format!("Failed responding to OAuth callback: {e}"))?;
|
|
||||||
return Ok(code);
|
|
||||||
}
|
}
|
||||||
Err(error) => {
|
result = rx.recv() => {
|
||||||
let _ = write_bad_request(&mut stream, &error).await;
|
if let Some(result) = result {
|
||||||
if error.starts_with("OAuth provider returned error:") {
|
return result;
|
||||||
return Err(error);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn handle_callback_connection(
|
||||||
|
mut stream: TcpStream,
|
||||||
|
expected_state: String,
|
||||||
|
app_base_url: String,
|
||||||
|
tx: tokio::sync::mpsc::Sender<CommandResult<String>>,
|
||||||
|
) {
|
||||||
|
let parsed = match tokio::time::timeout(
|
||||||
|
CALLBACK_READ_TIMEOUT,
|
||||||
|
parse_callback_request(&mut stream),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
Ok(parsed) => parsed,
|
||||||
|
Err(_) => return, // Idle speculative connection; drop it
|
||||||
|
};
|
||||||
|
|
||||||
|
match parsed {
|
||||||
|
Ok((state, code)) => {
|
||||||
|
if state != expected_state {
|
||||||
|
let _ = write_bad_request(&mut stream, "Invalid OAuth state").await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
let success_redirect = format!("{app_base_url}/login/oauth/success");
|
||||||
|
let result = match write_redirect(&mut stream, &success_redirect).await {
|
||||||
|
Ok(()) => Ok(code),
|
||||||
|
Err(e) => Err(format!("Failed responding to OAuth callback: {e}")),
|
||||||
|
};
|
||||||
|
let _ = tx.send(result).await;
|
||||||
|
}
|
||||||
|
Err(error) => {
|
||||||
|
let _ = write_bad_request(&mut stream, &error).await;
|
||||||
|
if error.starts_with("OAuth provider returned error:") {
|
||||||
|
let _ = tx.send(Err(error)).await;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
async fn parse_callback_request(stream: &mut TcpStream) -> CommandResult<(String, String)> {
|
async fn parse_callback_request(stream: &mut TcpStream) -> CommandResult<(String, String)> {
|
||||||
let target = read_http_target(stream).await?;
|
let target = read_http_target(stream).await?;
|
||||||
if !target.starts_with("/oauth/callback") {
|
if !target.starts_with("/oauth/callback") {
|
||||||
@@ -488,6 +525,37 @@ mod tests {
|
|||||||
assert!(err.contains("User denied"));
|
assert!(err.contains("User denied"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn receive_oauth_code_ignores_idle_speculative_connections() {
|
||||||
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
|
||||||
|
let addr = listener.local_addr().expect("local addr");
|
||||||
|
|
||||||
|
let server = tokio::spawn(async move {
|
||||||
|
receive_oauth_code(listener, "expected-state", "http://localhost:9444").await
|
||||||
|
});
|
||||||
|
|
||||||
|
// Browsers preconnect sockets that never carry a request; these must
|
||||||
|
// not block the connection carrying the real callback.
|
||||||
|
let _idle1 = TcpStream::connect(addr).await.expect("connect idle 1");
|
||||||
|
let _idle2 = TcpStream::connect(addr).await.expect("connect idle 2");
|
||||||
|
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||||
|
|
||||||
|
let mut client = TcpStream::connect(addr).await.expect("connect");
|
||||||
|
client
|
||||||
|
.write_all(
|
||||||
|
b"GET /oauth/callback?code=abc123&state=expected-state HTTP/1.1\r\nHost: localhost\r\n\r\n",
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("write");
|
||||||
|
|
||||||
|
let code = tokio::time::timeout(std::time::Duration::from_secs(2), server)
|
||||||
|
.await
|
||||||
|
.expect("idle connections must not block the real callback")
|
||||||
|
.expect("join")
|
||||||
|
.expect("should return code");
|
||||||
|
assert_eq!(code, "abc123");
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn receive_oauth_code_fails_fast_on_provider_error() {
|
async fn receive_oauth_code_fails_fast_on_provider_error() {
|
||||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
|
||||||
|
|||||||
Reference in New Issue
Block a user