use crate::core::spec::{UnifiedSecurityScheme, UnifiedSpec};
use crate::models::auth::{
AuthComplexity, AuthLocation, AuthSummary, OAuth2Flow, OAuth2Flows, RequirementOption,
SchemeRequirement, SchemeType, SecurityAnalysis, SecurityRequirement, SecuritySchemeDetails,
};
use anyhow::Result;
use std::collections::{HashMap, HashSet};
pub struct AuthDetector {
spec: UnifiedSpec,
}
impl AuthDetector {
pub fn new(spec: UnifiedSpec) -> Self {
Self { spec }
}
pub fn analyze(&self) -> Result<SecurityAnalysis> {
let schemes = self.extract_schemes()?;
let global_requirements = self.extract_global_requirements()?;
let operation_requirements = self.extract_operation_requirements()?;
let summary =
self.generate_summary(&schemes, &global_requirements, &operation_requirements)?;
Ok(SecurityAnalysis {
schemes,
global_requirements,
operation_requirements,
summary,
})
}
fn extract_schemes(&self) -> Result<HashMap<String, SecuritySchemeDetails>> {
let mut schemes = HashMap::new();
for (name, scheme) in &self.spec.security_schemes {
let details = self.convert_scheme_to_details(scheme)?;
schemes.insert(name.clone(), details);
}
Ok(schemes)
}
fn convert_scheme_to_details(
&self,
scheme: &UnifiedSecurityScheme,
) -> Result<SecuritySchemeDetails> {
let scheme_type = match scheme.scheme_type.as_str() {
"apiKey" => SchemeType::ApiKey,
"http" => SchemeType::Http,
"oauth2" => SchemeType::OAuth2,
"openIdConnect" => SchemeType::OpenIdConnect,
"mutualTLS" => SchemeType::MutualTls,
_ => SchemeType::Http, };
let location = scheme.location.as_ref().and_then(|loc| match loc.as_str() {
"query" => Some(AuthLocation::Query),
"header" => Some(AuthLocation::Header),
"cookie" => Some(AuthLocation::Cookie),
_ => None,
});
let flows = if scheme_type == SchemeType::OAuth2 {
Some(self.extract_oauth2_flows(scheme)?)
} else {
None
};
let bearer_format = if scheme_type == SchemeType::Http {
if let Some(http_scheme) = &scheme.scheme {
if http_scheme.to_lowercase() == "bearer" {
Some("JWT".to_string()) } else {
None
}
} else {
scheme.bearer_format.clone()
}
} else {
scheme.bearer_format.clone()
};
Ok(SecuritySchemeDetails {
scheme_type,
location,
name: scheme.name.clone(),
bearer_format,
flows,
openid_connect_url: scheme.openid_connect_url.clone(),
description: scheme.description.clone(),
})
}
fn extract_oauth2_flows(&self, scheme: &UnifiedSecurityScheme) -> Result<OAuth2Flows> {
let mut flows = OAuth2Flows {
implicit: None,
password: None,
client_credentials: None,
authorization_code: None,
device_code: None,
};
if let Some(flow_type) = &scheme.flow {
let flow = OAuth2Flow {
authorization_url: scheme.authorization_url.clone(),
token_url: scheme.token_url.clone(),
refresh_url: scheme.refresh_url.clone(),
scopes: scheme.scopes.clone().unwrap_or_default(),
device_url: None, };
match flow_type.as_str() {
"implicit" => flows.implicit = Some(flow),
"password" => flows.password = Some(flow),
"clientCredentials" | "client_credentials" => flows.client_credentials = Some(flow),
"authorizationCode" | "authorization_code" => flows.authorization_code = Some(flow),
_ => {}
}
}
Ok(flows)
}
fn extract_global_requirements(&self) -> Result<SecurityRequirement> {
let mut requirement = SecurityRequirement::default();
if let Some(security) = &self.spec.security {
for security_req in security {
let mut option_schemes = Vec::new();
for (scheme_name, scopes) in security_req {
option_schemes.push(SchemeRequirement {
name: scheme_name.clone(),
scopes: scopes.clone(),
});
}
if !option_schemes.is_empty() {
requirement.options.push(RequirementOption {
schemes: option_schemes,
});
}
}
}
Ok(requirement)
}
fn extract_operation_requirements(&self) -> Result<HashMap<String, SecurityRequirement>> {
let mut requirements = HashMap::new();
for (path, path_item) in &self.spec.paths {
for (method, operation) in &path_item.operations {
let operation_id = operation
.operation_id
.clone()
.unwrap_or_else(|| format!("{} {}", method.to_uppercase(), path));
if let Some(security) = &operation.security {
let mut requirement = SecurityRequirement::default();
for security_req in security {
let mut option_schemes = Vec::new();
for (scheme_name, scopes) in security_req {
option_schemes.push(SchemeRequirement {
name: scheme_name.clone(),
scopes: scopes.clone(),
});
}
if !option_schemes.is_empty() {
requirement.options.push(RequirementOption {
schemes: option_schemes,
});
}
}
requirements.insert(operation_id, requirement);
}
}
}
Ok(requirements)
}
fn generate_summary(
&self,
schemes: &HashMap<String, SecuritySchemeDetails>,
global: &SecurityRequirement,
operations: &HashMap<String, SecurityRequirement>,
) -> Result<AuthSummary> {
let total_schemes = schemes.len();
let mut required_schemes = HashSet::new();
let mut optional_schemes = HashSet::new();
let mut scheme_usage_count = HashMap::new();
for option in &global.options {
for scheme_req in &option.schemes {
required_schemes.insert(scheme_req.name.clone());
*scheme_usage_count
.entry(scheme_req.name.clone())
.or_insert(0) += 1;
}
}
for (_, req) in operations {
for option in &req.options {
for scheme_req in &option.schemes {
required_schemes.insert(scheme_req.name.clone());
*scheme_usage_count
.entry(scheme_req.name.clone())
.or_insert(0) += 1;
}
}
}
for scheme_name in schemes.keys() {
if !required_schemes.contains(scheme_name) {
optional_schemes.insert(scheme_name.clone());
}
}
let most_common_scheme = scheme_usage_count
.iter()
.max_by_key(|(_, count)| *count)
.map(|(name, _)| name.clone());
let complexity_score =
self.calculate_complexity(total_schemes, operations.len(), &required_schemes);
let total_operations = self.count_total_operations();
let operations_with_custom_auth = operations.len();
let operations_without_auth = total_operations
.saturating_sub(operations_with_custom_auth)
.saturating_sub(if global.options.is_empty() {
0
} else {
total_operations - operations_with_custom_auth
});
Ok(AuthSummary {
total_schemes,
required_schemes,
optional_schemes,
operations_with_custom_auth,
operations_without_auth,
most_common_scheme,
complexity_score,
})
}
fn calculate_complexity(
&self,
total_schemes: usize,
custom_operations: usize,
required_schemes: &HashSet<String>,
) -> AuthComplexity {
if required_schemes.is_empty() {
AuthComplexity::None
} else if total_schemes == 1 && custom_operations == 0 {
AuthComplexity::Simple
} else if total_schemes <= 3 && custom_operations <= 5 {
AuthComplexity::Moderate
} else {
AuthComplexity::Complex
}
}
fn count_total_operations(&self) -> usize {
self.spec
.paths
.values()
.map(|path| path.operations.len())
.sum()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_security_requirement_satisfaction() {
let mut requirement = SecurityRequirement::default();
requirement.options.push(RequirementOption {
schemes: vec![SchemeRequirement {
name: "api_key".to_string(),
scopes: vec![],
}],
});
requirement.options.push(RequirementOption {
schemes: vec![
SchemeRequirement {
name: "oauth2".to_string(),
scopes: vec!["read".to_string()],
},
SchemeRequirement {
name: "basic".to_string(),
scopes: vec![],
},
],
});
let mut available = HashSet::new();
available.insert("api_key".to_string());
assert!(requirement.is_satisfied_by(&available));
let mut available = HashSet::new();
available.insert("oauth2".to_string());
available.insert("basic".to_string());
assert!(requirement.is_satisfied_by(&available));
let mut available = HashSet::new();
available.insert("oauth2".to_string());
assert!(!requirement.is_satisfied_by(&available));
}
}