mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 18:07:47 +08:00
feat(security): harden gateway boundaries and usage policies
Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change. Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
@@ -1,28 +1,298 @@
|
||||
use thiserror::Error;
|
||||
use aether_contracts::redact_url_for_debug;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
const OAUTH_ERROR_BODY_EXCERPT_CHARS: usize = 500;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum OAuthError {
|
||||
#[error("unsupported oauth provider: {0}")]
|
||||
UnsupportedProvider(String),
|
||||
#[error("invalid oauth request: {0}")]
|
||||
InvalidRequest(String),
|
||||
#[error("oauth state is invalid or expired")]
|
||||
InvalidState,
|
||||
#[error("oauth provider returned HTTP {status_code}: {body_excerpt}")]
|
||||
// `body_excerpt` remains available to trusted callers for status
|
||||
// classification, but must not be rendered by the generic Error/Debug
|
||||
// paths: OAuth servers sometimes echo access tokens, authorization codes,
|
||||
// assertions, or client credentials in an error response.
|
||||
HttpStatus {
|
||||
status_code: u16,
|
||||
body_excerpt: String,
|
||||
},
|
||||
#[error("oauth provider returned invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
#[error("oauth transport failed: {0}")]
|
||||
Transport(String),
|
||||
#[error("oauth storage failed: {0}")]
|
||||
Storage(String),
|
||||
#[error("oauth encryption failed")]
|
||||
EncryptionUnavailable,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for OAuthError {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::UnsupportedProvider(detail) => write!(
|
||||
formatter,
|
||||
"unsupported oauth provider: {}",
|
||||
redact_oauth_error_detail(detail)
|
||||
),
|
||||
Self::InvalidRequest(detail) => write!(
|
||||
formatter,
|
||||
"invalid oauth request: {}",
|
||||
redact_oauth_error_detail(detail)
|
||||
),
|
||||
Self::InvalidState => formatter.write_str("oauth state is invalid or expired"),
|
||||
Self::HttpStatus { status_code, .. } => {
|
||||
write!(formatter, "oauth provider returned HTTP {status_code}")
|
||||
}
|
||||
Self::InvalidResponse(detail) => write!(
|
||||
formatter,
|
||||
"oauth provider returned invalid response: {}",
|
||||
redact_oauth_error_detail(detail)
|
||||
),
|
||||
Self::Transport(detail) => write!(
|
||||
formatter,
|
||||
"oauth transport failed: {}",
|
||||
redact_oauth_error_detail(detail)
|
||||
),
|
||||
Self::Storage(detail) => write!(
|
||||
formatter,
|
||||
"oauth storage failed: {}",
|
||||
redact_oauth_error_detail(detail)
|
||||
),
|
||||
Self::EncryptionUnavailable => formatter.write_str("oauth encryption failed"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for OAuthError {}
|
||||
|
||||
impl std::fmt::Debug for OAuthError {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::UnsupportedProvider(_) => formatter
|
||||
.debug_tuple("UnsupportedProvider")
|
||||
.field(&"[REDACTED]")
|
||||
.finish(),
|
||||
Self::InvalidRequest(_) => formatter
|
||||
.debug_tuple("InvalidRequest")
|
||||
.field(&"[REDACTED]")
|
||||
.finish(),
|
||||
Self::InvalidState => formatter.write_str("InvalidState"),
|
||||
Self::HttpStatus { status_code, .. } => formatter
|
||||
.debug_struct("HttpStatus")
|
||||
.field("status_code", status_code)
|
||||
.field("body_excerpt", &"[REDACTED]")
|
||||
.finish(),
|
||||
Self::InvalidResponse(_) => formatter
|
||||
.debug_tuple("InvalidResponse")
|
||||
.field(&"[REDACTED]")
|
||||
.finish(),
|
||||
Self::Transport(_) => formatter
|
||||
.debug_tuple("Transport")
|
||||
.field(&"[REDACTED]")
|
||||
.finish(),
|
||||
Self::Storage(_) => formatter
|
||||
.debug_tuple("Storage")
|
||||
.field(&"[REDACTED]")
|
||||
.finish(),
|
||||
Self::EncryptionUnavailable => formatter.write_str("EncryptionUnavailable"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Builds an error excerpt that is safe to persist or render in diagnostics.
|
||||
///
|
||||
/// Structured responses retain non-sensitive provider error codes/messages so
|
||||
/// callers can still classify `invalid_grant` and similar failures. Secret
|
||||
/// fields and secret-shaped values are removed before the size bound is
|
||||
/// applied. Unstructured bodies containing credential markers are replaced as
|
||||
/// a whole because their token boundaries cannot be determined reliably.
|
||||
pub fn redacted_oauth_error_body_excerpt(body: &str) -> String {
|
||||
let body = body.trim();
|
||||
if body.is_empty() {
|
||||
return "-".to_string();
|
||||
}
|
||||
|
||||
if let Ok(mut value) = serde_json::from_str::<Value>(body) {
|
||||
redact_oauth_error_json(&mut value);
|
||||
return value
|
||||
.to_string()
|
||||
.chars()
|
||||
.take(OAUTH_ERROR_BODY_EXCERPT_CHARS)
|
||||
.collect();
|
||||
}
|
||||
|
||||
if unstructured_body_may_contain_secret(body) {
|
||||
"[REDACTED upstream OAuth error body]".to_string()
|
||||
} else {
|
||||
body.chars().take(OAUTH_ERROR_BODY_EXCERPT_CHARS).collect()
|
||||
}
|
||||
}
|
||||
|
||||
fn redact_oauth_error_json(value: &mut Value) {
|
||||
match value {
|
||||
Value::Object(object) => {
|
||||
for (key, value) in object {
|
||||
if oauth_error_key_is_sensitive(key) {
|
||||
*value = json!("[REDACTED]");
|
||||
} else {
|
||||
redact_oauth_error_json(value);
|
||||
}
|
||||
}
|
||||
}
|
||||
Value::Array(items) => {
|
||||
for item in items {
|
||||
redact_oauth_error_json(item);
|
||||
}
|
||||
}
|
||||
Value::String(text) => {
|
||||
if oauth_error_value_is_safe_classification_code(text) {
|
||||
// Keep a small allowlist of non-secret provider error codes for classification.
|
||||
} else if oauth_error_value_looks_secret(text)
|
||||
|| unstructured_body_may_contain_secret(text)
|
||||
{
|
||||
*text = "[REDACTED]".to_string();
|
||||
} else {
|
||||
*text = redact_urls_in_text(text);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn oauth_error_key_is_sensitive(key: &str) -> bool {
|
||||
let normalized = key
|
||||
.chars()
|
||||
.filter(|ch| ch.is_ascii_alphanumeric())
|
||||
.collect::<String>()
|
||||
.to_ascii_lowercase();
|
||||
normalized.contains("token")
|
||||
|| normalized.contains("apikey")
|
||||
|| normalized.contains("password")
|
||||
|| normalized.contains("authorization")
|
||||
|| normalized.contains("secret")
|
||||
|| normalized.contains("clientsecret")
|
||||
|| normalized.contains("privatekey")
|
||||
|| normalized.contains("assertion")
|
||||
|| normalized.contains("credential")
|
||||
|| normalized.contains("cookie")
|
||||
|| normalized.contains("pkce")
|
||||
|| normalized.contains("verifier")
|
||||
|| normalized == "sessionkey"
|
||||
}
|
||||
|
||||
fn oauth_error_value_looks_secret(value: &str) -> bool {
|
||||
let value = value.trim();
|
||||
value.starts_with("Bearer ")
|
||||
|| value.starts_with("bearer ")
|
||||
|| value.starts_with("sk-")
|
||||
|| value.starts_with("sess-")
|
||||
|| value.starts_with("devin-session-token$")
|
||||
|| value.starts_with("ott$")
|
||||
|| value.starts_with("auth1_")
|
||||
|| (value.len() > 80
|
||||
&& value.split('.').count() == 3
|
||||
&& value.split('.').all(|segment| {
|
||||
!segment.is_empty()
|
||||
&& segment.bytes().all(|byte| {
|
||||
byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'=')
|
||||
})
|
||||
}))
|
||||
}
|
||||
|
||||
fn oauth_error_value_is_safe_classification_code(value: &str) -> bool {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
"refresh_token_reused" | "refresh_token_expired" | "invalid_refresh_token"
|
||||
)
|
||||
}
|
||||
|
||||
/// Redact dynamic OAuth error details before they reach `Display` consumers.
|
||||
/// Error details frequently originate in HTTP clients and may contain a full
|
||||
/// request URL or an upstream response body.
|
||||
fn redact_oauth_error_detail(detail: &str) -> String {
|
||||
let excerpt = redacted_oauth_error_body_excerpt(detail);
|
||||
redact_urls_in_text(&excerpt)
|
||||
.chars()
|
||||
.take(OAUTH_ERROR_BODY_EXCERPT_CHARS)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn redact_urls_in_text(text: &str) -> String {
|
||||
const URL_SCHEMES: [&str; 4] = ["http://", "https://", "ws://", "wss://"];
|
||||
let mut output = String::with_capacity(text.len());
|
||||
let mut cursor = 0;
|
||||
|
||||
while cursor < text.len() {
|
||||
let Some((relative_start, scheme)) = URL_SCHEMES
|
||||
.iter()
|
||||
.filter_map(|scheme| text[cursor..].find(scheme).map(|start| (start, *scheme)))
|
||||
.min_by_key(|(start, _)| *start)
|
||||
else {
|
||||
output.push_str(&text[cursor..]);
|
||||
break;
|
||||
};
|
||||
|
||||
let start = cursor + relative_start;
|
||||
output.push_str(&text[cursor..start]);
|
||||
let token_end = text[start..]
|
||||
.find(char::is_whitespace)
|
||||
.map(|offset| start + offset)
|
||||
.unwrap_or(text.len());
|
||||
let token = &text[start..token_end];
|
||||
let (url_token, suffix) = trim_url_suffix(token);
|
||||
if url_token.starts_with(scheme) {
|
||||
if url::Url::parse(url_token).is_ok() {
|
||||
output.push_str(&redact_url_for_debug(url_token));
|
||||
} else {
|
||||
output.push_str("[REDACTED URL]");
|
||||
}
|
||||
output.push_str(suffix);
|
||||
} else {
|
||||
output.push_str(token);
|
||||
}
|
||||
cursor = token_end;
|
||||
}
|
||||
|
||||
output
|
||||
}
|
||||
|
||||
fn trim_url_suffix(token: &str) -> (&str, &str) {
|
||||
let mut end = token.len();
|
||||
while end > 0 {
|
||||
let Some(ch) = token[..end].chars().next_back() else {
|
||||
break;
|
||||
};
|
||||
if matches!(
|
||||
ch,
|
||||
'.' | ',' | ';' | ':' | '!' | '?' | ')' | ']' | '}' | '\'' | '"'
|
||||
) {
|
||||
end -= ch.len_utf8();
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
(&token[..end], &token[end..])
|
||||
}
|
||||
|
||||
fn unstructured_body_may_contain_secret(value: &str) -> bool {
|
||||
let normalized = value.to_ascii_lowercase();
|
||||
[
|
||||
"access_token",
|
||||
"refresh_token",
|
||||
"id_token",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"authorization:",
|
||||
"authorization=",
|
||||
"client_secret",
|
||||
"secret=",
|
||||
"secret:",
|
||||
"password=",
|
||||
"password:",
|
||||
"assertion=",
|
||||
"session_token",
|
||||
"sessiontoken",
|
||||
]
|
||||
.iter()
|
||||
.any(|marker| normalized.contains(marker))
|
||||
|| oauth_error_value_looks_secret(value)
|
||||
}
|
||||
|
||||
impl OAuthError {
|
||||
pub fn invalid_request(detail: impl Into<String>) -> Self {
|
||||
Self::InvalidRequest(detail.into())
|
||||
@@ -36,3 +306,94 @@ impl OAuthError {
|
||||
Self::Transport(detail.into())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{redacted_oauth_error_body_excerpt, OAuthError};
|
||||
|
||||
#[test]
|
||||
fn oauth_error_body_excerpt_preserves_classification_and_redacts_secrets() {
|
||||
let excerpt = redacted_oauth_error_body_excerpt(
|
||||
r#"{
|
||||
"error": {
|
||||
"code": "invalid_grant",
|
||||
"message": "refresh token expired",
|
||||
"refresh_token": "refresh-body-canary",
|
||||
"nested": {"clientSecret": "client-secret-canary"}
|
||||
},
|
||||
"accessToken": "access-token-canary"
|
||||
}"#,
|
||||
);
|
||||
|
||||
assert!(excerpt.contains("invalid_grant"));
|
||||
assert!(excerpt.contains("refresh token expired"));
|
||||
assert!(!excerpt.contains("refresh-body-canary"));
|
||||
assert!(!excerpt.contains("client-secret-canary"));
|
||||
assert!(!excerpt.contains("access-token-canary"));
|
||||
assert!(excerpt.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_error_debug_and_display_do_not_render_upstream_body() {
|
||||
let error = OAuthError::HttpStatus {
|
||||
status_code: 401,
|
||||
body_excerpt: "authorization=Bearer oauth-error-canary".to_string(),
|
||||
};
|
||||
|
||||
let debug = format!("{error:?}");
|
||||
let display = error.to_string();
|
||||
assert!(!debug.contains("oauth-error-canary"));
|
||||
assert!(!display.contains("oauth-error-canary"));
|
||||
assert!(debug.contains("[REDACTED]"));
|
||||
assert_eq!(display, "oauth provider returned HTTP 401");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unstructured_oauth_error_with_secret_markers_is_replaced() {
|
||||
let excerpt = redacted_oauth_error_body_excerpt(
|
||||
"invalid request: refresh_token=plain-text-refresh-canary",
|
||||
);
|
||||
assert_eq!(excerpt, "[REDACTED upstream OAuth error body]");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_error_body_excerpt_preserves_long_refresh_rotation_message() {
|
||||
let body = r#"{"error":{"message":"Your refresh token has already been used to generate a new access token. Please try signing in again.","type":"invalid_request_error","param":null,"code":"refresh_token_reused"}}"#;
|
||||
let excerpt = redacted_oauth_error_body_excerpt(body);
|
||||
assert!(excerpt.contains("already been used to generate a new access token"));
|
||||
assert!(excerpt.contains("refresh_token_reused"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_oauth_error_display_redacts_credentials_and_url_queries() {
|
||||
let detail =
|
||||
"apiKey=sk-secret token=secret-token https://user:[email protected]?q=url-secret";
|
||||
for error in [
|
||||
OAuthError::UnsupportedProvider(detail.to_string()),
|
||||
OAuthError::InvalidRequest(detail.to_string()),
|
||||
OAuthError::InvalidResponse(detail.to_string()),
|
||||
OAuthError::Transport(detail.to_string()),
|
||||
OAuthError::Storage(detail.to_string()),
|
||||
] {
|
||||
let display = error.to_string();
|
||||
for secret in ["sk-secret", "secret-token", "user", "pass", "url-secret"] {
|
||||
assert!(
|
||||
!display.contains(secret),
|
||||
"display leaked {secret}: {display}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_oauth_error_display_redacts_standalone_url() {
|
||||
let error = OAuthError::invalid_response(
|
||||
"upstream request failed at https://user:[email protected]/path?code=secret",
|
||||
);
|
||||
let display = error.to_string();
|
||||
assert!(!display.contains("user"));
|
||||
assert!(!display.contains("pass"));
|
||||
assert!(!display.contains("code=secret"));
|
||||
assert!(display.contains("https://example.test/path"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct OAuthProviderMetadata {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
@@ -13,7 +13,27 @@ pub struct OAuthProviderMetadata {
|
||||
pub use_pkce: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
impl std::fmt::Debug for OAuthProviderMetadata {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthProviderMetadata")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("authorize_url", &"[REDACTED]")
|
||||
.field("token_url", &"[REDACTED]")
|
||||
.field("client_id", &self.client_id)
|
||||
.field(
|
||||
"client_secret",
|
||||
&self.client_secret.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("scopes", &self.scopes)
|
||||
.field("redirect_uri", &self.redirect_uri)
|
||||
.field("use_pkce", &self.use_pkce)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct OAuthAuthorizeRequest {
|
||||
pub state: String,
|
||||
pub code_challenge: Option<String>,
|
||||
@@ -21,7 +41,22 @@ pub struct OAuthAuthorizeRequest {
|
||||
pub login_hint: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
|
||||
impl std::fmt::Debug for OAuthAuthorizeRequest {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthAuthorizeRequest")
|
||||
.field("state", &"[REDACTED]")
|
||||
.field(
|
||||
"code_challenge",
|
||||
&self.code_challenge.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("prompt", &self.prompt)
|
||||
.field("has_login_hint", &self.login_hint.is_some())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, Serialize)]
|
||||
pub struct OAuthAuthorizeResponse {
|
||||
pub authorize_url: String,
|
||||
pub state: String,
|
||||
@@ -29,14 +64,39 @@ pub struct OAuthAuthorizeResponse {
|
||||
pub code_challenge: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
impl std::fmt::Debug for OAuthAuthorizeResponse {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthAuthorizeResponse")
|
||||
.field("authorize_url", &"[REDACTED]")
|
||||
.field("state", &"[REDACTED]")
|
||||
.field(
|
||||
"code_challenge",
|
||||
&self.code_challenge.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct OAuthCallback {
|
||||
pub code: String,
|
||||
pub state: String,
|
||||
pub scope: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)]
|
||||
impl std::fmt::Debug for OAuthCallback {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthCallback")
|
||||
.field("code", &"[REDACTED]")
|
||||
.field("state", &"[REDACTED]")
|
||||
.field("scope", &self.scope)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, Deserialize, Serialize)]
|
||||
pub struct OAuthDeviceAuthorization {
|
||||
pub device_code: String,
|
||||
pub user_code: String,
|
||||
@@ -45,3 +105,86 @@ pub struct OAuthDeviceAuthorization {
|
||||
pub expires_in: u64,
|
||||
pub interval: u64,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for OAuthDeviceAuthorization {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthDeviceAuthorization")
|
||||
.field("device_code", &"[REDACTED]")
|
||||
.field("user_code", &"[REDACTED]")
|
||||
.field("verification_uri", &self.verification_uri)
|
||||
.field("verification_uri_complete", &"[REDACTED]")
|
||||
.field("expires_in", &self.expires_in)
|
||||
.field("interval", &self.interval)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
OAuthAuthorizeRequest, OAuthAuthorizeResponse, OAuthCallback, OAuthDeviceAuthorization,
|
||||
OAuthProviderMetadata,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn oauth_flow_debug_output_redacts_capabilities_and_secrets() {
|
||||
let metadata = OAuthProviderMetadata {
|
||||
provider_type: "test".to_string(),
|
||||
display_name: "Test".to_string(),
|
||||
authorize_url: "https://idp.example/authorize".to_string(),
|
||||
token_url: "https://idp.example/token".to_string(),
|
||||
client_id: "public-client".to_string(),
|
||||
client_secret: Some("client-secret-canary".to_string()),
|
||||
scopes: vec!["openid".to_string()],
|
||||
redirect_uri: "https://gateway.example/callback".to_string(),
|
||||
use_pkce: true,
|
||||
};
|
||||
let request = OAuthAuthorizeRequest {
|
||||
state: "state-canary".to_string(),
|
||||
code_challenge: Some("challenge-canary".to_string()),
|
||||
prompt: None,
|
||||
login_hint: Some("login-hint-canary".to_string()),
|
||||
};
|
||||
let response = OAuthAuthorizeResponse {
|
||||
authorize_url:
|
||||
"https://idp.example/authorize?state=state-url-canary&code_challenge=challenge"
|
||||
.to_string(),
|
||||
state: "response-state-canary".to_string(),
|
||||
code_challenge: Some("response-challenge-canary".to_string()),
|
||||
};
|
||||
let callback = OAuthCallback {
|
||||
code: "authorization-code-canary".to_string(),
|
||||
state: "callback-state-canary".to_string(),
|
||||
scope: None,
|
||||
};
|
||||
let device = OAuthDeviceAuthorization {
|
||||
device_code: "device-code-canary".to_string(),
|
||||
user_code: "user-code-canary".to_string(),
|
||||
verification_uri: "https://idp.example/device".to_string(),
|
||||
verification_uri_complete: "https://idp.example/device?code=complete-canary"
|
||||
.to_string(),
|
||||
expires_in: 600,
|
||||
interval: 5,
|
||||
};
|
||||
|
||||
let debug = format!("{metadata:?} {request:?} {response:?} {callback:?} {device:?}");
|
||||
for secret in [
|
||||
"client-secret-canary",
|
||||
"state-canary",
|
||||
"challenge-canary",
|
||||
"login-hint-canary",
|
||||
"state-url-canary",
|
||||
"response-state-canary",
|
||||
"response-challenge-canary",
|
||||
"authorization-code-canary",
|
||||
"callback-state-canary",
|
||||
"device-code-canary",
|
||||
"user-code-canary",
|
||||
"complete-canary",
|
||||
] {
|
||||
assert!(!debug.contains(secret), "debug leaked {secret}");
|
||||
}
|
||||
assert!(debug.contains("[REDACTED]"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ mod pkce;
|
||||
mod registry;
|
||||
mod token;
|
||||
|
||||
pub use error::OAuthError;
|
||||
pub use error::{redacted_oauth_error_body_excerpt, OAuthError};
|
||||
pub use flow::{
|
||||
OAuthAuthorizeRequest, OAuthAuthorizeResponse, OAuthCallback, OAuthDeviceAuthorization,
|
||||
OAuthProviderMetadata,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use serde_json::Value;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct OAuthTokenSet {
|
||||
pub access_token: String,
|
||||
pub refresh_token: Option<String>,
|
||||
@@ -11,6 +11,26 @@ pub struct OAuthTokenSet {
|
||||
pub raw_payload: Option<Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for OAuthTokenSet {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthTokenSet")
|
||||
.field("access_token", &"<redacted>")
|
||||
.field(
|
||||
"refresh_token",
|
||||
&self.refresh_token.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field("token_type", &self.token_type)
|
||||
.field("scope", &self.scope)
|
||||
.field("expires_at_unix_secs", &self.expires_at_unix_secs)
|
||||
.field(
|
||||
"raw_payload",
|
||||
&self.raw_payload.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl OAuthTokenSet {
|
||||
pub fn from_token_payload(payload: Value) -> Option<Self> {
|
||||
let access_token = non_empty_string(payload.get("access_token"))
|
||||
@@ -109,4 +129,20 @@ mod tests {
|
||||
assert!(token.expires_at_unix_secs.is_some());
|
||||
assert_eq!(token.bearer_header_value(), "Bearer access");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn debug_output_redacts_tokens_and_raw_payload() {
|
||||
let token = OAuthTokenSet::from_token_payload(json!({
|
||||
"access_token": "access-secret-sentinel",
|
||||
"refresh_token": "refresh-secret-sentinel",
|
||||
"provider_secret": "raw-secret-sentinel"
|
||||
}))
|
||||
.expect("token should parse");
|
||||
|
||||
let debug = format!("{token:?}");
|
||||
assert!(!debug.contains("access-secret-sentinel"));
|
||||
assert!(!debug.contains("refresh-secret-sentinel"));
|
||||
assert!(!debug.contains("raw-secret-sentinel"));
|
||||
assert!(debug.contains("<redacted>"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct IdentityOAuthProviderConfig {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
@@ -20,14 +20,60 @@ pub struct IdentityOAuthProviderConfig {
|
||||
pub extra_config: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for IdentityOAuthProviderConfig {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("IdentityOAuthProviderConfig")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("authorization_url", &"[REDACTED]")
|
||||
.field("token_url", &"[REDACTED]")
|
||||
.field(
|
||||
"userinfo_url",
|
||||
&self.userinfo_url.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("client_id", &self.client_id)
|
||||
.field(
|
||||
"client_secret",
|
||||
&self.client_secret.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("scopes", &self.scopes)
|
||||
.field("redirect_uri", &self.redirect_uri)
|
||||
.field("frontend_callback_url", &self.frontend_callback_url)
|
||||
.field(
|
||||
"attribute_mapping",
|
||||
&self.attribute_mapping.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"extra_config",
|
||||
&self.extra_config.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct IdentityOAuthStartContext {
|
||||
pub state: String,
|
||||
pub code_challenge: Option<String>,
|
||||
pub network: OAuthNetworkContext,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for IdentityOAuthStartContext {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("IdentityOAuthStartContext")
|
||||
.field("state", &"[REDACTED]")
|
||||
.field(
|
||||
"code_challenge",
|
||||
&self.code_challenge.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("network", &self.network)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct IdentityOAuthExchangeContext {
|
||||
pub code: String,
|
||||
pub state: String,
|
||||
@@ -35,27 +81,75 @@ pub struct IdentityOAuthExchangeContext {
|
||||
pub network: OAuthNetworkContext,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for IdentityOAuthExchangeContext {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("IdentityOAuthExchangeContext")
|
||||
.field("code", &"[REDACTED]")
|
||||
.field("state", &"[REDACTED]")
|
||||
.field(
|
||||
"pkce_verifier",
|
||||
&self.pkce_verifier.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("network", &self.network)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct ExternalIdentity {
|
||||
pub provider_type: String,
|
||||
pub subject: String,
|
||||
pub email: Option<String>,
|
||||
pub email_verified: bool,
|
||||
pub username: Option<String>,
|
||||
pub display_name: Option<String>,
|
||||
pub avatar_url: Option<String>,
|
||||
pub raw: Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for ExternalIdentity {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ExternalIdentity")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("subject", &self.subject)
|
||||
.field("email", &self.email)
|
||||
.field("email_verified", &self.email_verified)
|
||||
.field("username", &self.username)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("avatar_url", &self.avatar_url)
|
||||
.field("raw", &"[REDACTED]")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct IdentityClaims {
|
||||
pub provider_type: String,
|
||||
pub subject: String,
|
||||
pub email: Option<String>,
|
||||
pub email_verified: bool,
|
||||
pub username: Option<String>,
|
||||
pub display_name: Option<String>,
|
||||
pub raw: Value,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for IdentityClaims {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("IdentityClaims")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("subject", &self.subject)
|
||||
.field("email", &self.email)
|
||||
.field("email_verified", &self.email_verified)
|
||||
.field("username", &self.username)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("raw", &"[REDACTED]")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait IdentityOAuthProvider: Send + Sync {
|
||||
fn provider_type(&self) -> &'static str;
|
||||
@@ -103,16 +197,34 @@ pub(crate) fn mapped_string(
|
||||
find_string(raw, mapped_key)
|
||||
}
|
||||
|
||||
pub(crate) fn mapped_bool(raw: &Value, mapping: Option<&Value>, logical_key: &str) -> Option<bool> {
|
||||
let mapped_key = match mapping
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|object| object.get(logical_key))
|
||||
{
|
||||
Some(value) => value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?,
|
||||
None => logical_key,
|
||||
};
|
||||
find_value(raw, mapped_key).and_then(Value::as_bool)
|
||||
}
|
||||
|
||||
pub(crate) fn find_string(raw: &Value, key: &str) -> Option<String> {
|
||||
find_value(raw, key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn find_value<'a>(raw: &'a Value, key: &str) -> Option<&'a Value> {
|
||||
let mut current = raw;
|
||||
for segment in key.split('.') {
|
||||
current = current.get(segment)?;
|
||||
}
|
||||
current
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
Some(current)
|
||||
}
|
||||
|
||||
pub(crate) fn form_headers() -> BTreeMap<String, String> {
|
||||
@@ -124,3 +236,128 @@ pub(crate) fn form_headers() -> BTreeMap<String, String> {
|
||||
("accept".to_string(), "application/json".to_string()),
|
||||
])
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
mapped_bool, ExternalIdentity, IdentityClaims, IdentityOAuthExchangeContext,
|
||||
IdentityOAuthProviderConfig, IdentityOAuthStartContext,
|
||||
};
|
||||
use crate::network::OAuthNetworkContext;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn identity_oauth_debug_output_redacts_credentials_and_raw_claims() {
|
||||
let config = IdentityOAuthProviderConfig {
|
||||
provider_type: "custom".to_string(),
|
||||
display_name: "Custom".to_string(),
|
||||
authorization_url: "https://idp.example/authorize".to_string(),
|
||||
token_url: "https://idp.example/token".to_string(),
|
||||
userinfo_url: Some("https://idp.example/userinfo".to_string()),
|
||||
client_id: "public-client".to_string(),
|
||||
client_secret: Some("identity-client-secret-canary".to_string()),
|
||||
scopes: vec!["openid".to_string()],
|
||||
redirect_uri: "https://gateway.example/callback".to_string(),
|
||||
frontend_callback_url: "https://app.example/callback".to_string(),
|
||||
attribute_mapping: None,
|
||||
extra_config: Some(json!({"secret": "identity-extra-canary"})),
|
||||
};
|
||||
let start = IdentityOAuthStartContext {
|
||||
state: "identity-state-canary".to_string(),
|
||||
code_challenge: Some("identity-challenge-canary".to_string()),
|
||||
network: OAuthNetworkContext::direct_identity(),
|
||||
};
|
||||
let exchange = IdentityOAuthExchangeContext {
|
||||
code: "identity-code-canary".to_string(),
|
||||
state: "identity-exchange-state-canary".to_string(),
|
||||
pkce_verifier: Some("identity-verifier-canary".to_string()),
|
||||
network: OAuthNetworkContext::direct_identity(),
|
||||
};
|
||||
let external = ExternalIdentity {
|
||||
provider_type: "custom".to_string(),
|
||||
subject: "subject".to_string(),
|
||||
email: None,
|
||||
email_verified: false,
|
||||
username: None,
|
||||
display_name: None,
|
||||
avatar_url: None,
|
||||
raw: json!({"access_token": "identity-raw-canary"}),
|
||||
};
|
||||
let claims = IdentityClaims {
|
||||
provider_type: "custom".to_string(),
|
||||
subject: "subject".to_string(),
|
||||
email: None,
|
||||
email_verified: false,
|
||||
username: None,
|
||||
display_name: None,
|
||||
raw: json!({"id_token": "identity-claims-canary"}),
|
||||
};
|
||||
|
||||
let debug = format!("{config:?} {start:?} {exchange:?} {external:?} {claims:?}");
|
||||
for secret in [
|
||||
"identity-client-secret-canary",
|
||||
"identity-extra-canary",
|
||||
"identity-state-canary",
|
||||
"identity-challenge-canary",
|
||||
"identity-code-canary",
|
||||
"identity-exchange-state-canary",
|
||||
"identity-verifier-canary",
|
||||
"identity-raw-canary",
|
||||
"identity-claims-canary",
|
||||
] {
|
||||
assert!(!debug.contains(secret), "debug leaked {secret}");
|
||||
}
|
||||
assert!(debug.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mapped_bool_accepts_only_an_explicit_json_boolean() {
|
||||
let raw = json!({
|
||||
"email_verified": true,
|
||||
"profile": {
|
||||
"verified": false,
|
||||
"string_verified": "true",
|
||||
"numeric_verified": 1
|
||||
}
|
||||
});
|
||||
|
||||
assert_eq!(mapped_bool(&raw, None, "email_verified"), Some(true));
|
||||
assert_eq!(
|
||||
mapped_bool(
|
||||
&raw,
|
||||
Some(&json!({"email_verified": "profile.verified"})),
|
||||
"email_verified"
|
||||
),
|
||||
Some(false)
|
||||
);
|
||||
assert_eq!(
|
||||
mapped_bool(
|
||||
&raw,
|
||||
Some(&json!({"email_verified": "profile.string_verified"})),
|
||||
"email_verified"
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
mapped_bool(
|
||||
&raw,
|
||||
Some(&json!({"email_verified": "profile.numeric_verified"})),
|
||||
"email_verified"
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(mapped_bool(&json!({}), None, "email_verified"), None);
|
||||
assert_eq!(
|
||||
mapped_bool(
|
||||
&raw,
|
||||
Some(&json!({"email_verified": true})),
|
||||
"email_verified"
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
mapped_bool(&raw, Some(&json!({"email_verified": ""})), "email_verified"),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use super::super::adapter::{find_string, form_headers, mapped_string};
|
||||
use crate::core::{OAuthAuthorizeResponse, OAuthError, OAuthTokenSet};
|
||||
use super::super::adapter::{find_string, form_headers, mapped_bool, mapped_string};
|
||||
use crate::core::{
|
||||
redacted_oauth_error_body_excerpt, OAuthAuthorizeResponse, OAuthError, OAuthTokenSet,
|
||||
};
|
||||
use crate::identity::{
|
||||
ExternalIdentity, IdentityClaims, IdentityOAuthExchangeContext, IdentityOAuthProvider,
|
||||
IdentityOAuthProviderConfig, IdentityOAuthStartContext,
|
||||
@@ -8,6 +10,16 @@ use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthNetworkContext};
|
||||
use async_trait::async_trait;
|
||||
use url::form_urlencoded;
|
||||
|
||||
const SERVER_MANAGED_AUTHORIZE_PARAMS: &[&str] = &[
|
||||
"response_type",
|
||||
"client_id",
|
||||
"redirect_uri",
|
||||
"state",
|
||||
"scope",
|
||||
"code_challenge",
|
||||
"code_challenge_method",
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct CustomOidcIdentityOAuthProvider;
|
||||
|
||||
@@ -24,6 +36,15 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider {
|
||||
) -> Result<OAuthAuthorizeResponse, OAuthError> {
|
||||
let mut url = url::Url::parse(&config.authorization_url)
|
||||
.map_err(|_| OAuthError::invalid_request("authorization_url must be absolute"))?;
|
||||
if url.query_pairs().any(|(name, _)| {
|
||||
SERVER_MANAGED_AUTHORIZE_PARAMS
|
||||
.iter()
|
||||
.any(|reserved| name.eq_ignore_ascii_case(reserved))
|
||||
}) {
|
||||
return Err(OAuthError::invalid_request(
|
||||
"authorization_url must not predefine server-managed OAuth parameters",
|
||||
));
|
||||
}
|
||||
{
|
||||
let mut query = url.query_pairs_mut();
|
||||
query.append_pair("response_type", "code");
|
||||
@@ -87,7 +108,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider {
|
||||
if !(200..300).contains(&response.status_code) {
|
||||
return Err(OAuthError::HttpStatus {
|
||||
status_code: response.status_code,
|
||||
body_excerpt: response.body_text.chars().take(500).collect(),
|
||||
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
|
||||
});
|
||||
}
|
||||
let payload = response
|
||||
@@ -130,7 +151,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider {
|
||||
if !(200..300).contains(&response.status_code) {
|
||||
return Err(OAuthError::HttpStatus {
|
||||
status_code: response.status_code,
|
||||
body_excerpt: response.body_text.chars().take(500).collect(),
|
||||
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
|
||||
});
|
||||
}
|
||||
let raw = response
|
||||
@@ -144,6 +165,8 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider {
|
||||
provider_type: config.provider_type.clone(),
|
||||
subject,
|
||||
email: mapped_string(&raw, config.attribute_mapping.as_ref(), "email"),
|
||||
email_verified: mapped_bool(&raw, config.attribute_mapping.as_ref(), "email_verified")
|
||||
.unwrap_or(false),
|
||||
username: mapped_string(&raw, config.attribute_mapping.as_ref(), "username"),
|
||||
display_name: mapped_string(&raw, config.attribute_mapping.as_ref(), "display_name")
|
||||
.or_else(|| find_string(&raw, "name")),
|
||||
@@ -160,6 +183,7 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider {
|
||||
Ok(IdentityClaims {
|
||||
provider_type: config.provider_type.clone(),
|
||||
subject: identity.subject,
|
||||
email_verified: identity.email.is_some() && identity.email_verified,
|
||||
email: identity.email,
|
||||
username: identity.username,
|
||||
display_name: identity.display_name,
|
||||
@@ -167,3 +191,123 @@ impl IdentityOAuthProvider for CustomOidcIdentityOAuthProvider {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::CustomOidcIdentityOAuthProvider;
|
||||
use crate::identity::{
|
||||
ExternalIdentity, IdentityOAuthProvider, IdentityOAuthProviderConfig,
|
||||
IdentityOAuthStartContext,
|
||||
};
|
||||
use crate::network::OAuthNetworkContext;
|
||||
use serde_json::json;
|
||||
|
||||
fn config() -> IdentityOAuthProviderConfig {
|
||||
IdentityOAuthProviderConfig {
|
||||
provider_type: "custom_oidc_work".to_string(),
|
||||
display_name: "Work OIDC".to_string(),
|
||||
authorization_url: "https://idp.example.test/authorize".to_string(),
|
||||
token_url: "https://idp.example.test/token".to_string(),
|
||||
userinfo_url: Some("https://idp.example.test/userinfo".to_string()),
|
||||
client_id: "client".to_string(),
|
||||
client_secret: None,
|
||||
scopes: vec!["openid".to_string(), "email".to_string()],
|
||||
redirect_uri: "https://gateway.example.test/callback".to_string(),
|
||||
frontend_callback_url: "https://app.example.test/callback".to_string(),
|
||||
attribute_mapping: None,
|
||||
extra_config: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn start_context() -> IdentityOAuthStartContext {
|
||||
IdentityOAuthStartContext {
|
||||
state: "server-state".to_string(),
|
||||
code_challenge: Some("server-challenge".to_string()),
|
||||
network: OAuthNetworkContext::direct_identity(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn custom_oidc_authorize_url_rejects_predefined_server_managed_parameters() {
|
||||
for name in [
|
||||
"response_type",
|
||||
"client_id",
|
||||
"redirect_uri",
|
||||
"state",
|
||||
"scope",
|
||||
"code_challenge",
|
||||
"code_challenge_method",
|
||||
] {
|
||||
let mut config = config();
|
||||
config.authorization_url =
|
||||
format!("https://idp.example.test/authorize?{name}=attacker");
|
||||
|
||||
assert!(CustomOidcIdentityOAuthProvider
|
||||
.build_authorize_url(&config, &start_context())
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn custom_oidc_authorize_url_preserves_non_oauth_tenant_parameters() {
|
||||
let mut config = config();
|
||||
config.authorization_url =
|
||||
"https://idp.example.test/authorize?tenant=workforce".to_string();
|
||||
|
||||
let response = CustomOidcIdentityOAuthProvider
|
||||
.build_authorize_url(&config, &start_context())
|
||||
.expect("tenant parameter should be preserved");
|
||||
let parsed = url::Url::parse(&response.authorize_url).expect("authorize URL");
|
||||
let params = parsed.query_pairs().collect::<Vec<_>>();
|
||||
|
||||
assert!(params
|
||||
.iter()
|
||||
.any(|(name, value)| name == "tenant" && value == "workforce"));
|
||||
assert_eq!(params.iter().filter(|(name, _)| name == "state").count(), 1);
|
||||
assert!(params
|
||||
.iter()
|
||||
.any(|(name, value)| name == "state" && value == "server-state"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn custom_oidc_propagates_an_explicit_verified_email_claim() {
|
||||
let claims = CustomOidcIdentityOAuthProvider
|
||||
.map_identity(
|
||||
&config(),
|
||||
ExternalIdentity {
|
||||
provider_type: "custom_oidc_work".to_string(),
|
||||
subject: "user-1".to_string(),
|
||||
email: Some("[email protected]".to_string()),
|
||||
email_verified: true,
|
||||
username: Some("user".to_string()),
|
||||
display_name: None,
|
||||
avatar_url: None,
|
||||
raw: json!({"email_verified": true}),
|
||||
},
|
||||
)
|
||||
.expect("identity should map");
|
||||
|
||||
assert!(claims.email_verified);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn custom_oidc_cannot_verify_a_missing_email() {
|
||||
let claims = CustomOidcIdentityOAuthProvider
|
||||
.map_identity(
|
||||
&config(),
|
||||
ExternalIdentity {
|
||||
provider_type: "custom_oidc_work".to_string(),
|
||||
subject: "user-1".to_string(),
|
||||
email: None,
|
||||
email_verified: true,
|
||||
username: Some("user".to_string()),
|
||||
display_name: None,
|
||||
avatar_url: None,
|
||||
raw: json!({"email_verified": true}),
|
||||
},
|
||||
)
|
||||
.expect("identity should map");
|
||||
|
||||
assert!(!claims.email_verified);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,6 +52,53 @@ impl IdentityOAuthProvider for LinuxDoIdentityOAuthProvider {
|
||||
config: &IdentityOAuthProviderConfig,
|
||||
identity: ExternalIdentity,
|
||||
) -> Result<IdentityClaims, OAuthError> {
|
||||
self.inner.map_identity(config, identity)
|
||||
let mut claims = self.inner.map_identity(config, identity)?;
|
||||
// Linux.do's OAuth user endpoint does not provide an OIDC-level guarantee
|
||||
// for the email verification claim, so it must not verify a local email.
|
||||
claims.email_verified = false;
|
||||
Ok(claims)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::LinuxDoIdentityOAuthProvider;
|
||||
use crate::identity::{ExternalIdentity, IdentityOAuthProvider, IdentityOAuthProviderConfig};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn linuxdo_does_not_promote_an_unverified_provider_assertion() {
|
||||
let provider = LinuxDoIdentityOAuthProvider::default();
|
||||
let config = IdentityOAuthProviderConfig {
|
||||
provider_type: "linuxdo".to_string(),
|
||||
display_name: "Linux.do".to_string(),
|
||||
authorization_url: "https://connect.linux.do/oauth2/authorize".to_string(),
|
||||
token_url: "https://connect.linux.do/oauth2/token".to_string(),
|
||||
userinfo_url: Some("https://connect.linux.do/api/user".to_string()),
|
||||
client_id: "client".to_string(),
|
||||
client_secret: None,
|
||||
scopes: vec![],
|
||||
redirect_uri: "https://gateway.example.test/callback".to_string(),
|
||||
frontend_callback_url: "https://app.example.test/callback".to_string(),
|
||||
attribute_mapping: None,
|
||||
extra_config: None,
|
||||
};
|
||||
let claims = provider
|
||||
.map_identity(
|
||||
&config,
|
||||
ExternalIdentity {
|
||||
provider_type: "linuxdo".to_string(),
|
||||
subject: "user-1".to_string(),
|
||||
email: Some("[email protected]".to_string()),
|
||||
email_verified: true,
|
||||
username: Some("user".to_string()),
|
||||
display_name: None,
|
||||
avatar_url: None,
|
||||
raw: json!({"email_verified": true}),
|
||||
},
|
||||
)
|
||||
.expect("identity should map");
|
||||
|
||||
assert!(!claims.email_verified);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,8 +5,8 @@ pub mod provider;
|
||||
|
||||
pub use core::{
|
||||
current_unix_secs, generate_oauth_nonce, generate_pkce_verifier, parse_oauth_callback_params,
|
||||
pkce_s256, OAuthAdapterRegistry, OAuthAuthorizeRequest, OAuthAuthorizeResponse, OAuthCallback,
|
||||
OAuthError, OAuthProviderMetadata, OAuthTokenSet,
|
||||
pkce_s256, redacted_oauth_error_body_excerpt, OAuthAdapterRegistry, OAuthAuthorizeRequest,
|
||||
OAuthAuthorizeResponse, OAuthCallback, OAuthError, OAuthProviderMetadata, OAuthTokenSet,
|
||||
};
|
||||
pub use network::{
|
||||
NetworkRequirement, OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse,
|
||||
|
||||
@@ -38,7 +38,7 @@ impl OAuthTimeouts {
|
||||
};
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct OAuthNetworkContext {
|
||||
pub policy: OAuthNetworkPolicy,
|
||||
pub requirement: NetworkRequirement,
|
||||
@@ -46,6 +46,18 @@ pub struct OAuthNetworkContext {
|
||||
pub timeouts: OAuthTimeouts,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for OAuthNetworkContext {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthNetworkContext")
|
||||
.field("policy", &self.policy)
|
||||
.field("requirement", &self.requirement)
|
||||
.field("has_proxy", &self.proxy.is_some())
|
||||
.field("timeouts", &self.timeouts)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl OAuthNetworkContext {
|
||||
pub fn direct_identity() -> Self {
|
||||
Self {
|
||||
@@ -70,3 +82,24 @@ impl OAuthNetworkContext {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::OAuthNetworkContext;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
|
||||
#[test]
|
||||
fn network_context_debug_output_does_not_expose_proxy_credentials() {
|
||||
let context = OAuthNetworkContext::provider_operation(Some(ProxySnapshot {
|
||||
url: Some("http://proxy-user:[email protected]:8080".to_string()),
|
||||
extra: Some(serde_json::json!({"authorization": "proxy-extra-canary"})),
|
||||
..ProxySnapshot::default()
|
||||
}));
|
||||
|
||||
let debug = format!("{context:?}");
|
||||
assert!(!debug.contains("proxy-user"));
|
||||
assert!(!debug.contains("proxy-password"));
|
||||
assert!(!debug.contains("proxy-extra-canary"));
|
||||
assert!(debug.contains("has_proxy: true"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
use crate::core::OAuthError;
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use aether_contracts::{redact_url_for_debug, ResolvedTransportProfile};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::OAuthNetworkContext;
|
||||
|
||||
const OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024;
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct OAuthHttpRequest {
|
||||
pub request_id: String,
|
||||
@@ -25,7 +27,7 @@ impl std::fmt::Debug for OAuthHttpRequest {
|
||||
.debug_struct("OAuthHttpRequest")
|
||||
.field("request_id", &self.request_id)
|
||||
.field("method", &self.method)
|
||||
.field("url", &self.url)
|
||||
.field("url", &redact_url_for_debug(&self.url))
|
||||
.field("header_names", &self.headers.keys().collect::<Vec<_>>())
|
||||
.field("content_type", &self.content_type)
|
||||
.field("has_json_body", &self.json_body.is_some())
|
||||
@@ -43,13 +45,24 @@ impl std::fmt::Debug for OAuthHttpRequest {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct OAuthHttpResponse {
|
||||
pub status_code: u16,
|
||||
pub body_text: String,
|
||||
pub json_body: Option<Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for OAuthHttpResponse {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("OAuthHttpResponse")
|
||||
.field("status_code", &self.status_code)
|
||||
.field("body_bytes_len", &self.body_text.len())
|
||||
.field("has_json_body", &self.json_body.is_some())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait OAuthHttpExecutor: Send + Sync {
|
||||
async fn execute(&self, request: OAuthHttpRequest) -> Result<OAuthHttpResponse, OAuthError>;
|
||||
@@ -81,15 +94,29 @@ impl OAuthHttpExecutor for ReqwestOAuthHttpExecutor {
|
||||
builder = builder.body(body_bytes.clone());
|
||||
}
|
||||
|
||||
let response = builder
|
||||
let mut response = builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| OAuthError::transport(err.to_string()))?;
|
||||
let status_code = response.status().as_u16();
|
||||
let body_text = response
|
||||
.text()
|
||||
if response
|
||||
.content_length()
|
||||
.is_some_and(|length| length > OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES as u64)
|
||||
{
|
||||
return Err(oauth_http_response_too_large());
|
||||
}
|
||||
let mut body = Vec::new();
|
||||
while let Some(chunk) = response
|
||||
.chunk()
|
||||
.await
|
||||
.map_err(|err| OAuthError::transport(err.to_string()))?;
|
||||
.map_err(|err| OAuthError::transport(err.to_string()))?
|
||||
{
|
||||
if chunk.len() > OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) {
|
||||
return Err(oauth_http_response_too_large());
|
||||
}
|
||||
body.extend_from_slice(&chunk);
|
||||
}
|
||||
let body_text = String::from_utf8_lossy(&body).to_string();
|
||||
let json_body = serde_json::from_str::<Value>(&body_text).ok();
|
||||
Ok(OAuthHttpResponse {
|
||||
status_code,
|
||||
@@ -98,3 +125,51 @@ impl OAuthHttpExecutor for ReqwestOAuthHttpExecutor {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn oauth_http_response_too_large() -> OAuthError {
|
||||
OAuthError::transport(format!(
|
||||
"OAuth response body exceeds {OAUTH_HTTP_RESPONSE_BODY_LIMIT_BYTES} bytes"
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{OAuthHttpRequest, OAuthHttpResponse};
|
||||
use crate::network::OAuthNetworkContext;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn response_debug_output_does_not_expose_token_payloads() {
|
||||
let response = OAuthHttpResponse {
|
||||
status_code: 200,
|
||||
body_text: "{\"access_token\":\"response-body-canary\"}".to_string(),
|
||||
json_body: Some(serde_json::json!({"refresh_token": "response-json-canary"})),
|
||||
};
|
||||
|
||||
let debug = format!("{response:?}");
|
||||
assert!(!debug.contains("response-body-canary"));
|
||||
assert!(!debug.contains("response-json-canary"));
|
||||
assert!(debug.contains("body_bytes_len"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_debug_redacts_url_credentials_and_query() {
|
||||
let request = OAuthHttpRequest {
|
||||
request_id: "request-1".into(),
|
||||
method: reqwest::Method::GET,
|
||||
url: "https://user:[email protected]/oauth?client_secret=url-secret".into(),
|
||||
headers: BTreeMap::from([("authorization".into(), "Bearer header-secret".into())]),
|
||||
content_type: None,
|
||||
json_body: None,
|
||||
body_bytes: None,
|
||||
network: OAuthNetworkContext::direct_identity(),
|
||||
transport_profile: None,
|
||||
};
|
||||
let debug = format!("{request:?}");
|
||||
assert!(!debug.contains("user"));
|
||||
assert!(!debug.contains("pass"));
|
||||
assert!(!debug.contains("url-secret"));
|
||||
assert!(!debug.contains("header-secret"));
|
||||
assert!(debug.contains("https://example.test/oauth"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -40,7 +40,7 @@ impl std::fmt::Debug for ProviderOAuthCookieAuthorizationInput {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct ProviderOAuthTransportContext {
|
||||
pub provider_id: String,
|
||||
pub provider_type: String,
|
||||
@@ -55,13 +55,57 @@ pub struct ProviderOAuthTransportContext {
|
||||
pub network: OAuthNetworkContext,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for ProviderOAuthTransportContext {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderOAuthTransportContext")
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("endpoint_id", &self.endpoint_id)
|
||||
.field("key_id", &self.key_id)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field(
|
||||
"decrypted_api_key",
|
||||
&self.decrypted_api_key.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field(
|
||||
"decrypted_auth_config",
|
||||
&self.decrypted_auth_config.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field(
|
||||
"provider_config",
|
||||
&self.provider_config.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field(
|
||||
"endpoint_config",
|
||||
&self.endpoint_config.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field(
|
||||
"key_config",
|
||||
&self.key_config.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field("network", &self.network)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct ProviderOAuthTokenSet {
|
||||
pub token_set: OAuthTokenSet,
|
||||
pub auth_config: Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for ProviderOAuthTokenSet {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderOAuthTokenSet")
|
||||
.field("token_set", &self.token_set)
|
||||
.field("auth_config", &"<redacted>")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct ProviderOAuthAccount {
|
||||
pub provider_type: String,
|
||||
pub access_token: String,
|
||||
@@ -70,6 +114,19 @@ pub struct ProviderOAuthAccount {
|
||||
pub identity: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderOAuthAccount {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderOAuthAccount")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("access_token", &"<redacted>")
|
||||
.field("auth_config", &"<redacted>")
|
||||
.field("expires_at_unix_secs", &self.expires_at_unix_secs)
|
||||
.field("identity_keys", &self.identity.keys().collect::<Vec<_>>())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl ProviderOAuthAccount {
|
||||
pub fn request_bearer_auth(&self) -> ProviderOAuthRequestAuth {
|
||||
ProviderOAuthRequestAuth::Header {
|
||||
@@ -79,7 +136,7 @@ impl ProviderOAuthAccount {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub enum ProviderOAuthRequestAuth {
|
||||
Header {
|
||||
name: String,
|
||||
@@ -93,7 +150,28 @@ pub enum ProviderOAuthRequestAuth {
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for ProviderOAuthRequestAuth {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Header { name, .. } => formatter
|
||||
.debug_struct("Header")
|
||||
.field("name", name)
|
||||
.field("value", &"<redacted>")
|
||||
.finish(),
|
||||
Self::Kiro {
|
||||
name, machine_id, ..
|
||||
} => formatter
|
||||
.debug_struct("Kiro")
|
||||
.field("name", name)
|
||||
.field("value", &"<redacted>")
|
||||
.field("auth_config", &"<redacted>")
|
||||
.field("machine_id", machine_id)
|
||||
.finish(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct ProviderOAuthImportInput {
|
||||
pub provider_type: String,
|
||||
pub name: Option<String>,
|
||||
@@ -102,7 +180,26 @@ pub struct ProviderOAuthImportInput {
|
||||
pub network: OAuthNetworkContext,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for ProviderOAuthImportInput {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderOAuthImportInput")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("name", &self.name)
|
||||
.field(
|
||||
"refresh_token",
|
||||
&self.refresh_token.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field(
|
||||
"raw_credentials",
|
||||
&self.raw_credentials.as_ref().map(|_| "<redacted>"),
|
||||
)
|
||||
.field("network", &self.network)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct ProviderOAuthAccountState {
|
||||
pub is_valid: bool,
|
||||
pub email: Option<String>,
|
||||
@@ -110,3 +207,95 @@ pub struct ProviderOAuthAccountState {
|
||||
pub invalid_reason: Option<String>,
|
||||
pub raw: Option<Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderOAuthAccountState {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderOAuthAccountState")
|
||||
.field("is_valid", &self.is_valid)
|
||||
.field("email", &self.email)
|
||||
.field("quota", &self.quota)
|
||||
.field("invalid_reason", &self.invalid_reason)
|
||||
.field("raw", &self.raw.as_ref().map(|_| "[REDACTED]"))
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
ProviderOAuthAccount, ProviderOAuthAccountState, ProviderOAuthImportInput,
|
||||
ProviderOAuthTokenSet, ProviderOAuthTransportContext,
|
||||
};
|
||||
use crate::core::OAuthTokenSet;
|
||||
use crate::network::OAuthNetworkContext;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn debug_output_redacts_provider_oauth_credentials() {
|
||||
let context = ProviderOAuthTransportContext {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_type: "generic".to_string(),
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
auth_type: None,
|
||||
decrypted_api_key: Some("api-key-secret-sentinel".to_string()),
|
||||
decrypted_auth_config: Some("auth-config-secret-sentinel".to_string()),
|
||||
provider_config: Some(json!({"client_secret": "provider-secret-sentinel"})),
|
||||
endpoint_config: None,
|
||||
key_config: None,
|
||||
network: OAuthNetworkContext::direct_identity(),
|
||||
};
|
||||
let token_set = ProviderOAuthTokenSet {
|
||||
token_set: OAuthTokenSet {
|
||||
access_token: "access-secret-sentinel".to_string(),
|
||||
refresh_token: Some("refresh-secret-sentinel".to_string()),
|
||||
token_type: None,
|
||||
scope: None,
|
||||
expires_at_unix_secs: None,
|
||||
raw_payload: None,
|
||||
},
|
||||
auth_config: json!({"password": "password-secret-sentinel"}),
|
||||
};
|
||||
let account = ProviderOAuthAccount {
|
||||
provider_type: "generic".to_string(),
|
||||
access_token: "account-secret-sentinel".to_string(),
|
||||
auth_config: json!({"client_secret": "account-config-secret-sentinel"}),
|
||||
expires_at_unix_secs: None,
|
||||
identity: BTreeMap::new(),
|
||||
};
|
||||
let import = ProviderOAuthImportInput {
|
||||
provider_type: "generic".to_string(),
|
||||
name: None,
|
||||
refresh_token: Some("import-refresh-secret-sentinel".to_string()),
|
||||
raw_credentials: Some(json!({"api_key": "import-raw-secret-sentinel"})),
|
||||
network: OAuthNetworkContext::direct_identity(),
|
||||
};
|
||||
let state = ProviderOAuthAccountState {
|
||||
is_valid: false,
|
||||
email: None,
|
||||
quota: None,
|
||||
invalid_reason: None,
|
||||
raw: Some(json!({"access_token": "probe-raw-secret-sentinel"})),
|
||||
};
|
||||
|
||||
let debug = format!("{context:?} {token_set:?} {account:?} {import:?} {state:?}");
|
||||
for secret in [
|
||||
"api-key-secret-sentinel",
|
||||
"auth-config-secret-sentinel",
|
||||
"provider-secret-sentinel",
|
||||
"access-secret-sentinel",
|
||||
"refresh-secret-sentinel",
|
||||
"password-secret-sentinel",
|
||||
"account-secret-sentinel",
|
||||
"account-config-secret-sentinel",
|
||||
"import-refresh-secret-sentinel",
|
||||
"import-raw-secret-sentinel",
|
||||
"probe-raw-secret-sentinel",
|
||||
] {
|
||||
assert!(!debug.contains(secret), "debug leaked {secret}");
|
||||
}
|
||||
assert!(debug.contains("<redacted>"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use crate::core::{current_unix_secs, OAuthAuthorizeResponse, OAuthError, OAuthTokenSet};
|
||||
use crate::core::{
|
||||
current_unix_secs, redacted_oauth_error_body_excerpt, OAuthAuthorizeResponse, OAuthError,
|
||||
OAuthTokenSet,
|
||||
};
|
||||
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest};
|
||||
use crate::provider::ProviderOAuthAdapter;
|
||||
use crate::provider::{
|
||||
@@ -17,6 +20,10 @@ use super::claude_code::{
|
||||
CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_REDIRECT_URI, CLAUDE_CODE_TOKEN_URL,
|
||||
};
|
||||
|
||||
pub const GEMINI_CLI_OAUTH_CLIENT_ID_ENV: &str = "AETHER_GEMINI_CLI_OAUTH_CLIENT_ID";
|
||||
pub const GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV: &str = "AETHER_GEMINI_CLI_OAUTH_CLIENT_SECRET";
|
||||
pub const ANTIGRAVITY_OAUTH_CLIENT_ID_ENV: &str = "AETHER_ANTIGRAVITY_OAUTH_CLIENT_ID";
|
||||
pub const ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV: &str = "AETHER_ANTIGRAVITY_OAUTH_CLIENT_SECRET";
|
||||
const CODEX_IDENTITY_FINGERPRINT_FIELD: &str = "codex_identity_fingerprint";
|
||||
const CODEX_IDENTITY_FINGERPRINT_VERSION: &str = "codex-persisted-fingerprint:v1";
|
||||
|
||||
@@ -53,7 +60,8 @@ pub struct GenericProviderOAuthTemplate {
|
||||
pub authorize_url: &'static str,
|
||||
pub token_url: &'static str,
|
||||
pub client_id: &'static str,
|
||||
pub client_secret: &'static str,
|
||||
pub client_id_env: Option<&'static str>,
|
||||
pub client_secret_env: Option<&'static str>,
|
||||
pub scopes: &'static [&'static str],
|
||||
pub redirect_uri: &'static str,
|
||||
pub use_pkce: bool,
|
||||
@@ -68,7 +76,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
authorize_url: CLAUDE_CODE_AUTHORIZE_URL,
|
||||
token_url: CLAUDE_CODE_TOKEN_URL,
|
||||
client_id: CLAUDE_CODE_CLIENT_ID,
|
||||
client_secret: "",
|
||||
client_id_env: None,
|
||||
client_secret_env: None,
|
||||
scopes: CLAUDE_CODE_OAUTH_SCOPES,
|
||||
redirect_uri: CLAUDE_CODE_REDIRECT_URI,
|
||||
use_pkce: true,
|
||||
@@ -81,7 +90,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
authorize_url: "https://auth.openai.com/oauth/authorize",
|
||||
token_url: "https://auth.openai.com/oauth/token",
|
||||
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
|
||||
client_secret: "",
|
||||
client_id_env: None,
|
||||
client_secret_env: None,
|
||||
scopes: &["openid", "email", "profile", "offline_access"],
|
||||
redirect_uri: "http://localhost:1455/auth/callback",
|
||||
use_pkce: true,
|
||||
@@ -94,7 +104,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
authorize_url: "https://auth.openai.com/oauth/authorize",
|
||||
token_url: "https://auth.openai.com/oauth/token",
|
||||
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
|
||||
client_secret: "",
|
||||
client_id_env: None,
|
||||
client_secret_env: None,
|
||||
scopes: &["openid", "email", "profile", "offline_access"],
|
||||
redirect_uri: "http://localhost:1455/auth/callback",
|
||||
use_pkce: true,
|
||||
@@ -107,7 +118,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url: "https://oauth2.googleapis.com/token",
|
||||
client_id: "681255809395-oo8ft2oprdrnp9e3aqf6av3hmdib135j.apps.googleusercontent.com",
|
||||
client_secret: "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",
|
||||
client_id_env: Some(GEMINI_CLI_OAUTH_CLIENT_ID_ENV),
|
||||
client_secret_env: Some(GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV),
|
||||
scopes: &[
|
||||
"https://www.googleapis.com/auth/cloud-platform",
|
||||
"https://www.googleapis.com/auth/userinfo.email",
|
||||
@@ -124,7 +136,8 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url: "https://oauth2.googleapis.com/token",
|
||||
client_id: "1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com",
|
||||
client_secret: "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",
|
||||
client_id_env: Some(ANTIGRAVITY_OAUTH_CLIENT_ID_ENV),
|
||||
client_secret_env: Some(ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV),
|
||||
scopes: &[
|
||||
"https://www.googleapis.com/auth/cloud-platform",
|
||||
"https://www.googleapis.com/auth/userinfo.email",
|
||||
@@ -139,10 +152,24 @@ pub const GENERIC_PROVIDER_OAUTH_TEMPLATES: &[GenericProviderOAuthTemplate] = &[
|
||||
},
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
#[derive(Clone)]
|
||||
pub struct GenericProviderOAuthAdapter {
|
||||
template: GenericProviderOAuthTemplate,
|
||||
token_url_override: Option<String>,
|
||||
client_id_override: Option<String>,
|
||||
client_secret_override: Option<String>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for GenericProviderOAuthAdapter {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GenericProviderOAuthAdapter")
|
||||
.field("provider_type", &self.template.provider_type)
|
||||
.field("has_token_url_override", &self.token_url_override.is_some())
|
||||
.field("client_id_env", &self.template.client_id_env)
|
||||
.field("client_secret_env", &self.template.client_secret_env)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl GenericProviderOAuthAdapter {
|
||||
@@ -150,6 +177,8 @@ impl GenericProviderOAuthAdapter {
|
||||
Self {
|
||||
template,
|
||||
token_url_override: None,
|
||||
client_id_override: None,
|
||||
client_secret_override: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -166,12 +195,52 @@ impl GenericProviderOAuthAdapter {
|
||||
self.with_token_url_override(token_url)
|
||||
}
|
||||
|
||||
#[doc(hidden)]
|
||||
pub fn with_oauth_credentials_for_tests(
|
||||
mut self,
|
||||
client_id: impl Into<String>,
|
||||
client_secret: impl Into<String>,
|
||||
) -> Self {
|
||||
self.client_id_override = Some(client_id.into());
|
||||
self.client_secret_override = Some(client_secret.into());
|
||||
self
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn without_oauth_client_secret_for_tests(mut self) -> Self {
|
||||
self.client_secret_override = Some(String::new());
|
||||
self
|
||||
}
|
||||
|
||||
fn token_url(&self) -> String {
|
||||
self.token_url_override
|
||||
.clone()
|
||||
.unwrap_or_else(|| self.template.token_url.to_string())
|
||||
}
|
||||
|
||||
fn client_id(&self) -> String {
|
||||
if let Some(value) = self.client_id_override.clone().and_then(non_empty_owned) {
|
||||
return value;
|
||||
}
|
||||
|
||||
self.template
|
||||
.client_id_env
|
||||
.and_then(non_empty_environment_value)
|
||||
.unwrap_or_else(|| self.template.client_id.to_string())
|
||||
}
|
||||
|
||||
fn client_secret(&self) -> Result<Option<String>, OAuthError> {
|
||||
let Some(env_name) = self.template.client_secret_env else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if let Some(value) = self.client_secret_override.clone() {
|
||||
return required_client_secret(env_name, non_empty_owned(value)).map(Some);
|
||||
}
|
||||
|
||||
required_client_secret(env_name, non_empty_environment_value(env_name)).map(Some)
|
||||
}
|
||||
|
||||
async fn exchange_grant(
|
||||
&self,
|
||||
executor: &dyn OAuthHttpExecutor,
|
||||
@@ -181,6 +250,8 @@ impl GenericProviderOAuthAdapter {
|
||||
state: Option<&str>,
|
||||
pkce_verifier: Option<&str>,
|
||||
) -> Result<ProviderOAuthTokenSet, OAuthError> {
|
||||
let client_id = self.client_id();
|
||||
let client_secret = self.client_secret()?;
|
||||
let scope = (!self.template.scopes.is_empty()).then(|| self.template.scopes.join(" "));
|
||||
let request_id = match grant_type {
|
||||
"authorization_code" => "provider-oauth:exchange-code".to_string(),
|
||||
@@ -196,11 +267,14 @@ impl GenericProviderOAuthAdapter {
|
||||
"grant_type".to_string(),
|
||||
Value::String(grant_type.to_string()),
|
||||
),
|
||||
(
|
||||
"client_id".to_string(),
|
||||
Value::String(self.template.client_id.to_string()),
|
||||
),
|
||||
("client_id".to_string(), Value::String(client_id.clone())),
|
||||
]);
|
||||
if let Some(client_secret) = client_secret.as_ref() {
|
||||
body.insert(
|
||||
"client_secret".to_string(),
|
||||
Value::String(client_secret.clone()),
|
||||
);
|
||||
}
|
||||
if grant_type == "authorization_code" {
|
||||
body.insert(
|
||||
"code".to_string(),
|
||||
@@ -247,7 +321,7 @@ impl GenericProviderOAuthAdapter {
|
||||
let form_body = {
|
||||
let mut form = form_urlencoded::Serializer::new(String::new());
|
||||
form.append_pair("grant_type", grant_type);
|
||||
form.append_pair("client_id", self.template.client_id);
|
||||
form.append_pair("client_id", &client_id);
|
||||
if grant_type == "authorization_code" {
|
||||
form.append_pair("redirect_uri", self.template.redirect_uri);
|
||||
form.append_pair("code", code_or_refresh_token);
|
||||
@@ -262,8 +336,8 @@ impl GenericProviderOAuthAdapter {
|
||||
form.append_pair("scope", scope);
|
||||
}
|
||||
}
|
||||
if !self.template.client_secret.trim().is_empty() {
|
||||
form.append_pair("client_secret", self.template.client_secret);
|
||||
if let Some(client_secret) = client_secret.as_deref() {
|
||||
form.append_pair("client_secret", client_secret);
|
||||
}
|
||||
form.finish().into_bytes()
|
||||
};
|
||||
@@ -340,12 +414,14 @@ impl ProviderOAuthAdapter for GenericProviderOAuthAdapter {
|
||||
state: &str,
|
||||
code_challenge: Option<&str>,
|
||||
) -> Result<OAuthAuthorizeResponse, OAuthError> {
|
||||
self.client_secret()?;
|
||||
let client_id = self.client_id();
|
||||
let mut url = url::Url::parse(self.template.authorize_url)
|
||||
.map_err(|_| OAuthError::invalid_request("authorize_url must be absolute"))?;
|
||||
{
|
||||
let mut query = url.query_pairs_mut();
|
||||
query.append_pair("response_type", "code");
|
||||
query.append_pair("client_id", self.template.client_id);
|
||||
query.append_pair("client_id", &client_id);
|
||||
query.append_pair("redirect_uri", self.template.redirect_uri);
|
||||
query.append_pair("state", state);
|
||||
if !self.template.scopes.is_empty() {
|
||||
@@ -472,6 +548,26 @@ impl ProviderOAuthAdapter for GenericProviderOAuthAdapter {
|
||||
}
|
||||
}
|
||||
|
||||
fn non_empty_environment_value(name: &str) -> Option<String> {
|
||||
std::env::var(name).ok().and_then(non_empty_owned)
|
||||
}
|
||||
|
||||
fn non_empty_owned(value: String) -> Option<String> {
|
||||
let value = value.trim();
|
||||
(!value.is_empty()).then(|| value.to_string())
|
||||
}
|
||||
|
||||
fn required_client_secret(
|
||||
env_name: &'static str,
|
||||
configured: Option<String>,
|
||||
) -> Result<String, OAuthError> {
|
||||
configured.ok_or_else(|| {
|
||||
OAuthError::invalid_request(format!(
|
||||
"{env_name} must be configured for this OAuth provider"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn template_for_provider_type(provider_type: &str) -> Option<GenericProviderOAuthTemplate> {
|
||||
let normalized = provider_type.trim();
|
||||
GENERIC_PROVIDER_OAUTH_TEMPLATES
|
||||
@@ -506,12 +602,7 @@ fn json_headers(provider_type: &str) -> BTreeMap<String, String> {
|
||||
}
|
||||
|
||||
fn truncate_body(body: &str) -> String {
|
||||
let body = body.trim();
|
||||
if body.is_empty() {
|
||||
"-".to_string()
|
||||
} else {
|
||||
body.chars().take(500).collect()
|
||||
}
|
||||
redacted_oauth_error_body_excerpt(body)
|
||||
}
|
||||
|
||||
fn secret_fingerprint(value: &str) -> String {
|
||||
@@ -790,8 +881,21 @@ fn value_to_string(value: &Value) -> Option<String> {
|
||||
|
||||
fn decode_jwt_claims(token: &str) -> Option<serde_json::Map<String, Value>> {
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024;
|
||||
|
||||
let payload = token.split('.').nth(1)?;
|
||||
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
|
||||
.saturating_add(2)
|
||||
.checked_div(3)
|
||||
.unwrap_or(usize::MAX)
|
||||
.saturating_mul(4);
|
||||
if payload.len() > max_encoded_len {
|
||||
return None;
|
||||
}
|
||||
let bytes = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?;
|
||||
if bytes.len() > MAX_UNVERIFIED_JWT_CLAIMS_BYTES {
|
||||
return None;
|
||||
}
|
||||
serde_json::from_slice::<Value>(&bytes)
|
||||
.ok()?
|
||||
.as_object()
|
||||
@@ -801,9 +905,12 @@ fn decode_jwt_claims(token: &str) -> Option<serde_json::Map<String, Value>> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
derive_codex_identity_fingerprint, enrich_generic_identity, template_for_provider_type,
|
||||
GenericProviderOAuthAdapter, CODEX_IDENTITY_FINGERPRINT_FIELD,
|
||||
decode_jwt_claims, derive_codex_identity_fingerprint, enrich_generic_identity,
|
||||
template_for_provider_type, GenericProviderOAuthAdapter,
|
||||
ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV, CODEX_IDENTITY_FINGERPRINT_FIELD,
|
||||
GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV,
|
||||
};
|
||||
use crate::core::OAuthError;
|
||||
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse};
|
||||
use crate::provider::ProviderOAuthAdapter;
|
||||
use crate::provider::{ProviderOAuthAccount, ProviderOAuthTransportContext};
|
||||
@@ -835,6 +942,36 @@ mod tests {
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn google_oauth_templates_reference_external_client_secrets() {
|
||||
let gemini = template_for_provider_type("gemini_cli").expect("gemini template");
|
||||
let antigravity = template_for_provider_type("antigravity").expect("antigravity template");
|
||||
|
||||
assert_eq!(
|
||||
gemini.client_secret_env,
|
||||
Some(GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV)
|
||||
);
|
||||
assert_eq!(
|
||||
antigravity.client_secret_env,
|
||||
Some(ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generic_adapter_debug_redacts_oauth_credentials() {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli")
|
||||
.expect("gemini adapter")
|
||||
.with_token_url_override("https://token.example.test/private-path")
|
||||
.with_oauth_credentials_for_tests("private-client-id", "private-client-secret");
|
||||
|
||||
let debug = format!("{adapter:?}");
|
||||
|
||||
assert!(!debug.contains("private-client-id"));
|
||||
assert!(!debug.contains("private-client-secret"));
|
||||
assert!(!debug.contains("private-path"));
|
||||
assert!(debug.contains(GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_identity_extracts_fedramp_workspace_claim() {
|
||||
let claims = json!({
|
||||
@@ -852,6 +989,19 @@ mod tests {
|
||||
assert_eq!(auth_config.get("is_fedramp"), Some(&json!(true)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generic_identity_rejects_oversized_jwt_claims_before_decode() {
|
||||
const MAX_UNVERIFIED_JWT_CLAIMS_BYTES: usize = 64 * 1024;
|
||||
let max_encoded_len = MAX_UNVERIFIED_JWT_CLAIMS_BYTES
|
||||
.saturating_add(2)
|
||||
.checked_div(3)
|
||||
.unwrap()
|
||||
.saturating_mul(4);
|
||||
let token = format!("header.{}.signature", "A".repeat(max_encoded_len + 1));
|
||||
|
||||
assert_eq!(decode_jwt_claims(&token), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_persisted_fingerprint_is_member_scoped_and_token_independent() {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type("codex")
|
||||
@@ -937,6 +1087,108 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn transport_context(provider_type: &str) -> ProviderOAuthTransportContext {
|
||||
ProviderOAuthTransportContext {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_type: provider_type.to_string(),
|
||||
endpoint_id: None,
|
||||
key_id: Some("key-1".to_string()),
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: None,
|
||||
endpoint_config: None,
|
||||
key_config: None,
|
||||
network: crate::network::OAuthNetworkContext::provider_operation(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn oauth_account(provider_type: &str) -> ProviderOAuthAccount {
|
||||
ProviderOAuthAccount {
|
||||
provider_type: provider_type.to_string(),
|
||||
access_token: "old-access-token".to_string(),
|
||||
auth_config: json!({
|
||||
"provider_type": provider_type,
|
||||
"refresh_token": "old-refresh-token",
|
||||
"updated_at": 1
|
||||
}),
|
||||
expires_at_unix_secs: Some(1),
|
||||
identity: BTreeMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn google_oauth_fails_closed_before_network_without_client_secret() {
|
||||
let seen_request = Arc::new(Mutex::new(None));
|
||||
let executor = StaticExecutor {
|
||||
seen_request: Arc::clone(&seen_request),
|
||||
response_payload: json!({}),
|
||||
};
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli")
|
||||
.expect("gemini adapter")
|
||||
.without_oauth_client_secret_for_tests();
|
||||
let ctx = transport_context("gemini_cli");
|
||||
|
||||
let authorize_error = adapter
|
||||
.build_authorize_url(&ctx, "state", None)
|
||||
.expect_err("authorization must fail without the configured secret");
|
||||
assert!(matches!(authorize_error, OAuthError::InvalidRequest(_)));
|
||||
|
||||
let refresh_error = adapter
|
||||
.refresh(&executor, &ctx, &oauth_account("gemini_cli"))
|
||||
.await
|
||||
.expect_err("refresh must fail without the configured secret");
|
||||
assert!(matches!(refresh_error, OAuthError::InvalidRequest(_)));
|
||||
assert!(
|
||||
seen_request.lock().expect("mutex should lock").is_none(),
|
||||
"credential validation must happen before the HTTP executor runs"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn google_oauth_injected_credentials_are_sent_in_token_form() {
|
||||
let seen_request = Arc::new(Mutex::new(None));
|
||||
let executor = StaticExecutor {
|
||||
seen_request: Arc::clone(&seen_request),
|
||||
response_payload: json!({
|
||||
"access_token": "new-access-token",
|
||||
"expires_in": 3600,
|
||||
}),
|
||||
};
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type("gemini_cli")
|
||||
.expect("gemini adapter")
|
||||
.with_oauth_credentials_for_tests("test-client-id", "test-client-secret");
|
||||
let ctx = transport_context("gemini_cli");
|
||||
|
||||
adapter
|
||||
.refresh(&executor, &ctx, &oauth_account("gemini_cli"))
|
||||
.await
|
||||
.expect("refresh should succeed");
|
||||
|
||||
let seen = seen_request
|
||||
.lock()
|
||||
.expect("mutex should lock")
|
||||
.clone()
|
||||
.expect("request should be captured");
|
||||
let form = String::from_utf8(seen.body_bytes.expect("form body should exist"))
|
||||
.expect("form body should be utf8");
|
||||
let fields = url::form_urlencoded::parse(form.as_bytes())
|
||||
.into_owned()
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
assert_eq!(
|
||||
fields.get("client_id").map(String::as_str),
|
||||
Some("test-client-id")
|
||||
);
|
||||
assert_eq!(
|
||||
fields.get("client_secret").map(String::as_str),
|
||||
Some("test-client-secret")
|
||||
);
|
||||
assert_eq!(
|
||||
fields.get("refresh_token").map(String::as_str),
|
||||
Some("old-refresh-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refresh_preserves_existing_metadata_when_refresh_token_is_not_rotated() {
|
||||
let seen_request = Arc::new(Mutex::new(None));
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use crate::core::{current_unix_secs, OAuthError};
|
||||
use crate::core::{current_unix_secs, redacted_oauth_error_body_excerpt, OAuthError};
|
||||
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest};
|
||||
use crate::provider::ProviderOAuthAdapter;
|
||||
use crate::provider::{
|
||||
@@ -18,7 +18,7 @@ pub const DEFAULT_SYSTEM_VERSION: &str = "other#unknown";
|
||||
const IDC_AMZ_USER_AGENT: &str =
|
||||
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct KiroAuthConfig {
|
||||
pub auth_method: Option<String>,
|
||||
pub refresh_token: Option<String>,
|
||||
@@ -36,6 +36,75 @@ pub struct KiroAuthConfig {
|
||||
pub access_token: Option<String>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for KiroAuthConfig {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("KiroAuthConfig")
|
||||
.field("auth_method", &self.auth_method)
|
||||
.field(
|
||||
"refresh_token",
|
||||
&self.refresh_token.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("expires_at", &self.expires_at)
|
||||
.field(
|
||||
"profile_arn",
|
||||
&self.profile_arn.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("region", &self.region)
|
||||
.field("auth_region", &self.auth_region)
|
||||
.field("api_region", &self.api_region)
|
||||
.field("client_id", &self.client_id.as_ref().map(|_| "[REDACTED]"))
|
||||
.field(
|
||||
"client_secret",
|
||||
&self.client_secret.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"machine_id",
|
||||
&self.machine_id.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("kiro_version", &self.kiro_version)
|
||||
.field("system_version", &self.system_version)
|
||||
.field("node_version", &self.node_version)
|
||||
.field(
|
||||
"access_token",
|
||||
&self.access_token.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns whether a value is safe to interpolate as one DNS label in a Kiro
|
||||
/// service hostname.
|
||||
///
|
||||
/// Region values are persisted in the encrypted auth configuration and can
|
||||
/// also come from an upstream OAuth response. They must never be treated as
|
||||
/// URL syntax: accepting `/`, `.`, `?`, `#`, `@`, or control characters here
|
||||
/// would let a crafted region redirect a token-bearing request to another
|
||||
/// origin. AWS region names are DNS labels, so an ASCII alphanumeric/hyphen
|
||||
/// allow-list is both stricter and forward-compatible with future regions.
|
||||
pub fn is_valid_kiro_region(value: &str) -> bool {
|
||||
let value = value.trim();
|
||||
!value.is_empty()
|
||||
&& value.len() <= 63
|
||||
&& !value.starts_with('-')
|
||||
&& !value.ends_with('-')
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'-')
|
||||
}
|
||||
|
||||
/// Trims a region and falls back to the known-safe default when it is not a
|
||||
/// single DNS label. The returned value is suitable for URL and `Host`
|
||||
/// header interpolation.
|
||||
pub fn normalize_kiro_region(value: &str) -> &str {
|
||||
let value = value.trim();
|
||||
if is_valid_kiro_region(value) {
|
||||
value
|
||||
} else {
|
||||
DEFAULT_REGION
|
||||
}
|
||||
}
|
||||
|
||||
impl KiroAuthConfig {
|
||||
pub fn from_json_value(value: &Value) -> Option<Self> {
|
||||
let object = value.as_object()?;
|
||||
@@ -93,19 +162,22 @@ impl KiroAuthConfig {
|
||||
}
|
||||
|
||||
pub fn effective_auth_region(&self) -> &str {
|
||||
self.auth_region
|
||||
.as_deref()
|
||||
.or(self.region.as_deref())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or(DEFAULT_REGION)
|
||||
for candidate in [self.auth_region.as_deref(), self.region.as_deref()] {
|
||||
let Some(candidate) = candidate else {
|
||||
continue;
|
||||
};
|
||||
let candidate = candidate.trim();
|
||||
if is_valid_kiro_region(candidate) {
|
||||
return candidate;
|
||||
}
|
||||
}
|
||||
DEFAULT_REGION
|
||||
}
|
||||
|
||||
pub fn effective_api_region(&self) -> &str {
|
||||
self.api_region
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(normalize_kiro_region)
|
||||
.unwrap_or(DEFAULT_REGION)
|
||||
}
|
||||
|
||||
@@ -304,7 +376,7 @@ impl KiroProviderOAuthAdapter {
|
||||
if !(200..300).contains(&response.status_code) {
|
||||
return Err(OAuthError::HttpStatus {
|
||||
status_code: response.status_code,
|
||||
body_excerpt: response.body_text.chars().take(500).collect(),
|
||||
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
|
||||
});
|
||||
}
|
||||
let payload = response
|
||||
@@ -383,7 +455,7 @@ impl KiroProviderOAuthAdapter {
|
||||
if !(200..300).contains(&response.status_code) {
|
||||
return Err(OAuthError::HttpStatus {
|
||||
status_code: response.status_code,
|
||||
body_excerpt: response.body_text.chars().take(500).collect(),
|
||||
body_excerpt: redacted_oauth_error_body_excerpt(&response.body_text),
|
||||
});
|
||||
}
|
||||
let payload = response
|
||||
@@ -648,7 +720,8 @@ fn secret_fingerprint(value: &str) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
generate_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter, IDC_AMZ_USER_AGENT,
|
||||
generate_kiro_machine_id, is_valid_kiro_region, normalize_kiro_region, KiroAuthConfig,
|
||||
KiroProviderOAuthAdapter, IDC_AMZ_USER_AGENT,
|
||||
};
|
||||
use crate::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse};
|
||||
use crate::provider::ProviderOAuthTransportContext;
|
||||
@@ -662,6 +735,39 @@ mod tests {
|
||||
response: serde_json::Value,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_auth_config_debug_output_redacts_all_token_material() {
|
||||
let config = KiroAuthConfig {
|
||||
auth_method: Some("idc".to_string()),
|
||||
refresh_token: Some("kiro-refresh-canary".to_string()),
|
||||
expires_at: Some(123),
|
||||
profile_arn: Some("kiro-profile-arn-canary".to_string()),
|
||||
region: Some("us-east-1".to_string()),
|
||||
auth_region: None,
|
||||
api_region: None,
|
||||
client_id: Some("kiro-client-id-canary".to_string()),
|
||||
client_secret: Some("kiro-client-secret-canary".to_string()),
|
||||
machine_id: Some("kiro-machine-id-canary".to_string()),
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: Some("kiro-access-canary".to_string()),
|
||||
};
|
||||
|
||||
let debug = format!("{config:?}");
|
||||
for secret in [
|
||||
"kiro-refresh-canary",
|
||||
"kiro-client-id-canary",
|
||||
"kiro-client-secret-canary",
|
||||
"kiro-access-canary",
|
||||
"kiro-profile-arn-canary",
|
||||
"kiro-machine-id-canary",
|
||||
] {
|
||||
assert!(!debug.contains(secret), "debug leaked {secret}");
|
||||
}
|
||||
assert!(debug.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthHttpExecutor for StaticExecutor {
|
||||
async fn execute(
|
||||
@@ -717,6 +823,106 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_region_values_that_can_escape_a_hostname() {
|
||||
for value in [
|
||||
"evil.example/",
|
||||
"evil.example\\",
|
||||
"evil.example?next=1",
|
||||
"evil.example#fragment",
|
||||
"evil@example",
|
||||
"us-east-1\r\nX-Injected: yes",
|
||||
"us.east.1",
|
||||
] {
|
||||
assert!(
|
||||
!is_valid_kiro_region(value),
|
||||
"region should be rejected: {value:?}"
|
||||
);
|
||||
assert_eq!(normalize_kiro_region(value), super::DEFAULT_REGION);
|
||||
}
|
||||
|
||||
for value in [
|
||||
"us-east-1",
|
||||
"us-gov-west-1",
|
||||
"us-iso-east-1",
|
||||
"eu-central-1",
|
||||
] {
|
||||
assert!(
|
||||
is_valid_kiro_region(value),
|
||||
"region should be accepted: {value}"
|
||||
);
|
||||
assert_eq!(normalize_kiro_region(value), value);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn effective_regions_fall_back_when_auth_config_is_malicious() {
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: None,
|
||||
refresh_token: None,
|
||||
expires_at: None,
|
||||
profile_arn: None,
|
||||
region: Some("eu-west-1".to_string()),
|
||||
auth_region: Some("evil.example/".to_string()),
|
||||
api_region: Some("evil.example/".to_string()),
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: None,
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
assert_eq!(auth_config.effective_auth_region(), "eu-west-1");
|
||||
assert_eq!(auth_config.effective_api_region(), super::DEFAULT_REGION);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refresh_urls_use_safe_default_for_malicious_auth_region() {
|
||||
let seen_request = Arc::new(Mutex::new(None));
|
||||
let executor = StaticExecutor {
|
||||
seen_request: Arc::clone(&seen_request),
|
||||
response: json!({
|
||||
"accessToken": "new-access-token",
|
||||
"refreshToken": "r".repeat(120),
|
||||
"expiresIn": 3600
|
||||
}),
|
||||
};
|
||||
let auth_config = KiroAuthConfig {
|
||||
auth_method: Some("social".to_string()),
|
||||
refresh_token: Some("r".repeat(120)),
|
||||
expires_at: Some(1),
|
||||
profile_arn: None,
|
||||
region: None,
|
||||
auth_region: Some("attacker.example/".to_string()),
|
||||
api_region: None,
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
machine_id: Some("machine".to_string()),
|
||||
kiro_version: None,
|
||||
system_version: None,
|
||||
node_version: None,
|
||||
access_token: None,
|
||||
};
|
||||
|
||||
KiroProviderOAuthAdapter::default()
|
||||
.refresh_auth_config(&executor, &test_ctx(), &auth_config)
|
||||
.await
|
||||
.expect("refresh should succeed");
|
||||
|
||||
let seen = seen_request
|
||||
.lock()
|
||||
.expect("mutex should lock")
|
||||
.clone()
|
||||
.expect("request should be captured");
|
||||
assert_eq!(
|
||||
seen.url,
|
||||
"https://prod.us-east-1.auth.desktop.kiro.dev/refreshToken"
|
||||
);
|
||||
assert!(!seen.url.contains("attacker.example"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn refreshes_social_auth_config_with_provider_adapter() {
|
||||
let seen_request = Arc::new(Mutex::new(None));
|
||||
|
||||
@@ -14,12 +14,14 @@ pub use claude_code::{
|
||||
pub use codex::CodexProviderOAuthAdapter;
|
||||
pub use generic::{
|
||||
derive_codex_identity_fingerprint, GenericProviderOAuthAdapter, GenericProviderOAuthTemplate,
|
||||
ANTIGRAVITY_OAUTH_CLIENT_ID_ENV, ANTIGRAVITY_OAUTH_CLIENT_SECRET_ENV,
|
||||
GEMINI_CLI_OAUTH_CLIENT_ID_ENV, GEMINI_CLI_OAUTH_CLIENT_SECRET_ENV,
|
||||
GENERIC_PROVIDER_OAUTH_TEMPLATES,
|
||||
};
|
||||
pub use kiro::{
|
||||
generate_kiro_machine_id, normalize_kiro_machine_id, KiroAuthConfig, KiroProviderOAuthAdapter,
|
||||
DEFAULT_KIRO_VERSION, DEFAULT_NODE_VERSION, DEFAULT_REGION, DEFAULT_SYSTEM_VERSION,
|
||||
KIRO_PROVIDER_TYPE,
|
||||
generate_kiro_machine_id, is_valid_kiro_region, normalize_kiro_machine_id,
|
||||
normalize_kiro_region, KiroAuthConfig, KiroProviderOAuthAdapter, DEFAULT_KIRO_VERSION,
|
||||
DEFAULT_NODE_VERSION, DEFAULT_REGION, DEFAULT_SYSTEM_VERSION, KIRO_PROVIDER_TYPE,
|
||||
};
|
||||
pub use windsurf::{
|
||||
WindsurfProviderOAuthAdapter, WINDSURF_CLIENT_ID, WINDSURF_PROVIDER_TYPE,
|
||||
|
||||
Reference in New Issue
Block a user