use crate::core::api::ApiError;
use crate::utils::security::validate_url;
use anyhow::{Context, Result};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use chrono::{DateTime, Duration, Utc};
use rand::Rng;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OAuth2Config {
pub provider: String,
pub client_id: String,
pub client_secret: Option<String>,
pub auth_url: String,
pub token_url: String,
pub redirect_uri: String,
pub scopes: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub additional_params: Option<HashMap<String, String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OAuth2Token {
pub access_token: String,
pub token_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub expires_at: Option<DateTime<Utc>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub refresh_token: Option<String>,
pub scopes: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AuthProfile {
pub name: String,
pub provider: String,
pub created_at: DateTime<Utc>,
#[serde(skip_serializing_if = "Option::is_none")]
pub last_used: Option<DateTime<Utc>>,
pub metadata: HashMap<String, Value>,
}
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
token_type: String,
#[serde(default)]
expires_in: Option<i64>,
#[serde(default)]
refresh_token: Option<String>,
#[serde(default)]
scope: Option<String>,
}
pub struct PkceParams {
pub verifier: String,
pub challenge: String,
}
pub fn generate_state() -> String {
let mut rng = rand::thread_rng();
let random_bytes: Vec<u8> = (0..32).map(|_| rng.gen::<u8>()).collect();
URL_SAFE_NO_PAD.encode(&random_bytes)
}
pub fn generate_pkce_pair() -> PkceParams {
let mut rng = rand::thread_rng();
let verifier_bytes: Vec<u8> = (0..32).map(|_| rng.gen::<u8>()).collect();
let verifier = URL_SAFE_NO_PAD.encode(&verifier_bytes);
let mut hasher = Sha256::new();
hasher.update(&verifier);
let challenge = URL_SAFE_NO_PAD.encode(hasher.finalize());
PkceParams {
verifier,
challenge,
}
}
pub fn build_auth_url(config: &OAuth2Config, state: &str, pkce_challenge: &str) -> Result<String> {
validate_url(&config.auth_url).context("Invalid auth URL - security validation failed")?;
let mut url = url::Url::parse(&config.auth_url).context("Invalid auth URL")?;
{
let mut params = url.query_pairs_mut();
params.append_pair("client_id", &config.client_id);
params.append_pair("redirect_uri", &config.redirect_uri);
params.append_pair("response_type", "code");
params.append_pair("state", state);
params.append_pair("code_challenge", pkce_challenge);
params.append_pair("code_challenge_method", "S256");
if !config.scopes.is_empty() {
params.append_pair("scope", &config.scopes.join(" "));
}
if let Some(additional) = &config.additional_params {
for (key, value) in additional {
params.append_pair(key, value);
}
}
}
Ok(url.to_string())
}
pub async fn exchange_code_for_tokens(
config: &OAuth2Config,
code: &str,
pkce_verifier: &str,
) -> Result<OAuth2Token> {
validate_url(&config.token_url).context("Invalid token URL - security validation failed")?;
let client = reqwest::Client::new();
let mut params = HashMap::new();
params.insert("grant_type", "authorization_code");
params.insert("code", code);
params.insert("redirect_uri", &config.redirect_uri);
params.insert("client_id", &config.client_id);
params.insert("code_verifier", pkce_verifier);
let mut request = client.post(&config.token_url);
if let Some(client_secret) = &config.client_secret {
if config.provider == "github" {
params.insert("client_secret", client_secret);
} else {
request = request.basic_auth(&config.client_id, Some(client_secret));
}
}
let response = request
.form(¶ms)
.header("Accept", "application/json")
.send()
.await
.context("Failed to exchange code for tokens")?;
if !response.status().is_success() {
let error_text = response.text().await?;
return Err(ApiError::AuthError(format!("Token exchange failed: {}", error_text)).into());
}
let token_response: TokenResponse = response
.json()
.await
.context("Failed to parse token response")?;
let expires_at = token_response
.expires_in
.map(|seconds| Utc::now() + Duration::seconds(seconds));
let scopes = token_response
.scope
.map(|s| s.split_whitespace().map(String::from).collect())
.unwrap_or_else(|| config.scopes.clone());
Ok(OAuth2Token {
access_token: token_response.access_token,
token_type: token_response.token_type,
expires_at,
refresh_token: token_response.refresh_token,
scopes,
})
}
pub async fn refresh_access_token(
config: &OAuth2Config,
refresh_token: &str,
) -> Result<OAuth2Token> {
validate_url(&config.token_url).context("Invalid token URL - security validation failed")?;
let client = reqwest::Client::new();
let mut params = HashMap::new();
params.insert("grant_type", "refresh_token");
params.insert("refresh_token", refresh_token);
params.insert("client_id", &config.client_id);
let mut request = client.post(&config.token_url);
if let Some(client_secret) = &config.client_secret {
request = request.basic_auth(&config.client_id, Some(client_secret));
}
let response = request
.form(¶ms)
.header("Accept", "application/json")
.send()
.await
.context("Failed to refresh token")?;
if !response.status().is_success() {
let error_text = response.text().await?;
return Err(ApiError::AuthError(format!("Token refresh failed: {}", error_text)).into());
}
let token_response: TokenResponse = response
.json()
.await
.context("Failed to parse refresh response")?;
let expires_at = token_response
.expires_in
.map(|seconds| Utc::now() + Duration::seconds(seconds));
let scopes = token_response
.scope
.map(|s| s.split_whitespace().map(String::from).collect())
.unwrap_or_else(|| config.scopes.clone());
Ok(OAuth2Token {
access_token: token_response.access_token,
token_type: token_response.token_type,
expires_at,
refresh_token: token_response
.refresh_token
.or(Some(refresh_token.to_string())),
scopes,
})
}
pub async fn oauth_login(config: OAuth2Config, profile_name: String) -> Result<()> {
use colored::*;
use std::time::Duration;
println!("🔐 {} OAuth Authentication", "Starting".bright_cyan());
println!(" Provider: {}", config.provider.bright_yellow());
println!(" Profile: {}", profile_name.bright_green());
let state = generate_state();
let pkce = generate_pkce_pair();
let (tx, rx) = std::sync::mpsc::channel();
let server_handle = crate::core::auth::server::start_callback_server(tx, state.clone()).await?;
let auth_url = build_auth_url(&config, &state, &pkce.challenge)?;
println!("\n🌐 Opening browser for authentication...");
println!(" If browser doesn't open, visit:");
println!(" {}", auth_url.bright_blue());
if let Err(e) = open::that(&auth_url) {
eprintln!(" ⚠️ Could not open browser: {}", e);
}
println!("\n⏳ Waiting for authorization (timeout: 5 minutes)...");
let auth_code = rx
.recv_timeout(Duration::from_secs(300))
.context("Authorization timeout - no response received")?;
drop(server_handle);
println!("✅ Authorization code received!");
println!("🔄 Exchanging code for tokens...");
let tokens = exchange_code_for_tokens(&config, &auth_code, &pkce.verifier).await?;
let profile = AuthProfile {
name: profile_name.clone(),
provider: config.provider.clone(),
created_at: Utc::now(),
last_used: None,
metadata: HashMap::new(),
};
crate::core::auth::token_store::store_profile(&profile)?;
crate::core::auth::token_store::store_tokens(&profile_name, &tokens)?;
crate::core::auth::providers::store_provider_config(&profile_name, &config)?;
println!("\n✅ {} successful!", "Authentication".bright_green());
println!(
" Profile '{}' created and ready to use",
profile_name.bright_yellow()
);
println!(
"\n Use it with: {}",
format!("mrapids run <operation> --profile {}", profile_name).bright_cyan()
);
Ok(())
}
pub async fn refresh_tokens(profile: &str) -> Result<OAuth2Token> {
use colored::*;
println!(
"🔄 Refreshing tokens for profile '{}'...",
profile.bright_yellow()
);
let tokens = crate::core::auth::token_store::load_tokens(profile)?;
let config = crate::core::auth::providers::load_provider_config(profile)?;
if let Some(refresh_token) = &tokens.refresh_token {
let new_tokens = refresh_access_token(&config, refresh_token).await?;
crate::core::auth::token_store::store_tokens(profile, &new_tokens)?;
println!("✅ Tokens refreshed successfully!");
Ok(new_tokens)
} else {
Err(ApiError::AuthError(format!(
"No refresh token available for profile '{}'",
profile
))
.into())
}
}
impl OAuth2Token {
pub fn is_expired(&self) -> bool {
crate::core::auth::is_token_expired(&self.expires_at)
}
pub fn auth_header(&self) -> String {
format!("{} {}", self.token_type, self.access_token)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_config() -> OAuth2Config {
OAuth2Config {
provider: "test-provider".to_string(),
client_id: "test-client-id".to_string(),
client_secret: Some("test-client-secret".to_string()),
auth_url: "https://auth.example.com/authorize".to_string(),
token_url: "https://auth.example.com/token".to_string(),
redirect_uri: "http://localhost:8080/callback".to_string(),
scopes: vec!["read".to_string(), "write".to_string()],
additional_params: None,
}
}
#[test]
fn test_oauth2_config_creation() {
let config = create_test_config();
assert_eq!(config.provider, "test-provider");
assert_eq!(config.client_id, "test-client-id");
assert!(config.client_secret.is_some());
assert_eq!(config.scopes.len(), 2);
}
#[test]
fn test_oauth2_config_serialization() {
let config = create_test_config();
let json = serde_json::to_string(&config).unwrap();
assert!(json.contains("test-client-id"));
assert!(json.contains("test-provider"));
}
#[test]
fn test_oauth2_config_deserialization() {
let json = r#"{
"provider": "github",
"client_id": "abc123",
"auth_url": "https://github.com/login/oauth/authorize",
"token_url": "https://github.com/login/oauth/access_token",
"redirect_uri": "http://localhost:8080",
"scopes": ["user", "repo"]
}"#;
let config: OAuth2Config = serde_json::from_str(json).unwrap();
assert_eq!(config.provider, "github");
assert_eq!(config.client_id, "abc123");
assert!(config.client_secret.is_none());
}
fn create_test_token() -> OAuth2Token {
OAuth2Token {
access_token: "test-access-token".to_string(),
token_type: "Bearer".to_string(),
expires_at: Some(Utc::now() + Duration::hours(1)),
refresh_token: Some("test-refresh-token".to_string()),
scopes: vec!["read".to_string()],
}
}
#[test]
fn test_oauth2_token_creation() {
let token = create_test_token();
assert_eq!(token.access_token, "test-access-token");
assert_eq!(token.token_type, "Bearer");
assert!(token.refresh_token.is_some());
}
#[test]
fn test_oauth2_token_auth_header() {
let token = create_test_token();
let header = token.auth_header();
assert_eq!(header, "Bearer test-access-token");
}
#[test]
fn test_oauth2_token_auth_header_custom_type() {
let mut token = create_test_token();
token.token_type = "CustomAuth".to_string();
let header = token.auth_header();
assert_eq!(header, "CustomAuth test-access-token");
}
#[test]
fn test_oauth2_token_not_expired() {
let token = create_test_token(); assert!(!token.is_expired());
}
#[test]
fn test_oauth2_token_expired() {
let token = OAuth2Token {
access_token: "test".to_string(),
token_type: "Bearer".to_string(),
expires_at: Some(Utc::now() - Duration::hours(1)), refresh_token: None,
scopes: vec![],
};
assert!(token.is_expired());
}
#[test]
fn test_oauth2_token_no_expiry() {
let token = OAuth2Token {
access_token: "test".to_string(),
token_type: "Bearer".to_string(),
expires_at: None, refresh_token: None,
scopes: vec![],
};
assert!(!token.is_expired());
}
#[test]
fn test_oauth2_token_serialization() {
let token = create_test_token();
let json = serde_json::to_string(&token).unwrap();
assert!(json.contains("test-access-token"));
assert!(json.contains("Bearer"));
}
#[test]
fn test_auth_profile_creation() {
let profile = AuthProfile {
name: "my-profile".to_string(),
provider: "github".to_string(),
created_at: Utc::now(),
last_used: None,
metadata: HashMap::new(),
};
assert_eq!(profile.name, "my-profile");
assert_eq!(profile.provider, "github");
assert!(profile.last_used.is_none());
}
#[test]
fn test_auth_profile_with_metadata() {
let mut metadata = HashMap::new();
metadata.insert("user_id".to_string(), Value::String("12345".to_string()));
metadata.insert(
"username".to_string(),
Value::String("testuser".to_string()),
);
let profile = AuthProfile {
name: "my-profile".to_string(),
provider: "github".to_string(),
created_at: Utc::now(),
last_used: Some(Utc::now()),
metadata,
};
assert!(profile.metadata.contains_key("user_id"));
assert!(profile.last_used.is_some());
}
#[test]
fn test_generate_state_length() {
let state = generate_state();
assert!(state.len() >= 40);
}
#[test]
fn test_generate_state_unique() {
let state1 = generate_state();
let state2 = generate_state();
assert_ne!(state1, state2, "States should be unique");
}
#[test]
fn test_generate_state_url_safe() {
let state = generate_state();
assert!(!state.contains('+'));
assert!(!state.contains('/'));
assert!(!state.contains('='));
}
#[test]
fn test_generate_pkce_pair() {
let pkce = generate_pkce_pair();
assert!(!pkce.verifier.is_empty());
assert!(!pkce.challenge.is_empty());
}
#[test]
fn test_generate_pkce_verifier_length() {
let pkce = generate_pkce_pair();
assert!(pkce.verifier.len() >= 40);
}
#[test]
fn test_generate_pkce_challenge_is_hash() {
let pkce = generate_pkce_pair();
assert!(pkce.challenge.len() >= 40);
}
#[test]
fn test_generate_pkce_pair_unique() {
let pkce1 = generate_pkce_pair();
let pkce2 = generate_pkce_pair();
assert_ne!(pkce1.verifier, pkce2.verifier);
assert_ne!(pkce1.challenge, pkce2.challenge);
}
#[test]
fn test_generate_pkce_url_safe() {
let pkce = generate_pkce_pair();
assert!(!pkce.verifier.contains('+'));
assert!(!pkce.verifier.contains('/'));
assert!(!pkce.challenge.contains('+'));
assert!(!pkce.challenge.contains('/'));
}
#[test]
fn test_build_auth_url_basic() {
let config = create_test_config();
let state = "test-state-123";
let challenge = "test-challenge-456";
let url = build_auth_url(&config, state, challenge).unwrap();
assert!(url.starts_with("https://auth.example.com/authorize"));
assert!(url.contains("client_id=test-client-id"));
assert!(url.contains("redirect_uri="));
assert!(url.contains("response_type=code"));
assert!(url.contains("state=test-state-123"));
assert!(url.contains("code_challenge=test-challenge-456"));
assert!(url.contains("code_challenge_method=S256"));
}
#[test]
fn test_build_auth_url_with_scopes() {
let config = create_test_config();
let url = build_auth_url(&config, "state", "challenge").unwrap();
assert!(url.contains("scope=read+write") || url.contains("scope=read%20write"));
}
#[test]
fn test_build_auth_url_with_additional_params() {
let mut config = create_test_config();
let mut additional = HashMap::new();
additional.insert("prompt".to_string(), "consent".to_string());
additional.insert("access_type".to_string(), "offline".to_string());
config.additional_params = Some(additional);
let url = build_auth_url(&config, "state", "challenge").unwrap();
assert!(url.contains("prompt=consent"));
assert!(url.contains("access_type=offline"));
}
#[test]
fn test_build_auth_url_no_scopes() {
let mut config = create_test_config();
config.scopes = vec![];
let url = build_auth_url(&config, "state", "challenge").unwrap();
assert!(url.contains("client_id=test-client-id"));
}
#[test]
fn test_build_auth_url_invalid_url() {
let mut config = create_test_config();
config.auth_url = "not-a-valid-url".to_string();
let result = build_auth_url(&config, "state", "challenge");
assert!(result.is_err());
}
#[test]
fn test_pkce_params_struct() {
let params = PkceParams {
verifier: "test-verifier".to_string(),
challenge: "test-challenge".to_string(),
};
assert_eq!(params.verifier, "test-verifier");
assert_eq!(params.challenge, "test-challenge");
}
}