mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
Merge remote-tracking branch 'origin/main' into fix/gemini-cli-v1internal
# Conflicts: # apps/aether-gateway/src/ai_serving/transport.rs # apps/aether-gateway/src/handlers/admin/provider/oauth/dispatch/batch/parse.rs # apps/aether-gateway/src/handlers/shared/catalog.rs # crates/aether-admin/src/provider/quota.rs # crates/aether-model-fetch/src/strategy.rs # crates/aether-provider-pool/src/lib.rs # crates/aether-provider-pool/src/service.rs # crates/aether-provider-transport/src/provider_types.rs # frontend/src/features/providers/components/ProviderDetailDrawer.vue # frontend/src/utils/__tests__/providerKeyQuota.spec.ts # frontend/src/utils/providerKeyQuota.ts # frontend/src/views/admin/PoolManagement.vue
This commit is contained in:
@@ -37,6 +37,7 @@ pub(crate) use self::provider::oauth::runtime::{
|
||||
refresh_provider_oauth_account_state_after_update,
|
||||
};
|
||||
pub(crate) use self::provider::ops::providers::actions::admin_provider_ops_local_action_response;
|
||||
pub(crate) use self::provider::ops::providers::store_admin_provider_ops_balance_cache;
|
||||
pub(crate) use self::provider::pool::config::admin_provider_pool_config;
|
||||
pub(crate) use self::provider::pool_admin::maybe_build_local_admin_pool_response;
|
||||
pub(crate) use self::provider::shared::payloads::{
|
||||
|
||||
@@ -59,6 +59,11 @@ pub(super) async fn build_admin_monitoring_redis_cache_categories_response(
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let mut categories = Vec::with_capacity(ADMIN_MONITORING_REDIS_CACHE_CATEGORIES.len());
|
||||
let mut total_keys = 0usize;
|
||||
let diagnostics = state
|
||||
.runtime_state()
|
||||
.redis_diagnostics()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(format!("redis diagnostics failed: {err}")))?;
|
||||
|
||||
for (key, name, pattern, description) in ADMIN_MONITORING_REDIS_CACHE_CATEGORIES {
|
||||
let count = list_admin_monitoring_namespaced_keys(state, pattern)
|
||||
@@ -81,6 +86,7 @@ pub(super) async fn build_admin_monitoring_redis_cache_categories_response(
|
||||
"backend": state.runtime_state().backend_kind().as_str(),
|
||||
"categories": categories,
|
||||
"total_keys": total_keys,
|
||||
"diagnostics": diagnostics,
|
||||
}
|
||||
}))
|
||||
.into_response())
|
||||
|
||||
@@ -1194,6 +1194,7 @@ async fn admin_monitoring_redis_keys_returns_local_payload_without_redis() {
|
||||
assert_eq!(payload["data"]["available"], json!(true));
|
||||
assert_eq!(payload["data"]["backend"], json!("memory"));
|
||||
assert_eq!(payload["data"]["total_keys"], json!(0));
|
||||
assert_eq!(payload["data"]["diagnostics"], serde_json::Value::Null);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -5,8 +5,11 @@ use super::local_monitoring_response;
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::AppState;
|
||||
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
|
||||
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
|
||||
use aether_data_contracts::repository::{
|
||||
candidates::RequestCandidateStatus, usage::UsageBodyCaptureState,
|
||||
};
|
||||
use axum::body::to_bytes;
|
||||
use base64::{engine::general_purpose::STANDARD as BASE64_STANDARD, Engine as _};
|
||||
use serde_json::json;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -530,6 +533,89 @@ async fn admin_monitoring_trace_request_exposes_failed_candidate_upstream_respon
|
||||
assert!(extra.get("provider_response").is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_decodes_connect_json_response_body_refs() {
|
||||
let mut candidate = sample_candidate(
|
||||
"cand-used",
|
||||
"request-connect",
|
||||
0,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(429),
|
||||
);
|
||||
candidate.extra_data = Some(json!({"cache_1h": true}));
|
||||
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![candidate]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
));
|
||||
let mut usage = sample_usage(
|
||||
"request-connect",
|
||||
"provider-1",
|
||||
"Windsurf",
|
||||
0,
|
||||
0.0,
|
||||
"failed",
|
||||
Some(429),
|
||||
100,
|
||||
);
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.response_headers = Some(json!({
|
||||
"content-type": "application/connect+json"
|
||||
}));
|
||||
let mut framed = Vec::new();
|
||||
framed.push(2);
|
||||
let payload = br#"{"error":{"code":"resource_exhausted","message":"quota exhausted"}}"#;
|
||||
framed.extend_from_slice(&(payload.len() as u32).to_be_bytes());
|
||||
framed.extend_from_slice(payload);
|
||||
usage.response_body = Some(json!(BASE64_STANDARD.encode(framed)));
|
||||
usage.response_body_ref = Some("usage://request/request-connect/response_body".to_string());
|
||||
usage.response_body_state = Some(UsageBodyCaptureState::Inline);
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||
request_candidates,
|
||||
usage_repository,
|
||||
)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let context = request_context(
|
||||
http::Method::GET,
|
||||
"/api/admin/monitoring/trace/request-connect",
|
||||
);
|
||||
|
||||
let response = local_monitoring_response(&state, &context)
|
||||
.await
|
||||
.expect("handler should not error")
|
||||
.expect("route should be handled locally");
|
||||
|
||||
assert_eq!(response.status(), http::StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
let upstream_response = &payload["candidates"][0]["extra_data"]["upstream_response"];
|
||||
assert_eq!(upstream_response["status_code"], json!(429));
|
||||
assert_eq!(
|
||||
upstream_response["body"]["error"]["code"],
|
||||
json!("resource_exhausted")
|
||||
);
|
||||
assert_eq!(
|
||||
upstream_response["body"]["error"]["message"],
|
||||
json!("quota exhausted")
|
||||
);
|
||||
assert_eq!(
|
||||
upstream_response["body_ref"],
|
||||
json!("usage://request/request-connect/response_body")
|
||||
);
|
||||
assert_eq!(upstream_response["body_state"], json!("inline"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_exposes_structured_ranking_metadata() {
|
||||
let mut candidate = sample_candidate(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
|
||||
@@ -25,10 +25,15 @@ use crate::handlers::admin::provider::oauth::runtime::{
|
||||
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, AdminProviderOAuthTemplate};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_oauth::core::OAuthError;
|
||||
use aether_oauth::provider::{
|
||||
ProviderOAuthImportInput, ProviderOAuthService, ProviderOAuthTransportContext,
|
||||
};
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
struct AdminProviderOAuthResolvedBatchImport {
|
||||
@@ -37,6 +42,16 @@ struct AdminProviderOAuthResolvedBatchImport {
|
||||
expires_at: Option<u64>,
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_batch_import_error(error: &OAuthError) -> String {
|
||||
match error {
|
||||
OAuthError::InvalidRequest(_) => "Windsurf 凭据验证失败: 请求参数无效".to_string(),
|
||||
OAuthError::HttpStatus { status_code, .. } => {
|
||||
format!("Windsurf 凭据验证失败: HTTP {status_code}")
|
||||
}
|
||||
_ => "Windsurf 凭据验证失败".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn estimate_admin_provider_oauth_batch_import_total(
|
||||
provider_type: &str,
|
||||
raw_credentials: &str,
|
||||
@@ -98,6 +113,61 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
if provider_type.eq_ignore_ascii_case("windsurf") {
|
||||
let token_for_import = refresh_token.or(access_token);
|
||||
let ctx = ProviderOAuthTransportContext {
|
||||
provider_id: String::new(),
|
||||
provider_type: provider_type.to_string(),
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: None,
|
||||
endpoint_config: None,
|
||||
key_config: None,
|
||||
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||
request_proxy.clone(),
|
||||
),
|
||||
};
|
||||
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||
let result = ProviderOAuthService::with_builtin_adapters()
|
||||
.import_credentials(
|
||||
&executor,
|
||||
&ctx,
|
||||
ProviderOAuthImportInput {
|
||||
provider_type: provider_type.to_string(),
|
||||
name: entry
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|raw| raw.get("name"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
refresh_token: token_for_import.map(ToOwned::to_owned),
|
||||
raw_credentials: entry.raw_credentials.clone(),
|
||||
network: ctx.network.clone(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(|error| sanitize_windsurf_batch_import_error(&error))?;
|
||||
let access_token = result.token_set.access_token.trim().to_string();
|
||||
if access_token.is_empty() {
|
||||
return Err("Windsurf 凭据验证返回缺少 apiKey/sessionToken".to_string());
|
||||
}
|
||||
let auth_config = result
|
||||
.auth_config
|
||||
.as_object()
|
||||
.cloned()
|
||||
.ok_or_else(|| "Windsurf 凭据验证返回缺少 auth_config".to_string())?;
|
||||
return Ok(AdminProviderOAuthResolvedBatchImport {
|
||||
access_token,
|
||||
auth_config,
|
||||
expires_at: result.token_set.expires_at_unix_secs,
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(refresh_token) = refresh_token {
|
||||
let Some(template) = template else {
|
||||
if provider_type_supports_access_token_import(provider_type) {
|
||||
@@ -228,6 +298,28 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
};
|
||||
|
||||
let template = admin_provider_oauth_template(provider_type);
|
||||
if template.is_none()
|
||||
&& !provider_type.eq_ignore_ascii_case("windsurf")
|
||||
&& !provider_type_supports_access_token_import(provider_type)
|
||||
{
|
||||
return Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success: 0,
|
||||
failed: entries.len(),
|
||||
results: entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, _)| {
|
||||
json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
|
||||
"replaced": false,
|
||||
})
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
}
|
||||
|
||||
let endpoint_resolution =
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, provider_type).await?;
|
||||
@@ -251,6 +343,25 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
let mut failed = 0usize;
|
||||
|
||||
for (index, entry) in entries.iter().enumerate() {
|
||||
if let Some(error) = entry.parse_error.as_ref() {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": error,
|
||||
"replaced": false,
|
||||
}));
|
||||
maybe_report_admin_provider_oauth_batch_import_progress(
|
||||
&mut progress,
|
||||
entries.len(),
|
||||
success,
|
||||
failed,
|
||||
&results,
|
||||
)
|
||||
.await;
|
||||
continue;
|
||||
}
|
||||
|
||||
let resolved_import = match resolve_admin_provider_oauth_batch_import_tokens(
|
||||
state,
|
||||
template,
|
||||
@@ -418,3 +529,33 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
results,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::sanitize_windsurf_batch_import_error;
|
||||
use aether_oauth::core::OAuthError;
|
||||
|
||||
#[test]
|
||||
fn windsurf_batch_import_error_redacts_http_body() {
|
||||
let error = OAuthError::HttpStatus {
|
||||
status_code: 401,
|
||||
body_excerpt: "sessionToken=devin-session-token$secret".to_string(),
|
||||
};
|
||||
|
||||
let detail = sanitize_windsurf_batch_import_error(&error);
|
||||
|
||||
assert_eq!(detail, "Windsurf 凭据验证失败: HTTP 401");
|
||||
assert!(!detail.contains("devin-session-token$secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_batch_import_error_redacts_provider_detail() {
|
||||
let error = OAuthError::invalid_response("apiKey=sk-secret token=secret-token");
|
||||
|
||||
let detail = sanitize_windsurf_batch_import_error(&error);
|
||||
|
||||
assert_eq!(detail, "Windsurf 凭据验证失败");
|
||||
assert!(!detail.contains("sk-secret"));
|
||||
assert!(!detail.contains("secret-token"));
|
||||
}
|
||||
}
|
||||
|
||||
+8
-1
@@ -8,7 +8,7 @@ use super::parse::{
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
build_admin_provider_oauth_backend_unavailable_response,
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_provider_id;
|
||||
@@ -60,6 +60,13 @@ pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
|
||||
"该 Provider 不是固定类型,无法使用 provider-oauth",
|
||||
));
|
||||
}
|
||||
if provider_type != "kiro"
|
||||
&& provider_type != "windsurf"
|
||||
&& admin_provider_oauth_template(&provider_type).is_none()
|
||||
{
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
|
||||
@@ -19,8 +19,10 @@ pub(super) struct AdminProviderOAuthBatchImportRequest {
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct AdminProviderOAuthBatchImportEntry {
|
||||
pub parse_error: Option<String>,
|
||||
pub refresh_token: Option<String>,
|
||||
pub access_token: Option<String>,
|
||||
pub raw_credentials: Option<serde_json::Value>,
|
||||
pub expires_at: Option<u64>,
|
||||
pub account_id: Option<String>,
|
||||
pub account_user_id: Option<String>,
|
||||
@@ -156,8 +158,10 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
access_token.as_deref(),
|
||||
);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
parse_error: None,
|
||||
refresh_token,
|
||||
access_token,
|
||||
raw_credentials: None,
|
||||
expires_at: None,
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
@@ -176,6 +180,8 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
}
|
||||
}
|
||||
serde_json::Value::Object(object) => {
|
||||
let is_grok = provider_type.trim().eq_ignore_ascii_case("grok");
|
||||
let is_windsurf = provider_type.trim().eq_ignore_ascii_case("windsurf");
|
||||
let refresh_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("refresh_token")
|
||||
@@ -186,12 +192,8 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
.get("access_token")
|
||||
.or_else(|| object.get("accessToken")),
|
||||
);
|
||||
let grok_token_alias = if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
object.get("token")
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let grok_cookie = if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
let grok_token_alias = if is_grok { object.get("token") } else { None };
|
||||
let grok_cookie = if is_grok {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
object.get("cookie").or_else(|| object.get("cookieHeader")),
|
||||
)
|
||||
@@ -214,9 +216,43 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
refresh_token.as_deref(),
|
||||
access_token.as_deref().or(session_token.as_deref()),
|
||||
);
|
||||
if refresh_token.is_none() && access_token.is_none() {
|
||||
let windsurf_api_key = is_windsurf
|
||||
.then(|| {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
object.get("api_key").or_else(|| object.get("apiKey")),
|
||||
)
|
||||
})
|
||||
.flatten();
|
||||
let windsurf_token = is_windsurf
|
||||
.then(|| {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("token")
|
||||
.or_else(|| object.get("auth_token"))
|
||||
.or_else(|| object.get("authToken")),
|
||||
)
|
||||
})
|
||||
.flatten();
|
||||
let windsurf_password = is_windsurf
|
||||
.then(|| coerce_admin_provider_oauth_import_str(object.get("password")))
|
||||
.flatten();
|
||||
let raw_credentials = if is_windsurf
|
||||
&& (windsurf_api_key.is_some()
|
||||
|| windsurf_token.is_some()
|
||||
|| windsurf_password.is_some())
|
||||
{
|
||||
Some(item.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if refresh_token.is_none() && access_token.is_none() && raw_credentials.is_none() {
|
||||
return None;
|
||||
}
|
||||
let refresh_token = if is_windsurf {
|
||||
refresh_token.or(windsurf_api_key).or(windsurf_token)
|
||||
} else {
|
||||
refresh_token
|
||||
};
|
||||
let expires_at =
|
||||
json_u64_value(object.get("expires_at").or_else(|| object.get("expiresAt")));
|
||||
let account_id = coerce_admin_provider_oauth_import_str(
|
||||
@@ -308,8 +344,10 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
.or_else(|| object.get("impersonate")),
|
||||
);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
parse_error: None,
|
||||
refresh_token,
|
||||
access_token,
|
||||
raw_credentials,
|
||||
expires_at,
|
||||
account_id,
|
||||
account_user_id,
|
||||
@@ -340,14 +378,17 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
}
|
||||
|
||||
if raw.starts_with('[') {
|
||||
if let Ok(serde_json::Value::Array(items)) = serde_json::from_str::<serde_json::Value>(raw)
|
||||
{
|
||||
return items
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, item)
|
||||
})
|
||||
.collect();
|
||||
match serde_json::from_str::<serde_json::Value>(raw) {
|
||||
Ok(serde_json::Value::Array(items)) => {
|
||||
return items
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, item)
|
||||
})
|
||||
.collect();
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(error) => return vec![parse_error_entry(format!("JSON 数组解析失败: {error}"))],
|
||||
}
|
||||
}
|
||||
|
||||
@@ -364,23 +405,61 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
raw.lines()
|
||||
.map(str::trim)
|
||||
.filter(|line| !line.is_empty() && !line.starts_with('#'))
|
||||
.filter_map(|line| {
|
||||
if line.starts_with('{') {
|
||||
return serde_json::from_str::<serde_json::Value>(line)
|
||||
.ok()
|
||||
.and_then(|value| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
|
||||
});
|
||||
.filter_map(|token| {
|
||||
if is_json_like_batch_line(token) {
|
||||
match serde_json::from_str::<serde_json::Value>(token) {
|
||||
Ok(value @ serde_json::Value::Object(_)) => {
|
||||
return extract_admin_provider_oauth_batch_import_entry(
|
||||
provider_type,
|
||||
&value,
|
||||
);
|
||||
}
|
||||
Ok(_) => {
|
||||
return Some(parse_error_entry(
|
||||
"JSON 行必须是账号对象,不能作为 raw token 导入".to_string(),
|
||||
));
|
||||
}
|
||||
Err(error) => {
|
||||
return Some(parse_error_entry(format!("JSON 行解析失败: {error}")));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
extract_admin_provider_oauth_batch_import_entry(
|
||||
provider_type,
|
||||
&serde_json::Value::String(line.to_string()),
|
||||
&serde_json::Value::String(token.to_string()),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn parse_error_entry(error: String) -> AdminProviderOAuthBatchImportEntry {
|
||||
AdminProviderOAuthBatchImportEntry {
|
||||
parse_error: Some(error),
|
||||
refresh_token: None,
|
||||
access_token: None,
|
||||
raw_credentials: None,
|
||||
expires_at: None,
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
pool_tier: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
account_name: None,
|
||||
sso_rw_token: None,
|
||||
cf_cookies: None,
|
||||
cf_clearance: None,
|
||||
user_agent: None,
|
||||
browser_profile: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_json_like_batch_line(line: &str) -> bool {
|
||||
let line = line.trim_start();
|
||||
line.starts_with('{') || line.starts_with('[')
|
||||
}
|
||||
|
||||
pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
provider_type: &str,
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
@@ -705,4 +784,115 @@ mod tests {
|
||||
Some(&json!("project-gemini-cli-2"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_windsurf_json_credentials_for_native_import() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"windsurf",
|
||||
r#"[
|
||||
{"api_key":"devin-session-token$abc","email":"[email protected]"},
|
||||
{"token":"firebase-id-token","name":"Browser Login"},
|
||||
{"email":"[email protected]","password":"secret"},
|
||||
{"access_token":"devin-session-token$alias","email":"[email protected]"}
|
||||
]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 4);
|
||||
assert_eq!(
|
||||
entries[0].refresh_token.as_deref(),
|
||||
Some("devin-session-token$abc")
|
||||
);
|
||||
assert_eq!(entries[0].email.as_deref(), Some("[email protected]"));
|
||||
assert_eq!(
|
||||
entries[0]
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("api_key")),
|
||||
Some(&json!("devin-session-token$abc"))
|
||||
);
|
||||
assert_eq!(
|
||||
entries[1]
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("token")),
|
||||
Some(&json!("firebase-id-token"))
|
||||
);
|
||||
assert_eq!(
|
||||
entries[2]
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("password")),
|
||||
Some(&json!("secret"))
|
||||
);
|
||||
assert_eq!(
|
||||
entries[3].access_token.as_deref(),
|
||||
Some("devin-session-token$alias")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_windsurf_json_lines_credentials_for_native_import() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"windsurf",
|
||||
r#"{"api_key":"devin-session-token$abc","email":"[email protected]"}
|
||||
{"token":"firebase-id-token","name":"Browser Login"}
|
||||
{"email":"[email protected]","password":"secret"}"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 3);
|
||||
assert_eq!(
|
||||
entries[0]
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("api_key")),
|
||||
Some(&json!("devin-session-token$abc"))
|
||||
);
|
||||
assert_eq!(
|
||||
entries[1]
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("token")),
|
||||
Some(&json!("firebase-id-token"))
|
||||
);
|
||||
assert_eq!(
|
||||
entries[2]
|
||||
.raw_credentials
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("password")),
|
||||
Some(&json!("secret"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_json_line_is_parse_error_not_token() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"windsurf",
|
||||
r#"{"email":"[email protected]","password":"secret""#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert!(entries[0].parse_error.is_some());
|
||||
assert!(entries[0].refresh_token.is_none());
|
||||
assert!(entries[0].access_token.is_none());
|
||||
assert!(entries[0].raw_credentials.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn json_like_line_after_token_is_parse_error_not_token() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"windsurf",
|
||||
"devin-session-token$abc\n[not-json",
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 2);
|
||||
assert!(entries[0].parse_error.is_none());
|
||||
assert_eq!(
|
||||
entries[0].refresh_token.as_deref(),
|
||||
Some("devin-session-token$abc")
|
||||
);
|
||||
assert!(entries[1].parse_error.is_some());
|
||||
assert!(entries[1].refresh_token.is_none());
|
||||
assert!(entries[1].access_token.is_none());
|
||||
assert!(entries[1].raw_credentials.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,7 +10,7 @@ use super::progress::{
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
build_admin_provider_oauth_backend_unavailable_response,
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_provider_id;
|
||||
@@ -124,6 +124,13 @@ pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_tas
|
||||
"该 Provider 不是固定类型,无法使用 provider-oauth",
|
||||
));
|
||||
}
|
||||
if provider_type != "kiro"
|
||||
&& provider_type != "windsurf"
|
||||
&& admin_provider_oauth_template(&provider_type).is_none()
|
||||
{
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
|
||||
+147
-3
@@ -12,6 +12,7 @@ use crate::GatewayError;
|
||||
use aether_data::repository::provider_oauth::{
|
||||
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
|
||||
};
|
||||
use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext};
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -28,6 +29,8 @@ const KIRO_SOCIAL_MANUAL_CALLBACK_PORT: u16 = 49153;
|
||||
const KIRO_SOCIAL_ALLOWED_CALLBACK_PORTS: &[u16] = &[
|
||||
3128, 4649, 6588, 8008, 9091, 49153, 50153, 51153, 52153, 53153,
|
||||
];
|
||||
const WINDSURF_BROWSER_AUTH_EXPIRES_IN_SECS: u64 = 600;
|
||||
const WINDSURF_BROWSER_AUTH_POLL_INTERVAL_SECS: u64 = 5;
|
||||
|
||||
fn normalize_kiro_device_auth_type(raw: Option<&str>) -> String {
|
||||
match raw
|
||||
@@ -119,6 +122,26 @@ fn build_kiro_social_authorization_url(
|
||||
)
|
||||
}
|
||||
|
||||
fn build_windsurf_authorization_url(authorize_url: &str, login_option: &str) -> String {
|
||||
let login_option = login_option.trim();
|
||||
if login_option.is_empty() {
|
||||
return authorize_url.to_string();
|
||||
}
|
||||
if let Ok(mut url) = Url::parse(authorize_url) {
|
||||
url.query_pairs_mut()
|
||||
.append_pair("login_option", login_option);
|
||||
return url.to_string();
|
||||
}
|
||||
let separator = if authorize_url.contains('?') {
|
||||
'&'
|
||||
} else {
|
||||
'?'
|
||||
};
|
||||
let mut serializer = form_urlencoded::Serializer::new(String::new());
|
||||
serializer.append_pair("login_option", login_option);
|
||||
format!("{authorize_url}{separator}{}", serializer.finish())
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
@@ -163,14 +186,14 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
));
|
||||
};
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
if provider_type != "kiro" {
|
||||
if provider_type != "kiro" && provider_type != "windsurf" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"设备授权仅支持 Kiro provider",
|
||||
"设备授权仅支持 Kiro / Windsurf provider",
|
||||
));
|
||||
}
|
||||
let endpoint_resolution =
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, "kiro").await?;
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
|
||||
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
|
||||
let request_proxy = state
|
||||
.resolve_admin_provider_oauth_operation_proxy_snapshot(
|
||||
@@ -184,6 +207,103 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
)
|
||||
.await;
|
||||
|
||||
if provider_type == "windsurf" {
|
||||
let session_id = generate_provider_oauth_nonce();
|
||||
let login_option = payload
|
||||
.login_option
|
||||
.as_deref()
|
||||
.or(payload.auth_type.as_deref())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("default")
|
||||
.to_ascii_lowercase();
|
||||
let ctx = ProviderOAuthTransportContext {
|
||||
provider_id: provider_id.clone(),
|
||||
provider_type: provider_type.clone(),
|
||||
endpoint_id: runtime_endpoint
|
||||
.as_ref()
|
||||
.map(|endpoint| endpoint.id.clone()),
|
||||
key_id: None,
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: provider.config.clone(),
|
||||
endpoint_config: runtime_endpoint
|
||||
.as_ref()
|
||||
.and_then(|endpoint| endpoint.config.clone()),
|
||||
key_config: None,
|
||||
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||
request_proxy.clone(),
|
||||
),
|
||||
};
|
||||
let mut authorization = match ProviderOAuthService::with_builtin_adapters()
|
||||
.build_authorize_url(&ctx, &session_id, None)
|
||||
{
|
||||
Ok(authorization) => authorization,
|
||||
Err(error) => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
format!("Windsurf 授权 URL 构建失败: {error}"),
|
||||
));
|
||||
}
|
||||
};
|
||||
authorization.authorize_url =
|
||||
build_windsurf_authorization_url(&authorization.authorize_url, &login_option);
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let session = StoredAdminProviderOAuthDeviceSession {
|
||||
provider_id: provider_id.clone(),
|
||||
region: String::new(),
|
||||
client_id: String::new(),
|
||||
client_secret: String::new(),
|
||||
device_code: String::new(),
|
||||
auth_type: Some("browser".to_string()),
|
||||
social_provider: Some(login_option.clone()),
|
||||
code_verifier: None,
|
||||
redirect_uri: Some("show-auth-token".to_string()),
|
||||
machine_id: Some(uuid::Uuid::new_v4().to_string().to_ascii_lowercase()),
|
||||
interval: WINDSURF_BROWSER_AUTH_POLL_INTERVAL_SECS,
|
||||
expires_at_unix_secs: now_unix_secs
|
||||
.saturating_add(WINDSURF_BROWSER_AUTH_EXPIRES_IN_SECS),
|
||||
status: "pending".to_string(),
|
||||
proxy_node_id: payload
|
||||
.proxy_node_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
created_at_unix_ms: now_unix_secs,
|
||||
key_id: None,
|
||||
email: None,
|
||||
replaced: false,
|
||||
error_msg: None,
|
||||
};
|
||||
if let Err(response) = state
|
||||
.save_provider_oauth_device_session(
|
||||
&session_id,
|
||||
&session,
|
||||
WINDSURF_BROWSER_AUTH_EXPIRES_IN_SECS
|
||||
.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Ok(response);
|
||||
}
|
||||
|
||||
return Ok(Json(json!({
|
||||
"session_id": session_id,
|
||||
"user_code": "",
|
||||
"verification_uri": "https://windsurf.com/windsurf/signin",
|
||||
"verification_uri_complete": authorization.authorize_url,
|
||||
"expires_in": WINDSURF_BROWSER_AUTH_EXPIRES_IN_SECS,
|
||||
"interval": WINDSURF_BROWSER_AUTH_POLL_INTERVAL_SECS,
|
||||
"auth_type": "browser",
|
||||
"login_option": login_option,
|
||||
"redirect_uri": "show-auth-token",
|
||||
"callback_required": true,
|
||||
}))
|
||||
.into_response());
|
||||
}
|
||||
|
||||
let auth_type = normalize_kiro_device_auth_type(payload.auth_type.as_deref());
|
||||
if let Some(social_provider) = kiro_social_provider_id(&auth_type) {
|
||||
let redirect_uri = match normalize_kiro_social_redirect_uri(payload.redirect_uri.as_deref())
|
||||
@@ -394,3 +514,27 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::build_windsurf_authorization_url;
|
||||
|
||||
#[test]
|
||||
fn windsurf_authorization_url_includes_login_option() {
|
||||
let url = build_windsurf_authorization_url(
|
||||
"https://windsurf.com/windsurf/signin?state=session-1",
|
||||
"github",
|
||||
);
|
||||
|
||||
let parsed = url::Url::parse(&url).expect("url should parse");
|
||||
let params = parsed
|
||||
.query_pairs()
|
||||
.map(|(key, value)| (key.to_string(), value.to_string()))
|
||||
.collect::<std::collections::BTreeMap<_, _>>();
|
||||
assert_eq!(params.get("state").map(String::as_str), Some("session-1"));
|
||||
assert_eq!(
|
||||
params.get("login_option").map(String::as_str),
|
||||
Some("github")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -26,6 +26,9 @@ use aether_data::repository::provider_oauth::StoredAdminProviderOAuthDeviceSessi
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_oauth::provider::{
|
||||
ProviderOAuthImportInput, ProviderOAuthService, ProviderOAuthTransportContext,
|
||||
};
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -86,6 +89,99 @@ fn kiro_social_poll_error_response(error: impl Into<String>) -> Response<Body> {
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn windsurf_browser_poll_error_response(error: impl Into<String>) -> Response<Body> {
|
||||
Json(json!({
|
||||
"status": "error",
|
||||
"error": error.into(),
|
||||
"replaced": false,
|
||||
}))
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_browser_poll_detail(detail: impl AsRef<str>) -> String {
|
||||
let detail = detail.as_ref().trim();
|
||||
if detail.is_empty() {
|
||||
return "-".to_string();
|
||||
}
|
||||
if contains_windsurf_sensitive_marker(detail) {
|
||||
"[REDACTED upstream error body]".to_string()
|
||||
} else {
|
||||
detail.chars().take(500).collect()
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_browser_poll_callback_error(error: &str, description: &str) -> String {
|
||||
let error = sanitize_windsurf_browser_poll_error_code(error);
|
||||
let description = sanitize_windsurf_browser_poll_detail(description);
|
||||
format!("{error}: {description}")
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_browser_poll_error_code(error: &str) -> String {
|
||||
let error = error.trim();
|
||||
if !error.is_empty()
|
||||
&& error.len() <= 80
|
||||
&& error
|
||||
.chars()
|
||||
.all(|ch| ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-' | '.'))
|
||||
{
|
||||
return error.to_string();
|
||||
}
|
||||
sanitize_windsurf_browser_poll_detail(error)
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_browser_poll_oauth_error(error: &aether_oauth::core::OAuthError) -> String {
|
||||
match error {
|
||||
aether_oauth::core::OAuthError::InvalidRequest(_) => {
|
||||
"Windsurf token 验证失败: 请求参数无效".to_string()
|
||||
}
|
||||
aether_oauth::core::OAuthError::HttpStatus { status_code, .. } => {
|
||||
format!("Windsurf token 验证失败: HTTP {status_code}")
|
||||
}
|
||||
_ => "Windsurf token 验证失败".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn contains_windsurf_sensitive_marker(value: &str) -> bool {
|
||||
let lowered = value.to_ascii_lowercase();
|
||||
[
|
||||
"token",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"sessiontoken",
|
||||
"firebase_id_token",
|
||||
"idtoken",
|
||||
"authorization",
|
||||
"password",
|
||||
"secret",
|
||||
"devin-session-token$",
|
||||
]
|
||||
.iter()
|
||||
.any(|marker| lowered.contains(marker))
|
||||
|| value.contains("sk-")
|
||||
}
|
||||
|
||||
fn secret_fingerprint(value: &str) -> Option<String> {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
use sha2::{Digest, Sha256};
|
||||
let digest = Sha256::digest(value.as_bytes());
|
||||
Some(
|
||||
digest[..8]
|
||||
.iter()
|
||||
.map(|byte| format!("{byte:02x}"))
|
||||
.collect::<String>(),
|
||||
)
|
||||
}
|
||||
|
||||
fn insert_secret_fingerprint(target: &mut serde_json::Map<String, Value>, key: &str, secret: &str) {
|
||||
if let Some(fingerprint) = secret_fingerprint(secret) {
|
||||
target.insert(key.to_string(), json!(fingerprint));
|
||||
}
|
||||
}
|
||||
|
||||
fn kiro_social_provider_from_login_option(login_option: Option<&str>) -> Option<&'static str> {
|
||||
match login_option
|
||||
.map(str::trim)
|
||||
@@ -322,8 +418,9 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
|
||||
"Provider 不存在",
|
||||
));
|
||||
};
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
let endpoint_resolution =
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, "kiro").await?;
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
|
||||
let endpoints = endpoint_resolution.endpoints;
|
||||
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
|
||||
let request_proxy = state
|
||||
@@ -338,6 +435,20 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
|
||||
)
|
||||
.await;
|
||||
|
||||
if provider_type == "windsurf" {
|
||||
return handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||
state,
|
||||
&provider,
|
||||
&endpoints,
|
||||
request_proxy,
|
||||
session_id,
|
||||
session,
|
||||
payload.callback_url.as_deref(),
|
||||
payload.token.as_deref(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
if kiro_device_session_is_social(&session) {
|
||||
return handle_admin_provider_oauth_kiro_social_device_poll(
|
||||
state,
|
||||
@@ -630,6 +741,276 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
|
||||
))
|
||||
}
|
||||
|
||||
fn windsurf_raw_api_key(value: &str) -> Option<&str> {
|
||||
let value = value.trim();
|
||||
if value.starts_with("devin-session-token$") || value.starts_with("sk-") {
|
||||
Some(value)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_admin_provider_oauth_windsurf_browser_device_poll(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
session_id: &str,
|
||||
mut session: StoredAdminProviderOAuthDeviceSession,
|
||||
callback_url: Option<&str>,
|
||||
token: Option<&str>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let callback_url = callback_url
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
let token = token.map(str::trim).filter(|value| !value.is_empty());
|
||||
if callback_url.is_none() && token.is_none() {
|
||||
return Ok(Json(json!({"status": "pending", "replaced": false})).into_response());
|
||||
}
|
||||
|
||||
let mut social_provider = session
|
||||
.social_provider
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let imported_token = if let Some(callback_url) = callback_url {
|
||||
let callback_params = parse_provider_oauth_callback_params(callback_url);
|
||||
if let Some(error) = callback_params.get("error").map(String::as_str) {
|
||||
let error_description = callback_params
|
||||
.get("error_description")
|
||||
.map(String::as_str)
|
||||
.unwrap_or("用户拒绝授权");
|
||||
let sanitized_error =
|
||||
sanitize_windsurf_browser_poll_callback_error(error, error_description);
|
||||
session.status = "error".to_string();
|
||||
session.error_msg = Some(sanitized_error.clone());
|
||||
let _ = state
|
||||
.save_provider_oauth_device_session(session_id, &session, 30)
|
||||
.await;
|
||||
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
|
||||
session_id,
|
||||
"error",
|
||||
windsurf_browser_poll_error_response(sanitized_error),
|
||||
));
|
||||
}
|
||||
let Some(callback_state) = callback_params
|
||||
.get("state")
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(windsurf_browser_poll_error_response("回调 URL 缺少 state"));
|
||||
};
|
||||
if callback_state != session_id {
|
||||
return Ok(windsurf_browser_poll_error_response(
|
||||
"回调 state 与会话不匹配",
|
||||
));
|
||||
}
|
||||
if let Some(provider) = callback_params
|
||||
.get("provider")
|
||||
.or_else(|| callback_params.get("login_option"))
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
social_provider = Some(provider.to_string());
|
||||
}
|
||||
let Some(callback_token) = callback_params
|
||||
.get("token")
|
||||
.or_else(|| callback_params.get("auth_token"))
|
||||
.or_else(|| callback_params.get("access_token"))
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(windsurf_browser_poll_error_response("回调 URL 缺少 token"));
|
||||
};
|
||||
callback_token.to_string()
|
||||
} else {
|
||||
let token = token.unwrap_or_default();
|
||||
if windsurf_raw_api_key(token).is_none() {
|
||||
return Ok(windsurf_browser_poll_error_response(
|
||||
"浏览器授权请提交包含 state 的回调 URL;纯 token 请使用导入授权",
|
||||
));
|
||||
}
|
||||
token.to_string()
|
||||
};
|
||||
|
||||
let mut raw_credentials = serde_json::Map::new();
|
||||
if windsurf_raw_api_key(&imported_token).is_some() {
|
||||
raw_credentials.insert("api_key".to_string(), json!(imported_token));
|
||||
} else {
|
||||
raw_credentials.insert("token".to_string(), json!(imported_token));
|
||||
}
|
||||
if let Some(social_provider) = social_provider.as_ref() {
|
||||
raw_credentials.insert("social_provider".to_string(), json!(social_provider));
|
||||
}
|
||||
|
||||
let ctx = ProviderOAuthTransportContext {
|
||||
provider_id: provider.id.clone(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: provider.config.clone(),
|
||||
endpoint_config: None,
|
||||
key_config: None,
|
||||
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||
request_proxy.clone(),
|
||||
),
|
||||
};
|
||||
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||
let result = match ProviderOAuthService::with_builtin_adapters()
|
||||
.import_credentials(
|
||||
&executor,
|
||||
&ctx,
|
||||
ProviderOAuthImportInput {
|
||||
provider_type: provider.provider_type.clone(),
|
||||
name: None,
|
||||
refresh_token: None,
|
||||
raw_credentials: Some(Value::Object(raw_credentials)),
|
||||
network: ctx.network.clone(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(error) => {
|
||||
let sanitized_error = sanitize_windsurf_browser_poll_oauth_error(&error);
|
||||
session.status = "error".to_string();
|
||||
session.error_msg = Some(sanitized_error.clone());
|
||||
let _ = state
|
||||
.save_provider_oauth_device_session(session_id, &session, 30)
|
||||
.await;
|
||||
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
|
||||
session_id,
|
||||
"error",
|
||||
windsurf_browser_poll_error_response(sanitized_error),
|
||||
));
|
||||
}
|
||||
};
|
||||
let access_token = result.token_set.access_token.trim().to_string();
|
||||
if access_token.is_empty() {
|
||||
return Ok(windsurf_browser_poll_error_response(
|
||||
"Windsurf token 验证返回缺少 apiKey/sessionToken",
|
||||
));
|
||||
}
|
||||
let mut auth_config = result.auth_config.as_object().cloned().unwrap_or_default();
|
||||
auth_config.insert("provider_type".to_string(), json!("windsurf"));
|
||||
auth_config.insert("auth_method".to_string(), json!("browser"));
|
||||
if let Some(social_provider) = social_provider.as_ref() {
|
||||
auth_config
|
||||
.entry("social_provider".to_string())
|
||||
.or_insert_with(|| json!(social_provider));
|
||||
}
|
||||
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&provider.id, &auth_config, None)
|
||||
.await
|
||||
{
|
||||
Ok(duplicate) => duplicate,
|
||||
Err(detail) => {
|
||||
return Ok(Json(json!({
|
||||
"status": "error",
|
||||
"error": detail,
|
||||
"replaced": false,
|
||||
}))
|
||||
.into_response());
|
||||
}
|
||||
};
|
||||
|
||||
let api_formats = provider_oauth_active_api_formats(endpoints);
|
||||
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
|
||||
let expires_at = result.token_set.expires_at_unix_secs;
|
||||
let email = auth_config
|
||||
.get("email")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let mut replaced = false;
|
||||
let persisted_key = if let Some(existing_key) = duplicate {
|
||||
replaced = true;
|
||||
match state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&existing_key,
|
||||
&provider.provider_type,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&api_formats,
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_name = email
|
||||
.as_deref()
|
||||
.map(|email| format!("windsurf_{email}"))
|
||||
.unwrap_or_else(|| format!("windsurf_{}", current_unix_secs()));
|
||||
match state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&provider.id,
|
||||
&provider.provider_type,
|
||||
&key_name,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&api_formats,
|
||||
key_proxy,
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth write unavailable",
|
||||
));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
provider.clone(),
|
||||
persisted_key.id.clone(),
|
||||
request_proxy.clone(),
|
||||
);
|
||||
|
||||
session.status = "authorized".to_string();
|
||||
session.key_id = Some(persisted_key.id.clone());
|
||||
session.email = email.clone();
|
||||
session.replaced = replaced;
|
||||
session.error_msg = None;
|
||||
let _ = state
|
||||
.save_provider_oauth_device_session(session_id, &session, 60)
|
||||
.await;
|
||||
|
||||
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
|
||||
session_id,
|
||||
"authorized",
|
||||
Json(json!({
|
||||
"status": "authorized",
|
||||
"key_id": persisted_key.id,
|
||||
"email": email,
|
||||
"replaced": replaced,
|
||||
}))
|
||||
.into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn handle_admin_provider_oauth_kiro_social_device_poll(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -814,7 +1195,7 @@ async fn handle_admin_provider_oauth_kiro_social_device_poll(
|
||||
.get("idToken")
|
||||
.or_else(|| token_result.get("id_token")),
|
||||
) {
|
||||
auth_config_object.insert("id_token".to_string(), json!(id_token));
|
||||
insert_secret_fingerprint(&mut auth_config_object, "id_token_fingerprint", &id_token);
|
||||
}
|
||||
if let Some(token_type) = json_non_empty_string(
|
||||
token_result
|
||||
@@ -921,3 +1302,18 @@ async fn handle_admin_provider_oauth_kiro_social_device_poll(
|
||||
.into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
#[test]
|
||||
fn windsurf_browser_poll_callback_error_redacts_sensitive_values() {
|
||||
let detail = super::sanitize_windsurf_browser_poll_callback_error(
|
||||
"access_denied",
|
||||
"bad token devin-session-token$secret and apiKey sk-secret",
|
||||
);
|
||||
|
||||
assert_eq!(detail, "access_denied: [REDACTED upstream error body]");
|
||||
assert!(!detail.contains("devin-session-token$secret"));
|
||||
assert!(!detail.contains("sk-secret"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ pub(super) struct AdminProviderOAuthDeviceAuthorizePayload {
|
||||
#[serde(default = "default_kiro_device_region")]
|
||||
pub(super) region: String,
|
||||
pub(super) auth_type: Option<String>,
|
||||
pub(super) login_option: Option<String>,
|
||||
pub(super) redirect_uri: Option<String>,
|
||||
pub(super) proxy_node_id: Option<String>,
|
||||
}
|
||||
@@ -20,6 +21,7 @@ pub(super) struct AdminProviderOAuthDeviceAuthorizePayload {
|
||||
pub(super) struct AdminProviderOAuthDevicePollPayload {
|
||||
pub(super) session_id: String,
|
||||
pub(super) callback_url: Option<String>,
|
||||
pub(super) token: Option<String>,
|
||||
}
|
||||
|
||||
pub(super) fn attach_admin_provider_oauth_device_poll_terminal_response(
|
||||
|
||||
@@ -25,6 +25,10 @@ use crate::handlers::admin::request::{
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_oauth::core::OAuthError;
|
||||
use aether_oauth::provider::{
|
||||
ProviderOAuthImportInput, ProviderOAuthService, ProviderOAuthTransportContext,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -39,6 +43,49 @@ struct AdminProviderOAuthSingleImportTokens {
|
||||
expires_at: Option<u64>,
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_import_error(error: &OAuthError) -> String {
|
||||
match error {
|
||||
OAuthError::InvalidRequest(_) => "Windsurf 凭据验证失败: 请求参数无效".to_string(),
|
||||
OAuthError::HttpStatus { status_code, .. } => {
|
||||
format!("Windsurf 凭据验证失败: HTTP {status_code}")
|
||||
}
|
||||
OAuthError::InvalidResponse(detail) => sanitize_windsurf_invalid_response_detail(detail)
|
||||
.unwrap_or_else(|| "Windsurf 凭据验证失败".to_string()),
|
||||
_ => "Windsurf 凭据验证失败".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_invalid_response_detail(detail: &str) -> Option<String> {
|
||||
let detail = detail.trim();
|
||||
if detail.eq_ignore_ascii_case("Auth1 response is not json") {
|
||||
return Some("Windsurf 凭据验证失败: Auth1 响应无法解析".to_string());
|
||||
}
|
||||
if detail.eq_ignore_ascii_case("Auth1 response missing token") {
|
||||
return Some("Windsurf 凭据验证失败: Auth1 响应缺少 token".to_string());
|
||||
}
|
||||
if detail.contains("WindsurfPostAuth response missing sessionToken")
|
||||
|| detail.contains("WindsurfPostAuth response is not json")
|
||||
|| (detail.contains("WindsurfPostAuth failed") && detail.contains("missing sessionToken"))
|
||||
{
|
||||
return Some("Windsurf 凭据验证失败: PostAuth 未返回 sessionToken".to_string());
|
||||
}
|
||||
if detail.contains("WindsurfPostAuth failed") {
|
||||
return Some("Windsurf 凭据验证失败: PostAuth 失败".to_string());
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn import_payload_has_windsurf_credentials(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> bool {
|
||||
import_payload_string(payload, "api_key", "apiKey").is_some()
|
||||
|| import_payload_string_any(payload, &["token", "auth_token", "authToken"]).is_some()
|
||||
|| import_payload_string(payload, "refresh_token", "refreshToken").is_some()
|
||||
|| import_payload_string(payload, "access_token", "accessToken").is_some()
|
||||
|| (import_payload_string_any(payload, &["email"]).is_some()
|
||||
&& import_payload_string_any(payload, &["password"]).is_some())
|
||||
}
|
||||
|
||||
fn import_payload_string(
|
||||
payload: &serde_json::Map<String, serde_json::Value>,
|
||||
snake_case: &str,
|
||||
@@ -266,6 +313,70 @@ async fn resolve_admin_provider_oauth_single_import_tokens(
|
||||
})
|
||||
}
|
||||
|
||||
async fn resolve_admin_provider_oauth_windsurf_single_import_tokens(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_type: &str,
|
||||
name: Option<String>,
|
||||
raw_payload: &serde_json::Map<String, serde_json::Value>,
|
||||
refresh_token: Option<&str>,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
) -> Result<AdminProviderOAuthSingleImportTokens, Response<Body>> {
|
||||
let ctx = ProviderOAuthTransportContext {
|
||||
provider_id: String::new(),
|
||||
provider_type: provider_type.to_string(),
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
auth_type: Some("oauth".to_string()),
|
||||
decrypted_api_key: None,
|
||||
decrypted_auth_config: None,
|
||||
provider_config: None,
|
||||
endpoint_config: None,
|
||||
key_config: None,
|
||||
network: aether_oauth::network::OAuthNetworkContext::provider_operation(
|
||||
request_proxy.clone(),
|
||||
),
|
||||
};
|
||||
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
|
||||
let service = ProviderOAuthService::with_builtin_adapters();
|
||||
let result = service
|
||||
.import_credentials(
|
||||
&executor,
|
||||
&ctx,
|
||||
ProviderOAuthImportInput {
|
||||
provider_type: provider_type.to_string(),
|
||||
name,
|
||||
refresh_token: refresh_token.map(ToOwned::to_owned),
|
||||
raw_credentials: Some(serde_json::Value::Object(raw_payload.clone())),
|
||||
network: ctx.network.clone(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
sanitize_windsurf_import_error(&error),
|
||||
)
|
||||
})?;
|
||||
let access_token = result.token_set.access_token.trim().to_string();
|
||||
if access_token.is_empty() {
|
||||
return Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Windsurf 凭据验证返回缺少 apiKey/sessionToken",
|
||||
));
|
||||
}
|
||||
let auth_config = result.auth_config.as_object().cloned().ok_or_else(|| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Windsurf 凭据验证返回缺少 auth_config",
|
||||
)
|
||||
})?;
|
||||
Ok(AdminProviderOAuthSingleImportTokens {
|
||||
access_token,
|
||||
auth_config,
|
||||
expires_at: result.token_set.expires_at_unix_secs,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
@@ -370,19 +481,47 @@ 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 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(value) => value,
|
||||
Err(response) => return Ok(response),
|
||||
let resolved_import = if provider_type == "windsurf" {
|
||||
if !import_payload_has_windsurf_credentials(&raw_payload) {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Windsurf 凭据不能为空",
|
||||
));
|
||||
}
|
||||
match resolve_admin_provider_oauth_windsurf_single_import_tokens(
|
||||
state,
|
||||
&provider_type,
|
||||
name.clone(),
|
||||
&raw_payload,
|
||||
refresh_token_input.as_deref(),
|
||||
request_proxy.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(value) => value,
|
||||
Err(response) => return Ok(response),
|
||||
}
|
||||
} else {
|
||||
if refresh_token_input.is_none() && access_token_input.is_none() {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Refresh Token 或 Access Token 不能为空",
|
||||
));
|
||||
}
|
||||
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(value) => value,
|
||||
Err(response) => return Ok(response),
|
||||
}
|
||||
};
|
||||
let AdminProviderOAuthSingleImportTokens {
|
||||
access_token,
|
||||
@@ -480,3 +619,47 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::sanitize_windsurf_import_error;
|
||||
use aether_oauth::core::OAuthError;
|
||||
|
||||
#[test]
|
||||
fn windsurf_import_error_redacts_http_body() {
|
||||
let error = OAuthError::HttpStatus {
|
||||
status_code: 400,
|
||||
body_excerpt: "token=secret-token password=secret-password".to_string(),
|
||||
};
|
||||
|
||||
let detail = sanitize_windsurf_import_error(&error);
|
||||
|
||||
assert_eq!(detail, "Windsurf 凭据验证失败: HTTP 400");
|
||||
assert!(!detail.contains("secret-token"));
|
||||
assert!(!detail.contains("secret-password"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_import_error_redacts_invalid_response_detail() {
|
||||
let error =
|
||||
OAuthError::invalid_response("RegisterUser failed with firebase_id_token=secret-token");
|
||||
|
||||
let detail = sanitize_windsurf_import_error(&error);
|
||||
|
||||
assert_eq!(detail, "Windsurf 凭据验证失败");
|
||||
assert!(!detail.contains("secret-token"));
|
||||
assert!(!detail.contains("firebase_id_token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_import_error_keeps_safe_post_auth_stage() {
|
||||
let error = OAuthError::invalid_response("WindsurfPostAuth response missing sessionToken");
|
||||
|
||||
let detail = sanitize_windsurf_import_error(&error);
|
||||
|
||||
assert_eq!(
|
||||
detail,
|
||||
"Windsurf 凭据验证失败: PostAuth 未返回 sessionToken"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -63,6 +63,12 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
|
||||
"该 Provider 不是固定类型,无法使用 provider-oauth",
|
||||
));
|
||||
}
|
||||
if provider_type == "windsurf" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Windsurf 请使用浏览器登录或导入凭据。",
|
||||
));
|
||||
}
|
||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
@@ -134,6 +140,12 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
|
||||
"Kiro 不支持 OAuth 授权,请使用导入授权。",
|
||||
));
|
||||
}
|
||||
if provider_type == "windsurf" {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Windsurf 请使用浏览器登录或导入凭据。",
|
||||
));
|
||||
}
|
||||
let Some(template) = admin_provider_oauth_template(&provider_type) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
|
||||
@@ -36,6 +36,13 @@ fn is_openai_provider_oauth_provider_type(value: Option<&serde_json::Value>) ->
|
||||
})
|
||||
}
|
||||
|
||||
fn is_windsurf_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("windsurf"))
|
||||
}
|
||||
|
||||
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>,
|
||||
@@ -106,6 +113,41 @@ fn match_codex_provider_oauth_identity(
|
||||
None
|
||||
}
|
||||
|
||||
fn match_windsurf_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> {
|
||||
let new_provider_type = new_auth_config.get("provider_type");
|
||||
let existing_provider_type = existing_auth_config.get("provider_type");
|
||||
if !is_windsurf_provider_oauth_provider_type(new_provider_type)
|
||||
&& !is_windsurf_provider_oauth_provider_type(existing_provider_type)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let new_account_id = normalize_provider_oauth_identity_value(new_auth_config.get("account_id"));
|
||||
let existing_account_id =
|
||||
normalize_provider_oauth_identity_value(existing_auth_config.get("account_id"));
|
||||
if let (Some(new_account_id), Some(existing_account_id)) =
|
||||
(new_account_id.as_deref(), existing_account_id.as_deref())
|
||||
{
|
||||
return Some(new_account_id == existing_account_id);
|
||||
}
|
||||
|
||||
let new_credential_fingerprint =
|
||||
normalize_provider_oauth_identity_value(new_auth_config.get("credential_fingerprint"));
|
||||
let existing_credential_fingerprint =
|
||||
normalize_provider_oauth_identity_value(existing_auth_config.get("credential_fingerprint"));
|
||||
if let (Some(new_fingerprint), Some(existing_fingerprint)) = (
|
||||
new_credential_fingerprint.as_deref(),
|
||||
existing_credential_fingerprint.as_deref(),
|
||||
) {
|
||||
return Some(new_fingerprint == existing_fingerprint);
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
fn is_codex_cross_plan_group_non_duplicate(
|
||||
new_auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
@@ -169,10 +211,17 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
) -> Result<Option<StoredProviderCatalogKey>, String> {
|
||||
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_credential_fingerprint =
|
||||
normalize_provider_oauth_identity_value(auth_config.get("credential_fingerprint"));
|
||||
let new_auth_method = normalize_provider_oauth_identity_value(auth_config.get("auth_method"));
|
||||
let new_kiro_provider = normalize_provider_oauth_identity_value(auth_config.get("provider"));
|
||||
|
||||
if new_email.is_none() && new_user_id.is_none() {
|
||||
if new_email.is_none()
|
||||
&& new_user_id.is_none()
|
||||
&& new_account_id.is_none()
|
||||
&& new_credential_fingerprint.is_none()
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
@@ -201,15 +250,28 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
normalize_provider_oauth_identity_value(existing_auth_config.get("auth_method"));
|
||||
let existing_kiro_provider =
|
||||
normalize_provider_oauth_identity_value(existing_auth_config.get("provider"));
|
||||
let is_windsurf = auth_config
|
||||
.get("provider_type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("windsurf"))
|
||||
|| existing_auth_config
|
||||
.get("provider_type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("windsurf"));
|
||||
|
||||
let mut is_duplicate = false;
|
||||
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 {
|
||||
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()
|
||||
&& windsurf_identity_match.is_none()
|
||||
&& !is_duplicate
|
||||
&& new_user_id.is_some()
|
||||
&& existing_user_id.is_some()
|
||||
@@ -220,7 +282,9 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
}
|
||||
|
||||
if codex_identity_match.is_none()
|
||||
&& windsurf_identity_match.is_none()
|
||||
&& !is_duplicate
|
||||
&& !is_windsurf
|
||||
&& new_email.is_some()
|
||||
&& existing_email.is_some()
|
||||
&& new_email == existing_email
|
||||
@@ -261,6 +325,12 @@ pub(crate) async fn find_duplicate_provider_oauth_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(|| {
|
||||
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();
|
||||
@@ -272,3 +342,91 @@ pub(crate) async fn find_duplicate_provider_oauth_key(
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::match_windsurf_provider_oauth_identity;
|
||||
use serde_json::{json, Map, Value};
|
||||
|
||||
fn auth_config(value: Value) -> Map<String, Value> {
|
||||
value.as_object().cloned().expect("auth config object")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_identity_matches_account_id_without_email() {
|
||||
let new_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"auth_method": "api_key",
|
||||
"account_id": "acct-ws-1"
|
||||
}));
|
||||
let existing_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"auth_method": "browser",
|
||||
"account_id": "acct-ws-1"
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
match_windsurf_provider_oauth_identity(&new_auth_config, &existing_auth_config),
|
||||
Some(true)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_identity_rejects_different_account_id() {
|
||||
let new_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"account_id": "acct-ws-1",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
let existing_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"account_id": "acct-ws-2",
|
||||
"email": "[email protected]"
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
match_windsurf_provider_oauth_identity(&new_auth_config, &existing_auth_config),
|
||||
Some(false)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_identity_matches_credential_fingerprint_without_profile() {
|
||||
let new_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"auth_method": "api_key",
|
||||
"credential_fingerprint": "abcdef0123456789"
|
||||
}));
|
||||
let existing_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"auth_method": "browser",
|
||||
"credential_fingerprint": "abcdef0123456789"
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
match_windsurf_provider_oauth_identity(&new_auth_config, &existing_auth_config),
|
||||
Some(true)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_identity_does_not_match_user_supplied_email_only() {
|
||||
let new_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"auth_method": "api_key",
|
||||
"email": "[email protected]",
|
||||
"email_verified": false
|
||||
}));
|
||||
let existing_auth_config = auth_config(json!({
|
||||
"provider_type": "windsurf",
|
||||
"auth_method": "api_key",
|
||||
"email": "[email protected]",
|
||||
"email_verified": false
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
match_windsurf_provider_oauth_identity(&new_auth_config, &existing_auth_config),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ use super::codex::refresh_codex_provider_quota_locally;
|
||||
use super::gemini_cli::refresh_gemini_cli_provider_quota_locally;
|
||||
use super::grok::refresh_grok_provider_quota_locally;
|
||||
use super::kiro::refresh_kiro_provider_quota_locally;
|
||||
use super::windsurf::refresh_windsurf_provider_quota_locally;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
@@ -41,6 +42,7 @@ const PROVIDER_QUOTA_REFRESH_HANDLERS: &[(&str, ProviderQuotaRefreshHandler)] =
|
||||
),
|
||||
("grok", refresh_grok_provider_quota_locally_boxed),
|
||||
("kiro", refresh_kiro_provider_quota_locally_boxed),
|
||||
("windsurf", refresh_windsurf_provider_quota_locally_boxed),
|
||||
];
|
||||
|
||||
pub(crate) async fn refresh_provider_pool_quota_locally(
|
||||
@@ -156,3 +158,19 @@ fn refresh_grok_provider_quota_locally_boxed<'a>(
|
||||
proxy_override,
|
||||
))
|
||||
}
|
||||
|
||||
fn refresh_windsurf_provider_quota_locally_boxed<'a>(
|
||||
state: &'a AdminAppState<'a>,
|
||||
provider: &'a StoredProviderCatalogProvider,
|
||||
endpoint: &'a StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> ProviderQuotaRefreshFuture<'a> {
|
||||
Box::pin(refresh_windsurf_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
endpoint,
|
||||
keys,
|
||||
proxy_override,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -6,3 +6,4 @@ pub(crate) mod gemini_cli;
|
||||
pub(crate) mod grok;
|
||||
pub(crate) mod kiro;
|
||||
pub(crate) mod shared;
|
||||
pub(crate) mod windsurf;
|
||||
|
||||
@@ -0,0 +1,632 @@
|
||||
use super::shared::{
|
||||
build_provider_quota_execution_plan, build_quota_snapshot_payload,
|
||||
default_provider_quota_execution_timeouts, execute_provider_quota_plan,
|
||||
extract_execution_error_message, persist_provider_quota_refresh_state,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_provider_pool::{
|
||||
build_windsurf_pool_model_configs_request_with_base_url,
|
||||
build_windsurf_pool_quota_request_with_base_url,
|
||||
build_windsurf_pool_rate_limit_request_with_base_url, ProviderPoolQuotaRequestSpec,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
async fn execute_windsurf_probe_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
spec: ProviderPoolQuotaRequestSpec,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
quota_kind: &str,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
let proxy = match proxy_override {
|
||||
Some(proxy) => Some(proxy.clone()),
|
||||
None => {
|
||||
state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await
|
||||
}
|
||||
};
|
||||
let timeouts = state
|
||||
.resolve_transport_execution_timeouts(transport)
|
||||
.or(Some(default_provider_quota_execution_timeouts(
|
||||
proxy.as_ref(),
|
||||
)));
|
||||
let plan = build_provider_quota_execution_plan(
|
||||
transport,
|
||||
spec,
|
||||
proxy,
|
||||
state.resolve_transport_profile(transport),
|
||||
timeouts,
|
||||
);
|
||||
|
||||
execute_provider_quota_plan(state, transport, plan, quota_kind).await
|
||||
}
|
||||
|
||||
async fn execute_windsurf_user_status_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_key: &str,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
let spec = build_windsurf_pool_quota_request_with_base_url(
|
||||
&transport.key.id,
|
||||
&transport.endpoint.base_url,
|
||||
api_key,
|
||||
);
|
||||
execute_windsurf_probe_plan(
|
||||
state,
|
||||
transport,
|
||||
spec,
|
||||
proxy_override,
|
||||
"windsurf:user_status",
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn execute_windsurf_model_configs_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_key: &str,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
let spec = build_windsurf_pool_model_configs_request_with_base_url(
|
||||
&transport.key.id,
|
||||
&transport.endpoint.base_url,
|
||||
api_key,
|
||||
);
|
||||
execute_windsurf_probe_plan(
|
||||
state,
|
||||
transport,
|
||||
spec,
|
||||
proxy_override,
|
||||
"windsurf:model_configs",
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn execute_windsurf_rate_limit_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_key: &str,
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
let spec = build_windsurf_pool_rate_limit_request_with_base_url(
|
||||
&transport.key.id,
|
||||
&transport.endpoint.base_url,
|
||||
api_key,
|
||||
);
|
||||
execute_windsurf_probe_plan(
|
||||
state,
|
||||
transport,
|
||||
spec,
|
||||
proxy_override,
|
||||
"windsurf:rate_limit",
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn merge_windsurf_probe_metadata(
|
||||
mut user_status_metadata: serde_json::Value,
|
||||
model_configs_metadata: Option<serde_json::Value>,
|
||||
rate_limit_metadata: Option<serde_json::Value>,
|
||||
) -> serde_json::Value {
|
||||
let Some(target) = user_status_metadata.as_object_mut() else {
|
||||
return user_status_metadata;
|
||||
};
|
||||
for metadata in [model_configs_metadata, rate_limit_metadata]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
{
|
||||
if let Some(source) = metadata.as_object() {
|
||||
for (key, value) in source {
|
||||
target.insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
user_status_metadata
|
||||
}
|
||||
|
||||
fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str, message: String) {
|
||||
let Some(target) = metadata.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
let warnings = target
|
||||
.entry("probe_warnings".to_string())
|
||||
.or_insert_with(|| serde_json::Value::Array(Vec::new()));
|
||||
if let Some(items) = warnings.as_array_mut() {
|
||||
items.push(json!({
|
||||
"probe": probe,
|
||||
"message": message,
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
fn build_windsurf_metadata_update(
|
||||
current_upstream_metadata: Option<&serde_json::Value>,
|
||||
patch: serde_json::Value,
|
||||
) -> serde_json::Value {
|
||||
let Some(patch_object) = patch.as_object() else {
|
||||
return json!({ "windsurf": patch });
|
||||
};
|
||||
let mut merged_bucket = current_upstream_metadata
|
||||
.and_then(|value| value.get("windsurf"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
for (key, value) in patch_object {
|
||||
merged_bucket.insert(key.clone(), value.clone());
|
||||
}
|
||||
json!({ "windsurf": merged_bucket })
|
||||
}
|
||||
|
||||
fn sanitize_windsurf_probe_detail(detail: impl AsRef<str>) -> String {
|
||||
let detail = detail.as_ref().trim();
|
||||
if detail.is_empty() {
|
||||
return "-".to_string();
|
||||
}
|
||||
if let Ok(mut value) = serde_json::from_str::<serde_json::Value>(detail) {
|
||||
redact_windsurf_sensitive_json(&mut value);
|
||||
return value.to_string().chars().take(500).collect();
|
||||
}
|
||||
if contains_windsurf_sensitive_marker(detail) {
|
||||
"[REDACTED upstream error body]".to_string()
|
||||
} else {
|
||||
detail.chars().take(500).collect()
|
||||
}
|
||||
}
|
||||
|
||||
fn redact_windsurf_sensitive_json(value: &mut serde_json::Value) {
|
||||
match value {
|
||||
serde_json::Value::Object(object) => {
|
||||
for (key, value) in object {
|
||||
if is_windsurf_sensitive_key(key) {
|
||||
*value = json!("[REDACTED]");
|
||||
} else {
|
||||
redact_windsurf_sensitive_json(value);
|
||||
}
|
||||
}
|
||||
}
|
||||
serde_json::Value::Array(items) => {
|
||||
for item in items {
|
||||
redact_windsurf_sensitive_json(item);
|
||||
}
|
||||
}
|
||||
serde_json::Value::String(text) if looks_like_windsurf_secret(text) => {
|
||||
*text = "[REDACTED]".to_string();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn is_windsurf_sensitive_key(key: &str) -> bool {
|
||||
let normalized = key
|
||||
.chars()
|
||||
.filter(|ch| ch.is_ascii_alphanumeric())
|
||||
.collect::<String>()
|
||||
.to_ascii_lowercase();
|
||||
normalized.contains("token")
|
||||
|| normalized.contains("apikey")
|
||||
|| normalized.contains("password")
|
||||
|| normalized.contains("authorization")
|
||||
|| normalized.contains("secret")
|
||||
}
|
||||
|
||||
fn looks_like_windsurf_secret(value: &str) -> bool {
|
||||
let value = value.trim();
|
||||
value.starts_with("devin-session-token$")
|
||||
|| value.starts_with("sk-")
|
||||
|| (value.len() > 80 && value.split('.').count() == 3)
|
||||
}
|
||||
|
||||
fn contains_windsurf_sensitive_marker(value: &str) -> bool {
|
||||
let lowered = value.to_ascii_lowercase();
|
||||
[
|
||||
"token",
|
||||
"api_key",
|
||||
"apikey",
|
||||
"sessiontoken",
|
||||
"firebase_id_token",
|
||||
"idtoken",
|
||||
"authorization",
|
||||
"password",
|
||||
"secret",
|
||||
"devin-session-token$",
|
||||
]
|
||||
.iter()
|
||||
.any(|marker| lowered.contains(marker))
|
||||
|| value.contains("sk-")
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_windsurf_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
proxy_override: Option<ProxySnapshot>,
|
||||
) -> Result<Option<serde_json::Value>, GatewayError> {
|
||||
let mut results = Vec::new();
|
||||
let mut success_count = 0usize;
|
||||
let mut failed_count = 0usize;
|
||||
|
||||
for key in keys {
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await?
|
||||
{
|
||||
Some(transport) => transport,
|
||||
None => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Provider transport snapshot unavailable",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let api_key = transport.key.decrypted_api_key.trim();
|
||||
if api_key.is_empty() {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "缺少 Windsurf apiKey/sessionToken",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
let result = match execute_windsurf_user_status_plan(
|
||||
state,
|
||||
&transport,
|
||||
api_key,
|
||||
proxy_override.as_ref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderQuotaExecutionOutcome::Response(result) => result,
|
||||
ProviderQuotaExecutionOutcome::Failure(detail) => {
|
||||
failed_count += 1;
|
||||
let detail = sanitize_windsurf_probe_detail(detail);
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": format!("GetUserStatus 请求执行失败: {detail}"),
|
||||
"status_code": 502,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let mut metadata_update = None::<serde_json::Value>;
|
||||
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
|
||||
quota_refresh_success_invalid_state(&key);
|
||||
let mut status = "error".to_string();
|
||||
let mut message = None::<String>;
|
||||
|
||||
if result.status_code == 200 {
|
||||
if let Some(body_json) = result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
{
|
||||
let mut windsurf_metadata =
|
||||
aether_admin::provider::quota::parse_windsurf_user_status_response(
|
||||
body_json,
|
||||
now_unix_secs,
|
||||
);
|
||||
if let Some(mut metadata) = windsurf_metadata.take() {
|
||||
let model_metadata = match execute_windsurf_model_configs_plan(
|
||||
state,
|
||||
&transport,
|
||||
api_key,
|
||||
proxy_override.as_ref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderQuotaExecutionOutcome::Response(model_result)
|
||||
if model_result.status_code == 200 =>
|
||||
{
|
||||
model_result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
.and_then(|body_json| {
|
||||
aether_admin::provider::quota::parse_windsurf_model_configs_response(
|
||||
body_json,
|
||||
now_unix_secs,
|
||||
)
|
||||
})
|
||||
}
|
||||
ProviderQuotaExecutionOutcome::Response(model_result) => {
|
||||
let detail = extract_execution_error_message(&model_result)
|
||||
.unwrap_or_else(|| format!("HTTP {}", model_result.status_code));
|
||||
let detail = sanitize_windsurf_probe_detail(detail);
|
||||
append_windsurf_probe_warning(
|
||||
&mut metadata,
|
||||
"model_configs",
|
||||
format!("GetCascadeModelConfigs 返回: {detail}"),
|
||||
);
|
||||
None
|
||||
}
|
||||
ProviderQuotaExecutionOutcome::Failure(detail) => {
|
||||
let detail = sanitize_windsurf_probe_detail(detail);
|
||||
append_windsurf_probe_warning(
|
||||
&mut metadata,
|
||||
"model_configs",
|
||||
format!("GetCascadeModelConfigs 执行失败: {detail}"),
|
||||
);
|
||||
None
|
||||
}
|
||||
};
|
||||
let rate_limit_metadata = match execute_windsurf_rate_limit_plan(
|
||||
state,
|
||||
&transport,
|
||||
api_key,
|
||||
proxy_override.as_ref(),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
ProviderQuotaExecutionOutcome::Response(rate_limit_result)
|
||||
if rate_limit_result.status_code == 200 =>
|
||||
{
|
||||
rate_limit_result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.json_body.as_ref())
|
||||
.and_then(|body_json| {
|
||||
aether_admin::provider::quota::parse_windsurf_rate_limit_response(
|
||||
body_json,
|
||||
now_unix_secs,
|
||||
)
|
||||
})
|
||||
}
|
||||
ProviderQuotaExecutionOutcome::Response(rate_limit_result) => {
|
||||
let detail = extract_execution_error_message(&rate_limit_result)
|
||||
.unwrap_or_else(|| format!("HTTP {}", rate_limit_result.status_code));
|
||||
let detail = sanitize_windsurf_probe_detail(detail);
|
||||
append_windsurf_probe_warning(
|
||||
&mut metadata,
|
||||
"rate_limit",
|
||||
format!("CheckUserMessageRateLimit 返回: {detail}"),
|
||||
);
|
||||
None
|
||||
}
|
||||
ProviderQuotaExecutionOutcome::Failure(detail) => {
|
||||
let detail = sanitize_windsurf_probe_detail(detail);
|
||||
append_windsurf_probe_warning(
|
||||
&mut metadata,
|
||||
"rate_limit",
|
||||
format!("CheckUserMessageRateLimit 执行失败: {detail}"),
|
||||
);
|
||||
None
|
||||
}
|
||||
};
|
||||
metadata = merge_windsurf_probe_metadata(
|
||||
metadata,
|
||||
model_metadata,
|
||||
rate_limit_metadata,
|
||||
);
|
||||
metadata_update = Some(build_windsurf_metadata_update(
|
||||
key.upstream_metadata.as_ref(),
|
||||
metadata,
|
||||
));
|
||||
status = "success".to_string();
|
||||
} else {
|
||||
status = "no_metadata".to_string();
|
||||
message = Some("响应中未包含 Windsurf 限额信息".to_string());
|
||||
}
|
||||
} else {
|
||||
status = "no_metadata".to_string();
|
||||
message = Some("无法解析 GetUserStatus 响应".to_string());
|
||||
}
|
||||
} else {
|
||||
let err_msg =
|
||||
extract_execution_error_message(&result).map(sanitize_windsurf_probe_detail);
|
||||
message = Some(match err_msg.as_deref() {
|
||||
Some(detail) if !detail.is_empty() => {
|
||||
format!(
|
||||
"GetUserStatus 返回状态码 {}: {}",
|
||||
result.status_code, detail
|
||||
)
|
||||
}
|
||||
_ => format!("GetUserStatus 返回状态码 {}", result.status_code),
|
||||
});
|
||||
let detail = err_msg
|
||||
.clone()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
|
||||
let mut metadata = serde_json::Map::new();
|
||||
metadata.insert("updated_at".to_string(), json!(now_unix_secs));
|
||||
metadata.insert("last_error".to_string(), json!(detail));
|
||||
match result.status_code {
|
||||
401 | 403 => {
|
||||
oauth_invalid_at_unix_secs = Some(now_unix_secs);
|
||||
oauth_invalid_reason =
|
||||
Some(format!("Windsurf token 无效或已被拒绝: {}", detail));
|
||||
metadata.insert("banned".to_string(), json!(result.status_code == 403));
|
||||
status = if result.status_code == 401 {
|
||||
"auth_invalid".to_string()
|
||||
} else {
|
||||
"forbidden".to_string()
|
||||
};
|
||||
}
|
||||
429 => {
|
||||
metadata.insert(
|
||||
"rate_limit".to_string(),
|
||||
json!({
|
||||
"limited": true,
|
||||
"message": metadata
|
||||
.get("last_error")
|
||||
.cloned()
|
||||
.unwrap_or_else(|| json!("rate limited")),
|
||||
}),
|
||||
);
|
||||
status = "rate_limited".to_string();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
metadata_update = Some(build_windsurf_metadata_update(
|
||||
key.upstream_metadata.as_ref(),
|
||||
serde_json::Value::Object(metadata),
|
||||
));
|
||||
}
|
||||
|
||||
if !persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key.id,
|
||||
metadata_update.as_ref(),
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
None,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
"key_id": key.id,
|
||||
"key_name": key.name,
|
||||
"status": "error",
|
||||
"message": "Key 状态写入失败",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
if status == "success" {
|
||||
success_count += 1;
|
||||
} else {
|
||||
failed_count += 1;
|
||||
}
|
||||
|
||||
let mut payload = serde_json::Map::new();
|
||||
payload.insert("key_id".to_string(), json!(key.id));
|
||||
payload.insert("key_name".to_string(), json!(key.name));
|
||||
payload.insert("status".to_string(), json!(status));
|
||||
if let Some(message) = message {
|
||||
payload.insert("message".to_string(), json!(message));
|
||||
}
|
||||
if result.status_code != 200 {
|
||||
payload.insert("status_code".to_string(), json!(result.status_code));
|
||||
}
|
||||
if let Some(metadata) = metadata_update
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("windsurf"))
|
||||
.cloned()
|
||||
{
|
||||
payload.insert("metadata".to_string(), metadata);
|
||||
}
|
||||
if let Some(quota_snapshot) = build_quota_snapshot_payload(
|
||||
"windsurf",
|
||||
key.status_snapshot.as_ref(),
|
||||
metadata_update.as_ref(),
|
||||
) {
|
||||
payload.insert("quota_snapshot".to_string(), quota_snapshot);
|
||||
}
|
||||
results.push(serde_json::Value::Object(payload));
|
||||
}
|
||||
|
||||
Ok(Some(json!({
|
||||
"success": success_count,
|
||||
"failed": failed_count,
|
||||
"total": success_count + failed_count,
|
||||
"results": results,
|
||||
"message": format!("已处理 {} 个 Key", success_count + failed_count),
|
||||
"auto_removed": 0,
|
||||
})))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn windsurf_probe_metadata_merges_user_status_models_and_rate_limit() {
|
||||
let metadata = super::merge_windsurf_probe_metadata(
|
||||
json!({
|
||||
"plan_name": "Pro",
|
||||
"daily_remaining_percent": 42.0,
|
||||
"updated_at": 1_770_000_000u64,
|
||||
}),
|
||||
Some(json!({
|
||||
"allowed_models_count": 2u64,
|
||||
"models": [
|
||||
{"model_uid": "claude-sonnet-4-5"},
|
||||
{"model_uid": "gpt-5-mini"}
|
||||
],
|
||||
"updated_at": 1_770_000_010u64,
|
||||
})),
|
||||
Some(json!({
|
||||
"rate_limit": {
|
||||
"limited": true,
|
||||
"messages_remaining": 0.0,
|
||||
"retry_after_ms": 60_000u64
|
||||
},
|
||||
"updated_at": 1_770_000_020u64,
|
||||
})),
|
||||
);
|
||||
|
||||
assert_eq!(metadata["plan_name"], json!("Pro"));
|
||||
assert_eq!(metadata["daily_remaining_percent"], json!(42.0));
|
||||
assert_eq!(metadata["allowed_models_count"], json!(2u64));
|
||||
assert_eq!(metadata["rate_limit"]["limited"], json!(true));
|
||||
assert_eq!(metadata["updated_at"], json!(1_770_000_020u64));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_probe_detail_redacts_sensitive_values() {
|
||||
let detail = super::sanitize_windsurf_probe_detail(
|
||||
r#"{"error":{"message":"bad"},"apiKey":"sk-secret","sessionToken":"devin-session-token$secret"}"#,
|
||||
);
|
||||
|
||||
assert!(detail.contains("[REDACTED]"));
|
||||
assert!(!detail.contains("sk-secret"));
|
||||
assert!(!detail.contains("devin-session-token$secret"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn windsurf_metadata_update_preserves_existing_bucket_fields() {
|
||||
let update = super::build_windsurf_metadata_update(
|
||||
Some(&json!({
|
||||
"windsurf": {
|
||||
"daily_remaining_percent": 0.0,
|
||||
"allowed_models_count": 3,
|
||||
"updated_at": 1u64
|
||||
}
|
||||
})),
|
||||
json!({
|
||||
"last_error": "HTTP 429",
|
||||
"rate_limit": {"limited": true},
|
||||
"updated_at": 2u64
|
||||
}),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
update.pointer("/windsurf/daily_remaining_percent"),
|
||||
Some(&json!(0.0))
|
||||
);
|
||||
assert_eq!(
|
||||
update.pointer("/windsurf/allowed_models_count"),
|
||||
Some(&json!(3))
|
||||
);
|
||||
assert_eq!(update.pointer("/windsurf/updated_at"), Some(&json!(2u64)));
|
||||
assert_eq!(
|
||||
update.pointer("/windsurf/rate_limit/limited"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -70,7 +70,7 @@ pub(super) async fn read_admin_provider_ops_balance_cache(
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn store_admin_provider_ops_balance_cache(
|
||||
pub(crate) async fn store_admin_provider_ops_balance_cache(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
payload: &Value,
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
use super::support::{AdminProviderOpsSaveConfigRequest, ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS};
|
||||
use super::support::{
|
||||
AdminProviderOpsQuotaAlertConfigRequest, AdminProviderOpsSaveConfigRequest,
|
||||
ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::ops as admin_provider_ops_pure;
|
||||
@@ -8,6 +11,10 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
const PROVIDER_OPS_QUOTA_ALERT_DEFAULT_FETCH_INTERVAL_SECS: u64 = 30;
|
||||
const PROVIDER_OPS_QUOTA_ALERT_MIN_FETCH_INTERVAL_SECS: u64 = 30;
|
||||
const PROVIDER_OPS_QUOTA_ALERT_MAX_FETCH_INTERVAL_SECS: u64 = 86_400;
|
||||
|
||||
pub(super) fn admin_provider_ops_config_object(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Option<&serde_json::Map<String, serde_json::Value>> {
|
||||
@@ -276,6 +283,7 @@ pub(super) fn build_admin_provider_ops_saved_config_value(
|
||||
)
|
||||
})
|
||||
.collect::<serde_json::Map<String, serde_json::Value>>();
|
||||
let quota_alert = normalize_admin_provider_ops_quota_alert(payload.quota_alert)?;
|
||||
|
||||
Ok(json!({
|
||||
"architecture_id": payload.architecture_id,
|
||||
@@ -287,9 +295,48 @@ pub(super) fn build_admin_provider_ops_saved_config_value(
|
||||
},
|
||||
"actions": actions,
|
||||
"schedule": payload.schedule,
|
||||
"quota_alert": quota_alert,
|
||||
}))
|
||||
}
|
||||
|
||||
fn normalize_admin_provider_ops_quota_alert(
|
||||
request: Option<AdminProviderOpsQuotaAlertConfigRequest>,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let Some(request) = request else {
|
||||
return Ok(default_admin_provider_ops_quota_alert());
|
||||
};
|
||||
let threshold_amount = request.threshold_amount.unwrap_or(0.0);
|
||||
if threshold_amount < 0.0 {
|
||||
return Err("quota_alert.threshold_amount 必须大于等于 0".to_string());
|
||||
}
|
||||
let fetch_interval_seconds = request
|
||||
.fetch_interval_seconds
|
||||
.unwrap_or(PROVIDER_OPS_QUOTA_ALERT_DEFAULT_FETCH_INTERVAL_SECS);
|
||||
if !(PROVIDER_OPS_QUOTA_ALERT_MIN_FETCH_INTERVAL_SECS
|
||||
..=PROVIDER_OPS_QUOTA_ALERT_MAX_FETCH_INTERVAL_SECS)
|
||||
.contains(&fetch_interval_seconds)
|
||||
{
|
||||
return Err(format!(
|
||||
"quota_alert.fetch_interval_seconds 必须在 {} 到 {} 秒之间",
|
||||
PROVIDER_OPS_QUOTA_ALERT_MIN_FETCH_INTERVAL_SECS,
|
||||
PROVIDER_OPS_QUOTA_ALERT_MAX_FETCH_INTERVAL_SECS
|
||||
));
|
||||
}
|
||||
Ok(json!({
|
||||
"enabled": request.enabled,
|
||||
"threshold_amount": threshold_amount,
|
||||
"fetch_interval_seconds": fetch_interval_seconds,
|
||||
}))
|
||||
}
|
||||
|
||||
fn default_admin_provider_ops_quota_alert() -> serde_json::Value {
|
||||
json!({
|
||||
"enabled": false,
|
||||
"threshold_amount": 0.0,
|
||||
"fetch_interval_seconds": PROVIDER_OPS_QUOTA_ALERT_DEFAULT_FETCH_INTERVAL_SECS,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn resolve_admin_provider_ops_base_url(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
@@ -356,5 +403,10 @@ pub(super) fn build_admin_provider_ops_config_payload(
|
||||
connector.and_then(|connector| connector.get("credentials")),
|
||||
),
|
||||
},
|
||||
"quota_alert": provider_ops_config
|
||||
.get("quota_alert")
|
||||
.filter(|value| value.is_object())
|
||||
.cloned()
|
||||
.unwrap_or_else(default_admin_provider_ops_quota_alert),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -4,4 +4,5 @@ mod config;
|
||||
mod routes;
|
||||
mod support;
|
||||
mod verify;
|
||||
pub(crate) use self::balance_cache::store_admin_provider_ops_balance_cache;
|
||||
pub(super) use self::routes::maybe_build_local_admin_provider_ops_providers_response;
|
||||
|
||||
@@ -33,6 +33,8 @@ pub(super) struct AdminProviderOpsSaveConfigRequest {
|
||||
pub(crate) actions: BTreeMap<String, AdminProviderOpsActionConfigRequest>,
|
||||
#[serde(default)]
|
||||
pub(crate) schedule: BTreeMap<String, String>,
|
||||
#[serde(default)]
|
||||
pub(crate) quota_alert: Option<AdminProviderOpsQuotaAlertConfigRequest>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -52,6 +54,16 @@ pub(super) struct AdminProviderOpsActionConfigRequest {
|
||||
pub(crate) config: serde_json::Map<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct AdminProviderOpsQuotaAlertConfigRequest {
|
||||
#[serde(default)]
|
||||
pub(crate) enabled: bool,
|
||||
#[serde(default, deserialize_with = "deserialize_optional_f64_from_number")]
|
||||
pub(crate) threshold_amount: Option<f64>,
|
||||
#[serde(default)]
|
||||
pub(crate) fetch_interval_seconds: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct AdminProviderOpsConnectRequest {
|
||||
#[serde(default)]
|
||||
@@ -71,3 +83,28 @@ fn default_admin_provider_ops_architecture_id() -> String {
|
||||
fn default_admin_provider_ops_action_enabled() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn deserialize_optional_f64_from_number<'de, D>(deserializer: D) -> Result<Option<f64>, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
let value = Option::<serde_json::Value>::deserialize(deserializer)?;
|
||||
match value {
|
||||
None | Some(serde_json::Value::Null) => Ok(None),
|
||||
Some(serde_json::Value::Number(number)) => number
|
||||
.as_f64()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(serde_json::Value::String(raw)) => raw
|
||||
.trim()
|
||||
.parse::<f64>()
|
||||
.ok()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number or numeric string")),
|
||||
Some(_) => Err(serde::de::Error::custom(
|
||||
"expected a finite number or numeric string",
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,7 +14,7 @@ use super::{provider_query_key_display_name, provider_query_provider_payload};
|
||||
use crate::ai_serving::{
|
||||
maybe_build_sync_finalize_outcome, GatewayControlDecision,
|
||||
ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, GEMINI_CHAT_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_CHAT_SYNC_FINALIZE_REPORT_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
|
||||
};
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::execution_runtime;
|
||||
@@ -1792,6 +1792,56 @@ async fn provider_query_execute_kiro_test_candidate(
|
||||
})
|
||||
}
|
||||
|
||||
async fn provider_query_finalize_windsurf_result(
|
||||
route_path: &str,
|
||||
trace_id: &str,
|
||||
requested_model: &str,
|
||||
mapped_model: &str,
|
||||
original_request_body: &Value,
|
||||
result: &aether_contracts::ExecutionResult,
|
||||
) -> Result<Option<Value>, GatewayError> {
|
||||
let decision = GatewayControlDecision::synthetic(
|
||||
route_path,
|
||||
Some("admin_proxy".to_string()),
|
||||
Some("provider_query_manage".to_string()),
|
||||
Some("test_model_failover".to_string()),
|
||||
Some("openai:chat".to_string()),
|
||||
);
|
||||
let payload = GatewaySyncReportRequest {
|
||||
trace_id: trace_id.to_string(),
|
||||
report_kind: OPENAI_CHAT_SYNC_FINALIZE_REPORT_KIND.to_string(),
|
||||
report_context: Some(json!({
|
||||
"client_api_format": "openai:chat",
|
||||
"provider_api_format": "openai:chat",
|
||||
"model": requested_model,
|
||||
"mapped_model": mapped_model,
|
||||
"needs_conversion": false,
|
||||
"has_envelope": true,
|
||||
"envelope_name": crate::provider_transport::windsurf::WINDSURF_ENVELOPE_NAME,
|
||||
"original_request_body": original_request_body,
|
||||
})),
|
||||
status_code: result.status_code,
|
||||
headers: result.headers.clone(),
|
||||
body_json: result.body.as_ref().and_then(|body| body.json_body.clone()),
|
||||
client_body_json: None,
|
||||
body_base64: result
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(|body| body.body_bytes_b64.clone()),
|
||||
telemetry: result.telemetry.clone(),
|
||||
};
|
||||
|
||||
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, &decision, &payload)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let bytes = to_bytes(outcome.response.into_body(), usize::MAX)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
serde_json::from_slice::<Value>(&bytes)
|
||||
.map(Some)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
fn provider_query_build_openai_image_test_request_body_for_route(
|
||||
payload: &Value,
|
||||
model: &str,
|
||||
@@ -2665,6 +2715,22 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
route_path,
|
||||
client_api_format,
|
||||
);
|
||||
if crate::provider_transport::is_windsurf_provider_transport(&transport)
|
||||
&& provider_query_normalize_api_format_alias(candidate.endpoint.api_format.as_str())
|
||||
== "openai:chat"
|
||||
{
|
||||
return provider_query_execute_windsurf_test_candidate(
|
||||
state,
|
||||
provider,
|
||||
candidate,
|
||||
payload,
|
||||
route_path,
|
||||
trace_id,
|
||||
transport,
|
||||
original_request_body,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
if !provider_query_transport_supports_model_test_execution(
|
||||
state,
|
||||
&transport,
|
||||
@@ -3133,6 +3199,182 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
})
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn provider_query_execute_windsurf_test_candidate(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
payload: &Value,
|
||||
route_path: &str,
|
||||
trace_id: &str,
|
||||
transport: AdminGatewayProviderTransportSnapshot,
|
||||
original_request_body: Value,
|
||||
) -> Result<ProviderQueryExecutionOutcome, GatewayError> {
|
||||
if let Some(_reason) =
|
||||
crate::provider_transport::local_windsurf_request_transport_unsupported_reason_with_network(
|
||||
&transport,
|
||||
)
|
||||
{
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
original_request_body,
|
||||
provider_query_standard_test_unsupported_reason(
|
||||
&transport,
|
||||
candidate.endpoint.api_format.as_str(),
|
||||
),
|
||||
));
|
||||
}
|
||||
|
||||
let incoming_request_headers = provider_query_extract_request_headers(payload);
|
||||
let request_body = original_request_body.clone();
|
||||
let request_model =
|
||||
provider_query_request_body_model(&request_body, &candidate.effective_model);
|
||||
let client_is_stream = request_body
|
||||
.get("stream")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let hard_requires_streaming = crate::ai_serving::force_upstream_streaming_for_provider(
|
||||
transport.provider.provider_type.as_str(),
|
||||
candidate.endpoint.api_format.as_str(),
|
||||
);
|
||||
let upstream_is_stream = crate::ai_serving::resolve_upstream_is_stream_from_endpoint_config(
|
||||
transport.endpoint.config.as_ref(),
|
||||
client_is_stream,
|
||||
hard_requires_streaming,
|
||||
);
|
||||
let Some((auth_header, auth_value)) =
|
||||
crate::provider_transport::windsurf::resolve_windsurf_cascade_auth(&transport).or_else(
|
||||
|| crate::provider_transport::auth::resolve_local_openai_bearer_auth(&transport),
|
||||
)
|
||||
else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
request_body,
|
||||
"Provider auth is unavailable for windsurf".to_string(),
|
||||
));
|
||||
};
|
||||
|
||||
let mut synthetic_request = http::Request::builder()
|
||||
.uri(route_path)
|
||||
.body(())
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
*synthetic_request.headers_mut() = incoming_request_headers;
|
||||
let (parts, _) = synthetic_request.into_parts();
|
||||
|
||||
let Some(provider_request_body) =
|
||||
crate::provider_transport::build_windsurf_cascade_request_body(
|
||||
&request_body,
|
||||
request_model,
|
||||
&auth_value,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
upstream_is_stream,
|
||||
)
|
||||
else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
request_body,
|
||||
"Provider request body could not be built for windsurf".to_string(),
|
||||
));
|
||||
};
|
||||
let Some(request_url) = crate::provider_transport::build_windsurf_cascade_upstream_url(
|
||||
transport.endpoint.base_url.as_str(),
|
||||
parts.uri.query(),
|
||||
) else {
|
||||
return Ok(provider_query_skipped_execution_outcome(
|
||||
provider_request_body,
|
||||
"Provider request URL is unavailable for windsurf".to_string(),
|
||||
));
|
||||
};
|
||||
let Some(request_headers) = crate::provider_transport::build_windsurf_cascade_headers(
|
||||
&parts.headers,
|
||||
&provider_request_body,
|
||||
&request_body,
|
||||
transport.endpoint.header_rules.as_ref(),
|
||||
&auth_header,
|
||||
&auth_value,
|
||||
upstream_is_stream,
|
||||
) else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "failed",
|
||||
skip_reason: None,
|
||||
error_message: Some("provider request headers build failed".to_string()),
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url,
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body: provider_request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
|
||||
let plan = ExecutionPlan {
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: Some(format!("provider-query-{}", candidate.key.id)),
|
||||
provider_name: Some(provider.name.clone()),
|
||||
provider_id: provider.id.clone(),
|
||||
endpoint_id: candidate.endpoint.id.clone(),
|
||||
key_id: candidate.key.id.clone(),
|
||||
method: "POST".to_string(),
|
||||
url: request_url.clone(),
|
||||
headers: request_headers.clone(),
|
||||
content_type: Some("application/connect+json".to_string()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(provider_request_body.clone()),
|
||||
stream: upstream_is_stream,
|
||||
client_api_format: "openai:chat".to_string(),
|
||||
provider_api_format: candidate.endpoint.api_format.clone(),
|
||||
model_name: Some(request_model.to_string()),
|
||||
proxy: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
|
||||
.await,
|
||||
transport_profile: state.resolve_transport_profile(&transport),
|
||||
timeouts: state.resolve_transport_execution_timeouts(&transport),
|
||||
};
|
||||
|
||||
let result = state
|
||||
.execute_execution_runtime_sync_plan(Some(trace_id), &plan)
|
||||
.await?;
|
||||
let response_body = if result.status_code < 400 {
|
||||
provider_query_finalize_windsurf_result(
|
||||
route_path,
|
||||
trace_id,
|
||||
request_model,
|
||||
request_model,
|
||||
&request_body,
|
||||
&result,
|
||||
)
|
||||
.await?
|
||||
} else {
|
||||
result.body.as_ref().and_then(|body| body.json_body.clone())
|
||||
};
|
||||
let missing_success_body = result.status_code < 400 && response_body.is_none();
|
||||
let did_fail = result.status_code >= 400 || missing_success_body;
|
||||
let error_message = if did_fail {
|
||||
provider_query_extract_error_message(&result).or_else(|| {
|
||||
missing_success_body.then(|| {
|
||||
format!(
|
||||
"Provider returned HTTP {} without a model-test response body",
|
||||
result.status_code
|
||||
)
|
||||
})
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Ok(ProviderQueryExecutionOutcome {
|
||||
status: if did_fail { "failed" } else { "success" },
|
||||
skip_reason: None,
|
||||
error_message,
|
||||
status_code: Some(result.status_code),
|
||||
latency_ms: result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
|
||||
request_url,
|
||||
request_headers,
|
||||
request_body: provider_request_body,
|
||||
response_headers: result.headers,
|
||||
response_body,
|
||||
})
|
||||
}
|
||||
|
||||
async fn build_admin_provider_query_kiro_failover_response(
|
||||
state: &AdminAppState<'_>,
|
||||
payload: &Value,
|
||||
|
||||
@@ -46,6 +46,22 @@ pub(super) fn provider_query_standard_test_unsupported_reason(
|
||||
api_format: &str,
|
||||
) -> String {
|
||||
let normalized_api_format = crate::ai_serving::normalize_api_format_alias(api_format);
|
||||
if crate::provider_transport::is_windsurf_provider_transport(transport)
|
||||
&& normalized_api_format == "openai:chat"
|
||||
{
|
||||
let reason =
|
||||
crate::provider_transport::local_windsurf_request_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
);
|
||||
return match reason {
|
||||
Some(reason) => format!(
|
||||
"{} ({reason})",
|
||||
provider_query_unsupported_test_api_format_message(api_format)
|
||||
),
|
||||
None => provider_query_unsupported_test_api_format_message(api_format),
|
||||
};
|
||||
}
|
||||
|
||||
let reason = match normalized_api_format.as_str() {
|
||||
"openai:chat" => {
|
||||
crate::provider_transport::policy::local_openai_chat_transport_unsupported_reason(
|
||||
@@ -294,6 +310,15 @@ pub(super) fn provider_query_transport_supports_model_test_execution(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
if crate::provider_transport::is_windsurf_provider_transport(transport)
|
||||
&& provider_query_normalize_api_format_alias(api_format) == "openai:chat"
|
||||
{
|
||||
return crate::provider_transport::local_windsurf_request_transport_unsupported_reason_with_network(
|
||||
transport,
|
||||
)
|
||||
.is_none();
|
||||
}
|
||||
|
||||
match provider_query_test_adapter_for_provider_api_format(
|
||||
transport.provider.provider_type.as_str(),
|
||||
api_format,
|
||||
|
||||
+130
-21
@@ -42,9 +42,9 @@ pub(super) fn provider_query_test_attempt_payload(
|
||||
"status_code": execution.status_code,
|
||||
"latency_ms": execution.latency_ms,
|
||||
"request_url": execution.request_url,
|
||||
"request_headers": provider_query_redact_diagnostic_headers(&execution.request_headers),
|
||||
"request_body": execution.request_body,
|
||||
"response_headers": provider_query_redact_diagnostic_headers(&execution.response_headers),
|
||||
"request_headers": redacted_provider_query_headers(&execution.request_headers),
|
||||
"request_body": redacted_provider_query_value(&execution.request_body),
|
||||
"response_headers": redacted_provider_query_headers(&execution.response_headers),
|
||||
"response_body": execution.response_body,
|
||||
})
|
||||
}
|
||||
@@ -172,34 +172,84 @@ fn provider_query_endpoint_route_payload(
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_redact_diagnostic_headers(
|
||||
headers: &BTreeMap<String, String>,
|
||||
) -> BTreeMap<String, String> {
|
||||
fn redacted_provider_query_headers(headers: &BTreeMap<String, String>) -> BTreeMap<String, String> {
|
||||
headers
|
||||
.iter()
|
||||
.map(|(name, value)| {
|
||||
if provider_query_header_is_sensitive(name) {
|
||||
(name.clone(), "<redacted>".to_string())
|
||||
.map(|(key, value)| {
|
||||
if provider_query_field_is_sensitive(key) {
|
||||
(key.clone(), "[REDACTED]".to_string())
|
||||
} else {
|
||||
(name.clone(), value.clone())
|
||||
(key.clone(), value.clone())
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn provider_query_header_is_sensitive(name: &str) -> bool {
|
||||
fn redacted_provider_query_value(value: &Value) -> Value {
|
||||
match value {
|
||||
Value::Object(object) => Value::Object(
|
||||
object
|
||||
.iter()
|
||||
.map(|(key, value)| {
|
||||
if provider_query_field_is_sensitive(key) {
|
||||
(key.clone(), Value::String("[REDACTED]".to_string()))
|
||||
} else {
|
||||
(key.clone(), redacted_provider_query_value(value))
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
Value::Array(items) => Value::Array(
|
||||
items
|
||||
.iter()
|
||||
.map(redacted_provider_query_value)
|
||||
.collect::<Vec<_>>(),
|
||||
),
|
||||
other => other.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_query_field_is_sensitive(key: &str) -> bool {
|
||||
let key = key.trim().to_ascii_lowercase();
|
||||
let normalized = key
|
||||
.chars()
|
||||
.filter(|ch| ch.is_ascii_alphanumeric())
|
||||
.collect::<String>();
|
||||
if matches!(
|
||||
normalized.as_str(),
|
||||
"maxtokens"
|
||||
| "maxoutputtokens"
|
||||
| "inputtokens"
|
||||
| "outputtokens"
|
||||
| "prompttokens"
|
||||
| "completiontokens"
|
||||
| "totaltokens"
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
matches!(
|
||||
name.trim().to_ascii_lowercase().as_str(),
|
||||
key.as_str(),
|
||||
"authorization"
|
||||
| "proxy-authorization"
|
||||
| "cookie"
|
||||
| "set-cookie"
|
||||
| "x-api-key"
|
||||
| "api_key"
|
||||
| "apikey"
|
||||
| "api-key"
|
||||
| "x-api-key"
|
||||
| "x-goog-api-key"
|
||||
| "anthropic-api-key"
|
||||
| "openai-api-key"
|
||||
)
|
||||
| "x-codeium-csrf-token"
|
||||
| "access_token"
|
||||
| "refresh_token"
|
||||
| "id_token"
|
||||
| "password"
|
||||
| "secret"
|
||||
) || normalized.ends_with("token")
|
||||
|| normalized.contains("secret")
|
||||
|| normalized.contains("apikey")
|
||||
|| normalized.contains("authorization")
|
||||
}
|
||||
|
||||
pub(super) fn provider_query_candidate_summary_payload(
|
||||
@@ -309,34 +359,93 @@ pub(super) fn provider_query_candidate_summary_payload(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use super::{redacted_provider_query_headers, redacted_provider_query_value};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
#[test]
|
||||
fn provider_query_diagnostic_headers_redact_credentials() {
|
||||
fn redacts_sensitive_provider_query_headers() {
|
||||
let headers = BTreeMap::from([
|
||||
("cookie".to_string(), "sso=secret".to_string()),
|
||||
("authorization".to_string(), "Bearer secret".to_string()),
|
||||
(
|
||||
"authorization".to_string(),
|
||||
"Bearer secret-token".to_string(),
|
||||
),
|
||||
("x-goog-api-key".to_string(), "secret".to_string()),
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
(
|
||||
"x-codeium-csrf-token".to_string(),
|
||||
"csrf-secret".to_string(),
|
||||
),
|
||||
]);
|
||||
|
||||
let redacted = provider_query_redact_diagnostic_headers(&headers);
|
||||
let redacted = redacted_provider_query_headers(&headers);
|
||||
|
||||
assert_eq!(
|
||||
redacted.get("cookie").map(String::as_str),
|
||||
Some("<redacted>")
|
||||
Some("[REDACTED]")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("authorization").map(String::as_str),
|
||||
Some("<redacted>")
|
||||
Some("[REDACTED]")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("x-goog-api-key").map(String::as_str),
|
||||
Some("<redacted>")
|
||||
Some("[REDACTED]")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("x-codeium-csrf-token").map(String::as_str),
|
||||
Some("[REDACTED]")
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.get("content-type").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redacts_sensitive_provider_query_request_body_fields() {
|
||||
let body = json!({
|
||||
"metadata": {
|
||||
"apiKey": "devin-session-token$secret",
|
||||
"ideName": "windsurf"
|
||||
},
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"stream": true
|
||||
});
|
||||
|
||||
let redacted = redacted_provider_query_value(&body);
|
||||
|
||||
assert_eq!(
|
||||
redacted.pointer("/metadata/apiKey"),
|
||||
Some(&json!("[REDACTED]"))
|
||||
);
|
||||
assert_eq!(
|
||||
redacted.pointer("/metadata/ideName"),
|
||||
Some(&json!("windsurf"))
|
||||
);
|
||||
assert_eq!(redacted.pointer("/stream"), Some(&json!(true)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn keeps_non_secret_token_count_fields_visible() {
|
||||
let body = json!({
|
||||
"maxTokens": 64,
|
||||
"usage": {
|
||||
"inputTokens": 10,
|
||||
"outputTokens": 2,
|
||||
"accessToken": "secret"
|
||||
}
|
||||
});
|
||||
|
||||
let redacted = redacted_provider_query_value(&body);
|
||||
|
||||
assert_eq!(redacted.pointer("/maxTokens"), Some(&json!(64)));
|
||||
assert_eq!(redacted.pointer("/usage/inputTokens"), Some(&json!(10)));
|
||||
assert_eq!(redacted.pointer("/usage/outputTokens"), Some(&json!(2)));
|
||||
assert_eq!(
|
||||
redacted.pointer("/usage/accessToken"),
|
||||
Some(&json!("[REDACTED]"))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -138,6 +138,13 @@ pub(crate) fn build_admin_provider_summary_value(
|
||||
.and_then(|cfg| cfg.get("simulated_cache_enabled"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let ops_quota_alert_enabled = provider_ops_config
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|cfg| cfg.get("quota_alert"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|cfg| cfg.get("enabled"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let billing_type = quota_snapshot
|
||||
.map(|quota| quota.billing_type.clone())
|
||||
.or_else(|| provider.billing_type.clone());
|
||||
@@ -197,6 +204,7 @@ pub(crate) fn build_admin_provider_summary_value(
|
||||
"ops_configured": ops_configured,
|
||||
"ops_architecture_id": ops_architecture_id,
|
||||
"kiro_simulated_cache_enabled": kiro_simulated_cache_enabled,
|
||||
"ops_quota_alert_enabled": ops_quota_alert_enabled,
|
||||
"created_at": endpoint_timestamp_or_now(provider.created_at_unix_ms, now_unix_secs),
|
||||
"updated_at": endpoint_timestamp_or_now(provider.updated_at_unix_secs, now_unix_secs),
|
||||
})
|
||||
|
||||
@@ -4,9 +4,9 @@ pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, Strin
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"custom" | "claude_code" | "kiro" | "codex" | "chatgpt_web" | "gemini_cli"
|
||||
| "antigravity" | "vertex_ai" | "grok" => Ok(normalized),
|
||||
| "antigravity" | "vertex_ai" | "grok" | "windsurf" => Ok(normalized),
|
||||
_ => Err(
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok"
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / chatgpt_web / gemini_cli / antigravity / vertex_ai / grok / windsurf"
|
||||
.to_string(),
|
||||
),
|
||||
}
|
||||
|
||||
@@ -49,7 +49,7 @@ use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
const ADMIN_SYSTEM_IMPORT_MAX_SIZE_BYTES: usize = 10 * 1024 * 1024;
|
||||
const ADMIN_SYSTEM_IMPORT_MAX_SIZE_BYTES: usize = 500 * 1024 * 1024;
|
||||
|
||||
fn invalid_request(detail: impl Into<String>) -> (http::StatusCode, Value) {
|
||||
(
|
||||
@@ -173,6 +173,9 @@ fn normalize_import_endpoint_format(value: &str) -> Result<String, String> {
|
||||
let normalized = match value.trim().to_ascii_lowercase().as_str() {
|
||||
"openai:cli" => "openai:responses",
|
||||
"openai:compact" => "openai:responses:compact",
|
||||
"openai_image" | "images" | "image" | "/v1/images/generations" | "/v1/images/edits" => {
|
||||
"openai:image"
|
||||
}
|
||||
"claude:chat" | "claude:cli" => "claude:messages",
|
||||
"gemini:chat" | "gemini:cli" => "gemini:generate_content",
|
||||
_ => value.trim(),
|
||||
@@ -956,7 +959,7 @@ impl<'a> AdminAppState<'a> {
|
||||
}
|
||||
|
||||
if request_body.len() > ADMIN_SYSTEM_DATA_IMPORT_MAX_SIZE_BYTES {
|
||||
return Ok(Err(invalid_request("请求体大小不能超过 20MB")));
|
||||
return Ok(Err(invalid_request("请求体大小不能超过 500MB")));
|
||||
}
|
||||
|
||||
let root = match serde_json::from_slice::<Value>(request_body) {
|
||||
@@ -1053,7 +1056,7 @@ impl<'a> AdminAppState<'a> {
|
||||
)));
|
||||
}
|
||||
if request_body.len() > ADMIN_SYSTEM_IMPORT_MAX_SIZE_BYTES {
|
||||
return Ok(Err(invalid_request("请求体大小不能超过 10MB")));
|
||||
return Ok(Err(invalid_request("请求体大小不能超过 500MB")));
|
||||
}
|
||||
|
||||
let parsed = routed!(parse_admin_system_config_import_request(request_body));
|
||||
@@ -1993,7 +1996,7 @@ impl<'a> AdminAppState<'a> {
|
||||
)));
|
||||
}
|
||||
if request_body.len() > ADMIN_SYSTEM_IMPORT_MAX_SIZE_BYTES {
|
||||
return Ok(Err(invalid_request("请求体大小不能超过 10MB")));
|
||||
return Ok(Err(invalid_request("请求体大小不能超过 500MB")));
|
||||
}
|
||||
|
||||
let root = match serde_json::from_slice::<Value>(request_body) {
|
||||
@@ -3062,6 +3065,10 @@ mod tests {
|
||||
for (raw, expected) in [
|
||||
("openai:cli", "openai:responses"),
|
||||
("openai:compact", "openai:responses:compact"),
|
||||
("openai_image", "openai:image"),
|
||||
("images", "openai:image"),
|
||||
("/v1/images/generations", "openai:image"),
|
||||
("/v1/images/edits", "openai:image"),
|
||||
("claude:chat", "claude:messages"),
|
||||
("claude:cli", "claude:messages"),
|
||||
("gemini:chat", "gemini:generate_content"),
|
||||
|
||||
@@ -9,7 +9,7 @@ mod proxy_nodes;
|
||||
mod templates;
|
||||
|
||||
const ADMIN_SYSTEM_DATA_EXPORT_VERSION: &str = "1.0";
|
||||
const ADMIN_SYSTEM_DATA_IMPORT_MAX_SIZE_BYTES: usize = 20 * 1024 * 1024;
|
||||
const ADMIN_SYSTEM_DATA_IMPORT_MAX_SIZE_BYTES: usize = 500 * 1024 * 1024;
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn upsert_system_config_json_value(
|
||||
|
||||
@@ -17,6 +17,7 @@ use crate::handlers::admin::system::shared::settings::{
|
||||
build_admin_system_stats_payload, current_aether_version, fetch_latest_admin_system_release,
|
||||
};
|
||||
use crate::handlers::admin::system::shared::smtp::build_admin_smtp_test_payload;
|
||||
use crate::important_notification::build_important_notification_test_payload;
|
||||
use crate::maintenance::{ManualUsageCleanupMode, ManualUsageCleanupOptions};
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::usage::UsageCleanupTargets;
|
||||
@@ -241,6 +242,16 @@ pub(super) async fn maybe_build_local_admin_core_system_response(
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("important_notification_test")
|
||||
&& request_method == http::Method::POST
|
||||
&& request_path == "/api/admin/system/important-notification/test"
|
||||
{
|
||||
return Ok(Some(
|
||||
Json(build_important_notification_test_payload(state, request_body).await?)
|
||||
.into_response(),
|
||||
));
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("cleanup") && request_method == http::Method::POST {
|
||||
return Ok(Some(attach_admin_audit_response(
|
||||
Json(build_admin_system_cleanup_payload(state).await?).into_response(),
|
||||
|
||||
@@ -33,6 +33,21 @@ fn admin_system_config_default_value(key: &str) -> Option<serde_json::Value> {
|
||||
admin_system_config_default_value_pure(key)
|
||||
}
|
||||
|
||||
fn legacy_admin_system_config_fallback_key(normalized_key: &str) -> Option<&'static str> {
|
||||
match normalized_key {
|
||||
"module.server_chan_push.enabled" => {
|
||||
Some("module.important_notification.server_chan_enabled")
|
||||
}
|
||||
"module.server_chan_push.send_key" => {
|
||||
Some("module.important_notification.server_chan_send_key")
|
||||
}
|
||||
"module.server_chan_push.template" => {
|
||||
Some("module.important_notification.server_chan_template")
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_system_configs_payload(
|
||||
entries: &[aether_data::repository::system::StoredSystemConfigEntry],
|
||||
) -> serde_json::Value {
|
||||
@@ -44,12 +59,14 @@ pub(crate) async fn build_admin_system_config_detail_payload(
|
||||
requested_key: &str,
|
||||
) -> Result<Result<serde_json::Value, (http::StatusCode, serde_json::Value)>, GatewayError> {
|
||||
let requested_key = requested_key.trim();
|
||||
let value = state
|
||||
.read_system_config_json_value(&normalize_admin_system_config_key(requested_key))
|
||||
.await?
|
||||
.or_else(|| {
|
||||
admin_system_config_default_value(&normalize_admin_system_config_key(requested_key))
|
||||
});
|
||||
let normalized_key = normalize_admin_system_config_key(requested_key);
|
||||
let mut value = state.read_system_config_json_value(&normalized_key).await?;
|
||||
if value.is_none() {
|
||||
if let Some(legacy_key) = legacy_admin_system_config_fallback_key(&normalized_key) {
|
||||
value = state.read_system_config_json_value(legacy_key).await?;
|
||||
}
|
||||
}
|
||||
let value = value.or_else(|| admin_system_config_default_value(&normalized_key));
|
||||
Ok(build_admin_system_config_detail_payload_pure(
|
||||
requested_key,
|
||||
value,
|
||||
|
||||
@@ -1,5 +1,11 @@
|
||||
use crate::bark_push::bark_push_configured;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::shared::{module_available_from_env, system_config_bool};
|
||||
use crate::important_notification::{
|
||||
important_notification_configured, IMPORTANT_NOTIFICATION_ENABLED_KEY,
|
||||
LEGACY_NOTIFICATION_EMAIL_ENABLED_KEY,
|
||||
};
|
||||
use crate::server_chan_push::server_chan_push_configured;
|
||||
use crate::system_features::ENABLE_MODEL_DIRECTIVES_CONFIG_KEY;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::system as admin_system_kernel;
|
||||
@@ -68,17 +74,41 @@ pub(crate) const ADMIN_MODULE_DEFINITIONS: &[AdminModuleDefinition] = &[
|
||||
admin_menu_order: 59,
|
||||
},
|
||||
AdminModuleDefinition {
|
||||
name: "notification_email",
|
||||
display_name: "异常通知",
|
||||
description: "为 5xx 异常发送邮件通知,可在模块管理中启用或禁用",
|
||||
name: "important_notification",
|
||||
display_name: "通知服务",
|
||||
description: "统一管理通知项、模板和推送服务选择,供后台任务和用户通知使用",
|
||||
category: "integration",
|
||||
env_key: "NOTIFICATION_EMAIL_AVAILABLE",
|
||||
env_key: "IMPORTANT_NOTIFICATION_AVAILABLE",
|
||||
default_available: true,
|
||||
admin_route: None,
|
||||
admin_menu_icon: Some("Mail"),
|
||||
admin_menu_group: Some("system"),
|
||||
admin_route: Some("/admin/notification-service"),
|
||||
admin_menu_icon: Some("BellRing"),
|
||||
admin_menu_group: None,
|
||||
admin_menu_order: 58,
|
||||
},
|
||||
AdminModuleDefinition {
|
||||
name: "server_chan_push",
|
||||
display_name: "Server 酱推送",
|
||||
description: "第三方推送服务,配置 Server 酱 Turbo SendKey 并测试微信推送",
|
||||
category: "integration",
|
||||
env_key: "SERVER_CHAN_PUSH_AVAILABLE",
|
||||
default_available: true,
|
||||
admin_route: Some("/admin/modules/server-chan"),
|
||||
admin_menu_icon: Some("Send"),
|
||||
admin_menu_group: Some("system"),
|
||||
admin_menu_order: 59,
|
||||
},
|
||||
AdminModuleDefinition {
|
||||
name: "bark_push",
|
||||
display_name: "Bark 推送",
|
||||
description: "第三方推送服务,配置 Bark Device Key 并测试 iOS 推送",
|
||||
category: "integration",
|
||||
env_key: "BARK_PUSH_AVAILABLE",
|
||||
default_available: true,
|
||||
admin_route: Some("/admin/modules/bark"),
|
||||
admin_menu_icon: Some("Send"),
|
||||
admin_menu_group: Some("system"),
|
||||
admin_menu_order: 59,
|
||||
},
|
||||
AdminModuleDefinition {
|
||||
name: "model_directives",
|
||||
display_name: "模型后缀参数",
|
||||
@@ -150,10 +180,17 @@ pub(crate) struct AdminModuleRuntimeState {
|
||||
oauth_providers: Vec<aether_data::repository::auth_modules::StoredOAuthProviderModuleConfig>,
|
||||
ldap_config: Option<aether_data::repository::auth_modules::StoredLdapModuleConfig>,
|
||||
gemini_files_has_capable_key: bool,
|
||||
smtp_configured: bool,
|
||||
important_notification_configured: bool,
|
||||
server_chan_push_configured: bool,
|
||||
bark_push_configured: bool,
|
||||
}
|
||||
|
||||
pub(crate) fn admin_module_by_name(name: &str) -> Option<&'static AdminModuleDefinition> {
|
||||
let name = if name == "notification_email" {
|
||||
"important_notification"
|
||||
} else {
|
||||
name
|
||||
};
|
||||
ADMIN_MODULE_DEFINITIONS
|
||||
.iter()
|
||||
.find(|module| module.name == name)
|
||||
@@ -170,11 +207,22 @@ pub(crate) fn admin_module_name_from_enabled_path(request_path: &str) -> Option<
|
||||
pub(crate) fn admin_module_enabled_config_key(module: &AdminModuleDefinition) -> String {
|
||||
if module.name == "model_directives" {
|
||||
ENABLE_MODEL_DIRECTIVES_CONFIG_KEY.to_string()
|
||||
} else if module.name == "important_notification" {
|
||||
IMPORTANT_NOTIFICATION_ENABLED_KEY.to_string()
|
||||
} else {
|
||||
format!("module.{}.enabled", module.name)
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_module_available(module: &AdminModuleDefinition) -> bool {
|
||||
if module.name == "important_notification" {
|
||||
let legacy_default =
|
||||
module_available_from_env("NOTIFICATION_EMAIL_AVAILABLE", module.default_available);
|
||||
return module_available_from_env(module.env_key, legacy_default);
|
||||
}
|
||||
module_available_from_env(module.env_key, module.default_available)
|
||||
}
|
||||
|
||||
pub(crate) fn oauth_module_config_is_valid(
|
||||
providers: &[aether_data::repository::auth_modules::StoredOAuthProviderModuleConfig],
|
||||
) -> bool {
|
||||
@@ -221,28 +269,17 @@ pub(crate) async fn build_admin_module_runtime_state(
|
||||
})
|
||||
};
|
||||
|
||||
let smtp_host = state.read_system_config_json_value("smtp_host").await?;
|
||||
let smtp_from_email = state
|
||||
.read_system_config_json_value("smtp_from_email")
|
||||
.await?;
|
||||
let smtp_configured = smtp_host
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.is_some()
|
||||
&& smtp_from_email
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.is_some();
|
||||
let notification_configured = important_notification_configured(state.app()).await?;
|
||||
let server_chan_configured = server_chan_push_configured(state.app()).await?;
|
||||
let bark_configured = bark_push_configured(state.app()).await?;
|
||||
|
||||
Ok(AdminModuleRuntimeState {
|
||||
oauth_providers,
|
||||
ldap_config,
|
||||
gemini_files_has_capable_key,
|
||||
smtp_configured,
|
||||
important_notification_configured: notification_configured,
|
||||
server_chan_push_configured: server_chan_configured,
|
||||
bark_push_configured: bark_configured,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -255,7 +292,9 @@ pub(crate) fn build_admin_module_validation_result(
|
||||
&runtime.oauth_providers,
|
||||
runtime.ldap_config.as_ref(),
|
||||
runtime.gemini_files_has_capable_key,
|
||||
runtime.smtp_configured,
|
||||
runtime.important_notification_configured,
|
||||
runtime.server_chan_push_configured,
|
||||
runtime.bark_push_configured,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -274,12 +313,19 @@ pub(crate) async fn build_admin_module_status_payload(
|
||||
module: &AdminModuleDefinition,
|
||||
runtime: &AdminModuleRuntimeState,
|
||||
) -> Result<serde_json::Value, GatewayError> {
|
||||
let available = module_available_from_env(module.env_key, module.default_available);
|
||||
let available = admin_module_available(module);
|
||||
let enabled = if available {
|
||||
let enabled = state
|
||||
let enabled_value = state
|
||||
.read_system_config_json_value(&admin_module_enabled_config_key(module))
|
||||
.await?;
|
||||
system_config_bool(enabled.as_ref(), false)
|
||||
let enabled_value = if module.name == "important_notification" && enabled_value.is_none() {
|
||||
state
|
||||
.read_system_config_json_value(LEGACY_NOTIFICATION_EMAIL_ENABLED_KEY)
|
||||
.await?
|
||||
} else {
|
||||
enabled_value
|
||||
};
|
||||
system_config_bool(enabled_value.as_ref(), false)
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
@@ -1,14 +1,10 @@
|
||||
use crate::email_delivery::{probe_smtp_connection, system_config_u16, SmtpDeliveryConfig};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::shared::{system_config_bool, system_config_string};
|
||||
use crate::GatewayError;
|
||||
use axum::body::Bytes;
|
||||
use base64::Engine;
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
use std::io::{BufRead, Write};
|
||||
use std::time::Duration;
|
||||
|
||||
const SMTP_TIMEOUT_SECS: u64 = 30;
|
||||
|
||||
#[derive(Debug, Default, Deserialize)]
|
||||
struct AdminSmtpTestRequest {
|
||||
@@ -53,12 +49,12 @@ pub(crate) async fn build_admin_smtp_test_payload(
|
||||
}));
|
||||
}
|
||||
|
||||
let result = tokio::task::spawn_blocking(move || test_smtp_connection_blocking(config))
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let result = probe_smtp_connection(config.into_delivery_config()).await;
|
||||
Ok(match result {
|
||||
Ok(()) => json!({ "success": true, "message": "SMTP 连接测试成功" }),
|
||||
Err(error) => json!({ "success": false, "message": translate_smtp_error(&error) }),
|
||||
Err(error) => {
|
||||
json!({ "success": false, "message": translate_smtp_error(&smtp_gateway_error_message(&error)) })
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -94,8 +90,8 @@ async fn resolve_admin_smtp_config(
|
||||
port: request
|
||||
.smtp_port
|
||||
.as_ref()
|
||||
.map(|value| system_config_u16(value, 587))
|
||||
.unwrap_or_else(|| system_config_u16_opt(smtp_port.as_ref(), 587)),
|
||||
.map(|value| system_config_u16(Some(value), 587))
|
||||
.unwrap_or_else(|| system_config_u16(smtp_port.as_ref(), 587)),
|
||||
user: request
|
||||
.smtp_user
|
||||
.as_ref()
|
||||
@@ -130,6 +126,21 @@ async fn resolve_admin_smtp_config(
|
||||
})
|
||||
}
|
||||
|
||||
impl ResolvedSmtpConfig {
|
||||
fn into_delivery_config(self) -> SmtpDeliveryConfig {
|
||||
SmtpDeliveryConfig {
|
||||
host: self.host.unwrap_or_default(),
|
||||
port: self.port,
|
||||
user: self.user,
|
||||
password: self.password,
|
||||
use_tls: self.use_tls,
|
||||
use_ssl: self.use_ssl,
|
||||
from_email: self.from_email.unwrap_or_default(),
|
||||
from_name: self.from_name,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn missing_smtp_fields(config: &ResolvedSmtpConfig) -> Vec<&'static str> {
|
||||
let mut fields = Vec::new();
|
||||
if config
|
||||
@@ -171,178 +182,13 @@ fn missing_smtp_fields(config: &ResolvedSmtpConfig) -> Vec<&'static str> {
|
||||
fields
|
||||
}
|
||||
|
||||
fn system_config_u16_opt(value: Option<&serde_json::Value>, default: u16) -> u16 {
|
||||
value
|
||||
.map(|value| system_config_u16(value, default))
|
||||
.unwrap_or(default)
|
||||
}
|
||||
|
||||
fn system_config_u16(value: &serde_json::Value, default: u16) -> u16 {
|
||||
match value {
|
||||
serde_json::Value::Number(value) => value
|
||||
.as_u64()
|
||||
.and_then(|value| u16::try_from(value).ok())
|
||||
.unwrap_or(default),
|
||||
serde_json::Value::String(value) => value.trim().parse::<u16>().unwrap_or(default),
|
||||
_ => default,
|
||||
fn smtp_gateway_error_message(error: &GatewayError) -> String {
|
||||
match error {
|
||||
GatewayError::Internal(message) => message.clone(),
|
||||
_ => format!("{error:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn build_tls_config() -> std::sync::Arc<rustls::ClientConfig> {
|
||||
let _ = rustls::crypto::ring::default_provider().install_default();
|
||||
let root_store =
|
||||
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
|
||||
std::sync::Arc::new(
|
||||
rustls::ClientConfig::builder()
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth(),
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_server_name(host: &str) -> Result<rustls::pki_types::ServerName<'static>, String> {
|
||||
let host = host.trim().trim_start_matches('[').trim_end_matches(']');
|
||||
if let Ok(ip) = host.parse::<std::net::IpAddr>() {
|
||||
return Ok(rustls::pki_types::ServerName::from(ip));
|
||||
}
|
||||
rustls::pki_types::ServerName::try_from(host.to_string()).map_err(|err| err.to_string())
|
||||
}
|
||||
|
||||
fn connect_tcp_stream(config: &ResolvedSmtpConfig) -> Result<std::net::TcpStream, String> {
|
||||
let host = config.host.as_deref().unwrap_or_default();
|
||||
let stream =
|
||||
std::net::TcpStream::connect((host, config.port)).map_err(|err| err.to_string())?;
|
||||
stream
|
||||
.set_read_timeout(Some(Duration::from_secs(SMTP_TIMEOUT_SECS)))
|
||||
.map_err(|err| err.to_string())?;
|
||||
stream
|
||||
.set_write_timeout(Some(Duration::from_secs(SMTP_TIMEOUT_SECS)))
|
||||
.map_err(|err| err.to_string())?;
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
fn wrap_tls_stream(
|
||||
stream: std::net::TcpStream,
|
||||
host: &str,
|
||||
) -> Result<rustls::StreamOwned<rustls::ClientConnection, std::net::TcpStream>, String> {
|
||||
let server_name = resolve_server_name(host)?;
|
||||
let connection = rustls::ClientConnection::new(build_tls_config(), server_name)
|
||||
.map_err(|err| err.to_string())?;
|
||||
Ok(rustls::StreamOwned::new(connection, stream))
|
||||
}
|
||||
|
||||
fn smtp_read_response<T: BufRead>(reader: &mut T) -> Result<(u16, String), String> {
|
||||
let mut message = String::new();
|
||||
let code = loop {
|
||||
let mut line = String::new();
|
||||
let bytes = reader.read_line(&mut line).map_err(|err| err.to_string())?;
|
||||
if bytes == 0 {
|
||||
return Err("smtp connection closed unexpectedly".to_string());
|
||||
}
|
||||
let trimmed = line.trim_end_matches(['\r', '\n']).to_string();
|
||||
if trimmed.len() < 3 {
|
||||
return Err("invalid smtp response".to_string());
|
||||
}
|
||||
let parsed_code = trimmed[..3].parse::<u16>().map_err(|err| err.to_string())?;
|
||||
let continuation = trimmed.as_bytes().get(3).copied() == Some(b'-');
|
||||
if !message.is_empty() {
|
||||
message.push('\n');
|
||||
}
|
||||
message.push_str(&trimmed);
|
||||
if !continuation {
|
||||
break parsed_code;
|
||||
}
|
||||
};
|
||||
Ok((code, message))
|
||||
}
|
||||
|
||||
fn smtp_expect<T: BufRead>(reader: &mut T, allowed_codes: &[u16]) -> Result<String, String> {
|
||||
let (code, message) = smtp_read_response(reader)?;
|
||||
if allowed_codes.contains(&code) {
|
||||
return Ok(message);
|
||||
}
|
||||
Err(format!("unexpected smtp response {code}: {message}"))
|
||||
}
|
||||
|
||||
fn smtp_write_line<T: Write>(writer: &mut T, line: &str) -> Result<(), String> {
|
||||
writer
|
||||
.write_all(line.as_bytes())
|
||||
.map_err(|err| err.to_string())?;
|
||||
writer.write_all(b"\r\n").map_err(|err| err.to_string())?;
|
||||
writer.flush().map_err(|err| err.to_string())
|
||||
}
|
||||
|
||||
fn smtp_send_command<S: std::io::Read + Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
command: &str,
|
||||
allowed_codes: &[u16],
|
||||
) -> Result<String, String> {
|
||||
smtp_write_line(reader.get_mut(), command)?;
|
||||
smtp_expect(reader, allowed_codes)
|
||||
}
|
||||
|
||||
fn smtp_authenticate<S: std::io::Read + Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
config: &ResolvedSmtpConfig,
|
||||
) -> Result<(), String> {
|
||||
let Some(username) = config
|
||||
.user
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
let password = config.password.as_deref().unwrap_or_default();
|
||||
smtp_send_command(reader, "AUTH LOGIN", &[334])?;
|
||||
smtp_send_command(
|
||||
reader,
|
||||
&base64::engine::general_purpose::STANDARD.encode(username.as_bytes()),
|
||||
&[334],
|
||||
)?;
|
||||
smtp_send_command(
|
||||
reader,
|
||||
&base64::engine::general_purpose::STANDARD.encode(password.as_bytes()),
|
||||
&[235],
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn smtp_probe<S: std::io::Read + Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
config: &ResolvedSmtpConfig,
|
||||
) -> Result<(), String> {
|
||||
smtp_send_command(reader, "EHLO aether.local", &[250])?;
|
||||
smtp_authenticate(reader, config)?;
|
||||
let _ = smtp_send_command(reader, "QUIT", &[221]);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn test_smtp_connection_blocking(config: ResolvedSmtpConfig) -> Result<(), String> {
|
||||
if config.use_ssl {
|
||||
let stream = connect_tcp_stream(&config)?;
|
||||
let tls_stream = wrap_tls_stream(stream, config.host.as_deref().unwrap_or_default())?;
|
||||
let mut reader = std::io::BufReader::new(tls_stream);
|
||||
smtp_expect(&mut reader, &[220])?;
|
||||
return smtp_probe(&mut reader, &config);
|
||||
}
|
||||
|
||||
let stream = connect_tcp_stream(&config)?;
|
||||
let mut reader = std::io::BufReader::new(stream);
|
||||
smtp_expect(&mut reader, &[220])?;
|
||||
smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
|
||||
if config.use_tls {
|
||||
smtp_send_command(&mut reader, "STARTTLS", &[220])?;
|
||||
let stream = reader.into_inner();
|
||||
let tls_stream = wrap_tls_stream(stream, config.host.as_deref().unwrap_or_default())?;
|
||||
let mut reader = std::io::BufReader::new(tls_stream);
|
||||
return smtp_probe(&mut reader, &config);
|
||||
}
|
||||
|
||||
smtp_authenticate(&mut reader, &config)?;
|
||||
let _ = smtp_send_command(&mut reader, "QUIT", &[221]);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn translate_smtp_error(error: &str) -> String {
|
||||
let error_lower = error.to_ascii_lowercase();
|
||||
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
use super::{
|
||||
decrypt_catalog_secret_with_fallbacks, escape_admin_email_template_html, json,
|
||||
read_admin_email_template_payload, render_admin_email_template_html, system_config_bool,
|
||||
system_config_string, system_config_u16, AppState, GatewayError,
|
||||
escape_admin_email_template_html, json, read_admin_email_template_payload,
|
||||
render_admin_email_template_html, system_config_string, AppState, GatewayError,
|
||||
AUTH_EMAIL_VERIFICATION_PREFIX, AUTH_EMAIL_VERIFIED_PREFIX, AUTH_EMAIL_VERIFIED_TTL_SECS,
|
||||
AUTH_SMTP_TIMEOUT_SECS,
|
||||
};
|
||||
use base64::Engine;
|
||||
use crate::email_delivery::{
|
||||
read_smtp_delivery_config, send_smtp_email, ComposedEmail, SmtpDeliveryConfig,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub(super) struct StoredAuthEmailVerificationCode {
|
||||
@@ -13,25 +13,8 @@ pub(super) struct StoredAuthEmailVerificationCode {
|
||||
pub(super) created_at: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct AuthSmtpConfig {
|
||||
pub(super) host: String,
|
||||
pub(super) port: u16,
|
||||
pub(super) user: Option<String>,
|
||||
pub(super) password: Option<String>,
|
||||
pub(super) use_tls: bool,
|
||||
pub(super) use_ssl: bool,
|
||||
pub(super) from_email: String,
|
||||
pub(super) from_name: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct AuthComposedEmail {
|
||||
pub(super) to_email: String,
|
||||
pub(super) subject: String,
|
||||
pub(super) html_body: String,
|
||||
pub(super) text_body: String,
|
||||
}
|
||||
pub(super) type AuthSmtpConfig = SmtpDeliveryConfig;
|
||||
pub(super) type AuthComposedEmail = ComposedEmail;
|
||||
|
||||
pub(super) fn auth_email_verification_key(email: &str) -> String {
|
||||
format!("{AUTH_EMAIL_VERIFICATION_PREFIX}{email}")
|
||||
@@ -84,25 +67,6 @@ fn render_auth_template_string(
|
||||
Ok(rendered)
|
||||
}
|
||||
|
||||
fn auth_encode_mime_header(value: &str) -> String {
|
||||
if value.is_ascii() {
|
||||
return value.to_string();
|
||||
}
|
||||
format!(
|
||||
"=?UTF-8?B?{}?=",
|
||||
base64::engine::general_purpose::STANDARD.encode(value.as_bytes())
|
||||
)
|
||||
}
|
||||
|
||||
fn auth_wrap_base64(value: &str) -> String {
|
||||
let mut wrapped = String::new();
|
||||
for chunk in value.as_bytes().chunks(76) {
|
||||
wrapped.push_str(std::str::from_utf8(chunk).unwrap_or_default());
|
||||
wrapped.push_str("\r\n");
|
||||
}
|
||||
wrapped
|
||||
}
|
||||
|
||||
fn auth_build_verification_text_body(
|
||||
app_name: &str,
|
||||
email: &str,
|
||||
@@ -114,244 +78,6 @@ fn auth_build_verification_text_body(
|
||||
)
|
||||
}
|
||||
|
||||
fn auth_build_tls_config() -> std::sync::Arc<rustls::ClientConfig> {
|
||||
let _ = rustls::crypto::ring::default_provider().install_default();
|
||||
let root_store =
|
||||
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
|
||||
let config = rustls::ClientConfig::builder()
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth();
|
||||
std::sync::Arc::new(config)
|
||||
}
|
||||
|
||||
fn auth_resolve_server_name(
|
||||
host: &str,
|
||||
) -> Result<rustls::pki_types::ServerName<'static>, GatewayError> {
|
||||
let host = host.trim().trim_start_matches('[').trim_end_matches(']');
|
||||
if let Ok(ip) = host.parse::<std::net::IpAddr>() {
|
||||
return Ok(rustls::pki_types::ServerName::from(ip));
|
||||
}
|
||||
rustls::pki_types::ServerName::try_from(host.to_string())
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
fn auth_connect_tcp_stream(config: &AuthSmtpConfig) -> Result<std::net::TcpStream, GatewayError> {
|
||||
let stream = std::net::TcpStream::connect((config.host.as_str(), config.port))
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
stream
|
||||
.set_read_timeout(Some(std::time::Duration::from_secs(AUTH_SMTP_TIMEOUT_SECS)))
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
stream
|
||||
.set_write_timeout(Some(std::time::Duration::from_secs(AUTH_SMTP_TIMEOUT_SECS)))
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
fn auth_wrap_tls_stream(
|
||||
stream: std::net::TcpStream,
|
||||
host: &str,
|
||||
) -> Result<rustls::StreamOwned<rustls::ClientConnection, std::net::TcpStream>, GatewayError> {
|
||||
let server_name = auth_resolve_server_name(host)?;
|
||||
let connection = rustls::ClientConnection::new(auth_build_tls_config(), server_name)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(rustls::StreamOwned::new(connection, stream))
|
||||
}
|
||||
|
||||
fn auth_smtp_read_response<T: std::io::BufRead>(
|
||||
reader: &mut T,
|
||||
) -> Result<(u16, String), GatewayError> {
|
||||
let mut message = String::new();
|
||||
let code = loop {
|
||||
let parsed_code;
|
||||
let continuation;
|
||||
let trimmed;
|
||||
{
|
||||
let mut line = String::new();
|
||||
let bytes = reader
|
||||
.read_line(&mut line)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
if bytes == 0 {
|
||||
return Err(GatewayError::Internal(
|
||||
"smtp connection closed unexpectedly".to_string(),
|
||||
));
|
||||
}
|
||||
trimmed = line.trim_end_matches(['\r', '\n']).to_string();
|
||||
if trimmed.len() < 3 {
|
||||
return Err(GatewayError::Internal("invalid smtp response".to_string()));
|
||||
}
|
||||
parsed_code = trimmed[..3]
|
||||
.parse::<u16>()
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
continuation = trimmed.as_bytes().get(3).copied() == Some(b'-');
|
||||
}
|
||||
if !message.is_empty() {
|
||||
message.push('\n');
|
||||
}
|
||||
message.push_str(&trimmed);
|
||||
if !continuation {
|
||||
break parsed_code;
|
||||
}
|
||||
};
|
||||
Ok((code, message))
|
||||
}
|
||||
|
||||
fn auth_smtp_expect<T: std::io::BufRead>(
|
||||
reader: &mut T,
|
||||
allowed_codes: &[u16],
|
||||
) -> Result<String, GatewayError> {
|
||||
let (code, message) = auth_smtp_read_response(reader)?;
|
||||
if allowed_codes.contains(&code) {
|
||||
return Ok(message);
|
||||
}
|
||||
Err(GatewayError::Internal(format!(
|
||||
"unexpected smtp response {code}: {message}"
|
||||
)))
|
||||
}
|
||||
|
||||
fn auth_smtp_write_line<T: std::io::Write>(writer: &mut T, line: &str) -> Result<(), GatewayError> {
|
||||
writer
|
||||
.write_all(line.as_bytes())
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
writer
|
||||
.write_all(b"\r\n")
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
writer
|
||||
.flush()
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
fn auth_smtp_send_command<S: std::io::Read + std::io::Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
command: &str,
|
||||
allowed_codes: &[u16],
|
||||
) -> Result<String, GatewayError> {
|
||||
auth_smtp_write_line(reader.get_mut(), command)?;
|
||||
auth_smtp_expect(reader, allowed_codes)
|
||||
}
|
||||
|
||||
fn auth_build_email_message(config: &AuthSmtpConfig, email: &AuthComposedEmail) -> String {
|
||||
let boundary = format!("aether-{}", uuid::Uuid::new_v4().simple());
|
||||
let text_body = auth_wrap_base64(
|
||||
&base64::engine::general_purpose::STANDARD.encode(email.text_body.as_bytes()),
|
||||
);
|
||||
let html_body = auth_wrap_base64(
|
||||
&base64::engine::general_purpose::STANDARD.encode(email.html_body.as_bytes()),
|
||||
);
|
||||
let from_header = if config.from_name.trim().is_empty() {
|
||||
format!("<{}>", config.from_email)
|
||||
} else {
|
||||
format!(
|
||||
"{} <{}>",
|
||||
auth_encode_mime_header(config.from_name.trim()),
|
||||
config.from_email
|
||||
)
|
||||
};
|
||||
format!(
|
||||
"From: {from_header}\r\nTo: <{to_email}>\r\nSubject: {subject}\r\nMIME-Version: 1.0\r\nContent-Type: multipart/alternative; boundary=\"{boundary}\"\r\n\r\n--{boundary}\r\nContent-Type: text/plain; charset=\"utf-8\"\r\nContent-Transfer-Encoding: base64\r\n\r\n{text_body}--{boundary}\r\nContent-Type: text/html; charset=\"utf-8\"\r\nContent-Transfer-Encoding: base64\r\n\r\n{html_body}--{boundary}--\r\n",
|
||||
to_email = email.to_email,
|
||||
subject = auth_encode_mime_header(&email.subject),
|
||||
)
|
||||
}
|
||||
|
||||
fn auth_smtp_authenticate<S: std::io::Read + std::io::Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
config: &AuthSmtpConfig,
|
||||
) -> Result<(), GatewayError> {
|
||||
let Some(username) = config
|
||||
.user
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
let password = config.password.as_deref().unwrap_or("");
|
||||
auth_smtp_send_command(reader, "AUTH LOGIN", &[334])?;
|
||||
auth_smtp_send_command(
|
||||
reader,
|
||||
&base64::engine::general_purpose::STANDARD.encode(username.as_bytes()),
|
||||
&[334],
|
||||
)?;
|
||||
auth_smtp_send_command(
|
||||
reader,
|
||||
&base64::engine::general_purpose::STANDARD.encode(password.as_bytes()),
|
||||
&[235],
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn auth_smtp_deliver_message<S: std::io::Read + std::io::Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
config: &AuthSmtpConfig,
|
||||
email: &AuthComposedEmail,
|
||||
) -> Result<(), GatewayError> {
|
||||
auth_smtp_send_command(
|
||||
reader,
|
||||
&format!("MAIL FROM:<{}>", config.from_email),
|
||||
&[250],
|
||||
)?;
|
||||
auth_smtp_send_command(
|
||||
reader,
|
||||
&format!("RCPT TO:<{}>", email.to_email),
|
||||
&[250, 251],
|
||||
)?;
|
||||
auth_smtp_send_command(reader, "DATA", &[354])?;
|
||||
let message = auth_build_email_message(config, email);
|
||||
reader
|
||||
.get_mut()
|
||||
.write_all(message.as_bytes())
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
reader
|
||||
.get_mut()
|
||||
.write_all(b"\r\n.\r\n")
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
reader
|
||||
.get_mut()
|
||||
.flush()
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let _ = auth_smtp_expect(reader, &[250])?;
|
||||
let _ = auth_smtp_send_command(reader, "QUIT", &[221]);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn auth_smtp_send_message<S: std::io::Read + std::io::Write>(
|
||||
reader: &mut std::io::BufReader<S>,
|
||||
config: &AuthSmtpConfig,
|
||||
email: &AuthComposedEmail,
|
||||
) -> Result<(), GatewayError> {
|
||||
auth_smtp_send_command(reader, "EHLO aether.local", &[250])?;
|
||||
auth_smtp_authenticate(reader, config)?;
|
||||
auth_smtp_deliver_message(reader, config, email)
|
||||
}
|
||||
|
||||
fn send_auth_email_blocking(
|
||||
config: AuthSmtpConfig,
|
||||
email: AuthComposedEmail,
|
||||
) -> Result<(), GatewayError> {
|
||||
if config.use_ssl {
|
||||
let stream = auth_connect_tcp_stream(&config)?;
|
||||
let tls_stream = auth_wrap_tls_stream(stream, &config.host)?;
|
||||
let mut reader = std::io::BufReader::new(tls_stream);
|
||||
let _ = auth_smtp_expect(&mut reader, &[220])?;
|
||||
return auth_smtp_send_message(&mut reader, &config, &email);
|
||||
}
|
||||
|
||||
let stream = auth_connect_tcp_stream(&config)?;
|
||||
let mut reader = std::io::BufReader::new(stream);
|
||||
let _ = auth_smtp_expect(&mut reader, &[220])?;
|
||||
let _ = auth_smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
|
||||
if config.use_tls {
|
||||
let _ = auth_smtp_send_command(&mut reader, "STARTTLS", &[220])?;
|
||||
let stream = reader.into_inner();
|
||||
let tls_stream = auth_wrap_tls_stream(stream, &config.host)?;
|
||||
let mut reader = std::io::BufReader::new(tls_stream);
|
||||
return auth_smtp_send_message(&mut reader, &config, &email);
|
||||
}
|
||||
|
||||
auth_smtp_authenticate(&mut reader, &config)?;
|
||||
auth_smtp_deliver_message(&mut reader, &config, &email)
|
||||
}
|
||||
|
||||
pub(super) async fn read_auth_email_verification_code(
|
||||
state: &AppState,
|
||||
email: &str,
|
||||
@@ -423,40 +149,7 @@ pub(super) async fn store_auth_email_verification_code(
|
||||
pub(super) async fn read_auth_smtp_config(
|
||||
state: &AppState,
|
||||
) -> Result<Option<AuthSmtpConfig>, GatewayError> {
|
||||
let smtp_host = state.read_system_config_json_value("smtp_host").await?;
|
||||
let smtp_from_email = state
|
||||
.read_system_config_json_value("smtp_from_email")
|
||||
.await?;
|
||||
let Some(host) = system_config_string(smtp_host.as_ref()) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(from_email) = system_config_string(smtp_from_email.as_ref()) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let smtp_port = state.read_system_config_json_value("smtp_port").await?;
|
||||
let smtp_user = state.read_system_config_json_value("smtp_user").await?;
|
||||
let smtp_password = state.read_system_config_json_value("smtp_password").await?;
|
||||
let smtp_use_tls = state.read_system_config_json_value("smtp_use_tls").await?;
|
||||
let smtp_use_ssl = state.read_system_config_json_value("smtp_use_ssl").await?;
|
||||
let smtp_from_name = state
|
||||
.read_system_config_json_value("smtp_from_name")
|
||||
.await?;
|
||||
|
||||
let password = system_config_string(smtp_password.as_ref()).map(|value| {
|
||||
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &value).unwrap_or(value)
|
||||
});
|
||||
|
||||
Ok(Some(AuthSmtpConfig {
|
||||
host,
|
||||
port: system_config_u16(smtp_port.as_ref(), 587),
|
||||
user: system_config_string(smtp_user.as_ref()),
|
||||
password,
|
||||
use_tls: system_config_bool(smtp_use_tls.as_ref(), true),
|
||||
use_ssl: system_config_bool(smtp_use_ssl.as_ref(), false),
|
||||
from_email,
|
||||
from_name: system_config_string(smtp_from_name.as_ref())
|
||||
.unwrap_or_else(|| "Aether".to_string()),
|
||||
}))
|
||||
read_smtp_delivery_config(state).await
|
||||
}
|
||||
|
||||
pub(super) async fn auth_email_app_name(state: &AppState) -> Result<String, GatewayError> {
|
||||
@@ -516,27 +209,20 @@ pub(super) async fn send_auth_email(
|
||||
if record_auth_email_delivery_for_tests(
|
||||
state,
|
||||
json!({
|
||||
"to_email": email.to_email,
|
||||
"subject": email.subject,
|
||||
"html_body": email.html_body,
|
||||
"text_body": email.text_body,
|
||||
"to_email": email.to_email.clone(),
|
||||
"subject": email.subject.clone(),
|
||||
"html_body": email.html_body.clone(),
|
||||
"text_body": email.text_body.clone(),
|
||||
}),
|
||||
) {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
tokio::task::spawn_blocking(move || send_auth_email_blocking(config, email))
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
send_smtp_email(config, email).await
|
||||
}
|
||||
|
||||
pub(super) async fn auth_registration_email_configured(
|
||||
state: &AppState,
|
||||
) -> Result<bool, GatewayError> {
|
||||
let smtp_host = state.read_system_config_json_value("smtp_host").await?;
|
||||
let smtp_from_email = state
|
||||
.read_system_config_json_value("smtp_from_email")
|
||||
.await?;
|
||||
Ok(system_config_string(smtp_host.as_ref()).is_some()
|
||||
&& system_config_string(smtp_from_email.as_ref()).is_some())
|
||||
Ok(read_smtp_delivery_config(state).await?.is_some())
|
||||
}
|
||||
|
||||
@@ -121,7 +121,6 @@ pub(super) const AUTH_REFRESH_TOKEN_EXPIRATION_DAYS: i64 = 7;
|
||||
pub(super) const AUTH_EMAIL_VERIFICATION_PREFIX: &str = "email:verification:";
|
||||
pub(super) const AUTH_EMAIL_VERIFIED_PREFIX: &str = "email:verified:";
|
||||
pub(super) const AUTH_EMAIL_VERIFIED_TTL_SECS: u64 = 3600;
|
||||
pub(super) const AUTH_SMTP_TIMEOUT_SECS: u64 = 30;
|
||||
|
||||
pub(crate) fn build_auth_json_response(
|
||||
status: http::StatusCode,
|
||||
|
||||
@@ -7,7 +7,9 @@ use axum::{
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
|
||||
use crate::handlers::shared::{deserialize_optional_json_patch, normalize_feature_settings};
|
||||
use crate::handlers::shared::{
|
||||
deserialize_optional_json_patch, normalize_user_self_feature_settings_update,
|
||||
};
|
||||
|
||||
use super::{
|
||||
auth_password_policy_level, build_auth_error_response, resolve_authenticated_local_user,
|
||||
@@ -65,12 +67,24 @@ pub(super) async fn handle_users_me_detail_put(
|
||||
let email = normalize_users_me_optional_non_empty_string(payload.email);
|
||||
let username = normalize_users_me_optional_non_empty_string(payload.username);
|
||||
let feature_settings = match payload.feature_settings {
|
||||
Some(value) => match normalize_feature_settings(value) {
|
||||
Ok(value) => Some(value),
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
Some(value) => {
|
||||
let current = match state.read_user_feature_settings(&auth.user.id).await {
|
||||
Ok(value) => value,
|
||||
Err(err) => {
|
||||
return build_auth_error_response(
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("user feature settings lookup failed: {err:?}"),
|
||||
false,
|
||||
)
|
||||
}
|
||||
};
|
||||
match normalize_user_self_feature_settings_update(value, current) {
|
||||
Ok(value) => Some(value),
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
None => None,
|
||||
};
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use aether_ai_serving::UPSTREAM_IS_STREAM_KEY;
|
||||
use aether_billing::{
|
||||
normalize_input_tokens_for_billing, normalize_total_input_context_for_cache_hit_rate,
|
||||
};
|
||||
@@ -314,7 +315,7 @@ fn users_me_usage_upstream_is_stream(item: &StoredRequestUsageAudit) -> bool {
|
||||
item.request_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| metadata.get("upstream_is_stream"))
|
||||
.and_then(|metadata| metadata.get(UPSTREAM_IS_STREAM_KEY))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.or_else(|| users_me_usage_headers_stream_flag(item.response_headers.as_ref()))
|
||||
.or_else(|| users_me_usage_infer_upstream_stream_from_captured_bodies(item))
|
||||
|
||||
@@ -1132,6 +1132,278 @@ fn build_chatgpt_web_quota_status_snapshot(
|
||||
}))
|
||||
}
|
||||
|
||||
fn windsurf_percent_quota_window_snapshot(
|
||||
metadata: &Map<String, Value>,
|
||||
code: &str,
|
||||
label: &str,
|
||||
remaining_percent_key: &str,
|
||||
reset_at_key: &str,
|
||||
observed_at_unix_secs: Option<u64>,
|
||||
) -> Option<Value> {
|
||||
let remaining_percent = metadata
|
||||
.get(remaining_percent_key)
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let reset_at = provider_quota_timestamp_unix_secs(metadata.get(reset_at_key));
|
||||
if remaining_percent.is_none() && reset_at.is_none() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let remaining_ratio = remaining_percent.map(|value| (value / 100.0).clamp(0.0, 1.0));
|
||||
let used_ratio = remaining_ratio.map(|value| (1.0 - value).clamp(0.0, 1.0));
|
||||
let reset_seconds = quota_window_reset_seconds(observed_at_unix_secs, reset_at);
|
||||
|
||||
Some(json!({
|
||||
"code": code,
|
||||
"label": label,
|
||||
"scope": "account",
|
||||
"unit": "percent",
|
||||
"used_ratio": used_ratio,
|
||||
"remaining_ratio": remaining_ratio,
|
||||
"reset_at": reset_at,
|
||||
"reset_seconds": reset_seconds,
|
||||
"is_exhausted": remaining_ratio.map(|value| value <= 1e-6),
|
||||
}))
|
||||
}
|
||||
|
||||
fn windsurf_count_quota_window_snapshot(
|
||||
metadata: &Map<String, Value>,
|
||||
code: &str,
|
||||
label: &str,
|
||||
used_key: &str,
|
||||
limit_key: &str,
|
||||
remaining_key: &str,
|
||||
) -> Option<Value> {
|
||||
let used = metadata
|
||||
.get(used_key)
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let limit = metadata
|
||||
.get(limit_key)
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64);
|
||||
let remaining = metadata
|
||||
.get(remaining_key)
|
||||
.and_then(admin_provider_quota_pure::coerce_json_f64)
|
||||
.or_else(|| limit.zip(used).map(|(limit, used)| (limit - used).max(0.0)));
|
||||
if used.is_none() && limit.is_none() && remaining.is_none() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let used_ratio = used
|
||||
.zip(limit)
|
||||
.and_then(|(used, limit)| (limit > 0.0).then_some((used / limit).clamp(0.0, 1.0)));
|
||||
let remaining_ratio = remaining.zip(limit).and_then(|(remaining, limit)| {
|
||||
(limit > 0.0).then_some((remaining / limit).clamp(0.0, 1.0))
|
||||
});
|
||||
|
||||
Some(json!({
|
||||
"code": code,
|
||||
"label": label,
|
||||
"scope": "account",
|
||||
"unit": "count",
|
||||
"used_ratio": used_ratio,
|
||||
"remaining_ratio": remaining_ratio,
|
||||
"used_value": used,
|
||||
"remaining_value": remaining,
|
||||
"limit_value": limit,
|
||||
"is_exhausted": remaining.is_some_and(|value| value <= 0.0),
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_windsurf_quota_status_snapshot(
|
||||
upstream_metadata: Option<&Value>,
|
||||
source: &str,
|
||||
) -> Option<Value> {
|
||||
let metadata = provider_quota_metadata_bucket(upstream_metadata, "windsurf")?;
|
||||
let observed_at_unix_secs = provider_quota_timestamp_unix_secs(metadata.get("updated_at"));
|
||||
let plan_type = metadata
|
||||
.get("plan_name")
|
||||
.or_else(|| metadata.get("plan_type"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let rate_limit = metadata
|
||||
.get("rate_limit")
|
||||
.cloned()
|
||||
.filter(|value| !value.is_null());
|
||||
let last_error = metadata
|
||||
.get("last_error")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let banned = metadata
|
||||
.get("banned")
|
||||
.or_else(|| metadata.get("is_banned"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||
== Some(true);
|
||||
let quarantined = metadata
|
||||
.get("quarantined")
|
||||
.or_else(|| metadata.get("is_quarantined"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||
== Some(true);
|
||||
|
||||
let mut windows = [
|
||||
windsurf_percent_quota_window_snapshot(
|
||||
metadata,
|
||||
"daily",
|
||||
"日",
|
||||
"daily_remaining_percent",
|
||||
"daily_reset_at",
|
||||
observed_at_unix_secs,
|
||||
),
|
||||
windsurf_percent_quota_window_snapshot(
|
||||
metadata,
|
||||
"weekly",
|
||||
"周",
|
||||
"weekly_remaining_percent",
|
||||
"weekly_reset_at",
|
||||
observed_at_unix_secs,
|
||||
),
|
||||
windsurf_count_quota_window_snapshot(
|
||||
metadata,
|
||||
"prompt",
|
||||
"Prompt",
|
||||
"prompt_used",
|
||||
"prompt_limit",
|
||||
"prompt_remaining",
|
||||
),
|
||||
windsurf_count_quota_window_snapshot(
|
||||
metadata,
|
||||
"flex",
|
||||
"Flex",
|
||||
"flex_used",
|
||||
"flex_limit",
|
||||
"flex_remaining",
|
||||
),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let mut rate_limit_cooling = false;
|
||||
let mut rate_limit_reset_seconds = None;
|
||||
let mut rate_limit_reason = None::<String>;
|
||||
if let Some(rate_limit) = rate_limit.as_ref() {
|
||||
if let Some(rate_limit_object) = rate_limit.as_object() {
|
||||
let retry_after_ms = rate_limit_object
|
||||
.get("retry_after_ms")
|
||||
.or_else(|| rate_limit_object.get("retryAfterMs"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_u64)
|
||||
.filter(|value| *value > 0);
|
||||
if let Some(retry_after_ms) = retry_after_ms {
|
||||
rate_limit_cooling = true;
|
||||
rate_limit_reset_seconds = Some(retry_after_ms.saturating_add(999) / 1000);
|
||||
rate_limit_reason = rate_limit_object
|
||||
.get("message")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
windows.push(json!({
|
||||
"code": "rate_limit",
|
||||
"label": "速率",
|
||||
"scope": "account",
|
||||
"unit": "count",
|
||||
"is_exhausted": false,
|
||||
"reset_seconds": rate_limit_reset_seconds,
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let allowed_models_count = metadata
|
||||
.get("allowed_models_count")
|
||||
.or_else(|| metadata.get("models_count"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_u64);
|
||||
|
||||
if windows.is_empty()
|
||||
&& plan_type.is_none()
|
||||
&& observed_at_unix_secs.is_none()
|
||||
&& rate_limit.is_none()
|
||||
&& allowed_models_count.is_none()
|
||||
&& !banned
|
||||
&& !quarantined
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let usage_ratio = quota_windows_usage_ratio(&windows);
|
||||
let reset_seconds = if rate_limit_cooling {
|
||||
rate_limit_reset_seconds.or_else(|| quota_windows_min_reset_seconds(&windows))
|
||||
} else {
|
||||
quota_windows_min_reset_seconds(&windows)
|
||||
};
|
||||
let reset_at = if rate_limit_cooling {
|
||||
None
|
||||
} else {
|
||||
quota_windows_min_reset_at(&windows)
|
||||
};
|
||||
let exhausted_by_window = windows.iter().filter_map(Value::as_object).any(|window| {
|
||||
window
|
||||
.get("code")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|code| {
|
||||
code.eq_ignore_ascii_case("daily") || code.eq_ignore_ascii_case("weekly")
|
||||
})
|
||||
&& window
|
||||
.get("is_exhausted")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||
.unwrap_or(false)
|
||||
});
|
||||
let exhausted = banned || quarantined || exhausted_by_window;
|
||||
let (code, label, reason) = if banned {
|
||||
(
|
||||
"banned",
|
||||
Some("账号已封禁"),
|
||||
last_error
|
||||
.clone()
|
||||
.or_else(|| Some("账号被 Windsurf 标记为不可用".to_string())),
|
||||
)
|
||||
} else if quarantined {
|
||||
(
|
||||
"quarantined",
|
||||
Some("账号隔离中"),
|
||||
last_error
|
||||
.clone()
|
||||
.or_else(|| Some("账号处于隔离状态".to_string())),
|
||||
)
|
||||
} else if rate_limit_cooling {
|
||||
(
|
||||
"cooldown",
|
||||
Some("冷却中"),
|
||||
last_error.clone().or(rate_limit_reason),
|
||||
)
|
||||
} else if exhausted {
|
||||
(
|
||||
"exhausted",
|
||||
Some("额度耗尽"),
|
||||
Some("额度窗口已耗尽".to_string()),
|
||||
)
|
||||
} else {
|
||||
("ok", None, last_error)
|
||||
};
|
||||
|
||||
Some(json!({
|
||||
"version": 2,
|
||||
"provider_type": "windsurf",
|
||||
"code": code,
|
||||
"label": label,
|
||||
"reason": reason,
|
||||
"freshness": "fresh",
|
||||
"source": source,
|
||||
"observed_at": observed_at_unix_secs,
|
||||
"exhausted": exhausted,
|
||||
"usage_ratio": usage_ratio,
|
||||
"updated_at": observed_at_unix_secs,
|
||||
"reset_at": reset_at,
|
||||
"reset_seconds": reset_seconds,
|
||||
"plan_type": plan_type,
|
||||
"allowed_models_count": allowed_models_count,
|
||||
"rate_limit": rate_limit.unwrap_or(Value::Null),
|
||||
"windows": windows,
|
||||
}))
|
||||
}
|
||||
|
||||
fn build_antigravity_quota_status_snapshot(
|
||||
upstream_metadata: Option<&Value>,
|
||||
source: &str,
|
||||
@@ -1516,6 +1788,7 @@ pub(crate) fn sync_provider_key_quota_status_snapshot(
|
||||
"codex" => build_codex_quota_status_snapshot(upstream_metadata, source),
|
||||
"kiro" => build_kiro_quota_status_snapshot(upstream_metadata, source),
|
||||
"chatgpt_web" => build_chatgpt_web_quota_status_snapshot(upstream_metadata, source),
|
||||
"windsurf" => build_windsurf_quota_status_snapshot(upstream_metadata, source),
|
||||
"antigravity" => build_antigravity_quota_status_snapshot(upstream_metadata, source),
|
||||
"grok" => build_grok_quota_status_snapshot(upstream_metadata, source),
|
||||
"gemini_cli" => build_gemini_cli_quota_status_snapshot(upstream_metadata, source),
|
||||
@@ -1552,6 +1825,12 @@ fn quota_snapshot_has_materialized_data(
|
||||
return false;
|
||||
}
|
||||
|
||||
if normalized_provider_type == "windsurf"
|
||||
&& windsurf_quota_snapshot_has_stale_cooldown(quota_snapshot)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
if quota_snapshot
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
@@ -1577,6 +1856,61 @@ fn quota_snapshot_has_materialized_data(
|
||||
})
|
||||
}
|
||||
|
||||
fn windsurf_quota_snapshot_has_stale_cooldown(quota_snapshot: &Map<String, Value>) -> bool {
|
||||
let code = quota_snapshot
|
||||
.get("code")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.unwrap_or_default();
|
||||
if !code.eq_ignore_ascii_case("cooldown") {
|
||||
return false;
|
||||
}
|
||||
|
||||
let rate_limit = quota_snapshot.get("rate_limit").and_then(Value::as_object);
|
||||
let retry_after_ms = rate_limit
|
||||
.and_then(|rate_limit| {
|
||||
rate_limit
|
||||
.get("retry_after_ms")
|
||||
.or_else(|| rate_limit.get("retryAfterMs"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_u64)
|
||||
})
|
||||
.unwrap_or(0);
|
||||
if retry_after_ms > 0 {
|
||||
return false;
|
||||
}
|
||||
|
||||
let has_positive_rate_limit_reset = quota_snapshot
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|windows| {
|
||||
windows.iter().filter_map(Value::as_object).any(|window| {
|
||||
window
|
||||
.get("code")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|code| code.eq_ignore_ascii_case("rate_limit"))
|
||||
&& window
|
||||
.get("reset_seconds")
|
||||
.or_else(|| window.get("reset_at"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_u64)
|
||||
.is_some_and(|value| value > 0)
|
||||
})
|
||||
});
|
||||
if has_positive_rate_limit_reset {
|
||||
return false;
|
||||
}
|
||||
|
||||
let exhausted = quota_snapshot
|
||||
.get("exhausted")
|
||||
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||
.unwrap_or(false);
|
||||
let has_capacity = rate_limit
|
||||
.and_then(|rate_limit| rate_limit.get("has_capacity"))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_bool)
|
||||
.unwrap_or(false);
|
||||
|
||||
has_capacity || !exhausted
|
||||
}
|
||||
|
||||
pub(crate) fn provider_key_status_snapshot_payload(
|
||||
key: &StoredProviderCatalogKey,
|
||||
provider_type: &str,
|
||||
@@ -2504,6 +2838,274 @@ mod tests {
|
||||
assert_eq!(windows[0].get("remaining_ratio"), Some(&json!(0.75)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_backfills_windsurf_daily_and_weekly_quota() {
|
||||
let mut key = sample_catalog_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"windsurf": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"plan_name": "Pro",
|
||||
"daily_remaining_percent": 40.0,
|
||||
"weekly_remaining_percent": 65.0,
|
||||
"daily_reset_at": 1_778_100_000u64,
|
||||
"weekly_reset_at": 1_778_600_000u64,
|
||||
"prompt_used": 12.0,
|
||||
"prompt_limit": 100.0,
|
||||
"prompt_remaining": 88.0,
|
||||
"flex_used": 3.0,
|
||||
"flex_limit": 10.0,
|
||||
"flex_remaining": 7.0,
|
||||
"allowed_models_count": 82
|
||||
}
|
||||
}));
|
||||
|
||||
let payload = provider_key_status_snapshot_payload(&key, "windsurf");
|
||||
let quota = payload
|
||||
.get("quota")
|
||||
.and_then(Value::as_object)
|
||||
.expect("quota snapshot should be object");
|
||||
let windows = quota
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
.expect("windsurf quota windows should exist");
|
||||
let daily = windows
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.find(|window| window.get("code") == Some(&json!("daily")))
|
||||
.expect("daily quota window should exist");
|
||||
let weekly = windows
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.find(|window| window.get("code") == Some(&json!("weekly")))
|
||||
.expect("weekly quota window should exist");
|
||||
|
||||
assert_eq!(quota.get("provider_type"), Some(&json!("windsurf")));
|
||||
assert_eq!(quota.get("code"), Some(&json!("ok")));
|
||||
assert_eq!(quota.get("plan_type"), Some(&json!("Pro")));
|
||||
assert_eq!(quota.get("usage_ratio"), Some(&json!(0.6)));
|
||||
assert_eq!(quota.get("reset_at"), Some(&json!(1_778_100_000u64)));
|
||||
assert_eq!(daily.get("remaining_ratio"), Some(&json!(0.4)));
|
||||
assert_eq!(daily.get("used_ratio"), Some(&json!(0.6)));
|
||||
assert_eq!(daily.get("reset_seconds"), Some(&json!(32_754u64)));
|
||||
assert_eq!(weekly.get("remaining_ratio"), Some(&json!(0.65)));
|
||||
assert_eq!(weekly.get("used_ratio"), Some(&json!(0.35)));
|
||||
assert_eq!(weekly.get("reset_seconds"), Some(&json!(532_754u64)));
|
||||
assert_eq!(quota.get("allowed_models_count"), Some(&json!(82)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_treats_windsurf_rate_limit_as_cooldown() {
|
||||
let mut key = sample_catalog_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"windsurf": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"daily_remaining_percent": 80.0,
|
||||
"rate_limit": {
|
||||
"limited": true,
|
||||
"retry_after_ms": 60_001u64,
|
||||
"message": "slow down"
|
||||
},
|
||||
"last_error": "slow down"
|
||||
}
|
||||
}));
|
||||
|
||||
let payload = provider_key_status_snapshot_payload(&key, "windsurf");
|
||||
let quota = payload
|
||||
.get("quota")
|
||||
.and_then(Value::as_object)
|
||||
.expect("quota snapshot should be object");
|
||||
let rate_window = quota
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
.and_then(|windows| {
|
||||
windows
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.find(|window| window.get("code") == Some(&json!("rate_limit")))
|
||||
})
|
||||
.expect("rate limit window should exist");
|
||||
|
||||
assert_eq!(quota.get("code"), Some(&json!("cooldown")));
|
||||
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
|
||||
assert_eq!(quota.get("reset_seconds"), Some(&json!(61u64)));
|
||||
assert_eq!(rate_window.get("is_exhausted"), Some(&json!(false)));
|
||||
assert_eq!(rate_window.get("reset_seconds"), Some(&json!(61u64)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_keeps_windsurf_capacity_probe_without_retry_after_ok() {
|
||||
let mut key = sample_catalog_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"windsurf": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"daily_remaining_percent": 100.0,
|
||||
"weekly_remaining_percent": 100.0,
|
||||
"rate_limit": {
|
||||
"limited": true,
|
||||
"has_capacity": false,
|
||||
"messages_remaining": 0.0,
|
||||
"max_messages": 100.0
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
let payload = provider_key_status_snapshot_payload(&key, "windsurf");
|
||||
let quota = payload
|
||||
.get("quota")
|
||||
.and_then(Value::as_object)
|
||||
.expect("quota snapshot should be object");
|
||||
let has_rate_limit_window =
|
||||
quota
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|windows| {
|
||||
windows
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.any(|window| window.get("code") == Some(&json!("rate_limit")))
|
||||
});
|
||||
|
||||
assert_eq!(quota.get("code"), Some(&json!("ok")));
|
||||
assert_eq!(quota.get("label"), Some(&Value::Null));
|
||||
assert_eq!(quota.get("exhausted"), Some(&json!(false)));
|
||||
assert_eq!(
|
||||
payload.pointer("/quota/rate_limit/limited"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
assert!(!has_rate_limit_window);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_refreshes_stale_windsurf_cooldown_when_probe_has_capacity(
|
||||
) {
|
||||
let mut key = sample_catalog_key();
|
||||
key.status_snapshot = Some(json!({
|
||||
"quota": {
|
||||
"version": 2,
|
||||
"provider_type": "windsurf",
|
||||
"code": "cooldown",
|
||||
"label": "冷却中",
|
||||
"exhausted": false,
|
||||
"windows": [
|
||||
{
|
||||
"code": "daily",
|
||||
"unit": "percent",
|
||||
"label": "日",
|
||||
"scope": "account",
|
||||
"remaining_ratio": 0.99,
|
||||
"is_exhausted": false
|
||||
},
|
||||
{
|
||||
"code": "rate_limit",
|
||||
"unit": "count",
|
||||
"label": "速率",
|
||||
"scope": "account",
|
||||
"is_exhausted": false,
|
||||
"reset_seconds": null
|
||||
}
|
||||
],
|
||||
"rate_limit": {
|
||||
"limited": true,
|
||||
"has_capacity": true,
|
||||
"messages_remaining": -1,
|
||||
"max_messages": -1
|
||||
}
|
||||
}
|
||||
}));
|
||||
key.upstream_metadata = Some(json!({
|
||||
"windsurf": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"daily_remaining_percent": 99.0,
|
||||
"weekly_remaining_percent": 100.0,
|
||||
"allowed_models_count": 118,
|
||||
"rate_limit": {
|
||||
"limited": true,
|
||||
"has_capacity": true,
|
||||
"messages_remaining": -1,
|
||||
"max_messages": -1
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
let payload = provider_key_status_snapshot_payload(&key, "windsurf");
|
||||
let quota = payload
|
||||
.get("quota")
|
||||
.and_then(Value::as_object)
|
||||
.expect("quota snapshot should be object");
|
||||
let has_rate_limit_window =
|
||||
quota
|
||||
.get("windows")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|windows| {
|
||||
windows
|
||||
.iter()
|
||||
.filter_map(Value::as_object)
|
||||
.any(|window| window.get("code") == Some(&json!("rate_limit")))
|
||||
});
|
||||
|
||||
assert_eq!(quota.get("code"), Some(&json!("ok")));
|
||||
assert_eq!(quota.get("label"), Some(&Value::Null));
|
||||
assert_eq!(quota.get("allowed_models_count"), Some(&json!(118u64)));
|
||||
assert!(!has_rate_limit_window);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_marks_windsurf_banned_and_quarantined_blocking() {
|
||||
let mut banned_key = sample_catalog_key();
|
||||
banned_key.upstream_metadata = Some(json!({
|
||||
"windsurf": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"banned": true,
|
||||
"reason": "forbidden"
|
||||
}
|
||||
}));
|
||||
let banned_payload = provider_key_status_snapshot_payload(&banned_key, "windsurf");
|
||||
|
||||
assert_eq!(
|
||||
banned_payload.pointer("/quota/code"),
|
||||
Some(&json!("banned"))
|
||||
);
|
||||
assert_eq!(
|
||||
banned_payload.pointer("/quota/exhausted"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
assert_eq!(
|
||||
banned_payload.pointer("/account/code"),
|
||||
Some(&json!("account_banned"))
|
||||
);
|
||||
assert_eq!(
|
||||
banned_payload.pointer("/account/blocked"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
|
||||
let mut quarantined_key = sample_catalog_key();
|
||||
quarantined_key.upstream_metadata = Some(json!({
|
||||
"windsurf": {
|
||||
"updated_at": 1_778_067_246u64,
|
||||
"quarantined": true
|
||||
}
|
||||
}));
|
||||
let quarantined_payload =
|
||||
provider_key_status_snapshot_payload(&quarantined_key, "windsurf");
|
||||
|
||||
assert_eq!(
|
||||
quarantined_payload.pointer("/quota/code"),
|
||||
Some(&json!("quarantined"))
|
||||
);
|
||||
assert_eq!(
|
||||
quarantined_payload.pointer("/quota/exhausted"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
assert_eq!(
|
||||
quarantined_payload.pointer("/account/code"),
|
||||
Some(&json!("account_quarantined"))
|
||||
);
|
||||
assert_eq!(
|
||||
quarantined_payload.pointer("/account/blocked"),
|
||||
Some(&json!(true))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_key_status_snapshot_payload_preserves_existing_materialized_quota_snapshot() {
|
||||
let mut key = sample_catalog_key();
|
||||
|
||||
@@ -38,7 +38,8 @@ pub(crate) use self::external_models::OFFICIAL_EXTERNAL_MODEL_PROVIDERS;
|
||||
pub(crate) use self::normalize::{
|
||||
deserialize_optional_json_patch, deserialize_optional_string_list_patch, ip_rules_allow,
|
||||
json_ip_rules_allow, normalize_feature_settings, normalize_ip_rules, normalize_json_array,
|
||||
normalize_json_object, normalize_string_list, parse_json_ip_rules,
|
||||
normalize_json_object, normalize_string_list, normalize_user_self_feature_settings_update,
|
||||
parse_json_ip_rules,
|
||||
};
|
||||
pub(crate) use self::payloads::{
|
||||
InternalGatewayAuthContextRequest, InternalGatewayExecuteRequest,
|
||||
|
||||
@@ -54,6 +54,7 @@ pub(crate) fn normalize_feature_settings(value: Option<Value>) -> Result<Option<
|
||||
Value::Null => Ok(None),
|
||||
Value::Object(ref mut settings) => {
|
||||
normalize_chat_pii_redaction_feature_settings(settings)?;
|
||||
normalize_notification_push_service_feature_settings(settings)?;
|
||||
if settings.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
@@ -64,6 +65,41 @@ pub(crate) fn normalize_feature_settings(value: Option<Value>) -> Result<Option<
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_user_self_feature_settings_update(
|
||||
value: Option<Value>,
|
||||
current: Option<Value>,
|
||||
) -> Result<Option<Value>, String> {
|
||||
let mut normalized = normalize_feature_settings(value)?;
|
||||
let current_notification_push_service = current
|
||||
.and_then(|value| match value {
|
||||
Value::Object(mut settings) => settings.remove("notification_push_service"),
|
||||
_ => None,
|
||||
})
|
||||
.and_then(|value| {
|
||||
let mut wrapper = Map::new();
|
||||
wrapper.insert("notification_push_service".to_string(), value);
|
||||
normalize_notification_push_service_feature_settings(&mut wrapper)
|
||||
.ok()
|
||||
.and_then(|_| wrapper.remove("notification_push_service"))
|
||||
});
|
||||
|
||||
match (&mut normalized, current_notification_push_service) {
|
||||
(Some(Value::Object(settings)), Some(value)) => {
|
||||
settings.insert("notification_push_service".to_string(), value);
|
||||
}
|
||||
(Some(Value::Object(settings)), None) => {
|
||||
settings.remove("notification_push_service");
|
||||
}
|
||||
(None, Some(value)) => {
|
||||
let mut settings = Map::new();
|
||||
settings.insert("notification_push_service".to_string(), value);
|
||||
normalized = Some(Value::Object(settings));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
Ok(normalized)
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_ip_rules(
|
||||
values: Option<Vec<String>>,
|
||||
) -> Result<Option<Vec<String>>, String> {
|
||||
@@ -326,9 +362,47 @@ fn normalize_chat_pii_redaction_feature_object(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn normalize_notification_push_service_feature_settings(
|
||||
settings: &mut Map<String, Value>,
|
||||
) -> Result<(), String> {
|
||||
let Some(value) = settings.get_mut("notification_push_service") else {
|
||||
return Ok(());
|
||||
};
|
||||
match value {
|
||||
Value::Null => {
|
||||
settings.remove("notification_push_service");
|
||||
Ok(())
|
||||
}
|
||||
Value::Object(feature) => {
|
||||
normalize_notification_push_service_feature_object(feature)?;
|
||||
if feature.is_empty() {
|
||||
settings.remove("notification_push_service");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
_ => Err("notification_push_service 必须是对象".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_notification_push_service_feature_object(
|
||||
feature: &mut Map<String, Value>,
|
||||
) -> Result<(), String> {
|
||||
for key in ["enabled"] {
|
||||
if let Some(value) = feature.get(key) {
|
||||
if !value.is_boolean() {
|
||||
return Err(format!("notification_push_service.{key} 必须是布尔值"));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{ip_rules_allow, json_ip_rules_allow, normalize_ip_rules, parse_json_ip_rules};
|
||||
use super::{
|
||||
ip_rules_allow, json_ip_rules_allow, normalize_feature_settings, normalize_ip_rules,
|
||||
normalize_user_self_feature_settings_update, parse_json_ip_rules,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::net::{IpAddr, Ipv4Addr};
|
||||
|
||||
@@ -358,6 +432,41 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_feature_settings_accepts_notification_push_service_permission() {
|
||||
let normalized = normalize_feature_settings(Some(json!({
|
||||
"notification_push_service": {"enabled": true}
|
||||
})))
|
||||
.expect("feature settings should normalize")
|
||||
.expect("feature settings should remain set");
|
||||
|
||||
assert_eq!(
|
||||
normalized["notification_push_service"]["enabled"],
|
||||
json!(true)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_self_feature_update_preserves_notification_push_permission() {
|
||||
let normalized = normalize_user_self_feature_settings_update(
|
||||
Some(json!({
|
||||
"chat_pii_redaction": {"enabled": true, "inject_model_instruction": false},
|
||||
"notification_push_service": {"enabled": false}
|
||||
})),
|
||||
Some(json!({
|
||||
"notification_push_service": {"enabled": true}
|
||||
})),
|
||||
)
|
||||
.expect("feature settings should normalize")
|
||||
.expect("feature settings should remain set");
|
||||
|
||||
assert_eq!(
|
||||
normalized["notification_push_service"]["enabled"],
|
||||
json!(true)
|
||||
);
|
||||
assert_eq!(normalized["chat_pii_redaction"]["enabled"], json!(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ip_rules_allow_applies_allow_rules_and_deny_overrides() {
|
||||
let rules = vec![
|
||||
|
||||
@@ -254,6 +254,11 @@ pub(crate) fn admin_proxy_local_requires_buffered_body(
|
||||
| (Some("system_manage"), http::Method::PUT, Some("config_set"))
|
||||
| (Some("system_manage"), http::Method::PUT, Some("email_template_set"))
|
||||
| (Some("system_manage"), http::Method::POST, Some("email_template_preview"))
|
||||
| (
|
||||
Some("system_manage"),
|
||||
http::Method::POST,
|
||||
Some("important_notification_test"),
|
||||
)
|
||||
| (
|
||||
Some("provider_models_manage"),
|
||||
http::Method::POST,
|
||||
|
||||
Reference in New Issue
Block a user