use anyhow::{Context, Result};
use clap::Args;
use colored::*;
use prettytable::{format, Cell, Row, Table};
use reqwest;
use std::collections::HashMap;
use std::time::Instant;
use crate::core::api::ApiError;
use crate::core::auth::selection::AuthSelector;
use crate::core::spec::UnifiedSpec;
use crate::models::auth::ValidationStatus;
#[derive(Debug, Args)]
pub struct ValidateCommand {
#[arg(short, long)]
pub scheme: Option<String>,
#[arg(long)]
pub spec: Option<String>,
#[arg(short = 'e', long)]
pub endpoint: Option<String>,
#[arg(short, long)]
pub verbose: bool,
#[arg(long)]
pub debug: bool,
#[arg(short, long)]
pub quick: bool,
}
#[derive(Debug)]
struct ValidationResult {
scheme_name: String,
status: ValidationStatus,
response_time: Option<u128>,
status_code: Option<u16>,
error_message: Option<String>,
suggestions: Vec<String>,
}
impl ValidateCommand {
pub async fn execute(&self) -> Result<()> {
println!("\n{}", "Authentication Validation".bold().cyan());
println!("{}", "═".repeat(60).cyan());
let spec = self.load_spec().await?;
let schemes = self.extract_schemes(&spec)?;
if schemes.is_empty() {
println!(
"{} No authentication schemes found in the specification",
"ℹ".blue()
);
return Ok(());
}
let mut selector = AuthSelector::new();
selector.discover_credentials(&schemes)?;
let schemes_to_validate = if let Some(scheme_name) = &self.scheme {
if !schemes.contains_key(scheme_name) {
return Err(ApiError::ValidationError(format!(
"Unknown authentication scheme: {}",
scheme_name
))
.into());
}
vec![scheme_name.clone()]
} else {
schemes.keys().cloned().collect()
};
let mut results = Vec::new();
for scheme_name in schemes_to_validate {
println!("\n{} Validating {}...", "→".yellow(), scheme_name.bold());
let result = if self.quick {
self.quick_validate(&scheme_name, &selector).await
} else {
self.full_validate(&scheme_name, &selector, &spec).await
};
results.push(result);
}
self.display_results(&results, &schemes)?;
let all_valid = results.iter().all(|r| {
matches!(
r.status,
ValidationStatus::Valid | ValidationStatus::NotValidated
)
});
if !all_valid {
return Err(ApiError::AuthError(
"Some authentication schemes failed validation".to_string(),
)
.into());
}
Ok(())
}
async fn load_spec(&self) -> Result<UnifiedSpec> {
let spec_path = self
.spec
.as_ref()
.map(|s| s.as_str())
.or_else(|| {
for path in &[
"openapi.yaml",
"openapi.json",
"swagger.yaml",
"swagger.json",
] {
if std::path::Path::new(path).exists() {
return Some(*path);
}
}
None
})
.context("No OpenAPI specification found. Use --spec to specify the file")?;
UnifiedSpec::from_file(spec_path).context("Failed to load OpenAPI specification")
}
fn extract_schemes(
&self,
spec: &UnifiedSpec,
) -> Result<HashMap<String, crate::models::auth::SecuritySchemeDetails>> {
let mut schemes = HashMap::new();
for (name, unified_scheme) in &spec.security_schemes {
let scheme_type = match unified_scheme.scheme_type.as_str() {
"apiKey" => crate::models::auth::SchemeType::ApiKey,
"http" => crate::models::auth::SchemeType::Http,
"oauth2" => crate::models::auth::SchemeType::OAuth2,
"openIdConnect" => crate::models::auth::SchemeType::OpenIdConnect,
"mutualTLS" => crate::models::auth::SchemeType::MutualTls,
_ => crate::models::auth::SchemeType::Http,
};
let location = unified_scheme
.location
.as_ref()
.and_then(|loc| match loc.as_str() {
"query" => Some(crate::models::auth::AuthLocation::Query),
"header" => Some(crate::models::auth::AuthLocation::Header),
"cookie" => Some(crate::models::auth::AuthLocation::Cookie),
_ => None,
});
schemes.insert(
name.clone(),
crate::models::auth::SecuritySchemeDetails {
scheme_type,
location,
name: unified_scheme.name.clone(),
bearer_format: unified_scheme.bearer_format.clone(),
flows: None, openid_connect_url: unified_scheme.openid_connect_url.clone(),
description: unified_scheme.description.clone(),
},
);
}
Ok(schemes)
}
async fn quick_validate(&self, scheme_name: &str, selector: &AuthSelector) -> ValidationResult {
let validation_result = selector.validate_selection(scheme_name);
match validation_result {
Ok(_) => ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::Valid,
response_time: None,
status_code: None,
error_message: None,
suggestions: vec!["Credentials are configured and appear valid".to_string()],
},
Err(e) => {
let mut suggestions = Vec::new();
let error_msg = e.to_string();
if error_msg.contains("not configured") {
suggestions.push(format!("Run: mrapids auth connect {}", scheme_name));
ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::NotValidated,
response_time: None,
status_code: None,
error_message: Some(error_msg),
suggestions,
}
} else {
suggestions.push("Check your credential configuration".to_string());
ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::Invalid(error_msg.clone()),
response_time: None,
status_code: None,
error_message: Some(error_msg),
suggestions,
}
}
}
}
}
async fn full_validate(
&self,
scheme_name: &str,
selector: &AuthSelector,
spec: &UnifiedSpec,
) -> ValidationResult {
let quick_result = self.quick_validate(scheme_name, selector).await;
if !matches!(quick_result.status, ValidationStatus::Valid) {
return quick_result;
}
let test_endpoint = if let Some(endpoint) = &self.endpoint {
endpoint.clone()
} else {
if let Some(endpoint) = self.find_test_endpoint(spec, scheme_name) {
endpoint
} else {
return ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::Valid,
response_time: None,
status_code: None,
error_message: None,
suggestions: vec![
"Credentials configured but no test endpoint available".to_string(),
"Use --endpoint to specify a test endpoint".to_string(),
],
};
}
};
if self.verbose {
println!(" Testing endpoint: {}", test_endpoint.cyan());
}
let start = Instant::now();
let client = reqwest::Client::new();
let mut request = client.get(&test_endpoint);
request = self.apply_test_auth(request, scheme_name);
match request.send().await {
Ok(response) => {
let elapsed = start.elapsed().as_millis();
let status = response.status();
if self.debug {
println!(" Response: {} in {}ms", status, elapsed);
}
if status.is_success() {
ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::Valid,
response_time: Some(elapsed),
status_code: Some(status.as_u16()),
error_message: None,
suggestions: vec!["Authentication successful".to_string()],
}
} else {
let body = response.text().await.unwrap_or_default();
let mut suggestions = Vec::new();
if status == reqwest::StatusCode::UNAUTHORIZED {
suggestions.push("Credentials may be invalid or expired".to_string());
suggestions
.push(format!("Try: mrapids auth connect {} --force", scheme_name));
} else if status == reqwest::StatusCode::FORBIDDEN {
suggestions
.push("Credentials valid but lack required permissions".to_string());
}
ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::Invalid(format!(
"HTTP {}: {}",
status,
if body.len() > 100 {
&body[..100]
} else {
&body
}
)),
response_time: Some(elapsed),
status_code: Some(status.as_u16()),
error_message: Some(format!(
"HTTP {}: {}",
status,
if body.len() > 100 {
&body[..100]
} else {
&body
}
)),
suggestions,
}
}
}
Err(e) => ValidationResult {
scheme_name: scheme_name.to_string(),
status: ValidationStatus::Invalid(e.to_string()),
response_time: None,
status_code: None,
error_message: Some(e.to_string()),
suggestions: vec![
"Check network connectivity".to_string(),
"Verify the endpoint URL is correct".to_string(),
],
},
}
}
fn find_test_endpoint(&self, spec: &UnifiedSpec, scheme_name: &str) -> Option<String> {
for (path, path_item) in &spec.paths {
for (method, operation) in &path_item.operations {
if method != "get" {
continue;
}
let uses_scheme = if let Some(security) = &operation.security {
security.iter().any(|req| req.contains_key(scheme_name))
} else if let Some(global_security) = &spec.security {
global_security
.iter()
.any(|req| req.contains_key(scheme_name))
} else {
false
};
if uses_scheme {
if let Some(server) = spec.servers.first() {
return Some(format!("{}{}", server.url, path));
}
}
}
}
None
}
fn apply_test_auth(
&self,
request: reqwest::RequestBuilder,
scheme_name: &str,
) -> reqwest::RequestBuilder {
if scheme_name.contains("bearer") || scheme_name.contains("jwt") {
if let Ok(token) = std::env::var("BEARER_TOKEN") {
return request.bearer_auth(token);
}
} else if scheme_name.contains("api_key") || scheme_name.contains("apikey") {
if let Ok(key) = std::env::var("API_KEY") {
return request.header("X-API-Key", key);
}
}
request
}
fn display_results(
&self,
results: &[ValidationResult],
schemes: &HashMap<String, crate::models::auth::SecuritySchemeDetails>,
) -> Result<()> {
println!("\n{}", "Validation Results".bold().green());
println!("{}", "─".repeat(60).green());
let mut table = Table::new();
table.set_format(*format::consts::FORMAT_NO_LINESEP_WITH_TITLE);
table.set_titles(Row::new(vec![
Cell::new("Scheme").style_spec("b"),
Cell::new("Status").style_spec("b"),
Cell::new("Response").style_spec("b"),
Cell::new("Details").style_spec("b"),
]));
let mut has_failures = false;
let mut has_warnings = false;
for result in results {
let status_display = match result.status {
ValidationStatus::Valid => "✓ Valid".green(),
ValidationStatus::Invalid(_) => {
has_failures = true;
"✗ Invalid".red()
}
ValidationStatus::NotValidated => {
has_warnings = true;
"○ Not Configured".yellow()
}
ValidationStatus::Expired => {
has_failures = true;
"⚠ Expired".yellow()
}
};
let response_display = if let Some(code) = result.status_code {
if let Some(time) = result.response_time {
format!("{} ({}ms)", code, time)
} else {
code.to_string()
}
} else {
"-".to_string()
};
let details = if let Some(error) = &result.error_message {
if error.len() > 50 {
format!("{}...", &error[..50])
} else {
error.clone()
}
} else if !result.suggestions.is_empty() {
result.suggestions[0].clone()
} else {
"-".to_string()
};
table.add_row(Row::new(vec![
Cell::new(&result.scheme_name),
Cell::new(&status_display.to_string()),
Cell::new(&response_display),
Cell::new(&details),
]));
}
table.printstd();
if has_failures || (has_warnings && self.verbose) {
println!("\n{}", "Troubleshooting".bold().yellow());
println!("{}", "─".repeat(60).yellow());
for result in results {
if matches!(
result.status,
ValidationStatus::Invalid(_) | ValidationStatus::NotValidated
) || (matches!(result.status, ValidationStatus::NotValidated) && self.verbose)
{
println!(
"\n{}: {}",
result.scheme_name.bold(),
match result.status {
ValidationStatus::Invalid(_) => "Invalid credentials",
ValidationStatus::NotValidated => "Not configured",
ValidationStatus::Expired => "Expired credentials",
_ => "Issue detected",
}
);
if let Some(error) = &result.error_message {
println!(" Error: {}", error.red());
}
for suggestion in &result.suggestions {
println!(" • {}", suggestion);
}
if let Some(scheme) = schemes.get(&result.scheme_name) {
match scheme.scheme_type {
crate::models::auth::SchemeType::ApiKey => {
println!(
" $ mrapids auth connect {} --auth-type api-key",
result.scheme_name
);
}
crate::models::auth::SchemeType::Http => {
if scheme.bearer_format.is_some() {
println!(
" $ mrapids auth connect {} --auth-type bearer",
result.scheme_name
);
} else {
println!(
" $ mrapids auth connect {} --auth-type basic",
result.scheme_name
);
}
}
crate::models::auth::SchemeType::OAuth2 => {
println!(
" $ mrapids auth connect {} --auth-type oauth2",
result.scheme_name
);
}
_ => {}
}
}
}
}
}
let valid_count = results
.iter()
.filter(|r| matches!(r.status, ValidationStatus::Valid))
.count();
let configured_count = results
.iter()
.filter(|r| !matches!(r.status, ValidationStatus::NotValidated))
.count();
let total = results.len();
println!("\n{}", "Summary".bold().cyan());
println!("{}", "─".repeat(60).cyan());
println!(" Total schemes: {}", total);
println!(" Configured: {}/{}", configured_count, total);
println!(
" Valid: {}/{}",
valid_count.to_string().green(),
configured_count
);
if valid_count == configured_count && configured_count > 0 {
println!(
"\n{} All configured authentication schemes are valid!",
"✓".green().bold()
);
} else if has_failures {
println!(
"\n{} Some authentication schemes need attention",
"⚠".yellow().bold()
);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validation_status_priority() {
assert!(matches!(
ValidationStatus::NotValidated,
ValidationStatus::NotValidated
));
assert!(matches!(
ValidationStatus::Invalid(String::from("test")),
ValidationStatus::Invalid(_)
));
assert!(matches!(ValidationStatus::Valid, ValidationStatus::Valid));
}
}