mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
feat(security): harden gateway boundaries and usage policies
Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change. Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
@@ -46,6 +46,8 @@ use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||
use tokio::sync::oneshot;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
fn codex_models_snapshot(
|
||||
api_key_id: &str,
|
||||
user_id: &str,
|
||||
@@ -234,6 +236,12 @@ fn codex_catalog_key(
|
||||
key_id: &str,
|
||||
allowed_models: &[&str],
|
||||
) -> StoredProviderCatalogKey {
|
||||
let bootstrap = AppState::new()
|
||||
.expect("bootstrap state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::disabled()
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
key_id.to_string(),
|
||||
provider_id.to_string(),
|
||||
@@ -245,7 +253,8 @@ fn codex_catalog_key(
|
||||
.expect("Codex key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!(["openai:responses"])),
|
||||
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "oauth-upstream-secret")
|
||||
bootstrap
|
||||
.seal_provider_catalog_key_api_key(provider_id, key_id, "oauth-upstream-secret")
|
||||
.expect("Codex test token should encrypt"),
|
||||
None,
|
||||
None,
|
||||
@@ -411,6 +420,84 @@ fn sample_gemini_video_task(
|
||||
}
|
||||
}
|
||||
|
||||
fn gemini_video_catalog_repository() -> Arc<InMemoryProviderCatalogReadRepository> {
|
||||
const PROVIDER_ID: &str = "provider-gemini-video-local-1";
|
||||
const ENDPOINT_ID: &str = "endpoint-gemini-video-local-1";
|
||||
const KEY_ID: &str = "key-gemini-video-local-1";
|
||||
let bootstrap = AppState::new()
|
||||
.expect("bootstrap state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::disabled()
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
);
|
||||
let provider = StoredProviderCatalogProvider::new(
|
||||
PROVIDER_ID.to_string(),
|
||||
"gemini-video".to_string(),
|
||||
Some("https://generativelanguage.googleapis.com".to_string()),
|
||||
"gemini".to_string(),
|
||||
)
|
||||
.expect("Gemini video provider should build")
|
||||
.with_transport_fields(
|
||||
true,
|
||||
false,
|
||||
false,
|
||||
None,
|
||||
Some(2),
|
||||
None,
|
||||
Some(20.0),
|
||||
None,
|
||||
None,
|
||||
);
|
||||
let endpoint = StoredProviderCatalogEndpoint::new(
|
||||
ENDPOINT_ID.to_string(),
|
||||
PROVIDER_ID.to_string(),
|
||||
"gemini:video".to_string(),
|
||||
Some("gemini".to_string()),
|
||||
Some("video".to_string()),
|
||||
true,
|
||||
)
|
||||
.expect("Gemini video endpoint should build")
|
||||
.with_transport_fields(
|
||||
"https://generativelanguage.googleapis.com".to_string(),
|
||||
None,
|
||||
None,
|
||||
Some(2),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("Gemini video endpoint transport should build");
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
KEY_ID.to_string(),
|
||||
PROVIDER_ID.to_string(),
|
||||
"prod".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("Gemini video key should build")
|
||||
.with_transport_fields(
|
||||
Some(json!(["gemini:video"])),
|
||||
bootstrap
|
||||
.seal_provider_catalog_key_api_key(PROVIDER_ID, KEY_ID, "sk-upstream-gemini-video")
|
||||
.expect("Gemini video api key should encrypt"),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("Gemini video key transport should build");
|
||||
Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![endpoint],
|
||||
vec![key],
|
||||
))
|
||||
}
|
||||
|
||||
struct PendingMinimalCandidateSelectionReadRepository;
|
||||
|
||||
impl PendingMinimalCandidateSelectionReadRepository {
|
||||
@@ -657,13 +744,19 @@ fn gateway_versioned_models_fail_closed_when_cached_auth_becomes_unusable_or_mis
|
||||
}
|
||||
|
||||
async fn run_versioned_models_auth_race_scenario() {
|
||||
let mut auth_race_snapshot = codex_models_snapshot(
|
||||
"key-codex-models-auth-race",
|
||||
"user-codex-models-auth-race",
|
||||
&["future-alias"],
|
||||
);
|
||||
// This scenario exercises API-key cache invalidation, not provider
|
||||
// allowlist resolution. Keep the catalog dependency absent so the warm
|
||||
// request remains on the local models route while the key is usable.
|
||||
auth_race_snapshot.user_allowed_providers = None;
|
||||
auth_race_snapshot.api_key_allowed_providers = None;
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-codex-models-auth-race")),
|
||||
codex_models_snapshot(
|
||||
"key-codex-models-auth-race",
|
||||
"user-codex-models-auth-race",
|
||||
&["future-alias"],
|
||||
),
|
||||
auth_race_snapshot,
|
||||
)]));
|
||||
let candidate_repository = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(
|
||||
Vec::new(),
|
||||
@@ -3477,7 +3570,10 @@ async fn gateway_does_not_locally_reject_image_model_name_on_chat_completions()
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_auth_api_key_data_reader_for_tests(auth_repository),
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository)
|
||||
.with_system_default_routing_group_for_tests(),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
@@ -3704,6 +3800,231 @@ async fn gateway_handles_gemini_operation_detail_without_hitting_fallback_probe(
|
||||
fallback_probe_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_hides_gemini_operation_detail_and_cancel_from_non_owner() {
|
||||
let fallback_probe_hits = Arc::new(Mutex::new(0usize));
|
||||
let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits);
|
||||
let fallback_probe = Router::new().route(
|
||||
"/{*path}",
|
||||
any(move |_request: Request| {
|
||||
let fallback_probe_hits_inner = Arc::clone(&fallback_probe_hits_clone);
|
||||
async move {
|
||||
*fallback_probe_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Json(json!({"proxied": true}))).into_response()
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let execution_runtime_hits = Arc::new(Mutex::new(0usize));
|
||||
let execution_runtime_hits_clone = Arc::clone(&execution_runtime_hits);
|
||||
let execution_runtime = Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(move |_request: Request| {
|
||||
let execution_runtime_hits_inner = Arc::clone(&execution_runtime_hits_clone);
|
||||
async move {
|
||||
*execution_runtime_hits_inner
|
||||
.lock()
|
||||
.expect("mutex should lock") += 1;
|
||||
Json(json!({
|
||||
"request_id": "unexpected-cross-user-cancel",
|
||||
"status_code": 200,
|
||||
"headers": {},
|
||||
"body": { "json_body": {} }
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-gemini-operation-non-owner")),
|
||||
unrestricted_models_snapshot(
|
||||
"key-gemini-operation-non-owner",
|
||||
"user-gemini-operation-non-owner",
|
||||
),
|
||||
)]));
|
||||
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
||||
repository
|
||||
.upsert(sample_gemini_video_task(
|
||||
"task-gemini-operation-owner",
|
||||
"opshort-owner-only",
|
||||
"user-gemini-operation-owner",
|
||||
"key-gemini-operation-owner",
|
||||
"operations/ext-owner-only",
|
||||
VideoTaskStatus::Submitted,
|
||||
))
|
||||
.await
|
||||
.expect("upsert should succeed");
|
||||
|
||||
let (_fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await;
|
||||
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||
let gateway = build_router_with_state(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests(
|
||||
auth_repository,
|
||||
Arc::clone(&repository),
|
||||
),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let detail_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/v1beta/operations/opshort-owner-only?key=sk-gemini-operation-non-owner"
|
||||
))
|
||||
.send()
|
||||
.await
|
||||
.expect("cross-user detail request should complete");
|
||||
assert_eq!(detail_response.status(), StatusCode::NOT_FOUND);
|
||||
assert_eq!(
|
||||
detail_response
|
||||
.json::<serde_json::Value>()
|
||||
.await
|
||||
.expect("json body should parse"),
|
||||
json!({ "detail": "Video task not found" })
|
||||
);
|
||||
|
||||
let cancel_response = client
|
||||
.post(format!(
|
||||
"{gateway_url}/v1beta/operations/opshort-owner-only:cancel"
|
||||
))
|
||||
.header("x-goog-api-key", "sk-gemini-operation-non-owner")
|
||||
.header(http::header::CONTENT_TYPE, "application/json")
|
||||
.body("{}")
|
||||
.send()
|
||||
.await
|
||||
.expect("cross-user cancel request should complete");
|
||||
assert_eq!(cancel_response.status(), StatusCode::NOT_FOUND);
|
||||
assert_eq!(
|
||||
cancel_response
|
||||
.json::<serde_json::Value>()
|
||||
.await
|
||||
.expect("json body should parse"),
|
||||
json!({ "detail": "Video task not found" })
|
||||
);
|
||||
|
||||
let stored = repository
|
||||
.find(VideoTaskLookupKey::Id("task-gemini-operation-owner"))
|
||||
.await
|
||||
.expect("task lookup should succeed")
|
||||
.expect("task should exist");
|
||||
assert_eq!(stored.status, VideoTaskStatus::Submitted);
|
||||
assert_eq!(
|
||||
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||
0
|
||||
);
|
||||
assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
fallback_probe_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_hides_gemini_video_file_from_non_owner_and_allows_owner_rotated_key() {
|
||||
let fallback_probe_hits = Arc::new(Mutex::new(0usize));
|
||||
let fallback_probe_hits_clone = Arc::clone(&fallback_probe_hits);
|
||||
let fallback_probe = Router::new().route(
|
||||
"/{*path}",
|
||||
any(move |_request: Request| {
|
||||
let fallback_probe_hits_inner = Arc::clone(&fallback_probe_hits_clone);
|
||||
async move {
|
||||
*fallback_probe_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Json(json!({"proxied": true}))).into_response()
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![
|
||||
(
|
||||
Some(hash_api_key("sk-gemini-video-file-foreign")),
|
||||
unrestricted_models_snapshot(
|
||||
"key-gemini-video-file-foreign",
|
||||
"user-gemini-video-file-foreign",
|
||||
),
|
||||
),
|
||||
(
|
||||
Some(hash_api_key("sk-gemini-video-file-owner")),
|
||||
unrestricted_models_snapshot(
|
||||
"key-gemini-video-file-owner-rotated",
|
||||
"user-gemini-video-file-owner",
|
||||
),
|
||||
),
|
||||
]));
|
||||
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
||||
let mut task = sample_gemini_video_task(
|
||||
"task-gemini-video-file-owner",
|
||||
"opshort-file-owner",
|
||||
"user-gemini-video-file-owner",
|
||||
"key-gemini-video-file-original",
|
||||
"operations/ext-file-owner",
|
||||
VideoTaskStatus::Completed,
|
||||
);
|
||||
task.provider_api_format = None;
|
||||
task.video_url = Some("https://8.8.8.8/video-owner.mp4".to_string());
|
||||
repository
|
||||
.upsert(task)
|
||||
.await
|
||||
.expect("upsert should succeed");
|
||||
|
||||
let (_unused_fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await;
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests(
|
||||
auth_repository,
|
||||
repository,
|
||||
),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("client should build");
|
||||
|
||||
let foreign_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/v1beta/files/aev_opshort-file-owner:download?alt=media&key=sk-gemini-video-file-foreign"
|
||||
))
|
||||
.send()
|
||||
.await
|
||||
.expect("foreign request should complete");
|
||||
assert_eq!(foreign_response.status(), StatusCode::NOT_FOUND);
|
||||
assert_eq!(
|
||||
foreign_response
|
||||
.json::<serde_json::Value>()
|
||||
.await
|
||||
.expect("json body should parse"),
|
||||
json!({"detail": "File not found"})
|
||||
);
|
||||
assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
let owner_response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/v1beta/files/aev_opshort-file-owner:download?alt=media&key=sk-gemini-video-file-owner"
|
||||
))
|
||||
.send()
|
||||
.await
|
||||
.expect("owner request should complete");
|
||||
assert_ne!(owner_response.status(), StatusCode::NOT_FOUND);
|
||||
assert_ne!(owner_response.status(), StatusCode::TEMPORARY_REDIRECT);
|
||||
assert_eq!(owner_response.status(), StatusCode::INTERNAL_SERVER_ERROR);
|
||||
assert_eq!(
|
||||
owner_response
|
||||
.json::<serde_json::Value>()
|
||||
.await
|
||||
.expect("json body should parse"),
|
||||
json!({"detail": "Service temporarily unavailable"})
|
||||
);
|
||||
assert_eq!(*fallback_probe_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
fallback_probe_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_lists_gemini_operations_without_hitting_fallback_probe() {
|
||||
let fallback_probe_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -3906,13 +4227,16 @@ async fn gateway_cancels_gemini_operation_without_hitting_fallback_probe() {
|
||||
|
||||
let (fallback_probe_url, fallback_probe_handle) = start_server(fallback_probe).await;
|
||||
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||
let provider_catalog_repository = gemini_video_catalog_repository();
|
||||
let gateway = build_router_with_state(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_and_video_task_repository_for_tests(
|
||||
auth_repository,
|
||||
crate::data::GatewayDataState::with_video_task_repository_and_provider_transport_for_tests(
|
||||
Arc::clone(&repository),
|
||||
),
|
||||
provider_catalog_repository,
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
)
|
||||
.with_auth_api_key_reader(auth_repository),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
@@ -3,6 +3,7 @@ use super::{
|
||||
sample_models_candidate_row, sample_provider, unrestricted_models_snapshot,
|
||||
InMemoryAuthApiKeySnapshotRepository, InMemoryMinimalCandidateSelectionReadRepository,
|
||||
InMemoryProviderCatalogReadRepository, InMemoryRequestCandidateRepository,
|
||||
InMemoryVideoTaskRepository, UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository,
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
};
|
||||
use crate::tests::{
|
||||
@@ -17,6 +18,330 @@ use aether_data::repository::usage::InMemoryUsageReadRepository;
|
||||
use aether_data::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot};
|
||||
use base64::Engine as _;
|
||||
|
||||
const INTERNAL_REPORT_CAPABILITY_FIELD: &str = "_aether_internal_report_capability";
|
||||
const INTERNAL_REPORT_CLIENT_KEY: &str = "sk-internal-report-capability";
|
||||
const INTERNAL_REPORT_USER_ID: &str = "user-internal-report-capability";
|
||||
const INTERNAL_REPORT_API_KEY_ID: &str = "api-key-internal-report-capability";
|
||||
|
||||
fn internal_report_planner_state() -> AppState {
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key(INTERNAL_REPORT_CLIENT_KEY)),
|
||||
unrestricted_models_snapshot(INTERNAL_REPORT_API_KEY_ID, INTERNAL_REPORT_USER_ID),
|
||||
)]));
|
||||
let candidate_repository =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
sample_models_candidate_row(
|
||||
"provider-internal-report",
|
||||
"openai",
|
||||
"openai:chat",
|
||||
"gpt-5",
|
||||
10,
|
||||
),
|
||||
]));
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-internal-report", "openai", 10)],
|
||||
vec![sample_endpoint(
|
||||
"endpoint-provider-internal-report",
|
||||
"provider-internal-report",
|
||||
"openai:chat",
|
||||
"https://api.openai.example",
|
||||
)],
|
||||
vec![sample_key(
|
||||
"key-provider-internal-report",
|
||||
"provider-internal-report",
|
||||
"openai:chat",
|
||||
"sk-upstream-openai",
|
||||
)],
|
||||
));
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
|
||||
AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
auth_repository,
|
||||
candidate_repository,
|
||||
provider_catalog_repository,
|
||||
request_candidate_repository,
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
fn internal_video_create_planner_state(api_format: &str, model: &str) -> AppState {
|
||||
let family = api_format
|
||||
.split_once(':')
|
||||
.map(|(family, _)| family)
|
||||
.expect("video api format should contain a family");
|
||||
let provider_id = format!("provider-internal-{family}-video");
|
||||
let mut candidate = sample_models_candidate_row(&provider_id, family, api_format, model, 10);
|
||||
candidate.endpoint_api_family = Some(family.to_string());
|
||||
candidate.endpoint_kind = Some("video".to_string());
|
||||
candidate.global_model_supports_streaming = Some(false);
|
||||
candidate.model_supports_streaming = Some(false);
|
||||
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key(INTERNAL_REPORT_CLIENT_KEY)),
|
||||
unrestricted_models_snapshot(INTERNAL_REPORT_API_KEY_ID, INTERNAL_REPORT_USER_ID),
|
||||
)]));
|
||||
let candidate_repository =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
candidate,
|
||||
]));
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider(&provider_id, family, 10)],
|
||||
vec![sample_endpoint(
|
||||
&format!("endpoint-{provider_id}"),
|
||||
&provider_id,
|
||||
api_format,
|
||||
if family == "gemini" {
|
||||
"https://generativelanguage.googleapis.com"
|
||||
} else {
|
||||
"https://api.openai.example"
|
||||
},
|
||||
)],
|
||||
vec![sample_key(
|
||||
&format!("key-{provider_id}"),
|
||||
&provider_id,
|
||||
api_format,
|
||||
"sk-upstream-video",
|
||||
)],
|
||||
));
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
|
||||
AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
auth_repository,
|
||||
candidate_repository,
|
||||
provider_catalog_repository,
|
||||
request_candidate_repository,
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
async fn internal_video_followup_planner_state(
|
||||
api_format: &str,
|
||||
task_id: &str,
|
||||
short_id: Option<&str>,
|
||||
external_task_id: &str,
|
||||
model: &str,
|
||||
) -> AppState {
|
||||
let family = api_format
|
||||
.split_once(':')
|
||||
.map(|(family, _)| family)
|
||||
.expect("video api format should contain a family");
|
||||
let provider_id = format!("provider-internal-{family}-video-followup");
|
||||
let endpoint_id = format!("endpoint-{provider_id}");
|
||||
let key_id = format!("key-{provider_id}");
|
||||
let repository = Arc::new(InMemoryVideoTaskRepository::default());
|
||||
repository
|
||||
.upsert(UpsertVideoTask {
|
||||
id: task_id.to_string(),
|
||||
short_id: short_id.map(ToOwned::to_owned),
|
||||
request_id: format!("request-{task_id}"),
|
||||
user_id: Some(INTERNAL_REPORT_USER_ID.to_string()),
|
||||
api_key_id: Some(INTERNAL_REPORT_API_KEY_ID.to_string()),
|
||||
username: Some("alice".to_string()),
|
||||
api_key_name: Some("default".to_string()),
|
||||
external_task_id: Some(external_task_id.to_string()),
|
||||
provider_id: Some(provider_id.clone()),
|
||||
endpoint_id: Some(endpoint_id.clone()),
|
||||
key_id: Some(key_id.clone()),
|
||||
client_api_format: Some(api_format.to_string()),
|
||||
provider_api_format: Some(api_format.to_string()),
|
||||
format_converted: false,
|
||||
model: Some(model.to_string()),
|
||||
prompt: Some("internal capability video".to_string()),
|
||||
original_request_body: Some(json!({
|
||||
"model": model,
|
||||
"prompt": "internal capability video",
|
||||
})),
|
||||
duration_seconds: Some(4),
|
||||
resolution: Some("720p".to_string()),
|
||||
aspect_ratio: Some("16:9".to_string()),
|
||||
size: Some("1280x720".to_string()),
|
||||
status: if family == "openai" {
|
||||
VideoTaskStatus::Completed
|
||||
} else {
|
||||
VideoTaskStatus::Submitted
|
||||
},
|
||||
progress_percent: if family == "openai" { 100 } else { 0 },
|
||||
progress_message: None,
|
||||
retry_count: 0,
|
||||
poll_interval_seconds: 10,
|
||||
next_poll_at_unix_secs: (family != "openai").then_some(1_700_000_010),
|
||||
poll_count: 0,
|
||||
max_poll_count: 360,
|
||||
created_at_unix_ms: 1_700_000_000_000,
|
||||
submitted_at_unix_secs: Some(1_700_000_000),
|
||||
completed_at_unix_secs: (family == "openai").then_some(1_700_000_100),
|
||||
updated_at_unix_secs: 1_700_000_000,
|
||||
error_code: None,
|
||||
error_message: None,
|
||||
video_url: None,
|
||||
request_metadata: None,
|
||||
})
|
||||
.await
|
||||
.expect("video task should seed");
|
||||
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key(INTERNAL_REPORT_CLIENT_KEY)),
|
||||
unrestricted_models_snapshot(INTERNAL_REPORT_API_KEY_ID, INTERNAL_REPORT_USER_ID),
|
||||
)]));
|
||||
let provider_catalog_repository = crate::tests::video::video_provider_catalog_repository(
|
||||
&provider_id,
|
||||
family,
|
||||
&endpoint_id,
|
||||
api_format,
|
||||
if family == "gemini" {
|
||||
"https://generativelanguage.googleapis.com"
|
||||
} else {
|
||||
"https://api.openai.example/v1"
|
||||
},
|
||||
&key_id,
|
||||
"sk-upstream-video",
|
||||
);
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let data_state = crate::data::GatewayDataState::with_video_task_provider_transport_and_request_candidate_repository_for_tests(
|
||||
repository,
|
||||
provider_catalog_repository,
|
||||
request_candidate_repository,
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
)
|
||||
.with_auth_api_key_reader(auth_repository);
|
||||
|
||||
AppState::new()
|
||||
.expect("state should build")
|
||||
.with_video_task_truth_source_mode(crate::VideoTaskTruthSourceMode::RustAuthoritative)
|
||||
.with_data_state_for_tests(data_state)
|
||||
}
|
||||
|
||||
async fn issue_internal_gateway_report_capability(
|
||||
client: &reqwest::Client,
|
||||
gateway_url: &str,
|
||||
endpoint: &str,
|
||||
trace_id: &str,
|
||||
method: &str,
|
||||
path: &str,
|
||||
request_headers: serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
) -> (String, serde_json::Value) {
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/{endpoint}"))
|
||||
.json(&json!({
|
||||
"trace_id": trace_id,
|
||||
"method": method,
|
||||
"path": path,
|
||||
"headers": request_headers,
|
||||
"body_json": body_json,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("planner request should succeed");
|
||||
let status = response.status();
|
||||
let payload: serde_json::Value = response.json().await.expect("planner body should parse");
|
||||
assert_eq!(
|
||||
status,
|
||||
StatusCode::OK,
|
||||
"planner should issue a report capability: {payload}"
|
||||
);
|
||||
let report_kind = payload["report_kind"]
|
||||
.as_str()
|
||||
.expect("planner should return a report kind")
|
||||
.to_string();
|
||||
let report_context = payload["report_context"].clone();
|
||||
assert!(
|
||||
report_context[INTERNAL_REPORT_CAPABILITY_FIELD]
|
||||
.as_str()
|
||||
.is_some_and(|value| !value.is_empty()),
|
||||
"planner should return an opaque report capability: {payload}"
|
||||
);
|
||||
(report_kind, report_context)
|
||||
}
|
||||
|
||||
async fn issue_openai_chat_report_capability(
|
||||
client: &reqwest::Client,
|
||||
gateway_url: &str,
|
||||
endpoint: &str,
|
||||
trace_id: &str,
|
||||
stream: bool,
|
||||
) -> (String, serde_json::Value) {
|
||||
issue_internal_gateway_report_capability(
|
||||
client,
|
||||
gateway_url,
|
||||
endpoint,
|
||||
trace_id,
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
json!({
|
||||
"content-type": "application/json",
|
||||
"x-api-key": INTERNAL_REPORT_CLIENT_KEY,
|
||||
}),
|
||||
json!({
|
||||
"model": "gpt-5",
|
||||
"messages": [],
|
||||
"stream": stream,
|
||||
}),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn post_internal_sync_report(
|
||||
client: &reqwest::Client,
|
||||
gateway_url: &str,
|
||||
trace_id: &str,
|
||||
report_kind: &str,
|
||||
report_context: serde_json::Value,
|
||||
) -> reqwest::Response {
|
||||
client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/report-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": trace_id,
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
},
|
||||
"body_json": {
|
||||
"id": "chatcmpl-internal-capability",
|
||||
"usage": {
|
||||
"input_tokens": 1,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
}
|
||||
}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("report request should succeed")
|
||||
}
|
||||
|
||||
async fn assert_internal_report_capability_rejected(response: reqwest::Response) {
|
||||
assert_eq!(response.status(), StatusCode::CONFLICT);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
payload,
|
||||
json!({
|
||||
"detail": "internal gateway report context does not carry a valid planner capability",
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
async fn assert_supplied_auth_context_rejected(response: reqwest::Response) {
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
payload,
|
||||
json!({
|
||||
"detail": "supplied auth_context is not accepted; authenticate through request headers",
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_internal_gateway_resolve_without_proxying_upstream() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -281,21 +606,11 @@ async fn gateway_handles_internal_gateway_execute_sync_locally_impl() {
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(EXECUTION_PATH_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some(EXECUTION_PATH_EXECUTION_RUNTIME_SYNC)
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["id"], "chatcmpl-local-execute-sync");
|
||||
assert_eq!(payload["object"], "chat.completion");
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
assert_eq!(
|
||||
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||
1
|
||||
0
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -460,22 +775,11 @@ async fn gateway_handles_internal_gateway_execute_stream_locally() {
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(EXECUTION_PATH_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some(EXECUTION_PATH_EXECUTION_RUNTIME_STREAM)
|
||||
);
|
||||
assert_eq!(
|
||||
strip_sse_keepalive_comments(&response.text().await.expect("body should read")),
|
||||
"data: one\n\ndata: [DONE]\n\n"
|
||||
);
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
assert_eq!(
|
||||
*execution_runtime_hits.lock().expect("mutex should lock"),
|
||||
1
|
||||
0
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -613,18 +917,25 @@ async fn gateway_handles_internal_gateway_report_sync_locally() {
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-report-sync",
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
assert_eq!(report_kind, "openai_chat_sync_success");
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/report-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-report-sync",
|
||||
"report_kind": "openai_chat_sync_success",
|
||||
"report_context": {
|
||||
"user_id": "user-report-sync",
|
||||
"api_key_id": "api-key-report-sync",
|
||||
},
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
@@ -667,18 +978,25 @@ async fn gateway_handles_internal_gateway_report_stream_locally() {
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-stream",
|
||||
"trace-internal-report-stream",
|
||||
true,
|
||||
)
|
||||
.await;
|
||||
assert_eq!(report_kind, "openai_chat_stream_success");
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/report-stream"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-report-stream",
|
||||
"report_kind": "openai_chat_stream_success",
|
||||
"report_context": {
|
||||
"user_id": "user-report-stream",
|
||||
"api_key_id": "api-key-report-stream",
|
||||
},
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "text/event-stream",
|
||||
@@ -698,6 +1016,181 @@ async fn gateway_handles_internal_gateway_report_stream_locally() {
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_internal_gateway_report_with_tampered_protected_context() {
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-report-tampered-context",
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
|
||||
for (field, forged_value) in [
|
||||
("user_id", json!("user-unrelated-victim")),
|
||||
("api_key_id", json!("api-key-unrelated-victim")),
|
||||
("provider_id", json!("provider-unrelated-victim")),
|
||||
("endpoint_id", json!("endpoint-unrelated-victim")),
|
||||
("key_id", json!("key-unrelated-victim")),
|
||||
("client_api_format", json!("gemini:video")),
|
||||
("task_id", json!("task-unrelated-victim")),
|
||||
("local_task_id", json!("local-task-unrelated-victim")),
|
||||
("local_short_id", json!("short-unrelated-victim")),
|
||||
("file_name", json!("files/unrelated-victim")),
|
||||
("file_key_id", json!("file-key-unrelated-victim")),
|
||||
] {
|
||||
let mut tampered_context = report_context.clone();
|
||||
tampered_context
|
||||
.as_object_mut()
|
||||
.expect("planner report context should be an object")
|
||||
.insert(field.to_string(), forged_value);
|
||||
let response = post_internal_sync_report(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"trace-internal-report-tampered-context",
|
||||
&report_kind,
|
||||
tampered_context,
|
||||
)
|
||||
.await;
|
||||
assert_internal_report_capability_rejected(response).await;
|
||||
}
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_internal_gateway_report_without_a_known_capability() {
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-report-missing-capability",
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
|
||||
let mut missing_capability = report_context.clone();
|
||||
missing_capability
|
||||
.as_object_mut()
|
||||
.expect("planner report context should be an object")
|
||||
.remove(INTERNAL_REPORT_CAPABILITY_FIELD);
|
||||
let response = post_internal_sync_report(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"trace-internal-report-missing-capability",
|
||||
&report_kind,
|
||||
missing_capability,
|
||||
)
|
||||
.await;
|
||||
assert_internal_report_capability_rejected(response).await;
|
||||
|
||||
let mut unknown_capability = report_context;
|
||||
unknown_capability[INTERNAL_REPORT_CAPABILITY_FIELD] =
|
||||
json!("00000000-0000-4000-8000-000000000000");
|
||||
let response = post_internal_sync_report(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"trace-internal-report-missing-capability",
|
||||
&report_kind,
|
||||
unknown_capability,
|
||||
)
|
||||
.await;
|
||||
assert_internal_report_capability_rejected(response).await;
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_internal_gateway_report_for_wrong_trace_or_scope() {
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-report-boundary",
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
|
||||
let response = post_internal_sync_report(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"trace-internal-report-wrong-trace",
|
||||
&report_kind,
|
||||
report_context.clone(),
|
||||
)
|
||||
.await;
|
||||
assert_internal_report_capability_rejected(response).await;
|
||||
|
||||
let response = post_internal_sync_report(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"trace-internal-report-boundary",
|
||||
"openai_image_sync_success",
|
||||
report_context,
|
||||
)
|
||||
.await;
|
||||
assert_internal_report_capability_rejected(response).await;
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_allows_internal_gateway_report_observation_fields() {
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, mut report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-report-observations",
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
let context = report_context
|
||||
.as_object_mut()
|
||||
.expect("planner report context should be an object");
|
||||
context.insert(
|
||||
"provider_response_headers".to_string(),
|
||||
json!({"x-request-id": "upstream-request-123"}),
|
||||
);
|
||||
context.insert(
|
||||
"client_response_headers".to_string(),
|
||||
json!({"content-type": "application/json"}),
|
||||
);
|
||||
context.insert("upstream_response".to_string(), json!({"status_code": 200}));
|
||||
context.insert("error_flow".to_string(), json!({"attempted": false}));
|
||||
|
||||
let response = post_internal_sync_report(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"trace-internal-report-observations",
|
||||
&report_kind,
|
||||
report_context,
|
||||
)
|
||||
.await;
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response
|
||||
.json::<serde_json::Value>()
|
||||
.await
|
||||
.expect("json body should parse"),
|
||||
json!({"ok": true})
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_internal_gateway_finalize_sync_locally() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -714,20 +1207,24 @@ async fn gateway_handles_internal_gateway_finalize_sync_locally() {
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let gateway = build_router_with_state(internal_report_planner_state());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (_report_kind, report_context) = issue_openai_chat_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"plan-sync",
|
||||
"trace-internal-finalize-sync",
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/finalize-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-finalize-sync",
|
||||
"report_kind": "openai_chat_sync_finalize",
|
||||
"report_context": {
|
||||
"user_id": "user-finalize-sync",
|
||||
"api_key_id": "api-key-finalize-sync",
|
||||
"client_api_format": "openai:chat",
|
||||
"provider_api_format": "openai:chat",
|
||||
},
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
@@ -781,24 +1278,37 @@ async fn gateway_handles_internal_gateway_finalize_sync_openai_video_locally() {
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let gateway = build_router_with_state(internal_video_create_planner_state(
|
||||
"openai:video",
|
||||
"sora-2",
|
||||
));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_internal_gateway_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-finalize-video",
|
||||
"POST",
|
||||
"/v1/videos",
|
||||
json!({
|
||||
"authorization": format!("Bearer {INTERNAL_REPORT_CLIENT_KEY}"),
|
||||
"content-type": "application/json",
|
||||
}),
|
||||
json!({
|
||||
"model": "sora-2",
|
||||
"prompt": "make a trailer",
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(report_kind, "openai_video_create_sync_finalize");
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/finalize-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-finalize-video",
|
||||
"report_kind": "openai_video_create_sync_finalize",
|
||||
"report_context": {
|
||||
"user_id": "user-finalize-video",
|
||||
"api_key_id": "api-key-finalize-video",
|
||||
"model": "sora-2",
|
||||
"local_task_id": "local-video-task-123",
|
||||
"local_created_at": 1712345678u64,
|
||||
"original_request_body": {
|
||||
"prompt": "make a trailer"
|
||||
}
|
||||
},
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
@@ -821,18 +1331,15 @@ async fn gateway_handles_internal_gateway_finalize_sync_openai_video_locally() {
|
||||
"true"
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
payload,
|
||||
json!({
|
||||
"id": "local-video-task-123",
|
||||
"object": "video",
|
||||
"status": "queued",
|
||||
"progress": 0,
|
||||
"created_at": 1712345678u64,
|
||||
"model": "sora-2",
|
||||
"prompt": "make a trailer",
|
||||
})
|
||||
);
|
||||
assert!(payload["id"]
|
||||
.as_str()
|
||||
.is_some_and(|value| !value.is_empty() && value != "vid-ext-123"));
|
||||
assert_eq!(payload["object"], "video");
|
||||
assert_eq!(payload["status"], "queued");
|
||||
assert_eq!(payload["progress"], 0);
|
||||
assert!(payload["created_at"].as_u64().is_some());
|
||||
assert_eq!(payload["model"], "sora-2");
|
||||
assert_eq!(payload["prompt"], "make a trailer");
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -855,20 +1362,34 @@ async fn gateway_handles_internal_gateway_finalize_sync_gemini_video_locally() {
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let gateway =
|
||||
build_router_with_state(internal_video_create_planner_state("gemini:video", "veo-3"));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_internal_gateway_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"plan-sync",
|
||||
"trace-internal-finalize-gemini-video",
|
||||
"POST",
|
||||
"/v1beta/models/veo-3:predictLongRunning",
|
||||
json!({
|
||||
"content-type": "application/json",
|
||||
"x-goog-api-key": INTERNAL_REPORT_CLIENT_KEY,
|
||||
}),
|
||||
json!({
|
||||
"prompt": "make a gemini trailer",
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(report_kind, "gemini_video_create_sync_finalize");
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/finalize-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-finalize-gemini-video",
|
||||
"report_kind": "gemini_video_create_sync_finalize",
|
||||
"report_context": {
|
||||
"user_id": "user-finalize-gemini-video",
|
||||
"api_key_id": "api-key-finalize-gemini-video",
|
||||
"model": "veo-3",
|
||||
"local_short_id": "gemini-short-123"
|
||||
},
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
@@ -892,14 +1413,11 @@ async fn gateway_handles_internal_gateway_finalize_sync_gemini_video_locally() {
|
||||
"true"
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
payload,
|
||||
json!({
|
||||
"name": "models/veo-3/operations/gemini-short-123",
|
||||
"done": false,
|
||||
"metadata": {},
|
||||
})
|
||||
);
|
||||
assert!(payload["name"]
|
||||
.as_str()
|
||||
.is_some_and(|value| value.starts_with("models/veo-3/operations/")));
|
||||
assert_eq!(payload["done"], false);
|
||||
assert_eq!(payload["metadata"], json!({}));
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -922,17 +1440,39 @@ async fn gateway_handles_internal_gateway_finalize_sync_openai_video_delete_loca
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let state = internal_video_followup_planner_state(
|
||||
"openai:video",
|
||||
"video-delete-123",
|
||||
None,
|
||||
"ext-video-delete-123",
|
||||
"sora-2",
|
||||
)
|
||||
.await;
|
||||
let gateway = build_router_with_state(state);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_internal_gateway_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"decision-sync",
|
||||
"trace-internal-finalize-video-delete",
|
||||
"DELETE",
|
||||
"/v1/videos/video-delete-123",
|
||||
json!({
|
||||
"authorization": format!("Bearer {INTERNAL_REPORT_CLIENT_KEY}"),
|
||||
"content-type": "application/json",
|
||||
}),
|
||||
json!({}),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(report_kind, "openai_video_delete_sync_finalize");
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/finalize-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-finalize-video-delete",
|
||||
"report_kind": "openai_video_delete_sync_finalize",
|
||||
"report_context": {
|
||||
"task_id": "video-delete-123"
|
||||
},
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
@@ -982,18 +1522,39 @@ async fn gateway_handles_internal_gateway_finalize_sync_gemini_video_cancel_loca
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router().expect("gateway should build");
|
||||
let state = internal_video_followup_planner_state(
|
||||
"gemini:video",
|
||||
"gemini-cancel-task-record",
|
||||
Some("gemini-cancel-123"),
|
||||
"operations/ext-gemini-cancel-123",
|
||||
"veo-3",
|
||||
)
|
||||
.await;
|
||||
let gateway = build_router_with_state(state);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let (report_kind, report_context) = issue_internal_gateway_report_capability(
|
||||
&client,
|
||||
&gateway_url,
|
||||
"plan-sync",
|
||||
"trace-internal-finalize-gemini-video-cancel",
|
||||
"POST",
|
||||
"/v1beta/models/veo-3/operations/gemini-cancel-123:cancel",
|
||||
json!({
|
||||
"content-type": "application/json",
|
||||
"x-goog-api-key": INTERNAL_REPORT_CLIENT_KEY,
|
||||
}),
|
||||
json!({}),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(report_kind, "gemini_video_cancel_sync_finalize");
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/api/internal/gateway/finalize-sync"))
|
||||
.json(&json!({
|
||||
"trace_id": "trace-internal-finalize-gemini-video-cancel",
|
||||
"report_kind": "gemini_video_cancel_sync_finalize",
|
||||
"report_context": {
|
||||
"task_id": "gemini-cancel-123",
|
||||
"model": "veo-3"
|
||||
},
|
||||
"report_kind": report_kind,
|
||||
"report_context": report_context,
|
||||
"status_code": 200,
|
||||
"headers": {
|
||||
"content-type": "application/json",
|
||||
@@ -1148,17 +1709,7 @@ async fn gateway_handles_internal_gateway_decision_sync_locally_with_supplied_au
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["action"], "execution_runtime_sync_decision");
|
||||
assert_eq!(payload["decision_kind"], "openai_chat_sync");
|
||||
assert_eq!(payload["provider_id"], "provider-1");
|
||||
assert_eq!(payload["endpoint_id"], "endpoint-provider-1");
|
||||
assert_eq!(payload["key_id"], "key-provider-1");
|
||||
assert_eq!(payload["provider_api_format"], "openai:chat");
|
||||
assert_eq!(payload["client_api_format"], "openai:chat");
|
||||
assert_eq!(payload["model_name"], "gpt-5");
|
||||
assert_eq!(payload["auth_context"], serde_json::Value::Null);
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -1267,10 +1818,7 @@ async fn gateway_internal_decision_sync_revalidates_supplied_auth_context_wallet
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["action"], "fallback_plan");
|
||||
assert_eq!(payload["auth_context"], serde_json::Value::Null);
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -1301,7 +1849,10 @@ async fn gateway_returns_internal_gateway_decision_sync_fallback_with_resolved_a
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_auth_api_key_data_reader_for_tests(auth_repository),
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_auth_api_key_reader_for_tests(auth_repository)
|
||||
.with_system_default_routing_group_for_tests(),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
@@ -1418,14 +1969,7 @@ async fn gateway_handles_internal_gateway_decision_stream_locally_with_supplied_
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["action"], "execution_runtime_stream_decision");
|
||||
assert_eq!(payload["decision_kind"], "openai_chat_stream");
|
||||
assert_eq!(payload["provider_id"], "provider-stream-1");
|
||||
assert_eq!(payload["endpoint_id"], "endpoint-provider-stream-1");
|
||||
assert_eq!(payload["key_id"], "key-provider-stream-1");
|
||||
assert_eq!(payload["auth_context"], serde_json::Value::Null);
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -1517,17 +2061,7 @@ async fn gateway_handles_internal_gateway_plan_sync_locally_with_supplied_auth_c
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["action"], "execution_runtime_sync");
|
||||
assert_eq!(payload["plan_kind"], "openai_chat_sync");
|
||||
assert_eq!(payload["plan"]["provider_id"], "provider-plan-sync-1");
|
||||
assert_eq!(
|
||||
payload["plan"]["endpoint_id"],
|
||||
"endpoint-provider-plan-sync-1"
|
||||
);
|
||||
assert_eq!(payload["plan"]["key_id"], "key-provider-plan-sync-1");
|
||||
assert_eq!(payload["auth_context"], serde_json::Value::Null);
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
@@ -1620,18 +2154,7 @@ async fn gateway_handles_internal_gateway_plan_stream_locally_with_supplied_auth
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["action"], "execution_runtime_stream");
|
||||
assert_eq!(payload["plan_kind"], "openai_chat_stream");
|
||||
assert_eq!(payload["plan"]["provider_id"], "provider-plan-stream-1");
|
||||
assert_eq!(
|
||||
payload["plan"]["endpoint_id"],
|
||||
"endpoint-provider-plan-stream-1"
|
||||
);
|
||||
assert_eq!(payload["plan"]["key_id"], "key-provider-plan-stream-1");
|
||||
assert_eq!(payload["plan"]["stream"], true);
|
||||
assert_eq!(payload["auth_context"], serde_json::Value::Null);
|
||||
assert_supplied_auth_context_rejected(response).await;
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
|
||||
@@ -5,6 +5,13 @@ use crate::tests::{
|
||||
use aether_data::repository::oauth_providers::{
|
||||
InMemoryOAuthProviderRepository, StoredOAuthProviderConfig,
|
||||
};
|
||||
use aether_data::repository::users::{
|
||||
InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserSessionRecord,
|
||||
};
|
||||
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
|
||||
use base64::Engine as _;
|
||||
use hmac::Mac as _;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
fn sample_identity_oauth_provider(provider_type: &str) -> StoredOAuthProviderConfig {
|
||||
StoredOAuthProviderConfig::new(
|
||||
@@ -28,6 +35,63 @@ fn sample_identity_oauth_provider(provider_type: &str) -> StoredOAuthProviderCon
|
||||
)
|
||||
}
|
||||
|
||||
fn oauth_login_cookie_name(state_nonce: &str) -> String {
|
||||
format!(
|
||||
"__Host-aether_oauth_login_{:x}",
|
||||
Sha256::digest(state_nonce.as_bytes())
|
||||
)
|
||||
}
|
||||
|
||||
fn oauth_state_from_authorize_location(location: &str) -> String {
|
||||
url::Url::parse(location)
|
||||
.expect("authorize location should be a URL")
|
||||
.query_pairs()
|
||||
.find_map(|(key, value)| (key == "state").then(|| value.into_owned()))
|
||||
.expect("authorize location should include state")
|
||||
}
|
||||
|
||||
fn cookie_pair_from_set_cookie(set_cookie: &str) -> String {
|
||||
set_cookie
|
||||
.split(';')
|
||||
.next()
|
||||
.expect("Set-Cookie should include a cookie pair")
|
||||
.to_string()
|
||||
}
|
||||
|
||||
fn build_oauth_test_access_token(
|
||||
user: &StoredUserAuthRecord,
|
||||
session_id: &str,
|
||||
expires_at: chrono::DateTime<chrono::Utc>,
|
||||
) -> String {
|
||||
let header = serde_json::json!({ "alg": "HS256", "typ": "JWT" });
|
||||
let payload = serde_json::json!({
|
||||
"exp": expires_at.timestamp(),
|
||||
"type": "access",
|
||||
"user_id": user.id,
|
||||
"role": user.role,
|
||||
"created_at": user.created_at.map(|value| value.to_rfc3339()),
|
||||
"session_id": session_id,
|
||||
});
|
||||
let header_segment = base64::engine::general_purpose::URL_SAFE_NO_PAD
|
||||
.encode(serde_json::to_vec(&header).expect("JWT header should serialize"));
|
||||
let payload_segment = base64::engine::general_purpose::URL_SAFE_NO_PAD
|
||||
.encode(serde_json::to_vec(&payload).expect("JWT payload should serialize"));
|
||||
let signing_input = format!("{header_segment}.{payload_segment}");
|
||||
let secret = std::env::var("JWT_SECRET_KEY")
|
||||
.ok()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string());
|
||||
let mut mac = hmac::Hmac::<Sha256>::new_from_slice(secret.as_bytes())
|
||||
.expect("test JWT secret should be valid");
|
||||
mac.update(signing_input.as_bytes());
|
||||
let signature = mac.finalize().into_bytes();
|
||||
format!(
|
||||
"{signing_input}.{}",
|
||||
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(signature)
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_serves_oauth_public_providers_locally_without_hitting_upstream() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -243,6 +307,294 @@ async fn gateway_serves_configured_oauth_provider_when_oauth_module_enabled() {
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_oauth_callback_without_browser_binding_without_consuming_state() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||
let upstream = Router::new().route(
|
||||
"/{*path}",
|
||||
any(move |_request: Request| {
|
||||
let upstream_hits_inner = Arc::clone(&upstream_hits_clone);
|
||||
async move {
|
||||
*upstream_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Body::from("unexpected upstream hit"))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![
|
||||
sample_identity_oauth_provider("linuxdo"),
|
||||
]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(repository)
|
||||
.with_system_config_values_for_tests(vec![(
|
||||
"module.oauth.enabled".to_string(),
|
||||
serde_json::json!(true),
|
||||
)]);
|
||||
let runtime_state = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
.with_runtime_state(runtime_state);
|
||||
|
||||
let binding_hash = format!("{:x}", Sha256::digest(b"correct-cookie"));
|
||||
let missing_cookie_state = crate::oauth::StoredIdentityOAuthState::login(
|
||||
"linuxdo",
|
||||
"device-csrf-test",
|
||||
Some("pkce-verifier".to_string()),
|
||||
Some(binding_hash.clone()),
|
||||
);
|
||||
let wrong_cookie_state = crate::oauth::StoredIdentityOAuthState::login(
|
||||
"linuxdo",
|
||||
"device-csrf-test",
|
||||
Some("pkce-verifier".to_string()),
|
||||
Some(binding_hash),
|
||||
);
|
||||
crate::oauth::save_identity_oauth_state(&state, &missing_cookie_state)
|
||||
.await
|
||||
.expect("missing-cookie state should save");
|
||||
crate::oauth::save_identity_oauth_state(&state, &wrong_cookie_state)
|
||||
.await
|
||||
.expect("wrong-cookie state should save");
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router_with_state(state.clone());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("client should build");
|
||||
|
||||
for (state_nonce, cookie_header) in [
|
||||
(missing_cookie_state.nonce.as_str(), None),
|
||||
(
|
||||
wrong_cookie_state.nonce.as_str(),
|
||||
Some(format!(
|
||||
"{}=wrong-cookie",
|
||||
oauth_login_cookie_name(&wrong_cookie_state.nonce)
|
||||
)),
|
||||
),
|
||||
] {
|
||||
let mut request = client.get(format!(
|
||||
"{gateway_url}/api/oauth/linuxdo/callback?code=provider-code&state={state_nonce}"
|
||||
));
|
||||
if let Some(cookie_header) = cookie_header.as_deref() {
|
||||
request = request.header(http::header::COOKIE, cookie_header);
|
||||
}
|
||||
let response = request
|
||||
.send()
|
||||
.await
|
||||
.expect("callback request should succeed");
|
||||
assert_eq!(response.status(), StatusCode::FOUND);
|
||||
let location = response
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("invalid callback should redirect");
|
||||
assert!(
|
||||
location.contains("invalid_state"),
|
||||
"unexpected location: {location}"
|
||||
);
|
||||
let expected_cookie_name = oauth_login_cookie_name(state_nonce);
|
||||
let clear_cookies = response
|
||||
.headers()
|
||||
.get_all(http::header::SET_COOKIE)
|
||||
.iter()
|
||||
.filter_map(|value| value.to_str().ok())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(clear_cookies.len(), 1);
|
||||
assert!(clear_cookies[0].starts_with(&format!("{expected_cookie_name}=;")));
|
||||
assert!(clear_cookies[0].contains("Max-Age=0"));
|
||||
}
|
||||
|
||||
for state_nonce in [&missing_cookie_state.nonce, &wrong_cookie_state.nonce] {
|
||||
assert!(state
|
||||
.runtime_kv_get(&crate::oauth::identity_oauth_state_storage_key(state_nonce))
|
||||
.await
|
||||
.expect("OAuth state lookup should succeed")
|
||||
.is_some());
|
||||
|
||||
let cookie_header = format!("{}=correct-cookie", oauth_login_cookie_name(state_nonce));
|
||||
let response = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={state_nonce}"
|
||||
))
|
||||
.header(http::header::COOKIE, &cookie_header)
|
||||
.send()
|
||||
.await
|
||||
.expect("bound callback request should succeed");
|
||||
assert_eq!(response.status(), StatusCode::FOUND);
|
||||
let location = response
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("denied callback should redirect");
|
||||
assert!(location.contains("authorization_denied"));
|
||||
assert!(state
|
||||
.runtime_kv_get(&crate::oauth::identity_oauth_state_storage_key(state_nonce))
|
||||
.await
|
||||
.expect("OAuth state lookup should succeed")
|
||||
.is_none());
|
||||
|
||||
let replay = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={state_nonce}"
|
||||
))
|
||||
.header(http::header::COOKIE, cookie_header)
|
||||
.send()
|
||||
.await
|
||||
.expect("replayed callback request should succeed");
|
||||
let replay_location = replay
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("replayed callback should redirect");
|
||||
assert!(replay_location.contains("invalid_state"));
|
||||
}
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_keeps_parallel_oauth_login_cookies_independent_and_clears_only_consumed_state() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||
let upstream = Router::new().route(
|
||||
"/{*path}",
|
||||
any(move |_request: Request| {
|
||||
let upstream_hits_inner = Arc::clone(&upstream_hits_clone);
|
||||
async move {
|
||||
*upstream_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Body::from("unexpected upstream hit"))
|
||||
}
|
||||
}),
|
||||
);
|
||||
let repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![
|
||||
sample_identity_oauth_provider("linuxdo"),
|
||||
]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(repository)
|
||||
.with_system_config_values_for_tests(vec![(
|
||||
"module.oauth.enabled".to_string(),
|
||||
serde_json::json!(true),
|
||||
)]);
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
.with_runtime_state(Arc::new(RuntimeState::memory(
|
||||
MemoryRuntimeStateConfig::default(),
|
||||
)));
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router_with_state(state);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("client should build");
|
||||
let first_authorize = client
|
||||
.get(format!("{gateway_url}/api/oauth/linuxdo/authorize"))
|
||||
.header("x-client-device-id", "parallel-device");
|
||||
let second_authorize = client
|
||||
.get(format!("{gateway_url}/api/oauth/linuxdo/authorize"))
|
||||
.header("x-client-device-id", "parallel-device");
|
||||
let (first_response, second_response) =
|
||||
tokio::join!(first_authorize.send(), second_authorize.send());
|
||||
let first_response = first_response.expect("first authorize request should succeed");
|
||||
let second_response = second_response.expect("second authorize request should succeed");
|
||||
assert_eq!(first_response.status(), StatusCode::FOUND);
|
||||
assert_eq!(second_response.status(), StatusCode::FOUND);
|
||||
|
||||
let first_location = first_response
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("first authorize location should exist");
|
||||
let second_location = second_response
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("second authorize location should exist");
|
||||
let first_state = oauth_state_from_authorize_location(first_location);
|
||||
let second_state = oauth_state_from_authorize_location(second_location);
|
||||
assert_ne!(first_state, second_state);
|
||||
|
||||
let first_cookie_name = oauth_login_cookie_name(&first_state);
|
||||
let second_cookie_name = oauth_login_cookie_name(&second_state);
|
||||
assert_ne!(first_cookie_name, second_cookie_name);
|
||||
let first_cookie = first_response
|
||||
.headers()
|
||||
.get(http::header::SET_COOKIE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(cookie_pair_from_set_cookie)
|
||||
.expect("first authorize response should set a login cookie");
|
||||
let second_cookie = second_response
|
||||
.headers()
|
||||
.get(http::header::SET_COOKIE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(cookie_pair_from_set_cookie)
|
||||
.expect("second authorize response should set a login cookie");
|
||||
assert!(first_cookie.starts_with(&format!("{first_cookie_name}=")));
|
||||
assert!(second_cookie.starts_with(&format!("{second_cookie_name}=")));
|
||||
|
||||
let combined_cookie_header = format!("{first_cookie}; {second_cookie}");
|
||||
let first_callback = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={first_state}"
|
||||
))
|
||||
.header(http::header::COOKIE, combined_cookie_header)
|
||||
.send()
|
||||
.await
|
||||
.expect("first callback request should succeed");
|
||||
assert_eq!(first_callback.status(), StatusCode::FOUND);
|
||||
let first_callback_location = first_callback
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("first callback location should exist");
|
||||
assert!(first_callback_location.contains("authorization_denied"));
|
||||
let first_clears = first_callback
|
||||
.headers()
|
||||
.get_all(http::header::SET_COOKIE)
|
||||
.iter()
|
||||
.filter_map(|value| value.to_str().ok())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(first_clears.len(), 1);
|
||||
assert!(first_clears[0].starts_with(&format!("{first_cookie_name}=;")));
|
||||
assert!(!first_clears[0].starts_with(&format!("{second_cookie_name}=;")));
|
||||
|
||||
let second_callback = client
|
||||
.get(format!(
|
||||
"{gateway_url}/api/oauth/linuxdo/callback?error=access_denied&state={second_state}"
|
||||
))
|
||||
.header(http::header::COOKIE, second_cookie)
|
||||
.send()
|
||||
.await
|
||||
.expect("second callback request should succeed");
|
||||
assert_eq!(second_callback.status(), StatusCode::FOUND);
|
||||
let second_callback_location = second_callback
|
||||
.headers()
|
||||
.get(http::header::LOCATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("second callback location should exist");
|
||||
assert!(second_callback_location.contains("authorization_denied"));
|
||||
let second_clears = second_callback
|
||||
.headers()
|
||||
.get_all(http::header::SET_COOKIE)
|
||||
.iter()
|
||||
.filter_map(|value| value.to_str().ok())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(second_clears.len(), 1);
|
||||
assert!(second_clears[0].starts_with(&format!("{second_cookie_name}=;")));
|
||||
assert!(!second_clears[0].starts_with(&format!("{first_cookie_name}=;")));
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_requires_auth_for_oauth_user_bindable_providers_without_hitting_upstream() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -277,6 +629,102 @@ async fn gateway_requires_auth_for_oauth_user_bindable_providers_without_hitting
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_oauth_account_lists_are_never_cacheable() {
|
||||
let provider_repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![
|
||||
sample_identity_oauth_provider("linuxdo"),
|
||||
]));
|
||||
let now = chrono::Utc::now();
|
||||
let user = StoredUserAuthRecord::new(
|
||||
"oauth-list-user".to_string(),
|
||||
Some("[email protected]".to_string()),
|
||||
true,
|
||||
"oauth-list-user".to_string(),
|
||||
Some("unused-password-hash".to_string()),
|
||||
"user".to_string(),
|
||||
"local".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
false,
|
||||
Some(now),
|
||||
Some(now),
|
||||
)
|
||||
.expect("test user should build");
|
||||
let session_id = "oauth-list-session";
|
||||
let device_id = "oauth-list-device";
|
||||
let session = StoredUserSessionRecord::new(
|
||||
session_id.to_string(),
|
||||
user.id.clone(),
|
||||
device_id.to_string(),
|
||||
None,
|
||||
StoredUserSessionRecord::hash_refresh_token("unused-refresh-token"),
|
||||
None,
|
||||
None,
|
||||
Some(now),
|
||||
Some(now + chrono::Duration::days(1)),
|
||||
None,
|
||||
None,
|
||||
Some("127.0.0.1".to_string()),
|
||||
Some("oauth-list-test".to_string()),
|
||||
Some(now),
|
||||
Some(now),
|
||||
)
|
||||
.expect("test session should build");
|
||||
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users([user.clone()]));
|
||||
let data_state = crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(
|
||||
provider_repository,
|
||||
)
|
||||
.with_system_config_values_for_tests(vec![(
|
||||
"module.oauth.enabled".to_string(),
|
||||
serde_json::json!(true),
|
||||
)])
|
||||
.with_user_reader(user_repository);
|
||||
let access_token =
|
||||
build_oauth_test_access_token(&user, session_id, now + chrono::Duration::hours(1));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
.with_auth_users_for_tests([user])
|
||||
.with_auth_session_for_tests(session);
|
||||
let gateway = build_router_with_state(state);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
for path in [
|
||||
"/api/user/oauth/bindable-providers",
|
||||
"/api/user/oauth/links",
|
||||
] {
|
||||
let response = client
|
||||
.get(format!("{gateway_url}{path}"))
|
||||
.header("authorization", format!("Bearer {access_token}"))
|
||||
.header("x-client-device-id", device_id)
|
||||
.send()
|
||||
.await
|
||||
.expect("OAuth account list request should succeed");
|
||||
assert_eq!(response.status(), StatusCode::OK, "{path}");
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(http::header::CACHE_CONTROL)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("no-store"),
|
||||
"{path}"
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(http::header::PRAGMA)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("no-cache"),
|
||||
"{path}"
|
||||
);
|
||||
}
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_requires_auth_for_oauth_user_bind_token_without_hitting_upstream() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -310,3 +758,123 @@ async fn gateway_requires_auth_for_oauth_user_bind_token_without_hitting_upstrea
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_oauth_bind_token_returns_authorize_url_and_stores_bound_browser_state() {
|
||||
let repository = Arc::new(InMemoryOAuthProviderRepository::seed(vec![
|
||||
sample_identity_oauth_provider("linuxdo"),
|
||||
]));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_oauth_provider_repository_for_tests(repository)
|
||||
.with_system_config_values_for_tests(vec![(
|
||||
"module.oauth.enabled".to_string(),
|
||||
serde_json::json!(true),
|
||||
)]);
|
||||
let now = chrono::Utc::now();
|
||||
let user = StoredUserAuthRecord::new(
|
||||
"oauth-bind-user".to_string(),
|
||||
Some("[email protected]".to_string()),
|
||||
true,
|
||||
"oauth-bind-user".to_string(),
|
||||
Some("unused-password-hash".to_string()),
|
||||
"user".to_string(),
|
||||
"local".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
false,
|
||||
Some(now),
|
||||
Some(now),
|
||||
)
|
||||
.expect("test user should build");
|
||||
let session_id = "oauth-bind-session";
|
||||
let device_id = "oauth-bind-device";
|
||||
let session = StoredUserSessionRecord::new(
|
||||
session_id.to_string(),
|
||||
user.id.clone(),
|
||||
device_id.to_string(),
|
||||
None,
|
||||
StoredUserSessionRecord::hash_refresh_token("unused-refresh-token"),
|
||||
None,
|
||||
None,
|
||||
Some(now),
|
||||
Some(now + chrono::Duration::days(1)),
|
||||
None,
|
||||
None,
|
||||
Some("127.0.0.1".to_string()),
|
||||
Some("oauth-bind-test".to_string()),
|
||||
Some(now),
|
||||
Some(now),
|
||||
)
|
||||
.expect("test session should build");
|
||||
let access_token =
|
||||
build_oauth_test_access_token(&user, session_id, now + chrono::Duration::hours(1));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state)
|
||||
.with_auth_users_for_tests([user.clone()])
|
||||
.with_auth_session_for_tests(session);
|
||||
let gateway = build_router_with_state(state.clone());
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/api/user/oauth/linuxdo/bind-token"))
|
||||
.header("authorization", format!("Bearer {access_token}"))
|
||||
.header("x-client-device-id", device_id)
|
||||
.header("user-agent", "AetherOAuthBindTest/1.0")
|
||||
.send()
|
||||
.await
|
||||
.expect("bind-token request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let set_cookie = response
|
||||
.headers()
|
||||
.get(http::header::SET_COOKIE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("bind-token response should set the browser binding cookie")
|
||||
.to_string();
|
||||
assert!(set_cookie.contains("Path=/"));
|
||||
assert!(set_cookie.contains("HttpOnly"));
|
||||
assert!(set_cookie.contains("SameSite=Lax"));
|
||||
let cookie_pair = cookie_pair_from_set_cookie(&set_cookie);
|
||||
let (cookie_name, browser_binding) = cookie_pair
|
||||
.split_once('=')
|
||||
.expect("browser binding cookie should contain a value");
|
||||
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert!(payload.get("bind_token").is_none());
|
||||
let authorize_url = payload["authorize_url"]
|
||||
.as_str()
|
||||
.expect("bind-token response should include authorize_url");
|
||||
let state_nonce = oauth_state_from_authorize_location(authorize_url);
|
||||
assert_eq!(cookie_name, oauth_login_cookie_name(&state_nonce));
|
||||
|
||||
let raw_state = state
|
||||
.runtime_kv_get(&crate::oauth::identity_oauth_state_storage_key(
|
||||
&state_nonce,
|
||||
))
|
||||
.await
|
||||
.expect("OAuth state lookup should succeed")
|
||||
.expect("OAuth bind state should be stored");
|
||||
assert!(crate::handlers::shared::runtime_secret_payload_is_sealed(
|
||||
&raw_state
|
||||
));
|
||||
assert!(!raw_state.contains("pkce_verifier"));
|
||||
let stored = crate::oauth::load_identity_oauth_state(&state, &state_nonce)
|
||||
.await
|
||||
.expect("OAuth state lookup should succeed")
|
||||
.expect("OAuth state should decrypt");
|
||||
assert_eq!(stored.mode, crate::oauth::IdentityOAuthStateMode::Bind);
|
||||
assert_eq!(stored.provider_type, "linuxdo");
|
||||
assert_eq!(stored.client_device_id, device_id);
|
||||
assert_eq!(stored.bind_user_id.as_deref(), Some(user.id.as_str()));
|
||||
assert_eq!(stored.bind_session_id.as_deref(), Some(session_id));
|
||||
let expected_binding_hash = format!("{:x}", Sha256::digest(browser_binding.as_bytes()));
|
||||
assert_eq!(
|
||||
stored.browser_binding_hash.as_deref(),
|
||||
Some(expected_binding_hash.as_str())
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
@@ -184,7 +184,7 @@ async fn gateway_exposes_frontdoor_manifest_without_proxying_upstream() {
|
||||
.any(|value| value == "/v1internal:streamGenerateContent"));
|
||||
assert_eq!(
|
||||
payload["rust_frontdoor"]["internal_gateway"]["status"],
|
||||
"rust_native_control_plane"
|
||||
"test_loopback_compatibility"
|
||||
);
|
||||
assert_eq!(
|
||||
payload["rust_frontdoor"]["internal_gateway"]["path_prefixes"][0],
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -329,7 +329,10 @@ async fn run_vscodex_gateway_integration() {
|
||||
.json()
|
||||
.await
|
||||
.expect("internal auth failure body should be JSON");
|
||||
assert_eq!(internal_denied_payload["detail"], "VS Codex 服务鉴权失败");
|
||||
assert_eq!(
|
||||
internal_denied_payload["detail"],
|
||||
"服务暂不可用,请稍后重试"
|
||||
);
|
||||
|
||||
let redirected = client
|
||||
.delete(format!(
|
||||
@@ -562,7 +565,7 @@ async fn run_vscodex_gateway_integration() {
|
||||
assert_eq!(disabled.status(), StatusCode::SERVICE_UNAVAILABLE);
|
||||
let disabled_payload: serde_json::Value =
|
||||
disabled.json().await.expect("disabled body should be JSON");
|
||||
assert_eq!(disabled_payload["detail"], "VS Codex 服务未启用");
|
||||
assert_eq!(disabled_payload["detail"], "服务暂不可用,请稍后重试");
|
||||
assert_eq!(
|
||||
captured_requests
|
||||
.lock()
|
||||
|
||||
Reference in New Issue
Block a user