mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
Fix OAuth token import and table filters
This commit is contained in:
+105
-38
@@ -1,3 +1,4 @@
|
||||
use super::super::token_import::build_codex_access_token_import_auth_config;
|
||||
use super::kiro_import::execute_admin_provider_oauth_kiro_batch_import;
|
||||
use super::parse::{
|
||||
apply_admin_provider_oauth_batch_import_hints, extract_admin_provider_oauth_batch_error_detail,
|
||||
@@ -22,12 +23,19 @@ use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, exchange_admin_provider_oauth_refresh_token,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
|
||||
use serde_json::json;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
struct AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: String,
|
||||
auth_config: Map<String, Value>,
|
||||
expires_at: Option<u64>,
|
||||
}
|
||||
|
||||
pub(super) fn estimate_admin_provider_oauth_batch_import_total(
|
||||
provider_type: &str,
|
||||
raw_credentials: &str,
|
||||
@@ -70,6 +78,90 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
}
|
||||
}
|
||||
|
||||
async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
provider_type: &str,
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
) -> Result<AdminProviderOAuthResolvedBatchImport, String> {
|
||||
let refresh_token = entry
|
||||
.refresh_token
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
let access_token = entry
|
||||
.access_token
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
if let Some(refresh_token) = refresh_token {
|
||||
let token_payload = match exchange_admin_provider_oauth_refresh_token(
|
||||
state,
|
||||
template,
|
||||
refresh_token,
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Err(response) => {
|
||||
let detail = extract_admin_provider_oauth_batch_error_detail(response).await;
|
||||
if provider_type.eq_ignore_ascii_case("codex") {
|
||||
if let Some(access_token) = access_token {
|
||||
let (auth_config, expires_at) = build_codex_access_token_import_auth_config(
|
||||
access_token,
|
||||
Some(refresh_token),
|
||||
entry.expires_at,
|
||||
Some(detail.as_str()),
|
||||
);
|
||||
return Ok(AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
}
|
||||
return Err(format!("Token 验证失败: {detail}"));
|
||||
}
|
||||
};
|
||||
|
||||
let (mut auth_config, access_token, returned_refresh_token, expires_at) =
|
||||
build_provider_oauth_auth_config_from_token_payload(provider_type, &token_payload);
|
||||
let Some(access_token) = access_token else {
|
||||
return Err("Token 刷新返回缺少 access_token".to_string());
|
||||
};
|
||||
|
||||
let refresh_token = returned_refresh_token
|
||||
.or_else(|| Some(refresh_token.to_string()))
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
if let Some(refresh_token) = refresh_token.as_ref() {
|
||||
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
|
||||
}
|
||||
return Ok(AdminProviderOAuthResolvedBatchImport {
|
||||
access_token,
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(access_token) = access_token {
|
||||
if !provider_type.eq_ignore_ascii_case("codex") {
|
||||
return Err("Access Token 导入仅支持 Codex Provider".to_string());
|
||||
}
|
||||
let (auth_config, expires_at) =
|
||||
build_codex_access_token_import_auth_config(access_token, None, entry.expires_at, None);
|
||||
return Ok(AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
|
||||
Err("Refresh Token 或 Access Token 不能为空".to_string())
|
||||
}
|
||||
|
||||
pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
@@ -145,24 +237,22 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
let mut failed = 0usize;
|
||||
|
||||
for (index, entry) in entries.iter().enumerate() {
|
||||
let token_payload = match exchange_admin_provider_oauth_refresh_token(
|
||||
let resolved_import = match resolve_admin_provider_oauth_batch_import_tokens(
|
||||
state,
|
||||
template,
|
||||
entry.refresh_token.as_str(),
|
||||
provider_type,
|
||||
entry,
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Err(response) => {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": format!(
|
||||
"Token 验证失败: {}",
|
||||
extract_admin_provider_oauth_batch_error_detail(response).await
|
||||
),
|
||||
"error": error,
|
||||
"replaced": false,
|
||||
}));
|
||||
maybe_report_admin_provider_oauth_batch_import_progress(
|
||||
@@ -176,34 +266,11 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let (mut auth_config, access_token, returned_refresh_token, expires_at) =
|
||||
build_provider_oauth_auth_config_from_token_payload(provider_type, &token_payload);
|
||||
let Some(access_token) = access_token else {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "Token 刷新返回缺少 access_token",
|
||||
"replaced": false,
|
||||
}));
|
||||
maybe_report_admin_provider_oauth_batch_import_progress(
|
||||
&mut progress,
|
||||
entries.len(),
|
||||
success,
|
||||
failed,
|
||||
&results,
|
||||
)
|
||||
.await;
|
||||
continue;
|
||||
};
|
||||
|
||||
let refresh_token = returned_refresh_token
|
||||
.or_else(|| Some(entry.refresh_token.clone()))
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
if let Some(refresh_token) = refresh_token.as_ref() {
|
||||
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
|
||||
}
|
||||
let AdminProviderOAuthResolvedBatchImport {
|
||||
access_token,
|
||||
mut auth_config,
|
||||
expires_at,
|
||||
} = resolved_import;
|
||||
apply_admin_provider_oauth_batch_import_hints(provider_type, entry, &mut auth_config);
|
||||
|
||||
let duplicate =
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use super::super::token_import::{import_tokens_from_raw_token, normalize_single_import_tokens};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::current_unix_secs;
|
||||
use crate::handlers::admin::provider::oauth::state::{current_unix_secs, json_u64_value};
|
||||
use axum::{
|
||||
body::{to_bytes, Body, Bytes},
|
||||
http,
|
||||
@@ -18,7 +19,9 @@ pub(super) struct AdminProviderOAuthBatchImportRequest {
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct AdminProviderOAuthBatchImportEntry {
|
||||
pub refresh_token: String,
|
||||
pub refresh_token: Option<String>,
|
||||
pub access_token: Option<String>,
|
||||
pub expires_at: Option<u64>,
|
||||
pub account_id: Option<String>,
|
||||
pub account_user_id: Option<String>,
|
||||
pub plan_type: Option<String>,
|
||||
@@ -73,8 +76,11 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
if refresh_token.is_empty() {
|
||||
None
|
||||
} else {
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(refresh_token);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token: refresh_token.to_string(),
|
||||
refresh_token,
|
||||
access_token,
|
||||
expires_at: None,
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
@@ -88,7 +94,19 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
object
|
||||
.get("refresh_token")
|
||||
.or_else(|| object.get("refreshToken")),
|
||||
)?;
|
||||
);
|
||||
let access_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("access_token")
|
||||
.or_else(|| object.get("accessToken")),
|
||||
);
|
||||
let (refresh_token, access_token) =
|
||||
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref());
|
||||
if refresh_token.is_none() && access_token.is_none() {
|
||||
return None;
|
||||
}
|
||||
let expires_at =
|
||||
json_u64_value(object.get("expires_at").or_else(|| object.get("expiresAt")));
|
||||
let account_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_id")
|
||||
@@ -121,6 +139,8 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
let email = coerce_admin_provider_oauth_import_str(object.get("email"));
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
expires_at,
|
||||
account_id,
|
||||
account_user_id,
|
||||
plan_type,
|
||||
@@ -163,13 +183,18 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
raw.lines()
|
||||
.map(str::trim)
|
||||
.filter(|line| !line.is_empty() && !line.starts_with('#'))
|
||||
.map(|refresh_token| AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token: refresh_token.to_string(),
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
.map(|token| {
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token);
|
||||
AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
expires_at: None,
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -291,3 +316,47 @@ pub(super) fn build_admin_provider_oauth_batch_task_state(
|
||||
"updated_at": updated_at,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::parse_admin_provider_oauth_batch_import_entries;
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
use serde_json::json;
|
||||
|
||||
fn unsigned_jwt(payload: serde_json::Value) -> String {
|
||||
let header = json!({"alg": "none", "typ": "JWT"});
|
||||
let encode = |value: serde_json::Value| {
|
||||
URL_SAFE_NO_PAD.encode(serde_json::to_vec(&value).expect("jwt json should serialize"))
|
||||
};
|
||||
format!("{}.{}.signature", encode(header), encode(payload))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_access_token_only_entry() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
r#"[{"accessToken":"at_1","expiresAt":2100000000,"accountId":"acc-1","email":"[email protected]"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("at_1"));
|
||||
assert_eq!(entries[0].expires_at, Some(2_100_000_000));
|
||||
assert_eq!(entries[0].account_id.as_deref(), Some("acc-1"));
|
||||
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_plain_jwt_line_as_access_token() {
|
||||
let token = unsigned_jwt(json!({
|
||||
"iss": "https://auth.openai.com",
|
||||
"aud": ["https://api.openai.com/v1"],
|
||||
"exp": 2_000_000_000u64,
|
||||
}));
|
||||
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(&token);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some(token.as_str()));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::provider_oauth_token_payload_expires_at_unix_secs;
|
||||
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::super::super::runtime::provider_oauth_runtime_endpoint_for_provider;
|
||||
use super::super::super::state::{
|
||||
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
|
||||
is_fixed_provider_type_for_provider_oauth, json_non_empty_string, json_u64_value,
|
||||
is_fixed_provider_type_for_provider_oauth, json_non_empty_string,
|
||||
};
|
||||
use super::shared::{
|
||||
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
|
||||
@@ -168,8 +169,8 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let expires_at = json_u64_value(token_payload.get("expires_in"))
|
||||
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
|
||||
let expires_at =
|
||||
provider_oauth_token_payload_expires_at_unix_secs(&token_payload, now_unix_secs);
|
||||
|
||||
let mut auth_config = serde_json::Map::new();
|
||||
auth_config.insert("provider_type".to_string(), json!(provider_type.clone()));
|
||||
|
||||
@@ -12,10 +12,17 @@ use super::super::runtime::{
|
||||
use super::super::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
exchange_admin_provider_oauth_refresh_token, is_fixed_provider_type_for_provider_oauth,
|
||||
json_u64_value,
|
||||
};
|
||||
use super::token_import::{
|
||||
build_codex_access_token_import_auth_config, normalize_single_import_tokens,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_import_provider_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::request::{
|
||||
AdminAppState, AdminProviderOAuthTemplate, AdminRequestContext,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -25,6 +32,125 @@ use axum::{
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
struct AdminProviderOAuthSingleImportTokens {
|
||||
access_token: String,
|
||||
auth_config: serde_json::Map<String, serde_json::Value>,
|
||||
expires_at: Option<u64>,
|
||||
}
|
||||
|
||||
fn import_payload_string(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
snake_case: &str,
|
||||
camel_case: &str,
|
||||
) -> Option<String> {
|
||||
payload
|
||||
.get(snake_case)
|
||||
.or_else(|| payload.get(camel_case))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn import_payload_u64(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
snake_case: &str,
|
||||
camel_case: &str,
|
||||
) -> Option<u64> {
|
||||
json_u64_value(payload.get(snake_case).or_else(|| payload.get(camel_case)))
|
||||
}
|
||||
|
||||
async fn resolve_admin_provider_oauth_single_import_tokens(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
provider_type: &str,
|
||||
refresh_token: Option<&str>,
|
||||
access_token: Option<&str>,
|
||||
imported_expires_at: Option<u64>,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
) -> Result<AdminProviderOAuthSingleImportTokens, Response<Body>> {
|
||||
if let Some(refresh_token) = refresh_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let token_payload = match state
|
||||
.exchange_admin_provider_oauth_refresh_token(
|
||||
template,
|
||||
refresh_token,
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Err(response) => {
|
||||
if provider_type.eq_ignore_ascii_case("codex") {
|
||||
if let Some(access_token) = access_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let (auth_config, expires_at) = build_codex_access_token_import_auth_config(
|
||||
access_token,
|
||||
Some(refresh_token),
|
||||
imported_expires_at,
|
||||
Some("Refresh Token 验证失败,已回退为 Access Token 导入"),
|
||||
);
|
||||
return Ok(AdminProviderOAuthSingleImportTokens {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
}
|
||||
return Err(response);
|
||||
}
|
||||
};
|
||||
|
||||
let (mut auth_config, access_token, returned_refresh_token, expires_at) =
|
||||
build_provider_oauth_auth_config_from_token_payload(provider_type, &token_payload);
|
||||
let Some(access_token) = access_token else {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"token refresh 返回缺少 access_token",
|
||||
));
|
||||
};
|
||||
let refresh_token = returned_refresh_token
|
||||
.or_else(|| Some(refresh_token.to_string()))
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
if let Some(refresh_token) = refresh_token.as_ref() {
|
||||
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
|
||||
}
|
||||
return Ok(AdminProviderOAuthSingleImportTokens {
|
||||
access_token,
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
|
||||
let Some(access_token) = access_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token 或 Access Token 不能为空",
|
||||
));
|
||||
};
|
||||
if !provider_type.eq_ignore_ascii_case("codex") {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Access Token 导入仅支持 Codex Provider",
|
||||
));
|
||||
}
|
||||
|
||||
let (auth_config, expires_at) =
|
||||
build_codex_access_token_import_auth_config(access_token, None, imported_expires_at, None);
|
||||
Ok(AdminProviderOAuthSingleImportTokens {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
@@ -54,17 +180,19 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
));
|
||||
}
|
||||
};
|
||||
let Some(refresh_token_input) = raw_payload
|
||||
.get("refresh_token")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
let refresh_token_input = import_payload_string(&raw_payload, "refresh_token", "refreshToken");
|
||||
let access_token_input = import_payload_string(&raw_payload, "access_token", "accessToken");
|
||||
let imported_expires_at = import_payload_u64(&raw_payload, "expires_at", "expiresAt");
|
||||
let (refresh_token_input, access_token_input) = normalize_single_import_tokens(
|
||||
refresh_token_input.as_deref(),
|
||||
access_token_input.as_deref(),
|
||||
);
|
||||
if refresh_token_input.is_none() && access_token_input.is_none() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token 不能为空",
|
||||
"Refresh Token 或 Access Token 不能为空",
|
||||
));
|
||||
};
|
||||
}
|
||||
let name = raw_payload
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
@@ -122,32 +250,30 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
.await;
|
||||
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
|
||||
|
||||
let token_payload = match state
|
||||
.exchange_admin_provider_oauth_refresh_token(
|
||||
template,
|
||||
refresh_token_input,
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
let resolved_import = match resolve_admin_provider_oauth_single_import_tokens(
|
||||
state,
|
||||
template,
|
||||
&provider_type,
|
||||
refresh_token_input.as_deref(),
|
||||
access_token_input.as_deref(),
|
||||
imported_expires_at,
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Ok(value) => value,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
|
||||
let (mut auth_config, access_token, returned_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 refresh 返回缺少 access_token",
|
||||
));
|
||||
};
|
||||
let refresh_token = returned_refresh_token
|
||||
.or_else(|| Some(refresh_token_input.to_string()))
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
if let Some(refresh_token) = refresh_token.as_ref() {
|
||||
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
|
||||
}
|
||||
let AdminProviderOAuthSingleImportTokens {
|
||||
access_token,
|
||||
auth_config,
|
||||
expires_at,
|
||||
} = resolved_import;
|
||||
let has_refresh_token = auth_config
|
||||
.get("refresh_token")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
|
||||
let api_formats = provider_oauth_active_api_formats(&endpoints);
|
||||
let duplicate = match state
|
||||
@@ -239,7 +365,11 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
"key_id": persisted_key.id,
|
||||
"provider_type": provider_type,
|
||||
"expires_at": expires_at,
|
||||
"has_refresh_token": refresh_token.is_some(),
|
||||
"has_refresh_token": has_refresh_token,
|
||||
"temporary": auth_config
|
||||
.get("access_token_import_temporary")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false),
|
||||
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
|
||||
"replaced": replaced,
|
||||
}))
|
||||
|
||||
@@ -27,6 +27,7 @@ mod kiro;
|
||||
mod refresh;
|
||||
mod start;
|
||||
mod tasks;
|
||||
mod token_import;
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
state: &AdminAppState<'_>,
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
use super::super::provisioning::build_provider_oauth_auth_config_from_token_payload;
|
||||
use super::super::state::json_u64_value;
|
||||
use base64::{
|
||||
engine::general_purpose::{URL_SAFE, URL_SAFE_NO_PAD},
|
||||
Engine as _,
|
||||
};
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
fn decode_base64_url_part(value: &str) -> Option<Vec<u8>> {
|
||||
URL_SAFE_NO_PAD
|
||||
.decode(value.as_bytes())
|
||||
.or_else(|_| URL_SAFE.decode(value.as_bytes()))
|
||||
.or_else(|_| {
|
||||
let mut padded = value.to_string();
|
||||
let remainder = padded.len() % 4;
|
||||
if remainder != 0 {
|
||||
padded.extend(std::iter::repeat_n('=', 4 - remainder));
|
||||
}
|
||||
URL_SAFE.decode(padded.as_bytes())
|
||||
})
|
||||
.ok()
|
||||
}
|
||||
|
||||
fn decode_unverified_jwt_json_part(part: &str) -> Option<Map<String, Value>> {
|
||||
let bytes = decode_base64_url_part(part)?;
|
||||
serde_json::from_slice::<Value>(&bytes)
|
||||
.ok()?
|
||||
.as_object()
|
||||
.cloned()
|
||||
}
|
||||
|
||||
pub(super) fn looks_like_access_token(token: &str) -> bool {
|
||||
let parts = token.trim().split('.').collect::<Vec<_>>();
|
||||
if parts.len() != 3 || parts.iter().any(|part| part.is_empty()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let Some(header) = decode_unverified_jwt_json_part(parts[0]) else {
|
||||
return false;
|
||||
};
|
||||
let Some(payload) = decode_unverified_jwt_json_part(parts[1]) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
let token_type = header
|
||||
.get("typ")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase();
|
||||
if !token_type.is_empty() && token_type != "jwt" && token_type != "at+jwt" {
|
||||
return false;
|
||||
}
|
||||
|
||||
["exp", "aud", "iss", "scope", "scp"]
|
||||
.iter()
|
||||
.any(|field| payload.contains_key(*field))
|
||||
}
|
||||
|
||||
pub(super) fn normalize_single_import_tokens(
|
||||
refresh_token: Option<&str>,
|
||||
access_token: Option<&str>,
|
||||
) -> (Option<String>, Option<String>) {
|
||||
let mut refresh_token = refresh_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let mut access_token = access_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
|
||||
if access_token.is_none()
|
||||
&& refresh_token
|
||||
.as_deref()
|
||||
.is_some_and(looks_like_access_token)
|
||||
{
|
||||
access_token = refresh_token.take();
|
||||
}
|
||||
|
||||
(refresh_token, access_token)
|
||||
}
|
||||
|
||||
pub(super) fn import_tokens_from_raw_token(token: &str) -> (Option<String>, Option<String>) {
|
||||
if looks_like_access_token(token) {
|
||||
(None, Some(token.trim().to_string()))
|
||||
} else {
|
||||
(Some(token.trim().to_string()), None)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn decode_access_token_expires_at(access_token: &str) -> Option<u64> {
|
||||
let payload = access_token.trim().split('.').nth(1)?;
|
||||
let claims = decode_unverified_jwt_json_part(payload)?;
|
||||
json_u64_value(claims.get("exp"))
|
||||
}
|
||||
|
||||
pub(super) fn build_codex_access_token_import_auth_config(
|
||||
access_token: &str,
|
||||
refresh_token: Option<&str>,
|
||||
imported_expires_at: Option<u64>,
|
||||
refresh_error: Option<&str>,
|
||||
) -> (Map<String, Value>, Option<u64>) {
|
||||
let token_payload = json!({
|
||||
"access_token": access_token,
|
||||
"token_type": "Bearer",
|
||||
});
|
||||
let (mut auth_config, _, _, _) =
|
||||
build_provider_oauth_auth_config_from_token_payload("codex", &token_payload);
|
||||
|
||||
let refresh_token = refresh_token
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
if let Some(refresh_token) = refresh_token {
|
||||
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
|
||||
}
|
||||
|
||||
auth_config.insert(
|
||||
"access_token_import_temporary".to_string(),
|
||||
json!(refresh_token.is_none()),
|
||||
);
|
||||
|
||||
if let Some(expires_at) = decode_access_token_expires_at(access_token).or(imported_expires_at) {
|
||||
auth_config.insert("expires_at".to_string(), json!(expires_at));
|
||||
}
|
||||
if let Some(refresh_error) = refresh_error
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
auth_config.insert(
|
||||
"refresh_token_import_error".to_string(),
|
||||
json!(refresh_error),
|
||||
);
|
||||
}
|
||||
|
||||
let expires_at = auth_config.get("expires_at").and_then(Value::as_u64);
|
||||
(auth_config, expires_at)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_codex_access_token_import_auth_config, decode_access_token_expires_at,
|
||||
looks_like_access_token, normalize_single_import_tokens,
|
||||
};
|
||||
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
|
||||
use serde_json::json;
|
||||
|
||||
fn unsigned_jwt(payload: serde_json::Value) -> String {
|
||||
let header = json!({"alg": "none", "typ": "JWT"});
|
||||
let encode = |value: serde_json::Value| {
|
||||
URL_SAFE_NO_PAD.encode(serde_json::to_vec(&value).expect("jwt json should serialize"))
|
||||
};
|
||||
format!("{}.{}.signature", encode(header), encode(payload))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_plain_jwt_access_token() {
|
||||
let token = unsigned_jwt(json!({
|
||||
"iss": "https://auth.openai.com",
|
||||
"aud": ["https://api.openai.com/v1"],
|
||||
"exp": 2_000_000_000u64,
|
||||
}));
|
||||
|
||||
assert!(looks_like_access_token(&token));
|
||||
let (refresh_token, access_token) = normalize_single_import_tokens(Some(&token), None);
|
||||
assert!(refresh_token.is_none());
|
||||
assert_eq!(access_token.as_deref(), Some(token.as_str()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_codex_temporary_auth_config_from_access_token() {
|
||||
let token = unsigned_jwt(json!({
|
||||
"exp": 2_000_000_000u64,
|
||||
"https://api.openai.com/profile": {
|
||||
"email": "[email protected]"
|
||||
},
|
||||
}));
|
||||
|
||||
let (auth_config, expires_at) =
|
||||
build_codex_access_token_import_auth_config(&token, None, None, None);
|
||||
|
||||
assert_eq!(expires_at, Some(2_000_000_000));
|
||||
assert_eq!(decode_access_token_expires_at(&token), Some(2_000_000_000));
|
||||
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
|
||||
assert_eq!(
|
||||
auth_config.get("access_token_import_temporary"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
assert!(auth_config.get("refresh_token").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_codex_auth_config_with_imported_expires_at_when_token_has_no_exp() {
|
||||
let token = unsigned_jwt(json!({
|
||||
"iss": "https://auth.openai.com",
|
||||
"aud": ["https://api.openai.com/v1"]
|
||||
}));
|
||||
|
||||
let (auth_config, expires_at) =
|
||||
build_codex_access_token_import_auth_config(&token, None, Some(2_100_000_000), None);
|
||||
|
||||
assert_eq!(expires_at, Some(2_100_000_000));
|
||||
assert_eq!(
|
||||
auth_config.get("expires_at"),
|
||||
Some(&json!(2_100_000_000u64))
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user