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:
elky
2026-09-04 03:45:52 +08:00
parent ddcbeb3ae9
commit 579f2c7cc1
1019 changed files with 190437 additions and 26080 deletions
+703 -63
View File
@@ -1,16 +1,24 @@
use crate::admin_api::AdminAppState;
use crate::{AppState, GatewayError};
use aether_contracts::{
ExecutionPlan, ExecutionResult, ExecutionTimeouts, RequestBody,
ExecutionPlan, ExecutionResult, ExecutionTimeouts, ProxySnapshot, RequestBody,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
};
use aether_oauth::core::OAuthError;
use aether_oauth::network::{OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse};
use aether_oauth::network::{
OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkPolicy, OAuthTimeouts,
};
use async_trait::async_trait;
use base64::{engine::general_purpose::STANDARD, Engine as _};
use flate2::read::{DeflateDecoder, GzDecoder};
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue, CONTENT_ENCODING, CONTENT_TYPE};
use std::collections::BTreeMap;
use std::io::Read;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
const OAUTH_RESPONSE_BODY_LIMIT_BYTES: usize = 4 * 1024 * 1024;
#[derive(Clone)]
pub(crate) struct GatewayOAuthHttpExecutor<'a> {
@@ -37,57 +45,421 @@ impl<'a> GatewayOAuthHttpExecutor<'a> {
#[async_trait]
impl<'a> OAuthHttpExecutor for GatewayOAuthHttpExecutor<'a> {
async fn execute(&self, request: OAuthHttpRequest) -> Result<OAuthHttpResponse, OAuthError> {
let body = if let Some(json_body) = request.json_body {
RequestBody::from_json(json_body)
} else {
RequestBody {
json_body: None,
body_bytes_b64: request.body_bytes.map(|bytes| STANDARD.encode(bytes)),
body_ref: None,
match request.network.policy {
OAuthNetworkPolicy::DirectOnly | OAuthNetworkPolicy::DirectOrSystemProxy => {
match identity_oauth_route(request.network.policy, request.network.proxy.as_ref())?
{
IdentityOAuthRoute::Direct => {
execute_direct_identity_oauth(&self.app, request).await
}
}
}
};
let timeouts = request.network.timeouts;
let mut headers = request.headers;
OAuthNetworkPolicy::ProviderOperationProxy => {
#[cfg(test)]
if self.app.execution_runtime_override_base_url().is_none()
&& identity_oauth_endpoint_policy(&self.app, &request.url)
== IdentityOAuthEndpointPolicy::ExplicitTestLoopback
{
return execute_direct_identity_oauth(&self.app, request).await;
}
execute_provider_oauth_via_runtime(&self.app, request).await
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum IdentityOAuthRoute {
Direct,
}
fn identity_oauth_route(
policy: OAuthNetworkPolicy,
proxy: Option<&ProxySnapshot>,
) -> Result<IdentityOAuthRoute, OAuthError> {
let Some(proxy) = proxy.filter(|proxy| proxy.enabled != Some(false)) else {
return Ok(IdentityOAuthRoute::Direct);
};
if policy == OAuthNetworkPolicy::DirectOnly {
return Err(OAuthError::transport(
"identity OAuth direct-only transport cannot use a proxy",
));
}
let has_proxy_url = proxy
.url
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if has_proxy_url {
return Err(OAuthError::transport(
"identity OAuth HTTP/SOCKS proxies are disabled; use a controlled tunnel or direct transport",
));
}
Err(OAuthError::transport(
"identity OAuth cannot use a configured proxy or tunnel; use direct transport",
))
}
async fn execute_provider_oauth_via_runtime(
app: &AppState,
request: OAuthHttpRequest,
) -> Result<OAuthHttpResponse, OAuthError> {
let plan = oauth_execution_plan(request, false);
let result = crate::execution_runtime::execute_execution_runtime_sync_plan(app, None, &plan)
.await
.map_err(gateway_error_to_oauth_error)?;
Ok(execution_result_to_oauth_response(&result))
}
fn oauth_execution_plan(request: OAuthHttpRequest, force_disable_redirects: bool) -> ExecutionPlan {
let OAuthHttpRequest {
request_id,
method,
url,
mut headers,
content_type,
json_body,
body_bytes,
network,
transport_profile,
} = request;
let body = if let Some(json_body) = json_body {
RequestBody::from_json(json_body)
} else {
RequestBody {
json_body: None,
body_bytes_b64: body_bytes.map(|bytes| STANDARD.encode(bytes)),
body_ref: None,
}
};
let timeouts = network.timeouts;
if force_disable_redirects {
headers.insert(
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
"false".to_string(),
);
} else {
headers
.entry(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string())
.or_insert_with(|| "true".to_string());
let plan = ExecutionPlan {
request_id: request.request_id,
candidate_id: None,
provider_name: Some("oauth".to_string()),
provider_id: String::new(),
endpoint_id: String::new(),
key_id: String::new(),
method: request.method.as_str().to_string(),
url: request.url,
headers,
content_type: request.content_type,
content_encoding: None,
body,
stream: false,
client_api_format: "oauth:exchange".to_string(),
provider_api_format: "oauth:exchange".to_string(),
model_name: Some("oauth-exchange".to_string()),
proxy: request.network.proxy,
transport_profile: request.transport_profile,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(timeouts.connect_ms),
read_ms: Some(timeouts.read_ms),
write_ms: Some(timeouts.write_ms),
pool_ms: Some(timeouts.connect_ms),
total_ms: Some(timeouts.total_ms),
..ExecutionTimeouts::default()
}),
};
let result =
crate::execution_runtime::execute_execution_runtime_sync_plan(&self.app, None, &plan)
.await
.map_err(gateway_error_to_oauth_error)?;
Ok(OAuthHttpResponse {
status_code: result.status_code,
body_text: execution_body_text(&result),
json_body: execution_json_body(&result),
})
}
let plan = ExecutionPlan {
request_id,
candidate_id: None,
provider_name: Some("oauth".to_string()),
provider_id: String::new(),
endpoint_id: String::new(),
key_id: String::new(),
method: method.as_str().to_string(),
url,
headers,
content_type,
content_encoding: None,
body,
stream: false,
client_api_format: "oauth:exchange".to_string(),
provider_api_format: "oauth:exchange".to_string(),
model_name: Some("oauth-exchange".to_string()),
proxy: network.proxy,
transport_profile,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(timeouts.connect_ms),
read_ms: Some(timeouts.read_ms),
write_ms: Some(timeouts.write_ms),
pool_ms: Some(timeouts.connect_ms),
total_ms: Some(timeouts.total_ms),
..ExecutionTimeouts::default()
}),
};
crate::execution_runtime::transport::with_upstream_response_body_limit(
&plan,
OAUTH_RESPONSE_BODY_LIMIT_BYTES,
)
}
#[derive(Debug)]
struct ResolvedIdentityOAuthEndpoint {
url: reqwest::Url,
host: String,
addrs: Vec<SocketAddr>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum IdentityOAuthEndpointPolicy {
PublicHttps,
#[cfg(test)]
ExplicitTestLoopback,
}
fn identity_oauth_endpoint_policy(app: &AppState, raw_url: &str) -> IdentityOAuthEndpointPolicy {
#[cfg(test)]
{
let parsed_url = reqwest::Url::parse(raw_url).ok();
let is_explicit_test_url = app
.provider_oauth_token_url_overrides
.lock()
.expect("provider oauth token URL overrides should lock")
.values()
.any(|value| {
parsed_url
.as_ref()
.is_some_and(|url| test_loopback_override_allows_url(value, url))
});
if is_explicit_test_url {
return IdentityOAuthEndpointPolicy::ExplicitTestLoopback;
}
}
let _ = (app, raw_url);
IdentityOAuthEndpointPolicy::PublicHttps
}
#[cfg(test)]
fn test_loopback_override_allows_url(override_url: &str, target: &reqwest::Url) -> bool {
let Ok(registered) = reqwest::Url::parse(override_url) else {
return false;
};
let Some(target_ip) = target
.host_str()
.and_then(|host| host.parse::<IpAddr>().ok())
.filter(|address| address.is_loopback())
else {
return false;
};
let same_origin = registered.scheme() == target.scheme()
&& registered
.host_str()
.and_then(|host| host.parse::<IpAddr>().ok())
== Some(target_ip)
&& registered.port_or_known_default() == target.port_or_known_default();
if !same_origin || registered.query().is_some() || registered.fragment().is_some() {
return registered == *target;
}
let registered_path = registered.path().trim_end_matches('/');
registered == *target
|| registered_path.is_empty()
|| target
.path()
.strip_prefix(registered_path)
.is_some_and(|suffix| suffix.starts_with('/'))
}
fn parse_identity_oauth_endpoint(
raw_url: &str,
policy: IdentityOAuthEndpointPolicy,
) -> Result<(reqwest::Url, String, u16), OAuthError> {
let url = reqwest::Url::parse(raw_url)
.map_err(|_| OAuthError::transport("identity OAuth endpoint URL is invalid"))?;
if !url.username().is_empty() || url.password().is_some() || url.fragment().is_some() {
return Err(OAuthError::transport(
"identity OAuth endpoint must not contain credentials or a fragment",
));
}
let host = url
.host_str()
.map(ToOwned::to_owned)
.ok_or_else(|| OAuthError::transport("identity OAuth endpoint is missing a host"))?;
match policy {
IdentityOAuthEndpointPolicy::PublicHttps if url.scheme() != "https" => {
return Err(OAuthError::transport(
"identity OAuth endpoint must use HTTPS",
));
}
#[cfg(test)]
IdentityOAuthEndpointPolicy::ExplicitTestLoopback => {
let is_loopback_literal = host
.parse::<IpAddr>()
.is_ok_and(|address| address.is_loopback());
if !matches!(url.scheme(), "http" | "https") || !is_loopback_literal {
return Err(OAuthError::transport(
"test identity OAuth endpoint must use a loopback IP literal",
));
}
}
_ => {}
}
let port = url
.port_or_known_default()
.ok_or_else(|| OAuthError::transport("identity OAuth endpoint is missing a port"))?;
Ok((url, host, port))
}
fn validate_identity_oauth_resolved_addrs(
addrs: &[SocketAddr],
policy: IdentityOAuthEndpointPolicy,
) -> Result<(), OAuthError> {
if addrs.is_empty() {
return Err(OAuthError::transport(
"identity OAuth endpoint DNS resolution returned no addresses",
));
}
#[cfg(test)]
if policy == IdentityOAuthEndpointPolicy::ExplicitTestLoopback {
if addrs.iter().all(|addr| addr.ip().is_loopback()) {
return Ok(());
}
return Err(OAuthError::transport(
"test identity OAuth endpoint must resolve only to loopback addresses",
));
}
if addrs
.iter()
.any(|addr| aether_http::is_private_or_reserved_ip(addr.ip()))
{
return Err(OAuthError::transport(
"identity OAuth endpoint resolves to a private or reserved address",
));
}
Ok(())
}
async fn resolve_identity_oauth_endpoint(
raw_url: &str,
policy: IdentityOAuthEndpointPolicy,
) -> Result<ResolvedIdentityOAuthEndpoint, OAuthError> {
let (url, host, port) = parse_identity_oauth_endpoint(raw_url, policy)?;
let addrs = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host.as_str(),
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| OAuthError::transport("identity OAuth endpoint DNS resolution failed"))?
};
validate_identity_oauth_resolved_addrs(&addrs, policy)?;
Ok(ResolvedIdentityOAuthEndpoint { url, host, addrs })
}
fn build_pinned_identity_oauth_client(
host: &str,
addrs: &[SocketAddr],
timeouts: OAuthTimeouts,
) -> Result<reqwest::Client, OAuthError> {
reqwest::Client::builder()
.no_proxy()
.redirect(identity_oauth_redirect_policy())
.connect_timeout(Duration::from_millis(timeouts.connect_ms))
.read_timeout(Duration::from_millis(timeouts.read_ms))
.timeout(Duration::from_millis(timeouts.total_ms))
.resolve_to_addrs(host, addrs)
.build()
.map_err(|_| OAuthError::transport("identity OAuth HTTP client initialization failed"))
}
fn identity_oauth_redirect_policy() -> reqwest::redirect::Policy {
reqwest::redirect::Policy::none()
}
fn identity_oauth_headers(
headers: &BTreeMap<String, String>,
content_type: Option<&str>,
) -> Result<HeaderMap, OAuthError> {
let mut result = HeaderMap::new();
for (name, value) in headers {
if name.eq_ignore_ascii_case(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER) {
continue;
}
let name = HeaderName::from_bytes(name.as_bytes())
.map_err(|_| OAuthError::transport("identity OAuth request has an invalid header"))?;
let value = HeaderValue::from_str(value)
.map_err(|_| OAuthError::transport("identity OAuth request has an invalid header"))?;
result.insert(name, value);
}
if !result.contains_key(CONTENT_TYPE) {
if let Some(content_type) = content_type {
let value = HeaderValue::from_str(content_type).map_err(|_| {
OAuthError::transport("identity OAuth request has an invalid content type")
})?;
result.insert(CONTENT_TYPE, value);
}
}
Ok(result)
}
async fn execute_direct_identity_oauth(
app: &AppState,
request: OAuthHttpRequest,
) -> Result<OAuthHttpResponse, OAuthError> {
let endpoint_policy = identity_oauth_endpoint_policy(app, &request.url);
let target = resolve_identity_oauth_endpoint(&request.url, endpoint_policy).await?;
let client = build_pinned_identity_oauth_client(
target.host.as_str(),
&target.addrs,
request.network.timeouts,
)?;
let headers = identity_oauth_headers(&request.headers, request.content_type.as_deref())?;
let mut builder = client.request(request.method, target.url).headers(headers);
if let Some(json_body) = request.json_body {
builder = builder.json(&json_body);
} else if let Some(body_bytes) = request.body_bytes {
builder = builder.body(body_bytes);
}
let response = builder
.send()
.await
.map_err(|_| OAuthError::transport("identity OAuth request failed"))?;
let status_code = response.status().as_u16();
let content_encoding = response
.headers()
.get(CONTENT_ENCODING)
.and_then(|value| value.to_str().ok())
.map(ToOwned::to_owned);
let body_bytes = collect_identity_oauth_response_body(response).await?;
let decoded = decode_response_bytes_with_limit(
&body_bytes,
content_encoding.as_deref(),
OAUTH_RESPONSE_BODY_LIMIT_BYTES,
)?
.unwrap_or(body_bytes);
Ok(OAuthHttpResponse {
status_code,
body_text: String::from_utf8_lossy(&decoded).to_string(),
json_body: serde_json::from_slice(&decoded).ok(),
})
}
async fn collect_identity_oauth_response_body(
response: reqwest::Response,
) -> Result<Vec<u8>, OAuthError> {
if response
.content_length()
.is_some_and(|length| length > OAUTH_RESPONSE_BODY_LIMIT_BYTES as u64)
{
return Err(oauth_response_too_large());
}
let mut body = Vec::new();
let mut stream = response.bytes_stream();
while let Some(chunk) = stream.next().await {
let chunk =
chunk.map_err(|_| OAuthError::transport("identity OAuth response body read failed"))?;
if chunk.len() > OAUTH_RESPONSE_BODY_LIMIT_BYTES.saturating_sub(body.len()) {
return Err(oauth_response_too_large());
}
body.extend_from_slice(&chunk);
}
Ok(body)
}
fn oauth_response_too_large() -> OAuthError {
OAuthError::transport(format!(
"OAuth response body exceeds {OAUTH_RESPONSE_BODY_LIMIT_BYTES} bytes"
))
}
fn execution_result_to_oauth_response(result: &ExecutionResult) -> OAuthHttpResponse {
OAuthHttpResponse {
status_code: result.status_code,
body_text: execution_body_text(result),
json_body: execution_json_body(result),
}
}
@@ -125,15 +497,28 @@ fn execution_body_bytes(
headers: &BTreeMap<String, String>,
body: &aether_contracts::ResponseBody,
) -> Option<Vec<u8>> {
let bytes = body
.body_bytes_b64
.as_deref()
.and_then(|value| STANDARD.decode(value).ok())?;
decode_response_bytes(&bytes, headers.get("content-encoding").map(String::as_str))
.or(Some(bytes))
let bytes = body.body_bytes_b64.as_deref().and_then(|value| {
crate::execution_runtime::transport::decode_base64_body_with_limit(
value,
OAUTH_RESPONSE_BODY_LIMIT_BYTES,
)
.ok()
})?;
decode_response_bytes_with_limit(
&bytes,
headers.get("content-encoding").map(String::as_str),
OAUTH_RESPONSE_BODY_LIMIT_BYTES,
)
.ok()
.flatten()
.or(Some(bytes))
}
fn decode_response_bytes(bytes: &[u8], content_encoding: Option<&str>) -> Option<Vec<u8>> {
fn decode_response_bytes_with_limit(
bytes: &[u8],
content_encoding: Option<&str>,
limit_bytes: usize,
) -> Result<Option<Vec<u8>>, OAuthError> {
match content_encoding
.map(str::trim)
.filter(|value| !value.is_empty())
@@ -142,20 +527,275 @@ fn decode_response_bytes(bytes: &[u8], content_encoding: Option<&str>) -> Option
{
Some("gzip") => {
let mut decoder = GzDecoder::new(bytes);
let mut out = Vec::new();
decoder.read_to_end(&mut out).ok()?;
Some(out)
read_oauth_response_decoder(&mut decoder, limit_bytes).map(Some)
}
Some("deflate") => {
let mut decoder = DeflateDecoder::new(bytes);
let mut out = Vec::new();
decoder.read_to_end(&mut out).ok()?;
Some(out)
read_oauth_response_decoder(&mut decoder, limit_bytes).map(Some)
}
_ => None,
_ => Ok(None),
}
}
fn read_oauth_response_decoder(
decoder: &mut impl Read,
limit_bytes: usize,
) -> Result<Vec<u8>, OAuthError> {
let read_limit = u64::try_from(limit_bytes)
.unwrap_or(u64::MAX)
.saturating_add(1);
let mut limited = decoder.take(read_limit);
let mut out = Vec::new();
limited
.read_to_end(&mut out)
.map_err(|_| OAuthError::transport("OAuth response body decompression failed"))?;
if out.len() > limit_bytes {
return Err(oauth_response_too_large());
}
Ok(out)
}
fn gateway_error_to_oauth_error(error: GatewayError) -> OAuthError {
OAuthError::Transport(error.into_message())
}
#[cfg(test)]
mod tests {
use super::{
decode_response_bytes_with_limit, execution_body_bytes, identity_oauth_endpoint_policy,
identity_oauth_redirect_policy, identity_oauth_route, oauth_execution_plan,
parse_identity_oauth_endpoint, resolve_identity_oauth_endpoint,
validate_identity_oauth_resolved_addrs, IdentityOAuthEndpointPolicy, IdentityOAuthRoute,
};
use aether_contracts::{
ProxySnapshot, ResponseBody, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
};
use aether_oauth::network::{
OAuthHttpRequest, OAuthNetworkContext, OAuthNetworkPolicy, OAuthTimeouts,
};
use std::collections::BTreeMap;
use std::io::Write;
use std::net::SocketAddr;
fn identity_request(proxy: Option<ProxySnapshot>) -> OAuthHttpRequest {
OAuthHttpRequest {
request_id: "identity-oauth:test".to_string(),
method: reqwest::Method::GET,
url: "https://oauth.example.test/userinfo".to_string(),
headers: BTreeMap::new(),
content_type: None,
json_body: None,
body_bytes: None,
network: OAuthNetworkContext {
policy: OAuthNetworkPolicy::DirectOrSystemProxy,
requirement: aether_oauth::network::NetworkRequirement::Optional,
proxy,
timeouts: OAuthTimeouts::DIRECT_DEFAULT,
},
transport_profile: None,
}
}
#[tokio::test]
async fn private_ip_literal_is_rejected_before_connect() {
let error = resolve_identity_oauth_endpoint(
"https://127.0.0.1/token",
IdentityOAuthEndpointPolicy::PublicHttps,
)
.await
.expect_err("loopback identity endpoint must be rejected");
assert!(error.to_string().contains("private or reserved"));
}
#[test]
fn public_https_endpoint_and_resolved_addresses_are_accepted_without_dns() {
let (url, host, port) = parse_identity_oauth_endpoint(
"https://oauth.example.test/token?flow=login",
IdentityOAuthEndpointPolicy::PublicHttps,
)
.expect("public HTTPS URL should parse");
let addrs = ["8.8.8.8:443".parse::<SocketAddr>().unwrap()];
assert_eq!(url.scheme(), "https");
assert_eq!(host, "oauth.example.test");
assert_eq!(port, 443);
validate_identity_oauth_resolved_addrs(&addrs, IdentityOAuthEndpointPolicy::PublicHttps)
.expect("controlled public address should pass validation");
}
#[test]
fn identity_oauth_rejects_credentials_fragments_and_any_private_dns_answer() {
assert!(parse_identity_oauth_endpoint(
"https://user:[email protected]/token",
IdentityOAuthEndpointPolicy::PublicHttps,
)
.is_err());
assert!(parse_identity_oauth_endpoint(
"https://example.test/token#secret",
IdentityOAuthEndpointPolicy::PublicHttps,
)
.is_err());
assert!(parse_identity_oauth_endpoint(
"http://example.test/token",
IdentityOAuthEndpointPolicy::PublicHttps,
)
.is_err());
let mixed = [
"8.8.8.8:443".parse::<SocketAddr>().unwrap(),
"10.0.0.4:443".parse::<SocketAddr>().unwrap(),
];
assert!(validate_identity_oauth_resolved_addrs(
&mixed,
IdentityOAuthEndpointPolicy::PublicHttps,
)
.is_err());
}
#[test]
fn explicit_test_endpoint_policy_only_allows_loopback_literals() {
let (_, host, port) = parse_identity_oauth_endpoint(
"http://127.0.0.1:32123/token",
IdentityOAuthEndpointPolicy::ExplicitTestLoopback,
)
.expect("explicit test loopback URL should parse");
assert_eq!(host, "127.0.0.1");
assert_eq!(port, 32123);
validate_identity_oauth_resolved_addrs(
&["127.0.0.1:32123".parse().unwrap()],
IdentityOAuthEndpointPolicy::ExplicitTestLoopback,
)
.expect("loopback resolution should be accepted for an explicit test endpoint");
assert!(parse_identity_oauth_endpoint(
"http://10.0.0.1/token",
IdentityOAuthEndpointPolicy::ExplicitTestLoopback,
)
.is_err());
assert!(parse_identity_oauth_endpoint(
"http://localhost/token",
IdentityOAuthEndpointPolicy::ExplicitTestLoopback,
)
.is_err());
assert!(validate_identity_oauth_resolved_addrs(
&["10.0.0.1:80".parse().unwrap()],
IdentityOAuthEndpointPolicy::ExplicitTestLoopback,
)
.is_err());
}
#[test]
fn test_loopback_policy_requires_a_registered_loopback_origin_and_path() {
let app = crate::AppState::new()
.expect("gateway state should build")
.with_provider_oauth_token_url_for_tests("codex", "http://127.0.0.1:32123/oauth")
.with_provider_oauth_token_url_for_tests("bad", "http://10.0.0.1/token");
assert_eq!(
identity_oauth_endpoint_policy(&app, "http://127.0.0.1:32123/oauth/token"),
IdentityOAuthEndpointPolicy::ExplicitTestLoopback
);
assert_eq!(
identity_oauth_endpoint_policy(&app, "http://127.0.0.1:32123/oauth2/token"),
IdentityOAuthEndpointPolicy::PublicHttps
);
assert_eq!(
identity_oauth_endpoint_policy(&app, "http://127.0.0.1:32124/oauth/token"),
IdentityOAuthEndpointPolicy::PublicHttps
);
assert_eq!(
identity_oauth_endpoint_policy(&app, "http://10.0.0.1/token"),
IdentityOAuthEndpointPolicy::PublicHttps
);
}
#[test]
fn identity_oauth_disables_redirects_for_direct_and_tunnel_requests() {
assert_eq!(
format!("{:?}", identity_oauth_redirect_policy()),
"Policy(None)"
);
let tunnel = ProxySnapshot {
enabled: Some(true),
mode: Some("tunnel".to_string()),
node_id: Some("node-1".to_string()),
..ProxySnapshot::default()
};
assert!(
identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, Some(&tunnel)).is_err()
);
let plan = oauth_execution_plan(identity_request(None), true);
assert_eq!(
plan.headers
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
.map(String::as_str),
Some("false")
);
assert_eq!(
crate::execution_runtime::transport::execution_plan_response_body_limit_bytes(&plan),
super::OAUTH_RESPONSE_BODY_LIMIT_BYTES
);
}
#[test]
fn identity_oauth_rejects_forward_proxy_snapshots() {
let proxy = ProxySnapshot {
enabled: Some(true),
mode: Some("http".to_string()),
url: Some("http://proxy.example.test:8080".to_string()),
..ProxySnapshot::default()
};
assert!(
identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, Some(&proxy)).is_err()
);
assert_eq!(
identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, None).unwrap(),
IdentityOAuthRoute::Direct
);
}
#[test]
fn identity_oauth_rejects_enabled_tunnel_even_when_node_id_is_present() {
let proxy = ProxySnapshot {
enabled: Some(true),
mode: Some("tunnel".to_string()),
node_id: Some("node-1".to_string()),
..ProxySnapshot::default()
};
assert!(
identity_oauth_route(OAuthNetworkPolicy::DirectOrSystemProxy, Some(&proxy)).is_err()
);
}
#[test]
fn oauth_response_decoder_rejects_decompression_bombs() {
let payload = b"123456789";
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
encoder
.write_all(payload)
.expect("gzip payload should write");
let encoded = encoder.finish().expect("gzip payload should finish");
let error = decode_response_bytes_with_limit(&encoded, Some("gzip"), 8)
.expect_err("decoded OAuth body above the limit must fail closed");
assert!(error.to_string().contains("exceeds"));
}
#[test]
fn oauth_execution_body_rejects_oversized_base64_before_decode() {
let encoded_limit =
crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit(
super::OAUTH_RESPONSE_BODY_LIMIT_BYTES,
);
let body = ResponseBody {
json_body: None,
body_bytes_b64: Some("A".repeat(encoded_limit + 1)),
};
assert!(execution_body_bytes(&BTreeMap::new(), &body).is_none());
}
}
+175 -68
View File
@@ -1,10 +1,16 @@
use crate::handlers::shared::{
decrypt_catalog_secret_with_fallbacks, module_available_from_env,
decrypt_or_migrate_identity_oauth_provider_client_secret, module_available_from_env,
system_config_bool as system_config_bool_with_default,
};
use crate::{AppState, GatewayError};
use aether_data::repository::oauth_providers::StoredOAuthProviderConfig;
use aether_data::repository::users::{StoredUserAuthRecord, StoredUserOAuthLinkSummary};
use aether_data::repository::oauth_providers::{
validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config,
validate_oauth_redirect_uri, StoredOAuthProviderConfig,
};
use aether_data::repository::users::{
BindUserOAuthLinkOutcome, BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome,
ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord, StoredUserOAuthLinkSummary,
};
use aether_oauth::identity::{IdentityClaims, IdentityOAuthProviderConfig};
use chrono::Utc;
use serde::Serialize;
@@ -15,6 +21,12 @@ const LINUXDO_AUTHORIZE_URL: &str = "https://connect.linux.do/oauth2/authorize";
const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token";
const LINUXDO_USERINFO_URL: &str = "https://connect.linux.do/api/user";
static IDENTITY_OAUTH_MUTATION_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
pub(crate) async fn lock_identity_oauth_mutation() -> tokio::sync::MutexGuard<'static, ()> {
IDENTITY_OAUTH_MUTATION_LOCK.lock().await
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub(crate) struct IdentityOAuthProviderSummary {
pub(crate) provider_type: String,
@@ -45,6 +57,7 @@ pub(crate) enum IdentityOAuthAccountError {
AlreadyBoundProvider,
LastOAuthBinding,
LastLoginMethod,
BindingSessionUnavailable,
Storage(String),
}
@@ -60,6 +73,7 @@ impl IdentityOAuthAccountError {
Self::AlreadyBoundProvider => "already_bound_provider",
Self::LastOAuthBinding => "last_oauth_binding",
Self::LastLoginMethod => "last_login_method",
Self::BindingSessionUnavailable => "invalid_state",
}
}
@@ -112,7 +126,9 @@ pub(crate) async fn get_enabled_identity_oauth_provider_config(
let Some(config) = config.filter(|config| config.is_enabled) else {
return Ok(None);
};
stored_provider_config_to_identity_config(state, config).map(Some)
stored_provider_config_to_identity_config(state, config)
.await
.map(Some)
}
async fn identity_oauth_module_enabled(state: &AppState) -> Result<bool, GatewayError> {
@@ -161,27 +177,39 @@ pub(crate) async fn resolve_identity_oauth_login_user(
state: &AppState,
claims: &IdentityClaims,
) -> Result<StoredUserAuthRecord, IdentityOAuthAccountError> {
let _mutation_guard = lock_identity_oauth_mutation().await;
let now = Utc::now();
if let Some(user) = state
let provider_enabled = state
.get_oauth_provider_config(&claims.provider_type)
.await
.map_err(|err| IdentityOAuthAccountError::Storage(format!("{err:?}")))?
.is_some_and(|provider| provider.is_enabled);
let verified_email = claims
.email_verified
.then(|| normalize_identity_email(claims.email.as_deref()))
.flatten();
match state
.data
.find_oauth_linked_user(&claims.provider_type, &claims.subject)
.resolve_enabled_oauth_linked_user(
&claims.provider_type,
&claims.subject,
claims.username.as_deref(),
claims.email.as_deref(),
None,
verified_email.as_deref(),
now,
provider_enabled,
)
.await
.map_err(repo_data_error)?
{
state
.data
.touch_oauth_link(
&claims.provider_type,
&claims.subject,
claims.username.as_deref(),
claims.email.as_deref(),
Some(claims.raw.clone()),
now,
)
.await
.map_err(repo_data_error)?;
return Ok(user);
ResolveOAuthLinkedUserOutcome::Linked(user) => return Ok(user),
ResolveOAuthLinkedUserOutcome::NotLinked => {}
ResolveOAuthLinkedUserOutcome::ProviderUnavailable => {
return Err(IdentityOAuthAccountError::ProviderUnavailable);
}
}
drop(_mutation_guard);
let email = normalize_identity_email(claims.email.as_deref());
if let Some(email) = email.as_deref() {
@@ -219,35 +247,44 @@ pub(crate) async fn resolve_identity_oauth_login_user(
.unwrap_or(10.0);
let username = unique_oauth_username(state, claims).await?;
let email_verified = email.is_some() && claims.email_verified;
let user = state
.data
.create_oauth_auth_user(email, username, now)
.create_oauth_auth_user(email, email_verified, username, now)
.await
.map_err(repo_data_error)?
.ok_or_else(|| IdentityOAuthAccountError::Storage("oauth user not created".to_string()))?;
match state
.initialize_auth_user_wallet(&user.id, initial_gift, false)
let owned_wallet_id = match state
.initialize_auth_user_wallet_with_outcome(&user.id, initial_gift, false)
.await
{
Ok(Some(_wallet)) => {}
Ok(Some(outcome)) => outcome.created.then(|| outcome.wallet.id),
Ok(None) => {
let _ = state.delete_local_auth_user(&user.id).await;
let _ = state
.rollback_provisional_auth_user_with_wallet(&user.id, None)
.await;
return Err(IdentityOAuthAccountError::ProviderUnavailable);
}
Err(err) => {
let _ = state.delete_local_auth_user(&user.id).await;
let _ = state
.rollback_provisional_auth_user_with_wallet(&user.id, None)
.await;
return Err(IdentityOAuthAccountError::Storage(format!("{err:?}")));
}
}
};
if let Err(err) = state
.assign_default_group_to_self_registered_user(&user.id)
.await
{
let _ = state.delete_local_auth_user(&user.id).await;
let _ = state
.rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref())
.await;
return Err(IdentityOAuthAccountError::Storage(format!("{err:?}")));
}
if let Err(err) = upsert_oauth_link(state, &user.id, claims, now).await {
let _ = state.delete_local_auth_user(&user.id).await;
if let Err(err) = bind_oauth_link_for_new_user(state, &user.id, claims, now).await {
let _ = state
.rollback_provisional_auth_user_with_wallet(&user.id, owned_wallet_id.as_deref())
.await;
return Err(err);
}
Ok(user)
@@ -257,29 +294,39 @@ pub(crate) async fn bind_identity_oauth_to_user(
state: &AppState,
user: &StoredUserAuthRecord,
claims: &IdentityClaims,
session_expectation: &BindUserOAuthLinkSessionExpectation,
) -> Result<(), IdentityOAuthAccountError> {
if user.auth_source.eq_ignore_ascii_case("ldap") {
return Err(IdentityOAuthAccountError::EmailIsLdap);
}
if let Some(owner) = state
.data
.find_oauth_link_owner(&claims.provider_type, &claims.subject)
.await
.map_err(repo_data_error)?
{
if owner != user.id {
let now = Utc::now();
match bind_oauth_link(state, &user.id, claims, now, Some(session_expectation)).await? {
BindUserOAuthLinkOutcome::Bound => {}
BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser
| BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider => {
return Err(IdentityOAuthAccountError::AlreadyBoundProvider);
}
BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser => {
return Err(IdentityOAuthAccountError::OAuthAlreadyBound);
}
BindUserOAuthLinkOutcome::SessionUnavailable => {
return Err(IdentityOAuthAccountError::BindingSessionUnavailable);
}
BindUserOAuthLinkOutcome::UserNotFound
| BindUserOAuthLinkOutcome::ProviderNotFound
| BindUserOAuthLinkOutcome::ProviderDisabled => {
return Err(IdentityOAuthAccountError::ProviderUnavailable);
}
}
if state
.data
.has_user_oauth_provider_link(&user.id, &claims.provider_type)
.await
.map_err(repo_data_error)?
{
return Err(IdentityOAuthAccountError::AlreadyBoundProvider);
if claims.email_verified {
if let Some(verified_email) = normalize_identity_email(claims.email.as_deref()) {
state
.data
.upgrade_oauth_email_verification_if_matches(&user.id, &verified_email, now)
.await
.map_err(repo_data_error)?;
}
}
upsert_oauth_link(state, &user.id, claims, Utc::now()).await?;
Ok(())
}
@@ -287,32 +334,58 @@ pub(crate) async fn unbind_identity_oauth(
state: &AppState,
user: &StoredUserAuthRecord,
provider_type: &str,
local_password_login_allowed: bool,
) -> Result<bool, IdentityOAuthAccountError> {
if user.auth_source.eq_ignore_ascii_case("ldap") {
return Err(IdentityOAuthAccountError::EmailIsLdap);
}
let link_count = state
let _mutation_guard = lock_identity_oauth_mutation().await;
let enabled_provider_types_snapshot = state
.list_oauth_provider_configs()
.await
.map_err(|err| IdentityOAuthAccountError::Storage(format!("{err:?}")))?
.into_iter()
.filter(|provider| provider.is_enabled)
.map(|provider| provider.provider_type)
.collect::<Vec<_>>();
let outcome = state
.data
.count_user_oauth_links(&user.id)
.delete_user_oauth_link(
&user.id,
provider_type.trim(),
local_password_login_allowed,
&enabled_provider_types_snapshot,
)
.await
.map_err(repo_data_error)?;
if user.auth_source.eq_ignore_ascii_case("oauth") && link_count <= 1 {
return Err(IdentityOAuthAccountError::LastOAuthBinding);
match outcome {
DeleteUserOAuthLinkOutcome::Deleted => Ok(true),
DeleteUserOAuthLinkOutcome::NotFound => Ok(false),
DeleteUserOAuthLinkOutcome::LastOAuthBinding => {
Err(IdentityOAuthAccountError::LastOAuthBinding)
}
DeleteUserOAuthLinkOutcome::LastLoginMethod => {
Err(IdentityOAuthAccountError::LastLoginMethod)
}
}
if !user.auth_source.eq_ignore_ascii_case("local") && link_count <= 1 {
return Err(IdentityOAuthAccountError::LastLoginMethod);
}
state
.data
.delete_user_oauth_link(&user.id, provider_type.trim())
.await
.map_err(repo_data_error)
}
fn stored_provider_config_to_identity_config(
async fn stored_provider_config_to_identity_config(
state: &AppState,
config: StoredOAuthProviderConfig,
) -> Result<IdentityOAuthProviderConfig, IdentityOAuthAccountError> {
validate_oauth_redirect_uri(config.redirect_uri.trim())
.map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?;
validate_oauth_frontend_callback_url(config.frontend_callback_url.trim())
.map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?;
validate_oauth_provider_endpoint_config(
&config.provider_type,
config.authorization_url_override.as_deref(),
config.token_url_override.as_deref(),
config.userinfo_url_override.as_deref(),
config.extra_config.as_ref(),
)
.map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?;
let defaults = identity_provider_defaults(&config.provider_type);
let authorization_url = config
.authorization_url_override
@@ -330,13 +403,9 @@ fn stored_provider_config_to_identity_config(
.userinfo_url_override
.clone()
.or_else(|| defaults.map(|defaults| defaults.2.to_string()));
let client_secret = match config.client_secret_encrypted.as_deref() {
Some(ciphertext) => Some(
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.ok_or(IdentityOAuthAccountError::ProviderUnavailable)?,
),
None => None,
};
let client_secret = decrypt_or_migrate_identity_oauth_provider_client_secret(state, &config)
.await
.map_err(|_| IdentityOAuthAccountError::ProviderUnavailable)?;
Ok(IdentityOAuthProviderConfig {
provider_type: config.provider_type,
@@ -421,27 +490,65 @@ async fn unique_oauth_username(
Ok(format!("oauth_{}", short_uuid()))
}
async fn upsert_oauth_link(
async fn bind_oauth_link(
state: &AppState,
user_id: &str,
claims: &IdentityClaims,
now: chrono::DateTime<Utc>,
) -> Result<(), IdentityOAuthAccountError> {
session_expectation: Option<&BindUserOAuthLinkSessionExpectation>,
) -> Result<BindUserOAuthLinkOutcome, IdentityOAuthAccountError> {
let _mutation_guard = lock_identity_oauth_mutation().await;
let provider_enabled = state
.get_oauth_provider_config(&claims.provider_type)
.await
.map_err(|err| IdentityOAuthAccountError::Storage(format!("{err:?}")))?
.is_some_and(|provider| provider.is_enabled);
if !provider_enabled {
return Ok(BindUserOAuthLinkOutcome::ProviderNotFound);
}
state
.data
.upsert_user_oauth_link(
.bind_user_oauth_link(
user_id,
&claims.provider_type,
&claims.subject,
claims.username.as_deref(),
claims.email.as_deref(),
Some(claims.raw.clone()),
None,
now,
provider_enabled,
session_expectation,
)
.await
.map_err(repo_data_error)
}
async fn bind_oauth_link_for_new_user(
state: &AppState,
user_id: &str,
claims: &IdentityClaims,
now: chrono::DateTime<Utc>,
) -> Result<(), IdentityOAuthAccountError> {
match bind_oauth_link(state, user_id, claims, now, None).await? {
BindUserOAuthLinkOutcome::Bound => Ok(()),
BindUserOAuthLinkOutcome::IdentityAlreadyBoundToUser
| BindUserOAuthLinkOutcome::IdentityBoundToAnotherUser => {
Err(IdentityOAuthAccountError::OAuthAlreadyBound)
}
BindUserOAuthLinkOutcome::UserAlreadyLinkedProvider => {
Err(IdentityOAuthAccountError::AlreadyBoundProvider)
}
BindUserOAuthLinkOutcome::SessionUnavailable => {
Err(IdentityOAuthAccountError::BindingSessionUnavailable)
}
BindUserOAuthLinkOutcome::UserNotFound
| BindUserOAuthLinkOutcome::ProviderNotFound
| BindUserOAuthLinkOutcome::ProviderDisabled => {
Err(IdentityOAuthAccountError::ProviderUnavailable)
}
}
}
fn normalize_identity_email(value: Option<&str>) -> Option<String> {
value
.map(str::trim)
+4 -4
View File
@@ -8,14 +8,14 @@ pub(crate) use http_executor::GatewayOAuthHttpExecutor;
pub(crate) use identity_repo::{
bind_identity_oauth_to_user, get_enabled_identity_oauth_provider_config,
list_bindable_identity_oauth_providers, list_enabled_identity_oauth_providers,
list_identity_oauth_links, resolve_identity_oauth_login_user, unbind_identity_oauth,
IdentityOAuthAccountError,
list_identity_oauth_links, lock_identity_oauth_mutation, resolve_identity_oauth_login_user,
unbind_identity_oauth, IdentityOAuthAccountError,
};
pub(crate) use provider_repo::ProviderOAuthRepository;
pub(crate) use proxy::{
resolve_identity_oauth_network_context, resolve_provider_oauth_operation_proxy_snapshot,
};
pub(crate) use state_store::{
consume_identity_oauth_state, save_identity_oauth_state, IdentityOAuthStateMode,
StoredIdentityOAuthState,
consume_identity_oauth_state, identity_oauth_state_storage_key, load_identity_oauth_state,
save_identity_oauth_state, IdentityOAuthStateMode, StoredIdentityOAuthState,
};
@@ -13,24 +13,6 @@ use aether_data_contracts::repository::provider_catalog::{
pub(crate) struct ProviderOAuthRepository;
impl ProviderOAuthRepository {
pub(crate) async fn update_provider_catalog_key_oauth_credentials(
state: &AdminAppState<'_>,
key_id: &str,
encrypted_api_key: &str,
encrypted_auth_config: Option<&str>,
expires_at_unix_secs: Option<u64>,
) -> Result<bool, GatewayError> {
state
.app()
.update_provider_catalog_key_oauth_credentials(
key_id,
encrypted_api_key,
encrypted_auth_config,
expires_at_unix_secs,
)
.await
}
pub(crate) async fn clear_provider_catalog_key_oauth_invalid_marker(
state: &AdminAppState<'_>,
key_id: &str,
+11 -4
View File
@@ -24,11 +24,18 @@ pub(crate) async fn resolve_provider_oauth_operation_proxy_snapshot(
temporary_proxy_node_id: Option<&str>,
configured_proxies: &[Option<&serde_json::Value>],
) -> Option<ProxySnapshot> {
if let Some(snapshot) = state
.resolve_admin_proxy_node_snapshot(temporary_proxy_node_id)
.await
if let Some(temporary_proxy_node_id) = temporary_proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
{
return Some(snapshot);
return state
.resolve_admin_proxy_node_snapshot(Some(temporary_proxy_node_id))
.await
.or_else(|| {
Some(crate::state::unavailable_proxy_snapshot(
"temporary_proxy_node_unavailable",
))
});
}
for proxy in configured_proxies {
+277 -8
View File
@@ -1,8 +1,11 @@
use crate::{AppState, GatewayError};
use aether_oauth::core::{current_unix_secs, generate_oauth_nonce};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
const IDENTITY_OAUTH_STATE_TTL_SECS: u64 = 10 * 60;
const IDENTITY_OAUTH_STATE_SECRET_PURPOSE: &str = "identity-oauth-state";
const IDENTITY_OAUTH_STATE_MAX_CLOCK_SKEW_SECS: u64 = 60;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
@@ -11,13 +14,15 @@ pub(crate) enum IdentityOAuthStateMode {
Bind,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[derive(Clone, PartialEq, Serialize, Deserialize)]
pub(crate) struct StoredIdentityOAuthState {
pub(crate) nonce: String,
pub(crate) provider_type: String,
pub(crate) mode: IdentityOAuthStateMode,
pub(crate) client_device_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) browser_binding_hash: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) pkce_verifier: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) bind_user_id: Option<String>,
@@ -26,17 +31,36 @@ pub(crate) struct StoredIdentityOAuthState {
pub(crate) created_at: u64,
}
impl std::fmt::Debug for StoredIdentityOAuthState {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("StoredIdentityOAuthState")
.field("nonce", &"[REDACTED]")
.field("provider_type", &self.provider_type)
.field("mode", &self.mode)
.field("client_device_id", &"[REDACTED]")
.field("browser_binding_hash", &"[REDACTED]")
.field("pkce_verifier", &"[REDACTED]")
.field("bind_user_id", &"[REDACTED]")
.field("bind_session_id", &"[REDACTED]")
.field("created_at", &self.created_at)
.finish()
}
}
impl StoredIdentityOAuthState {
pub(crate) fn login(
provider_type: impl Into<String>,
client_device_id: impl Into<String>,
pkce_verifier: Option<String>,
browser_binding_hash: Option<String>,
) -> Self {
Self {
nonce: generate_oauth_nonce(),
provider_type: provider_type.into(),
mode: IdentityOAuthStateMode::Login,
client_device_id: client_device_id.into(),
browser_binding_hash,
pkce_verifier,
bind_user_id: None,
bind_session_id: None,
@@ -48,6 +72,7 @@ impl StoredIdentityOAuthState {
provider_type: impl Into<String>,
client_device_id: impl Into<String>,
pkce_verifier: Option<String>,
browser_binding_hash: String,
user_id: impl Into<String>,
session_id: impl Into<String>,
) -> Self {
@@ -56,6 +81,7 @@ impl StoredIdentityOAuthState {
provider_type: provider_type.into(),
mode: IdentityOAuthStateMode::Bind,
client_device_id: client_device_id.into(),
browser_binding_hash: Some(browser_binding_hash),
pkce_verifier,
bind_user_id: Some(user_id.into()),
bind_session_id: Some(session_id.into()),
@@ -65,16 +91,43 @@ impl StoredIdentityOAuthState {
}
pub(crate) fn identity_oauth_state_storage_key(nonce: &str) -> String {
format!(
"identity_oauth_state:sha256:{:x}",
Sha256::digest(nonce.trim().as_bytes())
)
}
fn identity_oauth_state_secret_purpose(nonce: &str) -> String {
format!(
"{IDENTITY_OAUTH_STATE_SECRET_PURPOSE}:sha256:{:x}",
Sha256::digest(nonce.trim().as_bytes())
)
}
fn legacy_identity_oauth_state_storage_key(nonce: &str) -> String {
format!("identity_oauth_state:{}", nonce.trim())
}
fn is_generated_oauth_nonce(nonce: &str) -> bool {
let nonce = nonce.trim();
nonce.len() == 64
&& nonce
.bytes()
.all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
}
pub(crate) async fn save_identity_oauth_state(
state: &AppState,
record: &StoredIdentityOAuthState,
) -> Result<(), GatewayError> {
let key = identity_oauth_state_storage_key(&record.nonce);
let value =
let plaintext =
serde_json::to_string(record).map_err(|err| GatewayError::Internal(err.to_string()))?;
let purpose = identity_oauth_state_secret_purpose(&record.nonce);
let value = crate::handlers::shared::seal_runtime_secret_payload(state, &purpose, &plaintext)
.ok_or_else(|| {
GatewayError::Internal("identity OAuth state encryption unavailable".to_string())
})?;
state
.runtime_kv_setex(&key, &value, IDENTITY_OAUTH_STATE_TTL_SECS)
.await
@@ -85,10 +138,226 @@ pub(crate) async fn consume_identity_oauth_state(
nonce: &str,
) -> Result<Option<StoredIdentityOAuthState>, GatewayError> {
let key = identity_oauth_state_storage_key(nonce);
let raw = state.runtime_kv_getdel(&key).await?;
raw.map(|value| {
serde_json::from_str::<StoredIdentityOAuthState>(&value)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.transpose()
let raw = match state.runtime_kv_getdel(&key).await? {
Some(value) => Some(value),
None if is_generated_oauth_nonce(nonce) => {
state
.runtime_kv_getdel(&legacy_identity_oauth_state_storage_key(nonce))
.await?
}
None => None,
};
raw.map(|value| decode_identity_oauth_state(state, nonce, &value))
.transpose()
}
pub(crate) async fn load_identity_oauth_state(
state: &AppState,
nonce: &str,
) -> Result<Option<StoredIdentityOAuthState>, GatewayError> {
let key = identity_oauth_state_storage_key(nonce);
let raw = match state.runtime_kv_get(&key).await? {
Some(value) => Some(value),
None if is_generated_oauth_nonce(nonce) => {
state
.runtime_kv_get(&legacy_identity_oauth_state_storage_key(nonce))
.await?
}
None => None,
};
raw.map(|value| decode_identity_oauth_state(state, nonce, &value))
.transpose()
}
fn decode_identity_oauth_state(
state: &AppState,
expected_nonce: &str,
stored: &str,
) -> Result<StoredIdentityOAuthState, GatewayError> {
let expected_nonce = expected_nonce.trim();
let purpose = identity_oauth_state_secret_purpose(expected_nonce);
let plaintext = crate::handlers::shared::open_runtime_secret_payload(state, &purpose, stored)
// States created immediately before a rolling upgrade still live under the
// legacy key and were sealed with the fixed purpose. They remain safe to
// accept for their short TTL because the decoded record is checked against
// the callback nonce and all authority-bearing fields below.
.or_else(|| {
crate::handlers::shared::open_runtime_secret_payload(
state,
IDENTITY_OAUTH_STATE_SECRET_PURPOSE,
stored,
)
})
.ok_or_else(|| GatewayError::Internal("identity OAuth state is invalid".to_string()))?;
let record = serde_json::from_str::<StoredIdentityOAuthState>(&plaintext)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
validate_identity_oauth_state(expected_nonce, &record)?;
Ok(record)
}
fn validate_identity_oauth_state(
expected_nonce: &str,
record: &StoredIdentityOAuthState,
) -> Result<(), GatewayError> {
let invalid = || GatewayError::Internal("identity OAuth state is invalid".to_string());
if !is_generated_oauth_nonce(expected_nonce)
|| record.nonce != expected_nonce
|| !is_generated_oauth_nonce(&record.nonce)
|| record.provider_type.trim().is_empty()
|| record.provider_type != record.provider_type.trim().to_ascii_lowercase()
|| record.client_device_id.trim().is_empty()
|| !record
.browser_binding_hash
.as_deref()
.is_some_and(is_lower_hex_sha256)
|| !record
.pkce_verifier
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
{
return Err(invalid());
}
let mode_is_valid = match record.mode {
IdentityOAuthStateMode::Login => {
record.bind_user_id.is_none() && record.bind_session_id.is_none()
}
IdentityOAuthStateMode::Bind => {
record
.bind_user_id
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
&& record
.bind_session_id
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
}
};
let now = current_unix_secs();
let time_is_valid = record.created_at
<= now.saturating_add(IDENTITY_OAUTH_STATE_MAX_CLOCK_SKEW_SECS)
&& now.saturating_sub(record.created_at) <= IDENTITY_OAUTH_STATE_TTL_SECS;
if !mode_is_valid || !time_is_valid {
return Err(invalid());
}
Ok(())
}
fn is_lower_hex_sha256(value: &str) -> bool {
value.len() == 64
&& value
.bytes()
.all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
}
#[cfg(test)]
mod tests {
use super::{
decode_identity_oauth_state, identity_oauth_state_secret_purpose, StoredIdentityOAuthState,
};
use crate::{data::GatewayDataState, AppState};
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
fn state_with_encryption_key() -> AppState {
AppState::new()
.expect("test state should build")
.with_data_state_for_tests(
GatewayDataState::disabled()
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
}
#[test]
fn identity_oauth_state_ciphertext_is_bound_to_its_nonce() {
let state = state_with_encryption_key();
let record = StoredIdentityOAuthState::login(
"linuxdo",
"device-1",
Some("pkce-verifier".to_string()),
Some("a".repeat(64)),
);
let plaintext = serde_json::to_string(&record).expect("state should serialize");
let sealed = crate::handlers::shared::seal_runtime_secret_payload(
&state,
&identity_oauth_state_secret_purpose(&record.nonce),
&plaintext,
)
.expect("state should seal");
assert_eq!(
decode_identity_oauth_state(&state, &record.nonce, &sealed)
.expect("matching state should open"),
record
);
assert!(decode_identity_oauth_state(&state, &"b".repeat(64), &sealed).is_err());
}
#[test]
fn identity_oauth_state_debug_redacts_authorization_material() {
let record = StoredIdentityOAuthState::login(
"linuxdo",
"debug-secret-device",
Some("debug-secret-pkce".to_string()),
Some("debug-secret-binding-hash".to_string()),
);
let nonce = record.nonce.clone();
let rendered = format!("{record:?}");
for secret in [
nonce.as_str(),
"debug-secret-device",
"debug-secret-pkce",
"debug-secret-binding-hash",
] {
assert!(!rendered.contains(secret), "Debug output leaked {secret}");
}
assert!(rendered.contains("[REDACTED]"));
assert!(rendered.contains("linuxdo"));
}
#[test]
fn identity_oauth_state_accepts_valid_legacy_ciphertext_during_ttl_window() {
let state = state_with_encryption_key();
let record = StoredIdentityOAuthState::login(
"linuxdo",
"device-1",
Some("pkce-verifier".to_string()),
Some("a".repeat(64)),
);
let plaintext = serde_json::to_string(&record).expect("state should serialize");
let sealed = crate::handlers::shared::seal_runtime_secret_payload(
&state,
super::IDENTITY_OAUTH_STATE_SECRET_PURPOSE,
&plaintext,
)
.expect("legacy state should seal");
assert_eq!(
decode_identity_oauth_state(&state, &record.nonce, &sealed)
.expect("valid legacy state should open"),
record
);
assert!(decode_identity_oauth_state(&state, &"b".repeat(64), &sealed).is_err());
}
#[test]
fn identity_oauth_state_rejects_mode_field_confusion() {
let state = state_with_encryption_key();
let mut record = StoredIdentityOAuthState::login(
"linuxdo",
"device-1",
Some("pkce-verifier".to_string()),
Some("a".repeat(64)),
);
record.bind_user_id = Some("unexpected-user".to_string());
let plaintext = serde_json::to_string(&record).expect("state should serialize");
let sealed = crate::handlers::shared::seal_runtime_secret_payload(
&state,
&identity_oauth_state_secret_purpose(&record.nonce),
&plaintext,
)
.expect("state should seal");
assert!(decode_identity_oauth_state(&state, &record.nonce, &sealed).is_err());
}
}