fix: OAuth callback listener skips speculative/malformed browser connections

This commit is contained in:
2026-07-20 13:21:44 -06:00
parent f6bd02dc73
commit 13d31f850c
+55 -24
View File
@@ -322,9 +322,7 @@ async fn run_client_credentials_flow(
let access_token = response["access_token"] let access_token = response["access_token"]
.as_str() .as_str()
.ok_or_else(|| { .ok_or_else(|| anyhow!("Missing access_token in client_credentials response: {response}"))?
anyhow!("Missing access_token in client_credentials response: {response}")
})?
.to_string(); .to_string();
let expires_in = response["expires_in"] let expires_in = response["expires_in"]
.as_i64() .as_i64()
@@ -435,9 +433,8 @@ pub async fn prepare_oauth_access_token(
OAuthFlow::Pkce => refresh_oauth_token(client, provider, client_name, &tokens).await?, OAuthFlow::Pkce => refresh_oauth_token(client, provider, client_name, &tokens).await?,
OAuthFlow::ClientCredentials => { OAuthFlow::ClientCredentials => {
run_client_credentials_flow(provider, client_name).await?; run_client_credentials_flow(provider, client_name).await?;
load_oauth_tokens(client_name).ok_or_else(|| { load_oauth_tokens(client_name)
anyhow!("Token file missing after client_credentials refresh") .ok_or_else(|| anyhow!("Token file missing after client_credentials refresh"))?
})?
} }
} }
} else { } else {
@@ -506,22 +503,28 @@ fn listen_for_oauth_callback(redirect_uri: &str) -> Result<(String, String)> {
println!("Waiting for OAuth callback on {redirect_uri} ...\n"); println!("Waiting for OAuth callback on {redirect_uri} ...\n");
let listener = TcpListener::bind(format!("{host}:{port}"))?; let listener = TcpListener::bind(format!("{host}:{port}"))?;
let (mut stream, _) = listener.accept()?;
loop {
let (mut stream, _) = listener.accept()?;
let mut reader = BufReader::new(&stream); let mut reader = BufReader::new(&stream);
let mut request_line = String::new(); let mut request_line = String::new();
reader.read_line(&mut request_line)?; if reader.read_line(&mut request_line).is_err() || request_line.trim().is_empty() {
continue;
}
let request_path = request_line let Some(request_path) = request_line.split_whitespace().nth(1) else {
.split_whitespace() continue;
.nth(1) };
.ok_or_else(|| anyhow!("Malformed HTTP request from OAuth callback"))?;
let full_url = format!("http://{host}:{port}{request_path}"); let Ok(parsed) = format!("http://{host}:{port}{request_path}").parse::<Url>() else {
let parsed: Url = full_url.parse()?; continue;
};
if !parsed.path().starts_with(path) { if !parsed.path().starts_with(path) {
bail!("Unexpected callback path: {}", parsed.path()); let _ = stream.write_all(
b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
);
continue;
} }
let code = parsed let code = parsed
@@ -551,7 +554,8 @@ fn listen_for_oauth_callback(redirect_uri: &str) -> Result<(String, String)> {
); );
stream.write_all(response.as_bytes())?; stream.write_all(response.as_bytes())?;
Ok((code, returned_state)) return Ok((code, returned_state));
}
} }
pub fn get_oauth_provider(provider_type: &str) -> Option<Box<dyn OAuthProvider>> { pub fn get_oauth_provider(provider_type: &str) -> Option<Box<dyn OAuthProvider>> {
@@ -655,8 +659,8 @@ pub(crate) fn client_config_info(
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::client::{ModelData, ProviderModels};
use crate::client::openai_compatible::OpenAICompatibleConfig; use crate::client::openai_compatible::OpenAICompatibleConfig;
use crate::client::{ModelData, ProviderModels};
fn base_config() -> OAuthConfig { fn base_config() -> OAuthConfig {
OAuthConfig { OAuthConfig {
@@ -709,10 +713,19 @@ mod tests {
assert_eq!(merged.client_id, "user-id"); assert_eq!(merged.client_id, "user-id");
assert_eq!(merged.token_url, "https://user.example/token"); assert_eq!(merged.token_url, "https://user.example/token");
assert_eq!(merged.client_secret.as_deref(), Some("user-secret")); assert_eq!(merged.client_secret.as_deref(), Some("user-secret"));
assert_eq!(merged.authorize_url.as_deref(), Some("https://base.example/authorize")); assert_eq!(
merged.authorize_url.as_deref(),
Some("https://base.example/authorize")
);
assert_eq!(merged.redirect_port, Some(1234)); assert_eq!(merged.redirect_port, Some(1234));
assert_eq!(merged.scopes, vec!["c"]); assert_eq!(merged.scopes, vec!["c"]);
assert_eq!(merged.extra_authorize_params.get("plan").map(String::as_str), Some("user")); assert_eq!(
merged
.extra_authorize_params
.get("plan")
.map(String::as_str),
Some("user")
);
} }
#[test] #[test]
@@ -725,11 +738,23 @@ mod tests {
assert_eq!(merged.client_id, "user-id"); assert_eq!(merged.client_id, "user-id");
assert_eq!(merged.token_url, "https://user.example/token"); assert_eq!(merged.token_url, "https://user.example/token");
assert_eq!(merged.client_secret.as_deref(), Some("base-secret")); assert_eq!(merged.client_secret.as_deref(), Some("base-secret"));
assert_eq!(merged.authorize_url.as_deref(), Some("https://base.example/authorize")); assert_eq!(
merged.authorize_url.as_deref(),
Some("https://base.example/authorize")
);
assert_eq!(merged.redirect_port, Some(1234)); assert_eq!(merged.redirect_port, Some(1234));
assert_eq!(merged.scopes, vec!["a", "b"]); assert_eq!(merged.scopes, vec!["a", "b"]);
assert!(matches!(merged.token_request_format, Some(TokenRequestFormat::FormUrlEncoded))); assert!(matches!(
assert_eq!(merged.extra_authorize_params.get("plan").map(String::as_str), Some("base")); merged.token_request_format,
Some(TokenRequestFormat::FormUrlEncoded)
));
assert_eq!(
merged
.extra_authorize_params
.get("plan")
.map(String::as_str),
Some("base")
);
} }
#[test] #[test]
@@ -756,10 +781,16 @@ echo_pkce_in_token_exchange: true
assert_eq!(cfg.scopes.len(), 3); assert_eq!(cfg.scopes.len(), 3);
assert_eq!(cfg.redirect_port, Some(56121)); assert_eq!(cfg.redirect_port, Some(56121));
assert!(matches!(cfg.flow, OAuthFlow::Pkce)); assert!(matches!(cfg.flow, OAuthFlow::Pkce));
assert!(matches!(cfg.token_request_format, Some(TokenRequestFormat::FormUrlEncoded))); assert!(matches!(
cfg.token_request_format,
Some(TokenRequestFormat::FormUrlEncoded)
));
assert!(cfg.echo_pkce_in_token_exchange); assert!(cfg.echo_pkce_in_token_exchange);
assert!(cfg.include_state_in_token_exchange); assert!(cfg.include_state_in_token_exchange);
assert_eq!(cfg.extra_authorize_params.get("plan").map(String::as_str), Some("generic")); assert_eq!(
cfg.extra_authorize_params.get("plan").map(String::as_str),
Some("generic")
);
} }
#[test] #[test]