Merge remote-tracking branch 'origin/main' into codex/pool-key-bulk-management-20260714

# Conflicts:
#	apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs
#	frontend/src/api/endpoints/pool.ts
This commit is contained in:
MMEXA
2026-07-16 23:43:04 +08:00
1257 changed files with 80521 additions and 35495 deletions
@@ -0,0 +1,24 @@
[package]
name = "aether-data-sqlite"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
description = "SQLite repositories, pools, and migrations for Aether"
[dependencies]
aether-ai-formats.workspace = true
aether-data-contracts.workspace = true
aether-data-query.workspace = true
async-trait.workspace = true
chrono.workspace = true
chrono-tz.workspace = true
flate2.workspace = true
serde_json.workspace = true
sha2.workspace = true
sqlx = { workspace = true, features = ["sqlite", "runtime-tokio-rustls", "chrono", "migrate", "macros"] }
tracing.workspace = true
uuid.workspace = true
[dev-dependencies]
tokio.workspace = true
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,3 @@
ALTER TABLE management_tokens
ADD COLUMN permissions TEXT;
@@ -0,0 +1,46 @@
ALTER TABLE proxy_node_events
ADD COLUMN event_metadata TEXT;
CREATE TABLE IF NOT EXISTS proxy_node_metrics_1m (
node_id TEXT NOT NULL,
bucket_start_unix_secs INTEGER NOT NULL,
samples INTEGER NOT NULL DEFAULT 0,
uptime_samples INTEGER NOT NULL DEFAULT 0,
active_connections_sum INTEGER NOT NULL DEFAULT 0,
active_connections_max INTEGER NOT NULL DEFAULT 0,
heartbeat_rtt_ms_sum INTEGER NOT NULL DEFAULT 0,
heartbeat_rtt_ms_max INTEGER NOT NULL DEFAULT 0,
connect_errors_delta INTEGER NOT NULL DEFAULT 0,
disconnects_delta INTEGER NOT NULL DEFAULT 0,
error_events_delta INTEGER NOT NULL DEFAULT 0,
ws_in_bytes_delta INTEGER NOT NULL DEFAULT 0,
ws_out_bytes_delta INTEGER NOT NULL DEFAULT 0,
ws_in_frames_delta INTEGER NOT NULL DEFAULT 0,
ws_out_frames_delta INTEGER NOT NULL DEFAULT 0,
PRIMARY KEY (node_id, bucket_start_unix_secs)
);
CREATE TABLE IF NOT EXISTS proxy_node_metrics_1h (
node_id TEXT NOT NULL,
bucket_start_unix_secs INTEGER NOT NULL,
samples INTEGER NOT NULL DEFAULT 0,
uptime_samples INTEGER NOT NULL DEFAULT 0,
active_connections_sum INTEGER NOT NULL DEFAULT 0,
active_connections_max INTEGER NOT NULL DEFAULT 0,
heartbeat_rtt_ms_sum INTEGER NOT NULL DEFAULT 0,
heartbeat_rtt_ms_max INTEGER NOT NULL DEFAULT 0,
connect_errors_delta INTEGER NOT NULL DEFAULT 0,
disconnects_delta INTEGER NOT NULL DEFAULT 0,
error_events_delta INTEGER NOT NULL DEFAULT 0,
ws_in_bytes_delta INTEGER NOT NULL DEFAULT 0,
ws_out_bytes_delta INTEGER NOT NULL DEFAULT 0,
ws_in_frames_delta INTEGER NOT NULL DEFAULT 0,
ws_out_frames_delta INTEGER NOT NULL DEFAULT 0,
PRIMARY KEY (node_id, bucket_start_unix_secs)
);
CREATE INDEX IF NOT EXISTS idx_proxy_node_metrics_1m_bucket_start
ON proxy_node_metrics_1m (bucket_start_unix_secs);
CREATE INDEX IF NOT EXISTS idx_proxy_node_metrics_1h_bucket_start
ON proxy_node_metrics_1h (bucket_start_unix_secs);
@@ -0,0 +1,43 @@
CREATE TABLE IF NOT EXISTS background_task_runs (
id TEXT PRIMARY KEY,
task_key TEXT NOT NULL,
kind TEXT NOT NULL,
"trigger" TEXT NOT NULL,
status TEXT NOT NULL,
attempt INTEGER NOT NULL DEFAULT 0,
max_attempts INTEGER NOT NULL DEFAULT 0,
owner_instance TEXT,
progress_percent INTEGER NOT NULL DEFAULT 0,
progress_message TEXT,
payload_json TEXT,
result_json TEXT,
error_message TEXT,
cancel_requested INTEGER NOT NULL DEFAULT 0,
created_by TEXT,
created_at_unix_secs INTEGER NOT NULL,
started_at_unix_secs INTEGER,
finished_at_unix_secs INTEGER,
updated_at_unix_secs INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_background_task_runs_task_key
ON background_task_runs (task_key);
CREATE INDEX IF NOT EXISTS idx_background_task_runs_status
ON background_task_runs (status);
CREATE INDEX IF NOT EXISTS idx_background_task_runs_kind
ON background_task_runs (kind);
CREATE INDEX IF NOT EXISTS idx_background_task_runs_created_at
ON background_task_runs (created_at_unix_secs DESC);
CREATE TABLE IF NOT EXISTS background_task_events (
id TEXT PRIMARY KEY,
run_id TEXT NOT NULL,
event_type TEXT NOT NULL,
message TEXT NOT NULL,
payload_json TEXT,
created_at_unix_secs INTEGER NOT NULL,
FOREIGN KEY (run_id) REFERENCES background_task_runs(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_background_task_events_run_id
ON background_task_events (run_id, created_at_unix_secs ASC);
@@ -0,0 +1,101 @@
ALTER TABLE users ADD COLUMN allowed_providers_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN allowed_api_formats_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN allowed_models_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN rate_limit_mode TEXT NOT NULL DEFAULT 'system';
UPDATE users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS user_groups (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
normalized_name TEXT NOT NULL UNIQUE,
description TEXT,
priority INTEGER NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'inherit',
rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'inherit',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id TEXT NOT NULL REFERENCES user_groups(id) ON DELETE CASCADE,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
created_at INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
ON user_group_members (user_id);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
ON user_groups (priority DESC, name ASC, id ASC);
INSERT OR IGNORE INTO user_groups (
id,
name,
normalized_name,
description,
priority,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
created_at,
updated_at
)
VALUES (
'00000000-0000-0000-0000-000000000001',
'Default',
'default',
'Default group for all users',
0,
'unrestricted',
'unrestricted',
'unrestricted',
'system',
CAST(strftime('%s', 'now') AS INTEGER),
CAST(strftime('%s', 'now') AS INTEGER)
);
INSERT OR IGNORE INTO system_configs (
id,
key,
value,
description,
created_at,
updated_at
)
VALUES (
'00000000-0000-0000-0000-000000000002',
'default_user_group_id',
'"00000000-0000-0000-0000-000000000001"',
'Default user group',
CAST(strftime('%s', 'now') AS INTEGER),
CAST(strftime('%s', 'now') AS INTEGER)
);
INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at)
SELECT '00000000-0000-0000-0000-000000000001', id, CAST(strftime('%s', 'now') AS INTEGER)
FROM users
WHERE is_deleted = 0
AND LOWER(role) <> 'admin';
@@ -0,0 +1,29 @@
UPDATE users
SET allowed_providers_mode = 'unrestricted'
WHERE allowed_providers_mode = 'specific'
AND (
allowed_providers IS NULL
OR trim(allowed_providers) = ''
OR lower(trim(allowed_providers)) = 'null'
OR trim(allowed_providers) = '[]'
);
UPDATE users
SET allowed_api_formats_mode = 'unrestricted'
WHERE allowed_api_formats_mode = 'specific'
AND (
allowed_api_formats IS NULL
OR trim(allowed_api_formats) = ''
OR lower(trim(allowed_api_formats)) = 'null'
OR trim(allowed_api_formats) = '[]'
);
UPDATE users
SET allowed_models_mode = 'unrestricted'
WHERE allowed_models_mode = 'specific'
AND (
allowed_models IS NULL
OR trim(allowed_models) = ''
OR lower(trim(allowed_models)) = 'null'
OR trim(allowed_models) = '[]'
);
@@ -0,0 +1,14 @@
DELETE FROM user_group_members
WHERE user_id IN (
SELECT id
FROM users
WHERE LOWER(role) = 'admin'
)
AND (
group_id = '00000000-0000-0000-0000-000000000001'
OR group_id IN (
SELECT TRIM(value, '"')
FROM system_configs
WHERE key = 'default_user_group_id'
)
);
@@ -0,0 +1,37 @@
CREATE TABLE IF NOT EXISTS pool_member_scores (
id TEXT PRIMARY KEY,
pool_kind TEXT NOT NULL,
pool_id TEXT NOT NULL,
member_kind TEXT NOT NULL,
member_id TEXT NOT NULL,
capability TEXT NOT NULL,
scope_kind TEXT NOT NULL,
scope_id TEXT,
score REAL NOT NULL DEFAULT 0,
hard_state TEXT NOT NULL DEFAULT 'unknown',
score_version INTEGER NOT NULL DEFAULT 1,
score_reason TEXT NOT NULL,
last_ranked_at INTEGER,
last_scheduled_at INTEGER,
last_success_at INTEGER,
last_failure_at INTEGER,
failure_count INTEGER NOT NULL DEFAULT 0,
last_probe_attempt_at INTEGER,
last_probe_success_at INTEGER,
last_probe_failure_at INTEGER,
probe_failure_count INTEGER NOT NULL DEFAULT 0,
probe_status TEXT NOT NULL DEFAULT 'never',
updated_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS pool_member_scores_rank_idx
ON pool_member_scores (pool_kind, pool_id, capability, scope_kind, scope_id, hard_state, score DESC);
CREATE INDEX IF NOT EXISTS pool_member_scores_member_idx
ON pool_member_scores (pool_kind, pool_id, member_kind, member_id);
CREATE INDEX IF NOT EXISTS pool_member_scores_probe_idx
ON pool_member_scores (pool_kind, pool_id, probe_status, last_probe_success_at);
CREATE INDEX IF NOT EXISTS pool_member_scores_updated_at_idx
ON pool_member_scores (updated_at);
@@ -0,0 +1,2 @@
ALTER TABLE users ADD COLUMN feature_settings TEXT;
ALTER TABLE api_keys ADD COLUMN feature_settings TEXT;
@@ -0,0 +1,85 @@
ALTER TABLE payment_orders ADD COLUMN payment_provider TEXT;
ALTER TABLE payment_orders ADD COLUMN payment_channel TEXT;
ALTER TABLE payment_orders ADD COLUMN order_kind TEXT NOT NULL DEFAULT 'wallet_recharge';
ALTER TABLE payment_orders ADD COLUMN product_id TEXT;
ALTER TABLE payment_orders ADD COLUMN product_snapshot TEXT;
ALTER TABLE payment_orders ADD COLUMN fulfillment_status TEXT NOT NULL DEFAULT 'pending';
ALTER TABLE payment_orders ADD COLUMN fulfillment_error TEXT;
CREATE INDEX IF NOT EXISTS idx_payment_orders_kind_status
ON payment_orders (order_kind, status);
CREATE INDEX IF NOT EXISTS idx_payment_orders_product
ON payment_orders (product_id);
CREATE TABLE IF NOT EXISTS payment_gateway_configs (
provider TEXT PRIMARY KEY,
enabled INTEGER NOT NULL DEFAULT 0,
endpoint_url TEXT NOT NULL,
callback_base_url TEXT,
merchant_id TEXT NOT NULL,
merchant_key_encrypted TEXT,
pay_currency TEXT NOT NULL DEFAULT 'CNY',
usd_exchange_rate REAL NOT NULL DEFAULT 7.2,
min_recharge_usd REAL NOT NULL DEFAULT 1,
channels_json TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS billing_plans (
id TEXT PRIMARY KEY,
title TEXT NOT NULL,
description TEXT,
price_amount REAL NOT NULL,
price_currency TEXT NOT NULL DEFAULT 'CNY',
duration_unit TEXT NOT NULL,
duration_value INTEGER NOT NULL,
enabled INTEGER NOT NULL DEFAULT 1,
sort_order INTEGER NOT NULL DEFAULT 0,
max_active_per_user INTEGER NOT NULL DEFAULT 1,
entitlements_json TEXT NOT NULL,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_billing_plans_enabled_sort
ON billing_plans (enabled, sort_order);
CREATE TABLE IF NOT EXISTS user_plan_entitlements (
id TEXT PRIMARY KEY,
user_id TEXT NOT NULL,
plan_id TEXT NOT NULL,
payment_order_id TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'active',
starts_at INTEGER NOT NULL,
expires_at INTEGER NOT NULL,
entitlements_snapshot TEXT NOT NULL,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY(user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY(plan_id) REFERENCES billing_plans(id) ON DELETE RESTRICT,
FOREIGN KEY(payment_order_id) REFERENCES payment_orders(id) ON DELETE RESTRICT
);
CREATE INDEX IF NOT EXISTS idx_user_plan_entitlements_user_active
ON user_plan_entitlements (user_id, status, expires_at);
CREATE INDEX IF NOT EXISTS idx_user_plan_entitlements_order
ON user_plan_entitlements (payment_order_id);
CREATE TABLE IF NOT EXISTS entitlement_usage_ledgers (
id TEXT PRIMARY KEY,
user_entitlement_id TEXT NOT NULL,
user_id TEXT NOT NULL,
request_id TEXT NOT NULL,
amount_usd REAL NOT NULL,
balance_before REAL NOT NULL,
balance_after REAL NOT NULL,
usage_date TEXT NOT NULL,
created_at INTEGER NOT NULL,
UNIQUE (user_entitlement_id, request_id),
FOREIGN KEY(user_entitlement_id) REFERENCES user_plan_entitlements(id) ON DELETE CASCADE,
FOREIGN KEY(user_id) REFERENCES users(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_entitlement_usage_user_date
ON entitlement_usage_ledgers (user_id, usage_date);
@@ -0,0 +1,2 @@
ALTER TABLE billing_plans
ADD COLUMN purchase_limit_scope TEXT NOT NULL DEFAULT 'active_period';
@@ -0,0 +1,45 @@
CREATE TABLE IF NOT EXISTS routing_groups (
id TEXT PRIMARY KEY NOT NULL,
name TEXT NOT NULL,
description TEXT,
enabled INTEGER NOT NULL DEFAULT 1,
is_system_default INTEGER NOT NULL DEFAULT 0,
config_json TEXT NOT NULL,
version INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
published_at INTEGER,
UNIQUE (name)
);
CREATE INDEX IF NOT EXISTS routing_groups_system_default_idx
ON routing_groups (is_system_default, enabled);
CREATE TABLE IF NOT EXISTS routing_group_bindings (
id TEXT PRIMARY KEY NOT NULL,
group_id TEXT NOT NULL,
subject_type TEXT NOT NULL,
subject_id TEXT NOT NULL,
is_default INTEGER NOT NULL DEFAULT 0,
allow_explicit_select INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS routing_group_bindings_group_id_idx
ON routing_group_bindings (group_id);
CREATE INDEX IF NOT EXISTS routing_group_bindings_subject_idx
ON routing_group_bindings (subject_type, subject_id);
CREATE TABLE IF NOT EXISTS routing_group_versions (
id TEXT PRIMARY KEY NOT NULL,
group_id TEXT NOT NULL,
version INTEGER NOT NULL,
config_json TEXT NOT NULL,
created_at INTEGER NOT NULL,
created_by TEXT,
UNIQUE (group_id, version)
);
CREATE INDEX IF NOT EXISTS routing_group_versions_group_id_idx
ON routing_group_versions (group_id);
@@ -0,0 +1,34 @@
CREATE TABLE IF NOT EXISTS usage_counter_deltas (
id TEXT PRIMARY KEY NOT NULL,
request_id TEXT NOT NULL,
kind TEXT NOT NULL,
target_id TEXT NOT NULL,
request_count_delta INTEGER NOT NULL DEFAULT 0,
total_requests_delta INTEGER NOT NULL DEFAULT 0,
success_count_delta INTEGER NOT NULL DEFAULT 0,
error_count_delta INTEGER NOT NULL DEFAULT 0,
dns_failures_delta INTEGER NOT NULL DEFAULT 0,
stream_errors_delta INTEGER NOT NULL DEFAULT 0,
total_tokens_delta INTEGER NOT NULL DEFAULT 0,
total_cost_usd_delta REAL NOT NULL DEFAULT 0,
total_response_time_ms_delta INTEGER NOT NULL DEFAULT 0,
last_used_at_unix_secs INTEGER,
last_used_ip TEXT,
candidate_last_used_at_unix_secs INTEGER,
removed_last_used_at_unix_secs INTEGER,
usage_created_at_unix_secs INTEGER,
created_at INTEGER NOT NULL,
processed_at INTEGER
);
CREATE INDEX IF NOT EXISTS ix_usage_counter_deltas_unprocessed
ON usage_counter_deltas (created_at, id);
CREATE INDEX IF NOT EXISTS ix_usage_counter_deltas_processed
ON usage_counter_deltas (processed_at, created_at, id);
CREATE INDEX IF NOT EXISTS ix_usage_counter_deltas_request_kind
ON usage_counter_deltas (request_id, kind, target_id);
CREATE INDEX IF NOT EXISTS video_tasks_due_poll_idx
ON video_tasks (status, next_poll_at, updated_at);
CREATE INDEX IF NOT EXISTS idx_entitlement_usage_entitlement_date
ON entitlement_usage_ledgers (user_entitlement_id, usage_date);
@@ -0,0 +1,69 @@
ALTER TABLE users ADD COLUMN privacy_policy_accepted_version TEXT;
ALTER TABLE users ADD COLUMN privacy_policy_accepted_at INTEGER;
ALTER TABLE announcements ADD COLUMN requires_ack INTEGER NOT NULL DEFAULT 0;
CREATE TABLE IF NOT EXISTS user_invite_codes (
user_id TEXT PRIMARY KEY,
invite_code TEXT NOT NULL UNIQUE,
active INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY(user_id) REFERENCES users(id) ON DELETE CASCADE
);
CREATE TABLE IF NOT EXISTS user_referrals (
id TEXT PRIMARY KEY,
inviter_user_id TEXT NOT NULL,
invitee_user_id TEXT NOT NULL UNIQUE,
invite_code_snapshot TEXT NOT NULL,
source_json TEXT,
first_paid_order_id TEXT,
first_paid_at INTEGER,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY(inviter_user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY(invitee_user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY(first_paid_order_id) REFERENCES payment_orders(id) ON DELETE SET NULL
);
CREATE INDEX IF NOT EXISTS idx_user_referrals_inviter
ON user_referrals (inviter_user_id, created_at);
CREATE INDEX IF NOT EXISTS idx_user_referrals_created
ON user_referrals (created_at);
CREATE INDEX IF NOT EXISTS idx_user_referrals_invite_code
ON user_referrals (invite_code_snapshot);
CREATE TABLE IF NOT EXISTS referral_rewards (
id TEXT PRIMARY KEY,
referral_id TEXT NOT NULL,
inviter_user_id TEXT NOT NULL,
invitee_user_id TEXT NOT NULL,
reward_type TEXT NOT NULL,
trigger_point TEXT NOT NULL,
source_order_id TEXT,
idempotency_key TEXT NOT NULL UNIQUE,
amount_usd REAL NOT NULL,
status TEXT NOT NULL DEFAULT 'pending',
wallet_transaction_id TEXT,
reversed_amount_usd REAL NOT NULL DEFAULT 0,
pending_reversal_amount_usd REAL NOT NULL DEFAULT 0,
failure_reason TEXT,
admin_operator_id TEXT,
admin_note TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
FOREIGN KEY(referral_id) REFERENCES user_referrals(id) ON DELETE CASCADE,
FOREIGN KEY(inviter_user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY(invitee_user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY(source_order_id) REFERENCES payment_orders(id) ON DELETE SET NULL
);
CREATE INDEX IF NOT EXISTS idx_referral_rewards_inviter_status
ON referral_rewards (inviter_user_id, status, created_at);
CREATE INDEX IF NOT EXISTS idx_referral_rewards_inviter_created
ON referral_rewards (inviter_user_id, created_at);
CREATE INDEX IF NOT EXISTS idx_referral_rewards_created
ON referral_rewards (created_at);
CREATE INDEX IF NOT EXISTS idx_referral_rewards_source_order
ON referral_rewards (source_order_id);
@@ -0,0 +1 @@
ALTER TABLE oauth_providers ADD COLUMN icon_url TEXT;
@@ -0,0 +1,2 @@
CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_default_sort
ON provider_api_keys (provider_id, internal_priority, name, id);
@@ -0,0 +1 @@
ALTER TABLE api_keys ADD COLUMN ip_rules TEXT;
@@ -0,0 +1,18 @@
-- Usage is a historical fact table. Backfill nullable provider_id snapshots
-- from the unique provider name where the catalog row still exists.
UPDATE "usage"
SET provider_id = (
SELECT providers.id
FROM providers
WHERE providers.name = TRIM("usage".provider_name)
LIMIT 1
)
WHERE provider_id IS NULL
AND TRIM(COALESCE(provider_name, '')) <> ''
AND LOWER(TRIM(COALESCE(provider_name, ''))) NOT IN ('unknown', 'unknow', 'pending')
AND EXISTS (
SELECT 1
FROM providers
WHERE providers.name = TRIM("usage".provider_name)
);
@@ -0,0 +1,16 @@
CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_active_priority_id
ON provider_api_keys (provider_id, is_active, internal_priority, id);
CREATE INDEX IF NOT EXISTS pool_member_scores_scheduler_account_rank_idx
ON pool_member_scores (
pool_kind,
pool_id,
capability,
scope_kind,
scope_id,
hard_state,
score DESC,
last_ranked_at DESC,
member_id,
id
);
@@ -0,0 +1,2 @@
CREATE INDEX IF NOT EXISTS idx_provider_api_keys_provider_name_id
ON provider_api_keys (provider_id, name, id);
@@ -0,0 +1,177 @@
WITH endpoint_url_parts AS (
SELECT
e.id,
CASE
WHEN instr(e.base_url, '?') > 0 THEN rtrim(substr(e.base_url, 1, instr(e.base_url, '?') - 1), '/')
ELSE rtrim(e.base_url, '/')
END AS base_without_query,
CASE
WHEN instr(e.base_url, '?') > 0 THEN substr(e.base_url, instr(e.base_url, '?'))
ELSE ''
END AS query_suffix,
lower(trim(e.api_format)) AS normalized_api_format,
lower(trim(coalesce(e.custom_path, ''))) AS normalized_custom_path,
lower(CASE
WHEN instr(e.base_url, '?') > 0 THEN rtrim(substr(e.base_url, 1, instr(e.base_url, '?') - 1), '/')
ELSE rtrim(e.base_url, '/')
END) AS normalized_base,
lower(trim(coalesce(p.provider_type, ''))) AS provider_type
FROM provider_endpoints e
LEFT JOIN providers p ON p.id = e.provider_id
WHERE lower(trim(e.api_format)) IN (
'openai:chat',
'openai:responses',
'openai:responses:compact',
'openai:embedding',
'openai:rerank',
'openai:image',
'openai:video',
'jina:embedding',
'jina:rerank',
'claude:messages',
'gemini:generate_content',
'gemini:embedding',
'gemini:video'
)
),
endpoint_api_root_updates AS (
SELECT
id,
base_without_query
|| CASE
WHEN normalized_api_format IN ('gemini:generate_content', 'gemini:embedding', 'gemini:video')
THEN '/v1beta'
ELSE '/v1'
END
|| query_suffix AS next_base_url
FROM endpoint_url_parts
WHERE provider_type NOT IN (
'codex',
'chatgpt_web',
'claude_code',
'kiro',
'gemini_cli',
'vertex_ai',
'antigravity',
'grok',
'windsurf'
)
AND normalized_base NOT GLOB '*/v[0-9]'
AND normalized_base NOT GLOB '*/v[0-9][0-9]'
AND normalized_base NOT GLOB '*/v[0-9]/*'
AND normalized_base NOT GLOB '*/v[0-9][0-9]/*'
AND normalized_base NOT GLOB '*/v[0-9]beta*'
AND normalized_base NOT GLOB '*/v[0-9][0-9]beta*'
AND (
(
normalized_api_format IN ('gemini:generate_content', 'gemini:embedding', 'gemini:video')
AND normalized_custom_path LIKE '/v1beta/%'
)
OR (
normalized_api_format NOT IN ('gemini:generate_content', 'gemini:embedding', 'gemini:video')
AND normalized_custom_path LIKE '/v1/%'
)
OR normalized_custom_path = ''
)
)
UPDATE provider_endpoints
SET base_url = (
SELECT next_base_url
FROM endpoint_api_root_updates
WHERE endpoint_api_root_updates.id = provider_endpoints.id
)
WHERE id IN (SELECT id FROM endpoint_api_root_updates);
UPDATE provider_endpoints
SET custom_path = CASE
WHEN lower(trim(api_format)) = 'openai:chat'
AND lower(trim(coalesce(custom_path, ''))) = '/v1/chat/completions'
THEN NULL
WHEN lower(trim(api_format)) = 'openai:responses'
AND lower(trim(coalesce(custom_path, ''))) = '/v1/responses'
THEN NULL
WHEN lower(trim(api_format)) = 'openai:responses:compact'
AND lower(trim(coalesce(custom_path, ''))) = '/v1/responses/compact'
THEN NULL
WHEN lower(trim(api_format)) = 'claude:messages'
AND lower(trim(coalesce(custom_path, ''))) = '/v1/messages'
THEN NULL
WHEN lower(trim(api_format)) IN ('openai:embedding', 'jina:embedding')
AND lower(trim(coalesce(custom_path, ''))) = '/v1/embeddings'
THEN NULL
WHEN lower(trim(api_format)) IN ('openai:rerank', 'jina:rerank')
AND lower(trim(coalesce(custom_path, ''))) = '/v1/rerank'
THEN NULL
WHEN lower(trim(api_format)) = 'openai:image'
AND lower(trim(coalesce(custom_path, ''))) = '/v1/images/generations'
THEN NULL
WHEN lower(trim(api_format)) = 'openai:video'
AND lower(trim(coalesce(custom_path, ''))) = '/v1/videos'
THEN NULL
WHEN lower(trim(api_format)) = 'gemini:generate_content'
AND lower(trim(coalesce(custom_path, ''))) = '/v1beta/models/{model}:{action}'
THEN NULL
WHEN lower(trim(api_format)) = 'gemini:embedding'
AND lower(trim(coalesce(custom_path, ''))) IN ('/v1beta/models/{model}:embedcontent', '/v1beta/models/{model}:{action}')
THEN NULL
WHEN lower(trim(api_format)) = 'gemini:video'
AND lower(trim(coalesce(custom_path, ''))) = '/v1beta/models/{model}:predictlongrunning'
THEN NULL
WHEN lower(trim(api_format)) IN ('gemini:generate_content', 'gemini:embedding', 'gemini:video')
THEN '/' || substr(trim(custom_path), 9)
ELSE '/' || substr(trim(custom_path), 5)
END
WHERE lower(trim(api_format)) IN (
'openai:chat',
'openai:responses',
'openai:responses:compact',
'openai:embedding',
'openai:rerank',
'openai:image',
'openai:video',
'jina:embedding',
'jina:rerank',
'claude:messages',
'gemini:generate_content',
'gemini:embedding',
'gemini:video'
)
AND NOT EXISTS (
SELECT 1
FROM providers p
WHERE p.id = provider_endpoints.provider_id
AND lower(trim(coalesce(p.provider_type, ''))) IN (
'codex',
'chatgpt_web',
'claude_code',
'kiro',
'gemini_cli',
'vertex_ai',
'antigravity',
'grok',
'windsurf'
)
)
AND (
(
lower(trim(api_format)) IN ('gemini:generate_content', 'gemini:embedding', 'gemini:video')
AND lower(trim(coalesce(custom_path, ''))) LIKE '/v1beta/%'
)
OR (
lower(trim(api_format)) NOT IN ('gemini:generate_content', 'gemini:embedding', 'gemini:video')
AND lower(trim(coalesce(custom_path, ''))) LIKE '/v1/%'
)
)
AND (
lower(rtrim(CASE WHEN instr(base_url, '?') > 0 THEN substr(base_url, 1, instr(base_url, '?') - 1) ELSE base_url END, '/')) GLOB '*/v[0-9]'
OR lower(rtrim(CASE WHEN instr(base_url, '?') > 0 THEN substr(base_url, 1, instr(base_url, '?') - 1) ELSE base_url END, '/')) GLOB '*/v[0-9][0-9]'
OR lower(rtrim(CASE WHEN instr(base_url, '?') > 0 THEN substr(base_url, 1, instr(base_url, '?') - 1) ELSE base_url END, '/')) GLOB '*/v[0-9]/*'
OR lower(rtrim(CASE WHEN instr(base_url, '?') > 0 THEN substr(base_url, 1, instr(base_url, '?') - 1) ELSE base_url END, '/')) GLOB '*/v[0-9][0-9]/*'
OR lower(rtrim(CASE WHEN instr(base_url, '?') > 0 THEN substr(base_url, 1, instr(base_url, '?') - 1) ELSE base_url END, '/')) GLOB '*/v[0-9]beta*'
OR lower(rtrim(CASE WHEN instr(base_url, '?') > 0 THEN substr(base_url, 1, instr(base_url, '?') - 1) ELSE base_url END, '/')) GLOB '*/v[0-9][0-9]beta*'
);
UPDATE provider_endpoints
SET custom_path = NULL
WHERE custom_path IS NOT NULL
AND trim(custom_path) = '';
@@ -0,0 +1,26 @@
-- High-concurrency gateway read/cleanup paths.
-- SQLite remains single-node/lightweight but benefits from the same bounded scans.
CREATE INDEX IF NOT EXISTS idx_usage_created_id_desc
ON "usage" (created_at_unix_ms DESC, request_id ASC);
CREATE INDEX IF NOT EXISTS idx_usage_user_created_id_desc
ON "usage" (user_id, created_at_unix_ms DESC, request_id ASC);
CREATE INDEX IF NOT EXISTS idx_usage_api_format_created_id_desc
ON "usage" (api_format, created_at_unix_ms DESC, request_id ASC);
CREATE INDEX IF NOT EXISTS idx_usage_status_created_id_desc
ON "usage" (status, created_at_unix_ms DESC, request_id ASC);
CREATE INDEX IF NOT EXISTS idx_request_candidates_provider_created
ON request_candidates (provider_id, created_at DESC, id ASC);
CREATE INDEX IF NOT EXISTS idx_request_candidates_api_key_created
ON request_candidates (api_key_id, created_at ASC, id ASC);
CREATE INDEX IF NOT EXISTS idx_background_task_runs_status_created
ON background_task_runs (status, created_at_unix_secs DESC, updated_at_unix_secs DESC);
CREATE INDEX IF NOT EXISTS idx_background_task_runs_kind_created
ON background_task_runs (kind, created_at_unix_secs DESC, updated_at_unix_secs DESC);
@@ -0,0 +1,480 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::announcements::*;
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_eq, push_limit, push_limit_offset, WhereClause};
use crate::error::SqlResultExt;
use crate::SqlitePool;
const ANNOUNCEMENT_SELECT: &str = r#"
SELECT
a.id,
a.title,
a.content,
a.type,
a.priority,
a.is_active,
a.is_pinned,
a.requires_ack,
a.author_id,
u.username AS author_username,
a.start_time AS start_time_unix_secs,
a.end_time AS end_time_unix_secs,
a.created_at AS created_at_unix_ms,
a.updated_at AS updated_at_unix_secs
FROM announcements a
LEFT JOIN users u ON u.id = a.author_id
"#;
#[derive(Debug, Clone)]
pub struct SqliteAnnouncementRepository {
pool: SqlitePool,
}
impl SqliteAnnouncementRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn reload_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
self.find_by_id(announcement_id).await
}
fn apply_active_filter(
builder: &mut QueryBuilder<'_, Sqlite>,
where_clause: &mut WhereClause,
active_only: bool,
now_unix_secs: u64,
) -> Result<(), DataLayerError> {
if !active_only {
return Ok(());
}
let now = i64_from_u64(now_unix_secs, "announcements.now")?;
where_clause.push_next(builder);
builder
.push("a.is_active = 1 AND (a.start_time IS NULL OR a.start_time <= ")
.push_bind(now)
.push(") AND (a.end_time IS NULL OR a.end_time >= ")
.push_bind(now)
.push(")");
Ok(())
}
}
#[async_trait]
impl AnnouncementReadRepository for SqliteAnnouncementRepository {
async fn find_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(ANNOUNCEMENT_SELECT);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"a.id",
announcement_id.to_string(),
);
push_limit(&mut builder, 1);
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_announcement_row).transpose()
}
async fn list_announcements(
&self,
query: &AnnouncementListQuery,
) -> Result<StoredAnnouncementPage, DataLayerError> {
let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs);
let mut count_builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(a.id) AS total FROM announcements a");
let mut count_where = WhereClause::new();
Self::apply_active_filter(
&mut count_builder,
&mut count_where,
query.active_only,
now_unix_secs,
)?;
let total = count_builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?
.max(0) as u64;
let mut list_builder = QueryBuilder::<Sqlite>::new(ANNOUNCEMENT_SELECT);
let mut list_where = WhereClause::new();
Self::apply_active_filter(
&mut list_builder,
&mut list_where,
query.active_only,
now_unix_secs,
)?;
list_builder
.push(" ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC");
push_limit_offset(&mut list_builder, query.limit as i64, query.offset as i64);
let rows = list_builder
.build()
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_announcement_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredAnnouncementPage { items, total })
}
async fn count_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(a.id) AS total FROM announcements a");
let mut where_clause = WhereClause::new();
Self::apply_active_filter(&mut builder, &mut where_clause, true, now_unix_secs)?;
where_clause.push_next(&mut builder);
builder
.push("NOT EXISTS (SELECT 1 FROM announcement_reads r WHERE r.user_id = ")
.push_bind(user_id.to_string())
.push(" AND r.announcement_id = a.id)");
let total = builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?
.max(0) as u64;
Ok(total)
}
async fn list_required_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredAnnouncement>, DataLayerError> {
let rows = sqlx::query(&format!(
r#"
{ANNOUNCEMENT_SELECT}
WHERE a.is_active = 1
AND a.requires_ack = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
AND NOT EXISTS (
SELECT 1
FROM announcement_reads r
WHERE r.user_id = ?
AND r.announcement_id = a.id
)
ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
LIMIT ?
"#
))
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.bind(user_id)
.bind(limit as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_announcement_row).collect()
}
}
#[async_trait]
impl AnnouncementWriteRepository for SqliteAnnouncementRepository {
async fn create_announcement(
&self,
record: CreateAnnouncementRecord,
) -> Result<StoredAnnouncement, DataLayerError> {
record.validate()?;
let id = uuid::Uuid::new_v4().to_string();
let now = current_unix_secs() as i64;
sqlx::query(
r#"
INSERT INTO announcements (
id, title, content, type, priority, author_id, is_active, is_pinned,
requires_ack, start_time, end_time, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 1, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(record.title)
.bind(record.content)
.bind(record.kind)
.bind(record.priority)
.bind(record.author_id)
.bind(record.is_pinned)
.bind(record.requires_ack)
.bind(optional_i64_from_u64(
record.start_time_unix_secs,
"announcements.start_time",
)?)
.bind(optional_i64_from_u64(
record.end_time_unix_secs,
"announcements.end_time",
)?)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_id(&id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("created announcement missing".into()))
}
async fn update_announcement(
&self,
record: UpdateAnnouncementRecord,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
record.validate()?;
let id = record.announcement_id;
sqlx::query(
r#"
UPDATE announcements
SET title = COALESCE(?, title),
content = COALESCE(?, content),
type = COALESCE(?, type),
priority = COALESCE(?, priority),
is_active = COALESCE(?, is_active),
is_pinned = COALESCE(?, is_pinned),
requires_ack = COALESCE(?, requires_ack),
start_time = COALESCE(?, start_time),
end_time = COALESCE(?, end_time),
updated_at = ?
WHERE id = ?
"#,
)
.bind(record.title)
.bind(record.content)
.bind(record.kind)
.bind(record.priority)
.bind(record.is_active)
.bind(record.is_pinned)
.bind(record.requires_ack)
.bind(optional_i64_from_u64(
record.start_time_unix_secs,
"announcements.start_time",
)?)
.bind(optional_i64_from_u64(
record.end_time_unix_secs,
"announcements.end_time",
)?)
.bind(current_unix_secs() as i64)
.bind(&id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_id(&id).await
}
async fn delete_announcement(&self, announcement_id: &str) -> Result<bool, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM announcement_reads WHERE announcement_id = ?")
.bind(announcement_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let rows_affected = sqlx::query("DELETE FROM announcements WHERE id = ?")
.bind(announcement_id)
.execute(&mut *tx)
.await
.map_sql_err()?
.rows_affected();
tx.commit().await.map_sql_err()?;
Ok(rows_affected > 0)
}
async fn mark_announcement_as_read(
&self,
user_id: &str,
announcement_id: &str,
read_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
r#"
INSERT OR IGNORE INTO announcement_reads (id, user_id, announcement_id, read_at)
VALUES (?, ?, ?, ?)
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
.bind(user_id)
.bind(announcement_id)
.bind(i64_from_u64(
read_at_unix_secs,
"announcement_reads.read_at",
)?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn optional_i64_from_u64(
value: Option<u64>,
field_name: &str,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| i64_from_u64(value, field_name))
.transpose()
}
fn map_announcement_row(row: &SqliteRow) -> Result<StoredAnnouncement, DataLayerError> {
StoredAnnouncement::new(
row.try_get("id").map_sql_err()?,
row.try_get("title").map_sql_err()?,
row.try_get("content").map_sql_err()?,
row.try_get("type").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
row.try_get("is_pinned").map_sql_err()?,
row.try_get("requires_ack").map_sql_err()?,
row.try_get("author_id").map_sql_err()?,
row.try_get("author_username").map_sql_err()?,
row.try_get("start_time_unix_secs").map_sql_err()?,
row.try_get("end_time_unix_secs").map_sql_err()?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
)
}
#[cfg(test)]
mod tests {
use super::SqliteAnnouncementRepository;
use crate::run_migrations;
use aether_data_contracts::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, UpdateAnnouncementRecord,
};
#[tokio::test]
async fn sqlite_repository_reads_and_writes_announcements() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_announcement_user(&pool).await;
let repository = SqliteAnnouncementRepository::new(pool);
let created = repository
.create_announcement(CreateAnnouncementRecord {
title: "Initial".to_string(),
content: "Body".to_string(),
kind: "info".to_string(),
priority: 10,
is_pinned: true,
requires_ack: false,
author_id: "user-1".to_string(),
start_time_unix_secs: Some(100),
end_time_unix_secs: Some(300),
})
.await
.expect("announcement should create");
assert_eq!(created.author_username, Some("admin".to_string()));
assert!(created.is_active);
let page = repository
.list_announcements(&AnnouncementListQuery {
active_only: true,
offset: 0,
limit: 10,
now_unix_secs: Some(200),
})
.await
.expect("announcements should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].id, created.id);
let unread = repository
.count_unread_active_announcements("user-1", 200)
.await
.expect("unread count should load");
assert_eq!(unread, 1);
assert!(repository
.mark_announcement_as_read("user-1", &created.id, 210)
.await
.expect("read marker should insert"));
assert!(!repository
.mark_announcement_as_read("user-1", &created.id, 211)
.await
.expect("duplicate read marker should be ignored"));
assert_eq!(
repository
.count_unread_active_announcements("user-1", 200)
.await
.expect("unread count should reload"),
0
);
let updated = repository
.update_announcement(UpdateAnnouncementRecord {
announcement_id: created.id.clone(),
title: Some("Updated".to_string()),
content: None,
kind: None,
priority: Some(20),
is_active: Some(false),
is_pinned: Some(false),
requires_ack: Some(true),
start_time_unix_secs: None,
end_time_unix_secs: None,
})
.await
.expect("announcement should update")
.expect("announcement should exist");
assert_eq!(updated.title, "Updated");
assert!(!updated.is_active);
assert!(repository
.delete_announcement(&created.id)
.await
.expect("announcement should delete"));
assert!(repository
.find_by_id(&created.id)
.await
.expect("find should run")
.is_none());
}
async fn seed_announcement_user(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO users (
id, email, username, role, auth_source, email_verified, is_active, is_deleted, created_at, updated_at
)
VALUES ('user-1', 'admin@example.com', 'admin', 'admin', 'local', 1, 1, 0, 1, 1)
"#,
)
.execute(pool)
.await
.expect("user should seed");
}
}
@@ -0,0 +1,405 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use aether_data_contracts::repository::audit::*;
use aether_data_contracts::DataLayerError;
use crate::error::SqlResultExt;
use crate::SqlitePool;
#[derive(Debug, Clone)]
pub struct SqliteAuditLogReadRepository {
pool: SqlitePool,
}
impl SqliteAuditLogReadRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[async_trait]
impl AuditLogReadRepository for SqliteAuditLogReadRepository {
async fn list_admin_audit_logs(
&self,
query: &AuditLogListQuery,
) -> Result<StoredAdminAuditLogPage, DataLayerError> {
let total = sqlx::query_scalar::<_, i64>(
r#"
SELECT COUNT(*)
FROM audit_logs AS a
LEFT JOIN users AS u ON a.user_id = u.id
WHERE a.created_at >= ?
AND (? IS NULL OR LOWER(u.username) LIKE LOWER(?) ESCAPE '\')
AND (? IS NULL OR a.event_type = ?)
"#,
)
.bind(query.cutoff_unix_secs as i64)
.bind(query.username_pattern.as_deref())
.bind(query.username_pattern.as_deref())
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let rows = sqlx::query(
r#"
SELECT
a.id,
a.event_type,
a.user_id,
u.email AS user_email,
u.username AS user_username,
a.description,
a.ip_address,
a.status_code,
a.error_message,
a.event_metadata AS metadata,
a.created_at
FROM audit_logs AS a
LEFT JOIN users AS u ON a.user_id = u.id
WHERE a.created_at >= ?
AND (? IS NULL OR LOWER(u.username) LIKE LOWER(?) ESCAPE '\')
AND (? IS NULL OR a.event_type = ?)
ORDER BY a.created_at DESC
LIMIT ? OFFSET ?
"#,
)
.bind(query.cutoff_unix_secs as i64)
.bind(query.username_pattern.as_deref())
.bind(query.username_pattern.as_deref())
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.bind(query.limit as i64)
.bind(query.offset as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_sqlite_admin_audit_log_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredAdminAuditLogPage {
items,
total: total.max(0) as u64,
})
}
async fn list_admin_suspicious_activities(
&self,
cutoff_unix_secs: u64,
) -> Result<Vec<StoredSuspiciousActivity>, DataLayerError> {
let rows = sqlx::query(
r#"
SELECT id, event_type, user_id, description, ip_address, event_metadata AS metadata, created_at
FROM audit_logs
WHERE created_at >= ?
AND event_type IN (?, ?, ?, ?)
ORDER BY created_at DESC
LIMIT 100
"#,
)
.bind(cutoff_unix_secs as i64)
.bind(SUSPICIOUS_EVENT_TYPES[0])
.bind(SUSPICIOUS_EVENT_TYPES[1])
.bind(SUSPICIOUS_EVENT_TYPES[2])
.bind(SUSPICIOUS_EVENT_TYPES[3])
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter()
.map(map_sqlite_suspicious_activity_row)
.collect()
}
async fn read_admin_user_behavior_event_counts(
&self,
user_id: &str,
cutoff_unix_secs: u64,
) -> Result<std::collections::BTreeMap<String, u64>, DataLayerError> {
let rows = sqlx::query(
r#"
SELECT event_type, COUNT(*) AS count
FROM audit_logs
WHERE user_id = ?
AND created_at >= ?
GROUP BY event_type
"#,
)
.bind(user_id)
.bind(cutoff_unix_secs as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
Ok(rows
.iter()
.filter_map(|row| event_count_from_sqlite_row(row).ok())
.collect())
}
async fn list_user_audit_logs(
&self,
user_id: &str,
query: &AuditLogListQuery,
) -> Result<StoredUserAuditLogPage, DataLayerError> {
let total = sqlx::query_scalar::<_, i64>(
r#"
SELECT COUNT(*)
FROM audit_logs
WHERE user_id = ?
AND created_at >= ?
AND (? IS NULL OR event_type = ?)
"#,
)
.bind(user_id)
.bind(query.cutoff_unix_secs as i64)
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let rows = sqlx::query(
r#"
SELECT id, event_type, description, ip_address, status_code, created_at
FROM audit_logs
WHERE user_id = ?
AND created_at >= ?
AND (? IS NULL OR event_type = ?)
ORDER BY created_at DESC
LIMIT ? OFFSET ?
"#,
)
.bind(user_id)
.bind(query.cutoff_unix_secs as i64)
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.bind(query.limit as i64)
.bind(query.offset as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_sqlite_user_audit_log_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredUserAuditLogPage {
items,
total: total.max(0) as u64,
})
}
async fn delete_audit_logs_before(
&self,
cutoff_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
let deleted = sqlx::query(
r#"
DELETE FROM audit_logs
WHERE id IN (
SELECT id
FROM audit_logs
WHERE created_at < ?
ORDER BY created_at ASC, id ASC
LIMIT ?
)
"#,
)
.bind(cutoff_unix_secs.min(i64::MAX as u64) as i64)
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(deleted).unwrap_or(usize::MAX))
}
}
fn sqlite_created_at_unix_secs(row: &SqliteRow) -> Result<u64, DataLayerError> {
let value = row.try_get::<i64, _>("created_at").map_sql_err()?;
Ok(value.max(0) as u64)
}
fn map_sqlite_admin_audit_log_row(row: &SqliteRow) -> Result<StoredAdminAuditLog, DataLayerError> {
Ok(StoredAdminAuditLog {
id: row.try_get("id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
user_id: row.try_get("user_id").map_sql_err()?,
user_email: row.try_get("user_email").map_sql_err()?,
user_username: row.try_get("user_username").map_sql_err()?,
description: row.try_get("description").map_sql_err()?,
ip_address: row.try_get("ip_address").map_sql_err()?,
status_code: row.try_get("status_code").map_sql_err()?,
error_message: row.try_get("error_message").map_sql_err()?,
metadata: optional_json_from_text(row.try_get("metadata").map_sql_err()?)?,
created_at_unix_secs: sqlite_created_at_unix_secs(row)?,
})
}
fn map_sqlite_suspicious_activity_row(
row: &SqliteRow,
) -> Result<StoredSuspiciousActivity, DataLayerError> {
Ok(StoredSuspiciousActivity {
id: row.try_get("id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
user_id: row.try_get("user_id").map_sql_err()?,
description: row.try_get("description").map_sql_err()?,
ip_address: row.try_get("ip_address").map_sql_err()?,
metadata: optional_json_from_text(row.try_get("metadata").map_sql_err()?)?,
created_at_unix_secs: sqlite_created_at_unix_secs(row)?,
})
}
fn map_sqlite_user_audit_log_row(row: &SqliteRow) -> Result<StoredUserAuditLog, DataLayerError> {
Ok(StoredUserAuditLog {
id: row.try_get("id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
description: row.try_get("description").map_sql_err()?,
ip_address: row.try_get("ip_address").map_sql_err()?,
status_code: row.try_get("status_code").map_sql_err()?,
created_at_unix_secs: sqlite_created_at_unix_secs(row)?,
})
}
fn event_count_from_sqlite_row(row: &SqliteRow) -> Result<(String, u64), DataLayerError> {
let event_type = row.try_get("event_type").map_sql_err()?;
let count = row.try_get::<i64, _>("count").map_sql_err()?.max(0) as u64;
Ok((event_type, count))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::run_migrations;
#[tokio::test]
async fn sqlite_audit_log_repository_reads_monitoring_views() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_sqlite_audit_logs(&pool).await;
let repository = SqliteAuditLogReadRepository::new(pool);
let admin_page = repository
.list_admin_audit_logs(&AuditLogListQuery {
cutoff_unix_secs: 150,
username_pattern: Some("%ali%".to_string()),
event_type: Some("login_failed".to_string()),
limit: 10,
offset: 0,
})
.await
.expect("admin audit logs should read");
assert_eq!(admin_page.total, 1);
assert_eq!(admin_page.items[0].id, "audit-2");
assert_eq!(admin_page.items[0].user_username.as_deref(), Some("alice"));
assert_eq!(
admin_page.items[0]
.metadata
.as_ref()
.and_then(|value| value.get("risk"))
.and_then(|value| value.as_str()),
Some("high")
);
let suspicious = repository
.list_admin_suspicious_activities(150)
.await
.expect("suspicious activities should read");
assert_eq!(suspicious.len(), 1);
assert_eq!(suspicious[0].event_type, "login_failed");
let counts = repository
.read_admin_user_behavior_event_counts("user-1", 0)
.await
.expect("user behavior counts should read");
assert_eq!(counts.get("login_failed"), Some(&1));
assert_eq!(counts.get("request_success"), Some(&1));
let user_page = repository
.list_user_audit_logs(
"user-1",
&AuditLogListQuery {
cutoff_unix_secs: 0,
username_pattern: None,
event_type: Some("request_success".to_string()),
limit: 10,
offset: 0,
},
)
.await
.expect("user audit logs should read");
assert_eq!(user_page.total, 1);
assert_eq!(user_page.items[0].id, "audit-1");
assert_eq!(user_page.items[0].status_code, Some(200));
let deleted = repository
.delete_audit_logs_before(250, 1)
.await
.expect("audit cleanup should delete one old row");
assert_eq!(deleted, 1);
let user_page = repository
.list_user_audit_logs(
"user-1",
&AuditLogListQuery {
cutoff_unix_secs: 0,
username_pattern: None,
event_type: Some("request_success".to_string()),
limit: 10,
offset: 0,
},
)
.await
.expect("user audit logs should read after cleanup");
assert_eq!(user_page.total, 0);
}
async fn seed_sqlite_audit_logs(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO users (id, email, username, role, auth_source, created_at, updated_at)
VALUES
('user-1', 'alice@example.com', 'alice', 'user', 'local', 1, 1),
('user-2', 'bob@example.com', 'bob', 'user', 'local', 1, 1)
"#,
)
.execute(pool)
.await
.expect("users should insert");
sqlx::query(
r#"
INSERT INTO audit_logs (
id,
event_type,
user_id,
description,
ip_address,
event_metadata,
status_code,
created_at
)
VALUES
('audit-1', 'request_success', 'user-1', 'completed request', '127.0.0.1', NULL, 200, 100),
('audit-2', 'login_failed', 'user-1', 'failed login', '127.0.0.2', '{"risk":"high"}', 401, 200),
('audit-3', 'password_changed', 'user-2', 'other user changed password', '127.0.0.3', NULL, 200, 300)
"#,
)
.execute(pool)
.await
.expect("audit logs should insert");
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,305 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::auth_modules::*;
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_eq, push_limit, WhereClause};
use crate::error::SqlResultExt;
use crate::SqlitePool;
const OAUTH_PROVIDER_COLUMNS: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
redirect_uri
FROM oauth_providers
"#;
const LDAP_CONFIG_COLUMNS: &str = r#"
SELECT
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
FROM ldap_configs
"#;
#[derive(Debug, Clone)]
pub struct SqliteAuthModuleReadRepository {
pool: SqlitePool,
}
impl SqliteAuthModuleReadRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[derive(Debug, Clone)]
pub struct SqliteAuthModuleRepository {
pool: SqlitePool,
}
impl SqliteAuthModuleRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
async fn list_enabled_oauth_providers(
pool: &SqlitePool,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(OAUTH_PROVIDER_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(&mut builder, &mut where_clause, "is_enabled", true);
builder.push(" ORDER BY provider_type ASC");
let rows = builder.build().fetch_all(pool).await.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
}
async fn get_ldap_config(
pool: &SqlitePool,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(LDAP_CONFIG_COLUMNS);
builder.push(" ORDER BY id ASC");
push_limit(&mut builder, 1);
let row = builder.build().fetch_optional(pool).await.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
}
#[async_trait]
impl AuthModuleReadRepository for SqliteAuthModuleReadRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
list_enabled_oauth_providers(&self.pool).await
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
get_ldap_config(&self.pool).await
}
}
#[async_trait]
impl AuthModuleReadRepository for SqliteAuthModuleRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
list_enabled_oauth_providers(&self.pool).await
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
get_ldap_config(&self.pool).await
}
}
#[async_trait]
impl AuthModuleWriteRepository for SqliteAuthModuleRepository {
async fn upsert_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let now = now_unix_secs();
let updated = 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 ldap_configs
ORDER BY id ASC
LIMIT 1
)
"#,
)
.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)
.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
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn map_oauth_row(row: &SqliteRow) -> Result<StoredOAuthProviderModuleConfig, DataLayerError> {
StoredOAuthProviderModuleConfig::new(
row.try_get("provider_type").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("client_id").map_sql_err()?,
row.try_get("client_secret_encrypted").map_sql_err()?,
row.try_get("redirect_uri").map_sql_err()?,
)
}
fn map_ldap_row(row: &SqliteRow) -> Result<StoredLdapModuleConfig, DataLayerError> {
Ok(StoredLdapModuleConfig {
server_url: row.try_get("server_url").map_sql_err()?,
bind_dn: row.try_get("bind_dn").map_sql_err()?,
bind_password_encrypted: row.try_get("bind_password_encrypted").map_sql_err()?,
base_dn: row.try_get("base_dn").map_sql_err()?,
user_search_filter: row.try_get("user_search_filter").map_sql_err()?,
username_attr: row.try_get("username_attr").map_sql_err()?,
email_attr: row.try_get("email_attr").map_sql_err()?,
display_name_attr: row.try_get("display_name_attr").map_sql_err()?,
is_enabled: row.try_get("is_enabled").map_sql_err()?,
is_exclusive: row.try_get("is_exclusive").map_sql_err()?,
use_starttls: row.try_get("use_starttls").map_sql_err()?,
connect_timeout: row.try_get("connect_timeout").map_sql_err()?,
})
}
#[cfg(test)]
mod tests {
use super::SqliteAuthModuleRepository;
use aether_data_contracts::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
};
use crate::run_migrations;
#[tokio::test]
async fn sqlite_repository_reads_and_writes_auth_module_configs() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES
('github', 'GitHub', 'github-client', 'https://github.example.com/callback',
'https://frontend.example.com/callback', 1, 1, 1),
('disabled', 'Disabled', 'disabled-client', 'https://disabled.example.com/callback',
'https://frontend.example.com/callback', 0, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("oauth providers should seed");
let repository = SqliteAuthModuleRepository::new(pool);
let oauth = repository
.list_enabled_oauth_providers()
.await
.expect("oauth providers should load");
assert_eq!(oauth.len(), 1);
assert_eq!(oauth[0].provider_type, "github");
let ldap = StoredLdapModuleConfig {
server_url: "ldaps://ldap.example.com".to_string(),
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
bind_password_encrypted: Some("encrypted-password".to_string()),
base_dn: "dc=example,dc=com".to_string(),
user_search_filter: Some("(uid={username})".to_string()),
username_attr: Some("uid".to_string()),
email_attr: Some("mail".to_string()),
display_name_attr: Some("displayName".to_string()),
is_enabled: true,
is_exclusive: false,
use_starttls: true,
connect_timeout: Some(10),
};
let stored = repository
.upsert_ldap_config(&ldap)
.await
.expect("ldap should upsert")
.expect("ldap should be returned");
assert_eq!(stored.server_url, "ldaps://ldap.example.com");
let updated = repository
.upsert_ldap_config(&StoredLdapModuleConfig {
server_url: "ldap://ldap.example.com".to_string(),
..ldap
})
.await
.expect("ldap should update")
.expect("ldap should be returned");
assert_eq!(updated.server_url, "ldap://ldap.example.com");
}
}
@@ -0,0 +1,520 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::background_tasks::*;
use aether_data_query::{
push_ci_contains, push_eq, push_limit, push_limit_offset, SqlDialect, WhereClause,
};
use crate::error::SqlResultExt;
use crate::{DataLayerError, SqlitePool};
const RUN_COLUMNS: &str = r#"
SELECT
id,
task_key,
kind,
"trigger",
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
FROM background_task_runs
"#;
const EVENT_COLUMNS: &str = r#"
SELECT
id,
run_id,
event_type,
message,
payload_json,
created_at_unix_secs
FROM background_task_events
"#;
#[derive(Debug, Clone)]
pub struct SqliteBackgroundTaskRepository {
pool: SqlitePool,
}
impl SqliteBackgroundTaskRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
fn apply_run_filter(builder: &mut QueryBuilder<'_, Sqlite>, query: &BackgroundTaskListQuery) {
let mut where_clause = WhereClause::new();
if let Some(kind) = query.kind {
push_eq(builder, &mut where_clause, "kind", kind.as_database());
}
if let Some(status) = query.status {
push_eq(builder, &mut where_clause, "status", status.as_database());
}
if let Some(trigger) = query.trigger.as_deref() {
push_eq(
builder,
&mut where_clause,
&SqlDialect::Sqlite.quote_ident("trigger"),
trigger.to_string(),
);
}
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
push_ci_contains(
builder,
&mut where_clause,
SqlDialect::Sqlite,
"task_key",
task_key_substring,
);
}
}
}
#[async_trait]
impl BackgroundTaskReadRepository for SqliteBackgroundTaskRepository {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(RUN_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(&mut builder, &mut where_clause, "id", run_id.to_string());
push_limit(&mut builder, 1);
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_run_row).transpose()
}
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
let limit = query.limit.max(1);
let mut count_builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(id) AS total FROM background_task_runs");
Self::apply_run_filter(&mut count_builder, query);
let total = count_builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let mut builder = QueryBuilder::<Sqlite>::new(RUN_COLUMNS);
Self::apply_run_filter(&mut builder, query);
builder.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC");
push_limit_offset(
&mut builder,
i64_from_usize(limit, "run limit")?,
i64_from_usize(query.offset, "run offset")?,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let items = rows
.iter()
.map(map_run_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredBackgroundTaskRunPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let limit = limit.max(1);
let mut builder = QueryBuilder::<Sqlite>::new(EVENT_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"run_id",
run_id.to_string(),
);
builder.push(" ORDER BY created_at_unix_secs ASC, id ASC");
push_limit_offset(
&mut builder,
i64_from_usize(limit, "event limit")?,
i64_from_usize(offset, "event offset")?,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_event_row).collect()
}
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
let total = sqlx::query_scalar::<_, i64>("SELECT COUNT(id) FROM background_task_runs")
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let running_count = sqlx::query_scalar::<_, i64>(
"SELECT COUNT(id) FROM background_task_runs WHERE status = 'running'",
)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let status_rows = sqlx::query(
"SELECT status, COUNT(id) AS total FROM background_task_runs GROUP BY status",
)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let kind_rows =
sqlx::query("SELECT kind, COUNT(id) AS total FROM background_task_runs GROUP BY kind")
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let mut by_status = std::collections::BTreeMap::new();
for row in status_rows {
let key: String = row.try_get("status").map_sql_err()?;
let count: i64 = row.try_get("total").map_sql_err()?;
by_status.insert(key, u64::try_from(count).unwrap_or_default());
}
let mut by_kind = std::collections::BTreeMap::new();
for row in kind_rows {
let key: String = row.try_get("kind").map_sql_err()?;
let count: i64 = row.try_get("total").map_sql_err()?;
by_kind.insert(key, u64::try_from(count).unwrap_or_default());
}
Ok(BackgroundTaskSummary {
total: u64::try_from(total).unwrap_or_default(),
running_count: u64::try_from(running_count).unwrap_or_default(),
by_status,
by_kind,
})
}
}
#[async_trait]
impl BackgroundTaskWriteRepository for SqliteBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_runs (
id,
task_key,
kind,
"trigger",
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
ON CONFLICT(id) DO UPDATE SET
task_key = excluded.task_key,
kind = excluded.kind,
"trigger" = excluded."trigger",
status = excluded.status,
attempt = excluded.attempt,
max_attempts = excluded.max_attempts,
owner_instance = excluded.owner_instance,
progress_percent = excluded.progress_percent,
progress_message = excluded.progress_message,
payload_json = excluded.payload_json,
result_json = excluded.result_json,
error_message = excluded.error_message,
cancel_requested = excluded.cancel_requested,
created_by = excluded.created_by,
created_at_unix_secs = excluded.created_at_unix_secs,
started_at_unix_secs = excluded.started_at_unix_secs,
finished_at_unix_secs = excluded.finished_at_unix_secs,
updated_at_unix_secs = excluded.updated_at_unix_secs
"#,
)
.bind(&run.id)
.bind(&run.task_key)
.bind(run.kind.as_database())
.bind(&run.trigger)
.bind(run.status.as_database())
.bind(i64::from(run.attempt))
.bind(i64::from(run.max_attempts))
.bind(run.owner_instance.as_deref())
.bind(i32::from(run.progress_percent))
.bind(run.progress_message.as_deref())
.bind(run.payload_json.as_ref().map(serde_json::Value::to_string))
.bind(run.result_json.as_ref().map(serde_json::Value::to_string))
.bind(run.error_message.as_deref())
.bind(run.cancel_requested)
.bind(run.created_by.as_deref())
.bind(u64_to_i64(
run.created_at_unix_secs,
"created_at_unix_secs",
)?)
.bind(run.started_at_unix_secs.map(|value| value as i64))
.bind(run.finished_at_unix_secs.map(|value| value as i64))
.bind(u64_to_i64(
run.updated_at_unix_secs,
"updated_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_run(&run.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("background task run missing after upsert".to_string())
})
}
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let affected = sqlx::query(
"UPDATE background_task_runs SET cancel_requested = 1, updated_at_unix_secs = ? WHERE id = ?",
)
.bind(u64_to_i64(updated_at_unix_secs, "updated_at_unix_secs")?)
.bind(run_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(affected > 0)
}
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_events (
id, run_id, event_type, message, payload_json, created_at_unix_secs
) VALUES (?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
run_id = excluded.run_id,
event_type = excluded.event_type,
message = excluded.message,
payload_json = excluded.payload_json,
created_at_unix_secs = excluded.created_at_unix_secs
"#,
)
.bind(&event.id)
.bind(&event.run_id)
.bind(&event.event_type)
.bind(&event.message)
.bind(
event
.payload_json
.as_ref()
.map(serde_json::Value::to_string),
)
.bind(u64_to_i64(
event.created_at_unix_secs,
"created_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
let row = sqlx::query(&format!("{EVENT_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(&event.id)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
map_event_row(&row)
}
}
fn map_run_row(row: &SqliteRow) -> Result<StoredBackgroundTaskRun, DataLayerError> {
let kind: String = row.try_get("kind").map_sql_err()?;
let status: String = row.try_get("status").map_sql_err()?;
let attempt: i64 = row.try_get("attempt").map_sql_err()?;
let max_attempts: i64 = row.try_get("max_attempts").map_sql_err()?;
let progress_percent: i32 = row.try_get("progress_percent").map_sql_err()?;
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
let started_at_unix_secs: Option<i64> = row.try_get("started_at_unix_secs").map_sql_err()?;
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 {
id: row.try_get("id").map_sql_err()?,
task_key: row.try_get("task_key").map_sql_err()?,
kind: BackgroundTaskKind::from_database(&kind)?,
trigger: row.try_get("trigger").map_sql_err()?,
status: BackgroundTaskStatus::from_database(&status)?,
attempt: u32::try_from(attempt).unwrap_or_default(),
max_attempts: u32::try_from(max_attempts).unwrap_or_default(),
owner_instance: row.try_get("owner_instance").map_sql_err()?,
progress_percent: u16::try_from(progress_percent).unwrap_or_default(),
progress_message: row.try_get("progress_message").map_sql_err()?,
payload_json: parse_optional_json(row.try_get("payload_json").map_sql_err()?)?,
result_json: parse_optional_json(row.try_get("result_json").map_sql_err()?)?,
error_message: row.try_get("error_message").map_sql_err()?,
cancel_requested: row.try_get("cancel_requested").map_sql_err()?,
created_by: row.try_get("created_by").map_sql_err()?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
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(),
})
}
fn map_event_row(row: &SqliteRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
Ok(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()?,
message: row.try_get("message").map_sql_err()?,
payload_json: parse_optional_json(row.try_get("payload_json").map_sql_err()?)?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
})
}
fn parse_optional_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.map(|raw| {
serde_json::from_str::<serde_json::Value>(&raw).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid background task json payload: {err}"
))
})
})
.transpose()
}
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
fn u64_to_i64(value: u64, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
use crate::run_migrations;
#[tokio::test]
async fn sqlite_background_task_repository_round_trips_runs_and_events() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteBackgroundTaskRepository::new(pool);
let run = repository
.upsert_run(UpsertBackgroundTaskRun {
id: "run-1".to_string(),
task_key: "usage.cleanup".to_string(),
kind: BackgroundTaskKind::Scheduled,
trigger: "timer".to_string(),
status: BackgroundTaskStatus::Queued,
attempt: 0,
max_attempts: 3,
owner_instance: None,
progress_percent: 0,
progress_message: None,
payload_json: Some(json!({"partition": 7})),
result_json: None,
error_message: None,
cancel_requested: false,
created_by: Some("scheduler".to_string()),
created_at_unix_secs: 10,
started_at_unix_secs: None,
finished_at_unix_secs: None,
updated_at_unix_secs: 10,
})
.await
.expect("background task run should upsert");
assert_eq!(run.payload_json, Some(json!({"partition": 7})));
repository
.upsert_event(UpsertBackgroundTaskEvent {
id: "event-1".to_string(),
run_id: run.id.clone(),
event_type: "queued".to_string(),
message: "task queued".to_string(),
payload_json: Some(json!({"attempt": 0})),
created_at_unix_secs: 11,
})
.await
.expect("background task event should upsert");
let page = repository
.list_runs(&BackgroundTaskListQuery {
task_key_substring: Some("cleanup".to_string()),
kind: Some(BackgroundTaskKind::Scheduled),
status: Some(BackgroundTaskStatus::Queued),
trigger: Some("timer".to_string()),
offset: 0,
limit: 10,
})
.await
.expect("background task runs should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].id, "run-1");
let events = repository
.list_events("run-1", 0, 10)
.await
.expect("background task events should list");
assert_eq!(events.len(), 1);
assert_eq!(events[0].payload_json, Some(json!({"attempt": 0})));
assert!(repository
.request_cancel("run-1", 20)
.await
.expect("background task cancellation should update"));
let cancelled = repository
.find_run("run-1")
.await
.expect("background task run should load")
.expect("background task run should exist");
assert!(cancelled.cancel_requested);
assert_eq!(cancelled.updated_at_unix_secs, 20);
let summary = repository
.summarize_runs()
.await
.expect("background task summary should load");
assert_eq!(summary.total, 1);
assert_eq!(summary.by_status.get("queued"), Some(&1));
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,840 @@
use std::collections::{BTreeMap, BTreeSet};
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::candidates::{
request_candidate_lifecycle_would_regress, PublicHealthStatusCount, PublicHealthTimelineBucket,
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
StoredRequestCandidate, UpsertRequestCandidateRecord,
};
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_in, WhereClause};
use crate::error::SqlResultExt;
use crate::SqlitePool;
const CANDIDATE_COLUMNS: &str = r#"
SELECT
id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
candidate_index,
retry_index,
provider_id,
endpoint_id,
key_id,
status,
skip_reason,
is_cached,
status_code,
error_type,
error_message,
latency_ms,
concurrent_requests,
extra_data,
required_capabilities,
created_at AS created_at_unix_ms,
started_at AS started_at_unix_ms,
finished_at AS finished_at_unix_ms
FROM request_candidates
"#;
#[derive(Debug, Clone)]
pub struct SqliteRequestCandidateRepository {
pool: SqlitePool,
}
impl SqliteRequestCandidateRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn find_by_unique(
&self,
request_id: &str,
candidate_index: u32,
retry_index: u32,
) -> Result<Option<StoredRequestCandidate>, DataLayerError> {
let row = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? AND candidate_index = ? AND retry_index = ? LIMIT 1"
))
.bind(request_id)
.bind(to_i32(candidate_index)?)
.bind(to_i32(retry_index)?)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_candidate_row).transpose()
}
}
#[async_trait]
impl RequestCandidateReadRepository for SqliteRequestCandidateRepository {
async fn list_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? ORDER BY candidate_index ASC, retry_index ASC, created_at ASC"
))
.bind(request_id)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_attempted_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? \
AND (status IN ('streaming', 'success', 'failed', 'cancelled') \
OR (status = 'pending' AND started_at IS NOT NULL)) \
ORDER BY candidate_index ASC, retry_index ASC, created_at ASC"
))
.bind(request_id)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_recent(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} ORDER BY created_at DESC LIMIT ?"
))
.bind(limit_i64(limit, "recent request candidate limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_by_provider_id(
&self,
provider_id: &str,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE provider_id = ? ORDER BY created_at DESC LIMIT ?"
))
.bind(provider_id)
.bind(limit_i64(limit, "provider request candidate limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if endpoint_ids.is_empty() || limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_COLUMNS);
let mut where_clause = WhereClause::new();
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
.push(" AND status IN ('success', 'failed', 'skipped')")
.push(" ORDER BY created_at DESC LIMIT ")
.push_bind(limit_i64(limit, "finalized request candidate limit")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn count_finalized_statuses_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
) -> Result<Vec<PublicHealthStatusCount>, DataLayerError> {
if endpoint_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(
"SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates",
);
let mut where_clause = WhereClause::new();
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
.push(" AND status IN ('success', 'failed', 'skipped')")
.push(" GROUP BY endpoint_id, status");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(PublicHealthStatusCount {
endpoint_id: row.try_get("endpoint_id").map_sql_err()?,
status: RequestCandidateStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
count: u64::try_from(row.try_get::<i64, _>("count").map_sql_err()?).map_err(
|_| {
DataLayerError::UnexpectedValue(
"public health status count out of range".to_string(),
)
},
)?,
})
})
.collect()
}
async fn aggregate_finalized_timeline_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
if endpoint_ids.is_empty() || segments == 0 || until_unix_secs < since_unix_secs {
return Ok(Vec::new());
}
let since_ms = unix_secs_to_ms_i64(since_unix_secs)?;
let until_ms = unix_secs_to_ms_i64(until_unix_secs)?;
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_COLUMNS);
let mut where_clause = WhereClause::new();
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(since_ms)
.push(" AND created_at <= ")
.push_bind(until_ms)
.push(" AND status IN ('success', 'failed', 'skipped')");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
aggregate_timeline(
rows.iter()
.map(map_candidate_row)
.collect::<Result<Vec<_>, _>>()?,
since_unix_secs,
until_unix_secs,
segments,
)
}
}
#[async_trait]
impl RequestCandidateWriteRepository for SqliteRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.validate()?;
let existing = self
.find_by_unique(
&candidate.request_id,
candidate.candidate_index,
candidate.retry_index,
)
.await?;
let merged = merge_candidate(candidate, existing)?;
upsert_merged_candidate(&self.pool, &merged).await?;
Ok(merged)
}
async fn delete_created_before(
&self,
created_before_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
if limit == 0 {
return Ok(0);
}
let rows_affected = sqlx::query(
r#"
DELETE FROM request_candidates
WHERE id IN (
SELECT id
FROM request_candidates
WHERE created_at < ?
ORDER BY created_at ASC, id ASC
LIMIT ?
)
"#,
)
.bind(unix_secs_to_ms_i64(created_before_unix_secs)?)
.bind(limit_i64(limit, "request candidate delete limit")?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
async fn upsert_merged_candidate(
pool: &SqlitePool,
candidate: &StoredRequestCandidate,
) -> Result<(), DataLayerError> {
sqlx::query(
r#"
INSERT INTO request_candidates (
id, request_id, user_id, api_key_id, username, api_key_name,
candidate_index, retry_index, provider_id, endpoint_id, key_id, status,
skip_reason, is_cached, status_code, error_type, error_message, latency_ms,
concurrent_requests, extra_data, required_capabilities, created_at, started_at, finished_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(request_id, candidate_index, retry_index) DO UPDATE SET
user_id = excluded.user_id,
api_key_id = excluded.api_key_id,
username = excluded.username,
api_key_name = excluded.api_key_name,
provider_id = excluded.provider_id,
endpoint_id = excluded.endpoint_id,
key_id = excluded.key_id,
status = excluded.status,
skip_reason = excluded.skip_reason,
is_cached = excluded.is_cached,
status_code = excluded.status_code,
error_type = excluded.error_type,
error_message = excluded.error_message,
latency_ms = excluded.latency_ms,
concurrent_requests = excluded.concurrent_requests,
extra_data = excluded.extra_data,
required_capabilities = excluded.required_capabilities,
created_at = excluded.created_at,
started_at = excluded.started_at,
finished_at = excluded.finished_at
"#,
)
.bind(&candidate.id)
.bind(&candidate.request_id)
.bind(&candidate.user_id)
.bind(&candidate.api_key_id)
.bind(&candidate.username)
.bind(&candidate.api_key_name)
.bind(to_i32(candidate.candidate_index)?)
.bind(to_i32(candidate.retry_index)?)
.bind(&candidate.provider_id)
.bind(&candidate.endpoint_id)
.bind(&candidate.key_id)
.bind(status_to_database(candidate.status))
.bind(&candidate.skip_reason)
.bind(candidate.is_cached)
.bind(candidate.status_code.map(i32::from))
.bind(&candidate.error_type)
.bind(&candidate.error_message)
.bind(candidate.latency_ms.map(to_i32_u64).transpose()?)
.bind(candidate.concurrent_requests.map(to_i32).transpose()?)
.bind(json_to_string(&candidate.extra_data)?)
.bind(json_to_string(&candidate.required_capabilities)?)
.bind(u64_to_i64(
candidate.created_at_unix_ms,
"request candidate created_at",
)?)
.bind(optional_u64_to_i64(
candidate.started_at_unix_ms,
"request candidate started_at",
)?)
.bind(optional_u64_to_i64(
candidate.finished_at_unix_ms,
"request candidate finished_at",
)?)
.execute(pool)
.await
.map_sql_err()?;
Ok(())
}
fn merge_candidate(
candidate: UpsertRequestCandidateRecord,
existing: Option<StoredRequestCandidate>,
) -> Result<StoredRequestCandidate, DataLayerError> {
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|value| {
request_candidate_lifecycle_would_regress(value.status, candidate.status)
});
let merged_status = if preserve_existing_lifecycle {
existing
.as_ref()
.map(|value| value.status)
.unwrap_or(candidate.status)
} else {
candidate.status
};
let created_at_unix_ms = candidate
.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(candidate.started_at_unix_ms)
.or(candidate.finished_at_unix_ms)
.unwrap_or_else(current_unix_ms);
let id = existing
.as_ref()
.map(|value| value.id.clone())
.unwrap_or(candidate.id);
let extra_data = merge_json_objects(
existing.as_ref().and_then(|value| value.extra_data.clone()),
candidate.extra_data,
);
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())
}),
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())),
merged_status,
candidate.skip_reason.or_else(|| {
existing
.as_ref()
.and_then(|value| value.skip_reason.clone())
}),
candidate
.is_cached
.unwrap_or_else(|| existing.as_ref().is_some_and(|value| value.is_cached)),
if preserve_existing_lifecycle {
existing
.as_ref()
.and_then(|value| value.status_code.map(i32::from))
} else {
candidate.status_code.map(i32::from).or_else(|| {
existing
.as_ref()
.and_then(|value| value.status_code.map(i32::from))
})
},
if preserve_existing_lifecycle {
existing.as_ref().and_then(|value| value.error_type.clone())
} else {
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())
})
},
if preserve_existing_lifecycle {
match existing.as_ref().and_then(|value| value.latency_ms) {
Some(value) => Some(to_i32_u64(value)?),
None => None,
}
} else {
candidate.latency_ms.map(to_i32_u64).transpose()?.or(
match existing.as_ref().and_then(|value| value.latency_ms) {
Some(value) => Some(to_i32_u64(value)?),
None => None,
},
)
},
candidate.concurrent_requests.map(to_i32).transpose()?.or(
match existing
.as_ref()
.and_then(|value| value.concurrent_requests)
{
Some(value) => Some(to_i32(value)?),
None => None,
},
),
extra_data,
candidate.required_capabilities.or_else(|| {
existing
.as_ref()
.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))
.map(|value| u64_to_i64(value, "request candidate started_at"))
.transpose()?,
if preserve_existing_lifecycle {
existing
.as_ref()
.and_then(|value| value.finished_at_unix_ms)
} else {
candidate.finished_at_unix_ms.or_else(|| {
existing
.as_ref()
.and_then(|value| value.finished_at_unix_ms)
})
}
.map(|value| u64_to_i64(value, "request candidate finished_at"))
.transpose()?,
)
}
fn aggregate_timeline(
candidates: Vec<StoredRequestCandidate>,
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
let endpoint_ids = candidates
.iter()
.filter_map(|candidate| candidate.endpoint_id.clone())
.collect::<BTreeSet<_>>();
let span_ms = until_unix_secs
.saturating_sub(since_unix_secs)
.saturating_mul(1000)
.max(1);
let since_ms = since_unix_secs.saturating_mul(1000);
let mut buckets = BTreeMap::<(String, u32), PublicHealthTimelineBucket>::new();
for candidate in candidates {
let Some(endpoint_id) = candidate.endpoint_id.clone() else {
continue;
};
let offset = candidate.created_at_unix_ms.saturating_sub(since_ms);
let segment_idx = ((offset.saturating_mul(u64::from(segments))) / span_ms)
.min(u64::from(segments.saturating_sub(1))) as u32;
let bucket = buckets.entry((endpoint_id.clone(), segment_idx)).or_insert(
PublicHealthTimelineBucket {
endpoint_id,
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: Some(candidate.created_at_unix_ms),
max_created_at_unix_ms: Some(candidate.created_at_unix_ms),
},
);
bucket.total_count += 1;
if candidate.status == RequestCandidateStatus::Success {
bucket.success_count += 1;
}
if candidate.status == RequestCandidateStatus::Failed {
bucket.failed_count += 1;
}
bucket.min_created_at_unix_ms = bucket
.min_created_at_unix_ms
.map(|value| value.min(candidate.created_at_unix_ms));
bucket.max_created_at_unix_ms = bucket
.max_created_at_unix_ms
.map(|value| value.max(candidate.created_at_unix_ms));
}
for endpoint_id in endpoint_ids {
for segment_idx in 0..segments {
buckets.entry((endpoint_id.clone(), segment_idx)).or_insert(
PublicHealthTimelineBucket {
endpoint_id: endpoint_id.clone(),
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: None,
max_created_at_unix_ms: None,
},
);
}
}
Ok(buckets.into_values().collect())
}
fn map_candidate_row(row: &SqliteRow) -> Result<StoredRequestCandidate, DataLayerError> {
StoredRequestCandidate::new(
row.try_get("id").map_sql_err()?,
row.try_get("request_id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("candidate_index").map_sql_err()?,
row.try_get("retry_index").map_sql_err()?,
row.try_get("provider_id").map_sql_err()?,
row.try_get("endpoint_id").map_sql_err()?,
row.try_get("key_id").map_sql_err()?,
RequestCandidateStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
row.try_get("skip_reason").map_sql_err()?,
row.try_get("is_cached").map_sql_err()?,
row.try_get("status_code").map_sql_err()?,
row.try_get("error_type").map_sql_err()?,
row.try_get("error_message").map_sql_err()?,
row.try_get("latency_ms").map_sql_err()?,
row.try_get("concurrent_requests").map_sql_err()?,
parse_json(row.try_get("extra_data").ok().flatten())?,
parse_json(row.try_get("required_capabilities").ok().flatten())?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("started_at_unix_ms").map_sql_err()?,
row.try_get("finished_at_unix_ms").map_sql_err()?,
)
}
fn parse_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.filter(|value| !value.trim().is_empty())
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"request_candidates JSON field is invalid: {err}"
))
})
})
.transpose()
}
fn json_to_string(value: &Option<serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"request_candidates JSON field is unserializable: {err}"
))
})
})
.transpose()
}
fn merge_json_objects(
existing: Option<serde_json::Value>,
overlay: Option<serde_json::Value>,
) -> Option<serde_json::Value> {
match (existing, overlay) {
(
Some(serde_json::Value::Object(mut existing_object)),
Some(serde_json::Value::Object(overlay_object)),
) => {
existing_object.extend(overlay_object);
Some(serde_json::Value::Object(existing_object))
}
(_existing, Some(overlay)) => Some(overlay),
(existing, None) => existing,
}
}
fn status_to_database(status: RequestCandidateStatus) -> &'static str {
match status {
RequestCandidateStatus::Available => "available",
RequestCandidateStatus::Unused => "unused",
RequestCandidateStatus::Pending => "pending",
RequestCandidateStatus::Streaming => "streaming",
RequestCandidateStatus::Success => "success",
RequestCandidateStatus::Failed => "failed",
RequestCandidateStatus::Cancelled => "cancelled",
RequestCandidateStatus::Skipped => "skipped",
}
}
fn current_unix_ms() -> u64 {
chrono::Utc::now().timestamp_millis().max(0) as u64
}
fn unix_secs_to_ms_i64(value: u64) -> Result<i64, DataLayerError> {
let value = value.checked_mul(1000).ok_or_else(|| {
DataLayerError::UnexpectedValue("request candidate timestamp overflow".to_string())
})?;
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue("request candidate timestamp overflow".to_string())
})
}
fn limit_i64(value: usize, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {name}: {value}")))
}
fn to_i32(value: u32) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("request candidate value out of range: {value}"))
})
}
fn to_i32_u64(value: u64) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("request candidate value out of range: {value}"))
})
}
fn u64_to_i64(value: u64, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, DataLayerError> {
value.map(|value| u64_to_i64(value, name)).transpose()
}
#[cfg(test)]
mod tests {
use super::SqliteRequestCandidateRepository;
use crate::run_migrations;
use aether_data_contracts::repository::candidates::{
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
UpsertRequestCandidateRecord,
};
use serde_json::json;
#[tokio::test]
async fn sqlite_repository_writes_and_reads_request_candidates() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteRequestCandidateRepository::new(pool);
let created = repository
.upsert(sample_upsert(
"candidate-1",
RequestCandidateStatus::Pending,
Some(json!({"a": 1})),
1_000_000,
))
.await
.expect("candidate should insert");
assert_eq!(created.request_id, "request-1");
let updated = repository
.upsert(sample_upsert(
"candidate-replacement",
RequestCandidateStatus::Success,
Some(json!({"b": 2})),
1_000_500,
))
.await
.expect("candidate should update");
assert_eq!(updated.id, "candidate-1");
assert_eq!(updated.extra_data, Some(json!({"a": 1, "b": 2})));
let late_streaming = repository
.upsert(sample_upsert(
"candidate-late-streaming",
RequestCandidateStatus::Streaming,
Some(json!({"late": true})),
1_000_250,
))
.await
.expect("late streaming candidate should not regress terminal status");
assert_eq!(late_streaming.id, "candidate-1");
assert_eq!(late_streaming.status, RequestCandidateStatus::Success);
assert_eq!(late_streaming.finished_at_unix_ms, Some(1_000_502));
assert_eq!(
late_streaming.extra_data,
Some(json!({"a": 1, "b": 2, "late": true}))
);
assert_eq!(
repository
.list_by_request_id("request-1")
.await
.expect("request list should load")
.len(),
1
);
assert_eq!(
repository
.count_finalized_statuses_by_endpoint_ids_since(&["endpoint-1".to_string()], 900)
.await
.expect("status counts should load")[0]
.count,
1
);
assert_eq!(
repository
.aggregate_finalized_timeline_by_endpoint_ids_since(
&["endpoint-1".to_string()],
900,
1200,
3,
)
.await
.expect("timeline should load")
.len(),
3
);
assert_eq!(
repository
.delete_created_before(2_000, 10)
.await
.expect("old candidates should delete"),
1
);
}
fn sample_upsert(
id: &str,
status: RequestCandidateStatus,
extra_data: Option<serde_json::Value>,
created_at_unix_ms: u64,
) -> UpsertRequestCandidateRecord {
UpsertRequestCandidateRecord {
id: id.to_string(),
request_id: "request-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
username: Some("user".to_string()),
api_key_name: Some("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,
skip_reason: None,
is_cached: Some(false),
status_code: Some(200),
error_type: None,
error_message: None,
latency_ms: Some(123),
concurrent_requests: Some(2),
extra_data,
required_capabilities: Some(json!({"streaming": true})),
created_at_unix_ms: Some(created_at_unix_ms),
started_at_unix_ms: Some(created_at_unix_ms + 1),
finished_at_unix_ms: Some(created_at_unix_ms + 2),
}
}
}
@@ -0,0 +1,11 @@
use aether_data_contracts::DataLayerError;
pub(crate) trait SqlResultExt<T> {
fn map_sql_err(self) -> Result<T, DataLayerError>;
}
impl<T> SqlResultExt<T> for Result<T, sqlx::Error> {
fn map_sql_err(self) -> Result<T, DataLayerError> {
self.map_err(DataLayerError::sql)
}
}
@@ -0,0 +1,431 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
GeminiFileMappingStats, GeminiFileMappingWriteRepository, StoredGeminiFileMapping,
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
};
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, WhereClause};
use crate::error::SqlResultExt;
use crate::SqlitePool;
#[derive(Debug, Clone)]
pub struct SqliteGeminiFileMappingRepository {
pool: SqlitePool,
}
impl SqliteGeminiFileMappingRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn reload_by_file_name(
&self,
file_name: &str,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
self.find_by_file_name(file_name).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("gemini file mapping missing after write".to_string())
})
}
}
#[async_trait]
impl GeminiFileMappingReadRepository for SqliteGeminiFileMappingRepository {
async fn find_by_file_name(
&self,
file_name: &str,
) -> 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 file_name = ?
LIMIT 1
"#,
)
.bind(file_name)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
) -> Result<StoredGeminiFileMappingListPage, DataLayerError> {
let total = build_list_count_query(query)
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let rows = build_list_rows_query(query)
.build()
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows.iter().map(map_row).collect::<Result<Vec<_>, _>>()?;
Ok(StoredGeminiFileMappingListPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn summarize_mappings(
&self,
now_unix_secs: u64,
) -> Result<GeminiFileMappingStats, DataLayerError> {
let totals = sqlx::query(
r#"
SELECT
COUNT(*) AS total_mappings,
SUM(CASE WHEN expires_at > ? THEN 1 ELSE 0 END) AS active_mappings
FROM gemini_file_mappings
"#,
)
.bind(now_unix_secs as i64)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total_mappings =
usize::try_from(totals.try_get::<i64, _>("total_mappings").map_sql_err()?)
.unwrap_or_default();
let active_mappings = usize::try_from(
totals
.try_get::<Option<i64>, _>("active_mappings")
.map_sql_err()?
.unwrap_or(0),
)
.unwrap_or_default();
let by_mime_type_rows = sqlx::query(
r#"
SELECT
COALESCE(NULLIF(TRIM(mime_type), ''), 'unknown') AS mime_type,
COUNT(*) AS count
FROM gemini_file_mappings
WHERE expires_at > ?
GROUP BY COALESCE(NULLIF(TRIM(mime_type), ''), 'unknown')
ORDER BY mime_type ASC
"#,
)
.bind(now_unix_secs as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let by_mime_type = by_mime_type_rows
.iter()
.map(|row| {
Ok(GeminiFileMappingMimeTypeCount {
mime_type: row.try_get("mime_type").map_sql_err()?,
count: usize::try_from(row.try_get::<i64, _>("count").map_sql_err()?)
.unwrap_or_default(),
})
})
.collect::<Result<Vec<_>, DataLayerError>>()?;
Ok(GeminiFileMappingStats {
total_mappings,
active_mappings,
expired_mappings: total_mappings.saturating_sub(active_mappings),
by_mime_type,
})
}
}
#[async_trait]
impl GeminiFileMappingWriteRepository for SqliteGeminiFileMappingRepository {
async fn upsert(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
record.validate()?;
sqlx::query(
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 CONFLICT(file_name) DO UPDATE SET
key_id = excluded.key_id,
user_id = excluded.user_id,
display_name = excluded.display_name,
mime_type = excluded.mime_type,
source_hash = excluded.source_hash,
expires_at = excluded.expires_at
"#,
)
.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(&self.pool)
.await
.map_sql_err()?;
self.reload_by_file_name(&record.file_name).await
}
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)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn delete_by_id(
&self,
mapping_id: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let existing = 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 id = ?
LIMIT 1
"#,
)
.bind(mapping_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
let Some(existing) = existing else {
return Ok(None);
};
sqlx::query("DELETE FROM gemini_file_mappings WHERE id = ?")
.bind(mapping_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(Some(map_row(&existing)?))
}
async fn delete_expired_before(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE expires_at <= ?")
.bind(now_unix_secs as i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, Sqlite> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings");
let mut where_clause = WhereClause::new();
apply_list_filters(&mut builder, &mut where_clause, query);
builder
}
fn build_list_rows_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, Sqlite> {
let mut builder = QueryBuilder::<Sqlite>::new(
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
"#,
);
let mut where_clause = WhereClause::new();
apply_list_filters(&mut builder, &mut where_clause, query);
builder.push(" ORDER BY created_at DESC, file_name ASC");
push_limit_offset(
&mut builder,
i64::try_from(query.limit).unwrap_or(i64::MAX),
i64::try_from(query.offset).unwrap_or(i64::MAX),
);
builder
}
fn apply_list_filters(
builder: &mut QueryBuilder<'_, Sqlite>,
where_clause: &mut WhereClause,
query: &GeminiFileMappingListQuery,
) {
if !query.include_expired {
where_clause.push_next(builder);
builder.push("expires_at > ");
builder.push_bind(query.now_unix_secs as i64);
}
if let Some(search) = query
.search
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
push_ci_contains_any(
builder,
where_clause,
SqlDialect::Sqlite,
&["file_name", "COALESCE(display_name, '')"],
search,
);
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn map_row(row: &SqliteRow) -> Result<StoredGeminiFileMapping, DataLayerError> {
Ok(StoredGeminiFileMapping {
id: row.try_get("id").map_sql_err()?,
file_name: row.try_get("file_name").map_sql_err()?,
key_id: row.try_get("key_id").map_sql_err()?,
user_id: row.try_get("user_id").ok().flatten(),
display_name: row.try_get("display_name").ok().flatten(),
mime_type: row.try_get("mime_type").ok().flatten(),
source_hash: row.try_get("source_hash").ok().flatten(),
created_at_unix_ms: u64::try_from(
row.try_get::<i64, _>("created_at_unix_ms").map_sql_err()?,
)
.map_err(|_| {
DataLayerError::UnexpectedValue(
"gemini_file_mappings.created_at is invalid".to_string(),
)
})?,
expires_at_unix_secs: u64::try_from(
row.try_get::<i64, _>("expires_at_unix_secs")
.map_sql_err()?,
)
.map_err(|_| {
DataLayerError::UnexpectedValue(
"gemini_file_mappings.expires_at is invalid".to_string(),
)
})?,
})
}
#[cfg(test)]
mod tests {
use super::SqliteGeminiFileMappingRepository;
use crate::run_migrations;
use aether_data_contracts::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingReadRepository,
GeminiFileMappingWriteRepository, UpsertGeminiFileMappingRecord,
};
#[tokio::test]
async fn sqlite_repository_round_trips_gemini_file_mappings() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteGeminiFileMappingRepository::new(pool);
let created = repository
.upsert(UpsertGeminiFileMappingRecord {
id: "mapping-1".to_string(),
file_name: "files/example.png".to_string(),
key_id: "key-1".to_string(),
user_id: Some("user-1".to_string()),
display_name: Some("Example".to_string()),
mime_type: Some("image/png".to_string()),
source_hash: Some("hash-1".to_string()),
expires_at_unix_secs: 300,
})
.await
.expect("mapping should upsert");
assert_eq!(created.id, "mapping-1");
assert_eq!(created.mime_type, Some("image/png".to_string()));
let updated = repository
.upsert(UpsertGeminiFileMappingRecord {
id: "mapping-replacement".to_string(),
file_name: "files/example.png".to_string(),
key_id: "key-2".to_string(),
user_id: Some("user-2".to_string()),
display_name: Some("Updated".to_string()),
mime_type: Some("image/jpeg".to_string()),
source_hash: Some("hash-2".to_string()),
expires_at_unix_secs: 500,
})
.await
.expect("mapping should update");
assert_eq!(updated.id, "mapping-1");
assert_eq!(updated.key_id, "key-2");
let page = repository
.list_mappings(&GeminiFileMappingListQuery {
include_expired: false,
search: Some("updated".to_string()),
offset: 0,
limit: 10,
now_unix_secs: 400,
})
.await
.expect("mappings should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].file_name, "files/example.png");
let stats = repository
.summarize_mappings(400)
.await
.expect("stats should load");
assert_eq!(stats.total_mappings, 1);
assert_eq!(stats.active_mappings, 1);
assert_eq!(stats.by_mime_type[0].mime_type, "image/jpeg");
assert_eq!(
repository
.delete_expired_before(600)
.await
.expect("expired mappings should delete"),
1
);
assert!(repository
.find_by_file_name("files/example.png")
.await
.expect("find should run")
.is_none());
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,75 @@
//! SQLite repositories, pool primitives, and migrations.
mod announcements;
mod audit;
mod auth;
mod auth_modules;
mod background_tasks;
mod billing;
mod candidate_selection;
mod candidates;
mod error;
mod gemini_file_mappings;
mod global_models;
mod management_tokens;
mod migrations;
mod oauth_providers;
mod pool;
mod pool_scores;
mod provider_catalog;
mod proxy_nodes;
mod quota;
mod routing_profiles;
mod settlement;
mod usage;
mod users;
mod video_tasks;
mod wallet;
pub use aether_data_contracts::{DataLayerError, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
pub use announcements::SqliteAnnouncementRepository;
pub use audit::SqliteAuditLogReadRepository;
pub use auth::SqliteAuthApiKeyReadRepository;
pub use auth_modules::{SqliteAuthModuleReadRepository, SqliteAuthModuleRepository};
pub use background_tasks::SqliteBackgroundTaskRepository;
pub use billing::SqliteBillingReadRepository;
pub use candidate_selection::SqliteMinimalCandidateSelectionReadRepository;
pub use candidates::SqliteRequestCandidateRepository;
pub use gemini_file_mappings::SqliteGeminiFileMappingRepository;
pub use global_models::SqliteGlobalModelReadRepository;
pub use management_tokens::SqliteManagementTokenRepository;
pub use migrations::{pending_migrations, prepare_database_for_startup, run_migrations, MIGRATOR};
pub use oauth_providers::SqliteOAuthProviderRepository;
pub use pool::{SqlitePool, SqlitePoolConfig, SqlitePoolFactory};
pub use pool_scores::SqlitePoolMemberScoreRepository;
pub use provider_catalog::SqliteProviderCatalogReadRepository;
pub use proxy_nodes::SqliteProxyNodeReadRepository;
pub use quota::SqliteProviderQuotaRepository;
pub use routing_profiles::SqliteRoutingGroupRepository;
pub use settlement::SqliteSettlementRepository;
pub use usage::{SqliteUsageReadRepository, SqliteUsageWriteRepository};
pub use users::SqliteUserReadRepository;
pub use video_tasks::SqliteVideoTaskRepository;
pub use wallet::SqliteWalletReadRepository;
use sqlx::{sqlite::SqliteRow, Row};
pub fn sqlite_real(row: &SqliteRow, field: &str) -> Result<f64, DataLayerError> {
match row.try_get::<f64, _>(field) {
Ok(value) => Ok(value),
Err(real_err) => match row.try_get::<i64, _>(field) {
Ok(value) => Ok(value as f64),
Err(_) => Err(DataLayerError::sql(real_err)),
},
}
}
pub fn sqlite_optional_real(row: &SqliteRow, field: &str) -> Result<Option<f64>, DataLayerError> {
match row.try_get::<Option<f64>, _>(field) {
Ok(value) => Ok(value),
Err(real_err) => match row.try_get::<Option<i64>, _>(field) {
Ok(value) => Ok(value.map(|value| value as f64)),
Err(_) => Err(DataLayerError::sql(real_err)),
},
}
}
@@ -0,0 +1,593 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::management_tokens::{
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};
use crate::error::SqlResultExt;
use crate::SqlitePool;
#[derive(Debug, Clone)]
pub struct SqliteManagementTokenRepository {
pool: SqlitePool,
}
impl SqliteManagementTokenRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn get_token(
&self,
token_id: &str,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(TOKEN_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(&mut builder, &mut where_clause, "id", token_id.to_string());
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()
}
}
const TOKEN_COLUMNS: &str = r#"
SELECT
id,
user_id,
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
"#;
const TOKEN_WITH_USER_COLUMNS: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.permissions,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
"#;
#[async_trait]
impl ManagementTokenReadRepository for SqliteManagementTokenRepository {
async fn list_management_tokens(
&self,
query: &ManagementTokenListQuery,
) -> Result<StoredManagementTokenListPage, DataLayerError> {
let mut count_builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(mt.id) AS total FROM management_tokens mt");
let mut count_where = WhereClause::new();
apply_management_token_filters(&mut count_builder, &mut count_where, query);
let total = count_builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let mut list_builder = QueryBuilder::<Sqlite>::new(TOKEN_WITH_USER_COLUMNS);
let mut list_where = WhereClause::new();
apply_management_token_filters(&mut list_builder, &mut list_where, query);
list_builder.push(" ORDER BY mt.created_at DESC, mt.id DESC");
push_limit_offset(
&mut list_builder,
i64::try_from(query.limit).unwrap_or(i64::MAX),
i64::try_from(query.offset).unwrap_or(i64::MAX),
);
let rows = list_builder
.build()
.fetch_all(&self.pool)
.await
.map_sql_err()?;
Ok(StoredManagementTokenListPage {
items: rows
.iter()
.map(map_token_with_user_row)
.collect::<Result<Vec<_>, _>>()?,
total: usize::try_from(total.max(0)).unwrap_or(usize::MAX),
})
}
async fn get_management_token_with_user(
&self,
token_id: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(TOKEN_WITH_USER_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"mt.id",
token_id.to_string(),
);
push_limit(&mut builder, 1);
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_with_user_row).transpose()
}
async fn get_management_token_with_user_by_hash(
&self,
token_hash: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(TOKEN_WITH_USER_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"mt.token_hash",
token_hash.to_string(),
);
push_limit(&mut builder, 1);
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_with_user_row).transpose()
}
}
fn apply_management_token_filters(
builder: &mut QueryBuilder<'_, Sqlite>,
where_clause: &mut WhereClause,
query: &ManagementTokenListQuery,
) {
push_optional_eq(builder, where_clause, "mt.user_id", query.user_id.clone());
push_optional_eq(builder, where_clause, "mt.is_active", query.is_active);
}
#[async_trait]
impl ManagementTokenWriteRepository for SqliteManagementTokenRepository {
async fn create_management_token(
&self,
record: &CreateManagementTokenRecord,
) -> Result<StoredManagementToken, DataLayerError> {
record.validate()?;
let now = now_unix_secs();
sqlx::query(
r#"
INSERT INTO management_tokens (
id, user_id, token_hash, token_prefix, name, description, allowed_ips,
permissions, expires_at, is_active, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.id)
.bind(&record.user_id)
.bind(&record.token_hash)
.bind(record.token_prefix.as_deref())
.bind(&record.name)
.bind(record.description.as_deref())
.bind(json_to_string(record.allowed_ips.as_ref())?)
.bind(json_to_string(record.permissions.as_ref())?)
.bind(
record
.expires_at_unix_secs
.and_then(|value| i64::try_from(value).ok()),
)
.bind(record.is_active)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_err(|err| map_sqlite_write_error(err, Some(record.name.as_str())))?;
self.get_token(&record.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("created management token missing".to_string())
})
}
async fn update_management_token(
&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(&current.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();
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_sqlite_write_error(err, record.name.as_deref()))?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&record.token_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)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn set_management_token_active(
&self,
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)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(token_id).await
}
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
}
async fn record_management_token_usage(
&self,
token_id: &str,
last_used_ip: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let now = now_unix_secs();
let result = sqlx::query(
r#"
UPDATE management_tokens
SET last_used_at = ?,
last_used_ip = ?,
usage_count = COALESCE(usage_count, 0) + 1,
updated_at = ?
WHERE id = ?
"#,
)
.bind(now as i64)
.bind(last_used_ip)
.bind(now as i64)
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(token_id).await
}
}
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 json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid management token JSON field: {err}"
))
})
})
.transpose()
}
fn json_from_string(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid management token JSON field: {err}"
))
})
})
.transpose()
}
fn map_sqlite_write_error(err: sqlx::Error, requested_name: Option<&str>) -> DataLayerError {
let message = err.to_string();
if message.contains("management_tokens.user_id, management_tokens.name") {
return DataLayerError::InvalidInput(
requested_name
.map(|name| format!("已存在名为 '{}' 的 Token", name))
.unwrap_or_else(|| "Management Token 名称已存在".to_string()),
);
}
DataLayerError::sql(err)
}
fn map_token_row(row: &SqliteRow) -> Result<StoredManagementToken, DataLayerError> {
Ok(StoredManagementToken::new(
row.try_get("id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
)?
.with_display_fields(
row.try_get("description").map_sql_err()?,
row.try_get("token_prefix").map_sql_err()?,
json_from_string(row.try_get("allowed_ips").map_sql_err()?)?,
)
.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()?),
row.try_get("last_used_ip").map_sql_err()?,
u64::try_from(row.try_get::<i64, _>("usage_count").map_sql_err()?).unwrap_or(0),
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()?),
))
}
fn map_user_summary_row(
row: &SqliteRow,
) -> Result<StoredManagementTokenUserSummary, DataLayerError> {
StoredManagementTokenUserSummary::new(
row.try_get("user_row_id").map_sql_err()?,
row.try_get("user_email").map_sql_err()?,
row.try_get("user_username").map_sql_err()?,
row.try_get("user_role").map_sql_err()?,
)
}
fn map_token_with_user_row(
row: &SqliteRow,
) -> Result<StoredManagementTokenWithUser, DataLayerError> {
Ok(StoredManagementTokenWithUser::new(
map_token_row(row)?,
map_user_summary_row(row)?,
))
}
#[cfg(test)]
mod tests {
use super::SqliteManagementTokenRepository;
use crate::run_migrations;
use aether_data_contracts::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
StoredManagementTokenUserSummary, UpdateManagementTokenRecord,
};
#[tokio::test]
async fn sqlite_repository_round_trips_management_tokens() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO users (id, email, username, role, is_active, created_at, updated_at)
VALUES ('user-1', 'user-1@example.com', 'user-1', 'admin', 1, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("seed user should insert");
let repository = SqliteManagementTokenRepository::new(pool);
let user = StoredManagementTokenUserSummary::new(
"user-1".to_string(),
Some("user-1@example.com".to_string()),
"user-1".to_string(),
"admin".to_string(),
)
.expect("user summary should build");
let created = repository
.create_management_token(&CreateManagementTokenRecord {
id: "token-1".to_string(),
user_id: "user-1".to_string(),
user,
token_hash: "hash-1".to_string(),
token_prefix: Some("ae_1234".to_string()),
name: "primary".to_string(),
description: Some("primary token".to_string()),
allowed_ips: Some(serde_json::json!(["127.0.0.1"])),
permissions: Some(serde_json::json!(["admin:usage:read"])),
expires_at_unix_secs: Some(1_800_000_000),
is_active: true,
})
.await
.expect("token should create");
assert_eq!(created.name, "primary");
assert_eq!(
created.permissions,
Some(serde_json::json!(["admin:usage:read"]))
);
let page = repository
.list_management_tokens(&ManagementTokenListQuery {
user_id: Some("user-1".to_string()),
is_active: Some(true),
offset: 0,
limit: 10,
})
.await
.expect("tokens should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].token.id, "token-1");
let by_hash = repository
.get_management_token_with_user_by_hash("hash-1")
.await
.expect("hash lookup should succeed")
.expect("token should exist");
assert_eq!(by_hash.user.username, "user-1");
let updated = repository
.update_management_token(&UpdateManagementTokenRecord {
token_id: "token-1".to_string(),
name: Some("renamed".to_string()),
description: None,
clear_description: true,
allowed_ips: Some(serde_json::json!(["10.0.0.1"])),
clear_allowed_ips: false,
permissions: Some(serde_json::json!(["admin:usage:read", "admin:usage:write"])),
expires_at_unix_secs: None,
clear_expires_at: true,
is_active: Some(false),
})
.await
.expect("update should succeed")
.expect("token should exist");
assert_eq!(updated.name, "renamed");
assert!(!updated.is_active);
assert_eq!(updated.description, None);
assert_eq!(
updated.permissions,
Some(serde_json::json!(["admin:usage:read", "admin:usage:write"]))
);
assert_eq!(updated.expires_at_unix_secs, None);
let toggled = repository
.set_management_token_active("token-1", true)
.await
.expect("toggle should succeed")
.expect("token should exist");
assert!(toggled.is_active);
let regenerated = repository
.regenerate_management_token_secret(&RegenerateManagementTokenSecret {
token_id: "token-1".to_string(),
token_hash: "hash-2".to_string(),
token_prefix: Some("ae_5678".to_string()),
})
.await
.expect("regenerate should succeed")
.expect("token should exist");
assert_eq!(regenerated.token_prefix.as_deref(), Some("ae_5678"));
assert!(repository
.get_management_token_with_user_by_hash("hash-1")
.await
.expect("old hash lookup should succeed")
.is_none());
let used = repository
.record_management_token_usage("token-1", Some("127.0.0.1"))
.await
.expect("usage should record")
.expect("token should exist");
assert_eq!(used.usage_count, 1);
assert_eq!(used.last_used_ip.as_deref(), Some("127.0.0.1"));
assert!(repository
.delete_management_token("token-1")
.await
.expect("delete should succeed"));
}
}
@@ -0,0 +1,81 @@
use sqlx::{
migrate::{Migrate, MigrateError, Migrator},
SqlitePool,
};
use aether_data_contracts::PendingMigrationInfo;
pub static MIGRATOR: Migrator = sqlx::migrate!("./migrations");
pub async fn run_migrations(pool: &SqlitePool) -> Result<(), MigrateError> {
MIGRATOR.run(pool).await
}
pub async fn pending_migrations(
pool: &SqlitePool,
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
let mut conn = pool.acquire().await?;
let applied_migrations = match conn.list_applied_migrations().await {
Ok(applied_migrations) => applied_migrations,
Err(err) if is_missing_sqlx_migrations_table_error(&err) => Vec::new(),
Err(err) => return Err(err),
};
Ok(pending_migrations_from_applied(&applied_migrations))
}
pub async fn prepare_database_for_startup(
pool: &SqlitePool,
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
pending_migrations(pool).await
}
fn is_missing_sqlx_migrations_table_error(err: &MigrateError) -> bool {
let message = err.to_string().to_ascii_lowercase();
message.contains("_sqlx_migrations")
&& (message.contains("no such table")
|| message.contains("doesn't exist")
|| message.contains("does not exist")
|| message.contains("unknown table"))
}
fn pending_migrations_from_applied(
applied_migrations: &[sqlx::migrate::AppliedMigration],
) -> Vec<PendingMigrationInfo> {
let applied_versions = applied_migrations
.iter()
.map(|migration| migration.version)
.collect::<std::collections::HashSet<_>>();
MIGRATOR
.iter()
.filter(|migration| migration.migration_type.is_up_migration())
.filter(|migration| !applied_versions.contains(&migration.version))
.map(|migration| PendingMigrationInfo {
version: migration.version,
description: migration.description.to_string(),
})
.collect()
}
#[cfg(test)]
mod tests {
use super::{pending_migrations, run_migrations, MIGRATOR};
#[tokio::test]
async fn migrates_empty_database_and_clears_pending_set() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("in-memory sqlite pool");
let pending = pending_migrations(&pool).await.expect("pending migrations");
assert_eq!(pending.len(), MIGRATOR.iter().count());
assert!(!pending.is_empty());
run_migrations(&pool).await.expect("run sqlite migrations");
assert!(pending_migrations(&pool)
.await
.expect("pending migrations after run")
.is_empty());
}
}
@@ -0,0 +1,491 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::oauth_providers::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
UpsertOAuthProviderConfigRecord,
};
use aether_data_contracts::DataLayerError;
use aether_data_query::{push_eq, push_limit, WhereClause};
use crate::error::SqlResultExt;
use crate::SqlitePool;
#[derive(Debug, Clone)]
pub struct SqliteOAuthProviderRepository {
pool: SqlitePool,
}
impl SqliteOAuthProviderRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn get_provider(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(OAUTH_PROVIDER_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"provider_type",
provider_type.to_string(),
);
push_limit(&mut builder, 1);
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_oauth_provider_row).transpose()
}
}
const OAUTH_PROVIDER_COLUMNS: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
icon_url,
is_enabled,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM oauth_providers
"#;
const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
SELECT COUNT(DISTINCT users.id) AS locked_count
FROM users
JOIN user_oauth_links
ON users.id = user_oauth_links.user_id
WHERE users.is_active = 1
AND users.is_deleted = 0
AND user_oauth_links.provider_type = ?
AND (
(
users.auth_source = 'oauth'
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links other_links
JOIN oauth_providers other_provider
ON other_links.provider_type = other_provider.provider_type
WHERE other_links.user_id = users.id
AND other_links.provider_type <> ?
AND other_provider.is_enabled = 1
)
) OR (
? = 1
AND users.auth_source = 'local'
AND users.role <> 'admin'
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links other_links
JOIN oauth_providers other_provider
ON other_links.provider_type = other_provider.provider_type
WHERE other_links.user_id = users.id
AND other_links.provider_type <> ?
AND other_provider.is_enabled = 1
)
)
)
"#;
#[async_trait]
impl OAuthProviderReadRepository for SqliteOAuthProviderRepository {
async fn list_oauth_provider_configs(
&self,
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(OAUTH_PROVIDER_COLUMNS);
builder.push(" ORDER BY provider_type ASC");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_oauth_provider_row).collect()
}
async fn get_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
self.get_provider(provider_type).await
}
async fn count_locked_users_if_provider_disabled(
&self,
provider_type: &str,
ldap_exclusive: bool,
) -> Result<usize, DataLayerError> {
let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL)
.bind(provider_type)
.bind(provider_type)
.bind(ldap_exclusive)
.bind(provider_type)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let locked_count = row.try_get::<i64, _>("locked_count").map_sql_err()?;
usize::try_from(locked_count.max(0)).map_err(|_| {
DataLayerError::UnexpectedValue(
"oauth_providers.locked_user_count overflowed".to_string(),
)
})
}
}
#[async_trait]
impl OAuthProviderWriteRepository for SqliteOAuthProviderRepository {
async fn upsert_oauth_provider_config(
&self,
record: &UpsertOAuthProviderConfigRecord,
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
record.validate()?;
let now = now_unix_secs();
sqlx::query(
r#"
INSERT INTO oauth_providers (
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
icon_url,
is_enabled,
created_at,
updated_at
) VALUES (
?, ?, ?,
CASE ? WHEN 'set' THEN ? WHEN 'clear' THEN NULL ELSE NULL END,
?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?
)
ON CONFLICT(provider_type) DO UPDATE SET
display_name = excluded.display_name,
client_id = excluded.client_id,
client_secret_encrypted = CASE ?
WHEN 'set' THEN ?
WHEN 'clear' THEN NULL
ELSE oauth_providers.client_secret_encrypted
END,
authorization_url_override = excluded.authorization_url_override,
token_url_override = excluded.token_url_override,
userinfo_url_override = excluded.userinfo_url_override,
scopes = excluded.scopes,
redirect_uri = excluded.redirect_uri,
frontend_callback_url = excluded.frontend_callback_url,
attribute_mapping = excluded.attribute_mapping,
extra_config = excluded.extra_config,
icon_url = excluded.icon_url,
is_enabled = excluded.is_enabled,
updated_at = excluded.updated_at
"#,
)
.bind(&record.provider_type)
.bind(&record.display_name)
.bind(&record.client_id)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.bind(record.authorization_url_override.as_deref())
.bind(record.token_url_override.as_deref())
.bind(record.userinfo_url_override.as_deref())
.bind(scopes_to_json_string(record.scopes.as_ref())?)
.bind(&record.redirect_uri)
.bind(&record.frontend_callback_url)
.bind(json_to_string(record.attribute_mapping.as_ref())?)
.bind(json_to_string(record.extra_config.as_ref())?)
.bind(record.icon_url.as_deref())
.bind(record.is_enabled)
.bind(now as i64)
.bind(now as i64)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.execute(&self.pool)
.await
.map_sql_err()?;
self.get_provider(&record.provider_type)
.await?
.ok_or_else(|| {
DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string())
})
}
async fn delete_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?")
.bind(provider_type)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
}
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 json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("invalid OAuth provider JSON field: {err}"))
})
})
.transpose()
}
fn json_from_string(
value: Option<String>,
field_name: &str,
) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"{field_name} contains invalid JSON: {err}"
))
})
})
.transpose()
}
fn scopes_to_json_string(scopes: Option<&Vec<String>>) -> Result<Option<String>, DataLayerError> {
json_to_string(
scopes
.map(|items| {
serde_json::Value::Array(
items
.iter()
.cloned()
.map(serde_json::Value::String)
.collect(),
)
})
.as_ref(),
)
}
fn parse_scopes(value: Option<String>) -> Result<Option<Vec<String>>, DataLayerError> {
let Some(value) = json_from_string(value, "oauth_providers.scopes")? else {
return Ok(None);
};
parse_scopes_value(&value)
}
fn parse_scopes_value(value: &serde_json::Value) -> Result<Option<Vec<String>>, DataLayerError> {
match value {
serde_json::Value::Null => Ok(None),
serde_json::Value::Array(items) => parse_scopes_array(items).map(Some),
serde_json::Value::String(raw) => parse_embedded_scopes(raw),
_ => Err(DataLayerError::UnexpectedValue(
"oauth_providers.scopes is not a JSON array".to_string(),
)),
}
}
fn parse_embedded_scopes(raw: &str) -> Result<Option<Vec<String>>, DataLayerError> {
let raw = raw.trim();
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
return Ok(None);
}
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
return parse_scopes_value(&decoded);
}
Ok(Some(vec![raw.to_string()]))
}
fn parse_scopes_array(items: &[serde_json::Value]) -> Result<Vec<String>, DataLayerError> {
let mut scopes = Vec::with_capacity(items.len());
for item in items {
let Some(scope) = item.as_str() else {
return Err(DataLayerError::UnexpectedValue(
"oauth_providers.scopes contains non-string value".to_string(),
));
};
let scope = scope.trim();
if !scope.is_empty() {
scopes.push(scope.to_string());
}
}
Ok(scopes)
}
fn map_oauth_provider_row(row: &SqliteRow) -> Result<StoredOAuthProviderConfig, DataLayerError> {
Ok(StoredOAuthProviderConfig::new(
row.try_get("provider_type").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("client_id").map_sql_err()?,
row.try_get("redirect_uri").map_sql_err()?,
row.try_get("frontend_callback_url").map_sql_err()?,
)?
.with_config_fields(
row.try_get("client_secret_encrypted").map_sql_err()?,
row.try_get("authorization_url_override").map_sql_err()?,
row.try_get("token_url_override").map_sql_err()?,
row.try_get("userinfo_url_override").map_sql_err()?,
parse_scopes(row.try_get("scopes").map_sql_err()?)?,
json_from_string(
row.try_get("attribute_mapping").map_sql_err()?,
"oauth_providers.attribute_mapping",
)?,
json_from_string(
row.try_get("extra_config").map_sql_err()?,
"oauth_providers.extra_config",
)?,
row.try_get("icon_url").map_sql_err()?,
row.try_get("is_enabled").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()?),
))
}
#[cfg(test)]
mod tests {
use super::SqliteOAuthProviderRepository;
use crate::run_migrations;
use aether_data_contracts::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
UpsertOAuthProviderConfigRecord,
};
fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord {
UpsertOAuthProviderConfigRecord {
provider_type: provider_type.to_string(),
display_name: format!("{provider_type} display"),
client_id: format!("{provider_type}-client"),
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")),
token_url_override: Some(format!("https://{provider_type}.example.com/token")),
userinfo_url_override: None,
scopes: Some(vec!["openid".to_string(), "profile".to_string()]),
redirect_uri: format!("https://{provider_type}.example.com/redirect"),
frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(),
attribute_mapping: Some(serde_json::json!({"email": "email"})),
extra_config: Some(serde_json::json!({"team": true})),
icon_url: None,
is_enabled: true,
}
}
#[tokio::test]
async fn sqlite_repository_round_trips_oauth_provider_configs() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteOAuthProviderRepository::new(pool.clone());
let created = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
..sample_upsert("github")
})
.await
.expect("provider should upsert");
assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1"));
assert_eq!(
created.scopes,
Some(vec!["openid".to_string(), "profile".to_string()])
);
let updated = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
display_name: "GitHub".to_string(),
..sample_upsert("github")
})
.await
.expect("provider should update");
assert_eq!(updated.display_name, "GitHub");
assert_eq!(updated.client_secret_encrypted.as_deref(), Some("secret-1"));
let listed = repository
.list_oauth_provider_configs()
.await
.expect("providers should list");
assert_eq!(listed.len(), 1);
let fetched = repository
.get_oauth_provider_config("github")
.await
.expect("provider should fetch")
.expect("provider should exist");
assert_eq!(
fetched.attribute_mapping,
Some(serde_json::json!({"email": "email"}))
);
sqlx::query(
r#"
INSERT INTO users (
id, email, username, role, auth_source, is_active, is_deleted, created_at, updated_at
) VALUES
('user-oauth', 'oauth@example.com', 'oauth-user', 'user', 'oauth', 1, 0, 1, 1),
('user-local', 'local@example.com', 'local-user', 'user', 'local', 1, 0, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("users should seed");
sqlx::query(
r#"
INSERT INTO user_oauth_links (
id, user_id, provider_type, provider_user_id, linked_at
) VALUES
('link-1', 'user-oauth', 'github', 'gh-1', 1),
('link-2', 'user-local', 'github', 'gh-2', 1)
"#,
)
.execute(&pool)
.await
.expect("oauth links should seed");
assert_eq!(
repository
.count_locked_users_if_provider_disabled("github", false)
.await
.expect("locked users should count"),
1
);
assert_eq!(
repository
.count_locked_users_if_provider_disabled("github", true)
.await
.expect("locked users should count"),
2
);
assert!(repository
.delete_oauth_provider_config("github")
.await
.expect("provider should delete"));
}
}
@@ -0,0 +1,159 @@
use std::path::PathBuf;
use std::str::FromStr;
use std::time::Duration;
use crate::{DataLayerError, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions};
use sqlx::SqlitePool as SqlxSqlitePool;
pub type SqlitePool = SqlxSqlitePool;
pub type SqlitePoolConfig = SqlDatabaseConfig;
#[derive(Debug, Clone)]
pub struct SqlitePoolFactory {
config: SqlitePoolConfig,
}
impl SqlitePoolFactory {
pub fn new(config: SqlitePoolConfig) -> Result<Self, DataLayerError> {
if config.driver != DatabaseDriver::Sqlite {
return Err(DataLayerError::InvalidConfiguration(format!(
"sqlite pool requires sqlite driver, got {}",
config.driver
)));
}
config.validate()?;
Ok(Self { config })
}
pub fn config(&self) -> &SqlitePoolConfig {
&self.config
}
pub fn connect_options(&self) -> Result<SqliteConnectOptions, DataLayerError> {
ensure_sqlite_parent_dir(self.config.url.trim())?;
let is_memory = is_sqlite_memory_url(self.config.url.trim());
SqliteConnectOptions::from_str(self.config.url.trim())
.map(|options| {
let options = options
.create_if_missing(true)
.foreign_keys(true)
.statement_cache_capacity(self.config.pool.statement_cache_capacity);
if is_memory {
options
} else {
options.journal_mode(SqliteJournalMode::Wal)
}
})
.map_err(|err| {
DataLayerError::InvalidConfiguration(format!("invalid sqlite database url: {err}"))
})
}
pub fn connect_lazy(&self) -> Result<SqlitePool, DataLayerError> {
let SqlPoolConfig {
min_connections,
max_connections,
acquire_timeout_ms,
idle_timeout_ms,
max_lifetime_ms,
..
} = self.config.pool;
Ok(SqlitePoolOptions::new()
.min_connections(min_connections)
.max_connections(max_connections)
.acquire_timeout(Duration::from_millis(acquire_timeout_ms))
.idle_timeout(Duration::from_millis(idle_timeout_ms))
.max_lifetime(Duration::from_millis(max_lifetime_ms))
.connect_lazy_with(self.connect_options()?))
}
}
fn is_sqlite_memory_url(url: &str) -> bool {
matches!(url.trim(), "sqlite::memory:" | "sqlite://:memory:")
}
fn sqlite_file_path_from_url(url: &str) -> Option<PathBuf> {
let url = url.trim();
if is_sqlite_memory_url(url) {
return None;
}
let path = url
.strip_prefix("sqlite://")
.or_else(|| url.strip_prefix("sqlite:"))?;
if path.is_empty() {
return None;
}
Some(PathBuf::from(path))
}
fn ensure_sqlite_parent_dir(url: &str) -> Result<(), DataLayerError> {
let Some(path) = sqlite_file_path_from_url(url) else {
return Ok(());
};
let Some(parent) = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
else {
return Ok(());
};
std::fs::create_dir_all(parent).map_err(|err| {
DataLayerError::InvalidConfiguration(format!(
"failed to create sqlite database parent directory '{}': {err}",
parent.display()
))
})
}
#[cfg(test)]
mod tests {
use super::SqlitePoolFactory;
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use std::path::PathBuf;
#[tokio::test]
async fn factory_builds_lazy_pool_from_valid_config() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite://./data/aether.db".to_string(),
pool: SqlPoolConfig {
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
},
};
let factory = SqlitePoolFactory::new(config).expect("factory should build");
let _pool = factory.connect_lazy().expect("lazy pool should build");
}
#[tokio::test]
async fn factory_creates_parent_directory_for_file_database() {
let db_path = unique_temp_db_path();
let parent = db_path.parent().expect("temp db path should have parent");
let _ = std::fs::remove_dir_all(parent);
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: format!("sqlite://{}", db_path.display()),
pool: SqlPoolConfig::default(),
};
let factory = SqlitePoolFactory::new(config).expect("factory should build");
let _pool = factory.connect_lazy().expect("lazy pool should build");
assert!(parent.exists());
let _ = std::fs::remove_dir_all(parent);
}
fn unique_temp_db_path() -> PathBuf {
std::env::temp_dir()
.join(format!("aether-sqlite-{}", uuid::Uuid::new_v4()))
.join("nested")
.join("aether.db")
}
}
@@ -0,0 +1,640 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::pool_scores::*;
use aether_data_query::{push_eq, push_in, push_limit, push_limit_offset, WhereClause};
use crate::error::SqlResultExt;
use crate::{DataLayerError, SqlitePool};
const SCORE_COLUMNS: &str = r#"
SELECT
id,
pool_kind,
pool_id,
member_kind,
member_id,
capability,
scope_kind,
scope_id,
score,
hard_state,
score_version,
score_reason,
last_ranked_at,
last_scheduled_at,
last_success_at,
last_failure_at,
failure_count,
last_probe_attempt_at,
last_probe_success_at,
last_probe_failure_at,
probe_failure_count,
probe_status,
updated_at
FROM pool_member_scores
"#;
#[derive(Debug, Clone)]
pub struct SqlitePoolMemberScoreRepository {
pool: SqlitePool,
}
impl SqlitePoolMemberScoreRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn find_scores_by_identity(
&self,
identity: &PoolMemberIdentity,
scope: Option<&PoolScoreScope>,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"pool_kind",
identity.pool_kind.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"pool_id",
identity.pool_id.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"member_kind",
identity.member_kind.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"member_id",
identity.member_id.clone(),
);
if let Some(scope) = scope {
push_eq(
&mut builder,
&mut where_clause,
"capability",
scope.capability.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"scope_kind",
scope.scope_kind.clone(),
);
if let Some(scope_id) = &scope.scope_id {
push_eq(
&mut builder,
&mut where_clause,
"scope_id",
scope_id.clone(),
);
} else {
where_clause.push_next(&mut builder);
builder.push("scope_id IS NULL");
}
}
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect()
}
}
#[async_trait]
impl PoolScoreReadRepository for SqlitePoolMemberScoreRepository {
async fn list_ranked_pool_members(
&self,
query: &ListRankedPoolMembersQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"pool_kind",
query.pool_kind.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"pool_id",
query.pool_id.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"capability",
query.capability.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"scope_kind",
query.scope_kind.clone(),
);
if let Some(scope_id) = &query.scope_id {
push_eq(
&mut builder,
&mut where_clause,
"scope_id",
scope_id.clone(),
);
} else {
where_clause.push_next(&mut builder);
builder.push("scope_id IS NULL");
}
if !query.hard_states.is_empty() {
let states = query
.hard_states
.iter()
.map(|state| state.as_database())
.collect::<Vec<_>>();
push_in(&mut builder, &mut where_clause, "hard_state", &states);
}
if let Some(statuses) = &query.probe_statuses {
if !statuses.is_empty() {
let statuses = statuses
.iter()
.map(|status| status.as_database())
.collect::<Vec<_>>();
push_in(&mut builder, &mut where_clause, "probe_status", &statuses);
}
}
builder.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC");
push_limit_offset(
&mut builder,
i64_from_usize(query.limit.max(1), "pool score limit")?,
i64_from_usize(query.offset, "pool score offset")?,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect()
}
async fn list_pool_member_scores(
&self,
query: &ListPoolMemberScoresQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"pool_kind",
query.pool_kind.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"pool_id",
query.pool_id.clone(),
);
if let Some(capability) = &query.capability {
push_eq(
&mut builder,
&mut where_clause,
"capability",
capability.clone(),
);
}
if let Some(scope_kind) = &query.scope_kind {
push_eq(
&mut builder,
&mut where_clause,
"scope_kind",
scope_kind.clone(),
);
}
if let Some(scope_id) = &query.scope_id {
push_eq(
&mut builder,
&mut where_clause,
"scope_id",
scope_id.clone(),
);
}
if !query.hard_states.is_empty() {
let states = query
.hard_states
.iter()
.map(|state| state.as_database())
.collect::<Vec<_>>();
push_in(&mut builder, &mut where_clause, "hard_state", &states);
}
if let Some(statuses) = &query.probe_statuses {
if !statuses.is_empty() {
let statuses = statuses
.iter()
.map(|status| status.as_database())
.collect::<Vec<_>>();
push_in(&mut builder, &mut where_clause, "probe_status", &statuses);
}
}
builder.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC");
push_limit_offset(
&mut builder,
i64_from_usize(query.limit.max(1), "pool score limit")?,
i64_from_usize(query.offset, "pool score offset")?,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect()
}
async fn list_pool_member_probe_candidates(
&self,
query: &ListPoolMemberProbeCandidatesQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new();
push_eq(
&mut builder,
&mut where_clause,
"pool_kind",
query.pool_kind.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"pool_id",
query.pool_id.clone(),
);
if let Some(capability) = &query.capability {
push_eq(
&mut builder,
&mut where_clause,
"capability",
capability.clone(),
);
}
where_clause.push_next(&mut builder);
builder
.push("hard_state IN ('available','unknown','cooldown','quota_exhausted')")
.push(" AND (probe_status IN ('never','failed','stale')")
.push(" OR (probe_status = 'ok' AND (last_probe_success_at IS NULL OR last_probe_success_at <= ")
.push_bind(i64_from_u64(
query.stale_before_unix_secs,
"pool probe stale_before_unix_secs",
)?)
.push("))")
.push(" OR (probe_status = 'in_progress' AND (last_probe_attempt_at IS NULL OR last_probe_attempt_at <= ")
.push_bind(i64_from_u64(
query.stale_before_unix_secs,
"pool probe stale_before_unix_secs",
)?)
.push(")))")
.push(
r#"
ORDER BY
CASE
WHEN last_scheduled_at IS NOT NULL AND probe_status <> 'ok' THEN 0
WHEN hard_state = 'quota_exhausted' THEN 1
WHEN hard_state = 'unknown' THEN 2
WHEN probe_status = 'stale' THEN 3
ELSE 4
END ASC,
probe_failure_count DESC,
COALESCE(last_probe_success_at, 0) ASC,
COALESCE(last_scheduled_at, 0) DESC,
member_id ASC
"#,
);
push_limit(
&mut builder,
i64_from_usize(query.limit.max(1), "pool probe candidate limit")?,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect()
}
async fn get_pool_member_scores_by_ids(
&self,
query: &GetPoolMemberScoresByIdsQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
if query.ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new();
push_in(&mut builder, &mut where_clause, "id", &query.ids);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect()
}
}
#[async_trait]
impl PoolMemberScoreWriteRepository for SqlitePoolMemberScoreRepository {
async fn upsert_pool_member_score(
&self,
score: UpsertPoolMemberScore,
) -> Result<StoredPoolMemberScore, DataLayerError> {
score.validate()?;
let stored = score.into_stored();
let score_reason = serde_json::to_string(&stored.score_reason)
.map_err(|err| DataLayerError::InvalidInput(err.to_string()))?;
sqlx::query(
r#"
INSERT INTO pool_member_scores (
id, pool_kind, pool_id, member_kind, member_id, capability, scope_kind, scope_id,
score, hard_state, score_version, score_reason, last_ranked_at, last_scheduled_at,
last_success_at, last_failure_at, failure_count, last_probe_attempt_at,
last_probe_success_at, last_probe_failure_at, probe_failure_count, probe_status, updated_at
) VALUES (
?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?
)
ON CONFLICT(id) DO UPDATE SET
pool_kind = excluded.pool_kind,
pool_id = excluded.pool_id,
member_kind = excluded.member_kind,
member_id = excluded.member_id,
capability = excluded.capability,
scope_kind = excluded.scope_kind,
scope_id = excluded.scope_id,
score = excluded.score,
hard_state = excluded.hard_state,
score_version = excluded.score_version,
score_reason = excluded.score_reason,
last_ranked_at = excluded.last_ranked_at,
last_scheduled_at = COALESCE(excluded.last_scheduled_at, pool_member_scores.last_scheduled_at),
last_success_at = COALESCE(excluded.last_success_at, pool_member_scores.last_success_at),
last_failure_at = COALESCE(excluded.last_failure_at, pool_member_scores.last_failure_at),
failure_count = excluded.failure_count,
last_probe_attempt_at = COALESCE(excluded.last_probe_attempt_at, pool_member_scores.last_probe_attempt_at),
last_probe_success_at = COALESCE(excluded.last_probe_success_at, pool_member_scores.last_probe_success_at),
last_probe_failure_at = COALESCE(excluded.last_probe_failure_at, pool_member_scores.last_probe_failure_at),
probe_failure_count = excluded.probe_failure_count,
probe_status = excluded.probe_status,
updated_at = excluded.updated_at
"#,
)
.bind(stored.id.as_str())
.bind(stored.pool_kind.as_str())
.bind(stored.pool_id.as_str())
.bind(stored.member_kind.as_str())
.bind(stored.member_id.as_str())
.bind(stored.capability.as_str())
.bind(stored.scope_kind.as_str())
.bind(stored.scope_id.as_deref())
.bind(stored.score)
.bind(stored.hard_state.as_database())
.bind(i64_from_u64(stored.score_version, "pool score version")?)
.bind(score_reason)
.bind(i64_opt_from_u64(stored.last_ranked_at, "pool score last_ranked_at")?)
.bind(i64_opt_from_u64(stored.last_scheduled_at, "pool score last_scheduled_at")?)
.bind(i64_opt_from_u64(stored.last_success_at, "pool score last_success_at")?)
.bind(i64_opt_from_u64(stored.last_failure_at, "pool score last_failure_at")?)
.bind(i64_from_u64(stored.failure_count, "pool score failure_count")?)
.bind(i64_opt_from_u64(
stored.last_probe_attempt_at,
"pool score last_probe_attempt_at",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_success_at,
"pool score last_probe_success_at",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_failure_at,
"pool score last_probe_failure_at",
)?)
.bind(i64_from_u64(
stored.probe_failure_count,
"pool score probe_failure_count",
)?)
.bind(stored.probe_status.as_database())
.bind(i64_from_u64(stored.updated_at, "pool score updated_at")?)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(stored)
}
async fn mark_pool_member_probe_in_progress(
&self,
attempt: PoolMemberProbeAttempt,
) -> Result<usize, DataLayerError> {
let rows = self
.find_scores_by_identity(&attempt.identity, attempt.scope.as_ref())
.await?;
let count = rows.len();
for mut row in rows {
row.last_probe_attempt_at = Some(attempt.attempted_at);
row.probe_status = PoolMemberProbeStatus::InProgress;
row.score_reason =
merge_score_reason_patch(row.score_reason, attempt.score_reason_patch.clone());
row.updated_at = attempt.attempted_at;
self.upsert_pool_member_score(upsert_from_stored(row))
.await?;
}
Ok(count)
}
async fn record_pool_member_probe_result(
&self,
result: PoolMemberProbeResult,
) -> Result<usize, DataLayerError> {
let rows = self
.find_scores_by_identity(&result.identity, result.scope.as_ref())
.await?;
let count = rows.len();
for mut row in rows {
row.last_probe_attempt_at = Some(result.attempted_at);
row.probe_status = result.probe_status;
if result.succeeded {
row.last_probe_success_at = Some(result.attempted_at);
row.probe_failure_count = 0;
} else {
row.last_probe_failure_at = Some(result.attempted_at);
row.probe_failure_count = row.probe_failure_count.saturating_add(1);
}
if let Some(hard_state) = result.hard_state {
row.hard_state = hard_state;
}
row.score_reason =
merge_score_reason_patch(row.score_reason, result.score_reason_patch.clone());
row.updated_at = result.attempted_at;
self.upsert_pool_member_score(upsert_from_stored(row))
.await?;
}
Ok(count)
}
async fn record_pool_member_schedule_feedback(
&self,
feedback: PoolMemberScheduleFeedback,
) -> Result<usize, DataLayerError> {
let rows = self
.find_scores_by_identity(&feedback.identity, feedback.scope.as_ref())
.await?;
let count = rows.len();
for mut row in rows {
row.last_scheduled_at = Some(feedback.scheduled_at);
match feedback.succeeded {
Some(true) => row.last_success_at = Some(feedback.scheduled_at),
Some(false) => {
row.last_failure_at = Some(feedback.scheduled_at);
row.failure_count = row.failure_count.saturating_add(1);
}
None => {}
}
if let Some(hard_state) = feedback.hard_state {
row.hard_state = hard_state;
}
row.score = score_with_delta(row.score, feedback.score_delta);
row.score_reason =
merge_score_reason_patch(row.score_reason, feedback.score_reason_patch.clone());
row.updated_at = feedback.scheduled_at;
self.upsert_pool_member_score(upsert_from_stored(row))
.await?;
}
Ok(count)
}
async fn mark_pool_member_hard_state(
&self,
identity: &PoolMemberIdentity,
scope: Option<&PoolScoreScope>,
hard_state: PoolMemberHardState,
updated_at: u64,
) -> Result<usize, DataLayerError> {
let rows = self.find_scores_by_identity(identity, scope).await?;
let count = rows.len();
for mut row in rows {
row.hard_state = hard_state;
row.updated_at = updated_at;
self.upsert_pool_member_score(upsert_from_stored(row))
.await?;
}
Ok(count)
}
async fn delete_pool_member_scores_for_member(
&self,
identity: &PoolMemberIdentity,
) -> Result<usize, DataLayerError> {
let result = sqlx::query(
r#"
DELETE FROM pool_member_scores
WHERE pool_kind = ? AND pool_id = ? AND member_kind = ? AND member_id = ?
"#,
)
.bind(identity.pool_kind.as_str())
.bind(identity.pool_id.as_str())
.bind(identity.member_kind.as_str())
.bind(identity.member_id.as_str())
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() as usize)
}
}
fn map_score_row(row: &SqliteRow) -> Result<StoredPoolMemberScore, DataLayerError> {
let score_reason_raw: String = row.try_get("score_reason").map_sql_err()?;
Ok(StoredPoolMemberScore {
id: row.try_get("id").map_sql_err()?,
pool_kind: row.try_get("pool_kind").map_sql_err()?,
pool_id: row.try_get("pool_id").map_sql_err()?,
member_kind: row.try_get("member_kind").map_sql_err()?,
member_id: row.try_get("member_id").map_sql_err()?,
capability: row.try_get("capability").map_sql_err()?,
scope_kind: row.try_get("scope_kind").map_sql_err()?,
scope_id: row.try_get("scope_id").map_sql_err()?,
score: row.try_get("score").map_sql_err()?,
hard_state: PoolMemberHardState::from_database(
row.try_get::<String, _>("hard_state")
.map_sql_err()?
.as_str(),
)?,
score_version: u64_from_i64(
row.try_get("score_version").map_sql_err()?,
"pool_member_scores.score_version",
)?,
score_reason: serde_json::from_str(&score_reason_raw).unwrap_or(serde_json::Value::Null),
last_ranked_at: u64_opt_from_i64(
row.try_get("last_ranked_at").map_sql_err()?,
"pool_member_scores.last_ranked_at",
)?,
last_scheduled_at: u64_opt_from_i64(
row.try_get("last_scheduled_at").map_sql_err()?,
"pool_member_scores.last_scheduled_at",
)?,
last_success_at: u64_opt_from_i64(
row.try_get("last_success_at").map_sql_err()?,
"pool_member_scores.last_success_at",
)?,
last_failure_at: u64_opt_from_i64(
row.try_get("last_failure_at").map_sql_err()?,
"pool_member_scores.last_failure_at",
)?,
failure_count: u64_from_i64(
row.try_get("failure_count").map_sql_err()?,
"pool_member_scores.failure_count",
)?,
last_probe_attempt_at: u64_opt_from_i64(
row.try_get("last_probe_attempt_at").map_sql_err()?,
"pool_member_scores.last_probe_attempt_at",
)?,
last_probe_success_at: u64_opt_from_i64(
row.try_get("last_probe_success_at").map_sql_err()?,
"pool_member_scores.last_probe_success_at",
)?,
last_probe_failure_at: u64_opt_from_i64(
row.try_get("last_probe_failure_at").map_sql_err()?,
"pool_member_scores.last_probe_failure_at",
)?,
probe_failure_count: u64_from_i64(
row.try_get("probe_failure_count").map_sql_err()?,
"pool_member_scores.probe_failure_count",
)?,
probe_status: PoolMemberProbeStatus::from_database(
row.try_get::<String, _>("probe_status")
.map_sql_err()?
.as_str(),
)?,
updated_at: u64_from_i64(
row.try_get("updated_at").map_sql_err()?,
"pool_member_scores.updated_at",
)?,
})
}
fn upsert_from_stored(score: StoredPoolMemberScore) -> UpsertPoolMemberScore {
UpsertPoolMemberScore {
id: score.id,
identity: PoolMemberIdentity {
pool_kind: score.pool_kind,
pool_id: score.pool_id,
member_kind: score.member_kind,
member_id: score.member_id,
},
scope: PoolScoreScope {
capability: score.capability,
scope_kind: score.scope_kind,
scope_id: score.scope_id,
},
score: score.score,
hard_state: score.hard_state,
score_version: score.score_version,
score_reason: score.score_reason,
last_ranked_at: score.last_ranked_at,
last_scheduled_at: score.last_scheduled_at,
last_success_at: score.last_success_at,
last_failure_at: score.last_failure_at,
failure_count: score.failure_count,
last_probe_attempt_at: score.last_probe_attempt_at,
last_probe_success_at: score.last_probe_success_at,
last_probe_failure_at: score.last_probe_failure_at,
probe_failure_count: score.probe_failure_count,
probe_status: score.probe_status,
updated_at: score.updated_at,
}
}
fn i64_from_usize(value: usize, field: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds signed 64-bit range")))
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,217 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row, Sqlite};
use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
};
use aether_data_query::{DialectSql, SelectColumn, SelectQuery, SqlDialect};
use crate::error::SqlResultExt;
use crate::{sqlite_optional_real, sqlite_real, DataLayerError, SqlitePool};
fn quota_snapshot_select() -> SelectQuery<'static> {
SelectQuery::new("providers").select_columns([
SelectColumn::expr("id").alias("provider_id"),
SelectColumn::expr(
DialectSql::common("billing_type").with_postgres("CAST(billing_type AS TEXT)"),
)
.alias("billing_type"),
SelectColumn::expr(DialectSql::dialect(
"CAST(monthly_quota_usd AS DOUBLE PRECISION)",
"CAST(monthly_quota_usd AS REAL)",
))
.alias("monthly_quota_usd"),
SelectColumn::expr(DialectSql::dialect(
"CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION)",
"CAST(COALESCE(monthly_used_usd, 0) AS REAL)",
))
.alias("monthly_used_usd"),
SelectColumn::expr("quota_reset_day"),
SelectColumn::expr(DialectSql::dialect(
"CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT)",
"quota_last_reset_at",
))
.alias("quota_last_reset_at_unix_secs"),
SelectColumn::expr(DialectSql::dialect(
"CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT)",
"quota_expires_at",
))
.alias("quota_expires_at_unix_secs"),
SelectColumn::expr("is_active"),
])
}
#[derive(Debug, Clone)]
pub struct SqliteProviderQuotaRepository {
pool: SqlitePool,
}
impl SqliteProviderQuotaRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[async_trait]
impl ProviderQuotaReadRepository for SqliteProviderQuotaRepository {
async fn find_by_provider_id(
&self,
provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
let mut statement = quota_snapshot_select().statement::<Sqlite>(SqlDialect::Sqlite);
statement.where_eq("id", provider_id.to_string()).limit(1);
let row = statement
.finish()
.build()
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn find_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderQuotaSnapshot>, DataLayerError> {
if provider_ids.is_empty() {
return Ok(Vec::new());
}
let mut statement = quota_snapshot_select().statement::<Sqlite>(SqlDialect::Sqlite);
statement
.where_in("id", provider_ids)
.order_by_sql("id ASC");
let rows = statement
.finish()
.build()
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_row).collect()
}
}
#[async_trait]
impl ProviderQuotaWriteRepository for SqliteProviderQuotaRepository {
async fn reset_due(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let now = i64::try_from(now_unix_secs).map_err(|_| {
DataLayerError::InvalidInput("provider quota reset timestamp overflow".to_string())
})?;
let rows_affected = sqlx::query(
r#"
UPDATE providers
SET monthly_used_usd = 0.0,
quota_last_reset_at = ?,
updated_at = ?
WHERE billing_type = 'monthly_quota'
AND is_active = 1
AND (
quota_last_reset_at IS NULL
OR (? - quota_last_reset_at) >= (quota_reset_day * 86400)
)
"#,
)
.bind(now)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
fn map_row(row: &SqliteRow) -> Result<StoredProviderQuotaSnapshot, DataLayerError> {
StoredProviderQuotaSnapshot::new(
row.try_get("provider_id").map_sql_err()?,
row.try_get("billing_type").map_sql_err()?,
sqlite_optional_real(row, "monthly_quota_usd")?,
sqlite_real(row, "monthly_used_usd")?,
row.try_get("quota_reset_day").map_sql_err()?,
row.try_get("quota_last_reset_at_unix_secs").map_sql_err()?,
row.try_get("quota_expires_at_unix_secs").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
)
}
#[cfg(test)]
mod tests {
use super::SqliteProviderQuotaRepository;
use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository,
};
use crate::run_migrations;
#[tokio::test]
async fn sqlite_repository_reads_and_resets_provider_quotas() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_provider_quotas(&pool).await;
let repository = SqliteProviderQuotaRepository::new(pool);
let quota = repository
.find_by_provider_id("provider-1")
.await
.expect("quota should load")
.expect("quota should exist");
assert_eq!(quota.monthly_used_usd, 5.0);
let quota = repository
.find_by_provider_id("provider-null-used")
.await
.expect("quota with null usage should load")
.expect("quota with null usage should exist");
assert_eq!(quota.monthly_used_usd, 0.0);
let quotas = repository
.find_by_provider_ids(&["provider-2".to_string(), "provider-1".to_string()])
.await
.expect("quotas should load");
assert_eq!(
quotas
.iter()
.map(|quota| quota.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-1", "provider-2"]
);
let reset = repository
.reset_due(1_000 + 7 * 24 * 60 * 60)
.await
.expect("quota reset should run");
assert_eq!(reset, 1);
let quota = repository
.find_by_provider_id("provider-1")
.await
.expect("quota should reload")
.expect("quota should exist");
assert_eq!(quota.monthly_used_usd, 0.0);
assert_eq!(quota.quota_last_reset_at_unix_secs, Some(605_800));
}
async fn seed_provider_quotas(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO providers (
id, name, provider_type, billing_type, monthly_quota_usd, monthly_used_usd,
quota_reset_day, quota_last_reset_at, is_active, created_at, updated_at
)
VALUES
('provider-1', 'Provider One', 'openai', 'monthly_quota', 20.0, 5.0, 7, 1000, 1, 1, 1),
('provider-2', 'Provider Two', 'openai', 'payg', NULL, 1.5, NULL, NULL, 1, 1, 1),
('provider-null-used', 'Provider Null Used', 'openai', 'payg', NULL, NULL, NULL, NULL, 1, 1, 1)
"#,
)
.execute(pool)
.await
.expect("providers should seed");
}
}
@@ -0,0 +1,521 @@
use async_trait::async_trait;
use serde_json::Value;
use sqlx::{sqlite::SqliteRow, Row};
use aether_data_contracts::repository::routing_profiles::*;
use aether_data_contracts::DataLayerError;
use crate::error::SqlResultExt;
use crate::pool::SqlitePool;
const ROUTING_GROUP_SELECT: &str = r#"
SELECT
id,
name,
description,
enabled,
is_system_default,
config_json,
version,
created_at,
updated_at,
published_at
FROM routing_groups
"#;
const ROUTING_GROUP_BINDING_SELECT: &str = r#"
SELECT
id,
group_id,
subject_type,
subject_id,
is_default,
allow_explicit_select,
created_at,
updated_at
FROM routing_group_bindings
"#;
const ROUTING_GROUP_VERSION_SELECT: &str = r#"
SELECT
id,
group_id,
version,
config_json,
created_at,
created_by
FROM routing_group_versions
"#;
#[derive(Debug, Clone)]
pub struct SqliteRoutingGroupRepository {
pool: SqlitePool,
}
impl SqliteRoutingGroupRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn reload_group(&self, id: &str) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
self.find_routing_group(RoutingGroupLookupKey::Id(id)).await
}
async fn find_binding_by_id(
&self,
id: &str,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_binding_row).transpose()
}
}
#[async_trait]
impl RoutingGroupReadRepository for SqliteRoutingGroupRepository {
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
let rows = sqlx::query(&format!("{ROUTING_GROUP_SELECT} ORDER BY name ASC, id ASC"))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_group_row).collect()
}
async fn find_routing_group(
&self,
lookup: RoutingGroupLookupKey<'_>,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let row = match lookup {
RoutingGroupLookupKey::Id(id) => sqlx::query(&format!(
"{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1"
))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?,
RoutingGroupLookupKey::Name(name) => sqlx::query(&format!(
"{ROUTING_GROUP_SELECT} WHERE name = ? LIMIT 1"
))
.bind(name)
.fetch_optional(&self.pool)
.await
.map_sql_err()?,
RoutingGroupLookupKey::SystemDefault => sqlx::query(&format!(
"{ROUTING_GROUP_SELECT} WHERE is_system_default = 1 AND enabled = 1 ORDER BY updated_at DESC, id ASC LIMIT 1"
))
.fetch_optional(&self.pool)
.await
.map_sql_err()?,
};
row.as_ref().map(map_group_row).transpose()
}
async fn list_routing_group_bindings(
&self,
query: &RoutingGroupBindingQuery,
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
let rows = sqlx::query(&format!(
r#"
{ROUTING_GROUP_BINDING_SELECT}
WHERE (? IS NULL OR group_id = ?)
AND (? IS NULL OR subject_type = ?)
AND (? IS NULL OR subject_id = ?)
ORDER BY created_at ASC, id ASC
"#
))
.bind(query.group_id.as_deref())
.bind(query.group_id.as_deref())
.bind(query.subject_type.map(binding_subject_to_database))
.bind(query.subject_type.map(binding_subject_to_database))
.bind(query.subject_id.as_deref())
.bind(query.subject_id.as_deref())
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_binding_row).collect()
}
async fn list_routing_group_versions(
&self,
group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
let rows = sqlx::query(&format!(
"{ROUTING_GROUP_VERSION_SELECT} WHERE group_id = ? ORDER BY version DESC, created_at DESC, id ASC"
))
.bind(group_id)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_version_row).collect()
}
}
#[async_trait]
impl RoutingGroupWriteRepository for SqliteRoutingGroupRepository {
async fn create_routing_group(
&self,
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
sqlx::query(
r#"
INSERT INTO routing_groups (
id, name, description, enabled, is_system_default, config_json,
version, created_at, updated_at, published_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&group.id)
.bind(&group.name)
.bind(&group.description)
.bind(group.enabled)
.bind(group.is_system_default)
.bind(json_to_string(
&group.config_json,
"routing_groups.config_json",
)?)
.bind(group.version)
.bind(group.created_at)
.bind(group.updated_at)
.bind(group.published_at)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(group)
}
async fn update_routing_group(
&self,
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let Some(mut group) = self.reload_group(id).await? else {
return Ok(None);
};
apply_group_patch(&mut group, patch)?;
sqlx::query(
r#"
UPDATE routing_groups
SET name = ?,
description = ?,
enabled = ?,
is_system_default = ?,
config_json = ?,
version = ?,
updated_at = ?,
published_at = ?
WHERE id = ?
"#,
)
.bind(&group.name)
.bind(&group.description)
.bind(group.enabled)
.bind(group.is_system_default)
.bind(json_to_string(
&group.config_json,
"routing_groups.config_json",
)?)
.bind(group.version)
.bind(group.updated_at)
.bind(group.published_at)
.bind(id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(Some(group))
}
async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM routing_group_bindings WHERE group_id = ?")
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM routing_group_versions WHERE group_id = ?")
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let rows_affected = sqlx::query("DELETE FROM routing_groups WHERE id = ?")
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?
.rows_affected();
tx.commit().await.map_sql_err()?;
Ok(rows_affected > 0)
}
async fn create_routing_group_binding(
&self,
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
sqlx::query(
r#"
INSERT INTO routing_group_bindings (
id, group_id, subject_type, subject_id, is_default,
allow_explicit_select, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&binding.id)
.bind(&binding.group_id)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.bind(binding.is_default)
.bind(binding.allow_explicit_select)
.bind(binding.created_at)
.bind(binding.updated_at)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(binding)
}
async fn delete_routing_group_binding(&self, id: &str) -> Result<bool, DataLayerError> {
Ok(
sqlx::query("DELETE FROM routing_group_bindings WHERE id = ?")
.bind(id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected()
> 0,
)
}
async fn update_routing_group_binding(
&self,
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let Some(mut binding) = self.find_binding_by_id(id).await? else {
return Ok(None);
};
apply_binding_patch(&mut binding, patch)?;
sqlx::query(
r#"
UPDATE routing_group_bindings
SET group_id = ?,
subject_type = ?,
subject_id = ?,
is_default = ?,
allow_explicit_select = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(&binding.group_id)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.bind(binding.is_default)
.bind(binding.allow_explicit_select)
.bind(binding.updated_at)
.bind(id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(Some(binding))
}
async fn create_routing_group_version(
&self,
record: CreateRoutingGroupVersionRecord,
) -> Result<StoredRoutingGroupVersion, DataLayerError> {
let version = StoredRoutingGroupVersion::new(record)?;
sqlx::query(
r#"
INSERT INTO routing_group_versions (
id, group_id, version, config_json, created_at, created_by
)
VALUES (?, ?, ?, ?, ?, ?)
"#,
)
.bind(&version.id)
.bind(&version.group_id)
.bind(version.version)
.bind(json_to_string(
&version.config_json,
"routing_group_versions.config_json",
)?)
.bind(version.created_at)
.bind(&version.created_by)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(version)
}
}
fn map_group_row(row: &SqliteRow) -> Result<StoredRoutingGroup, DataLayerError> {
Ok(StoredRoutingGroup {
id: row.try_get("id").map_sql_err()?,
name: row.try_get("name").map_sql_err()?,
description: row.try_get("description").map_sql_err()?,
enabled: row.try_get("enabled").map_sql_err()?,
is_system_default: row.try_get("is_system_default").map_sql_err()?,
config_json: json_from_string(
row.try_get("config_json").map_sql_err()?,
"routing_groups.config_json",
)?,
version: row.try_get("version").map_sql_err()?,
created_at: row.try_get("created_at").map_sql_err()?,
updated_at: row.try_get("updated_at").map_sql_err()?,
published_at: row.try_get("published_at").map_sql_err()?,
})
}
fn map_binding_row(row: &SqliteRow) -> Result<StoredRoutingGroupBinding, DataLayerError> {
Ok(StoredRoutingGroupBinding {
id: row.try_get("id").map_sql_err()?,
group_id: row.try_get("group_id").map_sql_err()?,
subject_type: binding_subject_from_database(row.try_get("subject_type").map_sql_err()?)?,
subject_id: row.try_get("subject_id").map_sql_err()?,
is_default: row.try_get("is_default").map_sql_err()?,
allow_explicit_select: row.try_get("allow_explicit_select").map_sql_err()?,
created_at: row.try_get("created_at").map_sql_err()?,
updated_at: row.try_get("updated_at").map_sql_err()?,
})
}
fn map_version_row(row: &SqliteRow) -> Result<StoredRoutingGroupVersion, DataLayerError> {
Ok(StoredRoutingGroupVersion {
id: row.try_get("id").map_sql_err()?,
group_id: row.try_get("group_id").map_sql_err()?,
version: row.try_get("version").map_sql_err()?,
config_json: json_from_string(
row.try_get("config_json").map_sql_err()?,
"routing_group_versions.config_json",
)?,
created_at: row.try_get("created_at").map_sql_err()?,
created_by: row.try_get("created_by").map_sql_err()?,
})
}
fn json_to_string(value: &Value, field_name: &str) -> Result<String, DataLayerError> {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("{field_name} contains unserializable JSON: {err}"))
})
}
fn json_from_string(value: String, field_name: &str) -> Result<Value, DataLayerError> {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("{field_name} contains invalid JSON: {err}"))
})
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
use crate::run_migrations as run_sqlite_migrations;
#[tokio::test]
async fn sqlite_routing_group_repository_round_trips() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteRoutingGroupRepository::new(pool);
repository
.create_routing_group(CreateRoutingGroupRecord {
id: "routing-group-1".to_string(),
name: "default".to_string(),
description: Some("initial".to_string()),
enabled: true,
is_system_default: true,
config_json: json!({"allowed_models": ["gpt-*"]}),
version: 1,
created_at: 10,
updated_at: 10,
published_at: None,
})
.await
.expect("group should create");
let system_default = repository
.find_routing_group(RoutingGroupLookupKey::SystemDefault)
.await
.expect("group lookup should succeed")
.expect("system default should exist");
assert_eq!(system_default.id, "routing-group-1");
repository
.update_routing_group(
"routing-group-1",
UpdateRoutingGroupRecord {
description: Some(None),
version: Some(2),
updated_at: 20,
published_at: Some(Some(20)),
..UpdateRoutingGroupRecord::default()
},
)
.await
.expect("group should update");
let binding = repository
.create_routing_group_binding(CreateRoutingGroupBindingRecord {
id: "binding-1".to_string(),
group_id: "routing-group-1".to_string(),
subject_type: RoutingGroupBindingSubject::ApiKey,
subject_id: "api-key-1".to_string(),
is_default: true,
allow_explicit_select: true,
created_at: 10,
updated_at: 10,
})
.await
.expect("binding should create");
assert_eq!(binding.subject_type, RoutingGroupBindingSubject::ApiKey);
assert_eq!(
repository
.list_routing_group_bindings(&RoutingGroupBindingQuery {
group_id: Some("routing-group-1".to_string()),
subject_type: Some(RoutingGroupBindingSubject::ApiKey),
subject_id: Some("api-key-1".to_string()),
})
.await
.expect("bindings should list")
.len(),
1
);
repository
.create_routing_group_version(CreateRoutingGroupVersionRecord {
id: "version-1".to_string(),
group_id: "routing-group-1".to_string(),
version: 2,
config_json: json!({"allowed_models": ["gpt-*"]}),
created_at: 20,
created_by: Some("admin".to_string()),
})
.await
.expect("version should create");
assert_eq!(
repository
.list_routing_group_versions("routing-group-1")
.await
.expect("versions should list")
.len(),
1
);
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,783 @@
use super::{SqliteUsageReadRepository, SqliteUsageWriteRepository};
use crate::run_migrations;
use aether_data_contracts::repository::usage::{
UpsertUsageRecord, UsageAuditListQuery, UsageDailyHeatmapQuery,
UsageDashboardDailyBreakdownQuery, UsageDashboardSummaryQuery, UsageProviderPerformanceQuery,
UsageReadRepository, UsageTimeSeriesGranularity, UsageWriteRepository,
};
#[tokio::test]
async fn sqlite_provider_performance_can_skip_timeline() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
SqliteUsageWriteRepository::new(pool.clone())
.upsert(sample_usage(
"provider-performance",
"completed",
"pending",
1_000,
))
.await
.expect("usage should upsert");
let reader = SqliteUsageReadRepository::new(pool);
let mut query = UsageProviderPerformanceQuery {
created_from_unix_secs: 0,
created_until_unix_secs: 2_000,
granularity: UsageTimeSeriesGranularity::Hour,
tz_offset_minutes: 0,
limit: 1,
provider_id: None,
model: None,
api_format: None,
endpoint_kind: None,
is_stream: None,
has_format_conversion: None,
slow_threshold_ms: 10_000,
include_timeline: true,
};
let with_timeline = reader
.summarize_usage_provider_performance(&query)
.await
.expect("provider performance should load");
assert_eq!(with_timeline.summary.request_count, 1);
assert_eq!(with_timeline.providers.len(), 1);
assert_eq!(with_timeline.timeline.len(), 1);
query.include_timeline = false;
let without_timeline = reader
.summarize_usage_provider_performance(&query)
.await
.expect("provider performance without timeline should load");
assert_eq!(without_timeline.summary, with_timeline.summary);
assert_eq!(without_timeline.providers, with_timeline.providers);
assert!(without_timeline.timeline.is_empty());
}
#[tokio::test]
async fn sqlite_usage_write_repository_upserts_and_rebuilds_stats() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool.clone());
let record = repository
.upsert(sample_usage("request-1", "completed", "pending", 1_000))
.await
.expect("usage should upsert");
assert_eq!(record.request_id, "request-1");
assert_eq!(record.api_key_id.as_deref(), Some("api-key-1"));
assert_eq!(record.total_tokens, 7);
assert_eq!(record.cache_read_input_tokens, 2);
assert_eq!(
record.request_metadata.as_ref().unwrap()["trace_id"],
"trace-1"
);
assert_eq!(
record.request_metadata.as_ref().unwrap()["upstream_is_stream"],
true
);
let upstream_is_stream: Option<i64> =
sqlx::query_scalar("SELECT upstream_is_stream FROM \"usage\" WHERE request_id = ?")
.bind("request-1")
.fetch_one(&pool)
.await
.expect("usage stream mode should load");
assert_eq!(upstream_is_stream, Some(1));
let loaded = repository
.find_by_request_id("request-1")
.await
.expect("usage should load")
.expect("usage should exist");
assert_eq!(
loaded.provider_api_key_id.as_deref(),
Some("provider-key-1")
);
let stats = sqlx::query_as::<_, (i64, i64, f64, Option<i64>)>(
"SELECT total_requests, total_tokens, total_cost_usd, last_used_at FROM api_keys WHERE id = 'api-key-1'",
)
.fetch_one(&pool)
.await
.expect("api key stats should load");
assert_eq!(stats, (1, 7, 0.5, Some(1_000)));
let provider_stats = sqlx::query_as::<_, (i64, i64, i64, i64, f64, i64, Option<i64>)>(
"SELECT request_count, success_count, error_count, total_tokens, total_cost_usd, total_response_time_ms, last_used_at FROM provider_api_keys WHERE id = 'provider-key-1'",
)
.fetch_one(&pool)
.await
.expect("provider key stats should load");
assert_eq!(provider_stats, (1, 1, 0, 7, 0.5, 42, Some(1_000)));
}
#[tokio::test]
async fn sqlite_usage_write_repository_does_not_regress_void_usage() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool);
repository
.upsert(sample_usage("request-1", "failed", "void", 1_000))
.await
.expect("void usage should upsert");
let existing = repository
.upsert(sample_usage("request-1", "pending", "pending", 1_001))
.await
.expect("stale usage should be ignored");
assert_eq!(existing.status, "failed");
assert_eq!(existing.billing_status, "void");
assert_eq!(existing.updated_at_unix_secs, 1_000);
}
#[tokio::test]
async fn sqlite_usage_write_repository_does_not_reopen_void_failure_from_late_streaming() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool);
for (request_id, status_code) in [
("request-late-active", None),
("request-late-response-start", Some(200)),
] {
let mut failed = sample_usage(request_id, "failed", "void", 1_000);
failed.status_code = Some(503);
repository
.upsert(failed)
.await
.expect("failed usage should upsert");
let mut late_streaming = sample_usage(request_id, "streaming", "pending", 1_001);
late_streaming.status_code = status_code;
late_streaming.finalized_at_unix_secs = None;
let current = repository
.upsert(late_streaming)
.await
.expect("late streaming usage should be ignored");
assert_eq!(current.status, "failed");
assert_eq!(current.billing_status, "void");
assert_eq!(current.status_code, Some(503));
assert_eq!(current.finalized_at_unix_secs, Some(1_000));
}
}
#[tokio::test]
async fn sqlite_usage_write_repository_does_not_regress_terminal_usage_from_late_streaming() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool);
repository
.upsert(sample_usage("request-1", "completed", "pending", 1_000))
.await
.expect("terminal usage should upsert");
let mut late_streaming = sample_usage("request-1", "streaming", "pending", 1_001);
late_streaming.input_tokens = Some(0);
late_streaming.output_tokens = Some(0);
late_streaming.total_tokens = Some(0);
late_streaming.cache_read_input_tokens = Some(0);
late_streaming.cache_read_cost_usd = Some(0.0);
late_streaming.total_cost_usd = Some(0.0);
late_streaming.actual_total_cost_usd = Some(0.0);
late_streaming.response_time_ms = Some(9_999);
late_streaming.first_byte_time_ms = Some(9_999);
late_streaming.finalized_at_unix_secs = None;
let current = repository
.upsert(late_streaming)
.await
.expect("late streaming usage should not regress terminal usage");
assert_eq!(current.status, "completed");
assert_eq!(current.billing_status, "pending");
assert_eq!(current.total_tokens, 7);
assert_eq!(current.cache_read_input_tokens, 2);
assert_eq!(current.total_cost_usd, 0.5);
assert_eq!(current.actual_total_cost_usd, 0.4);
assert_eq!(current.response_time_ms, Some(42));
assert_eq!(current.first_byte_time_ms, Some(12));
assert_eq!(current.finalized_at_unix_secs, Some(1_000));
assert_eq!(current.updated_at_unix_secs, 1_000);
}
#[tokio::test]
async fn sqlite_usage_write_repository_preserves_streaming_response_start_from_late_active() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool);
repository
.upsert(sample_usage(
"request-late-active",
"streaming",
"pending",
1_000,
))
.await
.expect("response-start usage should upsert");
let mut late_active = sample_usage("request-late-active", "streaming", "pending", 1_001);
late_active.status_code = None;
late_active.response_time_ms = None;
late_active.first_byte_time_ms = None;
let current = repository
.upsert(late_active)
.await
.expect("late active usage should not clear response-start fields");
assert_eq!(current.status, "streaming");
assert_eq!(current.status_code, Some(200));
assert_eq!(current.response_time_ms, Some(42));
assert_eq!(current.first_byte_time_ms, Some(12));
}
#[tokio::test]
async fn sqlite_usage_write_repository_cleans_stale_pending_requests() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool.clone());
repository
.upsert(sample_usage("request-recovered", "streaming", "pending", 1))
.await
.expect("streaming usage should upsert");
repository
.upsert(sample_usage("request-failed", "pending", "pending", 1))
.await
.expect("pending usage should upsert");
sqlx::query(
r#"
INSERT INTO request_candidates (
id,
request_id,
candidate_index,
retry_index,
status,
is_cached,
created_at
) VALUES
('candidate-recovered', 'request-recovered', 0, 0, 'streaming', 0, 1),
('candidate-failed', 'request-failed', 0, 0, 'pending', 0, 1)
"#,
)
.execute(&pool)
.await
.expect("request candidates should seed");
let summary = repository
.cleanup_stale_pending_requests(2, 10, 5, 1)
.await
.expect("cleanup should run");
assert_eq!(summary.recovered, 1);
assert_eq!(summary.failed, 1);
let recovered = repository
.find_by_request_id("request-recovered")
.await
.expect("recovered usage should load")
.expect("recovered usage should exist");
assert_eq!(recovered.status, "completed");
assert_eq!(recovered.status_code, Some(200));
let failed = repository
.find_by_request_id("request-failed")
.await
.expect("failed usage should load")
.expect("failed usage should exist");
assert_eq!(failed.status, "failed");
assert_eq!(failed.status_code, Some(504));
assert_eq!(failed.billing_status, "void");
assert_eq!(failed.total_cost_usd, 0.0);
assert_eq!(failed.finalized_at_unix_secs, Some(10));
let candidate_statuses = sqlx::query_as::<_, (String, String, Option<i64>)>(
r#"
SELECT request_id, status, finished_at
FROM request_candidates
ORDER BY request_id
"#,
)
.fetch_all(&pool)
.await
.expect("candidate statuses should load");
assert_eq!(
candidate_statuses,
vec![
(
"request-failed".to_string(),
"failed".to_string(),
Some(10_000)
),
(
"request-recovered".to_string(),
"success".to_string(),
Some(10_000)
),
]
);
let snapshot = sqlx::query_as::<_, (String, Option<i64>)>(
"SELECT billing_status, finalized_at FROM usage_settlement_snapshots WHERE request_id = 'request-failed'",
)
.fetch_one(&pool)
.await
.expect("void settlement snapshot should load");
assert_eq!(snapshot, ("void".to_string(), Some(10)));
}
#[tokio::test]
async fn sqlite_usage_write_repository_cleanup_uses_failed_candidate_status_when_present() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let repository = SqliteUsageWriteRepository::new(pool.clone());
repository
.upsert(sample_usage(
"request-upstream-reset",
"pending",
"pending",
1,
))
.await
.expect("pending usage should upsert");
repository
.upsert(sample_usage("request-stuck", "pending", "pending", 1))
.await
.expect("pending usage should upsert");
// request-upstream-reset has a failed candidate carrying a concrete 502 status
// and a connection-reset message — cleanup should use them instead of 504.
// request-stuck has only a still-pending candidate, so cleanup should fall back to 504.
sqlx::query(
r#"
INSERT INTO request_candidates (
id,
request_id,
candidate_index,
retry_index,
status,
status_code,
error_message,
is_cached,
created_at,
started_at,
finished_at
) VALUES
('candidate-reset', 'request-upstream-reset', 0, 0, 'failed', 502, 'upstream connection reset by peer', 0, 1, 2, 3),
('candidate-stuck', 'request-stuck', 0, 0, 'pending', NULL, NULL, 0, 1, NULL, NULL)
"#,
)
.execute(&pool)
.await
.expect("request candidates should seed");
let summary = repository
.cleanup_stale_pending_requests(2, 10, 5, 5)
.await
.expect("cleanup should run");
assert_eq!(summary.recovered, 0);
assert_eq!(summary.failed, 2);
let reset = repository
.find_by_request_id("request-upstream-reset")
.await
.expect("upstream-reset usage should load")
.expect("upstream-reset usage should exist");
assert_eq!(reset.status, "failed");
assert_eq!(reset.status_code, Some(502));
assert_eq!(
reset.error_message.as_deref(),
Some("upstream connection reset by peer")
);
let stuck = repository
.find_by_request_id("request-stuck")
.await
.expect("stuck usage should load")
.expect("stuck usage should exist");
assert_eq!(stuck.status, "failed");
assert_eq!(stuck.status_code, Some(504));
assert!(stuck
.error_message
.as_deref()
.is_some_and(|message| message.contains("超过 5 分钟未完成")));
}
#[tokio::test]
async fn sqlite_usage_read_repository_reads_usage_contract_views() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_stats_targets(&pool).await;
let writer = SqliteUsageWriteRepository::new(pool.clone());
writer
.upsert(sample_usage("request-1", "completed", "settled", 1_000))
.await
.expect("usage should upsert");
writer
.upsert(sample_usage("request-2", "failed", "void", 1_010))
.await
.expect("usage should upsert");
let reader = SqliteUsageReadRepository::new(pool);
let loaded = reader
.find_by_request_id("request-1")
.await
.expect("usage should load")
.expect("usage should exist");
assert_eq!(loaded.total_tokens, 7);
assert_eq!(loaded.billing_status, "settled");
let listed = reader
.list_usage_audits(&UsageAuditListQuery {
provider_name: Some("Provider One".to_string()),
newest_first: true,
..UsageAuditListQuery::default()
})
.await
.expect("usage list should load");
assert_eq!(listed.len(), 2);
assert_eq!(listed[0].request_id, "request-2");
let summary = reader
.summarize_dashboard_usage(&UsageDashboardSummaryQuery {
created_from_unix_secs: 999,
created_until_unix_secs: 1_020,
user_id: Some("user-1".to_string()),
})
.await
.expect("dashboard summary should load");
assert_eq!(summary.total_requests, 2);
assert_eq!(summary.error_requests, 1);
assert_eq!(summary.total_tokens, 10);
}
#[tokio::test]
async fn sqlite_usage_daily_heatmap_reads_imported_daily_aggregates() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO stats_daily (
id, "date", total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, actual_total_cost, is_complete, created_at, updated_at
) VALUES (
'daily-1', 86400, 9, 8, 1, 10, 20, 3, 4, 1.25, 1.0, 1, 1, 1
);
INSERT INTO stats_user_daily (
id, user_id, username, "date", total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, created_at, updated_at
) VALUES (
'user-daily-1', 'user-1', 'user one', 86400, 5, 5, 0, 7, 8, 2, 1, 0.75, 1, 1
);
"#,
)
.execute(&pool)
.await
.expect("daily aggregates should seed");
let reader = SqliteUsageReadRepository::new(pool);
let admin = reader
.summarize_usage_daily_heatmap(&UsageDailyHeatmapQuery {
created_from_unix_secs: 0,
user_id: None,
admin_mode: true,
})
.await
.expect("admin heatmap should load");
assert_eq!(admin.len(), 1);
assert_eq!(admin[0].date, "1970-01-02");
assert_eq!(admin[0].requests, 9);
assert_eq!(admin[0].total_tokens, 37);
assert_eq!(admin[0].actual_total_cost_usd, 1.0);
let user = reader
.summarize_usage_daily_heatmap(&UsageDailyHeatmapQuery {
created_from_unix_secs: 0,
user_id: Some("user-1".to_string()),
admin_mode: false,
})
.await
.expect("user heatmap should load");
assert_eq!(user.len(), 1);
assert_eq!(user[0].date, "1970-01-02");
assert_eq!(user[0].requests, 5);
assert_eq!(user[0].total_tokens, 18);
assert_eq!(user[0].actual_total_cost_usd, 0.75);
}
#[tokio::test]
async fn sqlite_usage_totals_by_user_ids_reads_imported_user_daily_aggregates() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO stats_user_daily (
id, user_id, username, "date", total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, created_at, updated_at
) VALUES (
'user-daily-1', 'user-1', 'user one', 86400, 5, 5, 0, 7, 8, 2, 1, 0.75, 1, 1
);
INSERT INTO "usage" (
request_id, id, user_id, api_key_id, provider_name, model, total_tokens,
status, billing_status, created_at_unix_ms, updated_at_unix_secs
) VALUES
('raw-before-cutoff', 'usage-1', 'user-1', 'api-key-1', 'Provider One', 'model-1', 99,
'completed', 'settled', 90000, 90000),
('raw-after-cutoff', 'usage-2', 'user-1', 'api-key-1', 'Provider One', 'model-1', 7,
'completed', 'settled', 172800, 172800);
"#,
)
.execute(&pool)
.await
.expect("usage totals fixtures should seed");
let reader = SqliteUsageReadRepository::new(pool);
let totals = reader
.summarize_usage_totals_by_user_ids(&["user-1".to_string()])
.await
.expect("user totals should load");
assert_eq!(totals.len(), 1);
assert_eq!(totals[0].user_id, "user-1");
assert_eq!(totals[0].request_count, 6);
assert_eq!(totals[0].total_tokens, 25);
}
#[tokio::test]
async fn sqlite_dashboard_daily_stats_reads_imported_daily_aggregates() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO stats_daily (
id, "date", total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, actual_total_cost, is_complete, created_at, updated_at
) VALUES (
'daily-1', 86400, 9, 8, 1, 10, 20, 3, 4, 1.25, 1.0, 1, 1, 1
);
"#,
)
.execute(&pool)
.await
.expect("daily aggregates should seed");
let reader = SqliteUsageReadRepository::new(pool);
let summary = reader
.summarize_dashboard_usage(&UsageDashboardSummaryQuery {
created_from_unix_secs: 0,
created_until_unix_secs: 172800,
user_id: None,
})
.await
.expect("dashboard summary should load");
assert_eq!(summary.total_requests, 9);
assert_eq!(summary.total_tokens, 37);
let rows = reader
.list_dashboard_daily_breakdown(&UsageDashboardDailyBreakdownQuery {
created_from_unix_secs: 0,
created_until_unix_secs: 172800,
tz_offset_minutes: 480,
user_id: None,
})
.await
.expect("dashboard daily breakdown should load");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].date, "1970-01-02");
assert_eq!(rows[0].model, "aggregate");
assert_eq!(rows[0].requests, 9);
assert_eq!(rows[0].total_tokens, 37);
}
async fn seed_stats_targets(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO users (id, auth_source, created_at, updated_at)
VALUES ('user-1', 'local', 1, 1);
INSERT INTO api_keys (id, user_id, key_hash, created_at, updated_at)
VALUES ('api-key-1', 'user-1', 'hash-1', 1, 1);
INSERT INTO providers (id, name, provider_type, created_at, updated_at)
VALUES ('provider-1', 'Provider One', 'openai', 1, 1);
INSERT INTO provider_api_keys (id, provider_id, name, created_at, updated_at)
VALUES ('provider-key-1', 'provider-1', 'Provider Key One', 1, 1);
"#,
)
.execute(pool)
.await
.expect("stats targets should seed");
}
fn sample_usage(
request_id: &str,
status: &str,
billing_status: &str,
updated_at: u64,
) -> UpsertUsageRecord {
UpsertUsageRecord {
request_id: request_id.to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
username: Some("legacy-user".to_string()),
api_key_name: Some("legacy-key".to_string()),
provider_name: "Provider One".to_string(),
model: "model-1".to_string(),
target_model: Some("target-model".to_string()),
provider_id: Some("provider-1".to_string()),
provider_endpoint_id: Some("endpoint-1".to_string()),
provider_api_key_id: Some("provider-key-1".to_string()),
request_type: Some("chat".to_string()),
api_format: Some("openai".to_string()),
api_family: Some("chat".to_string()),
endpoint_kind: Some("chat".to_string()),
endpoint_api_format: Some("openai".to_string()),
provider_api_family: Some("chat".to_string()),
provider_endpoint_kind: Some("chat".to_string()),
has_format_conversion: Some(true),
is_stream: Some(false),
input_tokens: Some(2),
output_tokens: Some(3),
total_tokens: None,
cache_creation_input_tokens: None,
cache_creation_ephemeral_5m_input_tokens: Some(0),
cache_creation_ephemeral_1h_input_tokens: Some(0),
cache_read_input_tokens: Some(2),
cache_creation_cost_usd: Some(0.0),
cache_read_cost_usd: Some(0.1),
output_price_per_1m: Some(2.0),
total_cost_usd: Some(0.5),
actual_total_cost_usd: Some(0.4),
status_code: Some(200),
error_message: None,
error_category: None,
response_time_ms: Some(42),
first_byte_time_ms: Some(12),
status: status.to_string(),
billing_status: billing_status.to_string(),
request_headers: None,
request_body: None,
request_body_ref: None,
request_body_state: None,
provider_request_headers: None,
provider_request_body: None,
provider_request_body_ref: None,
provider_request_body_state: None,
response_headers: None,
response_body: None,
response_body_ref: None,
response_body_state: None,
client_response_headers: None,
client_response_body: None,
client_response_body_ref: None,
client_response_body_state: None,
candidate_id: Some("candidate-1".to_string()),
candidate_index: Some(1),
key_name: Some("key-one".to_string()),
planner_kind: Some("default".to_string()),
route_family: Some("chat".to_string()),
route_kind: Some("completion".to_string()),
execution_path: Some("remote".to_string()),
local_execution_runtime_miss_reason: None,
request_metadata: Some(serde_json::json!({
"trace_id": "trace-1",
"upstream_is_stream": true,
})),
finalized_at_unix_secs: Some(updated_at),
created_at_unix_ms: Some(updated_at),
updated_at_unix_secs: updated_at,
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,825 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::video_tasks::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
VideoTaskWriteRepository,
};
use aether_data_contracts::DataLayerError;
use crate::error::SqlResultExt;
use crate::SqlitePool;
const VIDEO_TASK_COLUMNS: &str = r#"
SELECT
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,
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 AS next_poll_at_unix_secs,
poll_count,
max_poll_count,
created_at AS created_at_unix_ms,
submitted_at AS submitted_at_unix_secs,
completed_at AS completed_at_unix_secs,
updated_at AS updated_at_unix_secs,
error_code,
error_message,
video_url,
request_metadata
FROM video_tasks
"#;
#[derive(Debug, Clone)]
pub struct SqliteVideoTaskRepository {
pool: SqlitePool,
}
impl SqliteVideoTaskRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn find_by_id(&self, id: &str) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!("{VIDEO_TASK_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(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,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!("{VIDEO_TASK_COLUMNS} WHERE short_id = ? LIMIT 1"))
.bind(short_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"
))
.bind(user_id)
.bind(external_task_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn reload_ids(&self, ids: &[String]) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(VIDEO_TASK_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for id in ids {
separated.push_bind(id);
}
}
builder.push(")");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut tasks = rows
.iter()
.map(map_video_task_row)
.collect::<Result<Vec<_>, _>>()?;
tasks.sort_by(|left, right| {
left.next_poll_at_unix_secs
.cmp(&right.next_poll_at_unix_secs)
.then_with(|| left.updated_at_unix_secs.cmp(&right.updated_at_unix_secs))
});
Ok(tasks)
}
}
#[async_trait]
impl VideoTaskReadRepository for SqliteVideoTaskRepository {
async fn find(
&self,
key: VideoTaskLookupKey<'_>,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
match key {
VideoTaskLookupKey::Id(id) => self.find_by_id(id).await,
VideoTaskLookupKey::ShortId(short_id) => self.find_by_short_id(short_id).await,
VideoTaskLookupKey::UserExternal {
user_id,
external_task_id,
} => self.find_by_user_external(user_id, external_task_id).await,
}
}
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE status IN ('pending', 'submitted', 'queued', 'processing') ORDER BY updated_at DESC LIMIT ?"
))
.bind(limit_i64(limit, "active video task limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_due(
&self,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE status IN ('submitted', 'queued', 'processing') AND next_poll_at IS NOT NULL AND next_poll_at <= ? AND poll_count < max_poll_count ORDER BY next_poll_at ASC, updated_at ASC LIMIT ?"
))
.bind(u64_to_i64(now_unix_secs, "video task now")?)
.bind(limit_i64(limit, "due video task limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_page(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(VIDEO_TASK_COLUMNS);
push_filter(&mut builder, filter, None);
builder
.push(" ORDER BY created_at DESC, updated_at DESC LIMIT ")
.push_bind(limit_i64(limit, "video task page limit")?)
.push(" OFFSET ")
.push_bind(limit_i64(offset, "video task page offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_page_summary(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
self.list_page(filter, offset, limit).await
}
async fn count(&self, filter: &VideoTaskQueryFilter) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new("SELECT COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
count_query(builder, &self.pool).await
}
async fn count_by_status(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<Vec<VideoTaskStatusCount>, DataLayerError> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT status, COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
builder.push(" GROUP BY status ORDER BY status ASC");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(VideoTaskStatusCount {
status: VideoTaskStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
count: count_value(row.try_get("total").map_sql_err()?)?,
})
})
.collect()
}
async fn count_distinct_users(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<u64, DataLayerError> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(DISTINCT user_id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
push_clause(&mut builder, "user_id IS NOT NULL");
push_clause(&mut builder, "user_id <> ''");
count_query(builder, &self.pool).await
}
async fn top_models(
&self,
filter: &VideoTaskQueryFilter,
limit: usize,
) -> Result<Vec<VideoTaskModelCount>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT model, COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
push_clause(&mut builder, "model IS NOT NULL");
push_clause(&mut builder, "model <> ''");
builder
.push(" GROUP BY model ORDER BY total DESC, model ASC LIMIT ")
.push_bind(limit_i64(limit, "video task top models limit")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(VideoTaskModelCount {
model: row.try_get("model").map_sql_err()?,
count: count_value(row.try_get("total").map_sql_err()?)?,
})
})
.collect()
}
async fn count_created_since(
&self,
filter: &VideoTaskQueryFilter,
created_since_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new("SELECT COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, Some(created_since_unix_secs));
count_query(builder, &self.pool).await
}
}
#[async_trait]
impl VideoTaskWriteRepository for SqliteVideoTaskRepository {
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
let id = task.id.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()))
}
async fn update_if_active(
&self,
task: UpsertVideoTask,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let id = task.id.clone();
let rows_affected = bind_task(sqlx::query(UPDATE_IF_ACTIVE_SQL), task, false, true)?
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
if rows_affected == 0 {
return Ok(None);
}
self.find_by_id(&id).await
}
async fn claim_due(
&self,
now_unix_secs: u64,
claim_until_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let due = self.list_due(now_unix_secs, limit).await?;
let ids = due.iter().map(|task| task.id.clone()).collect::<Vec<_>>();
for id in &ids {
sqlx::query(
"UPDATE video_tasks SET next_poll_at = ?, updated_at = MAX(updated_at, ?) WHERE id = ?",
)
.bind(u64_to_i64(claim_until_unix_secs, "video task claim_until")?)
.bind(u64_to_i64(now_unix_secs, "video task now")?)
.bind(id)
.execute(&self.pool)
.await
.map_sql_err()?;
}
self.reload_ids(&ids).await
}
}
const UPSERT_SQL: &str = 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,
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,
created_at, submitted_at, completed_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
short_id = excluded.short_id,
request_id = excluded.request_id,
user_id = excluded.user_id,
api_key_id = excluded.api_key_id,
username = excluded.username,
api_key_name = excluded.api_key_name,
external_task_id = excluded.external_task_id,
provider_id = excluded.provider_id,
endpoint_id = excluded.endpoint_id,
key_id = excluded.key_id,
client_api_format = excluded.client_api_format,
provider_api_format = excluded.provider_api_format,
format_converted = excluded.format_converted,
model = excluded.model,
prompt = excluded.prompt,
original_request_body = excluded.original_request_body,
duration_seconds = excluded.duration_seconds,
resolution = excluded.resolution,
aspect_ratio = excluded.aspect_ratio,
size = excluded.size,
status = excluded.status,
progress_percent = excluded.progress_percent,
progress_message = excluded.progress_message,
retry_count = excluded.retry_count,
poll_interval_seconds = excluded.poll_interval_seconds,
next_poll_at = excluded.next_poll_at,
poll_count = excluded.poll_count,
max_poll_count = excluded.max_poll_count,
video_url = excluded.video_url,
error_code = excluded.error_code,
error_message = excluded.error_message,
request_metadata = excluded.request_metadata,
created_at = excluded.created_at,
submitted_at = excluded.submitted_at,
completed_at = excluded.completed_at,
updated_at = excluded.updated_at
"#;
const UPDATE_IF_ACTIVE_SQL: &str = r#"
UPDATE video_tasks SET
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 = ?,
created_at = ?,
submitted_at = ?,
completed_at = ?,
updated_at = ?
WHERE id = ?
AND status IN ('pending', 'submitted', 'queued', 'processing')
"#;
fn bind_task<'q>(
query: sqlx::query::Query<'q, Sqlite, sqlx::sqlite::SqliteArguments<'q>>,
task: UpsertVideoTask,
include_insert_id: bool,
include_update_id: bool,
) -> Result<sqlx::query::Query<'q, Sqlite, sqlx::sqlite::SqliteArguments<'q>>, DataLayerError> {
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 {
query.bind(task.id.clone())
} else {
query
};
let bound = query
.bind(task.short_id)
.bind(task.request_id)
.bind(task.user_id)
.bind(task.api_key_id)
.bind(task.username)
.bind(task.api_key_name)
.bind(task.external_task_id)
.bind(task.provider_id)
.bind(task.endpoint_id)
.bind(task.key_id)
.bind(task.client_api_format)
.bind(task.provider_api_format)
.bind(task.format_converted)
.bind(task.model)
.bind(task.prompt)
.bind(original_request_body)
.bind(optional_u32_to_i32(
task.duration_seconds,
"video task duration_seconds",
)?)
.bind(task.resolution)
.bind(task.aspect_ratio)
.bind(task.size)
.bind(status_to_database(task.status))
.bind(i32::from(task.progress_percent))
.bind(task.progress_message)
.bind(u32_to_i32(task.retry_count, "video task retry_count")?)
.bind(u32_to_i32(
task.poll_interval_seconds,
"video task poll_interval_seconds",
)?)
.bind(optional_u64_to_i64(
task.next_poll_at_unix_secs,
"video task next_poll_at",
)?)
.bind(u32_to_i32(task.poll_count, "video task poll_count")?)
.bind(u32_to_i32(
task.max_poll_count,
"video task max_poll_count",
)?)
.bind(task.video_url)
.bind(task.error_code)
.bind(task.error_message)
.bind(request_metadata)
.bind(u64_to_i64(
task.created_at_unix_ms,
"video task created_at",
)?)
.bind(optional_u64_to_i64(
task.submitted_at_unix_secs,
"video task submitted_at",
)?)
.bind(optional_u64_to_i64(
task.completed_at_unix_secs,
"video task completed_at",
)?)
.bind(u64_to_i64(
task.updated_at_unix_secs,
"video task updated_at",
)?);
if include_update_id {
Ok(bound.bind(task.id))
} else {
Ok(bound)
}
}
fn push_filter<'args>(
builder: &mut QueryBuilder<'args, Sqlite>,
filter: &'args VideoTaskQueryFilter,
created_since_unix_secs: Option<u64>,
) {
if let Some(user_id) = filter.user_id.as_deref() {
push_clause(builder, "user_id = ");
builder.push_bind(user_id);
}
if let Some(status) = filter.status {
push_clause(builder, "status = ");
builder.push_bind(status_to_database(status));
}
if let Some(model_substring) = filter.model_substring.as_deref() {
push_clause(builder, "LOWER(model) LIKE ");
builder.push_bind(format!(
"%{}%",
escape_like_pattern(&model_substring.trim().to_ascii_lowercase())
));
builder.push(" ESCAPE '\\'");
}
if let Some(client_api_format) = filter.client_api_format.as_deref() {
push_clause(builder, "client_api_format = ");
builder.push_bind(client_api_format);
}
if let Some(created_since_unix_secs) = created_since_unix_secs {
push_clause(builder, "created_at >= ");
builder.push_bind(created_since_unix_secs as i64);
}
}
fn push_clause<'args>(builder: &mut QueryBuilder<'args, Sqlite>, clause: &str) {
let sql = builder.sql();
if sql.contains(" WHERE ") || sql.contains("\nWHERE ") {
builder.push(" AND ");
} else {
builder.push(" WHERE ");
}
builder.push(clause);
}
async fn count_query(
mut builder: QueryBuilder<'_, Sqlite>,
pool: &SqlitePool,
) -> Result<u64, DataLayerError> {
let row = builder.build().fetch_one(pool).await.map_sql_err()?;
count_value(row.try_get("total").map_sql_err()?)
}
fn count_value(value: i64) -> Result<u64, DataLayerError> {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("invalid video task count result: {value}"))
})
}
fn map_video_task_row(row: &SqliteRow) -> Result<StoredVideoTask, DataLayerError> {
StoredVideoTask::new(
row.try_get("id").map_sql_err()?,
row.try_get("short_id").map_sql_err()?,
row.try_get("request_id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("external_task_id").map_sql_err()?,
row.try_get("provider_id").map_sql_err()?,
row.try_get("endpoint_id").map_sql_err()?,
row.try_get("key_id").map_sql_err()?,
row.try_get("client_api_format").map_sql_err()?,
row.try_get("provider_api_format").map_sql_err()?,
row.try_get("format_converted").map_sql_err()?,
row.try_get("model").map_sql_err()?,
row.try_get("prompt").map_sql_err()?,
parse_json(row.try_get("original_request_body").ok().flatten())?,
row.try_get("duration_seconds").map_sql_err()?,
row.try_get("resolution").map_sql_err()?,
row.try_get("aspect_ratio").map_sql_err()?,
row.try_get("size").map_sql_err()?,
VideoTaskStatus::from_database(row.try_get::<String, _>("status").map_sql_err()?.as_str())?,
row.try_get("progress_percent").map_sql_err()?,
row.try_get("progress_message").map_sql_err()?,
row.try_get("retry_count").map_sql_err()?,
row.try_get("poll_interval_seconds").map_sql_err()?,
row.try_get("next_poll_at_unix_secs").map_sql_err()?,
row.try_get("poll_count").map_sql_err()?,
row.try_get("max_poll_count").map_sql_err()?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("submitted_at_unix_secs").map_sql_err()?,
row.try_get("completed_at_unix_secs").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
row.try_get("error_code").map_sql_err()?,
row.try_get("error_message").map_sql_err()?,
row.try_get("video_url").map_sql_err()?,
parse_json(row.try_get("request_metadata").ok().flatten())?,
)
}
fn parse_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.filter(|value| !value.trim().is_empty())
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("video task JSON field is invalid: {err}"))
})
})
.transpose()
}
fn json_to_string(value: &Option<serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"video task JSON field is unserializable: {err}"
))
})
})
.transpose()
}
fn status_to_database(status: VideoTaskStatus) -> &'static str {
match status {
VideoTaskStatus::Pending => "pending",
VideoTaskStatus::Submitted => "submitted",
VideoTaskStatus::Queued => "queued",
VideoTaskStatus::Processing => "processing",
VideoTaskStatus::Completed => "completed",
VideoTaskStatus::Failed => "failed",
VideoTaskStatus::Cancelled => "cancelled",
VideoTaskStatus::Expired => "expired",
VideoTaskStatus::Deleted => "deleted",
}
}
fn escape_like_pattern(value: &str) -> String {
value
.replace('\\', "\\\\")
.replace('%', "\\%")
.replace('_', "\\_")
}
fn limit_i64(value: usize, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {name}: {value}")))
}
fn u64_to_i64(value: u64, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, DataLayerError> {
value.map(|value| u64_to_i64(value, name)).transpose()
}
fn u32_to_i32(value: u32, name: &str) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u32_to_i32(value: Option<u32>, name: &str) -> Result<Option<i32>, DataLayerError> {
value.map(|value| u32_to_i32(value, name)).transpose()
}
#[cfg(test)]
mod tests {
use super::SqliteVideoTaskRepository;
use crate::run_migrations;
use aether_data_contracts::repository::video_tasks::{
UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository,
VideoTaskStatus, VideoTaskWriteRepository,
};
#[tokio::test]
async fn sqlite_repository_writes_and_reads_video_tasks() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteVideoTaskRepository::new(pool);
repository
.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("task should insert");
repository
.upsert(UpsertVideoTask {
user_id: Some("user-2".to_string()),
model: Some("veo-3-fast".to_string()),
client_api_format: Some("gemini:video".to_string()),
created_at_unix_ms: 260,
updated_at_unix_secs: 260,
status: VideoTaskStatus::Completed,
..sample_task("task-2", VideoTaskStatus::Completed, 260)
})
.await
.expect("task should insert");
assert!(repository
.find(VideoTaskLookupKey::ShortId("short-task-1"))
.await
.expect("short lookup should load")
.is_some());
assert!(repository
.find(VideoTaskLookupKey::UserExternal {
user_id: "user-1",
external_task_id: "ext-task-1",
})
.await
.expect("user/external lookup should load")
.is_some());
let due = repository
.list_due(100, 10)
.await
.expect("due tasks should load");
assert_eq!(due.len(), 1);
let claimed = repository
.claim_due(100, 130, 10)
.await
.expect("due tasks should claim");
assert_eq!(claimed.len(), 1);
assert_eq!(claimed[0].next_poll_at_unix_secs, Some(130));
let filter = VideoTaskQueryFilter {
user_id: Some("user-2".to_string()),
status: Some(VideoTaskStatus::Completed),
model_substring: Some("veo".to_string()),
client_api_format: Some("gemini:video".to_string()),
};
assert_eq!(
repository.count(&filter).await.expect("count should load"),
1
);
assert_eq!(
repository
.count_by_status(&filter)
.await
.expect("status counts should load")[0]
.count,
1
);
assert_eq!(
repository
.top_models(&filter, 10)
.await
.expect("top models should load")[0]
.model,
"veo-3-fast"
);
let updated = repository
.update_if_active(UpsertVideoTask {
status: VideoTaskStatus::Processing,
progress_percent: 50,
..sample_task("task-1", VideoTaskStatus::Processing, 150)
})
.await
.expect("active task should update")
.expect("active task should exist");
assert_eq!(updated.progress_percent, 50);
}
fn sample_task(
id: &str,
status: VideoTaskStatus,
updated_at_unix_secs: u64,
) -> UpsertVideoTask {
UpsertVideoTask {
id: id.to_string(),
short_id: Some(format!("short-{id}")),
request_id: format!("request-{id}"),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
username: Some("user".to_string()),
api_key_name: Some("primary".to_string()),
external_task_id: Some(format!("ext-{id}")),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("provider-key-1".to_string()),
client_api_format: Some("openai:video".to_string()),
provider_api_format: Some("openai:video".to_string()),
format_converted: false,
model: Some("sora-2".to_string()),
prompt: Some("hello".to_string()),
original_request_body: Some(serde_json::json!({"prompt": "hello"})),
duration_seconds: Some(4),
resolution: Some("720p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("1280x720".to_string()),
status,
progress_percent: 0,
progress_message: None,
retry_count: 0,
poll_interval_seconds: 10,
next_poll_at_unix_secs: Some(updated_at_unix_secs),
poll_count: 0,
max_poll_count: 360,
created_at_unix_ms: updated_at_unix_secs.saturating_sub(10),
submitted_at_unix_secs: Some(updated_at_unix_secs.saturating_sub(10)),
completed_at_unix_secs: None,
updated_at_unix_secs,
error_code: None,
error_message: None,
video_url: None,
request_metadata: Some(serde_json::json!({"request": id})),
}
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff