feat(providers): expand OAuth account management

Add Claude Code manual and cookie authorization, including redacted batch tasks. Harden OAuth imports, duplicate replacement, provider dialogs, and related account-management tests.
This commit is contained in:
elky
2026-07-27 15:53:28 +08:00
parent 531cf11025
commit 550cc36760
55 changed files with 4957 additions and 403 deletions
@@ -1,7 +1,8 @@
use super::super::helpers::admin_provider_oauth_key_name_from_auth_config;
use super::super::token_import::{
build_provider_access_token_import_auth_config, decode_access_token_expires_at,
provider_oauth_import_authorization_bearer_token, provider_type_supports_access_token_import,
is_claude_session_key, provider_oauth_import_authorization_bearer_token,
provider_type_supports_access_token_import, validate_claude_access_token_import,
};
use super::kiro_import::execute_admin_provider_oauth_kiro_batch_import;
use super::parse::{
@@ -56,6 +57,33 @@ fn sanitize_windsurf_batch_import_error(error: &OAuthError) -> String {
}
}
fn can_fallback_batch_refresh_to_access_token(
provider_type: &str,
access_token: Option<&str>,
) -> bool {
!provider_type.eq_ignore_ascii_case("claude_code")
&& access_token.is_some()
&& provider_type_supports_access_token_import(provider_type)
}
fn validate_batch_access_token_import(
provider_type: &str,
access_token: &str,
imported_expires_at: Option<u64>,
now_unix_secs: u64,
) -> Result<(), String> {
if !provider_type_supports_access_token_import(provider_type) {
return Err(
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider".to_string(),
);
}
if provider_type.eq_ignore_ascii_case("claude_code") {
validate_claude_access_token_import(access_token, imported_expires_at, now_unix_secs)
.map_err(str::to_string)?;
}
Ok(())
}
const CODEX_AGENT_IDENTITY_SAFE_FIELDS: &[(&str, &[&str])] = &[
("agent_runtime_id", &["agent_runtime_id", "agentRuntimeId"]),
(
@@ -249,6 +277,15 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty());
let is_claude = provider_type.eq_ignore_ascii_case("claude_code");
if is_claude
&& refresh_token
.into_iter()
.chain(access_token)
.any(is_claude_session_key)
{
return Err("Claude sessionKey 请使用 Cookie 授权,不能作为导入凭据".to_string());
}
if provider_type.eq_ignore_ascii_case("windsurf") {
let token_for_import = refresh_token.or(access_token);
@@ -307,7 +344,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
if let Some(refresh_token) = refresh_token {
let Some(template) = template else {
if provider_type_supports_access_token_import(provider_type) {
if can_fallback_batch_refresh_to_access_token(provider_type, access_token) {
if let Some(access_token) = access_token {
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
provider_type,
@@ -340,7 +377,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
Ok(payload) => payload,
Err(response) => {
let detail = extract_admin_provider_oauth_batch_error_detail(response).await;
if provider_type_supports_access_token_import(provider_type) {
if can_fallback_batch_refresh_to_access_token(provider_type, access_token) {
if let Some(access_token) = access_token {
let (auth_config, expires_at) =
build_provider_access_token_import_auth_config(
@@ -381,9 +418,17 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
}
if let Some(access_token) = access_token {
if !provider_type_supports_access_token_import(provider_type) {
return Err("Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider".to_string());
}
let now_unix_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
validate_batch_access_token_import(
provider_type,
access_token,
entry.expires_at,
now_unix_secs,
)?;
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
provider_type,
access_token,
@@ -734,7 +779,8 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
mod tests {
use super::super::parse::parse_admin_provider_oauth_batch_import_entries;
use super::{
codex_agent_identity_auth_config_from_import, sanitize_windsurf_batch_import_error,
can_fallback_batch_refresh_to_access_token, codex_agent_identity_auth_config_from_import,
sanitize_windsurf_batch_import_error, validate_batch_access_token_import,
};
use aether_oauth::core::OAuthError;
use serde_json::json;
@@ -763,6 +809,51 @@ mod tests {
assert!(!detail.contains("secret-token"));
}
#[test]
fn claude_batch_refresh_failure_never_falls_back_to_imported_access_token() {
assert!(!can_fallback_batch_refresh_to_access_token(
"claude_code",
Some("sk-ant-oat01-stale")
));
assert!(can_fallback_batch_refresh_to_access_token(
"codex",
Some("fallback-access-token")
));
}
#[test]
fn claude_batch_access_only_requires_oat_prefix_and_future_expiry() {
let now = 2_000_000_000;
assert!(validate_batch_access_token_import(
"claude_code",
"sk-ant-oat01-valid",
Some(now + 3600),
now,
)
.is_ok());
assert!(validate_batch_access_token_import(
"claude_code",
"not-an-oat",
Some(now + 3600),
now,
)
.is_err());
assert!(validate_batch_access_token_import(
"claude_code",
"sk-ant-oat01-missing-expiry",
None,
now,
)
.is_err());
assert!(validate_batch_access_token_import(
"claude_code",
"sk-ant-oat01-expired",
Some(now),
now,
)
.is_err());
}
#[test]
fn normalizes_codex_agent_identity_import_without_access_token() {
let entries = parse_admin_provider_oauth_batch_import_entries(
@@ -6,6 +6,7 @@ mod progress;
mod task;
pub(super) use orchestration::handle_admin_provider_oauth_batch_import;
pub(super) use parse::build_admin_provider_oauth_batch_task_state;
pub(super) use task::{
handle_admin_provider_oauth_start_agent_identity_import_task,
handle_admin_provider_oauth_start_batch_import_task,
@@ -1,6 +1,6 @@
use super::super::token_import::{
import_tokens_from_raw_token, normalize_provider_import_tokens,
normalize_provider_oauth_import_headers_from_object,
flatten_claude_code_credentials_payload, import_tokens_from_raw_token,
normalize_provider_import_tokens, normalize_provider_oauth_import_headers_from_object,
provider_oauth_import_authorization_bearer_token,
};
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
@@ -46,6 +46,10 @@ pub(super) struct AdminProviderOAuthBatchImportEntry {
pub request_headers: Option<BTreeMap<String, String>>,
pub user_agent: Option<String>,
pub browser_profile: Option<String>,
pub organization_uuid: Option<String>,
pub scopes: Option<serde_json::Value>,
pub subscription_type: Option<String>,
pub rate_limit_tier: Option<String>,
}
#[derive(Debug, Clone)]
@@ -239,10 +243,23 @@ fn extract_admin_provider_oauth_batch_import_entry(
request_headers: None,
user_agent: None,
browser_profile: None,
organization_uuid: None,
scopes: None,
subscription_type: None,
rate_limit_tier: None,
})
}
}
serde_json::Value::Object(object) => {
let is_claude = provider_type.trim().eq_ignore_ascii_case("claude_code");
let normalized_claude_object = if is_claude {
let mut normalized = object.clone();
flatten_claude_code_credentials_payload(&mut normalized);
Some(normalized)
} else {
None
};
let object = normalized_claude_object.as_ref().unwrap_or(object);
let is_grok = provider_type.trim().eq_ignore_ascii_case("grok");
let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf");
let is_codex_agent_identity = provider_type.trim().eq_ignore_ascii_case("codex")
@@ -271,6 +288,10 @@ fn extract_admin_provider_oauth_batch_import_entry(
request_headers: None,
user_agent: None,
browser_profile: None,
organization_uuid: None,
scopes: None,
subscription_type: None,
rate_limit_tier: None,
});
}
let refresh_token = coerce_admin_provider_oauth_import_str(
@@ -463,6 +484,39 @@ fn extract_admin_provider_oauth_batch_import_entry(
.or_else(|| object.get("browser"))
.or_else(|| object.get("impersonate")),
);
let organization_uuid = is_claude
.then(|| {
coerce_admin_provider_oauth_import_str(
object
.get("organization_uuid")
.or_else(|| object.get("organizationUuid"))
.or_else(|| object.get("org_uuid")),
)
})
.flatten();
let scopes = is_claude
.then(|| object.get("scopes"))
.flatten()
.filter(|value| value.is_array() || value.is_string())
.cloned();
let subscription_type = is_claude
.then(|| {
coerce_admin_provider_oauth_import_str(
object
.get("subscription_type")
.or_else(|| object.get("subscriptionType")),
)
})
.flatten();
let rate_limit_tier = is_claude
.then(|| {
coerce_admin_provider_oauth_import_str(
object
.get("rate_limit_tier")
.or_else(|| object.get("rateLimitTier")),
)
})
.flatten();
Some(AdminProviderOAuthBatchImportEntry {
parse_error: None,
refresh_token,
@@ -486,6 +540,10 @@ fn extract_admin_provider_oauth_batch_import_entry(
request_headers,
user_agent,
browser_profile,
organization_uuid,
scopes,
subscription_type,
rate_limit_tier,
})
}
_ => None,
@@ -718,6 +776,10 @@ fn parse_error_entry(error: String) -> AdminProviderOAuthBatchImportEntry {
request_headers: None,
user_agent: None,
browser_profile: None,
organization_uuid: None,
scopes: None,
subscription_type: None,
rate_limit_tier: None,
}
}
@@ -768,6 +830,29 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
}
return;
}
if provider_type == "claude_code" {
if let Some(organization_uuid) = entry.organization_uuid.as_ref() {
auth_config
.entry("org_uuid".to_string())
.or_insert_with(|| json!(organization_uuid));
}
if let Some(scopes) = entry.scopes.as_ref() {
auth_config
.entry("scopes".to_string())
.or_insert_with(|| scopes.clone());
}
if let Some(subscription_type) = entry.subscription_type.as_ref() {
auth_config
.entry("subscription_type".to_string())
.or_insert_with(|| json!(subscription_type));
}
if let Some(rate_limit_tier) = entry.rate_limit_tier.as_ref() {
auth_config
.entry("rate_limit_tier".to_string())
.or_insert_with(|| json!(rate_limit_tier));
}
return;
}
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
return;
}
@@ -878,7 +963,7 @@ pub(super) fn build_admin_provider_oauth_batch_import_response(
}))
}
pub(super) fn build_admin_provider_oauth_batch_task_state(
pub(in super::super) fn build_admin_provider_oauth_batch_task_state(
task_id: &str,
provider_id: &str,
provider_type: &str,
@@ -995,6 +1080,46 @@ mod tests {
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
}
#[test]
fn parses_claude_credentials_json_and_ignores_mcp_oauth() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"claude_code",
r#"{
"claudeAiOauth": {
"accessToken": "sk-ant-oat01-access",
"refreshToken": "sk-ant-ort01-refresh",
"expiresAt": 2100000000123,
"scopes": ["user:profile"],
"subscriptionType": "pro",
"rateLimitTier": "tier_1",
"organizationUuid": "org-123"
},
"mcpOAuth": {"accessToken": "must-not-be-imported"}
}"#,
);
assert_eq!(entries.len(), 1);
assert_eq!(
entries[0].access_token.as_deref(),
Some("sk-ant-oat01-access")
);
assert_eq!(
entries[0].refresh_token.as_deref(),
Some("sk-ant-ort01-refresh")
);
assert_eq!(entries[0].expires_at, Some(2_100_000_000));
assert_eq!(entries[0].organization_uuid.as_deref(), Some("org-123"));
assert_eq!(entries[0].scopes, Some(json!(["user:profile"])));
assert_eq!(entries[0].subscription_type.as_deref(), Some("pro"));
assert_eq!(entries[0].rate_limit_tier.as_deref(), Some("tier_1"));
let ignored = parse_admin_provider_oauth_batch_import_entries(
"claude_code",
r#"{"mcpOAuth":{"accessToken":"must-not-be-imported"}}"#,
);
assert!(ignored.is_empty());
}
#[test]
fn preserves_codex_agent_identity_entry_without_access_token() {
let entries = parse_admin_provider_oauth_batch_import_entries(
@@ -1,17 +1,8 @@
use super::super::super::duplicates::{
acquire_codex_oauth_account_locks, find_duplicate_provider_oauth_key,
release_codex_oauth_account_locks,
};
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::provisioning::{
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
update_existing_provider_oauth_catalog_key,
};
use super::super::super::runtime::{
resolve_provider_oauth_runtime_endpoints,
spawn_provider_oauth_account_state_refresh_after_update,
provider_oauth_key_proxy_value, provision_provider_oauth_token_payload_for_provider,
};
use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints;
use super::super::super::state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
is_fixed_provider_type_for_provider_oauth,
@@ -25,11 +16,8 @@ use crate::GatewayError;
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
response::Response,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_complete_provider(
state: &AdminAppState<'_>,
@@ -151,145 +139,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
Err(response) => return Ok(response),
};
let (auth_config, access_token, refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload);
let Some(access_token) = access_token else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let api_formats = provider_oauth_active_api_formats(&endpoints);
let codex_oauth_account_leases = if provider_type == "codex" {
match acquire_codex_oauth_account_locks(
state,
&provider_id,
&auth_config,
"provider-complete",
)
.await
{
Ok(leases) => leases,
Err(error) => {
return Ok(build_internal_control_error_response(
error.status_code(),
error.detail(),
));
}
}
} else {
Vec::new()
};
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
return Ok(build_internal_control_error_response(
if provider_type == "codex" {
http::StatusCode::CONFLICT
} else {
http::StatusCode::BAD_REQUEST
},
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
let update_result = state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&provider_type,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await;
match update_result {
Err(error) => {
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
return Err(error);
}
Ok(Some(key)) => key,
Ok(None) => {
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let name = payload
.name
.or_else(|| {
auth_config
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.unwrap_or_else(|| {
format!(
"账号_{}",
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
)
});
let create_result = state
.create_provider_oauth_catalog_key(
&provider_id,
&provider_type,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await;
match create_result {
Err(error) => {
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
return Err(error);
}
Ok(Some(key)) => key,
Ok(None) => {
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
release_codex_oauth_account_locks(state, codex_oauth_account_leases).await;
spawn_provider_oauth_account_state_refresh_after_update(
state.cloned_app(),
provider.clone(),
persisted_key.id.clone(),
request_proxy.clone(),
);
Ok(Json(json!({
"key_id": persisted_key.id,
"provider_type": provider_type,
"expires_at": expires_at,
"has_refresh_token": refresh_token.is_some(),
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
"replaced": replaced,
}))
.into_response())
provision_provider_oauth_token_payload_for_provider(
state,
&provider,
&endpoints,
&token_payload,
payload.name,
key_proxy,
request_proxy,
"provider-complete",
)
.await
}
@@ -0,0 +1,214 @@
use super::super::errors::build_internal_control_error_response;
use super::super::provisioning::{
provider_oauth_key_proxy_value, provision_provider_oauth_token_payload_for_provider,
};
use super::super::runtime::resolve_provider_oauth_runtime_endpoints;
use super::super::state::authorize_admin_provider_oauth_with_cookie;
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_cookie_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use axum::{body::Body, http, response::Response};
pub(super) const MAX_CLAUDE_COOKIE_AUTHORIZE_BODY_BYTES: usize = 32 * 1024;
pub(super) const MAX_CLAUDE_SESSION_KEY_BYTES: usize = 16 * 1024;
struct ClaudeCookieAuthorizeRequest {
session_key: String,
name: Option<String>,
proxy_node_id: Option<String>,
}
pub(super) async fn handle_admin_provider_oauth_cookie_authorize(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(super::super::state::build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) = admin_provider_oauth_cookie_provider_id(request_context.path()) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let payload = match parse_claude_cookie_authorize_request(request_body) {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "claude_code" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Cookie 授权仅支持 Claude Code Provider",
));
}
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),
&[
endpoint_resolution
.runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let key_proxy = provider_oauth_key_proxy_value(payload.proxy_node_id.as_deref());
let token_payload = match authorize_admin_provider_oauth_with_cookie(
state,
payload.session_key,
request_proxy.clone(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
provision_provider_oauth_token_payload_for_provider(
state,
&provider,
&endpoints,
&token_payload,
payload.name,
key_proxy,
request_proxy,
"cookie-authorize",
)
.await
}
fn parse_claude_cookie_authorize_request(
request_body: Option<&axum::body::Bytes>,
) -> Result<ClaudeCookieAuthorizeRequest, Response<Body>> {
let Some(request_body) = request_body else {
return Err(bad_cookie_request("请求体必须是合法的 JSON 对象"));
};
if request_body.len() > MAX_CLAUDE_COOKIE_AUTHORIZE_BODY_BYTES {
return Err(bad_cookie_request("Cookie 授权请求体过大"));
}
let payload = serde_json::from_slice::<serde_json::Value>(request_body)
.ok()
.and_then(|value| value.as_object().cloned())
.ok_or_else(|| bad_cookie_request("请求体必须是合法的 JSON 对象"))?;
let cookie = ["cookie", "session_key", "sessionKey"]
.into_iter()
.find_map(|key| payload.get(key).and_then(serde_json::Value::as_str))
.ok_or_else(|| bad_cookie_request("Cookie 不能为空"))?;
let session_key = normalize_claude_session_key(cookie)
.ok_or_else(|| bad_cookie_request("Cookie 中缺少有效的 sessionKey"))?;
Ok(ClaudeCookieAuthorizeRequest {
session_key,
name: optional_trimmed_string(&payload, "name"),
proxy_node_id: optional_trimmed_string(&payload, "proxy_node_id")
.or_else(|| optional_trimmed_string(&payload, "proxyNodeId")),
})
}
pub(super) fn normalize_claude_session_key(raw: &str) -> Option<String> {
let raw = raw.trim();
if raw.is_empty() || raw.len() > MAX_CLAUDE_SESSION_KEY_BYTES || raw.contains(['\r', '\n']) {
return None;
}
let cookie = raw
.split_once(':')
.filter(|(name, _)| name.trim().eq_ignore_ascii_case("cookie"))
.map(|(_, value)| value.trim())
.unwrap_or(raw);
if !cookie.contains('=') {
return valid_session_key_value(cookie).then(|| cookie.to_string());
}
let mut session_key = None;
for segment in cookie.split(';') {
let (name, value) = segment.trim().split_once('=')?;
if !name.trim().eq_ignore_ascii_case("sessionKey") {
continue;
}
if session_key.is_some() || !valid_session_key_value(value.trim()) {
return None;
}
session_key = Some(value.trim().to_string());
}
session_key
}
fn valid_session_key_value(value: &str) -> bool {
!value.is_empty()
&& value.len() <= MAX_CLAUDE_SESSION_KEY_BYTES
&& !value.contains(['\r', '\n', ';'])
&& http::HeaderValue::from_str(value).is_ok()
}
fn optional_trimmed_string(
payload: &serde_json::Map<String, serde_json::Value>,
key: &str,
) -> Option<String> {
payload
.get(key)
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn bad_cookie_request(detail: &'static str) -> Response<Body> {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, detail)
}
#[cfg(test)]
mod tests {
use super::normalize_claude_session_key;
#[test]
fn normalizes_supported_claude_cookie_inputs() {
for (input, expected) in [
("sk-ant-sid01-raw", "sk-ant-sid01-raw"),
("sessionKey=sk-ant-sid01-pair", "sk-ant-sid01-pair"),
(
"Cookie: other=value; sessionKey=sk-ant-sid01-header; theme=dark",
"sk-ant-sid01-header",
),
] {
assert_eq!(
normalize_claude_session_key(input).as_deref(),
Some(expected)
);
}
}
#[test]
fn rejects_ambiguous_or_unsafe_claude_cookie_inputs() {
for input in [
"",
"foo=bar",
"sessionKey=one; sessionKey=two",
"sessionKey=value\r\nx-leak: yes",
] {
assert!(
normalize_claude_session_key(input).is_none(),
"input={input:?}"
);
}
}
}
@@ -0,0 +1,682 @@
use super::super::errors::build_internal_control_error_response;
use super::super::provisioning::{
provider_oauth_key_proxy_value, provision_provider_oauth_token_payload_for_provider,
};
use super::super::runtime::resolve_provider_oauth_runtime_endpoints;
use super::super::state::{
authorize_admin_provider_oauth_with_cookie,
build_admin_provider_oauth_backend_unavailable_response,
};
use super::batch::build_admin_provider_oauth_batch_task_state;
use super::cookie::{normalize_claude_session_key, MAX_CLAUDE_SESSION_KEY_BYTES};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_cookie_task_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::task_runtime::{
append_event_with_logging, now_unix_secs, task_definition, update_run_status,
upsert_run_with_logging, TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT,
};
use crate::GatewayError;
use aether_data_contracts::repository::background_tasks::{
BackgroundTaskKind, BackgroundTaskStatus, UpsertBackgroundTaskRun,
};
use axum::{
body::{to_bytes, Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use futures_util::{stream, StreamExt};
use serde_json::{json, Value};
use std::collections::HashSet;
use std::time::{SystemTime, UNIX_EPOCH};
use tokio::task;
use uuid::Uuid;
const CLAUDE_COOKIE_TASK_IMPORT_KIND: &str = "cookie_authorize";
const CLAUDE_COOKIE_TASK_ID_PREFIX: &str = "claude-cookie-";
const MAX_CLAUDE_COOKIE_TASK_ENTRIES: usize = 20;
const MAX_CLAUDE_COOKIE_TASK_BODY_BYTES: usize = 768 * 1024;
const CLAUDE_COOKIE_AUTHORIZATION_CONCURRENCY: usize = 3;
const MAX_SAFE_ERROR_DETAIL_BYTES: usize = 512;
type ClaudeCookieTaskEntry = Result<String, String>;
struct ClaudeCookieTaskRequest {
entries: Vec<ClaudeCookieTaskEntry>,
proxy_node_id: Option<String>,
}
pub(super) async fn handle_admin_provider_oauth_start_cookie_task(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) = admin_provider_oauth_cookie_task_provider_id(request_context.path())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let payload = match parse_claude_cookie_task_request(request_body) {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "claude_code" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Cookie 授权仅支持 Claude Code Provider",
));
}
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),
&[
endpoint_resolution
.runtime_endpoint
.as_ref()
.and_then(|endpoint| endpoint.proxy.as_ref()),
provider.proxy.as_ref(),
],
)
.await;
let key_proxy = provider_oauth_key_proxy_value(payload.proxy_node_id.as_deref());
let task_id = format!("{CLAUDE_COOKIE_TASK_ID_PREFIX}{}", Uuid::new_v4());
let total = payload.entries.len();
let created_at = now_unix_secs();
let submitted_state = build_admin_provider_oauth_batch_task_state(
&task_id,
&provider_id,
&provider_type,
CLAUDE_COOKIE_TASK_IMPORT_KIND,
"submitted",
total,
0,
0,
0,
0,
0,
Some("任务已提交,等待执行"),
None,
Vec::new(),
created_at,
None,
None,
);
if state
.save_provider_oauth_batch_task_payload(&task_id, &submitted_state)
.await
.is_err()
{
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth batch task redis unavailable",
));
}
if state.has_background_task_data_writer() {
let max_attempts = task_definition(TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT)
.map(|item| item.retry_policy.max_attempts)
.unwrap_or(1);
let run = UpsertBackgroundTaskRun {
id: task_id.clone(),
task_key: TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT.to_string(),
kind: BackgroundTaskKind::OnDemand,
trigger: "manual".to_string(),
status: BackgroundTaskStatus::Queued,
attempt: 1,
max_attempts,
owner_instance: Some(state.app().tunnel.local_instance_id().to_string()),
progress_percent: 0,
progress_message: Some("Claude Cookie authorization queued".to_string()),
payload_json: Some(json!({
"provider_id": provider_id.clone(),
"provider_type": provider_type.clone(),
"import_kind": CLAUDE_COOKIE_TASK_IMPORT_KIND,
"total": total,
})),
result_json: None,
error_message: None,
cancel_requested: false,
created_by: Some("admin".to_string()),
created_at_unix_secs: created_at,
started_at_unix_secs: None,
finished_at_unix_secs: None,
updated_at_unix_secs: created_at,
};
let _ = upsert_run_with_logging(state.app(), run).await;
append_event_with_logging(
state.app(),
&task_id,
"queued",
"Claude Cookie authorization queued",
Some(json!({
"provider_id": provider_id.clone(),
"provider_type": provider_type.clone(),
"import_kind": CLAUDE_COOKIE_TASK_IMPORT_KIND,
"total": total,
})),
)
.await;
}
let task_state = state.cloned_app();
let task_id_for_worker = task_id.clone();
let provider_id_for_worker = provider_id.clone();
let provider_type_for_worker = provider_type.clone();
task::spawn(async move {
let started_at = current_unix_secs_or(created_at);
let task_admin_state = AdminAppState::new(&task_state);
save_cookie_task_state(
&task_admin_state,
&task_id_for_worker,
&provider_id_for_worker,
&provider_type_for_worker,
"processing",
total,
0,
0,
0,
0,
0,
Some("正在获取 Claude 授权"),
Vec::new(),
created_at,
started_at,
None,
)
.await;
let _ = update_run_status(
&task_state,
&task_id_for_worker,
BackgroundTaskStatus::Running,
Some(1),
Some("Claude Cookie authorization started".to_string()),
None,
None,
Some(started_at),
None,
)
.await;
append_event_with_logging(
&task_state,
&task_id_for_worker,
"running",
"Claude Cookie authorization started",
None,
)
.await;
let mut pending = stream::iter(payload.entries.into_iter().enumerate().map(
|(index, entry)| {
let proxy = request_proxy.clone();
let task_admin_state = &task_admin_state;
async move {
let result = match entry {
Ok(session_key) => authorize_admin_provider_oauth_with_cookie(
task_admin_state,
session_key,
proxy,
)
.await
.map_err(|_| "Claude Cookie 授权失败".to_string()),
Err(detail) => Err(detail),
};
(index, result)
}
},
))
.buffer_unordered(CLAUDE_COOKIE_AUTHORIZATION_CONCURRENCY);
let mut authorization_results = Vec::with_capacity(total);
while let Some(result) = pending.next().await {
authorization_results.push(result);
}
authorization_results.sort_by_key(|(index, _)| *index);
let mut success = 0usize;
let mut failed = 0usize;
let mut created_count = 0usize;
let mut replaced_count = 0usize;
let mut error_samples = Vec::new();
for (index, authorization_result) in authorization_results {
let result = match authorization_result {
Ok(token_payload) => match provision_provider_oauth_token_payload_for_provider(
&task_admin_state,
&provider,
&endpoints,
&token_payload,
None,
key_proxy.clone(),
request_proxy.clone(),
"cookie-authorize-batch",
)
.await
{
Ok(response) => cookie_task_item_from_response(index, response).await,
Err(_) => cookie_task_error(index, "provider oauth write unavailable"),
},
Err(detail) => cookie_task_error(index, detail.as_str()),
};
if result.get("status").and_then(Value::as_str) == Some("success") {
success += 1;
if result.get("replaced").and_then(Value::as_bool) == Some(true) {
replaced_count += 1;
} else {
created_count += 1;
}
} else {
failed += 1;
error_samples.push(result);
}
let processed = success.saturating_add(failed);
let message = format!("处理中 {processed}/{total}");
save_cookie_task_state(
&task_admin_state,
&task_id_for_worker,
&provider_id_for_worker,
&provider_type_for_worker,
"processing",
total,
processed,
success,
failed,
created_count,
replaced_count,
Some(message.as_str()),
error_samples.clone(),
created_at,
started_at,
None,
)
.await;
}
let finished_at = current_unix_secs_or(started_at);
let message = format!("授权完成:成功 {success},失败 {failed}");
save_cookie_task_state(
&task_admin_state,
&task_id_for_worker,
&provider_id_for_worker,
&provider_type_for_worker,
"completed",
total,
total,
success,
failed,
created_count,
replaced_count,
Some(message.as_str()),
error_samples,
created_at,
started_at,
Some(finished_at),
)
.await;
let _ = update_run_status(
&task_state,
&task_id_for_worker,
BackgroundTaskStatus::Succeeded,
Some(100),
Some(message),
Some(json!({
"provider_id": provider_id_for_worker,
"provider_type": provider_type_for_worker,
"import_kind": CLAUDE_COOKIE_TASK_IMPORT_KIND,
"total": total,
"success": success,
"failed": failed,
"created_count": created_count,
"replaced_count": replaced_count,
})),
None,
None,
Some(finished_at),
)
.await;
append_event_with_logging(
&task_state,
&task_id_for_worker,
"succeeded",
"Claude Cookie authorization completed",
None,
)
.await;
});
Ok(Json(submitted_state).into_response())
}
#[allow(clippy::too_many_arguments)]
async fn save_cookie_task_state(
state: &AdminAppState<'_>,
task_id: &str,
provider_id: &str,
provider_type: &str,
status: &str,
total: usize,
processed: usize,
success: usize,
failed: usize,
created_count: usize,
replaced_count: usize,
message: Option<&str>,
error_samples: Vec<Value>,
created_at: u64,
started_at: u64,
finished_at: Option<u64>,
) {
let task_state = build_admin_provider_oauth_batch_task_state(
task_id,
provider_id,
provider_type,
CLAUDE_COOKIE_TASK_IMPORT_KIND,
status,
total,
processed,
success,
failed,
created_count,
replaced_count,
message,
None,
error_samples,
created_at,
Some(started_at),
finished_at,
);
let _ = state
.save_provider_oauth_batch_task_payload(task_id, &task_state)
.await;
}
async fn cookie_task_item_from_response(index: usize, response: Response<Body>) -> Value {
let status = response.status();
let body = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES)
.await
.ok();
let payload = body
.as_deref()
.and_then(|body| serde_json::from_slice::<Value>(body).ok());
if status.is_success() {
let Some(payload) = payload else {
return cookie_task_error(index, "provider oauth write unavailable");
};
let Some(key_id) = payload.get("key_id").and_then(Value::as_str) else {
return cookie_task_error(index, "provider oauth write unavailable");
};
return json!({
"index": index,
"status": "success",
"key_id": key_id,
"email": payload.get("email").cloned().unwrap_or(Value::Null),
"replaced": payload.get("replaced").and_then(Value::as_bool).unwrap_or(false),
"error": Value::Null,
});
}
let detail = payload
.as_ref()
.and_then(|payload| payload.get("detail"))
.and_then(Value::as_str)
.and_then(safe_error_detail)
.unwrap_or("Claude 账号创建或更新失败");
cookie_task_error(index, detail)
}
fn safe_error_detail(detail: &str) -> Option<&str> {
let detail = detail.trim();
if detail.is_empty() || detail.len() > MAX_SAFE_ERROR_DETAIL_BYTES {
return None;
}
let normalized = detail.to_ascii_lowercase();
if normalized.contains("sessionkey")
|| normalized.contains("sk-ant-")
|| normalized.contains("cookie:")
{
return None;
}
Some(detail)
}
fn cookie_task_error(index: usize, detail: &str) -> Value {
json!({
"index": index,
"status": "error",
"error": detail,
"replaced": false,
})
}
fn parse_claude_cookie_task_request(
request_body: Option<&Bytes>,
) -> Result<ClaudeCookieTaskRequest, Response<Body>> {
let Some(request_body) = request_body else {
return Err(bad_cookie_task_request("请求体必须是合法的 JSON 对象"));
};
if request_body.len() > MAX_CLAUDE_COOKIE_TASK_BODY_BYTES {
return Err(bad_cookie_task_request("Cookie 授权请求体过大"));
}
let payload = serde_json::from_slice::<Value>(request_body)
.ok()
.and_then(|value| value.as_object().cloned())
.ok_or_else(|| bad_cookie_task_request("请求体必须是合法的 JSON 对象"))?;
let legacy_keys = ["cookie", "session_key", "sessionKey"];
let has_legacy_cookie = legacy_keys.iter().any(|key| payload.contains_key(*key));
let raw_entries = if let Some(cookies) = payload.get("cookies") {
if has_legacy_cookie {
return Err(bad_cookie_task_request("cookie 与 cookies 不能同时提供"));
}
let cookies = cookies
.as_array()
.ok_or_else(|| bad_cookie_task_request("cookies 必须是字符串数组"))?;
if cookies.is_empty() {
return Err(bad_cookie_task_request("Cookie 不能为空"));
}
if cookies.len() > MAX_CLAUDE_COOKIE_TASK_ENTRIES {
return Err(bad_cookie_task_request("Cookie 批量授权最多支持 20 条"));
}
cookies
.iter()
.map(|value| {
value
.as_str()
.map(ToOwned::to_owned)
.ok_or_else(|| bad_cookie_task_request("cookies 必须是字符串数组"))
})
.collect::<Result<Vec<_>, _>>()?
} else {
let raw = legacy_keys
.into_iter()
.find_map(|key| payload.get(key).and_then(Value::as_str))
.ok_or_else(|| bad_cookie_task_request("Cookie 不能为空"))?;
let entries = raw
.lines()
.map(str::trim)
.filter(|line| !line.is_empty())
.map(ToOwned::to_owned)
.collect::<Vec<_>>();
if entries.is_empty() {
return Err(bad_cookie_task_request("Cookie 不能为空"));
}
if entries.len() > MAX_CLAUDE_COOKIE_TASK_ENTRIES {
return Err(bad_cookie_task_request("Cookie 批量授权最多支持 20 条"));
}
entries
};
let mut seen_session_keys = HashSet::new();
let entries = raw_entries
.into_iter()
.map(|raw| {
let session_key =
normalize_claude_session_key(&raw).ok_or_else(|| "Cookie 格式无效".to_string())?;
if !seen_session_keys.insert(session_key.clone()) {
return Err("Cookie 重复".to_string());
}
Ok(session_key)
})
.collect();
Ok(ClaudeCookieTaskRequest {
entries,
proxy_node_id: optional_trimmed_string(&payload, "proxy_node_id")
.or_else(|| optional_trimmed_string(&payload, "proxyNodeId")),
})
}
fn optional_trimmed_string(payload: &serde_json::Map<String, Value>, key: &str) -> Option<String> {
payload
.get(key)
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn bad_cookie_task_request(detail: &'static str) -> Response<Body> {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, detail)
}
fn current_unix_secs_or(fallback: u64) -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(fallback)
}
#[cfg(test)]
mod tests {
use super::{
parse_claude_cookie_task_request, safe_error_detail, MAX_CLAUDE_COOKIE_TASK_BODY_BYTES,
MAX_CLAUDE_COOKIE_TASK_ENTRIES, MAX_CLAUDE_SESSION_KEY_BYTES,
};
use axum::body::{to_bytes, Bytes};
use serde_json::json;
#[test]
fn parses_canonical_and_multiline_cookie_batches_without_retaining_raw_headers() {
for payload in [
json!({
"cookies": [
"sessionKey=sk-ant-sid01-one",
"Cookie: theme=dark; sessionKey=sk-ant-sid01-two"
],
"proxy_node_id": "proxy-1"
}),
json!({
"cookie": "sessionKey=sk-ant-sid01-one\n\nCookie: sessionKey=sk-ant-sid01-two",
"proxyNodeId": "proxy-1"
}),
] {
let body = Bytes::from(payload.to_string());
let parsed =
parse_claude_cookie_task_request(Some(&body)).expect("cookie batch should parse");
assert_eq!(parsed.entries.len(), 2);
assert_eq!(parsed.entries[0].as_deref(), Ok("sk-ant-sid01-one"));
assert_eq!(parsed.entries[1].as_deref(), Ok("sk-ant-sid01-two"));
assert_eq!(parsed.proxy_node_id.as_deref(), Some("proxy-1"));
}
}
#[test]
fn keeps_invalid_cookie_lines_as_independent_sanitized_results() {
let body = Bytes::from(
json!({
"cookies": ["foo=bar", "sessionKey=valid", "Cookie: sessionKey=valid"]
})
.to_string(),
);
let parsed =
parse_claude_cookie_task_request(Some(&body)).expect("request should be accepted");
assert_eq!(parsed.entries.len(), 3);
assert_eq!(
parsed.entries[0].as_ref().expect_err("entry should fail"),
"Cookie 格式无效"
);
assert_eq!(parsed.entries[1].as_deref(), Ok("valid"));
assert_eq!(
parsed.entries[2].as_ref().expect_err("entry should fail"),
"Cookie 重复"
);
}
#[test]
fn accepts_twenty_maximum_length_session_keys_within_batch_body_limit() {
let cookies = (0..MAX_CLAUDE_COOKIE_TASK_ENTRIES)
.map(|index| {
let prefix = format!("{index:02}-");
format!(
"{prefix}{}",
"x".repeat(MAX_CLAUDE_SESSION_KEY_BYTES - prefix.len())
)
})
.collect::<Vec<_>>();
let body = Bytes::from(json!({"cookies": cookies}).to_string());
assert!(body.len() < MAX_CLAUDE_COOKIE_TASK_BODY_BYTES);
let parsed = parse_claude_cookie_task_request(Some(&body))
.expect("maximum valid batch should parse");
assert_eq!(parsed.entries.len(), MAX_CLAUDE_COOKIE_TASK_ENTRIES);
assert!(parsed.entries.iter().all(Result::is_ok));
}
#[tokio::test]
async fn rejects_ambiguous_or_oversized_cookie_batches_without_echoing_secrets() {
let too_many = vec!["sessionKey=value"; MAX_CLAUDE_COOKIE_TASK_ENTRIES + 1];
for payload in [
json!({"cookie": "sessionKey=secret", "cookies": ["sessionKey=other"]}),
json!({"cookies": too_many}),
json!({"cookies": []}),
] {
let body = Bytes::from(payload.to_string());
let response = match parse_claude_cookie_task_request(Some(&body)) {
Ok(_) => panic!("request should fail"),
Err(response) => response,
};
let response_body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body should read");
let text = String::from_utf8_lossy(&response_body);
assert!(!text.contains("secret"));
assert!(!text.contains("other"));
}
let oversized_body = Bytes::from(vec![b'x'; MAX_CLAUDE_COOKIE_TASK_BODY_BYTES + 1]);
let response = match parse_claude_cookie_task_request(Some(&oversized_body)) {
Ok(_) => panic!("oversized request should fail"),
Err(response) => response,
};
assert_eq!(response.status(), http::StatusCode::BAD_REQUEST);
}
#[test]
fn error_detail_filter_rejects_possible_cookie_or_token_leaks() {
assert_eq!(safe_error_detail("账号重复"), Some("账号重复"));
assert!(safe_error_detail("sessionKey=secret").is_none());
assert!(safe_error_detail("upstream sk-ant-oat01-secret").is_none());
assert!(safe_error_detail("Cookie: secret").is_none());
}
}
@@ -20,9 +20,10 @@ use super::super::state::{
use super::helpers::admin_provider_oauth_key_name_from_auth_config;
use super::token_import::{
build_provider_access_token_import_auth_config, decode_access_token_expires_at,
flatten_claude_code_credentials_payload, is_claude_session_key,
normalize_provider_import_tokens, normalize_provider_oauth_import_headers_from_object,
provider_oauth_import_authorization_bearer_token_from_object,
provider_type_supports_access_token_import,
provider_type_supports_access_token_import, validate_claude_access_token_import,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_import_provider_id;
use crate::handlers::admin::request::{
@@ -482,6 +483,29 @@ fn apply_single_import_hints(
}
return;
}
if provider_type == "claude_code" {
for (target, keys) in [
(
"org_uuid",
&["org_uuid", "organization_uuid", "organizationUuid"][..],
),
(
"subscription_type",
&["subscription_type", "subscriptionType"][..],
),
("rate_limit_tier", &["rate_limit_tier", "rateLimitTier"][..]),
] {
if let Some(value) = import_payload_string_any(payload, keys) {
auth_config
.entry(target.to_string())
.or_insert_with(|| json!(value));
}
}
if let Some(scopes) = payload.get("scopes").cloned() {
auth_config.entry("scopes".to_string()).or_insert(scopes);
}
return;
}
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
return;
}
@@ -612,7 +636,9 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
{
Ok(payload) => payload,
Err(response) => {
if provider_type_supports_access_token_import(provider_type) {
if !provider_type.eq_ignore_ascii_case("claude_code")
&& provider_type_supports_access_token_import(provider_type)
{
if let Some(access_token) = access_token
.map(str::trim)
.filter(|value| !value.is_empty())
@@ -666,10 +692,25 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
"Refresh Token 或 Access Token 不能为空",
));
};
if provider_type.eq_ignore_ascii_case("claude_code") {
let now_unix_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
if let Err(detail) =
validate_claude_access_token_import(access_token, imported_expires_at, now_unix_secs)
{
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
}
if !provider_type_supports_access_token_import(provider_type) {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider",
"Access Token 导入仅支持 Claude Code / Codex / ChatGPT Web / Grok Provider",
));
}
@@ -771,7 +812,7 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
let mut raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Ok(build_internal_control_error_response(
@@ -792,21 +833,6 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
} else {
None
};
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
let access_token_input = import_payload_string_any(
&raw_payload,
&[
"access_token",
"accessToken",
"sso_token",
"ssoToken",
"session_token",
"sessionToken",
],
)
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
let imported_expires_at =
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
let name = raw_payload
.get("name")
.and_then(serde_json::Value::as_str)
@@ -832,11 +858,41 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type == "claude_code" {
flatten_claude_code_credentials_payload(&mut raw_payload);
}
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
let access_token_input = import_payload_string_any(
&raw_payload,
&[
"access_token",
"accessToken",
"sso_token",
"ssoToken",
"session_token",
"sessionToken",
],
)
.or_else(|| provider_oauth_import_authorization_bearer_token_from_object(&raw_payload));
let imported_expires_at =
import_payload_u64_any(&raw_payload, &["expires_at", "expiresAt", "expired"]);
let (refresh_token_input, access_token_input) = normalize_provider_import_tokens(
&provider_type,
refresh_token_input.as_deref(),
access_token_input.as_deref(),
);
if provider_type == "claude_code"
&& refresh_token_input
.as_deref()
.into_iter()
.chain(access_token_input.as_deref())
.any(is_claude_session_key)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Claude sessionKey 请使用 Cookie 授权,不能作为导入凭据",
));
}
if !create_agent_identity && refresh_token_input.is_none() && access_token_input.is_none() {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
@@ -6,9 +6,11 @@ use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_agent_identity_import_task_provider_id,
admin_provider_oauth_batch_import_provider_id,
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id,
admin_provider_oauth_import_provider_id, admin_provider_oauth_refresh_key_id,
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
admin_provider_oauth_complete_provider_id, admin_provider_oauth_cookie_provider_id,
admin_provider_oauth_cookie_task_provider_id,
admin_provider_oauth_device_authorize_provider_id, admin_provider_oauth_import_provider_id,
admin_provider_oauth_refresh_key_id, admin_provider_oauth_start_key_id,
admin_provider_oauth_start_provider_id,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
@@ -21,6 +23,8 @@ use axum::{
mod batch;
mod complete;
mod cookie;
mod cookie_task;
mod device;
mod helpers;
mod import;
@@ -94,6 +98,12 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
));
}
if route_kind == Some("get_cookie_authorize_task_status") && *method == http::Method::GET {
return Ok(Some(
tasks::handle_admin_provider_oauth_cookie_task_status(state, request_context).await?,
));
}
if route_kind == Some("complete_key_oauth") && *method == http::Method::POST {
let response = complete::handle_admin_provider_oauth_complete_key(
state,
@@ -156,6 +166,38 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
)));
}
if route_kind == Some("cookie_authorize") && *method == http::Method::POST {
let response = cookie::handle_admin_provider_oauth_cookie_authorize(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(helpers::attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_cookie_authorized",
"authorize_provider_oauth_with_cookie",
"provider",
admin_provider_oauth_cookie_provider_id(request_context.path()),
)));
}
if route_kind == Some("start_cookie_authorize_task") && *method == http::Method::POST {
let response = cookie_task::handle_admin_provider_oauth_start_cookie_task(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(helpers::attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_cookie_task_started",
"start_provider_oauth_cookie_task",
"provider",
admin_provider_oauth_cookie_task_provider_id(request_context.path()),
)));
}
if route_kind == Some("batch_import_oauth") && *method == http::Method::POST {
let response =
batch::handle_admin_provider_oauth_batch_import(state, request_context, request_body)
@@ -226,7 +268,12 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
if matches!(
route_kind,
Some("refresh_key_oauth" | "import_refresh_token")
Some(
"refresh_key_oauth"
| "import_refresh_token"
| "cookie_authorize"
| "start_cookie_authorize_task",
)
) {
return Ok(Some(
build_admin_provider_oauth_backend_unavailable_response(),
@@ -1,7 +1,7 @@
use super::super::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_agent_identity_import_task_path,
admin_provider_oauth_batch_import_task_path,
admin_provider_oauth_batch_import_task_path, admin_provider_oauth_cookie_task_path,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
@@ -14,20 +14,41 @@ use axum::{
};
const PROVIDER_AGENT_IDENTITY_IMPORT_KIND: &str = "agent_identity";
const PROVIDER_OAUTH_BATCH_IMPORT_KIND: &str = "oauth_batch";
const PROVIDER_COOKIE_AUTHORIZE_IMPORT_KIND: &str = "cookie_authorize";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum ProviderOAuthTaskRouteKind {
BatchImport,
AgentIdentity,
CookieAuthorize,
}
fn provider_oauth_import_task_matches_route(
task_id: &str,
payload: &serde_json::Value,
agent_identity_only: bool,
route_kind: ProviderOAuthTaskRouteKind,
) -> bool {
let has_agent_prefix = task_id.starts_with("agent-identity-");
let has_cookie_prefix = task_id.starts_with("claude-cookie-");
let import_kind = payload
.get("import_kind")
.and_then(serde_json::Value::as_str);
if agent_identity_only {
has_agent_prefix && import_kind == Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND)
} else {
!has_agent_prefix && import_kind != Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND)
match route_kind {
ProviderOAuthTaskRouteKind::BatchImport => {
!has_agent_prefix
&& !has_cookie_prefix
&& matches!(
import_kind,
None | Some("") | Some(PROVIDER_OAUTH_BATCH_IMPORT_KIND)
)
}
ProviderOAuthTaskRouteKind::AgentIdentity => {
has_agent_prefix && import_kind == Some(PROVIDER_AGENT_IDENTITY_IMPORT_KIND)
}
ProviderOAuthTaskRouteKind::CookieAuthorize => {
has_cookie_prefix && import_kind == Some(PROVIDER_COOKIE_AUTHORIZE_IMPORT_KIND)
}
}
}
@@ -35,30 +56,63 @@ pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
handle_admin_provider_oauth_import_task_status(state, request_context, false).await
handle_admin_provider_oauth_import_task_status(
state,
request_context,
ProviderOAuthTaskRouteKind::BatchImport,
)
.await
}
pub(super) async fn handle_admin_provider_oauth_agent_identity_import_task_status(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
handle_admin_provider_oauth_import_task_status(state, request_context, true).await
handle_admin_provider_oauth_import_task_status(
state,
request_context,
ProviderOAuthTaskRouteKind::AgentIdentity,
)
.await
}
pub(super) async fn handle_admin_provider_oauth_cookie_task_status(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
handle_admin_provider_oauth_import_task_status(
state,
request_context,
ProviderOAuthTaskRouteKind::CookieAuthorize,
)
.await
}
async fn handle_admin_provider_oauth_import_task_status(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
agent_identity_only: bool,
route_kind: ProviderOAuthTaskRouteKind,
) -> Result<Response<Body>, GatewayError> {
let task_path = if agent_identity_only {
admin_provider_oauth_agent_identity_import_task_path(request_context.path())
} else {
admin_provider_oauth_batch_import_task_path(request_context.path())
let not_found_detail = match route_kind {
ProviderOAuthTaskRouteKind::BatchImport => "批量导入任务不存在或已过期",
ProviderOAuthTaskRouteKind::AgentIdentity => "Agent Identity 导入任务不存在或已过期",
ProviderOAuthTaskRouteKind::CookieAuthorize => "Cookie 授权任务不存在或已过期",
};
let task_path = match route_kind {
ProviderOAuthTaskRouteKind::BatchImport => {
admin_provider_oauth_batch_import_task_path(request_context.path())
}
ProviderOAuthTaskRouteKind::AgentIdentity => {
admin_provider_oauth_agent_identity_import_task_path(request_context.path())
}
ProviderOAuthTaskRouteKind::CookieAuthorize => {
admin_provider_oauth_cookie_task_path(request_context.path())
}
};
let Some((provider_id, task_id)) = task_path else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"批量导入任务不存在",
not_found_detail,
));
};
let payload = match state
@@ -69,7 +123,7 @@ async fn handle_admin_provider_oauth_import_task_status(
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"批量导入任务不存在或已过期",
not_found_detail,
));
}
Err(_) => {
@@ -79,10 +133,10 @@ async fn handle_admin_provider_oauth_import_task_status(
));
}
};
if !provider_oauth_import_task_matches_route(&task_id, &payload, agent_identity_only) {
if !provider_oauth_import_task_matches_route(&task_id, &payload, route_kind) {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"导入任务不存在或已过期",
not_found_detail,
));
}
let status = payload
@@ -91,20 +145,25 @@ async fn handle_admin_provider_oauth_import_task_status(
.map(ToOwned::to_owned)
.unwrap_or_default();
let response = Json(payload).into_response();
let (completed_event, failed_event, action, target_type) = if agent_identity_only {
(
let (completed_event, failed_event, action, target_type) = match route_kind {
ProviderOAuthTaskRouteKind::AgentIdentity => (
"admin_provider_oauth_agent_identity_import_completed_viewed",
"admin_provider_oauth_agent_identity_import_failed_viewed",
"view_provider_agent_identity_import_terminal_state",
"provider_agent_identity_import_task",
)
} else {
(
),
ProviderOAuthTaskRouteKind::CookieAuthorize => (
"admin_provider_oauth_cookie_task_completed_viewed",
"admin_provider_oauth_cookie_task_failed_viewed",
"view_provider_oauth_cookie_task_terminal_state",
"provider_oauth_cookie_task",
),
ProviderOAuthTaskRouteKind::BatchImport => (
"admin_provider_oauth_batch_task_completed_viewed",
"admin_provider_oauth_batch_task_failed_viewed",
"view_provider_oauth_batch_task_terminal_state",
"provider_oauth_batch_task",
)
),
};
Ok(match status.as_str() {
"completed" => attach_admin_audit_response(
@@ -127,38 +186,69 @@ async fn handle_admin_provider_oauth_import_task_status(
#[cfg(test)]
mod tests {
use super::provider_oauth_import_task_matches_route;
use super::{provider_oauth_import_task_matches_route, ProviderOAuthTaskRouteKind};
use serde_json::json;
#[test]
fn import_task_status_routes_are_bidirectionally_isolated() {
let agent_payload = json!({ "import_kind": "agent_identity" });
let batch_payload = json!({ "import_kind": "oauth_batch" });
let cookie_payload = json!({ "import_kind": "cookie_authorize" });
assert!(provider_oauth_import_task_matches_route(
"agent-identity-task-1",
&agent_payload,
true,
ProviderOAuthTaskRouteKind::AgentIdentity,
));
assert!(!provider_oauth_import_task_matches_route(
"agent-identity-task-1",
&agent_payload,
false,
ProviderOAuthTaskRouteKind::BatchImport,
));
assert!(provider_oauth_import_task_matches_route(
"batch-task-1",
&batch_payload,
false,
ProviderOAuthTaskRouteKind::BatchImport,
));
assert!(!provider_oauth_import_task_matches_route(
"batch-task-1",
&batch_payload,
true,
ProviderOAuthTaskRouteKind::AgentIdentity,
));
assert!(provider_oauth_import_task_matches_route(
"claude-cookie-task-1",
&cookie_payload,
ProviderOAuthTaskRouteKind::CookieAuthorize,
));
for route_kind in [
ProviderOAuthTaskRouteKind::BatchImport,
ProviderOAuthTaskRouteKind::AgentIdentity,
] {
assert!(!provider_oauth_import_task_matches_route(
"claude-cookie-task-1",
&cookie_payload,
route_kind,
));
}
for (task_id, payload) in [
("batch-task-1", &batch_payload),
("agent-identity-task-1", &agent_payload),
] {
assert!(!provider_oauth_import_task_matches_route(
task_id,
payload,
ProviderOAuthTaskRouteKind::CookieAuthorize,
));
}
assert!(!provider_oauth_import_task_matches_route(
"claude-cookie-task-1",
&batch_payload,
ProviderOAuthTaskRouteKind::CookieAuthorize,
));
assert!(provider_oauth_import_task_matches_route(
"legacy-batch-task",
&json!({}),
false,
ProviderOAuthTaskRouteKind::BatchImport,
));
}
}
@@ -100,6 +100,12 @@ pub(super) fn normalize_provider_import_tokens(
if provider_type == "grok" {
return (None, access_token.or(refresh_token));
}
if provider_type == "claude_code" {
if access_token.is_none() && refresh_token.as_deref().is_some_and(is_claude_access_token) {
return (None, refresh_token);
}
return (refresh_token, access_token);
}
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref())
}
@@ -210,10 +216,75 @@ pub(super) fn provider_oauth_import_authorization_bearer_token_from_object(
pub(super) fn provider_type_supports_access_token_import(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"codex" | "chatgpt_web" | "grok"
"claude_code" | "codex" | "chatgpt_web" | "grok"
)
}
pub(super) fn is_claude_access_token(value: &str) -> bool {
value.trim().starts_with("sk-ant-oat")
}
pub(super) fn is_claude_session_key(value: &str) -> bool {
value.trim().starts_with("sk-ant-sid")
}
pub(super) fn flatten_claude_code_credentials_payload(payload: &mut Map<String, Value>) {
let nested = payload
.get("claudeAiOauth")
.or_else(|| payload.get("claude_ai_oauth"))
.and_then(Value::as_object)
.cloned();
let Some(nested) = nested else {
return;
};
for (target, aliases) in [
("access_token", &["access_token", "accessToken"][..]),
("refresh_token", &["refresh_token", "refreshToken"][..]),
("scopes", &["scopes"][..]),
(
"subscription_type",
&["subscription_type", "subscriptionType"][..],
),
("rate_limit_tier", &["rate_limit_tier", "rateLimitTier"][..]),
(
"organization_uuid",
&["organization_uuid", "organizationUuid", "org_uuid"][..],
),
] {
if payload.contains_key(target) {
continue;
}
if let Some(value) = aliases.iter().find_map(|key| nested.get(*key)).cloned() {
payload.insert(target.to_string(), value);
}
}
if !payload.contains_key("expires_at") {
if let Some(expires_at) = json_u64_value(nested.get("expires_at")) {
payload.insert("expires_at".to_string(), json!(expires_at));
} else if let Some(expires_at_ms) = json_u64_value(nested.get("expiresAt")) {
payload.insert("expires_at".to_string(), json!(expires_at_ms / 1_000));
}
}
}
pub(super) fn validate_claude_access_token_import(
access_token: &str,
imported_expires_at: Option<u64>,
now_unix_secs: u64,
) -> Result<(), &'static str> {
if !is_claude_access_token(access_token) {
return Err("Claude Access Token 格式无效,请导入 sk-ant-oat 凭据");
}
if imported_expires_at.is_none_or(|expires_at| expires_at <= now_unix_secs) {
return Err(
"Claude Access Token 单独导入必须提供有效的未来 expires_at;建议导入完整 Claude credentials 或 Refresh Token",
);
}
Ok(())
}
pub(super) fn build_provider_access_token_import_auth_config(
provider_type: &str,
access_token: &str,
@@ -266,9 +337,10 @@ pub(super) fn build_provider_access_token_import_auth_config(
mod tests {
use super::{
build_provider_access_token_import_auth_config, decode_access_token_expires_at,
looks_like_access_token, normalize_provider_import_tokens,
normalize_provider_oauth_import_headers, normalize_single_import_tokens,
provider_oauth_import_authorization_bearer_token,
flatten_claude_code_credentials_payload, looks_like_access_token,
normalize_provider_import_tokens, normalize_provider_oauth_import_headers,
normalize_single_import_tokens, provider_oauth_import_authorization_bearer_token,
validate_claude_access_token_import,
};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
@@ -431,4 +503,64 @@ mod tests {
Some(&json!(2_200_000_000u64))
);
}
#[test]
fn flattens_only_claude_ai_oauth_credentials_and_converts_expiry_ms() {
let mut payload = json!({
"claudeAiOauth": {
"accessToken": "sk-ant-oat01-access",
"refreshToken": "sk-ant-ort01-refresh",
"expiresAt": 2_100_000_000_123u64,
"scopes": ["user:profile"],
"subscriptionType": "pro",
"rateLimitTier": "tier_1",
"organizationUuid": "org-123"
},
"mcpOAuth": {
"accessToken": "must-not-be-imported"
}
})
.as_object()
.cloned()
.expect("payload should be an object");
flatten_claude_code_credentials_payload(&mut payload);
assert_eq!(
payload.get("access_token"),
Some(&json!("sk-ant-oat01-access"))
);
assert_eq!(
payload.get("refresh_token"),
Some(&json!("sk-ant-ort01-refresh"))
);
assert_eq!(payload.get("expires_at"), Some(&json!(2_100_000_000u64)));
assert_eq!(payload.get("organization_uuid"), Some(&json!("org-123")));
assert_ne!(
payload.get("access_token"),
Some(&json!("must-not-be-imported"))
);
}
#[test]
fn validates_claude_access_token_prefix_and_future_expiry() {
assert!(validate_claude_access_token_import(
"sk-ant-oat01-access",
Some(2_100_000_000),
2_000_000_000,
)
.is_ok());
assert!(validate_claude_access_token_import(
"arbitrary-token",
Some(2_100_000_000),
2_000_000_000,
)
.is_err());
assert!(validate_claude_access_token_import(
"sk-ant-oat01-expired",
Some(1_900_000_000),
2_000_000_000,
)
.is_err());
}
}
@@ -8,6 +8,7 @@ use std::time::{Duration, SystemTime, UNIX_EPOCH};
use uuid::Uuid;
const CODEX_OAUTH_ACCOUNT_LOCK_TTL: Duration = Duration::from_secs(180);
const CLAUDE_OAUTH_ACCOUNT_LOCK_TTL: Duration = Duration::from_secs(180);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum CodexOAuthAccountLockError {
@@ -34,6 +35,31 @@ impl CodexOAuthAccountLockError {
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ClaudeOAuthAccountLockError {
MissingIdentity,
Contended,
Unavailable,
}
impl ClaudeOAuthAccountLockError {
pub(crate) const fn status_code(self) -> http::StatusCode {
match self {
Self::MissingIdentity => http::StatusCode::BAD_REQUEST,
Self::Contended => http::StatusCode::CONFLICT,
Self::Unavailable => http::StatusCode::SERVICE_UNAVAILABLE,
}
}
pub(crate) const fn detail(self) -> &'static str {
match self {
Self::MissingIdentity => "Claude 账号身份字段缺失,无法安全写入授权",
Self::Contended => "该 Claude 账号正在更新授权,请稍后重试",
Self::Unavailable => "Claude 账号授权锁暂不可用,请稍后重试",
}
}
}
fn normalize_codex_plan_group_for_provider_oauth(
plan_type: Option<&serde_json::Value>,
) -> Option<String> {
@@ -208,23 +234,96 @@ pub(crate) async fn acquire_codex_oauth_account_locks(
pub(crate) async fn release_codex_oauth_account_locks(
state: &AdminAppState<'_>,
leases: Vec<RuntimeLockLease>,
) {
release_provider_oauth_account_locks(state, leases).await;
}
pub(crate) async fn release_provider_oauth_account_locks(
state: &AdminAppState<'_>,
leases: Vec<RuntimeLockLease>,
) {
for lease in leases.into_iter().rev() {
match state.runtime_state().lock_release(&lease).await {
Ok(true) => {}
Ok(false) => tracing::warn!(
lock_key = %lease.key,
"gateway Codex OAuth account lock was not owned during release"
"gateway provider OAuth account lock was not owned during release"
),
Err(error) => tracing::warn!(
lock_key = %lease.key,
error = ?error,
"gateway Codex OAuth account lock release failed"
"gateway provider OAuth account lock release failed"
),
}
}
}
fn claude_oauth_account_lock_key(
provider_id: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
) -> Option<String> {
let (identity_kind, identity) = normalize_provider_oauth_identity_value_from_keys(
auth_config,
&["account_uuid", "accountUuid"],
)
.map(|value| ("account_uuid", value))
.or_else(|| {
normalize_provider_oauth_identity_value_from_keys(auth_config, &["email"])
.map(|value| ("email", value.to_ascii_lowercase()))
})?;
let mut digest = Sha256::new();
digest.update(provider_id.trim().as_bytes());
digest.update([0]);
digest.update(identity_kind.as_bytes());
digest.update([0]);
digest.update(identity.as_bytes());
Some(format!(
"provider_oauth_claude_account:{:x}",
digest.finalize()
))
}
pub(crate) async fn acquire_claude_oauth_account_lock(
state: &AdminAppState<'_>,
provider_id: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
operation: &str,
) -> Result<Vec<RuntimeLockLease>, ClaudeOAuthAccountLockError> {
let Some(lock_key) = claude_oauth_account_lock_key(provider_id, auth_config) else {
return Err(ClaudeOAuthAccountLockError::MissingIdentity);
};
let owner = format!(
"aether-gateway-claude-oauth-{}-{}",
operation.trim(),
Uuid::new_v4()
);
let lease = match state
.runtime_state()
.lock_try_acquire(
lock_key.as_str(),
owner.as_str(),
CLAUDE_OAUTH_ACCOUNT_LOCK_TTL,
)
.await
{
Ok(Some(lease)) => lease,
Ok(None) => return Err(ClaudeOAuthAccountLockError::Contended),
Err(error) => {
tracing::warn!(
provider_id = %provider_id,
lock_key = %lock_key,
operation,
error = ?error,
"gateway Claude OAuth account lock unavailable"
);
return Err(ClaudeOAuthAccountLockError::Unavailable);
}
};
state.app().data.clear_provider_catalog_cache();
Ok(vec![lease])
}
fn is_openai_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool {
value
.and_then(serde_json::Value::as_str)
@@ -242,6 +341,37 @@ fn is_windsurf_provider_oauth_provider_type(value: Option<&serde_json::Value>) -
.is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("windsurf"))
}
fn is_claude_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("claude_code"))
}
fn match_claude_provider_oauth_identity(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> Option<bool> {
if !is_claude_provider_oauth_provider_type(new_auth_config.get("provider_type"))
&& !is_claude_provider_oauth_provider_type(existing_auth_config.get("provider_type"))
{
return None;
}
let new_account_uuid = normalize_provider_oauth_identity_value_from_keys(
new_auth_config,
&["account_uuid", "accountUuid"],
);
let existing_account_uuid = normalize_provider_oauth_identity_value_from_keys(
existing_auth_config,
&["account_uuid", "accountUuid"],
);
match (new_account_uuid, existing_account_uuid) {
(Some(left), Some(right)) => Some(left == right),
_ => None,
}
}
fn match_codex_provider_oauth_identity(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
@@ -429,6 +559,10 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
let new_email = normalize_provider_oauth_identity_value(auth_config.get("email"));
let new_user_id = normalize_provider_oauth_identity_value(auth_config.get("user_id"));
let new_account_id = normalize_provider_oauth_identity_value(auth_config.get("account_id"));
let new_account_uuid = normalize_provider_oauth_identity_value_from_keys(
auth_config,
&["account_uuid", "accountUuid"],
);
let new_agent_runtime_id = normalize_provider_oauth_identity_value(
auth_config
.get("agent_runtime_id")
@@ -442,6 +576,7 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
if new_email.is_none()
&& new_user_id.is_none()
&& new_account_id.is_none()
&& new_account_uuid.is_none()
&& new_agent_runtime_id.is_none()
&& new_credential_fingerprint.is_none()
{
@@ -486,17 +621,22 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
.is_some_and(|value| value.eq_ignore_ascii_case("windsurf"));
let mut is_duplicate = false;
let claude_identity_match =
match_claude_provider_oauth_identity(auth_config, &existing_auth_config);
let codex_identity_match =
match_codex_provider_oauth_identity(auth_config, &existing_auth_config);
let windsurf_identity_match =
match_windsurf_provider_oauth_identity(auth_config, &existing_auth_config);
if let Some(codex_identity_match) = codex_identity_match {
if let Some(claude_identity_match) = claude_identity_match {
is_duplicate = claude_identity_match;
} else if let Some(codex_identity_match) = codex_identity_match {
is_duplicate = codex_identity_match;
} else if let Some(windsurf_identity_match) = windsurf_identity_match {
is_duplicate = windsurf_identity_match;
}
if codex_identity_match.is_none()
if claude_identity_match.is_none()
&& codex_identity_match.is_none()
&& windsurf_identity_match.is_none()
&& !is_duplicate
&& new_user_id.is_some()
@@ -507,7 +647,8 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
is_duplicate = true;
}
if codex_identity_match.is_none()
if claude_identity_match.is_none()
&& codex_identity_match.is_none()
&& windsurf_identity_match.is_none()
&& !is_duplicate
&& !is_windsurf
@@ -548,19 +689,20 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
if existing_provider_oauth_key_is_replaceable(&existing_key) {
return Ok(Some(existing_key));
}
let identifier =
normalize_provider_oauth_identity_value(auth_config.get("account_user_id"))
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id")))
.or_else(|| new_agent_runtime_id.clone())
.or_else(|| {
normalize_provider_oauth_identity_value(
auth_config.get("credential_fingerprint"),
)
.map(|value| format!("fingerprint:{value}"))
})
.or_else(|| new_email.clone())
.or_else(|| new_user_id.clone())
.unwrap_or_default();
let identifier = normalize_provider_oauth_identity_value_from_keys(
auth_config,
&["account_uuid", "accountUuid"],
)
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_user_id")))
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id")))
.or_else(|| new_agent_runtime_id.clone())
.or_else(|| {
normalize_provider_oauth_identity_value(auth_config.get("credential_fingerprint"))
.map(|value| format!("fingerprint:{value}"))
})
.or_else(|| new_email.clone())
.or_else(|| new_user_id.clone())
.unwrap_or_default();
return Err(format!(
"该 OAuth 账号 ({identifier}) 已存在于当前 Provider 中(名称: {})",
existing_key.name
@@ -573,9 +715,12 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
#[cfg(test)]
mod tests {
use super::{
acquire_codex_oauth_account_locks, codex_agent_identity_account_lock_keys,
match_codex_provider_oauth_identity, match_windsurf_provider_oauth_identity,
release_codex_oauth_account_locks, CodexOAuthAccountLockError,
acquire_claude_oauth_account_lock, acquire_codex_oauth_account_locks,
claude_oauth_account_lock_key, codex_agent_identity_account_lock_keys,
match_claude_provider_oauth_identity, match_codex_provider_oauth_identity,
match_windsurf_provider_oauth_identity, release_codex_oauth_account_locks,
release_provider_oauth_account_locks, ClaudeOAuthAccountLockError,
CodexOAuthAccountLockError,
};
use crate::handlers::admin::request::AdminAppState;
use crate::AppState;
@@ -604,6 +749,99 @@ mod tests {
);
}
#[test]
fn claude_identity_prefers_account_uuid_without_using_organization_uuid() {
let new_auth_config = auth_config(json!({
"provider_type": "claude_code",
"account_uuid": "account-1",
"org_uuid": "shared-org"
}));
let same_account = auth_config(json!({
"provider_type": "claude_code",
"accountUuid": "account-1",
"org_uuid": "other-org"
}));
let different_account = auth_config(json!({
"provider_type": "claude_code",
"account_uuid": "account-2",
"email": "[email protected]",
"org_uuid": "shared-org"
}));
let organization_only = auth_config(json!({
"provider_type": "claude_code",
"org_uuid": "shared-org"
}));
assert_eq!(
match_claude_provider_oauth_identity(&new_auth_config, &same_account),
Some(true)
);
assert_eq!(
match_claude_provider_oauth_identity(&new_auth_config, &different_account),
Some(false)
);
assert_eq!(
match_claude_provider_oauth_identity(&new_auth_config, &organization_only),
None
);
}
#[test]
fn claude_account_lock_prefers_uuid_and_falls_back_to_normalized_email() {
let with_uuid = auth_config(json!({
"account_uuid": "account-1",
"email": "[email protected]"
}));
let same_uuid_other_email = auth_config(json!({
"accountUuid": "account-1",
"email": "[email protected]"
}));
let email_only_uppercase = auth_config(json!({"email": "[email protected]"}));
let email_only_lowercase = auth_config(json!({"email": "[email protected]"}));
let uuid_key = claude_oauth_account_lock_key("provider-claude", &with_uuid)
.expect("uuid lock key should build");
assert_eq!(
Some(uuid_key.as_str()),
claude_oauth_account_lock_key("provider-claude", &same_uuid_other_email).as_deref()
);
assert_eq!(
claude_oauth_account_lock_key("provider-claude", &email_only_uppercase),
claude_oauth_account_lock_key("provider-claude", &email_only_lowercase)
);
assert!(uuid_key.starts_with("provider_oauth_claude_account:"));
assert!(!uuid_key.contains("account-1"));
assert!(claude_oauth_account_lock_key("provider-claude", &Map::new()).is_none());
}
#[tokio::test]
async fn concurrent_claude_writes_contend_on_the_same_account_lock() {
let app = AppState::new().expect("app state should build");
let state = AdminAppState::new(&app);
let config = auth_config(json!({
"provider_type": "claude_code",
"account_uuid": "account-shared",
"email": "[email protected]"
}));
let first =
acquire_claude_oauth_account_lock(&state, "provider-claude", &config, "first-test")
.await
.expect("first Claude lock should acquire");
let second =
acquire_claude_oauth_account_lock(&state, "provider-claude", &config, "second-test")
.await
.expect_err("second Claude lock should contend");
assert_eq!(second, ClaudeOAuthAccountLockError::Contended);
release_provider_oauth_account_locks(&state, first).await;
let retry =
acquire_claude_oauth_account_lock(&state, "provider-claude", &config, "retry-test")
.await
.expect("Claude lock should be reusable after release");
release_provider_oauth_account_locks(&state, retry).await;
}
#[test]
fn codex_agent_identity_matches_runtime_without_account_metadata() {
let new_auth_config = auth_config(json!({
@@ -1,3 +1,9 @@
use super::duplicates::{
acquire_claude_oauth_account_lock, acquire_codex_oauth_account_locks,
release_provider_oauth_account_locks,
};
use super::errors::build_internal_control_error_response;
use super::runtime::spawn_provider_oauth_account_state_refresh_after_update;
use super::state::{
decode_jwt_claims, enrich_admin_provider_oauth_auth_config, json_non_empty_string,
json_u64_value,
@@ -9,15 +15,22 @@ use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::request::AdminAppState;
use crate::provider_key_auth::provider_active_api_formats;
use crate::GatewayError;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::pool_scores::{
GetPoolMemberScoresByIdsQuery, PoolMemberIdentity,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_transport::{
grok_browser_transport_fingerprint_from_auth_config, provider_types::provider_type_is_fixed,
};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::{json, Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
@@ -101,6 +114,168 @@ pub(crate) fn build_provider_oauth_auth_config_from_token_payload(
(auth_config, access_token, refresh_token, expires_at)
}
pub(crate) async fn provision_provider_oauth_token_payload_for_provider(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
token_payload: &Value,
requested_name: Option<String>,
key_proxy: Option<Value>,
request_proxy: Option<ProxySnapshot>,
lock_operation: &'static str,
) -> Result<Response<Body>, GatewayError> {
let provider_id = provider.id.clone();
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
let (auth_config, access_token, refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, token_payload);
let Some(access_token) = access_token else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let api_formats = provider_oauth_active_api_formats(endpoints);
let oauth_account_leases = if provider_type == "codex" {
match acquire_codex_oauth_account_locks(state, &provider_id, &auth_config, lock_operation)
.await
{
Ok(leases) => leases,
Err(error) => {
return Ok(build_internal_control_error_response(
error.status_code(),
error.detail(),
));
}
}
} else if provider_type == "claude_code" {
match acquire_claude_oauth_account_lock(state, &provider_id, &auth_config, lock_operation)
.await
{
Ok(leases) => leases,
Err(error) => {
return Ok(build_internal_control_error_response(
error.status_code(),
error.detail(),
));
}
}
} else {
Vec::new()
};
let duplicate = match state
.find_duplicate_provider_oauth_key(&provider_id, &auth_config, None)
.await
{
Ok(duplicate) => duplicate,
Err(detail) => {
release_provider_oauth_account_locks(state, oauth_account_leases).await;
return Ok(build_internal_control_error_response(
if provider_type == "codex" {
http::StatusCode::CONFLICT
} else {
http::StatusCode::BAD_REQUEST
},
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
match state
.update_existing_provider_oauth_catalog_key(
&existing_key,
&provider_type,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await
{
Err(error) => {
release_provider_oauth_account_locks(state, oauth_account_leases).await;
return Err(error);
}
Ok(Some(key)) => key,
Ok(None) => {
release_provider_oauth_account_locks(state, oauth_account_leases).await;
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let name = requested_name
.or_else(|| {
auth_config
.get("email")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.unwrap_or_else(|| {
format!(
"账号_{}",
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
)
});
match state
.create_provider_oauth_catalog_key(
&provider_id,
&provider_type,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy,
expires_at,
)
.await
{
Err(error) => {
release_provider_oauth_account_locks(state, oauth_account_leases).await;
return Err(error);
}
Ok(Some(key)) => key,
Ok(None) => {
release_provider_oauth_account_locks(state, oauth_account_leases).await;
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
release_provider_oauth_account_locks(state, oauth_account_leases).await;
spawn_provider_oauth_account_state_refresh_after_update(
state.cloned_app(),
provider.clone(),
persisted_key.id.clone(),
request_proxy,
);
Ok(Json(json!({
"key_id": persisted_key.id,
"provider_type": provider_type,
"expires_at": expires_at,
"has_refresh_token": refresh_token.is_some(),
"temporary": refresh_token.is_none(),
"email": auth_config.get("email").cloned().unwrap_or(Value::Null),
"replaced": replaced,
}))
.into_response())
}
fn grok_oauth_catalog_key_fingerprint(
provider_type: &str,
auth_config: &Map<String, Value>,
@@ -3,8 +3,13 @@ use super::super::errors::{
};
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
use aether_contracts::ProxySnapshot;
use aether_oauth::provider::providers::GenericProviderOAuthAdapter;
use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext};
use aether_oauth::provider::providers::{
ClaudeCodeProviderOAuthAdapter, GenericProviderOAuthAdapter, CLAUDE_CODE_PROVIDER_TYPE,
CLAUDE_CODE_TOKEN_URL, CLAUDE_CODE_WEB_BASE_URL,
};
use aether_oauth::provider::{
ProviderOAuthCookieAuthorizationInput, ProviderOAuthService, ProviderOAuthTransportContext,
};
use axum::{body::Body, http, response::Response};
use std::sync::Arc;
@@ -140,3 +145,40 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
)
})
}
pub(crate) async fn authorize_admin_provider_oauth_with_cookie(
state: &AdminAppState<'_>,
session_key: String,
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let web_base_url =
state.provider_oauth_token_url("claude_code_cookie_base_url", CLAUDE_CODE_WEB_BASE_URL);
let token_url =
state.provider_oauth_token_url(CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_TOKEN_URL);
let service = ProviderOAuthService::new().with_adapter(Arc::new(
ClaudeCodeProviderOAuthAdapter::default().with_endpoint_overrides(web_base_url, token_url),
));
let ctx = provider_oauth_exchange_context(CLAUDE_CODE_PROVIDER_TYPE, proxy);
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let result = service
.authorize_with_cookie(
&executor,
&ctx,
ProviderOAuthCookieAuthorizationInput { session_key },
)
.await
.map_err(|error| {
let detail = if matches!(error, aether_oauth::core::OAuthError::InvalidRequest(_)) {
"Claude Cookie 格式无效"
} else {
"Claude Cookie 授权失败"
};
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, detail)
})?;
token_payload_from_provider_oauth_result(result).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Claude Cookie 授权返回缺少 access_token",
)
})
}
@@ -5,7 +5,8 @@ mod template;
pub(crate) use self::auth_config::enrich_admin_provider_oauth_auth_config;
pub(crate) use self::exchange::{
exchange_admin_provider_oauth_code, exchange_admin_provider_oauth_refresh_token,
authorize_admin_provider_oauth_with_cookie, exchange_admin_provider_oauth_code,
exchange_admin_provider_oauth_refresh_token,
};
pub(crate) use self::storage::build_provider_oauth_start_response;
pub(crate) use self::template::{
@@ -17,7 +17,7 @@ pub(crate) fn build_provider_oauth_start_response(
"authorization_url": authorization_url,
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url",
"instructions": "1) 打开 authorization_url 完成授权\n2) 复制授权页面显示的授权码或浏览器中的完整回调 URL\n3) 调用 complete 接口粘贴 callback_url",
})
}
@@ -20,9 +20,14 @@ pub(crate) fn admin_provider_oauth_template(
}
pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_json::Value> {
let service = aether_oauth::provider::ProviderOAuthService::with_builtin_adapters();
admin_provider_oauth_template_types()
.filter_map(|provider_type| admin_provider_oauth_template(provider_type))
.map(|template| {
let capabilities = service
.adapter(template.provider_type)
.ok()
.map(|adapter| adapter.capabilities());
json!({
"provider_type": template.provider_type,
"display_name": template.display_name,
@@ -31,6 +36,10 @@ pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_
"authorize_url": template.authorize_url,
"token_url": template.token_url,
"use_pkce": template.use_pkce,
"supports_authorization_code": capabilities.as_ref().is_some_and(|value| value.supports_authorization_code),
"supports_cookie_authorization": capabilities.as_ref().is_some_and(|value| value.supports_cookie_authorization),
"supports_refresh_token_import": capabilities.as_ref().is_some_and(|value| value.supports_refresh_token_import),
"supports_batch_import": capabilities.as_ref().is_some_and(|value| value.supports_batch_import),
})
})
.collect()
@@ -24,7 +24,9 @@ pub(crate) use self::oauth::{
admin_provider_oauth_agent_identity_import_task_provider_id,
admin_provider_oauth_batch_import_provider_id, admin_provider_oauth_batch_import_task_path,
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id,
admin_provider_oauth_complete_provider_id, admin_provider_oauth_cookie_provider_id,
admin_provider_oauth_cookie_task_path, admin_provider_oauth_cookie_task_provider_id,
admin_provider_oauth_device_authorize_provider_id,
admin_provider_oauth_device_poll_provider_id, admin_provider_oauth_import_provider_id,
admin_provider_oauth_refresh_key_id, admin_provider_oauth_start_key_id,
admin_provider_oauth_start_provider_id,
@@ -34,6 +34,33 @@ pub(crate) fn admin_provider_oauth_import_provider_id(request_path: &str) -> Opt
provider_oauth_provider_id_for_suffix(request_path, "/import-refresh-token")
}
pub(crate) fn admin_provider_oauth_cookie_provider_id(request_path: &str) -> Option<String> {
provider_oauth_provider_id_for_suffix(request_path, "/cookie-authorize")
}
pub(crate) fn admin_provider_oauth_cookie_task_provider_id(request_path: &str) -> Option<String> {
provider_oauth_provider_id_for_suffix(request_path, "/cookie-authorize/tasks")
}
pub(crate) fn admin_provider_oauth_cookie_task_path(
request_path: &str,
) -> Option<(String, String)> {
let suffix = request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/")
.unwrap_or(request_path.strip_prefix("/api/admin/provider-oauth/providers/")?);
let (provider_id, task_id) = suffix.split_once("/cookie-authorize/tasks/")?;
if provider_id.is_empty()
|| provider_id.contains('/')
|| task_id.is_empty()
|| task_id.contains('/')
|| !task_id.starts_with("claude-cookie-")
{
return None;
}
Some((provider_id.to_string(), task_id.to_string()))
}
pub(crate) fn admin_provider_oauth_batch_import_provider_id(request_path: &str) -> Option<String> {
provider_oauth_provider_id_for_suffix(request_path, "/batch-import")
}
@@ -110,6 +137,7 @@ mod tests {
use super::{
admin_provider_oauth_agent_identity_import_task_path,
admin_provider_oauth_agent_identity_import_task_provider_id,
admin_provider_oauth_cookie_task_path, admin_provider_oauth_cookie_task_provider_id,
};
#[test]
@@ -139,4 +167,28 @@ mod tests {
)
.is_none());
}
#[test]
fn parses_dedicated_claude_cookie_task_paths() {
assert_eq!(
admin_provider_oauth_cookie_task_provider_id(
"/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks",
)
.as_deref(),
Some("provider-claude")
);
assert_eq!(
admin_provider_oauth_cookie_task_path(
"/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks/claude-cookie-task-1",
),
Some((
"provider-claude".to_string(),
"claude-cookie-task-1".to_string(),
))
);
assert!(admin_provider_oauth_cookie_task_path(
"/api/admin/provider-oauth/providers/provider-claude/cookie-authorize/tasks/task-1",
)
.is_none());
}
}
@@ -555,6 +555,7 @@ impl<'a> AdminAppState<'a> {
json_body,
body_bytes,
network,
transport_profile: None,
};
let response = aether_oauth::network::OAuthHttpExecutor::execute(
&crate::oauth::GatewayOAuthHttpExecutor::new(*self),