refactor: split run_oauth_flow into pkce + client_credentials dispatchers

This commit is contained in:
2026-07-20 12:43:53 -06:00
parent 559107073d
commit 4669958bdd
+54
View File
@@ -176,6 +176,13 @@ pub struct OAuthTokens {
} }
pub async fn run_oauth_flow(provider: &dyn OAuthProvider, client_name: &str) -> Result<()> { pub async fn run_oauth_flow(provider: &dyn OAuthProvider, client_name: &str) -> Result<()> {
match provider.flow() {
OAuthFlow::Pkce => run_pkce_flow(provider, client_name).await,
OAuthFlow::ClientCredentials => run_client_credentials_flow(provider, client_name).await,
}
}
async fn run_pkce_flow(provider: &dyn OAuthProvider, client_name: &str) -> Result<()> {
let random_bytes: [u8; 32] = rand::random::<[u8; 32]>(); let random_bytes: [u8; 32] = rand::random::<[u8; 32]>();
let code_verifier = URL_SAFE_NO_PAD.encode(random_bytes); let code_verifier = URL_SAFE_NO_PAD.encode(random_bytes);
@@ -256,6 +263,10 @@ pub async fn run_oauth_flow(provider: &dyn OAuthProvider, client_name: &str) ->
if provider.include_state_in_token_exchange() { if provider.include_state_in_token_exchange() {
token_params.push(("state", state.as_str())); token_params.push(("state", state.as_str()));
} }
if provider.echo_pkce_in_token_exchange() {
token_params.push(("code_challenge", code_challenge.as_str()));
token_params.push(("code_challenge_method", "S256"));
}
let request = build_token_request(&client, provider, &token_params); let request = build_token_request(&client, provider, &token_params);
let response: Value = request.send().await?.json().await?; let response: Value = request.send().await?.json().await?;
@@ -291,6 +302,49 @@ pub async fn run_oauth_flow(provider: &dyn OAuthProvider, client_name: &str) ->
Ok(()) Ok(())
} }
async fn run_client_credentials_flow(
provider: &dyn OAuthProvider,
client_name: &str,
) -> Result<()> {
let client = ReqwestClient::new();
let scopes = provider.scopes();
let mut params: Vec<(&str, &str)> = vec![
("grant_type", "client_credentials"),
("client_id", provider.client_id()),
];
if !scopes.is_empty() {
params.push(("scope", scopes));
}
let request = build_token_request(&client, provider, &params);
let response: Value = request.send().await?.json().await?;
let access_token = response["access_token"]
.as_str()
.ok_or_else(|| {
anyhow!("Missing access_token in client_credentials response: {response}")
})?
.to_string();
let expires_in = response["expires_in"]
.as_i64()
.ok_or_else(|| anyhow!("Missing expires_in in client_credentials response: {response}"))?;
let expires_at = Utc::now().timestamp() + expires_in;
let tokens = OAuthTokens {
access_token,
refresh_token: None,
expires_at,
account_id: provider.extract_account_id(&response),
};
save_oauth_tokens(client_name, &tokens)?;
println!(
"Successfully authenticated client '{}' with {} via OAuth (client_credentials). Tokens saved.",
client_name,
provider.provider_name()
);
Ok(())
}
pub fn load_oauth_tokens(client_name: &str) -> Option<OAuthTokens> { pub fn load_oauth_tokens(client_name: &str) -> Option<OAuthTokens> {
let path = paths::token_file(client_name); let path = paths::token_file(client_name);
let content = fs::read_to_string(path).ok()?; let content = fs::read_to_string(path).ok()?;