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:
@@ -578,7 +578,7 @@ CREATE TABLE IF NOT EXISTS proxy_nodes (
|
||||
is_manual TINYINT(1) NOT NULL DEFAULT 0,
|
||||
proxy_url VARCHAR(500),
|
||||
proxy_username VARCHAR(255),
|
||||
proxy_password VARCHAR(500),
|
||||
proxy_password TEXT,
|
||||
created_at BIGINT NOT NULL,
|
||||
updated_at BIGINT NOT NULL,
|
||||
remote_config TEXT,
|
||||
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
CREATE TABLE IF NOT EXISTS usage_cost_reservations (
|
||||
`request_id` VARCHAR(128) NOT NULL,
|
||||
`subject_id` VARCHAR(128) NOT NULL,
|
||||
`reservation_token` VARCHAR(128) NOT NULL,
|
||||
`admitted_at` BIGINT NOT NULL,
|
||||
`reserved_cost_units` BIGINT NOT NULL,
|
||||
`actual_cost_units` BIGINT,
|
||||
`state` VARCHAR(20) NOT NULL,
|
||||
`reservation_expires_at` BIGINT NOT NULL,
|
||||
`retain_until` BIGINT NOT NULL,
|
||||
`finalized_at` BIGINT,
|
||||
`created_at` BIGINT NOT NULL,
|
||||
`updated_at` BIGINT NOT NULL,
|
||||
PRIMARY KEY (`reservation_token`),
|
||||
CONSTRAINT usage_cost_reservations_state_check
|
||||
CHECK (`state` IN ('reserved', 'finalized', 'released')),
|
||||
CONSTRAINT usage_cost_reservations_reserved_cost_units_check
|
||||
CHECK (`reserved_cost_units` >= 0),
|
||||
CONSTRAINT usage_cost_reservations_actual_cost_units_check
|
||||
CHECK (`actual_cost_units` IS NULL OR `actual_cost_units` >= 0),
|
||||
CONSTRAINT usage_cost_reservations_expiry_check
|
||||
CHECK (`reservation_expires_at` > `admitted_at`),
|
||||
CONSTRAINT usage_cost_reservations_retention_check
|
||||
CHECK (`retain_until` >= `reservation_expires_at`),
|
||||
CONSTRAINT usage_cost_reservations_lifecycle_check CHECK (
|
||||
(`state` = 'reserved' AND `actual_cost_units` IS NULL AND `finalized_at` IS NULL)
|
||||
OR (`state` = 'finalized' AND `actual_cost_units` IS NOT NULL AND `finalized_at` IS NOT NULL)
|
||||
OR (`state` = 'released' AND `actual_cost_units` IS NOT NULL
|
||||
AND `actual_cost_units` = 0 AND `finalized_at` IS NOT NULL)
|
||||
),
|
||||
KEY usage_cost_reservations_request_id_idx (`request_id`),
|
||||
KEY usage_cost_reservations_subject_admitted_at_idx (`subject_id`, `admitted_at`),
|
||||
KEY usage_cost_reservations_reservation_expires_at_idx (`reservation_expires_at`),
|
||||
KEY usage_cost_reservations_retain_until_token_idx (`retain_until`, `reservation_token`)
|
||||
);
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
CREATE TABLE IF NOT EXISTS usage_request_admissions (
|
||||
`request_id` VARCHAR(128) NOT NULL,
|
||||
`subject_id` VARCHAR(128) NOT NULL,
|
||||
`event_token` VARCHAR(128) NOT NULL,
|
||||
`admitted_at` BIGINT NOT NULL,
|
||||
`retain_until` BIGINT NOT NULL,
|
||||
`state` VARCHAR(20) NOT NULL,
|
||||
`released_at` BIGINT,
|
||||
`created_at` BIGINT NOT NULL,
|
||||
PRIMARY KEY (`event_token`),
|
||||
CONSTRAINT usage_request_admissions_retention_check
|
||||
CHECK (`retain_until` > `admitted_at`),
|
||||
CONSTRAINT usage_request_admissions_state_check
|
||||
CHECK (`state` IN ('active', 'released')),
|
||||
CONSTRAINT usage_request_admissions_lifecycle_check CHECK (
|
||||
(`state` = 'active' AND `released_at` IS NULL)
|
||||
OR (`state` = 'released' AND `released_at` IS NOT NULL
|
||||
AND `released_at` >= `admitted_at`)
|
||||
),
|
||||
KEY usage_request_admissions_subject_admitted_at_idx (`subject_id`, `admitted_at`),
|
||||
KEY usage_request_admissions_retain_until_token_idx (`retain_until`, `event_token`)
|
||||
);
|
||||
+12
@@ -0,0 +1,12 @@
|
||||
-- Preserve the already-published cost-ledger migration checksum. This
|
||||
-- follow-up also upgrades databases which ran the feature branch before user
|
||||
-- ownership was enforced. Keep one ALTER TABLE per migration because MySQL
|
||||
-- DDL implicitly commits.
|
||||
DELETE reservation
|
||||
FROM usage_cost_reservations AS reservation
|
||||
LEFT JOIN users AS app_user ON app_user.id = reservation.subject_id
|
||||
WHERE app_user.id IS NULL;
|
||||
|
||||
ALTER TABLE usage_cost_reservations
|
||||
ADD CONSTRAINT usage_cost_reservations_subject_id_fkey
|
||||
FOREIGN KEY (`subject_id`) REFERENCES users (`id`) ON DELETE CASCADE;
|
||||
+10
@@ -0,0 +1,10 @@
|
||||
-- Split from the cost-ledger foreign key so a MySQL implicit DDL commit cannot
|
||||
-- leave two table changes behind one dirty migration record.
|
||||
DELETE admission
|
||||
FROM usage_request_admissions AS admission
|
||||
LEFT JOIN users AS app_user ON app_user.id = admission.subject_id
|
||||
WHERE app_user.id IS NULL;
|
||||
|
||||
ALTER TABLE usage_request_admissions
|
||||
ADD CONSTRAINT usage_request_admissions_subject_id_fkey
|
||||
FOREIGN KEY (`subject_id`) REFERENCES users (`id`) ON DELETE CASCADE;
|
||||
+55
@@ -0,0 +1,55 @@
|
||||
-- A gateway transaction identifier may repeat across payment methods, but
|
||||
-- must never identify two orders in the same normalized method. MySQL commits
|
||||
-- persistent DDL implicitly, so reject historical conflicts before any
|
||||
-- persistent UPDATE or ALTER TABLE. CREATE/DROP TEMPORARY TABLE do not cause an
|
||||
-- implicit commit, and the leading DROP also makes a same-session retry safe.
|
||||
-- Diagnose with:
|
||||
-- SELECT LOWER(TRIM(payment_method)),
|
||||
-- CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin,
|
||||
-- COUNT(*)
|
||||
-- FROM payment_orders
|
||||
-- WHERE gateway_order_id IS NOT NULL
|
||||
-- GROUP BY LOWER(TRIM(payment_method)),
|
||||
-- CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin
|
||||
-- HAVING COUNT(*) > 1;
|
||||
DROP TEMPORARY TABLE IF EXISTS aether_payment_gateway_order_uniqueness_preflight;
|
||||
|
||||
CREATE TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight (
|
||||
conflict_marker TINYINT NOT NULL PRIMARY KEY
|
||||
);
|
||||
|
||||
INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker)
|
||||
VALUES (1);
|
||||
|
||||
-- Inserting the same marker fails on the first conflicting group. The grouping
|
||||
-- mirrors the values and collations used by the normalization and final index:
|
||||
-- payment methods use their existing column collation after LOWER/TRIM, while
|
||||
-- opaque gateway identifiers use MySQL 8's case-sensitive binary collation.
|
||||
INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker)
|
||||
SELECT 1
|
||||
FROM payment_orders
|
||||
WHERE gateway_order_id IS NOT NULL
|
||||
GROUP BY
|
||||
LOWER(TRIM(payment_method)),
|
||||
CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin
|
||||
HAVING COUNT(*) > 1
|
||||
LIMIT 1;
|
||||
|
||||
DROP TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight;
|
||||
|
||||
UPDATE payment_orders
|
||||
SET payment_method = LOWER(TRIM(payment_method))
|
||||
WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method));
|
||||
|
||||
UPDATE payment_callbacks
|
||||
SET payment_method = LOWER(TRIM(payment_method))
|
||||
WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method));
|
||||
|
||||
-- Gateway identifiers are opaque and case-sensitive. Changing the column
|
||||
-- collation and adding the unique index in one ALTER avoids a persistent
|
||||
-- intermediate schema if either operation fails.
|
||||
ALTER TABLE payment_orders
|
||||
MODIFY COLUMN gateway_order_id VARCHAR(128)
|
||||
CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin NULL,
|
||||
ADD UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id
|
||||
(payment_method, gateway_order_id);
|
||||
+5
@@ -0,0 +1,5 @@
|
||||
ALTER TABLE users
|
||||
ADD COLUMN security_version BIGINT NOT NULL DEFAULT 0;
|
||||
|
||||
ALTER TABLE user_sessions
|
||||
ADD COLUMN security_version BIGINT NOT NULL DEFAULT 0;
|
||||
+16
@@ -0,0 +1,16 @@
|
||||
UPDATE request_candidates
|
||||
SET
|
||||
username = NULL,
|
||||
api_key_name = NULL,
|
||||
extra_data = NULL,
|
||||
required_capabilities = NULL,
|
||||
error_message = NULL,
|
||||
error_type = NULL,
|
||||
skip_reason = NULL
|
||||
WHERE username IS NOT NULL
|
||||
OR api_key_name IS NOT NULL
|
||||
OR extra_data IS NOT NULL
|
||||
OR required_capabilities IS NOT NULL
|
||||
OR error_message IS NOT NULL
|
||||
OR error_type IS NOT NULL
|
||||
OR skip_reason IS NOT NULL;
|
||||
+124
@@ -0,0 +1,124 @@
|
||||
-- Historical HTTP captures predate the deny-by-default persistence policy and
|
||||
-- may contain credentials or request content. Keep the usage and projected
|
||||
-- billing facts, but remove every legacy copy of the raw exchange.
|
||||
UPDATE `usage` AS usage_rows
|
||||
LEFT JOIN usage_cost_reservations AS reservation
|
||||
ON reservation.state = 'reserved'
|
||||
AND REGEXP_LIKE(
|
||||
reservation.reservation_token,
|
||||
'^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$',
|
||||
'c'
|
||||
)
|
||||
AND reservation.request_id = usage_rows.request_id
|
||||
AND reservation.subject_id = usage_rows.user_id
|
||||
AND reservation.reservation_token = JSON_UNQUOTE(JSON_EXTRACT(
|
||||
CASE
|
||||
WHEN JSON_VALID(usage_rows.request_metadata) = 1
|
||||
THEN usage_rows.request_metadata
|
||||
ELSE JSON_OBJECT()
|
||||
END,
|
||||
'$.plan_usage_reservation_token'
|
||||
))
|
||||
SET usage_rows.error_message = NULL,
|
||||
usage_rows.error_category = CASE
|
||||
WHEN usage_rows.error_category IS NULL
|
||||
OR TRIM(usage_rows.error_category) = '' THEN NULL
|
||||
WHEN LOWER(TRIM(usage_rows.error_category)) IN (
|
||||
'auth', 'cancelled', 'client_error', 'http_error',
|
||||
'non_success_status', 'provider_error', 'rate_limit', 'redirect',
|
||||
'server_error', 'stream_missing_terminal_event',
|
||||
'stream_terminal_error', 'upstream_error'
|
||||
) THEN LOWER(TRIM(usage_rows.error_category))
|
||||
ELSE 'other_error'
|
||||
END,
|
||||
usage_rows.request_headers = NULL,
|
||||
usage_rows.request_body = NULL,
|
||||
usage_rows.provider_request_headers = NULL,
|
||||
usage_rows.provider_request_body = NULL,
|
||||
usage_rows.response_headers = NULL,
|
||||
usage_rows.response_body = NULL,
|
||||
usage_rows.client_response_headers = NULL,
|
||||
usage_rows.client_response_body = NULL,
|
||||
usage_rows.request_body_compressed = NULL,
|
||||
usage_rows.provider_request_body_compressed = NULL,
|
||||
usage_rows.response_body_compressed = NULL,
|
||||
usage_rows.client_response_body_compressed = NULL,
|
||||
usage_rows.request_metadata = CASE
|
||||
WHEN reservation.reservation_token IS NULL THEN NULL
|
||||
WHEN JSON_TYPE(JSON_EXTRACT(
|
||||
CASE
|
||||
WHEN JSON_VALID(usage_rows.request_metadata) = 1
|
||||
THEN usage_rows.request_metadata
|
||||
ELSE JSON_OBJECT()
|
||||
END,
|
||||
'$.plan_usage_reservation_deferred'
|
||||
)) = 'BOOLEAN'
|
||||
AND JSON_UNQUOTE(JSON_EXTRACT(
|
||||
CASE
|
||||
WHEN JSON_VALID(usage_rows.request_metadata) = 1
|
||||
THEN usage_rows.request_metadata
|
||||
ELSE JSON_OBJECT()
|
||||
END,
|
||||
'$.plan_usage_reservation_deferred'
|
||||
)) = 'true'
|
||||
THEN JSON_SET(
|
||||
JSON_OBJECT(
|
||||
'plan_usage_reservation_token', reservation.reservation_token
|
||||
),
|
||||
'$.plan_usage_reservation_deferred',
|
||||
JSON_EXTRACT('true', '$')
|
||||
)
|
||||
ELSE JSON_OBJECT(
|
||||
'plan_usage_reservation_token', reservation.reservation_token
|
||||
)
|
||||
END
|
||||
WHERE usage_rows.error_message IS NOT NULL
|
||||
OR usage_rows.error_category IS NOT NULL
|
||||
OR usage_rows.request_headers IS NOT NULL
|
||||
OR usage_rows.request_body IS NOT NULL
|
||||
OR usage_rows.provider_request_headers IS NOT NULL
|
||||
OR usage_rows.provider_request_body IS NOT NULL
|
||||
OR usage_rows.response_headers IS NOT NULL
|
||||
OR usage_rows.response_body IS NOT NULL
|
||||
OR usage_rows.client_response_headers IS NOT NULL
|
||||
OR usage_rows.client_response_body IS NOT NULL
|
||||
OR usage_rows.request_body_compressed IS NOT NULL
|
||||
OR usage_rows.provider_request_body_compressed IS NOT NULL
|
||||
OR usage_rows.response_body_compressed IS NOT NULL
|
||||
OR usage_rows.client_response_body_compressed IS NOT NULL
|
||||
OR usage_rows.request_metadata IS NOT NULL;
|
||||
|
||||
DELETE FROM usage_http_audits;
|
||||
|
||||
DELETE FROM usage_body_blobs;
|
||||
|
||||
-- All fields needed by billing reads are projected into typed columns on this
|
||||
-- table. The historical JSON documents include rule expressions, catalog
|
||||
-- snapshots, and arbitrary dimensions, so they are not retained.
|
||||
UPDATE usage_settlement_snapshots
|
||||
SET settlement_snapshot = NULL,
|
||||
billing_dimensions = NULL
|
||||
WHERE settlement_snapshot IS NOT NULL
|
||||
OR billing_dimensions IS NOT NULL;
|
||||
|
||||
UPDATE video_tasks
|
||||
SET original_request_body = NULL,
|
||||
converted_request_body = NULL,
|
||||
progress_message = NULL,
|
||||
error_message = NULL,
|
||||
request_metadata = NULL,
|
||||
video_url = NULL,
|
||||
video_urls = NULL,
|
||||
thumbnail_url = NULL,
|
||||
stored_video_path = NULL,
|
||||
webhook_url = NULL
|
||||
WHERE original_request_body IS NOT NULL
|
||||
OR converted_request_body IS NOT NULL
|
||||
OR progress_message IS NOT NULL
|
||||
OR error_message IS NOT NULL
|
||||
OR request_metadata IS NOT NULL
|
||||
OR video_url IS NOT NULL
|
||||
OR video_urls IS NOT NULL
|
||||
OR thumbnail_url IS NOT NULL
|
||||
OR stored_video_path IS NOT NULL
|
||||
OR webhook_url IS NOT NULL;
|
||||
+43
@@ -0,0 +1,43 @@
|
||||
UPDATE background_task_runs
|
||||
SET owner_instance = NULL,
|
||||
created_by = CASE
|
||||
WHEN LOWER(TRIM(created_by)) IN ('admin', 'scheduler', 'system')
|
||||
THEN LOWER(TRIM(created_by))
|
||||
ELSE NULL
|
||||
END,
|
||||
progress_message = NULL,
|
||||
payload_json = NULL,
|
||||
result_json = NULL,
|
||||
error_message = CASE
|
||||
WHEN status = 'failed' THEN 'background_task_failed'
|
||||
ELSE NULL
|
||||
END
|
||||
WHERE owner_instance IS NOT NULL
|
||||
OR created_by IS NOT NULL
|
||||
OR progress_message IS NOT NULL
|
||||
OR payload_json IS NOT NULL
|
||||
OR result_json IS NOT NULL
|
||||
OR error_message IS NOT NULL;
|
||||
|
||||
UPDATE background_task_events
|
||||
SET event_type = CASE
|
||||
WHEN event_type IN (
|
||||
'cancel_requested', 'failed', 'queued', 'running',
|
||||
'skipped', 'succeeded', 'worker_boot'
|
||||
) THEN event_type
|
||||
ELSE 'unclassified_event'
|
||||
END,
|
||||
message = CASE
|
||||
WHEN event_type IN (
|
||||
'cancel_requested', 'failed', 'queued', 'running',
|
||||
'skipped', 'succeeded', 'worker_boot'
|
||||
) THEN event_type
|
||||
ELSE 'unclassified_event'
|
||||
END,
|
||||
payload_json = NULL
|
||||
WHERE message <> event_type
|
||||
OR payload_json IS NOT NULL
|
||||
OR event_type NOT IN (
|
||||
'cancel_requested', 'failed', 'queued', 'running',
|
||||
'skipped', 'succeeded', 'worker_boot'
|
||||
);
|
||||
+13
@@ -0,0 +1,13 @@
|
||||
-- Stripe PaymentIntent client secrets are payment capabilities. New writes
|
||||
-- store only encrypted values; legacy plaintext cannot be migrated safely
|
||||
-- without the application encryption key, so remove it fail-closed.
|
||||
UPDATE payment_orders
|
||||
SET status = CASE
|
||||
WHEN LOWER(TRIM(payment_method)) = 'stripe' AND status = 'pending'
|
||||
THEN 'expired'
|
||||
ELSE status
|
||||
END,
|
||||
gateway_response = JSON_REMOVE(gateway_response, '$.client_secret')
|
||||
WHERE gateway_response IS NOT NULL
|
||||
AND JSON_VALID(gateway_response) = 1
|
||||
AND JSON_CONTAINS_PATH(gateway_response, 'one', '$.client_secret') = 1;
|
||||
+6
@@ -0,0 +1,6 @@
|
||||
-- Callback idempotency uses payload_hash and does not require the raw body.
|
||||
-- Older versions stored provider-controlled payloads that may contain payment
|
||||
-- capabilities or customer PII, so remove those legacy copies fail-closed.
|
||||
UPDATE payment_callbacks
|
||||
SET payload = NULL
|
||||
WHERE payload IS NOT NULL;
|
||||
+5
@@ -0,0 +1,5 @@
|
||||
-- Video prompts are user content and are not required for polling or billing.
|
||||
-- Remove historical copies now that new writes discard them before persistence.
|
||||
UPDATE video_tasks
|
||||
SET prompt = NULL
|
||||
WHERE prompt IS NOT NULL;
|
||||
+6
@@ -0,0 +1,6 @@
|
||||
-- Identity claims required by the application are stored in dedicated columns.
|
||||
-- Remove legacy provider-controlled userinfo JSON, which may contain unrelated
|
||||
-- PII or credentials and is not needed for authentication or account binding.
|
||||
UPDATE user_oauth_links
|
||||
SET extra_data = NULL
|
||||
WHERE extra_data IS NOT NULL;
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
-- Purpose-bound Fernet envelopes can exceed the former 500-character plaintext limit.
|
||||
ALTER TABLE proxy_nodes
|
||||
MODIFY COLUMN proxy_password TEXT NULL;
|
||||
+377
@@ -0,0 +1,377 @@
|
||||
SET @aether_drop_fact_user_fk_sql := IF(
|
||||
EXISTS (
|
||||
SELECT 1 FROM information_schema.TABLE_CONSTRAINTS
|
||||
WHERE CONSTRAINT_SCHEMA = DATABASE()
|
||||
AND TABLE_NAME = 'user_plan_entitlements'
|
||||
AND CONSTRAINT_NAME = 'user_plan_entitlements_user_id_fkey'
|
||||
AND CONSTRAINT_TYPE = 'FOREIGN KEY'
|
||||
),
|
||||
'ALTER TABLE user_plan_entitlements DROP FOREIGN KEY user_plan_entitlements_user_id_fkey',
|
||||
'DO 0'
|
||||
);
|
||||
PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql;
|
||||
EXECUTE aether_drop_fact_user_fk_stmt;
|
||||
DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt;
|
||||
|
||||
SET @aether_drop_fact_user_fk_sql := IF(
|
||||
EXISTS (
|
||||
SELECT 1 FROM information_schema.TABLE_CONSTRAINTS
|
||||
WHERE CONSTRAINT_SCHEMA = DATABASE()
|
||||
AND TABLE_NAME = 'entitlement_usage_ledgers'
|
||||
AND CONSTRAINT_NAME = 'entitlement_usage_ledgers_user_id_fkey'
|
||||
AND CONSTRAINT_TYPE = 'FOREIGN KEY'
|
||||
),
|
||||
'ALTER TABLE entitlement_usage_ledgers DROP FOREIGN KEY entitlement_usage_ledgers_user_id_fkey',
|
||||
'DO 0'
|
||||
);
|
||||
PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql;
|
||||
EXECUTE aether_drop_fact_user_fk_stmt;
|
||||
DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt;
|
||||
|
||||
SET @aether_drop_fact_user_fk_sql := IF(
|
||||
EXISTS (
|
||||
SELECT 1 FROM information_schema.TABLE_CONSTRAINTS
|
||||
WHERE CONSTRAINT_SCHEMA = DATABASE()
|
||||
AND TABLE_NAME = 'user_referrals'
|
||||
AND CONSTRAINT_NAME = 'user_referrals_inviter_user_id_fkey'
|
||||
AND CONSTRAINT_TYPE = 'FOREIGN KEY'
|
||||
),
|
||||
'ALTER TABLE user_referrals DROP FOREIGN KEY user_referrals_inviter_user_id_fkey',
|
||||
'DO 0'
|
||||
);
|
||||
PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql;
|
||||
EXECUTE aether_drop_fact_user_fk_stmt;
|
||||
DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt;
|
||||
|
||||
SET @aether_drop_fact_user_fk_sql := IF(
|
||||
EXISTS (
|
||||
SELECT 1 FROM information_schema.TABLE_CONSTRAINTS
|
||||
WHERE CONSTRAINT_SCHEMA = DATABASE()
|
||||
AND TABLE_NAME = 'user_referrals'
|
||||
AND CONSTRAINT_NAME = 'user_referrals_invitee_user_id_fkey'
|
||||
AND CONSTRAINT_TYPE = 'FOREIGN KEY'
|
||||
),
|
||||
'ALTER TABLE user_referrals DROP FOREIGN KEY user_referrals_invitee_user_id_fkey',
|
||||
'DO 0'
|
||||
);
|
||||
PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql;
|
||||
EXECUTE aether_drop_fact_user_fk_stmt;
|
||||
DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt;
|
||||
|
||||
SET @aether_drop_fact_user_fk_sql := IF(
|
||||
EXISTS (
|
||||
SELECT 1 FROM information_schema.TABLE_CONSTRAINTS
|
||||
WHERE CONSTRAINT_SCHEMA = DATABASE()
|
||||
AND TABLE_NAME = 'referral_rewards'
|
||||
AND CONSTRAINT_NAME = 'referral_rewards_inviter_user_id_fkey'
|
||||
AND CONSTRAINT_TYPE = 'FOREIGN KEY'
|
||||
),
|
||||
'ALTER TABLE referral_rewards DROP FOREIGN KEY referral_rewards_inviter_user_id_fkey',
|
||||
'DO 0'
|
||||
);
|
||||
PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql;
|
||||
EXECUTE aether_drop_fact_user_fk_stmt;
|
||||
DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt;
|
||||
|
||||
SET @aether_drop_fact_user_fk_sql := IF(
|
||||
EXISTS (
|
||||
SELECT 1 FROM information_schema.TABLE_CONSTRAINTS
|
||||
WHERE CONSTRAINT_SCHEMA = DATABASE()
|
||||
AND TABLE_NAME = 'referral_rewards'
|
||||
AND CONSTRAINT_NAME = 'referral_rewards_invitee_user_id_fkey'
|
||||
AND CONSTRAINT_TYPE = 'FOREIGN KEY'
|
||||
),
|
||||
'ALTER TABLE referral_rewards DROP FOREIGN KEY referral_rewards_invitee_user_id_fkey',
|
||||
'DO 0'
|
||||
);
|
||||
PREPARE aether_drop_fact_user_fk_stmt FROM @aether_drop_fact_user_fk_sql;
|
||||
EXECUTE aether_drop_fact_user_fk_stmt;
|
||||
DEALLOCATE PREPARE aether_drop_fact_user_fk_stmt;
|
||||
|
||||
UPDATE request_candidates
|
||||
SET username = NULL, api_key_name = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = request_candidates.user_id
|
||||
)
|
||||
AND (username IS NOT NULL OR api_key_name IS NOT NULL);
|
||||
|
||||
UPDATE request_candidates
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_name IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = request_candidates.api_key_id
|
||||
);
|
||||
|
||||
UPDATE video_tasks
|
||||
SET username = NULL, api_key_name = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = video_tasks.user_id
|
||||
)
|
||||
AND (username IS NOT NULL OR api_key_name IS NOT NULL);
|
||||
|
||||
UPDATE video_tasks
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_name IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = video_tasks.api_key_id
|
||||
);
|
||||
|
||||
UPDATE `usage`
|
||||
SET username = NULL, api_key_name = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = `usage`.user_id
|
||||
)
|
||||
AND (username IS NOT NULL OR api_key_name IS NOT NULL);
|
||||
|
||||
UPDATE `usage`
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_name IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = `usage`.api_key_id
|
||||
);
|
||||
|
||||
UPDATE stats_user_daily
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_summary
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_summary.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_model
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_model.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_provider
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_provider.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_api_format
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_api_format.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_model_provider
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_model_provider.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_cost_savings.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings_provider
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_cost_savings_provider.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings_model
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_cost_savings_model.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings_model_provider
|
||||
SET username = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = stats_user_daily_cost_savings_model_provider.user_id
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_daily_api_key
|
||||
SET api_key_name = NULL
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = stats_daily_api_key.api_key_id
|
||||
)
|
||||
AND api_key_name IS NOT NULL;
|
||||
|
||||
UPDATE user_plan_entitlements AS entitlement
|
||||
SET status = CASE WHEN status = 'active' THEN 'revoked' ELSE status END,
|
||||
expires_at = LEAST(expires_at, UNIX_TIMESTAMP()),
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = entitlement.user_id
|
||||
);
|
||||
|
||||
UPDATE wallets AS wallet
|
||||
SET status = 'disabled', updated_at = UNIX_TIMESTAMP()
|
||||
WHERE (wallet.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = wallet.user_id
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = wallet.api_key_id
|
||||
));
|
||||
|
||||
UPDATE user_referrals AS referral
|
||||
SET invite_code_snapshot = 'deleted-user',
|
||||
source_json = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = referral.inviter_user_id
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = referral.invitee_user_id
|
||||
);
|
||||
|
||||
UPDATE referral_rewards AS reward
|
||||
SET status = CASE
|
||||
WHEN status IN ('pending', 'failed', 'applying') THEN 'voided'
|
||||
ELSE status
|
||||
END,
|
||||
failure_reason = NULL,
|
||||
admin_note = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = reward.inviter_user_id
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = reward.invitee_user_id
|
||||
);
|
||||
|
||||
UPDATE audit_logs AS history
|
||||
SET description = 'deleted user event',
|
||||
ip_address = NULL,
|
||||
user_agent = NULL,
|
||||
event_metadata = NULL,
|
||||
error_message = NULL
|
||||
WHERE (history.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.user_id
|
||||
))
|
||||
OR (history.api_key_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = history.api_key_id
|
||||
));
|
||||
|
||||
UPDATE wallet_transactions AS history
|
||||
SET description = NULL
|
||||
WHERE EXISTS (
|
||||
SELECT 1
|
||||
FROM wallets AS wallet
|
||||
WHERE wallet.id = history.wallet_id
|
||||
AND (
|
||||
(wallet.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = wallet.user_id
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = wallet.api_key_id
|
||||
))
|
||||
)
|
||||
);
|
||||
|
||||
UPDATE wallet_transactions AS history
|
||||
SET description = NULL
|
||||
WHERE history.operator_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.operator_id
|
||||
);
|
||||
|
||||
UPDATE payment_orders AS history
|
||||
SET gateway_response = NULL
|
||||
WHERE (history.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.user_id
|
||||
))
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM wallets AS wallet
|
||||
WHERE wallet.id = history.wallet_id
|
||||
AND (
|
||||
(wallet.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = wallet.user_id
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = wallet.api_key_id
|
||||
))
|
||||
)
|
||||
);
|
||||
|
||||
UPDATE payment_callbacks AS history
|
||||
SET payload = NULL,
|
||||
error_message = NULL
|
||||
WHERE EXISTS (
|
||||
SELECT 1
|
||||
FROM payment_orders AS payment_order
|
||||
LEFT JOIN wallets AS wallet ON wallet.id = payment_order.wallet_id
|
||||
WHERE (
|
||||
payment_order.id = history.payment_order_id
|
||||
OR (history.order_no IS NOT NULL AND payment_order.order_no = history.order_no)
|
||||
)
|
||||
AND (
|
||||
(payment_order.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = payment_order.user_id
|
||||
))
|
||||
OR (wallet.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = wallet.user_id
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = wallet.api_key_id
|
||||
))
|
||||
)
|
||||
);
|
||||
|
||||
UPDATE refund_requests AS history
|
||||
SET reason = NULL,
|
||||
payout_reference = NULL,
|
||||
payout_proof = NULL,
|
||||
failure_reason = NULL
|
||||
WHERE (history.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.user_id
|
||||
))
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM wallets AS wallet
|
||||
WHERE wallet.id = history.wallet_id
|
||||
AND (
|
||||
(wallet.user_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = wallet.user_id
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM api_keys WHERE api_keys.id = wallet.api_key_id
|
||||
))
|
||||
)
|
||||
)
|
||||
OR (history.requested_by IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.requested_by
|
||||
))
|
||||
OR (history.approved_by IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.approved_by
|
||||
))
|
||||
OR (history.processed_by IS NOT NULL AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.processed_by
|
||||
));
|
||||
|
||||
UPDATE referral_rewards AS reward
|
||||
SET failure_reason = NULL,
|
||||
admin_note = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE reward.admin_operator_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = reward.admin_operator_id
|
||||
);
|
||||
|
||||
UPDATE redeem_code_batches AS history
|
||||
SET description = NULL
|
||||
WHERE history.created_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users WHERE users.id = history.created_by
|
||||
);
|
||||
+451
@@ -0,0 +1,451 @@
|
||||
-- Keep nullable/system-owned history rows intact while clearing values tied to
|
||||
-- deleted users or API keys. Raw payment callback payloads are always removed
|
||||
-- because they are not required for idempotency and may contain PII.
|
||||
UPDATE request_candidates
|
||||
SET username = NULL, api_key_name = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = request_candidates.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND (username IS NOT NULL OR api_key_name IS NOT NULL);
|
||||
|
||||
UPDATE request_candidates
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = request_candidates.api_key_id
|
||||
)
|
||||
AND api_key_name IS NOT NULL;
|
||||
|
||||
UPDATE video_tasks
|
||||
SET username = NULL, api_key_name = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = video_tasks.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND (username IS NOT NULL OR api_key_name IS NOT NULL);
|
||||
|
||||
UPDATE video_tasks
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = video_tasks.api_key_id
|
||||
)
|
||||
AND api_key_name IS NOT NULL;
|
||||
|
||||
UPDATE `usage`
|
||||
SET username = NULL, api_key_name = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = `usage`.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND (username IS NOT NULL OR api_key_name IS NOT NULL);
|
||||
|
||||
UPDATE `usage`
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = `usage`.api_key_id
|
||||
)
|
||||
AND api_key_name IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_summary
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_summary.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_model
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_model.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_provider
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_provider.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_api_format
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_api_format.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_model_provider
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_model_provider.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_cost_savings.user_id AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings_provider
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_cost_savings_provider.user_id
|
||||
AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings_model
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = stats_user_daily_cost_savings_model.user_id
|
||||
AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_user_daily_cost_savings_model_provider
|
||||
SET username = NULL
|
||||
WHERE user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM users
|
||||
WHERE users.id = stats_user_daily_cost_savings_model_provider.user_id
|
||||
AND users.is_deleted = 0
|
||||
)
|
||||
AND username IS NOT NULL;
|
||||
|
||||
UPDATE stats_daily_api_key
|
||||
SET api_key_name = NULL
|
||||
WHERE api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = stats_daily_api_key.api_key_id
|
||||
)
|
||||
AND api_key_name IS NOT NULL;
|
||||
|
||||
UPDATE user_plan_entitlements AS entitlement
|
||||
SET status = CASE WHEN status = 'active' THEN 'revoked' ELSE status END,
|
||||
expires_at = LEAST(expires_at, UNIX_TIMESTAMP()),
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = entitlement.user_id AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE user_referrals AS referral
|
||||
SET invite_code_snapshot = 'deleted-user',
|
||||
source_json = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = referral.inviter_user_id AND users.is_deleted = 0
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = referral.invitee_user_id AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE referral_rewards AS reward
|
||||
SET status = CASE
|
||||
WHEN status IN ('pending', 'failed', 'applying') THEN 'voided'
|
||||
ELSE status
|
||||
END,
|
||||
failure_reason = NULL,
|
||||
admin_note = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = reward.inviter_user_id AND users.is_deleted = 0
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = reward.invitee_user_id AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE referral_rewards AS reward
|
||||
SET failure_reason = NULL,
|
||||
admin_note = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE reward.admin_operator_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = reward.admin_operator_id AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE wallets AS wallet
|
||||
SET status = 'disabled', updated_at = UNIX_TIMESTAMP()
|
||||
WHERE (wallet.user_id IS NULL AND wallet.api_key_id IS NULL)
|
||||
OR (wallet.user_id IS NOT NULL AND wallet.api_key_id IS NOT NULL)
|
||||
OR (wallet.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = wallet.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = wallet.api_key_id
|
||||
));
|
||||
|
||||
UPDATE audit_logs AS history
|
||||
SET description = 'deleted user event',
|
||||
ip_address = NULL,
|
||||
user_agent = NULL,
|
||||
event_metadata = NULL,
|
||||
error_message = NULL
|
||||
WHERE (history.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR (history.api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = history.api_key_id
|
||||
));
|
||||
|
||||
UPDATE wallet_transactions AS history
|
||||
SET description = NULL
|
||||
WHERE EXISTS (
|
||||
SELECT 1
|
||||
FROM wallets AS wallet
|
||||
WHERE wallet.id = history.wallet_id
|
||||
AND ((wallet.user_id IS NULL AND wallet.api_key_id IS NULL)
|
||||
OR (wallet.user_id IS NOT NULL AND wallet.api_key_id IS NOT NULL)
|
||||
OR (wallet.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = wallet.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id
|
||||
AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = wallet.api_key_id
|
||||
)))
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM wallets AS wallet WHERE wallet.id = history.wallet_id
|
||||
)
|
||||
OR (history.operator_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.operator_id AND users.is_deleted = 0
|
||||
));
|
||||
|
||||
UPDATE payment_orders AS history
|
||||
SET gateway_response = NULL
|
||||
WHERE (history.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM wallets AS wallet
|
||||
WHERE wallet.id = history.wallet_id
|
||||
AND ((wallet.user_id IS NULL AND wallet.api_key_id IS NULL)
|
||||
OR (wallet.user_id IS NOT NULL AND wallet.api_key_id IS NOT NULL)
|
||||
OR (wallet.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = wallet.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id
|
||||
AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = wallet.api_key_id
|
||||
)))
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM wallets AS wallet WHERE wallet.id = history.wallet_id
|
||||
);
|
||||
|
||||
-- Raw provider payloads are never needed for idempotency and must be purged
|
||||
-- even when an old callback cannot be linked back to an order.
|
||||
UPDATE payment_callbacks
|
||||
SET payload = NULL
|
||||
WHERE payload IS NOT NULL;
|
||||
|
||||
UPDATE payment_callbacks AS history
|
||||
SET error_message = NULL
|
||||
WHERE history.error_message IS NOT NULL
|
||||
AND (
|
||||
NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM payment_orders AS payment_order
|
||||
WHERE payment_order.id = history.payment_order_id
|
||||
OR (history.order_no IS NOT NULL AND payment_order.order_no = history.order_no)
|
||||
)
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM payment_orders AS payment_order
|
||||
LEFT JOIN wallets AS wallet ON wallet.id = payment_order.wallet_id
|
||||
WHERE (payment_order.id = history.payment_order_id
|
||||
OR (history.order_no IS NOT NULL AND payment_order.order_no = history.order_no))
|
||||
AND ((payment_order.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = payment_order.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR wallet.id IS NULL
|
||||
OR (wallet.user_id IS NULL AND wallet.api_key_id IS NULL)
|
||||
OR (wallet.user_id IS NOT NULL AND wallet.api_key_id IS NOT NULL)
|
||||
OR (wallet.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = wallet.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id
|
||||
AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = wallet.api_key_id
|
||||
)))
|
||||
)
|
||||
);
|
||||
|
||||
UPDATE refund_requests AS history
|
||||
SET reason = NULL,
|
||||
payout_reference = NULL,
|
||||
payout_proof = NULL,
|
||||
failure_reason = NULL
|
||||
WHERE (history.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM wallets AS wallet
|
||||
WHERE wallet.id = history.wallet_id
|
||||
AND ((wallet.user_id IS NULL AND wallet.api_key_id IS NULL)
|
||||
OR (wallet.user_id IS NOT NULL AND wallet.api_key_id IS NOT NULL)
|
||||
OR (wallet.user_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = wallet.user_id AND users.is_deleted = 0
|
||||
))
|
||||
OR (wallet.api_key_id IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM api_keys
|
||||
JOIN users AS api_key_owner
|
||||
ON api_key_owner.id = api_keys.user_id
|
||||
AND api_key_owner.is_deleted = 0
|
||||
WHERE api_keys.id = wallet.api_key_id
|
||||
)))
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1 FROM wallets AS wallet WHERE wallet.id = history.wallet_id
|
||||
)
|
||||
OR (history.requested_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.requested_by AND users.is_deleted = 0
|
||||
))
|
||||
OR (history.approved_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.approved_by AND users.is_deleted = 0
|
||||
))
|
||||
OR (history.processed_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.processed_by AND users.is_deleted = 0
|
||||
));
|
||||
|
||||
UPDATE redeem_code_batches AS history
|
||||
SET description = NULL
|
||||
WHERE history.created_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = history.created_by AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE refund_requests
|
||||
SET requested_by = NULL
|
||||
WHERE requested_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = refund_requests.requested_by AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE refund_requests
|
||||
SET approved_by = NULL
|
||||
WHERE approved_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = refund_requests.approved_by AND users.is_deleted = 0
|
||||
);
|
||||
|
||||
UPDATE refund_requests
|
||||
SET processed_by = NULL
|
||||
WHERE processed_by IS NOT NULL
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM users
|
||||
WHERE users.id = refund_requests.processed_by AND users.is_deleted = 0
|
||||
);
|
||||
+13
@@ -0,0 +1,13 @@
|
||||
-- LDAP configuration is a database-wide singleton. Preserve the row selected by the legacy
|
||||
-- reader (the smallest id), remove historical duplicates, and let the database arbitrate
|
||||
-- concurrent first creation.
|
||||
DELETE FROM ldap_configs
|
||||
WHERE id <> (
|
||||
SELECT keep_id
|
||||
FROM (SELECT MIN(id) AS keep_id FROM ldap_configs) AS ldap_singleton_keeper
|
||||
);
|
||||
|
||||
ALTER TABLE ldap_configs
|
||||
ADD COLUMN singleton_key INT NOT NULL DEFAULT 1,
|
||||
ADD CONSTRAINT ldap_configs_singleton_key_check CHECK (singleton_key = 1),
|
||||
ADD UNIQUE KEY ldap_configs_singleton_key_key (singleton_key);
|
||||
+9
@@ -0,0 +1,9 @@
|
||||
ALTER TABLE proxy_nodes
|
||||
ADD COLUMN tunnel_generation VARCHAR(64) NULL AFTER id;
|
||||
|
||||
UPDATE proxy_nodes
|
||||
SET tunnel_generation = UUID()
|
||||
WHERE tunnel_generation IS NULL OR TRIM(tunnel_generation) = '';
|
||||
|
||||
ALTER TABLE proxy_nodes
|
||||
MODIFY COLUMN tunnel_generation VARCHAR(64) NOT NULL;
|
||||
+10
@@ -0,0 +1,10 @@
|
||||
-- A proxy endpoint has one stable node identity across manual and tunnel
|
||||
-- registrations. MySQL 8 performs ALTER TABLE atomically; if historical
|
||||
-- duplicates exist this migration fails without choosing or deleting a row.
|
||||
-- Diagnose with:
|
||||
-- SELECT ip, port, COUNT(*)
|
||||
-- FROM proxy_nodes
|
||||
-- GROUP BY ip, port
|
||||
-- HAVING COUNT(*) > 1;
|
||||
ALTER TABLE proxy_nodes
|
||||
ADD UNIQUE INDEX uq_proxy_node_ip_port (ip, port);
|
||||
+2
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE usage_counter_deltas
|
||||
ADD COLUMN target_tunnel_generation VARCHAR(64) NULL;
|
||||
+8
@@ -0,0 +1,8 @@
|
||||
-- Older gateway releases marked every OAuth email as verified without
|
||||
-- retaining verification provenance. Reset those claims conservatively; a
|
||||
-- later trusted OAuth assertion for the same normalized email can upgrade it.
|
||||
UPDATE users
|
||||
SET email_verified = 0,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE LOWER(TRIM(auth_source)) = 'oauth'
|
||||
AND email_verified = 1;
|
||||
@@ -3,9 +3,9 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::auth::{
|
||||
AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository,
|
||||
AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord,
|
||||
StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
|
||||
UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
|
||||
AuthApiKeyWriteRepository, CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord,
|
||||
CreateUserApiKeyRecord, StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord,
|
||||
StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
@@ -69,6 +69,21 @@ SELECT
|
||||
FROM api_keys
|
||||
"#;
|
||||
|
||||
const MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL: &[&str] = &[
|
||||
"UPDATE request_candidates SET api_key_name = NULL WHERE api_key_id = ?",
|
||||
"UPDATE video_tasks SET api_key_name = NULL WHERE api_key_id = ?",
|
||||
"UPDATE `usage` SET api_key_name = NULL WHERE api_key_id = ?",
|
||||
"UPDATE stats_daily_api_key SET api_key_name = NULL WHERE api_key_id = ?",
|
||||
"UPDATE audit_logs SET description = 'deleted API key event', ip_address = NULL, user_agent = NULL, event_metadata = NULL, error_message = NULL WHERE api_key_id = ?",
|
||||
"UPDATE wallet_transactions SET description = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)",
|
||||
"UPDATE payment_callbacks SET payload = NULL, error_message = NULL WHERE EXISTS (SELECT 1 FROM payment_orders AS history_order JOIN wallets AS history_wallet ON history_wallet.id = history_order.wallet_id WHERE history_wallet.api_key_id = ? AND (history_order.id = payment_callbacks.payment_order_id OR (payment_callbacks.order_no IS NOT NULL AND history_order.order_no = payment_callbacks.order_no)))",
|
||||
"UPDATE payment_orders SET gateway_response = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)",
|
||||
"UPDATE refund_requests SET reason = NULL, payout_reference = NULL, payout_proof = NULL, failure_reason = NULL WHERE wallet_id IN (SELECT id FROM wallets WHERE api_key_id = ?)",
|
||||
];
|
||||
|
||||
const MYSQL_DELETE_API_KEY_DEPENDENTS_SQL: &[&str] =
|
||||
&["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"];
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlAuthApiKeyReadRepository {
|
||||
pool: MysqlPool,
|
||||
@@ -111,6 +126,17 @@ impl MysqlAuthApiKeyReadRepository {
|
||||
record: CreateApiKeyInsertRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let now = current_unix_secs();
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let owner_exists: Option<String> =
|
||||
sqlx::query_scalar("SELECT id FROM users WHERE id = ? AND is_deleted = 0 FOR UPDATE")
|
||||
.bind(&record.user_id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if owner_exists.is_none() {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
}
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO api_keys (
|
||||
@@ -120,7 +146,7 @@ INSERT INTO api_keys (
|
||||
total_requests, total_tokens, total_cost_usd, is_standalone,
|
||||
created_at, updated_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&record.api_key_id)
|
||||
@@ -150,6 +176,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
&record.force_capabilities,
|
||||
"api_keys.force_capabilities",
|
||||
)?)
|
||||
.bind(optional_json_to_string(
|
||||
&record.feature_settings,
|
||||
"api_keys.feature_settings",
|
||||
)?)
|
||||
.bind(record.is_active)
|
||||
.bind(optional_i64_from_u64(
|
||||
record.expires_at_unix_secs,
|
||||
@@ -165,11 +195,25 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
.bind(record.is_standalone)
|
||||
.bind(now as i64)
|
||||
.bind(now as i64)
|
||||
.execute(&self.pool)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
self.reload_export_by_id(&record.api_key_id).await
|
||||
let reload_sql = format!("{EXPORT_COLUMNS}\nWHERE api_keys.id = ?\nLIMIT 1");
|
||||
let row = sqlx::query(&reload_sql)
|
||||
.bind(&record.api_key_id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(row) = row else {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"created api_keys row is missing: {}",
|
||||
record.api_key_id
|
||||
)));
|
||||
};
|
||||
let created = map_auth_api_key_export_row(&row)?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(Some(created))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,6 +230,7 @@ struct CreateApiKeyInsertRecord {
|
||||
rate_limit: Option<i32>,
|
||||
concurrent_limit: Option<i32>,
|
||||
force_capabilities: Option<serde_json::Value>,
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
is_active: bool,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
auto_delete_on_expiry: bool,
|
||||
@@ -337,6 +382,7 @@ impl AuthApiKeyReadRepository for MysqlAuthApiKeyReadRepository {
|
||||
if user_ids.is_empty() {
|
||||
return Ok(AuthApiKeyExportSummary::default());
|
||||
}
|
||||
let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?;
|
||||
|
||||
let mut builder = QueryBuilder::<MySql>::new(
|
||||
r#"
|
||||
@@ -345,7 +391,7 @@ SELECT
|
||||
SUM(CASE WHEN is_active = 1 AND (expires_at IS NULL OR expires_at >=
|
||||
"#,
|
||||
);
|
||||
builder.push_bind(now_unix_secs as i64);
|
||||
builder.push_bind(now_unix_secs);
|
||||
builder.push(
|
||||
r#") THEN 1 ELSE 0 END) AS active
|
||||
FROM api_keys
|
||||
@@ -429,6 +475,7 @@ WHERE id = ?
|
||||
rate_limit: Some(record.rate_limit),
|
||||
concurrent_limit: record.concurrent_limit,
|
||||
force_capabilities: record.force_capabilities,
|
||||
feature_settings: record.feature_settings,
|
||||
is_active: record.is_active,
|
||||
expires_at_unix_secs: record.expires_at_unix_secs,
|
||||
auto_delete_on_expiry: record.auto_delete_on_expiry,
|
||||
@@ -457,6 +504,7 @@ WHERE id = ?
|
||||
rate_limit: record.rate_limit,
|
||||
concurrent_limit: record.concurrent_limit,
|
||||
force_capabilities: record.force_capabilities,
|
||||
feature_settings: None,
|
||||
is_active: record.is_active,
|
||||
expires_at_unix_secs: record.expires_at_unix_secs,
|
||||
auto_delete_on_expiry: record.auto_delete_on_expiry,
|
||||
@@ -472,35 +520,41 @@ WHERE id = ?
|
||||
&self,
|
||||
record: UpdateUserApiKeyBasicRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let now = current_unix_secs() as i64;
|
||||
sqlx::query(
|
||||
self.update_user_api_key_basic_scoped(record, false).await
|
||||
}
|
||||
|
||||
async fn compare_and_swap_api_key_ciphertext(
|
||||
&self,
|
||||
mutation: &CompareAndSwapAuthApiKeyCiphertext,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET name = COALESCE(?, name),
|
||||
rate_limit = COALESCE(?, rate_limit),
|
||||
concurrent_limit = COALESCE(?, concurrent_limit),
|
||||
ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
SET key_encrypted = ?
|
||||
WHERE BINARY id = BINARY ?
|
||||
AND BINARY user_id = BINARY ?
|
||||
AND BINARY key_hash = BINARY ?
|
||||
AND is_standalone = ?
|
||||
AND BINARY key_encrypted = BINARY ?
|
||||
"#,
|
||||
)
|
||||
.bind(record.name.as_deref())
|
||||
.bind(record.rate_limit)
|
||||
.bind(record.concurrent_limit)
|
||||
.bind(record.ip_rules.is_some())
|
||||
.bind(json_string_from_nested_string_list(
|
||||
&record.ip_rules,
|
||||
"api_keys.ip_rules",
|
||||
)?)
|
||||
.bind(now)
|
||||
.bind(&record.api_key_id)
|
||||
.bind(&record.user_id)
|
||||
.bind(&mutation.key_encrypted)
|
||||
.bind(&mutation.api_key_id)
|
||||
.bind(&mutation.user_id)
|
||||
.bind(&mutation.key_hash)
|
||||
.bind(mutation.is_standalone)
|
||||
.bind(&mutation.expected_key_encrypted)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
self.reload_export_by_id(&record.api_key_id).await
|
||||
Ok(result.rows_affected() == 1)
|
||||
}
|
||||
|
||||
async fn update_user_api_key_basic_if_unlocked(
|
||||
&self,
|
||||
record: UpdateUserApiKeyBasicRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.update_user_api_key_basic_scoped(record, true).await
|
||||
}
|
||||
|
||||
async fn update_standalone_api_key_basic(
|
||||
@@ -511,7 +565,9 @@ WHERE id = ?
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET name = COALESCE(?, name),
|
||||
SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END,
|
||||
name = CASE WHEN ? THEN ? ELSE name END,
|
||||
force_capabilities = CASE WHEN ? THEN ? ELSE force_capabilities END,
|
||||
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
|
||||
concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END,
|
||||
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
|
||||
@@ -525,7 +581,15 @@ WHERE id = ?
|
||||
AND is_standalone = 1
|
||||
"#,
|
||||
)
|
||||
.bind(record.key_encrypted_present)
|
||||
.bind(record.key_encrypted.as_deref())
|
||||
.bind(record.name_present)
|
||||
.bind(record.name.as_deref())
|
||||
.bind(record.force_capabilities.is_some())
|
||||
.bind(optional_json_to_string(
|
||||
&record.force_capabilities.clone().flatten(),
|
||||
"api_keys.force_capabilities",
|
||||
)?)
|
||||
.bind(record.rate_limit_present)
|
||||
.bind(record.rate_limit)
|
||||
.bind(record.concurrent_limit_present)
|
||||
@@ -565,13 +629,143 @@ WHERE id = ?
|
||||
self.reload_export_by_id(&record.api_key_id).await
|
||||
}
|
||||
|
||||
async fn restore_api_key_if_matches(
|
||||
&self,
|
||||
expected: &StoredAuthApiKeyExportRecord,
|
||||
restored: &StoredAuthApiKeyExportRecord,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if restored.api_key_id != expected.api_key_id
|
||||
|| restored.user_id != expected.user_id
|
||||
|| restored.key_hash != expected.key_hash
|
||||
|| restored.is_standalone != expected.is_standalone
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let select_sql = format!("{EXPORT_COLUMNS} WHERE api_keys.id = ? LIMIT 1 FOR UPDATE");
|
||||
let row = sqlx::query(&select_sql)
|
||||
.bind(&expected.api_key_id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(row) = row else {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
};
|
||||
let current = map_auth_api_key_export_row(&row)?;
|
||||
if current != *expected {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET key_encrypted = ?,
|
||||
name = ?,
|
||||
allowed_providers = ?,
|
||||
allowed_api_formats = ?,
|
||||
allowed_models = ?,
|
||||
ip_rules = ?,
|
||||
rate_limit = ?,
|
||||
concurrent_limit = ?,
|
||||
force_capabilities = ?,
|
||||
feature_settings = ?,
|
||||
is_active = ?,
|
||||
expires_at = ?,
|
||||
auto_delete_on_expiry = ?,
|
||||
total_requests = ?,
|
||||
total_tokens = ?,
|
||||
total_cost_usd = ?,
|
||||
last_used_at = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND key_hash = ?
|
||||
AND is_standalone = ?
|
||||
"#,
|
||||
)
|
||||
.bind(restored.key_encrypted.as_deref())
|
||||
.bind(restored.name.as_deref())
|
||||
.bind(json_string_from_string_list(
|
||||
restored.allowed_providers.as_ref(),
|
||||
"api_keys.allowed_providers",
|
||||
)?)
|
||||
.bind(json_string_from_string_list(
|
||||
restored.allowed_api_formats.as_ref(),
|
||||
"api_keys.allowed_api_formats",
|
||||
)?)
|
||||
.bind(json_string_from_string_list(
|
||||
restored.allowed_models.as_ref(),
|
||||
"api_keys.allowed_models",
|
||||
)?)
|
||||
.bind(json_string_from_string_list(
|
||||
restored.ip_rules.as_ref(),
|
||||
"api_keys.ip_rules",
|
||||
)?)
|
||||
.bind(restored.rate_limit)
|
||||
.bind(restored.concurrent_limit)
|
||||
.bind(optional_json_to_string(
|
||||
&restored.force_capabilities,
|
||||
"api_keys.force_capabilities",
|
||||
)?)
|
||||
.bind(optional_json_to_string(
|
||||
&restored.feature_settings,
|
||||
"api_keys.feature_settings",
|
||||
)?)
|
||||
.bind(restored.is_active)
|
||||
.bind(optional_i64_from_u64(
|
||||
restored.expires_at_unix_secs,
|
||||
"api_keys.expires_at",
|
||||
)?)
|
||||
.bind(restored.auto_delete_on_expiry)
|
||||
.bind(i64_from_u64(
|
||||
restored.total_requests,
|
||||
"api_keys.total_requests",
|
||||
)?)
|
||||
.bind(i64_from_u64(
|
||||
restored.total_tokens,
|
||||
"api_keys.total_tokens",
|
||||
)?)
|
||||
.bind(restored.total_cost_usd)
|
||||
.bind(optional_i64_from_u64(
|
||||
restored.last_used_at_unix_secs,
|
||||
"api_keys.last_used_at",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(&restored.api_key_id)
|
||||
.bind(&restored.user_id)
|
||||
.bind(&restored.key_hash)
|
||||
.bind(restored.is_standalone)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() != 1 {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn set_user_api_key_active(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.set_active(api_key_id, Some(user_id), is_active, false)
|
||||
self.set_active(api_key_id, Some(user_id), is_active, false, false)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_active_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.set_active(api_key_id, Some(user_id), is_active, false, true)
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -580,7 +774,8 @@ WHERE id = ?
|
||||
api_key_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.set_active(api_key_id, None, is_active, true).await
|
||||
self.set_active(api_key_id, None, is_active, true, false)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_locked(
|
||||
@@ -615,26 +810,23 @@ WHERE id = ?
|
||||
api_key_id: &str,
|
||||
allowed_providers: Option<Vec<String>>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET allowed_providers = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
"#,
|
||||
self.set_user_api_key_allowed_providers_scoped(
|
||||
user_id,
|
||||
api_key_id,
|
||||
allowed_providers,
|
||||
false,
|
||||
)
|
||||
.bind(json_string_from_string_list(
|
||||
allowed_providers.as_ref(),
|
||||
"api_keys.allowed_providers",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_allowed_providers_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
allowed_providers: Option<Vec<String>>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.set_user_api_key_allowed_providers_scoped(user_id, api_key_id, allowed_providers, true)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_force_capabilities(
|
||||
@@ -643,26 +835,28 @@ WHERE id = ?
|
||||
api_key_id: &str,
|
||||
force_capabilities: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET force_capabilities = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
"#,
|
||||
self.set_user_api_key_force_capabilities_scoped(
|
||||
user_id,
|
||||
api_key_id,
|
||||
force_capabilities,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_force_capabilities_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
force_capabilities: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.set_user_api_key_force_capabilities_scoped(
|
||||
user_id,
|
||||
api_key_id,
|
||||
force_capabilities,
|
||||
true,
|
||||
)
|
||||
.bind(optional_json_to_string(
|
||||
&force_capabilities,
|
||||
"api_keys.force_capabilities",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_feature_settings(
|
||||
@@ -671,26 +865,18 @@ WHERE id = ?
|
||||
api_key_id: &str,
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET feature_settings = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
"#,
|
||||
)
|
||||
.bind(optional_json_to_string(
|
||||
&feature_settings,
|
||||
"api_keys.feature_settings",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, false)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_feature_settings_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
self.set_user_api_key_feature_settings_scoped(user_id, api_key_id, feature_settings, true)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_api_key_usage_totals(
|
||||
@@ -700,6 +886,11 @@ WHERE id = ?
|
||||
total_tokens: u64,
|
||||
total_cost_usd: f64,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
if !total_cost_usd.is_finite() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"api_keys.total_cost_usd is not finite".to_string(),
|
||||
));
|
||||
}
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
@@ -710,8 +901,8 @@ SET total_requests = ?,
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(total_requests as i64)
|
||||
.bind(total_tokens as i64)
|
||||
.bind(i64_from_u64(total_requests, "api_keys.total_requests")?)
|
||||
.bind(i64_from_u64(total_tokens, "api_keys.total_tokens")?)
|
||||
.bind(total_cost_usd)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
@@ -726,11 +917,21 @@ WHERE id = ?
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
self.delete_api_key(api_key_id, Some(user_id), false).await
|
||||
self.delete_api_key(api_key_id, Some(user_id), false, false)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn delete_user_api_key_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
self.delete_api_key(api_key_id, Some(user_id), false, true)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result<bool, DataLayerError> {
|
||||
self.delete_api_key(api_key_id, None, true).await
|
||||
self.delete_api_key(api_key_id, None, true, false).await
|
||||
}
|
||||
|
||||
async fn set_standalone_api_key_feature_settings(
|
||||
@@ -760,12 +961,65 @@ WHERE id = ?
|
||||
}
|
||||
|
||||
impl MysqlAuthApiKeyReadRepository {
|
||||
async fn update_user_api_key_basic_scoped(
|
||||
&self,
|
||||
record: UpdateUserApiKeyBasicRecord,
|
||||
require_unlocked: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET key_encrypted = CASE WHEN ? THEN ? ELSE key_encrypted END,
|
||||
name = CASE WHEN ? THEN ? ELSE name END,
|
||||
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
|
||||
concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END,
|
||||
ip_rules = CASE WHEN ? THEN ? ELSE ip_rules END,
|
||||
feature_settings = CASE WHEN ? THEN ? ELSE feature_settings END,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
AND (? = 0 OR is_locked = 0)
|
||||
"#,
|
||||
)
|
||||
.bind(record.key_encrypted_present)
|
||||
.bind(record.key_encrypted.as_deref())
|
||||
.bind(record.name_present)
|
||||
.bind(record.name.as_deref())
|
||||
.bind(record.rate_limit_present)
|
||||
.bind(record.rate_limit)
|
||||
.bind(record.concurrent_limit_present)
|
||||
.bind(record.concurrent_limit)
|
||||
.bind(record.ip_rules.is_some())
|
||||
.bind(json_string_from_nested_string_list(
|
||||
&record.ip_rules,
|
||||
"api_keys.ip_rules",
|
||||
)?)
|
||||
.bind(record.feature_settings.is_some())
|
||||
.bind(optional_json_to_string(
|
||||
&record.feature_settings.clone().flatten(),
|
||||
"api_keys.feature_settings",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(&record.api_key_id)
|
||||
.bind(&record.user_id)
|
||||
.bind(require_unlocked)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.reload_export_by_id(&record.api_key_id).await
|
||||
}
|
||||
|
||||
async fn set_active(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
user_id: Option<&str>,
|
||||
is_active: bool,
|
||||
is_standalone: bool,
|
||||
require_unlocked: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new("UPDATE api_keys SET is_active = ");
|
||||
builder
|
||||
@@ -779,7 +1033,115 @@ impl MysqlAuthApiKeyReadRepository {
|
||||
if let Some(user_id) = user_id {
|
||||
builder.push(" AND user_id = ").push_bind(user_id);
|
||||
}
|
||||
builder.build().execute(&self.pool).await.map_sql_err()?;
|
||||
if require_unlocked {
|
||||
builder.push(" AND is_locked = ").push_bind(false);
|
||||
}
|
||||
let result = builder.build().execute(&self.pool).await.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_allowed_providers_scoped(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
allowed_providers: Option<Vec<String>>,
|
||||
require_unlocked: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET allowed_providers = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
AND (? = 0 OR is_locked = 0)
|
||||
"#,
|
||||
)
|
||||
.bind(json_string_from_string_list(
|
||||
allowed_providers.as_ref(),
|
||||
"api_keys.allowed_providers",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.bind(require_unlocked)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_force_capabilities_scoped(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
force_capabilities: Option<serde_json::Value>,
|
||||
require_unlocked: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET force_capabilities = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
AND (? = 0 OR is_locked = 0)
|
||||
"#,
|
||||
)
|
||||
.bind(optional_json_to_string(
|
||||
&force_capabilities,
|
||||
"api_keys.force_capabilities",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.bind(require_unlocked)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
}
|
||||
|
||||
async fn set_user_api_key_feature_settings_scoped(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
require_unlocked: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE api_keys
|
||||
SET feature_settings = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND user_id = ?
|
||||
AND is_standalone = 0
|
||||
AND (? = 0 OR is_locked = 0)
|
||||
"#,
|
||||
)
|
||||
.bind(optional_json_to_string(
|
||||
&feature_settings,
|
||||
"api_keys.feature_settings",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.bind(require_unlocked)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.reload_export_by_id(api_key_id).await
|
||||
}
|
||||
|
||||
@@ -788,22 +1150,76 @@ impl MysqlAuthApiKeyReadRepository {
|
||||
api_key_id: &str,
|
||||
user_id: Option<&str>,
|
||||
is_standalone: bool,
|
||||
require_unlocked: bool,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new("DELETE FROM api_keys WHERE id = ");
|
||||
builder
|
||||
.push_bind(api_key_id)
|
||||
.push(" AND is_standalone = ")
|
||||
.push_bind(is_standalone);
|
||||
if let Some(user_id) = user_id {
|
||||
builder.push(" AND user_id = ").push_bind(user_id);
|
||||
}
|
||||
let rows_affected = builder
|
||||
.build()
|
||||
.execute(&self.pool)
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let matching_api_key = if let Some(user_id) = user_id {
|
||||
if require_unlocked {
|
||||
sqlx::query_scalar::<_, String>(
|
||||
"SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 AND is_locked = 0 FOR UPDATE",
|
||||
)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
} else {
|
||||
sqlx::query_scalar::<_, String>(
|
||||
"SELECT id FROM api_keys WHERE id = ? AND user_id = ? AND is_standalone = 0 FOR UPDATE",
|
||||
)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
}
|
||||
} else {
|
||||
sqlx::query_scalar::<_, String>(
|
||||
"SELECT id FROM api_keys WHERE id = ? AND is_standalone = 1 FOR UPDATE",
|
||||
)
|
||||
.bind(api_key_id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected > 0)
|
||||
};
|
||||
if matching_api_key.is_none() {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
sqlx::query(
|
||||
"UPDATE wallets SET status = 'disabled', updated_at = UNIX_TIMESTAMP() WHERE api_key_id = ? AND status <> 'disabled'",
|
||||
)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
for sql in MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL {
|
||||
sqlx::query(sql)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
for sql in MYSQL_DELETE_API_KEY_DEPENDENTS_SQL {
|
||||
sqlx::query(sql)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
let result = sqlx::query("DELETE FROM api_keys WHERE id = ? AND is_standalone = ?")
|
||||
.bind(api_key_id)
|
||||
.bind(is_standalone)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() != 1 {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(true)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -827,6 +1243,7 @@ async fn summarize_api_keys(
|
||||
is_standalone: bool,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<AuthApiKeyExportSummary, DataLayerError> {
|
||||
let now_unix_secs = i64_from_u64(now_unix_secs, "api_keys.summary_now")?;
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -836,7 +1253,7 @@ FROM api_keys
|
||||
WHERE is_standalone = ?
|
||||
"#,
|
||||
)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(now_unix_secs)
|
||||
.bind(is_standalone)
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
@@ -1037,7 +1454,36 @@ fn map_auth_api_key_export_row(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlAuthApiKeyReadRepository;
|
||||
use super::{
|
||||
MysqlAuthApiKeyReadRepository, MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL,
|
||||
MYSQL_DELETE_API_KEY_DEPENDENTS_SQL,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn api_key_delete_sql_preserves_ids_and_removes_private_snapshots() {
|
||||
for table in [
|
||||
"request_candidates",
|
||||
"video_tasks",
|
||||
"`usage`",
|
||||
"stats_daily_api_key",
|
||||
] {
|
||||
assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| {
|
||||
sql.starts_with(&format!("UPDATE {table} "))
|
||||
&& sql.contains("SET api_key_name = NULL")
|
||||
&& sql.ends_with("WHERE api_key_id = ?")
|
||||
}));
|
||||
}
|
||||
assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL
|
||||
.iter()
|
||||
.any(|sql| sql
|
||||
.starts_with("UPDATE audit_logs SET description = 'deleted API key event'")));
|
||||
assert!(MYSQL_ANONYMIZE_API_KEY_HISTORY_SQL.iter().any(|sql| sql
|
||||
.starts_with("UPDATE payment_callbacks SET payload = NULL, error_message = NULL")));
|
||||
assert_eq!(
|
||||
MYSQL_DELETE_API_KEY_DEPENDENTS_SQL,
|
||||
&["DELETE FROM api_key_provider_mappings WHERE api_key_id = ?"]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_builds_from_lazy_pool() {
|
||||
|
||||
@@ -3,7 +3,7 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::auth_modules::*;
|
||||
use aether_data_contracts::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
use aether_data_query::{push_eq, WhereClause};
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
@@ -35,6 +35,87 @@ SELECT
|
||||
FROM ldap_configs
|
||||
"#;
|
||||
|
||||
const UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL: &str = r#"
|
||||
UPDATE ldap_configs
|
||||
SET
|
||||
server_url = ?,
|
||||
bind_dn = ?,
|
||||
base_dn = ?,
|
||||
user_search_filter = ?,
|
||||
username_attr = ?,
|
||||
email_attr = ?,
|
||||
display_name_attr = ?,
|
||||
is_enabled = ?,
|
||||
is_exclusive = ?,
|
||||
use_starttls = ?,
|
||||
connect_timeout = ?,
|
||||
updated_at = GREATEST(updated_at + 1, ?)
|
||||
WHERE singleton_key = 1
|
||||
AND server_url <=> ?
|
||||
AND bind_dn <=> ?
|
||||
AND BINARY bind_password_encrypted <=> BINARY ?
|
||||
AND base_dn <=> ?
|
||||
AND user_search_filter <=> ?
|
||||
AND username_attr <=> ?
|
||||
AND email_attr <=> ?
|
||||
AND display_name_attr <=> ?
|
||||
AND is_enabled <=> ?
|
||||
AND is_exclusive <=> ?
|
||||
AND use_starttls <=> ?
|
||||
AND connect_timeout <=> ?
|
||||
"#;
|
||||
|
||||
const UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL: &str = r#"
|
||||
UPDATE ldap_configs
|
||||
SET
|
||||
server_url = ?,
|
||||
bind_dn = ?,
|
||||
bind_password_encrypted = ?,
|
||||
base_dn = ?,
|
||||
user_search_filter = ?,
|
||||
username_attr = ?,
|
||||
email_attr = ?,
|
||||
display_name_attr = ?,
|
||||
is_enabled = ?,
|
||||
is_exclusive = ?,
|
||||
use_starttls = ?,
|
||||
connect_timeout = ?,
|
||||
updated_at = GREATEST(updated_at + 1, ?)
|
||||
WHERE singleton_key = 1
|
||||
AND server_url <=> ?
|
||||
AND bind_dn <=> ?
|
||||
AND BINARY bind_password_encrypted <=> BINARY ?
|
||||
AND base_dn <=> ?
|
||||
AND user_search_filter <=> ?
|
||||
AND username_attr <=> ?
|
||||
AND email_attr <=> ?
|
||||
AND display_name_attr <=> ?
|
||||
AND is_enabled <=> ?
|
||||
AND is_exclusive <=> ?
|
||||
AND use_starttls <=> ?
|
||||
AND connect_timeout <=> ?
|
||||
"#;
|
||||
|
||||
const INSERT_LDAP_CONFIG_SQL: &str = r#"
|
||||
INSERT INTO ldap_configs (
|
||||
singleton_key,
|
||||
server_url,
|
||||
bind_dn,
|
||||
bind_password_encrypted,
|
||||
base_dn,
|
||||
user_search_filter,
|
||||
username_attr,
|
||||
email_attr,
|
||||
display_name_attr,
|
||||
is_enabled,
|
||||
is_exclusive,
|
||||
use_starttls,
|
||||
connect_timeout,
|
||||
created_at,
|
||||
updated_at
|
||||
) VALUES (1, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlAuthModuleReadRepository {
|
||||
pool: MysqlPool,
|
||||
@@ -72,8 +153,7 @@ async fn get_ldap_config(
|
||||
pool: &MysqlPool,
|
||||
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(LDAP_CONFIG_COLUMNS);
|
||||
builder.push(" ORDER BY id ASC");
|
||||
push_limit(&mut builder, 1);
|
||||
builder.push(" WHERE singleton_key = 1");
|
||||
let row = builder.build().fetch_optional(pool).await.map_sql_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
@@ -106,97 +186,208 @@ impl AuthModuleReadRepository for MysqlAuthModuleRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl AuthModuleWriteRepository for MysqlAuthModuleRepository {
|
||||
async fn upsert_ldap_config(
|
||||
async fn compare_and_swap_ldap_config(
|
||||
&self,
|
||||
config: &StoredLdapModuleConfig,
|
||||
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
expected: Option<&StoredLdapModuleConfig>,
|
||||
replacement: &StoredLdapModuleConfig,
|
||||
bind_password_update: &LdapBindPasswordUpdate,
|
||||
) -> Result<CompareAndSwapLdapConfigResult, DataLayerError> {
|
||||
let persisted =
|
||||
ldap_config_after_password_update(expected, replacement, bind_password_update)?;
|
||||
let now = now_unix_secs();
|
||||
let updated = sqlx::query(
|
||||
let Some(expected) = expected else {
|
||||
let insert = sqlx::query(INSERT_LDAP_CONFIG_SQL)
|
||||
.bind(&persisted.server_url)
|
||||
.bind(&persisted.bind_dn)
|
||||
.bind(persisted.bind_password_encrypted.as_deref())
|
||||
.bind(&persisted.base_dn)
|
||||
.bind(persisted.user_search_filter.as_deref())
|
||||
.bind(persisted.username_attr.as_deref())
|
||||
.bind(persisted.email_attr.as_deref())
|
||||
.bind(persisted.display_name_attr.as_deref())
|
||||
.bind(persisted.is_enabled)
|
||||
.bind(persisted.is_exclusive)
|
||||
.bind(persisted.use_starttls)
|
||||
.bind(persisted.connect_timeout)
|
||||
.bind(now as i64)
|
||||
.bind(now as i64)
|
||||
.execute(&self.pool)
|
||||
.await;
|
||||
return match insert {
|
||||
Ok(result) if result.rows_affected() == 1 => {
|
||||
Ok(CompareAndSwapLdapConfigResult::Applied(persisted))
|
||||
}
|
||||
Ok(_) => Ok(CompareAndSwapLdapConfigResult::Conflict),
|
||||
Err(error)
|
||||
if error
|
||||
.as_database_error()
|
||||
.is_some_and(|error| error.is_unique_violation()) =>
|
||||
{
|
||||
Ok(CompareAndSwapLdapConfigResult::Conflict)
|
||||
}
|
||||
Err(error) => Err(DataLayerError::sql(error)),
|
||||
};
|
||||
};
|
||||
|
||||
let rows_affected = match bind_password_update {
|
||||
LdapBindPasswordUpdate::Preserve => {
|
||||
sqlx::query(UPDATE_LDAP_CONFIG_PRESERVE_PASSWORD_SQL)
|
||||
.bind(&replacement.server_url)
|
||||
.bind(&replacement.bind_dn)
|
||||
.bind(&replacement.base_dn)
|
||||
.bind(replacement.user_search_filter.as_deref())
|
||||
.bind(replacement.username_attr.as_deref())
|
||||
.bind(replacement.email_attr.as_deref())
|
||||
.bind(replacement.display_name_attr.as_deref())
|
||||
.bind(replacement.is_enabled)
|
||||
.bind(replacement.is_exclusive)
|
||||
.bind(replacement.use_starttls)
|
||||
.bind(replacement.connect_timeout)
|
||||
.bind(now as i64)
|
||||
.bind(&expected.server_url)
|
||||
.bind(&expected.bind_dn)
|
||||
.bind(expected.bind_password_encrypted.as_deref())
|
||||
.bind(&expected.base_dn)
|
||||
.bind(expected.user_search_filter.as_deref())
|
||||
.bind(expected.username_attr.as_deref())
|
||||
.bind(expected.email_attr.as_deref())
|
||||
.bind(expected.display_name_attr.as_deref())
|
||||
.bind(expected.is_enabled)
|
||||
.bind(expected.is_exclusive)
|
||||
.bind(expected.use_starttls)
|
||||
.bind(expected.connect_timeout)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected()
|
||||
}
|
||||
LdapBindPasswordUpdate::Set(_) | LdapBindPasswordUpdate::Clear => {
|
||||
sqlx::query(UPDATE_LDAP_CONFIG_REPLACE_PASSWORD_SQL)
|
||||
.bind(&replacement.server_url)
|
||||
.bind(&replacement.bind_dn)
|
||||
.bind(persisted.bind_password_encrypted.as_deref())
|
||||
.bind(&replacement.base_dn)
|
||||
.bind(replacement.user_search_filter.as_deref())
|
||||
.bind(replacement.username_attr.as_deref())
|
||||
.bind(replacement.email_attr.as_deref())
|
||||
.bind(replacement.display_name_attr.as_deref())
|
||||
.bind(replacement.is_enabled)
|
||||
.bind(replacement.is_exclusive)
|
||||
.bind(replacement.use_starttls)
|
||||
.bind(replacement.connect_timeout)
|
||||
.bind(now as i64)
|
||||
.bind(&expected.server_url)
|
||||
.bind(&expected.bind_dn)
|
||||
.bind(expected.bind_password_encrypted.as_deref())
|
||||
.bind(&expected.base_dn)
|
||||
.bind(expected.user_search_filter.as_deref())
|
||||
.bind(expected.username_attr.as_deref())
|
||||
.bind(expected.email_attr.as_deref())
|
||||
.bind(expected.display_name_attr.as_deref())
|
||||
.bind(expected.is_enabled)
|
||||
.bind(expected.is_exclusive)
|
||||
.bind(expected.use_starttls)
|
||||
.bind(expected.connect_timeout)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected()
|
||||
}
|
||||
};
|
||||
if rows_affected == 1 {
|
||||
Ok(CompareAndSwapLdapConfigResult::Applied(persisted))
|
||||
} else {
|
||||
Ok(CompareAndSwapLdapConfigResult::Conflict)
|
||||
}
|
||||
}
|
||||
|
||||
async fn delete_ldap_config_if_matches(
|
||||
&self,
|
||||
expected: &StoredLdapModuleConfig,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let rows_affected = sqlx::query(
|
||||
r#"
|
||||
UPDATE ldap_configs
|
||||
SET
|
||||
server_url = ?,
|
||||
bind_dn = ?,
|
||||
bind_password_encrypted = ?,
|
||||
base_dn = ?,
|
||||
user_search_filter = ?,
|
||||
username_attr = ?,
|
||||
email_attr = ?,
|
||||
display_name_attr = ?,
|
||||
is_enabled = ?,
|
||||
is_exclusive = ?,
|
||||
use_starttls = ?,
|
||||
connect_timeout = ?,
|
||||
updated_at = ?
|
||||
WHERE id = (
|
||||
SELECT id FROM (
|
||||
SELECT id
|
||||
FROM ldap_configs
|
||||
ORDER BY id ASC
|
||||
LIMIT 1
|
||||
) selected_ldap_config
|
||||
)
|
||||
DELETE FROM ldap_configs
|
||||
WHERE singleton_key = 1
|
||||
AND server_url <=> ?
|
||||
AND bind_dn <=> ?
|
||||
AND BINARY bind_password_encrypted <=> BINARY ?
|
||||
AND base_dn <=> ?
|
||||
AND user_search_filter <=> ?
|
||||
AND username_attr <=> ?
|
||||
AND email_attr <=> ?
|
||||
AND display_name_attr <=> ?
|
||||
AND is_enabled <=> ?
|
||||
AND is_exclusive <=> ?
|
||||
AND use_starttls <=> ?
|
||||
AND connect_timeout <=> ?
|
||||
"#,
|
||||
)
|
||||
.bind(&config.server_url)
|
||||
.bind(&config.bind_dn)
|
||||
.bind(config.bind_password_encrypted.as_deref())
|
||||
.bind(&config.base_dn)
|
||||
.bind(config.user_search_filter.as_deref())
|
||||
.bind(config.username_attr.as_deref())
|
||||
.bind(config.email_attr.as_deref())
|
||||
.bind(config.display_name_attr.as_deref())
|
||||
.bind(config.is_enabled)
|
||||
.bind(config.is_exclusive)
|
||||
.bind(config.use_starttls)
|
||||
.bind(config.connect_timeout)
|
||||
.bind(now as i64)
|
||||
.bind(&expected.server_url)
|
||||
.bind(&expected.bind_dn)
|
||||
.bind(expected.bind_password_encrypted.as_deref())
|
||||
.bind(&expected.base_dn)
|
||||
.bind(expected.user_search_filter.as_deref())
|
||||
.bind(expected.username_attr.as_deref())
|
||||
.bind(expected.email_attr.as_deref())
|
||||
.bind(expected.display_name_attr.as_deref())
|
||||
.bind(expected.is_enabled)
|
||||
.bind(expected.is_exclusive)
|
||||
.bind(expected.use_starttls)
|
||||
.bind(expected.connect_timeout)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
if updated.rows_affected() == 0 {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO ldap_configs (
|
||||
server_url,
|
||||
bind_dn,
|
||||
bind_password_encrypted,
|
||||
base_dn,
|
||||
user_search_filter,
|
||||
username_attr,
|
||||
email_attr,
|
||||
display_name_attr,
|
||||
is_enabled,
|
||||
is_exclusive,
|
||||
use_starttls,
|
||||
connect_timeout,
|
||||
created_at,
|
||||
updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&config.server_url)
|
||||
.bind(&config.bind_dn)
|
||||
.bind(config.bind_password_encrypted.as_deref())
|
||||
.bind(&config.base_dn)
|
||||
.bind(config.user_search_filter.as_deref())
|
||||
.bind(config.username_attr.as_deref())
|
||||
.bind(config.email_attr.as_deref())
|
||||
.bind(config.display_name_attr.as_deref())
|
||||
.bind(config.is_enabled)
|
||||
.bind(config.is_exclusive)
|
||||
.bind(config.use_starttls)
|
||||
.bind(config.connect_timeout)
|
||||
.bind(now as i64)
|
||||
.bind(now as i64)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
|
||||
self.get_ldap_config().await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected == 1)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_ldap_bind_password(
|
||||
&self,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let rows_affected = sqlx::query(
|
||||
r#"
|
||||
UPDATE ldap_configs
|
||||
SET bind_password_encrypted = ?, updated_at = GREATEST(updated_at + 1, ?)
|
||||
WHERE singleton_key = 1
|
||||
AND BINARY bind_password_encrypted = BINARY ?
|
||||
"#,
|
||||
)
|
||||
.bind(replacement)
|
||||
.bind(now_unix_secs() as i64)
|
||||
.bind(expected)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected == 1)
|
||||
}
|
||||
}
|
||||
|
||||
fn ldap_config_after_password_update(
|
||||
expected: Option<&StoredLdapModuleConfig>,
|
||||
replacement: &StoredLdapModuleConfig,
|
||||
bind_password_update: &LdapBindPasswordUpdate,
|
||||
) -> Result<StoredLdapModuleConfig, DataLayerError> {
|
||||
let bind_password_encrypted = match bind_password_update {
|
||||
LdapBindPasswordUpdate::Preserve => expected
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidConfiguration(
|
||||
"LDAP bind password cannot be preserved while creating the singleton"
|
||||
.to_string(),
|
||||
)
|
||||
})?
|
||||
.bind_password_encrypted
|
||||
.clone(),
|
||||
LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()),
|
||||
LdapBindPasswordUpdate::Clear => None,
|
||||
};
|
||||
Ok(StoredLdapModuleConfig {
|
||||
bind_password_encrypted,
|
||||
..replacement.clone()
|
||||
})
|
||||
}
|
||||
|
||||
fn now_unix_secs() -> u64 {
|
||||
|
||||
@@ -209,8 +209,9 @@ impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
|
||||
impl BackgroundTaskWriteRepository for MysqlBackgroundTaskRepository {
|
||||
async fn upsert_run(
|
||||
&self,
|
||||
run: UpsertBackgroundTaskRun,
|
||||
mut run: UpsertBackgroundTaskRun,
|
||||
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
|
||||
run.sanitize_for_persistence();
|
||||
run.validate()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
@@ -309,8 +310,9 @@ ON DUPLICATE KEY UPDATE
|
||||
|
||||
async fn upsert_event(
|
||||
&self,
|
||||
event: UpsertBackgroundTaskEvent,
|
||||
mut event: UpsertBackgroundTaskEvent,
|
||||
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
|
||||
event.sanitize_for_persistence();
|
||||
event.validate()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
@@ -358,7 +360,7 @@ fn map_run_row(row: &MySqlRow) -> Result<StoredBackgroundTaskRun, DataLayerError
|
||||
let finished_at_unix_secs: Option<i64> = row.try_get("finished_at_unix_secs").map_sql_err()?;
|
||||
let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?;
|
||||
|
||||
Ok(StoredBackgroundTaskRun {
|
||||
let mut run = StoredBackgroundTaskRun {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
task_key: row.try_get("task_key").map_sql_err()?,
|
||||
kind: BackgroundTaskKind::from_database(&kind)?,
|
||||
@@ -381,12 +383,14 @@ fn map_run_row(row: &MySqlRow) -> Result<StoredBackgroundTaskRun, DataLayerError
|
||||
started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
|
||||
finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
|
||||
updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(),
|
||||
})
|
||||
};
|
||||
run.sanitize_persisted_data();
|
||||
Ok(run)
|
||||
}
|
||||
|
||||
fn map_event_row(row: &MySqlRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
|
||||
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
|
||||
Ok(StoredBackgroundTaskEvent {
|
||||
let mut event = StoredBackgroundTaskEvent {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
run_id: row.try_get("run_id").map_sql_err()?,
|
||||
event_type: row.try_get("event_type").map_sql_err()?,
|
||||
@@ -396,7 +400,9 @@ fn map_event_row(row: &MySqlRow) -> Result<StoredBackgroundTaskEvent, DataLayerE
|
||||
"payload_json",
|
||||
)?,
|
||||
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
|
||||
})
|
||||
};
|
||||
event.sanitize_persisted_data();
|
||||
Ok(event)
|
||||
}
|
||||
|
||||
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
|
||||
|
||||
@@ -4,8 +4,9 @@ use sqlx::{mysql::MySqlRow, Row};
|
||||
use aether_data_contracts::repository::billing::{
|
||||
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
|
||||
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
|
||||
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord,
|
||||
PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
|
||||
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
|
||||
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
|
||||
PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
|
||||
UserPlanEntitlementRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
@@ -637,6 +638,157 @@ LIMIT 1
|
||||
.transpose()
|
||||
}
|
||||
|
||||
async fn compare_and_swap_payment_gateway_secret(
|
||||
&self,
|
||||
update: &PaymentGatewaySecretCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE payment_gateway_configs
|
||||
SET merchant_key_encrypted = ?
|
||||
WHERE provider = ?
|
||||
AND BINARY merchant_key_encrypted = BINARY ?
|
||||
"#,
|
||||
)
|
||||
.bind(&update.merchant_key_encrypted)
|
||||
.bind(update.provider.trim().to_ascii_lowercase())
|
||||
.bind(&update.expected_merchant_key_encrypted)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() == 1)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_payment_gateway_config(
|
||||
&self,
|
||||
mutation: &PaymentGatewayConfigCasWriteInput,
|
||||
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, DataLayerError> {
|
||||
let input = &mutation.input;
|
||||
let provider = input.provider.trim().to_ascii_lowercase();
|
||||
let now = current_unix_secs_i64();
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
|
||||
if mutation.expected_existing {
|
||||
let current = sqlx::query(
|
||||
r#"
|
||||
SELECT merchant_key_encrypted
|
||||
FROM payment_gateway_configs
|
||||
WHERE provider = ?
|
||||
LIMIT 1
|
||||
FOR UPDATE
|
||||
"#,
|
||||
)
|
||||
.bind(&provider)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let current_secret = match current.as_ref() {
|
||||
Some(row) => row
|
||||
.try_get::<Option<String>, _>("merchant_key_encrypted")
|
||||
.map_sql_err()?,
|
||||
None => {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(AdminBillingMutationOutcome::NotFound);
|
||||
}
|
||||
};
|
||||
if current_secret != mutation.expected_merchant_key_encrypted {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(AdminBillingMutationOutcome::NotFound);
|
||||
}
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE payment_gateway_configs
|
||||
SET
|
||||
enabled = ?,
|
||||
endpoint_url = ?,
|
||||
callback_base_url = ?,
|
||||
merchant_id = ?,
|
||||
merchant_key_encrypted = CASE
|
||||
WHEN ? THEN merchant_key_encrypted
|
||||
ELSE ?
|
||||
END,
|
||||
pay_currency = ?,
|
||||
usd_exchange_rate = ?,
|
||||
min_recharge_usd = ?,
|
||||
channels_json = ?,
|
||||
updated_at = ?
|
||||
WHERE provider = ?
|
||||
"#,
|
||||
)
|
||||
.bind(input.enabled)
|
||||
.bind(&input.endpoint_url)
|
||||
.bind(input.callback_base_url.as_deref())
|
||||
.bind(&input.merchant_id)
|
||||
.bind(input.preserve_existing_secret)
|
||||
.bind(input.merchant_key_encrypted.as_deref())
|
||||
.bind(&input.pay_currency)
|
||||
.bind(input.usd_exchange_rate)
|
||||
.bind(input.min_recharge_usd)
|
||||
.bind(json_to_string(&input.channels_json)?)
|
||||
.bind(now)
|
||||
.bind(&provider)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
} else {
|
||||
let inserted = sqlx::query(
|
||||
r#"
|
||||
INSERT INTO payment_gateway_configs (
|
||||
provider, enabled, endpoint_url, callback_base_url, merchant_id,
|
||||
merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd,
|
||||
channels_json, created_at, updated_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&provider)
|
||||
.bind(input.enabled)
|
||||
.bind(&input.endpoint_url)
|
||||
.bind(input.callback_base_url.as_deref())
|
||||
.bind(&input.merchant_id)
|
||||
.bind(input.merchant_key_encrypted.as_deref())
|
||||
.bind(&input.pay_currency)
|
||||
.bind(input.usd_exchange_rate)
|
||||
.bind(input.min_recharge_usd)
|
||||
.bind(json_to_string(&input.channels_json)?)
|
||||
.bind(now)
|
||||
.bind(now)
|
||||
.execute(&mut *tx)
|
||||
.await;
|
||||
if let Err(err) = inserted {
|
||||
let unique = matches!(
|
||||
&err,
|
||||
sqlx::Error::Database(database_error) if database_error.is_unique_violation()
|
||||
);
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
if unique {
|
||||
return Ok(AdminBillingMutationOutcome::NotFound);
|
||||
}
|
||||
return Err(DataLayerError::sql(err));
|
||||
}
|
||||
}
|
||||
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
provider, enabled, endpoint_url, callback_base_url, merchant_id,
|
||||
merchant_key_encrypted, pay_currency, usd_exchange_rate, min_recharge_usd,
|
||||
channels_json, created_at AS created_at_unix_secs, updated_at AS updated_at_unix_secs
|
||||
FROM payment_gateway_configs
|
||||
WHERE provider = ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(&provider)
|
||||
.fetch_one(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let record = map_payment_gateway_config_mysql(&row)?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(AdminBillingMutationOutcome::Applied(record))
|
||||
}
|
||||
|
||||
async fn upsert_payment_gateway_config(
|
||||
&self,
|
||||
input: &PaymentGatewayConfigWriteInput,
|
||||
|
||||
@@ -832,12 +832,12 @@ fn map_candidate_selection_row(row: &MySqlRow) -> Result<CandidateSelectionRow,
|
||||
key_name: row.try_get("key_name").map_sql_err()?,
|
||||
key_auth_type: row.try_get("key_auth_type").map_sql_err()?,
|
||||
key_is_active: row.try_get("key_is_active").map_sql_err()?,
|
||||
key_api_formats: parse_string_list(
|
||||
parse_json(row.try_get("key_api_formats").ok().flatten())?,
|
||||
key_api_formats: parse_stored_key_policy_string_list(
|
||||
row.try_get("key_api_formats").map_sql_err()?,
|
||||
"provider_api_keys.api_formats",
|
||||
)?,
|
||||
key_allowed_models: parse_string_list(
|
||||
parse_json(row.try_get("key_allowed_models").ok().flatten())?,
|
||||
key_allowed_models: parse_stored_key_policy_string_list(
|
||||
row.try_get("key_allowed_models").map_sql_err()?,
|
||||
"provider_api_keys.allowed_models",
|
||||
)?,
|
||||
key_capabilities: parse_json(row.try_get("key_capabilities").ok().flatten())?,
|
||||
@@ -904,6 +904,82 @@ fn parse_string_list(
|
||||
parse_string_list_value(&value, field_name)
|
||||
}
|
||||
|
||||
fn parse_stored_key_policy_string_list(
|
||||
raw: Option<String>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let Some(raw) = raw else {
|
||||
return Ok(None);
|
||||
};
|
||||
let value = serde_json::from_str::<serde_json::Value>(&raw).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("{field_name} contains invalid JSON: {err}"))
|
||||
})?;
|
||||
parse_key_policy_string_list_value(&value, field_name)
|
||||
}
|
||||
|
||||
fn parse_key_policy_string_list_value(
|
||||
value: &serde_json::Value,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
match value {
|
||||
serde_json::Value::Null => Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains JSON null; use SQL NULL for an unset policy"
|
||||
))),
|
||||
serde_json::Value::Array(array) => {
|
||||
parse_key_policy_string_list_array(array, field_name).map(Some)
|
||||
}
|
||||
serde_json::Value::String(raw) => parse_embedded_key_policy_string_list(raw, field_name),
|
||||
_ => Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} is not a JSON array"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_embedded_key_policy_string_list(
|
||||
raw: &str,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains an empty string"
|
||||
)));
|
||||
}
|
||||
if raw.eq_ignore_ascii_case("null") {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains stringified JSON null; use SQL NULL for an unset policy"
|
||||
)));
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
return parse_key_policy_string_list_value(&decoded, field_name);
|
||||
}
|
||||
|
||||
Ok(Some(vec![raw.to_string()]))
|
||||
}
|
||||
|
||||
fn parse_key_policy_string_list_array(
|
||||
array: &[serde_json::Value],
|
||||
field_name: &str,
|
||||
) -> Result<Vec<String>, DataLayerError> {
|
||||
let mut items = Vec::with_capacity(array.len());
|
||||
for item in array {
|
||||
let Some(item) = item.as_str() else {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains a non-string item"
|
||||
)));
|
||||
};
|
||||
let item = item.trim();
|
||||
if item.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains an empty item"
|
||||
)));
|
||||
}
|
||||
items.push(item.to_string());
|
||||
}
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
fn parse_string_list_value(
|
||||
value: &serde_json::Value,
|
||||
field_name: &str,
|
||||
@@ -1108,7 +1184,8 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
api_format_page_query, pool_key_group_by_key_ids_query, pool_key_group_query,
|
||||
api_format_page_query, parse_stored_key_policy_string_list,
|
||||
pool_key_group_by_key_ids_query, pool_key_group_query,
|
||||
provider_model_mapping_api_format_covers, push_key_auth_channel_filter,
|
||||
requested_model_page_query, vertex_key_auth_channel_matches, ExactPageAccumulator,
|
||||
MysqlMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT,
|
||||
@@ -1132,6 +1209,25 @@ mod tests {
|
||||
assert!(sql.contains("LIMIT ? OFFSET ?"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_key_policy_never_degrades_to_unrestricted() {
|
||||
for raw in ["null", "\"null\"", "\"\"", "[\"openai:chat\",null]"] {
|
||||
assert!(parse_stored_key_policy_string_list(
|
||||
Some(raw.to_string()),
|
||||
"provider_api_keys.api_formats",
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
assert_eq!(
|
||||
parse_stored_key_policy_string_list(
|
||||
Some("[\"openai:chat\"]".to_string()),
|
||||
"provider_api_keys.api_formats",
|
||||
)
|
||||
.expect("valid key policy should parse"),
|
||||
Some(vec!["openai:chat".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn requested_model_page_query_filters_and_pages_before_fetch() {
|
||||
let query = requested_model_page_query("openai:chat", "gpt-5", 256, 256);
|
||||
|
||||
@@ -216,8 +216,9 @@ impl RequestCandidateReadRepository for MysqlRequestCandidateRepository {
|
||||
impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository {
|
||||
async fn upsert(
|
||||
&self,
|
||||
candidate: UpsertRequestCandidateRecord,
|
||||
mut candidate: UpsertRequestCandidateRecord,
|
||||
) -> Result<StoredRequestCandidate, DataLayerError> {
|
||||
candidate.sanitize_for_persistence();
|
||||
candidate.validate()?;
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
match upsert_candidate_in_transaction(&mut tx, candidate).await {
|
||||
@@ -234,12 +235,13 @@ impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository {
|
||||
|
||||
async fn upsert_many(
|
||||
&self,
|
||||
candidates: Vec<UpsertRequestCandidateRecord>,
|
||||
mut candidates: Vec<UpsertRequestCandidateRecord>,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if candidates.is_empty() {
|
||||
return Ok(0);
|
||||
}
|
||||
for candidate in &candidates {
|
||||
for candidate in &mut candidates {
|
||||
candidate.sanitize_for_persistence();
|
||||
candidate.validate()?;
|
||||
}
|
||||
|
||||
@@ -423,26 +425,8 @@ ON DUPLICATE KEY UPDATE
|
||||
THEN status_code
|
||||
ELSE COALESCE(VALUES(status_code), status_code)
|
||||
END,
|
||||
error_type = CASE
|
||||
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
|
||||
THEN error_type
|
||||
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
|
||||
THEN error_type
|
||||
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
|
||||
THEN error_type
|
||||
ELSE COALESCE(VALUES(error_type), error_type)
|
||||
END,
|
||||
error_message = CASE
|
||||
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
|
||||
THEN error_message
|
||||
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
|
||||
THEN error_message
|
||||
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
|
||||
THEN error_message
|
||||
ELSE COALESCE(VALUES(error_message), error_message)
|
||||
END,
|
||||
error_type = VALUES(error_type),
|
||||
error_message = NULL,
|
||||
latency_ms = CASE
|
||||
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
|
||||
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
|
||||
@@ -524,9 +508,10 @@ fn push_endpoint_in_clause<'args>(
|
||||
}
|
||||
|
||||
fn merge_candidate(
|
||||
candidate: UpsertRequestCandidateRecord,
|
||||
mut candidate: UpsertRequestCandidateRecord,
|
||||
existing: Option<StoredRequestCandidate>,
|
||||
) -> Result<StoredRequestCandidate, DataLayerError> {
|
||||
candidate.sanitize_for_persistence();
|
||||
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|value| {
|
||||
request_candidate_lifecycle_would_regress(value.status, candidate.status)
|
||||
});
|
||||
@@ -538,15 +523,11 @@ fn merge_candidate(
|
||||
} else {
|
||||
candidate.status
|
||||
};
|
||||
let created_at_unix_ms = candidate
|
||||
.created_at_unix_ms
|
||||
let created_at_unix_ms = existing
|
||||
.as_ref()
|
||||
.map(|value| value.created_at_unix_ms)
|
||||
.filter(|value| *value > 1000)
|
||||
.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.map(|value| value.created_at_unix_ms)
|
||||
.filter(|value| *value > 1000)
|
||||
})
|
||||
.or_else(|| candidate.created_at_unix_ms.filter(|value| *value > 1000))
|
||||
.or(candidate.started_at_unix_ms)
|
||||
.or(candidate.finished_at_unix_ms)
|
||||
.unwrap_or_else(current_unix_ms);
|
||||
@@ -561,35 +542,36 @@ fn merge_candidate(
|
||||
StoredRequestCandidate::new(
|
||||
id,
|
||||
candidate.request_id,
|
||||
candidate
|
||||
.user_id
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.user_id.clone())),
|
||||
candidate
|
||||
.api_key_id
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.api_key_id.clone())),
|
||||
candidate
|
||||
.username
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.username.clone())),
|
||||
candidate.api_key_name.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.api_key_name.clone())
|
||||
}),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.user_id.clone())
|
||||
.or(candidate.user_id),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.api_key_id.clone())
|
||||
.or(candidate.api_key_id),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.username.clone())
|
||||
.or(candidate.username),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.api_key_name.clone())
|
||||
.or(candidate.api_key_name),
|
||||
to_i32(candidate.candidate_index)?,
|
||||
to_i32(candidate.retry_index)?,
|
||||
candidate.provider_id.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.provider_id.clone())
|
||||
}),
|
||||
candidate.endpoint_id.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.endpoint_id.clone())
|
||||
}),
|
||||
candidate
|
||||
.key_id
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.key_id.clone())),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.provider_id.clone())
|
||||
.or(candidate.provider_id),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.endpoint_id.clone())
|
||||
.or(candidate.endpoint_id),
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.key_id.clone())
|
||||
.or(candidate.key_id),
|
||||
merged_status,
|
||||
candidate.skip_reason.or_else(|| {
|
||||
existing
|
||||
@@ -617,17 +599,7 @@ fn merge_candidate(
|
||||
.error_type
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.error_type.clone()))
|
||||
},
|
||||
if preserve_existing_lifecycle {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.error_message.clone())
|
||||
} else {
|
||||
candidate.error_message.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.error_message.clone())
|
||||
})
|
||||
},
|
||||
None,
|
||||
if preserve_existing_lifecycle {
|
||||
match existing.as_ref().and_then(|value| value.latency_ms) {
|
||||
Some(value) => Some(to_i32_u64(value)?),
|
||||
@@ -657,9 +629,10 @@ fn merge_candidate(
|
||||
.and_then(|value| value.required_capabilities.clone())
|
||||
}),
|
||||
u64_to_i64(created_at_unix_ms, "request candidate created_at")?,
|
||||
candidate
|
||||
.started_at_unix_ms
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.started_at_unix_ms))
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.started_at_unix_ms)
|
||||
.or(candidate.started_at_unix_ms)
|
||||
.map(|value| u64_to_i64(value, "request candidate started_at"))
|
||||
.transpose()?,
|
||||
if preserve_existing_lifecycle {
|
||||
@@ -915,7 +888,7 @@ mod tests {
|
||||
&request_id,
|
||||
"initial",
|
||||
RequestCandidateStatus::Pending,
|
||||
Some(json!({"initial": true})),
|
||||
Some(json!({"gateway_execution_runtime": true})),
|
||||
3_000_000,
|
||||
);
|
||||
initial.is_cached = Some(false);
|
||||
@@ -923,6 +896,18 @@ mod tests {
|
||||
.upsert(initial)
|
||||
.await
|
||||
.expect("initial candidate should insert");
|
||||
sqlx::query(
|
||||
"UPDATE request_candidates SET skip_reason = ?, error_type = ?, error_message = ?, extra_data = ?, required_capabilities = ? WHERE request_id = ?",
|
||||
)
|
||||
.bind("legacy skip reason with tenant-secret")
|
||||
.bind("legacy_error_type_with_token")
|
||||
.bind("Bearer legacy-secret")
|
||||
.bind(r#"{"gateway_execution_runtime":true,"request_body":{"password":"secret"}}"#)
|
||||
.bind(r#"{"streaming":true,"internal_capability":"secret"}"#)
|
||||
.bind(&request_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("legacy diagnostics should be injected for the conflict test");
|
||||
|
||||
const WRITERS: usize = 8;
|
||||
let barrier = std::sync::Arc::new(tokio::sync::Barrier::new(WRITERS));
|
||||
@@ -937,13 +922,22 @@ mod tests {
|
||||
} else {
|
||||
RequestCandidateStatus::Streaming
|
||||
};
|
||||
let mut extra_data = serde_json::Map::new();
|
||||
extra_data.insert(format!("writer_{writer}"), json!(writer));
|
||||
let extra_data = match writer {
|
||||
0 => json!({"stream_completed": true}),
|
||||
1 => json!({"cache_1h": true}),
|
||||
2 => json!({"first_byte_time_ms": 2}),
|
||||
3 => json!({"pool_key_index": 3}),
|
||||
4 => json!({"priority_slot": 4}),
|
||||
5 => json!({"ranking_index": 5}),
|
||||
6 => json!({"phase": "provider_request"}),
|
||||
7 => json!({"provider_api_format": "openai:responses"}),
|
||||
_ => unreachable!("writer index is bounded by WRITERS"),
|
||||
};
|
||||
let mut candidate = sample_upsert(
|
||||
&request_id,
|
||||
format!("writer-{writer}").as_str(),
|
||||
status,
|
||||
Some(serde_json::Value::Object(extra_data)),
|
||||
Some(extra_data),
|
||||
3_100_000 + u64::try_from(writer).expect("writer index should fit") * 10,
|
||||
);
|
||||
if writer != 0 {
|
||||
@@ -970,25 +964,60 @@ mod tests {
|
||||
assert_eq!(candidate.status, RequestCandidateStatus::Success);
|
||||
assert_eq!(candidate.latency_ms, Some(123));
|
||||
assert_eq!(candidate.finished_at_unix_ms, Some(3_100_002));
|
||||
let extra_data = candidate
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.expect("merged extra data should be an object");
|
||||
assert_eq!(extra_data.get("initial"), Some(&json!(true)));
|
||||
for writer in 0..WRITERS {
|
||||
assert_eq!(
|
||||
extra_data.get(format!("writer_{writer}").as_str()),
|
||||
Some(&json!(writer))
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
candidate.extra_data,
|
||||
Some(json!({
|
||||
"cache_1h": true,
|
||||
"first_byte_time_ms": 2,
|
||||
"gateway_execution_runtime": true,
|
||||
"phase": "provider_request",
|
||||
"pool_key_index": 3,
|
||||
"priority_slot": 4,
|
||||
"provider_api_format": "openai:responses",
|
||||
"ranking_index": 5,
|
||||
"stream_completed": true
|
||||
}))
|
||||
);
|
||||
let raw = sqlx::query(
|
||||
"SELECT skip_reason, error_type, error_message, extra_data, required_capabilities FROM request_candidates WHERE request_id = ?",
|
||||
)
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("raw candidate diagnostics should load");
|
||||
assert!(
|
||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_message")
|
||||
.expect("error_message should decode")
|
||||
.is_none()
|
||||
);
|
||||
assert_eq!(
|
||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "skip_reason")
|
||||
.expect("skip_reason should decode")
|
||||
.as_deref(),
|
||||
Some("unclassified_skip")
|
||||
);
|
||||
assert_eq!(
|
||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "error_type")
|
||||
.expect("error_type should decode")
|
||||
.as_deref(),
|
||||
Some("unclassified_error")
|
||||
);
|
||||
let raw_extra = sqlx::Row::try_get::<Option<String>, _>(&raw, "extra_data")
|
||||
.expect("extra_data should decode")
|
||||
.and_then(|value| serde_json::from_str::<serde_json::Value>(&value).ok());
|
||||
assert_eq!(raw_extra, candidate.extra_data);
|
||||
let raw_capabilities =
|
||||
sqlx::Row::try_get::<Option<String>, _>(&raw, "required_capabilities")
|
||||
.expect("required_capabilities should decode")
|
||||
.and_then(|value| serde_json::from_str::<serde_json::Value>(&value).ok());
|
||||
assert_eq!(raw_capabilities, Some(json!({"streaming": true})));
|
||||
|
||||
let batch_request_id = format!("candidate-batch-{}", uuid::Uuid::new_v4());
|
||||
let mut pending = sample_upsert(
|
||||
&batch_request_id,
|
||||
"batch-first",
|
||||
RequestCandidateStatus::Pending,
|
||||
Some(json!({"pending": true})),
|
||||
Some(json!({"gateway_execution_runtime": true})),
|
||||
4_000_000,
|
||||
);
|
||||
pending.is_cached = Some(false);
|
||||
@@ -996,7 +1025,7 @@ mod tests {
|
||||
&batch_request_id,
|
||||
"batch-second",
|
||||
RequestCandidateStatus::Streaming,
|
||||
Some(json!({"streaming": true})),
|
||||
Some(json!({"stream_completed": true})),
|
||||
4_000_100,
|
||||
);
|
||||
streaming.is_cached = None;
|
||||
@@ -1004,7 +1033,7 @@ mod tests {
|
||||
&batch_request_id,
|
||||
"batch-third",
|
||||
RequestCandidateStatus::Success,
|
||||
Some(json!({"success": true})),
|
||||
Some(json!({"cache_1h": true})),
|
||||
4_000_200,
|
||||
);
|
||||
success.is_cached = Some(true);
|
||||
@@ -1012,7 +1041,7 @@ mod tests {
|
||||
&batch_request_id,
|
||||
"batch-fourth",
|
||||
RequestCandidateStatus::Pending,
|
||||
Some(json!({"late": true})),
|
||||
Some(json!({"first_byte_time_ms": 42})),
|
||||
4_000_300,
|
||||
);
|
||||
late_pending.is_cached = None;
|
||||
@@ -1038,10 +1067,10 @@ mod tests {
|
||||
assert_eq!(
|
||||
batch_candidates[0].extra_data,
|
||||
Some(json!({
|
||||
"pending": true,
|
||||
"streaming": true,
|
||||
"success": true,
|
||||
"late": true
|
||||
"cache_1h": true,
|
||||
"first_byte_time_ms": 42,
|
||||
"gateway_execution_runtime": true,
|
||||
"stream_completed": true
|
||||
}))
|
||||
);
|
||||
|
||||
@@ -1082,7 +1111,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_candidate_keeps_terminal_status_when_streaming_arrives_late() {
|
||||
fn merge_candidate_preserves_first_identity_and_terminal_fact() {
|
||||
let existing = StoredRequestCandidate::new(
|
||||
"candidate-1".to_string(),
|
||||
"request-1".to_string(),
|
||||
@@ -1103,7 +1132,7 @@ mod tests {
|
||||
None,
|
||||
Some(123),
|
||||
None,
|
||||
Some(serde_json::json!({"terminal": true})),
|
||||
Some(serde_json::json!({"stream_completed": true})),
|
||||
None,
|
||||
1_000,
|
||||
Some(1_001),
|
||||
@@ -1115,24 +1144,24 @@ mod tests {
|
||||
UpsertRequestCandidateRecord {
|
||||
id: "candidate-late".to_string(),
|
||||
request_id: "request-1".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
api_key_id: Some("key-1".to_string()),
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
user_id: Some("attacker-user".to_string()),
|
||||
api_key_id: Some("attacker-api-key".to_string()),
|
||||
username: Some("mallory".to_string()),
|
||||
api_key_name: Some("attacker-key".to_string()),
|
||||
candidate_index: 0,
|
||||
retry_index: 0,
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("provider-key-1".to_string()),
|
||||
status: RequestCandidateStatus::Streaming,
|
||||
provider_id: Some("attacker-provider".to_string()),
|
||||
endpoint_id: Some("attacker-endpoint".to_string()),
|
||||
key_id: Some("attacker-provider-key".to_string()),
|
||||
status: RequestCandidateStatus::Failed,
|
||||
skip_reason: None,
|
||||
is_cached: Some(false),
|
||||
status_code: Some(200),
|
||||
error_type: None,
|
||||
error_message: None,
|
||||
error_message: Some("Bearer secret-token".to_string()),
|
||||
latency_ms: Some(9_999),
|
||||
concurrent_requests: None,
|
||||
extra_data: Some(serde_json::json!({"late": true})),
|
||||
extra_data: Some(serde_json::json!({"gateway_execution_runtime": true})),
|
||||
required_capabilities: None,
|
||||
created_at_unix_ms: Some(1_050),
|
||||
started_at_unix_ms: Some(1_051),
|
||||
@@ -1144,11 +1173,20 @@ mod tests {
|
||||
|
||||
assert_eq!(merged.id, "candidate-1");
|
||||
assert_eq!(merged.status, RequestCandidateStatus::Success);
|
||||
assert_eq!(merged.user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(merged.api_key_id.as_deref(), Some("key-1"));
|
||||
assert_eq!(merged.provider_id.as_deref(), Some("provider-1"));
|
||||
assert_eq!(merged.endpoint_id.as_deref(), Some("endpoint-1"));
|
||||
assert_eq!(merged.key_id.as_deref(), Some("provider-key-1"));
|
||||
assert!(merged.error_message.is_none());
|
||||
assert_eq!(merged.latency_ms, Some(123));
|
||||
assert_eq!(merged.finished_at_unix_ms, Some(1_123));
|
||||
assert_eq!(
|
||||
merged.extra_data,
|
||||
Some(serde_json::json!({"terminal": true, "late": true}))
|
||||
Some(serde_json::json!({
|
||||
"gateway_execution_runtime": true,
|
||||
"stream_completed": true
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -12,6 +12,19 @@ use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, Whe
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
const OWNER_GUARDED_UPSERT_SQL: &str = r#"
|
||||
INSERT INTO gemini_file_mappings (
|
||||
id, file_name, key_id, user_id, display_name, mime_type, source_hash,
|
||||
created_at, expires_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
display_name = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(display_name), display_name),
|
||||
mime_type = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(mime_type), mime_type),
|
||||
source_hash = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(source_hash), source_hash),
|
||||
expires_at = IF(BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id), VALUES(expires_at), expires_at)
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlGeminiFileMappingRepository {
|
||||
pool: MysqlPool,
|
||||
@@ -63,6 +76,85 @@ LIMIT 1
|
||||
row.as_ref().map(map_row).transpose()
|
||||
}
|
||||
|
||||
async fn find_active_by_file_name_for_user(
|
||||
&self,
|
||||
file_name: &str,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
file_name,
|
||||
key_id,
|
||||
user_id,
|
||||
display_name,
|
||||
mime_type,
|
||||
source_hash,
|
||||
created_at AS created_at_unix_ms,
|
||||
expires_at AS expires_at_unix_secs
|
||||
FROM gemini_file_mappings
|
||||
WHERE BINARY file_name = BINARY ?
|
||||
AND BINARY user_id = BINARY ?
|
||||
AND expires_at > ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(file_name)
|
||||
.bind(user_id)
|
||||
.bind(i64_from_u64(
|
||||
now_unix_secs,
|
||||
"gemini_file_mappings.owner_read_now",
|
||||
)?)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
row.as_ref().map(map_row).transpose()
|
||||
}
|
||||
|
||||
async fn find_active_by_file_name_for_owner(
|
||||
&self,
|
||||
file_name: &str,
|
||||
key_id: &str,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
file_name,
|
||||
key_id,
|
||||
user_id,
|
||||
display_name,
|
||||
mime_type,
|
||||
source_hash,
|
||||
created_at AS created_at_unix_ms,
|
||||
expires_at AS expires_at_unix_secs
|
||||
FROM gemini_file_mappings
|
||||
WHERE BINARY file_name = BINARY ?
|
||||
AND BINARY key_id = BINARY ?
|
||||
AND BINARY user_id = BINARY ?
|
||||
AND expires_at > ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(file_name)
|
||||
.bind(key_id)
|
||||
.bind(user_id)
|
||||
.bind(i64_from_u64(
|
||||
now_unix_secs,
|
||||
"gemini_file_mappings.owner_read_now",
|
||||
)?)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
row.as_ref().map(map_row).transpose()
|
||||
}
|
||||
|
||||
async fn list_mappings(
|
||||
&self,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
@@ -185,6 +277,60 @@ ON DUPLICATE KEY UPDATE
|
||||
self.reload_by_file_name(&record.file_name).await
|
||||
}
|
||||
|
||||
async fn upsert_if_owner_matches(
|
||||
&self,
|
||||
record: UpsertGeminiFileMappingRecord,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
|
||||
record.validate()?;
|
||||
let mut transaction = self.pool.begin().await.map_sql_err()?;
|
||||
sqlx::query(OWNER_GUARDED_UPSERT_SQL)
|
||||
.bind(&record.id)
|
||||
.bind(&record.file_name)
|
||||
.bind(&record.key_id)
|
||||
.bind(&record.user_id)
|
||||
.bind(&record.display_name)
|
||||
.bind(&record.mime_type)
|
||||
.bind(&record.source_hash)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(i64_from_u64(
|
||||
record.expires_at_unix_secs,
|
||||
"gemini_file_mappings.expires_at",
|
||||
)?)
|
||||
.execute(&mut *transaction)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
file_name,
|
||||
key_id,
|
||||
user_id,
|
||||
display_name,
|
||||
mime_type,
|
||||
source_hash,
|
||||
created_at AS created_at_unix_ms,
|
||||
expires_at AS expires_at_unix_secs
|
||||
FROM gemini_file_mappings
|
||||
WHERE file_name = ?
|
||||
LIMIT 1
|
||||
FOR UPDATE
|
||||
"#,
|
||||
)
|
||||
.bind(&record.file_name)
|
||||
.fetch_one(&mut *transaction)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let stored = map_row(&row)?;
|
||||
let owner_matches = stored.file_name == record.file_name
|
||||
&& stored.key_id == record.key_id
|
||||
&& stored.user_id == record.user_id;
|
||||
transaction.commit().await.map_sql_err()?;
|
||||
|
||||
Ok(owner_matches.then_some(stored))
|
||||
}
|
||||
|
||||
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
|
||||
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ?")
|
||||
.bind(file_name)
|
||||
@@ -195,6 +341,43 @@ ON DUPLICATE KEY UPDATE
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
async fn delete_by_file_name_for_user(
|
||||
&self,
|
||||
file_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let rows_affected =
|
||||
sqlx::query(
|
||||
"DELETE FROM gemini_file_mappings WHERE BINARY file_name = BINARY ? AND BINARY user_id = BINARY ?",
|
||||
)
|
||||
.bind(file_name)
|
||||
.bind(user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
async fn delete_by_file_name_for_owner(
|
||||
&self,
|
||||
file_name: &str,
|
||||
key_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let rows_affected = sqlx::query(
|
||||
"DELETE FROM gemini_file_mappings WHERE BINARY file_name = BINARY ? AND BINARY key_id = BINARY ? AND BINARY user_id = BINARY ?",
|
||||
)
|
||||
.bind(file_name)
|
||||
.bind(key_id)
|
||||
.bind(user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
async fn delete_by_id(
|
||||
&self,
|
||||
mapping_id: &str,
|
||||
@@ -282,6 +465,11 @@ fn apply_list_filters(
|
||||
where_clause: &mut WhereClause,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
) {
|
||||
if let Some(user_id) = query.user_id.as_deref() {
|
||||
where_clause.push_next(builder);
|
||||
builder.push("BINARY user_id = BINARY ");
|
||||
builder.push_bind(user_id.to_string());
|
||||
}
|
||||
if !query.include_expired {
|
||||
where_clause.push_next(builder);
|
||||
builder.push("expires_at > ");
|
||||
@@ -343,13 +531,17 @@ fn map_row(row: &MySqlRow) -> Result<StoredGeminiFileMapping, DataLayerError> {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository};
|
||||
use super::{
|
||||
build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository,
|
||||
OWNER_GUARDED_UPSERT_SQL,
|
||||
};
|
||||
use aether_data_contracts::repository::gemini_file_mappings::GeminiFileMappingListQuery;
|
||||
use sqlx::Execute;
|
||||
|
||||
#[test]
|
||||
fn list_query_uses_shared_mysql_filter_and_pagination_rendering() {
|
||||
let query = GeminiFileMappingListQuery {
|
||||
user_id: Some("user-1".to_string()),
|
||||
include_expired: false,
|
||||
search: Some(" Report ".to_string()),
|
||||
offset: 5,
|
||||
@@ -359,7 +551,9 @@ mod tests {
|
||||
|
||||
let mut count = build_list_count_query(&query);
|
||||
let count_sql = count.build().sql().to_string();
|
||||
assert!(count_sql.contains(" WHERE expires_at > ? AND (LOWER(file_name) LIKE ?"));
|
||||
assert!(count_sql.contains(
|
||||
" WHERE BINARY user_id = BINARY ? AND expires_at > ? AND (LOWER(file_name) LIKE ?"
|
||||
));
|
||||
assert!(!count_sql.contains("WHERE 1=1"));
|
||||
|
||||
let mut rows = build_list_rows_query(&query);
|
||||
@@ -368,6 +562,12 @@ mod tests {
|
||||
assert!(rows_sql.contains(" ORDER BY created_at DESC, file_name ASC LIMIT ? OFFSET ?"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn owner_guarded_upsert_uses_exact_file_key_and_user_identity() {
|
||||
let identity = "BINARY file_name = BINARY VALUES(file_name) AND BINARY key_id = BINARY VALUES(key_id) AND BINARY user_id <=> BINARY VALUES(user_id)";
|
||||
assert_eq!(OWNER_GUARDED_UPSERT_SQL.matches(identity).count(), 4);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_builds_from_lazy_pool() {
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with(
|
||||
|
||||
@@ -2,10 +2,10 @@ use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::management_tokens::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
|
||||
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
|
||||
UpdateManagementTokenRecord,
|
||||
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
|
||||
ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
|
||||
StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary,
|
||||
StoredManagementTokenWithUser, UpdateManagementTokenRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause};
|
||||
@@ -38,8 +38,147 @@ impl MysqlManagementTokenRepository {
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_token_row).transpose()
|
||||
}
|
||||
|
||||
async fn get_token_scoped(
|
||||
&self,
|
||||
token_id: &str,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(TOKEN_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", token_id.to_string());
|
||||
push_optional_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"user_id",
|
||||
expected_user_id.map(ToOwned::to_owned),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_token_row).transpose()
|
||||
}
|
||||
|
||||
async fn update_management_token_scoped(
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
record.validate()?;
|
||||
let allowed_ips = json_to_string(record.allowed_ips.as_ref())?;
|
||||
let permissions = json_to_string(record.permissions.as_ref())?;
|
||||
let now = now_unix_secs();
|
||||
|
||||
sqlx::query(UPDATE_MANAGEMENT_TOKEN_SQL)
|
||||
.bind(record.name.as_deref())
|
||||
.bind(record.clear_description)
|
||||
.bind(record.description.as_deref())
|
||||
.bind(record.clear_allowed_ips)
|
||||
.bind(allowed_ips)
|
||||
.bind(permissions)
|
||||
.bind(record.clear_expires_at)
|
||||
.bind(
|
||||
record
|
||||
.expires_at_unix_secs
|
||||
.and_then(|value| i64::try_from(value).ok()),
|
||||
)
|
||||
.bind(record.is_active)
|
||||
.bind(now as i64)
|
||||
.bind(&record.token_id)
|
||||
.bind(expected_user_id)
|
||||
.bind(expected_user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?;
|
||||
self.get_token_scoped(&record.token_id, expected_user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn delete_management_token_scoped(
|
||||
&self,
|
||||
token_id: &str,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
"DELETE FROM management_tokens WHERE id = ? AND (? IS NULL OR user_id = ?)",
|
||||
)
|
||||
.bind(token_id)
|
||||
.bind(expected_user_id)
|
||||
.bind(expected_user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
async fn set_management_token_active_scoped(
|
||||
&self,
|
||||
token_id: &str,
|
||||
expected_user_id: Option<&str>,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
"UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ? AND (? IS NULL OR user_id = ?)",
|
||||
)
|
||||
.bind(is_active)
|
||||
.bind(now_unix_secs() as i64)
|
||||
.bind(token_id)
|
||||
.bind(expected_user_id)
|
||||
.bind(expected_user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.get_token_scoped(token_id, expected_user_id).await
|
||||
}
|
||||
|
||||
async fn regenerate_management_token_secret_scoped(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE management_tokens
|
||||
SET token_hash = ?, token_prefix = ?, updated_at = ?
|
||||
WHERE id = ? AND (? IS NULL OR user_id = ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&mutation.token_hash)
|
||||
.bind(mutation.token_prefix.as_deref())
|
||||
.bind(now_unix_secs() as i64)
|
||||
.bind(&mutation.token_id)
|
||||
.bind(expected_user_id)
|
||||
.bind(expected_user_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.get_token_scoped(&mutation.token_id, expected_user_id)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
const UPDATE_MANAGEMENT_TOKEN_SQL: &str = r#"
|
||||
UPDATE management_tokens
|
||||
SET name = COALESCE(?, name),
|
||||
description = CASE WHEN ? THEN NULL ELSE COALESCE(?, description) END,
|
||||
allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END,
|
||||
permissions = COALESCE(?, permissions),
|
||||
expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END,
|
||||
is_active = COALESCE(?, is_active),
|
||||
updated_at = ?
|
||||
WHERE id = ? AND (? IS NULL OR user_id = ?)
|
||||
"#;
|
||||
|
||||
const TOKEN_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
@@ -59,6 +198,39 @@ SELECT
|
||||
FROM management_tokens
|
||||
"#;
|
||||
|
||||
const LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL: &str = r#"
|
||||
SELECT id
|
||||
FROM users
|
||||
WHERE id = ?
|
||||
AND is_active = TRUE
|
||||
AND is_deleted = FALSE
|
||||
AND LOWER(role) = 'admin'
|
||||
AND security_version = ?
|
||||
FOR UPDATE
|
||||
"#;
|
||||
|
||||
const LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
user_id,
|
||||
token_hash,
|
||||
name,
|
||||
description,
|
||||
token_prefix,
|
||||
allowed_ips,
|
||||
permissions,
|
||||
expires_at AS expires_at_unix_secs,
|
||||
last_used_at AS last_used_at_unix_secs,
|
||||
last_used_ip,
|
||||
COALESCE(usage_count, 0) AS usage_count,
|
||||
is_active,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM management_tokens
|
||||
WHERE id = ?
|
||||
FOR UPDATE
|
||||
"#;
|
||||
|
||||
const TOKEN_WITH_USER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
@@ -220,71 +392,29 @@ INSERT INTO management_tokens (
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
record.validate()?;
|
||||
let current = self.get_token(&record.token_id).await?;
|
||||
let Some(current) = current else {
|
||||
return Ok(None);
|
||||
};
|
||||
let name = record.name.as_deref().unwrap_or(¤t.name);
|
||||
let description = if record.clear_description {
|
||||
None
|
||||
} else {
|
||||
record
|
||||
.description
|
||||
.as_deref()
|
||||
.or(current.description.as_deref())
|
||||
};
|
||||
let allowed_ips = if record.clear_allowed_ips {
|
||||
None
|
||||
} else {
|
||||
record.allowed_ips.as_ref().or(current.allowed_ips.as_ref())
|
||||
};
|
||||
let permissions = record.permissions.as_ref().or(current.permissions.as_ref());
|
||||
let expires_at = if record.clear_expires_at {
|
||||
None
|
||||
} else {
|
||||
record.expires_at_unix_secs.or(current.expires_at_unix_secs)
|
||||
};
|
||||
let is_active = record.is_active.unwrap_or(current.is_active);
|
||||
let now = now_unix_secs();
|
||||
self.update_management_token_scoped(record, None).await
|
||||
}
|
||||
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE management_tokens
|
||||
SET name = ?,
|
||||
description = ?,
|
||||
allowed_ips = ?,
|
||||
permissions = ?,
|
||||
expires_at = ?,
|
||||
is_active = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(name)
|
||||
.bind(description)
|
||||
.bind(json_to_string(allowed_ips)?)
|
||||
.bind(json_to_string(permissions)?)
|
||||
.bind(expires_at.and_then(|value| i64::try_from(value).ok()))
|
||||
.bind(is_active)
|
||||
.bind(now as i64)
|
||||
.bind(&record.token_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.get_token(&record.token_id).await
|
||||
async fn update_management_token_for_user(
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
self.update_management_token_scoped(record, Some(user_id))
|
||||
.await
|
||||
}
|
||||
|
||||
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
|
||||
let result = sqlx::query("DELETE FROM management_tokens WHERE id = ?")
|
||||
.bind(token_id)
|
||||
.execute(&self.pool)
|
||||
self.delete_management_token_scoped(token_id, None).await
|
||||
}
|
||||
|
||||
async fn delete_management_token_for_user(
|
||||
&self,
|
||||
token_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
self.delete_management_token_scoped(token_id, Some(user_id))
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
async fn set_management_token_active(
|
||||
@@ -292,43 +422,135 @@ WHERE id = ?
|
||||
token_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
let result =
|
||||
sqlx::query("UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ?")
|
||||
.bind(is_active)
|
||||
.bind(now_unix_secs() as i64)
|
||||
.bind(token_id)
|
||||
.execute(&self.pool)
|
||||
self.set_management_token_active_scoped(token_id, None, is_active)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn set_management_token_active_for_user(
|
||||
&self,
|
||||
token_id: &str,
|
||||
user_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
self.set_management_token_active_scoped(token_id, Some(user_id), is_active)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn activate_management_token_if_matches(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let eligible_user =
|
||||
sqlx::query_scalar::<_, String>(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL)
|
||||
.bind(&mutation.expected_token.user_id)
|
||||
.bind(mutation.expected_user_security_version)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
if eligible_user.is_none() {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
self.get_token(token_id).await
|
||||
|
||||
let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL)
|
||||
.bind(&mutation.expected_token.id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let snapshot_matches = match locked.as_ref() {
|
||||
Some(row) => {
|
||||
let token_hash: String = row.try_get("token_hash").map_sql_err()?;
|
||||
let token = map_token_row(row)?;
|
||||
mutation.matches_locked_token_snapshot(&token, &token_hash)
|
||||
}
|
||||
None => false,
|
||||
};
|
||||
if !snapshot_matches {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE management_tokens
|
||||
SET is_active = TRUE, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND BINARY token_hash = BINARY ?
|
||||
AND is_active = FALSE
|
||||
AND (expires_at IS NULL OR expires_at > ?)
|
||||
"#,
|
||||
)
|
||||
.bind(now_unix_secs() as i64)
|
||||
.bind(&mutation.expected_token.id)
|
||||
.bind(&mutation.token_hash)
|
||||
.bind(i64::try_from(mutation.now_unix_secs).unwrap_or(i64::MAX))
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() != 1 {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn delete_inactive_management_token_if_matches(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let locked = sqlx::query(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL)
|
||||
.bind(&mutation.expected_token.id)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let snapshot_matches = match locked.as_ref() {
|
||||
Some(row) => {
|
||||
let token_hash: String = row.try_get("token_hash").map_sql_err()?;
|
||||
let token = map_token_row(row)?;
|
||||
mutation.matches_locked_token_snapshot(&token, &token_hash)
|
||||
}
|
||||
None => false,
|
||||
};
|
||||
if !snapshot_matches {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
let result = sqlx::query(
|
||||
"DELETE FROM management_tokens WHERE id = ? AND BINARY token_hash = BINARY ? AND is_active = FALSE",
|
||||
)
|
||||
.bind(&mutation.expected_token.id)
|
||||
.bind(&mutation.token_hash)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() != 1 {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn regenerate_management_token_secret(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
UPDATE management_tokens
|
||||
SET token_hash = ?, token_prefix = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&mutation.token_hash)
|
||||
.bind(mutation.token_prefix.as_deref())
|
||||
.bind(now_unix_secs() as i64)
|
||||
.bind(&mutation.token_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if result.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
self.get_token(&mutation.token_id).await
|
||||
self.regenerate_management_token_secret_scoped(mutation, None)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn regenerate_management_token_secret_for_user(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
self.regenerate_management_token_secret_scoped(mutation, Some(user_id))
|
||||
.await
|
||||
}
|
||||
|
||||
async fn record_management_token_usage(
|
||||
@@ -365,8 +587,18 @@ fn now_unix_secs() -> u64 {
|
||||
chrono::Utc::now().timestamp().max(0) as u64
|
||||
}
|
||||
|
||||
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
|
||||
value.and_then(|value| u64::try_from(value).ok())
|
||||
fn non_negative_u64(value: i64, field_name: &str) -> Result<u64, DataLayerError> {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"management_tokens.{field_name} must not be negative"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn optional_unix_secs(value: Option<i64>, field_name: &str) -> Result<Option<u64>, DataLayerError> {
|
||||
value
|
||||
.map(|value| non_negative_u64(value, field_name))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
|
||||
@@ -420,15 +652,30 @@ fn map_token_row(row: &MySqlRow) -> Result<StoredManagementToken, DataLayerError
|
||||
)
|
||||
.with_permissions(json_from_string(row.try_get("permissions").map_sql_err()?)?)
|
||||
.with_runtime_fields(
|
||||
optional_unix_secs(row.try_get("expires_at_unix_secs").map_sql_err()?),
|
||||
optional_unix_secs(row.try_get("last_used_at_unix_secs").map_sql_err()?),
|
||||
optional_unix_secs(
|
||||
row.try_get("expires_at_unix_secs").map_sql_err()?,
|
||||
"expires_at",
|
||||
)?,
|
||||
optional_unix_secs(
|
||||
row.try_get("last_used_at_unix_secs").map_sql_err()?,
|
||||
"last_used_at",
|
||||
)?,
|
||||
row.try_get("last_used_ip").map_sql_err()?,
|
||||
u64::try_from(row.try_get::<i64, _>("usage_count").map_sql_err()?).unwrap_or(0),
|
||||
non_negative_u64(
|
||||
row.try_get::<i64, _>("usage_count").map_sql_err()?,
|
||||
"usage_count",
|
||||
)?,
|
||||
row.try_get("is_active").map_sql_err()?,
|
||||
)
|
||||
.with_timestamps(
|
||||
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
|
||||
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
|
||||
optional_unix_secs(
|
||||
row.try_get("created_at_unix_ms").map_sql_err()?,
|
||||
"created_at",
|
||||
)?,
|
||||
optional_unix_secs(
|
||||
row.try_get("updated_at_unix_secs").map_sql_err()?,
|
||||
"updated_at",
|
||||
)?,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -454,7 +701,11 @@ fn map_token_with_user_row(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlManagementTokenRepository;
|
||||
use super::{
|
||||
non_negative_u64, optional_unix_secs, MysqlManagementTokenRepository,
|
||||
LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL, LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL,
|
||||
UPDATE_MANAGEMENT_TOKEN_SQL,
|
||||
};
|
||||
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||
|
||||
#[tokio::test]
|
||||
@@ -468,6 +719,60 @@ mod tests {
|
||||
let _repository = MysqlManagementTokenRepository::new(pool);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_install_activation_locks_admin_identity_and_token_snapshot() {
|
||||
for predicate in [
|
||||
"is_active = TRUE",
|
||||
"is_deleted = FALSE",
|
||||
"LOWER(role) = 'admin'",
|
||||
"security_version = ?",
|
||||
"FOR UPDATE",
|
||||
] {
|
||||
assert!(LOCK_ELIGIBLE_MANAGEMENT_TOKEN_ADMIN_SQL.contains(predicate));
|
||||
}
|
||||
for column in [
|
||||
"token_hash",
|
||||
"name",
|
||||
"description",
|
||||
"token_prefix",
|
||||
"allowed_ips",
|
||||
"permissions",
|
||||
"expires_at_unix_secs",
|
||||
"last_used_at_unix_secs",
|
||||
"last_used_ip",
|
||||
"usage_count",
|
||||
"is_active",
|
||||
"created_at_unix_ms",
|
||||
"updated_at_unix_secs",
|
||||
"FOR UPDATE",
|
||||
] {
|
||||
assert!(LOCK_MANAGEMENT_TOKEN_ACTIVATION_SNAPSHOT_SQL.contains(column));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_management_token_mapping_rejects_negative_integer_state() {
|
||||
assert!(optional_unix_secs(Some(-1), "expires_at").is_err());
|
||||
assert_eq!(
|
||||
optional_unix_secs(None, "expires_at").expect("SQL NULL should remain optional"),
|
||||
None
|
||||
);
|
||||
assert!(non_negative_u64(-1, "usage_count").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_management_token_updates_patch_only_explicit_fields() {
|
||||
for clause in [
|
||||
"name = COALESCE(?, name)",
|
||||
"allowed_ips = CASE WHEN ? THEN NULL ELSE COALESCE(?, allowed_ips) END",
|
||||
"permissions = COALESCE(?, permissions)",
|
||||
"expires_at = CASE WHEN ? THEN NULL ELSE COALESCE(?, expires_at) END",
|
||||
"is_active = COALESCE(?, is_active)",
|
||||
] {
|
||||
assert!(UPDATE_MANAGEMENT_TOKEN_SQL.contains(clause));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_management_token_pool_config_remains_driver_specific() {
|
||||
let config = SqlDatabaseConfig {
|
||||
|
||||
@@ -3,7 +3,7 @@ use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::oauth_providers::{
|
||||
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
@@ -115,6 +115,13 @@ WHERE users.is_active = 1
|
||||
)
|
||||
"#;
|
||||
|
||||
const COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL: &str = r#"
|
||||
UPDATE oauth_providers
|
||||
SET client_secret_encrypted = ?
|
||||
WHERE BINARY provider_type = BINARY ?
|
||||
AND BINARY client_secret_encrypted = BINARY ?
|
||||
"#;
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderReadRepository for MysqlOAuthProviderRepository {
|
||||
async fn list_oauth_provider_configs(
|
||||
@@ -158,12 +165,56 @@ impl OAuthProviderReadRepository for MysqlOAuthProviderRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderWriteRepository for MysqlOAuthProviderRepository {
|
||||
async fn upsert_oauth_provider_config(
|
||||
async fn upsert_oauth_provider_config_guarded(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
|
||||
ldap_exclusive: bool,
|
||||
force_disable: bool,
|
||||
_locked_users_snapshot: usize,
|
||||
) -> Result<UpsertOAuthProviderConfigOutcome, DataLayerError> {
|
||||
record.validate()?;
|
||||
let now = now_unix_secs();
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let existing_enabled: Option<bool> = if record.is_enabled || force_disable {
|
||||
None
|
||||
} else {
|
||||
sqlx::query_scalar::<_, String>(
|
||||
"SELECT provider_type FROM oauth_providers ORDER BY provider_type FOR UPDATE",
|
||||
)
|
||||
.fetch_all(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query_scalar("SELECT is_enabled FROM oauth_providers WHERE provider_type = ?")
|
||||
.bind(&record.provider_type)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
};
|
||||
if existing_enabled == Some(true) {
|
||||
let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL)
|
||||
.bind(&record.provider_type)
|
||||
.bind(&record.provider_type)
|
||||
.bind(ldap_exclusive)
|
||||
.bind(&record.provider_type)
|
||||
.fetch_one(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let affected_count =
|
||||
usize::try_from(row.try_get::<i64, _>("locked_count").map_sql_err()?.max(0))
|
||||
.map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.locked_user_count overflowed".to_string(),
|
||||
)
|
||||
})?;
|
||||
if affected_count > 0 {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(
|
||||
UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation {
|
||||
affected_count,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO oauth_providers (
|
||||
@@ -228,27 +279,65 @@ ON DUPLICATE KEY UPDATE
|
||||
.bind(now as i64)
|
||||
.bind(record.client_secret_encrypted.mode_name())
|
||||
.bind(record.client_secret_encrypted.value())
|
||||
.execute(&self.pool)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
self.get_provider(&record.provider_type)
|
||||
.await?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string())
|
||||
})
|
||||
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(&record.provider_type)
|
||||
.fetch_one(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let provider = map_oauth_provider_row(&row)?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(UpsertOAuthProviderConfigOutcome::Upserted(provider))
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config(
|
||||
async fn compare_and_swap_oauth_provider_client_secret(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?")
|
||||
let result = sqlx::query(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL)
|
||||
.bind(replacement)
|
||||
.bind(provider_type)
|
||||
.bind(expected)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
Ok(result.rows_affected() == 1)
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config_if_unlinked(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
has_links_snapshot: bool,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if has_links_snapshot {
|
||||
return Ok(false);
|
||||
}
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let provider_exists: Option<String> = sqlx::query_scalar(
|
||||
"SELECT provider_type FROM oauth_providers WHERE provider_type = ? FOR UPDATE",
|
||||
)
|
||||
.bind(provider_type)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if provider_exists.is_none() {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
}
|
||||
let result = sqlx::query(
|
||||
"DELETE FROM oauth_providers WHERE provider_type = ? AND NOT EXISTS (SELECT 1 FROM user_oauth_links WHERE user_oauth_links.provider_type = oauth_providers.provider_type)",
|
||||
)
|
||||
.bind(provider_type)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(result.rows_affected() == 1)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -378,7 +467,16 @@ fn map_oauth_provider_row(row: &MySqlRow) -> Result<StoredOAuthProviderConfig, D
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlOAuthProviderRepository;
|
||||
use super::{MysqlOAuthProviderRepository, COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL};
|
||||
|
||||
#[test]
|
||||
fn client_secret_cas_updates_only_the_secret_column() {
|
||||
assert!(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL
|
||||
.contains("SET client_secret_encrypted = ?"));
|
||||
assert!(COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL
|
||||
.contains("BINARY client_secret_encrypted = BINARY ?"));
|
||||
assert!(!COMPARE_AND_SWAP_OAUTH_PROVIDER_CLIENT_SECRET_SQL.contains("updated_at"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_builds_from_lazy_pool() {
|
||||
|
||||
@@ -30,13 +30,20 @@ impl MysqlPoolFactory {
|
||||
}
|
||||
|
||||
pub fn connect_options(&self) -> Result<MySqlConnectOptions, DataLayerError> {
|
||||
let ssl_mode = if self.config.pool.require_ssl {
|
||||
MySqlSslMode::Required
|
||||
} else {
|
||||
MySqlSslMode::Preferred
|
||||
};
|
||||
MySqlConnectOptions::from_str(self.config.url.trim())
|
||||
.map(|options| {
|
||||
// Preserve explicit VERIFY_CA/VERIFY_IDENTITY from the URL.
|
||||
// `require_ssl` is a minimum transport guarantee: upgrade
|
||||
// weaker modes to Required, never downgrade verification.
|
||||
let ssl_mode = if self.config.pool.require_ssl
|
||||
&& !matches!(
|
||||
options.get_ssl_mode(),
|
||||
MySqlSslMode::VerifyCa | MySqlSslMode::VerifyIdentity
|
||||
) {
|
||||
MySqlSslMode::Required
|
||||
} else {
|
||||
options.get_ssl_mode()
|
||||
};
|
||||
options
|
||||
.ssl_mode(ssl_mode)
|
||||
.statement_cache_capacity(self.config.pool.statement_cache_capacity)
|
||||
@@ -78,6 +85,52 @@ impl MysqlPoolFactory {
|
||||
mod tests {
|
||||
use super::MysqlPoolFactory;
|
||||
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||
use sqlx::mysql::MySqlSslMode;
|
||||
|
||||
fn ssl_mode(url: &str, require_ssl: bool) -> MySqlSslMode {
|
||||
MysqlPoolFactory::new(SqlDatabaseConfig {
|
||||
driver: DatabaseDriver::Mysql,
|
||||
url: url.to_string(),
|
||||
pool: SqlPoolConfig {
|
||||
require_ssl,
|
||||
..SqlPoolConfig::default()
|
||||
},
|
||||
})
|
||||
.expect("mysql config should build")
|
||||
.connect_options()
|
||||
.expect("mysql options should parse")
|
||||
.get_ssl_mode()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_explicit_mysql_verification_modes() {
|
||||
assert!(matches!(
|
||||
ssl_mode(
|
||||
"mysql://user:pass@localhost/aether?ssl-mode=VERIFY_IDENTITY",
|
||||
false
|
||||
),
|
||||
MySqlSslMode::VerifyIdentity
|
||||
));
|
||||
assert!(matches!(
|
||||
ssl_mode(
|
||||
"mysql://user:pass@localhost/aether?ssl-mode=VERIFY_CA",
|
||||
true
|
||||
),
|
||||
MySqlSslMode::VerifyCa
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn require_ssl_only_upgrades_weak_mysql_modes() {
|
||||
for mode in ["DISABLED", "PREFERRED", "REQUIRED"] {
|
||||
let url = format!("mysql://user:pass@localhost/aether?ssl-mode={mode}");
|
||||
assert!(matches!(ssl_mode(&url, true), MySqlSslMode::Required));
|
||||
}
|
||||
assert!(matches!(
|
||||
ssl_mode("mysql://user:pass@localhost/aether", false),
|
||||
MySqlSslMode::Preferred
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn factory_builds_lazy_pool_from_valid_config() {
|
||||
|
||||
@@ -9,9 +9,11 @@ use sqlx::{
|
||||
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate,
|
||||
ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate,
|
||||
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
@@ -550,6 +552,48 @@ WHERE id = ?
|
||||
self.reload_provider(&provider.id, "updated").await
|
||||
}
|
||||
|
||||
pub async fn compare_and_swap_provider_config(
|
||||
&self,
|
||||
update: &ProviderCatalogProviderConfigCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(&update.provider_id, "provider catalog provider_id")?;
|
||||
let expected_config =
|
||||
optional_json_to_string(&update.expected_config, "providers.expected_config")?;
|
||||
let config = optional_json_to_string(&update.config, "providers.config")?;
|
||||
let rows_affected = sqlx::query(
|
||||
r#"
|
||||
UPDATE providers
|
||||
SET config = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND config <=> ?
|
||||
"#,
|
||||
)
|
||||
.bind(config)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(&update.provider_id)
|
||||
.bind(expected_config)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected == 1)
|
||||
}
|
||||
|
||||
pub async fn compare_and_swap_provider_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(&update.record_id, "provider catalog provider_id")?;
|
||||
compare_and_swap_proxy_json(
|
||||
&self.pool,
|
||||
"SELECT proxy FROM providers WHERE id = ?",
|
||||
"UPDATE providers SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?",
|
||||
update,
|
||||
"providers.proxy",
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn delete_provider(&self, provider_id: &str) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(provider_id, "provider catalog provider_id")?;
|
||||
let rows_affected = sqlx::query("DELETE FROM providers WHERE id = ?")
|
||||
@@ -754,6 +798,21 @@ WHERE id = ?
|
||||
self.reload_endpoint(&endpoint.id, "updated").await
|
||||
}
|
||||
|
||||
pub async fn compare_and_swap_endpoint_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(&update.record_id, "provider catalog endpoint_id")?;
|
||||
compare_and_swap_proxy_json(
|
||||
&self.pool,
|
||||
"SELECT proxy FROM provider_endpoints WHERE id = ?",
|
||||
"UPDATE provider_endpoints SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?",
|
||||
update,
|
||||
"provider_endpoints.proxy",
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(endpoint_id, "provider catalog endpoint_id")?;
|
||||
let rows_affected = sqlx::query("DELETE FROM provider_endpoints WHERE id = ?")
|
||||
@@ -936,6 +995,46 @@ WHERE id = ?
|
||||
self.reload_key(&key.id, "updated").await
|
||||
}
|
||||
|
||||
pub async fn compare_and_swap_key_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(&update.record_id, "provider catalog key_id")?;
|
||||
compare_and_swap_proxy_json(
|
||||
&self.pool,
|
||||
"SELECT proxy FROM provider_api_keys WHERE id = ?",
|
||||
"UPDATE provider_api_keys SET proxy = ?, updated_at = ? WHERE id = ? AND BINARY proxy <=> BINARY ?",
|
||||
update,
|
||||
"provider_api_keys.proxy",
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn compare_and_swap_key_credentials(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyCredentialsCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(&update.key_id, "provider catalog key_id")?;
|
||||
validate_non_empty(
|
||||
&update.expected_provider_id,
|
||||
"provider catalog expected provider_id",
|
||||
)?;
|
||||
let rows_affected = sqlx::query(
|
||||
"UPDATE provider_api_keys SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND BINARY provider_id = BINARY ? AND BINARY COALESCE(api_key, encrypted_key) <=> BINARY ? AND BINARY auth_config <=> BINARY ?",
|
||||
)
|
||||
.bind(update.encrypted_api_key.as_deref())
|
||||
.bind(update.encrypted_auth_config.as_deref())
|
||||
.bind(&update.key_id)
|
||||
.bind(&update.expected_provider_id)
|
||||
.bind(update.expected_encrypted_api_key.as_deref())
|
||||
.bind(update.expected_encrypted_auth_config.as_deref())
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected == 1)
|
||||
}
|
||||
|
||||
pub async fn compare_and_update_key_admin_state(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyAdminCasUpdate,
|
||||
@@ -1354,51 +1453,18 @@ WHERE id = ?
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
pub async fn update_key_oauth_credentials(
|
||||
&self,
|
||||
key_id: &str,
|
||||
encrypted_api_key: &str,
|
||||
encrypted_auth_config: Option<&str>,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(key_id, "provider catalog key_id")?;
|
||||
validate_non_empty(encrypted_api_key, "provider catalog oauth api_key")?;
|
||||
let rows_affected = sqlx::query(
|
||||
r#"
|
||||
UPDATE provider_api_keys
|
||||
SET api_key = ?, auth_config = ?, expires_at = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(encrypted_api_key)
|
||||
.bind(encrypted_auth_config)
|
||||
.bind(optional_i64_from_u64(
|
||||
expires_at_unix_secs,
|
||||
"provider_api_keys.expires_at",
|
||||
)?)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(key_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
pub async fn update_key_oauth_runtime_state(
|
||||
&self,
|
||||
key_id: &str,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
oauth_invalid_reason: Option<&str>,
|
||||
encrypted_auth_config_update: Option<&str>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
validate_non_empty(key_id, "provider catalog key_id")?;
|
||||
let rows_affected = sqlx::query(
|
||||
r#"
|
||||
UPDATE provider_api_keys
|
||||
SET oauth_invalid_at = ?, oauth_invalid_reason = ?,
|
||||
auth_config = COALESCE(?, auth_config), updated_at = ?
|
||||
SET oauth_invalid_at = ?, oauth_invalid_reason = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
@@ -1407,7 +1473,6 @@ WHERE id = ?
|
||||
"provider_api_keys.oauth_invalid_at",
|
||||
)?)
|
||||
.bind(oauth_invalid_reason)
|
||||
.bind(encrypted_auth_config_update)
|
||||
.bind(updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64)
|
||||
.bind(key_id)
|
||||
.execute(&self.pool)
|
||||
@@ -1755,7 +1820,7 @@ WHERE id = ?
|
||||
update.expected_encrypted_auth_config.as_deref()
|
||||
{
|
||||
builder
|
||||
.push(" AND auth_config <=> ")
|
||||
.push(" AND BINARY auth_config <=> BINARY ")
|
||||
.push_bind(expected_encrypted_auth_config);
|
||||
}
|
||||
let rows_affected = builder
|
||||
@@ -1873,7 +1938,7 @@ SET health_by_format = ?, circuit_breaker_by_format = ?, updated_at = ?
|
||||
WHERE id = ?
|
||||
AND JSON_EXTRACT(health_by_format, '$') <=> CAST(? AS JSON)
|
||||
AND JSON_EXTRACT(circuit_breaker_by_format, '$') <=> CAST(? AS JSON)
|
||||
AND (? IS NULL OR auth_config <=> ?)
|
||||
AND (? IS NULL OR BINARY auth_config <=> BINARY ?)
|
||||
"#,
|
||||
)
|
||||
.bind(optional_json_to_string(
|
||||
@@ -2042,6 +2107,20 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
|
||||
Self::update_provider(self, provider).await
|
||||
}
|
||||
|
||||
async fn compare_and_swap_provider_config(
|
||||
&self,
|
||||
update: &ProviderCatalogProviderConfigCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::compare_and_swap_provider_config(self, update).await
|
||||
}
|
||||
|
||||
async fn compare_and_swap_provider_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::compare_and_swap_provider_proxy(self, update).await
|
||||
}
|
||||
|
||||
async fn delete_provider(&self, provider_id: &str) -> Result<bool, DataLayerError> {
|
||||
Self::delete_provider(self, provider_id).await
|
||||
}
|
||||
@@ -2077,6 +2156,13 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
|
||||
Self::update_endpoint(self, endpoint).await
|
||||
}
|
||||
|
||||
async fn compare_and_swap_endpoint_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::compare_and_swap_endpoint_proxy(self, update).await
|
||||
}
|
||||
|
||||
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, DataLayerError> {
|
||||
Self::delete_endpoint(self, endpoint_id).await
|
||||
}
|
||||
@@ -2095,6 +2181,20 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
|
||||
Self::update_key(self, key).await
|
||||
}
|
||||
|
||||
async fn compare_and_swap_key_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::compare_and_swap_key_proxy(self, update).await
|
||||
}
|
||||
|
||||
async fn compare_and_swap_key_credentials(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyCredentialsCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::compare_and_swap_key_credentials(self, update).await
|
||||
}
|
||||
|
||||
async fn compare_and_update_key_admin_state(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyAdminCasUpdate,
|
||||
@@ -2189,29 +2289,11 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
|
||||
Self::clear_key_oauth_invalid_marker(self, key_id).await
|
||||
}
|
||||
|
||||
async fn update_key_oauth_credentials(
|
||||
&self,
|
||||
key_id: &str,
|
||||
encrypted_api_key: &str,
|
||||
encrypted_auth_config: Option<&str>,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::update_key_oauth_credentials(
|
||||
self,
|
||||
key_id,
|
||||
encrypted_api_key,
|
||||
encrypted_auth_config,
|
||||
expires_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn update_key_oauth_runtime_state(
|
||||
&self,
|
||||
key_id: &str,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
oauth_invalid_reason: Option<&str>,
|
||||
encrypted_auth_config_update: Option<&str>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::update_key_oauth_runtime_state(
|
||||
@@ -2219,7 +2301,6 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
|
||||
key_id,
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
encrypted_auth_config_update,
|
||||
updated_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
@@ -2490,6 +2571,44 @@ fn optional_json_to_string(
|
||||
optional_json_ref_to_string(value.as_ref(), field_name)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_proxy_json(
|
||||
pool: &MysqlPool,
|
||||
select_sql: &'static str,
|
||||
update_sql: &'static str,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
field_name: &'static str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
// Legacy catalog rows may contain semantically identical JSON with Python-style
|
||||
// whitespace. Comparing a re-serialized serde_json::Value directly to a TEXT column
|
||||
// would make lazy credential migration conflict forever. Compare the parsed value first,
|
||||
// then fence the write against the exact raw bytes that were observed.
|
||||
// Outer None means the row does not exist; inner None is an existing SQL NULL proxy.
|
||||
let observed_raw: Option<Option<String>> = sqlx::query_scalar::<_, Option<String>>(select_sql)
|
||||
.bind(&update.record_id)
|
||||
.fetch_optional(pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(observed_raw) = observed_raw else {
|
||||
return Ok(false);
|
||||
};
|
||||
let observed = optional_json_from_string(observed_raw.clone(), field_name)?;
|
||||
if observed != update.expected_proxy {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let replacement = optional_json_to_string(&update.proxy, field_name)?;
|
||||
let rows_affected = sqlx::query(update_sql)
|
||||
.bind(replacement)
|
||||
.bind(current_unix_secs() as i64)
|
||||
.bind(&update.record_id)
|
||||
.bind(observed_raw)
|
||||
.execute(pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(rows_affected == 1)
|
||||
}
|
||||
|
||||
fn build_in_query<'a>(
|
||||
select_sql: &'static str,
|
||||
column: &'static str,
|
||||
@@ -3134,11 +3253,11 @@ fn map_key_row(row: &MySqlRow) -> Result<StoredProviderCatalogKey, DataLayerErro
|
||||
)?,
|
||||
);
|
||||
key.note = row.try_get("note").map_sql_err()?;
|
||||
key.auth_type_by_format = optional_json_from_string(
|
||||
let auth_type_by_format = optional_json_from_string(
|
||||
row.try_get("auth_type_by_format").map_sql_err()?,
|
||||
"provider_api_keys.auth_type_by_format",
|
||||
)?;
|
||||
key.allow_auth_channel_mismatch_formats = optional_json_from_string(
|
||||
let allow_auth_channel_mismatch_formats = optional_json_from_string(
|
||||
row.try_get("allow_auth_channel_mismatch_formats")
|
||||
.map_sql_err()?,
|
||||
"provider_api_keys.allow_auth_channel_mismatch_formats",
|
||||
@@ -3204,7 +3323,10 @@ fn map_key_row(row: &MySqlRow) -> Result<StoredProviderCatalogKey, DataLayerErro
|
||||
row.try_get("updated_at_unix_secs").map_sql_err()?,
|
||||
"provider_api_keys.updated_at",
|
||||
)?;
|
||||
Ok::<_, DataLayerError>(key)
|
||||
key.with_auth_channel_policy_fields(
|
||||
auth_type_by_format,
|
||||
allow_auth_channel_mismatch_formats,
|
||||
)
|
||||
})?
|
||||
}
|
||||
|
||||
@@ -3238,6 +3360,14 @@ mod tests {
|
||||
assert!(sql.contains("binary auth_config <=> binary ?"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn credential_cas_migrates_legacy_encrypted_key_with_binary_fence() {
|
||||
let source = include_str!("provider_catalog.rs");
|
||||
assert!(source.contains(
|
||||
"SET api_key = ?, encrypted_key = NULL, auth_config = ? WHERE id = ? AND BINARY provider_id = BINARY ? AND BINARY COALESCE(api_key, encrypted_key) <=> BINARY ? AND BINARY auth_config <=> BINARY ?"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn admin_credential_cas_has_atomic_rotation_guards() {
|
||||
let source = include_str!("provider_catalog.rs");
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,10 +1,14 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
use sqlx::{mysql::MySqlRow, Acquire, Row};
|
||||
|
||||
use aether_data_contracts::repository::settlement::{
|
||||
finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd,
|
||||
settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement,
|
||||
UsageSettlementInput, SETTLEMENT_EPSILON_USD,
|
||||
settlement_billing_status_for_usage_status, validate_wallet_settlement_values,
|
||||
ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput,
|
||||
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput,
|
||||
ReserveUsagePolicyRequestOutcome, SettlementWriteRepository, StoredUsagePolicyCostReservation,
|
||||
StoredUsagePolicyRequestAdmission, StoredUsageSettlement, UsagePolicyCostReservationState,
|
||||
UsagePolicyRequestAdmissionState, UsageSettlementInput, SETTLEMENT_EPSILON_USD,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
@@ -112,6 +116,142 @@ impl MysqlSettlementRepository {
|
||||
}
|
||||
}
|
||||
|
||||
fn usage_policy_cost_i64(value: u64, field: &str) -> Result<i64, DataLayerError> {
|
||||
i64::try_from(value)
|
||||
.map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds the integer range")))
|
||||
}
|
||||
|
||||
fn usage_policy_cost_u64(value: i64, field: &str) -> Result<u64, DataLayerError> {
|
||||
u64::try_from(value)
|
||||
.map_err(|_| DataLayerError::UnexpectedValue(format!("{field} must not be negative")))
|
||||
}
|
||||
|
||||
fn usage_policy_request_admission_from_mysql_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredUsagePolicyRequestAdmission, DataLayerError> {
|
||||
let state: String = row.try_get("state").map_sql_err()?;
|
||||
Ok(StoredUsagePolicyRequestAdmission {
|
||||
request_id: row.try_get("request_id").map_sql_err()?,
|
||||
subject_id: row.try_get("subject_id").map_sql_err()?,
|
||||
event_token: row.try_get("event_token").map_sql_err()?,
|
||||
admitted_at_unix_secs: usage_policy_cost_u64(
|
||||
row.try_get("admitted_at_unix_secs").map_sql_err()?,
|
||||
"usage policy request admitted_at",
|
||||
)?,
|
||||
retain_until_unix_secs: usage_policy_cost_u64(
|
||||
row.try_get("retain_until_unix_secs").map_sql_err()?,
|
||||
"usage policy request retain_until",
|
||||
)?,
|
||||
state: UsagePolicyRequestAdmissionState::parse(&state).ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"unknown usage policy request admission state {state}"
|
||||
))
|
||||
})?,
|
||||
released_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("released_at_unix_secs")
|
||||
.map_sql_err()?
|
||||
.map(|value| usage_policy_cost_u64(value, "usage policy request released_at"))
|
||||
.transpose()?,
|
||||
})
|
||||
}
|
||||
|
||||
const FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL: &str = r#"
|
||||
SELECT request_id, subject_id, event_token,
|
||||
admitted_at AS admitted_at_unix_secs,
|
||||
retain_until AS retain_until_unix_secs,
|
||||
state, released_at AS released_at_unix_secs
|
||||
FROM usage_request_admissions
|
||||
WHERE event_token = ?
|
||||
FOR UPDATE
|
||||
"#;
|
||||
|
||||
const USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL: &str =
|
||||
"SET TRANSACTION ISOLATION LEVEL READ COMMITTED";
|
||||
|
||||
fn usage_policy_cost_reservation_from_mysql_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredUsagePolicyCostReservation, DataLayerError> {
|
||||
let state: String = row.try_get("state").map_sql_err()?;
|
||||
Ok(StoredUsagePolicyCostReservation {
|
||||
request_id: row.try_get("request_id").map_sql_err()?,
|
||||
subject_id: row.try_get("subject_id").map_sql_err()?,
|
||||
reservation_token: row.try_get("reservation_token").map_sql_err()?,
|
||||
admitted_at_unix_secs: usage_policy_cost_u64(
|
||||
row.try_get("admitted_at_unix_secs").map_sql_err()?,
|
||||
"usage policy admitted_at",
|
||||
)?,
|
||||
reserved_cost_units: usage_policy_cost_u64(
|
||||
row.try_get("reserved_cost_units").map_sql_err()?,
|
||||
"usage policy reserved_cost_units",
|
||||
)?,
|
||||
actual_cost_units: row
|
||||
.try_get::<Option<i64>, _>("actual_cost_units")
|
||||
.map_sql_err()?
|
||||
.map(|value| usage_policy_cost_u64(value, "usage policy actual_cost_units"))
|
||||
.transpose()?,
|
||||
state: UsagePolicyCostReservationState::parse(&state).ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"unknown usage policy reservation state {state}"
|
||||
))
|
||||
})?,
|
||||
reservation_expires_at_unix_secs: usage_policy_cost_u64(
|
||||
row.try_get("reservation_expires_at_unix_secs")
|
||||
.map_sql_err()?,
|
||||
"usage policy reservation_expires_at",
|
||||
)?,
|
||||
retain_until_unix_secs: usage_policy_cost_u64(
|
||||
row.try_get("retain_until_unix_secs").map_sql_err()?,
|
||||
"usage policy retain_until",
|
||||
)?,
|
||||
finalized_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("finalized_at_unix_secs")
|
||||
.map_sql_err()?
|
||||
.map(|value| usage_policy_cost_u64(value, "usage policy finalized_at"))
|
||||
.transpose()?,
|
||||
})
|
||||
}
|
||||
|
||||
const FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL: &str = r#"
|
||||
SELECT
|
||||
request_id,
|
||||
subject_id,
|
||||
reservation_token,
|
||||
admitted_at AS admitted_at_unix_secs,
|
||||
reserved_cost_units,
|
||||
actual_cost_units,
|
||||
state,
|
||||
reservation_expires_at AS reservation_expires_at_unix_secs,
|
||||
retain_until AS retain_until_unix_secs,
|
||||
finalized_at AS finalized_at_unix_secs
|
||||
FROM usage_cost_reservations
|
||||
WHERE reservation_token = ?
|
||||
FOR UPDATE
|
||||
"#;
|
||||
|
||||
async fn lock_usage_policy_subject_mysql(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
subject_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let exists = sqlx::query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT id
|
||||
FROM users
|
||||
WHERE id = ?
|
||||
FOR UPDATE
|
||||
"#,
|
||||
)
|
||||
.bind(subject_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.is_some();
|
||||
Ok(exists)
|
||||
}
|
||||
|
||||
fn usage_policy_subject_missing() -> DataLayerError {
|
||||
DataLayerError::InvalidInput("usage policy subject does not exist".to_string())
|
||||
}
|
||||
|
||||
fn settlement_from_row(row: &MySqlRow) -> Result<StoredUsageSettlement, DataLayerError> {
|
||||
Ok(StoredUsageSettlement {
|
||||
request_id: row.try_get("request_id").map_sql_err()?,
|
||||
@@ -261,7 +401,12 @@ async fn consume_daily_quota_mysql(
|
||||
wallet_can_overdraft: bool,
|
||||
now_unix_secs: i64,
|
||||
) -> Result<DailyQuotaDebitResult, DataLayerError> {
|
||||
if total_cost_usd <= 0.0 {
|
||||
if !total_cost_usd.is_finite() || total_cost_usd < 0.0 {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"daily quota settlement cost must be finite and non-negative".to_string(),
|
||||
));
|
||||
}
|
||||
if total_cost_usd == 0.0 {
|
||||
return Ok(DailyQuotaDebitResult::default());
|
||||
}
|
||||
let rows = sqlx::query(
|
||||
@@ -335,8 +480,18 @@ WHERE user_entitlement_id = ?
|
||||
.fetch_one(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if !used.is_finite() || used < 0.0 {
|
||||
return Err(DataLayerError::UnexpectedValue(
|
||||
"daily quota usage ledger total is invalid".to_string(),
|
||||
));
|
||||
}
|
||||
let remaining = (grant.daily_quota_usd - used).max(0.0);
|
||||
total_remaining += remaining;
|
||||
if !total_remaining.is_finite() {
|
||||
return Err(DataLayerError::UnexpectedValue(
|
||||
"daily quota remaining total overflowed".to_string(),
|
||||
));
|
||||
}
|
||||
grants_with_remaining.push((grant, remaining));
|
||||
}
|
||||
let insufficient = (!allow_wallet_overage && total_remaining + 0.000_000_01 < total_cost_usd)
|
||||
@@ -386,6 +541,476 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
|
||||
#[async_trait]
|
||||
impl SettlementWriteRepository for MysqlSettlementRepository {
|
||||
async fn reserve_usage_policy_request(
|
||||
&self,
|
||||
input: ReserveUsagePolicyRequestInput,
|
||||
) -> Result<ReserveUsagePolicyRequestOutcome, DataLayerError> {
|
||||
input.validate()?;
|
||||
let admitted_at = usage_policy_cost_i64(
|
||||
input.admitted_at_unix_secs,
|
||||
"usage policy request admitted_at",
|
||||
)?;
|
||||
let retain_until = usage_policy_cost_i64(
|
||||
input.retain_until_unix_secs,
|
||||
"usage policy request retain_until",
|
||||
)?;
|
||||
let created_at = now_unix_secs()?;
|
||||
// Different subjects lock different `users` rows. Under InnoDB's default REPEATABLE READ,
|
||||
// two missing-token locking reads can retain compatible gap locks and then deadlock when
|
||||
// both transactions try to insert the same unique event token. READ COMMITTED removes
|
||||
// that gap-lock cycle while the subject row still serializes each subject's window count.
|
||||
let mut connection = self.pool.acquire().await.map_sql_err()?;
|
||||
sqlx::query(USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL)
|
||||
.execute(&mut *connection)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let mut tx = connection.begin().await.map_sql_err()?;
|
||||
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
|
||||
return Err(usage_policy_subject_missing());
|
||||
}
|
||||
|
||||
let existing_row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL)
|
||||
.bind(&input.event_token)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if let Some(row) = existing_row.as_ref() {
|
||||
let existing = usage_policy_request_admission_from_mysql_row(row)?;
|
||||
if existing.request_id != input.request_id || existing.subject_id != input.subject_id {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(ReserveUsagePolicyRequestOutcome::Conflict);
|
||||
}
|
||||
if existing.admitted_at_unix_secs != input.admitted_at_unix_secs {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"usage policy event_token must keep its original admitted_at".to_string(),
|
||||
));
|
||||
}
|
||||
sqlx::query(
|
||||
"UPDATE usage_request_admissions SET retain_until = GREATEST(retain_until, ?) WHERE event_token = ?",
|
||||
)
|
||||
.bind(retain_until)
|
||||
.bind(&input.event_token)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let outcome = match existing.state {
|
||||
UsagePolicyRequestAdmissionState::Active => {
|
||||
ReserveUsagePolicyRequestOutcome::Allowed
|
||||
}
|
||||
UsagePolicyRequestAdmissionState::Released => {
|
||||
ReserveUsagePolicyRequestOutcome::AlreadyReleased
|
||||
}
|
||||
};
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(outcome);
|
||||
}
|
||||
|
||||
for (window_index, window) in input.windows.iter().enumerate() {
|
||||
let used_requests = sqlx::query_scalar::<_, i64>(
|
||||
r#"
|
||||
SELECT CAST(COUNT(*) AS SIGNED)
|
||||
FROM usage_request_admissions
|
||||
WHERE subject_id = ?
|
||||
AND state = 'active'
|
||||
AND admitted_at >= ?
|
||||
AND admitted_at < ?
|
||||
"#,
|
||||
)
|
||||
.bind(&input.subject_id)
|
||||
.bind(usage_policy_cost_i64(
|
||||
window.starts_at_unix_secs,
|
||||
"usage policy request window start",
|
||||
)?)
|
||||
.bind(usage_policy_cost_i64(
|
||||
window.ends_at_unix_secs,
|
||||
"usage policy request window end",
|
||||
)?)
|
||||
.fetch_one(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let used_requests =
|
||||
usage_policy_cost_u64(used_requests, "usage policy request used_requests")?;
|
||||
if used_requests >= window.limit_requests {
|
||||
let outcome = ReserveUsagePolicyRequestOutcome::Rejected {
|
||||
window_index,
|
||||
limit_requests: window.limit_requests,
|
||||
used_requests,
|
||||
};
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(outcome);
|
||||
}
|
||||
}
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO usage_request_admissions (
|
||||
request_id, subject_id, event_token, admitted_at, retain_until,
|
||||
state, released_at, created_at
|
||||
) VALUES (?, ?, ?, ?, ?, 'active', NULL, ?)
|
||||
ON DUPLICATE KEY UPDATE event_token = VALUES(event_token)
|
||||
"#,
|
||||
)
|
||||
.bind(&input.request_id)
|
||||
.bind(&input.subject_id)
|
||||
.bind(&input.event_token)
|
||||
.bind(admitted_at)
|
||||
.bind(retain_until)
|
||||
.bind(created_at)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL)
|
||||
.bind(&input.event_token)
|
||||
.fetch_one(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let existing = usage_policy_request_admission_from_mysql_row(&row)?;
|
||||
if existing.request_id != input.request_id || existing.subject_id != input.subject_id {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(ReserveUsagePolicyRequestOutcome::Conflict);
|
||||
}
|
||||
if existing.admitted_at_unix_secs != input.admitted_at_unix_secs {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"usage policy event_token must keep its original admitted_at".to_string(),
|
||||
));
|
||||
}
|
||||
sqlx::query(
|
||||
"UPDATE usage_request_admissions SET retain_until = GREATEST(retain_until, ?) WHERE event_token = ?",
|
||||
)
|
||||
.bind(retain_until)
|
||||
.bind(&input.event_token)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let outcome = match existing.state {
|
||||
UsagePolicyRequestAdmissionState::Active => ReserveUsagePolicyRequestOutcome::Allowed,
|
||||
UsagePolicyRequestAdmissionState::Released => {
|
||||
ReserveUsagePolicyRequestOutcome::AlreadyReleased
|
||||
}
|
||||
};
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(outcome)
|
||||
}
|
||||
|
||||
async fn release_usage_policy_request_admission(
|
||||
&self,
|
||||
input: ReleaseUsagePolicyRequestAdmissionInput,
|
||||
) -> Result<Option<StoredUsagePolicyRequestAdmission>, DataLayerError> {
|
||||
input.validate()?;
|
||||
let released_at = usage_policy_cost_i64(
|
||||
input.released_at_unix_secs,
|
||||
"usage policy request released_at",
|
||||
)?;
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
}
|
||||
let row = sqlx::query(FIND_USAGE_POLICY_REQUEST_ADMISSION_MYSQL_SQL)
|
||||
.bind(&input.event_token)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(row) = row else {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
};
|
||||
let mut admission = usage_policy_request_admission_from_mysql_row(&row)?;
|
||||
if admission.request_id != input.request_id || admission.subject_id != input.subject_id {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
}
|
||||
if input.released_at_unix_secs < admission.admitted_at_unix_secs {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"usage policy released_at must not precede admitted_at".to_string(),
|
||||
));
|
||||
}
|
||||
if admission.state == UsagePolicyRequestAdmissionState::Active {
|
||||
sqlx::query(
|
||||
"UPDATE usage_request_admissions SET state = 'released', released_at = ? WHERE event_token = ? AND state = 'active'",
|
||||
)
|
||||
.bind(released_at)
|
||||
.bind(&input.event_token)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
admission.state = UsagePolicyRequestAdmissionState::Released;
|
||||
admission.released_at_unix_secs = Some(input.released_at_unix_secs);
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(Some(admission))
|
||||
}
|
||||
|
||||
async fn cleanup_usage_policy_request_admissions(
|
||||
&self,
|
||||
now_unix_secs: u64,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if batch_size == 0 {
|
||||
return Ok(0);
|
||||
}
|
||||
let now = usage_policy_cost_i64(now_unix_secs, "usage policy request cleanup timestamp")?;
|
||||
let limit = i64::try_from(batch_size).unwrap_or(i64::MAX);
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM usage_request_admissions
|
||||
WHERE retain_until <= ?
|
||||
ORDER BY retain_until, event_token
|
||||
LIMIT ?
|
||||
"#,
|
||||
)
|
||||
.bind(now)
|
||||
.bind(limit)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() as usize)
|
||||
}
|
||||
|
||||
async fn reserve_usage_policy_cost(
|
||||
&self,
|
||||
input: ReserveUsagePolicyCostInput,
|
||||
) -> Result<ReserveUsagePolicyCostOutcome, DataLayerError> {
|
||||
input.validate()?;
|
||||
let admitted_at =
|
||||
usage_policy_cost_i64(input.admitted_at_unix_secs, "usage policy admitted_at")?;
|
||||
let reservation_expires_at = usage_policy_cost_i64(
|
||||
input.reservation_expires_at_unix_secs,
|
||||
"usage policy reservation_expires_at",
|
||||
)?;
|
||||
let updated_at = now_unix_secs()?;
|
||||
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
|
||||
return Err(usage_policy_subject_missing());
|
||||
}
|
||||
let existing_row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL)
|
||||
.bind(&input.reservation_token)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let existing = existing_row
|
||||
.as_ref()
|
||||
.map(usage_policy_cost_reservation_from_mysql_row)
|
||||
.transpose()?;
|
||||
if let Some(existing) = existing.as_ref() {
|
||||
if existing.request_id != input.request_id || existing.subject_id != input.subject_id {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(ReserveUsagePolicyCostOutcome::Conflict);
|
||||
}
|
||||
if existing.state != UsagePolicyCostReservationState::Reserved {
|
||||
let outcome = ReserveUsagePolicyCostOutcome::AlreadyTerminal {
|
||||
state: existing.state,
|
||||
};
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(outcome);
|
||||
}
|
||||
if existing.admitted_at_unix_secs != input.admitted_at_unix_secs {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"usage policy reservation_token must keep its original admitted_at".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
let previous_reserved_cost_units = existing
|
||||
.as_ref()
|
||||
.map(|reservation| reservation.reserved_cost_units)
|
||||
.unwrap_or(0);
|
||||
let target_reserved_cost_units =
|
||||
previous_reserved_cost_units.max(input.reserved_cost_units);
|
||||
for (window_index, window) in input.windows.iter().enumerate() {
|
||||
let window_start =
|
||||
usage_policy_cost_i64(window.starts_at_unix_secs, "usage policy window start")?;
|
||||
let window_end =
|
||||
usage_policy_cost_i64(window.ends_at_unix_secs, "usage policy window end")?;
|
||||
let used_cost_units = sqlx::query_scalar::<_, i64>(
|
||||
r#"
|
||||
SELECT CAST(COALESCE(SUM(
|
||||
CASE
|
||||
WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0)
|
||||
WHEN state = 'reserved' AND reservation_expires_at > ? THEN reserved_cost_units
|
||||
ELSE 0
|
||||
END
|
||||
), 0) AS SIGNED)
|
||||
FROM usage_cost_reservations
|
||||
WHERE subject_id = ?
|
||||
AND admitted_at >= ?
|
||||
AND admitted_at < ?
|
||||
AND reservation_token <> ?
|
||||
"#,
|
||||
)
|
||||
.bind(admitted_at)
|
||||
.bind(&input.subject_id)
|
||||
.bind(window_start)
|
||||
.bind(window_end)
|
||||
.bind(&input.reservation_token)
|
||||
.fetch_one(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let used_cost_units =
|
||||
usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?;
|
||||
if used_cost_units
|
||||
.checked_add(target_reserved_cost_units)
|
||||
.is_none_or(|total| total > window.limit_cost_units)
|
||||
{
|
||||
let outcome = ReserveUsagePolicyCostOutcome::Rejected {
|
||||
window_index,
|
||||
limit_cost_units: window.limit_cost_units,
|
||||
used_cost_units,
|
||||
};
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(outcome);
|
||||
}
|
||||
}
|
||||
|
||||
let admitted_at = existing
|
||||
.as_ref()
|
||||
.map(|reservation| {
|
||||
usage_policy_cost_i64(
|
||||
reservation.admitted_at_unix_secs,
|
||||
"usage policy admitted_at",
|
||||
)
|
||||
})
|
||||
.transpose()?
|
||||
.unwrap_or(admitted_at);
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO usage_cost_reservations (
|
||||
request_id, subject_id, reservation_token, admitted_at,
|
||||
reserved_cost_units, actual_cost_units,
|
||||
state, reservation_expires_at, retain_until, finalized_at, created_at, updated_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, NULL, 'reserved', ?, ?, NULL, ?, ?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
reserved_cost_units = GREATEST(reserved_cost_units, VALUES(reserved_cost_units)),
|
||||
reservation_expires_at = GREATEST(
|
||||
reservation_expires_at,
|
||||
VALUES(reservation_expires_at)
|
||||
),
|
||||
retain_until = GREATEST(retain_until, VALUES(retain_until)),
|
||||
updated_at = VALUES(updated_at)
|
||||
"#,
|
||||
)
|
||||
.bind(&input.request_id)
|
||||
.bind(&input.subject_id)
|
||||
.bind(&input.reservation_token)
|
||||
.bind(admitted_at)
|
||||
.bind(usage_policy_cost_i64(
|
||||
target_reserved_cost_units,
|
||||
"usage policy reserved_cost_units",
|
||||
)?)
|
||||
.bind(reservation_expires_at)
|
||||
.bind(usage_policy_cost_i64(
|
||||
input.retain_until_unix_secs,
|
||||
"usage policy retain_until",
|
||||
)?)
|
||||
.bind(updated_at)
|
||||
.bind(updated_at)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(ReserveUsagePolicyCostOutcome::Allowed {
|
||||
reserved_cost_units: target_reserved_cost_units,
|
||||
additional_reserved_cost_units: target_reserved_cost_units
|
||||
.saturating_sub(previous_reserved_cost_units),
|
||||
})
|
||||
}
|
||||
|
||||
async fn reconcile_usage_policy_cost(
|
||||
&self,
|
||||
input: ReconcileUsagePolicyCostInput,
|
||||
) -> Result<Option<StoredUsagePolicyCostReservation>, DataLayerError> {
|
||||
input.validate()?;
|
||||
let actual_cost_units =
|
||||
usage_policy_cost_i64(input.actual_cost_units, "usage policy actual_cost_units")?;
|
||||
let finalized_at =
|
||||
usage_policy_cost_i64(input.finalized_at_unix_secs, "usage policy finalized_at")?;
|
||||
let updated_at = now_unix_secs()?;
|
||||
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
if !lock_usage_policy_subject_mysql(&mut tx, &input.subject_id).await? {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
}
|
||||
let row = sqlx::query(FIND_USAGE_POLICY_COST_RESERVATION_MYSQL_SQL)
|
||||
.bind(&input.reservation_token)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(row) = row else {
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
};
|
||||
let mut reservation = usage_policy_cost_reservation_from_mysql_row(&row)?;
|
||||
if reservation.request_id != input.request_id || reservation.subject_id != input.subject_id
|
||||
{
|
||||
// The token selects the row; audit identity must still match before the reservation
|
||||
// can be finalized.
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(None);
|
||||
}
|
||||
if reservation.state == UsagePolicyCostReservationState::Reserved {
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE usage_cost_reservations
|
||||
SET state = ?,
|
||||
actual_cost_units = ?,
|
||||
finalized_at = ?,
|
||||
updated_at = ?
|
||||
WHERE reservation_token = ?
|
||||
AND request_id = ?
|
||||
AND subject_id = ?
|
||||
AND state = 'reserved'
|
||||
"#,
|
||||
)
|
||||
.bind(input.terminal_state.as_str())
|
||||
.bind(actual_cost_units)
|
||||
.bind(finalized_at)
|
||||
.bind(updated_at)
|
||||
.bind(&input.reservation_token)
|
||||
.bind(&input.request_id)
|
||||
.bind(&input.subject_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
reservation.state = input.terminal_state;
|
||||
reservation.actual_cost_units = Some(input.actual_cost_units);
|
||||
reservation.finalized_at_unix_secs = Some(input.finalized_at_unix_secs);
|
||||
}
|
||||
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(Some(reservation))
|
||||
}
|
||||
|
||||
async fn cleanup_usage_policy_cost_reservations(
|
||||
&self,
|
||||
now_unix_secs: u64,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if batch_size == 0 {
|
||||
return Ok(0);
|
||||
}
|
||||
let now = usage_policy_cost_i64(now_unix_secs, "usage policy cleanup timestamp")?;
|
||||
let limit = i64::try_from(batch_size).unwrap_or(i64::MAX);
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM usage_cost_reservations
|
||||
WHERE retain_until <= ?
|
||||
ORDER BY retain_until, reservation_token
|
||||
LIMIT ?
|
||||
"#,
|
||||
)
|
||||
.bind(now)
|
||||
.bind(limit)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() as usize)
|
||||
}
|
||||
|
||||
async fn settle_usage(
|
||||
&self,
|
||||
input: UsageSettlementInput,
|
||||
@@ -465,7 +1090,7 @@ LIMIT 1
|
||||
let wallet_row = if let Some(api_key_id) = api_key_id {
|
||||
sqlx::query(
|
||||
r#"
|
||||
SELECT id, balance, gift_balance, limit_mode
|
||||
SELECT id, balance, gift_balance, total_consumed, limit_mode
|
||||
FROM wallets
|
||||
WHERE api_key_id = ?
|
||||
LIMIT 1
|
||||
@@ -486,7 +1111,7 @@ FOR UPDATE
|
||||
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
|
||||
sqlx::query(
|
||||
r#"
|
||||
SELECT id, balance, gift_balance, limit_mode
|
||||
SELECT id, balance, gift_balance, total_consumed, limit_mode
|
||||
FROM wallets
|
||||
WHERE user_id = ?
|
||||
LIMIT 1
|
||||
@@ -507,14 +1132,20 @@ FOR UPDATE
|
||||
let wallet_can_overdraft = wallet_row.is_some();
|
||||
let wallet_available_usd = match wallet_row.as_ref() {
|
||||
Some(row) => {
|
||||
let recharge_balance: f64 = row.try_get("balance").map_sql_err()?;
|
||||
let gift_balance: f64 = row.try_get("gift_balance").map_sql_err()?;
|
||||
let total_consumed: f64 = row.try_get("total_consumed").map_sql_err()?;
|
||||
validate_wallet_settlement_values(
|
||||
recharge_balance,
|
||||
gift_balance,
|
||||
total_consumed,
|
||||
0.0,
|
||||
)?;
|
||||
let limit_mode: String = row.try_get("limit_mode").map_sql_err()?;
|
||||
if limit_mode.eq_ignore_ascii_case("unlimited") {
|
||||
None
|
||||
} else {
|
||||
Some(finite_wallet_available_usd(
|
||||
row.try_get("balance").map_sql_err()?,
|
||||
row.try_get("gift_balance").map_sql_err()?,
|
||||
))
|
||||
Some(finite_wallet_available_usd(recharge_balance, gift_balance))
|
||||
}
|
||||
}
|
||||
None => Some(0.0),
|
||||
@@ -593,6 +1224,7 @@ FOR UPDATE
|
||||
let wallet_id: String = wallet_row.try_get("id").map_sql_err()?;
|
||||
let before_recharge: f64 = wallet_row.try_get("balance").map_sql_err()?;
|
||||
let before_gift: f64 = wallet_row.try_get("gift_balance").map_sql_err()?;
|
||||
let total_consumed: f64 = wallet_row.try_get("total_consumed").map_sql_err()?;
|
||||
let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?;
|
||||
let before_total = before_recharge + before_gift;
|
||||
let mut after_recharge = before_recharge;
|
||||
@@ -606,6 +1238,13 @@ FOR UPDATE
|
||||
(after_recharge, after_gift) =
|
||||
debit_plan.after_balances(before_recharge, before_gift);
|
||||
}
|
||||
let total_consumed_after = total_consumed + wallet_debit_cost_usd;
|
||||
validate_wallet_settlement_values(
|
||||
after_recharge,
|
||||
after_gift,
|
||||
total_consumed_after,
|
||||
0.0,
|
||||
)?;
|
||||
if final_billing_status == "settled" {
|
||||
sqlx::query(
|
||||
r#"
|
||||
@@ -613,14 +1252,14 @@ UPDATE wallets
|
||||
SET
|
||||
balance = ?,
|
||||
gift_balance = ?,
|
||||
total_consumed = COALESCE(total_consumed, 0) + ?,
|
||||
total_consumed = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(after_recharge)
|
||||
.bind(after_gift)
|
||||
.bind(wallet_debit_cost_usd)
|
||||
.bind(total_consumed_after)
|
||||
.bind(updated_at)
|
||||
.bind(&wallet_id)
|
||||
.execute(&mut *tx)
|
||||
@@ -719,10 +1358,11 @@ WHERE id = ?
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlSettlementRepository;
|
||||
use super::{MysqlSettlementRepository, USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL};
|
||||
use crate::run_migrations;
|
||||
use aether_data_contracts::repository::settlement::{
|
||||
SettlementWriteRepository, UsageSettlementInput,
|
||||
ReserveUsagePolicyRequestInput, ReserveUsagePolicyRequestOutcome,
|
||||
SettlementWriteRepository, UsagePolicyRequestWindow, UsageSettlementInput,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
@@ -736,6 +1376,89 @@ mod tests {
|
||||
let _repository = MysqlSettlementRepository::new(pool);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_admission_transactions_use_read_committed() {
|
||||
assert_eq!(
|
||||
USAGE_POLICY_REQUEST_TRANSACTION_ISOLATION_MYSQL_SQL,
|
||||
"SET TRANSACTION ISOLATION LEVEL READ COMMITTED"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn cross_subject_same_token_is_allowed_once_without_deadlock_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
else {
|
||||
eprintln!(
|
||||
"skipping mysql request admission race test because AETHER_TEST_MYSQL_URL is unset"
|
||||
);
|
||||
return;
|
||||
};
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new()
|
||||
.max_connections(2)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("mysql pool should connect");
|
||||
run_migrations(&pool)
|
||||
.await
|
||||
.expect("mysql migrations should run");
|
||||
cleanup_request_admission_rows(&pool).await;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO users (id, username, auth_source, created_at, updated_at)
|
||||
VALUES
|
||||
('admission-race-user-1', 'admission-race-user-1', 'local', 1, 1),
|
||||
('admission-race-user-2', 'admission-race-user-2', 'local', 1, 1)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("race users should seed");
|
||||
|
||||
let repository = MysqlSettlementRepository::new(pool.clone());
|
||||
let reserve = |request_id: &str, subject_id: &str| ReserveUsagePolicyRequestInput {
|
||||
request_id: request_id.to_string(),
|
||||
subject_id: subject_id.to_string(),
|
||||
event_token: "admission-race-token".to_string(),
|
||||
admitted_at_unix_secs: 100,
|
||||
retain_until_unix_secs: 1_000,
|
||||
windows: vec![UsagePolicyRequestWindow {
|
||||
starts_at_unix_secs: 0,
|
||||
ends_at_unix_secs: 1_000,
|
||||
limit_requests: 10,
|
||||
}],
|
||||
};
|
||||
let (first, second) = tokio::join!(
|
||||
repository.reserve_usage_policy_request(reserve(
|
||||
"admission-race-request-1",
|
||||
"admission-race-user-1"
|
||||
)),
|
||||
repository.reserve_usage_policy_request(reserve(
|
||||
"admission-race-request-2",
|
||||
"admission-race-user-2"
|
||||
))
|
||||
);
|
||||
let mut outcomes = vec![
|
||||
first.expect("first reserve should not deadlock"),
|
||||
second.expect("second reserve should not deadlock"),
|
||||
];
|
||||
outcomes.sort_by_key(|outcome| match outcome {
|
||||
ReserveUsagePolicyRequestOutcome::Allowed => 0,
|
||||
ReserveUsagePolicyRequestOutcome::Conflict => 1,
|
||||
_ => 2,
|
||||
});
|
||||
assert_eq!(
|
||||
outcomes,
|
||||
vec![
|
||||
ReserveUsagePolicyRequestOutcome::Allowed,
|
||||
ReserveUsagePolicyRequestOutcome::Conflict,
|
||||
]
|
||||
);
|
||||
|
||||
cleanup_request_admission_rows(&pool).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_repository_settles_once_and_enqueues_provider_delta_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
@@ -863,4 +1586,19 @@ WHERE request_id = 'settlement-request-1'
|
||||
.expect("settlement cleanup should succeed");
|
||||
}
|
||||
}
|
||||
|
||||
async fn cleanup_request_admission_rows(pool: &sqlx::MySqlPool) {
|
||||
sqlx::query(
|
||||
"DELETE FROM usage_request_admissions WHERE event_token = 'admission-race-token'",
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("admission race row cleanup should succeed");
|
||||
sqlx::query(
|
||||
"DELETE FROM users WHERE id IN ('admission-race-user-1', 'admission-race-user-2')",
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("admission race user cleanup should succeed");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,7 +6,9 @@ use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::usage::{
|
||||
strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure,
|
||||
sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence,
|
||||
sanitize_usage_request_metadata, usage_can_recover_terminal_failure,
|
||||
usage_error_category_for_status_code, usage_lifecycle_update_allowed,
|
||||
usage_request_metadata_client_family, PendingUsageCleanupSummary, StoredRequestUsageAudit,
|
||||
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardSummary,
|
||||
StoredUsageUserTotals, UpsertUsageRecord, UsageCleanupExecutionMode, UsageCleanupPreviewCounts,
|
||||
@@ -437,7 +439,6 @@ ON DUPLICATE KEY UPDATE
|
||||
const SELECT_STALE_PENDING_USAGE_BATCH_SQL: &str = r#"
|
||||
SELECT
|
||||
`usage`.request_id,
|
||||
`usage`.status,
|
||||
COALESCE(usage_settlement_snapshots.billing_status, `usage`.billing_status) AS billing_status
|
||||
FROM `usage`
|
||||
LEFT JOIN usage_settlement_snapshots
|
||||
@@ -1046,11 +1047,29 @@ impl UsageWriteRepository for MysqlUsageWriteRepository {
|
||||
&self,
|
||||
usage: UpsertUsageRecord,
|
||||
) -> Result<StoredRequestUsageAudit, DataLayerError> {
|
||||
let mut usage = strip_deprecated_usage_display_fields(usage);
|
||||
usage.validate()?;
|
||||
let prepared_capture = http_capture::prepare_usage_http_capture(&mut usage)?;
|
||||
// Auxiliary tables may receive only clear tombstones, never request or response content.
|
||||
let capture_usage = usage.clone();
|
||||
let mut usage = sanitize_usage_for_persistence(usage);
|
||||
usage.validate()?;
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
let existing = counters::lock_and_load_usage(&mut tx, &usage.request_id).await?;
|
||||
if let Some(existing) = existing.as_ref() {
|
||||
if !usage_lifecycle_update_allowed(
|
||||
&existing.status,
|
||||
&existing.billing_status,
|
||||
existing.updated_at_unix_secs,
|
||||
existing.finalized_at_unix_secs,
|
||||
&usage.status,
|
||||
&usage.billing_status,
|
||||
usage.updated_at_unix_secs,
|
||||
usage.finalized_at_unix_secs,
|
||||
) {
|
||||
let existing = existing.clone();
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return http_capture::hydrate_usage_body_refs(&self.pool, existing).await;
|
||||
}
|
||||
}
|
||||
let recovers_terminal_failure = existing.as_ref().is_some_and(|existing| {
|
||||
usage_can_recover_terminal_failure(
|
||||
&existing.status,
|
||||
@@ -1069,13 +1088,18 @@ impl UsageWriteRepository for MysqlUsageWriteRepository {
|
||||
}
|
||||
}
|
||||
|
||||
let mut capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
|
||||
let prepared_capture = http_capture::prepare_usage_http_capture(&mut capture_usage)?;
|
||||
let capture_update_allowed = recovers_terminal_failure
|
||||
|| http_capture::capture_update_allowed(existing.as_ref(), &usage.status);
|
||||
if capture_update_allowed {
|
||||
http_capture::apply_previous_metadata_tombstones(&mut usage, existing.as_ref());
|
||||
http_capture::apply_previous_metadata_tombstones(&mut capture_usage, existing.as_ref());
|
||||
usage.request_metadata =
|
||||
sanitize_usage_request_metadata(capture_usage.request_metadata.clone());
|
||||
}
|
||||
let prepared_snapshots = capture_update_allowed
|
||||
.then(|| snapshots::from_usage(&usage))
|
||||
// The control projection preserves safe typed routing and allow-listed billing facts.
|
||||
.then(|| snapshots::from_usage(&capture_usage))
|
||||
.transpose()?;
|
||||
bind_upsert(sqlx::query(UPSERT_USAGE_SQL), &usage)?
|
||||
.execute(&mut *tx)
|
||||
@@ -1230,7 +1254,7 @@ SET provider_api_keys.request_count = aggregated.request_count,
|
||||
&self,
|
||||
cutoff_unix_secs: u64,
|
||||
now_unix_secs: u64,
|
||||
timeout_minutes: u64,
|
||||
_timeout_minutes: u64,
|
||||
batch_size: usize,
|
||||
) -> Result<PendingUsageCleanupSummary, DataLayerError> {
|
||||
if batch_size == 0 {
|
||||
@@ -1264,7 +1288,6 @@ SET provider_api_keys.request_count = aggregated.request_count,
|
||||
.map(|row| {
|
||||
Ok(StalePendingUsageRow {
|
||||
request_id: row.try_get("request_id").map_sql_err()?,
|
||||
status: row.try_get("status").map_sql_err()?,
|
||||
billing_status: row.try_get("billing_status").map_sql_err()?,
|
||||
})
|
||||
})
|
||||
@@ -1280,7 +1303,8 @@ SET provider_api_keys.request_count = aggregated.request_count,
|
||||
UPDATE `usage`
|
||||
SET status = 'completed',
|
||||
status_code = 200,
|
||||
error_message = NULL
|
||||
error_message = NULL,
|
||||
error_category = NULL
|
||||
WHERE request_id = ?
|
||||
"#,
|
||||
)
|
||||
@@ -1308,11 +1332,8 @@ WHERE request_id = ?
|
||||
|
||||
let candidate_info =
|
||||
latest_failed_candidate_mysql(&mut tx, &row.request_id).await?;
|
||||
let (status_code, error_message) = resolve_stale_pending_failure(
|
||||
candidate_info.as_ref(),
|
||||
&row.status,
|
||||
timeout_minutes,
|
||||
);
|
||||
let status_code = resolve_stale_pending_status_code(candidate_info.as_ref());
|
||||
let error_category = usage_error_category_for_status_code(status_code);
|
||||
let status_code_i64 = i64::from(status_code);
|
||||
if row.billing_status == "pending" {
|
||||
sqlx::query(
|
||||
@@ -1320,7 +1341,8 @@ WHERE request_id = ?
|
||||
UPDATE `usage`
|
||||
SET status = 'failed',
|
||||
status_code = ?,
|
||||
error_message = ?,
|
||||
error_message = NULL,
|
||||
error_category = ?,
|
||||
billing_status = 'void',
|
||||
finalized_at = ?,
|
||||
total_cost_usd = 0,
|
||||
@@ -1329,7 +1351,7 @@ WHERE request_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(status_code_i64)
|
||||
.bind(&error_message)
|
||||
.bind(error_category)
|
||||
.bind(to_i64(now_unix_secs, "usage finalized_at")?)
|
||||
.bind(&row.request_id)
|
||||
.execute(&mut *tx)
|
||||
@@ -1347,12 +1369,13 @@ WHERE request_id = ?
|
||||
UPDATE `usage`
|
||||
SET status = 'failed',
|
||||
status_code = ?,
|
||||
error_message = ?
|
||||
error_message = NULL,
|
||||
error_category = ?
|
||||
WHERE request_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(status_code_i64)
|
||||
.bind(&error_message)
|
||||
.bind(error_category)
|
||||
.bind(&row.request_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
@@ -1364,7 +1387,8 @@ WHERE request_id = ?
|
||||
UPDATE request_candidates
|
||||
SET status = 'failed',
|
||||
finished_at = ?,
|
||||
error_message = '请求超时(服务器可能已重启)'
|
||||
error_type = 'internal',
|
||||
error_message = NULL
|
||||
WHERE request_id = ?
|
||||
AND status IN ('pending', 'streaming')
|
||||
"#,
|
||||
@@ -1451,7 +1475,6 @@ WHERE request_id = ?
|
||||
|
||||
struct StalePendingUsageRow {
|
||||
request_id: String,
|
||||
status: String,
|
||||
billing_status: String,
|
||||
}
|
||||
|
||||
@@ -1561,29 +1584,14 @@ ON DUPLICATE KEY UPDATE
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn stale_pending_error_message(status: &str, timeout_minutes: u64) -> String {
|
||||
format!("请求超时: 状态 '{status}' 超过 {timeout_minutes} 分钟未完成")
|
||||
}
|
||||
|
||||
struct FailedCandidateCleanupInfo {
|
||||
status_code: Option<u16>,
|
||||
error_message: Option<String>,
|
||||
}
|
||||
|
||||
fn resolve_stale_pending_failure(
|
||||
candidate: Option<&FailedCandidateCleanupInfo>,
|
||||
status: &str,
|
||||
timeout_minutes: u64,
|
||||
) -> (u16, String) {
|
||||
match candidate {
|
||||
Some(info) => (
|
||||
info.status_code.unwrap_or(502),
|
||||
info.error_message
|
||||
.clone()
|
||||
.unwrap_or_else(|| stale_pending_error_message(status, timeout_minutes)),
|
||||
),
|
||||
None => (504, stale_pending_error_message(status, timeout_minutes)),
|
||||
}
|
||||
fn resolve_stale_pending_status_code(candidate: Option<&FailedCandidateCleanupInfo>) -> u16 {
|
||||
candidate
|
||||
.and_then(|info| info.status_code)
|
||||
.unwrap_or(if candidate.is_some() { 502 } else { 504 })
|
||||
}
|
||||
|
||||
async fn latest_failed_candidate_mysql(
|
||||
@@ -1592,7 +1600,7 @@ async fn latest_failed_candidate_mysql(
|
||||
) -> Result<Option<FailedCandidateCleanupInfo>, DataLayerError> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT status_code, error_message
|
||||
SELECT status_code
|
||||
FROM request_candidates
|
||||
WHERE request_id = ?
|
||||
AND status IN ('failed', 'cancelled')
|
||||
@@ -1615,15 +1623,7 @@ LIMIT 1
|
||||
.try_get::<Option<i64>, _>("status_code")
|
||||
.map_sql_err()?
|
||||
.and_then(|value| u16::try_from(value).ok());
|
||||
let error_message = row
|
||||
.try_get::<Option<String>, _>("error_message")
|
||||
.map_sql_err()?
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
Ok(Some(FailedCandidateCleanupInfo {
|
||||
status_code,
|
||||
error_message,
|
||||
}))
|
||||
Ok(Some(FailedCandidateCleanupInfo { status_code }))
|
||||
}
|
||||
|
||||
fn bind_upsert<'q>(
|
||||
|
||||
@@ -1,11 +1,8 @@
|
||||
use std::io::Write;
|
||||
|
||||
use aether_data_contracts::repository::usage::{
|
||||
parse_usage_body_ref, usage_body_ref, UsageBodyField, UsageCleanupExecutionMode,
|
||||
UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow,
|
||||
UsageCleanupExecutionMode, UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets,
|
||||
UsageCleanupWindow,
|
||||
};
|
||||
use chrono::{DateTime, Utc};
|
||||
use flate2::{write::GzEncoder, Compression};
|
||||
use serde_json::Value;
|
||||
use sqlx::Row;
|
||||
use tracing::warn;
|
||||
@@ -66,17 +63,6 @@ OR EXISTS (
|
||||
)
|
||||
"#;
|
||||
|
||||
const INLINE_OR_COMPRESSED_BODY_PREDICATE: &str = r#"
|
||||
request_body IS NOT NULL
|
||||
OR response_body IS NOT NULL
|
||||
OR provider_request_body IS NOT NULL
|
||||
OR client_response_body IS NOT NULL
|
||||
OR request_body_compressed IS NOT NULL
|
||||
OR response_body_compressed IS NOT NULL
|
||||
OR provider_request_body_compressed IS NOT NULL
|
||||
OR client_response_body_compressed IS NOT NULL
|
||||
"#;
|
||||
|
||||
const HEADER_PREDICATE: &str = r#"
|
||||
request_headers IS NOT NULL
|
||||
OR response_headers IS NOT NULL
|
||||
@@ -116,6 +102,20 @@ OR request_body_compressed IS NOT NULL
|
||||
OR response_body_compressed IS NOT NULL
|
||||
OR provider_request_body_compressed IS NOT NULL
|
||||
OR client_response_body_compressed IS NOT NULL
|
||||
OR EXISTS (
|
||||
SELECT 1 FROM usage_body_blobs
|
||||
WHERE usage_body_blobs.request_id = `usage`.request_id
|
||||
)
|
||||
OR EXISTS (
|
||||
SELECT 1 FROM usage_http_audits
|
||||
WHERE usage_http_audits.request_id = `usage`.request_id
|
||||
AND (
|
||||
usage_http_audits.request_body_ref IS NOT NULL
|
||||
OR usage_http_audits.provider_request_body_ref IS NOT NULL
|
||||
OR usage_http_audits.response_body_ref IS NOT NULL
|
||||
OR usage_http_audits.client_response_body_ref IS NOT NULL
|
||||
)
|
||||
)
|
||||
OR (
|
||||
request_metadata IS NOT NULL
|
||||
AND JSON_VALID(request_metadata)
|
||||
@@ -136,43 +136,6 @@ struct CleanupRow {
|
||||
request_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct BodyRow {
|
||||
id: String,
|
||||
request_id: String,
|
||||
request_body: Option<Value>,
|
||||
request_body_compressed: Option<Vec<u8>>,
|
||||
provider_request_body: Option<Value>,
|
||||
provider_request_body_compressed: Option<Vec<u8>>,
|
||||
response_body: Option<Value>,
|
||||
response_body_compressed: Option<Vec<u8>>,
|
||||
client_response_body: Option<Value>,
|
||||
client_response_body_compressed: Option<Vec<u8>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct DetachedRefs {
|
||||
request_body_ref: Option<String>,
|
||||
provider_request_body_ref: Option<String>,
|
||||
response_body_ref: Option<String>,
|
||||
client_response_body_ref: Option<String>,
|
||||
}
|
||||
|
||||
impl DetachedRefs {
|
||||
fn any_present(&self) -> bool {
|
||||
self.request_body_ref.is_some()
|
||||
|| self.provider_request_body_ref.is_some()
|
||||
|| self.response_body_ref.is_some()
|
||||
|| self.client_response_body_ref.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
struct DetachedBlob {
|
||||
body_ref: String,
|
||||
body_field: &'static str,
|
||||
payload_gzip: Vec<u8>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum BodyCleanupKind {
|
||||
Raw,
|
||||
@@ -188,10 +151,6 @@ impl BodyCleanupKind {
|
||||
Self::All => ALL_BODY_PREDICATE,
|
||||
}
|
||||
}
|
||||
|
||||
fn clears_detached(self) -> bool {
|
||||
self != Self::Raw
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_usage(
|
||||
@@ -263,12 +222,19 @@ pub(crate) async fn cleanup_usage(
|
||||
};
|
||||
let detail_newer_than = detail_body_newer_than(window, targets);
|
||||
let legacy_body_refs_migrated = if targets.detail_body {
|
||||
migrate_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await?
|
||||
purge_legacy_body_refs(pool, window.detail_cutoff, detail_newer_than, batch_size).await?
|
||||
} else {
|
||||
0
|
||||
};
|
||||
let body_externalized = if targets.detail_body {
|
||||
externalize_detail_bodies(pool, window.detail_cutoff, detail_newer_than, batch_size).await?
|
||||
cleanup_body_fields(
|
||||
pool,
|
||||
window.detail_cutoff,
|
||||
detail_newer_than,
|
||||
batch_size,
|
||||
BodyCleanupKind::All,
|
||||
)
|
||||
.await?
|
||||
} else {
|
||||
0
|
||||
};
|
||||
@@ -291,6 +257,8 @@ pub(crate) async fn cleanup_usage(
|
||||
header_cleaned,
|
||||
keys_cleaned,
|
||||
records_deleted,
|
||||
cost_reservations_deleted: 0,
|
||||
request_admissions_deleted: 0,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -611,14 +579,13 @@ WHERE id = ?
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if kind.clears_detached() {
|
||||
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
|
||||
.bind(&row.request_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
|
||||
.bind(&row.request_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE usage_http_audits
|
||||
SET request_body_ref = NULL,
|
||||
provider_request_body_ref = NULL,
|
||||
@@ -628,13 +595,12 @@ SET request_body_ref = NULL,
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE request_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&row.request_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
delete_empty_http_audit(&mut tx, &row.request_id).await?;
|
||||
}
|
||||
)
|
||||
.bind(&row.request_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
delete_empty_http_audit(&mut tx, &row.request_id).await?;
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
total = total.saturating_add(row_count);
|
||||
@@ -670,14 +636,14 @@ WHERE request_id = ?
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn migrate_legacy_body_refs(
|
||||
async fn purge_legacy_body_refs(
|
||||
pool: &MysqlPool,
|
||||
cutoff: DateTime<Utc>,
|
||||
newer_than: Option<DateTime<Utc>>,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if invalid_window(cutoff, newer_than) {
|
||||
warn!(%cutoff, ?newer_than, "MySQL usage legacy body-ref migration skipped due to invalid window");
|
||||
warn!(%cutoff, ?newer_than, "MySQL usage legacy body-ref purge skipped due to invalid window");
|
||||
return Ok(0);
|
||||
}
|
||||
let mut total = 0usize;
|
||||
@@ -695,7 +661,7 @@ async fn migrate_legacy_body_refs(
|
||||
}
|
||||
let row_count = rows.len();
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let mut migrated = 0usize;
|
||||
let mut purged = 0usize;
|
||||
for row in rows {
|
||||
let metadata: Option<String> =
|
||||
sqlx::query_scalar("SELECT request_metadata FROM `usage` WHERE id = ? LIMIT 1")
|
||||
@@ -704,14 +670,9 @@ async fn migrate_legacy_body_refs(
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.flatten();
|
||||
let Some((refs, metadata)) =
|
||||
legacy_body_ref_plan(&row.request_id, metadata.as_deref())?
|
||||
else {
|
||||
let Some(metadata) = legacy_body_ref_purge_plan(metadata.as_deref())? else {
|
||||
continue;
|
||||
};
|
||||
if refs.any_present() {
|
||||
upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?;
|
||||
}
|
||||
let updated = sqlx::query(
|
||||
r#"
|
||||
UPDATE `usage`
|
||||
@@ -726,23 +687,23 @@ WHERE id = ?
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
purge_detached_body_capture(&mut tx, &row.request_id).await?;
|
||||
if updated > 0 {
|
||||
migrated += 1;
|
||||
purged += 1;
|
||||
}
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
total = total.saturating_add(migrated);
|
||||
if row_count < batch_size || migrated == 0 {
|
||||
total = total.saturating_add(purged);
|
||||
if row_count < batch_size || purged == 0 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(total)
|
||||
}
|
||||
|
||||
fn legacy_body_ref_plan(
|
||||
request_id: &str,
|
||||
fn legacy_body_ref_purge_plan(
|
||||
metadata: Option<&str>,
|
||||
) -> Result<Option<(DetachedRefs, Option<String>)>, DataLayerError> {
|
||||
) -> Result<Option<Option<String>>, DataLayerError> {
|
||||
let Some(metadata) = metadata else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -752,30 +713,16 @@ fn legacy_body_ref_plan(
|
||||
let Value::Object(mut object) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
let mut refs = DetachedRefs::default();
|
||||
let mut removed = false;
|
||||
for field in [
|
||||
UsageBodyField::RequestBody,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
UsageBodyField::ResponseBody,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
for key in [
|
||||
"request_body_ref",
|
||||
"provider_request_body_ref",
|
||||
"response_body_ref",
|
||||
"client_response_body_ref",
|
||||
] {
|
||||
let Some(value) = object.remove(field.as_ref_key()) else {
|
||||
continue;
|
||||
};
|
||||
removed = true;
|
||||
let parsed = value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(parse_usage_body_ref)
|
||||
.filter(|(parsed_request_id, parsed_field)| {
|
||||
parsed_request_id == request_id && *parsed_field == field
|
||||
})
|
||||
.map(|(parsed_request_id, parsed_field)| {
|
||||
usage_body_ref(&parsed_request_id, parsed_field)
|
||||
});
|
||||
set_ref(&mut refs, field, parsed);
|
||||
if object.remove(key).is_some() {
|
||||
removed = true;
|
||||
}
|
||||
}
|
||||
if !removed {
|
||||
return Ok(None);
|
||||
@@ -791,277 +738,35 @@ fn legacy_body_ref_plan(
|
||||
})?,
|
||||
)
|
||||
};
|
||||
Ok(Some((refs, metadata)))
|
||||
Ok(Some(metadata))
|
||||
}
|
||||
|
||||
async fn externalize_detail_bodies(
|
||||
pool: &MysqlPool,
|
||||
cutoff: DateTime<Utc>,
|
||||
newer_than: Option<DateTime<Utc>>,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if invalid_window(cutoff, newer_than) {
|
||||
warn!(%cutoff, ?newer_than, "MySQL usage body externalization skipped due to invalid window");
|
||||
return Ok(0);
|
||||
}
|
||||
let batch_size = batch_size.clamp(1, 25);
|
||||
let mut total = 0usize;
|
||||
loop {
|
||||
let rows = fetch_body_rows(pool, cutoff, newer_than, batch_size).await?;
|
||||
if rows.is_empty() {
|
||||
break;
|
||||
}
|
||||
let row_count = rows.len();
|
||||
let mut externalized = 0usize;
|
||||
for row in rows {
|
||||
let (blobs, refs) = build_detached_bodies(&row)?;
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
for blob in blobs {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip)
|
||||
VALUES (?, ?, ?, ?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
request_id = VALUES(request_id),
|
||||
body_field = VALUES(body_field),
|
||||
payload_gzip = VALUES(payload_gzip),
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
"#,
|
||||
)
|
||||
.bind(blob.body_ref)
|
||||
.bind(&row.request_id)
|
||||
.bind(blob.body_field)
|
||||
.bind(blob.payload_gzip)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
if refs.any_present() {
|
||||
upsert_http_audit_refs(&mut tx, &row.request_id, &refs).await?;
|
||||
}
|
||||
let updated = sqlx::query(
|
||||
r#"
|
||||
UPDATE `usage`
|
||||
SET request_body = NULL,
|
||||
response_body = NULL,
|
||||
provider_request_body = NULL,
|
||||
client_response_body = NULL,
|
||||
request_body_compressed = NULL,
|
||||
response_body_compressed = NULL,
|
||||
provider_request_body_compressed = NULL,
|
||||
client_response_body_compressed = NULL
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(row.id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
tx.commit().await.map_sql_err()?;
|
||||
if updated > 0 {
|
||||
externalized += 1;
|
||||
}
|
||||
}
|
||||
total = total.saturating_add(externalized);
|
||||
if row_count < batch_size || externalized == 0 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(total)
|
||||
}
|
||||
|
||||
async fn fetch_body_rows(
|
||||
pool: &MysqlPool,
|
||||
cutoff: DateTime<Utc>,
|
||||
newer_than: Option<DateTime<Utc>>,
|
||||
batch_size: usize,
|
||||
) -> Result<Vec<BodyRow>, DataLayerError> {
|
||||
let newer_than = newer_than.map(|value| value.timestamp());
|
||||
let sql = format!(
|
||||
r#"
|
||||
SELECT id,
|
||||
request_id,
|
||||
CAST(request_body AS CHAR) AS request_body,
|
||||
request_body_compressed,
|
||||
CAST(provider_request_body AS CHAR) AS provider_request_body,
|
||||
provider_request_body_compressed,
|
||||
CAST(response_body AS CHAR) AS response_body,
|
||||
response_body_compressed,
|
||||
CAST(client_response_body AS CHAR) AS client_response_body,
|
||||
client_response_body_compressed
|
||||
FROM `usage`
|
||||
WHERE created_at_unix_ms < ?
|
||||
AND (? IS NULL OR created_at_unix_ms >= ?)
|
||||
AND ({INLINE_OR_COMPRESSED_BODY_PREDICATE})
|
||||
ORDER BY created_at_unix_ms ASC, id ASC
|
||||
LIMIT ?
|
||||
"#
|
||||
);
|
||||
sqlx::query(&sql)
|
||||
.bind(cutoff.timestamp())
|
||||
.bind(newer_than)
|
||||
.bind(newer_than)
|
||||
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
Ok(BodyRow {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
request_id: row.try_get("request_id").map_sql_err()?,
|
||||
request_body: parse_optional_json(row.try_get("request_body").map_sql_err()?)?,
|
||||
request_body_compressed: row.try_get("request_body_compressed").map_sql_err()?,
|
||||
provider_request_body: parse_optional_json(
|
||||
row.try_get("provider_request_body").map_sql_err()?,
|
||||
)?,
|
||||
provider_request_body_compressed: row
|
||||
.try_get("provider_request_body_compressed")
|
||||
.map_sql_err()?,
|
||||
response_body: parse_optional_json(row.try_get("response_body").map_sql_err()?)?,
|
||||
response_body_compressed: row.try_get("response_body_compressed").map_sql_err()?,
|
||||
client_response_body: parse_optional_json(
|
||||
row.try_get("client_response_body").map_sql_err()?,
|
||||
)?,
|
||||
client_response_body_compressed: row
|
||||
.try_get("client_response_body_compressed")
|
||||
.map_sql_err()?,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn parse_optional_json(raw: Option<String>) -> Result<Option<Value>, DataLayerError> {
|
||||
raw.map(|raw| {
|
||||
serde_json::from_str(&raw).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("invalid inline usage body JSON: {err}"))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn build_detached_bodies(
|
||||
row: &BodyRow,
|
||||
) -> Result<(Vec<DetachedBlob>, DetachedRefs), DataLayerError> {
|
||||
let mut blobs = Vec::new();
|
||||
let mut refs = DetachedRefs::default();
|
||||
add_detached_body(
|
||||
&mut blobs,
|
||||
&mut refs,
|
||||
&row.request_id,
|
||||
UsageBodyField::RequestBody,
|
||||
row.request_body.as_ref(),
|
||||
row.request_body_compressed.as_deref(),
|
||||
)?;
|
||||
add_detached_body(
|
||||
&mut blobs,
|
||||
&mut refs,
|
||||
&row.request_id,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
row.provider_request_body.as_ref(),
|
||||
row.provider_request_body_compressed.as_deref(),
|
||||
)?;
|
||||
add_detached_body(
|
||||
&mut blobs,
|
||||
&mut refs,
|
||||
&row.request_id,
|
||||
UsageBodyField::ResponseBody,
|
||||
row.response_body.as_ref(),
|
||||
row.response_body_compressed.as_deref(),
|
||||
)?;
|
||||
add_detached_body(
|
||||
&mut blobs,
|
||||
&mut refs,
|
||||
&row.request_id,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
row.client_response_body.as_ref(),
|
||||
row.client_response_body_compressed.as_deref(),
|
||||
)?;
|
||||
Ok((blobs, refs))
|
||||
}
|
||||
|
||||
fn add_detached_body(
|
||||
blobs: &mut Vec<DetachedBlob>,
|
||||
refs: &mut DetachedRefs,
|
||||
request_id: &str,
|
||||
field: UsageBodyField,
|
||||
raw: Option<&Value>,
|
||||
compressed: Option<&[u8]>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let payload_gzip = match raw {
|
||||
Some(value) => Some(compress_json(value)?),
|
||||
None => compressed.map(ToOwned::to_owned),
|
||||
};
|
||||
let Some(payload_gzip) = payload_gzip else {
|
||||
return Ok(());
|
||||
};
|
||||
let body_ref = usage_body_ref(request_id, field);
|
||||
blobs.push(DetachedBlob {
|
||||
body_ref: body_ref.clone(),
|
||||
body_field: field.as_storage_field(),
|
||||
payload_gzip,
|
||||
});
|
||||
set_ref(refs, field, Some(body_ref));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn compress_json(value: &Value) -> Result<Vec<u8>, DataLayerError> {
|
||||
let bytes = serde_json::to_vec(value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}"))
|
||||
})?;
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6));
|
||||
encoder.write_all(&bytes).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}"))
|
||||
})?;
|
||||
encoder.finish().map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}"))
|
||||
})
|
||||
}
|
||||
|
||||
fn set_ref(refs: &mut DetachedRefs, field: UsageBodyField, value: Option<String>) {
|
||||
match field {
|
||||
UsageBodyField::RequestBody => refs.request_body_ref = value,
|
||||
UsageBodyField::ProviderRequestBody => refs.provider_request_body_ref = value,
|
||||
UsageBodyField::ResponseBody => refs.response_body_ref = value,
|
||||
UsageBodyField::ClientResponseBody => refs.client_response_body_ref = value,
|
||||
}
|
||||
}
|
||||
|
||||
async fn upsert_http_audit_refs(
|
||||
async fn purge_detached_body_capture(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
request_id: &str,
|
||||
refs: &DetachedRefs,
|
||||
) -> Result<(), DataLayerError> {
|
||||
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
|
||||
.bind(request_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO usage_http_audits (
|
||||
request_id,
|
||||
request_body_ref,
|
||||
provider_request_body_ref,
|
||||
response_body_ref,
|
||||
client_response_body_ref,
|
||||
body_capture_mode
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, 'ref_backed')
|
||||
ON DUPLICATE KEY UPDATE
|
||||
request_body_ref = COALESCE(VALUES(request_body_ref), request_body_ref),
|
||||
provider_request_body_ref = COALESCE(VALUES(provider_request_body_ref), provider_request_body_ref),
|
||||
response_body_ref = COALESCE(VALUES(response_body_ref), response_body_ref),
|
||||
client_response_body_ref = COALESCE(VALUES(client_response_body_ref), client_response_body_ref),
|
||||
body_capture_mode = 'ref_backed',
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
UPDATE usage_http_audits
|
||||
SET request_body_ref = NULL,
|
||||
provider_request_body_ref = NULL,
|
||||
response_body_ref = NULL,
|
||||
client_response_body_ref = NULL,
|
||||
body_capture_mode = 'none',
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
WHERE request_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(request_id)
|
||||
.bind(refs.request_body_ref.as_deref())
|
||||
.bind(refs.provider_request_body_ref.as_deref())
|
||||
.bind(refs.response_body_ref.as_deref())
|
||||
.bind(refs.client_response_body_ref.as_deref())
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(())
|
||||
delete_empty_http_audit(tx, request_id).await
|
||||
}
|
||||
|
||||
async fn cleanup_expired_api_keys(
|
||||
@@ -1127,29 +832,21 @@ ORDER BY expires_at ASC, id ASC
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io::Read;
|
||||
|
||||
use flate2::read::GzDecoder;
|
||||
use serde_json::json;
|
||||
|
||||
use super::{compress_json, legacy_body_ref_plan};
|
||||
use super::{legacy_body_ref_purge_plan, DETAIL_BODY_PREDICATE};
|
||||
|
||||
#[test]
|
||||
fn mysql_cleanup_legacy_body_ref_plan_preserves_unrelated_metadata() {
|
||||
fn mysql_cleanup_legacy_body_ref_purge_preserves_unrelated_metadata() {
|
||||
let metadata = json!({
|
||||
"trace": "kept",
|
||||
"request_body_ref": "usage://request/request-1/request_body",
|
||||
"response_body_ref": "usage://request/other/response_body"
|
||||
})
|
||||
.to_string();
|
||||
let (refs, metadata) = legacy_body_ref_plan("request-1", Some(&metadata))
|
||||
let metadata = legacy_body_ref_purge_plan(Some(&metadata))
|
||||
.expect("legacy plan should build")
|
||||
.expect("legacy refs should be present");
|
||||
assert_eq!(
|
||||
refs.request_body_ref.as_deref(),
|
||||
Some("usage://request/request-1/request_body")
|
||||
);
|
||||
assert!(refs.response_body_ref.is_none());
|
||||
assert_eq!(
|
||||
serde_json::from_str::<serde_json::Value>(
|
||||
metadata.as_deref().expect("trace metadata should remain")
|
||||
@@ -1160,18 +857,9 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_cleanup_gzip_payload_round_trips() {
|
||||
let value = json!({"hello": "world"});
|
||||
let payload = compress_json(&value).expect("body should compress");
|
||||
let mut decoder = GzDecoder::new(payload.as_slice());
|
||||
let mut decoded = Vec::new();
|
||||
decoder
|
||||
.read_to_end(&mut decoded)
|
||||
.expect("body should decompress");
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<serde_json::Value>(&decoded)
|
||||
.expect("decoded body should be JSON"),
|
||||
value
|
||||
);
|
||||
fn mysql_detail_cleanup_includes_detached_capture() {
|
||||
assert!(DETAIL_BODY_PREDICATE.contains("usage_body_blobs"));
|
||||
assert!(DETAIL_BODY_PREDICATE.contains("usage_http_audits"));
|
||||
assert!(!DETAIL_BODY_PREDICATE.contains("payload_gzip"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,6 +24,7 @@ SELECT
|
||||
id,
|
||||
kind,
|
||||
target_id,
|
||||
target_tunnel_generation,
|
||||
request_count_delta,
|
||||
total_requests_delta,
|
||||
success_count_delta,
|
||||
@@ -49,6 +50,7 @@ struct DeltaRow {
|
||||
id: String,
|
||||
kind: String,
|
||||
target_id: String,
|
||||
target_tunnel_generation: Option<String>,
|
||||
request_count_delta: i64,
|
||||
total_requests_delta: i64,
|
||||
success_count_delta: i64,
|
||||
@@ -71,7 +73,10 @@ struct Aggregates {
|
||||
provider_api_keys: BTreeMap<String, ProviderApiKeyUsageDelta>,
|
||||
models: BTreeMap<String, ModelUsageDelta>,
|
||||
provider_monthly: BTreeMap<String, f64>,
|
||||
proxy_nodes: BTreeMap<String, ProxyNodeCounterDelta>,
|
||||
// Keep the node incarnation in the aggregation key. A node id can be
|
||||
// reused after deletion, so a bare id would route old deltas to the new
|
||||
// node.
|
||||
proxy_nodes: BTreeMap<(String, String), ProxyNodeCounterDelta>,
|
||||
management_tokens: BTreeMap<String, ManagementTokenCounterDelta>,
|
||||
api_key_last_used: BTreeMap<String, ApiKeyLastUsedDelta>,
|
||||
}
|
||||
@@ -142,16 +147,28 @@ impl Aggregates {
|
||||
.or_default() += row.total_cost_usd_delta;
|
||||
}
|
||||
KIND_PROXY_NODE => {
|
||||
let entry = aggregates
|
||||
.proxy_nodes
|
||||
.entry(row.target_id.clone())
|
||||
.or_insert(ProxyNodeCounterDelta {
|
||||
let Some(tunnel_generation) = row
|
||||
.target_tunnel_generation
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
else {
|
||||
// Legacy rows have no identity fence. Mark them
|
||||
// processed without applying them to any node.
|
||||
continue;
|
||||
};
|
||||
let aggregate_key = (row.target_id.clone(), tunnel_generation.clone());
|
||||
let entry = aggregates.proxy_nodes.entry(aggregate_key).or_insert(
|
||||
ProxyNodeCounterDelta {
|
||||
node_id: row.target_id.clone(),
|
||||
expected_tunnel_generation: Some(tunnel_generation),
|
||||
total_requests_delta: 0,
|
||||
failed_requests_delta: 0,
|
||||
dns_failures_delta: 0,
|
||||
stream_errors_delta: 0,
|
||||
});
|
||||
},
|
||||
);
|
||||
entry.total_requests_delta += row.total_requests_delta;
|
||||
entry.failed_requests_delta += row.error_count_delta;
|
||||
entry.dns_failures_delta += row.dns_failures_delta;
|
||||
@@ -241,8 +258,8 @@ pub(super) async fn flush(
|
||||
for (target_id, delta) in &aggregates.provider_monthly {
|
||||
apply_provider_monthly(&mut tx, target_id, *delta).await?;
|
||||
}
|
||||
for (target_id, delta) in &aggregates.proxy_nodes {
|
||||
apply_proxy_node(&mut tx, target_id, delta).await?;
|
||||
for ((target_id, tunnel_generation), delta) in &aggregates.proxy_nodes {
|
||||
apply_proxy_node(&mut tx, target_id, tunnel_generation, delta).await?;
|
||||
}
|
||||
for (target_id, delta) in &aggregates.management_tokens {
|
||||
apply_management_token(&mut tx, target_id, delta).await?;
|
||||
@@ -283,9 +300,42 @@ pub(super) async fn enqueue_proxy_node(
|
||||
if delta.is_noop() {
|
||||
return Ok(false);
|
||||
}
|
||||
let Some(expected_tunnel_generation) = delta
|
||||
.expected_tunnel_generation
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.filter(|value| value.len() <= 64)
|
||||
.map(ToOwned::to_owned)
|
||||
else {
|
||||
// A bare id is not an identity fence. Reject it instead of rebinding
|
||||
// the delta to whichever incarnation currently owns that id.
|
||||
return Ok(false);
|
||||
};
|
||||
let node_id = delta.node_id.trim().to_string();
|
||||
let request_id = format!("proxy_node:{node_id}:{}", uuid::Uuid::new_v4());
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
// Keep the parent lookup lock-free because flush claims outbox rows before
|
||||
// updating proxy_nodes. The generation is stored in the outbox row and is
|
||||
// checked again by flush, so a concurrent id reuse can only discard this
|
||||
// delta, never apply it to the replacement row.
|
||||
let tunnel_generation: Option<String> = sqlx::query_scalar(
|
||||
"SELECT tunnel_generation FROM proxy_nodes WHERE id = ? AND BINARY tunnel_generation = BINARY ? LIMIT 1",
|
||||
)
|
||||
.bind(&node_id)
|
||||
.bind(&expected_tunnel_generation)
|
||||
.fetch_optional(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(_tunnel_generation) = tunnel_generation
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
else {
|
||||
tx.rollback().await.map_sql_err()?;
|
||||
return Ok(false);
|
||||
};
|
||||
insert_delta(
|
||||
&mut tx,
|
||||
DeltaInsert {
|
||||
@@ -296,6 +346,7 @@ pub(super) async fn enqueue_proxy_node(
|
||||
error_count_delta: delta.failed_requests_delta,
|
||||
dns_failures_delta: delta.dns_failures_delta,
|
||||
stream_errors_delta: delta.stream_errors_delta,
|
||||
target_tunnel_generation: Some(&expected_tunnel_generation),
|
||||
..DeltaInsert::default()
|
||||
},
|
||||
)
|
||||
@@ -731,6 +782,7 @@ struct DeltaInsert<'a> {
|
||||
request_id: &'a str,
|
||||
kind: &'a str,
|
||||
target_id: &'a str,
|
||||
target_tunnel_generation: Option<&'a str>,
|
||||
request_count_delta: i64,
|
||||
total_requests_delta: i64,
|
||||
success_count_delta: i64,
|
||||
@@ -759,18 +811,20 @@ async fn insert_delta(
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO usage_counter_deltas (
|
||||
id, request_id, kind, target_id, request_count_delta, total_requests_delta,
|
||||
id, request_id, kind, target_id, target_tunnel_generation,
|
||||
request_count_delta, total_requests_delta,
|
||||
success_count_delta, error_count_delta, dns_failures_delta, stream_errors_delta,
|
||||
total_tokens_delta, total_cost_usd_delta, total_response_time_ms_delta,
|
||||
last_used_at_unix_secs, last_used_ip, candidate_last_used_at_unix_secs,
|
||||
removed_last_used_at_unix_secs, usage_created_at_unix_secs, created_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(uuid::Uuid::new_v4().to_string())
|
||||
.bind(request_id)
|
||||
.bind(input.kind)
|
||||
.bind(target_id)
|
||||
.bind(input.target_tunnel_generation)
|
||||
.bind(input.request_count_delta)
|
||||
.bind(input.total_requests_delta)
|
||||
.bind(input.success_count_delta)
|
||||
@@ -814,6 +868,7 @@ fn map_row(row: &sqlx::mysql::MySqlRow) -> Result<DeltaRow, DataLayerError> {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
kind: row.try_get("kind").map_sql_err()?,
|
||||
target_id: row.try_get("target_id").map_sql_err()?,
|
||||
target_tunnel_generation: row.try_get("target_tunnel_generation").map_sql_err()?,
|
||||
request_count_delta: row.try_get("request_count_delta").map_sql_err()?,
|
||||
total_requests_delta: row.try_get("total_requests_delta").map_sql_err()?,
|
||||
success_count_delta: row.try_get("success_count_delta").map_sql_err()?,
|
||||
@@ -997,9 +1052,10 @@ async fn apply_provider_monthly(
|
||||
async fn apply_proxy_node(
|
||||
tx: &mut sqlx::Transaction<'_, MySql>,
|
||||
target_id: &str,
|
||||
tunnel_generation: &str,
|
||||
delta: &ProxyNodeCounterDelta,
|
||||
) -> Result<(), DataLayerError> {
|
||||
if target_id.trim().is_empty() || delta.is_noop() {
|
||||
if target_id.trim().is_empty() || tunnel_generation.trim().is_empty() || delta.is_noop() {
|
||||
return Ok(());
|
||||
}
|
||||
sqlx::query(
|
||||
@@ -1010,7 +1066,7 @@ SET total_requests = total_requests + GREATEST(?, 0),
|
||||
dns_failures = dns_failures + GREATEST(?, 0),
|
||||
stream_errors = stream_errors + GREATEST(?, 0),
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
WHERE id = ? AND BINARY tunnel_generation = BINARY ?
|
||||
"#,
|
||||
)
|
||||
.bind(delta.total_requests_delta)
|
||||
@@ -1019,6 +1075,7 @@ WHERE id = ?
|
||||
.bind(delta.stream_errors_delta)
|
||||
.bind(current_unix_secs())
|
||||
.bind(target_id)
|
||||
.bind(tunnel_generation)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
use std::io::{Read, Write};
|
||||
|
||||
use aether_data_contracts::repository::usage::{
|
||||
parse_usage_body_ref, usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord,
|
||||
UsageBodyCaptureState, UsageBodyField,
|
||||
canonical_usage_body_ref_for, parse_usage_body_ref, read_decompressed_usage_json,
|
||||
usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState,
|
||||
UsageBodyField,
|
||||
};
|
||||
use flate2::{read::GzDecoder, write::GzEncoder, Compression};
|
||||
use flate2::read::GzDecoder;
|
||||
use serde_json::{Map, Value};
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
@@ -28,9 +27,7 @@ pub(crate) struct PreparedUsageHttpCapture {
|
||||
|
||||
#[derive(Debug)]
|
||||
struct PreparedBody {
|
||||
field: UsageBodyField,
|
||||
payload_gzip: Option<Vec<u8>>,
|
||||
clear_existing: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
@@ -58,15 +55,6 @@ struct HttpAuditStates {
|
||||
client_response_body_state: Option<UsageBodyCaptureState>,
|
||||
}
|
||||
|
||||
impl HttpAuditStates {
|
||||
fn any_present(&self) -> bool {
|
||||
self.request_body_state.is_some()
|
||||
|| self.provider_request_body_state.is_some()
|
||||
|| self.response_body_state.is_some()
|
||||
|| self.client_response_body_state.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn capture_update_allowed(
|
||||
previous: Option<&StoredRequestUsageAudit>,
|
||||
incoming_status: &str,
|
||||
@@ -141,26 +129,10 @@ pub(crate) fn prepare_usage_http_capture(
|
||||
.then_some(usage.client_response_body.as_ref())
|
||||
.flatten();
|
||||
|
||||
let request_body = prepare_body(
|
||||
UsageBodyField::RequestBody,
|
||||
request_body_value,
|
||||
clear_request,
|
||||
)?;
|
||||
let provider_request_body = prepare_body(
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
provider_request_body_value,
|
||||
clear_provider_request,
|
||||
)?;
|
||||
let response_body = prepare_body(
|
||||
UsageBodyField::ResponseBody,
|
||||
response_body_value,
|
||||
clear_response,
|
||||
)?;
|
||||
let client_response_body = prepare_body(
|
||||
UsageBodyField::ClientResponseBody,
|
||||
client_response_body_value,
|
||||
clear_client_response,
|
||||
)?;
|
||||
let request_body = prepare_body(request_body_value)?;
|
||||
let provider_request_body = prepare_body(provider_request_body_value)?;
|
||||
let response_body = prepare_body(response_body_value)?;
|
||||
let client_response_body = prepare_body(client_response_body_value)?;
|
||||
|
||||
let refs = HttpAuditRefs {
|
||||
request_body_ref: resolved_write_ref(
|
||||
@@ -276,29 +248,13 @@ pub(crate) fn prepare_usage_http_capture(
|
||||
})
|
||||
}
|
||||
|
||||
fn prepare_body(
|
||||
field: UsageBodyField,
|
||||
value: Option<&Value>,
|
||||
clear_existing: bool,
|
||||
) -> Result<PreparedBody, DataLayerError> {
|
||||
Ok(PreparedBody {
|
||||
field,
|
||||
payload_gzip: value.map(compress_json).transpose()?,
|
||||
clear_existing,
|
||||
})
|
||||
}
|
||||
|
||||
fn compress_json(value: &Value) -> Result<Vec<u8>, DataLayerError> {
|
||||
let bytes = serde_json::to_vec(value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to serialize usage body: {err}"))
|
||||
})?;
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6));
|
||||
encoder.write_all(&bytes).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to gzip usage body: {err}"))
|
||||
})?;
|
||||
encoder.finish().map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to finish usage body gzip: {err}"))
|
||||
})
|
||||
fn prepare_body(value: Option<&Value>) -> Result<PreparedBody, DataLayerError> {
|
||||
if value.is_some() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"usage body persistence is disabled".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(PreparedBody { payload_gzip: None })
|
||||
}
|
||||
|
||||
fn resolved_write_ref(
|
||||
@@ -308,9 +264,7 @@ fn resolved_write_ref(
|
||||
has_blob: bool,
|
||||
) -> Option<String> {
|
||||
explicit_ref
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field))
|
||||
.or_else(|| has_blob.then(|| usage_body_ref(request_id, field)))
|
||||
}
|
||||
|
||||
@@ -378,23 +332,63 @@ pub(crate) async fn sync_usage_http_capture(
|
||||
request_id: &str,
|
||||
prepared: &PreparedUsageHttpCapture,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for body in [
|
||||
let bodies = [
|
||||
&prepared.request_body,
|
||||
&prepared.provider_request_body,
|
||||
&prepared.response_body,
|
||||
&prepared.client_response_body,
|
||||
] {
|
||||
sync_body(tx, request_id, body).await?;
|
||||
];
|
||||
let contains_capture = prepared.request_headers.is_some()
|
||||
|| prepared.provider_request_headers.is_some()
|
||||
|| prepared.response_headers.is_some()
|
||||
|| prepared.client_response_headers.is_some()
|
||||
|| prepared.refs.any_present()
|
||||
|| bodies.iter().any(|body| body.payload_gzip.is_some())
|
||||
|| prepared.capture_mode != "none";
|
||||
if contains_capture {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"usage HTTP capture persistence is disabled".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
sqlx::query("DELETE FROM usage_http_audits WHERE request_id = ?")
|
||||
.bind(request_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query("DELETE FROM usage_body_blobs WHERE request_id = ?")
|
||||
.bind(request_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE `usage`
|
||||
SET request_headers = NULL,
|
||||
request_body = NULL,
|
||||
provider_request_headers = NULL,
|
||||
provider_request_body = NULL,
|
||||
response_headers = NULL,
|
||||
response_body = NULL,
|
||||
client_response_headers = NULL,
|
||||
client_response_body = NULL,
|
||||
request_body_compressed = NULL,
|
||||
provider_request_body_compressed = NULL,
|
||||
response_body_compressed = NULL,
|
||||
client_response_body_compressed = NULL
|
||||
WHERE request_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(request_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let headers_present = prepared.request_headers.is_some()
|
||||
|| prepared.provider_request_headers.is_some()
|
||||
|| prepared.response_headers.is_some()
|
||||
|| prepared.client_response_headers.is_some();
|
||||
if !headers_present
|
||||
&& !prepared.refs.any_present()
|
||||
&& !prepared.states.any_present()
|
||||
&& prepared.capture_mode == "none"
|
||||
{
|
||||
if !headers_present && !prepared.refs.any_present() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
@@ -512,67 +506,6 @@ ON DUPLICATE KEY UPDATE
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn sync_body(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
request_id: &str,
|
||||
body: &PreparedBody,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let body_ref = usage_body_ref(request_id, body.field);
|
||||
if body.clear_existing || body.payload_gzip.is_some() {
|
||||
sqlx::query(clear_legacy_body_sql(body.field))
|
||||
.bind(request_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
if body.clear_existing {
|
||||
sqlx::query("DELETE FROM usage_body_blobs WHERE body_ref = ?")
|
||||
.bind(body_ref)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
return Ok(());
|
||||
}
|
||||
if let Some(payload_gzip) = body.payload_gzip.as_deref() {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip)
|
||||
VALUES (?, ?, ?, ?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
request_id = VALUES(request_id),
|
||||
body_field = VALUES(body_field),
|
||||
payload_gzip = VALUES(payload_gzip),
|
||||
updated_at = UNIX_TIMESTAMP()
|
||||
"#,
|
||||
)
|
||||
.bind(body_ref)
|
||||
.bind(request_id)
|
||||
.bind(body.field.as_storage_field())
|
||||
.bind(payload_gzip)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn clear_legacy_body_sql(field: UsageBodyField) -> &'static str {
|
||||
match field {
|
||||
UsageBodyField::RequestBody => {
|
||||
"UPDATE `usage` SET request_body = NULL, request_body_compressed = NULL WHERE request_id = ?"
|
||||
}
|
||||
UsageBodyField::ProviderRequestBody => {
|
||||
"UPDATE `usage` SET provider_request_body = NULL, provider_request_body_compressed = NULL WHERE request_id = ?"
|
||||
}
|
||||
UsageBodyField::ResponseBody => {
|
||||
"UPDATE `usage` SET response_body = NULL, response_body_compressed = NULL WHERE request_id = ?"
|
||||
}
|
||||
UsageBodyField::ClientResponseBody => {
|
||||
"UPDATE `usage` SET client_response_body = NULL, client_response_body_compressed = NULL WHERE request_id = ?"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn hydrate_usage_row(
|
||||
row: &MySqlRow,
|
||||
usage: &mut StoredRequestUsageAudit,
|
||||
@@ -690,8 +623,7 @@ fn resolved_read_ref(
|
||||
has_compressed: bool,
|
||||
) -> Option<String> {
|
||||
audit_ref
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|body_ref| canonical_usage_body_ref_for(&body_ref, request_id, field))
|
||||
.or_else(|| has_compressed.then(|| usage_body_ref(request_id, field)))
|
||||
.or_else(|| metadata_body_ref(metadata, request_id, field))
|
||||
}
|
||||
@@ -704,13 +636,7 @@ fn metadata_body_ref(
|
||||
metadata
|
||||
.and_then(|metadata| metadata.get(field.as_ref_key()))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(parse_usage_body_ref)
|
||||
.filter(|(parsed_request_id, parsed_field)| {
|
||||
parsed_request_id == request_id && *parsed_field == field
|
||||
})
|
||||
.map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field))
|
||||
.and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field))
|
||||
}
|
||||
|
||||
fn optional_state(
|
||||
@@ -752,7 +678,11 @@ pub(crate) async fn hydrate_usage_body_refs(
|
||||
let Some(body_ref) = usage.body_ref(field) else {
|
||||
continue;
|
||||
};
|
||||
let value = resolve_body_ref(pool, body_ref).await?;
|
||||
let Some(body_ref) = canonical_usage_body_ref_for(body_ref, &usage.request_id, field)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let value = resolve_body_ref(pool, &body_ref).await?;
|
||||
match field {
|
||||
UsageBodyField::RequestBody => usage.request_body = value,
|
||||
UsageBodyField::ProviderRequestBody => usage.provider_request_body = value,
|
||||
@@ -767,19 +697,22 @@ pub(crate) async fn resolve_body_ref(
|
||||
pool: &MysqlPool,
|
||||
body_ref: &str,
|
||||
) -> Result<Option<Value>, DataLayerError> {
|
||||
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let canonical_ref = usage_body_ref(&request_id, field);
|
||||
if let Some(payload_gzip) = sqlx::query_scalar::<_, Vec<u8>>(
|
||||
"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? LIMIT 1",
|
||||
"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = ? AND request_id = ? AND body_field = ? LIMIT 1",
|
||||
)
|
||||
.bind(body_ref)
|
||||
.bind(&canonical_ref)
|
||||
.bind(&request_id)
|
||||
.bind(field.as_storage_field())
|
||||
.fetch_optional(pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return inflate_json(&payload_gzip).map(Some);
|
||||
}
|
||||
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let (inline_column, compressed_column) = usage_body_sql_columns(field);
|
||||
let row = sqlx::query(&format!(
|
||||
"SELECT CAST({inline_column} AS CHAR) AS inline_body, {compressed_column} AS compressed_body FROM `usage` WHERE request_id = ? LIMIT 1"
|
||||
@@ -819,11 +752,7 @@ fn usage_body_sql_columns(field: UsageBodyField) -> (&'static str, &'static str)
|
||||
}
|
||||
|
||||
fn inflate_json(bytes: &[u8]) -> Result<Value, DataLayerError> {
|
||||
let mut decoder = GzDecoder::new(bytes);
|
||||
let mut decoded = Vec::new();
|
||||
decoder.read_to_end(&mut decoded).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to decompress usage body: {err}"))
|
||||
})?;
|
||||
let decoded = read_decompressed_usage_json(GzDecoder::new(bytes))?;
|
||||
serde_json::from_slice(&decoded).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to decode usage body JSON: {err}"))
|
||||
})
|
||||
|
||||
@@ -152,6 +152,124 @@ fn mysql_usage_upsert_guards_candidate_identity_metadata_and_routing_from_late_l
|
||||
.contains("OR (status = 'streaming' AND VALUES(status) = 'pending')"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_stale_terminal_event_is_a_full_transaction_noop_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
else {
|
||||
eprintln!("skipping MySQL stale terminal test because AETHER_TEST_MYSQL_URL is unset");
|
||||
return;
|
||||
};
|
||||
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("mysql test pool should connect");
|
||||
run_migrations(&pool)
|
||||
.await
|
||||
.expect("mysql migrations should run");
|
||||
let suffix = unique_suffix();
|
||||
let user_id = format!("stale-user-{suffix}");
|
||||
let api_key_id = format!("stale-api-key-{suffix}");
|
||||
let provider_id = format!("stale-provider-{suffix}");
|
||||
let provider_key_id = format!("stale-provider-key-{suffix}");
|
||||
let request_id = format!("stale-request-{suffix}");
|
||||
seed_stats_targets(&pool, &user_id, &api_key_id, &provider_id, &provider_key_id).await;
|
||||
let repository = MysqlUsageWriteRepository::new(pool.clone());
|
||||
|
||||
let mut newer = sample_usage(
|
||||
&request_id,
|
||||
&user_id,
|
||||
&api_key_id,
|
||||
&provider_id,
|
||||
&provider_key_id,
|
||||
"completed",
|
||||
"pending",
|
||||
2_000,
|
||||
);
|
||||
newer.candidate_id = Some("candidate-new".to_string());
|
||||
newer.route_kind = Some("route-new".to_string());
|
||||
repository
|
||||
.upsert(newer)
|
||||
.await
|
||||
.expect("newer terminal usage should upsert");
|
||||
|
||||
let counter_rows_before: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("counter rows should count");
|
||||
let routing_before: (Option<String>, Option<String>) = sqlx::query_as(
|
||||
"SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?",
|
||||
)
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("routing snapshot should load");
|
||||
let settlement_before: (String, Option<f64>) = sqlx::query_as(
|
||||
"SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?",
|
||||
)
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("settlement snapshot should load");
|
||||
|
||||
let mut stale = sample_usage(
|
||||
&request_id,
|
||||
&user_id,
|
||||
&api_key_id,
|
||||
&provider_id,
|
||||
&provider_key_id,
|
||||
"failed",
|
||||
"void",
|
||||
1_999,
|
||||
);
|
||||
stale.status_code = Some(503);
|
||||
stale.total_cost_usd = Some(99.0);
|
||||
stale.actual_total_cost_usd = Some(98.0);
|
||||
stale.candidate_id = Some("candidate-stale".to_string());
|
||||
stale.route_kind = Some("route-stale".to_string());
|
||||
let stored = repository
|
||||
.upsert(stale)
|
||||
.await
|
||||
.expect("stale terminal usage should be ignored");
|
||||
|
||||
assert_eq!(stored.status, "completed");
|
||||
assert_eq!(stored.billing_status, "pending");
|
||||
assert_eq!(stored.status_code, Some(200));
|
||||
assert_eq!(stored.total_cost_usd, 0.5);
|
||||
assert_eq!(stored.routing_candidate_id(), Some("candidate-new"));
|
||||
assert_eq!(stored.routing_route_kind(), Some("route-new"));
|
||||
assert_eq!(stored.updated_at_unix_secs, 2_000);
|
||||
|
||||
let counter_rows_after: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM usage_counter_deltas WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("counter rows should count");
|
||||
let routing_after: (Option<String>, Option<String>) = sqlx::query_as(
|
||||
"SELECT candidate_id, route_kind FROM usage_routing_snapshots WHERE request_id = ?",
|
||||
)
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("routing snapshot should load");
|
||||
let settlement_after: (String, Option<f64>) = sqlx::query_as(
|
||||
"SELECT billing_status, billing_total_cost_usd FROM usage_settlement_snapshots WHERE request_id = ?",
|
||||
)
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("settlement snapshot should load");
|
||||
assert_eq!(counter_rows_after, counter_rows_before);
|
||||
assert_eq!(routing_after, routing_before);
|
||||
assert_eq!(settlement_after, settlement_before);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_usage_write_repository_upserts_and_flushes_counters_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
@@ -470,7 +588,7 @@ async fn mysql_concurrent_same_request_upserts_enqueue_counters_once_when_url_is
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
|
||||
async fn mysql_usage_http_capture_is_not_persisted_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
@@ -516,29 +634,17 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
|
||||
.upsert(rich)
|
||||
.await
|
||||
.expect("MySQL canonical capture should upsert");
|
||||
assert_eq!(
|
||||
stored.request_headers,
|
||||
Some(serde_json::json!({"x-client": "one"}))
|
||||
);
|
||||
assert_eq!(
|
||||
stored.request_body,
|
||||
Some(serde_json::json!({"request": true}))
|
||||
);
|
||||
assert_eq!(
|
||||
stored.request_body_state,
|
||||
Some(UsageBodyCaptureState::Reference)
|
||||
);
|
||||
assert_eq!(
|
||||
stored.request_body_ref.as_deref(),
|
||||
Some(format!("usage://request/{request_id}/request_body").as_str())
|
||||
);
|
||||
assert!(stored.request_headers.is_none());
|
||||
assert!(stored.request_body.is_none());
|
||||
assert!(stored.request_body_state.is_none());
|
||||
assert!(stored.request_body_ref.is_none());
|
||||
let blob_count: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("MySQL canonical blobs should count");
|
||||
assert_eq!(blob_count, 2);
|
||||
assert_eq!(blob_count, 0);
|
||||
let legacy_body: Option<String> =
|
||||
sqlx::query_scalar("SELECT CAST(request_body AS CHAR) FROM `usage` WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
@@ -561,8 +667,8 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
|
||||
.upsert(sparse)
|
||||
.await
|
||||
.expect("MySQL sparse capture should upsert");
|
||||
assert_eq!(sparse_stored.request_headers, stored.request_headers);
|
||||
assert_eq!(sparse_stored.request_body, stored.request_body);
|
||||
assert!(sparse_stored.request_headers.is_none());
|
||||
assert!(sparse_stored.request_body.is_none());
|
||||
|
||||
let mut clear = sample_usage(
|
||||
&request_id,
|
||||
@@ -582,11 +688,8 @@ async fn mysql_usage_http_capture_round_trips_when_url_is_set() {
|
||||
.expect("MySQL explicit none should clear");
|
||||
assert!(cleared.request_body.is_none());
|
||||
assert!(cleared.request_body_ref.is_none());
|
||||
assert_eq!(
|
||||
cleared.request_body_state,
|
||||
Some(UsageBodyCaptureState::None)
|
||||
);
|
||||
assert_eq!(cleared.provider_request_body, stored.provider_request_body);
|
||||
assert!(cleared.request_body_state.is_none());
|
||||
assert!(cleared.provider_request_body.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -992,16 +1095,27 @@ async fn mysql_usage_cleanup_executes_when_url_is_set() {
|
||||
.await
|
||||
.expect("MySQL detail cleanup should succeed");
|
||||
assert!(summary.body_externalized >= 1);
|
||||
let body_ref: String =
|
||||
sqlx::query_scalar("SELECT request_body_ref FROM usage_http_audits WHERE request_id = ?")
|
||||
let stored_body: Option<String> =
|
||||
sqlx::query_scalar("SELECT CAST(request_body AS CHAR) FROM `usage` WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("externalized body ref should load");
|
||||
assert_eq!(
|
||||
body_ref,
|
||||
format!("usage://request/{request_id}/request_body")
|
||||
);
|
||||
.expect("purged body should load");
|
||||
assert!(stored_body.is_none());
|
||||
let body_blobs: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM usage_body_blobs WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("purged body blobs should count");
|
||||
assert_eq!(body_blobs, 0);
|
||||
let body_audits: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM usage_http_audits WHERE request_id = ?")
|
||||
.bind(&request_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("purged body refs should count");
|
||||
assert_eq!(body_audits, 0);
|
||||
|
||||
let headers_only = UsageCleanupTargets {
|
||||
detail_body: false,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -72,6 +72,22 @@ impl MysqlVideoTaskRepository {
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
async fn find_by_id_for_user(
|
||||
&self,
|
||||
id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredVideoTask>, DataLayerError> {
|
||||
let row = sqlx::query(&format!(
|
||||
"{VIDEO_TASK_COLUMNS} WHERE BINARY id = BINARY ? AND BINARY user_id = BINARY ? LIMIT 1"
|
||||
))
|
||||
.bind(id)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
async fn find_by_short_id(
|
||||
&self,
|
||||
short_id: &str,
|
||||
@@ -84,13 +100,29 @@ impl MysqlVideoTaskRepository {
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
async fn find_by_short_id_for_user(
|
||||
&self,
|
||||
short_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredVideoTask>, DataLayerError> {
|
||||
let row = sqlx::query(&format!(
|
||||
"{VIDEO_TASK_COLUMNS} WHERE BINARY short_id = BINARY ? AND BINARY user_id = BINARY ? LIMIT 1"
|
||||
))
|
||||
.bind(short_id)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
async fn find_by_user_external(
|
||||
&self,
|
||||
user_id: &str,
|
||||
external_task_id: &str,
|
||||
) -> Result<Option<StoredVideoTask>, DataLayerError> {
|
||||
let row = sqlx::query(&format!(
|
||||
"{VIDEO_TASK_COLUMNS} WHERE user_id = ? AND external_task_id = ? LIMIT 1"
|
||||
"{VIDEO_TASK_COLUMNS} WHERE BINARY user_id = BINARY ? AND BINARY external_task_id = BINARY ? LIMIT 1"
|
||||
))
|
||||
.bind(user_id)
|
||||
.bind(external_task_id)
|
||||
@@ -117,6 +149,26 @@ impl VideoTaskReadRepository for MysqlVideoTaskRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn find_for_user(
|
||||
&self,
|
||||
key: VideoTaskLookupKey<'_>,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredVideoTask>, DataLayerError> {
|
||||
match key {
|
||||
VideoTaskLookupKey::Id(id) => self.find_by_id_for_user(id, user_id).await,
|
||||
VideoTaskLookupKey::ShortId(short_id) => {
|
||||
self.find_by_short_id_for_user(short_id, user_id).await
|
||||
}
|
||||
VideoTaskLookupKey::UserExternal {
|
||||
user_id: lookup_user_id,
|
||||
external_task_id,
|
||||
} if lookup_user_id == user_id => {
|
||||
self.find_by_user_external(user_id, external_task_id).await
|
||||
}
|
||||
VideoTaskLookupKey::UserExternal { .. } => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
|
||||
if limit == 0 {
|
||||
return Ok(Vec::new());
|
||||
@@ -258,15 +310,21 @@ impl VideoTaskReadRepository for MysqlVideoTaskRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl VideoTaskWriteRepository for MysqlVideoTaskRepository {
|
||||
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
|
||||
async fn upsert(&self, mut task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
|
||||
task.sanitize_for_persistence();
|
||||
let id = task.id.clone();
|
||||
bind_task(sqlx::query(UPSERT_SQL), task, true, false)?
|
||||
let expected_identity = task.clone();
|
||||
bind_task(sqlx::query(upsert_sql()), task, true, false)?
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
self.find_by_id(&id)
|
||||
.await?
|
||||
.ok_or_else(|| DataLayerError::UnexpectedValue("upserted video task missing".into()))
|
||||
let stored = self.find_by_id(&id).await?.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"video task {id} conflicts with persisted immutable identity"
|
||||
))
|
||||
})?;
|
||||
stored.ensure_immutable_identity_matches(&expected_identity)?;
|
||||
Ok(stored)
|
||||
}
|
||||
|
||||
async fn update_if_active(
|
||||
@@ -368,7 +426,76 @@ FOR UPDATE SKIP LOCKED
|
||||
}
|
||||
}
|
||||
|
||||
const UPSERT_SQL: &str = r#"
|
||||
const IMMUTABLE_IDENTITY_MATCH_SQL: &str = r#"BINARY id <=> BINARY VALUES(id)
|
||||
AND BINARY short_id <=> BINARY VALUES(short_id)
|
||||
AND BINARY request_id <=> BINARY VALUES(request_id)
|
||||
AND BINARY user_id <=> BINARY VALUES(user_id)
|
||||
AND BINARY api_key_id <=> BINARY VALUES(api_key_id)
|
||||
AND BINARY external_task_id <=> BINARY VALUES(external_task_id)
|
||||
AND BINARY provider_id <=> BINARY VALUES(provider_id)
|
||||
AND BINARY endpoint_id <=> BINARY VALUES(endpoint_id)
|
||||
AND BINARY key_id <=> BINARY VALUES(key_id)
|
||||
AND BINARY client_api_format <=> BINARY VALUES(client_api_format)
|
||||
AND BINARY provider_api_format <=> BINARY VALUES(provider_api_format)
|
||||
AND format_converted <=> VALUES(format_converted)
|
||||
AND BINARY model <=> BINARY VALUES(model)
|
||||
AND duration_seconds <=> VALUES(duration_seconds)
|
||||
AND BINARY resolution <=> BINARY VALUES(resolution)
|
||||
AND BINARY aspect_ratio <=> BINARY VALUES(aspect_ratio)
|
||||
AND BINARY size <=> BINARY VALUES(size)"#;
|
||||
|
||||
const UPSERT_UPDATE_COLUMNS: &[&str] = &[
|
||||
"short_id",
|
||||
"request_id",
|
||||
"user_id",
|
||||
"api_key_id",
|
||||
"username",
|
||||
"api_key_name",
|
||||
"external_task_id",
|
||||
"provider_id",
|
||||
"endpoint_id",
|
||||
"key_id",
|
||||
"client_api_format",
|
||||
"provider_api_format",
|
||||
"format_converted",
|
||||
"model",
|
||||
"prompt",
|
||||
"original_request_body",
|
||||
"duration_seconds",
|
||||
"resolution",
|
||||
"aspect_ratio",
|
||||
"size",
|
||||
"status",
|
||||
"progress_percent",
|
||||
"progress_message",
|
||||
"retry_count",
|
||||
"poll_interval_seconds",
|
||||
"next_poll_at",
|
||||
"poll_count",
|
||||
"max_poll_count",
|
||||
"video_url",
|
||||
"error_code",
|
||||
"error_message",
|
||||
"request_metadata",
|
||||
"submitted_at",
|
||||
"completed_at",
|
||||
"updated_at",
|
||||
];
|
||||
|
||||
fn upsert_sql() -> &'static str {
|
||||
static SQL: std::sync::OnceLock<String> = std::sync::OnceLock::new();
|
||||
SQL.get_or_init(|| {
|
||||
let guarded_updates = UPSERT_UPDATE_COLUMNS
|
||||
.iter()
|
||||
.map(|column| {
|
||||
format!(
|
||||
" {column} = IF(({IMMUTABLE_IDENTITY_MATCH_SQL}), VALUES({column}), {column})"
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(",\n");
|
||||
format!(
|
||||
r#"
|
||||
INSERT INTO video_tasks (
|
||||
id, short_id, request_id, user_id, api_key_id, username, api_key_name,
|
||||
external_task_id, provider_id, endpoint_id, key_id, client_api_format,
|
||||
@@ -380,43 +507,11 @@ INSERT INTO video_tasks (
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
short_id = VALUES(short_id),
|
||||
request_id = VALUES(request_id),
|
||||
user_id = VALUES(user_id),
|
||||
api_key_id = VALUES(api_key_id),
|
||||
username = VALUES(username),
|
||||
api_key_name = VALUES(api_key_name),
|
||||
external_task_id = VALUES(external_task_id),
|
||||
provider_id = VALUES(provider_id),
|
||||
endpoint_id = VALUES(endpoint_id),
|
||||
key_id = VALUES(key_id),
|
||||
client_api_format = VALUES(client_api_format),
|
||||
provider_api_format = VALUES(provider_api_format),
|
||||
format_converted = VALUES(format_converted),
|
||||
model = VALUES(model),
|
||||
prompt = VALUES(prompt),
|
||||
original_request_body = VALUES(original_request_body),
|
||||
duration_seconds = VALUES(duration_seconds),
|
||||
resolution = VALUES(resolution),
|
||||
aspect_ratio = VALUES(aspect_ratio),
|
||||
size = VALUES(size),
|
||||
status = VALUES(status),
|
||||
progress_percent = VALUES(progress_percent),
|
||||
progress_message = VALUES(progress_message),
|
||||
retry_count = VALUES(retry_count),
|
||||
poll_interval_seconds = VALUES(poll_interval_seconds),
|
||||
next_poll_at = VALUES(next_poll_at),
|
||||
poll_count = VALUES(poll_count),
|
||||
max_poll_count = VALUES(max_poll_count),
|
||||
video_url = VALUES(video_url),
|
||||
error_code = VALUES(error_code),
|
||||
error_message = VALUES(error_message),
|
||||
request_metadata = VALUES(request_metadata),
|
||||
created_at = VALUES(created_at),
|
||||
submitted_at = VALUES(submitted_at),
|
||||
completed_at = VALUES(completed_at),
|
||||
updated_at = VALUES(updated_at)
|
||||
"#;
|
||||
{guarded_updates}
|
||||
"#
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
const UPDATE_IF_ACTIVE_SQL: &str = r#"
|
||||
UPDATE video_tasks SET
|
||||
@@ -452,20 +547,38 @@ UPDATE video_tasks SET
|
||||
error_code = ?,
|
||||
error_message = ?,
|
||||
request_metadata = ?,
|
||||
created_at = ?,
|
||||
created_at = COALESCE(created_at, ?),
|
||||
submitted_at = ?,
|
||||
completed_at = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
AND status IN ('pending', 'submitted', 'queued', 'processing')
|
||||
AND BINARY short_id <=> BINARY ?
|
||||
AND BINARY request_id <=> BINARY ?
|
||||
AND BINARY user_id <=> BINARY ?
|
||||
AND BINARY api_key_id <=> BINARY ?
|
||||
AND BINARY external_task_id <=> BINARY ?
|
||||
AND BINARY provider_id <=> BINARY ?
|
||||
AND BINARY endpoint_id <=> BINARY ?
|
||||
AND BINARY key_id <=> BINARY ?
|
||||
AND BINARY client_api_format <=> BINARY ?
|
||||
AND BINARY provider_api_format <=> BINARY ?
|
||||
AND format_converted <=> ?
|
||||
AND BINARY model <=> BINARY ?
|
||||
AND duration_seconds <=> ?
|
||||
AND BINARY resolution <=> BINARY ?
|
||||
AND BINARY aspect_ratio <=> BINARY ?
|
||||
AND BINARY size <=> BINARY ?
|
||||
"#;
|
||||
|
||||
fn bind_task<'q>(
|
||||
query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>,
|
||||
task: UpsertVideoTask,
|
||||
mut task: UpsertVideoTask,
|
||||
include_insert_id: bool,
|
||||
include_update_id: bool,
|
||||
) -> Result<sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
|
||||
task.sanitize_for_persistence();
|
||||
let identity = task.clone();
|
||||
let original_request_body = json_to_string(&task.original_request_body)?;
|
||||
let request_metadata = json_to_string(&task.request_metadata)?;
|
||||
let query = if include_insert_id {
|
||||
@@ -535,19 +648,45 @@ fn bind_task<'q>(
|
||||
"video task updated_at",
|
||||
)?);
|
||||
if include_update_id {
|
||||
Ok(bound.bind(task.id))
|
||||
bind_identity_guard(bound.bind(task.id), identity)
|
||||
} else {
|
||||
Ok(bound)
|
||||
}
|
||||
}
|
||||
|
||||
fn bind_identity_guard<'q>(
|
||||
query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>,
|
||||
identity: UpsertVideoTask,
|
||||
) -> Result<sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
|
||||
Ok(query
|
||||
.bind(identity.short_id)
|
||||
.bind(identity.request_id)
|
||||
.bind(identity.user_id)
|
||||
.bind(identity.api_key_id)
|
||||
.bind(identity.external_task_id)
|
||||
.bind(identity.provider_id)
|
||||
.bind(identity.endpoint_id)
|
||||
.bind(identity.key_id)
|
||||
.bind(identity.client_api_format)
|
||||
.bind(identity.provider_api_format)
|
||||
.bind(identity.format_converted)
|
||||
.bind(identity.model)
|
||||
.bind(optional_u32_to_i32(
|
||||
identity.duration_seconds,
|
||||
"video task duration_seconds",
|
||||
)?)
|
||||
.bind(identity.resolution)
|
||||
.bind(identity.aspect_ratio)
|
||||
.bind(identity.size))
|
||||
}
|
||||
|
||||
fn push_filter<'args>(
|
||||
builder: &mut QueryBuilder<'args, MySql>,
|
||||
filter: &'args VideoTaskQueryFilter,
|
||||
created_since_unix_secs: Option<u64>,
|
||||
) {
|
||||
if let Some(user_id) = filter.user_id.as_deref() {
|
||||
push_clause(builder, "user_id = ");
|
||||
push_clause(builder, "BINARY user_id = BINARY ");
|
||||
builder.push_bind(user_id);
|
||||
}
|
||||
if let Some(status) = filter.status {
|
||||
@@ -706,13 +845,86 @@ fn optional_u32_to_i32(value: Option<u32>, name: &str) -> Result<Option<i32>, Da
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlVideoTaskRepository;
|
||||
use super::{
|
||||
upsert_sql, MysqlVideoTaskRepository, IMMUTABLE_IDENTITY_MATCH_SQL, UPDATE_IF_ACTIVE_SQL,
|
||||
UPSERT_UPDATE_COLUMNS,
|
||||
};
|
||||
use crate::run_migrations;
|
||||
use aether_data_contracts::repository::video_tasks::{
|
||||
UpsertVideoTask, VideoTaskStatus, VideoTaskWriteRepository,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
#[test]
|
||||
fn mysql_write_sql_atomically_guards_immutable_identity() {
|
||||
assert!(IMMUTABLE_IDENTITY_MATCH_SQL.contains("BINARY id <=> BINARY VALUES(id)"));
|
||||
for column in [
|
||||
"short_id",
|
||||
"request_id",
|
||||
"user_id",
|
||||
"api_key_id",
|
||||
"external_task_id",
|
||||
"provider_id",
|
||||
"endpoint_id",
|
||||
"key_id",
|
||||
"client_api_format",
|
||||
"provider_api_format",
|
||||
"model",
|
||||
"resolution",
|
||||
"aspect_ratio",
|
||||
"size",
|
||||
] {
|
||||
assert!(
|
||||
IMMUTABLE_IDENTITY_MATCH_SQL
|
||||
.contains(&format!("BINARY {column} <=> BINARY VALUES({column})")),
|
||||
"upsert identity predicate should guard {column}"
|
||||
);
|
||||
}
|
||||
for column in ["format_converted", "duration_seconds"] {
|
||||
assert!(
|
||||
IMMUTABLE_IDENTITY_MATCH_SQL.contains(&format!("{column} <=> VALUES({column})")),
|
||||
"upsert identity predicate should guard {column}"
|
||||
);
|
||||
}
|
||||
|
||||
let upsert = upsert_sql();
|
||||
assert!(!UPSERT_UPDATE_COLUMNS.contains(&"created_at"));
|
||||
for column in UPSERT_UPDATE_COLUMNS {
|
||||
assert!(
|
||||
upsert.contains(&format!("{column} = IF((BINARY id <=> BINARY VALUES(id)")),
|
||||
"upsert assignment should be conditional for {column}"
|
||||
);
|
||||
}
|
||||
for column in [
|
||||
"short_id",
|
||||
"request_id",
|
||||
"user_id",
|
||||
"api_key_id",
|
||||
"external_task_id",
|
||||
"provider_id",
|
||||
"endpoint_id",
|
||||
"key_id",
|
||||
"client_api_format",
|
||||
"provider_api_format",
|
||||
"model",
|
||||
"resolution",
|
||||
"aspect_ratio",
|
||||
"size",
|
||||
] {
|
||||
assert!(
|
||||
UPDATE_IF_ACTIVE_SQL.contains(&format!("BINARY {column} <=> BINARY ?")),
|
||||
"active update should guard {column}"
|
||||
);
|
||||
}
|
||||
for column in ["format_converted", "duration_seconds"] {
|
||||
assert!(
|
||||
UPDATE_IF_ACTIVE_SQL.contains(&format!("{column} <=> ?")),
|
||||
"active update should guard {column}"
|
||||
);
|
||||
}
|
||||
assert!(UPDATE_IF_ACTIVE_SQL.contains("created_at = COALESCE(created_at, ?)"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_builds_from_lazy_pool() {
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -75,7 +75,7 @@ fn mysql_wallet_admin_builders_cover_filters_ordering_and_mapping_columns() {
|
||||
let order_sql = compact_sql(admin_payment_order_list_builder(&order_query, 100, 8, 6).sql());
|
||||
assert!(order_sql.contains("payment_method = ?"));
|
||||
assert!(order_sql.contains(
|
||||
"CASE WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at < ? THEN 'expired'"
|
||||
"CASE WHEN status = 'pending' AND expires_at IS NOT NULL AND expires_at <= ? THEN 'expired'"
|
||||
));
|
||||
assert!(order_sql.contains("ORDER BY created_at DESC, id DESC LIMIT ? OFFSET ?"));
|
||||
|
||||
@@ -114,6 +114,66 @@ fn compact_sql(sql: &str) -> String {
|
||||
sql.split_whitespace().collect::<Vec<_>>().join(" ")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_gateway_order_uniqueness_preflights_before_persistent_changes() {
|
||||
const UNIQUENESS_MIGRATION: &str = include_str!(
|
||||
"../../migrations/20260821120000_enforce_payment_gateway_order_uniqueness.sql"
|
||||
);
|
||||
|
||||
let executable_migration = UNIQUENESS_MIGRATION
|
||||
.lines()
|
||||
.filter(|line| !line.trim_start().starts_with("--"))
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
let migration = compact_sql(&executable_migration);
|
||||
let initial_cleanup = migration
|
||||
.find("DROP TEMPORARY TABLE IF EXISTS aether_payment_gateway_order_uniqueness_preflight")
|
||||
.expect("migration should clean up a same-session failed preflight");
|
||||
let create_preflight = migration
|
||||
.find("CREATE TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight")
|
||||
.expect("migration should create a non-persistent conflict guard");
|
||||
let seed_preflight = migration
|
||||
.find("INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) VALUES (1)")
|
||||
.expect("migration should seed the duplicate-key conflict guard");
|
||||
let conflict_probe = migration
|
||||
.find("INSERT INTO aether_payment_gateway_order_uniqueness_preflight (conflict_marker) SELECT 1 FROM payment_orders")
|
||||
.expect("migration should reject normalized historical conflicts");
|
||||
let final_cleanup = migration
|
||||
.find("DROP TEMPORARY TABLE aether_payment_gateway_order_uniqueness_preflight;")
|
||||
.expect("migration should remove the successful preflight guard");
|
||||
let first_persistent_update = migration
|
||||
.find("UPDATE payment_orders SET payment_method")
|
||||
.expect("migration should normalize payment order methods");
|
||||
let alter = migration
|
||||
.find(
|
||||
"MODIFY COLUMN gateway_order_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_bin NULL",
|
||||
)
|
||||
.expect("migration should enforce binary gateway identifiers");
|
||||
let unique = migration
|
||||
.find("ADD UNIQUE INDEX uq_payment_orders_payment_method_gateway_order_id")
|
||||
.expect("migration should create composite uniqueness");
|
||||
|
||||
assert!(migration.contains("GROUP BY LOWER(TRIM(payment_method)), CONVERT(gateway_order_id USING utf8mb4) COLLATE utf8mb4_0900_bin HAVING COUNT(*) > 1 LIMIT 1"));
|
||||
assert!(migration.contains("WHERE BINARY payment_method <> BINARY LOWER(TRIM(payment_method))"));
|
||||
let preflight = &migration[..final_cleanup];
|
||||
assert!(!preflight.contains("UPDATE payment_orders"));
|
||||
assert!(!preflight.contains("UPDATE payment_callbacks"));
|
||||
assert!(!preflight.contains("ALTER TABLE"));
|
||||
assert!(
|
||||
initial_cleanup < create_preflight
|
||||
&& create_preflight < seed_preflight
|
||||
&& seed_preflight < conflict_probe
|
||||
&& conflict_probe < final_cleanup
|
||||
&& final_cleanup < first_persistent_update
|
||||
&& first_persistent_update < alter,
|
||||
"the conflict probe must finish before any persistent UPDATE or ALTER"
|
||||
);
|
||||
assert!(
|
||||
alter < unique && !migration.contains("CREATE UNIQUE INDEX"),
|
||||
"the collation and unique index must be one atomic ALTER TABLE"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_wallet_read_repository_reads_wallet_contract_views() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
|
||||
Reference in New Issue
Block a user