use std::{
panic::AssertUnwindSafe,
sync::Arc,
time::{Duration, Instant},
};
use std::sync::{
LazyLock, OnceLock,
atomic::{AtomicBool, Ordering},
};
use anyhow::Result;
use dashmap::DashMap;
use futures::FutureExt;
use http::StatusCode;
use kingfisher_core::ValidationOutcome;
use liquid::Object;
use liquid_core::{Value, ValueView};
use reqwest::Client;
#[cfg(test)]
use reqwest::{
header,
header::{HeaderMap, HeaderValue},
};
use rustc_hash::FxHashMap;
use tokio::{sync::Notify, time};
use tracing::warn;
use crate::{
cli::global::TlsMode,
location::OffsetSpan,
matcher::{OwnedBlobMatch, SerializableCaptures},
provider_endpoints::{
ProviderEndpointOverrides, endpoint_var_names, hydrate_endpoint_globals_for_rule,
},
rules::rule::Validation,
validation_body::{self},
};
use crate::validation_rate_limit::should_rate_limit_validation;
pub use kingfisher_rules::TlsMode as RuleTlsMode;
pub use kingfisher_scanner::validation::CachedResponse;
pub use kingfisher_scanner::validation::aws;
pub use kingfisher_scanner::validation::http_validation as httpvalidation;
pub use kingfisher_scanner::validation::mysql::validate_mysql;
pub use kingfisher_scanner::validation::postgres::validate_postgres;
pub use kingfisher_scanner::validation::{
azure, coinbase, gcp, jdbc, jwt, mongodb, mysql, postgres,
};
pub(crate) mod candidates;
pub mod utils;
fn truncate_to_char_boundary(s: &mut String, max_len: usize) {
if s.len() <= max_len {
return;
}
let mut new_len = max_len;
while new_len > 0 && !s.is_char_boundary(new_len) {
new_len -= 1;
}
s.truncate(new_len);
}
fn truncate_preview(body: &str, max_len: usize) -> String {
if max_len == 0 || body.len() <= max_len {
return body.to_string();
}
let mut preview = body.to_string();
truncate_to_char_boundary(&mut preview, max_len);
preview
}
static USER_AGENT_SUFFIX: OnceLock<String> = OnceLock::new();
const BROWSER_USER_AGENT: &str = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) \
AppleWebKit/537.36 (KHTML, like Gecko) \
Chrome/140.0.0.0 Safari/537.36";
fn build_user_agent() -> String {
let base = format!("kingfisher/{}", env!("CARGO_PKG_VERSION"));
if let Some(suffix) = USER_AGENT_SUFFIX.get() {
format!("{base} {suffix} {BROWSER_USER_AGENT}")
} else {
format!("{base} {BROWSER_USER_AGENT}")
}
}
pub static GLOBAL_USER_AGENT: LazyLock<String> = LazyLock::new(build_user_agent);
pub fn set_user_agent_suffix<S: Into<String>>(suffix: Option<S>) {
if let Some(suffix) = suffix {
let trimmed = suffix.into().trim().to_string();
if trimmed.is_empty() {
return;
}
let _ = USER_AGENT_SUFFIX.set(trimmed.clone());
kingfisher_scanner::validation::set_user_agent_suffix(Some(trimmed));
}
}
#[derive(Clone)]
pub struct ValidationClients {
strict: Client,
credential_uri: Client,
credential_uri_lax: Client,
lax: Client,
pub global_mode: TlsMode,
pub allow_internal_ips: bool,
}
impl ValidationClients {
pub fn new(global_mode: TlsMode, allow_internal_ips: bool) -> anyhow::Result<Self> {
let timeout = std::time::Duration::from_secs(30);
let strict = Client::builder()
.user_agent(GLOBAL_USER_AGENT.as_str())
.danger_accept_invalid_certs(false)
.redirect(reqwest::redirect::Policy::none())
.timeout(timeout)
.build()?;
let lax = Client::builder()
.user_agent(GLOBAL_USER_AGENT.as_str())
.danger_accept_invalid_certs(true)
.redirect(reqwest::redirect::Policy::none())
.timeout(timeout)
.build()?;
let credential_uri = build_credential_uri_client(timeout, false)?;
let credential_uri_lax = build_credential_uri_client(timeout, true)?;
Ok(Self {
strict,
credential_uri,
credential_uri_lax,
lax,
global_mode,
allow_internal_ips,
})
}
pub fn client_for_rule(&self, rule_tls_mode: Option<kingfisher_rules::TlsMode>) -> &Client {
match self.global_mode {
TlsMode::Off => &self.lax,
TlsMode::Lax => {
let rule_wants_lax = matches!(rule_tls_mode, Some(kingfisher_rules::TlsMode::Lax));
if rule_wants_lax { &self.lax } else { &self.strict }
}
TlsMode::Strict => &self.strict,
}
}
pub fn credential_uri_client(
&self,
rule_tls_mode: Option<kingfisher_rules::TlsMode>,
) -> &Client {
if self.should_use_lax(rule_tls_mode) {
&self.credential_uri_lax
} else {
&self.credential_uri
}
}
pub fn should_use_lax(&self, rule_tls_mode: Option<kingfisher_rules::TlsMode>) -> bool {
match self.global_mode {
TlsMode::Off => true,
TlsMode::Lax => matches!(rule_tls_mode, Some(kingfisher_rules::TlsMode::Lax)),
TlsMode::Strict => false,
}
}
}
pub(crate) fn build_credential_uri_client(timeout: Duration, use_lax_tls: bool) -> Result<Client> {
Ok(Client::builder()
.danger_accept_invalid_certs(use_lax_tls)
.redirect(reqwest::redirect::Policy::none())
.timeout(timeout)
.build()?)
}
type Cache = kingfisher_scanner::validation::Cache;
pub(crate) fn validation_context_values(m: &OwnedBlobMatch) -> Vec<(String, Option<String>)> {
let mut values: std::collections::BTreeMap<String, Option<String>> =
utils::process_captures(&m.captures)
.into_iter()
.map(|(name, value, ..)| (name, Some(value)))
.collect();
for dep in m.rule.syntax().depends_on_rule.iter().flatten() {
if dep.variable.eq_ignore_ascii_case("TOKEN") {
continue;
}
let variable = dep.variable.to_uppercase();
if let Some(value) = m.dependent_captures.get(&variable) {
values.insert(variable, Some(value.clone()));
} else {
values.entry(variable).or_insert(None);
}
}
values.into_iter().collect()
}
pub(crate) fn skip_ambiguous_dependencies(m: &mut OwnedBlobMatch) -> bool {
if m.ambiguous_dependencies.is_empty() {
return false;
}
let details = m
.ambiguous_dependencies
.iter()
.map(|(variable, count)| format!("{variable}: {count} distinct candidates"))
.collect::<Vec<_>>()
.join(", ");
m.validation_success = false;
m.validation_response_status = StatusCode::PRECONDITION_REQUIRED;
m.validation_response_body = validation_body::from_string(format!(
"Validation skipped - ambiguous dependency: {details}. Narrow the rule's within window or provide explicit context."
));
m.refresh_validation_outcome();
true
}
fn validation_dedup_key(m: &OwnedBlobMatch) -> [u8; 32] {
let mut hasher = blake3::Hasher::new();
hasher.update(b"kingfisher.validation-dedup.v1\0");
hash_key_part(&mut hasher, m.rule.syntax().id.as_bytes());
let capture_value = if matches!(&m.rule.syntax().validation, Some(Validation::CredentialUri)) {
m.captures
.captures
.iter()
.find(|capture| capture.name.is_some_and(|name| name.eq_ignore_ascii_case("URI")))
.or_else(|| m.captures.captures.first())
.map(|capture| capture.raw_value())
} else {
m.captures.captures.first().map(|capture| capture.raw_value())
};
if let Some(val) = capture_value {
hash_key_part(&mut hasher, val.as_bytes());
}
for (variable, count) in &m.ambiguous_dependencies {
hash_key_part(&mut hasher, variable.as_bytes());
hash_key_part(&mut hasher, &count.to_le_bytes());
}
for (variable, value) in validation_context_values(m) {
hash_key_part(&mut hasher, variable.as_bytes());
if let Some(value) = value {
hash_key_part(&mut hasher, value.as_bytes());
} else {
hash_key_part(&mut hasher, b"<missing>");
}
}
*hasher.finalize().as_bytes()
}
fn hash_key_part(hasher: &mut blake3::Hasher, part: &[u8]) {
hasher.update(&(part.len() as u64).to_le_bytes());
hasher.update(part);
}
static VALIDATION_CACHE: OnceLock<DashMap<[u8; 32], CachedResponse>> = OnceLock::new();
struct InFlightValidation {
notify: Notify,
completed: AtomicBool,
}
static IN_FLIGHT: OnceLock<DashMap<[u8; 32], Arc<InFlightValidation>>> = OnceLock::new();
fn cache_validation_result(fp: [u8; 32], m: &OwnedBlobMatch) {
VALIDATION_CACHE.get_or_init(DashMap::new).insert(
fp,
CachedResponse {
body: m.validation_response_body.clone(),
status: m.validation_response_status,
is_valid: m.validation_success,
outcome: m.validation_outcome,
timestamp: Instant::now(),
},
);
}
fn clear_in_flight_validation(fp: [u8; 32]) {
if let Some((_, in_flight)) = IN_FLIGHT.get_or_init(DashMap::new).remove(&fp) {
in_flight.completed.store(true, Ordering::Release);
in_flight.notify.notify_waiters();
}
}
pub fn init_validation_caches() {
VALIDATION_CACHE.set(DashMap::new()).ok();
IN_FLIGHT.set(DashMap::new()).ok();
aws::set_aws_validation_concurrency(15);
}
pub fn clear_validation_caches() {
if let Some(c) = VALIDATION_CACHE.get() {
c.clear();
c.shrink_to_fit();
}
if let Some(c) = IN_FLIGHT.get() {
c.clear();
c.shrink_to_fit();
}
}
pub fn set_skip_aws_account_ids<I, S>(ids: I)
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
aws::set_aws_skip_account_ids(ids);
}
#[cfg(test)]
use kingfisher_scanner::validation::aws::validate_aws_credential_pair;
pub(crate) use kingfisher_scanner::validation::credential_uri::{
CredentialUriTarget, classify_credential_uri, is_parseable_credential_uri,
};
pub use kingfisher_scanner::validation::credential_uri::{
is_parseable_mongodb_uri, is_parseable_mysql_uri, is_parseable_postgres_uri,
};
#[cfg(test)]
use kingfisher_scanner::validation::credential_uri::{
received_basic_auth_challenge, validate_http_credential_uri,
};
#[allow(clippy::type_complexity)]
pub fn collect_variables_and_dependencies(
matches: &[OwnedBlobMatch],
) -> (FxHashMap<String, Vec<(String, OffsetSpan)>>, FxHashMap<String, Vec<String>>) {
let mut variable_map: FxHashMap<String, Vec<(String, OffsetSpan)>> = FxHashMap::default();
let mut missing_deps: FxHashMap<String, Vec<String>> = FxHashMap::default();
for m in matches {
let rule_id = m.rule.syntax().id.clone();
for dependency in m.rule.syntax().depends_on_rule.iter().flatten() {
if dependency.within.is_some()
|| m.dependent_captures.contains_key(&dependency.variable.to_uppercase())
{
let variable = dependency.variable.to_uppercase();
if let Some(value) = m.dependent_captures.get(&variable) {
variable_map
.entry(variable)
.or_default()
.push((value.clone(), m.matching_input_offset_span));
} else if !dependency.optional {
missing_deps
.entry(rule_id.clone())
.or_default()
.push(dependency.rule_id.clone());
}
continue;
}
let dependency_rule_id = &dependency.rule_id;
let matching_dependencies: Vec<_> =
matches.iter().filter(|x| x.rule.syntax().id == *dependency_rule_id).collect();
if !matching_dependencies.is_empty() {
for other_match in matching_dependencies {
let matching_input = other_match
.captures
.captures
.first()
.expect("Expected at least one capture");
variable_map.entry(dependency.variable.to_uppercase()).or_default().push((
matching_input.raw_value().to_string(),
other_match.matching_input_offset_span,
));
}
} else if !dependency.optional {
missing_deps.entry(rule_id.clone()).or_default().push(dependency.rule_id.clone());
}
}
}
(variable_map, missing_deps)
}
#[allow(clippy::too_many_arguments)]
pub async fn validate_single_match(
m: &mut OwnedBlobMatch,
parser: &liquid::Parser,
clients: &ValidationClients,
dependent_variables: &FxHashMap<String, Vec<(String, OffsetSpan)>>,
missing_dependencies: &FxHashMap<String, Vec<String>>,
cache: &Cache,
validation_timeout: Duration,
validation_retries: u32,
rate_limiter: Option<&crate::validation_rate_limit::ValidationRateLimiter>,
provider_endpoints: &ProviderEndpointOverrides,
max_body_len: usize,
) {
if candidates::ready(m) {
candidates::run(m, validation_timeout, |mut attempt, remaining| async move {
let mut globals = Object::new();
populate_globals_from_captures(
&mut globals,
&utils::process_captures(&attempt.captures),
);
for dep in attempt.rule.syntax().depends_on_rule.iter().flatten() {
let name = dep.variable.to_uppercase();
if name != "TOKEN"
&& let Some(value) = attempt.dependent_captures.get(&name)
{
globals.insert(name.into(), Value::scalar(value.clone()));
}
}
hydrate_endpoint_globals_for_rule(attempt.rule.id(), &mut globals);
provider_endpoints.apply_scan_overrides(&mut globals);
for name in endpoint_var_names() {
if let Some(value) = globals.get(*name).and_then(|v| v.as_scalar()) {
attempt
.dependent_captures
.entry((*name).to_string())
.or_insert_with(|| value.to_kstr().to_string());
}
}
validate_resolved_match(
&mut attempt,
parser,
clients,
dependent_variables,
missing_dependencies,
cache,
remaining,
validation_retries,
rate_limiter,
provider_endpoints,
max_body_len,
)
.await;
attempt
})
.await;
} else {
validate_resolved_match(
m,
parser,
clients,
dependent_variables,
missing_dependencies,
cache,
validation_timeout,
validation_retries,
rate_limiter,
provider_endpoints,
max_body_len,
)
.await;
}
}
#[allow(clippy::too_many_arguments)]
async fn validate_resolved_match(
m: &mut OwnedBlobMatch,
parser: &liquid::Parser,
clients: &ValidationClients,
dependent_variables: &FxHashMap<String, Vec<(String, OffsetSpan)>>,
missing_dependencies: &FxHashMap<String, Vec<String>>,
cache: &Cache,
validation_timeout: Duration,
validation_retries: u32,
rate_limiter: Option<&crate::validation_rate_limit::ValidationRateLimiter>,
provider_endpoints: &ProviderEndpointOverrides,
max_body_len: usize,
) {
if !m.rule.syntax().is_authoritative() {
m.validation_success = false;
m.validation_response_status = StatusCode::CONTINUE;
m.validation_response_body = None;
m.validation_outcome = ValidationOutcome::NotAttempted;
return;
}
if skip_ambiguous_dependencies(m) {
return;
}
let fp = validation_dedup_key(m);
let mut owns_in_flight = false;
let timeout_result = time::timeout(
validation_timeout,
AssertUnwindSafe(
timed_validate_single_match(
&mut owns_in_flight,
m,
parser,
clients,
dependent_variables,
missing_dependencies,
cache,
validation_timeout,
validation_retries,
rate_limiter,
provider_endpoints,
max_body_len,
)
.boxed(),
)
.catch_unwind(),
)
.await;
match timeout_result {
Ok(Ok(())) => {}
Ok(Err(_panic_payload)) => {
warn!(
rule_id = %m.rule.syntax().id,
"validator panicked; marking match as failed",
);
m.validation_success = false;
m.validation_response_body = validation_body::from_string(format!(
"Validation panicked for rule {}",
m.rule.syntax().id
));
m.validation_response_status = StatusCode::INTERNAL_SERVER_ERROR;
m.validation_outcome = ValidationOutcome::Unavailable;
if owns_in_flight {
cache_validation_result(fp, m);
clear_in_flight_validation(fp);
}
}
Err(_) => {
m.validation_success = false;
m.validation_response_body = validation_body::from_string(format!(
"Validation timed out after {} seconds",
validation_timeout.as_secs()
));
m.validation_response_status = StatusCode::REQUEST_TIMEOUT;
m.validation_outcome = ValidationOutcome::Unavailable;
if owns_in_flight {
cache_validation_result(fp, m);
clear_in_flight_validation(fp);
}
}
}
}
#[allow(clippy::too_many_arguments)]
async fn timed_validate_single_match(
owns_in_flight: &mut bool,
m: &mut OwnedBlobMatch,
parser: &liquid::Parser,
clients: &ValidationClients,
_dependent_variables: &FxHashMap<String, Vec<(String, OffsetSpan)>>,
missing_dependencies: &FxHashMap<String, Vec<String>>,
_cache: &Cache,
validation_timeout: Duration,
validation_retries: u32,
rate_limiter: Option<&crate::validation_rate_limit::ValidationRateLimiter>,
provider_endpoints: &ProviderEndpointOverrides,
max_body_len: usize,
) {
let rule_tls_mode = m.rule.tls_mode();
let client =
if m.rule.syntax().depends_on_rule.iter().flatten().any(|dep| dep.verify_candidates) {
clients.credential_uri_client(rule_tls_mode)
} else {
clients.client_for_rule(rule_tls_mode)
};
let use_lax_tls = clients.should_use_lax(rule_tls_mode);
let fp = validation_dedup_key(m);
if let Some(entry) = VALIDATION_CACHE.get_or_init(DashMap::new).get(&fp) {
m.validation_success = entry.is_valid;
m.validation_response_body = entry.body.clone();
m.validation_response_status = entry.status;
m.validation_outcome = entry.outcome;
return;
}
let in_flight =
Arc::new(InFlightValidation { notify: Notify::new(), completed: AtomicBool::new(false) });
let wait = match IN_FLIGHT.get_or_init(DashMap::new).entry(fp) {
dashmap::mapref::entry::Entry::Occupied(entry) => Some(entry.get().clone()),
dashmap::mapref::entry::Entry::Vacant(entry) => {
entry.insert(Arc::clone(&in_flight));
*owns_in_flight = true;
None
}
};
if let Some(wait) = wait {
loop {
let notified = wait.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if wait.completed.load(Ordering::Acquire) {
break;
}
notified.await;
}
if let Some(entry) = VALIDATION_CACHE.get().unwrap().get(&fp) {
m.validation_success = entry.is_valid;
m.validation_response_body = entry.body.clone();
m.validation_response_status = entry.status;
m.validation_outcome = entry.outcome;
}
return;
}
let commit_and_return = |m: &OwnedBlobMatch| {
cache_validation_result(fp, m);
clear_in_flight_validation(fp);
};
if let Some(missing) = missing_dependencies.get(&m.rule.syntax().id)
&& !missing.is_empty()
&& m.rule.syntax().depends_on_rule.iter().flatten().any(|dep| {
!dep.optional
&& !m.dependent_captures.contains_key(&dep.variable.to_uppercase())
&& missing.contains(&dep.rule_id)
})
{
m.validation_success = false;
m.validation_response_body = validation_body::from_string(format!(
"Validation skipped - missing dependent rules: {}",
missing.join(", ")
));
m.validation_response_status = StatusCode::PRECONDITION_REQUIRED;
m.validation_outcome = ValidationOutcome::Skipped;
commit_and_return(m);
return;
}
let match_re_result = m.rule.syntax().as_anchored_regex();
let mut captured_values: Vec<(String, String, usize, usize)> = match match_re_result {
Ok(_) => utils::process_captures(&m.captures),
Err(e) => {
m.validation_success = false;
m.validation_response_body =
validation_body::from_string(format!("Regex error: {}", e));
m.validation_response_status = StatusCode::INTERNAL_SERVER_ERROR;
m.validation_outcome = ValidationOutcome::Unavailable;
commit_and_return(m);
return;
}
};
for dep in m.rule.syntax().depends_on_rule.iter().flatten() {
if dep.variable.eq_ignore_ascii_case("TOKEN") {
continue;
}
let dep_name = dep.variable.to_uppercase();
if let Some(value) = m.dependent_captures.get(&dep_name).cloned() {
captured_values.push((
dep_name.clone(),
value,
m.matching_input_offset_span.start,
m.matching_input_offset_span.end,
));
continue;
}
}
let mut globals = Object::new();
populate_globals_from_captures(&mut globals, &captured_values);
hydrate_endpoint_globals_for_rule(m.rule.id(), &mut globals);
provider_endpoints.apply_scan_overrides(&mut globals);
for (k, v, ..) in &captured_values {
if k.eq_ignore_ascii_case("TOKEN") {
continue;
}
m.dependent_captures.entry(k.to_uppercase()).or_insert_with(|| v.clone());
}
for endpoint_var in endpoint_var_names() {
if let Some(value) = globals.get(*endpoint_var).and_then(|v| v.as_scalar()) {
m.dependent_captures
.entry((*endpoint_var).to_string())
.or_insert_with(|| value.to_kstr().to_string());
}
}
{
let rule_syntax = m.rule.syntax();
if let (Some(limiter), Some(validation)) = (rate_limiter, rule_syntax.validation.as_ref())
&& should_rate_limit_validation(validation)
{
limiter.wait_for_rule(m.rule.id()).await;
}
}
let result = kingfisher_scanner::validation::ValidationEngine::new(client, parser)
.credential_uri_client(clients.credential_uri_client(rule_tls_mode))
.timeout(validation_timeout)
.retries(validation_retries)
.allow_internal_ips(clients.allow_internal_ips)
.use_lax_tls(use_lax_tls)
.validate(&m.rule, &globals)
.await;
m.validation_outcome = result.outcome;
m.validation_success = result.outcome.is_verified_active();
m.validation_response_status = result
.http_status
.and_then(|s| StatusCode::from_u16(s).ok())
.unwrap_or(StatusCode::CONTINUE);
let body = if result.response_body.is_empty() {
result.reason.map(|reason| format!("Validation {:?}", reason)).unwrap_or_default()
} else {
result.response_body
};
let response_is_html = matches!(m.rule.syntax().validation.as_ref(), Some(Validation::Http(config)) if config.request.response_is_html);
m.validation_response_body = validation_body::from_string(if response_is_html {
utils::format_response_body_for_display(&body, max_body_len, true)
} else {
truncate_preview(&body, max_body_len)
});
commit_and_return(m);
}
pub(crate) use kingfisher_scanner::validation::engine::is_aws_session_token_rule;
fn populate_globals_from_captures(
globals: &mut Object,
captured_values: &[(String, String, usize, usize)],
) {
let mut best_token: Option<&String> = None;
for (k, v, ..) in captured_values {
if k.eq_ignore_ascii_case("TOKEN") {
if best_token.is_none_or(|best| v.len() >= best.len()) {
best_token = Some(v);
}
} else {
globals.insert(k.to_uppercase().into(), Value::scalar(v.clone()));
}
}
if let Some(token) = best_token {
globals.insert("TOKEN".into(), Value::scalar(token.clone()));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::rules::rule::{Confidence, DependsOnRule, Rule, RuleSyntax};
#[test]
fn credential_uri_classifier_normalizes_supported_database_schemes() {
assert_eq!(
classify_credential_uri(
"POSTGRESQL://alice:hunter2@db.internal:5432/app",
Some("POSTGRESQL")
),
CredentialUriTarget::Postgres(
"postgresql://alice:hunter2@db.internal:5432/app".to_string()
)
);
assert_eq!(
classify_credential_uri(
"MARIADB://alice:hunter2@db.internal:3306/app",
Some("MARIADB")
),
CredentialUriTarget::MySQL("mysql://alice:hunter2@db.internal:3306/app".to_string())
);
assert!(is_parseable_credential_uri(
"mongodb://alice:hunter2@mongo.internal:27017/app",
Some("mongodb")
));
}
#[test]
fn credential_uri_classifier_accepts_http_basic_auth_uris() {
assert_eq!(
classify_credential_uri("https://alice:hunter2@service.internal", Some("https")),
CredentialUriTarget::Http("https://alice:hunter2@service.internal".to_string())
);
assert!(is_parseable_credential_uri(
"https://alice:hunter2@service.internal",
Some("https")
));
assert!(!is_parseable_credential_uri(
"postgresql://alice:hunter2@db.internal:70000/app",
Some("postgresql")
));
let malformed_https =
classify_credential_uri("https://alice@service.internal", Some("https"));
assert_eq!(malformed_https.scheme(), "https");
assert!(!malformed_https.is_parseable());
}
#[tokio::test]
async fn credential_uri_client_does_not_follow_redirects() {
use axum::{Router, response::Redirect, routing::get};
let app = Router::new()
.route("/challenge", get(|| async { Redirect::temporary("/target") }))
.route("/target", get(|| async { "redirected" }));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
for use_lax_tls in [false, true] {
let client = build_credential_uri_client(Duration::from_secs(5), use_lax_tls).unwrap();
let response = client.get(format!("http://{address}/challenge")).send().await.unwrap();
assert_eq!(response.status(), StatusCode::TEMPORARY_REDIRECT);
assert_eq!(response.url().path(), "/challenge");
}
server.abort();
}
#[test]
fn http_credential_uri_requires_basic_challenge_before_authentication() {
let mut headers = HeaderMap::new();
headers
.append(header::WWW_AUTHENTICATE, HeaderValue::from_static("Digest realm=\"example\""));
assert!(!received_basic_auth_challenge(StatusCode::OK, &headers));
assert!(!received_basic_auth_challenge(StatusCode::UNAUTHORIZED, &headers));
headers.append(
header::WWW_AUTHENTICATE,
HeaderValue::from_static("Digest realm=\"example\", basic = \"metadata\""),
);
assert!(!received_basic_auth_challenge(StatusCode::UNAUTHORIZED, &headers));
headers
.append(header::WWW_AUTHENTICATE, HeaderValue::from_static("Basic realm=\"example\""));
assert!(received_basic_auth_challenge(StatusCode::UNAUTHORIZED, &headers));
}
#[tokio::test]
async fn http_credential_uri_rejects_plaintext_transport() {
let client = Client::builder().redirect(reqwest::redirect::Policy::none()).build().unwrap();
let result = validate_http_credential_uri(
"http://alice:hunter2@service.internal/health",
&client,
Duration::from_secs(5),
0,
true,
)
.await
.unwrap_err();
assert!(result.to_string().contains("requires HTTPS"));
}
#[tokio::test]
async fn http_credential_uri_rejects_ambiguous_or_invalid_userinfo() {
let client = Client::builder().redirect(reqwest::redirect::Policy::none()).build().unwrap();
let cases = [
(
"https://alice%3Aadmin:hunter2@service.invalid/protected",
"usernames cannot contain ':'",
),
("https://alice%FF:hunter2@service.invalid/protected", "username is not valid UTF-8"),
("https://alice:hunter%FF@service.invalid/protected", "password is not valid UTF-8"),
];
for (uri, expected_error) in cases {
let error = validate_http_credential_uri(uri, &client, Duration::from_secs(5), 0, true)
.await
.unwrap_err();
assert!(error.to_string().contains(expected_error), "{error}");
}
}
async fn run_https_basic_auth_flow(
authenticated_status: StatusCode,
mode: TlsMode,
rule_mode: Option<kingfisher_rules::TlsMode>,
) -> ((bool, StatusCode, String), Vec<Option<String>>) {
use rcgen::{CertifiedKey, generate_simple_self_signed};
use rustls::{ServerConfig, pki_types::PrivatePkcs8KeyDer};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio_rustls::TlsAcceptor;
let CertifiedKey { cert, signing_key } =
generate_simple_self_signed(vec!["127.0.0.1".to_string()]).unwrap();
let cert_der = cert.der().clone();
let key = PrivatePkcs8KeyDer::from(signing_key.serialize_der());
let provider = rustls::crypto::aws_lc_rs::default_provider();
let config = ServerConfig::builder_with_provider(Arc::new(provider))
.with_safe_default_protocol_versions()
.unwrap()
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key.into())
.unwrap();
let acceptor = TlsAcceptor::from(Arc::new(config));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let mut authorization_headers = Vec::new();
for status in [StatusCode::UNAUTHORIZED, authenticated_status] {
let (socket, _) = listener.accept().await.unwrap();
let Ok(mut stream) = acceptor.accept(socket).await else {
break;
};
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
loop {
let read = stream.read(&mut buffer).await.unwrap();
if read == 0 {
break;
}
request.extend_from_slice(&buffer[..read]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
let request = String::from_utf8(request).unwrap();
authorization_headers.push(
request
.lines()
.find_map(|line| line.strip_prefix("authorization: "))
.map(str::to_string),
);
let status_line = match status {
StatusCode::OK => "200 OK",
StatusCode::UNAUTHORIZED => "401 Unauthorized",
other => panic!("unsupported test status: {other}"),
};
let challenge = if status == StatusCode::UNAUTHORIZED {
"WWW-Authenticate: Basic realm=\"test\"\r\n"
} else {
""
};
let response = format!(
"HTTP/1.1 {status_line}\r\n{challenge}Content-Length: 0\r\nConnection: close\r\n\r\n"
);
stream.write_all(response.as_bytes()).await.unwrap();
stream.shutdown().await.unwrap();
}
authorization_headers
});
let clients = ValidationClients::new(mode, true).unwrap();
let client = clients.credential_uri_client(rule_mode);
let result = validate_http_credential_uri(
&format!("https://alice:hunter2@{address}/protected"),
client,
Duration::from_secs(5),
0,
true,
)
.await;
let authorization_headers = server.await.unwrap();
let result = result.unwrap_or_else(|error| {
assert!(authorization_headers.is_empty());
(false, StatusCode::BAD_GATEWAY, error.to_string())
});
(result, authorization_headers)
}
#[tokio::test]
async fn validation_clients_send_user_agent() {
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::method};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(|req: &Request| {
req.headers.get(header::USER_AGENT).and_then(|v| v.to_str().ok())
== Some(GLOBAL_USER_AGENT.as_str())
})
.respond_with(ResponseTemplate::new(200))
.expect(3)
.mount(&server)
.await;
for mode in [TlsMode::Strict, TlsMode::Lax, TlsMode::Off] {
let clients = ValidationClients::new(mode, true).unwrap();
let response = clients.client_for_rule(None).get(server.uri()).send().await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
}
#[tokio::test]
async fn http_credential_uri_sends_basic_auth_only_after_https_challenge() {
let ((valid, status, _), authorization_headers) =
run_https_basic_auth_flow(StatusCode::OK, TlsMode::Off, None).await;
assert!(valid);
assert_eq!(status, StatusCode::OK);
assert_eq!(authorization_headers[0], None);
assert_eq!(authorization_headers[1].as_deref(), Some("Basic YWxpY2U6aHVudGVyMg=="));
}
#[tokio::test]
async fn http_credential_uri_rejects_untrusted_tls_without_opt_in() {
for (mode, rule_mode) in [
(TlsMode::Strict, None),
(TlsMode::Strict, Some(kingfisher_rules::TlsMode::Lax)),
(TlsMode::Lax, None),
] {
let ((valid, status, message), headers) =
run_https_basic_auth_flow(StatusCode::OK, mode, rule_mode).await;
assert!(!valid);
assert_eq!(status, StatusCode::BAD_GATEWAY);
assert!(message.contains("challenge request failed"));
assert!(headers.is_empty());
}
}
#[tokio::test]
async fn http_credential_uri_honors_rule_lax_opt_in() {
let ((valid, _, _), headers) = run_https_basic_auth_flow(
StatusCode::OK,
TlsMode::Lax,
Some(kingfisher_rules::TlsMode::Lax),
)
.await;
assert!(valid);
assert_eq!(headers[0], None);
assert!(headers[1].is_some());
}
#[tokio::test]
async fn http_credential_uri_treats_repeated_unauthorized_as_inactive() {
let ((valid, status, _), authorization_headers) =
run_https_basic_auth_flow(StatusCode::UNAUTHORIZED, TlsMode::Off, None).await;
assert!(!valid);
assert_eq!(status, StatusCode::UNAUTHORIZED);
assert_eq!(authorization_headers[0], None);
assert!(authorization_headers[1].is_some());
}
fn aws_rule(id: &str, secret_dependency: bool) -> Rule {
Rule::new(RuleSyntax {
name: id.to_string(),
id: id.to_string(),
pattern: "(secret)".to_string(),
min_entropy: 0.0,
confidence: Confidence::Low,
visible: true,
examples: Vec::new(),
negative_examples: Vec::new(),
references: Vec::new(),
validation: Some(Validation::AWS),
revocation: None,
depends_on_rule: secret_dependency
.then(|| DependsOnRule {
rule_id: "private.aws.secret".to_string(),
verify_candidates: false,
variable: "AWS_SECRET_ACCESS_KEY".to_string(),
optional: false,
within: None,
})
.into_iter()
.map(Some)
.collect(),
pattern_requirements: None,
tls_mode: None,
path: None,
betterleaks_filter: None,
betterleaks_secret_group: None,
authoritative: true,
vectorscan_compatible: true,
})
}
#[tokio::test]
async fn waiter_timeout_preserves_owner_and_cached_result() {
let mut matched = OwnedBlobMatch {
rule: Arc::new(aws_rule("test.waiter-timeout", false)),
blob_id: crate::blob::BlobId::new(b"waiter-timeout"),
finding_fingerprint: 1,
matching_input_offset_span: OffsetSpan::from_range(0..6),
captures: crate::matcher::SerializableCaptures { captures: Default::default() },
validation_response_body: None,
validation_response_status: StatusCode::CONTINUE,
validation_success: false,
validation_outcome: ValidationOutcome::NotAttempted,
calculated_entropy: 0.0,
is_base64: false,
dependent_captures: Default::default(),
ambiguous_dependencies: Default::default(),
dependency_candidates: Default::default(),
};
let variables = FxHashMap::default();
let fp = validation_dedup_key(&matched);
let owner = Arc::new(InFlightValidation {
notify: Notify::new(),
completed: AtomicBool::new(false),
});
IN_FLIGHT.get_or_init(DashMap::new).insert(fp, owner.clone());
let parser = liquid::ParserBuilder::with_stdlib().build().unwrap();
let clients = ValidationClients::new(TlsMode::Strict, true).unwrap();
validate_single_match(
&mut matched,
&parser,
&clients,
&variables,
&FxHashMap::default(),
&Cache::default(),
Duration::from_millis(10),
0,
None,
&ProviderEndpointOverrides::default(),
1024,
)
.await;
assert_eq!(matched.validation_response_status, StatusCode::REQUEST_TIMEOUT);
assert!(Arc::ptr_eq(IN_FLIGHT.get().unwrap().get(&fp).unwrap().value(), &owner));
assert!(!owner.completed.load(Ordering::Acquire));
assert!(!VALIDATION_CACHE.get().unwrap().contains_key(&fp));
matched.validation_success = true;
matched.validation_response_status = StatusCode::OK;
matched.refresh_validation_outcome();
cache_validation_result(fp, &matched);
clear_in_flight_validation(fp);
matched.validation_success = false;
validate_single_match(
&mut matched,
&parser,
&clients,
&variables,
&FxHashMap::default(),
&Cache::default(),
Duration::from_millis(10),
0,
None,
&ProviderEndpointOverrides::default(),
1024,
)
.await;
assert!(matched.validation_success);
assert_eq!(matched.validation_response_status, StatusCode::OK);
VALIDATION_CACHE.get().unwrap().remove(&fp);
}
#[test]
fn populate_globals_prefers_longest_token() {
let captured_values = vec![
("TOKEN".to_string(), "short".to_string(), 0usize, 5usize),
("BODY".to_string(), "body".to_string(), 0usize, 4usize),
("TOKEN".to_string(), "longervalue".to_string(), 0usize, 11usize),
];
let mut globals = Object::new();
populate_globals_from_captures(&mut globals, &captured_values);
assert_eq!(globals.get("TOKEN"), Some(Value::scalar("longervalue")).as_ref());
assert_eq!(globals.get("BODY"), Some(Value::scalar("body")).as_ref());
}
#[test]
fn populate_globals_handles_missing_token() {
let captured_values = vec![("CHECKSUM".to_string(), "123456".to_string(), 0usize, 6usize)];
let mut globals = Object::new();
populate_globals_from_captures(&mut globals, &captured_values);
assert!(globals.get("TOKEN").is_none());
assert_eq!(globals.get("CHECKSUM"), Some(Value::scalar("123456")).as_ref());
}
#[test]
fn associated_betterleaks_component_wins_over_other_blob_values() {
use crate::{
blob::BlobId,
matcher::{SerializableCapture, SerializableCaptures},
util::intern,
};
use smallvec::smallvec;
let mut syntax = aws_rule("betterleaks.primary", false).syntax().clone();
syntax.validation = None;
syntax.depends_on_rule = vec![Some(DependsOnRule {
rule_id: "betterleaks.component".into(),
verify_candidates: false,
variable: "COMPONENT".into(),
optional: false,
within: Some("5L".into()),
})];
let mut primary = OwnedBlobMatch {
rule: Arc::new(Rule::new(syntax)),
blob_id: BlobId::new(b"associated-component"),
finding_fingerprint: 1,
matching_input_offset_span: OffsetSpan::from_range(20..27),
captures: SerializableCaptures {
captures: smallvec![SerializableCapture {
name: Some(intern("TOKEN")),
match_number: 1,
start: 20,
end: 27,
value: "primary".into(),
}],
},
validation_response_body: None,
validation_response_status: StatusCode::CONTINUE,
validation_success: false,
validation_outcome: ValidationOutcome::NotAttempted,
calculated_entropy: 0.0,
is_base64: false,
dependent_captures: std::collections::BTreeMap::new(),
ambiguous_dependencies: Default::default(),
dependency_candidates: Default::default(),
};
primary.dependent_captures.insert("COMPONENT".into(), "associated".into());
let (variables, missing) = collect_variables_and_dependencies(&[primary]);
assert_eq!(variables["COMPONENT"][0].0, "associated");
assert!(missing.is_empty());
}
#[test]
fn betterleaks_aws_session_token_rule_is_detected_by_its_secret_dependency() {
let rule = aws_rule("betterleaks.aws-session-token", true);
assert!(is_aws_session_token_rule(&rule));
}
#[tokio::test]
async fn shared_aws_validation_skips_canaries_before_network_access() {
let result =
validate_aws_credential_pair("AKIAXYZDQCEN4B6JSJQI", "not-a-real-secret-key", None)
.await;
assert!(!result.is_valid);
assert_eq!(result.status, StatusCode::PRECONDITION_REQUIRED);
assert_eq!(result.outcome, ValidationOutcome::Skipped);
assert_eq!(result.account_id.as_deref(), Some("534261010715"));
assert!(result.message.starts_with("(skip list entry)"));
}
#[test]
fn truncate_to_char_boundary_handles_multibyte_characters() {
let max_len = 2048;
let mut body = "a".repeat(max_len);
body.push('é');
truncate_to_char_boundary(&mut body, max_len);
assert_eq!(body.len(), max_len);
assert!(body.is_char_boundary(body.len()));
assert!(body.ends_with('a'));
}
#[test]
fn truncate_skipped_when_max_body_len_is_zero() {
let original_len = 4096;
let body = "x".repeat(original_len);
let preview = truncate_preview(&body, 0);
assert_eq!(preview.len(), original_len);
}
#[test]
fn truncate_applies_custom_max_body_len() {
let body = "y".repeat(5000);
let preview = truncate_preview(&body, 1024);
assert_eq!(preview.len(), 1024);
}
mod tls_mode_tests {
use super::*;
#[test]
fn validation_clients_new_creates_both_clients() {
let clients = ValidationClients::new(TlsMode::Strict, false).unwrap();
assert_eq!(clients.global_mode, TlsMode::Strict);
let clients_lax = ValidationClients::new(TlsMode::Lax, false).unwrap();
assert_eq!(clients_lax.global_mode, TlsMode::Lax);
let clients_off = ValidationClients::new(TlsMode::Off, false).unwrap();
assert_eq!(clients_off.global_mode, TlsMode::Off);
}
#[test]
fn client_for_rule_strict_mode_always_returns_strict_client() {
let clients = ValidationClients::new(TlsMode::Strict, false).unwrap();
let client1 = clients.client_for_rule(None);
let client2 = clients.client_for_rule(Some(kingfisher_rules::TlsMode::Lax));
let client3 = clients.client_for_rule(Some(kingfisher_rules::TlsMode::Strict));
assert!(std::ptr::eq(client1, client2));
assert!(std::ptr::eq(client2, client3));
}
#[test]
fn client_for_rule_off_mode_always_returns_lax_client() {
let clients = ValidationClients::new(TlsMode::Off, false).unwrap();
let client1 = clients.client_for_rule(None);
let client2 = clients.client_for_rule(Some(kingfisher_rules::TlsMode::Lax));
let client3 = clients.client_for_rule(Some(kingfisher_rules::TlsMode::Strict));
assert!(std::ptr::eq(client1, client2));
assert!(std::ptr::eq(client2, client3));
}
#[test]
fn client_for_rule_lax_mode_respects_rule_preference() {
let clients = ValidationClients::new(TlsMode::Lax, false).unwrap();
let strict_client = clients.client_for_rule(None);
let lax_client = clients.client_for_rule(Some(kingfisher_rules::TlsMode::Lax));
assert!(std::ptr::eq(clients.client_for_rule(None), strict_client));
assert!(std::ptr::eq(
clients.client_for_rule(Some(kingfisher_rules::TlsMode::Strict)),
strict_client
));
assert!(std::ptr::eq(
clients.client_for_rule(Some(kingfisher_rules::TlsMode::Lax)),
lax_client
));
assert!(!std::ptr::eq(strict_client, lax_client));
}
#[test]
fn should_use_lax_off_mode_always_returns_true() {
let clients = ValidationClients::new(TlsMode::Off, false).unwrap();
assert!(clients.should_use_lax(None));
assert!(clients.should_use_lax(Some(kingfisher_rules::TlsMode::Strict)));
assert!(clients.should_use_lax(Some(kingfisher_rules::TlsMode::Lax)));
}
#[test]
fn builtin_database_validators_opt_into_lax_tls() {
let rules = kingfisher_rules::get_builtin_rules(None).unwrap();
for id in ["betterleaks.mongodb-connection-string", "betterleaks.jwt"] {
let tls_mode = rules.rules.get(id).unwrap().tls_mode;
assert_eq!(tls_mode, Some(kingfisher_rules::TlsMode::Lax), "{id}");
assert!(
!ValidationClients::new(TlsMode::Strict, false)
.unwrap()
.should_use_lax(tls_mode),
"{id} must stay strict under the default --tls-mode strict"
);
assert!(
ValidationClients::new(TlsMode::Lax, false).unwrap().should_use_lax(tls_mode),
"{id} must go lax under --tls-mode lax"
);
}
let private_key = rules.rules.get("betterleaks.private-key").unwrap();
assert_eq!(private_key.tls_mode, None);
assert!(
!ValidationClients::new(TlsMode::Lax, false)
.unwrap()
.should_use_lax(private_key.tls_mode)
);
}
#[test]
fn should_use_lax_strict_mode_always_returns_false() {
let clients = ValidationClients::new(TlsMode::Strict, false).unwrap();
assert!(!clients.should_use_lax(None));
assert!(!clients.should_use_lax(Some(kingfisher_rules::TlsMode::Strict)));
assert!(!clients.should_use_lax(Some(kingfisher_rules::TlsMode::Lax)));
}
#[test]
fn should_use_lax_lax_mode_respects_rule_preference() {
let clients = ValidationClients::new(TlsMode::Lax, false).unwrap();
assert!(!clients.should_use_lax(None));
assert!(!clients.should_use_lax(Some(kingfisher_rules::TlsMode::Strict)));
assert!(clients.should_use_lax(Some(kingfisher_rules::TlsMode::Lax)));
}
}
}