From ecc16673eb591e26d5fb3b8680ce14786b11a37f Mon Sep 17 00:00:00 2001 From: elky Date: Thu, 10 Sep 2026 08:14:58 +0800 Subject: [PATCH] fix: harden concurrency limits and high-RPM runtime paths Bound request, stream, queue, and shutdown resource lifetimes. Reduce scheduler and Redis hot-path work and isolate database maintenance. Include regression coverage, load probes, and concurrency audit results. --- .env.example | 53 +- Cargo.lock | 7 + README.md | 14 +- .../examples/execution_runtime_harness.rs | 7 +- .../examples/tunnel_runtime_harness.rs | 7 +- .../planner/candidate_materialization.rs | 118 +- .../src/ai_serving/planner/state/scheduler.rs | 67 +- apps/aether-gateway/src/data/state/catalog.rs | 13 + apps/aether-gateway/src/data/state/testkit.rs | 27 +- .../src/dispatch/pool_scheduler.rs | 121 +- .../execution_runtime/attempt_cancellation.rs | 2 + .../execution_runtime/attempt_lifecycle.rs | 2 + .../src/execution_runtime/mod.rs | 1 + .../stream/capture_budget.rs | 251 +++ .../src/execution_runtime/stream/execution.rs | 692 ++++++- .../src/execution_runtime/stream/mod.rs | 3 + .../stream/usage_fallback.rs | 1163 +++++++++++ .../src/execution_runtime/stream_pump.rs | 782 +++----- .../execution_runtime/stream_read_timeout.rs | 187 ++ .../src/execution_runtime/sync/execution.rs | 2 + .../src/execution_runtime/transport.rs | 78 +- .../admin/provider/pool/runtime/keys.rs | 14 - .../admin/provider/pool/runtime/mod.rs | 3 +- .../admin/provider/pool/runtime/reads.rs | 654 +++++- .../handlers/admin/system/shared/update.rs | 1 + .../src/handlers/proxy/body_buffer.rs | 31 +- apps/aether-gateway/src/handlers/proxy/mod.rs | 116 +- .../handlers/proxy/websocket/live/audit.rs | 4 + .../src/handlers/shared/provider_pool.rs | 3 +- apps/aether-gateway/src/headers.rs | 431 +++- apps/aether-gateway/src/main.rs | 229 ++- apps/aether-gateway/src/maintenance/mod.rs | 14 +- .../aether-gateway/src/maintenance/runtime.rs | 3 +- .../maintenance/runtime/pool_quota_probe.rs | 688 ++++++- .../src/orchestration/effects.rs | 4 +- .../src/request_candidate_queue.rs | 18 + apps/aether-gateway/src/request_lifecycle.rs | 117 +- apps/aether-gateway/src/router.rs | 17 +- .../src/scheduler/candidate/mod.rs | 3 + .../src/scheduler/candidate/runtime.rs | 107 +- .../candidate/tests/concurrency_wait.rs | 548 +++++ .../src/scheduler/candidate/tests/mod.rs | 1 + apps/aether-gateway/src/shutdown_tests.rs | 209 ++ apps/aether-gateway/src/state/app.rs | 33 +- apps/aether-gateway/src/state/core.rs | 380 ++++ apps/aether-gateway/src/state/integrations.rs | 2 +- .../src/state/runtime/candidate_queries.rs | 10 + apps/aether-gateway/src/testkit.rs | 11 + .../src/tests/architecture/workspace_tiers.rs | 3 +- apps/aether-gateway/src/tunnel/mod.rs | 31 + apps/aether-tunnel/src/app.rs | 1 + apps/aether-tunnel/src/main.rs | 7 +- .../formats/gemini/generate_content/stream.rs | 268 ++- .../formats/src/formats/openai/chat/stream.rs | 642 +++++- .../shared/stream_core/format_matrix.rs | 18 +- .../stream_core/terminal_observation_tests.rs | 282 +++ crates/aether-billing/Cargo.toml | 3 + crates/aether-billing/src/event_enrichment.rs | 540 ++++- .../adapters/postgres/src/candidates.rs | 104 +- .../aether-data/adapters/postgres/src/lib.rs | 2 +- .../adapters/postgres/src/migrations.rs | 4 +- .../aether-data/adapters/postgres/src/pool.rs | 235 ++- .../adapters/postgres/src/settlement.rs | 400 +++- .../aether-data/adapters/postgres/src/tx.rs | 8 + .../adapters/postgres/src/usage/mod.rs | 161 +- .../postgres/src/usage/preparation.rs | 327 +++ .../adapters/postgres/src/usage/tests.rs | 140 +- .../src/repository/candidates/types.rs | 94 +- .../src/repository/usage/capture_memory.rs | 490 +++++ .../contracts/src/repository/usage/mod.rs | 6 + .../contracts/src/repository/usage/policy.rs | 1 + .../contracts/src/repository/usage/types.rs | 207 +- .../src/backend/maintenance/postgres.rs | 27 +- .../runtime/src/backend/postgres.rs | 34 + .../src/backend/stats/postgres_daily/mod.rs | 6 + .../src/backend/stats/postgres_hourly/mod.rs | 6 + .../runtime/src/backend/wallet/postgres.rs | 8 + .../src/lifecycle/backfill/postgres.rs | 2 +- .../src/repository/candidates/memory.rs | 216 +- .../src/repository/usage/memory/tests.rs | 14 + .../runtime/src/repository/usage/mod.rs | 1 + crates/aether-gateway/frontdoor/Cargo.toml | 9 +- crates/aether-gateway/frontdoor/src/body.rs | 378 +++- .../frontdoor/src/connection.rs | 229 +++ .../frontdoor/src/connection_tests.rs | 411 ++++ crates/aether-gateway/frontdoor/src/lib.rs | 6 +- crates/aether-runtime/base/src/lib.rs | 3 +- crates/aether-runtime/base/src/tracing.rs | 276 ++- .../aether-runtime/base/src/tracing/writer.rs | 1048 ++++++++++ .../base/tests/blocked_stdout_logging.rs | 253 +++ .../base/tests/nonblocking_logging.rs | 354 ++++ .../aether-runtime/base/tests/root_logging.rs | 66 +- crates/aether-runtime/state/src/lib.rs | 618 +++++- crates/aether-runtime/state/src/memory.rs | 937 ++++++++- .../aether-runtime/state/src/redis/client.rs | 276 ++- .../state/src/redis/dead_letter_transfer.lua | 44 + .../src/redis/dead_letter_transfer_tests.rs | 500 +++++ crates/aether-runtime/state/src/redis/mod.rs | 1 + .../aether-runtime/state/src/redis/runtime.rs | 211 +- .../state/src/redis/score_window.lua | 36 + .../aether-runtime/state/src/redis/stream.rs | 240 ++- .../src/redis/stream_owned_reply_tests.rs | 410 ++++ .../state/src/redis/stream_receive_tests.rs | 860 ++++++++ .../state/src/redis/usage_cleanup.rs | 175 ++ .../state/src/redis/usage_copy.lua | 22 + .../state/src/redis/usage_copy_commit.lua | 35 + .../src/redis/usage_limit_cleanup_tests.rs | 1125 +++++++++++ .../aether-runtime/state/src/score_window.rs | 25 + crates/aether-scheduler-core/src/health.rs | 65 + crates/aether-testing/integration/Cargo.toml | 1 + .../src/bin/capacity_curve_baseline.rs | 280 ++- .../src/bin/dependency_pressure_baseline.rs | 7 +- .../src/bin/failure_recovery_baseline.rs | 7 +- .../src/bin/gateway_pressure_seed.rs | 81 +- .../src/bin/gateway_tunnel_stream_baseline.rs | 7 +- .../src/bin/llm_stream_stability_baseline.rs | 7 +- .../src/bin/mock_openai_upstream.rs | 104 +- .../bin/multi_instance_admission_baseline.rs | 7 +- .../multi_instance_owner_relay_baseline.rs | 7 +- .../src/bin/single_instance_baseline.rs | 7 +- .../bin/usage_aux_counter_hotspot_baseline.rs | 7 +- .../src/bin/usage_counter_hotspot_baseline.rs | 8 +- .../bin/usage_settlement_hotspot_baseline.rs | 56 +- .../src/bin/gateway_pressure_probe.rs | 70 +- .../src/bin/redis_worker_baseline.rs | 7 +- .../src/bin/runtime_redis_pressure.rs | 7 +- crates/aether-testing/testkit/src/gateway.rs | 21 +- crates/aether-testing/testkit/src/tunnel.rs | 63 +- .../aether-usage/runtime/src/body_capture.rs | 4 + crates/aether-usage/runtime/src/config.rs | 41 + .../runtime/src/dead_letter_encoding.rs | 541 +++++ crates/aether-usage/runtime/src/event.rs | 772 ++++++- .../runtime/src/event_capture_budget.rs | 102 + crates/aether-usage/runtime/src/event_wire.rs | 999 ++++++++++ crates/aether-usage/runtime/src/executor.rs | 36 +- crates/aether-usage/runtime/src/lib.rs | 6 + crates/aether-usage/runtime/src/queue.rs | 577 +++++- .../runtime/src/queue_read_budget.rs | 373 ++++ crates/aether-usage/runtime/src/record.rs | 44 + crates/aether-usage/runtime/src/runtime.rs | 1774 +++++++++++++++-- .../src/runtime_queue_payload_tests.rs | 432 ++++ .../runtime/src/runtime_shutdown_tests.rs | 366 ++++ crates/aether-usage/runtime/src/settlement.rs | 90 +- .../runtime/src/settlement_reuse_tests.rs | 517 +++++ crates/aether-usage/runtime/src/shutdown.rs | 99 + crates/aether-usage/runtime/src/worker.rs | 1235 +++++++++++- .../runtime/src/worker_dead_letter_tests.rs | 449 +++++ crates/aether-usage/runtime/src/write.rs | 3 + .../concurrency-design-audit-2026-09-09.md | 526 +++++ 149 files changed, 27963 insertions(+), 1926 deletions(-) create mode 100644 apps/aether-gateway/src/execution_runtime/stream/capture_budget.rs create mode 100644 apps/aether-gateway/src/execution_runtime/stream/usage_fallback.rs create mode 100644 apps/aether-gateway/src/execution_runtime/stream_read_timeout.rs create mode 100644 apps/aether-gateway/src/scheduler/candidate/tests/concurrency_wait.rs create mode 100644 apps/aether-gateway/src/shutdown_tests.rs create mode 100644 crates/aether-ai/formats/src/formats/shared/stream_core/terminal_observation_tests.rs create mode 100644 crates/aether-data/adapters/postgres/src/usage/preparation.rs create mode 100644 crates/aether-data/contracts/src/repository/usage/capture_memory.rs create mode 100644 crates/aether-gateway/frontdoor/src/connection.rs create mode 100644 crates/aether-gateway/frontdoor/src/connection_tests.rs create mode 100644 crates/aether-runtime/base/src/tracing/writer.rs create mode 100644 crates/aether-runtime/base/tests/blocked_stdout_logging.rs create mode 100644 crates/aether-runtime/base/tests/nonblocking_logging.rs create mode 100644 crates/aether-runtime/state/src/redis/dead_letter_transfer.lua create mode 100644 crates/aether-runtime/state/src/redis/dead_letter_transfer_tests.rs create mode 100644 crates/aether-runtime/state/src/redis/score_window.lua create mode 100644 crates/aether-runtime/state/src/redis/stream_owned_reply_tests.rs create mode 100644 crates/aether-runtime/state/src/redis/stream_receive_tests.rs create mode 100644 crates/aether-runtime/state/src/redis/usage_cleanup.rs create mode 100644 crates/aether-runtime/state/src/redis/usage_copy.lua create mode 100644 crates/aether-runtime/state/src/redis/usage_copy_commit.lua create mode 100644 crates/aether-runtime/state/src/redis/usage_limit_cleanup_tests.rs create mode 100644 crates/aether-runtime/state/src/score_window.rs create mode 100644 crates/aether-usage/runtime/src/dead_letter_encoding.rs create mode 100644 crates/aether-usage/runtime/src/event_capture_budget.rs create mode 100644 crates/aether-usage/runtime/src/event_wire.rs create mode 100644 crates/aether-usage/runtime/src/queue_read_budget.rs create mode 100644 crates/aether-usage/runtime/src/runtime_queue_payload_tests.rs create mode 100644 crates/aether-usage/runtime/src/runtime_shutdown_tests.rs create mode 100644 crates/aether-usage/runtime/src/settlement_reuse_tests.rs create mode 100644 crates/aether-usage/runtime/src/shutdown.rs create mode 100644 crates/aether-usage/runtime/src/worker_dead_letter_tests.rs create mode 100644 docs/operations/concurrency-design-audit-2026-09-09.md diff --git a/.env.example b/.env.example index d95533644..f71c8e2f8 100644 --- a/.env.example +++ b/.env.example @@ -83,11 +83,58 @@ ADMIN_USERNAME=admin123456 # PostgreSQL 连接池配置(默认每核 4 条、总池至少 32 条且最多 100 条;多实例部署应显式分配每实例预算) # AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=12 # AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=80 +# 普通 PostgreSQL 连接默认语句超时 30 秒、锁等待超时 3 秒;0 显式关闭。 +# 这是单条 SQL 的期限,不是整个事务总期限;迁移和历史 backfill 使用独立连接放宽。 +# AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS=30000 +# AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS=3000 # AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS=2048 +# 所有监听分片共用的入站 TCP 连接上限,包含握手、空闲 keep-alive 和升级后的 WebSocket。 +# 未设置或 0 时按请求上限 + WebSocket 上限推导,最大 65536;显式配置也受 FD 余量限制。 +# 已知 FD soft limit 时最多 max(1, (FD - 256) / 2),不是整个进程的 FD/内存保证。 +# 满额的新连接在 HTTP 解析前关闭,不排队创建任务,也不会返回 HTTP 429/503。 +# AETHER_GATEWAY_MAX_HTTP_CONNECTIONS=4096 +# 停机先等待 HTTP 请求,再排空本地用量写入;以下期限单位为毫秒。 +# 进程管理器的强杀期限应覆盖两阶段之和,再预留至少 10 秒收尾。 +# AETHER_GATEWAY_HTTP_SHUTDOWN_TIMEOUT_MS=30000 +# AETHER_GATEWAY_USAGE_SHUTDOWN_TIMEOUT_MS=30000 +# 每客户端、每 origin 的上游空闲连接缓存;不限制活动请求或流持续时间。 +# AETHER_GATEWAY_UPSTREAM_POOL_MAX_IDLE_PER_HOST=32 +# AETHER_GATEWAY_UPSTREAM_POOL_IDLE_TIMEOUT_MS=15000 # AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB=256 -# 请求体完整读取总超时默认关闭;确需限制时配置 1000-600000 毫秒的非零值。 -# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=0 -# 单请求解压后 Payload 上限(MiB),默认 256;显式设为 0 才表示不限制。 +# 请求体按实际缓冲增长申请额度,解压同时计入输入和输出;额度不足返回 503。 +# 请求体完整读取总超时默认 120000 毫秒;非零值限制在 1000-600000,显式 0 关闭。 +# AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS=120000 +# 上游流首包后空闲超时默认 300000 毫秒;执行配置 read_ms 优先,显式 0 关闭。 +# AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS=300000 +# 流式响应诊断捕获共享预算默认 128 MiB,包含 provider/client 分配;不足时截断审计副本。 +# 显式 0 关闭此类捕获;不限制协议解析、终态编码和 usage 队列的总内存。 +# AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES=134217728 +# usage 诊断正文共享预算,按 JSON 堆内存估算,默认 128 MiB。 +# 覆盖终态队列 seed、Redis 解码事件、数据库写入 DTO 及正文副本;额度随正文释放。 +# 不足或显式 0 时先保留计费事实再舍弃正文;已有清空/禁用状态保留,其余标为截断。 +# 不包含原始 Redis 批次、解码临时分配、序列化和压缩结果、协议观察缓冲或进程总内存。 +# AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES=134217728 +# 新 usage 队列消息的完整 JSON payload 上限,默认 1 MiB;显式 0 非法。 +# 超限先保留计费事实并舍弃诊断字段;仍超限或无法保留计费语义则拒绝入队,终态尝试受限落库,失败即明确失败。 +# 不限制存量 Redis 消息、整个读取批次、DLQ 或进程总内存。 +# usage_runtime_queue_payload_* 降级/拒绝计数包含入队和重试预校验的编码尝试,不代表唯一事件数。 +# AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES=1048576 +# usage worker 读取/重领共用的逻辑 payload 预留:全进程默认 128 MiB,单批目标 8 MiB。 +# 按当前消息上限推导 COUNT,默认最多 8 条;预留覆盖整批处理和确认,额度不足等待。 +# 当前消息上限不能超过总预留额度;单批目标不足一条时仍读一条。0 或非法值回退默认。 +# 历史/其他生产者的大消息继续处理并计数,不是 RESP、实际堆内存或 DLQ 的硬上限。 +# AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES=134217728 +# AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES=8388608 +# DLQ 原文和最坏 JSON 编码独立预留默认 64 MiB,后台编码及写入最多同时 4 个任务。 +# 入场前一次预留,额度占满或单条超预算立即失败并保留 pending 原消息,后续重领。 +# 编码失败不阻塞同批其余正常消息;批次仍报告失败,只有成功项会被确认。 +# JSON 按字符串最多 6 倍转义保守估算;不截原账务字段,不包含 Redis 命令/连接副本或 RSS。 +# 0/非法值回退默认;bytes 最大约 4 GiB,jobs 最大 128。超大存量可能需调高额度后恢复。 +# AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES=67108864 +# AETHER_USAGE_DLQ_ENCODING_MAX_JOBS=4 +# 内置 Redis 死信转移要求 Redis 7+ 及 EVAL/TYPE/XPENDING/XADD/XACK/XDEL 权限。 +# stream 与 DLQ 不能同名;Cluster 还要求两键同 slot,现有默认键未自动迁移。 +# 单请求解压后 Payload 上限(MiB),默认 256;显式 0 仍受 256 MiB 硬上限保护。 # AETHER_MAX_REQUEST_BODY_MB=256 # AETHER_GATEWAY_SECURITY_CACHE_TTL_MS=1000 # AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB=64 diff --git a/Cargo.lock b/Cargo.lock index a2171f728..d83fc53c1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -116,6 +116,7 @@ name = "aether-billing" version = "0.1.0" dependencies = [ "aether-data-contracts", + "aether-runtime-state", "aether-usage-runtime", "async-trait", "serde", @@ -370,8 +371,12 @@ dependencies = [ "bytes", "futures-util", "http", + "http-body-util", + "hyper", + "hyper-util", "serde_json", "tokio", + "tokio-util", "tower", "tracing", "tracing-subscriber", @@ -421,6 +426,7 @@ dependencies = [ "aether-data", "aether-data-contracts", "aether-gateway", + "aether-runtime", "aether-runtime-state", "aether-testkit", "async-stream", @@ -5274,6 +5280,7 @@ dependencies = [ "bytes", "futures-core", "futures-sink", + "futures-util", "pin-project-lite", "tokio", ] diff --git a/README.md b/README.md index 818bcf9ba..30d4a5d77 100644 --- a/README.md +++ b/README.md @@ -121,9 +121,17 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙 - `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}` - `DATABASE_URL`:PostgreSQL 连接串,例如 `postgresql://USER:PASSWORD@HOST:5432/aether` - `AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS` / `AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS`:数据库连接池手动覆盖值;未配置时 PostgreSQL 按每核 `4` 条自动推导,总池范围为 `32-100`。该预算按进程计算,多实例部署应按数据库连接上限显式分配 +- `AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS` / `AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS`:普通数据库连接的单条 SQL / 锁等待期限,默认 `30000` / `3000` 毫秒,显式 `0` 关闭;不是整个事务总期限。迁移与历史 backfill 使用独立连接放宽,事务可通过局部设置覆盖 +- `AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES`:usage 诊断正文共享预算,默认 `134217728`(128 MiB),按 JSON 堆内存估算,覆盖进入终态队列的 seed、Redis 解码后的事件、数据库写入 DTO 及其正文副本。额度不足或显式 `0` 时先保留计费事实,再舍弃诊断正文;已有清空或禁用状态保持不变,其余标记截断。预算随正文保留到释放,后台构建或压缩不会因调用方取消而提前归还额度。该额度不覆盖原始 Redis 批次、解码临时分配、序列化及压缩结果、协议观察缓冲或进程总内存;可通过 `usage_runtime_event_capture_memory_*` 指标观察 +- `AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES`:新增 usage 队列消息的完整 JSON payload 上限,默认 `1048576`(1 MiB),按序列化后的 UTF-8 字节计算,显式 `0` 非法。超限先保留计费事实并舍弃诊断字段;仍超限或无法保留计费语义时拒绝入队,终态消息尝试受限数据库落库,失败则明确失败,不继续 Redis 重试。该限制不覆盖存量 Redis 消息、整个读取批次、DLQ 或进程总内存。`usage_runtime_queue_payload_*` 导出上限及进程级降级、拒绝编码尝试次数,包含入队和重试预校验,不代表唯一事件数;`usage_runtime_enqueue_retry_permanent_failure_total` 记录永久输入错误导致的重试拒绝或终止 +- `AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES` / `AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES`:usage worker 读取和重领共用的进程级逻辑 payload 预留,默认总额 `134217728`(128 MiB)、单批目标 `8388608`(8 MiB)。按当前 `QUEUE_PAYLOAD_MAX_BYTES` 推导实际 COUNT,默认最多读取 8 条,自动扩容使用实际 COUNT 判断批次是否读满。预留覆盖读取、整批处理和确认,额度不足等待;取消/失败释放。单批目标至少允许一条,当前 payload 上限大于总额时读取报配置错误。`0` 或非法值回退默认,过大值收敛到约 4 GiB 的有效总额。收到消息后按全部字段值长度缩减多余预留;历史消息、其他生产者使用更高上限或额外字段可能超出估算,仍继续原计费流程并记录 `usage_runtime_queue_read_oversized_*`。`usage_runtime_queue_read_*` 同时导出预留、等待与累计字段字节;该预留不是 RESP 解码、连接缓冲容量、字段结构、诊断 JSON、DLQ 或进程 RSS 的硬上限,旧公开 Vec 读取接口不携带处理阶段预留 +- `AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES` / `AETHER_USAGE_DLQ_ENCODING_MAX_JOBS`:死信原文和 JSON 编码独立共享预留,默认 `67108864`(64 MiB)、最多 `4` 个后台编码及写入任务。根据原始字段、ID、错误字符串及 JSON 最坏 6 倍转义一次预留;预算占满或单条超总额时立即失败,worker 保留原消息等待重领,不截断账务原文。编码失败会继续处理同批其他消息,只确认成功项,批次末尾仍报告失败;存储转移失败则停止该批后续处理。取消编码等待不会提前归还仍在后台使用的额度。`0`/非法值回退默认,bytes 最大约 4 GiB,jobs 最大 128;超大存量消息可能需要调高总额后恢复。`usage_runtime_dlq_encoding_*` 导出额度、在途任务、拒绝和编码尝试次数;不包含字段结构、字符串额外容量、Redis 命令/连接副本或进程 RSS。内置 Redis/Memory worker 将死信追加、源 ACK 和删除作为一次原子转移,同一源 stream、消费组及 pending ID 的并发或重试只追加一次;Redis 要求 7+ 及 `EVAL/TYPE/XPENDING/XADD/XACK/XDEL` 权限,Cluster 两键须同 slot(当前默认键不自动迁移)。源和 DLQ 不能同名。源已不在 PEL 时不宣称已归档;外部 ACK/trim/delete 及多消费组仍有原来的删除语义。公开 `push_dead_letter` 仍为追加接口,未实现新原子 trait 方法的外部后端沿用追加后 ACK,仍可能重复归档 - `AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS`:单实例请求并发上限;未配置时按 CPU 自动推导(基础范围 `512-65536`),低文件描述符预算时会进一步下调 -- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB` -- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:可选的请求体完整读取超时;默认或显式设为 `0` 时关闭,非零值限制在 `1000-600000ms` +- `AETHER_GATEWAY_MAX_HTTP_CONNECTIONS`:二进制入口全部监听分片共用的入站 TCP 连接上限,包含握手、空闲 keep-alive 和 HTTP 升级后仍存活的 socket。未设置或 `0` 时使用请求上限与 WebSocket 上限之和;自动及显式值均最多 `65536`,已知 FD soft limit 时进一步限制为 `max(1, (FD - 256) / 2)`。接入后立即尝试取得额度,满额时关闭新连接,不创建 HTTP 处理任务、不等待额度,不返回 HTTP 状态码;取消、解析失败和连接释放归还,WebSocket 升级不会提前归还。HTTP/2 多流共用一个 TCP 许可,原请求和 WebSocket 准入仍独立有效。`gateway_http_connections_*` 导出配置上限、当前数、高水位、拒绝数及 accept 错误数。该限制不包含 kernel backlog、上游、Redis 或数据库连接,也不是整个进程 FD/内存硬上限。临时 accept 错误重试,资源类错误退避一秒后重试,避免单次错误停止监听 +- `AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB`:单实例同时读取和解压请求体的加权内存预算,默认 `256MB`;压缩和未知长度上传按实际缓冲增长申请额度,解压时计入同时存活的输入和输出。额度不足返回 `503`;接近单请求上限的压缩上传需要为输入和解压输出预留额外预算 +- `AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS`:请求体完整读取超时,默认 `120000ms`;显式设为 `0` 时关闭,非零值限制在 `1000-600000ms` +- `AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS`:上游流首包后的空闲超时,默认 `300000ms`;请求执行配置中的 `read_ms` 优先,显式 `0` 关闭对应超时。网关生成的 keepalive 不会重置计时 +- `AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES`:进程内流式响应诊断捕获的共享字节预算,默认 `134217728`(128 MiB);包含 provider/client 捕获容量和扩容时的新旧分配。额度不足时仅截断审计副本,显式 `0` 关闭此类捕获;协议解析、客户端传输和计费观察继续执行。该预算不包含协议解析缓冲、终态编码及 usage 队列副本,不是进程总内存上限 - `AETHER_MAX_REQUEST_BODY_MB`:单请求解压后请求体上限,默认 `256MB`;显式设为 `0` 表示不再收紧默认值,但仍受 `256MB` 安全硬上限约束 - `AETHER_MAX_INTERNAL_BUFFERED_BODY_MB`:heartbeat、管理探测等内部整包响应体上限,默认 `64MB`;显式设为 `0` 表示不再收紧默认值,但仍受 `256MB` 安全硬上限约束 - `AETHER_TUNNEL_NODE_STATUS_QUEUE_CAPACITY`:隧道节点状态上报队列容量,默认 `1024`;满载时拒绝新事件,避免控制面故障导致无界内存增长 @@ -144,6 +152,8 @@ Aether Tunnel 是配套的正向代理节点,部署在海外 VPS 上,为墙 - `RUST_LOG`:Rust 日志过滤,例如 `aether_gateway=info`、`aether_gateway=debug,sqlx=warn` - `DB_PASSWORD` / `REDIS_PASSWORD`:Docker Compose 后端密码,首次安装时分别随机生成;手工部署必须替换示例占位值,不要互相复用 +运行日志由独立后台线程写入 stdout 和文件,每个输出队列最多 4096 条、保留正文最多 8 MiB(包含正在写入的记录),单条最多 256 KiB。队列满、正文预算不足或单条超限时整条丢弃,不等待日志设备;`Both` 两个输出独立降级。`logging_stdout_*` 和 `logging_file_*` 指标记录丢弃和写入错误,网关指标沿用其命名空间前缀。正常退出时日志最多等待 2 秒排空;这不是请求优雅排空或整个进程退出期限。日志格式化仍在调用线程执行,日志预算不包含格式化临时内存,运行日志也不能作为可靠计费账本。 + ### S3 备份离线恢复 先从 S3 下载完整的 `.json.zst.aes256gcm` 对象,再使用原始的完整 S3 object key 做认证解密。恢复工具只验证并输出本地 JSON,不会直接写数据库;数据库导入仍应在维护窗口通过管理端完成。 diff --git a/apps/aether-gateway/examples/execution_runtime_harness.rs b/apps/aether-gateway/examples/execution_runtime_harness.rs index 7cc050093..58fc6149e 100644 --- a/apps/aether-gateway/examples/execution_runtime_harness.rs +++ b/apps/aether-gateway/examples/execution_runtime_harness.rs @@ -71,8 +71,13 @@ struct Args { distributed_request_command_timeout_ms: u64, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { let _ = rustls::crypto::ring::default_provider().install_default(); init_service_runtime(ServiceRuntimeConfig::new( diff --git a/apps/aether-gateway/examples/tunnel_runtime_harness.rs b/apps/aether-gateway/examples/tunnel_runtime_harness.rs index 70015ffb0..2d1c5d24d 100644 --- a/apps/aether-gateway/examples/tunnel_runtime_harness.rs +++ b/apps/aether-gateway/examples/tunnel_runtime_harness.rs @@ -88,8 +88,13 @@ struct Args { distributed_request_command_timeout_ms: u64, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_service_runtime(ServiceRuntimeConfig::new( "aether-tunnel-standalone", "aether_gateway=info", diff --git a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs index 8e1b740f0..0369f5118 100644 --- a/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs +++ b/apps/aether-gateway/src/ai_serving/planner/candidate_materialization.rs @@ -996,7 +996,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> { ); if page_is_exact_auth_api_key_concurrency_limited(&page) { - if self.wait_for_auth_api_key_concurrency_retry().await { + if self.wait_for_auth_api_key_concurrency_retry().await? { continue; } self.persist_final_auth_api_key_concurrency_skips(page.skipped_candidates) @@ -1087,20 +1087,23 @@ impl<'a> RequestedModelAttemptPageCursor<'a> { } } - async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> bool { + async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> Result { let now = Instant::now(); let deadline = *self .auth_api_key_concurrency_wait_deadline .get_or_insert(now + AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET); - if now >= deadline { - return false; + if !crate::scheduler::candidate::wait_for_auth_api_key_concurrency_retry( + self.state.app(), + Some(&self.auth_snapshot), + deadline, + AUTH_API_KEY_CONCURRENCY_RETRY_DELAY, + ) + .await? + { + return Ok(false); } - - let sleep_duration = - AUTH_API_KEY_CONCURRENCY_RETRY_DELAY.min(deadline.saturating_duration_since(now)); - tokio::time::sleep(sleep_duration).await; self.page_cursor.restart_scan(); - true + Ok(true) } async fn persist_final_auth_api_key_concurrency_skips( @@ -2289,6 +2292,103 @@ mod tests { candidate } + #[tokio::test] + async fn auth_concurrency_wait_paged_scan_retries_once_at_original_deadline() { + let now = current_unix_ms(); + let active = serde_json::from_value(json!({ + "id": "active-candidate", + "request_id": "active-request", + "api_key_id": "api-key-1", + "candidate_index": 0, + "retry_index": 0, + "status": "pending", + "is_cached": false, + "created_at_unix_ms": now, + "started_at_unix_ms": now + })) + .expect("active candidate should build"); + let repository = Arc::new(InMemoryRequestCandidateRepository::seed([active])); + let app = AppState::new() + .expect("state should build") + .with_data_state_for_tests( + GatewayDataState::with_request_candidate_repository_for_tests(repository), + ); + let mut auth_snapshot = sample_auth_snapshot(); + auth_snapshot.api_key_concurrent_limit = Some(1); + let page_cursor = LocalCandidatePreselectionPageCursor::new( + PlannerAppState::new(&app), + &crate::system_features::ModelDirectivePolicySnapshot::default(), + "openai:chat", + "gpt-5", + None, + true, + None, + &auth_snapshot, + None, + None, + None, + false, + LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel, + true, + Some("trace-auth-wait"), + ) + .await; + let mut cursor = RequestedModelAttemptPageCursor { + state: PlannerAppState::new(&app), + trace_id: "trace-auth-wait".to_string(), + client_api_format: "openai:chat".to_string(), + requested_model: "gpt-5".to_string(), + auth_snapshot, + client_session_affinity: None, + required_capabilities: None, + routing_policy: None, + sticky_session_token: None, + request_auth_channel: None, + skipped_user_id: "user-1".to_string(), + skipped_api_key_id: "api-key-1".to_string(), + skipped_required_capabilities: None, + skipped_error_context: "test auth wait", + record_runtime_miss_diagnostic: false, + resolution_mode: LocalCandidateResolutionMode::Standard, + decorate_skipped_candidate: Arc::new(identity_skipped_candidate), + page_cursor, + pending_items: VecDeque::new(), + skipped_provider_ids: BTreeSet::new(), + skipped_endpoint_ids: BTreeSet::new(), + skipped_credential_ids: BTreeSet::new(), + candidate_count: 0, + next_candidate_index: 0, + remembered_affinity: false, + scheduler_cache_affinity_enabled: false, + auth_api_key_concurrency_wait_deadline: None, + deferred_error: None, + }; + + let started = Instant::now(); + let mut scan_restarts = 0; + while cursor + .wait_for_auth_api_key_concurrency_retry() + .await + .expect("auth wait should succeed") + { + scan_restarts += 1; + } + assert_eq!( + scan_restarts, 1, + "blocked polls must not restart page scans" + ); + assert!(started.elapsed() >= AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET); + let original_deadline = cursor.auth_api_key_concurrency_wait_deadline; + assert!(!cursor + .wait_for_auth_api_key_concurrency_retry() + .await + .expect("expired auth wait should succeed")); + assert_eq!( + cursor.auth_api_key_concurrency_wait_deadline, + original_deadline + ); + } + #[tokio::test] async fn pool_group_keys_are_not_persisted_as_available_before_attempt() { let repository = Arc::new(InMemoryRequestCandidateRepository::default()); diff --git a/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs b/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs index 724e77a6e..b1c559166 100644 --- a/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs +++ b/apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs @@ -1,9 +1,7 @@ use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate}; use std::time::Duration; -use tokio::time::Instant; use super::{GatewayAuthApiKeySnapshot, PlannerAppState}; -use crate::clock::current_unix_secs; use crate::constants::{ API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS, API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS, }; @@ -97,11 +95,13 @@ impl<'a> PlannerAppState<'a> { ), GatewayError, > { - let wait_timeout = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS); - let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1)); - let wait_deadline = Instant::now() + wait_timeout; - let mut attempt_now_unix_secs = now_unix_secs; - loop { + crate::scheduler::candidate::select_with_auth_concurrency_wait( + self.app(), + auth_snapshot, + now_unix_secs, + Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS), + Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS), + |attempt_now_unix_secs| async move { let result = crate::scheduler::candidate::list_selectable_candidates_with_skip_reasons_for_request_operation( self.app().data.as_ref(), self.app(), @@ -118,21 +118,13 @@ impl<'a> PlannerAppState<'a> { ) .await?; - if !crate::scheduler::candidate::is_exact_all_skipped_by_auth_limit( + let auth_limit_blocked = crate::scheduler::candidate::is_exact_all_skipped_by_auth_limit( &result.0, &result.1, - ) { - return Ok(result); - } - - let now = Instant::now(); - if now >= wait_deadline { - return Ok(result); - } - - let remaining = wait_deadline.duration_since(now); - tokio::time::sleep(wait_interval.min(remaining)).await; - attempt_now_unix_secs = current_unix_secs(); - } + ); + Ok((result, auth_limit_blocked)) + }, + ) + .await } #[allow(clippy::too_many_arguments)] @@ -178,13 +170,14 @@ impl<'a> PlannerAppState<'a> { now_unix_secs: u64, ordering_config: SchedulerOrderingConfig, ) -> Result, GatewayError> { - let wait_timeout = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS); - let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1)); - let wait_deadline = Instant::now() + wait_timeout; - let mut attempt_now_unix_secs = now_unix_secs; - - loop { - let (result, auth_limit_blocked) = crate::scheduler::candidate::list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal( + crate::scheduler::candidate::select_with_auth_concurrency_wait( + self.app(), + auth_snapshot, + now_unix_secs, + Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS), + Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS), + |attempt_now_unix_secs| { + crate::scheduler::candidate::list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal( self.app().data.as_ref(), self.app(), candidate_api_format, @@ -195,20 +188,8 @@ impl<'a> PlannerAppState<'a> { attempt_now_unix_secs, ordering_config, ) - .await?; - - if !auth_limit_blocked { - return Ok(result); - } - - let now = Instant::now(); - if now >= wait_deadline { - return Ok(result); - } - - let remaining = wait_deadline.duration_since(now); - tokio::time::sleep(wait_interval.min(remaining)).await; - attempt_now_unix_secs = current_unix_secs(); - } + }, + ) + .await } } diff --git a/apps/aether-gateway/src/data/state/catalog.rs b/apps/aether-gateway/src/data/state/catalog.rs index 43ecf7a90..ddedda015 100644 --- a/apps/aether-gateway/src/data/state/catalog.rs +++ b/apps/aether-gateway/src/data/state/catalog.rs @@ -81,6 +81,19 @@ impl GatewayDataState { } } + pub(crate) async fn list_recent_runtime_request_candidates( + &self, + limit: usize, + ) -> Result, DataLayerError> { + match &self.request_candidate_reader { + Some(repository) => repository + .list_recent_runtime(limit) + .await + .map(sanitize_request_candidate_rows), + None => Ok(Vec::new()), + } + } + pub(crate) async fn list_finalized_request_candidates_by_endpoint_ids_since( &self, endpoint_ids: &[String], diff --git a/apps/aether-gateway/src/data/state/testkit.rs b/apps/aether-gateway/src/data/state/testkit.rs index f1cf3856f..19e32bc1f 100644 --- a/apps/aether-gateway/src/data/state/testkit.rs +++ b/apps/aether-gateway/src/data/state/testkit.rs @@ -19,9 +19,14 @@ use aether_data::repository::management_tokens::{ }; use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode}; use aether_data::repository::proxy_nodes::{ProxyNodeReadRepository, ProxyNodeWriteRepository}; +use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository; use aether_data::repository::users::{ InMemoryUserReadRepository, StoredUserAuthRecord, UserReadRepository, }; +use aether_data_contracts::repository::routing_profiles::{ + StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion, +}; +use aether_routing_core::RoutingGroupConfig; use sha2::{Digest, Sha256}; use super::{GatewayDataConfig, GatewayDataState}; @@ -142,6 +147,24 @@ impl GatewayDataState { provider_catalog_repository; let usage_reader: Arc = usage_repository.clone(); let usage_writer: Arc = usage_repository; + let routing_groups = Arc::new(InMemoryRoutingGroupRepository::seed( + [StoredRoutingGroup { + id: "system-default".to_string(), + name: "system-default".to_string(), + description: Some("pressure harness routing strategy".to_string()), + enabled: true, + is_system_default: true, + sort_order: 0, + config_json: serde_json::to_value(RoutingGroupConfig::default()) + .expect("default routing config should serialize"), + version: 1, + created_at: 1, + updated_at: 1, + published_at: Some(1), + }], + std::iter::empty::(), + std::iter::empty::(), + )); Self { config: GatewayDataConfig::disabled().with_encryption_key(encryption_key), @@ -174,8 +197,8 @@ impl GatewayDataState { pool_score_writer: None, provider_quota_reader: None, provider_quota_writer: None, - routing_group_reader: None, - routing_group_writer: None, + routing_group_reader: Some(routing_groups.clone()), + routing_group_writer: Some(routing_groups), usage_reader: Some(usage_reader), usage_writer: Some(usage_writer), user_reader: None, diff --git a/apps/aether-gateway/src/dispatch/pool_scheduler.rs b/apps/aether-gateway/src/dispatch/pool_scheduler.rs index 1b35ef641..6338b93ff 100644 --- a/apps/aether-gateway/src/dispatch/pool_scheduler.rs +++ b/apps/aether-gateway/src/dispatch/pool_scheduler.rs @@ -35,7 +35,6 @@ use crate::ai_serving::{ SkippedLocalExecutionCandidate, }; use crate::clock::current_unix_ms; -use crate::handlers::shared::provider_pool::read_admin_provider_pool_runtime_state; use crate::handlers::shared::provider_pool::{ admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value, }; @@ -44,6 +43,9 @@ use crate::handlers::shared::provider_pool::{ read_admin_provider_pool_key_cooldown_reason, AdminProviderPoolConfig, AdminProviderPoolRuntimeState, AdminProviderPoolSchedulingPreset, }; +use crate::handlers::shared::provider_pool::{ + read_provider_pool_scheduling_runtime_state, read_provider_pool_sticky_bound_key_id, +}; use crate::handlers::shared::{parse_catalog_auth_config_json, provider_key_health_summary}; use crate::maintenance::spawn_pool_quota_probe_replenish_for_request; use crate::orchestration::LocalExecutionCandidateMetadata; @@ -141,7 +143,7 @@ async fn schedule_pool_page_candidates( AdminProviderPoolRuntimeState::default() } else { let runtime_started_at = std::time::Instant::now(); - let runtime = read_admin_provider_pool_runtime_state( + let runtime = read_provider_pool_scheduling_runtime_state( state.app().runtime_state.as_ref(), provider_id.as_str(), &key_ids, @@ -982,15 +984,13 @@ impl<'a> PoolKeyCursor<'a> { if !admin_provider_pool_cache_affinity_enabled(&pool_config) { return None; } - let runtime = read_admin_provider_pool_runtime_state( + let sticky_key_id = read_provider_pool_sticky_bound_key_id( self.state.app().runtime_state.as_ref(), self.group.candidate.provider_id.as_str(), - &[], &pool_config, self.sticky_session_token.as_deref(), ) - .await; - let sticky_key_id = runtime.sticky_bound_key_id?; + .await?; if self .routing_overlay .as_ref() @@ -2028,7 +2028,9 @@ mod tests { }; use crate::data::GatewayDataState; use crate::handlers::shared::provider_pool::{ - admin_provider_pool_cache_affinity_enabled, record_admin_provider_pool_error, + admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value, + read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state, + record_admin_provider_pool_error, record_admin_provider_pool_success, AdminProviderPoolRuntimeState, }; use crate::orchestration::LocalExecutionCandidateMetadata; @@ -2060,6 +2062,111 @@ mod tests { use std::collections::{BTreeMap, BTreeSet, VecDeque}; use std::sync::Arc; + #[tokio::test] + async fn scheduling_runtime_preserves_pool_ranking_and_cost_rejections() { + let runtime = aether_runtime_state::RuntimeState::memory( + aether_runtime_state::MemoryRuntimeStateConfig::default(), + ); + let writer_config = admin_provider_pool_config_from_config_value(Some(&json!({ + "pool_advanced": { + "cost_limit_per_key_tokens": 100, + "scheduling_presets": [ + {"preset": "cache_affinity", "enabled": true}, + {"preset": "latency_first", "enabled": true} + ] + } + }))) + .expect("writer pool config"); + for (key_id, cost, latency) in [("key-a", 100, 10), ("key-b", 20, 100)] { + record_admin_provider_pool_success( + &runtime, + "provider-pool", + key_id, + &writer_config, + Some(key_id), + cost, + Some(latency), + ) + .await; + } + let key_ids = vec!["key-a".to_string(), "key-b".to_string()]; + for (preset, cost_limit) in [ + ("cache_affinity", None), + ("priority_first", None), + ("latency_first", None), + ("cost_first", None), + ("quota_balanced", None), + ("latency_first", Some(100)), + ] { + let provider_config = json!({ + "pool_advanced": { + "cost_limit_per_key_tokens": cost_limit, + "scheduling_presets": [{"preset": preset, "enabled": true}] + } + }); + let pool_config = admin_provider_pool_config_from_config_value(Some(&provider_config)) + .expect("reader pool config"); + let admin = read_admin_provider_pool_runtime_state( + &runtime, + "provider-pool", + &key_ids, + &pool_config, + Some("key-a"), + ) + .await; + let scheduling = read_provider_pool_scheduling_runtime_state( + &runtime, + "provider-pool", + &key_ids, + &pool_config, + Some("key-a"), + ) + .await; + let run = |snapshot| { + let candidates = key_ids + .iter() + .map(|key_id| { + sample_eligible_candidate( + "provider-pool", + "endpoint-1", + key_id, + 10, + Some(provider_config.clone()), + ) + }) + .collect(); + let (scheduled, skipped) = apply_local_execution_pool_scheduler_with_runtime_map( + candidates, + &BTreeMap::from([("provider-pool".to_string(), snapshot)]), + &BTreeMap::new(), + ); + ( + scheduled + .into_iter() + .map(|item| item.candidate.key_id) + .collect::>(), + skipped + .into_iter() + .map(|item| (item.candidate.key_id, item.skip_reason)) + .collect::>(), + ) + }; + let expected = run(admin); + let actual = run(scheduling); + assert_eq!( + actual, expected, + "preset: {preset}, cost limit: {cost_limit:?}" + ); + if cost_limit.is_some() { + assert_eq!(actual.0, vec!["key-b"]); + assert_eq!( + actual.1, + vec![("key-a".to_string(), "pool_cost_limit_reached")] + ); + } + } + } + #[test] fn pool_scheduler_groups_interleaved_candidates_and_reorders_internal_keys() { let pool_first = sample_eligible_candidate( diff --git a/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs b/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs index 2edf0c4cf..0a1d388af 100644 --- a/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs +++ b/apps/aether-gateway/src/execution_runtime/attempt_cancellation.rs @@ -234,7 +234,9 @@ impl Drop for AttemptCancellationGuard { ); return; }; + let usage_producer = state.usage_runtime.track_producer(); handle.spawn(async move { + let _usage_producer = usage_producer; settle_cancelled_attempt(state, armed, error_type, error_message).await; }); } diff --git a/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs index 46ff852bd..6c9d80623 100644 --- a/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs +++ b/apps/aether-gateway/src/execution_runtime/attempt_lifecycle.rs @@ -674,8 +674,10 @@ impl ExecutionAttemptLifecycle { let billing_void = settlement.billing.is_void(); let usage_runtime = Arc::clone(&state.usage_runtime); let usage_data = Arc::clone(state.usage_lifecycle_data_state()); + let usage_producer = usage_runtime.track_producer(); self.stage_guard .await_detachable_stage(self.trace_id.as_str(), "usage_terminal", async move { + let _usage_producer = usage_producer; usage_runtime .record_stream_terminal( usage_data.as_ref(), diff --git a/apps/aether-gateway/src/execution_runtime/mod.rs b/apps/aether-gateway/src/execution_runtime/mod.rs index b9e48c3a8..afe70b7a5 100644 --- a/apps/aether-gateway/src/execution_runtime/mod.rs +++ b/apps/aether-gateway/src/execution_runtime/mod.rs @@ -20,6 +20,7 @@ mod response_header_rules; mod server; pub(crate) mod stream; mod stream_pump; +mod stream_read_timeout; pub(crate) mod submission; pub(crate) mod sync; pub(crate) mod transport; diff --git a/apps/aether-gateway/src/execution_runtime/stream/capture_budget.rs b/apps/aether-gateway/src/execution_runtime/stream/capture_budget.rs new file mode 100644 index 000000000..1d66fb563 --- /dev/null +++ b/apps/aether-gateway/src/execution_runtime/stream/capture_budget.rs @@ -0,0 +1,251 @@ +use std::ops::Deref; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, LazyLock}; + +const DEFAULT_STREAM_CAPTURE_MEMORY_BUDGET_BYTES: usize = 128 * 1024 * 1024; +const STREAM_CAPTURE_MEMORY_BUDGET_ENV: &str = "AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES"; + +static STREAM_CAPTURE_BUDGET: LazyLock> = LazyLock::new(|| { + StreamCaptureBudget::new( + std::env::var(STREAM_CAPTURE_MEMORY_BUDGET_ENV) + .ok() + .and_then(|value| value.trim().parse().ok()) + .unwrap_or(DEFAULT_STREAM_CAPTURE_MEMORY_BUDGET_BYTES), + ) +}); + +#[derive(Debug)] +pub(super) struct StreamCaptureBudget { + available: AtomicUsize, +} + +impl StreamCaptureBudget { + pub(super) fn new(bytes: usize) -> Arc { + Arc::new(Self { + available: AtomicUsize::new(bytes), + }) + } + + fn reserve_up_to(&self, wanted: usize, minimum: usize) -> usize { + let mut available = self.available.load(Ordering::Relaxed); + loop { + let reserved = wanted.min(available); + if reserved < minimum { + return 0; + } + match self.available.compare_exchange_weak( + available, + available - reserved, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => return reserved, + Err(current) => available = current, + } + } + } + + fn release(&self, bytes: usize) { + self.available.fetch_add(bytes, Ordering::Relaxed); + } +} + +/// Only retained diagnostic bytes belong here. Protocol and billing observers +/// must consume the original chunks independently of capture admission. +#[derive(Debug)] +pub(super) struct StreamBodyCapture { + bytes: Vec, + budget: Arc, +} + +impl Default for StreamBodyCapture { + fn default() -> Self { + Self::with_budget(Arc::clone(&STREAM_CAPTURE_BUDGET)) + } +} + +impl StreamBodyCapture { + pub(super) fn with_budget(budget: Arc) -> Self { + Self { + bytes: Vec::new(), + budget, + } + } + + pub(super) fn append(&mut self, chunk: &[u8], limit: usize, truncated: &mut bool) { + if chunk.is_empty() || *truncated { + return; + } + let wanted_len = self.bytes.len().saturating_add(chunk.len()).min(limit); + if wanted_len > self.bytes.capacity() { + // Keep the old allocation charged until its replacement has been + // allocated and copied, including their overlap during growth. + let wanted_capacity = wanted_len + .max(self.bytes.capacity().saturating_mul(2)) + .min(limit); + let reserved = self + .budget + .reserve_up_to(wanted_capacity, self.bytes.capacity().saturating_add(1)); + if reserved > 0 { + let mut replacement = Vec::new(); + if replacement.try_reserve_exact(reserved).is_ok() { + let extra = replacement.capacity().saturating_sub(reserved); + if extra == 0 || self.budget.reserve_up_to(extra, extra) == extra { + replacement.extend_from_slice(&self.bytes); + let old = std::mem::replace(&mut self.bytes, replacement); + let old_capacity = old.capacity(); + drop(old); + self.budget.release(old_capacity); + } else { + drop(replacement); + self.budget.release(reserved); + } + } else { + self.budget.release(reserved); + } + } + } + let keep = wanted_len + .min(self.bytes.capacity()) + .saturating_sub(self.bytes.len()); + self.bytes.extend_from_slice(&chunk[..keep]); + // Once bytes are omitted, never append a later suffix to this prefix. + *truncated = keep < chunk.len(); + } +} + +impl Deref for StreamBodyCapture { + type Target = [u8]; + + fn deref(&self) -> &Self::Target { + &self.bytes + } +} + +impl Drop for StreamBodyCapture { + fn drop(&mut self) { + let bytes = std::mem::take(&mut self.bytes); + let capacity = bytes.capacity(); + drop(bytes); + self.budget.release(capacity); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn stream_capture_budget_is_shared_and_released_on_drop() { + let budget = StreamCaptureBudget::new(12); + let mut provider = StreamBodyCapture::with_budget(Arc::clone(&budget)); + let mut client = StreamBodyCapture::with_budget(Arc::clone(&budget)); + let mut provider_truncated = false; + let mut client_truncated = false; + provider.append(b"12345678", 64, &mut provider_truncated); + client.append(b"abcdefgh", 64, &mut client_truncated); + assert_eq!(&*provider, b"12345678"); + assert_eq!(&*client, b"abcd"); + assert!(!provider_truncated); + assert!(client_truncated); + assert_eq!(budget.available.load(Ordering::Relaxed), 0); + drop(provider); + client.append(b"later", 64, &mut client_truncated); + assert_eq!(&*client, b"abcd"); + drop(client); + assert_eq!(budget.available.load(Ordering::Relaxed), 12); + } + + #[test] + fn stream_capture_budget_charges_capacity_and_reallocation_overlap() { + let budget = StreamCaptureBudget::new(16); + let mut capture = StreamBodyCapture::with_budget(Arc::clone(&budget)); + let mut truncated = false; + capture.append(b"1234", 64, &mut truncated); + capture.append(b"5", 64, &mut truncated); + assert_eq!(capture.bytes.capacity(), 8); + assert_eq!(budget.available.load(Ordering::Relaxed), 8); + capture.append(b"6789", 64, &mut truncated); + assert_eq!(&*capture, b"12345678"); + assert!(truncated); + assert_eq!(budget.available.load(Ordering::Relaxed), 8); + drop(capture); + assert_eq!(budget.available.load(Ordering::Relaxed), 16); + } + + #[test] + fn stream_capture_budget_zero_disables_capture_without_allocating() { + let budget = StreamCaptureBudget::new(0); + let mut capture = StreamBodyCapture::with_budget(budget); + let mut truncated = false; + capture.append(b"data", 64, &mut truncated); + assert!(capture.is_empty()); + assert_eq!(capture.bytes.capacity(), 0); + assert!(truncated); + } + + #[test] + fn stream_capture_budget_exhaustion_uses_existing_spare_capacity() { + let budget = StreamCaptureBudget::new(14); + let mut capture = StreamBodyCapture::with_budget(budget); + let mut truncated = false; + capture.append(b"1234", 64, &mut truncated); + capture.append(b"5", 64, &mut truncated); + assert_eq!(capture.bytes.capacity(), 8); + capture.append(b"6789", 64, &mut truncated); + assert_eq!(&*capture, b"12345678"); + assert_eq!(capture.bytes.capacity(), 8); + assert!(truncated); + } + + #[test] + fn stream_capture_budget_local_limit_keeps_a_contiguous_prefix() { + let budget = StreamCaptureBudget::new(128); + let mut capture = StreamBodyCapture::with_budget(Arc::clone(&budget)); + let mut truncated = false; + capture.append(b"abcdef", 3, &mut truncated); + assert_eq!(&*capture, b"abc"); + assert!(truncated); + assert_eq!(budget.available.load(Ordering::Relaxed), 125); + } + + #[test] + fn stream_capture_budget_concurrent_growth_and_drop_never_exceeds_capacity() { + const LIMIT: usize = 256; + const THREADS: usize = 8; + let budget = StreamCaptureBudget::new(LIMIT); + let barrier = std::sync::Barrier::new(THREADS); + let held = AtomicUsize::new(0); + std::thread::scope(|scope| { + for index in 0..THREADS { + let budget = &budget; + let barrier = &barrier; + let held = &held; + scope.spawn(move || { + for _ in 0..32 { + let mut capture = StreamBodyCapture::with_budget(Arc::clone(budget)); + let mut truncated = false; + barrier.wait(); + capture.append(&[1; 16], LIMIT, &mut truncated); + capture.append(&[2; 48], LIMIT, &mut truncated); + held.fetch_add(capture.bytes.capacity(), Ordering::SeqCst); + barrier.wait(); + if index == 0 { + let retained = held.load(Ordering::SeqCst); + assert!(retained <= LIMIT); + assert_eq!(retained + budget.available.load(Ordering::Relaxed), LIMIT,); + } + barrier.wait(); + drop(capture); + barrier.wait(); + if index == 0 { + assert_eq!(budget.available.load(Ordering::Relaxed), LIMIT); + held.store(0, Ordering::SeqCst); + } + barrier.wait(); + } + }); + } + }); + } +} diff --git a/apps/aether-gateway/src/execution_runtime/stream/execution.rs b/apps/aether-gateway/src/execution_runtime/stream/execution.rs index 41dc8944f..d9d30a4c6 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/execution.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/execution.rs @@ -42,6 +42,7 @@ use tokio::time::MissedTickBehavior; use tokio_util::codec::{FramedRead, LinesCodec}; use tracing::{debug, info, warn}; +use super::capture_budget::StreamBodyCapture; use super::commit_policy::{ anthropic_error_status_code, find_sse_record_boundary, StreamCommitGate, StreamCommitPolicy, StreamPrecommitObservation, @@ -53,6 +54,7 @@ use super::error::{ stream_client_error_status_code_for_upstream_status, synthetic_error_response_headers, StreamPrefetchInspection, }; +use super::usage_fallback::StreamUsageFallback; #[path = "execution_failures.rs"] mod execution_failures; use self::execution_failures::{ @@ -95,17 +97,20 @@ use crate::execution_runtime::kiro_web_search::maybe_execute_kiro_web_search_str use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry; #[cfg(test)] use crate::execution_runtime::remote_compat::post_stream_plan_to_remote_execution_runtime; +use crate::execution_runtime::stream_read_timeout::{ + await_stream_idle_read, resolve_stream_idle_timeout, stream_idle_timeout_message, +}; use crate::execution_runtime::submission::{ resolve_core_error_background_report_kind, resolve_local_sync_error_status_code, strip_utf8_bom_and_ws, submit_local_core_error_or_sync_finalize, }; use crate::execution_runtime::transport::{ - decode_base64_body_with_limit, execute_stream_plan_via_local_tunnel, format_hyper_error_chain, - format_upstream_request_error, format_wreq_upstream_request_error, - record_manual_proxy_request_failure, record_manual_proxy_request_success, - record_manual_proxy_stream_error, stream_first_byte_timeout_message, - DirectSyncExecutionRuntime, DirectUpstreamResponse, DirectUpstreamStreamExecution, - ExecutionRuntimeTransportError, + decode_base64_body_with_limit, direct_upstream_response_byte_stream, + execute_stream_plan_via_local_tunnel, format_hyper_error_chain, format_upstream_request_error, + format_wreq_upstream_request_error, record_manual_proxy_request_failure, + record_manual_proxy_request_success, record_manual_proxy_stream_error, + stream_first_byte_timeout_message, DirectSyncExecutionRuntime, DirectUpstreamResponse, + DirectUpstreamStreamExecution, ExecutionRuntimeTransportError, }; use crate::execution_runtime::windsurf::maybe_execute_windsurf_stream; use crate::execution_runtime::{ @@ -481,7 +486,9 @@ async fn record_sync_terminal_usage_with_handoff_after_spawn( let (context_seed, payload_seed) = build_sync_terminal_usage_seeds(plan, report_context, payload); let state = state.clone(); + let usage_producer = state.usage_runtime.track_producer(); let task = tokio::spawn(async move { + let _usage_producer = usage_producer; before_dispatch.await; state .usage_runtime @@ -1040,12 +1047,39 @@ fn append_stream_capture_bytes( } } +fn append_budgeted_stream_capture_bytes( + buffer: &mut StreamBodyCapture, + chunk: &[u8], + max_bytes: usize, + truncated: &mut bool, +) { + buffer.append(chunk, max_bytes, truncated); +} + +struct StreamUsageObservationBuffer { + line: Vec, + fallback: StreamUsageFallback, + recovered_usage_after_parser_error: bool, +} + +impl StreamUsageObservationBuffer { + fn new(record_limit: usize) -> Self { + Self { + line: Vec::new(), + fallback: StreamUsageFallback::new(record_limit), + recovered_usage_after_parser_error: false, + } + } +} + fn observe_stream_usage_bytes( observer: &mut StreamingStandardTerminalObserver, report_context: &Value, - buffered: &mut Vec, + buffer: &mut StreamUsageObservationBuffer, chunk: &[u8], ) { + buffer.fallback.observe(report_context, chunk); + let buffered = &mut buffer.line; if chunk.is_empty() || observer .latest_summary() @@ -1065,7 +1099,7 @@ fn observe_stream_usage_bytes( observer.disable_with_error(format!( "stream usage event exceeded {SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES} bytes" )); - buffered.clear(); + *buffered = Vec::new(); return; } buffered.extend_from_slice(&remaining[..line_part_len]); @@ -1074,7 +1108,7 @@ fn observe_stream_usage_bytes( let line = std::mem::take(buffered); if let Err(_err) = observer.push_line(report_context, line) { observer.disable_with_error("stream usage parsing failed"); - buffered.clear(); + *buffered = Vec::new(); return; } } @@ -1084,12 +1118,13 @@ fn observe_stream_usage_bytes( fn finalize_stream_usage_observer( observer: &mut Option, report_context: Option<&Value>, - buffered: &mut Vec, + buffer: &mut StreamUsageObservationBuffer, ) -> Option { let (Some(observer), Some(report_context)) = (observer.as_mut(), report_context) else { return None; }; + let buffered = &mut buffer.line; if !buffered.is_empty() { let line = std::mem::take(buffered); if let Err(_err) = observer.push_line(report_context, line) { @@ -1097,13 +1132,53 @@ fn finalize_stream_usage_observer( } } - match observer.finish(report_context) { + let mut summary = match observer.finish(report_context) { Ok(summary) => summary, Err(_err) => { observer.disable_with_error("stream usage parsing failed"); observer.latest_summary().cloned() } + }; + let mut fallback_usage = buffer.fallback.finish(report_context); + let fallback_tier = buffer.fallback.take_service_tier(); + if let Some(summary) = summary.as_mut() { + if summary.parser_error.is_some() && fallback_usage.is_some() { + // A disabled parser can retain an earlier usage snapshot. Later + // complete fallback events remain authoritative even when their + // signal score is unchanged or an explicit zero reduces it. + summary.standardized_usage = fallback_usage.take(); + buffer.recovered_usage_after_parser_error = true; + } + if summary.provider_actual_service_tier.is_none() { + summary.provider_actual_service_tier = fallback_tier.clone(); + } } + let fallback = fallback_usage.map(|usage| ExecutionStreamTerminalSummary { + standardized_usage: Some(usage), + provider_actual_service_tier: summary.is_none().then_some(fallback_tier).flatten(), + ..ExecutionStreamTerminalSummary::default() + }); + merge_stream_terminal_summary(summary, fallback) +} + +fn merge_observed_stream_terminal_summary( + current: Option, + observed: Option, + usage_buffer: &StreamUsageObservationBuffer, +) -> Option { + let recovered_usage = usage_buffer + .recovered_usage_after_parser_error + .then(|| { + observed + .as_ref() + .and_then(|summary| summary.standardized_usage.clone()) + }) + .flatten(); + let mut summary = merge_stream_terminal_summary(current, observed); + if let (Some(summary), Some(usage)) = (summary.as_mut(), recovered_usage) { + summary.standardized_usage = Some(usage); + } + summary } fn merge_stream_terminal_summary( @@ -1694,43 +1769,6 @@ fn should_use_direct_sse_passthrough( type DirectUpstreamByteStream = BoxStream<'static, Result>; -fn direct_upstream_response_byte_stream( - prefetched_body: VecDeque>, - response: DirectUpstreamResponse, -) -> DirectUpstreamByteStream { - let response_stream = match response { - DirectUpstreamResponse::Reqwest(response) => response - .bytes_stream() - .map(|item| item.map_err(|err| format_upstream_request_error(&err))) - .boxed(), - DirectUpstreamResponse::HyperH2c(response) => response - .into_body() - .into_data_stream() - .map(|item| item.map_err(|err| format_hyper_error_chain(&err))) - .boxed(), - DirectUpstreamResponse::BrowserWreq(response) => response - .bytes_stream() - .map(|item| item.map_err(|err| format_wreq_upstream_request_error(&err))) - .boxed(), - DirectUpstreamResponse::LocalTunnel(mut response) => stream! { - loop { - match response.next_chunk().await { - Ok(Some(chunk)) => yield Ok(chunk), - Ok(None) => break, - Err(err) => { - yield Err(err); - break; - } - } - } - } - .boxed(), - }; - futures_stream::iter(prefetched_body) - .chain(response_stream) - .boxed() -} - async fn await_direct_passthrough_first_item( future: F, started_at: Instant, @@ -1762,7 +1800,7 @@ async fn forward_direct_passthrough_client_chunk( client_stream_completion_tracker: &mut ClientVisibleStreamCompletionTracker, observe_stream_completion: bool, client_stream_bytes: &mut u64, - buffered_body: &mut Vec, + buffered_body: &mut StreamBodyCapture, client_body_truncated: &mut bool, max_stream_body_buffer_bytes: usize, stream_started_at: Instant, @@ -1775,7 +1813,7 @@ async fn forward_direct_passthrough_client_chunk( if chunk.is_empty() { return false; } - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( buffered_body, chunk.as_ref(), max_stream_body_buffer_bytes, @@ -1840,11 +1878,11 @@ struct DirectPassthroughFinalizerCore { headers: BTreeMap, stream_usage_report_context: Option, stream_usage_observer: Option, - stream_usage_observer_buffered: Vec, + stream_usage_observer_buffered: StreamUsageObservationBuffer, provider_error_inspection: ProviderStreamErrorInspection, max_stream_body_buffer_bytes: usize, - provider_buffered_body: Vec, - buffered_body: Vec, + provider_buffered_body: StreamBodyCapture, + buffered_body: StreamBodyCapture, provider_body_truncated: bool, client_body_truncated: bool, client_stream_completion_tracker: ClientVisibleStreamCompletionTracker, @@ -1990,7 +2028,7 @@ impl DirectPassthroughFinalizer { core.provider_stream_bytes = core .provider_stream_bytes .saturating_add(u64::try_from(chunk.len()).unwrap_or(u64::MAX)); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut core.provider_buffered_body, chunk.as_ref(), core.max_stream_body_buffer_bytes, @@ -2028,7 +2066,7 @@ impl DirectPassthroughFinalizer { return; } let core = self.core_mut(); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut core.buffered_body, chunk.as_ref(), core.max_stream_body_buffer_bytes, @@ -2096,7 +2134,9 @@ impl DirectPassthroughFinalizer { // client disconnect or an execution timeout may cancel this body // future while terminal admission is backpressured; the handoff must // continue independently so the usage row cannot remain streaming. + let usage_producer = core.state.usage_runtime.track_producer(); let task = tokio::spawn(async move { + let _usage_producer = usage_producer; core.finalize(downstream_dropped).await; }); if let Err(_err) = task.await { @@ -2117,7 +2157,9 @@ impl Drop for DirectPassthroughFinalizer { }; observe_gateway_stage_ms("stream_finalizer_enqueue", 0); if let Ok(handle) = tokio::runtime::Handle::try_current() { + let usage_producer = core.state.usage_runtime.track_producer(); handle.spawn(async move { + let _usage_producer = usage_producer; core.finalize(true).await; }); } @@ -2543,6 +2585,7 @@ struct DirectPassthroughInlineBodyState { upstream_control_filter: Option, upstream_started_at: Instant, stream_first_byte_timeout: Option, + stream_idle_timeout: Option, observed_first_body_poll: bool, observed_first_client_yield: bool, upstream_done: bool, @@ -2559,6 +2602,7 @@ impl DirectPassthroughInlineBodyState { upstream_started_at: Instant, stream_first_byte_timeout: Option, ) -> Self { + let stream_idle_timeout = resolve_stream_idle_timeout(&finalizer.core().plan); Self { finalizer: Some(finalizer), upstream: Some(direct_upstream_response_byte_stream( @@ -2568,6 +2612,7 @@ impl DirectPassthroughInlineBodyState { upstream_control_filter: Some(SseControlBlockFilter::default()), upstream_started_at, stream_first_byte_timeout, + stream_idle_timeout, observed_first_body_poll: false, observed_first_client_yield: false, upstream_done: false, @@ -2723,7 +2768,28 @@ impl DirectPassthroughInlineBodyState { } } } else { - upstream.next().await + match await_stream_idle_read(upstream.next(), self.stream_idle_timeout).await { + Ok(item) => item, + Err(timeout) => { + self.upstream.take(); + if let Some(finalizer) = self.finalizer.as_mut() { + if finalizer.terminal_failure().is_none() + && !finalizer + .core() + .client_stream_completion_tracker + .successful_completion() + { + finalizer.set_terminal_failure(build_stream_transport_failure_report( + "read_timeout", + stream_idle_timeout_message(timeout), + 504, + )); + } + finalizer.core_mut()._provider_pool_in_flight_guard.take(); + } + None + } + } } } @@ -2800,7 +2866,12 @@ impl Drop for DirectPassthroughInlineBodyState { if let Some(finalizer) = self.finalizer.take() { observe_gateway_stage_ms("stream_finalizer_enqueue", 0); if let Ok(handle) = tokio::runtime::Handle::try_current() { + let usage_producer = finalizer + .core + .as_ref() + .map(|core| core.state.usage_runtime.track_producer()); handle.spawn(async move { + let _usage_producer = usage_producer; let mut finalizer = finalizer; finalizer.finalize(true).await; }); @@ -2901,6 +2972,7 @@ async fn execute_stream_from_direct_passthrough( started_at: upstream_started_at, response_observation, stream_first_byte_timeout, + stream_idle_timeout, upstream_target_permit, } = execution; @@ -3036,11 +3108,13 @@ async fn execute_stream_from_direct_passthrough( headers: headers_for_report, stream_usage_report_context, stream_usage_observer, - stream_usage_observer_buffered: Vec::new(), + stream_usage_observer_buffered: StreamUsageObservationBuffer::new( + max_stream_body_buffer_bytes, + ), provider_error_inspection: ProviderStreamErrorInspection::default(), max_stream_body_buffer_bytes, - provider_buffered_body: Vec::new(), - buffered_body: Vec::new(), + provider_buffered_body: StreamBodyCapture::default(), + buffered_body: StreamBodyCapture::default(), provider_body_truncated: false, client_body_truncated: false, client_stream_completion_tracker: ClientVisibleStreamCompletionTracker::default(), @@ -3102,7 +3176,9 @@ async fn execute_stream_from_direct_passthrough( let candidate_id_for_report = candidate_id.clone(); let provider_pool_in_flight_guard_for_report = in_flight_guard; record_stream_pre_first_byte_spawn(); + let usage_producer = state_for_report.usage_runtime.track_producer(); tokio::spawn(async move { + let _usage_producer = usage_producer; let mut stage_trace_for_report = stage_trace_for_report; let _stream_total_guard = StageElapsedGuard::from_started_at("stream_total", stream_started_at_for_report); @@ -3118,10 +3194,11 @@ async fn execute_stream_from_direct_passthrough( let mut stream_usage_observer = stream_usage_report_context .as_ref() .map(|_| StreamingStandardTerminalObserver::default()); - let mut stream_usage_observer_buffered = Vec::new(); + let mut stream_usage_observer_buffered = + StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes); let mut provider_error_inspection = ProviderStreamErrorInspection::default(); - let mut provider_buffered_body = Vec::new(); - let mut buffered_body = Vec::new(); + let mut provider_buffered_body = StreamBodyCapture::default(); + let mut buffered_body = StreamBodyCapture::default(); let mut provider_body_truncated = false; let mut client_body_truncated = false; let mut upstream_control_filter = Some(SseControlBlockFilter::default()); @@ -3180,7 +3257,20 @@ async fn execute_stream_from_direct_passthrough( downstream_dropped = true; break; } - item = upstream.next() => item, + result = await_stream_idle_read(upstream.next(), stream_idle_timeout) => { + match result { + Ok(item) => item, + Err(timeout) => { + if terminal_failure.is_none() + && !client_stream_completion_tracker.successful_completion() { + terminal_failure = Some(build_stream_transport_failure_report( + "read_timeout", stream_idle_timeout_message(timeout), 504, + )); + } + break; + } + } + }, } }; @@ -3304,7 +3394,7 @@ async fn execute_stream_from_direct_passthrough( provider_stream_bytes = provider_stream_bytes .saturating_add(u64::try_from(provider_chunk.len()).unwrap_or(u64::MAX)); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut provider_buffered_body, provider_chunk.as_ref(), max_stream_body_buffer_bytes, @@ -5165,7 +5255,7 @@ enum SseTerminalPolicy { } #[derive(Default)] -struct ClientVisibleStreamCompletionTracker { +pub(crate) struct ClientVisibleStreamCompletionTracker { line_buffer: Vec, event_type: Option, data_payload: String, @@ -5175,14 +5265,23 @@ struct ClientVisibleStreamCompletionTracker { discarded_line_nonempty: bool, skip_next_lf: bool, completed: bool, + successfully_completed: bool, } impl ClientVisibleStreamCompletionTracker { - fn observe_chunk(&mut self, chunk: &[u8]) -> bool { + pub(crate) fn observe_chunk(&mut self, chunk: &[u8]) -> bool { self.observe_chunk_terminal_end(chunk); self.completed } + pub(crate) fn successful_completion(&self) -> bool { + self.successfully_completed + } + + pub(crate) fn observed_terminal(&self) -> bool { + self.completed + } + fn observe_chunk_terminal_end(&mut self, chunk: &[u8]) -> Option { self.observe_chunk_terminal_end_with_policy(chunk, SseTerminalPolicy::AnyKnown) } @@ -5270,6 +5369,9 @@ impl ClientVisibleStreamCompletionTracker { if line.is_empty() { self.completed = self.current_event_is_terminal(policy); + if self.completed { + self.successfully_completed = self.current_event_is_successful(); + } self.reset_current_event(); self.record_bytes = 0; return; @@ -5334,6 +5436,37 @@ impl ClientVisibleStreamCompletionTracker { self.data_payload.clear(); self.has_data_payload = false; } + + fn current_event_is_successful(&self) -> bool { + let payload = self + .has_data_payload + .then(|| serde_json::from_str::(&self.data_payload).ok()) + .flatten(); + let payload_type = payload + .as_ref() + .and_then(|value| value.get("type")) + .and_then(Value::as_str); + if [self.event_type.as_deref(), payload_type] + .into_iter() + .flatten() + .any(|kind| matches!(kind, "response.failed" | "response.incomplete" | "error")) + { + return false; + } + if payload + .as_ref() + .and_then(|value| value.pointer("/response/status")) + .and_then(Value::as_str) + .is_some_and(|status| status != "completed") + { + return false; + } + self.data_payload == "[DONE]" + || matches!( + payload_type.or(self.event_type.as_deref()), + Some("message_stop" | "response.completed") + ) + } } fn is_terminal_sse_event_type(event_type: &str) -> bool { @@ -7148,13 +7281,15 @@ async fn execute_stream_from_frame_stream_with_retry_scope( let stage_trace_for_report = stage_trace; let request_diagnostics_for_report = current_request_diagnostics(); let provider_pool_in_flight_guard_for_report = in_flight_guard; + let usage_producer = state_for_report.usage_runtime.track_producer(); tokio::spawn(async move { + let _usage_producer = usage_producer; let mut stage_trace_for_report = stage_trace_for_report; let _stream_total_guard = StageElapsedGuard::from_started_at("stream_total", stream_started_at_for_report); let _provider_pool_in_flight_guard = provider_pool_in_flight_guard_for_report; - let mut provider_buffered_body = Vec::new(); - let mut buffered_body = Vec::new(); + let mut provider_buffered_body = StreamBodyCapture::default(); + let mut buffered_body = StreamBodyCapture::default(); let mut provider_body_truncated = false; let mut client_body_truncated = false; let mut private_stream_normalizer = if sync_json_stream_bridge_active_for_report { @@ -7178,15 +7313,16 @@ async fn execute_stream_from_frame_stream_with_retry_scope( .as_ref() .filter(|_| !sync_json_stream_bridge_active_for_report) .map(|_| StreamingStandardTerminalObserver::default()); - let mut stream_usage_observer_buffered = Vec::new(); + let mut stream_usage_observer_buffered = + StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes); let mut provider_error_inspection = ProviderStreamErrorInspection::default(); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut provider_buffered_body, &provider_prefetched_body_for_report, max_stream_body_buffer_bytes, &mut provider_body_truncated, ); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut buffered_body, &prefetched_body_for_report, max_stream_body_buffer_bytes, @@ -7413,6 +7549,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } } + // These buffers restore parser/rewriter state above. Audit capture owns + // its budgeted copies; retaining semantic prefetch duplicates for the + // rest of the stream would bypass the capture memory limit. + drop(provider_prefetched_body_for_report); + drop(prefetched_body_for_report); + if terminal_failure.is_none() && !reached_eof { loop { let draining_after_anthropic_stop = anthropic_post_stop_drain_started_at.is_some(); @@ -7593,7 +7735,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( u64::try_from(chunk.len()).unwrap_or(u64::MAX), Ordering::Relaxed, ); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut provider_buffered_body, &chunk, max_stream_body_buffer_bytes, @@ -7694,7 +7836,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( continue; } - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut buffered_body, &rewritten_chunk, max_stream_body_buffer_bytes, @@ -7885,7 +8027,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } } if !rewritten_chunk.is_empty() { - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut buffered_body, &rewritten_chunk, max_stream_body_buffer_bytes, @@ -7964,7 +8106,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( } match finish_result { Ok(flushed_chunk) if !flushed_chunk.is_empty() => { - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut buffered_body, &flushed_chunk, max_stream_body_buffer_bytes, @@ -8050,7 +8192,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope( Ok(error_event) => { let error_event_len = u64::try_from(error_event.len()).unwrap_or(u64::MAX); - append_stream_capture_bytes( + append_budgeted_stream_capture_bytes( &mut buffered_body, error_event.as_ref(), max_stream_body_buffer_bytes, @@ -8098,13 +8240,15 @@ async fn execute_stream_from_frame_stream_with_retry_scope( idle_monitor_done.store(true, Ordering::Relaxed); idle_monitor_handle.abort(); - stream_terminal_summary = merge_stream_terminal_summary( + let observed_terminal_summary = finalize_stream_usage_observer( + &mut stream_usage_observer, + stream_usage_report_context.as_ref(), + &mut stream_usage_observer_buffered, + ); + stream_terminal_summary = merge_observed_stream_terminal_summary( stream_terminal_summary, - finalize_stream_usage_observer( - &mut stream_usage_observer, - stream_usage_report_context.as_ref(), - &mut stream_usage_observer_buffered, - ), + observed_terminal_summary, + &stream_usage_observer_buffered, ); if downstream_dropped && client_visible_stream_completed && terminal_failure.is_none() { @@ -9487,11 +9631,13 @@ mod tests { )]), stream_usage_report_context: None, stream_usage_observer: None, - stream_usage_observer_buffered: Vec::new(), + stream_usage_observer_buffered: super::StreamUsageObservationBuffer::new( + super::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, + ), provider_error_inspection: ProviderStreamErrorInspection::default(), max_stream_body_buffer_bytes: super::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, - provider_buffered_body: Vec::new(), - buffered_body: Vec::new(), + provider_buffered_body: super::StreamBodyCapture::default(), + buffered_body: super::StreamBodyCapture::default(), provider_body_truncated: false, client_body_truncated: false, client_stream_completion_tracker: ClientVisibleStreamCompletionTracker::default(), @@ -9527,6 +9673,7 @@ mod tests { upstream_control_filter: Some(super::SseControlBlockFilter::default()), upstream_started_at: Instant::now(), stream_first_byte_timeout: None, + stream_idle_timeout: None, observed_first_body_poll: false, observed_first_client_yield: false, upstream_done: false, @@ -9536,6 +9683,372 @@ mod tests { } } + #[tokio::test] + async fn stream_capture_budget_exhaustion_preserves_inline_bytes_and_terminal_usage() { + use super::super::capture_budget::{StreamBodyCapture, StreamCaptureBudget}; + + let chunks = [ + Bytes::from_static(b"data: {\"id\":\"x\",\"model\":\"gpt\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"hello\"},\"finish_reason\":null}]}\n\n"), + Bytes::from_static(b"data: {\"id\":\"x\",\"model\":\"gpt\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":11,\"completion_tokens\":7,\"prompt_tokens_details\":{\"cached_tokens\":3}}}\n\n"), + Bytes::from_static(b"data: [DONE]\n\n"), + ]; + for budget_bytes in [0, 64] { + let budget = StreamCaptureBudget::new(budget_bytes); + let mut state = direct_anthropic_inline_state( + "capture-budget-inline", + chunks.iter().cloned().map(Ok).collect(), + ); + let core = state.finalizer.as_mut().unwrap().core_mut(); + core.requires_anthropic_message_stop = false; + core.plan.provider_api_format = "openai:chat".to_string(); + core.plan.client_api_format = "openai:chat".to_string(); + core.stream_usage_report_context = Some(json!({ + "provider_api_format": "openai:chat", "client_api_format": "openai:chat" + })); + core.stream_usage_observer = Some(super::StreamingStandardTerminalObserver::default()); + core.provider_buffered_body = StreamBodyCapture::with_budget(Arc::clone(&budget)); + core.buffered_body = StreamBodyCapture::with_budget(budget); + for expected in &chunks { + let (actual, next) = state.next_item().await.expect("streamed chunk"); + assert_eq!(actual.unwrap(), *expected); + state = next; + } + let core = state.finalizer.as_mut().unwrap().core_mut(); + assert!(core.terminal_failure.is_none()); + assert!(core.client_visible_stream_completed); + assert!(core.provider_body_truncated); + assert!(core.client_body_truncated); + assert!(core.provider_buffered_body.len() + core.buffered_body.len() <= budget_bytes); + let summary = super::finalize_stream_usage_observer( + &mut core.stream_usage_observer, + core.stream_usage_report_context.as_ref(), + &mut core.stream_usage_observer_buffered, + ) + .unwrap(); + assert!(summary.observed_finish); + assert!(summary.parser_error.is_none()); + let payload = super::build_stream_usage_payload( + "capture-budget-inline".to_string(), + "openai_chat_stream".to_string(), + core.stream_usage_report_context.clone(), + 200, + BTreeMap::new(), + &core.provider_buffered_body, + core.provider_body_truncated, + &core.buffered_body, + core.client_body_truncated, + Some(summary), + None, + ); + let seed = aether_usage_runtime::build_stream_terminal_usage_payload_seed(&payload); + let usage = seed.standardized_usage.unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + assert_eq!(usage.cache_read_tokens, 3); + assert_eq!( + payload.provider_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + assert_eq!( + payload.client_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + discard_direct_test_finalizer(&mut state); + } + } + + #[test] + fn stream_capture_fallback_after_disabled_observer_updates_tokens_zero_cache_and_tier() { + let context = json!({"provider_api_format": "openai:chat"}); + let mut observer = Some(super::StreamingStandardTerminalObserver::default()); + let mut buffer = + super::StreamUsageObservationBuffer::new(super::BASIC_STREAM_BODY_ANALYSIS_LIMIT_BYTES); + super::observe_stream_usage_bytes(observer.as_mut().unwrap(), &context, &mut buffer, + b"data: {\"id\":\"x\",\"model\":\"gpt\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":10,\"prompt_tokens_details\":{\"cached_tokens\":30}}}\n\n"); + let oversized = format!( + "data: {{\"content\":\"{}\"}}\n\n", + "x".repeat(super::SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES) + ); + for part in oversized.as_bytes().chunks(4096) { + super::observe_stream_usage_bytes( + observer.as_mut().unwrap(), + &context, + &mut buffer, + part, + ); + } + super::observe_stream_usage_bytes(observer.as_mut().unwrap(), &context, &mut buffer, + b"data: {\"\\u0075sage\":{\"prompt_tokens\":100,\"completion_tokens\":500,\"prompt_tokens_details\":{\"cached_tokens\":0}},\"service_tier\":\"priority\"}\n\n"); + let summary = + super::finalize_stream_usage_observer(&mut observer, Some(&context), &mut buffer) + .unwrap(); + assert!(summary + .parser_error + .as_deref() + .unwrap() + .contains("exceeded")); + assert!(summary.observed_finish); + assert_eq!( + summary.provider_actual_service_tier.as_deref(), + Some("priority") + ); + let usage = summary.standardized_usage.as_ref().unwrap(); + assert_eq!(usage.input_tokens, 100); + assert_eq!(usage.output_tokens, 500); + assert_eq!(usage.cache_read_tokens, 0); + + let eof_summary = ExecutionStreamTerminalSummary { + standardized_usage: Some(StandardizedUsage { + input_tokens: 100, + output_tokens: 10, + cache_read_tokens: 30, + ..StandardizedUsage::new() + }), + response_id: Some("authoritative-eof-id".to_string()), + finish_reason: Some("stop".to_string()), + observed_finish: true, + ..ExecutionStreamTerminalSummary::default() + }; + let merged = super::merge_observed_stream_terminal_summary( + Some(eof_summary), + Some(summary), + &buffer, + ) + .unwrap(); + assert_eq!(merged.response_id.as_deref(), Some("authoritative-eof-id")); + assert_eq!(merged.finish_reason.as_deref(), Some("stop")); + assert!(merged.observed_finish); + assert!(merged.parser_error.as_deref().unwrap().contains("exceeded")); + let usage = merged.standardized_usage.unwrap(); + assert_eq!(usage.input_tokens, 100); + assert_eq!(usage.output_tokens, 500); + assert_eq!(usage.cache_read_tokens, 0); + } + + #[test] + fn stream_capture_budget_zero_preserves_conversion_bytes_and_usage() { + use super::super::capture_budget::{StreamBodyCapture, StreamCaptureBudget}; + + let context = json!({ + "provider_api_format": "openai:chat", + "client_api_format": "claude:messages", + "needs_conversion": true, + }); + let chunks: [&[u8]; 3] = [ + b"data: {\"id\":\"x\",\"model\":\"gpt\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hello\"},\"finish_reason\":null}]}\n\n", + b"data: {\"id\":\"x\",\"model\":\"gpt\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":11,\"completion_tokens\":7,\"prompt_tokens_details\":{\"cached_tokens\":3}}}\n\n", + b"data: [DONE]\n\n", + ]; + let mut expected = None; + for bytes in [32 * 1024, 0] { + let budget = StreamCaptureBudget::new(bytes); + let mut provider = StreamBodyCapture::with_budget(Arc::clone(&budget)); + let mut client = StreamBodyCapture::with_budget(budget); + let mut provider_truncated = false; + let mut client_truncated = false; + let mut observer = Some(super::StreamingStandardTerminalObserver::default()); + let mut buffer = super::StreamUsageObservationBuffer::new(32 * 1024); + let mut rewriter = super::maybe_build_stream_response_rewriter(Some(&context)).unwrap(); + let mut delivered = Vec::new(); + for chunk in chunks { + provider.append(chunk, 32 * 1024, &mut provider_truncated); + super::observe_stream_usage_bytes( + observer.as_mut().unwrap(), + &context, + &mut buffer, + chunk, + ); + let output = rewriter.push_chunk(chunk).unwrap(); + client.append(&output, 32 * 1024, &mut client_truncated); + delivered.extend(output); + } + let tail = rewriter.finish().unwrap(); + client.append(&tail, 32 * 1024, &mut client_truncated); + delivered.extend(tail); + let summary = + super::finalize_stream_usage_observer(&mut observer, Some(&context), &mut buffer) + .unwrap(); + assert!(summary.observed_finish); + assert!(summary.parser_error.is_none()); + let payload = super::build_stream_usage_payload( + "capture-budget-conversion".to_string(), + "claude_chat_stream".to_string(), + Some(context.clone()), + 200, + BTreeMap::new(), + &provider, + provider_truncated, + &client, + client_truncated, + Some(summary), + None, + ); + let seed = aether_usage_runtime::build_stream_terminal_usage_payload_seed(&payload); + let usage = seed.standardized_usage.unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + assert_eq!(usage.cache_read_tokens, 3); + assert!(String::from_utf8_lossy(&delivered).contains("message_stop")); + if let Some(expected) = &expected { + assert_eq!(&delivered, expected); + assert_eq!( + payload.provider_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + assert_eq!( + payload.client_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + } else { + expected = Some(delivered); + } + } + } + + #[test] + fn stream_capture_budget_zero_preserves_sync_json_bridge_terminal_usage_and_tier() { + use super::super::capture_budget::{StreamBodyCapture, StreamCaptureBudget}; + + let response = json!({ + "id": "chatcmpl-capture", "object": "chat.completion", "model": "gpt", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 11, "completion_tokens": 7, "total_tokens": 18, + "prompt_tokens_details": {"cached_tokens": 3}}, + "service_tier": "priority", + }); + let outcome = super::maybe_bridge_standard_sync_json_to_stream( + &response, + "openai:chat", + "openai:chat", + None, + ) + .unwrap() + .unwrap(); + let delivered = outcome.sse_body; + assert!(String::from_utf8_lossy(&delivered).contains("hello")); + assert!(String::from_utf8_lossy(&delivered).contains("[DONE]")); + let budget = StreamCaptureBudget::new(0); + let mut provider = StreamBodyCapture::with_budget(Arc::clone(&budget)); + let mut client = StreamBodyCapture::with_budget(budget); + let mut provider_truncated = false; + let mut client_truncated = false; + provider.append( + &serde_json::to_vec(&response).unwrap(), + 32 * 1024, + &mut provider_truncated, + ); + client.append(&delivered, 32 * 1024, &mut client_truncated); + let summary = outcome.terminal_summary.unwrap(); + assert!(summary.observed_finish); + assert!(summary.parser_error.is_none()); + assert_eq!( + summary.provider_actual_service_tier.as_deref(), + Some("priority") + ); + let payload = super::build_stream_usage_payload( + "capture-budget-sync-bridge".to_string(), + "openai_chat_stream".to_string(), + None, + 200, + BTreeMap::new(), + &provider, + provider_truncated, + &client, + client_truncated, + Some(summary), + None, + ); + assert!(payload.provider_body_base64.is_none()); + assert!(payload.client_body_base64.is_none()); + assert_eq!( + payload.provider_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + let seed = aether_usage_runtime::build_stream_terminal_usage_payload_seed(&payload); + let usage = seed.standardized_usage.unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + assert_eq!(usage.cache_read_tokens, 3); + assert_eq!( + seed.provider_actual_service_tier.as_deref(), + Some("priority") + ); + } + + #[tokio::test] + async fn direct_inline_idle_timeout_after_first_chunk_emits_terminal_read_timeout() { + let message_start = Bytes::from_static( + b"event: message_start\ndata: {\"type\":\"message_start\",\"message\":{}}\n\n", + ); + let mut state = direct_anthropic_inline_state("req-inline-idle-timeout", Vec::new()); + state.stream_idle_timeout = Some(Duration::from_millis(5)); + state.upstream = Some( + futures_util::stream::iter(vec![Ok(message_start.clone())]) + .chain(futures_util::stream::pending()) + .boxed(), + ); + let (first, state) = state.next_item().await.expect("first chunk should stream"); + assert_eq!(first.expect("first chunk"), message_start); + let (error, mut state) = tokio::time::timeout(Duration::from_secs(1), state.next_item()) + .await + .expect("idle timeout should complete") + .expect("terminal error should stream"); + let error = + String::from_utf8(error.expect("terminal error should encode").to_vec()).unwrap(); + assert!(error.starts_with("event: error\n")); + let failure = state + .finalizer + .as_ref() + .unwrap() + .terminal_failure() + .unwrap(); + assert_eq!(failure.error_type, "read_timeout"); + assert_eq!(failure.status_code, 504); + assert!( + state.upstream.is_none(), + "timeout should drop the upstream before settlement" + ); + discard_direct_test_finalizer(&mut state); + } + + #[tokio::test] + async fn direct_inline_idle_timeout_preserves_successful_protocol_completion() { + for terminal in [ + "data: [DONE]\n\n", + "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\"}}\n\n", + ] { + let mut state = direct_anthropic_inline_state("req-inline-idle-completed", Vec::new()); + state.finalizer.as_mut().unwrap().core_mut().requires_anthropic_message_stop = false; + state.stream_idle_timeout = Some(Duration::from_millis(5)); + state.upstream = Some(futures_util::stream::iter(vec![Ok(Bytes::from(terminal))]) + .chain(futures_util::stream::pending()).boxed()); + let (first, mut state) = state.next_item().await.expect("terminal chunk should stream"); + assert_eq!(first.unwrap(), Bytes::from(terminal)); + assert!(state.finalizer.as_ref().unwrap().core().client_stream_completion_tracker.successful_completion()); + let item = tokio::time::timeout(Duration::from_secs(1), state.next_upstream_item()) + .await.expect("teardown idle should finish"); + assert!(item.is_none()); + assert!(state.upstream.is_none()); + assert!(state.finalizer.as_ref().unwrap().terminal_failure().is_none(), + "successful protocol terminal must not become a read timeout"); + discard_direct_test_finalizer(&mut state); + } + } + + #[test] + fn idle_timeout_completion_tracker_does_not_treat_failure_as_success() { + for terminal in [ + "event: response.failed\ndata: {\"type\":\"response.failed\"}\n\n", + "event: response.incomplete\ndata: {\"type\":\"response.incomplete\"}\n\n", + "event: error\ndata: {\"type\":\"error\"}\n\n", + "event: response.completed\ndata: {\"type\":\"response.failed\"}\n\n", + ] { + let mut tracker = ClientVisibleStreamCompletionTracker::default(); + assert!(tracker.observe_chunk(terminal.as_bytes())); + assert!(!tracker.successful_completion()); + } + } + async fn execute_native_anthropic_prefetch_stream( request_id: &str, chunks: Vec, @@ -10607,11 +11120,13 @@ mod tests { headers: BTreeMap::new(), stream_usage_report_context: None, stream_usage_observer: None, - stream_usage_observer_buffered: Vec::new(), + stream_usage_observer_buffered: super::StreamUsageObservationBuffer::new( + super::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, + ), provider_error_inspection: ProviderStreamErrorInspection::default(), max_stream_body_buffer_bytes: super::DEFAULT_USAGE_RESPONSE_BODY_CAPTURE_LIMIT_BYTES, - provider_buffered_body: Vec::new(), - buffered_body: Vec::new(), + provider_buffered_body: super::StreamBodyCapture::default(), + buffered_body: super::StreamBodyCapture::default(), provider_body_truncated: false, client_body_truncated: false, client_stream_completion_tracker: ClientVisibleStreamCompletionTracker::default(), @@ -10639,6 +11154,7 @@ mod tests { upstream_control_filter: None, upstream_started_at: stream_started_at, stream_first_byte_timeout: None, + stream_idle_timeout: None, observed_first_body_poll: true, observed_first_client_yield: false, upstream_done: false, diff --git a/apps/aether-gateway/src/execution_runtime/stream/mod.rs b/apps/aether-gateway/src/execution_runtime/stream/mod.rs index c45389853..ded0cb59a 100644 --- a/apps/aether-gateway/src/execution_runtime/stream/mod.rs +++ b/apps/aether-gateway/src/execution_runtime/stream/mod.rs @@ -1,7 +1,10 @@ +mod capture_budget; mod commit_policy; mod error; mod execution; +mod usage_fallback; pub(crate) use execution::{ execute_execution_runtime_stream, execute_execution_runtime_stream_with_retry_scope, + ClientVisibleStreamCompletionTracker, }; diff --git a/apps/aether-gateway/src/execution_runtime/stream/usage_fallback.rs b/apps/aether-gateway/src/execution_runtime/stream/usage_fallback.rs new file mode 100644 index 000000000..8f6e3a18a --- /dev/null +++ b/apps/aether-gateway/src/execution_runtime/stream/usage_fallback.rs @@ -0,0 +1,1163 @@ +use aether_contracts::StandardizedUsage; +use aether_data_contracts::repository::usage::extract_provider_actual_service_tier_from_response; +use aether_usage_runtime::{map_usage, map_usage_from_response}; +use serde::de::{IgnoredAny, MapAccess, SeqAccess, Visitor}; +use serde::{Deserialize, Deserializer}; +use serde_json::Value; + +use super::commit_policy::find_sse_record_boundary; + +/// Retain only the current protocol record and the latest usage signal. This +/// preserves the old captured-body billing fallback when audit capture stops. +pub(super) struct StreamUsageFallback { + record: Vec, + limit: usize, + dropping_record: bool, + json_body: Option, + latest_usage: Option, + latest_service_tier: Option, + claude_usage: Value, + claude_usage_observed: bool, + #[cfg(test)] + copied_record_bytes: usize, +} + +impl StreamUsageFallback { + pub(super) fn new(limit: usize) -> Self { + Self { + record: Vec::new(), + limit, + dropping_record: false, + json_body: None, + latest_usage: None, + latest_service_tier: None, + claude_usage: Value::Object(serde_json::Map::new()), + claude_usage_observed: false, + #[cfg(test)] + copied_record_bytes: 0, + } + } + + pub(super) fn observe(&mut self, context: &Value, chunk: &[u8]) { + self.observe_inner::(context, chunk); + } + + fn observe_inner( + &mut self, + context: &Value, + chunk: &[u8], + ) { + if self.json_body.is_none() { + self.json_body = chunk + .iter() + .find(|byte| !byte.is_ascii_whitespace()) + .map(|byte| matches!(byte, b'{' | b'[')); + } + if self.json_body == Some(true) { + if self.record.len().saturating_add(chunk.len()) > self.limit { + self.record = Vec::new(); + self.dropping_record = true; + } else if !self.dropping_record { + self.append_record(chunk); + } + return; + } + let mut remaining = chunk; + while !remaining.is_empty() { + if BORROW_COMPLETE_RECORDS && self.record.is_empty() && !self.dropping_record { + if let Some((end, separator)) = find_sse_record_boundary(remaining) { + let mut consumed = end + separator; + // The buffered path processes CR and LF separately, so a + // final CR already ends the record and leaves its LF behind. + // Preserve that boundary and its contribution to the next limit. + if remaining[..consumed].ends_with(b"\r\n") { + consumed -= 1; + } + if consumed <= self.limit { + self.observe_record(context, &remaining[..consumed]); + remaining = &remaining[consumed..]; + continue; + } + } + } + // Only incomplete or oversized records need the original carry path. + let part_len = remaining + .iter() + .position(|byte| matches!(byte, b'\r' | b'\n')) + .map_or(remaining.len(), |index| index + 1); + let (part, rest) = remaining.split_at(part_len); + remaining = rest; + let scan_start = self.record.len().saturating_sub(3); + if self.record.len().saturating_add(part.len()) > self.limit { + self.dropping_record = true; + let suffix = self.record.len().saturating_sub(3); + self.record = self.record[suffix..].to_vec(); + } + if self.dropping_record { + let boundary = find_sse_record_boundary(&self.record).is_some() + || self + .record + .last() + .is_some_and(|byte| matches!(byte, b'\r' | b'\n')) + && matches!(part, b"\n" | b"\r" | b"\r\n") + && !(self.record.last() == Some(&b'\r') && part == b"\n"); + self.record = part[part.len().saturating_sub(3)..].to_vec(); + if boundary { + self.record = Vec::new(); + self.dropping_record = false; + } + continue; + } + self.append_record(part); + if find_sse_record_boundary(&self.record[scan_start..]).is_some() { + let record = std::mem::take(&mut self.record); + self.observe_record(context, &record); + } + } + } + + pub(super) fn finish(&mut self, context: &Value) -> Option { + let record = std::mem::take(&mut self.record); + if !self.dropping_record && !record.is_empty() { + self.observe_record(context, &record); + } + self.latest_usage.take() + } + + pub(super) fn take_service_tier(&mut self) -> Option { + self.latest_service_tier.take() + } + + fn append_record(&mut self, bytes: &[u8]) { + #[cfg(test)] + { + self.copied_record_bytes += bytes.len(); + } + let required = self.record.len().saturating_add(bytes.len()); + if required > self.record.capacity() { + let capacity = required + .max(self.record.capacity().saturating_mul(2)) + .min(self.limit); + self.record.reserve_exact(capacity - self.record.len()); + } + self.record.extend_from_slice(bytes); + } + + fn observe_record(&mut self, context: &Value, record: &[u8]) { + // Content-only events do not need another JSON parse. + let image_response = is_openai_image_api(provider_format(context)); + if !record.contains(&b'\\') + && !record.windows(5).any(|part| { + matches!(part, b"usage" | b"_tier" | b"speed") + || image_response && matches!(part, b"\"data" | b"resul") + }) + { + return; + } + let Ok(record) = std::str::from_utf8(record) else { + return; + }; + let mut data_lines = record + .split(['\r', '\n']) + .filter_map(|line| line.trim().strip_prefix("data:").map(str::trim)); + let Some(first) = data_lines.next() else { + self.observe_json(context, record); + return; + }; + let Some(second) = data_lines.next() else { + self.observe_json(context, first); + return; + }; + let mut payload = String::with_capacity(record.len()); + payload.push_str(first); + for line in std::iter::once(second).chain(data_lines) { + payload.push('\n'); + payload.push_str(line); + } + self.observe_json(context, &payload); + } + + fn observe_json(&mut self, context: &Value, json: &str) { + // serde ignores text/image/output fields instead of allocating another + // copy of a large completed response just to recover its usage object. + let Ok(envelope) = serde_json::from_str::(json) else { + return; + }; + let provider_format = provider_format(context); + let image_count = is_openai_image_api(provider_format) + .then(|| envelope.image_count()) + .flatten(); + let envelope = envelope.into_value(); + if let Some(tier) = extract_provider_actual_service_tier_from_response(Some(&envelope)) { + self.latest_service_tier = Some(tier); + } + // Anthropic alone sends partial usage. Keep only fields the mapper reads; + // unknown fields must not accumulate across the lifetime of the stream. + let provider_family = provider_format.split(':').next().unwrap_or_default().trim(); + if provider_family.eq_ignore_ascii_case("claude") + || provider_family.eq_ignore_ascii_case("anthropic") + { + let event_type = envelope.get("type").and_then(Value::as_str); + let raw_usage = match event_type { + Some("message_start") => { + self.claude_usage = Value::Object(serde_json::Map::new()); + self.claude_usage_observed = false; + envelope.pointer("/message/usage") + } + Some("message_delta") => envelope.get("usage"), + _ => None, + }; + if let Some(raw_usage) = raw_usage.and_then(Value::as_object) { + self.claude_usage_observed |= !raw_usage.is_empty(); + merge_claude_usage_projection(&mut self.claude_usage, raw_usage); + let usage = map_usage(&self.claude_usage, provider_format); + if usage.has_token_signal() || self.claude_usage_observed { + self.latest_usage = Some(usage); + } + return; + } + } + let mut usage = map_usage_from_response(&envelope, provider_format); + if let Some(image_count) = image_count { + usage.request_count = image_count; + usage + .dimensions + .insert("image_count".to_owned(), Value::from(image_count)); + } + if usage.has_token_signal() || contains_explicit_usage(&envelope) { + self.latest_usage = Some(usage); + } + } +} + +fn provider_format(context: &Value) -> &str { + [ + "provider_stream_event_api_format", + "provider_stream_api_format", + "provider_api_format", + ] + .into_iter() + .filter_map(|field| context.get(field).and_then(Value::as_str)) + .map(str::trim) + .find(|value| !value.is_empty()) + .unwrap_or_default() +} + +fn is_openai_image_api(api_format: &str) -> bool { + let mut parts = api_format.split(':').map(str::trim); + parts + .next() + .is_some_and(|part| part.eq_ignore_ascii_case("openai")) + && parts + .next() + .is_some_and(|part| part.eq_ignore_ascii_case("image")) +} + +fn merge_claude_usage_projection(target: &mut Value, incoming: &serde_json::Map) { + let target = target.as_object_mut().expect("Claude usage is an object"); + for key in [ + "input_tokens", + "output_tokens", + "cache_creation_input_tokens", + "cache_read_input_tokens", + "total_tokens", + ] { + if let Some(value) = incoming.get(key) { + target.insert(key.to_owned(), usage_integer_or_null(value)); + } + } + if let Some(value) = incoming.get("cache_creation") { + let projected = value.as_object().map(|object| { + ["ephemeral_5m_input_tokens", "ephemeral_1h_input_tokens"] + .into_iter() + .filter_map(|key| { + object + .get(key) + .map(|value| (key.to_owned(), usage_integer_or_null(value))) + }) + .collect::>() + }); + // The original merge replaces this whole subobject, including an empty + // or malformed replacement. Do not merge its leaves with older values. + target.insert( + "cache_creation".to_owned(), + projected.map(Value::Object).unwrap_or(Value::Null), + ); + } +} + +fn usage_integer_or_null(value: &Value) -> Value { + value.as_i64().map(Value::from).unwrap_or(Value::Null) +} + +#[derive(Default, Deserialize)] +#[serde(default)] +struct UsageEnvelope { + #[serde(rename = "type")] + event_type: Option, + #[serde(deserialize_with = "deserialize_present_usage")] + usage: Option, + #[serde( + rename = "usageMetadata", + deserialize_with = "deserialize_present_usage" + )] + usage_metadata: Option, + service_tier: Option, + speed: Option, + #[serde(deserialize_with = "deserialize_nested_usage")] + response: Option>, + #[serde(deserialize_with = "deserialize_nested_usage")] + message: Option>, + #[serde(deserialize_with = "deserialize_nested_usage")] + item: Option>, + #[serde(deserialize_with = "deserialize_usage_array")] + candidates: Option>, + #[serde(deserialize_with = "deserialize_usage_array")] + chunks: Option>, + data: ImageResponseCount, + result: ImageResponseCount, +} + +fn deserialize_present_usage<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result, D::Error> { + // A present null stops the mapper's search through nested or older usage. + // Missing fields still use the envelope's default None. + Value::deserialize(deserializer).map(Some) +} + +#[derive(Default)] +enum ImageResponseCount { + #[default] + Empty, + Array(usize), + Single, +} + +impl<'de> Deserialize<'de> for ImageResponseCount { + fn deserialize>(deserializer: D) -> Result { + struct ImageCountVisitor; + + impl<'de> Visitor<'de> for ImageCountVisitor { + type Value = ImageResponseCount; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("an image response result") + } + + fn visit_seq>( + self, + mut sequence: A, + ) -> Result { + let mut count = 0; + while sequence.next_element::()?.is_some() { + count += 1; + } + Ok(ImageResponseCount::Array(count)) + } + + fn visit_map>(self, mut map: A) -> Result { + let mut nonempty = false; + while map.next_entry::()?.is_some() { + nonempty = true; + } + Ok(if nonempty { + ImageResponseCount::Single + } else { + ImageResponseCount::Empty + }) + } + + fn visit_str(self, text: &str) -> Result { + Ok(if text.trim().is_empty() { + ImageResponseCount::Empty + } else { + ImageResponseCount::Single + }) + } + + fn visit_unit(self) -> Result { + Ok(ImageResponseCount::Empty) + } + + fn visit_bool(self, _: bool) -> Result { + Ok(ImageResponseCount::Empty) + } + + fn visit_i64(self, _: i64) -> Result { + Ok(ImageResponseCount::Empty) + } + + fn visit_u64(self, _: u64) -> Result { + Ok(ImageResponseCount::Empty) + } + + fn visit_f64(self, _: f64) -> Result { + Ok(ImageResponseCount::Empty) + } + } + + deserializer.deserialize_any(ImageCountVisitor) + } +} + +enum UsageValue { + Object(UsageEnvelope), + Array(Vec), + Ignored, +} + +impl<'de> Deserialize<'de> for UsageValue { + fn deserialize>(deserializer: D) -> Result { + struct UsageValueVisitor; + + impl<'de> Visitor<'de> for UsageValueVisitor { + type Value = UsageValue; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("an optional usage envelope") + } + + fn visit_map>(self, map: A) -> Result { + UsageEnvelope::deserialize(serde::de::value::MapAccessDeserializer::new(map)) + .map(UsageValue::Object) + } + + fn visit_seq>( + self, + mut sequence: A, + ) -> Result { + let mut first = None; + let mut envelopes = Vec::new(); + while let Some(value) = sequence.next_element::()? { + let envelope = match value { + UsageValue::Object(envelope) => envelope, + UsageValue::Array(_) | UsageValue::Ignored => UsageEnvelope::default(), + }; + if first.is_none() && envelopes.is_empty() { + first = Some(envelope); + } else if !envelope.is_empty() { + // Only candidates[0] is positional in the usage mapper. + // Keep that slot even if empty; subsequent empty slots + // cannot contribute usage, explicit-zero signals or tier. + if let Some(first) = first.take() { + envelopes.push(first); + } + envelopes.push(envelope); + } + } + if let Some(first) = first.filter(|first| !first.is_empty()) { + envelopes.push(first); + } + Ok(UsageValue::Array(envelopes)) + } + + fn visit_unit(self) -> Result { + Ok(UsageValue::Ignored) + } + + fn visit_bool(self, _: bool) -> Result { + Ok(UsageValue::Ignored) + } + + fn visit_i64(self, _: i64) -> Result { + Ok(UsageValue::Ignored) + } + + fn visit_u64(self, _: u64) -> Result { + Ok(UsageValue::Ignored) + } + + fn visit_f64(self, _: f64) -> Result { + Ok(UsageValue::Ignored) + } + + fn visit_str(self, _: &str) -> Result { + Ok(UsageValue::Ignored) + } + } + + deserializer.deserialize_any(UsageValueVisitor) + } +} + +fn deserialize_nested_usage<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result>, D::Error> { + match UsageValue::deserialize(deserializer)? { + UsageValue::Object(envelope) => Ok(Some(Box::new(envelope))), + UsageValue::Array(_) | UsageValue::Ignored => Ok(None), + } +} + +fn deserialize_usage_array<'de, D: Deserializer<'de>>( + deserializer: D, +) -> Result>, D::Error> { + match UsageValue::deserialize(deserializer)? { + UsageValue::Array(envelopes) => Ok(Some(envelopes)), + UsageValue::Object(_) | UsageValue::Ignored => Ok(None), + } +} + +impl UsageEnvelope { + fn image_count(&self) -> Option { + match self.data { + ImageResponseCount::Array(count) if count > 0 => Some(count as i64), + _ => match self.result { + ImageResponseCount::Array(count) if count > 0 => Some(count as i64), + ImageResponseCount::Single => Some(1), + _ => None, + }, + } + } + + fn is_empty(&self) -> bool { + self.event_type.is_none() + && self.usage.is_none() + && self.usage_metadata.is_none() + && self.service_tier.is_none() + && self.speed.is_none() + && [&self.response, &self.message, &self.item] + .into_iter() + .all(|value| value.as_ref().is_none_or(|value| value.is_empty())) + && [&self.candidates, &self.chunks].into_iter().all(|values| { + values + .as_ref() + .is_none_or(|values| values.iter().all(Self::is_empty)) + }) + } + + fn into_value(self) -> Value { + let mut object = serde_json::Map::new(); + if let Some(event_type) = self.event_type { + object.insert("type".to_string(), Value::String(event_type)); + } + for (key, value) in [ + ("usage", self.usage), + ("usageMetadata", self.usage_metadata), + ("service_tier", self.service_tier), + ("speed", self.speed), + ] { + if let Some(value) = value { + object.insert(key.to_string(), value); + } + } + for (key, value) in [ + ("response", self.response), + ("message", self.message), + ("item", self.item), + ] { + if let Some(value) = value { + object.insert(key.to_string(), value.into_value()); + } + } + for (key, values) in [("candidates", self.candidates), ("chunks", self.chunks)] { + if let Some(values) = values.filter(|values| !values.is_empty()) { + object.insert( + key.to_string(), + Value::Array(values.into_iter().map(Self::into_value).collect()), + ); + } + } + Value::Object(object) + } +} + +fn contains_explicit_usage(value: &Value) -> bool { + ["usage", "usageMetadata"].into_iter().any(|key| { + value + .get(key) + .and_then(Value::as_object) + .is_some_and(|value| !value.is_empty()) + }) || ["response", "message", "item"] + .into_iter() + .any(|key| value.get(key).is_some_and(contains_explicit_usage)) + || ["candidates", "chunks"].into_iter().any(|key| { + value + .get(key) + .and_then(Value::as_array) + .is_some_and(|values| values.iter().any(contains_explicit_usage)) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn stream_capture_fallback_preserves_present_null_usage_lookup_barriers() { + let previous = json!({ + "usage": {"prompt_tokens": 10, "completion_tokens": 20}, + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20} + }); + let cases = [ + ( + "openai:chat", + json!({"chunks": [ + {"usage": {"prompt_tokens": 10, "completion_tokens": 20}}, + {"usage": null} + ]}), + ), + ( + "openai:responses", + json!({"usage": null, "response": { + "usage": {"input_tokens": 10, "output_tokens": 20} + }}), + ), + ( + "gemini:generate_content", + json!({"chunks": [ + {"usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20}}, + {"usageMetadata": null} + ]}), + ), + ( + "gemini:generate_content", + json!({"usageMetadata": null, "response": { + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20} + }}), + ), + ( + "gemini:generate_content", + json!({"candidates": [ + {"usageMetadata": null}, + {"usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20}} + ]}), + ), + ]; + for (format, original) in cases { + let context = json!({"provider_api_format": format}); + let expected = map_usage_from_response(&original, format); + assert_eq!(expected.input_tokens, 0); + assert_eq!(expected.output_tokens, 0); + assert!(contains_explicit_usage(&original)); + let projected = serde_json::from_value::(original.clone()) + .unwrap() + .into_value(); + assert_eq!(map_usage_from_response(&projected, format), expected); + assert!(contains_explicit_usage(&projected)); + + let json = serde_json::to_vec(&original).unwrap(); + let sse = format!("data: {original}\r\n\r\n"); + for chunk_size in [1, 3, 17, sse.len()] { + let mut fallback = StreamUsageFallback::new(1024); + for chunk in json.chunks(chunk_size) { + fallback.observe(&context, chunk); + } + assert_eq!(fallback.finish(&context), Some(expected.clone())); + + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, format!("data: {previous}\n\n").as_bytes()); + assert_eq!(fallback.latest_usage.as_ref().unwrap().output_tokens, 20); + for chunk in sse.as_bytes().chunks(chunk_size) { + fallback.observe(&context, chunk); + } + assert_eq!(fallback.finish(&context), Some(expected.clone())); + } + } + } + + #[test] + fn stream_capture_fallback_distinguishes_missing_and_null_usage_without_new_signal() { + let missing = serde_json::from_value::(json!({})).unwrap(); + assert!(missing.usage.is_none()); + assert!(missing.usage_metadata.is_none()); + assert!(missing.is_empty()); + for (format, key, tokens) in [ + ( + "openai:chat", + "usage", + json!({"prompt_tokens": 10, "completion_tokens": 20}), + ), + ( + "gemini:generate_content", + "usageMetadata", + json!({"promptTokenCount": 10, "candidatesTokenCount": 20}), + ), + ] { + let context = json!({"provider_api_format": format}); + let null = json!({key: null}); + let projected = serde_json::from_value::(null.clone()).unwrap(); + assert!(!projected.is_empty()); + assert_eq!(projected.into_value(), null); + assert!(!contains_explicit_usage(&null)); + let valid = json!({"response": {key: tokens}}); + let expected = map_usage_from_response(&valid, format); + assert_eq!(expected.output_tokens, 20); + let mut fallback = StreamUsageFallback::new(1024); + let sse = format!("data: {valid}\n\ndata: {null}\n\n"); + for chunk in sse.as_bytes().chunks(3) { + fallback.observe(&context, chunk); + } + assert_eq!(fallback.finish(&context), Some(expected)); + } + } + + #[test] + fn stream_capture_fallback_image_counts_match_response_mapper_without_retaining_images() { + let cases = [ + json!({}), + json!({"unknown": [1, 2, 3]}), + json!({"data": null, "result": null}), + json!({"data": [], "result": []}), + json!({"data": {}, "result": {}}), + json!({"data": "ignored", "result": " \n\t "}), + json!({"data": true, "result": 3.5}), + json!({"data": false, "result": true}), + json!({"data": [null, false, 1], "result": [1, 2]}), + json!({"data": [], "result": [null, false]}), + json!({"data": [null], "result": [1, 2, 3]}), + json!({"data": {"ignored": 1}, "result": {"url": "image"}}), + json!({"data": 1, "result": " image "}), + json!({"result": "escaped\nimage"}), + json!({"response": {"data": [1, 2]}, "chunks": [{"result": [1, 2]}]}), + json!({"data": [{"b64_json": "x".repeat(64 * 1024)}, null]}), + ]; + for mut original in cases { + original.as_object_mut().unwrap().insert( + "usage".to_owned(), + json!({"prompt_tokens": 10, "completion_tokens": 20}), + ); + let bytes = serde_json::to_vec(&original).unwrap(); + let projected = serde_json::from_slice::(&bytes).unwrap(); + let image_count = projected.image_count(); + let compact = projected.into_value(); + assert!(compact.get("data").is_none()); + assert!(compact.get("result").is_none()); + assert!(serde_json::to_vec(&compact).unwrap().len() < 256); + for format in [ + " OpenAI : Image ", + "openai:chat", + "gemini:generate_content", + "unknown:api", + ] { + let expected = map_usage_from_response(&original, format); + if is_openai_image_api(format) { + assert_eq!( + image_count.map(Value::from).as_ref(), + expected.dimensions.get("image_count") + ); + } + let context = json!({"provider_api_format": format}); + let sse = format!("data: {original}\r\n\r\n"); + for chunk_size in [3, 17, sse.len()] { + for input in [bytes.as_slice(), sse.as_bytes()] { + let mut fallback = StreamUsageFallback::new(128 * 1024); + for chunk in input.chunks(chunk_size) { + fallback.observe(&context, chunk); + } + assert_eq!(fallback.finish(&context), Some(expected.clone())); + } + } + } + } + } + + #[test] + fn stream_capture_fallback_image_only_records_keep_count_signal() { + for original in [ + json!({"data": [null, null]}), + json!({"result": [null, null, null]}), + json!({"result": {"url": "image"}}), + json!({"result": "image"}), + ] { + let context = json!({"provider_api_format": "openai:image"}); + let expected = map_usage_from_response(&original, "openai:image"); + let json = serde_json::to_vec(&original).unwrap(); + let sse = format!("data: {original}\n\n"); + for input in [json.as_slice(), sse.as_bytes()] { + for chunk_size in [1, 3, input.len()] { + let mut fallback = StreamUsageFallback::new(1024); + for chunk in input.chunks(chunk_size) { + fallback.observe(&context, chunk); + } + assert_eq!(fallback.finish(&context), Some(expected.clone())); + } + } + let context = json!({"provider_api_format": "openai:chat"}); + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, sse.as_bytes()); + assert_eq!(fallback.finish(&context), None); + } + } + + #[test] + fn stream_capture_fallback_claude_projection_matches_original_raw_merge() { + for provider_format in ["claude:messages", " Anthropic:messages "] { + let context = json!({"provider_api_format": provider_format}); + let mut fallback = StreamUsageFallback::new(64 * 1024); + let mut raw_merged = serde_json::Map::new(); + let mut expected_usage = None; + let mut expected_tier = None; + let events = [ + json!({"type": "message_start", "message": {"usage": {"unknown": [1, 2]}}}), + json!({"type": "message_delta", "usage": {}}), + json!({"type": "message_delta", "usage": { + "input_tokens": 100, "output_tokens": 10, "total_tokens": 110, + "cache_creation_input_tokens": 7, "cache_read_input_tokens": 30, + "cache_creation": {"ephemeral_5m_input_tokens": 5, "ephemeral_1h_input_tokens": 2}, + "speed": " FAST " + }}), + json!({"type": "message_delta", "usage": { + "input_tokens": 0, "output_tokens": 500, "total_tokens": 500, + "cache_read_input_tokens": 0, "cache_creation_input_tokens": 0, + "cache_creation": {"ephemeral_1h_input_tokens": 3}, + "service_tier": "priority" + }}), + json!({"type": "message_delta", "usage": { + "input_tokens": "12", "output_tokens": null, "total_tokens": "500", + "cache_read_input_tokens": [10], "cache_creation_input_tokens": {"nested": 20}, + "cache_creation": {}, "speed": "standard" + }}), + json!({"type": "message_delta", "usage": { + "input_tokens": -1, "output_tokens": u64::MAX, "total_tokens": u64::MAX, + "cache_creation": {"ephemeral_5m_input_tokens": 1.5, "ephemeral_1h_input_tokens": false} + }}), + json!({"type": "message_delta", "usage": {"cache_creation": [7]}}), + json!({"type": "message_start", "message": {"usage": {}}}), + json!({"type": "message_delta", "usage": {"unknown_only": false}}), + json!({"type": "message_delta", "usage": null}), + ]; + for event in events { + if let Some(tier) = extract_provider_actual_service_tier_from_response(Some(&event)) + { + expected_tier = Some(tier); + } + let raw_usage = match event.get("type").and_then(Value::as_str) { + Some("message_start") => { + raw_merged.clear(); + event.pointer("/message/usage") + } + Some("message_delta") => event.get("usage"), + _ => None, + }; + let original = if let Some(raw_usage) = raw_usage.and_then(Value::as_object) { + raw_merged.extend(raw_usage.clone()); + json!({"usage": raw_merged}) + } else { + event.clone() + }; + let mapped = map_usage_from_response(&original, provider_format); + if mapped.has_token_signal() || contains_explicit_usage(&original) { + expected_usage = Some(mapped); + } + fallback.observe(&context, format!("data: {event}\n\n").as_bytes()); + assert_eq!(fallback.latest_usage, expected_usage, "event={event}"); + assert_eq!(fallback.latest_service_tier, expected_tier, "event={event}"); + } + } + } + + #[test] + fn stream_capture_fallback_claude_unknown_fields_do_not_accumulate() { + let context = json!({"provider_api_format": "claude:messages"}); + let mut fallback = StreamUsageFallback::new(16 * 1024); + for index in 0..256 { + let event = json!({"type": "message_delta", "usage": { + format!("unknown_{index}"): "x".repeat(4096), + "output_tokens": index, + "cache_read_input_tokens": 0 + }}); + fallback.observe(&context, format!("data: {event}\n\n").as_bytes()); + let retained = fallback.claude_usage.as_object().unwrap(); + assert_eq!(retained.len(), 2); + assert_eq!(retained.get("output_tokens"), Some(&json!(index))); + assert_eq!(retained.get("cache_read_input_tokens"), Some(&json!(0))); + assert_eq!(fallback.record.capacity(), 0); + } + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.output_tokens, 255); + assert_eq!(usage.cache_read_tokens, 0); + } + + #[test] + fn stream_capture_fallback_array_projection_skips_empty_envelopes_without_shifting_first() { + let mut empty = vec![ + Value::Null, + json!(false), + json!(3), + json!({"content": "unused"}), + ]; + empty.extend(std::iter::repeat_n(Value::Null, 10_000)); + let empty_response = json!({"candidates": empty, "chunks": empty}); + let projected = serde_json::from_value::(empty_response).unwrap(); + assert!(projected.candidates.as_ref().unwrap().is_empty()); + assert!(projected.chunks.as_ref().unwrap().is_empty()); + + let original = json!({ + "candidates": [null, {"usageMetadata": {"promptTokenCount": 99}}], + "chunks": [null, {"content": "unused"}, {"usage": { + "input_tokens": 10, "output_tokens": 7, "cache_read_input_tokens": 0 + }, "service_tier": "priority"}, null, {}, {"usageMetadata": { + "promptTokenCount": 12, "candidatesTokenCount": 0 + }, "speed": "fast"}, null] + }); + let projected = serde_json::from_value::(original.clone()).unwrap(); + assert_eq!(projected.candidates.as_ref().unwrap().len(), 2); + assert_eq!(projected.chunks.as_ref().unwrap().len(), 3); + let compact = projected.into_value(); + for format in [ + "openai:chat", + "openai:responses", + "openai:image", + "gemini:generate_content", + "claude:messages", + "unknown:format", + ] { + assert_eq!( + map_usage_from_response(&compact, format), + map_usage_from_response(&original, format) + ); + } + assert_eq!( + contains_explicit_usage(&compact), + contains_explicit_usage(&original) + ); + assert_eq!( + extract_provider_actual_service_tier_from_response(Some(&compact)), + extract_provider_actual_service_tier_from_response(Some(&original)) + ); + } + + #[test] + fn stream_capture_fallback_complete_records_borrow_transport_bytes() { + let context = json!({"provider_api_format": "openai:chat"}); + let chunk = format!( + "data: {{\"content\":\"{}\"}}\n\ndata: {{\"usage\":{{\"prompt_tokens\":10,\"completion_tokens\":20}}}}\n\n", + "x".repeat(32 * 1024) + ); + let mut borrowed = StreamUsageFallback::new(64 * 1024); + let mut buffered = StreamUsageFallback::new(64 * 1024); + borrowed.observe(&context, chunk.as_bytes()); + buffered.observe_inner::(&context, chunk.as_bytes()); + assert_eq!(borrowed.copied_record_bytes, 0); + assert_eq!(buffered.copied_record_bytes, chunk.len()); + assert_eq!(borrowed.latest_usage, buffered.latest_usage); + assert_eq!(borrowed.finish(&context).unwrap().output_tokens, 20); + } + + #[test] + fn stream_capture_fallback_borrowed_records_match_buffered_limits_and_split_endings() { + let context = json!({"provider_api_format": "openai:chat"}); + for ending in ["\n", "\r\n", "\r"] { + let first = format!( + "data: {{\"usage\":{{\"prompt_tokens\":100,\"completion_tokens\":10}},\"content\":\"{}\"}}{ending}{ending}", + "x".repeat(80) + ); + let input = format!( + ": comment{ending}{ending}{first}data: {{\"\\u0075sage\":{ending}data: {{\"completion_tokens\":20}},\"service_tier\":\"priority\"}}{ending}{ending}data: [DONE]{ending}{ending}data: {{\"usage\":{{\"completion_tokens\":30}}}}" + ); + for limit in [8, 64, first.len() - 1, first.len(), first.len() + 1, 512] { + for chunk_size in [1, 2, 3, 7, 16, 64, input.len()] { + let mut borrowed = StreamUsageFallback::new(limit); + let mut buffered = StreamUsageFallback::new(limit); + for chunk in input.as_bytes().chunks(chunk_size) { + borrowed.observe(&context, chunk); + buffered.observe_inner::(&context, chunk); + assert_eq!( + borrowed.record, buffered.record, + "ending={ending:?} limit={limit} chunk={chunk_size}" + ); + assert_eq!( + borrowed.dropping_record, buffered.dropping_record, + "ending={ending:?} limit={limit} chunk={chunk_size}" + ); + assert_eq!( + borrowed.latest_usage, buffered.latest_usage, + "ending={ending:?} limit={limit} chunk={chunk_size}" + ); + assert_eq!( + borrowed.latest_service_tier, buffered.latest_service_tier, + "ending={ending:?} limit={limit} chunk={chunk_size}" + ); + } + assert_eq!(borrowed.finish(&context), buffered.finish(&context)); + assert_eq!(borrowed.take_service_tier(), buffered.take_service_tier()); + } + } + } + } + + #[test] + fn stream_capture_fallback_keeps_latest_split_record_usage_and_releases_capacity() { + let context = json!({"provider_api_format": "openai:chat"}); + let mut fallback = StreamUsageFallback::new(1024); + let events = b"data: {\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":10}}\r\n\r\ndata: {\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":20,\"prompt_tokens_details\":{\"cached_tokens\":30}}}\r\n\r\n"; + for part in events.chunks(7) { + fallback.observe(&context, part); + } + let usage = fallback.finish(&context).expect("last usage signal"); + assert_eq!(usage.input_tokens, 100); + assert_eq!(usage.output_tokens, 20); + assert_eq!(usage.cache_read_tokens, 30); + assert_eq!(fallback.record.capacity(), 0); + } + + #[test] + fn stream_capture_fallback_recovers_after_an_oversized_record() { + let context = json!({"provider_api_format": "openai:chat"}); + for ending in ["\n", "\r\n", "\r"] { + for chunk_size in 1..=16 { + let mut fallback = StreamUsageFallback::new(96); + let events = format!( + "data: {}{ending}{ending}data: {{\"usage\":{{\"completion_tokens\":12}}}}{ending}{ending}", + "x".repeat(128), + ); + for chunk in events.as_bytes().chunks(chunk_size) { + fallback.observe(&context, chunk); + } + assert_eq!( + fallback.finish(&context).unwrap().output_tokens, + 12, + "ending={ending:?}, chunk_size={chunk_size}" + ); + assert_eq!(fallback.record.capacity(), 0); + } + } + } + + #[test] + fn stream_capture_fallback_handles_nested_gemini_usage() { + let context = json!({"provider_api_format": "gemini:generate_content"}); + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, b"data: {\"response\":{\"usageMetadata\":{\"promptTokenCount\":11,\"candidatesTokenCount\":7,\"cachedContentTokenCount\":3}}}\n\n"); + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + assert_eq!(usage.cache_read_tokens, 3); + } + + #[test] + fn stream_capture_fallback_merges_anthropic_partial_fields_and_preserves_explicit_zero() { + for format in ["claude:messages", " Anthropic:messages "] { + let context = json!({"provider_api_format": format}); + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, b"data: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":100,\"cache_read_input_tokens\":30,\"cache_creation_input_tokens\":7}}}\n\n"); + fallback.observe(&context, b"data: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":20,\"cache_read_input_tokens\":0}}\n\n"); + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.input_tokens, 100); + assert_eq!(usage.output_tokens, 20); + assert_eq!(usage.cache_read_tokens, 0); + assert_eq!(usage.cache_creation_tokens, 7); + } + } + + #[test] + fn stream_capture_fallback_latest_complete_snapshot_can_reset_cache_to_zero() { + let context = json!({"provider_api_format": "openai:chat"}); + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, b"data: {\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":10,\"prompt_tokens_details\":{\"cached_tokens\":30}}}\n\n"); + fallback.observe(&context, b"data: {\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":20,\"prompt_tokens_details\":{\"cached_tokens\":0}}}\n\n"); + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.output_tokens, 20); + assert_eq!(usage.cache_read_tokens, 0); + } + + #[test] + fn stream_capture_fallback_accepts_escaped_keys_multiline_and_eof_records() { + let context = json!({ + "provider_stream_event_api_format": null, + "provider_stream_api_format": " openai:chat ", + "provider_api_format": "gemini:generate_content" + }); + let mut fallback = StreamUsageFallback::new(1024); + let event = b"data: {\"\\u0075sage\":\r\ndata: {\"prompt_tokens\":11,\"completion_tokens\":7},\"service_tier\":\" PRIORITY \"}"; + for chunk in event.chunks(3) { + fallback.observe(&context, chunk); + } + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + assert_eq!(fallback.take_service_tier().as_deref(), Some("priority")); + assert_eq!(fallback.record.capacity(), 0); + } + + #[test] + fn stream_capture_fallback_json_whitespace_does_not_flush_an_incomplete_response() { + let context = json!({"provider_api_format": "openai:chat"}); + let mut fallback = StreamUsageFallback::new(1024); + let response = b"{\n\n\"choices\": [],\n\n\"usage\": {\"prompt_tokens\":11,\"completion_tokens\":7}\n}"; + for chunk in response.chunks(5) { + fallback.observe(&context, chunk); + } + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + assert_eq!(fallback.record.capacity(), 0); + } + + #[test] + fn stream_capture_fallback_record_boundaries_accept_all_sse_endings_and_fragment_sizes() { + let context = json!({"provider_api_format": "openai:chat"}); + for ending in ["\n", "\r\n", "\r"] { + for chunk_size in 1..=16 { + let mut fallback = StreamUsageFallback::new(1024); + let events = format!( + "data: {{\"usage\":{{\"prompt_tokens\":100,\"completion_tokens\":10}}}}{ending}{ending}data: {{\"usage\":{{\"prompt_tokens\":100,\"completion_tokens\":20}}}}{ending}{ending}data: [DONE]{ending}{ending}", + ); + for chunk in events.as_bytes().chunks(chunk_size) { + fallback.observe(&context, chunk); + } + let usage = fallback.finish(&context).unwrap(); + assert_eq!( + usage.input_tokens, 100, + "ending={ending:?}, chunk_size={chunk_size}" + ); + assert_eq!( + usage.output_tokens, 20, + "ending={ending:?}, chunk_size={chunk_size}" + ); + assert_eq!(fallback.record.capacity(), 0); + } + } + } + + #[test] + fn stream_capture_fallback_ignores_non_object_envelopes_and_large_unknown_content() { + let context = json!({"provider_api_format": "openai:chat"}); + let mut fallback = StreamUsageFallback::new(128 * 1024); + let body = json!({ + "usage": {"prompt_tokens": 11, "completion_tokens": 7}, + "message": "provider diagnostic text", + "response": false, + "item": 42, + "content": "x".repeat(64 * 1024), + "chunks": [null, false, 7, "diagnostic", {"content": "unused"}], + }); + let bytes = serde_json::to_vec(&body).unwrap(); + let compact = serde_json::from_slice::(&bytes) + .unwrap() + .into_value(); + assert!(compact.get("content").is_none()); + assert!(compact.get("message").is_none()); + assert!(compact.get("response").is_none()); + assert!(serde_json::to_vec(&compact).unwrap().len() < 256); + fallback.observe(&context, &bytes); + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.input_tokens, 11); + assert_eq!(usage.output_tokens, 7); + + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, b"data: {\"chunks\":[null,\"ignored\",{\"usage\":{\"prompt_tokens\":21,\"completion_tokens\":14}},false]}\n\n"); + let usage = fallback.finish(&context).unwrap(); + assert_eq!(usage.input_tokens, 21); + assert_eq!(usage.output_tokens, 14); + } + + #[test] + fn stream_capture_fallback_preserves_candidate_indices_for_gemini_mapping() { + let context = json!({"provider_api_format": "gemini:generate_content"}); + let response = json!({ + "candidates": [null, {"usageMetadata": { + "promptTokenCount": 11, "candidatesTokenCount": 7 + }}] + }); + let expected = map_usage_from_response(&response, "gemini:generate_content"); + let mut fallback = StreamUsageFallback::new(1024); + fallback.observe(&context, &serde_json::to_vec(&response).unwrap()); + assert_eq!(fallback.finish(&context).unwrap(), expected); + assert_eq!(expected.input_tokens, 0); + assert_eq!(expected.output_tokens, 0); + } +} diff --git a/apps/aether-gateway/src/execution_runtime/stream_pump.rs b/apps/aether-gateway/src/execution_runtime/stream_pump.rs index ab4e2e408..96b240a78 100644 --- a/apps/aether-gateway/src/execution_runtime/stream_pump.rs +++ b/apps/aether-gateway/src/execution_runtime/stream_pump.rs @@ -12,7 +12,6 @@ use async_stream::stream; use axum::body::Bytes; use base64::Engine as _; use futures_util::{Stream, StreamExt}; -use http_body_util::BodyExt; use serde_json::Value; use tracing::warn; @@ -21,9 +20,14 @@ use crate::ai_serving::api::{ normalize_provider_private_report_context, StreamingStandardTerminalObserver, }; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; +use crate::execution_runtime::stream::ClientVisibleStreamCompletionTracker; +use crate::execution_runtime::stream_read_timeout::{ + await_stream_idle_read, stream_idle_timeout_message, +}; use crate::execution_runtime::transport::{ append_upstream_response_body_chunk, decode_response_body_bytes, - stream_first_byte_timeout_message, DirectUpstreamResponse, + direct_upstream_response_byte_stream, stream_first_byte_timeout_message, + DirectUpstreamResponse, }; use crate::execution_runtime::DirectUpstreamStreamExecution; use crate::GatewayError; @@ -31,6 +35,15 @@ use crate::GatewayError; const STREAM_USAGE_OBSERVER_MAX_LINE_BYTES: usize = 1024 * 1024; const UPSTREAM_STREAM_READ_ERROR_MESSAGE: &str = "Upstream response stream failed"; +fn upstream_stream_error_category(response: &DirectUpstreamResponse) -> &'static str { + match response { + DirectUpstreamResponse::Reqwest(_) => "reqwest_body_read_failed", + DirectUpstreamResponse::HyperH2c(_) => "hyper_body_read_failed", + DirectUpstreamResponse::BrowserWreq(_) => "browser_body_read_failed", + DirectUpstreamResponse::LocalTunnel(_) => "tunnel_body_read_failed", + } +} + pub(crate) fn build_direct_execution_frame_stream( execution: DirectUpstreamStreamExecution, ) -> impl Stream> + Send + 'static { @@ -49,9 +62,11 @@ pub(crate) fn build_direct_execution_frame_stream( started_at, response_observation, stream_first_byte_timeout, + stream_idle_timeout, upstream_target_permit, } = execution; let _upstream_target_permit = upstream_target_permit; + let upstream_error_category = upstream_stream_error_category(&response); let mut observer_context = stream_summary_report_context; if observer_context @@ -74,6 +89,7 @@ pub(crate) fn build_direct_execution_frame_stream( let mut private_stream_normalizer = maybe_build_provider_private_stream_normalizer(Some(&observer_context)); let mut stream_terminal_observer = StreamingStandardTerminalObserver::default(); + let mut stream_completion = ClientVisibleStreamCompletionTracker::default(); let mut observer_buffered = Vec::new(); if should_buffer_non_stream_response( @@ -87,6 +103,7 @@ pub(crate) fn build_direct_execution_frame_stream( response, started_at, stream_first_byte_timeout, + stream_idle_timeout, ) .await { @@ -166,6 +183,7 @@ pub(crate) fn build_direct_execution_frame_stream( ttfb_ms, upstream_bytes, first_byte_timeout, + idle_timeout, }) => { match encode_headers_frame( status_code, @@ -180,6 +198,8 @@ pub(crate) fn build_direct_execution_frame_stream( } let error_frame = if let Some(timeout) = first_byte_timeout { encode_first_byte_timeout_frame(timeout) + } else if let Some(timeout) = idle_timeout { + encode_idle_timeout_frame(timeout) } else { encode_error_frame(message) }; @@ -228,7 +248,9 @@ pub(crate) fn build_direct_execution_frame_stream( let mut prefetched_body_failed = false; for item in prefetched_body { match item { + Ok(chunk) if chunk.is_empty() => continue, Ok(chunk) => { + stream_completion.observe_chunk(&chunk); if ttfb_ms.is_none() { ttfb_ms = Some(started_at.elapsed().as_millis() as u64); } @@ -280,328 +302,97 @@ pub(crate) fn build_direct_execution_frame_stream( } } if !prefetched_body_failed { - match response { - DirectUpstreamResponse::Reqwest(response) => { - let mut bytes_stream = response.bytes_stream(); - loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - bytes_stream.next(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - match encode_first_byte_timeout_frame(timeout) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - break; - } + let mut bytes_stream = direct_upstream_response_byte_stream(VecDeque::new(), response); + loop { + let item = if ttfb_ms.is_none() { + match await_stream_first_byte( + bytes_stream.next(), started_at, stream_first_byte_timeout, + ).await { + Ok(item) => item, + Err(timeout) => { + match encode_first_byte_timeout_frame(timeout) { + Ok(frame) => yield Ok(frame), + Err(err) => yield Err(err), } - } else { - bytes_stream.next().await - }; - let Some(item) = item else { break; - }; - match item { - Ok(chunk) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - if !first_chunk_telemetry_emitted { - match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - first_chunk_telemetry_emitted = true; - } - upstream_bytes += chunk.len() as u64; - observe_stream_chunk( - &mut stream_terminal_observer, - &normalized_observer_context, - private_stream_normalizer.as_mut(), - &mut observer_buffered, - chunk.as_ref(), - ); - match encode_data_frame(&chunk) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - } - Err(_err) => { - let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - status_code, - upstream_bytes, - error_category = "reqwest_body_read_failed", - "upstream body stream read error" - ); - match encode_error_frame(message) { - Ok(frame) => yield Ok(frame), - Err(encode_err) => { - yield Err(encode_err); - return; - } - } - break; - } } } - } - DirectUpstreamResponse::HyperH2c(response) => { - let mut bytes_stream = response.into_body().into_data_stream(); - loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - bytes_stream.next(), - started_at, - stream_first_byte_timeout, - ) - .await + } else { + match await_stream_idle_read(bytes_stream.next(), stream_idle_timeout).await { + Ok(item) => item, + Err(timeout) => { + drop(bytes_stream); + if stream_completion.successful_completion() + || (!stream_completion.observed_terminal() + && stream_terminal_observer.latest_summary().is_some_and(|summary| { + summary.observed_finish && summary.parser_error.is_none() + && summary.finish_reason.as_deref() != Some("error") + })) { - Ok(item) => item, - Err(timeout) => { - match encode_first_byte_timeout_frame(timeout) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - break; - } - } - } else { - bytes_stream.next().await - }; - let Some(item) = item else { - break; - }; - match item { - Ok(chunk) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - if !first_chunk_telemetry_emitted { - match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - first_chunk_telemetry_emitted = true; - } - upstream_bytes += chunk.len() as u64; - observe_stream_chunk( - &mut stream_terminal_observer, - &normalized_observer_context, - private_stream_normalizer.as_mut(), - &mut observer_buffered, - chunk.as_ref(), - ); - match encode_data_frame(&chunk) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - } - Err(_err) => { - let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - status_code, - upstream_bytes, - error_category = "hyper_body_read_failed", - "upstream body stream read error" - ); - match encode_error_frame(message) { - Ok(frame) => yield Ok(frame), - Err(encode_err) => { - yield Err(encode_err); - return; - } - } break; } + if stream_terminal_observer.latest_summary().is_some_and(|summary| { + summary.observed_finish && summary.parser_error.is_some() + }) { + // The terminal summary carries the original provider failure. + break; + } + match encode_idle_timeout_frame(timeout) { + Ok(frame) => yield Ok(frame), + Err(err) => yield Err(err), + } + break; } } - } - DirectUpstreamResponse::BrowserWreq(response) => { - let mut bytes_stream = response.bytes_stream(); - loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - bytes_stream.next(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - match encode_first_byte_timeout_frame(timeout) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - break; - } - } - } else { - bytes_stream.next().await - }; - let Some(item) = item else { - break; - }; - match item { - Ok(chunk) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - if !first_chunk_telemetry_emitted { - match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - first_chunk_telemetry_emitted = true; - } - upstream_bytes += chunk.len() as u64; - observe_stream_chunk( - &mut stream_terminal_observer, - &normalized_observer_context, - private_stream_normalizer.as_mut(), - &mut observer_buffered, - chunk.as_ref(), - ); - match encode_data_frame(&chunk) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - } - Err(_err) => { - let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - status_code, - upstream_bytes, - error_category = "browser_body_read_failed", - "upstream body stream read error" - ); - match encode_error_frame(message) { - Ok(frame) => yield Ok(frame), - Err(encode_err) => { - yield Err(encode_err); - return; - } - } - break; - } + }; + let Some(item) = item else { break }; + match item { + Ok(chunk) => { + stream_completion.observe_chunk(&chunk); + if ttfb_ms.is_none() { + ttfb_ms = Some(started_at.elapsed().as_millis() as u64); } - } - } - DirectUpstreamResponse::LocalTunnel(mut response) => loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - response.next_chunk(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - match encode_first_byte_timeout_frame(timeout) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - break; - } - } - } else { - response.next_chunk().await - }; - match item { - Ok(Some(chunk)) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - if !first_chunk_telemetry_emitted { - match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) { - Ok(frame) => yield Ok(frame), - Err(err) => { - yield Err(err); - return; - } - } - first_chunk_telemetry_emitted = true; - } - upstream_bytes += chunk.len() as u64; - observe_stream_chunk( - &mut stream_terminal_observer, - &normalized_observer_context, - private_stream_normalizer.as_mut(), - &mut observer_buffered, - chunk.as_ref(), - ); - match encode_data_frame(&chunk) { + if !first_chunk_telemetry_emitted { + match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) { Ok(frame) => yield Ok(frame), Err(err) => { yield Err(err); return; } } + first_chunk_telemetry_emitted = true; } - Ok(None) => break, - Err(_message) => { - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - status_code, - upstream_bytes, - error_category = "tunnel_body_read_failed", - "upstream body stream read error" - ); - match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) { - Ok(frame) => yield Ok(frame), - Err(encode_err) => { - yield Err(encode_err); - return; - } + upstream_bytes += chunk.len() as u64; + observe_stream_chunk( + &mut stream_terminal_observer, + &normalized_observer_context, + private_stream_normalizer.as_mut(), + &mut observer_buffered, + chunk.as_ref(), + ); + match encode_data_frame(&chunk) { + Ok(frame) => yield Ok(frame), + Err(err) => { + yield Err(err); + return; } - break; } } + Err(_) => { + warn!( + event_name = "stream_pump_body_read_error", + log_type = "ops", + status_code, + upstream_bytes, + error_category = upstream_error_category, + "upstream body stream read error" + ); + match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) { + Ok(frame) => yield Ok(frame), + Err(err) => yield Err(err), + } + break; + } } } } @@ -704,6 +495,22 @@ fn encode_first_byte_timeout_frame(timeout: Duration) -> Result }) } +fn encode_idle_timeout_frame(timeout: Duration) -> Result { + encode_stream_frame_ndjson(&StreamFrame { + frame_type: StreamFrameType::Error, + payload: StreamFramePayload::Error { + error: ExecutionError { + kind: ExecutionErrorKind::ReadTimeout, + phase: ExecutionPhase::StreamRead, + message: stream_idle_timeout_message(timeout), + upstream_status: None, + retryable: true, + failover_recommended: true, + }, + }, + }) +} + async fn await_stream_first_byte( future: F, started_at: Instant, @@ -737,6 +544,7 @@ struct BufferedUpstreamBodyError { ttfb_ms: Option, upstream_bytes: u64, first_byte_timeout: Option, + idle_timeout: Option, } fn append_buffered_upstream_body_chunk( @@ -752,6 +560,7 @@ fn append_buffered_upstream_body_chunk( ttfb_ms, upstream_bytes: *upstream_bytes, first_byte_timeout: None, + idle_timeout: None, } }) } @@ -816,16 +625,52 @@ fn should_buffer_non_stream_response( } async fn buffer_non_sse_upstream_body( - mut prefetched_body: VecDeque>, + prefetched_body: VecDeque>, response: DirectUpstreamResponse, started_at: Instant, stream_first_byte_timeout: Option, + stream_idle_timeout: Option, ) -> Result { let mut body_bytes = Vec::new(); let mut upstream_bytes = 0u64; let mut ttfb_ms = None; - - while let Some(item) = prefetched_body.pop_front() { + let upstream_error_category = upstream_stream_error_category(&response); + let mut bytes_stream = direct_upstream_response_byte_stream(prefetched_body, response); + loop { + let item = if ttfb_ms.is_none() { + match await_stream_first_byte( + bytes_stream.next(), + started_at, + stream_first_byte_timeout, + ) + .await + { + Ok(item) => item, + Err(timeout) => { + return Err(BufferedUpstreamBodyError { + message: stream_first_byte_timeout_message(timeout), + ttfb_ms, + upstream_bytes, + first_byte_timeout: Some(timeout), + idle_timeout: None, + }) + } + } + } else { + match await_stream_idle_read(bytes_stream.next(), stream_idle_timeout).await { + Ok(item) => item, + Err(timeout) => { + return Err(BufferedUpstreamBodyError { + message: stream_idle_timeout_message(timeout), + ttfb_ms, + upstream_bytes, + first_byte_timeout: None, + idle_timeout: Some(timeout), + }) + } + } + }; + let Some(item) = item else { break }; match item { Ok(chunk) => { if ttfb_ms.is_none() { @@ -838,246 +683,24 @@ async fn buffer_non_sse_upstream_body( &mut upstream_bytes, )?; } - Err(_message) => { + Err(_) => { + warn!( + event_name = "stream_pump_body_read_error", + log_type = "ops", + upstream_bytes, + error_category = upstream_error_category, + "upstream body stream read error" + ); return Err(BufferedUpstreamBodyError { message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(), ttfb_ms, upstream_bytes, first_byte_timeout: None, + idle_timeout: None, }); } } } - - match response { - DirectUpstreamResponse::Reqwest(response) => { - let mut bytes_stream = response.bytes_stream(); - loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - bytes_stream.next(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - return Err(BufferedUpstreamBodyError { - message: stream_first_byte_timeout_message(timeout), - ttfb_ms, - upstream_bytes, - first_byte_timeout: Some(timeout), - }); - } - } - } else { - bytes_stream.next().await - }; - let Some(item) = item else { - break; - }; - match item { - Ok(chunk) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - append_buffered_upstream_body_chunk( - &mut body_bytes, - &chunk, - ttfb_ms, - &mut upstream_bytes, - )?; - } - Err(_err) => { - let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - upstream_bytes, - error_category = "reqwest_body_read_failed", - "upstream body stream read error" - ); - return Err(BufferedUpstreamBodyError { - message, - ttfb_ms, - upstream_bytes, - first_byte_timeout: None, - }); - } - } - } - } - DirectUpstreamResponse::HyperH2c(response) => { - let mut bytes_stream = response.into_body().into_data_stream(); - loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - bytes_stream.next(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - return Err(BufferedUpstreamBodyError { - message: stream_first_byte_timeout_message(timeout), - ttfb_ms, - upstream_bytes, - first_byte_timeout: Some(timeout), - }); - } - } - } else { - bytes_stream.next().await - }; - let Some(item) = item else { - break; - }; - match item { - Ok(chunk) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - append_buffered_upstream_body_chunk( - &mut body_bytes, - &chunk, - ttfb_ms, - &mut upstream_bytes, - )?; - } - Err(_err) => { - let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - upstream_bytes, - error_category = "hyper_body_read_failed", - "upstream body stream read error" - ); - return Err(BufferedUpstreamBodyError { - message, - ttfb_ms, - upstream_bytes, - first_byte_timeout: None, - }); - } - } - } - } - DirectUpstreamResponse::BrowserWreq(response) => { - let mut bytes_stream = response.bytes_stream(); - loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - bytes_stream.next(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - return Err(BufferedUpstreamBodyError { - message: stream_first_byte_timeout_message(timeout), - ttfb_ms, - upstream_bytes, - first_byte_timeout: Some(timeout), - }); - } - } - } else { - bytes_stream.next().await - }; - let Some(item) = item else { - break; - }; - match item { - Ok(chunk) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - append_buffered_upstream_body_chunk( - &mut body_bytes, - &chunk, - ttfb_ms, - &mut upstream_bytes, - )?; - } - Err(_err) => { - let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(); - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - upstream_bytes, - error_category = "browser_body_read_failed", - "upstream body stream read error" - ); - return Err(BufferedUpstreamBodyError { - message, - ttfb_ms, - upstream_bytes, - first_byte_timeout: None, - }); - } - } - } - } - DirectUpstreamResponse::LocalTunnel(mut response) => loop { - let item = if ttfb_ms.is_none() { - match await_stream_first_byte( - response.next_chunk(), - started_at, - stream_first_byte_timeout, - ) - .await - { - Ok(item) => item, - Err(timeout) => { - return Err(BufferedUpstreamBodyError { - message: stream_first_byte_timeout_message(timeout), - ttfb_ms, - upstream_bytes, - first_byte_timeout: Some(timeout), - }); - } - } - } else { - response.next_chunk().await - }; - match item { - Ok(Some(chunk)) => { - if ttfb_ms.is_none() { - ttfb_ms = Some(started_at.elapsed().as_millis() as u64); - } - append_buffered_upstream_body_chunk( - &mut body_bytes, - &chunk, - ttfb_ms, - &mut upstream_bytes, - )?; - } - Ok(None) => break, - Err(_message) => { - warn!( - event_name = "stream_pump_body_read_error", - log_type = "ops", - upstream_bytes, - error_category = "tunnel_body_read_failed", - "upstream body stream read error" - ); - return Err(BufferedUpstreamBodyError { - message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(), - ttfb_ms, - upstream_bytes, - first_byte_timeout: None, - }); - } - } - }, - } - Ok(BufferedUpstreamBody { body_bytes, ttfb_ms, @@ -1605,6 +1228,89 @@ mod tests { assert_eq!(error.get("failover_recommended"), Some(&Value::Bool(true))); } + #[tokio::test] + async fn direct_execution_frame_stream_enforces_idle_timeout_after_first_byte() { + for (content_type, first_chunk, expect_timeout, provider_format) in [ + ("text/event-stream", "data: hello\n\n", true, "openai:chat"), + ("application/json", "{\"message\":", true, "openai:chat"), + ("text/event-stream", "data: [DONE]\n\n", false, "openai:chat"), + ("text/event-stream", "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\"}}\n\n", false, "openai:responses"), + ("text/event-stream", "event: response.incomplete\ndata: {\"type\":\"response.incomplete\",\"response\":{}}\n\n", true, "openai:responses"), + ] { + let listener = crate::test_support::bind_loopback_listener().await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = [0_u8; 4096]; + socket.read(&mut request).await.unwrap(); + let response = if content_type == "application/json" { + format!("HTTP/1.1 200 OK\r\ncontent-type: {content_type}\r\ncontent-length: 1024\r\n\r\n{first_chunk}") + } else { + format!( + "HTTP/1.1 200 OK\r\ncontent-type: {content_type}\r\ntransfer-encoding: chunked\r\n\r\n{:x}\r\n{first_chunk}\r\n", + first_chunk.len(), + ) + }; + socket.write_all(response.as_bytes()).await.unwrap(); + socket.flush().await.unwrap(); + tokio::time::sleep(Duration::from_secs(5)).await; + }); + let execution = DirectSyncExecutionRuntime::new() + .execute_stream(&ExecutionPlan { + request_id: "req-stream-idle-timeout".into(), + candidate_id: Some("cand-stream-idle-timeout".into()), + provider_name: Some("openai".into()), + provider_id: "prov-1".into(), + endpoint_id: "ep-1".into(), + key_id: "key-1".into(), + method: "POST".into(), + url: format!("http://{addr}/chat"), + headers: BTreeMap::from([("content-type".into(), "application/json".into())]), + content_type: Some("application/json".into()), + content_encoding: None, + body: RequestBody::from_json(serde_json::json!({"stream": true})), + stream: true, + client_api_format: "openai:chat".into(), + provider_api_format: provider_format.into(), + model_name: Some("gpt-5".into()), + proxy: None, + transport_profile: None, + timeouts: Some(ExecutionTimeouts { + first_byte_ms: Some(1_000), + read_ms: Some(10), + ..ExecutionTimeouts::default() + }), + }) + .await + .expect("stream response headers"); + let frames = tokio::time::timeout( + Duration::from_secs(1), + build_direct_execution_frame_stream(execution).collect::>(), + ) + .await; + server.abort(); + let frames = frames + .expect("idle timeout must terminate both SSE and buffered JSON") + .into_iter() + .map(|line| serde_json::from_slice::(&line.unwrap()).unwrap()) + .collect::>(); + let errors = frames + .iter() + .filter(|frame| frame["type"] == "error") + .collect::>(); + assert_eq!(errors.len(), usize::from(expect_timeout)); + if expect_timeout { + assert_eq!(errors[0]["payload"]["error"]["kind"], "read_timeout"); + assert_eq!(errors[0]["payload"]["error"]["phase"], "stream_read"); + } + assert!(frames.iter().any(|frame| frame["type"] == "eof")); + assert_eq!( + frames.iter().any(|frame| frame["type"] == "data"), + content_type == "text/event-stream" + ); + } + } + #[tokio::test] async fn direct_execution_frame_stream_emits_telemetry_before_first_data_frame() { let listener = crate::test_support::bind_loopback_listener() diff --git a/apps/aether-gateway/src/execution_runtime/stream_read_timeout.rs b/apps/aether-gateway/src/execution_runtime/stream_read_timeout.rs new file mode 100644 index 000000000..f19de10c3 --- /dev/null +++ b/apps/aether-gateway/src/execution_runtime/stream_read_timeout.rs @@ -0,0 +1,187 @@ +use std::future::Future; +use std::time::Duration; + +use aether_contracts::ExecutionPlan; +use axum::body::Bytes; +use futures_util::{Stream, StreamExt}; + +const STREAM_IDLE_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS"; +const DEFAULT_STREAM_IDLE_TIMEOUT_MS: u64 = 300_000; + +pub(crate) fn resolve_stream_idle_timeout(plan: &ExecutionPlan) -> Option { + if !plan.stream { + return None; + } + let configured = std::env::var(STREAM_IDLE_TIMEOUT_MS_ENV).ok(); + stream_idle_timeout_from_config( + plan.timeouts.as_ref().and_then(|timeouts| timeouts.read_ms), + configured.as_deref(), + ) +} + +fn stream_idle_timeout_from_config( + read_ms: Option, + configured: Option<&str>, +) -> Option { + let timeout_ms = read_ms + .or_else(|| configured.and_then(|value| value.trim().parse::().ok())) + .unwrap_or(DEFAULT_STREAM_IDLE_TIMEOUT_MS); + // Zero explicitly disables the idle limit for providers with long silent reasoning phases. + (timeout_ms > 0).then(|| Duration::from_millis(timeout_ms)) +} + +pub(crate) fn stream_idle_timeout_message(timeout: Duration) -> String { + format!( + "provider stream idle read timeout after {} ms", + timeout.as_millis() + ) +} + +pub(crate) async fn await_stream_idle_read( + future: impl Future, + timeout: Option, +) -> Result { + match timeout { + Some(timeout) => tokio::time::timeout(timeout, future) + .await + .map_err(|_| timeout), + None => Ok(future.await), + } +} + +pub(crate) fn skip_empty_upstream_chunks( + upstream: impl Stream> + Send + 'static, +) -> impl Stream> + Send { + async_stream::stream! { + tokio::pin!(upstream); + while let Some(item) = upstream.next().await { + match item { + Ok(chunk) if chunk.is_empty() => { + // Empty frames are not progress; yield so an always-ready source cannot + // monopolize the executor or prevent its enclosing timeout from firing. + tokio::task::yield_now().await; + } + item => yield item, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; + + #[test] + fn idle_timeout_configuration_preserves_provider_override_and_explicit_disable() { + assert_eq!( + stream_idle_timeout_from_config(None, None), + Some(Duration::from_secs(300)) + ); + assert_eq!( + stream_idle_timeout_from_config(None, Some(" invalid ")), + Some(Duration::from_secs(300)) + ); + assert_eq!( + stream_idle_timeout_from_config(None, Some(" 600000 ")), + Some(Duration::from_secs(600)) + ); + assert_eq!( + stream_idle_timeout_from_config(Some(120_000), Some("600000")), + Some(Duration::from_secs(120)) + ); + assert_eq!( + stream_idle_timeout_from_config(Some(0), Some("600000")), + None + ); + assert_eq!(stream_idle_timeout_from_config(None, Some("0")), None); + } + + #[tokio::test] + async fn idle_timeout_cancels_the_pending_upstream_read() { + struct DropMarker(Arc); + impl Drop for DropMarker { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + let dropped = Arc::new(AtomicBool::new(false)); + let marker = DropMarker(Arc::clone(&dropped)); + let outcome = await_stream_idle_read( + async move { + let _marker = marker; + std::future::pending::<()>().await; + }, + Some(Duration::from_millis(5)), + ) + .await; + assert_eq!(outcome, Err(Duration::from_millis(5))); + assert!(dropped.load(Ordering::SeqCst)); + } + + #[tokio::test] + async fn idle_timeout_allows_progressing_stream_to_outlive_one_timeout() { + for _ in 0..3 { + assert_eq!( + await_stream_idle_read( + async { + tokio::time::sleep(Duration::from_millis(10)).await; + 1 + }, + Some(Duration::from_millis(25)) + ) + .await, + Ok(1) + ); + } + assert_eq!( + await_stream_idle_read( + async { + tokio::time::sleep(Duration::from_millis(30)).await; + 2 + }, + None + ) + .await, + Ok(2) + ); + } + + #[tokio::test] + async fn empty_upstream_chunks_do_not_reset_idle_timeout() { + let upstream = futures_util::stream::repeat(Ok::<_, ()>(Bytes::new())); + let filtered = skip_empty_upstream_chunks(upstream); + tokio::pin!(filtered); + let outcome = tokio::time::timeout( + Duration::from_secs(1), + await_stream_idle_read(filtered.next(), Some(Duration::from_millis(5))), + ) + .await + .expect("empty ready chunks must yield to the idle timer"); + assert_eq!(outcome, Err(Duration::from_millis(5))); + } + + #[tokio::test] + async fn downstream_keepalive_ticks_do_not_reset_pending_upstream_idle_timeout() { + let read = await_stream_idle_read( + std::future::pending::<()>(), + Some(Duration::from_millis(30)), + ); + tokio::pin!(read); + let mut keepalive = tokio::time::interval(Duration::from_millis(2)); + let mut ticks = 0; + loop { + tokio::select! { + result = &mut read => { + assert_eq!(result, Err(Duration::from_millis(30))); + assert!(ticks > 0); + break; + } + _ = keepalive.tick() => { ticks += 1; } + } + } + } +} diff --git a/apps/aether-gateway/src/execution_runtime/sync/execution.rs b/apps/aether-gateway/src/execution_runtime/sync/execution.rs index fe72956be..37abb8c18 100644 --- a/apps/aether-gateway/src/execution_runtime/sync/execution.rs +++ b/apps/aether-gateway/src/execution_runtime/sync/execution.rs @@ -270,7 +270,9 @@ impl Drop for SyncAttemptTerminalGuard { let candidate_started_unix_ms = self.candidate_started_unix_ms; let candidate_started_at = self.candidate_started_at; if let Ok(handle) = tokio::runtime::Handle::try_current() { + let usage_producer = state.usage_runtime.track_producer(); handle.spawn(async move { + let _usage_producer = usage_producer; record_sync_attempt_forced_terminal_state( state, plan, diff --git a/apps/aether-gateway/src/execution_runtime/transport.rs b/apps/aether-gateway/src/execution_runtime/transport.rs index 9247fd5f0..5cfeb3cf4 100644 --- a/apps/aether-gateway/src/execution_runtime/transport.rs +++ b/apps/aether-gateway/src/execution_runtime/transport.rs @@ -36,7 +36,7 @@ use hyper::body::Incoming as HyperIncomingBody; use hyper::client::conn::http2::SendRequest as HyperH2cSendRequest; use hyper_util::client::legacy::connect::HttpConnector; use hyper_util::client::legacy::Client as HyperLegacyClient; -use hyper_util::rt::{TokioExecutor, TokioIo}; +use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; use reqwest::redirect::Policy; use serde::Serialize; @@ -50,6 +50,7 @@ use tokio::sync::OnceCell as TokioOnceCell; use crate::ai_serving::api::extract_provider_private_stream_error_body; #[cfg(test)] use crate::execution_runtime::remote_compat::execute_sync_plan_via_remote_execution_runtime; +use crate::execution_runtime::stream_read_timeout::resolve_stream_idle_timeout; use crate::execution_runtime::windsurf::maybe_execute_windsurf_sync; use crate::frontdoor_loop_guard::{ configured_gateway_frontdoor_base_url, gateway_frontdoor_self_loop_guard_error, @@ -109,7 +110,7 @@ const DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_PREWARM_SYNC_CLIENTS"; const DEFAULT_H2_TARGET_STREAMS_PER_CLIENT: usize = 8; const DEFAULT_HTTP1_TARGET_STREAMS_PER_CLIENT: usize = 512; -const DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: usize = 512; +const DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: usize = 32; const DEFAULT_DIRECT_H2C_TARGET_STREAMS_PER_CLIENT: usize = 128; const DEFAULT_DIRECT_H2C_SENDER_SELECT_WINDOW: usize = 4; const MAX_DIRECT_H2C_DRIVER_RUNTIME_THREADS: usize = 16; @@ -382,6 +383,7 @@ static DIRECT_H2C_SENDER_CACHE: LazyLock< static DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: LazyLock = LazyLock::new(|| { env_positive_usize(DIRECT_H2C_POOL_MAX_IDLE_PER_HOST_ENV) .unwrap_or(DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST) + .min(1024) }); static DIRECT_H2C_SENDER_SELECT_WINDOW: LazyLock = LazyLock::new(|| { @@ -1198,6 +1200,42 @@ pub(crate) enum DirectUpstreamResponse { LocalTunnel(tunnel::DirectRelayResponse), } +pub(crate) fn direct_upstream_response_byte_stream( + prefetched_body: VecDeque>, + response: DirectUpstreamResponse, +) -> futures_util::stream::BoxStream<'static, Result> { + let response_stream = match response { + DirectUpstreamResponse::Reqwest(response) => response + .bytes_stream() + .map(|item| item.map_err(|err| format_upstream_request_error(&err))) + .boxed(), + DirectUpstreamResponse::HyperH2c(response) => response + .into_body() + .into_data_stream() + .map(|item| item.map_err(|err| format_hyper_error_chain(&err))) + .boxed(), + DirectUpstreamResponse::BrowserWreq(response) => response + .bytes_stream() + .map(|item| item.map_err(|err| format_wreq_upstream_request_error(&err))) + .boxed(), + DirectUpstreamResponse::LocalTunnel(mut response) => async_stream::stream! { + loop { + match response.next_chunk().await { + Ok(Some(chunk)) => yield Ok(chunk), + Ok(None) => break, + Err(err) => { + yield Err(err); + break; + } + } + } + } + .boxed(), + }; + let upstream = futures_util::stream::iter(prefetched_body).chain(response_stream); + crate::execution_runtime::stream_read_timeout::skip_empty_upstream_chunks(upstream).boxed() +} + pub(crate) struct DirectUpstreamStreamExecution { pub(crate) request_id: String, pub(crate) candidate_id: Option, @@ -1214,6 +1252,7 @@ pub(crate) struct DirectUpstreamStreamExecution { pub(crate) started_at: Instant, pub(crate) response_observation: ExecutionResponseObservation, pub(crate) stream_first_byte_timeout: Option, + pub(crate) stream_idle_timeout: Option, pub(crate) upstream_target_permit: Option, } @@ -1347,6 +1386,7 @@ impl DirectSyncExecutionRuntime { request_order_id, }, stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan), + stream_idle_timeout: resolve_stream_idle_timeout(plan), upstream_target_permit: None, }) } @@ -1494,6 +1534,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel( request_order_id, }, stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan), + stream_idle_timeout: resolve_stream_idle_timeout(plan), upstream_target_permit: None, })) } @@ -2593,6 +2634,8 @@ fn build_direct_h2c_client_from_cache_key( builder.http2_only(true); builder.http2_adaptive_window(true); builder.pool_max_idle_per_host(cache_key.pool_max_idle_per_host); + builder.pool_timer(TokioTimer::new()); + builder.pool_idle_timeout(Duration::from_millis(upstream_pool_idle_timeout_ms())); builder.build(connector) } @@ -2854,8 +2897,18 @@ async fn send_via_browser_wreq_transport( let profile = plan.transport_profile.as_ref().ok_or_else(|| { ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new()) })?; + let mut client_timeouts = plan.timeouts.clone(); + if plan.stream { + if let Some(timeouts) = client_timeouts.as_mut() { + // Streamed responses use the shared idle reader; sync collectors retain + // their existing client read timeout. Zero explicitly disables either. + if apply_request_total_timeout || timeouts.read_ms == Some(0) { + timeouts.read_ms = None; + } + } + } let client = build_browser_wreq_client( - plan.timeouts.as_ref(), + client_timeouts.as_ref(), plan.proxy.as_ref(), profile, transport_controls, @@ -4214,6 +4267,7 @@ fn build_direct_reqwest_client_from_cache_key( &HttpClientConfig { connect_timeout_ms: cache_key.connect_timeout_ms, pool_max_idle_per_host: Some(direct_reqwest_pool_max_idle_per_host()), + pool_idle_timeout_ms: Some(upstream_pool_idle_timeout_ms()), ..HttpClientConfig::default() }, ); @@ -4233,12 +4287,22 @@ fn build_direct_reqwest_client_from_cache_key( } fn direct_reqwest_pool_max_idle_per_host() -> usize { - const DEFAULT_MAX_IDLE_PER_HOST: usize = 1024; + const DEFAULT_MAX_IDLE_PER_HOST: usize = 32; std::env::var("AETHER_GATEWAY_UPSTREAM_POOL_MAX_IDLE_PER_HOST") .ok() .and_then(|value| value.trim().parse::().ok()) .filter(|value| *value > 0) .unwrap_or(DEFAULT_MAX_IDLE_PER_HOST) + .min(1024) +} + +fn upstream_pool_idle_timeout_ms() -> u64 { + std::env::var("AETHER_GATEWAY_UPSTREAM_POOL_IDLE_TIMEOUT_MS") + .ok() + .and_then(|value| value.trim().parse::().ok()) + .filter(|value| *value > 0) + .unwrap_or(15_000) + .min(300_000) } pub(crate) fn direct_reqwest_client_cache_metric_samples() -> Vec { @@ -4585,7 +4649,11 @@ pub(crate) fn build_browser_wreq_client( ) -> Result { let emulation = browser_wreq_emulation_from_profile(transport_profile)?; let proxy_url = resolve_proxy_url(proxy)?; - let mut builder = wreq::Client::builder().no_proxy().emulation(emulation); + let mut builder = wreq::Client::builder() + .no_proxy() + .emulation(emulation) + .pool_max_idle_per_host(direct_reqwest_pool_max_idle_per_host()) + .pool_idle_timeout(Duration::from_millis(upstream_pool_idle_timeout_ms())); if proxy_url.is_none() { builder = builder.dns_resolver(ExecutionSafeDnsResolver); } diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs index 71ac5890c..71c3711a6 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/keys.rs @@ -40,20 +40,6 @@ pub(super) fn pool_stream_timeout_key(provider_id: &str, key_id: &str) -> String format!("ap:{provider_id}:stream_timeout:{key_id}") } -pub(super) fn parse_pool_cost_member(member: &str) -> u64 { - member - .rsplit_once(':') - .and_then(|(_, suffix)| suffix.parse::().ok()) - .unwrap_or(0) -} - -pub(super) fn parse_pool_latency_member(member: &str) -> u64 { - member - .rsplit_once(':') - .and_then(|(_, suffix)| suffix.parse::().ok()) - .unwrap_or(0) -} - pub(super) fn pool_cooldown_keys(provider_id: &str, key_ids: &[String]) -> Vec { key_ids .iter() diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/mod.rs b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/mod.rs index c0fe997bc..0862e93eb 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/mod.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/mod.rs @@ -12,7 +12,8 @@ pub(crate) use self::mutations::{ pub(crate) use self::reads::{ read_admin_provider_pool_cooldown_count, read_admin_provider_pool_cooldown_counts, read_admin_provider_pool_cooldown_key_ids, read_admin_provider_pool_key_cooldown_reason, - read_admin_provider_pool_runtime_state, + read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state, + read_provider_pool_sticky_bound_key_id, }; pub(crate) use self::status::build_admin_provider_pool_status_payload; pub(crate) use self::writes::{ diff --git a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/reads.rs b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/reads.rs index 89b424157..d1e4b2dd4 100644 --- a/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/reads.rs +++ b/apps/aether-gateway/src/handlers/admin/provider/pool/runtime/reads.rs @@ -1,7 +1,6 @@ use super::keys::{ - parse_pool_cost_member, parse_pool_latency_member, pool_cooldown_index_key, pool_cooldown_key, - pool_cooldown_keys, pool_cost_keys, pool_latency_keys, pool_lru_key, pool_sticky_key, - pool_sticky_pattern, + pool_cooldown_index_key, pool_cooldown_key, pool_cooldown_keys, pool_cost_keys, + pool_latency_keys, pool_lru_key, pool_sticky_key, pool_sticky_pattern, }; use crate::handlers::admin::provider::pool::config::admin_provider_pool_cache_affinity_enabled; use crate::handlers::admin::provider::shared::support::{ @@ -12,7 +11,8 @@ use crate::maintenance::PoolQuotaProbeWorkerConfig; use crate::provider_pool_demand::{ provider_pool_burst_pending, read_provider_pool_demand_snapshot, }; -use aether_runtime_state::{DataLayerError, RuntimeState}; +use aether_pool_core::{normalize_enabled_pool_presets, PoolSchedulingPreset}; +use aether_runtime_state::{DataLayerError, RuntimeState, ScoreWindowU64Stats}; use futures_util::future::join_all; use std::collections::{BTreeMap, BTreeSet}; use std::time::{SystemTime, UNIX_EPOCH}; @@ -48,6 +48,103 @@ fn bounded_runtime_window_metric_key_ids(key_ids: &[String], limit: usize) -> &[ &key_ids[..end] } +#[derive(Clone, Copy, PartialEq, Eq)] +enum PoolRuntimeReadPurpose { + Admin, + Scheduling, +} + +fn scheduling_window_metrics(pool_config: &AdminProviderPoolConfig) -> (bool, bool) { + let presets = pool_config + .scheduling_presets + .iter() + .map(|preset| PoolSchedulingPreset { + preset: preset.preset.clone(), + enabled: preset.enabled, + mode: preset.mode.clone(), + }) + .collect::>(); + let active = normalize_enabled_pool_presets(&presets); + let cost = pool_config.cost_limit_per_key_tokens.is_some() + || active + .iter() + .any(|preset| matches!(preset.as_str(), "cost_first" | "quota_balanced")); + let latency = active.iter().any(|preset| preset == "latency_first"); + (cost, latency) +} + +async fn read_window_stats( + runtime: &RuntimeState, + keys: &[String], + min_score: f64, +) -> Vec { + let aggregates = match runtime.score_window_u64_stats_by_min(keys, min_score).await { + Ok(values) => values, + Err(err) => { + warn!( + "gateway provider pool: bounded window aggregation failed, using exact range reads: {err:?}" + ); + vec![None; keys.len()] + } + }; + join_all(keys.iter().zip(aggregates).map(|(key, stats)| async move { + match stats { + Some(stats) => stats, + // Large windows and failed aggregation retain the original exact + // read. A missing aggregate must never be treated as zero cost. + None => { + let members = runtime + .score_range_by_min(key, min_score) + .await + .unwrap_or_default(); + ScoreWindowU64Stats::from_members(members.iter().map(String::as_str)) + } + } + })) + .await +} + +pub(crate) async fn read_provider_pool_sticky_bound_key_id( + runtime: &RuntimeState, + provider_id: &str, + pool_config: &AdminProviderPoolConfig, + sticky_session_token: Option<&str>, +) -> Option { + if pool_config.sticky_session_ttl_seconds == 0 + || !admin_provider_pool_cache_affinity_enabled(pool_config) + { + return None; + } + let sticky_session_token = sticky_session_token + .map(str::trim) + .filter(|value| !value.is_empty())?; + let sticky_key = pool_sticky_key(provider_id, sticky_session_token); + let bound_key_id = runtime.kv_get(&sticky_key).await.ok().flatten()?; + let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id); + match runtime.kv_exists(&cooldown_key).await { + Ok(false) => { + let _ = runtime + .key_expire( + &sticky_key, + std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds), + ) + .await; + Some(bound_key_id) + } + Ok(true) => { + let _ = runtime.kv_delete(&sticky_key).await; + None + } + Err(err) => { + warn!( + "gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}", + err + ); + Some(bound_key_id) + } + } +} + pub(crate) async fn read_admin_provider_pool_cooldown_counts( runtime: &RuntimeState, provider_ids: &[String], @@ -71,9 +168,52 @@ pub(crate) async fn read_admin_provider_pool_runtime_state( pool_config: &AdminProviderPoolConfig, sticky_session_token: Option<&str>, ) -> AdminProviderPoolRuntimeState { + read_provider_pool_runtime_state( + runtime, + provider_id, + key_ids, + pool_config, + sticky_session_token, + PoolRuntimeReadPurpose::Admin, + ) + .await +} + +pub(crate) async fn read_provider_pool_scheduling_runtime_state( + runtime: &RuntimeState, + provider_id: &str, + key_ids: &[String], + pool_config: &AdminProviderPoolConfig, + sticky_session_token: Option<&str>, +) -> AdminProviderPoolRuntimeState { + read_provider_pool_runtime_state( + runtime, + provider_id, + key_ids, + pool_config, + sticky_session_token, + PoolRuntimeReadPurpose::Scheduling, + ) + .await +} + +async fn read_provider_pool_runtime_state( + runtime: &RuntimeState, + provider_id: &str, + key_ids: &[String], + pool_config: &AdminProviderPoolConfig, + sticky_session_token: Option<&str>, + purpose: PoolRuntimeReadPurpose, +) -> AdminProviderPoolRuntimeState { + let include_admin_metrics = purpose == PoolRuntimeReadPurpose::Admin; let mut state = AdminProviderPoolRuntimeState::default(); let cooldown_keys = pool_cooldown_keys(provider_id, key_ids); - let metric_key_limit = pool_runtime_window_metric_key_limit(); + let metric_key_limit = if include_admin_metrics { + pool_runtime_window_metric_key_limit() + } else { + key_ids.len() + }; + // The admin display cap must not hide a candidate's strict cost limit. let metric_key_ids = bounded_runtime_window_metric_key_ids(key_ids, metric_key_limit); if metric_key_ids.len() < key_ids.len() { info!( @@ -86,44 +226,33 @@ pub(crate) async fn read_admin_provider_pool_runtime_state( "gateway limited admin pool runtime cost/latency window reads" ); } - let cost_keys = pool_cost_keys(provider_id, metric_key_ids); - let latency_keys = pool_latency_keys(provider_id, metric_key_ids); + let (load_cost, load_latency) = if include_admin_metrics { + (true, true) + } else { + scheduling_window_metrics(pool_config) + }; + let cost_keys = if load_cost { + pool_cost_keys(provider_id, metric_key_ids) + } else { + Vec::new() + }; + let latency_keys = if load_latency { + pool_latency_keys(provider_id, metric_key_ids) + } else { + Vec::new() + }; let sticky_sessions_enabled = pool_config.sticky_session_ttl_seconds > 0 && admin_provider_pool_cache_affinity_enabled(pool_config); - if let Some(sticky_session_token) = sticky_session_token - .map(str::trim) - .filter(|value| !value.is_empty()) - .filter(|_| sticky_sessions_enabled) - { - let sticky_key = pool_sticky_key(provider_id, sticky_session_token); - if let Ok(Some(bound_key_id)) = runtime.kv_get(&sticky_key).await { - let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id); - match runtime.kv_exists(&cooldown_key).await { - Ok(false) => { - let _ = runtime - .key_expire( - &sticky_key, - std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds), - ) - .await; - state.sticky_bound_key_id = Some(bound_key_id); - } - Ok(true) => { - let _ = runtime.kv_delete(&sticky_key).await; - } - Err(err) => { - warn!( - "gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}", - err - ); - state.sticky_bound_key_id = Some(bound_key_id); - } - } - } - } + state.sticky_bound_key_id = read_provider_pool_sticky_bound_key_id( + runtime, + provider_id, + pool_config, + sticky_session_token, + ) + .await; - if sticky_sessions_enabled { + if include_admin_metrics && sticky_sessions_enabled { let sticky_keys = runtime .scan_keys(&pool_sticky_pattern(provider_id), 200) .await @@ -161,23 +290,25 @@ pub(crate) async fn read_admin_provider_pool_runtime_state( .unwrap_or_default(); } - let probe_config = PoolQuotaProbeWorkerConfig::from_env(); - let demand_snapshot = read_provider_pool_demand_snapshot( - runtime, - provider_id, - key_ids.len(), - probe_config.max_keys_per_provider, - ) - .await; - state.provider_in_flight = demand_snapshot.in_flight; - state.provider_ema_in_flight = demand_snapshot.ema_in_flight; - state.provider_desired_hot = if pool_config.probing_enabled { - demand_snapshot.desired_hot - } else { - 0 - }; - state.provider_burst_pending = - pool_config.probing_enabled && provider_pool_burst_pending(runtime, provider_id).await; + if include_admin_metrics || pool_config.probing_enabled { + let probe_config = PoolQuotaProbeWorkerConfig::from_env(); + let demand_snapshot = read_provider_pool_demand_snapshot( + runtime, + provider_id, + key_ids.len(), + probe_config.max_keys_per_provider, + ) + .await; + state.provider_in_flight = demand_snapshot.in_flight; + state.provider_ema_in_flight = demand_snapshot.ema_in_flight; + state.provider_desired_hot = if pool_config.probing_enabled { + demand_snapshot.desired_hot + } else { + 0 + }; + state.provider_burst_pending = + pool_config.probing_enabled && provider_pool_burst_pending(runtime, provider_id).await; + } if !cooldown_keys.is_empty() { let cooldown_reasons = runtime @@ -190,12 +321,14 @@ pub(crate) async fn read_admin_provider_pool_runtime_state( { if let Some(reason) = reason { state.cooldown_reason_by_key.insert(key_id.clone(), reason); - if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await { - if let Ok(ttl_seconds) = u64::try_from(ttl) { - if ttl_seconds > 0 { - state - .cooldown_ttl_by_key - .insert(key_id.clone(), ttl_seconds); + if include_admin_metrics { + if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await { + if let Ok(ttl_seconds) = u64::try_from(ttl) { + if ttl_seconds > 0 { + state + .cooldown_ttl_by_key + .insert(key_id.clone(), ttl_seconds); + } } } } @@ -205,42 +338,24 @@ pub(crate) async fn read_admin_provider_pool_runtime_state( let now = current_unix_secs(); let cost_window_start = now.saturating_sub(pool_config.cost_window_seconds) as f64; - let cost_results = join_all( - cost_keys - .iter() - .map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)), - ) - .await; - for (key_id, members) in metric_key_ids.iter().zip(cost_results) { - let total = members - .unwrap_or_default() - .iter() - .map(|member| parse_pool_cost_member(member)) - .sum::(); - if total > 0 { - state.cost_window_usage_by_key.insert(key_id.clone(), total); + let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64; + let (cost_results, latency_results) = tokio::join!( + read_window_stats(runtime, &cost_keys, cost_window_start), + read_window_stats(runtime, &latency_keys, latency_window_start), + ); + for (key_id, stats) in metric_key_ids.iter().zip(cost_results) { + if stats.sum > 0 { + state + .cost_window_usage_by_key + .insert(key_id.clone(), stats.sum); } } - let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64; - let latency_results = join_all( - latency_keys - .iter() - .map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)), - ) - .await; - for (key_id, members) in metric_key_ids.iter().zip(latency_results) { - let samples = members - .unwrap_or_default() - .iter() - .map(|member| parse_pool_latency_member(member)) - .filter(|value| *value > 0) - .collect::>(); - if samples.is_empty() { + for (key_id, stats) in metric_key_ids.iter().zip(latency_results) { + if stats.positive_count == 0 { continue; } - let total = samples.iter().sum::() as f64; - let average = total / samples.len() as f64; + let average = stats.sum as f64 / stats.positive_count as f64; if average.is_finite() && average >= 0.0 { state.latency_avg_ms_by_key.insert(key_id.clone(), average); } @@ -300,7 +415,358 @@ pub(crate) async fn read_admin_provider_pool_key_cooldown_reason( #[cfg(test)] mod tests { - use super::bounded_runtime_window_metric_key_ids; + use super::super::keys::{pool_cooldown_key, pool_cost_key, pool_latency_key, pool_sticky_key}; + use super::{ + bounded_runtime_window_metric_key_ids, current_unix_secs, + read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state, + read_provider_pool_sticky_bound_key_id, + }; + use crate::handlers::admin::provider::pool::config::admin_provider_pool_config_from_config_value; + use crate::handlers::admin::provider::shared::support::AdminProviderPoolConfig; + use aether_runtime_state::{MemoryRuntimeStateConfig, RedisClientConfig, RuntimeState}; + use aether_test_support::ManagedRedisServer; + use serde_json::json; + use std::time::Duration; + + fn config(value: serde_json::Value) -> AdminProviderPoolConfig { + admin_provider_pool_config_from_config_value(Some(&json!({ "pool_advanced": value }))) + .expect("pool config") + } + + async fn seed_window_metrics(runtime: &RuntimeState, provider_id: &str, key_id: &str) { + let now = current_unix_secs() as f64; + for (key, member, timestamp) in [ + (pool_cost_key(provider_id, key_id), "current:70", now), + (pool_cost_key(provider_id, key_id), "earlier:30", now - 1.0), + ( + pool_cost_key(provider_id, key_id), + "expired:999", + now - 20_000.0, + ), + (pool_latency_key(provider_id, key_id), "first:10", now), + ( + pool_latency_key(provider_id, key_id), + "second:30", + now - 1.0, + ), + ] { + runtime + .score_set(&key, member, timestamp) + .await + .expect("seed window"); + } + } + + async fn admin_command_count(runtime: &RuntimeState) -> u64 { + runtime + .redis_diagnostics() + .await + .expect("diagnostics") + .expect("Redis runtime") + .lanes + .into_iter() + .find(|lane| lane.lane == "admin") + .expect("admin lane") + .command_count + } + + #[tokio::test] + async fn scheduling_runtime_aggregates_bounded_windows_and_falls_back_for_large_windows() { + let redis = match ManagedRedisServer::start().await { + Ok(server) => server, + Err(err) if err.to_string().contains("No such file or directory") => { + eprintln!("skipping redis-backed scheduling runtime test: {err}"); + return; + } + Err(err) => panic!("start Redis: {err}"), + }; + let runtime = RuntimeState::redis( + RedisClientConfig { + url: redis.redis_url().to_string(), + key_prefix: Some("pool-window-aggregation-test".to_string()), + }, + Some(1_000), + ) + .await + .expect("runtime Redis"); + let keys = vec!["bounded".to_string(), "large".to_string()]; + let now = current_unix_secs() as f64; + for (key_id, count) in [(&keys[0], 512), (&keys[1], 2048)] { + let cost_key = pool_cost_key("pool", key_id); + for index in 0..count { + runtime + .score_set(&cost_key, &format!("{index}:100"), now) + .await + .expect("seed cost window"); + } + runtime + .score_set(&cost_key, "expired:9999999", now - 20_000.0) + .await + .expect("expired cost"); + for (member, score) in [("first:10", now), ("second:30", now), ("zero:0", now)] { + runtime + .score_set(&pool_latency_key("pool", key_id), member, score) + .await + .expect("seed latency"); + } + } + let pool_config = config(json!({ + "cost_limit_per_key_tokens": 50_000, + "cost_window_seconds": 600, + "latency_window_seconds": 600, + "scheduling_presets": [{"preset": "latency_first", "enabled": true}] + })); + let scheduled = read_provider_pool_scheduling_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + None, + ) + .await; + assert_eq!( + scheduled.cost_window_usage_by_key.get("bounded"), + Some(&51_200) + ); + assert_eq!( + scheduled.cost_window_usage_by_key.get("large"), + Some(&204_800) + ); + assert_eq!(scheduled.latency_avg_ms_by_key.get("bounded"), Some(&20.0)); + assert_eq!(scheduled.latency_avg_ms_by_key.get("large"), Some(&20.0)); + + runtime + .score_remove_by_score(&pool_cost_key("pool", "large"), f64::INFINITY) + .await + .expect("reset window"); + runtime + .score_set(&pool_cost_key("pool", "large"), "after-reset:75", now) + .await + .expect("post-reset cost"); + let reset = read_provider_pool_scheduling_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + None, + ) + .await; + assert_eq!(reset.cost_window_usage_by_key.get("large"), Some(&75)); + + runtime + .kv_set(&pool_cost_key("pool", "large"), "wrong-type", None) + .await + .expect("simulate invalid metric key"); + let partial = read_provider_pool_scheduling_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + None, + ) + .await; + assert_eq!( + partial.cost_window_usage_by_key.get("bounded"), + Some(&51_200), + "one failed aggregate must not discard another key's strict cost check" + ); + } + + #[tokio::test] + async fn scheduling_runtime_skips_admin_scan_and_unused_window_queries() { + let redis = match ManagedRedisServer::start().await { + Ok(server) => server, + Err(err) if err.to_string().contains("No such file or directory") => { + eprintln!("skipping redis-backed scheduling runtime test: {err}"); + return; + } + Err(err) => panic!("Redis server should start: {err}"), + }; + let runtime = RuntimeState::redis( + RedisClientConfig { + url: redis.redis_url().to_string(), + key_prefix: Some("scheduling-runtime-reads".to_string()), + }, + Some(2_000), + ) + .await + .expect("Redis runtime"); + let pool_config = config(json!({})); + let keys = vec!["ready".to_string(), "cooling".to_string()]; + seed_window_metrics(&runtime, "pool", "ready").await; + for session in ["current", "other"] { + runtime + .kv_set( + &pool_sticky_key("pool", session), + "ready".to_string(), + Some(Duration::from_secs(60)), + ) + .await + .expect("seed sticky session"); + } + runtime + .kv_set( + &pool_cooldown_key("pool", "cooling"), + "rate_limit".to_string(), + Some(Duration::from_secs(60)), + ) + .await + .expect("seed cooldown"); + + let before = admin_command_count(&runtime).await; + let scheduled = read_provider_pool_scheduling_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + Some("current"), + ) + .await; + let after = admin_command_count(&runtime).await; + + assert_eq!( + after - before, + 1, + "only the diagnostics INFO may use the admin lane" + ); + assert_eq!(scheduled.sticky_bound_key_id.as_deref(), Some("ready")); + assert_eq!( + scheduled + .cooldown_reason_by_key + .get("cooling") + .map(String::as_str), + Some("rate_limit") + ); + assert!(scheduled.cooldown_ttl_by_key.is_empty()); + assert!(scheduled.cost_window_usage_by_key.is_empty()); + assert!(scheduled.latency_avg_ms_by_key.is_empty()); + assert_eq!(scheduled.total_sticky_sessions, 0); + + let admin = read_admin_provider_pool_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + Some("current"), + ) + .await; + assert_eq!(admin.total_sticky_sessions, 2); + assert_eq!(admin.sticky_sessions_by_key.get("ready"), Some(&2)); + assert_eq!(admin.cost_window_usage_by_key.get("ready"), Some(&100)); + assert_eq!(admin.latency_avg_ms_by_key.get("ready"), Some(&20.0)); + assert!(admin + .cooldown_ttl_by_key + .get("cooling") + .is_some_and(|ttl| *ttl > 0)); + } + + #[tokio::test] + async fn scheduling_runtime_loads_only_metrics_used_by_enabled_strategies() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let keys = vec!["key".to_string()]; + seed_window_metrics(&runtime, "pool", "key").await; + for (value, expected_cost, expected_latency) in [ + (json!({}), false, false), + (json!({"cost_limit_per_key_tokens": 100}), true, false), + (json!({"cost_limit_per_key_tokens": 0}), true, false), + ( + json!({"scheduling_presets": [{"preset": "cost_first", "enabled": true}]}), + true, + false, + ), + ( + json!({"scheduling_presets": [{"preset": "quota_balanced", "enabled": true}]}), + true, + false, + ), + ( + json!({"scheduling_presets": [{"preset": "latency_first", "enabled": true}]}), + false, + true, + ), + ( + json!({"scheduling_presets": [ + {"preset": "cost_first", "enabled": false}, + {"preset": "latency_first", "enabled": false} + ]}), + false, + false, + ), + ] { + let pool_config = config(value.clone()); + let snapshot = read_provider_pool_scheduling_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + None, + ) + .await; + assert_eq!( + snapshot.cost_window_usage_by_key.get("key").copied(), + expected_cost.then_some(100), + "config: {value}" + ); + assert_eq!( + snapshot.latency_avg_ms_by_key.get("key").copied(), + expected_latency.then_some(20.0), + "config: {value}" + ); + } + } + + #[tokio::test] + async fn scheduling_runtime_checks_cost_for_candidates_beyond_admin_display_limit() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let keys = (0..513) + .map(|index| format!("key-{index}")) + .collect::>(); + seed_window_metrics(&runtime, "pool", &keys[512]).await; + let pool_config = config(json!({ "cost_limit_per_key_tokens": 100 })); + let snapshot = read_provider_pool_scheduling_runtime_state( + &runtime, + "pool", + &keys, + &pool_config, + None, + ) + .await; + assert_eq!( + snapshot.cost_window_usage_by_key.get(&keys[512]), + Some(&100) + ); + } + + #[tokio::test] + async fn scheduling_sticky_lookup_invalidates_a_cooled_down_binding() { + let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + let pool_config = config(json!({})); + let sticky_key = pool_sticky_key("pool", "session"); + runtime + .kv_set(&sticky_key, "key".to_string(), None) + .await + .expect("sticky session"); + runtime + .kv_set( + &pool_cooldown_key("pool", "key"), + "rate_limit".to_string(), + None, + ) + .await + .expect("cooldown"); + assert!(read_provider_pool_sticky_bound_key_id( + &runtime, + "pool", + &pool_config, + Some("session") + ) + .await + .is_none()); + assert!(!runtime + .kv_exists(&sticky_key) + .await + .expect("sticky existence")); + } #[test] fn runtime_window_metric_key_ids_are_bounded() { diff --git a/apps/aether-gateway/src/handlers/admin/system/shared/update.rs b/apps/aether-gateway/src/handlers/admin/system/shared/update.rs index ec7b8b5a5..c4992153b 100644 --- a/apps/aether-gateway/src/handlers/admin/system/shared/update.rs +++ b/apps/aether-gateway/src/handlers/admin/system/shared/update.rs @@ -1937,6 +1937,7 @@ pub(crate) async fn start_admin_system_rollback_task( } fn request_process_restart() -> ! { + let _ = aether_runtime::shutdown_logging(std::time::Duration::from_secs(2)); std::process::exit(RESTART_EXIT_CODE); } diff --git a/apps/aether-gateway/src/handlers/proxy/body_buffer.rs b/apps/aether-gateway/src/handlers/proxy/body_buffer.rs index 94dc3fefc..ab8dd2a79 100644 --- a/apps/aether-gateway/src/handlers/proxy/body_buffer.rs +++ b/apps/aether-gateway/src/handlers/proxy/body_buffer.rs @@ -225,6 +225,9 @@ impl RequestBodyBufferError { RequestBodyNormalizationError::RequestBodyTooLarge { .. } => { "request_body_too_large" } + RequestBodyNormalizationError::BodyBufferOverloaded { .. } => { + "request_body_buffer_overloaded" + } }, Self::TooLarge { .. } => "request_body_too_large", Self::Overloaded { .. } => "request_body_buffer_overloaded", @@ -292,15 +295,37 @@ pub(super) async fn buffer_and_normalize_request_body( .await .map_err(RequestBodyBufferError::from)?; let elapsed_ms = buffered.elapsed().as_millis() as u64; + let retained_input_capacity = buffered + .requested_bytes() + .saturating_sub(buffered.bytes().len()); let normalized = buffered - .try_map(|body| { - crate::headers::normalize_request_body_headers_and_bytes_with_limit( + .try_map_with_budget(|body, memory| { + crate::headers::normalize_request_body_headers_and_bytes_with_budget( headers, body, policy.effective_max_bytes(), + &mut |requested_bytes| { + let requested_bytes = requested_bytes.saturating_add(retained_input_capacity); + memory.try_reserve_bytes(requested_bytes).map_err(|_| { + RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes, + budget_bytes: policy.budget_bytes(), + } + }) + }, ) }) - .map_err(RequestBodyBufferError::Normalization)?; + .map_err(|error| match error { + RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes, + budget_bytes, + } => RequestBodyBufferError::Overloaded { + requested_bytes, + budget_bytes, + timeout_ms: 0, + }, + error => RequestBodyBufferError::Normalization(error), + })?; info!( event_name = "frontdoor_request_body_buffer_completed", log_type = "event", diff --git a/apps/aether-gateway/src/handlers/proxy/mod.rs b/apps/aether-gateway/src/handlers/proxy/mod.rs index c464ff9a2..45e7e7d3c 100644 --- a/apps/aether-gateway/src/handlers/proxy/mod.rs +++ b/apps/aether-gateway/src/handlers/proxy/mod.rs @@ -1035,11 +1035,10 @@ pub(crate) async fn proxy_request( ConnectInfo(remote_addr): ConnectInfo, request: Request, ) -> Result, GatewayError> { - crate::request_lifecycle::run_request(Box::pin(proxy_request_inner( - state, - remote_addr, - request, - ))) + crate::request_lifecycle::run_request_with_usage( + state.usage_runtime.clone(), + Box::pin(proxy_request_inner(state, remote_addr, request)), + ) .await } @@ -3269,6 +3268,113 @@ mod tests { )); } + #[tokio::test] + async fn request_body_buffer_allows_parallel_compressed_uploads() { + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + encoder.write_all(br#"{"model":"test"}"#).unwrap(); + let encoded = Bytes::from(encoder.finish().unwrap()); + let budget_bytes = 2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES; + let budget = Arc::new(Semaphore::new(2)); + let policy = RequestBodyBufferPolicy::for_tests_with_budget( + budget_bytes as u64, + Duration::from_secs(1), + Duration::from_millis(50), + budget_bytes, + Arc::clone(&budget), + ); + let mut headers = HeaderMap::new(); + headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip")); + headers.insert(header::CONTENT_LENGTH, HeaderValue::from(encoded.len())); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (finish_tx, finish_rx) = tokio::sync::oneshot::channel(); + let first_policy = policy.clone(); + let mut first_headers = headers.clone(); + let first_encoded = encoded.clone(); + let first = async move { + let stream = async_stream::stream! { + let middle = first_encoded.len() / 2; + yield Ok::<_, std::io::Error>(first_encoded.slice(..middle)); + let _ = started_tx.send(()); + let _ = finish_rx.await; + yield Ok(first_encoded.slice(middle..)); + }; + buffer_and_normalize_request_body( + &mut Some(Body::from_stream(stream)), + &mut first_headers, + "test owns body", + "trace-compressed-first", + &Method::POST, + "/v1/responses", + "test", + first_policy, + ) + .await + }; + let second = async move { + started_rx.await.unwrap(); + let result = buffer_and_normalize_request_body( + &mut Some(Body::from(encoded)), + &mut headers, + "test owns body", + "trace-compressed-second", + &Method::POST, + "/v1/responses", + "test", + policy, + ) + .await; + let _ = finish_tx.send(()); + result + }; + let (first, second) = tokio::time::timeout(Duration::from_secs(2), async { + tokio::join!(first, second) + }) + .await + .expect("concurrent compressed requests should finish"); + assert_eq!(first.unwrap().as_ref(), br#"{"model":"test"}"#); + assert_eq!(second.unwrap().as_ref(), br#"{"model":"test"}"#); + assert_eq!(budget.available_permits(), 2); + } + + #[tokio::test] + async fn request_body_buffer_rejects_decompression_growth_when_budget_is_busy() { + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + encoder.write_all(&vec![b'a'; 100_000]).unwrap(); + let encoded = encoder.finish().unwrap(); + let budget_bytes = 2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES; + let budget = Arc::new(Semaphore::new(2)); + let held = Arc::clone(&budget).acquire_owned().await.unwrap(); + let mut headers = HeaderMap::new(); + headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip")); + headers.insert(header::CONTENT_LENGTH, HeaderValue::from(encoded.len())); + let result = buffer_and_normalize_request_body( + &mut Some(Body::from(encoded)), + &mut headers, + "test owns body", + "trace-decompression-overload", + &Method::POST, + "/v1/responses", + "test", + RequestBodyBufferPolicy::for_tests_with_budget( + budget_bytes as u64, + Duration::from_secs(1), + Duration::from_secs(1), + budget_bytes, + Arc::clone(&budget), + ), + ) + .await + .unwrap_err(); + assert!(matches!( + result, + RequestBodyBufferError::Overloaded { timeout_ms: 0, .. } + )); + assert_eq!(result.http_status(), http::StatusCode::SERVICE_UNAVAILABLE); + assert_eq!(budget.available_permits(), 1); + drop(held); + assert_eq!(budget.available_permits(), 2); + } + #[tokio::test] async fn request_body_buffer_times_out_instead_of_waiting_forever() { let stream = async_stream::stream! { diff --git a/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs b/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs index 2c724408f..89f788467 100644 --- a/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs +++ b/apps/aether-gateway/src/handlers/proxy/websocket/live/audit.rs @@ -508,7 +508,9 @@ async fn persist_live_audit_event( let usage_runtime = std::sync::Arc::clone(&state.usage_runtime); let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state()); let write_request_id = request_id.clone(); + let usage_producer = usage_runtime.track_producer(); let task = tokio::spawn(async move { + let _usage_producer = usage_producer; if tokio::time::timeout( LIVE_AUDIT_WRITE_HARD_TIMEOUT, usage_runtime.record_terminal_event_direct(usage_data.as_ref(), event), @@ -572,7 +574,9 @@ fn spawn_live_audit_event_detached(state: &AppState, event: UsageEvent, audit_sc }; let usage_runtime = std::sync::Arc::clone(&state.usage_runtime); let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state()); + let usage_producer = usage_runtime.track_producer(); runtime.spawn(async move { + let _usage_producer = usage_producer; if tokio::time::timeout( LIVE_AUDIT_WRITE_HARD_TIMEOUT, usage_runtime.record_terminal_event_direct(usage_data.as_ref(), event), diff --git a/apps/aether-gateway/src/handlers/shared/provider_pool.rs b/apps/aether-gateway/src/handlers/shared/provider_pool.rs index 64c6aedb9..53893bcbb 100644 --- a/apps/aether-gateway/src/handlers/shared/provider_pool.rs +++ b/apps/aether-gateway/src/handlers/shared/provider_pool.rs @@ -3,7 +3,8 @@ pub(crate) use super::super::admin::provider::pool::config::{ }; pub(crate) use super::super::admin::provider::pool::runtime::{ admin_provider_pool_key_terminal_error_reason, read_admin_provider_pool_key_cooldown_reason, - read_admin_provider_pool_runtime_state, record_admin_provider_pool_error, + read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state, + read_provider_pool_sticky_bound_key_id, record_admin_provider_pool_error, record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success, release_admin_provider_pool_key_lease, }; diff --git a/apps/aether-gateway/src/headers.rs b/apps/aether-gateway/src/headers.rs index 36e5b0757..060d554db 100644 --- a/apps/aether-gateway/src/headers.rs +++ b/apps/aether-gateway/src/headers.rs @@ -402,9 +402,21 @@ pub(crate) enum RequestBodyNormalizationError { InvalidBodyFraming, AmbiguousBodyFraming, UnsupportedContentEncoding(String), - DecodeFailed { encoding: String, reason: String }, - DecompressedBodyTooLarge { encoding: String, limit_bytes: u64 }, - RequestBodyTooLarge { limit_bytes: u64 }, + DecodeFailed { + encoding: String, + reason: String, + }, + DecompressedBodyTooLarge { + encoding: String, + limit_bytes: u64, + }, + RequestBodyTooLarge { + limit_bytes: u64, + }, + BodyBufferOverloaded { + requested_bytes: usize, + budget_bytes: usize, + }, } impl RequestBodyNormalizationError { @@ -428,6 +440,9 @@ impl RequestBodyNormalizationError { Self::RequestBodyTooLarge { limit_bytes } => { format!("Request body exceeds {limit_bytes} bytes") } + Self::BodyBufferOverloaded { .. } => { + "Request body buffering capacity is temporarily exhausted".to_string() + } } } @@ -440,6 +455,7 @@ impl RequestBodyNormalizationError { Self::UnsupportedContentEncoding(_) | Self::DecodeFailed { .. } => { http::StatusCode::BAD_REQUEST } + Self::BodyBufferOverloaded { .. } => http::StatusCode::SERVICE_UNAVAILABLE, } } } @@ -468,6 +484,10 @@ impl fmt::Display for RequestBodyNormalizationError { Self::RequestBodyTooLarge { limit_bytes } => { write!(f, "request body exceeds {limit_bytes} bytes") } + Self::BodyBufferOverloaded { requested_bytes, budget_bytes } => write!( + f, + "request body buffering needs {requested_bytes} bytes of a {budget_bytes} byte budget" + ), } } } @@ -489,9 +509,24 @@ pub(crate) fn normalize_request_body_headers_and_bytes_with_limit( headers: &mut http::HeaderMap, body_bytes: Bytes, limit_bytes: u64, +) -> Result { + normalize_request_body_headers_and_bytes_with_budget( + headers, + body_bytes, + limit_bytes, + &mut |_| Ok(()), + ) +} + +pub(crate) fn normalize_request_body_headers_and_bytes_with_budget( + headers: &mut http::HeaderMap, + body_bytes: Bytes, + limit_bytes: u64, + budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>, ) -> Result { let body_was_encoded = !request_content_encodings(headers).is_empty(); - let decoded = decoded_request_body_bytes_with_limit(headers, body_bytes.as_ref(), limit_bytes)?; + let decoded = + decoded_request_body_bytes_with_budget(headers, body_bytes.as_ref(), limit_bytes, budget)?; if !body_was_encoded { return Ok(body_bytes); } @@ -533,6 +568,15 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>( headers: &http::HeaderMap, body_bytes: &'a [u8], limit: u64, +) -> Result, RequestBodyNormalizationError> { + decoded_request_body_bytes_with_budget(headers, body_bytes, limit, &mut |_| Ok(())) +} + +fn decoded_request_body_bytes_with_budget<'a>( + headers: &http::HeaderMap, + body_bytes: &'a [u8], + limit: u64, + budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>, ) -> Result, RequestBodyNormalizationError> { validate_request_body_framing(headers)?; let encodings = request_content_encodings(headers); @@ -543,11 +587,21 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>( return Ok(Cow::Borrowed(body_bytes)); } - let mut decoded = body_bytes.to_vec(); + let mut decoded = Cow::Borrowed(body_bytes); for encoding in encodings.iter().rev() { - decoded = decode_single_request_body_with_limit(encoding, decoded.as_slice(), limit)?; + let retained_input_bytes = body_bytes.len().saturating_add(match &decoded { + Cow::Borrowed(_) => 0, + Cow::Owned(bytes) => bytes.capacity(), + }); + decoded = Cow::Owned(decode_single_request_body_with_budget( + encoding, + decoded.as_ref(), + limit, + retained_input_bytes, + budget, + )?); } - Ok(Cow::Owned(decoded)) + Ok(decoded) } fn request_content_encodings(headers: &http::HeaderMap) -> Vec { @@ -631,11 +685,56 @@ fn decode_single_request_body_with_limit( encoding: &str, body_bytes: &[u8], limit: u64, +) -> Result, RequestBodyNormalizationError> { + decode_single_request_body_with_budget( + encoding, + body_bytes, + limit, + body_bytes.len(), + &mut |_| Ok(()), + ) +} + +fn decode_single_request_body_with_budget( + encoding: &str, + body_bytes: &[u8], + limit: u64, + retained_input_bytes: usize, + budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>, ) -> Result, RequestBodyNormalizationError> { match encoding { - "gzip" | "x-gzip" => decode_gzip_body_with_limit(encoding, body_bytes, limit), - "deflate" => decode_deflate_body_with_limit(encoding, body_bytes, limit), - "zstd" => decode_zstd_body_with_limit(encoding, body_bytes, limit), + "gzip" | "x-gzip" => { + let mut decoder = GzDecoder::new(body_bytes); + read_request_decoder_to_end_with_budget( + encoding, + &mut decoder, + limit, + retained_input_bytes, + budget, + ) + } + "deflate" => decode_deflate_body_with_budget( + encoding, + body_bytes, + limit, + retained_input_bytes, + budget, + ), + "zstd" => { + let mut decoder = zstd::stream::read::Decoder::new(body_bytes).map_err(|err| { + RequestBodyNormalizationError::DecodeFailed { + encoding: encoding.to_string(), + reason: err.to_string(), + } + })?; + read_request_decoder_to_end_with_budget( + encoding, + &mut decoder, + limit, + retained_input_bytes, + budget, + ) + } _ => Err(RequestBodyNormalizationError::UnsupportedContentEncoding( encoding.to_string(), )), @@ -669,20 +768,48 @@ fn decode_deflate_body_with_limit( encoding: &str, body_bytes: &[u8], limit: u64, +) -> Result, RequestBodyNormalizationError> { + decode_deflate_body_with_budget(encoding, body_bytes, limit, body_bytes.len(), &mut |_| { + Ok(()) + }) +} + +fn decode_deflate_body_with_budget( + encoding: &str, + body_bytes: &[u8], + limit: u64, + retained_input_bytes: usize, + budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>, ) -> Result, RequestBodyNormalizationError> { let mut zlib_decoder = ZlibDecoder::new(body_bytes); - match read_request_decoder_to_end_with_limit(encoding, &mut zlib_decoder, limit) { + match read_request_decoder_to_end_with_budget( + encoding, + &mut zlib_decoder, + limit, + retained_input_bytes, + budget, + ) { Ok(decoded) => Ok(decoded), - Err(err @ RequestBodyNormalizationError::DecompressedBodyTooLarge { .. }) => Err(err), - Err(zlib_error) => { + Err(zlib_error @ RequestBodyNormalizationError::DecodeFailed { .. }) => { let mut raw_decoder = DeflateDecoder::new(body_bytes); - read_request_decoder_to_end_with_limit(encoding, &mut raw_decoder, limit).map_err( - |raw_error| RequestBodyNormalizationError::DecodeFailed { - encoding: encoding.to_string(), - reason: format!("{zlib_error}; raw deflate fallback failed: {raw_error}"), - }, + read_request_decoder_to_end_with_budget( + encoding, + &mut raw_decoder, + limit, + retained_input_bytes, + budget, ) + .map_err(|raw_error| match raw_error { + RequestBodyNormalizationError::DecodeFailed { .. } => { + RequestBodyNormalizationError::DecodeFailed { + encoding: encoding.to_string(), + reason: format!("{zlib_error}; raw deflate fallback failed: {raw_error}"), + } + } + error => error, + }) } + Err(error) => Err(error), } } @@ -719,21 +846,70 @@ fn read_request_decoder_to_end_with_limit( decoder: &mut impl Read, limit: u64, ) -> Result, RequestBodyNormalizationError> { - let mut limited = decoder.take(limit.saturating_add(1)); + read_request_decoder_to_end_with_budget(encoding, decoder, limit, 0, &mut |_| Ok(())) +} + +fn read_request_decoder_to_end_with_budget( + encoding: &str, + decoder: &mut impl Read, + limit: u64, + retained_input_bytes: usize, + budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>, +) -> Result, RequestBodyNormalizationError> { + let capacity_limit = usize::try_from(limit).unwrap_or(usize::MAX); + let mut scratch = [0_u8; 8 * 1024]; let mut out = Vec::new(); - limited - .read_to_end(&mut out) - .map_err(|err| RequestBodyNormalizationError::DecodeFailed { - encoding: encoding.to_string(), - reason: err.to_string(), + loop { + let remaining = limit.saturating_sub(out.len() as u64).saturating_add(1); + let read_limit = scratch + .len() + .min(usize::try_from(remaining).unwrap_or(usize::MAX)); + let read = decoder.read(&mut scratch[..read_limit]).map_err(|err| { + RequestBodyNormalizationError::DecodeFailed { + encoding: encoding.to_string(), + reason: err.to_string(), + } })?; - if out.len() as u64 > limit { - return Err(RequestBodyNormalizationError::DecompressedBodyTooLarge { - encoding: encoding.to_string(), - limit_bytes: limit, - }); + if read == 0 { + return Ok(out); + } + let next_len = out.len().saturating_add(read); + if next_len as u64 > limit { + return Err(RequestBodyNormalizationError::DecompressedBodyTooLarge { + encoding: encoding.to_string(), + limit_bytes: limit, + }); + } + if next_len > out.capacity() { + let mut capacity = out + .capacity() + .saturating_mul(2) + .max(next_len) + .min(capacity_limit); + // The encoded body and previous decoding layer remain alive during growth. + match budget(retained_input_bytes.saturating_add(capacity)) { + Ok(()) => {} + Err(RequestBodyNormalizationError::BodyBufferOverloaded { .. }) + if capacity > next_len => + { + // Rejected reservations leave the budget unchanged. Spare capacity + // must not reject a body whose actual bytes still fit. + capacity = next_len; + budget(retained_input_bytes.saturating_add(capacity))?; + } + Err(error) => return Err(error), + } + out.try_reserve_exact(capacity.saturating_sub(out.len())) + .map_err(|err| RequestBodyNormalizationError::DecodeFailed { + encoding: encoding.to_string(), + reason: err.to_string(), + })?; + if out.capacity() > capacity { + budget(retained_input_bytes.saturating_add(out.capacity()))?; + } + } + out.extend_from_slice(&scratch[..read]); } - Ok(out) } pub(crate) fn header_equals( @@ -1223,6 +1399,14 @@ mod tests { #[test] fn request_body_normalization_error_maps_http_status() { + assert_eq!( + RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes: 2, + budget_bytes: 1, + } + .http_status(), + http::StatusCode::SERVICE_UNAVAILABLE + ); assert_eq!( RequestBodyNormalizationError::RequestBodyTooLarge { limit_bytes: 1 }.http_status(), http::StatusCode::PAYLOAD_TOO_LARGE @@ -1250,6 +1434,193 @@ mod tests { ); } + #[test] + fn budgeted_normalization_accounts_for_encoded_input_and_output_capacity() { + let payload = vec![b'a'; 150_000]; + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + encoder.write_all(&payload).expect("gzip payload"); + let encoded = encoder.finish().expect("gzip finish"); + let encoded_len = encoded.len(); + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CONTENT_ENCODING, + HeaderValue::from_static("gzip"), + ); + let mut reservations = Vec::new(); + + let decoded = super::normalize_request_body_headers_and_bytes_with_budget( + &mut headers, + encoded.into(), + 256 * 1024, + &mut |bytes| { + reservations.push(bytes); + Ok(()) + }, + ) + .expect("budgeted gzip should decode"); + + assert_eq!(decoded.as_ref(), payload.as_slice()); + assert!(!headers.contains_key(http::header::CONTENT_ENCODING)); + assert_eq!(reservations[0], encoded_len + 8 * 1024); + assert!(reservations.windows(2).all(|pair| pair[0] <= pair[1])); + assert!(reservations.last().copied().unwrap() >= encoded_len + payload.len()); + assert!(reservations.last().copied().unwrap() <= encoded_len + payload.len() * 2); + assert!( + reservations.len() <= 8, + "output growth should remain geometric" + ); + } + + #[test] + fn budgeted_normalization_accounts_for_retained_encoding_layers() { + let payload = b"small chained request body"; + let mut inner = GzEncoder::new(Vec::new(), Compression::default()); + inner.write_all(payload).expect("inner gzip payload"); + let inner = inner.finish().expect("inner gzip finish"); + let mut outer = GzEncoder::new(Vec::new(), Compression::default()); + outer.write_all(&inner).expect("outer gzip payload"); + let encoded = outer.finish().expect("outer gzip finish"); + let encoded_len = encoded.len(); + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CONTENT_ENCODING, + HeaderValue::from_static("gzip, gzip"), + ); + let mut reservations = Vec::new(); + + let decoded = super::normalize_request_body_headers_and_bytes_with_budget( + &mut headers, + encoded.into(), + 256, + &mut |bytes| { + reservations.push(bytes); + Ok(()) + }, + ) + .expect("chained gzip should decode"); + + assert_eq!(decoded.as_ref(), payload); + assert_eq!( + reservations, + vec![ + encoded_len + inner.len(), + encoded_len + inner.len() + payload.len(), + ] + ); + } + + #[test] + fn budgeted_decoder_stops_before_collecting_when_capacity_is_exhausted() { + let source = vec![b'a'; 150_000]; + let mut decoder = std::io::Cursor::new(source); + let error = super::read_request_decoder_to_end_with_budget( + "test", + &mut decoder, + 256 * 1024, + 100, + &mut |requested_bytes| { + Err(RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes, + budget_bytes: 100, + }) + }, + ) + .expect_err("budget rejection must stop output growth"); + + assert_eq!(decoder.position(), 8 * 1024); + assert_eq!( + error, + RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes: 100 + 8 * 1024, + budget_bytes: 100, + } + ); + } + + #[test] + fn budgeted_decoder_accepts_actual_output_when_geometric_growth_does_not_fit() { + let source = vec![b'a'; 100_000]; + let mut decoder = std::io::Cursor::new(&source); + let retained_input_bytes = 100; + let budget_bytes = retained_input_bytes + source.len(); + let mut reserved_bytes = retained_input_bytes; + let mut rejected_growth = false; + + let decoded = super::read_request_decoder_to_end_with_budget( + "test", + &mut decoder, + 256 * 1024, + retained_input_bytes, + &mut |requested_bytes| { + if requested_bytes > budget_bytes { + rejected_growth = true; + return Err(RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes, + budget_bytes, + }); + } + reserved_bytes = reserved_bytes.max(requested_bytes); + Ok(()) + }, + ) + .expect("actual output within the budget should finish decoding"); + + assert!( + rejected_growth, + "test must exercise oversized spare capacity" + ); + assert_eq!(decoded, source); + assert_eq!(reserved_bytes, budget_bytes); + assert_eq!(decoded.capacity() + retained_input_bytes, budget_bytes); + } + + #[test] + fn budgeted_deflate_preserves_capacity_and_size_rejections() { + let payload = [b'a'; 128]; + let mut wrapped = ZlibEncoder::new(Vec::new(), Compression::default()); + wrapped.write_all(&payload).expect("zlib payload"); + let wrapped = wrapped.finish().expect("zlib finish"); + let mut raw = DeflateEncoder::new(Vec::new(), Compression::default()); + raw.write_all(&payload).expect("raw deflate payload"); + let raw = raw.finish().expect("raw deflate finish"); + + for encoded in [wrapped, raw] { + let mut headers = HeaderMap::new(); + headers.insert( + http::header::CONTENT_ENCODING, + HeaderValue::from_static("deflate"), + ); + let rejection = RequestBodyNormalizationError::BodyBufferOverloaded { + requested_bytes: 300, + budget_bytes: 200, + }; + let error = super::normalize_request_body_headers_and_bytes_with_budget( + &mut headers, + encoded.clone().into(), + 256, + &mut |_| Err(rejection.clone()), + ) + .expect_err("capacity rejection must remain an overload error"); + assert_eq!(error, rejection); + assert!(headers.contains_key(http::header::CONTENT_ENCODING)); + + let error = super::normalize_request_body_headers_and_bytes_with_budget( + &mut headers, + encoded.into(), + 64, + &mut |_| Ok(()), + ) + .expect_err("size rejection must remain a size error"); + assert!(matches!( + error, + RequestBodyNormalizationError::DecompressedBodyTooLarge { + limit_bytes: 64, + .. + } + )); + } + } + #[test] fn check_request_content_length_allows_missing_or_within_limit() { let empty = HeaderMap::new(); diff --git a/apps/aether-gateway/src/main.rs b/apps/aether-gateway/src/main.rs index 6afdcc016..b6d3cda52 100644 --- a/apps/aether-gateway/src/main.rs +++ b/apps/aether-gateway/src/main.rs @@ -18,6 +18,7 @@ use hyper_util::{ server::conn::auto::Builder as HyperServerBuilder, service::TowerToHyperService, }; +use tokio_util::sync::CancellationToken; use tower::{Service as _, ServiceExt as _}; use tracing::{debug, info, warn}; @@ -127,6 +128,7 @@ use aether_gateway::{ FrontdoorCorsConfig, FrontdoorUserRpmConfig, GatewayDataConfig, UsageRuntimeConfig, VideoTaskTruthSourceMode, }; +use aether_gateway_frontdoor::{http_connection_limit, HttpConnectionBudget}; use aether_runtime::{ init_service_runtime, FileLoggingConfig, LogDestination, LogFormat, LogRotation, ServiceRuntimeConfig, @@ -1007,6 +1009,13 @@ struct GatewayUsageArgs { )] queue_stream_maxlen: usize, + #[arg( + long, + env = "AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES", + default_value_t = 1024 * 1024 + )] + queue_payload_max_bytes: usize, + #[arg( long, env = "AETHER_GATEWAY_USAGE_QUEUE_BATCH_SIZE", @@ -1220,6 +1229,7 @@ impl GatewayUsageArgs { consumer_group: self.queue_group.trim().to_string(), dlq_stream_key: self.queue_dlq_stream_key.trim().to_string(), stream_maxlen: self.queue_stream_maxlen.max(1), + queue_payload_max_bytes: self.queue_payload_max_bytes, consumer_batch_size: self.queue_batch_size.max(1), consumer_block_ms: self.queue_block_ms.max(1), reclaim_idle_ms: self.queue_reclaim_idle_ms.max(1), @@ -1486,6 +1496,22 @@ struct Args { /// Maximum number of HTTP/1 request header fields. http_max_headers: usize, + #[arg( + long, + env = "AETHER_GATEWAY_HTTP_SHUTDOWN_TIMEOUT_MS", + default_value_t = 30_000 + )] + /// Grace period for HTTP requests and upgraded connections before forced close. + http_shutdown_timeout_ms: u64, + + #[arg( + long, + env = "AETHER_GATEWAY_USAGE_SHUTDOWN_TIMEOUT_MS", + default_value_t = 30_000 + )] + /// Additional time for request finalizers and local usage buffers to persist. + usage_shutdown_timeout_ms: u64, + /// 容器内健康检查入口:根据当前 bind 端口探测本地 /health。 #[arg(long, hide = true, default_value_t = false)] healthcheck: bool, @@ -1567,6 +1593,11 @@ struct Args { #[arg(long, env = "AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS")] max_in_flight_requests: Option, + /// Maximum accepted HTTP TCP connections across all listener shards, including upgrades. + /// Unset or 0 follows request plus WebSocket capacity, bounded by the FD allowance. + #[arg(long, env = "AETHER_GATEWAY_MAX_HTTP_CONNECTIONS")] + max_http_connections: Option, + /// Maximum number of long-lived public WebSocket connections. When unset, /// this follows `max_in_flight_requests` while remaining an independent /// gate. Set `AETHER_GATEWAY_MAX_WEBSOCKET_CONNECTIONS` to override it. @@ -1851,10 +1882,12 @@ fn gateway_listeners( async fn serve_gateway_router( listeners: Vec, router: axum::Router, + connection_budget: Arc, http2_max_concurrent_streams: u32, http_header_read_timeout_ms: u64, http_header_max_bytes: usize, http_max_headers: usize, + shutdown: CancellationToken, ) -> Result<(), Box> { let http2_max_concurrent_streams = gateway_http2_max_concurrent_streams(http2_max_concurrent_streams); @@ -1865,23 +1898,41 @@ async fn serve_gateway_router( let mut servers = tokio::task::JoinSet::new(); for listener in listeners { let router = router.clone(); + let connection_budget = Arc::clone(&connection_budget); + let shutdown = shutdown.clone(); servers.spawn(async move { serve_gateway_listener( listener, router, + connection_budget, http2_max_concurrent_streams, http_header_read_timeout_ms, http_header_max_bytes, http_max_headers, + shutdown, ) .await }); } - if let Some(result) = servers.join_next().await { - servers.abort_all(); - let serve_result = result - .map_err(|err| std::io::Error::other(format!("gateway listener task failed: {err}")))?; - serve_result?; + let mut failure = None; + while let Some(result) = servers.join_next().await { + let result = result.unwrap_or_else(|err| { + Err(std::io::Error::other(format!( + "gateway listener task failed: {err}" + ))) + }); + if let Err(error) = result { + failure.get_or_insert(error); + shutdown.cancel(); + connection_budget.force_close(); + } + } + if let Some(error) = failure { + return Err(error.into()); + } + // Hyper hands upgrades to application tasks; their IO still owns this budget. + while connection_budget.snapshot().in_flight != 0 { + tokio::time::sleep(std::time::Duration::from_millis(10)).await; } Ok(()) } @@ -1889,14 +1940,26 @@ async fn serve_gateway_router( async fn serve_gateway_listener( listener: tokio::net::TcpListener, router: axum::Router, + connection_budget: Arc, http2_max_concurrent_streams: u32, http_header_read_timeout_ms: u64, http_header_max_bytes: usize, http_max_headers: usize, + shutdown: CancellationToken, ) -> Result<(), std::io::Error> { let mut make_service = router.into_make_service_with_connect_info::(); + let mut connections = tokio::task::JoinSet::new(); loop { - let (io, remote_addr) = listener.accept().await?; + let (io, remote_addr) = tokio::select! { + biased; + _ = shutdown.cancelled() => break, + _ = connections.join_next(), if !connections.is_empty() => continue, + accepted = connection_budget.accept(&listener) => accepted, + }; + let Ok(io) = connection_budget.try_admit(io) else { + tokio::task::yield_now().await; + continue; + }; let tower_service = make_service .call(remote_addr) .await @@ -1909,7 +1972,9 @@ async fn serve_gateway_listener( }); let io = TokioIo::new(io); - tokio::spawn(async move { + let shutdown = shutdown.clone(); + let connection_budget = Arc::clone(&connection_budget); + connections.spawn(async move { let mut builder = HyperServerBuilder::new(TokioExecutor::new()); // Hyper's HTTP/1 header timer is opt-in when using the custom // connection builder. Configure both protocol parsers explicitly: @@ -1938,17 +2003,34 @@ async fn serve_gateway_listener( // the service so a peer cannot hold a socket open while dribbling // protocol bytes or an initial header block. Once the gate opens, // request and response bodies remain fully streaming. - let connection_result = drive_gateway_connection( - builder.serve_connection_with_upgrades(io, hyper_service), - first_request_gate, - std::time::Duration::from_millis(http_header_read_timeout_ms), - ) - .await; + let connection = builder.serve_connection_with_upgrades(io, hyper_service); + tokio::pin!(connection); + let draining_connection = async { + tokio::select! { + result = &mut connection => result, + _ = shutdown.cancelled() => { + connection.as_mut().graceful_shutdown(); + connection.await + } + } + }; + let connection_result = tokio::select! { + biased; + _ = connection_budget.wait_for_forced_close() => Ok(()), + result = drive_gateway_connection( + draining_connection, + first_request_gate, + std::time::Duration::from_millis(http_header_read_timeout_ms), + ) => result, + }; if let Err(err) = connection_result { tracing::trace!(error = ?err, "gateway connection closed with error"); } }); } + drop(listener); + while connections.join_next().await.is_some() {} + Ok(()) } fn resolve_local_http_base_url(app_port: u16) -> Result { @@ -2065,11 +2147,14 @@ fn validate_deployment_topology( } fn main() -> Result<(), Box> { - tokio::runtime::Builder::new_multi_thread() + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + let runtime = tokio::runtime::Builder::new_multi_thread() .enable_all() .thread_stack_size(GATEWAY_TOKIO_WORKER_STACK_SIZE_BYTES) - .build()? - .block_on(run()) + .build()?; + let result = runtime.block_on(run()); + aether_usage_runtime::shutdown_usage_background_runtime(std::time::Duration::from_secs(5)); + result } async fn run() -> Result<(), Box> { @@ -2133,6 +2218,13 @@ async fn run() -> Result<(), Box> { .max_websocket_connections .filter(|limit| *limit > 0) .unwrap_or(request_concurrency_limit); + let http_connection_limit = http_connection_limit( + args.max_http_connections, + request_concurrency_limit, + websocket_connection_limit, + soft_fd_limit(), + ); + let http_connection_budget = Arc::new(HttpConnectionBudget::new(http_connection_limit)); let distributed_websocket_connection_limit = match args.distributed_websocket_connection_limit { Some(limit) if limit > 0 => Some(limit), Some(_) => None, @@ -2323,7 +2415,8 @@ async fn run() -> Result<(), Box> { } state = state .with_request_concurrency_limit(request_concurrency_limit) - .with_websocket_connection_limit(websocket_connection_limit); + .with_websocket_connection_limit(websocket_connection_limit) + .with_http_connection_budget(Arc::clone(&http_connection_budget)); if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) { let distributed_gate = state .runtime_state() @@ -2467,6 +2560,7 @@ async fn run() -> Result<(), Box> { let listeners = gateway_listeners(bind_addr, listen_backlog, listener_shards)?; let public_base_url = resolve_local_http_base_url(app_port)?; let frontdoor_health_url = format!("{public_base_url}/_gateway/health"); + let shutdown_state = state.clone(); let api_router = build_router_with_state(state); // Compose the final router: API routes + optional static file serving. @@ -2486,6 +2580,7 @@ async fn run() -> Result<(), Box> { app_port, listen_backlog, listener_shards, + max_http_connections = http_connection_limit, http2_max_concurrent_streams = gateway_http2_max_concurrent_streams(args.http2_max_concurrent_streams), public_url = %public_base_url, healthcheck_url = %frontdoor_health_url, @@ -2493,18 +2588,61 @@ async fn run() -> Result<(), Box> { "aether-gateway ready" ); - serve_gateway_router( - listeners, - router, - args.http2_max_concurrent_streams, - args.http_header_read_timeout_ms, - args.http_header_max_bytes, - args.http_max_headers, - ) - .await?; + let shutdown = CancellationToken::new(); + let serve_result = { + let server = serve_gateway_router( + listeners, + router, + Arc::clone(&http_connection_budget), + args.http2_max_concurrent_streams, + args.http_header_read_timeout_ms, + args.http_header_max_bytes, + args.http_max_headers, + shutdown.clone(), + ); + tokio::pin!(server); + tokio::select! { + result = &mut server => result, + signal = aether_runtime::wait_for_shutdown_signal() => { + signal?; + info!("shutdown signal received, draining gateway requests"); + shutdown.cancel(); + match tokio::time::timeout( + std::time::Duration::from_millis(args.http_shutdown_timeout_ms), + &mut server, + ).await { + Ok(result) => result, + Err(_) => { + warn!( + event_name = "gateway_http_shutdown_deadline", + connections = http_connection_budget.snapshot().in_flight, + "HTTP drain deadline reached; closing remaining sockets" + ); + http_connection_budget.force_close(); + match tokio::time::timeout(std::time::Duration::from_secs(5), &mut server).await { + Ok(result) => result, + Err(_) => Err(std::io::Error::new(std::io::ErrorKind::TimedOut, + "gateway connection tasks did not stop after forced close").into()), + } + } + } + } + } + }; + let usage_result = shutdown_state + .shutdown_usage_runtime(std::time::Duration::from_millis( + args.usage_shutdown_timeout_ms, + )) + .await; if let Some(background_tasks) = background_tasks { background_tasks.shutdown().await; } + serve_result?; + usage_result?; + info!( + event_name = "gateway_shutdown_complete", + "gateway local persistence drained" + ); Ok(()) } @@ -3437,6 +3575,10 @@ fn pending_backfills_error( #[cfg(test)] mod tests { + mod shutdown { + include!("shutdown_tests.rs"); + } + use super::{ automatic_gateway_request_concurrency_for_capacity, automatic_gateway_request_concurrency_for_parallelism, automatic_sql_pool_config, @@ -3478,6 +3620,8 @@ mod tests { http_header_read_timeout_ms: DEFAULT_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS, http_header_max_bytes: DEFAULT_GATEWAY_HTTP_HEADER_MAX_BYTES, http_max_headers: DEFAULT_GATEWAY_HTTP_MAX_HEADERS, + http_shutdown_timeout_ms: 30_000, + usage_shutdown_timeout_ms: 30_000, healthcheck: false, healthcheck_timeout_ms: 3_000, deployment_topology: DeploymentTopologyArg::SingleNode, @@ -3492,6 +3636,7 @@ mod tests { video_task_poller_batch_size: 32, video_task_store_path: None, max_in_flight_requests: None, + max_http_connections: None, max_websocket_connections: None, distributed_request_limit: None, distributed_websocket_connection_limit: None, @@ -3532,6 +3677,7 @@ mod tests { queue_group: "usage_consumers".to_string(), queue_dlq_stream_key: "usage:events:dlq".to_string(), queue_stream_maxlen: 200_000, + queue_payload_max_bytes: 1024 * 1024, queue_batch_size: 128, queue_block_ms: 500, queue_reclaim_idle_ms: 60_000, @@ -3925,6 +4071,37 @@ mod tests { } } + #[test] + fn gateway_usage_queue_payload_limit_preserves_cli_override_and_rejects_zero() { + let command = ::command(); + let argument = command + .get_arguments() + .find(|argument| argument.get_id() == "queue_payload_max_bytes") + .expect("usage payload argument must be registered"); + assert_eq!( + argument.get_env(), + Some(std::ffi::OsStr::new( + "AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES" + )) + ); + assert_eq!(argument.get_default_values()[0].to_str(), Some("1048576")); + let args = Args::try_parse_from(["aether-gateway", "--queue-payload-max-bytes", "32768"]) + .expect("explicit usage payload limit should parse"); + let config = args.usage.to_config(4, 8, Some(4)); + assert_eq!(config.queue_payload_max_bytes, 32_768); + assert!(config.validate().is_ok()); + + let mut args = test_args(); + assert_eq!( + args.usage.to_config(4, 8, Some(4)).queue_payload_max_bytes, + 1024 * 1024 + ); + args.usage.queue_payload_max_bytes = 0; + let config = args.usage.to_config(4, 8, Some(4)); + assert_eq!(config.queue_payload_max_bytes, 0); + assert!(config.validate().is_err()); + } + #[test] fn gateway_usage_queue_workers_manual_override_wins_and_is_capped() { let mut args = test_args(); diff --git a/apps/aether-gateway/src/maintenance/mod.rs b/apps/aether-gateway/src/maintenance/mod.rs index 868fafe0a..5f90f7987 100644 --- a/apps/aether-gateway/src/maintenance/mod.rs +++ b/apps/aether-gateway/src/maintenance/mod.rs @@ -26,11 +26,11 @@ pub(crate) use runtime::{ start_manual_usage_cleanup_task, start_proxy_upgrade_rollout, AccountSelfCheckRunSummary, AdminCleanupRunRecord, AdminCleanupTaskKind, AdminStatsRebuildSummary, AdminSystemCleanupSummary, ManualUsageCleanupError, ManualUsageCleanupMode, - ManualUsageCleanupOptions, OAuthTokenRefreshRunSummary, PoolQuotaProbeRunSummary, - PoolQuotaProbeWorkerConfig, ProviderCheckinRunSummary, ProviderQuotaAlertRunSummary, - ProxyUpgradeRolloutCancelSummary, ProxyUpgradeRolloutConflictClearSummary, - ProxyUpgradeRolloutNodeActionSummary, ProxyUpgradeRolloutProbeConfig, - ProxyUpgradeRolloutSkippedRestoreSummary, ProxyUpgradeRolloutStatus, - ProxyUpgradeRolloutTrackedNodeState, UsageCounterFlushRuntimeMetrics, - UsageCounterFlushWorkerConfig, + ManualUsageCleanupOptions, OAuthTokenRefreshRunSummary, PoolQuotaProbeReplenishCoordinator, + PoolQuotaProbeRunSummary, PoolQuotaProbeWorkerConfig, ProviderCheckinRunSummary, + ProviderQuotaAlertRunSummary, ProxyUpgradeRolloutCancelSummary, + ProxyUpgradeRolloutConflictClearSummary, ProxyUpgradeRolloutNodeActionSummary, + ProxyUpgradeRolloutProbeConfig, ProxyUpgradeRolloutSkippedRestoreSummary, + ProxyUpgradeRolloutStatus, ProxyUpgradeRolloutTrackedNodeState, + UsageCounterFlushRuntimeMetrics, UsageCounterFlushWorkerConfig, }; diff --git a/apps/aether-gateway/src/maintenance/runtime.rs b/apps/aether-gateway/src/maintenance/runtime.rs index 66d104b65..b5532fa79 100644 --- a/apps/aether-gateway/src/maintenance/runtime.rs +++ b/apps/aether-gateway/src/maintenance/runtime.rs @@ -84,7 +84,8 @@ pub(crate) use pool_quota_probe::{ perform_pool_quota_probe_once, perform_pool_quota_probe_once_for_provider_with_config, perform_pool_quota_probe_once_with_config, pool_quota_probe_target_count, select_pool_quota_probe_key_ids, spawn_pool_quota_probe_replenish_for_request, - spawn_pool_quota_probe_worker, PoolQuotaProbeRunSummary, PoolQuotaProbeWorkerConfig, + spawn_pool_quota_probe_worker, PoolQuotaProbeReplenishCoordinator, PoolQuotaProbeRunSummary, + PoolQuotaProbeWorkerConfig, }; pub(crate) use pool_score_rebuild::{ ensure_provider_key_pool_scores_for_keys, perform_pool_score_rebuild_once, diff --git a/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs b/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs index f71c9031d..7db3602d3 100644 --- a/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs +++ b/apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs @@ -1,4 +1,6 @@ -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::{BTreeMap, BTreeSet, HashMap}; +use std::future::Future; +use std::sync::{Arc, Mutex}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use aether_data_contracts::repository::pool_scores::{ @@ -47,6 +49,166 @@ const POOL_QUOTA_PROBE_BURST_RETRY_GUARD_SECONDS: u64 = 15; const POOL_QUOTA_PROBE_AUTO_MIN_INTERVAL_SECONDS: u64 = 30; const POOL_QUOTA_PROBE_AUTO_MAX_INTERVAL_SECONDS: u64 = 10 * 60; const POOL_QUOTA_PROBE_AUTO_MAX_PRESSURE: u64 = 64; +const POOL_QUOTA_PROBE_LOCAL_MAX_PROVIDERS: usize = 1024; + +#[derive(Debug)] +pub(crate) struct PoolQuotaProbeReplenishCoordinator { + capacity: usize, + state: Mutex, +} + +#[derive(Debug, Default)] +struct PoolQuotaProbeReplenishState { + providers: HashMap, + started_total: u64, + coalesced_total: u64, + capacity_rejected_total: u64, +} + +#[derive(Debug)] +struct PoolQuotaProbeReplenishEntry { + identity: Arc<()>, + pending: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct PoolQuotaProbeReplenishSnapshot { + pub(crate) capacity: usize, + pub(crate) active: usize, + pub(crate) started_total: u64, + pub(crate) coalesced_total: u64, + pub(crate) capacity_rejected_total: u64, +} + +impl Default for PoolQuotaProbeReplenishCoordinator { + fn default() -> Self { + Self::new(POOL_QUOTA_PROBE_LOCAL_MAX_PROVIDERS) + } +} + +impl PoolQuotaProbeReplenishCoordinator { + fn new(capacity: usize) -> Self { + Self { + capacity, + state: Mutex::new(PoolQuotaProbeReplenishState::default()), + } + } + + pub(crate) fn snapshot(&self) -> PoolQuotaProbeReplenishSnapshot { + let state = self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + PoolQuotaProbeReplenishSnapshot { + capacity: self.capacity, + active: state.providers.len(), + started_total: state.started_total, + coalesced_total: state.coalesced_total, + capacity_rejected_total: state.capacity_rejected_total, + } + } + + fn request(self: &Arc, provider_id: String) -> Option { + let mut state = self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if let Some(entry) = state.providers.get_mut(&provider_id) { + entry.pending = true; + state.coalesced_total = state.coalesced_total.saturating_add(1); + return None; + } + if state.providers.len() >= self.capacity { + // Replenishment is best effort; the periodic base scan remains available. + state.capacity_rejected_total = state.capacity_rejected_total.saturating_add(1); + return None; + } + let identity = Arc::new(()); + state.providers.insert( + provider_id.clone(), + PoolQuotaProbeReplenishEntry { + identity: Arc::clone(&identity), + pending: true, + }, + ); + state.started_total = state.started_total.saturating_add(1); + Some(PoolQuotaProbeReplenishGuard { + coordinator: Arc::clone(self), + provider_id, + identity, + finished: false, + }) + } + + fn spawn( + self: &Arc, + provider_id: String, + mut replenish: F, + ) -> Option> + where + F: FnMut() -> Fut + Send + 'static, + Fut: Future + Send + 'static, + { + // Own the guard before spawn so cancellation before the first poll also cleans up. + let mut guard = self.request(provider_id)?; + Some(tokio::spawn(async move { + while guard.next_pass() { + replenish().await; + } + })) + } +} + +struct PoolQuotaProbeReplenishGuard { + coordinator: Arc, + provider_id: String, + identity: Arc<()>, + finished: bool, +} + +impl PoolQuotaProbeReplenishGuard { + fn next_pass(&mut self) -> bool { + let mut state = self + .coordinator + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if let Some(entry) = state + .providers + .get_mut(&self.provider_id) + .filter(|entry| Arc::ptr_eq(&entry.identity, &self.identity)) + { + if entry.pending { + entry.pending = false; + return true; + } + // Check for a follow-up and release ownership in one critical section. + state.providers.remove(&self.provider_id); + } + self.finished = true; + false + } +} + +impl Drop for PoolQuotaProbeReplenishGuard { + fn drop(&mut self) { + if self.finished { + return; + } + let mut state = self + .coordinator + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if state + .providers + .get(&self.provider_id) + .is_some_and(|entry| Arc::ptr_eq(&entry.identity, &self.identity)) + { + state.providers.remove(&self.provider_id); + } + } +} #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum PoolQuotaProbeMode { @@ -1672,38 +1834,62 @@ pub(crate) fn spawn_pool_quota_probe_replenish_for_request( return None; } - Some(tokio::spawn(async move { - let runtime = state.runtime_state.clone(); - mark_probe_burst_pending(runtime.as_ref(), &provider_id).await; - let lease = - acquire_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), &provider_id).await; - if lease.is_none() { - return; - } + let coordinator = Arc::clone(&state.pool_quota_probe_replenish); + coordinator.spawn(provider_id.clone(), move || { + run_pool_quota_probe_replenish(state.clone(), provider_id.clone()) + }) +} - let config = PoolQuotaProbeWorkerConfig::from_env(); - loop { - let pending = runtime - .kv_take(&probe_burst_pending_key(&provider_id)) - .await - .ok() - .flatten() - .is_some(); - if !pending { - break; - } - - match perform_pool_quota_probe_once_for_provider_with_mode( +async fn run_pool_quota_probe_replenish(state: AppState, provider_id: String) { + let runtime = state.runtime_state.clone(); + let runtime = runtime.as_ref(); + let config = PoolQuotaProbeWorkerConfig::from_env(); + run_pool_quota_probe_replenish_with( + runtime, + &provider_id, + || { + perform_pool_quota_probe_once_for_provider_with_mode( &state, &provider_id, config, PoolQuotaProbeMode::Burst, ) - .await - { + }, + |lease| release_pool_quota_probe_burst_trigger_lock(runtime, Some(lease)), + ) + .await; +} + +async fn run_pool_quota_probe_replenish_with( + runtime: &RuntimeState, + provider_id: &str, + mut probe: Probe, + mut release: Release, +) where + Probe: FnMut() -> ProbeFuture, + ProbeFuture: Future>, + Release: FnMut(RuntimeLockLease) -> ReleaseFuture, + ReleaseFuture: Future, +{ + mark_probe_burst_pending(runtime, provider_id).await; + loop { + let Some(lease) = acquire_pool_quota_probe_burst_trigger_lock(runtime, provider_id).await + else { + return; + }; + let recheck_after_release = loop { + let pending = match runtime.kv_take(&probe_burst_pending_key(provider_id)).await { + Ok(pending) => pending.is_some(), + Err(_) => break false, + }; + if !pending { + break true; + } + + match probe().await { Ok(summary) => { if summary.providers_busy > 0 { - mark_probe_burst_pending(runtime.as_ref(), &provider_id).await; + mark_probe_burst_pending(runtime, provider_id).await; tokio::time::sleep(Duration::from_millis(250)).await; continue; } @@ -1717,17 +1903,30 @@ pub(crate) fn spawn_pool_quota_probe_replenish_for_request( } } - let still_pending = runtime - .kv_exists(&probe_burst_pending_key(&provider_id)) + match runtime + .kv_exists(&probe_burst_pending_key(provider_id)) .await - .unwrap_or(false); - if !still_pending { - break; + { + Ok(true) => {} + Ok(false) => break true, + Err(_) => break false, } - } + }; - release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), lease).await; - })) + release(lease).await; + // A different instance may publish after the final pending check and fail + // to acquire our old lease. Recheck after release, then acquire a fresh token + // before consuming that signal. Read failures terminate instead of spinning. + if !recheck_after_release + || !runtime + .kv_exists(&probe_burst_pending_key(provider_id)) + .await + .unwrap_or(false) + { + return; + } + tokio::task::yield_now().await; + } } pub(crate) fn spawn_pool_quota_probe_worker( @@ -1764,7 +1963,430 @@ pub(crate) fn spawn_pool_quota_probe_worker( #[cfg(test)] mod tests { use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + + use aether_runtime_state::MemoryRuntimeStateConfig; use serde_json::json; + use tokio::sync::Notify; + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn pool_quota_probe_local_coalesces_before_spawn_and_keeps_one_follow_up() { + let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(8)); + let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let probe_calls = Arc::new(AtomicUsize::new(0)); + let release_calls = Arc::new(AtomicUsize::new(0)); + let started = Arc::new(Notify::new()); + let finish_first = Arc::new(Notify::new()); + let leader = coordinator + .spawn("provider".to_string(), { + let runtime = Arc::clone(&runtime); + let probe_calls = Arc::clone(&probe_calls); + let release_calls = Arc::clone(&release_calls); + let started = Arc::clone(&started); + let finish_first = Arc::clone(&finish_first); + move || { + let runtime = Arc::clone(&runtime); + let probe_calls = Arc::clone(&probe_calls); + let release_calls = Arc::clone(&release_calls); + let started = Arc::clone(&started); + let finish_first = Arc::clone(&finish_first); + async move { + run_pool_quota_probe_replenish_with( + runtime.as_ref(), + "provider", + || async { + if probe_calls.fetch_add(1, Ordering::AcqRel) == 0 { + started.notify_one(); + finish_first.notified().await; + } + Ok(PoolQuotaProbeRunSummary::empty()) + }, + |lease| { + release_calls.fetch_add(1, Ordering::AcqRel); + release_pool_quota_probe_burst_trigger_lock( + runtime.as_ref(), + Some(lease), + ) + }, + ) + .await; + } + } + }) + .expect("one leader"); + tokio::time::timeout(Duration::from_secs(2), started.notified()) + .await + .expect("first probe starts"); + let barrier = Arc::new(tokio::sync::Barrier::new(65)); + let mut triggers = tokio::task::JoinSet::new(); + for _ in 0..64 { + let coordinator = Arc::clone(&coordinator); + let barrier = Arc::clone(&barrier); + triggers.spawn(async move { + barrier.wait().await; + assert!(coordinator + .spawn("provider".to_string(), || async { + panic!("a duplicate trigger must not spawn work") + }) + .is_none()); + }); + } + barrier.wait().await; + while let Some(result) = triggers.join_next().await { + result.expect("concurrent trigger"); + } + assert_eq!(coordinator.snapshot().started_total, 1); + assert_eq!(coordinator.snapshot().coalesced_total, 64); + assert_eq!(probe_calls.load(Ordering::Acquire), 1); + assert!( + !runtime + .kv_exists(&probe_burst_pending_key("provider")) + .await + .expect("pending read"), + "local duplicates must not each write Redis pending while the leader is running" + ); + finish_first.notify_one(); + tokio::time::timeout(Duration::from_secs(2), leader) + .await + .expect("leader finishes") + .expect("leader task"); + assert_eq!(probe_calls.load(Ordering::Acquire), 2); + assert_eq!( + release_calls.load(Ordering::Acquire), + 2, + "64 retriggers produce one additional Redis lock/drain cycle" + ); + assert_eq!(coordinator.snapshot().active, 0); + } + + #[tokio::test] + async fn pool_quota_probe_local_different_providers_run_independently() { + let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(2)); + let started = Arc::new(tokio::sync::Semaphore::new(0)); + let finish = Arc::new(tokio::sync::Semaphore::new(0)); + let mut tasks = Vec::new(); + for provider_id in ["provider-a", "provider-b"] { + tasks.push( + coordinator + .spawn(provider_id.to_string(), { + let started = Arc::clone(&started); + let finish = Arc::clone(&finish); + move || { + let started = Arc::clone(&started); + let finish = Arc::clone(&finish); + async move { + started.add_permits(1); + finish.acquire().await.expect("finish signal").forget(); + } + } + }) + .expect("independent provider leader"), + ); + } + tokio::time::timeout(Duration::from_secs(2), started.acquire_many(2)) + .await + .expect("both providers start") + .expect("started permits") + .forget(); + assert_eq!(coordinator.snapshot().active, 2); + finish.add_permits(2); + for task in tasks { + task.await.expect("provider finishes"); + } + assert_eq!(coordinator.snapshot().active, 0); + } + + #[test] + fn pool_quota_probe_local_exit_handoff_keeps_exactly_one_owner() { + let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1)); + for _ in 0..100 { + let mut first = coordinator + .request("provider".to_string()) + .expect("first owner"); + assert!(first.next_pass()); + let barrier = Arc::new(std::sync::Barrier::new(2)); + let (mut first, continues, replacement) = std::thread::scope(|scope| { + let first_barrier = Arc::clone(&barrier); + let exit = scope.spawn(move || { + first_barrier.wait(); + let continues = first.next_pass(); + (first, continues) + }); + let trigger = scope.spawn(|| { + barrier.wait(); + coordinator.request("provider".to_string()) + }); + let (first, continues) = exit.join().expect("exit thread"); + (first, continues, trigger.join().expect("trigger thread")) + }); + assert_ne!( + continues, + replacement.is_some(), + "the signal is consumed by exactly one owner" + ); + assert_eq!(coordinator.snapshot().active, 1); + if continues { + assert!(!first.next_pass()); + } + drop(first); + if replacement.is_some() { + assert_eq!( + coordinator.snapshot().active, + 1, + "old cleanup must not delete the replacement" + ); + } + drop(replacement); + assert_eq!(coordinator.snapshot().active, 0); + } + } + + #[tokio::test] + async fn pool_quota_probe_local_abort_and_panic_release_admission() { + let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1)); + let unpolled = coordinator + .spawn("provider".to_string(), || async { + std::future::pending::<()>().await; + }) + .expect("unpolled owner"); + unpolled.abort(); + assert!(unpolled.await.expect_err("aborted").is_cancelled()); + assert_eq!(coordinator.snapshot().active, 0); + + let started = Arc::new(Notify::new()); + let running = coordinator + .spawn("provider".to_string(), { + let started = Arc::clone(&started); + move || { + let started = Arc::clone(&started); + async move { + started.notify_one(); + std::future::pending::<()>().await; + } + } + }) + .expect("running owner"); + started.notified().await; + assert!(coordinator.request("provider".to_string()).is_none()); + running.abort(); + assert!(running.await.expect_err("aborted").is_cancelled()); + assert_eq!(coordinator.snapshot().active, 0); + + let panicked = coordinator + .spawn("provider".to_string(), || async { + panic!("probe panicked") + }) + .expect("panic owner"); + assert!(panicked.await.expect_err("probe panic").is_panic()); + assert_eq!(coordinator.snapshot().active, 0); + coordinator + .spawn("provider".to_string(), || std::future::ready(())) + .expect("later trigger can run") + .await + .expect("recovered probe"); + assert_eq!(coordinator.snapshot().active, 0); + } + + #[test] + fn pool_quota_probe_local_capacity_is_bounded_and_completed_keys_are_removed() { + let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(2)); + let first = coordinator.request("a".to_string()).expect("first"); + let second = coordinator.request("b".to_string()).expect("second"); + assert!(coordinator.request("c".to_string()).is_none()); + assert!(coordinator.request("a".to_string()).is_none()); + assert_eq!(coordinator.snapshot().capacity_rejected_total, 1); + assert_eq!(coordinator.snapshot().coalesced_total, 1); + drop(first); + let replacement = coordinator + .request("c".to_string()) + .expect("freed capacity"); + drop((second, replacement)); + for index in 0..1000 { + drop( + coordinator + .request(format!("provider-{index}")) + .expect("new provider"), + ); + assert_eq!(coordinator.snapshot().active, 0); + } + assert_eq!(coordinator.snapshot().started_total, 1003); + } + + #[tokio::test] + async fn pool_quota_probe_local_state_clones_share_only_the_same_runtime_binding() { + let state = AppState::new().expect("state"); + let cloned = state.clone(); + assert!(Arc::ptr_eq( + &state.pool_quota_probe_replenish, + &cloned.pool_quota_probe_replenish + )); + let rebound_same = cloned.with_runtime_state(Arc::clone(&state.runtime_state)); + assert!(Arc::ptr_eq( + &state.pool_quota_probe_replenish, + &rebound_same.pool_quota_probe_replenish + )); + let rebound = state + .clone() + .with_runtime_state(Arc::new(RuntimeState::memory( + MemoryRuntimeStateConfig::default(), + ))); + assert!(!Arc::ptr_eq( + &state.pool_quota_probe_replenish, + &rebound.pool_quota_probe_replenish + )); + let first = state + .pool_quota_probe_replenish + .request("provider".to_string()) + .expect("first runtime"); + assert!(rebound_same + .pool_quota_probe_replenish + .request("provider".to_string()) + .is_none()); + let second = rebound + .pool_quota_probe_replenish + .request("provider".to_string()) + .expect("other runtime is independent"); + drop((first, second)); + } + + #[tokio::test] + async fn pool_quota_probe_replenish_rechecks_remote_pending_after_unlock() { + let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let first = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1)); + let second = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1)); + let before_unlock = Arc::new(Notify::new()); + let finish_unlock = Arc::new(Notify::new()); + let first_probes = Arc::new(AtomicUsize::new(0)); + let first_releases = Arc::new(AtomicUsize::new(0)); + let first_task = first + .spawn("provider".to_string(), { + let runtime = Arc::clone(&runtime); + let before_unlock = Arc::clone(&before_unlock); + let finish_unlock = Arc::clone(&finish_unlock); + let first_probes = Arc::clone(&first_probes); + let first_releases = Arc::clone(&first_releases); + move || { + let runtime = Arc::clone(&runtime); + let before_unlock = Arc::clone(&before_unlock); + let finish_unlock = Arc::clone(&finish_unlock); + let first_probes = Arc::clone(&first_probes); + let first_releases = Arc::clone(&first_releases); + async move { + run_pool_quota_probe_replenish_with( + runtime.as_ref(), + "provider", + || { + first_probes.fetch_add(1, Ordering::AcqRel); + std::future::ready(Ok(PoolQuotaProbeRunSummary::empty())) + }, + |lease| { + let runtime = Arc::clone(&runtime); + let before_unlock = Arc::clone(&before_unlock); + let finish_unlock = Arc::clone(&finish_unlock); + let first_release = + first_releases.fetch_add(1, Ordering::AcqRel) == 0; + async move { + if first_release { + before_unlock.notify_one(); + finish_unlock.notified().await; + } + release_pool_quota_probe_burst_trigger_lock( + runtime.as_ref(), + Some(lease), + ) + .await; + } + }, + ) + .await; + } + } + }) + .expect("first instance leader"); + tokio::time::timeout(Duration::from_secs(2), before_unlock.notified()) + .await + .expect("first instance drained but still owns Redis lease"); + assert_eq!(first_probes.load(Ordering::Acquire), 1); + assert!(!runtime + .kv_exists(&probe_burst_pending_key("provider")) + .await + .expect("drained pending")); + let second_task = second.spawn("provider".to_string(), { + let runtime = Arc::clone(&runtime); + move || { + let runtime = Arc::clone(&runtime); + async move { + run_pool_quota_probe_replenish_with( + runtime.as_ref(), "provider", + || async { panic!("second instance must not consume pending without the Redis lease") }, + |lease| release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), Some(lease)), + ).await; + } + } + }).expect("independent second instance leader"); + second_task + .await + .expect("second instance leaves pending for the current owner"); + assert_eq!(second.snapshot().active, 0); + assert!(runtime + .kv_exists(&probe_burst_pending_key("provider")) + .await + .expect("new pending signal")); + finish_unlock.notify_one(); + tokio::time::timeout(Duration::from_secs(2), first_task) + .await + .expect("handoff drains") + .expect("first instance finishes"); + assert_eq!( + first_probes.load(Ordering::Acquire), + 2, + "the cross-instance exit signal must trigger a second probe" + ); + assert_eq!(first_releases.load(Ordering::Acquire), 2); + assert_eq!(first.snapshot().active, 0); + assert!(!runtime + .kv_exists(&probe_burst_pending_key("provider")) + .await + .expect("all pending consumed")); + let lease = acquire_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), "provider").await; + assert!(lease.is_some(), "the replacement Redis lease is released"); + release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), lease).await; + } + + #[tokio::test] + async fn pool_quota_probe_replenish_public_spawn_returns_none_for_merged_triggers() { + let repository = Arc::new( + aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository::seed( + Vec::new(), + Vec::new(), + Vec::new(), + ), + ); + let state = AppState::new().expect("state").with_data_state_for_tests( + crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(repository), + ); + let leader = + spawn_pool_quota_probe_replenish_for_request(state.clone(), "provider".to_string()) + .expect("leader handle"); + for _ in 0..64 { + assert!(spawn_pool_quota_probe_replenish_for_request( + state.clone(), + "provider".to_string() + ) + .is_none()); + } + leader.await.expect("leader completes"); + assert_eq!(state.pool_quota_probe_replenish.snapshot().started_total, 1); + assert_eq!( + state.pool_quota_probe_replenish.snapshot().coalesced_total, + 64 + ); + assert_eq!(state.pool_quota_probe_replenish.snapshot().active, 0); + spawn_pool_quota_probe_replenish_for_request(state.clone(), "provider".to_string()) + .expect("later leader handle") + .await + .expect("later leader finishes"); + } #[test] fn worker_error_score_reason_drops_runtime_error_details() { diff --git a/apps/aether-gateway/src/orchestration/effects.rs b/apps/aether-gateway/src/orchestration/effects.rs index d43f23058..dd74c3c69 100644 --- a/apps/aether-gateway/src/orchestration/effects.rs +++ b/apps/aether-gateway/src/orchestration/effects.rs @@ -973,7 +973,7 @@ async fn record_adaptive_rate_limit_effect( let _effect_guard = effect_lock.lock().await; let observed_at_unix_secs = current_unix_secs(); let current_rpm = state - .read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT) + .read_recent_runtime_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT) .await .ok() .map(|recent_candidates| { @@ -1123,7 +1123,7 @@ async fn record_adaptive_success_effect( return; } let Some(recent_candidates) = state - .read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT) + .read_recent_runtime_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT) .await .ok() else { diff --git a/apps/aether-gateway/src/request_candidate_queue.rs b/apps/aether-gateway/src/request_candidate_queue.rs index 365577481..97574ba26 100644 --- a/apps/aether-gateway/src/request_candidate_queue.rs +++ b/apps/aether-gateway/src/request_candidate_queue.rs @@ -643,6 +643,24 @@ impl RequestCandidateQueueRuntime { } } + pub(crate) fn pending_writes(&self) -> usize { + [ + self.metrics.pending_current.load(Ordering::Acquire), + self.metrics + .priority_pending_current + .load(Ordering::Acquire), + self.metrics.active_pending_current.load(Ordering::Acquire), + self.metrics + .terminal_pending_current + .load(Ordering::Acquire), + self.metrics + .terminal_barrier_pending + .load(Ordering::Acquire), + ] + .into_iter() + .fold(0_usize, usize::saturating_add) + } + pub(crate) fn metric_samples(&self) -> Vec { vec![ MetricSample::new( diff --git a/apps/aether-gateway/src/request_lifecycle.rs b/apps/aether-gateway/src/request_lifecycle.rs index 107b1f7ea..bd52557c5 100644 --- a/apps/aether-gateway/src/request_lifecycle.rs +++ b/apps/aether-gateway/src/request_lifecycle.rs @@ -5,6 +5,7 @@ use std::sync::Arc; use std::task::{Context, Poll}; use aether_routing_core::RoutingExecutionPolicy; +use aether_usage_runtime::{UsageProducerGuard, UsageRuntime}; use axum::body::{Body, Bytes, HttpBody}; use http::Response; use http_body::{Frame, SizeHint}; @@ -29,24 +30,49 @@ pub(crate) fn cancel_on_client_disconnect() -> bool { .unwrap_or(false) } +#[cfg(test)] pub(crate) async fn run_request(future: F) -> Result, GatewayError> +where + F: Future, GatewayError>> + Send + 'static, +{ + run_tracked_request(future, None).await +} + +pub(crate) async fn run_request_with_usage( + usage: Arc, + future: F, +) -> Result, GatewayError> +where + F: Future, GatewayError>> + Send + 'static, +{ + run_tracked_request(future, Some(Arc::new(usage.track_producer()))).await +} + +async fn run_tracked_request( + future: F, + producer: Option>, +) -> Result, GatewayError> where F: Future, GatewayError>> + Send + 'static, { let cancel = Arc::new(AtomicBool::new(true)); let diagnostics = Arc::new(RequestDiagnostics::default()); let cancel_for_response = Arc::clone(&cancel); + let producer_for_request = producer.clone(); let future = CANCEL_ON_CLIENT_DISCONNECT.scope( Arc::clone(&cancel), scope_request_diagnostics_with(Some(Arc::clone(&diagnostics)), async move { let response = future.await?; - if cancel_for_response.load(Ordering::Acquire) { + let complete_on_disconnect = !cancel_for_response.load(Ordering::Acquire); + if !complete_on_disconnect && producer.is_none() { return Ok(response); } Ok(response.map(|body| { Body::new(CompleteOnDisconnectBody { body: Some(body), diagnostics, + complete_on_disconnect, + producer, }) })) }), @@ -54,6 +80,7 @@ where CompleteOnDisconnectRequest { future: Some(Box::pin(future)), cancel, + producer: producer_for_request, } .await } @@ -64,6 +91,7 @@ where { future: Option>>, cancel: Arc, + producer: Option>, } impl Future for CompleteOnDisconnectRequest @@ -97,7 +125,9 @@ where if let (Some(future), Ok(runtime)) = (self.future.take(), tokio::runtime::Handle::try_current()) { + let producer = self.producer.take(); runtime.spawn(async move { + let _producer = producer; if let Ok(response) = future.await { drain_body(response.into_body()).await; } @@ -109,6 +139,9 @@ where struct CompleteOnDisconnectBody { body: Option, diagnostics: Arc, + complete_on_disconnect: bool, + // Drop the body first so its terminal handoff registers before this guard ends. + producer: Option>, } impl HttpBody for CompleteOnDisconnectBody { @@ -125,6 +158,7 @@ impl HttpBody for CompleteOnDisconnectBody { let result = Pin::new(body).poll_frame(context); if matches!(result, Poll::Ready(None | Some(Err(_)))) { self.body.take(); + self.producer.take(); } result } @@ -143,13 +177,20 @@ impl HttpBody for CompleteOnDisconnectBody { impl Drop for CompleteOnDisconnectBody { fn drop(&mut self) { + if !self.complete_on_disconnect { + return; + } let Some(body) = self.body.take().filter(|body| !body.is_end_stream()) else { return; }; if let Ok(runtime) = tokio::runtime::Handle::try_current() { + let producer = self.producer.take(); runtime.spawn(scope_request_diagnostics_with( Some(Arc::clone(&self.diagnostics)), - drain_body(body), + async move { + let _producer = producer; + drain_body(body).await; + }, )); } } @@ -297,6 +338,78 @@ mod tests { assert!(sender.is_closed()); } + #[tokio::test] + async fn usage_shutdown_waits_for_a_disconnected_request_before_headers() { + let usage = Arc::new(UsageRuntime::disabled()); + let (started_tx, started_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel::<()>(); + let request = tokio::spawn(run_request_with_usage(usage.clone(), async move { + configure_client_disconnect(RoutingExecutionPolicy::default()); + started_tx.send(()).unwrap(); + release_rx.await.unwrap(); + Ok(Response::new(Body::empty())) + })); + started_rx.await.unwrap(); + request.abort(); + assert!(request.await.unwrap_err().is_cancelled()); + assert!(usage.shutdown(Duration::from_millis(30)).await.is_err()); + assert_eq!(usage.metrics_snapshot().producers_in_flight, 1); + release_tx.send(()).unwrap(); + usage.shutdown(Duration::from_secs(1)).await.unwrap(); + assert_eq!(usage.metrics_snapshot().producers_in_flight, 0); + } + + #[tokio::test] + async fn usage_shutdown_waits_for_disconnected_body_drain() { + let usage = Arc::new(UsageRuntime::disabled()); + let (sender, receiver) = mpsc::channel::>(1); + let response = run_request_with_usage(usage.clone(), async move { + configure_client_disconnect(RoutingExecutionPolicy::default()); + Ok(Response::new(Body::from_stream(stream::unfold( + receiver, + |mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) }, + )))) + }) + .await + .unwrap(); + drop(response); + assert!(usage.shutdown(Duration::from_millis(30)).await.is_err()); + sender.send(Ok(Bytes::from_static(b"last"))).await.unwrap(); + drop(sender); + usage.shutdown(Duration::from_secs(1)).await.unwrap(); + assert_eq!(usage.metrics_snapshot().producers_in_flight, 0); + } + + #[tokio::test] + async fn tracked_bodies_release_shutdown_on_cancellation_or_eof() { + for cancel_on_client_disconnect in [false, true] { + let usage = Arc::new(UsageRuntime::disabled()); + let (sender, receiver) = mpsc::channel::>(1); + let response = run_request_with_usage(usage.clone(), async move { + configure_client_disconnect(RoutingExecutionPolicy { + cancel_on_client_disconnect, + ..Default::default() + }); + Ok(Response::new(Body::from_stream(stream::unfold( + receiver, + |mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) }, + )))) + }) + .await + .unwrap(); + let mut body = response.into_body(); + if cancel_on_client_disconnect { + drop(body); + assert!(sender.is_closed()); + } else { + drop(sender); + assert!(body.frame().await.is_none()); + assert_eq!(usage.metrics_snapshot().producers_in_flight, 0); + } + usage.shutdown(Duration::from_secs(1)).await.unwrap(); + } + } + #[tokio::test] async fn connected_response_preserves_headers_size_hint_and_trailers() { let response = run_request(async { diff --git a/apps/aether-gateway/src/router.rs b/apps/aether-gateway/src/router.rs index 3c78123db..e7d48f818 100644 --- a/apps/aether-gateway/src/router.rs +++ b/apps/aether-gateway/src/router.rs @@ -18,6 +18,7 @@ use tower::{Service as _, ServiceExt}; use tower_http::services::{ServeDir, ServeFile}; use tracing::warn; +use aether_gateway_frontdoor::{http_connection_limit, HttpConnectionBudget}; use aether_runtime::{prometheus_response, ConcurrencyError}; use aether_runtime_state::RuntimeSemaphoreError; @@ -244,9 +245,23 @@ pub(crate) enum RequestAdmissionError { pub async fn serve_tcp(bind: &str) -> Result<(), Box> { let listener = tokio::net::TcpListener::bind(bind).await?; let router = build_router()?; + let configured_connection_limit = std::env::var("AETHER_GATEWAY_MAX_HTTP_CONNECTIONS") + .ok() + .and_then(|value| value.trim().parse::().ok()); + // This compatibility entry point has no configured request capacities or FD probe. + let connection_budget = Arc::new(HttpConnectionBudget::new(http_connection_limit( + configured_connection_limit, + 2048, + 2048, + None, + ))); let mut make_service = router.into_make_service_with_connect_info::(); loop { - let (io, remote_addr) = listener.accept().await?; + let (io, remote_addr) = connection_budget.accept(&listener).await; + let Ok(io) = connection_budget.try_admit(io) else { + tokio::task::yield_now().await; + continue; + }; let tower_service = make_service .call(remote_addr) .await diff --git a/apps/aether-gateway/src/scheduler/candidate/mod.rs b/apps/aether-gateway/src/scheduler/candidate/mod.rs index 7a53b268f..f9b8f7177 100644 --- a/apps/aether-gateway/src/scheduler/candidate/mod.rs +++ b/apps/aether-gateway/src/scheduler/candidate/mod.rs @@ -32,6 +32,9 @@ use regex::Regex; use sha2::{Digest, Sha256}; use std::collections::BTreeMap; +pub(crate) use self::runtime::{ + select_with_auth_concurrency_wait, wait_for_auth_api_key_concurrency_retry, +}; pub(crate) use self::selection::{ is_auth_api_key_concurrency_limit_skip_reason, SchedulerSkippedCandidate, API_KEY_CONCURRENCY_LIMIT_SKIP_REASON, AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON, diff --git a/apps/aether-gateway/src/scheduler/candidate/runtime.rs b/apps/aether-gateway/src/scheduler/candidate/runtime.rs index 25bb44260..51a2003f3 100644 --- a/apps/aether-gateway/src/scheduler/candidate/runtime.rs +++ b/apps/aether-gateway/src/scheduler/candidate/runtime.rs @@ -1,4 +1,5 @@ use std::collections::{BTreeMap, BTreeSet}; +use std::future::Future; use aether_admin::provider::{ pool as admin_provider_pool_pure, status as admin_provider_status_pure, @@ -10,7 +11,8 @@ use aether_scheduler_core::{ candidate_is_selectable_with_runtime_state, candidate_runtime_skip_reason_with_state, effective_provider_key_rpm_limit, CandidateRuntimeSelectabilityInput, }; -use std::time::{SystemTime, UNIX_EPOCH}; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; +use tokio::time::Instant; use crate::data::auth::GatewayAuthApiKeySnapshot; use crate::GatewayError; @@ -109,25 +111,92 @@ pub(super) fn auth_snapshot_concurrency_limit_reached( snapshot: &CandidateRuntimeSelectionSnapshot, now_unix_secs: u64, ) -> bool { - auth_snapshot - .and_then(|snapshot| { - usize::try_from(snapshot.api_key_concurrent_limit?) - .ok() - .and_then(|limit| { - if limit == 0 { - return None; - } - Some((snapshot.api_key_id.as_str(), limit)) - }) - }) - .is_some_and(|(api_key_id, limit)| { - auth_api_key_concurrency_limit_reached( - &snapshot.recent_candidates, - now_unix_secs, - api_key_id, - limit, + auth_snapshot_concurrency_limit(auth_snapshot).is_some_and(|(api_key_id, limit)| { + auth_api_key_concurrency_limit_reached( + &snapshot.recent_candidates, + now_unix_secs, + api_key_id, + limit, + ) + }) +} + +fn auth_snapshot_concurrency_limit( + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, +) -> Option<(&str, usize)> { + let snapshot = auth_snapshot?; + let limit = usize::try_from(snapshot.api_key_concurrent_limit?).ok()?; + (limit > 0).then_some((snapshot.api_key_id.as_str(), limit)) +} + +async fn read_auth_api_key_concurrency_limit_reached( + state: &(impl SchedulerRuntimeState + ?Sized), + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, +) -> Result { + let Some((api_key_id, limit)) = auth_snapshot_concurrency_limit(auth_snapshot) else { + return Ok(false); + }; + let recent_candidates = state.read_recent_request_candidates(128).await?; + Ok(auth_api_key_concurrency_limit_reached( + &recent_candidates, + crate::clock::current_unix_secs(), + api_key_id, + limit, + )) +} + +/// A retry always rebuilds candidates, including at the deadline. Only the +/// intervening polls omit catalog, quota and ranking work while auth is blocked. +pub(crate) async fn wait_for_auth_api_key_concurrency_retry( + state: &(impl SchedulerRuntimeState + ?Sized), + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, + deadline: Instant, + poll_interval: Duration, +) -> Result { + if Instant::now() >= deadline { + return Ok(false); + } + let poll_interval = poll_interval.max(Duration::from_millis(1)); + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + tokio::time::sleep(poll_interval.min(remaining)).await; + if Instant::now() >= deadline + || !read_auth_api_key_concurrency_limit_reached(state, auth_snapshot).await? + { + return Ok(true); + } + } +} + +pub(crate) async fn select_with_auth_concurrency_wait( + state: &(impl SchedulerRuntimeState + ?Sized), + auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, + now_unix_secs: u64, + wait_timeout: Duration, + poll_interval: Duration, + mut select: Select, +) -> Result +where + Select: FnMut(u64) -> Selection, + Selection: Future>, +{ + let deadline = Instant::now() + wait_timeout; + let mut attempt_now_unix_secs = now_unix_secs; + loop { + let (result, auth_limit_blocked) = select(attempt_now_unix_secs).await?; + if !auth_limit_blocked + || !wait_for_auth_api_key_concurrency_retry( + state, + auth_snapshot, + deadline, + poll_interval, ) - }) + .await? + { + return Ok(result); + } + attempt_now_unix_secs = crate::clock::current_unix_secs(); + } } pub(super) fn is_candidate_selectable( diff --git a/apps/aether-gateway/src/scheduler/candidate/tests/concurrency_wait.rs b/apps/aether-gateway/src/scheduler/candidate/tests/concurrency_wait.rs new file mode 100644 index 000000000..1514ffe1b --- /dev/null +++ b/apps/aether-gateway/src/scheduler/candidate/tests/concurrency_wait.rs @@ -0,0 +1,548 @@ +use std::collections::VecDeque; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Mutex; +use std::time::Duration; + +use aether_data::DataLayerError; +use aether_data_contracts::repository::candidate_selection::{ + StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery, + StoredRequestedModelCandidateRowsQuery, +}; +use aether_data_contracts::repository::candidates::{ + RequestCandidateStatus, StoredRequestCandidate, +}; +use aether_data_contracts::repository::provider_catalog::{ + StoredProviderCatalogKey, StoredProviderCatalogProvider, +}; +use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot; +use aether_scheduler_core::{SchedulerAffinityTarget, SchedulerMinimalCandidateSelectionCandidate}; +use async_trait::async_trait; +use tokio::sync::Notify; + +use crate::data::auth::GatewayAuthApiKeySnapshot; +use crate::data::candidate_selection::MinimalCandidateSelectionRowSource; +use crate::scheduler::config::SchedulerOrderingConfig; +use crate::scheduler::state::SchedulerRuntimeState; +use crate::GatewayError; + +use super::super::{ + is_exact_all_skipped_by_auth_limit, + list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal, + list_selectable_candidates_with_skip_reasons, select_with_auth_concurrency_wait, + SchedulerSkippedCandidate, +}; +use super::support::{sample_auth_snapshot, sample_provider, sample_row}; + +const POLL_INTERVAL: Duration = Duration::from_millis(2); + +enum RecentReadAction { + Keep, + Release, + ReleaseAndReplaceCandidate, + ReleaseAndRemoveCandidates, + Fail, +} + +struct CountingState { + rows: Mutex>, + recent: Mutex>, + recent_actions: Mutex>, + row_reads: AtomicUsize, + format_reads: AtomicUsize, + provider_reads: AtomicUsize, + key_reads: AtomicUsize, + quota_reads: AtomicUsize, + recent_reads: AtomicUsize, + row_error_at: Option, + first_row_delay: Duration, + poll_observed: Notify, +} + +impl CountingState { + fn blocked() -> Self { + Self { + rows: Mutex::new(vec![sample_row()]), + recent: Mutex::new(vec![active_candidate()]), + recent_actions: Mutex::new(VecDeque::new()), + row_reads: AtomicUsize::new(0), + format_reads: AtomicUsize::new(0), + provider_reads: AtomicUsize::new(0), + key_reads: AtomicUsize::new(0), + quota_reads: AtomicUsize::new(0), + recent_reads: AtomicUsize::new(0), + row_error_at: None, + first_row_delay: Duration::ZERO, + poll_observed: Notify::new(), + } + } + + fn on_recent_reads(self, actions: impl IntoIterator) -> Self { + *self.recent_actions.lock().unwrap() = actions.into_iter().collect(); + self + } +} + +fn active_candidate() -> StoredRequestCandidate { + let now_ms = i64::try_from(crate::clock::current_unix_secs() * 1000).unwrap(); + StoredRequestCandidate::new( + "active-candidate".to_string(), + "active-request".to_string(), + Some("user-1".to_string()), + Some("api-key-1".to_string()), + None, + None, + 0, + 0, + Some("provider-1".to_string()), + Some("endpoint-1".to_string()), + Some("key-1".to_string()), + RequestCandidateStatus::Streaming, + None, + false, + None, + None, + None, + None, + None, + None, + None, + now_ms, + Some(now_ms), + None, + ) + .unwrap() +} + +fn limited_auth() -> GatewayAuthApiKeySnapshot { + let mut auth = sample_auth_snapshot("api-key-1"); + auth.api_key_concurrent_limit = Some(1); + auth +} + +#[async_trait] +impl MinimalCandidateSelectionRowSource for CountingState { + async fn read_minimal_candidate_selection_rows_for_api_format_and_global_model( + &self, + _api_format: &str, + _global_model_name: &str, + ) -> Result, DataLayerError> { + panic!("requested-model selection should use its paged query") + } + + async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model( + &self, + _api_format: &str, + _requested_model_name: &str, + ) -> Result, DataLayerError> { + panic!("requested-model selection should use its paged query") + } + + async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page( + &self, + query: &StoredRequestedModelCandidateRowsQuery, + ) -> Result, DataLayerError> { + let read = self.row_reads.fetch_add(1, Ordering::SeqCst) + 1; + if read == 1 && !self.first_row_delay.is_zero() { + tokio::time::sleep(self.first_row_delay).await; + } + if self.row_error_at == Some(read) { + return Err(DataLayerError::Postgres( + "candidate query failed".to_string(), + )); + } + Ok(self + .rows + .lock() + .unwrap() + .iter() + .filter(|row| { + row.endpoint_api_format == query.api_format + && row.global_model_name == query.requested_model_name + }) + .skip(query.offset as usize) + .take(query.limit as usize) + .cloned() + .collect()) + } + + async fn read_minimal_candidate_selection_rows_for_api_format( + &self, + _api_format: &str, + ) -> Result, DataLayerError> { + self.format_reads.fetch_add(1, Ordering::SeqCst); + Ok(self.rows.lock().unwrap().clone()) + } + + async fn read_pool_key_candidate_rows_for_group( + &self, + _query: &StoredPoolKeyCandidateRowsQuery, + ) -> Result, DataLayerError> { + panic!("the test has no provider pool") + } +} + +#[async_trait] +impl SchedulerRuntimeState for CountingState { + async fn read_provider_quota_snapshot( + &self, + _provider_id: &str, + ) -> Result, GatewayError> { + self.quota_reads.fetch_add(1, Ordering::SeqCst); + Ok(None) + } + + async fn read_provider_catalog_providers_by_ids( + &self, + provider_ids: &[String], + ) -> Result, GatewayError> { + self.provider_reads.fetch_add(1, Ordering::SeqCst); + Ok(provider_ids + .iter() + .map(|id| sample_provider(id, None)) + .collect()) + } + + async fn read_provider_catalog_keys_by_ids( + &self, + _key_ids: &[String], + ) -> Result, GatewayError> { + self.key_reads.fetch_add(1, Ordering::SeqCst); + Ok(Vec::new()) + } + + async fn read_recent_request_candidates( + &self, + limit: usize, + ) -> Result, GatewayError> { + assert_eq!( + limit, 128, + "polls must use the same sample as full selection" + ); + let read = self.recent_reads.fetch_add(1, Ordering::SeqCst) + 1; + let action = self + .recent_actions + .lock() + .unwrap() + .pop_front() + .unwrap_or(RecentReadAction::Keep); + if matches!(action, RecentReadAction::Fail) { + return Err(GatewayError::Internal( + "recent candidates failed".to_string(), + )); + } + if matches!( + action, + RecentReadAction::Release + | RecentReadAction::ReleaseAndReplaceCandidate + | RecentReadAction::ReleaseAndRemoveCandidates + ) { + for candidate in self.recent.lock().unwrap().iter_mut() { + candidate.status = RequestCandidateStatus::Success; + candidate.finished_at_unix_ms = Some(crate::clock::current_unix_secs() * 1000); + } + } + match action { + RecentReadAction::ReleaseAndReplaceCandidate => { + self.rows.lock().unwrap()[0].key_id = "replacement-key".to_string(); + } + RecentReadAction::ReleaseAndRemoveCandidates => self.rows.lock().unwrap().clear(), + _ => {} + } + if read > 1 { + self.poll_observed.notify_one(); + } + Ok(self.recent.lock().unwrap().clone()) + } + + fn provider_key_rpm_reset_at(&self, _key_id: &str, _now_unix_secs: u64) -> Option { + None + } + + fn read_cached_scheduler_affinity_target( + &self, + _cache_key: &str, + _ttl: Duration, + ) -> Option { + None + } + + fn scheduler_affinity_epoch(&self) -> u64 { + 0 + } + + fn remember_scheduler_affinity_target( + &self, + _cache_key: &str, + _target: SchedulerAffinityTarget, + _ttl: Duration, + _max_entries: usize, + ) { + } + + fn remember_scheduler_affinity_target_for_epoch( + &self, + _cache_key: &str, + _target: SchedulerAffinityTarget, + _ttl: Duration, + _max_entries: usize, + _expected_epoch: Option, + ) -> bool { + true + } +} + +type Selection = ( + Vec, + Vec, +); + +async fn select_requested_model( + state: &CountingState, + auth: Option<&GatewayAuthApiKeySnapshot>, + timeout: Duration, +) -> Result { + select_with_auth_concurrency_wait( + state, + auth, + crate::clock::current_unix_secs(), + timeout, + POLL_INTERVAL, + |now| async move { + let result = list_selectable_candidates_with_skip_reasons( + state, + state, + "openai:chat", + "gpt-4.1", + false, + None, + auth, + None, + now, + false, + SchedulerOrderingConfig::default(), + ) + .await?; + let blocked = is_exact_all_skipped_by_auth_limit(&result.0, &result.1); + Ok((result, blocked)) + }, + ) + .await +} + +#[tokio::test] +async fn concurrent_blocked_selectors_only_prepare_at_start_and_deadline() { + let state = CountingState::blocked(); + let auth = limited_auth(); + let outcomes = futures_util::future::join_all( + (0..8).map(|_| select_requested_model(&state, Some(&auth), Duration::from_millis(80))), + ) + .await; + + for outcome in outcomes { + let (selected, skipped) = outcome.unwrap(); + assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped)); + } + assert_eq!(state.row_reads.load(Ordering::SeqCst), 16); + assert_eq!(state.provider_reads.load(Ordering::SeqCst), 32); + assert_eq!(state.key_reads.load(Ordering::SeqCst), 16); + assert_eq!(state.quota_reads.load(Ordering::SeqCst), 16); + assert!(state.recent_reads.load(Ordering::SeqCst) > 16); + assert_eq!(state.format_reads.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn released_auth_slot_rebuilds_changed_candidates() { + let state = CountingState::blocked().on_recent_reads([ + RecentReadAction::Keep, + RecentReadAction::Keep, + RecentReadAction::ReleaseAndReplaceCandidate, + ]); + let auth = limited_auth(); + let (selected, skipped) = select_requested_model(&state, Some(&auth), Duration::from_secs(1)) + .await + .unwrap(); + + assert!(skipped.is_empty()); + assert_eq!(selected.len(), 1); + assert_eq!(selected[0].key_id, "replacement-key"); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 2); + assert_eq!(state.recent_reads.load(Ordering::SeqCst), 4); +} + +#[tokio::test] +async fn released_auth_slot_does_not_reuse_removed_candidates() { + let state = CountingState::blocked().on_recent_reads([ + RecentReadAction::Keep, + RecentReadAction::ReleaseAndRemoveCandidates, + ]); + let (selected, skipped) = + select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1)) + .await + .unwrap(); + + assert!(selected.is_empty()); + assert!(skipped.is_empty()); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn lightweight_poll_errors_are_propagated_without_another_full_query() { + let state = + CountingState::blocked().on_recent_reads([RecentReadAction::Keep, RecentReadAction::Fail]); + let error = select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1)) + .await + .unwrap_err(); + + assert!( + matches!(error, GatewayError::Internal(message) if message == "recent candidates failed") + ); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 1); + assert_eq!(state.recent_reads.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn recovered_selection_errors_are_propagated() { + let mut state = CountingState::blocked() + .on_recent_reads([RecentReadAction::Keep, RecentReadAction::Release]); + state.row_error_at = Some(2); + let error = select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1)) + .await + .unwrap_err(); + + assert!( + matches!(error, GatewayError::Internal(message) if message.contains("candidate query failed")) + ); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 2); +} + +#[tokio::test] +async fn first_selection_time_consumes_the_wait_budget() { + let mut state = CountingState::blocked(); + state.first_row_delay = Duration::from_millis(30); + let (selected, skipped) = + select_requested_model(&state, Some(&limited_auth()), Duration::from_millis(5)) + .await + .unwrap(); + + assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped)); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 1); + assert_eq!(state.recent_reads.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn cancellation_stops_polling_and_candidate_queries() { + let state = CountingState::blocked(); + let auth = limited_auth(); + let mut selection = Box::pin(select_requested_model( + &state, + Some(&auth), + Duration::from_secs(1), + )); + tokio::select! { + result = &mut selection => panic!("selection completed before cancellation: {result:?}"), + _ = state.poll_observed.notified() => {} + } + drop(selection); + let recent_reads = state.recent_reads.load(Ordering::SeqCst); + tokio::time::sleep(Duration::from_millis(10)).await; + + assert!(recent_reads >= 2); + assert_eq!(state.recent_reads.load(Ordering::SeqCst), recent_reads); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn no_model_capability_selection_does_not_reenumerate_during_polls() { + let state = CountingState::blocked(); + let auth = limited_auth(); + let selected = select_with_auth_concurrency_wait( + &state, + Some(&auth), + crate::clock::current_unix_secs(), + Duration::from_millis(80), + POLL_INTERVAL, + |now| { + list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal( + &state, + &state, + "openai:chat", + "cache_1h", + false, + Some(&auth), + None, + now, + SchedulerOrderingConfig::default(), + ) + }, + ) + .await + .unwrap(); + + assert!(selected.is_empty()); + assert_eq!(state.format_reads.load(Ordering::SeqCst), 2); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 2); + assert!(state.recent_reads.load(Ordering::SeqCst) > 2); +} + +#[tokio::test] +async fn absent_or_disabled_auth_limits_do_not_poll() { + for limit in [None, Some(0), Some(-1)] { + let state = CountingState::blocked(); + let mut auth = limited_auth(); + auth.api_key_concurrent_limit = limit; + let (selected, _) = select_requested_model(&state, Some(&auth), Duration::from_secs(1)) + .await + .unwrap(); + assert_eq!(selected.len(), 1); + assert_eq!(state.row_reads.load(Ordering::SeqCst), 1); + assert_eq!(state.recent_reads.load(Ordering::SeqCst), 0); + } + let state = CountingState::blocked(); + assert_eq!( + select_requested_model(&state, None, Duration::from_secs(1)) + .await + .unwrap() + .0 + .len(), + 1 + ); + assert_eq!(state.recent_reads.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn auth_wait_keeps_candidate_row_and_lifecycle_counting_semantics() { + let state = CountingState::blocked(); + let mut duplicate = active_candidate(); + duplicate.id = "second-attempt-same-request".to_string(); + state.recent.lock().unwrap().push(duplicate); + let mut auth = limited_auth(); + auth.api_key_concurrent_limit = Some(2); + let (selected, skipped) = select_requested_model(&state, Some(&auth), Duration::ZERO) + .await + .unwrap(); + assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped)); + + let state = CountingState::blocked(); + state.recent.lock().unwrap()[0].finished_at_unix_ms = + Some(crate::clock::current_unix_secs() * 1000); + assert_eq!( + select_requested_model(&state, Some(&limited_auth()), Duration::ZERO) + .await + .unwrap() + .0 + .len(), + 1 + ); + + let state = CountingState::blocked(); + state.recent.lock().unwrap()[0].started_at_unix_ms = + Some(crate::clock::current_unix_secs().saturating_sub(301) * 1000); + assert_eq!( + select_requested_model(&state, Some(&limited_auth()), Duration::ZERO) + .await + .unwrap() + .0 + .len(), + 1 + ); +} diff --git a/apps/aether-gateway/src/scheduler/candidate/tests/mod.rs b/apps/aether-gateway/src/scheduler/candidate/tests/mod.rs index 9e2866002..55eb4f8ac 100644 --- a/apps/aether-gateway/src/scheduler/candidate/tests/mod.rs +++ b/apps/aether-gateway/src/scheduler/candidate/tests/mod.rs @@ -1,4 +1,5 @@ mod affinity; +mod concurrency_wait; mod model; mod required_capability; mod selection; diff --git a/apps/aether-gateway/src/shutdown_tests.rs b/apps/aether-gateway/src/shutdown_tests.rs new file mode 100644 index 000000000..4c322c4d0 --- /dev/null +++ b/apps/aether-gateway/src/shutdown_tests.rs @@ -0,0 +1,209 @@ +use std::sync::Arc; +use std::time::Duration; + +use axum::{body::Body, routing::get, Router}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::sync::Notify; +use tokio_util::sync::CancellationToken; + +use super::super::{serve_gateway_router, HttpConnectionBudget}; + +async fn start( + router: Router, +) -> ( + std::net::SocketAddr, + CancellationToken, + Arc, + tokio::task::JoinHandle>, +) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let shutdown = CancellationToken::new(); + let budget = Arc::new(HttpConnectionBudget::new(16)); + let stop = shutdown.clone(); + let shared = Arc::clone(&budget); + let server = tokio::spawn(async move { + serve_gateway_router( + vec![listener], + router, + shared, + 16, + 10_000, + 32_768, + 100, + stop, + ) + .await + .map_err(|error| error.to_string()) + }); + (address, shutdown, budget, server) +} + +async fn within(future: impl std::future::Future) -> T { + tokio::time::timeout(Duration::from_secs(3), future) + .await + .expect("shutdown deadline") +} + +#[tokio::test] +async fn gateway_shutdown_drains_in_flight_http1_and_http2_responses() { + for http2 in [false, true] { + let started = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + let handler_started = Arc::clone(&started); + let handler_release = Arc::clone(&release); + let router = Router::new().route( + "/", + get(move || { + let started = Arc::clone(&handler_started); + let release = Arc::clone(&handler_release); + async move { + started.notify_one(); + release.notified().await; + "complete response" + } + }), + ); + let (address, shutdown, budget, server) = start(router).await; + let client = if http2 { + reqwest::Client::builder().http2_prior_knowledge() + } else { + reqwest::Client::builder().http1_only() + } + .build() + .unwrap(); + let request = tokio::spawn(async move { + client + .get(format!("http://{address}/")) + .send() + .await + .unwrap() + .text() + .await + .unwrap() + }); + within(started.notified()).await; + shutdown.cancel(); + tokio::time::sleep(Duration::from_millis(20)).await; + assert!(!server.is_finished()); + assert!(!request.is_finished()); + assert!(TcpStream::connect(address).await.is_err()); + release.notify_one(); + assert_eq!(within(request).await.unwrap(), "complete response"); + within(server).await.unwrap().unwrap(); + assert_eq!(budget.snapshot().in_flight, 0); + } +} + +#[tokio::test] +async fn gateway_shutdown_closes_idle_protocol_detection_connections() { + let (address, shutdown, budget, server) = start(Router::new()).await; + let mut peer = TcpStream::connect(address).await.unwrap(); + within(async { + while budget.snapshot().in_flight == 0 { + tokio::task::yield_now().await; + } + }) + .await; + shutdown.cancel(); + within(server).await.unwrap().unwrap(); + assert_eq!(within(peer.read(&mut [0_u8; 1])).await.unwrap(), 0); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[tokio::test] +async fn gateway_shutdown_force_cancels_a_handler_without_socket_io() { + struct HandlerDrop(Arc); + impl Drop for HandlerDrop { + fn drop(&mut self) { + self.0.notify_one(); + } + } + for http2 in [false, true] { + let started = Arc::new(Notify::new()); + let dropped = Arc::new(Notify::new()); + let request_started = Arc::clone(&started); + let request_dropped = Arc::clone(&dropped); + let router = Router::new().route( + "/", + get(move || { + let started = Arc::clone(&request_started); + let dropped = Arc::clone(&request_dropped); + async move { + let _guard = HandlerDrop(dropped); + started.notify_one(); + std::future::pending::<&'static str>().await + } + }), + ); + let (address, shutdown, budget, server) = start(router).await; + let client = if http2 { + reqwest::Client::builder().http2_prior_knowledge() + } else { + reqwest::Client::builder().http1_only() + } + .build() + .unwrap(); + let request = + tokio::spawn(async move { client.get(format!("http://{address}/")).send().await }); + within(started.notified()).await; + shutdown.cancel(); + budget.force_close(); + within(server).await.unwrap().unwrap(); + within(dropped.notified()).await; + assert!(within(request).await.unwrap().is_err()); + assert_eq!(budget.snapshot().in_flight, 0); + } +} + +#[tokio::test] +async fn gateway_shutdown_force_closes_upgraded_io_and_waits_for_release() { + let upgraded_done = Arc::new(Notify::new()); + let done = Arc::clone(&upgraded_done); + let router = Router::new().route( + "/", + get(move |mut request: axum::extract::Request| { + let done = Arc::clone(&done); + async move { + let upgrade = hyper::upgrade::on(&mut request); + tokio::spawn(async move { + let upgraded = upgrade.await.unwrap(); + let mut io = hyper_util::rt::TokioIo::new(upgraded); + let error = io.read_u8().await.unwrap_err(); + assert_eq!(error.kind(), std::io::ErrorKind::ConnectionAborted); + drop(io); + done.notify_one(); + }); + axum::http::Response::builder() + .status(101) + .header("connection", "upgrade") + .header("upgrade", "echo") + .body(Body::empty()) + .unwrap() + } + }), + ); + let (address, shutdown, budget, server) = start(router).await; + let mut peer = TcpStream::connect(address).await.unwrap(); + peer.write_all( + b"GET / HTTP/1.1\r\nHost: localhost\r\nConnection: upgrade\r\nUpgrade: echo\r\n\r\n", + ) + .await + .unwrap(); + let mut header = Vec::new(); + within(async { + while !header.ends_with(b"\r\n\r\n") { + header.push(peer.read_u8().await.unwrap()); + } + }) + .await; + assert!(header.starts_with(b"HTTP/1.1 101")); + shutdown.cancel(); + tokio::time::sleep(Duration::from_millis(20)).await; + assert!(!server.is_finished()); + budget.force_close(); + within(upgraded_done.notified()).await; + within(server).await.unwrap().unwrap(); + assert_eq!(budget.snapshot().in_flight, 0); +} diff --git a/apps/aether-gateway/src/state/app.rs b/apps/aether-gateway/src/state/app.rs index 3a8a63ce8..1851b7c0c 100644 --- a/apps/aether-gateway/src/state/app.rs +++ b/apps/aether-gateway/src/state/app.rs @@ -8,6 +8,7 @@ use aether_data::repository::users::StoredUserGroup; use aether_data_contracts::repository::billing::UserDailyQuotaAvailabilityRecord; use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot; use aether_data_contracts::repository::usage::UsageCounterHealthSnapshot; +use aether_gateway_frontdoor::HttpConnectionBudget; use aether_runtime::ConcurrencyGate; use aether_runtime_state::{RuntimeSemaphore, RuntimeState}; use dashmap::DashMap; @@ -35,6 +36,7 @@ use super::{ }; const MIN_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 1_000; +const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_000; const MAX_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 600_000; const REQUEST_BODY_READ_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS"; const DEFAULT_REQUEST_BODY_BUFFER_BUDGET_MB: usize = 256; @@ -193,7 +195,9 @@ fn optional_env_duration_ms(key: &str, min_ms: u64, max_ms: u64) -> Option, min_ms: u64, max_ms: u64) -> Option { - let parsed = raw?.trim().parse::().ok()?; + let parsed = raw + .and_then(|value| value.trim().parse::().ok()) + .unwrap_or(DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS); if parsed == 0 { return None; } @@ -388,6 +392,7 @@ pub struct AppState { pub(crate) video_task_poller: Option, pub(crate) frontdoor_runtime_guards: Arc, pub(crate) request_body_buffer_budget: Arc, + pub(crate) http_connection_budget: Option>, pub(crate) request_gate: Option>, pub(crate) websocket_connection_gate: Option>, pub(crate) auth_snapshot_load_gate: Option>, @@ -456,6 +461,8 @@ pub struct AppState { Arc>, pub(crate) admin_monitoring_error_stats_reset_at: Arc>>, pub(crate) provider_delete_tasks: Arc>>, + pub(crate) pool_quota_probe_replenish: + Arc, #[cfg(test)] pub(crate) turnstile_siteverify_url_override: Option, #[cfg(test)] @@ -533,20 +540,22 @@ mod tests { }; #[test] - fn request_body_read_timeout_parser_defaults_to_disabled() { - assert_eq!( - parse_optional_duration_ms( - None, - MIN_REQUEST_BODY_READ_TIMEOUT_MS, - MAX_REQUEST_BODY_READ_TIMEOUT_MS, - ), - None - ); + fn request_body_read_timeout_parser_uses_a_finite_default() { + for value in [None, Some(""), Some("invalid"), Some("-1")] { + assert_eq!( + parse_optional_duration_ms( + value, + MIN_REQUEST_BODY_READ_TIMEOUT_MS, + MAX_REQUEST_BODY_READ_TIMEOUT_MS, + ), + Some(Duration::from_millis(DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS)), + ); + } } #[test] - fn request_body_read_timeout_parser_disables_zero_and_invalid_values() { - for value in ["", "invalid", "-1", "0", " 0 "] { + fn request_body_read_timeout_parser_only_disables_explicit_zero() { + for value in ["0", " 0 "] { assert_eq!( parse_optional_duration_ms( Some(value), diff --git a/apps/aether-gateway/src/state/core.rs b/apps/aether-gateway/src/state/core.rs index 44731f40d..de6ced1df 100644 --- a/apps/aether-gateway/src/state/core.rs +++ b/apps/aether-gateway/src/state/core.rs @@ -13,6 +13,7 @@ use aether_data::repository::proxy_nodes::{ use aether_data_contracts::repository::usage::{ UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot, }; +use aether_gateway_frontdoor::{HttpConnectionBudget, HttpConnectionBudgetSnapshot}; use aether_http::{apply_http_client_config, HttpClientConfig}; use aether_runtime::{ service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, MetricKind, @@ -358,6 +359,7 @@ impl AppState { request_body_buffer_budget: Arc::new(tokio::sync::Semaphore::new( frontdoor_runtime_guards.request_body_buffer_budget_permits, )), + http_connection_budget: None, request_gate: None, websocket_connection_gate: None, auth_snapshot_load_gate: frontdoor_runtime_guards @@ -438,6 +440,9 @@ impl AppState { local_execution_runtime_miss_diagnostics: Arc::new(DashMap::new()), admin_monitoring_error_stats_reset_at: Arc::new(StdMutex::new(None)), provider_delete_tasks: Arc::new(StdMutex::new(HashMap::new())), + pool_quota_probe_replenish: Arc::new( + crate::maintenance::PoolQuotaProbeReplenishCoordinator::default(), + ), #[cfg(test)] turnstile_siteverify_url_override: None, #[cfg(test)] @@ -614,6 +619,11 @@ impl AppState { self } + pub fn with_http_connection_budget(mut self, budget: Arc) -> Self { + self.http_connection_budget = Some(budget); + self + } + pub fn with_request_concurrency_limit(mut self, limit: usize) -> Self { let limit = limit.max(1); self.request_gate = Some(Arc::new(ConcurrencyGate::new("gateway_requests", limit))); @@ -635,6 +645,10 @@ impl AppState { } pub fn with_runtime_state(mut self, runtime_state: Arc) -> Self { + if !Arc::ptr_eq(&self.runtime_state, &runtime_state) { + self.pool_quota_probe_replenish = + Arc::new(crate::maintenance::PoolQuotaProbeReplenishCoordinator::default()); + } self.runtime_state = runtime_state; self.admin_security_blacklist_cache.clear(); self.admin_security_whitelist_cache.clear(); @@ -1666,6 +1680,9 @@ impl AppState { .unwrap_or(u64::MAX), ), ]); + if let Some(budget) = &self.http_connection_budget { + samples.extend(http_connection_metric_samples(&budget.snapshot())); + } if let Some(snapshot) = self.request_concurrency_snapshot() { samples.extend(snapshot.to_metric_samples("gateway_requests")); } @@ -1820,10 +1837,44 @@ impl AppState { &self.task_supervisor_metrics.snapshot(), )); samples.extend(crate::tokio_metrics::gateway_tokio_runtime_metric_samples()); + samples.extend(aether_runtime::logging_metric_samples()); samples.extend( crate::execution_runtime::transport::direct_reqwest_client_cache_metric_samples(), ); samples.extend(self.upstream_target_admission.metric_samples()); + let probe_replenish = self.pool_quota_probe_replenish.snapshot(); + samples.extend([ + MetricSample::new( + "pool_quota_probe_replenish_provider_capacity", + "Maximum providers with an active local request-triggered probe task.", + MetricKind::Gauge, + probe_replenish.capacity as u64, + ), + MetricSample::new( + "pool_quota_probe_replenish_active_providers", + "Providers with an active local request-triggered probe task.", + MetricKind::Gauge, + probe_replenish.active as u64, + ), + MetricSample::new( + "pool_quota_probe_replenish_started_total", + "Local request-triggered provider probe tasks admitted.", + MetricKind::Counter, + probe_replenish.started_total, + ), + MetricSample::new( + "pool_quota_probe_replenish_coalesced_total", + "Request triggers merged into an existing local provider probe task.", + MetricKind::Counter, + probe_replenish.coalesced_total, + ), + MetricSample::new( + "pool_quota_probe_replenish_capacity_rejected_total", + "Request-triggered probe tasks skipped because the local provider limit was reached.", + MetricKind::Counter, + probe_replenish.capacity_rejected_total, + ), + ]); samples.extend(crate::cache::candidate_page_cache_metric_samples()); samples.extend(crate::stage_metrics::gateway_stage_metric_samples()); samples.extend(self.tunnel.metric_samples()); @@ -2118,6 +2169,36 @@ impl AppState { state } + pub async fn shutdown_usage_runtime( + &self, + timeout: Duration, + ) -> Result<(), aether_data_contracts::DataLayerError> { + tokio::time::timeout(timeout, async { + let local_queue = self.runtime_state.is_memory().then(|| { + let queue: Arc = self.runtime_state.clone(); + queue + }); + self.usage_runtime + .shutdown_with_local_queue(timeout, local_queue) + .await?; + while self + .request_candidate_queue + .as_ref() + .is_some_and(|queue| queue.pending_writes() != 0) + { + tokio::time::sleep(Duration::from_millis(10)).await; + } + Ok(()) + }) + .await + .map_err(|_| { + aether_data_contracts::DataLayerError::TimedOut( + "gateway usage or candidate persistence did not drain before shutdown deadline" + .to_string(), + ) + })? + } + pub fn spawn_background_tasks(&self) -> crate::task_runtime::TaskSupervisor { let background_state = self.background_worker_state(); let mut supervisor = @@ -2277,6 +2358,41 @@ fn database_bounded_auth_load_limit( }) } +fn http_connection_metric_samples(snapshot: &HttpConnectionBudgetSnapshot) -> Vec { + vec![ + MetricSample::new( + "gateway_http_connections_limit", + "Maximum admitted inbound TCP connections across gateway listeners.", + MetricKind::Gauge, + u64::try_from(snapshot.limit).unwrap_or(u64::MAX), + ), + MetricSample::new( + "gateway_http_connections_in_flight", + "Currently admitted inbound TCP connections, including upgraded connections.", + MetricKind::Gauge, + u64::try_from(snapshot.in_flight).unwrap_or(u64::MAX), + ), + MetricSample::new( + "gateway_http_connections_high_watermark", + "Highest simultaneous admitted inbound TCP connection count.", + MetricKind::Gauge, + u64::try_from(snapshot.high_watermark).unwrap_or(u64::MAX), + ), + MetricSample::new( + "gateway_http_connections_rejected_total", + "Inbound TCP connections closed because the connection budget was full.", + MetricKind::Counter, + snapshot.rejected_total, + ), + MetricSample::new( + "gateway_http_connections_accept_errors_total", + "Listener accept errors retried with bounded backoff.", + MetricKind::Counter, + snapshot.accept_errors_total, + ), + ] +} + fn task_supervisor_metric_samples( snapshot: &aether_task_runtime::TaskSupervisorMetricsSnapshot, ) -> Vec { @@ -3190,6 +3306,18 @@ fn usage_runtime_metric_samples( snapshot: &usage::UsageRuntimeMetricsSnapshot, ) -> Vec { vec![ + MetricSample::new( + "usage_runtime_shutdown_started", "Whether local usage admission is closed for shutdown.", + MetricKind::Gauge, u64::from(snapshot.shutdown_started), + ), + MetricSample::new( + "usage_runtime_producers_in_flight", "Tracked requests and finalizers that can still submit usage events.", + MetricKind::Gauge, snapshot.producers_in_flight as u64, + ), + MetricSample::new( + "usage_runtime_delayed_lifecycle_pending", "Delayed lifecycle events still held in this process.", + MetricKind::Gauge, snapshot.delayed_lifecycle_pending as u64, + ), MetricSample::new( "usage_runtime_enabled", "Whether the gateway usage runtime is enabled.", @@ -3322,6 +3450,132 @@ fn usage_runtime_metric_samples( MetricKind::Gauge, u64::from(snapshot.retry_deferred_lifecycle_events), ), + MetricSample::new( + "usage_runtime_queue_payload_max_bytes", + "Maximum serialized JSON payload bytes for a new usage queue message.", + MetricKind::Gauge, + snapshot.queue_payload_max_bytes as u64, + ), + MetricSample::new( + "usage_runtime_queue_payload_downgraded_total", + "Process-wide enqueue and retry validation encoding attempts that omitted diagnostic data after exceeding the payload limit; not unique events.", + MetricKind::Counter, + snapshot.queue_payload_downgraded_total, + ), + MetricSample::new( + "usage_runtime_queue_payload_rejected_total", + "Process-wide enqueue and retry validation encoding attempts rejected because the payload exceeded its limit or billing facts could not be preserved; not unique events.", + MetricKind::Counter, + snapshot.queue_payload_rejected_total, + ), + MetricSample::new( + "usage_runtime_queue_read_payload_budget_bytes", + "Process-wide payload reservation limit for usage worker reads and reclaims; not a wire or heap limit.", + MetricKind::Gauge, + snapshot.queue_read_payload_budget_bytes as u64, + ), + MetricSample::new( + "usage_runtime_queue_read_batch_payload_bytes", + "Target payload reservation per usage worker batch, allowing at least one configured maximum payload.", + MetricKind::Gauge, + snapshot.queue_read_batch_payload_bytes as u64, + ), + MetricSample::new( + "usage_runtime_queue_read_payload_reserved_bytes", + "Process-wide logical payload bytes reserved by usage worker reads, reclaims and unprocessed batches.", + MetricKind::Gauge, + snapshot.queue_read_payload_reserved_bytes as u64, + ), + MetricSample::new( + "usage_runtime_queue_read_payload_waiters", + "Usage workers currently waiting for shared payload reservation capacity.", + MetricKind::Gauge, + snapshot.queue_read_payload_waiters as u64, + ), + MetricSample::new( + "usage_runtime_queue_read_payload_wait_total", + "Process-wide usage worker batch reservation attempts that had to wait for capacity.", + MetricKind::Counter, + snapshot.queue_read_payload_wait_total, + ), + MetricSample::new( + "usage_runtime_queue_read_actual_field_bytes_total", + "Cumulative key and value bytes observed in reserved usage worker batches, including repeated reclaims; excludes Redis framing and allocations.", + MetricKind::Counter, + snapshot.queue_read_actual_field_bytes_total, + ), + MetricSample::new( + "usage_runtime_queue_read_oversized_entries_total", + "Observed usage entries whose combined field value bytes exceed the consumer payload estimate, including repeated reclaims.", + MetricKind::Counter, + snapshot.queue_read_oversized_entries_total, + ), + MetricSample::new( + "usage_runtime_queue_read_oversized_batches_total", + "Observed usage batches whose combined field value bytes exceed their initial reservation; entries remain eligible for processing.", + MetricKind::Counter, + snapshot.queue_read_oversized_batches_total, + ), + MetricSample::new( + "usage_runtime_event_capture_memory_budget_bytes", + "Process-wide diagnostic JSON heap estimate budget for retained usage events.", + MetricKind::Gauge, + snapshot.event_capture_memory_budget_bytes as u64, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_budget_bytes", + "Process-wide logical raw string and JSON reservation limit for dead letter encoding; excludes Redis buffers.", + MetricKind::Gauge, + snapshot.dlq_encoding_budget_bytes as u64, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_max_jobs", + "Maximum admitted dead letter encoding and write jobs per process.", + MetricKind::Gauge, + snapshot.dlq_encoding_max_jobs as u64, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_reserved_bytes", + "Logical raw string and JSON bytes reserved by admitted dead letter jobs.", + MetricKind::Gauge, + snapshot.dlq_encoding_reserved_bytes as u64, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_active_jobs", + "Admitted dead letter jobs awaiting or performing encoding or queue writes.", + MetricKind::Gauge, + snapshot.dlq_encoding_active_jobs as u64, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_capacity_rejected_total", + "Dead letter attempts rejected because encoding byte or job capacity was occupied; source remains pending.", + MetricKind::Counter, + snapshot.dlq_encoding_capacity_rejected_total, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_oversized_rejected_total", + "Dead letter attempts rejected because the conservative encoding reservation exceeded the total budget or overflowed.", + MetricKind::Counter, + snapshot.dlq_encoding_oversized_rejected_total, + ), + MetricSample::new( + "usage_runtime_dlq_encoding_encoded_total", + "Completed dead letter JSON encodings, including repeated attempts; not successful archives.", + MetricKind::Counter, + snapshot.dlq_encoding_encoded_total, + ), + MetricSample::new( + "usage_runtime_event_capture_memory_retained_bytes", + "Estimated diagnostic JSON heap retained by budgeted usage events and their clones.", + MetricKind::Gauge, + snapshot.event_capture_memory_retained_bytes as u64, + ), + MetricSample::new( + "usage_runtime_event_capture_memory_downgraded_total", + "Usage event diagnostic captures omitted after memory budget exhaustion.", + MetricKind::Counter, + snapshot.event_capture_memory_downgraded_total, + ), MetricSample::new( "usage_runtime_terminal_submission_limit", "Maximum concurrent end-to-end terminal usage submissions.", @@ -3700,6 +3954,12 @@ fn usage_runtime_metric_samples( MetricKind::Counter, snapshot.enqueue_retry_failed_total, ), + MetricSample::new( + "usage_runtime_enqueue_retry_permanent_failure_total", + "Usage enqueue retry submissions rejected or queued events terminated because of permanent input errors.", + MetricKind::Counter, + snapshot.enqueue_retry_permanent_failure_total, + ), MetricSample::new( "usage_runtime_enqueue_retry_closed_or_unavailable_total", "Total usage events rejected because the local enqueue dispatcher was full, closed, or unavailable.", @@ -3972,6 +4232,126 @@ mod tests { assert_eq!(database_bounded_auth_load_limit(Some(64), None), Some(64)); } + #[tokio::test] + async fn http_connection_budget_is_optional_and_shared_by_app_state_clones() { + let state = AppState::new().expect("app state should build"); + assert!(state.http_connection_budget.is_none()); + let budget = Arc::new(super::HttpConnectionBudget::new(1)); + let state = state.with_http_connection_budget(Arc::clone(&budget)); + let cloned = state.clone(); + assert!(Arc::ptr_eq( + cloned.http_connection_budget.as_ref().unwrap(), + &budget, + )); + + let connection = budget.try_admit(()).expect("first connection admitted"); + assert!(cloned + .http_connection_budget + .as_ref() + .unwrap() + .try_admit(()) + .is_err()); + let samples = state.collect_metric_samples().await; + for (name, expected) in [ + ("gateway_http_connections_limit", 1), + ("gateway_http_connections_in_flight", 1), + ("gateway_http_connections_high_watermark", 1), + ("gateway_http_connections_rejected_total", 1), + ("gateway_http_connections_accept_errors_total", 0), + ] { + assert_eq!( + samples + .iter() + .find(|sample| sample.name == name) + .unwrap() + .value, + expected, + ); + } + drop(connection); + assert_eq!(budget.snapshot().in_flight, 0); + } + + #[test] + fn http_connection_metrics_export_all_budget_counters() { + let samples = super::http_connection_metric_samples(&super::HttpConnectionBudgetSnapshot { + limit: 4096, + in_flight: 11, + high_watermark: 30, + rejected_total: 7, + accept_errors_total: 3, + }); + for (name, kind, value) in [ + ("gateway_http_connections_limit", MetricKind::Gauge, 4096), + ("gateway_http_connections_in_flight", MetricKind::Gauge, 11), + ( + "gateway_http_connections_high_watermark", + MetricKind::Gauge, + 30, + ), + ( + "gateway_http_connections_rejected_total", + MetricKind::Counter, + 7, + ), + ( + "gateway_http_connections_accept_errors_total", + MetricKind::Counter, + 3, + ), + ] { + let matching = samples + .iter() + .filter(|sample| sample.name == name) + .collect::>(); + assert_eq!(matching.len(), 1); + assert_eq!((matching[0].kind, matching[0].value), (kind, value)); + } + } + + #[test] + fn usage_runtime_metrics_export_queue_payload_limit_and_attempt_counters() { + let mut snapshot = crate::usage::UsageRuntimeMetricsSnapshot::default(); + snapshot.queue_payload_max_bytes = 1024 * 1024; + snapshot.queue_payload_downgraded_total = 11; + snapshot.queue_payload_rejected_total = 3; + snapshot.enqueue_retry_permanent_failure_total = 2; + let samples = usage_runtime_metric_samples(&snapshot); + for (name, kind, value) in [ + ( + "usage_runtime_queue_payload_max_bytes", + MetricKind::Gauge, + 1024 * 1024, + ), + ( + "usage_runtime_queue_payload_downgraded_total", + MetricKind::Counter, + 11, + ), + ( + "usage_runtime_queue_payload_rejected_total", + MetricKind::Counter, + 3, + ), + ( + "usage_runtime_enqueue_retry_permanent_failure_total", + MetricKind::Counter, + 2, + ), + ] { + let matching = samples + .iter() + .filter(|sample| sample.name == name) + .collect::>(); + assert_eq!( + matching.len(), + 1, + "each payload metric must be emitted once" + ); + assert_eq!((matching[0].kind, matching[0].value), (kind, value)); + } + } + #[test] fn usage_runtime_metrics_export_first_byte_batch_counters() { let mut snapshot = crate::usage::UsageRuntimeMetricsSnapshot::default(); diff --git a/apps/aether-gateway/src/state/integrations.rs b/apps/aether-gateway/src/state/integrations.rs index 8d61f0a77..db94e5786 100644 --- a/apps/aether-gateway/src/state/integrations.rs +++ b/apps/aether-gateway/src/state/integrations.rs @@ -671,7 +671,7 @@ impl SchedulerRuntimeState for AppState { &self, limit: usize, ) -> Result, GatewayError> { - AppState::read_recent_request_candidates(self, limit).await + AppState::read_recent_runtime_request_candidates(self, limit).await } fn provider_key_rpm_reset_at(&self, key_id: &str, now_unix_secs: u64) -> Option { diff --git a/apps/aether-gateway/src/state/runtime/candidate_queries.rs b/apps/aether-gateway/src/state/runtime/candidate_queries.rs index eef5ef3f4..0b69201f9 100644 --- a/apps/aether-gateway/src/state/runtime/candidate_queries.rs +++ b/apps/aether-gateway/src/state/runtime/candidate_queries.rs @@ -109,6 +109,16 @@ impl AppState { .map_err(|err| GatewayError::Internal(err.to_string())) } + pub(crate) async fn read_recent_runtime_request_candidates( + &self, + limit: usize, + ) -> Result, GatewayError> { + self.data + .list_recent_runtime_request_candidates(limit) + .await + .map_err(|err| GatewayError::Internal(err.to_string())) + } + pub(crate) async fn upsert_request_candidate( &self, mut candidate: candidates::UpsertRequestCandidateRecord, diff --git a/apps/aether-gateway/src/testkit.rs b/apps/aether-gateway/src/testkit.rs index 04d40c6e8..95ec6910d 100644 --- a/apps/aether-gateway/src/testkit.rs +++ b/apps/aether-gateway/src/testkit.rs @@ -19,6 +19,17 @@ use sha2::{Digest, Sha256}; use crate::data::GatewayDataState; use crate::AppState; +pub use crate::tunnel::build_tunnel_pressure_router; + +pub async fn gateway_metric_samples( + state: &AppState, +) -> Result, String> { + if !state.prewarm_metric_snapshot().await { + return Err("gateway harness metric refresh timed out".to_string()); + } + Ok(state.metric_samples().await) +} + #[derive(Debug, Clone)] pub struct OpenAiChatPressureTarget { pub base_url: String, diff --git a/apps/aether-gateway/src/tests/architecture/workspace_tiers.rs b/apps/aether-gateway/src/tests/architecture/workspace_tiers.rs index 5f808e06f..7687bcbe9 100644 --- a/apps/aether-gateway/src/tests/architecture/workspace_tiers.rs +++ b/apps/aether-gateway/src/tests/architecture/workspace_tiers.rs @@ -183,7 +183,8 @@ fn gateway_tunnel_protocol_path_is_a_thin_compatibility_facade() { fn frontdoor_owns_bounded_request_body_buffering() { let frontdoor = read_workspace_file("crates/aether-gateway/frontdoor/src/body.rs"); assert!(frontdoor.contains("acquire_many_owned")); - assert!(frontdoor.contains("to_bytes(body, body_limit)")); + assert!(frontdoor.contains("body.into_data_stream()")); + assert!(frontdoor.contains("try_reserve_bytes")); assert!(frontdoor.contains("BodyBufferReservation")); let gateway = read_workspace_file("apps/aether-gateway/src/handlers/proxy/body_buffer.rs"); diff --git a/apps/aether-gateway/src/tunnel/mod.rs b/apps/aether-gateway/src/tunnel/mod.rs index bde158308..661d066af 100644 --- a/apps/aether-gateway/src/tunnel/mod.rs +++ b/apps/aether-gateway/src/tunnel/mod.rs @@ -59,6 +59,37 @@ pub use embedded::{ ControlPlaneClient as TunnelControlPlaneClient, }; +#[cfg(feature = "testkit")] +pub fn build_tunnel_pressure_router( + state: TunnelRuntimeState, + instance_id: &str, + secret: &[u8], +) -> Result { + validate_tunnel_relay_auth_secret(secret)?; + let directory = TunnelAttachmentDirectory::from_parts(instance_id, None::, 90); + let state = state.with_relay_auth( + instance_id, + Some(secret.to_vec()), + Arc::clone(&directory.runtime_state), + ); + let runtime_router = build_tunnel_runtime_router_with_state(state.clone()); + let mut gateway = AppState::new().map_err(|error| error.to_string())?; + gateway.tunnel = EmbeddedTunnelState { + inner: state, + attachment_directory: directory, + relay_auth_secret: Ok(Arc::from(secret)), + }; + // HTTP relay requests must pass the same authentication and verified spool + // preparation as the gateway before reaching the embedded local dispatcher. + Ok(axum::Router::new() + .route( + TUNNEL_RELAY_PATH_PATTERN, + axum::routing::post(relay_request), + ) + .with_state(gateway) + .fallback_service(runtime_router)) +} + const DEFAULT_ATTACHMENT_TTL_SECS: u64 = 90; const TUNNEL_ATTACHMENT_KEY_PREFIX: &str = "tunnel.attachments."; const TUNNEL_ATTACHMENT_REDIS_KEY_PREFIX: &str = "tunnel:attachments:"; diff --git a/apps/aether-tunnel/src/app.rs b/apps/aether-tunnel/src/app.rs index 8b6906411..ebd4f99c4 100644 --- a/apps/aether-tunnel/src/app.rs +++ b/apps/aether-tunnel/src/app.rs @@ -480,6 +480,7 @@ async fn diagnostics_metrics( AxumState(diagnostics): AxumState, ) -> impl axum::response::IntoResponse { let mut samples = diagnostics.state.metric_samples().await; + samples.extend(aether_runtime::logging_metric_samples()); let servers = diagnostics.server_contexts.lock().await.clone(); for server in servers { samples.extend(server.metric_samples()); diff --git a/apps/aether-tunnel/src/main.rs b/apps/aether-tunnel/src/main.rs index ee621be81..d2b01e23d 100644 --- a/apps/aether-tunnel/src/main.rs +++ b/apps/aether-tunnel/src/main.rs @@ -53,8 +53,13 @@ fn build_command() -> clap::Command { .subcommand_negates_reqs(true) } +fn main() -> anyhow::Result<()> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> anyhow::Result<()> { +async fn run() -> anyhow::Result<()> { rustls::crypto::ring::default_provider() .install_default() .map_err(|_| anyhow::anyhow!("Failed to install rustls CryptoProvider"))?; diff --git a/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs b/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs index 5f5875ab2..de5b46e6d 100644 --- a/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs +++ b/crates/aether-ai/formats/src/formats/gemini/generate_content/stream.rs @@ -24,6 +24,8 @@ struct GeminiProviderToolResultState { #[derive(Default)] pub struct GeminiProviderState { + terminal_observation_only: bool, + observed_tool_calls: bool, response_id: Option, model: Option, started: bool, @@ -37,6 +39,13 @@ pub struct GeminiProviderState { } impl GeminiProviderState { + pub(crate) fn terminal_observation() -> Self { + Self { + terminal_observation_only: true, + ..Self::default() + } + } + fn identity(&self, report_context: &Value) -> (String, String) { resolve_identity( self.response_id.as_deref(), @@ -133,14 +142,21 @@ impl GeminiProviderState { let Some(part_object) = part.as_object() else { continue; }; - let reasoning_signature = part_object - .get("thoughtSignature") - .or_else(|| part_object.get("thought_signature")) - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(ToOwned::to_owned); + let reasoning_signature = if self.terminal_observation_only { + None + } else { + part_object + .get("thoughtSignature") + .or_else(|| part_object.get("thought_signature")) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(ToOwned::to_owned) + }; if let Some(text) = render_gemini_part_as_text(part_object) { + if self.terminal_observation_only { + continue; + } let is_reasoning = part_object .get("thought") .and_then(Value::as_bool) @@ -193,6 +209,9 @@ impl GeminiProviderState { .or_else(|| part_object.get("function_response")) .and_then(Value::as_object) { + if self.terminal_observation_only { + continue; + } let tool_use_id = function_response .get("id") .and_then(Value::as_str) @@ -240,6 +259,9 @@ impl GeminiProviderState { else { if let Some(content_part) = canonical_content_part_from_gemini_part(part_object) { + if self.terminal_observation_only { + continue; + } let should_emit = self .content_parts .get(&index) @@ -260,6 +282,10 @@ impl GeminiProviderState { } continue; }; + if self.terminal_observation_only { + self.observed_tool_calls = true; + continue; + } let tool_state = self.tool_calls.entry(index).or_default(); tool_state.call_id = function_call .get("id") @@ -329,7 +355,7 @@ impl GeminiProviderState { if let Some(finish_reason) = candidate_object.get("finishReason").and_then(Value::as_str) { - let has_tool_calls = !self.tool_calls.is_empty(); + let has_tool_calls = self.observed_tool_calls || !self.tool_calls.is_empty(); let mut finish_reason = normalize_openai_finish_reason(map_gemini_stream_finish_reason(finish_reason)); if has_tool_calls && finish_reason.as_deref().is_none_or(|value| value == "stop") { @@ -936,6 +962,232 @@ mod tests { format!("data: {}\n", value).into_bytes() } + fn terminal_frames(frames: Vec) -> Vec { + frames + .into_iter() + .filter(|frame| { + matches!( + frame.event, + CanonicalStreamEvent::Start + | CanonicalStreamEvent::Finish { .. } + | CanonicalStreamEvent::UnknownEvent(_) + ) + }) + .collect() + } + + fn observation_record(parts: Vec, finish_reason: Option<&str>) -> Value { + let mut record = json!({ + "responseId": "resp_observation", + "modelVersion": "gemini-2.5-pro", + "candidates": [{"content": {"parts": parts}}], + "usageMetadata": { + "promptTokenCount": 22, + "cachedContentTokenCount": 7, + "candidatesTokenCount": 13, + "thoughtsTokenCount": 5, + "totalTokenCount": 40 + } + }); + if let Some(finish_reason) = finish_reason { + record["candidates"][0]["finishReason"] = json!(finish_reason); + } + record + } + + fn assert_terminal_record_matches( + normal: &mut GeminiProviderState, + observer: &mut GeminiProviderState, + record: Value, + ) -> Vec { + let context = json!({"mapped_model": "fallback-model"}); + let expected = terminal_frames( + normal + .push_line(&context, data_line(record.clone())) + .expect("normal provider parser"), + ); + let actual = observer + .push_line(&context, data_line(record)) + .expect("terminal provider parser"); + assert_eq!( + actual, expected, + "terminal frames must retain their full payloads" + ); + actual + } + + fn assert_observer_has_no_content_buffers(observer: &GeminiProviderState) { + assert!(observer.text_parts.is_empty()); + assert!(observer.reasoning_parts.is_empty()); + assert!(observer.reasoning_signatures.is_empty()); + assert!(observer.content_parts.is_empty()); + assert!(observer.tool_calls.is_empty()); + assert!(observer.tool_results.is_empty()); + } + + #[test] + fn gemini_terminal_observation_does_not_retain_long_stream_content() { + let mut normal = GeminiProviderState::default(); + let mut observer = GeminiProviderState::terminal_observation(); + let media = "YQ==".repeat(4_096); + for index in 0..32 { + let text = format!("record-{index}:{}", "x".repeat(16_384)); + assert_terminal_record_matches( + &mut normal, + &mut observer, + observation_record( + vec![ + json!({"text": text}), + json!({"text": text, "thought": true, "thoughtSignature": media}), + json!({"functionCall": {"id": format!("call-{index}"), "name": "lookup", "args": {"value": text}}}), + json!({"functionResponse": {"name": "lookup", "response": {"result": text}}}), + json!({"inlineData": {"mimeType": "image/png", "data": media}}), + json!({"inline_data": {"mime_type": "audio/wav", "data": media}}), + json!({"inlineData": {"mimeType": "application/pdf", "data": media}}), + ], + None, + ), + ); + assert_observer_has_no_content_buffers(&observer); + assert!(observer.observed_tool_calls); + } + assert!(!normal.text_parts.is_empty()); + assert!(!normal.reasoning_parts.is_empty()); + assert!(!normal.reasoning_signatures.is_empty()); + assert!(!normal.content_parts.is_empty()); + assert!(!normal.tool_calls.is_empty()); + assert!(normal + .tool_results + .values() + .any(|state| state.content.len() > 16_384)); + let frames = assert_terminal_record_matches( + &mut normal, + &mut observer, + observation_record(Vec::new(), Some("STOP")), + ); + assert!(frames.iter().any(|frame| matches!( + &frame.event, + CanonicalStreamEvent::Finish { finish_reason: Some(reason), usage: Some(usage) } + if reason == "tool_calls" && usage.input_tokens == 22 && usage.output_tokens == 18 + && usage.cache_read_tokens == 7 && usage.reasoning_tokens == 5 && usage.total_tokens == 40 + ))); + assert_observer_has_no_content_buffers(&observer); + } + + #[test] + fn gemini_terminal_observation_matches_known_and_unknown_part_classification() { + let parts = vec![ + Value::Null, + json!({}), + json!({"text": ""}), + json!({"text": 17}), + json!({"text": "", "thought_signature": "signature"}), + json!({"thoughtSignature": "signature"}), + json!({"executableCode": {"language": "python", "code": "print(1)"}}), + json!({"codeExecutionResult": {"output": "1"}}), + json!({"functionResponse": {}}), + json!({"function_response": {"response": ["result"]}}), + json!({"functionCall": {}}), + json!({"functionCall": null}), + json!({"inlineData": {"mimeType": "image/png", "data": "YQ=="}}), + json!({"inline_data": {"mime_type": "audio/wav", "data": "YQ=="}}), + json!({"inlineData": {"mimeType": "application/pdf", "data": "YQ=="}}), + json!({"inlineData": {"mimeType": "", "data": "YQ=="}}), + json!({"inlineData": {"mimeType": "image/png", "data": ""}}), + json!({"file_data": {"file_uri": "gs://test/file", "mime_type": "application/pdf"}}), + json!({"fileData": {"fileUri": ""}}), + json!({"futurePart": {"kept": true}}), + ]; + for part in parts { + let mut normal = GeminiProviderState::default(); + let mut observer = GeminiProviderState::terminal_observation(); + assert_terminal_record_matches( + &mut normal, + &mut observer, + json!({"responseId": "outer", "response": observation_record(vec![part], Some("STOP"))}), + ); + assert_eq!(observer.response_id.as_deref(), Some("resp_observation")); + assert_observer_has_no_content_buffers(&observer); + } + } + + #[test] + fn gemini_terminal_observation_preserves_finish_reasons_and_eof() { + for has_tool in [false, true] { + for reason in [ + None, + Some("STOP"), + Some("MAX_TOKENS"), + Some("RECITATION"), + Some("FUTURE_REASON"), + ] { + let mut normal = GeminiProviderState::default(); + let mut observer = GeminiProviderState::terminal_observation(); + let part = if has_tool { + json!({"functionCall": {"name": "lookup", "args": {"query": "test"}}}) + } else { + json!({"text": "partial output"}) + }; + assert_terminal_record_matches( + &mut normal, + &mut observer, + observation_record(vec![part], None), + ); + if let Some(reason) = reason { + assert_terminal_record_matches( + &mut normal, + &mut observer, + observation_record(Vec::new(), Some(reason)), + ); + } + let context = json!({}); + assert_eq!( + observer.finish(&context).expect("observer EOF"), + terminal_frames(normal.finish(&context).expect("normal EOF")) + ); + assert!(observer + .finish(&context) + .expect("idempotent EOF") + .is_empty()); + assert_observer_has_no_content_buffers(&observer); + } + } + } + + #[test] + fn gemini_terminal_observation_preserves_failure_payloads_with_usage() { + for reason in [ + "MALFORMED_FUNCTION_CALL", + "UNEXPECTED_TOOL_CALL", + "TOO_MANY_TOOL_CALLS", + "MISSING_THOUGHT_SIGNATURE", + "MALFORMED_RESPONSE", + ] { + for content in [ + Value::Null, + json!({}), + json!({"parts": []}), + json!({"parts": [{"text": ""}]}), + ] { + let mut normal = GeminiProviderState::default(); + let mut observer = GeminiProviderState::terminal_observation(); + let mut record = observation_record(Vec::new(), Some(reason)); + record["candidates"][0]["content"] = content; + record["candidates"][0]["finishMessage"] = json!("provider failure details"); + let frames = assert_terminal_record_matches(&mut normal, &mut observer, record); + assert!(frames.iter().any(|frame| matches!( + &frame.event, + CanonicalStreamEvent::UnknownEvent(payload) + if payload["response"]["error"]["code"] == reason + && payload["response"]["error"]["message"] == "provider failure details" + && payload["response"]["usage"]["input_tokens"] == 22 + ))); + assert!(observer.finish(&json!({})).expect("failed EOF").is_empty()); + assert_observer_has_no_content_buffers(&observer); + } + } + } + #[test] fn gemini_provider_state_emits_unknown_events_for_unknown_parts() { let mut state = GeminiProviderState::default(); diff --git a/crates/aether-ai/formats/src/formats/openai/chat/stream.rs b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs index 22dc3e502..a5a29524e 100644 --- a/crates/aether-ai/formats/src/formats/openai/chat/stream.rs +++ b/crates/aether-ai/formats/src/formats/openai/chat/stream.rs @@ -1,6 +1,7 @@ use std::collections::{BTreeMap, BTreeSet}; use serde_json::{json, Map, Value}; +use sha2::{Digest, Sha256}; use crate::formats::openai::namespace::NamespaceToolAliases; use crate::formats::openai::responses::{ @@ -33,6 +34,7 @@ struct OpenAIChatProviderToolState { #[derive(Default)] pub struct OpenAIChatProviderState { + terminal_only: bool, response_id: Option, model: Option, actual_service_tier: Option, @@ -59,6 +61,7 @@ struct OpenAIResponsesProviderToolResultState { #[derive(Default)] pub struct OpenAIResponsesProviderState { + terminal_only: bool, response_id: Option, model: Option, actual_service_tier: Option, @@ -71,11 +74,24 @@ pub struct OpenAIResponsesProviderState { tool_results: BTreeMap, tool_index_by_key: BTreeMap, image_item_keys: BTreeSet, - opaque_completed_item_keys: BTreeSet, + opaque_completed_item_keys: BTreeSet, last_tool_index: Option, } +#[derive(PartialEq, Eq, PartialOrd, Ord)] +enum OpenAIResponsesOutputItemKey { + Full(String), + Digest([u8; 32]), +} + impl OpenAIChatProviderState { + pub(crate) fn terminal_observation() -> Self { + Self { + terminal_only: true, + ..Self::default() + } + } + pub(crate) fn actual_service_tier(&self) -> Option<&str> { self.actual_service_tier.as_deref() } @@ -243,12 +259,14 @@ impl OpenAIChatProviderState { recognized_delta = true; if !content.is_empty() { self.ensure_started(report_context, &mut out); - let (id, model) = self.identity(report_context); - out.push(CanonicalStreamFrame { - id, - model, - event: CanonicalStreamEvent::TextDelta(content.to_string()), - }); + if !self.terminal_only { + let (id, model) = self.identity(report_context); + out.push(CanonicalStreamFrame { + id, + model, + event: CanonicalStreamEvent::TextDelta(content.to_string()), + }); + } } } else if delta.contains_key("content") { recognized_delta = true; @@ -258,12 +276,16 @@ impl OpenAIChatProviderState { recognized_delta = true; if !reasoning_content.is_empty() { self.ensure_started(report_context, &mut out); - let (id, model) = self.identity(report_context); - out.push(CanonicalStreamFrame { - id, - model, - event: CanonicalStreamEvent::ReasoningDelta(reasoning_content.to_string()), - }); + if !self.terminal_only { + let (id, model) = self.identity(report_context); + out.push(CanonicalStreamFrame { + id, + model, + event: CanonicalStreamEvent::ReasoningDelta( + reasoning_content.to_string(), + ), + }); + } } } else if delta.contains_key("reasoning_content") { recognized_delta = true; @@ -272,68 +294,71 @@ impl OpenAIChatProviderState { if let Some(tool_calls) = delta.get("tool_calls").and_then(Value::as_array) { recognized_delta = true; self.ensure_started(report_context, &mut out); - let (id, model) = self.identity(report_context); - for tool_call in tool_calls { - let Some(tool_call_object) = tool_call.as_object() else { - continue; - }; - let index = tool_call_object - .get("index") - .and_then(Value::as_u64) - .map(|value| value as usize) - .unwrap_or(0); - let state = self.tool_calls.entry(index).or_default(); - if let Some(call_id) = tool_call_object.get("id").and_then(Value::as_str) { - state.id = Some(call_id.to_string()); - } - let mut arguments = None; - if let Some(function) = - tool_call_object.get("function").and_then(Value::as_object) - { - if let Some(name) = function.get("name").and_then(Value::as_str) { - state.name = Some(name.to_string()); + if !self.terminal_only { + let (id, model) = self.identity(report_context); + for tool_call in tool_calls { + let Some(tool_call_object) = tool_call.as_object() else { + continue; + }; + let index = tool_call_object + .get("index") + .and_then(Value::as_u64) + .map(|value| value as usize) + .unwrap_or(0); + let state = self.tool_calls.entry(index).or_default(); + if let Some(call_id) = tool_call_object.get("id").and_then(Value::as_str) { + state.id = Some(call_id.to_string()); } - arguments = function - .get("arguments") - .and_then(Value::as_str) - .filter(|arguments| !arguments.is_empty()); - } - if !state.started_emitted { - if let Some(arguments) = arguments { - state.pending_arguments.push_str(arguments); - } - if let (Some(call_id), Some(name)) = (state.id.clone(), state.name.clone()) + let mut arguments = None; + if let Some(function) = + tool_call_object.get("function").and_then(Value::as_object) { - out.push(CanonicalStreamFrame { - id: id.clone(), - model: model.clone(), - event: CanonicalStreamEvent::ToolCallStart { - index, - call_id, - name, - }, - }); - state.started_emitted = true; - if !state.pending_arguments.is_empty() { + if let Some(name) = function.get("name").and_then(Value::as_str) { + state.name = Some(name.to_string()); + } + arguments = function + .get("arguments") + .and_then(Value::as_str) + .filter(|arguments| !arguments.is_empty()); + } + if !state.started_emitted { + if let Some(arguments) = arguments { + state.pending_arguments.push_str(arguments); + } + if let (Some(call_id), Some(name)) = + (state.id.clone(), state.name.clone()) + { out.push(CanonicalStreamFrame { id: id.clone(), model: model.clone(), - event: CanonicalStreamEvent::ToolCallArgumentsDelta { + event: CanonicalStreamEvent::ToolCallStart { index, - arguments: std::mem::take(&mut state.pending_arguments), + call_id, + name, }, }); + state.started_emitted = true; + if !state.pending_arguments.is_empty() { + out.push(CanonicalStreamFrame { + id: id.clone(), + model: model.clone(), + event: CanonicalStreamEvent::ToolCallArgumentsDelta { + index, + arguments: std::mem::take(&mut state.pending_arguments), + }, + }); + } } + } else if let Some(arguments) = arguments { + out.push(CanonicalStreamFrame { + id: id.clone(), + model: model.clone(), + event: CanonicalStreamEvent::ToolCallArgumentsDelta { + index, + arguments: arguments.to_string(), + }, + }); } - } else if let Some(arguments) = arguments { - out.push(CanonicalStreamFrame { - id: id.clone(), - model: model.clone(), - event: CanonicalStreamEvent::ToolCallArgumentsDelta { - index, - arguments: arguments.to_string(), - }, - }); } } } else if delta.contains_key("tool_calls") { @@ -389,6 +414,13 @@ impl OpenAIChatProviderState { } impl OpenAIResponsesProviderState { + pub(crate) fn terminal_observation() -> Self { + Self { + terminal_only: true, + ..Self::default() + } + } + pub(crate) fn actual_service_tier(&self) -> Option<&str> { self.actual_service_tier.as_deref() } @@ -494,6 +526,10 @@ impl OpenAIResponsesProviderState { if text.is_empty() { return; } + if self.terminal_only { + self.ensure_started(report_context, out); + return; + } self.text_parts.entry(key).or_default().push_str(text); self.ensure_started(report_context, out); let (id, model) = self.identity(report_context); @@ -511,6 +547,12 @@ impl OpenAIResponsesProviderState { key: String, text: &str, ) { + if self.terminal_only { + if !text.is_empty() { + self.ensure_started(report_context, out); + } + return; + } let missing = { let current = self.text_parts.entry(key).or_default(); let missing = if text.starts_with(current.as_str()) { @@ -543,6 +585,12 @@ impl OpenAIResponsesProviderState { out: &mut Vec, reasoning: &str, ) { + if self.terminal_only { + if !reasoning.is_empty() { + self.ensure_started(report_context, out); + } + return; + } let missing = if reasoning.starts_with(&self.reasoning) { reasoning[self.reasoning.len()..].to_string() } else if self.reasoning == reasoning { @@ -573,6 +621,10 @@ impl OpenAIResponsesProviderState { if text.is_empty() { return; } + if self.terminal_only { + self.ensure_started(report_context, out); + return; + } let missing = { let current = self.reasoning_parts.entry(summary_index).or_default(); let missing = if text.starts_with(current.as_str()) { @@ -608,6 +660,9 @@ impl OpenAIResponsesProviderState { out: &mut Vec, index: usize, ) { + if self.terminal_only { + return; + } let (id, model) = self.identity(report_context); let Some(state) = self.tool_calls.get_mut(&index) else { return; @@ -806,12 +861,13 @@ impl OpenAIResponsesProviderState { if let Some(name) = incoming_chat_name { state.name = name; } - let completed_arguments = item - .get("arguments") - .and_then(Value::as_str) - .unwrap_or_default() - .to_string(); - Self::merge_tool_call_arguments(state, &completed_arguments); + if !self.terminal_only { + let completed_arguments = item + .get("arguments") + .and_then(Value::as_str) + .unwrap_or_default(); + Self::merge_tool_call_arguments(state, completed_arguments); + } self.emit_ready_function_call(report_context, out, index); } @@ -832,10 +888,14 @@ impl OpenAIResponsesProviderState { .filter(|value| !value.is_empty()) .unwrap_or("custom_tool") .to_string(); - let arguments = tool_arguments_from_maybe_json_string( - item.get("input").or_else(|| item.get("arguments")), - "input", - ); + let arguments = if self.terminal_only { + String::new() + } else { + tool_arguments_from_maybe_json_string( + item.get("input").or_else(|| item.get("arguments")), + "input", + ) + }; self.emit_generic_tool_call_item(report_context, out, item, output_index, name, arguments); } @@ -852,16 +912,20 @@ impl OpenAIResponsesProviderState { "shell_call" => "shell", _ => return, }; - let arguments = tool_arguments_from_named_fields( - item, - &[ - "action", - "environment", - "status", - "created_by", - "max_output_length", - ], - ); + let arguments = if self.terminal_only { + String::new() + } else { + tool_arguments_from_named_fields( + item, + &[ + "action", + "environment", + "status", + "created_by", + "max_output_length", + ], + ) + }; self.emit_generic_tool_call_item( report_context, out, @@ -882,7 +946,11 @@ impl OpenAIResponsesProviderState { if item.get("type").and_then(Value::as_str) != Some("apply_patch_call") { return; } - let arguments = tool_arguments_from_named_fields(item, &["operation", "status"]); + let arguments = if self.terminal_only { + String::new() + } else { + tool_arguments_from_named_fields(item, &["operation", "status"]) + }; self.emit_generic_tool_call_item( report_context, out, @@ -903,10 +971,14 @@ impl OpenAIResponsesProviderState { if item.get("type").and_then(Value::as_str) != Some("computer_call") { return; } - let arguments = tool_arguments_from_named_fields( - item, - &["action", "actions", "pending_safety_checks", "status"], - ); + let arguments = if self.terminal_only { + String::new() + } else { + tool_arguments_from_named_fields( + item, + &["action", "actions", "pending_safety_checks", "status"], + ) + }; self.emit_generic_tool_call_item( report_context, out, @@ -941,10 +1013,20 @@ impl OpenAIResponsesProviderState { .unwrap_or(state.call_id.as_str()) .to_string(); state.name = name; - Self::merge_tool_call_arguments(state, &arguments); + if !self.terminal_only { + Self::merge_tool_call_arguments(state, &arguments); + } self.emit_ready_tool_call(report_context, out, index); } + fn tool_result_content(&self, value: Option<&Value>) -> String { + if self.terminal_only { + String::new() + } else { + openai_tool_result_content_from_value(value) + } + } + fn emit_missing_tool_result( &mut self, report_context: &Value, @@ -955,6 +1037,9 @@ impl OpenAIResponsesProviderState { content: &str, ) { self.ensure_started(report_context, out); + if self.terminal_only { + return; + } let state = self.tool_results.entry(index).or_default(); let missing = if !state.emitted { content.to_string() @@ -1005,7 +1090,7 @@ impl OpenAIResponsesProviderState { Some(format!("function_call_output:{tool_use_id}")), output_index, ); - let content = openai_tool_result_content_from_value( + let content = self.tool_result_content( item.get("output") .or_else(|| item.get("content")) .or_else(|| item.get("delta")), @@ -1047,7 +1132,7 @@ impl OpenAIResponsesProviderState { .to_string(); let index = self.tool_index_for_key(Some(format!("{item_type}:{tool_use_id}")), output_index); - let content = openai_tool_result_content_from_value( + let content = self.tool_result_content( item.get("output") .or_else(|| item.get("content")) .or_else(|| item.get("delta")), @@ -1110,6 +1195,24 @@ impl OpenAIResponsesProviderState { if item.get("type").and_then(Value::as_str) != Some("reasoning") { return; } + if self.terminal_only { + if item + .get("summary") + .and_then(Value::as_array) + .is_some_and(|summary| { + summary.iter().any(|part| { + part.get("type").and_then(Value::as_str) == Some("summary_text") + && part + .get("text") + .and_then(Value::as_str) + .is_some_and(|text| !text.is_empty()) + }) + }) + { + self.ensure_started(report_context, out); + } + return; + } let mut completed_reasoning = String::new(); for raw_summary in item .get("summary") @@ -1159,6 +1262,10 @@ impl OpenAIResponsesProviderState { if !has_image_payload { return; } + if self.terminal_only { + self.ensure_started(report_context, out); + return; + } let index = output_index.unwrap_or(self.image_item_keys.len()); let key = item .get("id") @@ -1180,6 +1287,42 @@ impl OpenAIResponsesProviderState { }); } + fn retained_output_item_key(&self, item: &Map) -> OpenAIResponsesOutputItemKey { + let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default(); + if self.terminal_only { + struct DigestWriter(Sha256); + + impl std::io::Write for DigestWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0.update(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + // Hash exactly the normal key's bytes, without retaining or first + // serializing an entire encrypted/opaque output item into a string. + let mut writer = DigestWriter(Sha256::new()); + writer.0.update(item_type.as_bytes()); + if let Some(item_id) = item.get("id").and_then(Value::as_str) { + writer.0.update(b":id:"); + writer.0.update(item_id.as_bytes()); + } else if let Some(content) = item.get("encrypted_content").and_then(Value::as_str) { + writer.0.update(b":encrypted_content:"); + writer.0.update(content.as_bytes()); + } else { + writer.0.update(b":"); + serde_json::to_writer(&mut writer, item) + .expect("JSON value serialization into a digest cannot fail"); + } + return OpenAIResponsesOutputItemKey::Digest(writer.0.finalize().into()); + } + OpenAIResponsesOutputItemKey::Full(Self::output_item_key(item)) + } + fn output_item_key(item: &Map) -> String { let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default(); if let Some(item_id) = item.get("id").and_then(Value::as_str) { @@ -1290,10 +1433,13 @@ impl OpenAIResponsesProviderState { } if final_item { - self.opaque_completed_item_keys - .insert(Self::output_item_key(item)); + let key = self.retained_output_item_key(item); + self.opaque_completed_item_keys.insert(key); } self.ensure_started(report_context, out); + if self.terminal_only { + return; + } let (id, model) = self.identity(report_context); out.push(CanonicalStreamFrame { id, @@ -1325,7 +1471,7 @@ impl OpenAIResponsesProviderState { if !self.emit_output_item(report_context, out, item, Some(output_index), true) && !self .opaque_completed_item_keys - .contains(&Self::output_item_key(item)) + .contains(&self.retained_output_item_key(item)) { out.push(self.unknown_frame(report_context, Value::Object(item.clone()))); } @@ -1504,6 +1650,9 @@ impl OpenAIResponsesProviderState { .map(|value| value as usize) .unwrap_or(0); self.ensure_started(report_context, &mut out); + if self.terminal_only { + return Ok(out); + } self.reasoning.push_str(piece); self.reasoning_parts .entry(summary_index) @@ -1543,6 +1692,9 @@ impl OpenAIResponsesProviderState { ); } self.ensure_started(report_context, &mut out); + if self.terminal_only { + return Ok(out); + } let (id, model) = self.identity(report_context); out.push(CanonicalStreamFrame { id, @@ -1595,7 +1747,9 @@ impl OpenAIResponsesProviderState { .unwrap_or("custom_tool") .to_string(); } - state.arguments.push_str(delta); + if !self.terminal_only { + state.arguments.push_str(delta); + } self.emit_ready_tool_call(report_context, &mut out, index); } "response.custom_tool_call_input.done" => { @@ -1623,11 +1777,13 @@ impl OpenAIResponsesProviderState { .unwrap_or("custom_tool") .to_string(); } - let arguments = tool_arguments_from_maybe_json_string( - Some(&Value::String(input.to_string())), - "input", - ); - Self::merge_tool_call_arguments(state, &arguments); + if !self.terminal_only { + let arguments = tool_arguments_from_maybe_json_string( + Some(&Value::String(input.to_string())), + "input", + ); + Self::merge_tool_call_arguments(state, &arguments); + } self.emit_ready_tool_call(report_context, &mut out, index); } "response.function_call_arguments.delta" => { @@ -1654,7 +1810,9 @@ impl OpenAIResponsesProviderState { if let Some(call_id) = value.get("call_id").and_then(Value::as_str) { state.call_id = call_id.to_string(); } - state.arguments.push_str(delta); + if !self.terminal_only { + state.arguments.push_str(delta); + } self.emit_ready_function_call(report_context, &mut out, index); } "response.function_call_arguments.done" => { @@ -1722,7 +1880,9 @@ impl OpenAIResponsesProviderState { if let Some(name) = incoming_chat_name { state.name = name; } - Self::merge_tool_call_arguments(state, arguments); + if !self.terminal_only { + Self::merge_tool_call_arguments(state, arguments); + } self.emit_ready_function_call(report_context, &mut out, index); } "response.function_call_output.delta" | "response.function_call_output.done" => { @@ -1743,7 +1903,7 @@ impl OpenAIResponsesProviderState { Some(format!("function_call_output:{tool_use_id}")), output_index, ); - let content = openai_tool_result_content_from_value( + let content = self.tool_result_content( value .get("delta") .or_else(|| value.get("output")) @@ -1781,7 +1941,7 @@ impl OpenAIResponsesProviderState { .map(|value| value as usize); let index = self .tool_index_for_key(Some(format!("{item_type}:{tool_use_id}")), output_index); - let content = openai_tool_result_content_from_value( + let content = self.tool_result_content( value .get("delta") .or_else(|| value.get("output")) @@ -3754,6 +3914,276 @@ mod tests { format!("data: {}\n", value).into_bytes() } + fn terminal_frames(frames: Vec) -> Vec { + frames + .into_iter() + .filter(|frame| { + matches!( + frame.event, + CanonicalStreamEvent::Start + | CanonicalStreamEvent::UnknownEvent(_) + | CanonicalStreamEvent::Finish { .. } + ) + }) + .collect() + } + + fn assert_no_responses_observation_content(state: &OpenAIResponsesProviderState) { + assert!(state.text_parts.is_empty()); + assert_eq!(state.reasoning.capacity(), 0); + assert!(state.reasoning_parts.is_empty()); + assert!(state + .tool_calls + .values() + .all(|tool| tool.arguments.capacity() == 0)); + assert!(state.tool_results.is_empty()); + assert!(state.image_item_keys.is_empty()); + assert!(state + .opaque_completed_item_keys + .iter() + .all(|key| matches!(key, OpenAIResponsesOutputItemKey::Digest(_)))); + } + + #[test] + fn openai_responses_terminal_observation_does_not_retain_long_stream_content() { + let context = json!({"provider_api_format": "openai:responses"}); + let mut observed = OpenAIResponsesProviderState::terminal_observation(); + let mut full = OpenAIResponsesProviderState::default(); + let initial = json!({ + "type": "response.output_item.added", "output_index": 0, + "item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "read", "arguments": ""} + }); + assert_eq!( + observed.push_event(&context, &initial).unwrap(), + terminal_frames(full.push_event(&context, &initial).unwrap()) + ); + let text = "x".repeat(1024); + let events = [ + json!({"type": "response.output_text.delta", "delta": text, "output_index": 3}), + json!({"type": "response.reasoning_summary_text.delta", "delta": text, "summary_index": 0}), + json!({"type": "response.function_call_arguments.delta", "delta": text, "output_index": 0}), + json!({"type": "response.custom_tool_call_input.delta", "delta": text, "output_index": 1}), + json!({"type": "response.function_call_output.delta", "delta": text, "output_index": 2, "call_id": "call_1"}), + ]; + for iteration in 0..2048 { + for event in &events { + let frames = observed.push_event(&context, event).unwrap(); + assert!(frames.is_empty()); + if iteration < 32 { + assert_eq!( + frames, + terminal_frames(full.push_event(&context, event).unwrap()) + ); + } + } + assert_no_responses_observation_content(&observed); + } + assert_eq!(observed.tool_calls.len(), 2); + assert_eq!(observed.tool_index_by_key, full.tool_index_by_key); + assert!(full.text_parts.values().any(|text| text.len() == 32 * 1024)); + assert_eq!(full.reasoning.len(), 32 * 1024); + assert_eq!(full.reasoning_parts[&0].len(), 32 * 1024); + assert_eq!(full.tool_calls[&0].arguments.len(), 32 * 1024); + assert!(!full.tool_results.is_empty()); + assert_eq!( + observed.finish(&context).unwrap(), + full.finish(&context).unwrap() + ); + } + + #[test] + fn openai_responses_terminal_observation_skips_completed_content_and_images() { + let context = json!({}); + let text = "x".repeat(64 * 1024); + let items = [ + json!({"type": "message", "content": [{"type": "output_text", "text": text}]}), + json!({"type": "reasoning", "summary": [{"type": "summary_text", "text": text}]}), + json!({"type": "custom_tool_call", "input": text, "name": "custom"}), + json!({"type": "shell_call", "action": {"command": text}}), + json!({"type": "apply_patch_call", "operation": {"patch": text}}), + json!({"type": "computer_call", "action": {"keys": [text]}}), + json!({"type": "function_call_output", "output": text}), + json!({"type": "custom_tool_call_output", "output": {"text": text}}), + json!({"type": "image_generation_call", "result": text, "status": "completed"}), + ]; + let mut observed = OpenAIResponsesProviderState::terminal_observation(); + let mut full = OpenAIResponsesProviderState::default(); + for (index, item) in items.iter().enumerate() { + let event = + json!({"type": "response.output_item.done", "output_index": index, "item": item}); + assert_eq!( + observed.push_event(&context, &event).unwrap(), + terminal_frames(full.push_event(&context, &event).unwrap()) + ); + assert_no_responses_observation_content(&observed); + } + let completed = json!({"type": "response.completed", "response": { + "id": "resp_complete", "model": "model", "service_tier": "Priority", + "output": items, "usage": {"input_tokens": 100, "output_tokens": 20, "input_tokens_details": {"cached_tokens": 0}} + }}); + assert_eq!( + observed.push_event(&context, &completed).unwrap(), + terminal_frames(full.push_event(&context, &completed).unwrap()) + ); + assert_eq!(observed.actual_service_tier(), Some("priority")); + assert_no_responses_observation_content(&observed); + } + + #[test] + fn openai_responses_terminal_observation_hashes_opaque_keys_without_changing_deduplication() { + let context = json!({}); + let text = "x".repeat(64 * 1024); + let items = [ + json!({"type": "future_item", "id": "id:with:separators", "encrypted_content": text}), + json!({"type": "compaction", "encrypted_content": text}), + json!({"type": "future_item", "payload": {"text": text, "escaped": "\n\"\\"}}), + ]; + let mut observed = OpenAIResponsesProviderState::terminal_observation(); + let mut full = OpenAIResponsesProviderState::default(); + for item in &items { + let object = item.as_object().unwrap(); + let normal_key = OpenAIResponsesProviderState::output_item_key(object); + let expected: [u8; 32] = Sha256::digest(normal_key.as_bytes()).into(); + let OpenAIResponsesOutputItemKey::Digest(actual) = + observed.retained_output_item_key(object) + else { + panic!("observation must retain only an opaque item digest"); + }; + assert_eq!(actual, expected); + let event = json!({"type": "response.output_item.done", "item": item}); + for _ in 0..2 { + assert_eq!( + observed.push_event(&context, &event).unwrap(), + terminal_frames(full.push_event(&context, &event).unwrap()) + ); + } + } + assert_eq!(observed.opaque_completed_item_keys.len(), items.len()); + assert_no_responses_observation_content(&observed); + let mut final_items = items.to_vec(); + final_items.push(json!({"type": "future_item", "id": "new_item"})); + let final_event = + json!({"type": "response.completed", "response": {"output": final_items}}); + let frames = observed.push_event(&context, &final_event).unwrap(); + assert_eq!( + frames, + terminal_frames(full.push_event(&context, &final_event).unwrap()) + ); + assert_eq!( + frames + .iter() + .filter(|frame| matches!(frame.event, CanonicalStreamEvent::UnknownEvent(_))) + .count(), + 1 + ); + } + + #[test] + fn openai_responses_terminal_observation_preserves_tool_identity_validation() { + let context = json!({"original_request_body": {"tools": [{ + "type": "namespace", "name": "reports", "description": "Reporting tools", "tools": [{ + "type": "function", "name": "write_report", "parameters": {"type": "object"} + }] + }]}}); + let expected_alias = NamespaceToolAliases::from_report_context(&context) + .chat_name("reports", "write_report") + .expect("fixture must contain a valid namespace tool") + .to_owned(); + let mut observed = OpenAIResponsesProviderState::terminal_observation(); + let mut full = OpenAIResponsesProviderState::default(); + let events = [ + json!({"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 0, "delta": "{"}), + json!({"type": "response.output_item.added", "output_index": 0, "item": { + "type": "function_call", "id": "fc_1", "namespace": "reports", "name": "write_report", "arguments": "" + }}), + json!({"type": "response.function_call_arguments.done", "item_id": "fc_1", "call_id": "call_1", "namespace": "reports", "arguments": "{}"}), + json!({"type": "response.function_call_output.delta", "call_id": "call_1", "output_index": 7, "delta": "result"}), + json!({"type": "response.function_call_arguments.done", "item_id": "fc_other", "name": "ordinary", "arguments": "{}"}), + json!({"type": "response.function_call_arguments.done", "item_id": "fc_1", "namespace": "missing", "arguments": "{}"}), + json!({"type": "response.output_item.done", "item": { + "type": "function_call", "id": "invalid", "name": "read", "caller": {"type": "future"}, "arguments": "{}" + }}), + json!({"type": "response.output_item.done", "item": {"content": "missing type"}}), + json!({"type": "response.future.delta", "payload": "unsupported"}), + ]; + let mut unknowns = 0; + for (step, event) in events.into_iter().enumerate() { + let frames = observed.push_event(&context, &event).unwrap(); + assert_eq!( + frames, + terminal_frames(full.push_event(&context, &event).unwrap()) + ); + unknowns += frames + .iter() + .filter(|frame| matches!(frame.event, CanonicalStreamEvent::UnknownEvent(_))) + .count(); + assert_eq!(observed.tool_index_by_key, full.tool_index_by_key); + assert_eq!(observed.last_tool_index, full.last_tool_index); + assert_eq!(observed.tool_calls.len(), full.tool_calls.len()); + for (index, tool) in &observed.tool_calls { + assert_eq!(tool.name, full.tool_calls[index].name); + assert_eq!(tool.call_id, full.tool_calls[index].call_id); + } + if step == 1 || step == 2 { + assert_eq!(unknowns, 0); + assert_eq!(observed.tool_calls[&0].name, expected_alias); + assert_eq!( + observed.tool_calls[&0].call_id, + if step == 1 { "" } else { "call_1" } + ); + } + assert_no_responses_observation_content(&observed); + } + assert_eq!(unknowns, 4); + assert_eq!( + observed.finish(&context).unwrap(), + full.finish(&context).unwrap() + ); + } + + #[test] + fn openai_chat_terminal_observation_does_not_buffer_arguments_before_identity() { + let context = json!({}); + let mut observed = OpenAIChatProviderState::terminal_observation(); + let mut full = OpenAIChatProviderState::default(); + let delta = json!({"choices": [{"delta": { + "content": "x".repeat(1024), "reasoning_content": "r".repeat(1024), + "tool_calls": [{"index": 0, "function": {"arguments": "a".repeat(1024)}}] + }}]}); + for iteration in 0..2048 { + let frames = observed + .push_line(&context, data_line(delta.clone())) + .unwrap(); + if iteration < 32 { + assert_eq!( + frames, + terminal_frames(full.push_line(&context, data_line(delta.clone())).unwrap()) + ); + } + assert!(observed.tool_calls.is_empty()); + } + assert_eq!(full.tool_calls[&0].pending_arguments.len(), 32 * 1024); + for event in [ + json!({"choices": [{"delta": {"tool_calls": [{"index": 0, "id": "call_1", "function": {"name": "read", "arguments": "end"}}]}}]}), + json!({"choices": [{"delta": {"future": "unknown"}}]}), + json!({"choices": [{"delta": {}, "finish_reason": "tool_calls"}]}), + json!({"choices": [], "usage": {"prompt_tokens": 100, "completion_tokens": 20, "prompt_tokens_details": {"cached_tokens": 0}}, "service_tier": "Flex"}), + ] { + assert_eq!( + observed + .push_line(&context, data_line(event.clone())) + .unwrap(), + terminal_frames(full.push_line(&context, data_line(event)).unwrap()) + ); + } + assert_eq!(observed.actual_service_tier(), Some("flex")); + assert!(observed.tool_calls.is_empty()); + assert_eq!( + observed.finish(&context).unwrap(), + full.finish(&context).unwrap() + ); + } + fn response_sequence_numbers(sse: &str) -> Vec { let mut sequence_numbers = Vec::new(); for payload in sse.lines().filter_map(|line| line.strip_prefix("data: ")) { diff --git a/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs b/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs index a702631ea..9351e6a25 100644 --- a/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs +++ b/crates/aether-ai/formats/src/formats/shared/stream_core/format_matrix.rs @@ -423,10 +423,26 @@ impl TerminalStreamParser { { return Some(Self::OpenAIImage(OpenAiImageStreamTerminalState::default())); } - ProviderStreamParser::for_api_format(provider_api_format).map(Self::Standard) + let provider = match ProviderStreamParser::for_api_format(provider_api_format)? { + ProviderStreamParser::OpenAIChat(_) => { + ProviderStreamParser::OpenAIChat(OpenAIChatProviderState::terminal_observation()) + } + ProviderStreamParser::OpenAIResponses(_) => ProviderStreamParser::OpenAIResponses( + OpenAIResponsesProviderState::terminal_observation(), + ), + ProviderStreamParser::Gemini(_) => { + ProviderStreamParser::Gemini(GeminiProviderState::terminal_observation()) + } + provider @ ProviderStreamParser::Claude(_) => provider, + }; + Some(Self::Standard(provider)) } } +#[cfg(test)] +#[path = "terminal_observation_tests.rs"] +mod terminal_observation_tests; + enum ProviderStreamParser { OpenAIChat(OpenAIChatProviderState), OpenAIResponses(OpenAIResponsesProviderState), diff --git a/crates/aether-ai/formats/src/formats/shared/stream_core/terminal_observation_tests.rs b/crates/aether-ai/formats/src/formats/shared/stream_core/terminal_observation_tests.rs new file mode 100644 index 000000000..ddf076bd2 --- /dev/null +++ b/crates/aether-ai/formats/src/formats/shared/stream_core/terminal_observation_tests.rs @@ -0,0 +1,282 @@ +use serde_json::{json, Value}; + +use super::{ProviderStreamParser, StreamingStandardTerminalObserver, TerminalStreamParser}; + +fn full_observer(context: &Value) -> StreamingStandardTerminalObserver { + StreamingStandardTerminalObserver { + provider: Some(TerminalStreamParser::Standard( + ProviderStreamParser::for_api_format(context["provider_api_format"].as_str().unwrap()) + .unwrap(), + )), + ..Default::default() + } +} + +// Compare every input prefix, including EOF, to catch changes in identity and terminal timing. +fn assert_summaries_match(context: &Value, events: &[Value]) { + let structured = context["provider_api_format"] + .as_str() + .unwrap() + .starts_with("openai:responses"); + for end in 0..=events.len() { + let mut compact = StreamingStandardTerminalObserver::default(); + let mut full = full_observer(context); + let mut via_event = StreamingStandardTerminalObserver::default(); + for (index, event) in events[..end].iter().enumerate() { + let line = format!("data: {event}\n").into_bytes(); + compact.push_line(context, line.clone()).unwrap(); + full.push_line(context, line).unwrap(); + assert_eq!( + compact.latest_summary(), + full.latest_summary(), + "prefix {index}: {event}" + ); + if structured { + via_event.push_event(context, event).unwrap(); + assert_eq!(via_event.latest_summary(), full.latest_summary()); + } + } + let expected = full.finish(context).unwrap(); + assert_eq!(compact.finish(context).unwrap(), expected, "EOF at {end}"); + if structured { + assert_eq!(via_event.finish(context).unwrap(), expected); + } + } +} + +fn context(format: &str) -> Value { + json!({"provider_api_format": format, "client_api_format": format, "mapped_model": "test-model"}) +} + +fn completed(output: Vec) -> Value { + json!({"type":"response.completed","response":{ + "id":"resp-final","model":"final-model","status":"completed","service_tier":" PRIORITY ", + "output":output,"usage":{"input_tokens":11,"output_tokens":7,"total_tokens":18, + "input_tokens_details":{"cached_tokens":0},"output_tokens_details":{"reasoning_tokens":3}} + }}) +} + +#[test] +fn responses_content_snapshots_and_deltas_preserve_every_summary() { + let events = vec![ + json!({"type":"response.output_text.delta","delta":""}), + json!({"type":"response.output_text.delta","delta":"he","output_index":0}), + json!({"type":"response.created","response":{"id":"late-id","model":"late-model"}}), + json!({"type":"response.output_text.delta","delta":{"text":"hello"},"output_index":0}), + json!({"type":"response.output_text.done","text":"hello","output_index":0}), + json!({"type":"response.content_part.added","content_index":1,"part":{"type":"output_text","text":"another"}}), + json!({"type":"response.content_part.done","content_index":1,"part":{"type":"output_text","text":"another part"}}), + json!({"type":"response.refusal.delta","delta":"refuse"}), + json!({"type":"response.refusal.done","refusal":"refused"}), + json!({"type":"response.audio.transcript.delta","delta":"audio"}), + json!({"type":"response.audio.transcript.done","transcript":"audio transcript"}), + json!({"type":"response.reasoning_summary_text.delta","delta":"think","summary_index":0}), + json!({"type":"response.reasoning_text.delta","delta":" again","summary_index":0}), + json!({"type":"response.reasoning_summary_part.added","summary_index":1,"part":{"type":"summary_text","text":"second"}}), + json!({"type":"response.reasoning_summary_part.done","summary_index":1,"part":{"type":"summary_text","text":"second thought"}}), + json!({"type":"response.reasoning_summary_text.done","summary_index":0,"text":"think again"}), + json!({"type":"response.reasoning_text.done","text":""}), + completed(vec![ + json!({"type":"message","content":[{"type":"output_text","text":"hello"},{"type":"refusal","refusal":"refused"}]}), + json!({"type":"reasoning","summary":[{"type":"summary_text","text":"think again"},{"type":"summary_text","text":"second thought"}]}), + ]), + ]; + for format in ["openai:responses", "openai:responses:compact"] { + assert_summaries_match(&context(format), &events); + for event in &events { + assert_summaries_match(&context(format), std::slice::from_ref(event)); + } + } +} + +#[test] +fn responses_tools_and_unknown_execution_fields_preserve_summaries() { + let calls = vec![ + json!({"type":"function_call","call_id":"call-0","name":"lookup","arguments":"{\"x\":1}"}), + json!({"type":"custom_tool_call","call_id":"call-1","name":"custom","input":"raw input"}), + json!({"type":"shell_call","call_id":"call-2","action":{"commands":["pwd"]}}), + json!({"type":"local_shell_call","call_id":"call-3","action":{"command":["pwd"]}}), + json!({"type":"apply_patch_call","call_id":"call-4","operation":{"type":"update_file","path":"a","diff":"+b"}}), + json!({"type":"computer_call","call_id":"call-5","action":{"type":"click","x":1,"y":2}}), + json!({"type":"function_call","call_id":"bad-caller","name":"lookup","arguments":"{}","caller":{"type":"direct"}}), + json!({"type":"function_call","call_id":"bad-namespace","namespace":42,"name":"lookup","arguments":"{}"}), + ]; + let mut events = vec![ + json!({"type":"response.function_call_arguments.delta","output_index":0,"delta":"{"}), + json!({"type":"response.function_call_arguments.delta","output_index":0,"call_id":"call-0","delta":"\"x\":1}"}), + json!({"type":"response.function_call_arguments.done","output_index":0,"item":{"call_id":"call-0","name":"lookup","arguments":"{\"x\":1}"}}), + json!({"type":"response.function_call_arguments.done","output_index":0,"namespace":{},"arguments":"{}"}), + json!({"type":"response.custom_tool_call_input.delta","output_index":1,"delta":"raw"}), + json!({"type":"response.custom_tool_call_input.done","output_index":1,"input":"raw input"}), + ]; + for (index, item) in calls.iter().enumerate() { + events.push(json!({"type":"response.output_item.added","output_index":index,"item":item})); + events.push(json!({"type":"response.output_item.done","output_index":index,"item":item})); + } + for kind in [ + "function_call", + "custom_tool_call", + "shell_call", + "local_shell_call", + "apply_patch_call", + "computer_call", + ] { + events.push(json!({"type":format!("response.{kind}_output.delta"),"call_id":"result","delta":"result"})); + events.push(json!({"type":format!("response.{kind}_output.done"),"call_id":"result","output":{"ok":true}})); + events.push(json!({"type":"response.output_item.done","item":{"type":format!("{kind}_output"),"call_id":"result","output":"result complete"}})); + } + events.push(completed(calls)); + assert_summaries_match(&context("openai:responses"), &events); +} + +#[test] +fn responses_opaque_dedup_and_images_preserve_unknown_counts() { + let items = vec![ + json!({"type":"future_item","id":"stable-id","payload":"first"}), + json!({"type":"future_item","encrypted_content":"encrypted","payload":"second"}), + json!({"type":"future_item","payload":{"no_id":true}}), + json!({"type":"image_generation_call","result":"image data","status":"completed"}), + json!({"type":"reasoning","summary":[],"encrypted_content":"reasoning"}), + ]; + let mut events = Vec::new(); + for item in &items { + events.push(json!({"type":"response.output_item.added","item":item})); + events.push(json!({"type":"response.output_item.done","item":item})); + events.push(json!({"type":"response.output_item.done","item":item})); + } + events.push(json!({"type":"response.output_item.done","item":{"missing_type":true}})); + let mut output = items; + output[0]["payload"] = json!("changed body, same id"); + output[1]["payload"] = json!("changed body, same encrypted content"); + output.push(json!({"type":"new_future_item","payload":"never emitted"})); + events.push(completed(output)); + let ctx = context("openai:responses"); + assert_summaries_match(&ctx, &events); + let mut observer = StreamingStandardTerminalObserver::default(); + for event in &events { + observer.push_event(&ctx, event).unwrap(); + } + assert_eq!( + observer.finish(&ctx).unwrap().unwrap().unknown_event_count, + 2 + ); +} + +#[test] +fn responses_namespace_validation_keeps_existing_tool_identity() { + let mut ctx = context("openai:responses"); + ctx["original_request_body"] = json!({"tools":[{ + "type":"namespace","name":"search","description":"Search tools","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}] + }]}); + let events = [ + json!({"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"call","namespace":"search","name":"lookup","arguments":"{"}}), + json!({"type":"response.function_call_arguments.delta","output_index":0,"delta":"}"}), + json!({"type":"response.function_call_arguments.done","output_index":0,"namespace":"search","arguments":"{}"}), + json!({"type":"response.function_call_arguments.done","output_index":0,"namespace":"missing","arguments":"{}"}), + completed(vec![ + json!({"type":"function_call","call_id":"call","namespace":"search","name":"lookup","arguments":"{}"}), + ]), + ]; + assert_summaries_match(&ctx, &events); + let mut observer = StreamingStandardTerminalObserver::default(); + for event in &events { + observer.push_event(&ctx, event).unwrap(); + } + let summary = observer.finish(&ctx).unwrap().unwrap(); + assert_eq!(summary.unknown_event_count, 1); + assert_eq!(summary.finish_reason.as_deref(), Some("tool_calls")); +} + +#[test] +fn responses_errors_zero_usage_and_event_type_lines_preserve_summaries() { + for terminal in [ + json!({"type":"response.failed","response":{"status":"failed","error":{"message":"failed"},"usage":{"input_tokens":0,"output_tokens":0}}}), + json!({"type":"error","error":{"type":"server_error","message":"failed"}}), + json!({"type":"response.incomplete","response":{"status":"incomplete","incomplete_details":{"reason":"max_output_tokens"},"usage":{"input_tokens":2,"output_tokens":0}}}), + json!({"type":"response.done","response":{"status":"completed","usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}}), + json!({"type":"response.completed","response":null}), + ] { + let ctx = context("openai:responses"); + let events = vec![ + json!({"type":"response.future"}), + json!({"type":"ping"}), + terminal, + ]; + assert_summaries_match(&ctx, &events); + let mut compact = StreamingStandardTerminalObserver::default(); + let mut full = full_observer(&ctx); + for mut event in events { + let kind = event.as_object_mut().unwrap().remove("type").unwrap(); + for line in [ + format!("event: {}\n", kind.as_str().unwrap()), + format!("data: {event}\n"), + "\n".to_string(), + ] { + compact.push_line(&ctx, line.as_bytes().to_vec()).unwrap(); + full.push_line(&ctx, line.into_bytes()).unwrap(); + assert_eq!(compact.latest_summary(), full.latest_summary()); + } + } + assert_eq!(compact.finish(&ctx).unwrap(), full.finish(&ctx).unwrap()); + } +} + +#[test] +fn chat_delayed_tool_identity_and_usage_only_preserve_summaries() { + let ctx = context("openai:chat"); + assert_summaries_match( + &ctx, + &[ + json!({"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{"}}]}}]}), + json!({"id":"late-id","model":"late-model","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call","function":{"name":"lookup","arguments":"}"}}]}}]}), + json!({"choices":[{"delta":{"content":"text","reasoning_content":"reason"}}],"service_tier":"priority"}), + json!({"choices":[{"delta":{},"finish_reason":"tool_calls"}]}), + json!({"choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}), + ], + ); + for event in [ + json!({"usage":{"prompt_tokens":0,"completion_tokens":0,"total_tokens":0}}), + json!({"choices":[{"delta":{"tool_calls":"malformed"}}]}), + json!({"choices":[{"delta":{"tool_calls":[null,{}, {"function":null}]}}]}), + json!({"choices":[{"delta":{},"finish_reason":"future_reason"}]}), + json!({"choices":[{"delta":{"future_content":true}}]}), + ] { + assert_summaries_match(&ctx, &[event]); + } +} + +#[test] +fn gemini_content_tools_and_errors_preserve_summaries() { + let ctx = context("gemini:generate_content"); + let parts = vec![ + json!({"text":"text"}), + json!({"text":"reason","thought":true,"thoughtSignature":"sig"}), + json!({"functionCall":{"id":"call","name":"lookup","args":{"x":1}}}), + json!({"functionResponse":{"id":"call","name":"lookup","response":{"ok":true}}}), + json!({"inlineData":{"mimeType":"image/png","data":"aW1hZ2U="}}), + json!({"futureContent":"unknown"}), + ]; + let mut events = Vec::new(); + for part in &parts { + let event = json!({"candidates":[{"content":{"parts":[part]}}]}); + events.push(event.clone()); + events.push(event.clone()); + assert_summaries_match(&ctx, &[event]); + } + events.push(json!({"responseId":"late-id","modelVersion":"late-model","candidates":[{"content":{"parts":parts},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":5,"candidatesTokenCount":3,"totalTokenCount":8}})); + assert_summaries_match(&ctx, &events); + for reason in [ + "MALFORMED_FUNCTION_CALL", + "SAFETY", + "MAX_TOKENS", + "FUTURE_REASON", + ] { + assert_summaries_match( + &ctx, + &[ + json!({"response":{"candidates":[{"content":{"parts":[{"text":"partial"}]}}]}}), + json!({"candidates":[{"content":{"parts":[]},"finishReason":reason}],"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0}}), + ], + ); + } +} diff --git a/crates/aether-billing/Cargo.toml b/crates/aether-billing/Cargo.toml index adf92f681..edcfeb781 100644 --- a/crates/aether-billing/Cargo.toml +++ b/crates/aether-billing/Cargo.toml @@ -14,3 +14,6 @@ serde.workspace = true serde_json.workspace = true thiserror.workspace = true tokio.workspace = true + +[dev-dependencies] +aether-runtime-state.workspace = true diff --git a/crates/aether-billing/src/event_enrichment.rs b/crates/aether-billing/src/event_enrichment.rs index d2267a0f9..d61eb526f 100644 --- a/crates/aether-billing/src/event_enrichment.rs +++ b/crates/aether-billing/src/event_enrichment.rs @@ -500,9 +500,14 @@ fn build_settlement_snapshot( #[cfg(test)] mod tests { + use std::sync::Arc; + use aether_data_contracts::repository::billing::StoredBillingModelContext; use aether_data_contracts::repository::usage::UsageBodyCaptureState; - use aether_usage_runtime::{UsageEvent, UsageEventData, UsageEventType}; + use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; + use aether_usage_runtime::{ + UsageEvent, UsageEventData, UsageEventType, UsageQueue, UsageRuntimeConfig, + }; use async_trait::async_trait; use serde_json::json; use serde_json::Value; @@ -539,6 +544,539 @@ mod tests { } } + fn wire_billing_lookup(pricing: Option, request_price: Option) -> TestLookup { + TestLookup { + name_context: Some( + StoredBillingModelContext::new( + "provider-wire".to_string(), + Some("pay_as_you_go".to_string()), + Some("key-wire".to_string()), + None, + Some(5), + "global-model-wire".to_string(), + "wire-model".to_string(), + None, + request_price, + pricing, + Some("model-wire".to_string()), + Some("wire-model".to_string()), + None, + None, + None, + ) + .expect("wire billing context"), + ), + model_id_context: None, + } + } + + fn wire_billing_event(request_id: &str) -> UsageEvent { + UsageEvent::new( + UsageEventType::Completed, + request_id, + UsageEventData { + user_id: Some("user-wire".to_string()), + api_key_id: Some("key-wire".to_string()), + provider_name: "OpenAI".to_string(), + provider_id: Some("provider-wire".to_string()), + provider_api_key_id: Some("key-wire".to_string()), + model: "gpt-5.6-sol".to_string(), + target_model: Some("gpt-5.6-sol".to_string()), + request_type: Some("chat".to_string()), + api_format: Some("openai:responses".to_string()), + endpoint_api_format: Some("openai:responses".to_string()), + input_tokens: Some(1_000), + output_tokens: Some(100), + total_tokens: Some(1_100), + cache_creation_input_tokens: Some(0), + cache_creation_ephemeral_5m_input_tokens: Some(0), + cache_creation_ephemeral_1h_input_tokens: Some(0), + cache_read_input_tokens: Some(0), + status_code: Some(200), + first_byte_time_ms: Some(12), + response_time_ms: Some(30), + request_headers: Some(json!({"x-audit": "request"})), + provider_request_headers: Some(json!({"x-audit": "provider request"})), + response_headers: Some(json!({"x-audit": "provider response"})), + client_response_headers: Some(json!({"x-audit": "client response"})), + provider_request_body: Some(json!({"model": "gpt-5.6-sol"})), + provider_request_body_state: Some(UsageBodyCaptureState::Inline), + response_body: Some(json!({"service_tier": "default", "output": "x".repeat(8192)})), + response_body_state: Some(UsageBodyCaptureState::Inline), + request_metadata: Some(json!({ + "usage_available": true, + "usage_pricing_available": true, + "api_key_is_standalone": true, + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000", + "plan_usage_reservation_deferred": true + })), + ..UsageEventData::default() + }, + ) + } + + fn wire_billing_result(event: &UsageEvent) -> Value { + let mut data = serde_json::to_value(&event.data).expect("serialized billing event"); + let object = data.as_object_mut().expect("event data object"); + for key in [ + "request_body", + "provider_request_body", + "response_body", + "client_response_body", + "request_body_state", + "provider_request_body_state", + "response_body_state", + "client_response_body_state", + "request_headers", + "provider_request_headers", + "response_headers", + "client_response_headers", + ] { + object.remove(key); + } + if let Some(metadata) = object + .get_mut("request_metadata") + .and_then(Value::as_object_mut) + { + metadata.retain(|key, _| { + matches!( + key.as_str(), + "usage_available" + | "usage_pricing_available" + | "api_key_is_standalone" + | "plan_usage_reservation_token" + | "plan_usage_reservation_deferred" + | "cancelled_request_fee" + | "dimensions" + | "billing_dimensions" + | "billing_snapshot" + | "settlement_snapshot" + | "rate_multiplier" + | "is_free_tier" + | "settlement_snapshot_schema_version" + ) + }); + for key in ["billing_snapshot", "settlement_snapshot"] { + if let Some(snapshot) = metadata.get_mut(key).and_then(Value::as_object_mut) { + snapshot.remove("calculated_at"); + } + } + } + json!({ + "event_type": event.event_type, + "request_id": event.request_id, + "timestamp_ms": event.timestamp_ms, + "data": data, + }) + } + + async fn assert_wire_billing_equivalent( + lookup: &TestLookup, + original: UsageEvent, + ) -> UsageEvent { + const LIMIT: usize = 4096; + let original_fields = original.to_stream_fields().expect("legacy full envelope"); + assert!(original_fields["payload"].len() > LIMIT); + let provider_body_present = original + .data + .provider_request_body + .as_ref() + .is_some_and(|body| !body.is_null()); + let queue = UsageQueue::new( + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), + UsageRuntimeConfig { + enabled: true, + queue_payload_max_bytes: LIMIT, + consumer_block_ms: 1, + ..UsageRuntimeConfig::default() + }, + ) + .expect("bounded billing queue"); + queue.ensure_consumer_group().await.expect("billing group"); + queue + .enqueue(&original) + .await + .expect("diagnostic projection should fit"); + let entries = queue + .read_group("billing-wire-reader") + .await + .expect("billing queue read"); + assert_eq!(entries.len(), 1); + assert!(entries[0].fields["payload"].len() <= LIMIT); + // The legacy consumer also prepares request facts after decoding the full envelope. + let mut original = + UsageEvent::from_stream_fields(&original_fields).expect("legacy consumer event"); + let mut queued = + UsageEvent::from_stream_fields(&entries[0].fields).expect("projected event"); + assert!(original.data.response_body.is_some()); + assert_eq!( + original.data.provider_request_body.is_some(), + provider_body_present + ); + assert!(queued.data.response_body.is_none()); + assert!(queued.data.request_headers.is_none()); + assert!(queued.data.provider_request_headers.is_none()); + assert!(queued.data.response_headers.is_none()); + assert!(queued.data.client_response_headers.is_none()); + assert_eq!( + queued.data.response_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + + enrich_usage_event_with_billing(lookup, &mut original) + .await + .expect("original billing"); + enrich_usage_event_with_billing(lookup, &mut queued) + .await + .expect("projected billing"); + assert_eq!(wire_billing_result(&queued), wire_billing_result(&original)); + queued + } + + #[tokio::test] + async fn wire_projection_preserves_openai_requested_tier_and_effective_cache_ttl() { + let lookup = wire_billing_lookup( + Some(json!({ + "tiers": [{"up_to": null, "input_price_per_1m": 5.0, + "output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25, + "cache_read_price_per_1m": 0.5, + "cache_ttl_pricing": [{"ttl_minutes": 60, + "cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}], + "processing_tiers": {"priority": {"price_multiplier": 2.0}} + })), + None, + ); + let mut event = wire_billing_event("wire-openai-tier"); + event.data.cache_creation_input_tokens = Some(100); + event.data.provider_request_body = Some(json!({ + "model": "gpt-5.6-sol", "service_tier": "priority", "reasoning": {"effort": "high"} + })); + event.data.request_metadata.as_mut().unwrap()["provider_service_tier"] = json!("flex"); + event.data.request_metadata.as_mut().unwrap()["provider_actual_service_tier"] = + json!("flex"); + let queued = assert_wire_billing_equivalent(&lookup, event).await; + let metadata = queued.data.request_metadata.as_ref().unwrap(); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "flex"); + assert_eq!(metadata["provider_reasoning_effort"], "high"); + assert_eq!(metadata["billing_dimensions"]["cache_ttl_minutes"], 30); + assert_eq!( + metadata["billing_dimensions"]["billing_processing_tier"], + "priority" + ); + assert!(queued.data.total_cost_usd.unwrap() > 0.0); + } + + #[tokio::test] + async fn wire_projection_preserves_non_object_body_authority_and_null_decode_semantics() { + let lookup = wire_billing_lookup( + Some(json!({ + "tiers": [{"up_to": null, "input_price_per_1m": 5.0, + "output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25, + "cache_read_price_per_1m": 0.5, + "cache_ttl_pricing": [{"ttl_minutes": 60, + "cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}], + "processing_tiers": {"priority": {"price_multiplier": 2.0}} + })), + None, + ); + for (kind, body) in [ + ("string", json!("not an object")), + ("array", json!([{"service_tier": "flex"}])), + ("number", json!(42)), + ("boolean", json!(false)), + ("null", Value::Null), + ] { + for state in [ + Some(UsageBodyCaptureState::Inline), + Some(UsageBodyCaptureState::Reference), + None, + ] { + let mut event = wire_billing_event(&format!("wire-{kind}-{state:?}")); + event.data.input_tokens = Some(1_000_000); + event.data.output_tokens = Some(0); + event.data.total_tokens = Some(1_000_000); + event.data.cache_creation_input_tokens = Some(1_000_000); + event.data.provider_request_body = Some(body.clone()); + event.data.provider_request_body_state = state; + let metadata = event.data.request_metadata.as_mut().unwrap(); + metadata["provider_service_tier"] = json!("priority"); + metadata["provider_reasoning_effort"] = json!("high"); + metadata["provider_cache_ttl_minutes"] = json!(60); + + let queued = assert_wire_billing_equivalent(&lookup, event).await; + let metadata = queued.data.request_metadata.as_ref().unwrap(); + assert!(queued.data.provider_request_body.is_none()); + assert_eq!(metadata["provider_cache_ttl_minutes"], 60); + assert_eq!(metadata["billing_dimensions"]["cache_ttl_minutes"], 60); + let expected_tier = if body.is_null() && state.is_some() { + Some("priority") + } else { + None + }; + assert_eq!( + usage_event_processing_tiers(&queued.data) + .requested + .as_deref(), + expected_tier, + "{kind} with {state:?}" + ); + assert_eq!( + queued.data.total_cost_usd, + Some(if expected_tier.is_some() { + 200.0 + } else { + 100.0 + }), + "{kind} with {state:?}" + ); + if body.is_null() { + // Option decodes JSON null as absent, so the old capture marker remains. + assert_eq!(queued.data.provider_request_body_state, state); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_reasoning_effort"], "high"); + } else { + assert_eq!( + queued.data.provider_request_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + assert!(metadata.get("provider_service_tier").is_none()); + assert!(metadata.get("provider_reasoning_effort").is_none()); + } + } + } + } + + #[tokio::test] + async fn wire_projection_preserves_raw_body_ttl_with_non_authoritative_capture_states() { + let lookup = wire_billing_lookup( + Some(json!({ + "tiers": [{"up_to": null, "input_price_per_1m": 5.0, + "output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25, + "cache_read_price_per_1m": 0.5, + "cache_ttl_pricing": [{"ttl_minutes": 60, + "cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}] + })), + None, + ); + for state in [ + UsageBodyCaptureState::Disabled, + UsageBodyCaptureState::Unavailable, + UsageBodyCaptureState::Truncated, + ] { + let mut event = wire_billing_event(&format!("wire-capture-state-{state:?}")); + event.data.input_tokens = Some(1_000_000); + event.data.output_tokens = Some(0); + event.data.total_tokens = Some(1_000_000); + event.data.cache_creation_input_tokens = Some(1_000_000); + event.data.provider_request_body = Some(json!({ + "model": "gpt-5.6-sol", "prompt_cache_options": {"ttl": "30m"} + })); + event.data.provider_request_body_state = Some(state); + event.data.request_metadata.as_mut().unwrap()["provider_cache_ttl_minutes"] = json!(60); + + let queued = assert_wire_billing_equivalent(&lookup, event).await; + assert_eq!(queued.data.provider_request_body_state, Some(state)); + assert!(queued.data.provider_request_body.is_none()); + assert_eq!(queued.data.total_cost_usd, Some(6.25)); + assert_eq!( + queued.data.request_metadata.as_ref().unwrap()["billing_dimensions"] + ["cache_ttl_minutes"], + 30 + ); + } + } + + #[tokio::test] + async fn wire_projection_rejects_typed_none_when_omitting_raw_ttl_would_change_billing() { + let lookup = wire_billing_lookup( + Some(json!({ + "tiers": [{"up_to": null, "input_price_per_1m": 5.0, + "output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25, + "cache_read_price_per_1m": 0.5, + "cache_ttl_pricing": [{"ttl_minutes": 60, + "cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}] + })), + None, + ); + let mut event = wire_billing_event("wire-typed-none-ttl"); + event.data.input_tokens = Some(1_000_000); + event.data.output_tokens = Some(0); + event.data.total_tokens = Some(1_000_000); + event.data.cache_creation_input_tokens = Some(1_000_000); + event.data.provider_request_body = Some(json!({ + "model": "gpt-5.6-sol", "prompt_cache_options": {"ttl": "30m"} + })); + event.data.provider_request_body_state = Some(UsageBodyCaptureState::None); + event.data.request_metadata.as_mut().unwrap()["provider_cache_ttl_minutes"] = json!(60); + let original_fields = event.to_stream_fields().expect("legacy full envelope"); + assert!(original_fields["payload"].len() > 4096); + let mut legacy = UsageEvent::from_stream_fields(&original_fields).expect("legacy consumer"); + assert_eq!( + legacy.data.provider_request_body_state, + Some(UsageBodyCaptureState::None) + ); + assert!(legacy.data.provider_request_body.is_some()); + assert!(legacy + .data + .request_metadata + .as_ref() + .unwrap() + .get("provider_cache_ttl_minutes") + .is_none()); + enrich_usage_event_with_billing(&lookup, &mut legacy) + .await + .expect("legacy billing"); + assert_eq!(legacy.data.total_cost_usd, Some(6.25)); + + let queue = UsageQueue::new( + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), + UsageRuntimeConfig { + enabled: true, + queue_payload_max_bytes: 4096, + consumer_block_ms: 1, + ..UsageRuntimeConfig::default() + }, + ) + .expect("bounded billing queue"); + queue.ensure_consumer_group().await.expect("billing group"); + assert!(matches!( + queue.enqueue(&event).await, + Err(aether_data_contracts::DataLayerError::InvalidInput(_)) + )); + assert_eq!(event.to_stream_fields().unwrap(), original_fields); + assert!(queue + .read_group("typed-none-reader") + .await + .unwrap() + .is_empty()); + + // The unchanged source event remains usable by the terminal direct-write fallback. + enrich_usage_event_with_billing(&lookup, &mut event) + .await + .expect("direct fallback billing"); + assert_eq!(wire_billing_result(&event), wire_billing_result(&legacy)); + } + + #[tokio::test] + async fn wire_projection_preserves_claude_cache_ttl_and_explicit_zero_segments() { + let lookup = wire_billing_lookup( + Some(json!({ + "tiers": [{"up_to": null, "input_price_per_1m": 3.0, + "output_price_per_1m": 15.0, "cache_creation_price_per_1m": 3.75, + "cache_read_price_per_1m": 0.3, + "cache_ttl_pricing": [{"ttl_minutes": 60, + "cache_creation_price_per_1m": 6.0, "cache_read_price_per_1m": 0.6}]}] + })), + None, + ); + let mut event = wire_billing_event("wire-claude-cache"); + event.data.model = "claude-sonnet-4-6".to_string(); + event.data.target_model = Some("claude-sonnet-4-6".to_string()); + event.data.api_format = Some("claude:chat".to_string()); + event.data.endpoint_api_format = Some("claude:chat".to_string()); + event.data.provider_request_body = None; + event.data.provider_request_body_state = Some(UsageBodyCaptureState::Disabled); + event.data.cache_creation_input_tokens = Some(200); + event.data.cache_creation_ephemeral_5m_input_tokens = Some(0); + event.data.cache_creation_ephemeral_1h_input_tokens = Some(100); + event.data.request_metadata.as_mut().unwrap()["provider_cache_ttl_minutes"] = json!(60); + let queued = assert_wire_billing_equivalent(&lookup, event).await; + assert_eq!( + queued.data.cache_creation_ephemeral_5m_input_tokens, + Some(0) + ); + assert_eq!(queued.data.cache_read_input_tokens, Some(0)); + let dimensions = &queued.data.request_metadata.as_ref().unwrap()["billing_dimensions"]; + assert_eq!(dimensions["cache_ttl_minutes"], 60); + assert_eq!(dimensions["cache_creation_ephemeral_1h_tokens"], 100); + assert_eq!(dimensions["cache_creation_uncategorized_tokens"], 100); + } + + #[tokio::test] + async fn wire_projection_preserves_unknown_zero_error_and_cancellation_billing() { + let lookup = wire_billing_lookup(None, Some(0.02)); + for mode in ["unknown", "unpriced", "zero", "error_present", "cancelled"] { + let mut event = wire_billing_event(mode); + event.data.input_tokens = Some(0); + event.data.output_tokens = Some(0); + event.data.total_tokens = Some(0); + match mode { + "unknown" => { + event.data.input_tokens = None; + event.data.output_tokens = None; + event.data.total_tokens = None; + event.data.request_metadata.as_mut().unwrap()["usage_available"] = json!(false); + } + "unpriced" => { + event.data.input_tokens = Some(12); + event.data.output_tokens = Some(3); + event.data.total_tokens = Some(15); + event.data.request_metadata.as_mut().unwrap()["usage_pricing_available"] = + json!(false); + } + "error_present" => event.data.error_message = Some(String::new()), + "cancelled" => event.event_type = UsageEventType::Cancelled, + _ => {} + } + let queued = assert_wire_billing_equivalent(&lookup, event).await; + match mode { + "unknown" => { + assert_eq!(queued.data.input_tokens, None); + assert_eq!(queued.data.total_cost_usd, None); + } + "unpriced" => { + assert_eq!(queued.data.input_tokens, Some(12)); + assert_eq!(queued.data.total_cost_usd, None); + } + "error_present" => { + assert_eq!(queued.data.error_message.as_deref(), Some("")); + assert_eq!(queued.data.total_cost_usd, Some(0.0)); + } + "zero" | "cancelled" => assert_eq!(queued.data.total_cost_usd, Some(0.02)), + _ => unreachable!(), + } + } + } + + #[tokio::test] + async fn wire_projection_preserves_image_matrix_dimensions_and_request_count() { + let lookup = wire_billing_lookup( + Some(json!({ + "image_output_price_default": 0.01, + "image_output_prices": {"1536x1024": {"medium": 0.041, "high": 0.165}} + })), + Some(0.02), + ); + let mut event = wire_billing_event("wire-image"); + event.data.request_type = Some("image".to_string()); + event.data.api_format = Some("openai:image".to_string()); + event.data.endpoint_api_format = Some("openai:image".to_string()); + event.data.input_tokens = Some(0); + event.data.output_tokens = Some(0); + event.data.total_tokens = Some(0); + event.data.request_metadata.as_mut().unwrap()["dimensions"] = json!({ + "image_count": 2, "image_size": "1536x1024", "image_quality": "medium", + "image_output_format": "png" + }); + let queued = assert_wire_billing_equivalent(&lookup, event).await; + let metadata = queued.data.request_metadata.as_ref().unwrap(); + assert_eq!(metadata["billing_dimensions"]["image_count"], 2); + assert_eq!(metadata["billing_dimensions"]["request_count"], 2); + assert_eq!( + metadata["billing_dimensions"]["image_price_key"], + "1536x1024:medium" + ); + assert_eq!( + metadata["billing_snapshot"]["cost_breakdown"]["image_output_cost"], + 0.082 + ); + assert_eq!( + metadata["billing_snapshot"]["cost_breakdown"]["request_cost"], + 0.04 + ); + } + #[tokio::test] async fn unmetered_session_audit_does_not_fabricate_tokens_or_request_cost() { let lookup = TestLookup { diff --git a/crates/aether-data/adapters/postgres/src/candidates.rs b/crates/aether-data/adapters/postgres/src/candidates.rs index 44855e46c..3c86bfcbe 100644 --- a/crates/aether-data/adapters/postgres/src/candidates.rs +++ b/crates/aether-data/adapters/postgres/src/candidates.rs @@ -67,6 +67,21 @@ GROUP BY FLOOR(EXTRACT(EPOCH FROM (created_at - TO_TIMESTAMP($2))) / $4)::BIGINT "#; +const RUNTIME_CANDIDATE_COLUMNS: &str = r#" +SELECT + id, request_id, user_id, api_key_id, + NULL::text AS username, NULL::text AS api_key_name, + candidate_index, retry_index, provider_id, endpoint_id, key_id, status, + NULL::text AS skip_reason, is_cached, status_code, + NULL::text AS error_type, NULL::text AS error_message, + latency_ms, concurrent_requests, + NULL::jsonb AS extra_data, NULL::jsonb AS required_capabilities, + CAST(EXTRACT(EPOCH FROM created_at) * 1000 AS BIGINT) AS created_at_unix_ms, + CAST(EXTRACT(EPOCH FROM started_at) * 1000 AS BIGINT) AS started_at_unix_ms, + CAST(EXTRACT(EPOCH FROM finished_at) * 1000 AS BIGINT) AS finished_at_unix_ms +FROM request_candidates +"#; + const UPSERT_SQL_TEMPLATE: &str = r#" INSERT INTO request_candidates ( id, @@ -561,12 +576,29 @@ impl SqlxRequestCandidateReadRepository { pub async fn list_recent( &self, limit: usize, + ) -> Result, DataLayerError> { + self.list_recent_with_columns(limit, candidate_columns()) + .await + } + + pub async fn list_recent_runtime( + &self, + limit: usize, + ) -> Result, DataLayerError> { + self.list_recent_with_columns(limit, RUNTIME_CANDIDATE_COLUMNS) + .await + } + + async fn list_recent_with_columns( + &self, + limit: usize, + columns: &'static str, ) -> Result, DataLayerError> { if limit == 0 { return Ok(Vec::new()); } - let mut builder = QueryBuilder::::new(candidate_columns()); + let mut builder = QueryBuilder::::new(columns); builder.push(" ORDER BY created_at DESC"); push_limit( &mut builder, @@ -1068,6 +1100,13 @@ impl RequestCandidateReadRepository for SqlxRequestCandidateReadRepository { Self::list_recent(self, limit).await } + async fn list_recent_runtime( + &self, + limit: usize, + ) -> Result, DataLayerError> { + Self::list_recent_runtime(self, limit).await + } + async fn list_finalized_by_endpoint_ids_since( &self, endpoint_ids: &[String], @@ -1699,4 +1738,67 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW()) .await .expect("candidate NUL test rows should clean up"); } + + #[tokio::test] + #[ignore = "requires isolated AETHER_TEST_DATABASE_URL; uses a connection-local table"] + async fn live_postgres_candidate_runtime_projection_preserves_metadata_and_admin_rows() { + let database_url = + std::env::var("AETHER_TEST_DATABASE_URL").expect("isolated test database"); + let pool = sqlx::postgres::PgPoolOptions::new() + .max_connections(1) + .connect(&database_url) + .await + .unwrap(); + sqlx::query( + r#" +CREATE TEMP TABLE request_candidates ( + id text, request_id text, user_id text, api_key_id text, + username text, api_key_name text, candidate_index integer, retry_index integer, + provider_id text, endpoint_id text, key_id text, status text, skip_reason text, + is_cached boolean, status_code integer, error_type text, error_message text, + latency_ms integer, concurrent_requests integer, extra_data jsonb, required_capabilities jsonb, + created_at timestamptz, started_at timestamptz, finished_at timestamptz +) +"#, + ) + .execute(&pool) + .await + .unwrap(); + for index in 0..3 { + sqlx::query( + r#" +INSERT INTO request_candidates VALUES ( + $1, $2, 'user', 'api-key', NULL, NULL, $3, 0, + 'provider', 'endpoint', 'key', 'failed', NULL, false, 500, + 'upstream_error', 'admin diagnostic', 20, 17, $4, '{"vision": true}'::jsonb, + TO_TIMESTAMP(100 + $3), TO_TIMESTAMP(101 + $3), TO_TIMESTAMP(102 + $3) +) +"#, + ) + .bind(format!("candidate-{index}")) + .bind(format!("request-{index}")) + .bind(index) + .bind(json!({"upstream_response": {"body": "x".repeat(32_768)}})) + .execute(&pool) + .await + .unwrap(); + } + let repository = SqlxRequestCandidateReadRepository::new(pool.clone()); + let full = repository.list_recent(2).await.unwrap(); + let runtime = repository.list_recent_runtime(2).await.unwrap(); + assert_eq!( + runtime, + full.iter() + .map(|row| row.runtime_snapshot()) + .collect::>() + ); + assert_eq!(runtime[0].id, "candidate-2"); + assert_eq!(runtime[0].concurrent_requests, Some(17)); + assert!(runtime[0].extra_data.is_none()); + assert!(runtime[0].error_message.is_none()); + assert!(full[0].extra_data.is_some()); + assert_eq!(repository.list_recent(2).await.unwrap(), full); + assert!(repository.list_recent_runtime(0).await.unwrap().is_empty()); + pool.close().await; + } } diff --git a/crates/aether-data/adapters/postgres/src/lib.rs b/crates/aether-data/adapters/postgres/src/lib.rs index 80d052939..967c31e0c 100644 --- a/crates/aether-data/adapters/postgres/src/lib.rs +++ b/crates/aether-data/adapters/postgres/src/lib.rs @@ -52,7 +52,7 @@ pub use migrations::{ run_migrations_with_bootstrap, BootstrapFuture, PostgresMigrationBootstrap, POSTGRES_MIGRATOR, }; pub use oauth_providers::SqlxOAuthProviderRepository; -pub use pool::{PostgresPool, PostgresPoolFactory}; +pub use pool::{acquire_postgres_migration_connection, PostgresPool, PostgresPoolFactory}; pub use pool_scores::PostgresPoolMemberScoreRepository; pub use provider_catalog::SqlxProviderCatalogReadRepository; pub use proxy_nodes::SqlxProxyNodeRepository; diff --git a/crates/aether-data/adapters/postgres/src/migrations.rs b/crates/aether-data/adapters/postgres/src/migrations.rs index 94a6563ef..2525e060d 100644 --- a/crates/aether-data/adapters/postgres/src/migrations.rs +++ b/crates/aether-data/adapters/postgres/src/migrations.rs @@ -93,7 +93,7 @@ pub async fn run_migrations_with_bootstrap( pool: &PgPool, bootstrap: &dyn PostgresMigrationBootstrap, ) -> Result<(), MigrateError> { - let mut conn = pool.acquire().await?; + let mut conn = crate::pool::acquire_postgres_migration_connection(pool).await?; if POSTGRES_MIGRATOR.locking { conn.lock().await?; @@ -132,7 +132,7 @@ pub async fn prepare_database_for_startup_with_bootstrap( pool: &PgPool, bootstrap: &dyn PostgresMigrationBootstrap, ) -> Result, MigrateError> { - let mut conn = pool.acquire().await?; + let mut conn = crate::pool::acquire_postgres_migration_connection(pool).await?; if POSTGRES_MIGRATOR.locking { conn.lock().await?; diff --git a/crates/aether-data/adapters/postgres/src/pool.rs b/crates/aether-data/adapters/postgres/src/pool.rs index 544e76f24..835d2f7d5 100644 --- a/crates/aether-data/adapters/postgres/src/pool.rs +++ b/crates/aether-data/adapters/postgres/src/pool.rs @@ -4,6 +4,70 @@ use sqlx::PgPool; use std::str::FromStr; use std::time::Duration; +const STATEMENT_TIMEOUT_ENV: &str = "AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS"; +const LOCK_TIMEOUT_ENV: &str = "AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS"; + +#[derive(Debug, Clone, Copy)] +struct PostgresSessionTimeouts { + statement_ms: u32, + lock_ms: u32, +} + +impl PostgresSessionTimeouts { + fn from_env() -> Result { + Ok(Self { + statement_ms: read_timeout_env(STATEMENT_TIMEOUT_ENV, 30_000)?, + lock_ms: read_timeout_env(LOCK_TIMEOUT_ENV, 3_000)?, + }) + } + + fn apply(self, options: PgConnectOptions) -> PgConnectOptions { + options.options([ + ("statement_timeout", self.statement_ms), + ("lock_timeout", self.lock_ms), + ]) + } +} + +fn parse_timeout_ms(name: &str, value: &str) -> Result { + value + .trim() + .parse::() + .ok() + .filter(|value| *value <= i32::MAX as u32) + .ok_or_else(|| { + DataLayerError::InvalidConfiguration(format!( + "{name} must be milliseconds in 0..=2147483647 (0 disables the timeout)" + )) + }) +} + +fn read_timeout_env(name: &str, default: u32) -> Result { + match std::env::var(name) { + Ok(value) => parse_timeout_ms(name, &value), + Err(std::env::VarError::NotPresent) => Ok(default), + Err(std::env::VarError::NotUnicode(_)) => Err(DataLayerError::InvalidConfiguration( + format!("{name} must contain a valid integer"), + )), + } +} + +/// Migration and historical backfill connections are discarded on every exit path, +/// including cancellation, so their relaxed deadlines cannot escape into request work. +pub async fn acquire_postgres_migration_connection( + pool: &PgPool, +) -> Result, sqlx::Error> { + let mut conn = pool.acquire().await?; + conn.close_on_drop(); + sqlx::query("SET statement_timeout = 0") + .execute(&mut *conn) + .await?; + sqlx::query("SET lock_timeout = 0") + .execute(&mut *conn) + .await?; + Ok(conn) +} + fn connect_options(config: &PostgresPoolConfig) -> Result { config.validate()?; let options = PgConnectOptions::from_str(config.database_url.trim()).map_err(|err| { @@ -33,12 +97,16 @@ pub type PostgresPool = PgPool; #[derive(Debug, Clone)] pub struct PostgresPoolFactory { config: PostgresPoolConfig, + timeouts: PostgresSessionTimeouts, } impl PostgresPoolFactory { pub fn new(config: PostgresPoolConfig) -> Result { config.validate()?; - Ok(Self { config }) + Ok(Self { + config, + timeouts: PostgresSessionTimeouts::from_env()?, + }) } pub fn config(&self) -> &PostgresPoolConfig { @@ -46,7 +114,7 @@ impl PostgresPoolFactory { } pub fn connect_lazy(&self) -> Result { - let options = connect_options(&self.config)?; + let options = self.timeouts.apply(connect_options(&self.config)?); Ok(PgPoolOptions::new() .min_connections(self.config.min_connections) .max_connections(self.config.max_connections) @@ -59,10 +127,171 @@ impl PostgresPoolFactory { #[cfg(test)] mod tests { - use super::{connect_options, PostgresPoolFactory}; + use super::{connect_options, parse_timeout_ms, PostgresPoolFactory, PostgresSessionTimeouts}; use crate::PostgresPoolConfig; use sqlx::postgres::PgSslMode; + #[tokio::test] + async fn migration_connection_future_is_send() { + fn assert_send(_: impl Send) {} + + let pool = sqlx::postgres::PgPoolOptions::new() + .connect_lazy("postgres://localhost/aether") + .unwrap(); + assert_send(super::acquire_postgres_migration_connection(&pool)); + } + + #[test] + fn validates_session_timeout_milliseconds() { + assert_eq!(parse_timeout_ms("timeout", "0").unwrap(), 0); + assert_eq!(parse_timeout_ms("timeout", " 3000 ").unwrap(), 3_000); + assert_eq!( + parse_timeout_ms("timeout", "2147483647").unwrap(), + i32::MAX as u32 + ); + for invalid in ["", "-1", "3s", "2147483648", "4294967296"] { + assert!(parse_timeout_ms("timeout", invalid).is_err()); + } + } + + #[test] + fn session_deadlines_preserve_unrelated_connection_options() { + let options = PostgresSessionTimeouts { + statement_ms: 30_000, + lock_ms: 3_000, + } + .apply(sqlx::postgres::PgConnectOptions::new().options([("search_path", "audit")])); + assert_eq!( + options.get_options(), + Some("-c search_path=audit -c statement_timeout=30000 -c lock_timeout=3000") + ); + } + + #[tokio::test] + #[ignore = "requires an isolated AETHER_TEST_DATABASE_URL"] + async fn live_session_deadlines_rollback_transactions_and_isolate_migration_overrides() { + use crate::error::SqlxResultExt; + use crate::{PostgresTransactionOptions, PostgresTransactionRunner}; + + let factory = PostgresPoolFactory { + config: PostgresPoolConfig { + database_url: std::env::var("AETHER_TEST_DATABASE_URL").expect("test database URL"), + min_connections: 0, + max_connections: 2, + ..PostgresPoolConfig::default() + }, + timeouts: PostgresSessionTimeouts { + statement_ms: 100, + lock_ms: 40, + }, + }; + let pool = factory.connect_lazy().unwrap(); + let table = format!("deadline_test_{}", uuid::Uuid::new_v4().simple()); + sqlx::query(&format!( + "CREATE TABLE {table} (id INTEGER PRIMARY KEY, value INTEGER NOT NULL)" + )) + .execute(&pool) + .await + .unwrap(); + sqlx::query(&format!("INSERT INTO {table} VALUES (1, 0)")) + .execute(&pool) + .await + .unwrap(); + let mut blocker = pool.begin().await.unwrap(); + sqlx::query(&format!("UPDATE {table} SET value = 7 WHERE id = 1")) + .execute(&mut *blocker) + .await + .unwrap(); + let runner = PostgresTransactionRunner::new(pool.clone()); + let insert = format!("INSERT INTO {table} VALUES (2, 2)"); + let update = format!("UPDATE {table} SET value = 9 WHERE id = 1"); + let started = std::time::Instant::now(); + let error = runner + .run_read_write(|tx| { + Box::pin(async move { + sqlx::query(&insert) + .execute(&mut **tx) + .await + .map_postgres_err()?; + sqlx::query(&update) + .execute(&mut **tx) + .await + .map_postgres_err()?; + Ok(()) + }) + }) + .await + .unwrap_err(); + assert!(error.to_string().contains("SQLSTATE 55P03"), "{error}"); + assert!(started.elapsed() < std::time::Duration::from_secs(2)); + blocker.rollback().await.unwrap(); + assert_eq!( + sqlx::query_scalar::<_, i64>(&format!("SELECT COUNT(*) FROM {table}")) + .fetch_one(&pool) + .await + .unwrap(), + 1 + ); + assert_eq!( + sqlx::query_scalar::<_, i32>(&format!("SELECT value FROM {table} WHERE id = 1")) + .fetch_one(&pool) + .await + .unwrap(), + 0 + ); + + let error = sqlx::query("SELECT pg_sleep(0.3)") + .execute(&pool) + .await + .map_postgres_err() + .unwrap_err(); + assert!(error.to_string().contains("SQLSTATE 57014"), "{error}"); + runner + .run( + PostgresTransactionOptions { + statement_timeout_ms: Some(1_000), + ..PostgresTransactionOptions::read_write() + }, + |tx| { + Box::pin(async move { + sqlx::query("SELECT pg_sleep(0.15)") + .execute(&mut **tx) + .await + .map_postgres_err()?; + Ok(()) + }) + }, + ) + .await + .unwrap(); + + let mut migration = super::acquire_postgres_migration_connection(&pool) + .await + .unwrap(); + sqlx::query("SELECT pg_sleep(0.15)") + .execute(&mut *migration) + .await + .unwrap(); + drop(migration); + for _ in 0..2 { + let configured: i64 = sqlx::query_scalar( + "SELECT setting::BIGINT FROM pg_settings WHERE name = 'statement_timeout'", + ) + .fetch_one(&pool) + .await + .unwrap(); + assert_eq!( + configured, 100, + "relaxed/local overrides must not leak into pooled requests" + ); + } + sqlx::query(&format!("DROP TABLE {table}")) + .execute(&pool) + .await + .unwrap(); + pool.close().await; + } + fn ssl_mode(url: &str, require_ssl: bool) -> PgSslMode { connect_options(&PostgresPoolConfig { database_url: url.to_string(), diff --git a/crates/aether-data/adapters/postgres/src/settlement.rs b/crates/aether-data/adapters/postgres/src/settlement.rs index cd457c96f..fba7d0bd5 100644 --- a/crates/aether-data/adapters/postgres/src/settlement.rs +++ b/crates/aether-data/adapters/postgres/src/settlement.rs @@ -1,5 +1,5 @@ use async_trait::async_trait; -use sqlx::{PgPool, Row}; +use sqlx::{PgPool, Postgres, QueryBuilder, Row}; use aether_data_contracts::repository::settlement::{ finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd, @@ -297,6 +297,76 @@ fn usage_policy_subject_missing() -> DataLayerError { DataLayerError::InvalidInput("usage policy subject does not exist".to_string()) } +fn usage_policy_window_aggregate_query( + windows: impl Iterator, + aggregate: &str, +) -> Result<(QueryBuilder<'static, Postgres>, i64, i64), DataLayerError> { + let mut builder = QueryBuilder::new("SELECT "); + let mut earliest = i64::MAX; + let mut latest = i64::MIN; + for (index, (start, end)) in windows.enumerate() { + let start = usage_policy_cost_i64(start, "usage policy window start")?; + let end = usage_policy_cost_i64(end, "usage policy window end")?; + earliest = earliest.min(start); + latest = latest.max(end); + if index > 0 { + builder.push(", "); + } + builder + .push("COALESCE(") + .push(aggregate) + .push(" FILTER (WHERE admitted_at >= TO_TIMESTAMP(") + .push_bind(start) + .push("::double precision) AND admitted_at < TO_TIMESTAMP(") + .push_bind(end) + .push("::double precision)), 0)::BIGINT"); + } + Ok((builder, earliest, latest)) +} + +async fn usage_policy_request_window_counts( + tx: &mut sqlx::Transaction<'_, Postgres>, + input: &ReserveUsagePolicyRequestInput, +) -> Result { + let (mut query, earliest, latest) = usage_policy_window_aggregate_query( + input + .windows + .iter() + .map(|window| (window.starts_at_unix_secs, window.ends_at_unix_secs)), + "COUNT(*)", + )?; + // The subject lock protects all windows. One bounded history scan replaces + // repeated scans of overlapping windows without approximating their counts. + query + .push(" FROM usage_request_admissions WHERE subject_id = ") + .push_bind(input.subject_id.clone()) + .push(" AND state = 'active' AND admitted_at >= TO_TIMESTAMP(") + .push_bind(earliest) + .push("::double precision) AND admitted_at < TO_TIMESTAMP(") + .push_bind(latest) + .push("::double precision)"); + query.build().fetch_one(&mut **tx).await.map_postgres_err() +} + +async fn usage_policy_cost_window_totals( + tx: &mut sqlx::Transaction<'_, Postgres>, + input: &ReserveUsagePolicyCostInput, +) -> Result { + let (mut query, earliest, latest) = usage_policy_window_aggregate_query( + input.windows.iter().map(|window| (window.starts_at_unix_secs, window.ends_at_unix_secs)), + "SUM(CASE WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0) ELSE reserved_cost_units END)", + )?; + query.push(" FROM usage_cost_reservations WHERE subject_id = ") + .push_bind(input.subject_id.clone()) + .push(" AND admitted_at >= TO_TIMESTAMP(").push_bind(earliest) + .push("::double precision) AND admitted_at < TO_TIMESTAMP(").push_bind(latest) + .push("::double precision) AND reservation_token <> ").push_bind(input.reservation_token.clone()) + .push(" AND (state = 'finalized' OR (state = 'reserved' AND reservation_expires_at > TO_TIMESTAMP(") + .push_bind(usage_policy_cost_i64(input.admitted_at_unix_secs, "usage policy admitted_at")?) + .push("::double precision)))"); + query.build().fetch_one(&mut **tx).await.map_postgres_err() +} + fn settlement_from_row( row: &sqlx::postgres::PgRow, ) -> Result { @@ -479,6 +549,8 @@ async fn consume_daily_quota_postgres( return Ok(DailyQuotaDebitResult::default()); } let now = chrono::Utc::now(); + // Serialize each entitlement's debits. Read the shared plan's current overage policy + // from this statement's snapshot without locking every subscriber's plan row. let entitlement_rows = sqlx::query( r#" SELECT @@ -494,7 +566,7 @@ WHERE user_plan_entitlements.user_id = $1 ORDER BY user_plan_entitlements.expires_at ASC, user_plan_entitlements.created_at ASC, user_plan_entitlements.id ASC -FOR UPDATE +FOR UPDATE OF user_plan_entitlements "#, ) .bind(user_id) @@ -655,31 +727,10 @@ WHERE event_token = $1 }); } + let window_counts = usage_policy_request_window_counts(tx, &input).await?; for (window_index, window) in input.windows.iter().enumerate() { - let used_requests = sqlx::query_scalar::<_, i64>( - r#" -SELECT COUNT(*)::BIGINT -FROM usage_request_admissions -WHERE subject_id = $1 - AND state = 'active' - AND admitted_at >= TO_TIMESTAMP($2::double precision) - AND admitted_at < TO_TIMESTAMP($3::double precision) - "#, - ) - .bind(&input.subject_id) - .bind(usage_policy_cost_i64( - window.starts_at_unix_secs, - "usage policy request window start", - )?) - .bind(usage_policy_cost_i64( - window.ends_at_unix_secs, - "usage policy request window end", - )?) - .fetch_one(&mut **tx) - .await - .map_postgres_err()?; let used_requests = usage_policy_cost_u64( - used_requests, + window_counts.try_get(window_index).map_postgres_err()?, "usage policy request used_requests", )?; if used_requests >= window.limit_requests { @@ -884,43 +935,12 @@ WHERE retain_until <= TO_TIMESTAMP($1::double precision) .unwrap_or(0); let target_reserved_cost_units = previous_reserved_cost_units.max(input.reserved_cost_units); + let window_totals = usage_policy_cost_window_totals(tx, &input).await?; for (window_index, window) in input.windows.iter().enumerate() { - let used_cost_units = sqlx::query_scalar::<_, i64>( - r#" -SELECT COALESCE(SUM( - CASE - WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0) - WHEN state = 'reserved' AND reservation_expires_at > TO_TIMESTAMP($4::double precision) - THEN reserved_cost_units - ELSE 0 - END -), 0)::BIGINT -FROM usage_cost_reservations -WHERE subject_id = $1 - AND admitted_at >= TO_TIMESTAMP($2::double precision) - AND admitted_at < TO_TIMESTAMP($3::double precision) - AND reservation_token <> $5 - "#, - ) - .bind(&input.subject_id) - .bind(usage_policy_cost_i64( - window.starts_at_unix_secs, - "usage policy window start", - )?) - .bind(usage_policy_cost_i64( - window.ends_at_unix_secs, - "usage policy window end", - )?) - .bind(usage_policy_cost_i64( - input.admitted_at_unix_secs, - "usage policy admitted_at", - )?) - .bind(&input.reservation_token) - .fetch_one(&mut **tx) - .await - .map_postgres_err()?; - let used_cost_units = - usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?; + let used_cost_units = usage_policy_cost_u64( + window_totals.try_get(window_index).map_postgres_err()?, + "usage policy used_cost_units", + )?; if used_cost_units .checked_add(target_reserved_cost_units) .is_none_or(|total| total > window.limit_cost_units) @@ -1415,6 +1435,264 @@ WHERE id = $1 #[cfg(test)] mod tests { + use futures_util::FutureExt; + use std::panic::AssertUnwindSafe; + + async fn isolated_settlement_test_pool() -> (sqlx::PgPool, String) { + let database_url = std::env::var("AETHER_TEST_DATABASE_URL") + .expect("AETHER_TEST_DATABASE_URL must point at the test database"); + let schema = format!("settlement_test_{}", uuid::Uuid::new_v4().simple()); + let options = database_url + .parse::() + .expect("test database URL should parse") + .options([("search_path", schema.as_str())]); + let pool = sqlx::postgres::PgPoolOptions::new() + .max_connections(2) + .connect_with(options) + .await + .expect("test database should connect"); + sqlx::query(&format!("CREATE SCHEMA {schema}")) + .execute(&pool) + .await + .expect("isolated settlement schema should be created"); + // Separate connections must see the same fixture, so pg_temp cannot be used here. + for table in [ + "billing_plans", + "user_plan_entitlements", + "entitlement_usage_ledgers", + "users", + "usage_request_admissions", + "usage_cost_reservations", + ] { + sqlx::query(&format!( + "CREATE TABLE {table} (LIKE public.{table} INCLUDING ALL)" + )) + .execute(&pool) + .await + .expect("isolated settlement table should be created"); + } + (pool, schema) + } + + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] + async fn live_usage_policy_window_aggregates_preserve_exact_admission_and_idempotency() { + use super::*; + use aether_data_contracts::repository::settlement::{ + UsagePolicyCostWindow, UsagePolicyRequestWindow, + }; + + let (pool, schema) = isolated_settlement_test_pool().await; + let result = AssertUnwindSafe(async { + sqlx::query("INSERT INTO users (id, username, email_verified) VALUES ('subject', 'subject', false)") + .execute(&pool).await.unwrap(); + sqlx::raw_sql("INSERT INTO usage_request_admissions (request_id, subject_id, event_token, admitted_at, retain_until, state, released_at) VALUES + ('older', 'subject', 'older', TO_TIMESTAMP(50), TO_TIMESTAMP(500), 'active', NULL), + ('start', 'subject', 'start', TO_TIMESTAMP(100), TO_TIMESTAMP(500), 'active', NULL), + ('inside', 'subject', 'inside', TO_TIMESTAMP(150), TO_TIMESTAMP(500), 'active', NULL), + ('end', 'subject', 'end', TO_TIMESTAMP(200), TO_TIMESTAMP(500), 'active', NULL), + ('released', 'subject', 'released', TO_TIMESTAMP(150), TO_TIMESTAMP(500), 'released', TO_TIMESTAMP(170))") + .execute(&pool).await.unwrap(); + let repo = SqlxSettlementRepository::new(pool.clone()); + let mut request = ReserveUsagePolicyRequestInput { + request_id: "new".to_string(), subject_id: "subject".to_string(), event_token: "new".to_string(), + admitted_at_unix_secs: 175, retain_until_unix_secs: 500, + windows: vec![ + UsagePolicyRequestWindow { starts_at_unix_secs: 100, ends_at_unix_secs: 200, limit_requests: 2 }, + UsagePolicyRequestWindow { starts_at_unix_secs: 0, ends_at_unix_secs: 300, limit_requests: 4 }, + ], + }; + assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(), + ReserveUsagePolicyRequestOutcome::Rejected { window_index: 0, limit_requests: 2, used_requests: 2 }); + request.windows[0].limit_requests = 3; + assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(), + ReserveUsagePolicyRequestOutcome::Rejected { window_index: 1, limit_requests: 4, used_requests: 4 }); + request.windows[1].limit_requests = 5; + assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(), ReserveUsagePolicyRequestOutcome::Allowed); + assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(), ReserveUsagePolicyRequestOutcome::Allowed); + assert_eq!(sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM usage_request_admissions WHERE event_token = 'new'") + .fetch_one(&pool).await.unwrap(), 1); + repo.release_usage_policy_request_admission(ReleaseUsagePolicyRequestAdmissionInput { + request_id: request.request_id.clone(), subject_id: request.subject_id.clone(), event_token: request.event_token.clone(), released_at_unix_secs: 180, + }).await.unwrap(); + assert_eq!(repo.reserve_usage_policy_request(request).await.unwrap(), ReserveUsagePolicyRequestOutcome::AlreadyReleased); + + sqlx::raw_sql("INSERT INTO usage_cost_reservations (request_id, subject_id, reservation_token, admitted_at, reserved_cost_units, actual_cost_units, state, reservation_expires_at, retain_until, finalized_at) VALUES + ('older', 'subject', 'older', TO_TIMESTAMP(50), 99, 11, 'finalized', TO_TIMESTAMP(160), TO_TIMESTAMP(500), TO_TIMESTAMP(160)), + ('start', 'subject', 'start', TO_TIMESTAMP(100), 99, 7, 'finalized', TO_TIMESTAMP(160), TO_TIMESTAMP(500), TO_TIMESTAMP(160)), + ('inside', 'subject', 'inside', TO_TIMESTAMP(150), 5, NULL, 'reserved', TO_TIMESTAMP(300), TO_TIMESTAMP(500), NULL), + ('expired', 'subject', 'expired', TO_TIMESTAMP(150), 99, NULL, 'reserved', TO_TIMESTAMP(175), TO_TIMESTAMP(500), NULL), + ('end', 'subject', 'end', TO_TIMESTAMP(200), 99, 13, 'finalized', TO_TIMESTAMP(300), TO_TIMESTAMP(500), TO_TIMESTAMP(250)), + ('released', 'subject', 'released', TO_TIMESTAMP(150), 99, 0, 'released', TO_TIMESTAMP(300), TO_TIMESTAMP(500), TO_TIMESTAMP(170))") + .execute(&pool).await.unwrap(); + let mut cost = ReserveUsagePolicyCostInput { + request_id: "cost".to_string(), subject_id: "subject".to_string(), reservation_token: "cost".to_string(), + admitted_at_unix_secs: 175, reserved_cost_units: 3, reservation_expires_at_unix_secs: 400, retain_until_unix_secs: 500, + windows: vec![ + UsagePolicyCostWindow { window_id: "short".to_string(), starts_at_unix_secs: 100, ends_at_unix_secs: 200, limit_cost_units: 14 }, + UsagePolicyCostWindow { window_id: "long".to_string(), starts_at_unix_secs: 0, ends_at_unix_secs: 300, limit_cost_units: 38 }, + ], + }; + assert_eq!(repo.reserve_usage_policy_cost(cost.clone()).await.unwrap(), + ReserveUsagePolicyCostOutcome::Rejected { window_index: 0, limit_cost_units: 14, used_cost_units: 12 }); + cost.windows[0].limit_cost_units = 15; + assert_eq!(repo.reserve_usage_policy_cost(cost.clone()).await.unwrap(), + ReserveUsagePolicyCostOutcome::Rejected { window_index: 1, limit_cost_units: 38, used_cost_units: 36 }); + cost.windows[1].limit_cost_units = 39; + let allowed = repo.reserve_usage_policy_cost(cost.clone()).await.unwrap(); + assert!(matches!(allowed, ReserveUsagePolicyCostOutcome::Allowed { .. }), "{allowed:?}"); + let repeated = repo.reserve_usage_policy_cost(cost.clone()).await.unwrap(); + assert!(matches!(repeated, ReserveUsagePolicyCostOutcome::Allowed { .. }), "{repeated:?}"); + cost.reserved_cost_units = 4; + assert_eq!(repo.reserve_usage_policy_cost(cost).await.unwrap(), + ReserveUsagePolicyCostOutcome::Rejected { window_index: 0, limit_cost_units: 15, used_cost_units: 12 }); + + sqlx::query("DELETE FROM usage_request_admissions").execute(&pool).await.unwrap(); + let make_request = |id: &str| ReserveUsagePolicyRequestInput { + request_id: id.to_string(), subject_id: "subject".to_string(), event_token: id.to_string(), + admitted_at_unix_secs: 175, retain_until_unix_secs: 500, + windows: vec![UsagePolicyRequestWindow { starts_at_unix_secs: 0, ends_at_unix_secs: 300, limit_requests: 1 }], + }; + let (first, second) = tokio::join!(repo.reserve_usage_policy_request(make_request("race-a")), repo.reserve_usage_policy_request(make_request("race-b"))); + let outcomes = [first.unwrap(), second.unwrap()]; + assert_eq!(outcomes.iter().filter(|outcome| matches!(outcome, ReserveUsagePolicyRequestOutcome::Allowed)).count(), 1); + assert_eq!(outcomes.iter().filter(|outcome| matches!(outcome, ReserveUsagePolicyRequestOutcome::Rejected { used_requests: 1, .. })).count(), 1); + }).catch_unwind().await; + sqlx::query(&format!("DROP SCHEMA {schema} CASCADE")) + .execute(&pool) + .await + .unwrap(); + pool.close().await; + if let Err(panic) = result { + std::panic::resume_unwind(panic); + } + } + + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] + async fn live_daily_quota_serializes_each_entitlement_without_locking_shared_plan() { + let (pool, schema) = isolated_settlement_test_pool().await; + let result = AssertUnwindSafe(async { + let grant = serde_json::json!([{ + "type": "daily_quota", + "daily_quota_usd": 10.0, + "reset_timezone": "UTC", + "allow_wallet_overage": false, + }]); + sqlx::query( + "INSERT INTO billing_plans (id, title, price_amount, duration_unit, duration_value, entitlements_json, created_at, updated_at) VALUES ('shared-plan', 'Shared plan', 10, 'month', 1, $1, NOW(), NOW())", + ) + .bind(&grant) + .execute(&pool) + .await + .expect("shared plan should insert"); + for user_id in ["user-a", "user-b"] { + sqlx::query( + "INSERT INTO user_plan_entitlements (id, user_id, plan_id, payment_order_id, starts_at, expires_at, entitlements_snapshot, created_at, updated_at) VALUES ($1, $1, 'shared-plan', $1, NOW() - INTERVAL '1 hour', NOW() + INTERVAL '1 day', $2, NOW(), NOW())", + ) + .bind(user_id) + .bind(&grant) + .execute(&pool) + .await + .expect("user entitlement should insert"); + } + + let mut first = pool.begin().await.expect("first transaction should start"); + let first_debit = super::consume_daily_quota_postgres( + &mut first, "user-a", "request-a", 7.0, Some(0.0), false, + ) + .await + .expect("first user should consume quota"); + assert_eq!(first_debit.debited_usd, 7.0); + assert!(!first_debit.insufficient); + + let mut second = pool.begin().await.expect("second transaction should start"); + sqlx::query("SET LOCAL lock_timeout = '500ms'") + .execute(&mut *second) + .await + .expect("lock timeout should be configured"); + let second_debit = super::consume_daily_quota_postgres( + &mut second, "user-b", "request-b", 2.0, Some(0.0), false, + ) + .await + .expect("another user's quota must not wait for the shared plan"); + assert_eq!(second_debit.debited_usd, 2.0); + assert!(!second_debit.insufficient); + second.commit().await.expect("second debit should commit"); + + let mut same_user = pool.begin().await.expect("contending transaction should start"); + sqlx::query("SET LOCAL lock_timeout = '500ms'") + .execute(&mut *same_user) + .await + .expect("lock timeout should be configured"); + let blocked = super::consume_daily_quota_postgres( + &mut same_user, "user-a", "request-a-next", 2.0, Some(0.0), false, + ) + .await + .expect_err("the same entitlement must remain locked until commit"); + assert!(blocked.to_string().contains("SQLSTATE 55P03"), "{blocked}"); + same_user.rollback().await.expect("blocked transaction should roll back"); + first.commit().await.expect("first debit should commit"); + + let mut next = pool.begin().await.expect("next transaction should start"); + let next_debit = super::consume_daily_quota_postgres( + &mut next, "user-a", "request-a-next", 2.0, Some(0.0), false, + ) + .await + .expect("same user should consume the remaining quota after commit"); + assert_eq!(next_debit.debited_usd, 2.0); + assert!(!next_debit.insufficient); + next.commit().await.expect("next debit should commit"); + let balance: (f64, f64) = sqlx::query_as( + "SELECT balance_before::double precision, balance_after::double precision FROM entitlement_usage_ledgers WHERE request_id = 'request-a-next'", + ) + .fetch_one(&pool) + .await + .expect("next debit ledger should exist"); + assert_eq!(balance, (3.0, 1.0)); + + let mut held = pool.begin().await.expect("quota transaction should start"); + super::consume_daily_quota_postgres( + &mut held, "user-a", "request-policy-before", 0.5, Some(0.0), false, + ) + .await + .expect("quota transaction should retain its entitlement lock"); + let mut edit = pool.begin().await.expect("plan edit transaction should start"); + sqlx::query("SET LOCAL lock_timeout = '500ms'") + .execute(&mut *edit) + .await + .expect("plan edit timeout should be configured"); + sqlx::query( + "UPDATE billing_plans SET entitlements_json = jsonb_set(entitlements_json, '{0,allow_wallet_overage}', 'true'::jsonb) WHERE id = 'shared-plan'", + ) + .execute(&mut *edit) + .await + .expect("plan configuration edits must not wait for usage settlement"); + edit.commit().await.expect("plan edit should commit"); + held.rollback().await.expect("held quota debit should roll back"); + + let mut after_edit = pool.begin().await.expect("fresh transaction should start"); + let updated_policy = super::consume_daily_quota_postgres( + &mut after_edit, "user-a", "request-policy-after", 2.0, Some(5.0), true, + ) + .await + .expect("fresh quota read should use current plan configuration"); + assert!(!updated_policy.insufficient); + assert_eq!(updated_policy.debited_usd, 1.0); + after_edit.rollback().await.expect("policy verification should roll back"); + }) + .catch_unwind() + .await; + sqlx::query(&format!("DROP SCHEMA {schema} CASCADE")) + .execute(&pool) + .await + .expect("isolated settlement schema should be removed"); + pool.close().await; + if let Err(panic) = result { + std::panic::resume_unwind(panic); + } + } + #[test] fn finalize_usage_billing_sql_does_not_require_usage_updated_at_column() { assert!(!super::FINALIZE_USAGE_BILLING_SQL.contains("updated_at")); diff --git a/crates/aether-data/adapters/postgres/src/tx.rs b/crates/aether-data/adapters/postgres/src/tx.rs index a31a0d7ea..c6471b892 100644 --- a/crates/aether-data/adapters/postgres/src/tx.rs +++ b/crates/aether-data/adapters/postgres/src/tx.rs @@ -34,6 +34,14 @@ impl PostgresTransactionOptions { } } + pub fn maintenance() -> Self { + Self { + mode: TransactionMode::ReadWrite, + statement_timeout_ms: Some(300_000), + lock_timeout_ms: Some(30_000), + } + } + pub fn validate(&self) -> Result<(), DataLayerError> { if matches!(self.statement_timeout_ms, Some(0)) { return Err(DataLayerError::InvalidConfiguration( diff --git a/crates/aether-data/adapters/postgres/src/usage/mod.rs b/crates/aether-data/adapters/postgres/src/usage/mod.rs index 84360990d..23f7615d3 100644 --- a/crates/aether-data/adapters/postgres/src/usage/mod.rs +++ b/crates/aether-data/adapters/postgres/src/usage/mod.rs @@ -34,7 +34,7 @@ use sqlx::{ PgPool, Postgres, QueryBuilder, Row, }; use std::collections::{BTreeMap, BTreeSet}; -use std::io::Write; +use std::io::{BufWriter, Write}; use uuid::Uuid; use crate::{ @@ -58,6 +58,9 @@ use aether_data_contracts::repository::usage::{ use aether_data_contracts::DataLayerError; pub mod cleanup; +mod preparation; + +use preparation::prepare_usage_in_background; // Legacy inline body columns on public.usage are deprecated. Keep the threshold at zero so // newly captured bodies always spill to usage_body_blobs and resolve through usage_http_audits. @@ -2121,12 +2124,9 @@ impl PreparedPendingUsage { )); } - // Keep the capture input separate from the accounting row. The persistence sanitizer - // intentionally removes HTTP bodies/headers/states, but the pending batch still needs - // those values to populate the canonical audit/blob tables. - let capture_usage = usage.clone(); - let usage = sanitize_usage_for_persistence(usage); - let prepared = prepare_usage_upsert_context(&capture_usage)?; + // Prepare captures before the accounting sanitizer removes HTTP bodies/headers/states. + let (usage, prepared) = prepare_usage_for_persistence(usage); + let prepared = prepared?; let input_tokens = usage .input_tokens .map(to_i32) @@ -8450,10 +8450,10 @@ ORDER BY "usage".user_id ASC usage: UpsertUsageRecord, ) -> Result { usage.validate()?; - // `usage` is the sanitized accounting projection; prepare the auxiliary capture and - // snapshots from the original event so typed `none` markers can clear prior facts. - let capture_usage = usage.clone(); - let usage = sanitize_usage_for_persistence(usage); + // Move the event before cloning or compressing captures, and do not hold a connection + // while preparing them. Stale lifecycle updates still ignore preparation errors below. + let (usage, prepared) = + prepare_usage_in_background(move || Ok(prepare_usage_for_persistence(usage))).await?; self.tx_runner .run_read_write(|tx| { Box::pin(async move { @@ -8519,7 +8519,7 @@ ORDER BY "usage".user_id ASC clear_provider_request_body, clear_response_body, clear_client_response_body, - } = prepare_usage_upsert_context(&capture_usage)?; + } = prepared?; let capture_update_allowed = recovers_terminal_failure || usage_capture_update_allowed( previous_usage.as_ref().map(|stored| { @@ -8938,33 +8938,36 @@ ORDER BY "usage".user_id ASC return Ok(()); } - let mut request_id_counts = BTreeMap::::new(); - for usage in &usages { - *request_id_counts - .entry(usage.request_id.clone()) - .or_default() += 1; - } - - // Duplicate request IDs must retain the caller's exact sequential merge order. They are - // uncommon in lifecycle batches, so keep them on the canonical single-row path. - let mut batch_rows = Vec::<(usize, PreparedPendingUsage)>::new(); - let mut fallback_rows = Vec::<(usize, UpsertUsageRecord)>::new(); - for (sequence, usage) in usages.into_iter().enumerate() { - let original_usage = usage.clone(); - let prepared = PreparedPendingUsage::try_from_usage(usage)?; - if request_id_counts - .get(&prepared.usage.request_id) - .copied() - .unwrap_or_default() - == 1 - { - batch_rows.push((sequence, prepared)); - } else { - // Preserve capture markers for the canonical fallback; that path performs the - // sanitized bind only after preparing the auxiliary audit/blob state. - fallback_rows.push((sequence, original_usage)); + let (batch_rows, mut fallback_rows) = prepare_usage_in_background(move || { + let mut request_id_counts = BTreeMap::::new(); + for usage in &usages { + *request_id_counts + .entry(usage.request_id.clone()) + .or_default() += 1; } - } + + // Duplicate request IDs must retain the caller's exact sequential merge order. + let mut batch_rows = Vec::<(usize, PreparedPendingUsage)>::new(); + let mut fallback_rows = Vec::<(usize, UpsertUsageRecord)>::new(); + for (sequence, usage) in usages.into_iter().enumerate() { + let duplicate = request_id_counts + .get(&usage.request_id) + .copied() + .unwrap_or_default() + > 1; + let original_usage = duplicate.then(|| usage.clone()); + let prepared = PreparedPendingUsage::try_from_usage(usage)?; + if let Some(original_usage) = original_usage { + // Preserve capture markers for the canonical fallback, including validation + // of every row before starting the batch transaction. + fallback_rows.push((sequence, original_usage)); + } else { + batch_rows.push((sequence, prepared)); + } + } + Ok((batch_rows, fallback_rows)) + }) + .await?; let mut inserted_request_ids = BTreeSet::::new(); if !batch_rows.is_empty() { @@ -10294,7 +10297,7 @@ RETURNING pub async fn rebuild_api_key_usage_stats(&self) -> Result { self.tx_runner - .run_read_write(|tx| { + .run(crate::PostgresTransactionOptions::maintenance(), |tx| { Box::pin(async move { sqlx::query(RESET_API_KEY_USAGE_STATS_SQL) .execute(&mut **tx) @@ -10313,7 +10316,7 @@ RETURNING pub async fn rebuild_provider_api_key_usage_stats(&self) -> Result { self.tx_runner - .run_read_write(|tx| { + .run(crate::PostgresTransactionOptions::maintenance(), |tx| { Box::pin(async move { sqlx::query(RESET_PROVIDER_API_KEY_USAGE_STATS_SQL) .execute(&mut **tx) @@ -12359,24 +12362,35 @@ fn prepare_usage_body_storage(value: Option<&Value>) -> Result ( + UpsertUsageRecord, + Result, +) { + // Capture controls and accounting metadata have different sanitizers. Move the + // large payloads out before copying the metadata needed by both projections. + let request_body = usage.request_body.take(); + let provider_request_body = usage.provider_request_body.take(); + let response_body = usage.response_body.take(); + let client_response_body = usage.client_response_body.take(); + let request_headers = usage.request_headers.take(); + let provider_request_headers = usage.provider_request_headers.take(); + let response_headers = usage.response_headers.take(); + let client_response_headers = usage.client_response_headers.take(); + let mut capture = usage.clone(); + capture.request_body = request_body; + capture.provider_request_body = provider_request_body; + capture.response_body = response_body; + capture.client_response_body = client_response_body; + capture.request_headers = request_headers; + capture.provider_request_headers = provider_request_headers; + capture.response_headers = response_headers; + capture.client_response_headers = client_response_headers; + capture.capture_retention = std::mem::take(&mut usage.capture_retention); + let capture = sanitize_usage_capture_controls_for_persistence(capture); + let prepared = prepare_usage_upsert_context(&capture); + (sanitize_usage_for_persistence(usage), prepared) +} + fn prepare_usage_upsert_context( usage: &UpsertUsageRecord, ) -> Result { - let usage = sanitize_usage_capture_controls_for_persistence(usage.clone()); - let usage = &usage; let replace_client_request_body_facts = request_body_capture_replaces_derived_facts( usage.request_body.as_ref(), usage.request_body_state, diff --git a/crates/aether-data/adapters/postgres/src/usage/preparation.rs b/crates/aether-data/adapters/postgres/src/usage/preparation.rs new file mode 100644 index 000000000..215472769 --- /dev/null +++ b/crates/aether-data/adapters/postgres/src/usage/preparation.rs @@ -0,0 +1,327 @@ +use std::sync::{Arc, OnceLock}; +use std::time::Duration; + +use aether_data_contracts::DataLayerError; +use tokio::sync::Semaphore; + +static USAGE_PREPARATION_EXECUTOR: OnceLock = OnceLock::new(); + +pub(super) async fn prepare_usage_in_background( + prepare: impl FnOnce() -> Result + Send + 'static, +) -> Result { + USAGE_PREPARATION_EXECUTOR + .get_or_init(|| { + UsagePreparationExecutor::new(4, 32, Duration::from_secs(1), Duration::from_secs(30)) + }) + .run(prepare) + .await +} + +struct UsagePreparationExecutor { + workers: Arc, + admitted: Arc, + queue_timeout: Duration, + execution_timeout: Duration, +} + +impl UsagePreparationExecutor { + fn new( + workers: usize, + admitted: usize, + queue_timeout: Duration, + execution_timeout: Duration, + ) -> Self { + Self { + workers: Arc::new(Semaphore::new(workers)), + admitted: Arc::new(Semaphore::new(admitted)), + queue_timeout, + execution_timeout, + } + } + + async fn run( + &self, + prepare: impl FnOnce() -> Result + Send + 'static, + ) -> Result { + // Bound both running work and callers retaining input while waiting for a worker. + // These limits count tasks, not bytes in the caller's original usage records. + let admitted = self.admitted.clone().try_acquire_owned().map_err(|_| { + DataLayerError::TimedOut("usage preparation capacity exhausted".to_string()) + })?; + let worker = tokio::time::timeout(self.queue_timeout, self.workers.clone().acquire_owned()) + .await + .map_err(|_| { + DataLayerError::TimedOut( + "timed out waiting for usage preparation worker".to_string(), + ) + })? + .map_err(|_| { + DataLayerError::TimedOut("usage preparation workers unavailable".to_string()) + })?; + + // The closure owns both permits even if its caller times out or is cancelled. It + // prepares input only; detached completion must never begin a database transaction. + let mut task = tokio::task::spawn_blocking(move || { + let _admitted = admitted; + let _worker = worker; + prepare() + }); + match tokio::time::timeout(self.execution_timeout, &mut task).await { + Ok(result) => result.map_err(|error| { + DataLayerError::TimedOut(format!("usage preparation worker failed: {error}")) + })?, + Err(_) => { + // This cancels work still queued in Tokio; running blocking work keeps its + // permits until it actually exits, since abort cannot stop a blocking thread. + task.abort(); + Err(DataLayerError::TimedOut( + "timed out preparing usage storage".to_string(), + )) + } + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicBool, Ordering}; + use std::sync::mpsc; + + use super::*; + + fn executor( + admitted: usize, + queue_timeout: Duration, + execution_timeout: Duration, + ) -> Arc { + Arc::new(UsagePreparationExecutor::new( + 1, + admitted, + queue_timeout, + execution_timeout, + )) + } + + async fn wait_for_worker_release(executor: &UsagePreparationExecutor) { + tokio::time::timeout(Duration::from_secs(2), async { + while executor.workers.available_permits() != 1 + || executor.admitted.available_permits() == 0 + { + tokio::task::yield_now().await; + } + }) + .await + .expect("finished blocking work should release its permits"); + } + + #[tokio::test(flavor = "current_thread")] + async fn preparation_runs_off_the_runtime_thread_and_preserves_errors() { + let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2)); + let runtime_thread = std::thread::current().id(); + executor + .run(move || { + assert_ne!(std::thread::current().id(), runtime_thread); + Ok(()) + }) + .await + .expect("preparation should succeed"); + + let error = executor + .run(|| Err::<(), _>(DataLayerError::InvalidInput("bad usage".to_string()))) + .await + .expect_err("input errors must reach the caller"); + assert!(matches!(error, DataLayerError::InvalidInput(message) if message == "bad usage")); + } + + #[tokio::test] + async fn saturated_admission_rejects_work_without_running_it() { + let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let first_executor = executor.clone(); + let first = tokio::spawn(async move { + first_executor + .run(move || { + let _ = started_tx.send(()); + let _ = release_rx.recv(); + Ok(()) + }) + .await + }); + started_rx.await.expect("first job should start"); + + let error = executor + .run(|| -> Result<(), DataLayerError> { panic!("rejected work must not execute") }) + .await + .expect_err("admission should fail immediately"); + assert!(matches!(error, DataLayerError::TimedOut(message) if message.contains("capacity"))); + release_tx + .send(()) + .expect("first job should still be alive"); + first.await.unwrap().unwrap(); + } + + #[tokio::test] + async fn blocking_worker_failure_is_retryable_and_releases_capacity() { + let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2)); + let error = executor + .run(|| -> Result<(), DataLayerError> { panic!("simulated worker failure") }) + .await + .expect_err("worker failure must reach the caller"); + assert!( + matches!(error, DataLayerError::TimedOut(message) if message.contains("worker failed")) + ); + executor + .run(|| Ok(())) + .await + .expect("failed workers should release capacity"); + } + + #[tokio::test] + async fn waiting_for_a_worker_has_a_deadline_and_never_starts_expired_work() { + let executor = executor(2, Duration::from_millis(20), Duration::from_secs(2)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let first_executor = executor.clone(); + let first = tokio::spawn(async move { + first_executor + .run(move || { + let _ = started_tx.send(()); + let _ = release_rx.recv(); + Ok(()) + }) + .await + }); + started_rx.await.expect("first job should start"); + let ran = Arc::new(AtomicBool::new(false)); + let work_ran = ran.clone(); + let error = executor + .run(move || { + work_ran.store(true, Ordering::SeqCst); + Ok(()) + }) + .await + .expect_err("the queued job should time out"); + assert!(matches!(error, DataLayerError::TimedOut(message) if message.contains("waiting"))); + assert_eq!(executor.admitted.available_permits(), 1); + release_tx + .send(()) + .expect("first job should still be alive"); + first.await.unwrap().unwrap(); + assert!(!ran.load(Ordering::SeqCst)); + } + + #[tokio::test] + async fn cancellation_keeps_permits_until_running_blocking_work_exits() { + let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let first_executor = executor.clone(); + let first = tokio::spawn(async move { + first_executor + .run(move || { + let _ = started_tx.send(()); + let _ = release_rx.recv(); + Ok(()) + }) + .await + }); + started_rx.await.expect("first job should start"); + first.abort(); + assert!(first.await.unwrap_err().is_cancelled()); + assert_eq!(executor.workers.available_permits(), 0); + assert!(matches!( + executor.run(|| Ok(())).await, + Err(DataLayerError::TimedOut(_)) + )); + release_tx + .send(()) + .expect("blocking work should outlive cancellation"); + wait_for_worker_release(&executor).await; + executor + .run(|| Ok(())) + .await + .expect("the executor should recover"); + } + + #[tokio::test] + async fn cancellation_keeps_capture_budget_until_blocking_input_is_dropped() { + use aether_data_contracts::repository::usage::{ + usage_json_heap_estimate, UpsertUsageRecord, UsageCaptureMemoryBudget, + }; + + let mut usage: UpsertUsageRecord = serde_json::from_value(serde_json::json!({ + "request_id": "req-cancelled-preparation", + "provider_name": "test", + "model": "test", + "status": "completed", + "billing_status": "pending", + "updated_at_unix_secs": 100, + "request_body": {"content": "retained".repeat(1024)} + })) + .unwrap(); + let bytes = std::mem::size_of::() + + usage_json_heap_estimate(usage.request_body.as_ref().unwrap()); + let budget = Arc::new(UsageCaptureMemoryBudget::new(bytes)); + assert!(usage.capture_retention.reserve(Arc::clone(&budget), bytes)); + let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let first_executor = Arc::clone(&executor); + let first = tokio::spawn(async move { + first_executor + .run(move || { + let _ = started_tx.send(()); + let _ = release_rx.recv(); + drop(usage); + Ok(()) + }) + .await + }); + started_rx.await.unwrap(); + first.abort(); + assert!(first.await.unwrap_err().is_cancelled()); + assert_eq!(budget.retained_bytes(), bytes); + release_tx.send(()).unwrap(); + wait_for_worker_release(&executor).await; + assert_eq!(budget.retained_bytes(), 0); + } + + #[tokio::test] + async fn execution_timeout_keeps_permits_until_running_blocking_work_exits() { + let executor = executor(1, Duration::from_secs(1), Duration::from_millis(20)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = mpsc::channel(); + let first_executor = executor.clone(); + let first = tokio::spawn(async move { + first_executor + .run(move || { + let _ = started_tx.send(()); + let _ = release_rx.recv(); + Ok(()) + }) + .await + }); + started_rx.await.expect("first job should start"); + let error = first + .await + .unwrap() + .expect_err("running work should time out"); + assert!( + matches!(error, DataLayerError::TimedOut(message) if message.contains("preparing")) + ); + assert_eq!(executor.workers.available_permits(), 0); + assert!(matches!( + executor.run(|| Ok(())).await, + Err(DataLayerError::TimedOut(_)) + )); + release_tx + .send(()) + .expect("blocking work should outlive timeout"); + wait_for_worker_release(&executor).await; + executor + .run(|| Ok(())) + .await + .expect("the executor should recover"); + } +} diff --git a/crates/aether-data/adapters/postgres/src/usage/tests.rs b/crates/aether-data/adapters/postgres/src/usage/tests.rs index a279f053e..082765ae7 100644 --- a/crates/aether-data/adapters/postgres/src/usage/tests.rs +++ b/crates/aether-data/adapters/postgres/src/usage/tests.rs @@ -8,7 +8,7 @@ use super::{ attach_usage_routing_snapshot_metadata, attach_usage_settlement_pricing_snapshot_metadata, clear_previous_request_body_facts, inflate_usage_json_value, prepare_request_metadata_for_body_storage, prepare_usage_body_storage, - prepare_usage_upsert_context, push_postgres_usage_websocket_filter, + prepare_usage_for_persistence, push_postgres_usage_websocket_filter, request_body_capture_replaces_derived_facts, resolved_read_usage_body_ref, resolved_write_usage_body_ref, split_dashboard_daily_aggregate_range, split_dashboard_hourly_aggregate_range, usage_body_capture_state_for_storage, usage_body_ref, @@ -39,6 +39,7 @@ fn fast_clear_usage_record( terminal_service_tier: Option<&str>, ) -> UpsertUsageRecord { UpsertUsageRecord { + capture_retention: Default::default(), request_id: request_id.to_string(), user_id: None, api_key_id: None, @@ -204,7 +205,26 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() { let repository = SqlxUsageReadRepository::new(factory.connect_lazy().unwrap()); crate::run_migrations(repository.pool()).await.unwrap(); - for batch in [false, true] { + for write_mode in 0..3 { + use aether_data_contracts::repository::usage::{ + usage_json_heap_estimate, UsageCaptureMemoryBudget, + }; + let batch = write_mode != 0; + let budget = Arc::new(UsageCaptureMemoryBudget::new(4 * 1024 * 1024)); + let retain_capture = |usage: &mut UpsertUsageRecord| { + let bytes = [ + usage.request_body.as_ref(), + usage.provider_request_body.as_ref(), + usage.response_body.as_ref(), + usage.client_response_body.as_ref(), + ] + .into_iter() + .flatten() + .map(|body| std::mem::size_of::() + usage_json_heap_estimate(body)) + .sum(); + assert!(usage.capture_retention.reserve(Arc::clone(&budget), bytes)); + bytes + }; let request_id = format!("req-full-capture-{}", uuid::Uuid::new_v4().simple()); let now_unix_secs = Utc::now().timestamp() as u64; let mut pending = fast_clear_usage_record( @@ -225,14 +245,18 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() { pending.response_body_state = Some(UsageBodyCaptureState::Inline); pending.client_response_body = Some(json!("pending client response")); pending.client_response_body_state = Some(UsageBodyCaptureState::Inline); + let pending_bytes = retain_capture(&mut pending); if batch { - repository - .upsert_pending_many(vec![pending.clone()]) - .await - .unwrap(); + let records = if write_mode == 2 { + vec![pending.clone(), pending.clone()] + } else { + vec![pending.clone()] + }; + repository.upsert_pending_many(records).await.unwrap(); } else { repository.upsert(pending.clone()).await.unwrap(); } + assert_eq!(budget.retained_bytes(), pending_bytes); for (field, expected) in [ (UsageBodyField::RequestBody, pending.request_body.as_ref()), ( @@ -274,7 +298,10 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() { terminal.response_body_state = Some(UsageBodyCaptureState::Inline); terminal.client_response_body = Some(json!({"output": "final response"})); terminal.client_response_body_state = Some(UsageBodyCaptureState::Inline); + let terminal_bytes = retain_capture(&mut terminal); repository.upsert(terminal.clone()).await.unwrap(); + assert_eq!(budget.retained_bytes(), pending_bytes + terminal_bytes); + assert_eq!(budget.downgraded_total(), 0); let stored = repository .find_by_request_id_shallow(&request_id) @@ -330,6 +357,9 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() { .execute(repository.pool()) .await .unwrap(); + drop(pending); + drop(terminal); + assert_eq!(budget.retained_bytes(), 0); } } @@ -2411,6 +2441,7 @@ async fn validates_upsert_before_hitting_database() { let repository = SqlxUsageReadRepository::new(pool); let result = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "".to_string(), user_id: None, api_key_id: None, @@ -4324,6 +4355,96 @@ fn prepare_usage_body_storage_compresses_large_payloads() { ); } +#[test] +fn prepare_usage_body_storage_streams_json_shapes_into_compatible_gzip() { + for payload in [ + serde_json::Value::Null, + json!(false), + json!(42), + json!(["quoted\"text", "line\nbreak", "\u{4e2d}\u{6587}", null]), + json!({ + "content": "escaped\n\"\\value".repeat(32 * 1024), + "nested": {"values": [true, null, 1.25, -7]} + }), + ] { + let storage = prepare_usage_body_storage(Some(&payload)).expect("body should compress"); + assert!(storage.inline_json.is_none()); + let compressed = storage + .detached_blob_bytes + .expect("body should be detached"); + assert_eq!( + inflate_usage_json_value(&compressed).expect("body should remain readable"), + payload + ); + } +} + +#[test] +fn managed_capture_preparation_moves_bodies_without_a_second_reservation() { + use aether_data_contracts::repository::usage::{ + sanitize_usage_for_persistence, usage_json_heap_estimate, UsageCaptureMemoryBudget, + }; + + let mut usage = fast_clear_usage_record( + "req-managed-capture", + "managed-capture", + 100, + true, + UsageBodyCaptureState::Inline, + Some("priority"), + ); + let bodies = [ + json!({"messages": [{"role": "user", "content": "request".repeat(4096)}]}), + json!({"input": "provider request".repeat(4096), "service_tier": "priority"}), + json!({"output": "provider response".repeat(4096)}), + json!({"output": "client response".repeat(4096)}), + ]; + usage.request_body = Some(bodies[0].clone()); + usage.provider_request_body = Some(bodies[1].clone()); + usage.response_body = Some(bodies[2].clone()); + usage.client_response_body = Some(bodies[3].clone()); + usage.request_body_state = Some(UsageBodyCaptureState::Inline); + usage.response_body_state = Some(UsageBodyCaptureState::Inline); + usage.client_response_body_state = Some(UsageBodyCaptureState::Inline); + usage.request_headers = Some(json!({"content-type": "application/json"})); + usage.cache_read_input_tokens = Some(0); + usage.total_cost_usd = Some(0.25); + usage.actual_total_cost_usd = Some(0.125); + let expected_accounting = sanitize_usage_for_persistence(usage.clone()); + let bytes = [ + usage.request_body.as_ref(), + usage.provider_request_body.as_ref(), + usage.response_body.as_ref(), + usage.client_response_body.as_ref(), + ] + .into_iter() + .flatten() + .map(|body| std::mem::size_of::() + usage_json_heap_estimate(body)) + .sum(); + let budget = Arc::new(UsageCaptureMemoryBudget::new(bytes)); + assert!(usage.capture_retention.reserve(Arc::clone(&budget), bytes)); + + let (accounting, prepared) = prepare_usage_for_persistence(usage); + let prepared = prepared.expect("managed capture should prepare without cloning bodies"); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 0); + assert_eq!(accounting, expected_accounting); + for (storage, expected) in [ + prepared.request_body_storage, + prepared.provider_request_body_storage, + prepared.response_body_storage, + prepared.client_response_body_storage, + ] + .into_iter() + .zip(bodies) + { + assert_eq!( + inflate_usage_json_value(storage.detached_blob_bytes.as_deref().unwrap()).unwrap(), + expected + ); + } +} + #[test] fn usage_body_capture_state_for_storage_marks_detached_bodies_as_reference() { let payload = json!({"message": "hello"}); @@ -4413,7 +4534,8 @@ fn explicit_none_capture_drops_residual_body_ref_and_incoming_fast_metadata_befo "provider_request_body_ref": "usage://request/req-none-residual/provider_request_body" })); - let prepared = prepare_usage_upsert_context(&usage).expect("usage should prepare"); + let (_, prepared) = prepare_usage_for_persistence(usage); + let prepared = prepared.expect("usage should prepare"); assert!(prepared.clear_provider_request_body); assert!(!prepared.provider_request_body_storage.has_detached_blob()); assert_eq!(prepared.http_audit_refs.provider_request_body_ref, None); @@ -4805,6 +4927,7 @@ fn attach_usage_http_audit_body_refs_adds_missing_metadata_without_overwriting_e fn usage_routing_snapshot_from_usage_only_activates_for_routing_metadata() { let snapshot = usage_routing_snapshot_from_usage( &UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-123".to_string(), user_id: None, api_key_id: None, @@ -4904,6 +5027,7 @@ fn usage_routing_snapshot_from_usage_only_activates_for_routing_metadata() { let empty_snapshot = usage_routing_snapshot_from_usage( &UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-124".to_string(), user_id: None, api_key_id: None, @@ -4982,6 +5106,7 @@ fn usage_routing_snapshot_from_usage_only_activates_for_routing_metadata() { fn usage_routing_snapshot_from_usage_prefers_typed_routing_fields_without_metadata() { let snapshot = usage_routing_snapshot_from_usage( &UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-typed-routing-1".to_string(), user_id: None, api_key_id: None, @@ -5116,6 +5241,7 @@ fn attach_usage_routing_snapshot_metadata_adds_missing_keys_without_overwriting_ fn usage_settlement_pricing_snapshot_from_usage_extracts_typed_billing_fields() { let snapshot = usage_settlement_pricing_snapshot_from_usage( &UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-125".to_string(), user_id: None, api_key_id: None, diff --git a/crates/aether-data/contracts/src/repository/candidates/types.rs b/crates/aether-data/contracts/src/repository/candidates/types.rs index 37b4b5b37..ec4787faa 100644 --- a/crates/aether-data/contracts/src/repository/candidates/types.rs +++ b/crates/aether-data/contracts/src/repository/candidates/types.rs @@ -224,6 +224,36 @@ pub struct StoredRequestCandidate { } impl StoredRequestCandidate { + /// Scheduling needs identity, status, counters and times, without diagnostic payloads. + pub fn runtime_snapshot(&self) -> Self { + Self { + id: self.id.clone(), + request_id: self.request_id.clone(), + user_id: self.user_id.clone(), + api_key_id: self.api_key_id.clone(), + username: None, + api_key_name: None, + candidate_index: self.candidate_index, + retry_index: self.retry_index, + provider_id: self.provider_id.clone(), + endpoint_id: self.endpoint_id.clone(), + key_id: self.key_id.clone(), + status: self.status, + skip_reason: None, + is_cached: self.is_cached, + status_code: self.status_code, + error_type: None, + error_message: None, + latency_ms: self.latency_ms, + concurrent_requests: self.concurrent_requests, + extra_data: None, + required_capabilities: None, + created_at_unix_ms: self.created_at_unix_ms, + started_at_unix_ms: self.started_at_unix_ms, + finished_at_unix_ms: self.finished_at_unix_ms, + } + } + pub fn sanitize_for_persistence(&mut self) { self.username = None; self.api_key_name = None; @@ -691,6 +721,19 @@ pub trait RequestCandidateReadRepository: Send + Sync { limit: usize, ) -> Result, crate::DataLayerError>; + /// Same ordering and limit as `list_recent`, omitting diagnostic fields. + async fn list_recent_runtime( + &self, + limit: usize, + ) -> Result, crate::DataLayerError> { + Ok(self + .list_recent(limit) + .await? + .iter() + .map(StoredRequestCandidate::runtime_snapshot) + .collect()) + } + async fn list_by_provider_id( &self, provider_id: &str, @@ -849,9 +892,7 @@ pub fn sanitize_request_candidate_extra_data_for_persistence( extra_data: Option, ) -> Option { let object = extra_data.as_ref()?.as_object()?; - let mut sanitized = sanitize_request_candidate_extra_data(extra_data.clone()) - .and_then(|value| value.as_object().cloned()) - .unwrap_or_default(); + let mut sanitized = sanitize_candidate_extra_data_object(object); for (key, fields) in [ ("upstream_response", &["headers", "body"][..]), ("error_flow", &["message"][..]), @@ -884,7 +925,10 @@ pub fn sanitize_request_candidate_extra_data_for_persistence( }; let mut summary = sanitized .remove(key) - .and_then(|value| value.as_object().cloned()) + .and_then(|value| match value { + serde_json::Value::Object(object) => Some(object), + _ => None, + }) .unwrap_or_default(); for field in fields { if let Some(value) = diagnostic.get(*field).filter(|value| !value.is_null()) { @@ -907,70 +951,72 @@ pub fn sanitize_request_candidate_extra_data( let serde_json::Value::Object(object) = extra_data? else { return None; }; + let sanitized = sanitize_candidate_extra_data_object(&object); + (!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized)) +} + +fn sanitize_candidate_extra_data_object( + object: &serde_json::Map, +) -> serde_json::Map { let mut sanitized = serde_json::Map::new(); for field in ["gateway_execution_runtime", "stream_completed", "cache_1h"] { - insert_candidate_bool(&object, &mut sanitized, field); + insert_candidate_bool(object, &mut sanitized, field); } for field in ["first_byte_time_ms", "pool_key_index"] { - insert_candidate_u64(&object, &mut sanitized, field); + insert_candidate_u64(object, &mut sanitized, field); } - insert_candidate_i64(&object, &mut sanitized, "priority_slot"); - insert_candidate_u64(&object, &mut sanitized, "ranking_index"); + insert_candidate_i64(object, &mut sanitized, "priority_slot"); + insert_candidate_u64(object, &mut sanitized, "ranking_index"); - insert_candidate_known_string(&object, &mut sanitized, "phase", sanitize_candidate_phase); + insert_candidate_known_string(object, &mut sanitized, "phase", sanitize_candidate_phase); for field in [ "client_api_format", "provider_api_format", "client_contract", "provider_contract", ] { - insert_candidate_known_string( - &object, - &mut sanitized, - field, - sanitize_candidate_api_format, - ); + insert_candidate_known_string(object, &mut sanitized, field, sanitize_candidate_api_format); } insert_candidate_known_string( - &object, + object, &mut sanitized, "execution_strategy", sanitize_candidate_execution_strategy, ); insert_candidate_known_string( - &object, + object, &mut sanitized, "conversion_mode", sanitize_candidate_conversion_mode, ); insert_candidate_known_string( - &object, + object, &mut sanitized, "ranking_mode", sanitize_candidate_ranking_mode, ); insert_candidate_known_string( - &object, + object, &mut sanitized, "priority_mode", sanitize_candidate_priority_mode, ); insert_candidate_known_string( - &object, + object, &mut sanitized, "promoted_by", sanitize_candidate_promotion_reason, ); insert_candidate_known_string( - &object, + object, &mut sanitized, "demoted_by", sanitize_candidate_demotion_reason, ); - insert_candidate_known_string(&object, &mut sanitized, "source", sanitize_candidate_source); + insert_candidate_known_string(object, &mut sanitized, "source", sanitize_candidate_source); insert_candidate_known_string( - &object, + object, &mut sanitized, "execution_path", sanitize_candidate_execution_path, @@ -1030,7 +1076,7 @@ pub fn sanitize_request_candidate_extra_data( sanitized.insert("pool_group_exhaustion".to_string(), exhaustion); } - (!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized)) + sanitized } pub fn sanitize_request_candidate_required_capabilities( diff --git a/crates/aether-data/contracts/src/repository/usage/capture_memory.rs b/crates/aether-data/contracts/src/repository/usage/capture_memory.rs new file mode 100644 index 000000000..f7df6f63a --- /dev/null +++ b/crates/aether-data/contracts/src/repository/usage/capture_memory.rs @@ -0,0 +1,490 @@ +use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; +use std::sync::Arc; + +use serde_json::{Map, Value}; + +use super::types::UsageBodyCaptureState; + +/// Shared accounting for the estimated heap retained by diagnostic JSON bodies. +#[doc(hidden)] +#[derive(Debug)] +pub struct UsageCaptureMemoryBudget { + limit: usize, + retained: AtomicUsize, + downgraded_total: AtomicU64, +} + +impl UsageCaptureMemoryBudget { + pub fn new(limit: usize) -> Self { + Self { + limit, + retained: AtomicUsize::new(0), + downgraded_total: AtomicU64::new(0), + } + } + + fn try_reserve(&self, bytes: usize) -> bool { + if bytes == 0 { + return true; + } + self.retained + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |retained| { + retained + .checked_add(bytes) + .filter(|next| *next <= self.limit) + }) + .is_ok() + } + + fn release(&self, bytes: usize) { + if bytes != 0 { + self.retained.fetch_sub(bytes, Ordering::AcqRel); + } + } + + fn record_downgrade(&self) { + self.downgraded_total.fetch_add(1, Ordering::Relaxed); + } + + pub fn retained_bytes(&self) -> usize { + self.retained.load(Ordering::Acquire) + } + + pub fn downgraded_total(&self) -> u64 { + self.downgraded_total.load(Ordering::Relaxed) + } + + pub fn snapshot(&self) -> (usize, usize, u64) { + (self.limit, self.retained_bytes(), self.downgraded_total()) + } +} + +/// Non-serialized ownership of the diagnostic JSON heap estimate. +#[doc(hidden)] +#[derive(Debug, Default)] +pub struct UsageCaptureRetention { + budget: Option>, + bytes: usize, +} + +// Runtime accounting does not participate in value or wire equality. +impl PartialEq for UsageCaptureRetention { + fn eq(&self, _other: &Self) -> bool { + true + } +} + +impl UsageCaptureRetention { + pub fn reserve(&mut self, budget: Arc, bytes: usize) -> bool { + if self + .budget + .as_ref() + .is_some_and(|current| Arc::ptr_eq(current, &budget)) + { + if bytes > self.bytes && !budget.try_reserve(bytes - self.bytes) { + budget.record_downgrade(); + return false; + } + if bytes < self.bytes { + budget.release(self.bytes - bytes); + } + self.bytes = bytes; + return true; + } + if !budget.try_reserve(bytes) { + budget.record_downgrade(); + return false; + } + *self = Self { + budget: Some(budget), + bytes, + }; + true + } + + pub fn clear(&mut self, budget: Arc) { + *self = Self { + budget: Some(budget), + bytes: 0, + }; + } + + pub fn clone_for_bodies(&self, estimate: impl FnOnce() -> usize) -> (Self, bool) { + let Some(budget) = &self.budget else { + return (Self::default(), true); + }; + let mut retention = Self::default(); + if retention.reserve(Arc::clone(budget), estimate()) { + (retention, true) + } else { + retention.clear(Arc::clone(budget)); + (retention, false) + } + } +} + +impl Drop for UsageCaptureRetention { + fn drop(&mut self) { + if let Some(budget) = &self.budget { + budget.release(self.bytes); + } + } +} + +// serde_json::Map does not expose its backing allocation capacity. This charges a +// conservative per-entry estimate, not an allocator or process RSS measurement. +#[doc(hidden)] +pub fn usage_json_heap_estimate(value: &Value) -> usize { + match value { + Value::String(value) => value.capacity(), + Value::Array(values) => values.iter().fold( + values + .capacity() + .saturating_mul(std::mem::size_of::()), + |bytes, value| bytes.saturating_add(usage_json_heap_estimate(value)), + ), + Value::Object(values) => values.iter().fold( + values.len().saturating_mul( + 4 * (std::mem::size_of::() + + std::mem::size_of::() + + std::mem::size_of::()), + ), + |bytes, (key, value)| { + bytes + .saturating_add(key.capacity()) + .saturating_add(usage_json_heap_estimate(value)) + }, + ), + Value::Null | Value::Bool(_) | Value::Number(_) => 0, + } +} + +/// Marks an omitted diagnostic body using the existing capture metadata shape. +#[doc(hidden)] +pub fn mark_usage_capture_memory_omitted(metadata: &mut Option, key: &str) { + let source_bytes = metadata + .as_ref() + .and_then(|metadata| metadata.get("body_capture")) + .and_then(|capture| capture.get(key)) + .and_then(|entry| entry.get("source_bytes")) + .and_then(Value::as_u64); + let Some(metadata) = metadata + .get_or_insert_with(|| Value::Object(Map::with_capacity(1))) + .as_object_mut() + else { + return; + }; + let Some(body_capture) = metadata + .entry("body_capture".to_owned()) + .or_insert_with(|| Value::Object(Map::with_capacity(1))) + .as_object_mut() + else { + return; + }; + let mut entry = Map::with_capacity(3 + usize::from(source_bytes.is_some())); + entry.insert( + "state".to_owned(), + Value::String(UsageBodyCaptureState::Truncated.as_str().to_owned()), + ); + entry.insert("stored_bytes".to_owned(), Value::from(0)); + if let Some(source_bytes) = source_bytes { + entry.insert("source_bytes".to_owned(), Value::from(source_bytes)); + } + entry.insert( + "reason".to_owned(), + Value::String("usage_event_memory_budget_exceeded".to_owned()), + ); + body_capture.insert(key.to_owned(), Value::Object(entry)); +} + +#[cfg(test)] +mod tests { + use super::super::types::UpsertUsageRecord; + use super::*; + use serde_json::json; + + fn upsert_with_diagnostic_bodies() -> UpsertUsageRecord { + serde_json::from_str( + r#"{ + "request_id": "retained-request", + "provider_name": "openai", + "model": "test-model", + "status": "completed", + "billing_status": "settled", + "updated_at_unix_secs": 123, + "input_tokens": 100, + "output_tokens": 500, + "total_tokens": 600, + "cache_creation_input_tokens": 7, + "cache_creation_ephemeral_5m_input_tokens": 7, + "cache_creation_ephemeral_1h_input_tokens": 0, + "cache_read_input_tokens": 0, + "total_cost_usd": 1.25, + "actual_total_cost_usd": 0.75, + "cache_creation_cost_usd": 0.05, + "cache_read_cost_usd": 0.0, + "status_code": 200, + "error_message": "preserved diagnostic classification", + "request_headers": {"x-request": "original"}, + "provider_request_headers": {"x-provider-request": "original"}, + "response_headers": {"x-response": "original"}, + "client_response_headers": {"x-client-response": "original"}, + "request_body": {"text": "original request body"}, + "provider_request_body": {"text": "original provider request body"}, + "response_body": {"text": "original provider response body"}, + "client_response_body": {"text": "original client response body"}, + "request_body_ref": "usage://retained-request/request", + "provider_request_body_ref": "usage://retained-request/provider_request", + "response_body_ref": "usage://retained-request/response", + "client_response_body_ref": "usage://retained-request/client_response", + "request_body_state": "inline", + "provider_request_body_state": "reference", + "response_body_state": "inline", + "client_response_body_state": "truncated", + "request_metadata": { + "trace_id": "unchanged", + "body_capture": { + "request": {"state": "inline", "source_bytes": 100}, + "provider_request": {"state": "reference", "source_bytes": 200}, + "response": {"state": "inline", "source_bytes": 300}, + "client_response": {"state": "truncated", "source_bytes": 400} + } + } + }"#, + ) + .expect("valid usage write fixture") + } + + fn upsert_body_estimate(record: &UpsertUsageRecord) -> usize { + [ + &record.request_body, + &record.provider_request_body, + &record.response_body, + &record.client_response_body, + ] + .into_iter() + .flatten() + .map(|body| std::mem::size_of::() + usage_json_heap_estimate(body)) + .sum() + } + + #[test] + fn usage_capture_memory_upsert_serde_skips_retention_and_preserves_value_equality() { + let mut source = upsert_with_diagnostic_bodies(); + let weight = upsert_body_estimate(&source); + let budget = Arc::new(UsageCaptureMemoryBudget::new(weight)); + assert!(source + .capture_retention + .reserve(Arc::clone(&budget), weight)); + let mut serialized = serde_json::to_value(&source).unwrap(); + assert!(serialized.get("capture_retention").is_none()); + serialized["capture_retention"] = json!({"bytes": usize::MAX}); + let roundtrip: UpsertUsageRecord = serde_json::from_value(serialized).unwrap(); + assert_eq!(source, roundtrip); + let unmanaged_clone = roundtrip.clone(); + assert_eq!(source, unmanaged_clone); + assert!(unmanaged_clone.request_body.is_some()); + assert!(unmanaged_clone.provider_request_body.is_some()); + assert!(unmanaged_clone.response_body.is_some()); + assert!(unmanaged_clone.client_response_body.is_some()); + assert_eq!(budget.retained_bytes(), weight); + assert_eq!(budget.downgraded_total(), 0); + drop((roundtrip, unmanaged_clone)); + assert_eq!(budget.retained_bytes(), weight); + drop(source); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn usage_capture_memory_upsert_clone_reserves_for_all_four_deep_copies() { + let mut source = upsert_with_diagnostic_bodies(); + let original = serde_json::to_value(&source).unwrap(); + let weight = upsert_body_estimate(&source); + let budget = Arc::new(UsageCaptureMemoryBudget::new(weight * 2)); + assert!(source + .capture_retention + .reserve(Arc::clone(&budget), weight)); + let cloned = source.clone(); + assert_eq!(source, cloned); + assert_eq!(budget.retained_bytes(), weight * 2); + assert_eq!(budget.downgraded_total(), 0); + for (source_body, cloned_body) in [ + (&source.request_body, &cloned.request_body), + (&source.provider_request_body, &cloned.provider_request_body), + (&source.response_body, &cloned.response_body), + (&source.client_response_body, &cloned.client_response_body), + ] { + let source_text = source_body.as_ref().unwrap()["text"].as_str().unwrap(); + let cloned_text = cloned_body.as_ref().unwrap()["text"].as_str().unwrap(); + assert_eq!(source_text, cloned_text); + assert_ne!(source_text.as_ptr(), cloned_text.as_ptr()); + } + assert_eq!(serde_json::to_value(&source).unwrap(), original); + drop(source); + assert_eq!(budget.retained_bytes(), weight); + assert_eq!(serde_json::to_value(&cloned).unwrap(), original); + drop(cloned); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn usage_capture_memory_upsert_clone_over_budget_only_omits_four_bodies() { + let mut source = upsert_with_diagnostic_bodies(); + let original = serde_json::to_value(&source).unwrap(); + let weight = upsert_body_estimate(&source); + let budget = Arc::new(UsageCaptureMemoryBudget::new(weight)); + assert!(source + .capture_retention + .reserve(Arc::clone(&budget), weight)); + let cloned = source.clone(); + assert_eq!(budget.retained_bytes(), weight); + assert_eq!(budget.downgraded_total(), 1); + assert_eq!(serde_json::to_value(&source).unwrap(), original); + + let mut expected = original; + for (body, state, key, source_bytes) in [ + ("request_body", "request_body_state", "request", 100), + ( + "provider_request_body", + "provider_request_body_state", + "provider_request", + 200, + ), + ("response_body", "response_body_state", "response", 300), + ( + "client_response_body", + "client_response_body_state", + "client_response", + 400, + ), + ] { + expected[body] = Value::Null; + expected[state] = json!("truncated"); + expected["request_metadata"]["body_capture"][key] = json!({ + "state": "truncated", + "stored_bytes": 0, + "source_bytes": source_bytes, + "reason": "usage_event_memory_budget_exceeded" + }); + } + assert_eq!(serde_json::to_value(&cloned).unwrap(), expected); + assert_eq!(cloned.output_tokens, Some(500)); + assert_eq!(cloned.cache_read_input_tokens, Some(0)); + assert_eq!(cloned.cache_creation_ephemeral_1h_input_tokens, Some(0)); + assert_eq!(cloned.total_cost_usd, Some(1.25)); + assert_eq!(cloned.actual_total_cost_usd, Some(0.75)); + drop(cloned); + assert_eq!(budget.retained_bytes(), weight); + drop(source); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn usage_capture_memory_upsert_clone_preserves_explicit_clearing_states() { + for state in [ + UsageBodyCaptureState::None, + UsageBodyCaptureState::Disabled, + UsageBodyCaptureState::Unavailable, + ] { + let mut source = upsert_with_diagnostic_bodies(); + source.request_body_state = Some(state); + source.provider_request_body_state = Some(state); + source.response_body_state = Some(state); + source.client_response_body_state = Some(state); + let original = serde_json::to_value(&source).unwrap(); + let weight = upsert_body_estimate(&source); + let budget = Arc::new(UsageCaptureMemoryBudget::new(weight)); + assert!(source + .capture_retention + .reserve(Arc::clone(&budget), weight)); + let cloned = source.clone(); + let mut expected = original.clone(); + for body in [ + "request_body", + "provider_request_body", + "response_body", + "client_response_body", + ] { + expected[body] = Value::Null; + } + assert_eq!(serde_json::to_value(&cloned).unwrap(), expected); + assert_eq!(serde_json::to_value(&source).unwrap(), original); + assert_eq!(budget.retained_bytes(), weight); + assert_eq!(budget.downgraded_total(), 1); + drop(cloned); + assert_eq!(budget.retained_bytes(), weight); + drop(source); + assert_eq!(budget.retained_bytes(), 0); + } + } + + #[test] + fn usage_capture_memory_metadata_preserves_source_bytes_and_unrelated_metadata() { + let mut metadata = Some(json!({ + "trace_id": "unchanged", + "body_capture": { + "request": {"state": "complete", "source_bytes": 42, "stored_bytes": 42, "extra": true}, + "response": {"state": "complete", "source_bytes": 7} + } + })); + mark_usage_capture_memory_omitted(&mut metadata, "request"); + assert_eq!( + metadata, + Some(json!({ + "trace_id": "unchanged", + "body_capture": { + "request": { + "state": "truncated", + "source_bytes": 42, + "stored_bytes": 0, + "reason": "usage_event_memory_budget_exceeded" + }, + "response": {"state": "complete", "source_bytes": 7} + } + })) + ); + } + + #[test] + fn usage_capture_memory_metadata_creates_missing_objects_and_replaces_entries() { + for mut metadata in [ + None, + Some(json!({})), + Some(json!({"body_capture": {}})), + Some(json!({"body_capture": {"request": null}})), + Some(json!({"body_capture": {"request": "legacy"}})), + Some(json!({"body_capture": {"request": {"source_bytes": "42"}}})), + Some(json!({"body_capture": {"request": {"source_bytes": -1}}})), + ] { + mark_usage_capture_memory_omitted(&mut metadata, "request"); + assert_eq!( + metadata, + Some(json!({"body_capture": {"request": { + "state": "truncated", + "stored_bytes": 0, + "reason": "usage_event_memory_budget_exceeded" + }}})) + ); + } + } + + #[test] + fn usage_capture_memory_metadata_preserves_existing_non_object_containers() { + for metadata in [ + Value::Null, + json!(false), + json!(7), + json!("legacy"), + json!([]), + json!({"body_capture": null}), + json!({"body_capture": false}), + json!({"body_capture": 7}), + json!({"body_capture": "legacy"}), + json!({"body_capture": []}), + ] { + let mut actual = Some(metadata.clone()); + mark_usage_capture_memory_omitted(&mut actual, "request"); + assert_eq!(actual, Some(metadata)); + } + } +} diff --git a/crates/aether-data/contracts/src/repository/usage/mod.rs b/crates/aether-data/contracts/src/repository/usage/mod.rs index d0b19c6fc..ed2377a3a 100644 --- a/crates/aether-data/contracts/src/repository/usage/mod.rs +++ b/crates/aether-data/contracts/src/repository/usage/mod.rs @@ -1,8 +1,14 @@ +mod capture_memory; mod compression; mod metadata_policy; mod policy; mod types; +#[doc(hidden)] +pub use capture_memory::{ + mark_usage_capture_memory_omitted, usage_json_heap_estimate, UsageCaptureMemoryBudget, + UsageCaptureRetention, +}; pub use compression::{read_decompressed_usage_json, MAX_DECOMPRESSED_USAGE_JSON_BYTES}; pub use metadata_policy::*; pub use policy::*; diff --git a/crates/aether-data/contracts/src/repository/usage/policy.rs b/crates/aether-data/contracts/src/repository/usage/policy.rs index 6bd29e166..d5a4eadab 100644 --- a/crates/aether-data/contracts/src/repository/usage/policy.rs +++ b/crates/aether-data/contracts/src/repository/usage/policy.rs @@ -673,6 +673,7 @@ mod tests { fn usage_with_http_capture() -> UpsertUsageRecord { UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-sensitive-capture".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("key-1".to_string()), diff --git a/crates/aether-data/contracts/src/repository/usage/types.rs b/crates/aether-data/contracts/src/repository/usage/types.rs index de4d0f933..a2b88639a 100644 --- a/crates/aether-data/contracts/src/repository/usage/types.rs +++ b/crates/aether-data/contracts/src/repository/usage/types.rs @@ -54,11 +54,11 @@ pub fn extract_provider_reasoning_effort_from_body(value: Option<&Value>) -> Opt } fn normalize_provider_reasoning_effort(value: &str) -> Option { - let normalized = value.trim().to_ascii_lowercase(); - if normalized.is_empty() || normalized.len() > 64 { + let value = value.trim(); + if value.is_empty() || value.len() > 64 { return None; } - Some(normalized) + Some(value.to_ascii_lowercase()) } pub fn extract_provider_service_tier_from_body(value: Option<&Value>) -> Option { @@ -112,11 +112,11 @@ pub fn extract_provider_actual_service_tier_from_response(value: Option<&Value>) } pub fn normalize_provider_service_tier(value: &str) -> Option { - let normalized = value.trim().to_ascii_lowercase(); - if normalized.is_empty() || normalized.len() > 64 { + let value = value.trim(); + if value.is_empty() || value.len() > 64 { return None; } - Some(normalized) + Some(value.to_ascii_lowercase()) } /// Resolves a provider processing tier exclusively from the final upstream request. @@ -1980,7 +1980,7 @@ pub trait UsageReadRepository: Send + Sync { /// Request/response headers and bodies here are capture inputs that the repository persists into /// the dedicated HTTP audit/body stores. Deprecated mirror columns on `public.usage` remain in the /// schema for compatibility only and are not the intended long-term destination for new writes. -#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] +#[derive(Debug, PartialEq, serde::Serialize, serde::Deserialize)] pub struct UpsertUsageRecord { pub request_id: String, pub user_id: Option, @@ -2055,6 +2055,142 @@ pub struct UpsertUsageRecord { pub finalized_at_unix_secs: Option, pub created_at_unix_ms: Option, pub updated_at_unix_secs: u64, + #[doc(hidden)] + #[serde(skip)] + pub capture_retention: super::UsageCaptureRetention, +} + +impl Clone for UpsertUsageRecord { + fn clone(&self) -> Self { + let (capture_retention, retain_bodies) = self.capture_retention.clone_for_bodies(|| { + [ + &self.request_body, + &self.provider_request_body, + &self.response_body, + &self.client_response_body, + ] + .into_iter() + .flatten() + .fold(0usize, |bytes, body| { + bytes + .saturating_add(std::mem::size_of::()) + .saturating_add(super::usage_json_heap_estimate(body)) + }) + }); + let mut cloned = Self { + request_id: self.request_id.clone(), + user_id: self.user_id.clone(), + api_key_id: self.api_key_id.clone(), + username: self.username.clone(), + api_key_name: self.api_key_name.clone(), + provider_name: self.provider_name.clone(), + model: self.model.clone(), + target_model: self.target_model.clone(), + provider_id: self.provider_id.clone(), + provider_endpoint_id: self.provider_endpoint_id.clone(), + provider_api_key_id: self.provider_api_key_id.clone(), + request_type: self.request_type.clone(), + api_format: self.api_format.clone(), + api_family: self.api_family.clone(), + endpoint_kind: self.endpoint_kind.clone(), + endpoint_api_format: self.endpoint_api_format.clone(), + provider_api_family: self.provider_api_family.clone(), + provider_endpoint_kind: self.provider_endpoint_kind.clone(), + has_format_conversion: self.has_format_conversion, + is_stream: self.is_stream, + input_tokens: self.input_tokens, + output_tokens: self.output_tokens, + total_tokens: self.total_tokens, + cache_creation_input_tokens: self.cache_creation_input_tokens, + cache_creation_ephemeral_5m_input_tokens: self.cache_creation_ephemeral_5m_input_tokens, + cache_creation_ephemeral_1h_input_tokens: self.cache_creation_ephemeral_1h_input_tokens, + cache_read_input_tokens: self.cache_read_input_tokens, + cache_creation_cost_usd: self.cache_creation_cost_usd, + cache_read_cost_usd: self.cache_read_cost_usd, + output_price_per_1m: self.output_price_per_1m, + total_cost_usd: self.total_cost_usd, + actual_total_cost_usd: self.actual_total_cost_usd, + status_code: self.status_code, + error_message: self.error_message.clone(), + error_category: self.error_category.clone(), + response_time_ms: self.response_time_ms, + first_byte_time_ms: self.first_byte_time_ms, + status: self.status.clone(), + billing_status: self.billing_status.clone(), + request_headers: self.request_headers.clone(), + request_body: retain_bodies.then(|| self.request_body.clone()).flatten(), + request_body_ref: self.request_body_ref.clone(), + request_body_state: self.request_body_state, + provider_request_headers: self.provider_request_headers.clone(), + provider_request_body: retain_bodies + .then(|| self.provider_request_body.clone()) + .flatten(), + provider_request_body_ref: self.provider_request_body_ref.clone(), + provider_request_body_state: self.provider_request_body_state, + response_headers: self.response_headers.clone(), + response_body: retain_bodies.then(|| self.response_body.clone()).flatten(), + response_body_ref: self.response_body_ref.clone(), + response_body_state: self.response_body_state, + client_response_headers: self.client_response_headers.clone(), + client_response_body: retain_bodies + .then(|| self.client_response_body.clone()) + .flatten(), + client_response_body_ref: self.client_response_body_ref.clone(), + client_response_body_state: self.client_response_body_state, + candidate_id: self.candidate_id.clone(), + candidate_index: self.candidate_index, + key_name: self.key_name.clone(), + planner_kind: self.planner_kind.clone(), + route_family: self.route_family.clone(), + route_kind: self.route_kind.clone(), + execution_path: self.execution_path.clone(), + local_execution_runtime_miss_reason: self.local_execution_runtime_miss_reason.clone(), + request_metadata: self.request_metadata.clone(), + finalized_at_unix_secs: self.finalized_at_unix_secs, + created_at_unix_ms: self.created_at_unix_ms, + updated_at_unix_secs: self.updated_at_unix_secs, + capture_retention, + }; + if !retain_bodies { + for (present, key, state) in [ + ( + self.request_body.is_some(), + "request", + &mut cloned.request_body_state, + ), + ( + self.provider_request_body.is_some(), + "provider_request", + &mut cloned.provider_request_body_state, + ), + ( + self.response_body.is_some(), + "response", + &mut cloned.response_body_state, + ), + ( + self.client_response_body.is_some(), + "client_response", + &mut cloned.client_response_body_state, + ), + ] { + if present + && !matches!( + state, + Some( + UsageBodyCaptureState::None + | UsageBodyCaptureState::Disabled + | UsageBodyCaptureState::Unavailable + ) + ) + { + *state = Some(UsageBodyCaptureState::Truncated); + super::mark_usage_capture_memory_omitted(&mut cloned.request_metadata, key); + } + } + } + cloned + } } impl UpsertUsageRecord { @@ -2501,15 +2637,59 @@ fn parse_timestamp(value: i64, field_name: &str) -> Result Option, + normalize_provider_service_tier, + ] { + assert_eq!(normalize(" \t\r\n"), None); + assert_eq!(normalize(" HIGH\n"), Some("high".to_string())); + assert_eq!(normalize(&"A".repeat(64)), Some("a".repeat(64))); + assert_eq!( + normalize(&format!(" \t{}\n", "A".repeat(64))), + Some("a".repeat(64)) + ); + assert_eq!(normalize(&"A".repeat(65)), None); + } + } + + #[test] + fn provider_fact_normalization_preserves_non_ascii_case_and_byte_count() { + for normalize in [ + normalize_provider_reasoning_effort as fn(&str) -> Option, + normalize_provider_service_tier, + ] { + let accepted = format!("{}A", "\u{00c9}".repeat(31)); + assert_eq!( + normalize(&accepted), + Some(format!("{}a", "\u{00c9}".repeat(31))) + ); + assert_eq!( + normalize(&"\u{00c9}".repeat(32)), + Some("\u{00c9}".repeat(32)) + ); + assert_eq!(normalize(&format!("{}A", "\u{00c9}".repeat(32))), None); + assert_eq!(normalize("\u{2003}FAST\u{2003}"), Some("fast".to_string())); + } + } + + #[test] + fn provider_fact_normalization_rejects_large_input_before_copying() { + let oversized = "A".repeat(4 * 1024 * 1024); + assert_eq!(normalize_provider_reasoning_effort(&oversized), None); + assert_eq!(normalize_provider_service_tier(&oversized), None); + } + fn sample_usage() -> StoredRequestUsageAudit { StoredRequestUsageAudit::new( "usage-1".to_string(), @@ -2700,6 +2880,7 @@ mod tests { #[test] fn rejects_invalid_upsert_payload() { let mut record = UpsertUsageRecord { + capture_retention: Default::default(), request_id: "".to_string(), user_id: None, api_key_id: None, diff --git a/crates/aether-data/runtime/src/backend/maintenance/postgres.rs b/crates/aether-data/runtime/src/backend/maintenance/postgres.rs index 3b8d93bed..ee48f3e3b 100644 --- a/crates/aether-data/runtime/src/backend/maintenance/postgres.rs +++ b/crates/aether-data/runtime/src/backend/maintenance/postgres.rs @@ -16,17 +16,36 @@ impl PostgresBackend { table_names: &[&str], ) -> Result { let mut summary = DatabaseMaintenanceSummary::default(); + if table_names.is_empty() { + return Ok(summary); + } + // VACUUM cannot run inside a transaction. Discard this connection on every + // exit path so its longer session deadlines never leak into request queries. + let mut conn = self.pool().acquire().await.map_postgres_err()?; + conn.close_on_drop(); + sqlx::query("SET statement_timeout = '5min'") + .execute(&mut *conn) + .await + .map_postgres_err()?; + sqlx::query("SET lock_timeout = '30s'") + .execute(&mut *conn) + .await + .map_postgres_err()?; for table_name in table_names { let table_name = maintenance_identifier(table_name)?; summary.attempted += 1; let statement = format!("VACUUM ANALYZE \"{table_name}\""); - if sqlx::raw_sql(&statement) - .execute(self.pool()) + match sqlx::query(&statement) + .execute(&mut *conn) .await .map_postgres_err() - .is_ok() { - summary.succeeded += 1; + Ok(_) => summary.succeeded += 1, + Err(error) => tracing::warn!( + table_name, + error = %error, + "PostgreSQL table maintenance failed" + ), } } Ok(summary) diff --git a/crates/aether-data/runtime/src/backend/postgres.rs b/crates/aether-data/runtime/src/backend/postgres.rs index e9b52932a..3eb3b4aa4 100644 --- a/crates/aether-data/runtime/src/backend/postgres.rs +++ b/crates/aether-data/runtime/src/backend/postgres.rs @@ -274,6 +274,40 @@ mod tests { use super::PostgresBackend; use crate::driver::postgres::{PostgresLeaseRunnerConfig, PostgresPoolConfig}; + #[tokio::test] + async fn maintenance_and_aggregation_futures_are_send() { + fn assert_send(_: impl Send) {} + + let backend = PostgresBackend::from_config(PostgresPoolConfig { + database_url: "postgres://localhost/aether".to_string(), + min_connections: 0, + ..PostgresPoolConfig::default() + }) + .unwrap(); + let now = chrono::Utc::now(); + let daily = crate::StatsDailyAggregationInput { + target_day_utc: now, + aggregated_at: now, + }; + let hourly = crate::StatsHourlyAggregationInput { + target_hour_utc: now, + aggregated_at: now, + }; + let wallet = crate::WalletDailyUsageAggregationInput { + billing_date: "2026-09-09".to_string(), + billing_timezone: "UTC".to_string(), + window_start_unix_secs: 0, + window_end_unix_secs: 86_400, + aggregated_at_unix_secs: 86_400, + }; + + // Drop without polling: these are compile-time checks for spawned workers. + assert_send(backend.run_table_maintenance(&["usage"])); + assert_send(backend.aggregate_stats_daily(&daily)); + assert_send(backend.aggregate_stats_hourly(&hourly)); + assert_send(backend.aggregate_wallet_daily_usage(&wallet)); + } + #[tokio::test] async fn backend_retains_config_and_pool() { let config = PostgresPoolConfig { diff --git a/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs b/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs index 605adfc79..479c20582 100644 --- a/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs +++ b/crates/aether-data/runtime/src/backend/stats/postgres_daily/mod.rs @@ -66,6 +66,12 @@ async fn perform_stats_aggregation_for_day( ) -> Result { let day_end_utc = day_start_utc + chrono::Duration::days(1); let mut tx = pool.begin().await?; + sqlx::query("SET LOCAL statement_timeout = '5min'") + .execute(&mut *tx) + .await?; + sqlx::query("SET LOCAL lock_timeout = '30s'") + .execute(&mut *tx) + .await?; let aggregate_row = sqlx::query(SELECT_STATS_DAILY_AGGREGATE_SQL) .bind(day_start_utc) .bind(day_end_utc) diff --git a/crates/aether-data/runtime/src/backend/stats/postgres_hourly/mod.rs b/crates/aether-data/runtime/src/backend/stats/postgres_hourly/mod.rs index f921eb14a..b81c9d081 100644 --- a/crates/aether-data/runtime/src/backend/stats/postgres_hourly/mod.rs +++ b/crates/aether-data/runtime/src/backend/stats/postgres_hourly/mod.rs @@ -65,6 +65,12 @@ async fn perform_stats_hourly_aggregation_for_hour( ) -> Result { let hour_end = hour_utc + chrono::Duration::hours(1); let mut tx = pool.begin().await?; + sqlx::query("SET LOCAL statement_timeout = '5min'") + .execute(&mut *tx) + .await?; + sqlx::query("SET LOCAL lock_timeout = '30s'") + .execute(&mut *tx) + .await?; let row = sqlx::query(SELECT_STATS_HOURLY_AGGREGATE_SQL) .bind(hour_utc) diff --git a/crates/aether-data/runtime/src/backend/wallet/postgres.rs b/crates/aether-data/runtime/src/backend/wallet/postgres.rs index 816a466fb..3c70980d9 100644 --- a/crates/aether-data/runtime/src/backend/wallet/postgres.rs +++ b/crates/aether-data/runtime/src/backend/wallet/postgres.rs @@ -104,6 +104,14 @@ impl PostgresBackend { let window_end = unix_secs_to_utc(input.window_end_unix_secs, "window_end")?; let aggregated_at = unix_secs_to_utc(input.aggregated_at_unix_secs, "aggregated_at")?; let mut tx = self.pool().begin().await.map_postgres_err()?; + sqlx::query("SET LOCAL statement_timeout = '5min'") + .execute(&mut *tx) + .await + .map_postgres_err()?; + sqlx::query("SET LOCAL lock_timeout = '30s'") + .execute(&mut *tx) + .await + .map_postgres_err()?; let aggregated_wallets = sqlx::query(UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL) .bind(window_start) diff --git a/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs b/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs index a9bad0950..fa01331f4 100644 --- a/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs +++ b/crates/aether-data/runtime/src/lifecycle/backfill/postgres.rs @@ -57,7 +57,7 @@ pub(super) struct AppliedBackfill { } pub async fn run_backfills(pool: &PgPool) -> Result<(), MigrateError> { - let mut conn = pool.acquire().await?; + let mut conn = aether_data_postgres::acquire_postgres_migration_connection(pool).await?; if BACKFILL_MIGRATOR.locking { conn.lock().await?; diff --git a/crates/aether-data/runtime/src/repository/candidates/memory.rs b/crates/aether-data/runtime/src/repository/candidates/memory.rs index a293bb742..8d7487a89 100644 --- a/crates/aether-data/runtime/src/repository/candidates/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidates/memory.rs @@ -45,7 +45,54 @@ fn merge_extra_data( #[derive(Debug, Default)] pub struct InMemoryRequestCandidateRepository { - by_id: RwLock>, + rows: RwLock, +} + +#[derive(Debug, Default)] +struct CandidateRows { + by_id: BTreeMap, + by_request: BTreeMap>, + by_created: BTreeSet<(std::cmp::Reverse, String)>, +} + +impl CandidateRows { + fn remove(&mut self, id: &str) -> Option { + let row = self.by_id.remove(id)?; + self.by_created + .remove(&(std::cmp::Reverse(row.created_at_unix_ms), row.id.clone())); + if let Some(ids) = self.by_request.get_mut(&row.request_id) { + ids.remove(id); + if ids.is_empty() { + self.by_request.remove(&row.request_id); + } + } + Some(row) + } + + fn insert(&mut self, row: StoredRequestCandidate) -> &StoredRequestCandidate { + // Keep all indexes behind one lock and sanitize every insertion. Reads + // can clone these records without rebuilding their diagnostic JSON. + let row = sanitize_stored_candidate(row); + self.remove(&row.id); + self.by_request + .entry(row.request_id.clone()) + .or_default() + .insert(row.id.clone()); + self.by_created + .insert((std::cmp::Reverse(row.created_at_unix_ms), row.id.clone())); + self.by_id + .entry(row.id.clone()) + .insert_entry(row) + .into_mut() + } + + fn for_request(&self, request_id: &str) -> impl Iterator { + self.by_request + .get(request_id) + .into_iter() + .flatten() + .filter_map(|id| self.by_id.get(id)) + } } impl InMemoryRequestCandidateRepository { @@ -53,12 +100,12 @@ impl InMemoryRequestCandidateRepository { where I: IntoIterator, { - let mut by_id = BTreeMap::new(); - for item in items.into_iter().map(sanitize_stored_candidate) { - by_id.insert(item.id.clone(), item); + let mut rows = CandidateRows::default(); + for item in items { + rows.insert(item); } Self { - by_id: RwLock::new(by_id), + rows: RwLock::new(rows), } } } @@ -70,13 +117,11 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { request_id: &str, ) -> Result, DataLayerError> { let mut rows = self - .by_id + .rows .read() .expect("request candidate repository lock") - .values() - .filter(|row| row.request_id == request_id) + .for_request(request_id) .cloned() - .map(sanitize_stored_candidate) .collect::>(); rows.sort_by(|left, right| { left.candidate_index @@ -95,17 +140,14 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { return Ok(Vec::new()); } - let mut rows = self - .by_id - .read() - .expect("request candidate repository lock") - .values() + let rows = self.rows.read().expect("request candidate repository lock"); + Ok(rows + .by_created + .iter() + .take(limit) + .filter_map(|(_, id)| rows.by_id.get(id)) .cloned() - .map(sanitize_stored_candidate) - .collect::>(); - rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms)); - rows.truncate(limit); - Ok(rows) + .collect()) } async fn list_by_provider_id( @@ -118,19 +160,33 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { } let mut rows = self - .by_id + .rows .read() .expect("request candidate repository lock") + .by_id .values() .filter(|row| row.provider_id.as_deref() == Some(provider_id)) .cloned() - .map(sanitize_stored_candidate) .collect::>(); rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms)); rows.truncate(limit); Ok(rows) } + async fn list_recent_runtime( + &self, + limit: usize, + ) -> Result, DataLayerError> { + let rows = self.rows.read().expect("request candidate repository lock"); + Ok(rows + .by_created + .iter() + .take(limit) + .filter_map(|(_, id)| rows.by_id.get(id)) + .map(StoredRequestCandidate::runtime_snapshot) + .collect()) + } + async fn list_finalized_by_endpoint_ids_since( &self, endpoint_ids: &[String], @@ -143,9 +199,10 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { let endpoint_ids = endpoint_ids.iter().cloned().collect::>(); let mut rows = self - .by_id + .rows .read() .expect("request candidate repository lock") + .by_id .values() .filter(|row| { row.endpoint_id @@ -160,7 +217,6 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { ) }) .cloned() - .map(sanitize_stored_candidate) .collect::>(); rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms)); rows.truncate(limit); @@ -179,9 +235,10 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { let endpoint_ids = endpoint_ids.iter().cloned().collect::>(); let mut counts = BTreeMap::<(String, &'static str), u64>::new(); for row in self - .by_id + .rows .read() .expect("request candidate repository lock") + .by_id .values() { let Some(endpoint_id) = row.endpoint_id.as_ref() else { @@ -243,9 +300,10 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository { let mut buckets = BTreeMap::<(String, u32), PublicHealthTimelineBucket>::new(); for row in self - .by_id + .rows .read() .expect("request candidate repository lock") + .by_id .values() { let Some(endpoint_id) = row.endpoint_id.as_ref() else { @@ -319,19 +377,17 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { candidate.sanitize_for_persistence(); candidate.validate()?; - let mut by_id = self - .by_id + let mut rows = self + .rows .write() .expect("request candidate repository lock"); - let existing = by_id - .values() + let existing = rows + .for_request(&candidate.request_id) .find(|row| { - row.request_id == candidate.request_id - && row.candidate_index == candidate.candidate_index + row.candidate_index == candidate.candidate_index && row.retry_index == candidate.retry_index }) - .cloned() - .map(sanitize_stored_candidate); + .cloned(); let preserve_existing_lifecycle = existing.as_ref().is_some_and(|row| { request_candidate_lifecycle_would_regress(row.status, candidate.status) @@ -449,10 +505,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { .or_else(|| existing.as_ref().and_then(|row| row.finished_at_unix_ms)) }, }; - let stored = sanitize_stored_candidate(stored); - - by_id.insert(stored.id.clone(), stored.clone()); - Ok(stored) + Ok(rows.insert(stored).clone()) } async fn delete_created_before( @@ -464,11 +517,12 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { return Ok(0); } - let mut by_id = self - .by_id + let mut rows = self + .rows .write() .expect("request candidate repository lock"); - let mut ids = by_id + let mut ids = rows + .by_id .values() .filter(|row| row.created_at_unix_ms < created_before_unix_secs * 1000) .map(|row| (row.created_at_unix_ms, row.id.clone())) @@ -477,7 +531,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository { let mut deleted = 0usize; for (_, id) in ids.into_iter().take(limit) { - if by_id.remove(&id).is_some() { + if rows.remove(&id).is_some() { deleted += 1; } } @@ -546,6 +600,79 @@ mod tests { assert_eq!(rows[1].request_id, "req-1"); } + #[tokio::test] + async fn request_index_tracks_replaced_ids_and_removes_empty_requests() { + let repository = InMemoryRequestCandidateRepository::seed([ + sample_candidate("same-id", "old-request", 100), + sample_candidate("same-id", "new-request", 200), + sample_candidate("other-id", "new-request", 300), + ]); + assert!(repository + .list_by_request_id("old-request") + .await + .unwrap() + .is_empty()); + assert_eq!( + repository + .list_by_request_id("new-request") + .await + .unwrap() + .len(), + 2 + ); + assert_eq!(repository.delete_created_before(1, 1).await.unwrap(), 1); + let remaining = repository.list_by_request_id("new-request").await.unwrap(); + assert_eq!(remaining.len(), 1); + assert_eq!(remaining[0].id, "other-id"); + assert_eq!(repository.delete_created_before(1, 1).await.unwrap(), 1); + let rows = repository.rows.read().unwrap(); + assert!(rows.by_id.is_empty()); + assert!(rows.by_request.is_empty()); + assert!(rows.by_created.is_empty()); + } + + #[tokio::test] + async fn recent_index_preserves_equal_timestamp_order_and_replaced_dates() { + let repository = InMemoryRequestCandidateRepository::seed([ + sample_candidate("b", "req-b", 400), + sample_candidate("a", "req-a", 200), + sample_candidate("c", "req-c", 200), + sample_candidate("b", "req-b", 100), + ]); + let recent = repository.list_recent(2).await.unwrap(); + assert_eq!( + recent.iter().map(|row| row.id.as_str()).collect::>(), + ["a", "c"] + ); + assert_eq!(repository.list_recent(0).await.unwrap().len(), 0); + assert_eq!(repository.rows.read().unwrap().by_created.len(), 3); + } + + #[tokio::test] + async fn runtime_reads_keep_metadata_without_diagnostic_payloads() { + let mut candidate = sample_candidate("candidate", "request", 100); + candidate.extra_data = Some(json!({"upstream_response": {"body": "x".repeat(32_768)}})); + candidate.error_message = Some("diagnostic detail".into()); + candidate.required_capabilities = Some(json!({"vision": true})); + candidate.concurrent_requests = Some(17); + let repository = InMemoryRequestCandidateRepository::seed([candidate]); + let full = repository.list_recent(1).await.unwrap(); + let runtime = repository.list_recent_runtime(1).await.unwrap(); + assert_eq!( + runtime, + full.iter() + .map(StoredRequestCandidate::runtime_snapshot) + .collect::>() + ); + assert_eq!(runtime[0].concurrent_requests, Some(17)); + assert!(runtime[0].extra_data.is_none()); + assert!(runtime[0].error_message.is_none()); + assert!(runtime[0].required_capabilities.is_none()); + assert!(full[0].extra_data.is_some()); + assert_eq!(repository.list_recent(1).await.unwrap(), full); + assert!(repository.list_recent_runtime(0).await.unwrap().is_empty()); + } + #[tokio::test] async fn lists_recent_request_candidates_in_descending_created_order() { let repository = InMemoryRequestCandidateRepository::seed(vec![ @@ -602,10 +729,11 @@ mod tests { { let stored = repository - .by_id + .rows .read() .expect("request candidate repository lock"); let candidate = stored + .by_id .get("cand-raw") .expect("seeded candidate should exist"); assert_eq!( @@ -628,10 +756,10 @@ mod tests { bypassed_candidate.id = "cand-bypassed".to_string(); bypassed_candidate.request_id = "req-bypassed".to_string(); repository - .by_id + .rows .write() .expect("request candidate repository lock") - .insert(bypassed_candidate.id.clone(), bypassed_candidate); + .insert(bypassed_candidate); let rows = repository .list_recent(10) diff --git a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs index ab0114bc7..17bd0a7e9 100644 --- a/crates/aether-data/runtime/src/repository/usage/memory/tests.rs +++ b/crates/aether-data/runtime/src/repository/usage/memory/tests.rs @@ -66,6 +66,7 @@ fn sample_usage(request_id: &str, created_at_unix_ms: i64) -> StoredRequestUsage fn sample_upsert_usage_record(request_id: &str) -> UpsertUsageRecord { UpsertUsageRecord { + capture_retention: Default::default(), request_id: request_id.to_string(), user_id: None, api_key_id: None, @@ -535,6 +536,7 @@ async fn stale_pending_update_does_not_regress_finalized_usage() { let repository = InMemoryUsageReadRepository::default(); repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-finalized-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("api-key-1".to_string()), @@ -608,6 +610,7 @@ async fn stale_pending_update_does_not_regress_finalized_usage() { repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-finalized-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("api-key-1".to_string()), @@ -695,6 +698,7 @@ async fn upsert_allows_completed_recovery_after_void_failure() { let repository = InMemoryUsageReadRepository::default(); repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-recover-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("api-key-1".to_string()), @@ -768,6 +772,7 @@ async fn upsert_allows_completed_recovery_after_void_failure() { repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-recover-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("api-key-1".to_string()), @@ -1009,6 +1014,7 @@ async fn stale_pending_update_does_not_regress_streaming_usage() { let repository = InMemoryUsageReadRepository::default(); repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-streaming-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("api-key-1".to_string()), @@ -1084,6 +1090,7 @@ async fn stale_pending_update_does_not_regress_streaming_usage() { repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-streaming-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("api-key-1".to_string()), @@ -1339,6 +1346,7 @@ async fn upsert_writes_usage_record() { let repository = InMemoryUsageReadRepository::default(); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-upsert-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("key-1".to_string()), @@ -1430,6 +1438,7 @@ async fn upsert_defaults_created_at_to_second_timestamp() { let repository = InMemoryUsageReadRepository::default(); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-upsert-ms-default".to_string(), user_id: None, api_key_id: None, @@ -1509,6 +1518,7 @@ async fn upsert_does_not_backfill_legacy_output_price_from_request_metadata() { let repository = InMemoryUsageReadRepository::default(); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-upsert-price-metadata".to_string(), user_id: None, api_key_id: None, @@ -1591,6 +1601,7 @@ async fn upsert_does_not_backfill_typed_body_refs_from_request_metadata() { let repository = InMemoryUsageReadRepository::default(); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-upsert-body-ref-metadata".to_string(), user_id: None, api_key_id: None, @@ -1673,6 +1684,7 @@ async fn upsert_keeps_typed_routing_fields_out_of_request_metadata() { let repository = InMemoryUsageReadRepository::default(); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-upsert-routing-metadata".to_string(), user_id: None, api_key_id: None, @@ -1771,6 +1783,7 @@ async fn upsert_does_not_persist_legacy_display_columns_for_new_rows() { let repository = InMemoryUsageReadRepository::default(); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-upsert-display-columns".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("key-1".to_string()), @@ -1921,6 +1934,7 @@ async fn upsert_preserves_existing_legacy_display_columns_when_new_write_omits_t }]); let stored = repository .upsert(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-existing-display-columns".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("key-1".to_string()), diff --git a/crates/aether-data/runtime/src/repository/usage/mod.rs b/crates/aether-data/runtime/src/repository/usage/mod.rs index abd7d9128..cac259b07 100644 --- a/crates/aether-data/runtime/src/repository/usage/mod.rs +++ b/crates/aether-data/runtime/src/repository/usage/mod.rs @@ -57,6 +57,7 @@ mod tests { #[test] fn strip_deprecated_usage_display_fields_clears_legacy_display_columns() { let usage = strip_deprecated_usage_display_fields(UpsertUsageRecord { + capture_retention: Default::default(), request_id: "req-1".to_string(), user_id: Some("user-1".to_string()), api_key_id: Some("key-1".to_string()), diff --git a/crates/aether-gateway/frontdoor/Cargo.toml b/crates/aether-gateway/frontdoor/Cargo.toml index c256e7208..f0e373aa8 100644 --- a/crates/aether-gateway/frontdoor/Cargo.toml +++ b/crates/aether-gateway/frontdoor/Cargo.toml @@ -10,13 +10,18 @@ description = "HTTP frontdoor middleware and request lifecycle primitives for Ae aether-ai-formats.workspace = true axum.workspace = true bytes.workspace = true +futures-util.workspace = true http.workspace = true -tokio.workspace = true +tokio = { workspace = true, features = ["io-util"] } +tokio-util = { workspace = true, features = ["rt"] } tracing.workspace = true uuid.workspace = true [dev-dependencies] -futures-util.workspace = true +http-body-util = "0.1" +hyper = { version = "1", features = ["client", "server", "http1", "http2"] } +hyper-util = { version = "0.1", features = ["tokio"] } serde_json.workspace = true +tokio = { workspace = true, features = ["test-util"] } tower = { version = "0.5", features = ["util"] } tracing-subscriber.workspace = true diff --git a/crates/aether-gateway/frontdoor/src/body.rs b/crates/aether-gateway/frontdoor/src/body.rs index a58d9a920..2927af805 100644 --- a/crates/aether-gateway/frontdoor/src/body.rs +++ b/crates/aether-gateway/frontdoor/src/body.rs @@ -1,14 +1,13 @@ //! Bounded request-body buffering for frontdoor adapters. //! -//! The policy reserves weighted memory before reading a body and holds the -//! reservation through the caller's normalization callback. This keeps body -//! buffering independent from gateway business routing while preventing a -//! burst of compressed requests from bypassing the memory budget. +//! The policy grows weighted reservations as bytes arrive and holds them through +//! normalization. Growth never waits while retaining a partial buffer, so +//! concurrent uploads cannot deadlock while competing for the remaining budget. -use axum::body::{to_bytes, Body}; +use axum::body::Body; use bytes::Bytes; +use futures_util::StreamExt; use http::{header, HeaderMap, StatusCode}; -use std::error::Error as StdError; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::sync::{OwnedSemaphorePermit, Semaphore}; @@ -127,7 +126,12 @@ impl BodyBufferPolicy { } pub fn reservation_bytes(&self, headers: &HeaderMap) -> usize { - reservation_bytes(headers, self.max_bytes, self.budget_bytes) + reservation_bytes( + headers, + self.max_bytes, + self.budget_bytes, + self.permit_bytes, + ) } pub fn reservation_permits(&self, reservation_bytes: usize) -> u32 { @@ -168,67 +172,134 @@ impl BodyBufferPolicy { }; Ok(BodyBufferReservation { - permit, + memory: BodyBufferBudget { + permit, + budget: Arc::clone(&self.budget), + budget_bytes: self.budget_bytes, + permit_bytes: self.permit_bytes, + requested_bytes, + }, max_bytes: effective_max_bytes, read_timeout: self.read_timeout, - requested_bytes, }) } } #[derive(Debug)] pub struct BodyBufferReservation { - permit: OwnedSemaphorePermit, + memory: BodyBufferBudget, max_bytes: u64, read_timeout: Option, +} + +#[derive(Debug)] +pub struct BodyBufferBudget { + permit: OwnedSemaphorePermit, + budget: Arc, + budget_bytes: usize, + permit_bytes: usize, requested_bytes: usize, } +impl BodyBufferBudget { + /// Reserve a new high-water mark before retaining or decoding more bytes. + /// Never queue for growth while another partial request may hold the rest. + pub fn try_reserve_bytes(&mut self, requested_bytes: usize) -> Result<(), BodyBufferError> { + if requested_bytes > self.budget_bytes { + return Err(self.overloaded(requested_bytes)); + } + let permits = reservation_permits(requested_bytes, self.permit_bytes) as usize; + let additional = permits.saturating_sub(self.permit.num_permits()); + if additional > 0 { + let permit = Arc::clone(&self.budget) + .try_acquire_many_owned(additional as u32) + .map_err(|_| self.overloaded(requested_bytes))?; + self.permit.merge(permit); + } + self.requested_bytes = self.requested_bytes.max(requested_bytes); + Ok(()) + } + + fn overloaded(&self, requested_bytes: usize) -> BodyBufferError { + BodyBufferError::Overloaded { + requested_bytes, + budget_bytes: self.budget_bytes, + timeout_ms: 0, + } + } +} + impl BodyBufferReservation { pub fn requested_bytes(&self) -> usize { - self.requested_bytes + self.memory.requested_bytes } pub async fn collect(self, body: Body) -> Result { let Self { - permit, + mut memory, max_bytes, read_timeout, - requested_bytes, } = self; let started_at = Instant::now(); - let body_limit = usize::try_from(max_bytes).unwrap_or(usize::MAX); - let collected = match read_timeout { - Some(read_timeout) => { - match tokio::time::timeout(read_timeout, to_bytes(body, body_limit)).await { - Ok(result) => result, - Err(_) => { - return Err(BodyBufferError::Timeout { - timeout_ms: duration_millis(read_timeout), - }); + let collect = async { + let mut stream = body.into_data_stream(); + let mut bytes = Vec::new(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(|error| BodyBufferError::ReadFailed { + message: error.to_string(), + })?; + let length = + bytes + .len() + .checked_add(chunk.len()) + .ok_or(BodyBufferError::TooLarge { + limit_bytes: max_bytes, + })?; + if length as u64 > max_bytes { + return Err(BodyBufferError::TooLarge { + limit_bytes: max_bytes, + }); + } + if length > bytes.capacity() { + let mut capacity = if length <= DEFAULT_BODY_BUFFER_PERMIT_BYTES { + length + } else { + bytes + .capacity() + .saturating_mul(2) + .max(length) + .min(usize::try_from(max_bytes).unwrap_or(usize::MAX)) + }; + if memory.try_reserve_bytes(capacity).is_err() { + memory.try_reserve_bytes(length)?; + capacity = length; + } + // Account for geometric growth, falling back to the bytes needed under load. + bytes + .try_reserve_exact(capacity - bytes.len()) + .map_err(|error| BodyBufferError::ReadFailed { + message: error.to_string(), + })?; + if bytes.capacity() > capacity { + memory.try_reserve_bytes(bytes.capacity())?; } } + bytes.extend_from_slice(&chunk); } - None => to_bytes(body, body_limit).await, - }; - let bytes = match collected { - Ok(bytes) => bytes, - Err(error) if collection_exceeded_limit(&error) => { - return Err(BodyBufferError::TooLarge { - limit_bytes: max_bytes, - }); - } - Err(error) => { - return Err(BodyBufferError::ReadFailed { - message: error.to_string(), - }); - } + Ok(Bytes::from(bytes)) }; + let bytes = match read_timeout { + Some(read_timeout) => tokio::time::timeout(read_timeout, collect) + .await + .map_err(|_| BodyBufferError::Timeout { + timeout_ms: duration_millis(read_timeout), + }), + None => Ok(collect.await), + }??; Ok(BufferedBody { bytes, - permit: Some(permit), - requested_bytes, + memory, elapsed: started_at.elapsed(), }) } @@ -237,8 +308,7 @@ impl BodyBufferReservation { #[derive(Debug)] pub struct BufferedBody { bytes: Bytes, - permit: Option, - requested_bytes: usize, + memory: BodyBufferBudget, elapsed: Duration, } @@ -248,19 +318,28 @@ impl BufferedBody { } pub fn requested_bytes(&self) -> usize { - self.requested_bytes + self.memory.requested_bytes } pub fn elapsed(&self) -> Duration { self.elapsed } - /// Apply normalization while retaining the memory permit until the - /// callback completes. + /// Retain the permit for a callback that does not grow the buffered payload. pub fn try_map(self, map: impl FnOnce(Bytes) -> Result) -> Result { - let Self { bytes, permit, .. } = self; - let result = map(bytes); - drop(permit); + self.try_map_with_budget(|bytes, _| map(bytes)) + } + + /// The callback must account for decoded buffers before allocating them. + pub fn try_map_with_budget( + self, + map: impl FnOnce(Bytes, &mut BodyBufferBudget) -> Result, + ) -> Result { + let Self { + bytes, mut memory, .. + } = self; + let result = map(bytes, &mut memory); + drop(memory); result } } @@ -389,23 +468,15 @@ fn invalid_body_headers(message: &str) -> BodyBufferError { } } -fn reservation_bytes(headers: &HeaderMap, max_bytes: u64, budget_bytes: usize) -> usize { +fn reservation_bytes( + headers: &HeaderMap, + max_bytes: u64, + budget_bytes: usize, + permit_bytes: usize, +) -> usize { let reservation_ceiling = usize::try_from(max_bytes) .unwrap_or(usize::MAX) .min(budget_bytes); - let encoded = headers - .get_all(header::CONTENT_ENCODING) - .iter() - .any(|value| { - value.to_str().map_or(true, |value| { - value.split(',').map(str::trim).any(|encoding| { - !encoding.is_empty() && !encoding.eq_ignore_ascii_case("identity") - }) - }) - }); - if encoded { - return reservation_ceiling; - } declared_content_length(headers) .ok() .flatten() @@ -414,7 +485,7 @@ fn reservation_bytes(headers: &HeaderMap, max_bytes: u64, budget_bytes: usize) - .unwrap_or(usize::MAX) .min(reservation_ceiling) }) - .unwrap_or(reservation_ceiling) + .unwrap_or_else(|| permit_bytes.min(reservation_ceiling)) } fn reservation_permits(reservation_bytes: usize, permit_bytes: usize) -> u32 { @@ -429,17 +500,6 @@ fn duration_millis(duration: Duration) -> u64 { u64::try_from(duration.as_millis()).unwrap_or(u64::MAX) } -fn collection_exceeded_limit(error: &(dyn StdError + 'static)) -> bool { - let mut current = Some(error); - while let Some(error) = current { - if error.to_string().contains("length limit exceeded") { - return true; - } - current = error.source(); - } - false -} - #[cfg(test)] mod tests { use super::{BodyBufferError, BodyBufferPolicy, DEFAULT_BODY_BUFFER_PERMIT_BYTES}; @@ -579,9 +639,9 @@ mod tests { let reservation = policy .reserve(&headers) .await - .expect("encoded unlimited body should reserve the available budget"); - assert_eq!(reservation.requested_bytes(), 4); - assert_eq!(budget.available_permits(), 0); + .expect("encoded unlimited body should reserve its initial chunk"); + assert_eq!(reservation.requested_bytes(), 1); + assert_eq!(budget.available_permits(), 3); let error = reservation .collect(Body::from(Bytes::from_static(b"01234"))) @@ -707,4 +767,172 @@ mod tests { .expect_err("exhausted budget should fail closed"); assert!(matches!(error, BodyBufferError::Overloaded { .. })); } + + #[tokio::test] + async fn small_compressed_and_unknown_length_requests_share_the_budget() { + let budget_bytes = 256 * 1024 * 1024; + let budget = Arc::new(Semaphore::new( + budget_bytes / DEFAULT_BODY_BUFFER_PERMIT_BYTES, + )); + let policy = policy( + budget_bytes as u64, + Duration::from_secs(1), + Arc::clone(&budget), + ); + let mut headers = HeaderMap::new(); + headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip")); + headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("1024")); + let compressed = policy.reserve(&headers).await.unwrap(); + let unknown = policy.reserve(&HeaderMap::new()).await.unwrap(); + assert_eq!(compressed.requested_bytes(), 1024); + assert_eq!(unknown.requested_bytes(), DEFAULT_BODY_BUFFER_PERMIT_BYTES); + assert_eq!(budget.available_permits(), 4094); + drop((compressed, unknown)); + assert_eq!(budget.available_permits(), 4096); + } + + #[tokio::test] + async fn unknown_length_body_grows_its_reservation_and_holds_it_until_normalized() { + let budget = Arc::new(Semaphore::new(8)); + let policy = BodyBufferPolicy::with_permit_bytes( + 8, + Duration::from_secs(1), + Duration::from_secs(1), + 8, + 1, + Arc::clone(&budget), + ); + let reservation = policy.reserve(&HeaderMap::new()).await.unwrap(); + assert_eq!(budget.available_permits(), 7); + let body = Body::from_stream(stream::iter([ + Ok::<_, std::io::Error>(Bytes::from_static(b"ab")), + Ok(Bytes::from_static(b"cd")), + ])); + let buffered = reservation.collect(body).await.unwrap(); + assert_eq!(buffered.bytes().as_ref(), b"abcd"); + assert_eq!(budget.available_permits(), 4); + buffered + .try_map_with_budget(|_, memory| { + memory.try_reserve_bytes(7)?; + assert_eq!(budget.available_permits(), 1); + Ok::<_, BodyBufferError>(()) + }) + .unwrap(); + assert_eq!(budget.available_permits(), 8); + } + + #[tokio::test] + async fn partial_upload_growth_rejects_without_waiting_and_releases_its_budget() { + let budget = Arc::new(Semaphore::new(2)); + let policy = BodyBufferPolicy::with_permit_bytes( + 2, + Duration::from_secs(60), + Duration::from_secs(60), + 2, + 1, + Arc::clone(&budget), + ); + let first = policy.reserve(&HeaderMap::new()).await.unwrap(); + let second = policy.reserve(&HeaderMap::new()).await.unwrap(); + let error = + tokio::time::timeout(Duration::from_millis(100), first.collect(Body::from("ab"))) + .await + .expect("growth must not wait while holding a partial reservation") + .unwrap_err(); + assert_eq!( + error, + BodyBufferError::Overloaded { + requested_bytes: 2, + budget_bytes: 2, + timeout_ms: 0, + } + ); + assert_eq!(budget.available_permits(), 1); + let buffered = second.collect(Body::from("ab")).await.unwrap(); + assert_eq!(budget.available_permits(), 0); + drop(buffered); + assert_eq!(budget.available_permits(), 2); + } + + #[tokio::test] + async fn normalization_budget_failure_releases_all_upload_permits() { + let budget = Arc::new(Semaphore::new(4)); + let policy = BodyBufferPolicy::with_permit_bytes( + 4, + Duration::from_secs(1), + Duration::from_secs(1), + 4, + 1, + Arc::clone(&budget), + ); + let buffered = policy + .reserve(&HeaderMap::new()) + .await + .unwrap() + .collect(Body::from("ab")) + .await + .unwrap(); + let result = buffered.try_map_with_budget(|_, memory| memory.try_reserve_bytes(5)); + assert!(matches!( + result, + Err(BodyBufferError::Overloaded { + requested_bytes: 5, + .. + }) + )); + assert_eq!(budget.available_permits(), 4); + } + + #[tokio::test] + async fn upload_growth_uses_available_budget_without_requiring_geometric_headroom() { + let budget = Arc::new(Semaphore::new(128)); + let held = Arc::clone(&budget).acquire_many_owned(32).await.unwrap(); + let policy = BodyBufferPolicy::with_permit_bytes( + 128 * 1024, + Duration::from_secs(1), + Duration::from_secs(1), + 128 * 1024, + 1024, + Arc::clone(&budget), + ); + let body = Body::from_stream(stream::iter([ + Ok::<_, std::io::Error>(Bytes::from(vec![b'a'; 70_000])), + Ok(Bytes::from(vec![b'b'; 20_000])), + ])); + let buffered = policy + .reserve(&HeaderMap::new()) + .await + .unwrap() + .collect(body) + .await + .unwrap(); + assert_eq!(buffered.bytes().len(), 90_000); + assert_eq!(buffered.requested_bytes(), 90_000); + drop((buffered, held)); + assert_eq!(budget.available_permits(), 128); + } + + #[tokio::test] + async fn cancellation_releases_partial_upload_budget() { + let budget = Arc::new(Semaphore::new(4)); + let policy = BodyBufferPolicy::with_permit_bytes( + 4, + Duration::from_secs(60), + Duration::from_secs(1), + 4, + 1, + Arc::clone(&budget), + ); + let reservation = policy.reserve(&HeaderMap::new()).await.unwrap(); + let body = Body::from_stream( + stream::once(async { Ok::<_, std::io::Error>(Bytes::from_static(b"ab")) }) + .chain(stream::pending()), + ); + assert!( + tokio::time::timeout(Duration::from_millis(10), reservation.collect(body)) + .await + .is_err() + ); + assert_eq!(budget.available_permits(), 4); + } } diff --git a/crates/aether-gateway/frontdoor/src/connection.rs b/crates/aether-gateway/frontdoor/src/connection.rs new file mode 100644 index 000000000..f60b4ab69 --- /dev/null +++ b/crates/aether-gateway/frontdoor/src/connection.rs @@ -0,0 +1,229 @@ +use std::future::Future; +use std::io; +use std::net::SocketAddr; +use std::pin::Pin; +use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; +use std::sync::Arc; +use std::task::{Context, Poll}; +use std::time::Duration; + +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; +use tokio_util::sync::{CancellationToken, WaitForCancellationFutureOwned}; + +const MAX_HTTP_CONNECTIONS: usize = 65_536; +const FD_RESERVE: usize = 256; + +/// Selects a per-server incoming TCP limit. The FD allowance leaves room for +/// upstream sockets and process infrastructure; it is not a complete FD budget. +pub fn http_connection_limit( + configured: Option, + request_limit: usize, + websocket_limit: usize, + fd_soft_limit: Option, +) -> usize { + let configured = configured + .filter(|limit| *limit > 0) + .unwrap_or_else(|| request_limit.saturating_add(websocket_limit)) + .clamp(1, MAX_HTTP_CONNECTIONS); + let fd_allowance = fd_soft_limit + .map(|limit| (limit.saturating_sub(FD_RESERVE) / 2).max(1)) + .unwrap_or(MAX_HTTP_CONNECTIONS); + configured.min(fd_allowance) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct HttpConnectionBudgetSnapshot { + pub limit: usize, + pub in_flight: usize, + pub high_watermark: usize, + pub rejected_total: u64, + pub accept_errors_total: u64, +} + +/// Share one budget across listeners. Admission belongs to the underlying IO, +/// so HTTP/1 upgrades keep their permit and HTTP/2 streams share one permit. +#[derive(Debug)] +pub struct HttpConnectionBudget { + limit: usize, + permits: Arc, + in_flight: AtomicUsize, + high_watermark: AtomicUsize, + rejected_total: AtomicU64, + accept_errors_total: AtomicU64, + shutdown: CancellationToken, +} + +impl HttpConnectionBudget { + pub fn new(limit: usize) -> Self { + let limit = limit.clamp(1, MAX_HTTP_CONNECTIONS.min(Semaphore::MAX_PERMITS)); + Self { + limit, + permits: Arc::new(Semaphore::new(limit)), + in_flight: AtomicUsize::new(0), + high_watermark: AtomicUsize::new(0), + rejected_total: AtomicU64::new(0), + accept_errors_total: AtomicU64::new(0), + shutdown: CancellationToken::new(), + } + } + + /// Admit after accepting. Waiting for a permit before accept can let idle + /// reuseport listeners monopolize permits needed by a busy listener. + pub fn try_admit(self: &Arc, io: T) -> Result, ()> { + let permit = Arc::clone(&self.permits).try_acquire_owned().map_err(|_| { + self.rejected_total.fetch_add(1, Ordering::Relaxed); + })?; + let in_flight = self.in_flight.fetch_add(1, Ordering::Relaxed) + 1; + self.high_watermark.fetch_max(in_flight, Ordering::Relaxed); + Ok(AdmittedConnection { + io, + read_shutdown: Box::pin(self.shutdown.clone().cancelled_owned()), + write_shutdown: Box::pin(self.shutdown.clone().cancelled_owned()), + _permit: ConnectionPermit { + budget: Arc::clone(self), + _permit: permit, + }, + }) + } + + pub fn snapshot(&self) -> HttpConnectionBudgetSnapshot { + HttpConnectionBudgetSnapshot { + limit: self.limit, + in_flight: self.in_flight.load(Ordering::Relaxed), + high_watermark: self.high_watermark.load(Ordering::Relaxed), + rejected_total: self.rejected_total.load(Ordering::Relaxed), + accept_errors_total: self.accept_errors_total.load(Ordering::Relaxed), + } + } + + /// End the drain deadline for all sockets, including upgraded connections. + pub fn force_close(&self) { + self.permits.close(); + self.shutdown.cancel(); + } + + pub async fn wait_for_forced_close(&self) { + self.shutdown.cancelled().await; + } + + pub async fn accept(&self, listener: &TcpListener) -> (TcpStream, SocketAddr) { + self.accept_with(|| listener.accept()).await + } + + async fn accept_with(&self, mut accept: A) -> T + where + A: FnMut() -> F, + F: Future>, + { + loop { + match accept().await { + Ok(connection) => return connection, + Err(error) => { + self.accept_errors_total.fetch_add(1, Ordering::Relaxed); + // Match Axum's listener behavior: failed peers can be retried + // immediately; resource failures such as EMFILE need backoff. + if matches!( + error.kind(), + io::ErrorKind::ConnectionRefused + | io::ErrorKind::ConnectionAborted + | io::ErrorKind::ConnectionReset + ) { + continue; + } + tracing::error!( + event_name = "http_connection_accept_failed", + error = %error, + "HTTP listener accept failed; retrying after one second" + ); + tokio::time::sleep(Duration::from_secs(1)).await; + } + } + } + } +} + +#[derive(Debug)] +struct ConnectionPermit { + budget: Arc, + _permit: OwnedSemaphorePermit, +} + +impl Drop for ConnectionPermit { + fn drop(&mut self) { + self.budget.in_flight.fetch_sub(1, Ordering::Relaxed); + } +} + +#[derive(Debug)] +pub struct AdmittedConnection { + // Close the socket before returning its permit, including upgrade teardown. + io: T, + read_shutdown: Pin>, + write_shutdown: Pin>, + _permit: ConnectionPermit, +} + +impl AsyncRead for AdmittedConnection { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &mut ReadBuf<'_>, + ) -> Poll> { + if self.read_shutdown.as_mut().poll(cx).is_ready() { + return Poll::Ready(Err(shutdown_error())); + } + Pin::new(&mut self.io).poll_read(cx, buf) + } +} + +impl AsyncWrite for AdmittedConnection { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + if self.write_shutdown.as_mut().poll(cx).is_ready() { + return Poll::Ready(Err(shutdown_error())); + } + Pin::new(&mut self.io).poll_write(cx, buf) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if self.write_shutdown.as_mut().poll(cx).is_ready() { + return Poll::Ready(Err(shutdown_error())); + } + Pin::new(&mut self.io).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.io).poll_shutdown(cx) + } + + fn is_write_vectored(&self) -> bool { + self.io.is_write_vectored() + } + + fn poll_write_vectored( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bufs: &[io::IoSlice<'_>], + ) -> Poll> { + if self.write_shutdown.as_mut().poll(cx).is_ready() { + return Poll::Ready(Err(shutdown_error())); + } + Pin::new(&mut self.io).poll_write_vectored(cx, bufs) + } +} + +fn shutdown_error() -> io::Error { + io::Error::new( + io::ErrorKind::ConnectionAborted, + "gateway shutdown deadline reached", + ) +} + +#[cfg(test)] +#[path = "connection_tests.rs"] +mod tests; diff --git a/crates/aether-gateway/frontdoor/src/connection_tests.rs b/crates/aether-gateway/frontdoor/src/connection_tests.rs new file mode 100644 index 000000000..f253b0b61 --- /dev/null +++ b/crates/aether-gateway/frontdoor/src/connection_tests.rs @@ -0,0 +1,411 @@ +use std::collections::VecDeque; +use std::convert::Infallible; +use std::sync::atomic::AtomicBool; + +use bytes::Bytes; +use http::{Request, Response, StatusCode}; +use http_body_util::{BodyExt, Empty}; +use hyper::body::Incoming; +use hyper::service::service_fn; +use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; + +use super::*; + +async fn within(future: impl Future) -> T { + tokio::time::timeout(Duration::from_secs(3), future) + .await + .expect("connection test exceeded its deadline") +} + +async fn tcp_pair() -> (TcpStream, TcpStream) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let client = TcpStream::connect(listener.local_addr().unwrap()) + .await + .unwrap(); + let (server, _) = listener.accept().await.unwrap(); + (client, server) +} + +#[tokio::test] +async fn forced_shutdown_wakes_independent_read_and_write_waiters() { + let budget = Arc::new(HttpConnectionBudget::new(2)); + let (io, _peer) = tokio::io::duplex(1); + let io = budget.try_admit(io).unwrap(); + let (mut reader, mut writer) = tokio::io::split(io); + writer.write_all(b"a").await.unwrap(); + let reading = tokio::spawn(async move { reader.read_u8().await }); + let writing = tokio::spawn(async move { writer.write_all(b"b").await }); + tokio::task::yield_now().await; + budget.force_close(); + assert_eq!( + within(reading).await.unwrap().unwrap_err().kind(), + io::ErrorKind::ConnectionAborted + ); + assert_eq!( + within(writing).await.unwrap().unwrap_err().kind(), + io::ErrorKind::ConnectionAborted + ); + assert_eq!(budget.snapshot().in_flight, 0); + assert!(budget.try_admit(tokio::io::empty()).is_err()); +} + +#[tokio::test] +async fn forced_shutdown_before_first_poll_closes_read_write_and_flush() { + let budget = Arc::new(HttpConnectionBudget::new(1)); + let (io, _peer) = tokio::io::duplex(16); + let mut io = budget.try_admit(io).unwrap(); + budget.force_close(); + assert_eq!( + io.read_u8().await.unwrap_err().kind(), + io::ErrorKind::ConnectionAborted + ); + assert_eq!( + io.write_all(b"a").await.unwrap_err().kind(), + io::ErrorKind::ConnectionAborted + ); + assert_eq!( + io.flush().await.unwrap_err().kind(), + io::ErrorKind::ConnectionAborted + ); + let buffers = [io::IoSlice::new(b"a")]; + assert_eq!( + io.write_vectored(&buffers).await.unwrap_err().kind(), + io::ErrorKind::ConnectionAborted + ); + io.shutdown().await.unwrap(); + drop(io); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[test] +fn http_connection_limits_apply_to_auto_explicit_zero_and_fd_bounds() { + for (configured, requests, websockets, fd_limit, expected) in [ + (None, 0, 0, None, 1), + (None, 3, 5, None, 8), + (Some(0), 3, 5, None, 8), + (Some(9), 3, 5, None, 9), + (Some(usize::MAX), 0, 0, None, MAX_HTTP_CONNECTIONS), + (None, usize::MAX, usize::MAX, None, MAX_HTTP_CONNECTIONS), + (None, 512, 512, Some(1_024), 384), + (Some(500), 3, 5, Some(1_024), 384), + (Some(32), 3, 5, Some(1_024), 32), + (Some(500), 3, 5, Some(260), 2), + (Some(500), 3, 5, Some(258), 1), + (Some(500), 3, 5, Some(0), 1), + ] { + assert_eq!( + http_connection_limit(configured, requests, websockets, fd_limit), + expected, + ); + } + assert_eq!(HttpConnectionBudget::new(0).snapshot().limit, 1); + assert_eq!( + HttpConnectionBudget::new(usize::MAX).snapshot().limit, + MAX_HTTP_CONNECTIONS, + ); +} + +#[test] +fn http_connection_budget_drops_io_before_returning_capacity() { + struct DropProbe { + budget: Arc, + observed: Arc, + } + impl Drop for DropProbe { + fn drop(&mut self) { + self.observed + .store(self.budget.snapshot().in_flight == 1, Ordering::Relaxed); + } + } + let budget = Arc::new(HttpConnectionBudget::new(1)); + let admitted_dropped = Arc::new(AtomicBool::new(false)); + let rejected_dropped = Arc::new(AtomicBool::new(false)); + let admitted = budget + .try_admit(DropProbe { + budget: Arc::clone(&budget), + observed: Arc::clone(&admitted_dropped), + }) + .unwrap(); + assert!(budget + .try_admit(DropProbe { + budget: Arc::clone(&budget), + observed: Arc::clone(&rejected_dropped), + }) + .is_err()); + assert!(rejected_dropped.load(Ordering::Relaxed)); + assert_eq!(budget.snapshot().in_flight, 1); + drop(admitted); + assert!(admitted_dropped.load(Ordering::Relaxed)); + assert_eq!(budget.snapshot().in_flight, 0); + assert_eq!(budget.snapshot().high_watermark, 1); + assert_eq!(budget.snapshot().rejected_total, 1); + drop(budget.try_admit(()).unwrap()); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[tokio::test] +async fn http_connection_budget_forwards_read_write_vectored_and_half_close() { + let budget = Arc::new(HttpConnectionBudget::new(1)); + let (mut peer, io) = tokio::io::duplex(32); + let vectored = io.is_write_vectored(); + let mut admitted = budget.try_admit(io).unwrap(); + assert_eq!(admitted.is_write_vectored(), vectored); + let written = admitted + .write_vectored(&[io::IoSlice::new(b"ab"), io::IoSlice::new(b"cd")]) + .await + .unwrap(); + assert!((1..=4).contains(&written)); + admitted.write_all(&b"abcd"[written..]).await.unwrap(); + admitted.flush().await.unwrap(); + let mut message = [0; 4]; + peer.read_exact(&mut message).await.unwrap(); + assert_eq!(&message, b"abcd"); + + admitted.shutdown().await.unwrap(); + assert_eq!(peer.read(&mut [0; 1]).await.unwrap(), 0); + assert_eq!(budget.snapshot().in_flight, 1); + peer.write_all(b"reply").await.unwrap(); + let mut reply = [0; 5]; + admitted.read_exact(&mut reply).await.unwrap(); + assert_eq!(&reply, b"reply"); + drop(admitted); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[tokio::test] +async fn http_connection_budget_two_tcp_listeners_share_capacity_and_recover() { + let first_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let second_listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let budget = Arc::new(HttpConnectionBudget::new(1)); + let first_client = TcpStream::connect(first_listener.local_addr().unwrap()) + .await + .unwrap(); + let (first_io, _) = within(budget.accept(&first_listener)).await; + let first = budget.try_admit(first_io).unwrap(); + + let mut rejected_client = TcpStream::connect(second_listener.local_addr().unwrap()) + .await + .unwrap(); + let (second_io, _) = within(budget.accept(&second_listener)).await; + assert!(Arc::clone(&budget).try_admit(second_io).is_err()); + let rejected_read = within(rejected_client.read(&mut [0; 1])).await; + assert!( + matches!(rejected_read, Ok(0)) + || matches!(rejected_read, Err(ref error) if matches!( + error.kind(), io::ErrorKind::ConnectionReset | io::ErrorKind::ConnectionAborted + )) + ); + assert_eq!(budget.snapshot().in_flight, 1); + assert_eq!(budget.snapshot().rejected_total, 1); + + drop((first, first_client, rejected_client)); + let replacement_client = TcpStream::connect(second_listener.local_addr().unwrap()) + .await + .unwrap(); + let (replacement_io, _) = within(budget.accept(&second_listener)).await; + let replacement = budget.try_admit(replacement_io).unwrap(); + assert_eq!(budget.snapshot().in_flight, 1); + assert_eq!(budget.snapshot().high_watermark, 1); + drop((replacement, replacement_client)); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[tokio::test] +async fn http_connection_budget_task_cancelled_before_first_poll_returns_permit() { + let (client, server) = tcp_pair().await; + let budget = Arc::new(HttpConnectionBudget::new(1)); + let admitted = budget.try_admit(server).unwrap(); + let polled = Arc::new(AtomicBool::new(false)); + let polled_by_task = Arc::clone(&polled); + let task = tokio::spawn(async move { + let _io = admitted; + polled_by_task.store(true, Ordering::Relaxed); + std::future::pending::<()>().await; + }); + task.abort(); + assert!(within(task).await.unwrap_err().is_cancelled()); + assert!(!polled.load(Ordering::Relaxed)); + assert_eq!(budget.snapshot().in_flight, 0); + drop(client); +} + +#[tokio::test] +async fn http_connection_budget_header_timeout_and_parse_failure_return_permit() { + for malformed in [false, true] { + let (mut client, server_io) = tcp_pair().await; + let budget = Arc::new(HttpConnectionBudget::new(1)); + let admitted = budget.try_admit(server_io).unwrap(); + let requests = Arc::new(AtomicUsize::new(0)); + let requests_seen = Arc::clone(&requests); + let server = tokio::spawn(async move { + let service = service_fn(move |_: Request| { + requests_seen.fetch_add(1, Ordering::Relaxed); + async { Ok::<_, Infallible>(Response::new(Empty::::new())) } + }); + hyper::server::conn::http1::Builder::new() + .timer(TokioTimer::new()) + .header_read_timeout(Duration::from_millis(20)) + .serve_connection(TokioIo::new(admitted), service) + .await + }); + if malformed { + client.write_all(b"invalid request\r\n\r\n").await.unwrap(); + } + let _connection_result = within(server).await.unwrap(); + assert_eq!(requests.load(Ordering::Relaxed), 0); + assert_eq!(budget.snapshot().in_flight, 0); + drop(client); + } +} + +#[tokio::test] +async fn http_connection_budget_h1_upgrade_keeps_permit_after_connection_future_finishes() { + let (mut client, server_io) = tcp_pair().await; + let budget = Arc::new(HttpConnectionBudget::new(1)); + let admitted = budget.try_admit(server_io).unwrap(); + let (upgrade_finished_tx, upgrade_finished_rx) = tokio::sync::oneshot::channel(); + let upgrade_finished_tx = Arc::new(std::sync::Mutex::new(Some(upgrade_finished_tx))); + let server = tokio::spawn(async move { + let service = service_fn(move |mut request: Request| { + let on_upgrade = hyper::upgrade::on(&mut request); + let finished = upgrade_finished_tx.lock().unwrap().take().unwrap(); + tokio::spawn(async move { + let mut upgraded = TokioIo::new(on_upgrade.await.unwrap()); + let mut message = [0; 4]; + upgraded.read_exact(&mut message).await.unwrap(); + upgraded.write_all(&message).await.unwrap(); + assert_eq!(upgraded.read(&mut [0; 1]).await.unwrap(), 0); + drop(upgraded); + let _ = finished.send(()); + }); + async { + Ok::<_, Infallible>( + Response::builder() + .status(StatusCode::SWITCHING_PROTOCOLS) + .header("connection", "upgrade") + .header("upgrade", "echo") + .body(Empty::::new()) + .unwrap(), + ) + } + }); + hyper::server::conn::http1::Builder::new() + .serve_connection(TokioIo::new(admitted), service) + .with_upgrades() + .await + }); + client + .write_all( + b"GET / HTTP/1.1\r\nHost: localhost\r\nConnection: upgrade\r\nUpgrade: echo\r\n\r\n", + ) + .await + .unwrap(); + let response = within(async { + let mut headers = Vec::new(); + while !headers.ends_with(b"\r\n\r\n") { + assert!(headers.len() < 4096); + headers.push(client.read_u8().await.unwrap()); + } + headers + }) + .await; + assert!(response.starts_with(b"HTTP/1.1 101")); + within(server).await.unwrap().unwrap(); + assert_eq!(budget.snapshot().in_flight, 1); + assert!(budget.try_admit(()).is_err()); + client.write_all(b"ping").await.unwrap(); + let mut echoed = [0; 4]; + within(client.read_exact(&mut echoed)).await.unwrap(); + assert_eq!(&echoed, b"ping"); + drop(client); + within(upgrade_finished_rx).await.unwrap(); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[tokio::test] +async fn http_connection_budget_h2_parallel_streams_share_one_socket_permit() { + let (client_io, server_io) = tcp_pair().await; + let budget = Arc::new(HttpConnectionBudget::new(1)); + let admitted = budget.try_admit(server_io).unwrap(); + let requests = Arc::new(AtomicUsize::new(0)); + let concurrent = Arc::new(tokio::sync::Barrier::new(3)); + let server_requests = Arc::clone(&requests); + let server_concurrent = Arc::clone(&concurrent); + let server = tokio::spawn(async move { + let service = service_fn(move |_: Request| { + server_requests.fetch_add(1, Ordering::Relaxed); + let concurrent = Arc::clone(&server_concurrent); + async move { + concurrent.wait().await; + Ok::<_, Infallible>(Response::new(Empty::::new())) + } + }); + hyper::server::conn::http2::Builder::new(TokioExecutor::new()) + .max_concurrent_streams(2) + .serve_connection(TokioIo::new(admitted), service) + .await + }); + let (sender, connection) = within( + hyper::client::conn::http2::Builder::new(TokioExecutor::new()) + .handshake::<_, Empty>(TokioIo::new(client_io)), + ) + .await + .unwrap(); + let client_driver = tokio::spawn(connection); + let mut requests_in_flight = tokio::task::JoinSet::new(); + for path in ["first", "second"] { + let mut sender = sender.clone(); + requests_in_flight.spawn(async move { + let response = sender + .send_request( + Request::builder() + .uri(format!("http://localhost/{path}")) + .body(Empty::::new()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + response.into_body().collect().await.unwrap(); + }); + } + within(concurrent.wait()).await; + assert_eq!(requests.load(Ordering::Relaxed), 2); + assert_eq!(budget.snapshot().in_flight, 1); + assert_eq!(budget.snapshot().high_watermark, 1); + within(async { + while let Some(result) = requests_in_flight.join_next().await { + result.unwrap(); + } + }) + .await; + assert_eq!(budget.snapshot().in_flight, 1); + drop(sender); + client_driver.abort(); + let _ = within(client_driver).await; + let _ = within(server).await.unwrap(); + assert_eq!(budget.snapshot().in_flight, 0); +} + +#[tokio::test(start_paused = true)] +async fn http_connection_accept_retries_peer_errors_and_backs_off_resource_errors() { + let budget = HttpConnectionBudget::new(1); + let mut attempts = VecDeque::from([ + Err(io::Error::from(io::ErrorKind::ConnectionAborted)), + Err(io::Error::from(io::ErrorKind::ConnectionReset)), + Err(io::Error::other("injected file descriptor exhaustion")), + Err(io::Error::other("injected temporary accept failure")), + Ok(42), + ]); + let started = tokio::time::Instant::now(); + let accepted = budget + .accept_with(|| std::future::ready(attempts.pop_front().unwrap())) + .await; + assert_eq!(accepted, 42); + assert!(attempts.is_empty()); + assert_eq!(started.elapsed(), Duration::from_secs(2)); + assert_eq!(budget.snapshot().accept_errors_total, 4); + assert_eq!(budget.snapshot().in_flight, 0); + assert_eq!(budget.snapshot().rejected_total, 0); +} diff --git a/crates/aether-gateway/frontdoor/src/lib.rs b/crates/aether-gateway/frontdoor/src/lib.rs index 4adf89dc5..feeabd876 100644 --- a/crates/aether-gateway/frontdoor/src/lib.rs +++ b/crates/aether-gateway/frontdoor/src/lib.rs @@ -1,11 +1,15 @@ pub mod body; +mod connection; pub mod middleware; mod request_id; pub use body::{ - BodyBufferError, BodyBufferPolicy, BodyBufferReservation, BufferedBody, + BodyBufferBudget, BodyBufferError, BodyBufferPolicy, BodyBufferReservation, BufferedBody, DEFAULT_BODY_BUFFER_PERMIT_BYTES, }; +pub use connection::{ + http_connection_limit, AdmittedConnection, HttpConnectionBudget, HttpConnectionBudgetSnapshot, +}; pub use middleware::access_log::{ access_log_middleware, sanitize_access_log_path, should_downgrade_access_log, diff --git a/crates/aether-runtime/base/src/lib.rs b/crates/aether-runtime/base/src/lib.rs index 452333bdb..712325fc5 100644 --- a/crates/aether-runtime/base/src/lib.rs +++ b/crates/aether-runtime/base/src/lib.rs @@ -34,5 +34,6 @@ pub use queue::{ pub use redaction::{summarize_text_payload, TextPayloadSummary}; pub use shutdown::wait_for_shutdown_signal; pub use tracing::{ - init_reloadable_service_tracing, init_reloadable_tracing, LogFormat, LogReloader, + init_reloadable_service_tracing, init_reloadable_tracing, logging_metric_samples, + shutdown_logging, LogFormat, LogReloader, LogShutdownGuard, }; diff --git a/crates/aether-runtime/base/src/tracing.rs b/crates/aether-runtime/base/src/tracing.rs index 057b1fc65..8817d0606 100644 --- a/crates/aether-runtime/base/src/tracing.rs +++ b/crates/aether-runtime/base/src/tracing.rs @@ -21,6 +21,11 @@ use crate::config::ServiceRuntimeConfig; use crate::error::RuntimeBootstrapError; use crate::observability::{FileLoggingConfig, LogDestination, LogRotation}; +mod writer; + +pub use writer::{logging_metric_samples, shutdown_logging, LogShutdownGuard}; +use writer::{register_log_workers, LogWorker, NonBlockingLogWriter}; + static TRACING_INIT: OnceLock> = OnceLock::new(); pub type LogReloader = Box; @@ -385,17 +390,12 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo .unwrap_or_else(|_| config.default_log_filter.into()); let identity = RuntimeLogIdentity::from_config(&config); - let (file_writer, startup_cleanup_warning) = - if config.observability.log_destination.needs_file_sink() { - let Some(file_logging) = config.observability.file_logging.clone() else { - return Err("file logging requires a configured log directory".to_string()); - }; - let (writer, startup_cleanup_warning) = - RollingFileMakeWriter::new(config.service_name, file_logging)?; - (Some(writer), startup_cleanup_warning) - } else { - (None, None) - }; + let RuntimeLogWriters { + stdout_writer, + file_writer, + workers, + startup_cleanup_warning, + } = RuntimeLogWriters::new(&config)?; let init_result = match ( config.observability.log_format, @@ -403,16 +403,26 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo ) { (LogFormat::Pretty, LogDestination::Stdout) => tracing_subscriber::registry() .with(filter) - .with(tracing_subscriber::fmt::layer().event_format( - PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()), - )) + .with( + tracing_subscriber::fmt::layer() + .event_format(PrettyRuntimeEventFormatter::new( + identity.clone(), + stdout_supports_ansi(), + )) + .with_writer( + stdout_writer.clone().expect("stdout writer should exist"), + ), + ) .try_init(), (LogFormat::Json, LogDestination::Stdout) => tracing_subscriber::registry() .with(filter) .with( tracing_subscriber::fmt::layer() .json() - .event_format(JsonRuntimeEventFormatter::new(identity.clone())), + .event_format(JsonRuntimeEventFormatter::new(identity.clone())) + .with_writer( + stdout_writer.clone().expect("stdout writer should exist"), + ), ) .try_init(), (LogFormat::Pretty, LogDestination::File) => tracing_subscriber::registry() @@ -435,9 +445,16 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo .try_init(), (LogFormat::Pretty, LogDestination::Both) => tracing_subscriber::registry() .with(filter) - .with(tracing_subscriber::fmt::layer().event_format( - PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()), - )) + .with( + tracing_subscriber::fmt::layer() + .event_format(PrettyRuntimeEventFormatter::new( + identity.clone(), + stdout_supports_ansi(), + )) + .with_writer( + stdout_writer.clone().expect("stdout writer should exist"), + ), + ) .with( tracing_subscriber::fmt::layer() .with_ansi(false) @@ -450,7 +467,10 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo .with( tracing_subscriber::fmt::layer() .json() - .event_format(JsonRuntimeEventFormatter::new(identity.clone())), + .event_format(JsonRuntimeEventFormatter::new(identity.clone())) + .with_writer( + stdout_writer.clone().expect("stdout writer should exist"), + ), ) .with( tracing_subscriber::fmt::layer() @@ -463,6 +483,7 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo .map_err(|err| err.to_string()); if init_result.is_ok() { + register_log_workers(workers); if let Some(warning) = startup_cleanup_warning.as_ref() { emit_log_cleanup_warning("startup", warning.log_dir.as_path(), &warning.error); } @@ -497,39 +518,35 @@ pub fn init_reloadable_service_tracing( let (filter_layer, reload_handle) = reload::Layer::new(filter); let identity = RuntimeLogIdentity::from_config(&config); - let (file_writer, startup_cleanup_warning) = - if config.observability.log_destination.needs_file_sink() { - let Some(file_logging) = config.observability.file_logging.clone() else { - return Err(RuntimeBootstrapError::Tracing( - "file logging requires a configured log directory".to_string(), - )); - }; - let (writer, startup_cleanup_warning) = - RollingFileMakeWriter::new(config.service_name, file_logging) - .map_err(RuntimeBootstrapError::Tracing)?; - (Some(writer), startup_cleanup_warning) - } else { - (None, None) - }; + let RuntimeLogWriters { + stdout_writer, + file_writer, + workers, + startup_cleanup_warning, + } = RuntimeLogWriters::new(&config).map_err(RuntimeBootstrapError::Tracing)?; match ( config.observability.log_format, config.observability.log_destination, ) { - (LogFormat::Pretty, LogDestination::Stdout) => { - tracing_subscriber::registry() - .with(filter_layer) - .with(tracing_subscriber::fmt::layer().event_format( - PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()), - )) - .try_init() - } + (LogFormat::Pretty, LogDestination::Stdout) => tracing_subscriber::registry() + .with(filter_layer) + .with( + tracing_subscriber::fmt::layer() + .event_format(PrettyRuntimeEventFormatter::new( + identity.clone(), + stdout_supports_ansi(), + )) + .with_writer(stdout_writer.clone().expect("stdout writer should exist")), + ) + .try_init(), (LogFormat::Json, LogDestination::Stdout) => tracing_subscriber::registry() .with(filter_layer) .with( tracing_subscriber::fmt::layer() .json() - .event_format(JsonRuntimeEventFormatter::new(identity.clone())), + .event_format(JsonRuntimeEventFormatter::new(identity.clone())) + .with_writer(stdout_writer.clone().expect("stdout writer should exist")), ) .try_init(), (LogFormat::Pretty, LogDestination::File) => tracing_subscriber::registry() @@ -550,26 +567,30 @@ pub fn init_reloadable_service_tracing( .with_writer(file_writer.clone().expect("file writer should exist")), ) .try_init(), - (LogFormat::Pretty, LogDestination::Both) => { - tracing_subscriber::registry() - .with(filter_layer) - .with(tracing_subscriber::fmt::layer().event_format( - PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()), - )) - .with( - tracing_subscriber::fmt::layer() - .with_ansi(false) - .event_format(PrettyRuntimeEventFormatter::new(identity.clone(), false)) - .with_writer(file_writer.clone().expect("file writer should exist")), - ) - .try_init() - } + (LogFormat::Pretty, LogDestination::Both) => tracing_subscriber::registry() + .with(filter_layer) + .with( + tracing_subscriber::fmt::layer() + .event_format(PrettyRuntimeEventFormatter::new( + identity.clone(), + stdout_supports_ansi(), + )) + .with_writer(stdout_writer.clone().expect("stdout writer should exist")), + ) + .with( + tracing_subscriber::fmt::layer() + .with_ansi(false) + .event_format(PrettyRuntimeEventFormatter::new(identity.clone(), false)) + .with_writer(file_writer.clone().expect("file writer should exist")), + ) + .try_init(), (LogFormat::Json, LogDestination::Both) => tracing_subscriber::registry() .with(filter_layer) .with( tracing_subscriber::fmt::layer() .json() - .event_format(JsonRuntimeEventFormatter::new(identity.clone())), + .event_format(JsonRuntimeEventFormatter::new(identity.clone())) + .with_writer(stdout_writer.clone().expect("stdout writer should exist")), ) .with( tracing_subscriber::fmt::layer() @@ -581,6 +602,7 @@ pub fn init_reloadable_service_tracing( } .map_err(|err| RuntimeBootstrapError::Tracing(err.to_string()))?; + register_log_workers(workers); if let Some(warning) = startup_cleanup_warning.as_ref() { emit_log_cleanup_warning("startup", warning.log_dir.as_path(), &warning.error); } @@ -595,6 +617,54 @@ pub fn init_reloadable_service_tracing( })) } +struct RuntimeLogWriters { + stdout_writer: Option, + file_writer: Option, + workers: Vec, + startup_cleanup_warning: Option, +} + +impl RuntimeLogWriters { + fn new(config: &ServiceRuntimeConfig) -> Result { + let (file_sink, startup_cleanup_warning) = + if config.observability.log_destination.needs_file_sink() { + let file_logging = config + .observability + .file_logging + .clone() + .ok_or("file logging requires a configured log directory")?; + let (sink, warning) = + RollingFileMakeWriter::new(config.service_name, file_logging)?; + (Some(sink), warning) + } else { + (None, None) + }; + let mut workers = Vec::with_capacity(2); + let stdout_writer = if config.observability.log_destination != LogDestination::File { + let (writer, worker) = NonBlockingLogWriter::new("stdout", io::stdout()) + .map_err(|err| format!("failed to start stdout log writer: {err}"))?; + workers.push(worker); + Some(writer) + } else { + None + }; + let file_writer = if let Some(sink) = file_sink { + let (writer, worker) = NonBlockingLogWriter::new("file", sink.make_writer()) + .map_err(|err| format!("failed to start file log writer: {err}"))?; + workers.push(worker); + Some(writer) + } else { + None + }; + Ok(Self { + stdout_writer, + file_writer, + workers, + startup_cleanup_warning, + }) + } +} + #[derive(Debug, Clone)] struct RollingFileMakeWriter { sink: Arc, @@ -692,7 +762,10 @@ impl RollingFileSink { } fn write(&self, buf: &[u8]) -> io::Result { - let now = Local::now(); + self.write_at(buf, Local::now()) + } + + fn write_at(&self, buf: &[u8], now: DateTime) -> io::Result { let mut state = self .state .lock() @@ -833,8 +906,19 @@ fn spawn_log_cleanup_task(service_name: &'static str, config: FileLoggingConfig) let interval = Duration::from_secs(6 * 60 * 60); loop { tokio::time::sleep(interval).await; - if let Err(err) = cleanup_log_files(service_name, &config) { - emit_log_cleanup_warning("background", config.dir.as_path(), &err); + let cleanup_config = config.clone(); + match tokio::task::spawn_blocking(move || { + cleanup_log_files(service_name, &cleanup_config) + }) + .await + { + Ok(Ok(_)) => {} + Ok(Err(err)) => { + emit_log_cleanup_warning("background", config.dir.as_path(), &err); + } + Err(err) => { + emit_log_cleanup_warning("background", config.dir.as_path(), &err); + } } } }); @@ -1028,6 +1112,78 @@ mod tests { ); } + #[test] + fn rolling_file_sink_rotates_without_mixing_bucket_contents() { + for rotation in [LogRotation::Hourly, LogRotation::Daily] { + let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4())); + let config = FileLoggingConfig::new(&dir, rotation, 7, 30); + let (sink, _) = RollingFileSink::new("runtime-test", config).expect("sink should open"); + let before = Local + .with_ymd_and_hms(2026, 4, 4, 23, 59, 59) + .single() + .expect("timestamp should build"); + let after = before + chrono::Duration::seconds(2); + + assert_eq!(sink.write_at(b"before\n", before).unwrap(), 7); + assert_eq!(sink.write_at(b"after\n", after).unwrap(), 6); + sink.flush().unwrap(); + for (instant, expected) in [(before, "before\n"), (after, "after\n")] { + let path = + bucketed_log_path(&dir, "runtime-test", &log_bucket_key(rotation, instant)); + assert_eq!(fs::read_to_string(&path).unwrap(), expected); + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt as _; + assert_eq!( + fs::metadata(path).unwrap().permissions().mode() & 0o777, + 0o600 + ); + } + } + drop(sink); + fs::remove_dir_all(&dir).unwrap(); + } + } + + #[cfg(unix)] + #[test] + fn failed_rotation_preserves_old_file_and_can_retry_safely() { + use std::os::unix::fs::symlink; + + let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4())); + let config = FileLoggingConfig::new(&dir, LogRotation::Daily, 7, 30); + let (sink, _) = RollingFileSink::new("runtime-test", config).unwrap(); + let before = Local + .with_ymd_and_hms(2026, 4, 4, 12, 0, 0) + .single() + .unwrap(); + let after = before + chrono::Duration::days(1); + let old_bucket = log_bucket_key(LogRotation::Daily, before); + let old_path = bucketed_log_path(&dir, "runtime-test", &old_bucket); + let new_path = bucketed_log_path( + &dir, + "runtime-test", + &log_bucket_key(LogRotation::Daily, after), + ); + let victim = dir.join("victim.txt"); + fs::write(&victim, b"unchanged").unwrap(); + sink.write_at(b"before\n", before).unwrap(); + symlink(&victim, &new_path).unwrap(); + + assert!(sink.write_at(b"rejected\n", after).is_err()); + assert_eq!(sink.state.lock().unwrap().current_bucket, old_bucket); + assert_eq!(fs::read(&victim).unwrap(), b"unchanged"); + assert_eq!(fs::read(&old_path).unwrap(), b"before\n"); + + fs::remove_file(&new_path).unwrap(); + sink.write_at(b"after\n", after).unwrap(); + sink.flush().unwrap(); + assert_eq!(fs::read(&new_path).unwrap(), b"after\n"); + assert_eq!(fs::read(&old_path).unwrap(), b"before\n"); + drop(sink); + fs::remove_dir_all(&dir).unwrap(); + } + #[test] fn cleanup_log_files_removes_matching_files_on_disk() { let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4())); diff --git a/crates/aether-runtime/base/src/tracing/writer.rs b/crates/aether-runtime/base/src/tracing/writer.rs new file mode 100644 index 000000000..47e8a9b6d --- /dev/null +++ b/crates/aether-runtime/base/src/tracing/writer.rs @@ -0,0 +1,1048 @@ +use std::io::{self, Write}; +use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}; +use std::sync::{Arc, Condvar, Mutex, OnceLock}; +use std::time::{Duration, Instant}; + +use tokio::sync::mpsc; +use tracing_subscriber::fmt::writer::MakeWriter; + +use crate::metrics::{MetricKind, MetricSample}; + +const QUEUE_CAPACITY: usize = 4096; +const BYTE_LIMIT: usize = 8 * 1024 * 1024; +const EVENT_LIMIT: usize = 256 * 1024; +const LOCAL_WORKER_DROP_TIMEOUT: Duration = Duration::from_millis(100); + +static LOG_WORKERS: OnceLock> = OnceLock::new(); + +#[derive(Clone, Copy)] +struct Limits { + queue_capacity: usize, + bytes: usize, + event_bytes: usize, +} + +impl Default for Limits { + fn default() -> Self { + Self { + queue_capacity: QUEUE_CAPACITY, + bytes: BYTE_LIMIT, + event_bytes: EVENT_LIMIT, + } + } +} + +struct State { + destination: &'static str, + limits: Limits, + accepting: AtomicBool, + running: AtomicBool, + retained_bytes: AtomicUsize, + retained_events: AtomicUsize, + accepted: AtomicU64, + dropped_full: AtomicU64, + dropped_bytes: AtomicU64, + dropped_oversize: AtomicU64, + dropped_closed: AtomicU64, + write_errors: AtomicU64, + worker_panics: AtomicU64, + shutdown_timeouts: AtomicU64, +} + +impl State { + fn new(destination: &'static str, limits: Limits) -> Self { + Self { + destination, + limits, + accepting: AtomicBool::new(true), + running: AtomicBool::new(true), + retained_bytes: AtomicUsize::new(0), + retained_events: AtomicUsize::new(0), + accepted: AtomicU64::new(0), + dropped_full: AtomicU64::new(0), + dropped_bytes: AtomicU64::new(0), + dropped_oversize: AtomicU64::new(0), + dropped_closed: AtomicU64::new(0), + write_errors: AtomicU64::new(0), + worker_panics: AtomicU64::new(0), + shutdown_timeouts: AtomicU64::new(0), + } + } + + fn reserve(self: &Arc, bytes: usize) -> Option { + self.retained_bytes + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |retained| { + retained + .checked_add(bytes) + .filter(|next| *next <= self.limits.bytes) + }) + .ok()?; + self.retained_events.fetch_add(1, Ordering::AcqRel); + Some(RetainedBytes { + state: Arc::clone(self), + bytes, + }) + } +} + +struct RetainedBytes { + state: Arc, + bytes: usize, +} + +impl Drop for RetainedBytes { + fn drop(&mut self) { + self.state + .retained_bytes + .fetch_sub(self.bytes, Ordering::AcqRel); + self.state.retained_events.fetch_sub(1, Ordering::AcqRel); + } +} + +struct BufferedEvent { + bytes: Box<[u8]>, + // Field order frees the allocation before making its budget available again. + _retained: RetainedBytes, +} + +enum Message { + Event(BufferedEvent), + Wake, +} + +#[derive(Default)] +struct Completion { + result: Mutex>, + changed: Condvar, +} + +impl Completion { + fn finish(&self, success: bool) { + *self + .result + .lock() + .unwrap_or_else(|error| error.into_inner()) = Some(success); + self.changed.notify_all(); + } + + fn wait(&self, timeout: Duration) -> Option { + let result = self + .result + .lock() + .unwrap_or_else(|error| error.into_inner()); + let (result, _) = self + .changed + .wait_timeout_while(result, timeout, |result| result.is_none()) + .unwrap_or_else(|error| error.into_inner()); + *result + } +} + +struct CompletionGuard { + state: Arc, + completion: Arc, + success: bool, +} + +impl Drop for CompletionGuard { + fn drop(&mut self) { + self.state.accepting.store(false, Ordering::Release); + self.state.running.store(false, Ordering::Release); + self.completion.finish(self.success); + } +} + +#[derive(Clone)] +pub(super) struct NonBlockingLogWriter { + sender: mpsc::Sender, + state: Arc, +} + +impl NonBlockingLogWriter { + pub(super) fn new( + destination: &'static str, + sink: impl Write + Send + 'static, + ) -> io::Result<(Self, LogWorker)> { + Self::with_limits(destination, sink, Limits::default()) + } + + fn with_limits( + destination: &'static str, + sink: impl Write + Send + 'static, + limits: Limits, + ) -> io::Result<(Self, LogWorker)> { + if limits.queue_capacity == 0 || limits.bytes == 0 || limits.event_bytes == 0 { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "log writer limits must be positive", + )); + } + let (sender, receiver) = mpsc::channel(limits.queue_capacity); + let state = Arc::new(State::new(destination, limits)); + let completion = Arc::new(Completion::default()); + let worker_state = Arc::clone(&state); + let worker_completion = Arc::clone(&completion); + // The thread owns blocking I/O. No join is attempted when a sink stalls. + std::thread::Builder::new() + .name(format!("aether-log-{destination}")) + .spawn(move || { + let mut completed = CompletionGuard { + state: Arc::clone(&worker_state), + completion: worker_completion, + success: false, + }; + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + run_worker(receiver, sink, &worker_state) + })); + completed.success = match result { + Ok(success) => success, + Err(_) => { + worker_state.worker_panics.fetch_add(1, Ordering::Relaxed); + false + } + }; + })?; + Ok(( + Self { + sender: sender.clone(), + state: Arc::clone(&state), + }, + LogWorker { + sender, + state, + completion, + }, + )) + } + + fn enqueue(&self, bytes: &[u8]) { + if bytes.is_empty() { + return; + } + if !self.state.accepting.load(Ordering::Acquire) { + self.state.dropped_closed.fetch_add(1, Ordering::Relaxed); + return; + } + if bytes.len() > self.state.limits.event_bytes { + self.state.dropped_oversize.fetch_add(1, Ordering::Relaxed); + return; + } + let Some(retained) = self.state.reserve(bytes.len()) else { + self.state.dropped_bytes.fetch_add(1, Ordering::Relaxed); + return; + }; + if !self.state.accepting.load(Ordering::Acquire) { + self.state.dropped_closed.fetch_add(1, Ordering::Relaxed); + return; + } + let event = BufferedEvent { + bytes: Box::from(bytes), + _retained: retained, + }; + self.enqueue_reserved(event); + } + + fn enqueue_reserved(&self, event: BufferedEvent) { + // Do not reserve channel slots: an unpublished reservation could prevent + // the shutdown Wake from reaching an otherwise idle receiver. + match self.sender.try_send(Message::Event(event)) { + Ok(()) => { + self.state.accepted.fetch_add(1, Ordering::Relaxed); + } + Err(mpsc::error::TrySendError::Full(_)) => { + self.state.dropped_full.fetch_add(1, Ordering::Relaxed); + } + Err(mpsc::error::TrySendError::Closed(_)) => { + self.state.dropped_closed.fetch_add(1, Ordering::Relaxed); + } + } + } +} + +impl Write for NonBlockingLogWriter { + fn write(&mut self, bytes: &[u8]) -> io::Result { + // The fmt layer supplies one complete formatted event to write_all. + // A rejection must not cause retries or synchronous stderr diagnostics. + self.enqueue(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + // A producer flush cannot wait for a blocked destination. Shutdown owns it. + Ok(()) + } +} + +impl<'a> MakeWriter<'a> for NonBlockingLogWriter { + type Writer = Self; + + fn make_writer(&'a self) -> Self::Writer { + self.clone() + } +} + +fn run_worker(mut receiver: mpsc::Receiver, mut sink: impl Write, state: &State) -> bool { + loop { + if !state.accepting.load(Ordering::Acquire) { + receiver.close(); + } + match receiver.blocking_recv() { + Some(Message::Event(event)) => { + if sink.write_all(&event.bytes).is_err() { + state.write_errors.fetch_add(1, Ordering::Relaxed); + } + } + Some(Message::Wake) => receiver.close(), + None => break, + } + } + match sink.flush() { + Ok(()) => true, + Err(_) => { + state.write_errors.fetch_add(1, Ordering::Relaxed); + false + } + } +} + +pub(super) struct LogWorker { + sender: mpsc::Sender, + state: Arc, + completion: Arc, +} + +impl LogWorker { + fn close(&self) { + self.state.accepting.store(false, Ordering::Release); + // A full queue already wakes the worker, which checks accepting each turn. + let _ = self.sender.try_send(Message::Wake); + } + + fn wait(&self, timeout: Duration) -> bool { + match self.completion.wait(timeout) { + Some(success) => success, + None => { + self.state.shutdown_timeouts.fetch_add(1, Ordering::Relaxed); + false + } + } + } +} + +impl Drop for LogWorker { + fn drop(&mut self) { + self.close(); + self.wait(LOCAL_WORKER_DROP_TIMEOUT); + } +} + +pub(super) fn register_log_workers(workers: Vec) { + // Failed or duplicate initialization drops only its own local workers. + let _ = LOG_WORKERS.set(workers); +} + +fn shutdown_workers(workers: &[LogWorker], timeout: Duration) -> bool { + let started = Instant::now(); + for worker in workers { + worker.close(); + } + let mut finished = true; + for worker in workers { + finished &= worker.wait(timeout.saturating_sub(started.elapsed())); + } + finished +} + +/// Stops accepting logs and waits within one total deadline for drain and flush. +/// A true result reports completed workers and a successful final flush; previous +/// write failures may already have lost events and remain in write error metrics. +/// This does not promise fsync durability or cancel a destination's blocked I/O. +/// Producers already preparing an event may release their reservations afterward. +pub fn shutdown_logging(timeout: Duration) -> bool { + LOG_WORKERS + .get() + .is_none_or(|workers| shutdown_workers(workers, timeout)) +} + +#[must_use = "keep the logging guard alive until service shutdown finishes"] +pub struct LogShutdownGuard; + +impl LogShutdownGuard { + pub fn new() -> Self { + Self + } +} + +impl Default for LogShutdownGuard { + fn default() -> Self { + Self::new() + } +} + +impl Drop for LogShutdownGuard { + fn drop(&mut self) { + shutdown_logging(Duration::from_secs(2)); + } +} + +#[derive(Default)] +struct Snapshot { + workers: u64, + byte_limit: u64, + event_limit: u64, + queue_capacity: u64, + retained_bytes: u64, + retained_events: u64, + accepted: u64, + dropped_full: u64, + dropped_bytes: u64, + dropped_oversize: u64, + dropped_closed: u64, + write_errors: u64, + worker_panics: u64, + shutdown_timeouts: u64, + accepting: u64, + running: u64, +} + +impl Snapshot { + fn add(&mut self, state: &State) { + self.workers += 1; + self.byte_limit += state.limits.bytes as u64; + self.event_limit = self.event_limit.max(state.limits.event_bytes as u64); + self.queue_capacity += state.limits.queue_capacity as u64; + self.retained_bytes += state.retained_bytes.load(Ordering::Relaxed) as u64; + self.retained_events += state.retained_events.load(Ordering::Relaxed) as u64; + self.accepted += state.accepted.load(Ordering::Relaxed); + self.dropped_full += state.dropped_full.load(Ordering::Relaxed); + self.dropped_bytes += state.dropped_bytes.load(Ordering::Relaxed); + self.dropped_oversize += state.dropped_oversize.load(Ordering::Relaxed); + self.dropped_closed += state.dropped_closed.load(Ordering::Relaxed); + self.write_errors += state.write_errors.load(Ordering::Relaxed); + self.worker_panics += state.worker_panics.load(Ordering::Relaxed); + self.shutdown_timeouts += state.shutdown_timeouts.load(Ordering::Relaxed); + self.accepting += u64::from(state.accepting.load(Ordering::Relaxed)); + self.running += u64::from(state.running.load(Ordering::Relaxed)); + } +} + +fn metric_samples(workers: &[LogWorker]) -> Vec { + let mut stdout = Snapshot::default(); + let mut file = Snapshot::default(); + let mut other = Snapshot::default(); + for worker in workers { + match worker.state.destination { + "stdout" => stdout.add(&worker.state), + "file" => file.add(&worker.state), + _ => other.add(&worker.state), + } + } + let mut samples = Vec::new(); + macro_rules! append { + ($prefix:literal, $snapshot:ident) => { + if $snapshot.workers > 0 { + macro_rules! metric { + ($field:ident, $suffix:literal, $help:literal, $kind:ident) => { + samples.push(MetricSample::new( + concat!($prefix, $suffix), + $help, + MetricKind::$kind, + $snapshot.$field, + )); + }; + } + metric!( + byte_limit, + "_byte_limit", + "Log buffer byte limit, including writes in progress.", + Gauge + ); + metric!( + event_limit, + "_event_byte_limit", + "Maximum bytes accepted in one formatted log event.", + Gauge + ); + metric!( + queue_capacity, + "_queue_capacity", + "Maximum queued log events excluding the active write.", + Gauge + ); + metric!( + retained_bytes, + "_retained_bytes", + "Owned log bytes including producer reservations and active writes.", + Gauge + ); + metric!( + retained_events, + "_retained_events", + "Owned log events including producer reservations and active writes.", + Gauge + ); + metric!( + accepted, + "_accepted_total", + "Log events accepted for background writing.", + Counter + ); + metric!( + dropped_full, + "_dropped_full_total", + "Log events dropped because the queue was full.", + Counter + ); + metric!( + dropped_bytes, + "_dropped_bytes_total", + "Log events dropped because the byte budget was full.", + Counter + ); + metric!( + dropped_oversize, + "_dropped_oversize_total", + "Log events exceeding the per-event byte limit.", + Counter + ); + metric!( + dropped_closed, + "_dropped_closed_total", + "Log events dropped because the writer was closing or closed.", + Counter + ); + metric!( + write_errors, + "_write_errors_total", + "Background log write or flush errors.", + Counter + ); + metric!( + worker_panics, + "_worker_panics_total", + "Background log worker panics.", + Counter + ); + metric!( + shutdown_timeouts, + "_shutdown_timeouts_total", + "Log shutdown waits that exceeded their deadline.", + Counter + ); + metric!( + accepting, + "_accepting_workers", + "Log workers accepting new events.", + Gauge + ); + metric!( + running, + "_running_workers", + "Log workers not yet finished draining and flushing.", + Gauge + ); + } + }; + } + append!("logging_stdout", stdout); + append!("logging_file", file); + append!("logging_other", other); + samples +} + +pub fn logging_metric_samples() -> Vec { + LOG_WORKERS + .get() + .map_or_else(Vec::new, |workers| metric_samples(workers)) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::mpsc as std_mpsc; + use tracing_subscriber::prelude::*; + + use super::super::{ + JsonRuntimeEventFormatter, PrettyRuntimeEventFormatter, RuntimeLogIdentity, + }; + + #[derive(Clone, Default)] + struct Buffer { + events: Arc>>>, + flushes: Arc, + } + + impl Write for Buffer { + fn write(&mut self, bytes: &[u8]) -> io::Result { + self.events.lock().unwrap().push(bytes.to_vec()); + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + self.flushes.fetch_add(1, Ordering::SeqCst); + Ok(()) + } + } + + struct Release(Arc<(Mutex, Condvar)>); + + impl Release { + fn open(&self) { + *self.0 .0.lock().unwrap() = true; + self.0 .1.notify_all(); + } + } + + impl Drop for Release { + fn drop(&mut self) { + self.open(); + } + } + + struct BlockedSink { + release: Arc<(Mutex, Condvar)>, + entered: Option>, + buffer: Buffer, + } + + impl Write for BlockedSink { + fn write(&mut self, bytes: &[u8]) -> io::Result { + if let Some(entered) = self.entered.take() { + entered.send(()).unwrap(); + } + let guard = self.release.0.lock().unwrap(); + drop(self.release.1.wait_while(guard, |open| !*open).unwrap()); + self.buffer.write(bytes) + } + + fn flush(&mut self) -> io::Result<()> { + self.buffer.flush() + } + } + + fn blocked_sink() -> (BlockedSink, Release, std_mpsc::Receiver<()>, Buffer) { + let release = Arc::new((Mutex::new(false), Condvar::new())); + let (entered, observed) = std_mpsc::channel(); + let buffer = Buffer::default(); + ( + BlockedSink { + release: Arc::clone(&release), + entered: Some(entered), + buffer: buffer.clone(), + }, + Release(release), + observed, + buffer, + ) + } + + fn assert_released(state: &State) { + assert_eq!(state.retained_bytes.load(Ordering::Acquire), 0); + assert_eq!(state.retained_events.load(Ordering::Acquire), 0); + assert!(!state.running.load(Ordering::Acquire)); + } + + #[test] + fn nonblocking_log_writer_slow_sink_does_not_block_concurrent_producers() { + let (sink, release, entered, buffer) = blocked_sink(); + let (mut writer, worker) = NonBlockingLogWriter::with_limits( + "stdout", + sink, + Limits { + queue_capacity: 4, + bytes: 1024, + event_bytes: 8, + }, + ) + .unwrap(); + writer.write_all(b"first\n").unwrap(); + entered.recv_timeout(Duration::from_secs(2)).unwrap(); + let (finished, completed) = std_mpsc::channel(); + let threads = (0..8) + .map(|_| { + let mut writer = writer.clone(); + let finished = finished.clone(); + std::thread::spawn(move || { + for _ in 0..100 { + writer.write_all(b"event\n").unwrap(); + } + finished.send(()).unwrap(); + }) + }) + .collect::>(); + for _ in 0..8 { + completed + .recv_timeout(Duration::from_secs(2)) + .expect("producers must finish while sink is blocked"); + } + for thread in threads { + thread.join().unwrap(); + } + assert_eq!(writer.state.retained_bytes.load(Ordering::Acquire), 30); + assert_eq!(writer.state.accepted.load(Ordering::Relaxed), 5); + assert_eq!(writer.state.dropped_full.load(Ordering::Relaxed), 796); + release.open(); + assert!(shutdown_workers(&[worker], Duration::from_secs(2))); + assert_eq!(buffer.events.lock().unwrap().len(), 5); + assert_released(&writer.state); + } + + #[test] + fn nonblocking_log_writer_byte_limit_includes_active_write_and_rejects_whole_events() { + let (sink, release, entered, buffer) = blocked_sink(); + let (mut writer, worker) = NonBlockingLogWriter::with_limits( + "file", + sink, + Limits { + queue_capacity: 8, + bytes: 12, + event_bytes: 8, + }, + ) + .unwrap(); + writer.write_all(b"12345678").unwrap(); + entered.recv_timeout(Duration::from_secs(2)).unwrap(); + writer.write_all(b"12345").unwrap(); + assert_eq!(writer.state.dropped_bytes.load(Ordering::Relaxed), 1); + writer.write_all(b"123456789").unwrap(); + assert_eq!(writer.state.dropped_oversize.load(Ordering::Relaxed), 1); + writer.write_all(b"1234").unwrap(); + assert_eq!(writer.state.retained_bytes.load(Ordering::Acquire), 12); + release.open(); + assert!(shutdown_workers(&[worker], Duration::from_secs(2))); + assert_eq!( + *buffer.events.lock().unwrap(), + [b"12345678".to_vec(), b"1234".to_vec()] + ); + assert_released(&writer.state); + writer.write_all(b"closed").unwrap(); + assert_eq!(writer.state.dropped_closed.load(Ordering::Relaxed), 1); + } + + #[test] + fn nonblocking_log_writer_shutdown_rejects_a_producer_still_preparing_an_event() { + let buffer = Buffer::default(); + let (writer, worker) = NonBlockingLogWriter::new("stdout", buffer.clone()).unwrap(); + let retained = writer.state.reserve(6).unwrap(); + worker.close(); + assert!(worker.wait(Duration::from_secs(2))); + assert_eq!(writer.state.retained_bytes.load(Ordering::Acquire), 6); + writer.enqueue_reserved(BufferedEvent { + bytes: Box::from(&b"raced\n"[..]), + _retained: retained, + }); + assert!(buffer.events.lock().unwrap().is_empty()); + assert_eq!(writer.state.accepted.load(Ordering::Relaxed), 0); + assert_eq!(writer.state.dropped_closed.load(Ordering::Relaxed), 1); + assert_released(&writer.state); + } + + #[test] + fn nonblocking_log_writer_concurrent_shutdown_drains_every_accepted_event() { + let buffer = Buffer::default(); + let (mut writer, worker) = NonBlockingLogWriter::new("stdout", buffer.clone()).unwrap(); + writer.write_all(b"before\n").unwrap(); + let barrier = Arc::new(std::sync::Barrier::new(9)); + let producers = (0..8) + .map(|_| { + let mut writer = writer.clone(); + let barrier = Arc::clone(&barrier); + std::thread::spawn(move || { + barrier.wait(); + for _ in 0..128 { + writer.write_all(b"concurrent\n").unwrap(); + } + }) + }) + .collect::>(); + barrier.wait(); + worker.close(); + for producer in producers { + producer.join().unwrap(); + } + assert!(worker.wait(Duration::from_secs(2))); + let accepted = writer.state.accepted.load(Ordering::Relaxed); + let dropped = writer.state.dropped_closed.load(Ordering::Relaxed); + assert_eq!(accepted + dropped, 1 + 8 * 128); + assert_eq!(buffer.events.lock().unwrap().len() as u64, accepted); + assert_released(&writer.state); + } + + #[test] + fn nonblocking_log_writer_real_formatters_preserve_concurrent_event_boundaries() { + for json in [false, true] { + let buffer = Buffer::default(); + let (writer, worker) = NonBlockingLogWriter::new("stdout", buffer.clone()).unwrap(); + let identity = RuntimeLogIdentity { + service: "concurrent-log-test", + node_role: Some("gateway".to_owned()), + instance_id: Some("test-1".to_owned()), + }; + let dispatch = if json { + tracing::Dispatch::new( + tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .json() + .with_writer(writer.clone()) + .event_format(JsonRuntimeEventFormatter::new(identity)), + ), + ) + } else { + tracing::Dispatch::new( + tracing_subscriber::registry().with( + tracing_subscriber::fmt::layer() + .with_writer(writer.clone()) + .event_format(PrettyRuntimeEventFormatter::new(identity, false)), + ), + ) + }; + let producers = (0..8_u64) + .map(|producer| { + let dispatch = dispatch.clone(); + std::thread::spawn(move || { + tracing::dispatcher::with_default(&dispatch, || { + for sequence in 0..64_u64 { + let event_id = producer * 64 + sequence; + tracing::info!( + target: "concurrent_log_test", + event_id, + enabled = true, + ratio = 1.5_f64, + note = "first\nsecond", + "event-{event_id:03} \"quoted\" \\" + ); + } + }); + }) + }) + .collect::>(); + for producer in producers { + producer.join().unwrap(); + } + assert!(shutdown_workers(&[worker], Duration::from_secs(2))); + let events = buffer.events.lock().unwrap(); + assert_eq!(events.len(), 512, "each write must contain one whole event"); + let mut observed = std::collections::BTreeSet::new(); + for event in events.iter() { + let text = std::str::from_utf8(event).unwrap(); + assert_eq!(text.lines().count(), 1, "no partial or merged records"); + let event_id = if json { + let payload: serde_json::Value = serde_json::from_str(text).unwrap(); + assert_eq!(payload["service"], "concurrent-log-test"); + assert_eq!(payload["node_role"], "gateway"); + assert_eq!(payload["instance_id"], "test-1"); + assert_eq!(payload["fields"]["enabled"], true); + assert_eq!(payload["fields"]["ratio"], 1.5); + assert_eq!(payload["fields"]["note"], "first\nsecond"); + let event_id = payload["fields"]["event_id"].as_u64().unwrap(); + assert_eq!( + payload["fields"]["message"], + format!("event-{event_id:03} \"quoted\" \\") + ); + event_id + } else { + assert!(!text.contains('\u{1b}')); + assert!(text.contains("enabled=true ratio=1.5 note=\"first\\nsecond\"")); + let event_id = text + .split_once("event_id=") + .unwrap() + .1 + .split_whitespace() + .next() + .unwrap() + .parse::() + .unwrap(); + assert!(text.contains(&format!(" - event-{event_id:03} \"quoted\" \\"))); + event_id + }; + assert!(observed.insert(event_id), "duplicate event {event_id}"); + } + assert_eq!(observed, (0..512).collect()); + assert_eq!(writer.state.accepted.load(Ordering::Relaxed), 512); + assert_released(&writer.state); + } + } + + #[test] + fn nonblocking_log_writer_shutdown_closes_all_sinks_before_waiting() { + for (slow_destination, healthy_destination) in [("stdout", "file"), ("file", "stdout")] { + let (sink, release, entered, _) = blocked_sink(); + let (mut slow, slow_worker) = + NonBlockingLogWriter::new(slow_destination, sink).unwrap(); + let buffer = Buffer::default(); + let (mut healthy, healthy_worker) = + NonBlockingLogWriter::new(healthy_destination, buffer.clone()).unwrap(); + slow.write_all(b"stalled\n").unwrap(); + entered.recv_timeout(Duration::from_secs(2)).unwrap(); + healthy.write_all(b"healthy\n").unwrap(); + let workers = [slow_worker, healthy_worker]; + assert!(!shutdown_workers(&workers, Duration::from_millis(50))); + assert!(!healthy.state.accepting.load(Ordering::Acquire)); + assert!(workers[1].wait(Duration::from_secs(2))); + assert_eq!(*buffer.events.lock().unwrap(), [b"healthy\n".to_vec()]); + assert_eq!(buffer.flushes.load(Ordering::SeqCst), 1); + assert_eq!(slow.state.shutdown_timeouts.load(Ordering::Relaxed), 1); + release.open(); + assert!(shutdown_workers(&workers, Duration::from_secs(2))); + assert_released(&slow.state); + assert_released(&healthy.state); + } + } + + #[test] + fn nonblocking_log_writer_blocked_flush_respects_deadline_and_other_sink_drains() { + struct BlockedFlushSink { + release: Arc<(Mutex, Condvar)>, + entered: std_mpsc::Sender<()>, + buffer: Buffer, + } + + impl Write for BlockedFlushSink { + fn write(&mut self, bytes: &[u8]) -> io::Result { + self.buffer.write(bytes) + } + + fn flush(&mut self) -> io::Result<()> { + self.entered.send(()).unwrap(); + let guard = self.release.0.lock().unwrap(); + drop(self.release.1.wait_while(guard, |open| !*open).unwrap()); + self.buffer.flush() + } + } + + let release = Release(Arc::new((Mutex::new(false), Condvar::new()))); + let (entered, observed) = std_mpsc::channel(); + let flushed_buffer = Buffer::default(); + let (mut flushing, flushing_worker) = NonBlockingLogWriter::new( + "file", + BlockedFlushSink { + release: Arc::clone(&release.0), + entered, + buffer: flushed_buffer.clone(), + }, + ) + .unwrap(); + let healthy_buffer = Buffer::default(); + let (mut healthy, healthy_worker) = + NonBlockingLogWriter::new("stdout", healthy_buffer.clone()).unwrap(); + flushing.write_all(b"before flush\n").unwrap(); + flushing_worker.close(); + observed.recv_timeout(Duration::from_secs(2)).unwrap(); + assert_eq!(flushed_buffer.flushes.load(Ordering::SeqCst), 0); + assert_eq!(flushing.state.retained_bytes.load(Ordering::Acquire), 0); + healthy.write_all(b"while other sink flushes\n").unwrap(); + let (finished, completion) = std_mpsc::channel(); + let shutdown = std::thread::spawn(move || { + let workers = [flushing_worker, healthy_worker]; + let success = shutdown_workers(&workers, Duration::from_millis(50)); + assert!(finished.send((success, workers)).is_ok()); + }); + let (success, workers) = completion + .recv_timeout(Duration::from_secs(2)) + .expect("shutdown must return while the destination is still flushing"); + shutdown.join().unwrap(); + assert!(!success); + assert!(workers[1].wait(Duration::from_secs(2))); + assert_eq!( + *healthy_buffer.events.lock().unwrap(), + [b"while other sink flushes\n".to_vec()] + ); + assert_eq!(healthy_buffer.flushes.load(Ordering::SeqCst), 1); + assert_eq!(flushing.state.shutdown_timeouts.load(Ordering::Relaxed), 1); + assert!(flushing.state.running.load(Ordering::Acquire)); + release.open(); + assert!(shutdown_workers(&workers, Duration::from_secs(2))); + assert_eq!(flushed_buffer.flushes.load(Ordering::SeqCst), 1); + assert_released(&flushing.state); + assert_released(&healthy.state); + } + + struct FailingSink { + fail_flush: bool, + } + + impl Write for FailingSink { + fn write(&mut self, _bytes: &[u8]) -> io::Result { + Err(io::Error::other("write failed")) + } + fn flush(&mut self) -> io::Result<()> { + if self.fail_flush { + Err(io::Error::other("flush failed")) + } else { + Ok(()) + } + } + } + + #[test] + fn nonblocking_log_writer_io_errors_do_not_block_producers_or_leak_budget() { + for fail_flush in [false, true] { + let (mut writer, worker) = + NonBlockingLogWriter::new("file", FailingSink { fail_flush }).unwrap(); + writer.write_all(b"one\n").unwrap(); + writer.write_all(b"two\n").unwrap(); + assert_eq!( + shutdown_workers(&[worker], Duration::from_secs(2)), + !fail_flush + ); + assert_eq!( + writer.state.write_errors.load(Ordering::Relaxed), + 2 + u64::from(fail_flush) + ); + assert_released(&writer.state); + } + } + + #[test] + fn nonblocking_log_writer_panic_releases_pending_budget_and_marks_worker_finished() { + struct PanicSink; + impl Write for PanicSink { + fn write(&mut self, _bytes: &[u8]) -> io::Result { + panic!("sink panic") + } + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + let (mut writer, worker) = NonBlockingLogWriter::new("file", PanicSink).unwrap(); + for _ in 0..8 { + writer.write_all(b"event\n").unwrap(); + } + assert!(!shutdown_workers(&[worker], Duration::from_secs(2))); + assert_eq!(writer.state.worker_panics.load(Ordering::Relaxed), 1); + assert_released(&writer.state); + } + + #[test] + fn nonblocking_log_writer_failed_initialization_drop_wakes_and_flushes_idle_worker() { + let buffer = Buffer::default(); + let (writer, worker) = NonBlockingLogWriter::new("stdout", buffer.clone()).unwrap(); + let completion = Arc::clone(&worker.completion); + drop(worker); + assert_eq!(completion.wait(Duration::from_secs(2)), Some(true)); + assert_eq!(buffer.flushes.load(Ordering::SeqCst), 1); + assert_released(&writer.state); + } + + #[test] + fn nonblocking_log_writer_metrics_use_unique_destination_names() { + let (_, stdout) = NonBlockingLogWriter::new("stdout", io::sink()).unwrap(); + let (_, file) = NonBlockingLogWriter::new("file", io::sink()).unwrap(); + let workers = [stdout, file]; + let samples = metric_samples(&workers); + let names = samples + .iter() + .map(|sample| sample.name) + .collect::>(); + assert_eq!(names.len(), samples.len()); + assert!(samples + .iter() + .any(|sample| sample.name == "logging_stdout_byte_limit" + && sample.value == BYTE_LIMIT as u64)); + assert!(samples + .iter() + .any(|sample| sample.name == "logging_file_byte_limit" + && sample.value == BYTE_LIMIT as u64)); + assert!(shutdown_workers(&workers, Duration::from_secs(2))); + } +} diff --git a/crates/aether-runtime/base/tests/blocked_stdout_logging.rs b/crates/aether-runtime/base/tests/blocked_stdout_logging.rs new file mode 100644 index 000000000..2d206abe5 --- /dev/null +++ b/crates/aether-runtime/base/tests/blocked_stdout_logging.rs @@ -0,0 +1,253 @@ +use std::fs; +use std::io::{self, Read}; +use std::path::{Path, PathBuf}; +use std::process::{Child, Command, ExitStatus, Stdio}; +use std::thread::JoinHandle; +use std::time::{Duration, Instant}; + +use aether_runtime::{ + init_service_runtime, logging_metric_samples, FileLoggingConfig, LogDestination, LogFormat, + LogRotation, LogShutdownGuard, ServiceRuntimeConfig, +}; + +const CHILD_ENV: &str = "AETHER_TEST_BLOCKED_STDOUT_CHILD"; +const DIRECTORY_ENV: &str = "AETHER_TEST_BLOCKED_STDOUT_DIR"; +const SERVICE_NAME: &str = "blocked-stdout-test"; +const EVENT_NAME: &str = "blocked_stdout_probe"; +const TEST_NAME: &str = "blocked_stdout_does_not_block_file_logs_or_process_exit"; + +#[test] +fn blocked_stdout_does_not_block_file_logs_or_process_exit() { + if std::env::var_os(CHILD_ENV).is_some() { + run_child_scenario(); + eprintln!("blocked stdout guard returned"); + // The scenario returns normally and drops its guard. Skip libtest's own + // stdout report, while still exercising Rust's standard exit cleanup. + std::process::exit(0); + } + + let directory = TestDirectory::new(); + let mut command = Command::new(std::env::current_exe().expect("test executable")); + command + .args(["--exact", TEST_NAME, "--nocapture", "--quiet"]) + .env(CHILD_ENV, "1") + .env(DIRECTORY_ENV, &directory.0) + .env_remove("RUST_LOG") + .env_remove("NO_COLOR") + .env_remove("FORCE_COLOR"); + let mut child = BlockedStdoutChild::spawn(&mut command).expect("logging child should start"); + let (status, timed_out, stderr) = child + .wait(Duration::from_secs(8)) + .expect("logging child should be reaped"); + let stderr = String::from_utf8_lossy(&stderr); + assert!( + !timed_out, + "blocked stdout prevented process exit within 8 seconds: {stderr}" + ); + assert!(status.success(), "logging child failed: {status}: {stderr}"); + assert!( + stderr.contains("blocked stdout saturated") + && stderr.contains("blocked stdout guard returned"), + "child did not reach saturation and return from its guard: {stderr}" + ); + + let records = read_file_records(&directory.0); + assert!( + !records.is_empty(), + "healthy file destination received no logs" + ); + let final_markers: Vec<_> = records + .iter() + .filter(|record| record["fields"]["phase"] == "final") + .collect(); + assert_eq!( + final_markers.len(), + 1, + "file lost or duplicated final marker" + ); + let final_marker = final_markers[0]; + assert_eq!(final_marker["fields"]["event_name"], EVENT_NAME); + assert!( + final_marker["fields"]["stdout_dropped_full"] + .as_u64() + .expect("stdout queue drop counter") + + final_marker["fields"]["stdout_dropped_bytes"] + .as_u64() + .expect("stdout byte drop counter") + > 0, + "file marker must prove stdout saturation" + ); + assert_eq!(records.last().unwrap()["fields"]["phase"], "final"); +} + +fn run_child_scenario() { + let _shutdown = LogShutdownGuard::new(); + let directory = PathBuf::from(std::env::var_os(DIRECTORY_ENV).expect("child log directory")); + init_service_runtime( + ServiceRuntimeConfig::new(SERVICE_NAME, "info") + .with_log_destination(LogDestination::Both) + .with_log_format(LogFormat::Json) + .with_file_logging(FileLoggingConfig::new(directory, LogRotation::Daily, 7, 30)), + ) + .expect("both log destinations should initialize"); + + let payload = "x".repeat(384); + let flood_deadline = Instant::now() + Duration::from_secs(3); + let mut emitted = 0; + while emitted < 20_000 && Instant::now() < flood_deadline { + tracing::info!( + event_name = EVENT_NAME, + phase = "flood", + sequence = emitted, + payload = %payload, + "fill unread stdout" + ); + emitted += 1; + if emitted % 64 == 0 && stdout_dropped_events() > 0 { + break; + } + } + assert!(stdout_dropped_events() > 0, "stdout queue did not saturate"); + + let file_deadline = Instant::now() + Duration::from_secs(1); + while metric("logging_file_retained_bytes") > 0 && Instant::now() < file_deadline { + std::thread::sleep(Duration::from_millis(5)); + } + assert_eq!( + metric("logging_file_retained_bytes"), + 0, + "file writer stalled" + ); + assert_eq!(metric("logging_file_write_errors_total"), 0); + let dropped_full = metric("logging_stdout_dropped_full_total"); + let dropped_bytes = metric("logging_stdout_dropped_bytes_total"); + eprintln!( + "blocked stdout saturated: full={dropped_full} bytes={dropped_bytes} emitted={emitted}" + ); + tracing::info!( + event_name = EVENT_NAME, + phase = "final", + stdout_dropped_full = dropped_full, + stdout_dropped_bytes = dropped_bytes, + "healthy file final marker" + ); +} + +fn metric(name: &str) -> u64 { + logging_metric_samples() + .into_iter() + .find(|sample| sample.name == name) + .unwrap_or_else(|| panic!("missing logging metric: {name}")) + .value +} + +fn stdout_dropped_events() -> u64 { + metric("logging_stdout_dropped_full_total") + metric("logging_stdout_dropped_bytes_total") +} + +fn read_file_records(directory: &Path) -> Vec { + let mut paths: Vec<_> = fs::read_dir(directory) + .expect("log directory") + .map(|entry| entry.expect("log entry").path()) + .filter(|path| { + path.file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| name.starts_with(SERVICE_NAME) && name.ends_with(".log")) + }) + .collect(); + paths.sort(); + let mut records = Vec::new(); + for path in paths { + let contents = fs::read_to_string(path).expect("UTF-8 file logs"); + assert!(contents.ends_with('\n'), "partial final file record"); + for line in contents.lines() { + assert!(line.len() < 1024, "test event exceeded 1 KiB"); + records.push(serde_json::from_str(line).expect("complete JSON file record")); + } + } + records +} + +struct BlockedStdoutChild { + child: Child, + stderr_reader: Option>>>, +} + +impl BlockedStdoutChild { + fn spawn(command: &mut Command) -> io::Result { + let child = command + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn()?; + let mut guarded = Self { + child, + stderr_reader: None, + }; + let mut stderr = guarded.child.stderr.take().expect("child stderr pipe"); + guarded.stderr_reader = Some(std::thread::Builder::new().spawn(move || { + let mut captured = Vec::new(); + let mut chunk = [0u8; 1024]; + loop { + let count = stderr.read(&mut chunk)?; + if count == 0 { + return Ok(captured); + } + let retained = count.min((16 * 1024usize).saturating_sub(captured.len())); + captured.extend_from_slice(&chunk[..retained]); + } + })?); + Ok(guarded) + } + + fn wait(&mut self, timeout: Duration) -> io::Result<(ExitStatus, bool, Vec)> { + let deadline = Instant::now() + timeout; + let (status, timed_out) = loop { + if let Some(status) = self.child.try_wait()? { + break (status, false); + } + if Instant::now() >= deadline { + let _ = self.child.kill(); + break (self.child.wait()?, true); + } + std::thread::sleep(Duration::from_millis(10)); + }; + // Keep the stdout read end open and completely unread until the child + // has exited or been killed. Closing it earlier would unblock writes. + drop(self.child.stdout.take()); + let stderr = self + .stderr_reader + .take() + .expect("stderr reader") + .join() + .map_err(|_| io::Error::other("stderr reader panicked"))??; + Ok((status, timed_out, stderr)) + } +} + +impl Drop for BlockedStdoutChild { + fn drop(&mut self) { + let _ = self.child.kill(); + let _ = self.child.wait(); + drop(self.child.stdout.take()); + if let Some(reader) = self.stderr_reader.take() { + let _ = reader.join(); + } + } +} + +struct TestDirectory(PathBuf); + +impl TestDirectory { + fn new() -> Self { + let path = + std::env::temp_dir().join(format!("aether-blocked-stdout-{}", uuid::Uuid::new_v4())); + fs::create_dir(&path).expect("test directory"); + Self(path) + } +} + +impl Drop for TestDirectory { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.0); + } +} diff --git a/crates/aether-runtime/base/tests/nonblocking_logging.rs b/crates/aether-runtime/base/tests/nonblocking_logging.rs new file mode 100644 index 000000000..7f754e51d --- /dev/null +++ b/crates/aether-runtime/base/tests/nonblocking_logging.rs @@ -0,0 +1,354 @@ +use std::fs; +use std::io::Read; +use std::path::{Path, PathBuf}; +use std::process::{Command, Output, Stdio}; +use std::time::{Duration, Instant}; + +use aether_runtime::{ + init_reloadable_service_tracing, init_service_runtime, FileLoggingConfig, LogDestination, + LogFormat, LogRotation, LogShutdownGuard, ServiceRuntimeConfig, +}; + +const CASE_ENV: &str = "AETHER_TEST_NONBLOCKING_LOGGING_CASE"; +const DIRECTORY_ENV: &str = "AETHER_TEST_NONBLOCKING_LOGGING_DIR"; +const EVENT_NAME: &str = "nonblocking_logging_probe"; +const SERVICE_NAME: &str = "nonblocking-logging-test"; +const RECORD_COUNT: usize = 32; +const FIELD_VALUE: &str = "quote\" newline\n backslash\\ \u{4e2d}\u{6587}"; + +#[test] +fn nonblocking_logging_entrypoints_reload_and_guard_drain() { + if let Ok(scenario) = std::env::var(CASE_ENV) { + run_scenario(&scenario); + return; + } + + let root = TestDirectory::new(); + for entrypoint in ["standard", "reloadable"] { + for destination in ["stdout", "file", "both"] { + for format in ["pretty", "json"] { + let scenario = format!("{entrypoint}-{destination}-{format}"); + let directory = root.0.join(&scenario); + fs::create_dir(&directory).expect("scenario directory"); + let mut command = Command::new(std::env::current_exe().expect("test executable")); + command + .args([ + "--exact", + "nonblocking_logging_entrypoints_reload_and_guard_drain", + "--nocapture", + "--quiet", + ]) + .env(CASE_ENV, &scenario) + .env(DIRECTORY_ENV, &directory) + .env_remove("RUST_LOG") + .env_remove("NO_COLOR") + .env_remove("FORCE_COLOR"); + let output = run_subprocess(&mut command); + let stdout = String::from_utf8(output.stdout).expect("UTF-8 stdout"); + let stderr = String::from_utf8(output.stderr).expect("UTF-8 stderr"); + assert!(output.status.success(), "{scenario}: {stdout}\n{stderr}"); + let file_output = read_log_files(&directory); + verify_output( + &stdout, + entrypoint, + format, + destination != "file", + &scenario, + ); + verify_output( + &file_output, + entrypoint, + format, + destination != "stdout", + &scenario, + ); + } + } + } +} + +fn run_scenario(scenario: &str) { + let parts: Vec<_> = scenario.split('-').collect(); + let [entrypoint, destination, format] = parts.as_slice() else { + panic!("invalid logging scenario: {scenario}"); + }; + let _shutdown = LogShutdownGuard::new(); + let destination = match *destination { + "stdout" => LogDestination::Stdout, + "file" => LogDestination::File, + "both" => LogDestination::Both, + other => panic!("unknown destination: {other}"), + }; + let mut config = ServiceRuntimeConfig::new(SERVICE_NAME, "info") + .with_node_role("integration") + .with_instance_id("logging-child") + .with_log_destination(destination) + .with_log_format(match *format { + "pretty" => LogFormat::Pretty, + "json" => LogFormat::Json, + other => panic!("unknown format: {other}"), + }); + if matches!(destination, LogDestination::File | LogDestination::Both) { + config = config.with_file_logging(FileLoggingConfig::new( + PathBuf::from(std::env::var_os(DIRECTORY_ENV).expect("scenario log directory")), + LogRotation::Daily, + 7, + 30, + )); + } + let reload = match *entrypoint { + "standard" => { + init_service_runtime(config).expect("standard logging initializes"); + None + } + "reloadable" => Some( + init_reloadable_service_tracing("info", config) + .expect("reloadable logging initializes"), + ), + other => panic!("unknown entrypoint: {other}"), + }; + + tracing::debug!( + event_name = EVENT_NAME, + phase = "initial_hidden", + "filtered debug" + ); + tracing::info!(event_name = EVENT_NAME, phase = "initial", "initial event"); + if let Some(reload) = reload { + reload("debug"); + tracing::debug!( + event_name = EVENT_NAME, + phase = "reloaded_debug", + "visible debug" + ); + let invalid_filter = "nonblocking_logging=not-a-level"; + assert!(tracing_subscriber::EnvFilter::try_new(invalid_filter).is_err()); + reload(invalid_filter); + tracing::debug!( + event_name = EVENT_NAME, + phase = "invalid_reload_unchanged", + "still debug" + ); + reload("error"); + tracing::info!( + event_name = EVENT_NAME, + phase = "error_filter_hidden", + "filtered info" + ); + tracing::error!( + event_name = EVENT_NAME, + phase = "reloaded_error", + "visible error" + ); + reload("info"); + } + for sequence in 0..RECORD_COUNT { + tracing::info!( + event_name = EVENT_NAME, + phase = "record", + sequence = sequence as u64, + value = FIELD_VALUE, + "complete record" + ); + } + tracing::info!( + event_name = EVENT_NAME, + phase = "tail", + "final event before guard drop" + ); + // Returning drops the guard. The parent verifies the tail after process exit. +} + +fn verify_output(output: &str, entrypoint: &str, format: &str, enabled: bool, scenario: &str) { + let lines: Vec<_> = output + .lines() + .filter(|line| line.contains(EVENT_NAME)) + .collect(); + if !enabled { + assert!( + lines.is_empty(), + "unexpected destination output in {scenario}: {output}" + ); + return; + } + let expected_count = RECORD_COUNT + 2 + usize::from(entrypoint == "reloadable") * 3; + assert_eq!( + lines.len(), + expected_count, + "missing or duplicate records in {scenario}: {output}" + ); + assert!( + !output.contains('\u{1b}'), + "redirected/file output must not contain ANSI: {scenario}" + ); + assert!( + !output.contains("initial_hidden"), + "initial filter failed: {scenario}" + ); + assert!( + !output.contains("error_filter_hidden"), + "reloaded filter failed: {scenario}" + ); + let mut phases = Vec::new(); + let mut sequences = Vec::new(); + for line in lines { + if format == "json" { + let record: serde_json::Value = serde_json::from_str(line) + .unwrap_or_else(|error| panic!("incomplete JSON in {scenario}: {error}: {line}")); + assert_eq!(record["service"], SERVICE_NAME); + assert_eq!(record["node_role"], "integration"); + assert_eq!(record["instance_id"], "logging-child"); + let phase = record["fields"]["phase"].as_str().expect("event phase"); + phases.push(phase.to_string()); + if phase == "record" { + assert_eq!(record["fields"]["value"], FIELD_VALUE); + sequences.push( + record["fields"]["sequence"] + .as_u64() + .expect("record sequence"), + ); + } + } else { + assert!( + line.contains(" | INFO") || line.contains(" | DEBUG") || line.contains(" | ERROR"), + "incomplete Pretty record in {scenario}: {line}" + ); + let phase = [ + "initial", + "reloaded_debug", + "invalid_reload_unchanged", + "reloaded_error", + "record", + "tail", + ] + .into_iter() + .find(|phase| line.contains(&format!("phase=\"{phase}\""))) + .expect("complete Pretty phase field"); + phases.push(phase.to_string()); + if phase == "record" { + let expected_value = format!("value={FIELD_VALUE:?}"); + assert!( + line.contains(&expected_value), + "incomplete Pretty value in {scenario}: {line}" + ); + let sequence = line + .split_whitespace() + .find_map(|field| field.strip_prefix("sequence=")) + .expect("complete Pretty sequence field") + .parse::() + .expect("sequence number"); + sequences.push(sequence); + } + } + } + assert_eq!(phases.first().map(String::as_str), Some("initial")); + assert_eq!( + phases.last().map(String::as_str), + Some("tail"), + "guard lost tail event: {scenario}" + ); + for phase in [ + "reloaded_debug", + "invalid_reload_unchanged", + "reloaded_error", + ] { + assert_eq!( + phases + .iter() + .filter(|value| value.as_str() == phase) + .count(), + usize::from(entrypoint == "reloadable"), + "reload phase {phase} in {scenario}" + ); + } + assert_eq!( + sequences, + (0..RECORD_COUNT as u64).collect::>(), + "records must remain complete and ordered in {scenario}" + ); +} + +fn read_log_files(directory: &Path) -> String { + let mut paths: Vec<_> = fs::read_dir(directory) + .expect("log directory") + .map(|entry| entry.expect("log directory entry").path()) + .filter(|path| { + path.file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| name.starts_with(SERVICE_NAME) && name.ends_with(".log")) + }) + .collect(); + paths.sort(); + paths + .into_iter() + .map(|path| fs::read_to_string(path).expect("UTF-8 file log")) + .collect() +} + +fn run_subprocess(command: &mut Command) -> Output { + let mut child = command + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .expect("logging subprocess"); + let stdout = child.stdout.take().expect("stdout pipe"); + let stderr = child.stderr.take().expect("stderr pipe"); + let stdout_reader = std::thread::spawn(move || read_pipe(stdout)); + let stderr_reader = std::thread::spawn(move || read_pipe(stderr)); + let deadline = Instant::now() + Duration::from_secs(10); + let mut timed_out = false; + let status = loop { + if let Some(status) = child.try_wait().expect("child status") { + break status; + } + if Instant::now() >= deadline { + timed_out = true; + let _ = child.kill(); + break child.wait().expect("reap timed out child"); + } + std::thread::sleep(Duration::from_millis(10)); + }; + let output = Output { + status, + stdout: stdout_reader.join().expect("stdout reader"), + stderr: stderr_reader.join().expect("stderr reader"), + }; + assert!( + !timed_out, + "logging subprocess timed out: {}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + output +} + +fn read_pipe(mut pipe: impl Read) -> Vec { + const MAX_CAPTURE_BYTES: usize = 4 * 1024 * 1024; + let mut captured = Vec::new(); + let mut chunk = [0u8; 8192]; + loop { + let count = pipe.read(&mut chunk).expect("drain child pipe"); + if count == 0 { + return captured; + } + let retained = count.min(MAX_CAPTURE_BYTES.saturating_sub(captured.len())); + captured.extend_from_slice(&chunk[..retained]); + } +} + +struct TestDirectory(PathBuf); + +impl TestDirectory { + fn new() -> Self { + let path = + std::env::temp_dir().join(format!("aether-nonblocking-logs-{}", uuid::Uuid::new_v4())); + fs::create_dir(&path).expect("test directory"); + Self(path) + } +} + +impl Drop for TestDirectory { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.0); + } +} diff --git a/crates/aether-runtime/base/tests/root_logging.rs b/crates/aether-runtime/base/tests/root_logging.rs index 5965ebe04..0b4d7d656 100644 --- a/crates/aether-runtime/base/tests/root_logging.rs +++ b/crates/aether-runtime/base/tests/root_logging.rs @@ -1,13 +1,15 @@ #![cfg(target_os = "linux")] use std::fs; +use std::io::Read; use std::os::unix::fs::{MetadataExt as _, PermissionsExt as _}; use std::path::PathBuf; -use std::process::Command; +use std::process::{Command, Output, Stdio}; +use std::time::{Duration, Instant}; use aether_runtime::{ - init_reloadable_service_tracing, init_service_runtime, FileLoggingConfig, LogDestination, - LogFormat, LogRotation, ServiceRuntimeConfig, + init_reloadable_service_tracing, init_service_runtime, shutdown_logging, FileLoggingConfig, + LogDestination, LogFormat, LogRotation, ServiceRuntimeConfig, }; #[test] @@ -23,7 +25,9 @@ fn root_appends_to_existing_logs_without_changing_ownership() { for format in ["pretty", "json"] { for owner in ["0", "1000", "65532", "new"] { let scenario = format!("{entrypoint}-{destination}-{format}-{owner}"); - let output = Command::new(std::env::current_exe().expect("test executable")) + let mut command = + Command::new(std::env::current_exe().expect("test executable")); + command .args([ "--ignored", "--exact", @@ -31,9 +35,8 @@ fn root_appends_to_existing_logs_without_changing_ownership() { "--nocapture", ]) .env("AETHER_TEST_ROOT_LOGGING_CASE", &scenario) - .env_remove("RUST_LOG") - .output() - .expect("root logging subprocess"); + .env_remove("RUST_LOG"); + let output = run_subprocess(&mut command); let stdout = String::from_utf8_lossy(&output.stdout); let stderr = String::from_utf8_lossy(&output.stderr); assert!(output.status.success(), "{scenario}: {stdout}\n{stderr}"); @@ -49,6 +52,51 @@ fn root_appends_to_existing_logs_without_changing_ownership() { } } +fn run_subprocess(command: &mut Command) -> Output { + let mut child = command + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .expect("root logging subprocess"); + let mut stdout = child.stdout.take().expect("stdout pipe"); + let mut stderr = child.stderr.take().expect("stderr pipe"); + let stdout_reader = std::thread::spawn(move || { + let mut bytes = Vec::new(); + stdout.read_to_end(&mut bytes).expect("read child stdout"); + bytes + }); + let stderr_reader = std::thread::spawn(move || { + let mut bytes = Vec::new(); + stderr.read_to_end(&mut bytes).expect("read child stderr"); + bytes + }); + let deadline = Instant::now() + Duration::from_secs(10); + let mut timed_out = false; + let status = loop { + if let Some(status) = child.try_wait().expect("child status") { + break status; + } + if Instant::now() >= deadline { + timed_out = true; + let _ = child.kill(); + break child.wait().expect("reap timed out child"); + } + std::thread::sleep(Duration::from_millis(10)); + }; + let output = Output { + status, + stdout: stdout_reader.join().expect("stdout reader"), + stderr: stderr_reader.join().expect("stderr reader"), + }; + assert!( + !timed_out, + "root logging subprocess timed out: {}\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + output +} + fn run_scenario(scenario: &str) { assert_eq!(unsafe { libc::geteuid() }, 0); assert_eq!(unsafe { libc::getegid() }, 0); @@ -117,6 +165,10 @@ fn run_scenario(scenario: &str) { other => panic!("unknown entrypoint: {other}"), }; tracing::info!("root logging ready"); + assert!( + shutdown_logging(Duration::from_secs(2)), + "root file logging should drain" + ); let metadata = fs::metadata(&log_file).expect("written log file"); assert_eq!(metadata.uid(), expected_owner); diff --git a/crates/aether-runtime/state/src/lib.rs b/crates/aether-runtime/state/src/lib.rs index 22ea91650..a73ddb163 100644 --- a/crates/aether-runtime/state/src/lib.rs +++ b/crates/aether-runtime/state/src/lib.rs @@ -1,6 +1,7 @@ mod error; mod memory; pub mod redis; +mod score_window; use std::collections::BTreeMap; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; @@ -17,6 +18,7 @@ use async_trait::async_trait; pub use error::DataLayerError; use memory::MemoryRuntimeBackend; pub use memory::MemoryRuntimeStateConfig; +pub use score_window::{ScoreWindowU64Stats, SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT}; use tokio::task::JoinHandle; use tracing::warn; use uuid::Uuid; @@ -621,6 +623,32 @@ impl RuntimeState { } } + /// Aggregate at most 512 timestamped `prefix:u64` members per key without + /// transferring their history. `None` requires an exact full-range fallback; + /// it never represents an empty or cached window. + pub async fn score_window_u64_stats_by_min( + &self, + keys: &[String], + min_score: f64, + ) -> Result>, DataLayerError> { + if !min_score.is_finite() { + return Err(DataLayerError::InvalidInput( + "runtime window minimum score must be finite".to_string(), + )); + } + match self.backend.as_ref() { + RuntimeStateBackend::Memory(memory) => { + Ok(memory.score_window_u64_stats_by_min(keys, min_score).await) + } + RuntimeStateBackend::Redis(redis) => { + redis + .runtime + .score_window_u64_stats_by_min(keys, min_score) + .await + } + } + } + pub async fn score_remove_by_score( &self, key: &str, @@ -1046,6 +1074,26 @@ pub struct RuntimeQueueEntry { pub fields: BTreeMap, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RuntimeQueueReclaimPage { + /// Resume the next reclaim scan here; `0-0` marks the end of the current scan. + pub next_start_id: String, + pub entries: Vec, + pub deleted_ids: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RuntimeQueueTransferOutcome { + Transferred { + destination_id: String, + acked: usize, + deleted: usize, + }, + /// No pending entry was present. This does not assert that it was archived: + /// another consumer, deletion, or retention policy may have removed it. + NotPending, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub struct RuntimeQueueStats { pub stream_length: u64, @@ -1085,6 +1133,45 @@ fn validate_runtime_queue_reclaim_config( Ok(()) } +pub(crate) fn validate_runtime_queue_transfer( + source: &str, + group: &str, + entry_id: &str, + destination: &str, + destination_fields: &BTreeMap, +) -> Result<(), DataLayerError> { + validate_runtime_queue_name(source, "runtime queue source stream")?; + validate_runtime_queue_name(group, "runtime queue group")?; + validate_runtime_queue_name(destination, "runtime queue destination stream")?; + if source == destination { + return Err(DataLayerError::InvalidInput( + "runtime queue transfer source and destination must differ".to_string(), + )); + } + if destination_fields.is_empty() { + return Err(DataLayerError::InvalidInput( + "runtime queue transfer destination fields cannot be empty".to_string(), + )); + } + let canonical_u64 = |value: &str| { + !value.is_empty() + && value.bytes().all(|byte| byte.is_ascii_digit()) + && (value.len() == 1 || !value.starts_with('0')) + && value.parse::().is_ok() + }; + if !entry_id + .split_once('-') + .is_some_and(|(milliseconds, sequence)| { + canonical_u64(milliseconds) && canonical_u64(sequence) + }) + { + return Err(DataLayerError::InvalidInput( + "runtime queue transfer entry id must be a canonical u64-u64 stream id".to_string(), + )); + } + Ok(()) +} + #[async_trait] pub trait RuntimeQueueStore: Send + Sync { async fn ensure_consumer_group( @@ -1119,6 +1206,39 @@ pub trait RuntimeQueueStore: Send + Sync { config: RuntimeQueueReclaimConfig, ) -> Result, DataLayerError>; + /// Existing queue backends can retain their complete-scan behavior without implementing paging. + async fn claim_stale_page( + &self, + stream: &str, + group: &str, + consumer: &str, + start_id: &str, + config: RuntimeQueueReclaimConfig, + ) -> Result { + Ok(RuntimeQueueReclaimPage { + next_start_id: "0-0".to_string(), + entries: self + .claim_stale(stream, group, consumer, start_id, config) + .await?, + deleted_ids: Vec::new(), + }) + } + + /// Atomically append caller-supplied fields, acknowledge the pending source entry, and + /// delete that source ID. Repeated calls must not append when the entry is no longer pending. + /// `None` means unsupported and has no side effects; callers may explicitly retain their + /// existing non-atomic fallback for third-party queue implementations. + async fn try_transfer_pending_to_stream( + &self, + _source: &str, + _group: &str, + _entry_id: &str, + _destination: &str, + _destination_fields: &BTreeMap, + ) -> Result, DataLayerError> { + Ok(None) + } + async fn ack(&self, stream: &str, group: &str, ids: &[String]) -> Result; @@ -1242,6 +1362,20 @@ impl RuntimeQueueStore for RuntimeState { start_id: &str, config: RuntimeQueueReclaimConfig, ) -> Result, DataLayerError> { + Ok(self + .claim_stale_page(stream, group, consumer, start_id, config) + .await? + .entries) + } + + async fn claim_stale_page( + &self, + stream: &str, + group: &str, + consumer: &str, + start_id: &str, + config: RuntimeQueueReclaimConfig, + ) -> Result { validate_runtime_queue_name(stream, "runtime queue stream")?; validate_runtime_queue_name(group, "runtime queue group")?; validate_runtime_queue_name(consumer, "runtime queue consumer")?; @@ -1250,32 +1384,76 @@ impl RuntimeQueueStore for RuntimeState { match self.backend.as_ref() { RuntimeStateBackend::Memory(memory) => { memory - .queue_claim_stale(stream, group, consumer, start_id, config) + .queue_claim_stale_page(stream, group, consumer, start_id, config) .await } - RuntimeStateBackend::Redis(redis) => Ok(redis - .stream - .claim_stale( - &RedisStreamName(stream.to_string()), - &RedisConsumerGroup(group.to_string()), - &RedisConsumerName(consumer.to_string()), - start_id, - RedisStreamReclaimConfig { - min_idle_ms: config.min_idle_ms, - count: config.count, - }, - ) - .await? - .entries - .into_iter() - .map(|entry| RuntimeQueueEntry { - id: entry.id, - fields: entry.fields, + RuntimeStateBackend::Redis(redis) => { + let page = redis + .stream + .claim_stale( + &RedisStreamName(stream.to_string()), + &RedisConsumerGroup(group.to_string()), + &RedisConsumerName(consumer.to_string()), + start_id, + RedisStreamReclaimConfig { + min_idle_ms: config.min_idle_ms, + count: config.count, + }, + ) + .await?; + Ok(RuntimeQueueReclaimPage { + next_start_id: page.next_start_id, + entries: page + .entries + .into_iter() + .map(|entry| RuntimeQueueEntry { + id: entry.id, + fields: entry.fields, + }) + .collect(), + deleted_ids: page.deleted_ids, }) - .collect()), + } } } + async fn try_transfer_pending_to_stream( + &self, + source: &str, + group: &str, + entry_id: &str, + destination: &str, + destination_fields: &BTreeMap, + ) -> Result, DataLayerError> { + validate_runtime_queue_transfer(source, group, entry_id, destination, destination_fields)?; + let outcome = match self.backend.as_ref() { + RuntimeStateBackend::Memory(memory) => { + memory + .queue_transfer_pending_to_stream( + source, + group, + entry_id, + destination, + destination_fields, + ) + .await? + } + RuntimeStateBackend::Redis(redis) => { + redis + .stream + .try_transfer_pending_to_stream( + source, + group, + entry_id, + destination, + destination_fields, + ) + .await? + } + }; + Ok(Some(outcome)) + } + async fn ack( &self, stream: &str, @@ -1817,6 +1995,18 @@ mod tests { use std::process::{Child, Command, Stdio}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; + mod stream_receive { + include!("redis/stream_receive_tests.rs"); + } + + mod dead_letter_transfer { + include!("redis/dead_letter_transfer_tests.rs"); + } + + mod usage_limit_cleanup { + include!("redis/usage_limit_cleanup_tests.rs"); + } + #[tokio::test] async fn memory_kv_expires_entries() { let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default()); @@ -2722,6 +2912,131 @@ mod tests { } } + #[tokio::test] + async fn redis_large_stream_batches_preserve_fields_across_read_reclaim_and_ack() { + let Some(redis) = TestRedisServer::start().await else { + return; + }; + for protocol in ["resp2", "resp3"] { + let runtime = RuntimeState::redis( + RedisClientConfig { + url: format!("{}?protocol={protocol}", redis.redis_url), + key_prefix: Some(format!("large-batch-{protocol}")), + }, + Some(5_000), + ) + .await + .expect("large batch runtime should connect"); + let stream = "usage:large-batch"; + let group = "workers"; + RuntimeQueueStore::ensure_consumer_group(&runtime, stream, group, "0-0") + .await + .unwrap(); + let payload = format!( + "{}\r\n\"escaped\"\\\u{4e2d}\u{6587}", + "x".repeat(512 * 1024) + ); + let mut expected = BTreeMap::new(); + for sequence in 0..24 { + let fields = BTreeMap::from([ + ("payload".to_string(), payload.clone()), + ("sequence".to_string(), sequence.to_string()), + ("legacy_marker".to_string(), "preserve exactly".to_string()), + ]); + let id = + RuntimeQueueStore::append_fields_with_maxlen(&runtime, stream, &fields, None) + .await + .unwrap(); + expected.insert(id, sequence.to_string()); + } + let mut readers = tokio::task::JoinSet::new(); + for index in 0..3 { + let runtime = runtime.clone(); + readers.spawn(async move { + RuntimeQueueStore::read_group( + &runtime, + stream, + group, + &format!("reader-{index}"), + 8, + Some(1), + ) + .await + .unwrap() + }); + } + let mut delivered = std::collections::BTreeSet::new(); + while let Some(entries) = readers.join_next().await { + let entries = entries.unwrap(); + assert_eq!(entries.len(), 8); + for entry in entries { + assert_eq!(entry.fields.len(), 3); + assert_eq!(entry.fields["payload"].as_bytes(), payload.as_bytes()); + assert_eq!(entry.fields["sequence"], expected[&entry.id]); + assert_eq!(entry.fields["legacy_marker"], "preserve exactly"); + assert!(delivered.insert(entry.id)); + } + } + assert_eq!(delivered.len(), 24); + let stats = RuntimeQueueStore::stats(&runtime, stream, Some(group)) + .await + .unwrap(); + assert_eq!(stats.group_pending, 24); + assert_eq!(stats.group_lag, Some(0)); + + tokio::time::sleep(Duration::from_millis(20)).await; + let mut reclaimed = std::collections::BTreeSet::new(); + while reclaimed.len() < 24 { + let entries = RuntimeQueueStore::claim_stale( + &runtime, + stream, + group, + "retry-consumer", + "0-0", + RuntimeQueueReclaimConfig { + min_idle_ms: 1, + count: 5, + }, + ) + .await + .unwrap(); + assert!(!entries.is_empty()); + assert!(entries.len() <= 5); + let mut ids = Vec::new(); + for entry in entries { + assert_eq!(entry.fields.len(), 3); + assert_eq!(entry.fields["payload"].as_bytes(), payload.as_bytes()); + assert_eq!(entry.fields["sequence"], expected[&entry.id]); + assert_eq!(entry.fields["legacy_marker"], "preserve exactly"); + assert!(reclaimed.insert(entry.id.clone())); + ids.push(entry.id); + } + assert_eq!( + RuntimeQueueStore::ack(&runtime, stream, group, &ids) + .await + .unwrap(), + ids.len() + ); + assert_eq!( + RuntimeQueueStore::delete(&runtime, stream, &ids) + .await + .unwrap(), + ids.len() + ); + } + assert_eq!(reclaimed, delivered); + let stats = RuntimeQueueStore::stats(&runtime, stream, Some(group)) + .await + .unwrap(); + assert_eq!(stats.stream_length, 0); + assert_eq!(stats.group_pending, 0); + assert_eq!(stats.group_lag, Some(0)); + eprintln!( + "verified {protocol}: 24 large records, 3 readers, read/reclaim/ack complete" + ); + } + } + #[tokio::test] async fn redis_connection_manager_recovers_after_restart() { let Some(mut redis) = TestRedisServer::start().await else { @@ -2776,6 +3091,269 @@ mod tests { assert_kv_score_and_queue_contract(&redis_runtime).await; } + #[tokio::test] + async fn runtime_backends_share_bounded_score_window_aggregation() { + let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default()); + assert_bounded_score_window_aggregation(&memory).await; + + let Some((_server, runtime)) = redis_runtime_for_test("score-window").await else { + return; + }; + assert_bounded_score_window_aggregation(&runtime).await; + } + + async fn assert_bounded_score_window_aggregation(runtime: &RuntimeState) { + let keys = (0..35) + .map(|index| format!("window:{index}")) + .collect::>(); + for (member, score) in [ + ("expired:999", 99.999), + ("boundary:7", 100.0), + ("recent:9007199254740993", 101.0), + ("nested:prefix:+00012", 102.0), + ("zero:0", 103.0), + ("invalid:1.5", 104.0), + ("invalid:-1", 105.0), + ("invalid:18446744073709551616", 106.0), + ("invalid: 12", 107.0), + ("missing-separator", 108.0), + ] { + runtime + .score_set(&keys[0], member, score) + .await + .expect("seed values"); + } + runtime + .score_set(&keys[1], "max:18446744073709551615", 100.0) + .await + .expect("seed max"); + runtime + .score_set(&keys[2], "max:18446744073709551615", 100.0) + .await + .expect("seed overflow"); + runtime + .score_set(&keys[2], "additional:2", 100.0) + .await + .expect("seed overflow addition"); + for index in 0..SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT { + runtime + .score_set(&keys[3], &format!("{index}:3"), 100.0) + .await + .expect("seed bounded window"); + } + runtime + .score_set(&keys[3], "expired:9999", 0.0) + .await + .expect("seed expired sample"); + let stats = runtime + .score_window_u64_stats_by_min(&keys, 100.0) + .await + .expect("aggregate"); + assert_eq!( + stats.len(), + keys.len(), + "pipeline batches preserve key order" + ); + assert_eq!( + stats[0], + Some(ScoreWindowU64Stats { + sum: 9_007_199_254_741_012, + positive_count: 3 + }) + ); + assert_eq!( + stats[1], + Some(ScoreWindowU64Stats { + sum: u64::MAX, + positive_count: 1 + }) + ); + assert_eq!( + stats[2], + Some(ScoreWindowU64Stats { + sum: u64::MAX, + positive_count: 2 + }) + ); + assert_eq!( + stats[3], + Some(ScoreWindowU64Stats { + sum: 1536, + positive_count: 512 + }) + ); + assert!(stats[4..] + .iter() + .all(|stats| *stats == Some(ScoreWindowU64Stats::default()))); + + runtime + .score_set(&keys[3], "overflowing-window:11", 101.0) + .await + .expect("exceed server limit"); + let stats = runtime + .score_window_u64_stats_by_min(&keys[3..4], 100.0) + .await + .expect("bounded fallback"); + assert_eq!( + stats, + vec![None], + "oversized windows require the full exact read" + ); + let members = runtime + .score_range_by_min(&keys[3], 100.0) + .await + .expect("full window"); + assert_eq!( + ScoreWindowU64Stats::from_members(members.iter().map(String::as_str)).sum, + 1547 + ); + runtime + .score_remove(&keys[3], "overflowing-window:11") + .await + .expect("remove newest"); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys[3..4], 100.0) + .await + .expect("read after remove")[0] + .unwrap() + .sum, + 1536 + ); + runtime + .score_set(&keys[3], "0:3", 99.0) + .await + .expect("move sample outside window"); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys[3..4], 100.0) + .await + .expect("read changed score")[0] + .unwrap() + .sum, + 1533 + ); + assert!(runtime + .score_window_u64_stats_by_min(&[], 100.0) + .await + .expect("empty query") + .is_empty()); + assert!(runtime + .score_window_u64_stats_by_min(&keys, f64::NAN) + .await + .is_err()); + } + + #[tokio::test] + async fn redis_score_window_aggregation_reloads_scripts_without_caching_old_cost() { + let Some((server, runtime)) = redis_runtime_for_test("score-window-reload").await else { + return; + }; + let keys = vec!["reload:cost".to_string()]; + runtime + .score_set(&keys[0], "first:7", 100.0) + .await + .expect("first cost"); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys, 100.0) + .await + .expect("first aggregate")[0] + .unwrap() + .sum, + 7 + ); + let client = ::redis::Client::open(server.redis_url.as_str()).expect("test Redis client"); + let mut connection = client + .get_multiplexed_async_connection() + .await + .expect("test connection"); + ::redis::cmd("SCRIPT") + .arg("FLUSH") + .query_async::<()>(&mut connection) + .await + .expect("flush scripts"); + runtime + .score_set(&keys[0], "second:11", 101.0) + .await + .expect("new cost"); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys, 100.0) + .await + .expect("reload aggregate")[0] + .unwrap() + .sum, + 18 + ); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys, 101.0) + .await + .expect("changed window")[0] + .unwrap() + .sum, + 11 + ); + runtime + .key_expire(&keys[0], Duration::ZERO) + .await + .expect("expire window"); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys, 100.0) + .await + .expect("expired aggregate"), + vec![Some(ScoreWindowU64Stats::default())] + ); + } + + #[tokio::test] + async fn redis_score_window_aggregation_observes_completed_concurrent_writes() { + let Some((_server, runtime)) = redis_runtime_for_test("score-window-concurrent").await + else { + return; + }; + let writer_runtime = runtime.clone(); + let (written_tx, mut written_rx) = tokio::sync::mpsc::channel(8); + let writer = tokio::spawn(async move { + for index in 1..=128_u64 { + writer_runtime + .score_set("concurrent:cost", &format!("{index}:2"), 100.0) + .await + .expect("concurrent write"); + written_tx.send(index).await.expect("notify reader"); + } + }); + let keys = vec!["concurrent:cost".to_string()]; + while let Some(written) = written_rx.recv().await { + let stats = runtime + .score_window_u64_stats_by_min(&keys, 100.0) + .await + .expect("concurrent aggregate")[0] + .unwrap(); + assert!( + stats.positive_count >= written, + "completed writes must not be hidden by a stale aggregate" + ); + assert_eq!( + stats.sum, + stats.positive_count * 2, + "one script observes one consistent window" + ); + } + writer.await.expect("writer task"); + assert_eq!( + runtime + .score_window_u64_stats_by_min(&keys, 100.0) + .await + .expect("final aggregate")[0] + .unwrap() + .sum, + 256 + ); + } + #[tokio::test] async fn runtime_backends_reject_invalid_shared_inputs() { let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default()); diff --git a/crates/aether-runtime/state/src/memory.rs b/crates/aether-runtime/state/src/memory.rs index 4fdddd6a7..072e2bef3 100644 --- a/crates/aether-runtime/state/src/memory.rs +++ b/crates/aether-runtime/state/src/memory.rs @@ -6,8 +6,11 @@ use std::time::{Duration, Instant}; use tokio::sync::Mutex; -use crate::UsageLimitCheck; -use crate::{DataLayerError, RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats}; +use crate::{ + DataLayerError, RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueReclaimPage, + RuntimeQueueStats, RuntimeQueueTransferOutcome, +}; +use crate::{ScoreWindowU64Stats, UsageLimitCheck, SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT}; const MEMORY_RATE_LIMIT_COUNTER_SHARD_COUNT: usize = 64; const MEMORY_RATE_LIMIT_COUNTER_PRUNE_INTERVAL: u64 = 256; @@ -883,6 +886,29 @@ impl MemoryRuntimeBackend { .unwrap_or_default() } + pub(crate) async fn score_window_u64_stats_by_min( + &self, + keys: &[String], + min_score: f64, + ) -> Vec> { + let mut scores = self.scores.lock().await; + keys.iter() + .map(|key| { + prune_memory_key(&mut scores, key, Instant::now()); + let members = scores + .get(key) + .into_iter() + .flat_map(|entry| entry.scores.iter()) + .filter(|(_, score)| **score >= min_score) + .take(SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT + 1) + .map(|(member, _)| member.as_str()) + .collect::>(); + (members.len() <= SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT) + .then(|| ScoreWindowU64Stats::from_members(members)) + }) + .collect() + } + pub(crate) async fn score_remove_by_score(&self, key: &str, max_score: f64) -> usize { let mut scores = self.scores.lock().await; prune_memory_key(&mut scores, key, Instant::now()); @@ -937,12 +963,13 @@ impl MemoryRuntimeBackend { fields: BTreeMap, maxlen: Option, ) -> String { + // Stream IDs must follow insertion order, including when appenders wait for this lock. + let mut queues = self.queues.lock().await; let sequence = self .queue_seq .fetch_add(1, Ordering::Relaxed) .saturating_add(1); let id = format!("{sequence}-0"); - let mut queues = self.queues.lock().await; prune_memory_key(&mut queues, stream, Instant::now()); let stream_state = queues.entry(stream.to_string()).or_default(); stream_state.entries.push_back(MemoryQueuedEntry { @@ -1019,14 +1046,12 @@ impl MemoryRuntimeBackend { let now = Instant::now(); let mut delivered = Vec::new(); let last_delivered_sequence = group_state.last_delivered_sequence; - let queued_entries = stream_state + for queued in stream_state .entries .iter() .filter(|entry| entry.sequence > last_delivered_sequence) .take(count.max(1)) - .cloned() - .collect::>(); - for queued in queued_entries { + { group_state.last_delivered_sequence = queued.sequence; group_state.pending.insert( queued.entry.id.clone(), @@ -1054,7 +1079,8 @@ impl MemoryRuntimeBackend { } } - pub(crate) async fn queue_claim_stale( + #[cfg(test)] + async fn queue_claim_stale( &self, stream: &str, group: &str, @@ -1062,6 +1088,20 @@ impl MemoryRuntimeBackend { start_id: &str, config: RuntimeQueueReclaimConfig, ) -> Result, DataLayerError> { + Ok(self + .queue_claim_stale_page(stream, group, consumer, start_id, config) + .await? + .entries) + } + + pub(crate) async fn queue_claim_stale_page( + &self, + stream: &str, + group: &str, + consumer: &str, + start_id: &str, + config: RuntimeQueueReclaimConfig, + ) -> Result { let start_sequence = parse_memory_stream_sequence(start_id)?; let min_idle = Duration::from_millis(config.min_idle_ms.max(1)); let now = Instant::now(); @@ -1086,16 +1126,106 @@ impl MemoryRuntimeBackend { .collect::>(); let mut ids = ids; ids.sort_by_key(|(sequence, _)| *sequence); + let count = config.count.max(1); + let next_start_id = ids + .get(count) + .map(|(_, id)| id.clone()) + .unwrap_or_else(|| "0-0".to_string()); let mut claimed = Vec::new(); - for (_, id) in ids.into_iter().take(config.count.max(1)) { + for (_, id) in ids.into_iter().take(count) { if let Some(pending) = group_state.pending.get_mut(&id) { pending.consumer = consumer.to_string(); pending.delivered_at = now; claimed.push(pending.entry.clone()); } } - Ok(claimed) + Ok(RuntimeQueueReclaimPage { + next_start_id, + entries: claimed, + deleted_ids: Vec::new(), + }) + } + + pub(crate) async fn queue_transfer_pending_to_stream( + &self, + source: &str, + group: &str, + entry_id: &str, + destination: &str, + destination_fields: &BTreeMap, + ) -> Result { + crate::validate_runtime_queue_transfer( + source, + group, + entry_id, + destination, + destination_fields, + )?; + // Match ordinary memory append ownership without copying a large payload while locked. + let destination_fields = destination_fields.clone(); + let mut queues = self.queues.lock().await; + let now = Instant::now(); + prune_memory_key(&mut queues, source, now); + let source_state = queues.get(source).ok_or_else(|| { + DataLayerError::InvalidInput(format!("runtime queue stream {source} does not exist")) + })?; + let group_state = source_state.groups.get(group).ok_or_else(|| { + DataLayerError::InvalidInput(format!( + "runtime queue group {group} does not exist for stream {source}" + )) + })?; + // Memory trimming/deletion already removes PEL entries. Absence here cannot prove + // archival, and must not delete an unread entry or append another dead letter. + if !group_state.pending.contains_key(entry_id) { + return Ok(RuntimeQueueTransferOutcome::NotPending); + } + + let previous_sequence = self + .queue_seq + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |sequence| { + sequence.checked_add(1) + }) + .map_err(|_| { + DataLayerError::UnexpectedValue("runtime queue sequence exhausted".to_string()) + })?; + let sequence = previous_sequence + 1; + let destination_id = format!("{sequence}-0"); + prune_memory_key(&mut queues, destination, now); + queues + .entry(destination.to_string()) + .or_default() + .entries + .push_back(MemoryQueuedEntry { + sequence, + entry: RuntimeQueueEntry { + id: destination_id.clone(), + fields: destination_fields, + }, + }); + + // No await occurs between archive creation and source removal. Cancellation can only + // happen while waiting for the mutex, so it cannot leave a half-completed transfer. + let source_state = queues.get_mut(source).expect("validated source stream"); + let acked = usize::from( + source_state + .groups + .get_mut(group) + .expect("validated source group") + .pending + .remove(entry_id) + .is_some(), + ); + let before = source_state.entries.len(); + source_state + .entries + .retain(|entry| entry.entry.id != entry_id); + remove_pending_from_all_groups(source_state, entry_id); + Ok(RuntimeQueueTransferOutcome::Transferred { + destination_id, + acked, + deleted: before.saturating_sub(source_state.entries.len()), + }) } pub(crate) async fn queue_ack( @@ -1449,6 +1579,793 @@ fn unix_time_ms() -> u64 { mod tests { use super::*; + fn memory_queue_test_fields(index: usize) -> BTreeMap { + BTreeMap::from([ + ( + "payload".to_string(), + format!("record-{index}:{}\n\"\\\u{03bb}", "payload".repeat(8_192)), + ), + ("kind".to_string(), format!("event-{index}")), + (String::new(), String::new()), + ]) + } + + async fn age_memory_queue_pending(backend: &MemoryRuntimeBackend, stream: &str) { + let stale = Instant::now() + .checked_sub(Duration::from_secs(1)) + .expect("test clock should support one second of history"); + let mut queues = backend.queues.lock().await; + for group in queues + .get_mut(stream) + .expect("test stream") + .groups + .values_mut() + { + for pending in group.pending.values_mut() { + pending.delivered_at = stale; + } + } + } + + async fn memory_queue_transfer_fixture( + backend: &MemoryRuntimeBackend, + count: usize, + ) -> Vec { + backend + .queue_ensure_consumer_group("transfer:source", "workers", "0-0") + .await + .expect("source group"); + for index in 0..count { + backend + .queue_append("transfer:source", memory_queue_test_fields(index), None) + .await; + } + backend + .queue_read("transfer:source", "workers", "reader", count, None) + .await + .expect("pending source entries") + } + + #[tokio::test] + async fn memory_queue_transfer_preserves_fields_and_only_removes_the_target_entry() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 3).await; + backend + .queue_ensure_consumer_group("transfer:source", "other-workers", "0-0") + .await + .expect("second source group"); + backend + .queue_read("transfer:source", "other-workers", "reader", 3, None) + .await + .expect("second group pending entries"); + let mut archived_fields = entries[1].fields.clone(); + archived_fields.insert("source_id".to_string(), entries[1].id.clone()); + archived_fields.insert( + "error".to_string(), + "invalid payload\noriginal retained".to_string(), + ); + let outcome = backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entries[1].id, + "transfer:archive", + &archived_fields, + ) + .await + .expect("atomic transfer"); + let RuntimeQueueTransferOutcome::Transferred { + destination_id, + acked, + deleted, + } = outcome + else { + panic!("pending entry should transfer"); + }; + assert_eq!((acked, deleted), (1, 1)); + let queues = backend.queues.lock().await; + let source = &queues["transfer:source"]; + let remaining = source + .entries + .iter() + .map(|entry| &entry.entry) + .collect::>(); + assert_eq!(remaining, [&entries[0], &entries[2]]); + for group in ["workers", "other-workers"] { + let pending = &source.groups[group].pending; + assert_eq!(pending.len(), 2); + assert!(pending.contains_key(&entries[0].id)); + assert!(pending.contains_key(&entries[2].id)); + } + let archive = &queues["transfer:archive"]; + assert_eq!(archive.entries.len(), 1); + assert_eq!(archive.entries[0].entry.id, destination_id); + assert_eq!(archive.entries[0].entry.fields, archived_fields); + } + + #[tokio::test] + async fn memory_queue_transfer_concurrent_and_repeated_attempts_archive_once() { + let backend = std::sync::Arc::new(MemoryRuntimeBackend::new( + MemoryRuntimeStateConfig::default(), + )); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + let fields = std::sync::Arc::new(entries[0].fields.clone()); + let mut tasks = tokio::task::JoinSet::new(); + for _ in 0..16 { + let backend = std::sync::Arc::clone(&backend); + let fields = std::sync::Arc::clone(&fields); + let entry_id = entries[0].id.clone(); + tasks.spawn(async move { + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entry_id, + "transfer:archive", + &fields, + ) + .await + .expect("transfer attempt") + }); + } + let mut transferred = 0; + let mut not_pending = 0; + while let Some(outcome) = tasks.join_next().await { + match outcome.expect("transfer task") { + RuntimeQueueTransferOutcome::Transferred { acked, deleted, .. } => { + assert_eq!((acked, deleted), (1, 1)); + transferred += 1; + } + RuntimeQueueTransferOutcome::NotPending => not_pending += 1, + } + } + assert_eq!((transferred, not_pending), (1, 15)); + assert_eq!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entries[0].id, + "transfer:archive", + &fields, + ) + .await + .expect("sequential retry"), + RuntimeQueueTransferOutcome::NotPending + ); + let stats = backend + .queue_stats("transfer:source", Some("workers")) + .await; + assert_eq!((stats.stream_length, stats.group_pending), (0, 0)); + assert_eq!( + backend + .queue_stats("transfer:archive", None) + .await + .stream_length, + 1 + ); + } + + #[tokio::test] + async fn memory_queue_transfer_retry_after_lost_response_does_not_archive_twice() { + async fn commit_then_lose_response( + backend: &MemoryRuntimeBackend, + entry: &RuntimeQueueEntry, + ) -> Result { + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entry.id, + "transfer:archive", + &entry.fields, + ) + .await?; + Err(DataLayerError::TimedOut( + "transfer response lost after commit".to_string(), + )) + } + + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + assert!(matches!( + commit_then_lose_response(&backend, &entries[0]).await, + Err(DataLayerError::TimedOut(_)) + )); + assert_eq!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entries[0].id, + "transfer:archive", + &entries[0].fields, + ) + .await + .expect("retry after lost response"), + RuntimeQueueTransferOutcome::NotPending + ); + let queues = backend.queues.lock().await; + assert!(queues["transfer:source"].entries.is_empty()); + assert!(queues["transfer:source"].groups["workers"] + .pending + .is_empty()); + assert_eq!(queues["transfer:archive"].entries.len(), 1); + assert_eq!( + queues["transfer:archive"].entries[0].entry.fields, + entries[0].fields + ); + } + + #[tokio::test] + async fn memory_queue_transfer_rejects_invalid_input_before_any_mutation() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + let entry = &entries[0]; + for invalid_id in [ + "", + "1", + "1-", + "-1-0", + "+1-0", + "01-0", + "1-00", + "1-+0", + "1-0x", + "1-0-0", + " 1-0", + "1-0 ", + "\u{0661}-0", + "18446744073709551616-0", + "1-18446744073709551616", + ] { + assert!( + matches!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + invalid_id, + "transfer:archive", + &entry.fields, + ) + .await, + Err(DataLayerError::InvalidInput(_)) + ), + "invalid entry id {invalid_id:?}" + ); + } + for (source, group, destination) in [ + ("", "workers", "transfer:archive"), + ("transfer:source", " ", "transfer:archive"), + ("transfer:source", "workers", ""), + ("transfer:source", "workers", "transfer:source"), + ("missing-source", "workers", "transfer:archive"), + ("transfer:source", "missing-group", "transfer:archive"), + ] { + assert!(matches!( + backend + .queue_transfer_pending_to_stream( + source, + group, + &entry.id, + destination, + &entry.fields + ) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + } + assert!(matches!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entry.id, + "transfer:archive", + &BTreeMap::new(), + ) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + let queues = backend.queues.lock().await; + assert_eq!(queues.len(), 1); + assert_eq!(queues["transfer:source"].entries[0].entry, *entry); + assert!(queues["transfer:source"].groups["workers"] + .pending + .contains_key(&entry.id)); + assert_eq!(backend.queue_seq.load(Ordering::Acquire), 1); + } + + #[tokio::test] + async fn memory_queue_transfer_archive_failure_keeps_source_pending() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + backend.queue_seq.store(u64::MAX, Ordering::Release); + assert!(matches!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entries[0].id, + "transfer:archive", + &entries[0].fields, + ) + .await, + Err(DataLayerError::UnexpectedValue(_)) + )); + let queues = backend.queues.lock().await; + assert_eq!(queues.len(), 1); + assert_eq!(queues["transfer:source"].entries[0].entry, entries[0]); + assert!(queues["transfer:source"].groups["workers"] + .pending + .contains_key(&entries[0].id)); + } + + #[tokio::test] + async fn memory_queue_transfer_cancelled_lock_wait_has_no_side_effects() { + use std::future::Future; + use std::task::Poll; + + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + let queues = backend.queues.lock().await; + let mut transfer = Box::pin(backend.queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entries[0].id, + "transfer:archive", + &entries[0].fields, + )); + std::future::poll_fn(|cx| { + assert!(transfer.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + drop(transfer); + assert_eq!(backend.queue_seq.load(Ordering::Acquire), 1); + assert_eq!(queues.len(), 1); + assert_eq!(queues["transfer:source"].entries[0].entry, entries[0]); + assert!(queues["transfer:source"].groups["workers"] + .pending + .contains_key(&entries[0].id)); + drop(queues); + assert_eq!( + backend + .queue_stats("transfer:archive", None) + .await + .stream_length, + 0 + ); + } + + #[tokio::test] + async fn memory_queue_transfer_not_pending_does_not_archive_or_delete_unread_or_acked_entries() + { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + let unread_id = backend + .queue_append("transfer:source", memory_queue_test_fields(1), None) + .await; + backend + .queue_ack("transfer:source", "workers", &[entries[0].id.clone()]) + .await + .expect("ack without deleting"); + for id in [ + entries[0].id.as_str(), + unread_id.as_str(), + "1-1", + "0-0", + "18446744073709551615-18446744073709551615", + ] { + assert_eq!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + id, + "transfer:archive", + &entries[0].fields, + ) + .await + .expect("valid but non-pending entry id"), + RuntimeQueueTransferOutcome::NotPending + ); + } + let stats = backend + .queue_stats("transfer:source", Some("workers")) + .await; + assert_eq!( + (stats.stream_length, stats.group_pending, stats.group_lag), + (2, 0, Some(1)) + ); + assert_eq!( + backend + .queue_stats("transfer:archive", None) + .await + .stream_length, + 0 + ); + assert_eq!(backend.queue_seq.load(Ordering::Acquire), 2); + } + + #[tokio::test] + async fn memory_queue_transfer_retains_existing_trimmed_pending_semantics() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let entries = memory_queue_transfer_fixture(&backend, 1).await; + backend + .queue_append("transfer:source", memory_queue_test_fields(1), Some(1)) + .await; + // Memory retention already removed both the original entry and its PEL copy. + // NotPending does not claim that the retained caller copy was archived elsewhere. + assert_eq!( + backend + .queue_transfer_pending_to_stream( + "transfer:source", + "workers", + &entries[0].id, + "transfer:archive", + &entries[0].fields, + ) + .await + .expect("trimmed entry"), + RuntimeQueueTransferOutcome::NotPending + ); + let stats = backend + .queue_stats("transfer:source", Some("workers")) + .await; + assert_eq!((stats.stream_length, stats.group_pending), (1, 0)); + assert_eq!( + backend + .queue_stats("transfer:archive", None) + .await + .stream_length, + 0 + ); + } + + #[tokio::test] + async fn memory_queue_waiting_append_does_not_reserve_an_out_of_order_sequence() { + use std::future::Future; + use std::task::Poll; + + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let stream = "queue:append-order"; + backend + .queue_ensure_consumer_group(stream, "workers", "0-0") + .await + .expect("consumer group"); + let lock = backend.queues.lock().await; + let mut first = Box::pin(backend.queue_append( + stream, + BTreeMap::from([("payload".to_string(), "first".to_string())]), + None, + )); + let mut second = Box::pin(backend.queue_append( + stream, + BTreeMap::from([("payload".to_string(), "second".to_string())]), + None, + )); + std::future::poll_fn(|cx| { + assert!(first.as_mut().poll(cx).is_pending()); + assert!(second.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + assert_eq!( + backend.queue_seq.load(Ordering::Acquire), + 0, + "an appender must own the insertion lock before assigning a stream sequence" + ); + drop(lock); + let (first_id, second_id) = tokio::join!(first, second); + assert_eq!(first_id, "1-0"); + assert_eq!(second_id, "2-0"); + for (expected_id, expected_payload) in [(first_id, "first"), (second_id, "second")] { + let entries = backend + .queue_read(stream, "workers", "reader", 1, None) + .await + .expect("ordered delivery"); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].id, expected_id); + assert_eq!(entries[0].fields["payload"], expected_payload); + } + assert!(backend + .queue_read(stream, "workers", "reader", 1, None) + .await + .expect("all entries delivered exactly once") + .is_empty()); + } + + #[tokio::test] + async fn memory_queue_reclaim_page_advances_and_rescans_after_reaching_the_end() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let stream = "queue:reclaim-page"; + backend + .queue_ensure_consumer_group(stream, "workers", "0-0") + .await + .expect("consumer group"); + for index in 0..4 { + backend + .queue_append(stream, memory_queue_test_fields(index), None) + .await; + } + let expected = backend + .queue_read(stream, "workers", "reader", 4, None) + .await + .expect("initial delivery"); + age_memory_queue_pending(&backend, stream).await; + backend + .queues + .lock() + .await + .get_mut(stream) + .unwrap() + .groups + .get_mut("workers") + .unwrap() + .pending + .get_mut(&expected[0].id) + .unwrap() + .delivered_at = Instant::now(); + let config = RuntimeQueueReclaimConfig { + min_idle_ms: 500, + count: 1, + }; + let mut cursor = "0-0".to_string(); + for index in 1..4 { + let page = backend + .queue_claim_stale_page(stream, "workers", "reclaimer", &cursor, config) + .await + .expect("reclaim page"); + assert_eq!(page.entries.as_slice(), &expected[index..index + 1]); + assert!(page.deleted_ids.is_empty()); + cursor = page.next_start_id; + assert_eq!( + cursor, + expected + .get(index + 1) + .map_or("0-0", |entry| entry.id.as_str()) + ); + } + let stale = Instant::now().checked_sub(Duration::from_secs(1)).unwrap(); + backend + .queues + .lock() + .await + .get_mut(stream) + .unwrap() + .groups + .get_mut("workers") + .unwrap() + .pending + .get_mut(&expected[0].id) + .unwrap() + .delivered_at = stale; + let page = backend + .queue_claim_stale_page(stream, "workers", "reclaimer", &cursor, config) + .await + .expect("next scan rechecks the earlier fresh entry"); + assert_eq!(page.entries.as_slice(), &expected[..1]); + assert_eq!(page.next_start_id, "0-0"); + assert!(page.deleted_ids.is_empty()); + } + + #[tokio::test] + async fn memory_queue_read_batches_preserve_fields_and_independent_ownership() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let stream = "queue:read-ownership"; + backend + .queue_ensure_consumer_group(stream, "workers", "0-0") + .await + .expect("consumer group"); + let mut expected = Vec::new(); + for index in 0..3 { + let fields = memory_queue_test_fields(index); + let id = backend.queue_append(stream, fields.clone(), None).await; + expected.push(RuntimeQueueEntry { id, fields }); + } + + let mut first_batch = backend + .queue_read(stream, "workers", "consumer-a", 2, None) + .await + .expect("first batch"); + assert_eq!(first_batch.as_slice(), &expected[..2]); + let stats = backend.queue_stats(stream, Some("workers")).await; + assert_eq!(stats.group_pending, 2); + assert_eq!(stats.group_lag, Some(1)); + first_batch[0].id.clear(); + first_batch[0].fields.get_mut("payload").unwrap().clear(); + first_batch[0].fields.remove("kind"); + first_batch[1].fields.clear(); + + let second_batch = backend + .queue_read(stream, "workers", "consumer-a", 2, None) + .await + .expect("second batch"); + assert_eq!(second_batch.as_slice(), &expected[2..]); + assert!(backend + .queue_read(stream, "workers", "consumer-a", 2, None) + .await + .expect("all entries have been delivered") + .is_empty()); + + age_memory_queue_pending(&backend, stream).await; + let mut reclaimed = backend + .queue_claim_stale( + stream, + "workers", + "consumer-b", + "0-0", + RuntimeQueueReclaimConfig { + min_idle_ms: 500, + count: 1, + }, + ) + .await + .expect("bounded reclaim"); + assert_eq!(reclaimed.as_slice(), &expected[..1]); + reclaimed[0].fields.clear(); + age_memory_queue_pending(&backend, stream).await; + assert_eq!( + backend + .queue_claim_stale( + stream, + "workers", + "consumer-c", + "0-0", + RuntimeQueueReclaimConfig { + min_idle_ms: 500, + count: 3 + }, + ) + .await + .expect("reclaim still owns original fields"), + expected + ); + + backend + .queue_ensure_consumer_group(stream, "later-group", "0-0") + .await + .expect("independent consumer group"); + assert_eq!( + backend + .queue_read(stream, "later-group", "consumer-d", 3, None) + .await + .expect("stream still owns original fields"), + expected + ); + } + + #[tokio::test] + async fn memory_queue_read_ack_and_delete_preserve_pending_group_semantics() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let stream = "queue:ack-delete"; + let mut expected = Vec::new(); + for index in 0..3 { + let fields = memory_queue_test_fields(index); + let id = backend.queue_append(stream, fields.clone(), None).await; + expected.push(RuntimeQueueEntry { id, fields }); + } + for group in ["workers-a", "workers-b"] { + backend + .queue_ensure_consumer_group(stream, group, "0-0") + .await + .expect("consumer group"); + assert_eq!( + backend + .queue_read(stream, group, "reader", 3, None) + .await + .expect("read batch"), + expected + ); + } + assert_eq!( + backend + .queue_ack(stream, "workers-a", std::slice::from_ref(&expected[0].id)) + .await + .expect("ack only the first group"), + 1 + ); + assert_eq!( + backend + .queue_delete(stream, &[expected[1].id.clone(), "missing-0".to_string()]) + .await, + 1 + ); + age_memory_queue_pending(&backend, stream).await; + for (group, wanted) in [ + ("workers-a", vec![expected[2].clone()]), + ("workers-b", vec![expected[0].clone(), expected[2].clone()]), + ] { + let stats = backend.queue_stats(stream, Some(group)).await; + assert_eq!(stats.stream_length, 2); + assert_eq!(stats.group_pending, wanted.len() as u64); + assert_eq!( + backend + .queue_claim_stale( + stream, + group, + "reclaimer", + "0-0", + RuntimeQueueReclaimConfig { + min_idle_ms: 500, + count: 3 + }, + ) + .await + .expect("deleted entries cannot be reclaimed"), + wanted + ); + } + assert_eq!( + backend + .queue_delete(stream, &[expected[0].id.clone(), expected[2].id.clone()]) + .await, + 2 + ); + for group in ["workers-a", "workers-b"] { + let stats = backend.queue_stats(stream, Some(group)).await; + assert_eq!(stats.stream_length, 0); + assert_eq!(stats.group_pending, 0); + } + } + + #[tokio::test] + async fn memory_queue_read_retains_returned_fields_after_pending_trim() { + let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); + let stream = "queue:pending-trim"; + backend + .queue_ensure_consumer_group(stream, "workers", "0-0") + .await + .expect("consumer group"); + let mut expected = Vec::new(); + for index in 0..2 { + let fields = memory_queue_test_fields(index); + let id = backend.queue_append(stream, fields.clone(), Some(2)).await; + expected.push(RuntimeQueueEntry { id, fields }); + } + let delivered = backend + .queue_read(stream, "workers", "reader", 2, None) + .await + .expect("read batch before trim"); + let next_fields = memory_queue_test_fields(2); + let next_id = backend + .queue_append(stream, next_fields.clone(), Some(2)) + .await; + assert_eq!( + delivered, expected, + "trimming must not invalidate returned entries" + ); + age_memory_queue_pending(&backend, stream).await; + assert_eq!( + backend + .queue_claim_stale( + stream, + "workers", + "reclaimer", + "0-0", + RuntimeQueueReclaimConfig { + min_idle_ms: 500, + count: 2 + }, + ) + .await + .expect("trimmed entry is removed from the PEL"), + vec![expected[1].clone()] + ); + assert_eq!( + backend + .queue_read(stream, "workers", "reader", 2, None) + .await + .expect("read the remaining new entry"), + vec![RuntimeQueueEntry { + id: next_id, + fields: next_fields + }] + ); + } + #[tokio::test] async fn rate_limit_shard_amortizes_expired_entry_cleanup() { let backend = MemoryRuntimeBackend::new(MemoryRuntimeStateConfig::default()); diff --git a/crates/aether-runtime/state/src/redis/client.rs b/crates/aether-runtime/state/src/redis/client.rs index 42069b582..83381f140 100644 --- a/crates/aether-runtime/state/src/redis/client.rs +++ b/crates/aether-runtime/state/src/redis/client.rs @@ -1,9 +1,12 @@ use crate::error::RedisResultExt; use crate::redis::RedisKeyspace; use crate::DataLayerError; +use std::future::Future; +use std::pin::Pin; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use std::time::Duration; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; use tracing::info; pub(crate) type RedisClient = redis::Client; @@ -141,8 +144,8 @@ pub(crate) struct RedisConnectionRouter { fast: RedisManagedConnection, stream: Arc>, stream_next: Arc, - blocking_stream: Arc>, - blocking_stream_next: Arc, + blocking_stream: Arc, + usage_cleanup: Arc, admin: RedisManagedConnection, metrics: Arc, } @@ -152,7 +155,7 @@ impl std::fmt::Debug for RedisConnectionRouter { f.debug_struct("RedisConnectionRouter") .field("lanes", &["fast", "stream", "blocking_stream", "admin"]) .field("stream_lanes", &self.stream.len()) - .field("blocking_stream_lanes", &self.blocking_stream.len()) + .field("blocking_stream_lanes", &self.blocking_stream.capacity) .finish() } } @@ -182,7 +185,7 @@ impl RedisConnectionRouter { ) .await?; let stream_lanes = stream.len(); - let blocking_stream_lanes = blocking_stream.len(); + let blocking_stream_lanes = blocking_stream.capacity; info!( redis_lanes = "fast,stream,blocking_stream,admin", redis_stream_lanes = stream_lanes, @@ -194,7 +197,11 @@ impl RedisConnectionRouter { stream: Arc::new(stream), stream_next: Arc::new(AtomicUsize::new(0)), blocking_stream: Arc::new(blocking_stream), - blocking_stream_next: Arc::new(AtomicUsize::new(0)), + usage_cleanup: Arc::new(RedisBlockingStreamPool::new( + client, + command_timeout_ms, + vec![None, None], + )), admin, metrics: Arc::new(RedisConnectionMetrics::default()), }) @@ -208,13 +215,26 @@ impl RedisConnectionRouter { self.stream[index].clone() } RedisConnectionLane::BlockingStream => { - let index = next_lane_index(&self.blocking_stream_next, self.blocking_stream.len()); - self.blocking_stream[index].clone() + unreachable!("blocking stream commands require an exclusive connection lease") } RedisConnectionLane::Admin => self.admin.clone(), } } + pub(crate) async fn blocking_stream_connection( + &self, + ) -> Result { + self.blocking_stream.checkout().await + } + + // WATCH/MULTI state must never share a multiplexed connection with other callers. + // These lazy leases own their drivers, so cancellation closes an unfinished transaction. + pub(crate) async fn usage_cleanup_connection( + &self, + ) -> Result { + self.usage_cleanup.checkout().await + } + pub(crate) fn record_error(&self, lane: RedisConnectionLane) { self.metrics .for_lane(lane) @@ -257,6 +277,114 @@ impl RedisConnectionRouter { } } +struct RedisBlockingConnection { + connection: redis::aio::MultiplexedConnection, + driver: Pin + Send>>, +} + +impl RedisBlockingConnection { + async fn query(&mut self, command: &redis::Cmd) -> Result { + // Drive the connection inside its owning query, so cancellation drops the + // socket immediately instead of leaving a spawned driver with an old BLOCK. + tokio::select! { + result = command.query_async(&mut self.connection) => result.map_redis_err(), + () = self.driver.as_mut() => Err(DataLayerError::Redis( + "runtime redis blocking stream connection driver terminated".to_string(), + )), + } + } +} + +struct RedisBlockingStreamPool { + client: RedisClient, + command_timeout_ms: Option, + capacity: usize, + available: Mutex>>, + permits: Arc, +} + +impl RedisBlockingStreamPool { + fn new( + client: RedisClient, + command_timeout_ms: Option, + available: Vec>, + ) -> Self { + let capacity = available.len(); + Self { + client, + command_timeout_ms, + capacity, + available: Mutex::new(available), + permits: Arc::new(Semaphore::new(capacity)), + } + } + + async fn checkout(self: &Arc) -> Result { + let permit = Arc::clone(&self.permits) + .acquire_owned() + .await + .map_err(|_| { + DataLayerError::Redis("runtime redis blocking stream pool closed".to_string()) + })?; + let connection = self + .available + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .pop() + .expect("blocking stream permit must have an available slot"); + Ok(RedisBlockingStreamLease { + pool: Arc::clone(self), + connection, + reusable: false, + _permit: permit, + }) + } +} + +pub(crate) struct RedisBlockingStreamLease { + pool: Arc, + connection: Option, + reusable: bool, + _permit: OwnedSemaphorePermit, +} + +impl RedisBlockingStreamLease { + pub(crate) async fn query( + &mut self, + command: &redis::Cmd, + ) -> Result { + self.reusable = false; + if self.connection.is_none() { + self.connection = Some( + connect_blocking_stream_lane(&self.pool.client, self.pool.command_timeout_ms) + .await?, + ); + } + self.connection + .as_mut() + .expect("blocking stream connection initialized") + .query(command) + .await + } + + pub(crate) fn recycle(&mut self) { + self.reusable = true; + } +} + +impl Drop for RedisBlockingStreamLease { + fn drop(&mut self) { + let connection = self.connection.take().filter(|_| self.reusable); + self.pool + .available + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .push(connection); + // The permit is released after the slot is restored. An uncompleted + // query has already dropped both its connection and its owned driver. + } +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)] pub struct RedisLaneDiagnostics { pub lane: &'static str, @@ -406,21 +534,46 @@ async fn connect_blocking_stream_lanes( client: &RedisClient, command_timeout_ms: Option, requested_lanes: Option, -) -> Result, DataLayerError> { +) -> Result { let lane_count = blocking_stream_lane_count(requested_lanes)?; let mut lanes = Vec::with_capacity(lane_count); for _ in 0..lane_count { - lanes.push( - connect_lane( - client, - connection_manager_config(command_timeout_ms), - RedisConnectionLane::BlockingStream, - command_timeout_ms, - ) - .await?, - ); + lanes.push(Some( + connect_blocking_stream_lane(client, command_timeout_ms).await?, + )); } - Ok(lanes) + Ok(RedisBlockingStreamPool::new( + client.clone(), + command_timeout_ms, + lanes, + )) +} + +async fn connect_blocking_stream_lane( + client: &RedisClient, + command_timeout_ms: Option, +) -> Result { + let connect = client.create_multiplexed_tokio_connection(); + let result = if let Some(timeout_ms) = command_timeout_ms { + tokio::time::timeout(Duration::from_millis(timeout_ms), connect) + .await + .map_err(|_| { + DataLayerError::TimedOut(format!( + "runtime redis blocking_stream lane connection exceeded {timeout_ms}ms timeout" + )) + })? + } else { + connect.await + }; + let (connection, driver) = result.map_err(|err| { + DataLayerError::Redis(format!( + "failed to initialize runtime redis blocking_stream lane: {err}" + )) + })?; + Ok(RedisBlockingConnection { + connection, + driver: Box::pin(driver), + }) } fn blocking_stream_lane_count(requested_lanes: Option) -> Result { @@ -480,10 +633,12 @@ async fn connect_lane( mod tests { use super::{ blocking_stream_lane_count, default_blocking_stream_lane_count, next_lane_index, - stream_lane_count, RedisClientConfig, RedisClientFactory, RedisLaneMetrics, - DEFAULT_STREAM_LANES, MAX_BLOCKING_STREAM_LANES_CAP, REDIS_COMMAND_LATENCY_BUCKETS_MS, + stream_lane_count, RedisBlockingStreamPool, RedisClientConfig, RedisClientFactory, + RedisLaneMetrics, DEFAULT_STREAM_LANES, MAX_BLOCKING_STREAM_LANES_CAP, + REDIS_COMMAND_LATENCY_BUCKETS_MS, }; use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::Arc; use std::time::Duration; #[test] @@ -543,6 +698,85 @@ mod tests { assert!(blocking_stream_lane_count(Some(0)).is_err()); } + fn empty_blocking_pool(capacity: usize) -> Arc { + Arc::new(RedisBlockingStreamPool::new( + redis::Client::open("redis://127.0.0.1/0").expect("lazy client"), + None, + (0..capacity).map(|_| None).collect(), + )) + } + + #[tokio::test] + async fn blocking_stream_pool_cancelled_checkout_preserves_owner_and_capacity() { + use std::future::Future; + use std::task::Poll; + + let pool = empty_blocking_pool(1); + let owner = pool.checkout().await.expect("first lease"); + let mut waiting = Box::pin(pool.checkout()); + std::future::poll_fn(|cx| { + assert!(matches!(waiting.as_mut().poll(cx), Poll::Pending)); + Poll::Ready(()) + }) + .await; + drop(waiting); + assert_eq!(pool.permits.available_permits(), 0); + assert!(pool.available.lock().unwrap().is_empty()); + + drop(owner); + assert_eq!(pool.permits.available_permits(), 1); + let replacement = pool + .checkout() + .await + .expect("cancelled owner slot restored"); + assert!(replacement.connection.is_none()); + let panic = tokio::spawn(async move { + let _lease = replacement; + panic!("test owner panic"); + }); + assert!(panic.await.expect_err("owner panicked").is_panic()); + assert_eq!(pool.permits.available_permits(), 1); + assert_eq!(pool.available.lock().unwrap().len(), 1); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn blocking_stream_pool_concurrent_checkouts_stay_within_capacity() { + let pool = empty_blocking_pool(3); + let active = Arc::new(AtomicUsize::new(0)); + let peak = Arc::new(AtomicUsize::new(0)); + let barrier = Arc::new(tokio::sync::Barrier::new(32)); + let mut tasks = tokio::task::JoinSet::new(); + for _ in 0..32 { + let pool = Arc::clone(&pool); + let active = Arc::clone(&active); + let peak = Arc::clone(&peak); + let barrier = Arc::clone(&barrier); + tasks.spawn(async move { + barrier.wait().await; + for _ in 0..32 { + let lease = pool.checkout().await.expect("bounded lease"); + let concurrent = active.fetch_add(1, Ordering::SeqCst) + 1; + peak.fetch_max(concurrent, Ordering::SeqCst); + assert!(concurrent <= 3); + tokio::task::yield_now().await; + active.fetch_sub(1, Ordering::SeqCst); + drop(lease); + } + }); + } + tokio::time::timeout(Duration::from_secs(5), async { + while let Some(result) = tasks.join_next().await { + result.expect("checkout task"); + } + }) + .await + .expect("all checkouts complete without losing capacity"); + assert!(peak.load(Ordering::SeqCst) <= 3); + assert_eq!(active.load(Ordering::SeqCst), 0); + assert_eq!(pool.permits.available_permits(), 3); + assert_eq!(pool.available.lock().unwrap().len(), 3); + } + #[test] fn stream_lane_count_uses_fixed_default() { assert_eq!(stream_lane_count(), DEFAULT_STREAM_LANES); diff --git a/crates/aether-runtime/state/src/redis/dead_letter_transfer.lua b/crates/aether-runtime/state/src/redis/dead_letter_transfer.lua new file mode 100644 index 000000000..9e568cdf9 --- /dev/null +++ b/crates/aether-runtime/state/src/redis/dead_letter_transfer.lua @@ -0,0 +1,44 @@ +-- Redis scripts do not roll back errors. Complete all predictable checks before XADD. +if #KEYS ~= 2 or #ARGV < 4 or (#ARGV - 2) % 2 ~= 0 then + return redis.error_reply('ERR invalid pending transfer arguments') +end +if KEYS[1] == KEYS[2] then + return redis.error_reply('ERR pending transfer source and destination must differ') +end +if type(redis.acl_check_cmd) ~= 'function' then + return redis.error_reply('ERR pending transfer requires Redis 7 or later for ACL preflight') +end + +-- The exact range bounds keep work independent of the size of the PEL. +local pending = redis.call('XPENDING', KEYS[1], ARGV[1], ARGV[2], ARGV[2], 1) +if #pending == 0 then + return {0, '', 0, 0} +end +if pending[1][1] ~= ARGV[2] then + return redis.error_reply('ERR pending transfer requires an exact canonical entry ID') +end +local destination_type = redis.call('TYPE', KEYS[2]).ok +if destination_type ~= 'none' and destination_type ~= 'stream' then + return redis.error_reply('WRONGTYPE pending transfer destination must be a stream') +end + +local append = {'XADD', KEYS[2], '*'} +for index = 3, #ARGV do + append[#append + 1] = ARGV[index] +end +if not redis.acl_check_cmd(unpack(append)) then + return redis.error_reply('NOPERM pending transfer requires XADD permission') +end +if not redis.acl_check_cmd('XACK', KEYS[1], ARGV[1], ARGV[2]) then + return redis.error_reply('NOPERM pending transfer requires XACK permission') +end +if not redis.acl_check_cmd('XDEL', KEYS[1], ARGV[2]) then + return redis.error_reply('NOPERM pending transfer requires XDEL permission') +end + +-- PEL membership, stream types and ACLs cannot change between these commands. +-- A trimmed source body is still recoverable from the caller's retained fields. +local destination_id = redis.call(unpack(append)) +local acked = redis.call('XACK', KEYS[1], ARGV[1], ARGV[2]) +local deleted = redis.call('XDEL', KEYS[1], ARGV[2]) +return {1, destination_id, acked, deleted} diff --git a/crates/aether-runtime/state/src/redis/dead_letter_transfer_tests.rs b/crates/aether-runtime/state/src/redis/dead_letter_transfer_tests.rs new file mode 100644 index 000000000..c275666bc --- /dev/null +++ b/crates/aether-runtime/state/src/redis/dead_letter_transfer_tests.rs @@ -0,0 +1,500 @@ +use super::*; + +type TransferTestConnection = ::redis::aio::MultiplexedConnection; + +const TRANSFER_GROUP: &str = "transfer-workers"; +const TRANSFER_USER: &str = "transfer-worker"; + +async fn transfer_runtime( + protocol: &str, +) -> Option<(TestRedisServer, RuntimeState, TransferTestConnection)> { + let Some(server) = TestRedisServer::start().await else { + eprintln!( + "dead letter transfer {protocol} skipped: isolated Redis fixture unavailable; check AETHER_REDIS_SERVER_BIN" + ); + return None; + }; + let mut admin = ::redis::Client::open(server.redis_url.clone()) + .expect("transfer admin client") + .get_multiplexed_async_connection() + .await + .expect("transfer admin connection"); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg(TRANSFER_USER) + .arg("on") + .arg(">transfer-test-password") + .arg("~*") + .arg("+@all") + .query_async::<()>(&mut admin) + .await + .expect("transfer test user"); + let runtime = RuntimeState::redis_with_blocking_stream_lanes( + RedisClientConfig { + url: format!( + "redis://{TRANSFER_USER}:transfer-test-password@127.0.0.1:{}/5?protocol={protocol}", + server.port + ), + key_prefix: Some(format!("transfer-{protocol}")), + }, + Some(5_000), + Some(4), + ) + .await + .expect("authenticated transfer runtime"); + ::redis::cmd("SELECT") + .arg(5) + .query_async::<()>(&mut admin) + .await + .expect("transfer admin database"); + eprintln!( + "dead letter transfer fixture ready: protocol={protocol} db=5 port={} authenticated=true", + server.port + ); + Some((server, runtime, admin)) +} + +fn transfer_source_fields() -> BTreeMap { + BTreeMap::from([ + ( + "payload".to_string(), + "malformed\r\n\"quoted\"\\\u{4e2d}\u{6587}\0".to_string(), + ), + ("legacy".to_string(), "retain every field".to_string()), + ( + String::new(), + "empty field name is valid Redis data".to_string(), + ), + ]) +} + +fn transfer_archive_fields(entry: &RuntimeQueueEntry) -> BTreeMap { + BTreeMap::from([ + ( + "payload".to_string(), + serde_json::json!({ + "entry_id": entry.id, + "fields": entry.fields, + "error": "invalid record\r\n\"details\"\\\u{4e2d}\u{6587}" + }) + .to_string(), + ), + ("archive_version".to_string(), "1".to_string()), + ]) +} + +async fn seed_transfer_entry(runtime: &RuntimeState, source: &str) -> RuntimeQueueEntry { + RuntimeQueueStore::ensure_consumer_group(runtime, source, TRANSFER_GROUP, "0-0") + .await + .expect("source consumer group"); + let id = RuntimeQueueStore::append_fields_with_maxlen( + runtime, + source, + &transfer_source_fields(), + None, + ) + .await + .expect("source append"); + let mut entries = + RuntimeQueueStore::read_group(runtime, source, TRANSFER_GROUP, "owner", 1, None) + .await + .expect("pending source entry"); + assert_eq!(entries.len(), 1); + let entry = entries.pop().expect("one entry"); + assert_eq!(entry.id, id); + assert_eq!(entry.fields, transfer_source_fields()); + entry +} + +async fn transfer_entries( + admin: &mut TransferTestConnection, + stream: &str, +) -> Vec { + let rows = ::redis::cmd("XRANGE") + .arg(stream) + .arg("-") + .arg("+") + .query_async::<::redis::streams::StreamRangeReply>(admin) + .await + .expect("inspect transfer stream"); + rows.ids + .into_iter() + .map(|row| RuntimeQueueEntry { + id: row.id, + fields: row + .map + .into_iter() + .map(|(field, value)| { + let value = + ::redis::from_redis_value::(&value).expect("string field value"); + (field, value) + }) + .collect(), + }) + .collect() +} + +async fn transfer_pending( + admin: &mut TransferTestConnection, + source: &str, +) -> Vec<(String, String, u64, u64)> { + ::redis::cmd("XPENDING") + .arg(source) + .arg(TRANSFER_GROUP) + .arg("-") + .arg("+") + .arg(16) + .query_async(admin) + .await + .expect("inspect transfer pending entries") +} + +async fn transfer_entry( + runtime: &RuntimeState, + source: &str, + entry: &RuntimeQueueEntry, + destination: &str, + fields: &BTreeMap, +) -> Result { + RuntimeQueueStore::try_transfer_pending_to_stream( + runtime, + source, + TRANSFER_GROUP, + &entry.id, + destination, + fields, + ) + .await + .map(|outcome| outcome.expect("Redis implements atomic transfer")) +} + +async fn assert_transfer_source_unchanged( + admin: &mut TransferTestConnection, + source: &str, + entry: &RuntimeQueueEntry, +) { + assert_eq!(transfer_entries(admin, source).await, [entry.clone()]); + let pending = transfer_pending(admin, source).await; + assert_eq!(pending.len(), 1); + assert_eq!(pending[0].0, entry.id); + assert_eq!(pending[0].1, "owner"); +} + +async fn assert_transfer_completed( + admin: &mut TransferTestConnection, + source: &str, + destination: &str, + fields: &BTreeMap, +) { + assert!(transfer_entries(admin, source).await.is_empty()); + assert!(transfer_pending(admin, source).await.is_empty()); + let archived = transfer_entries(admin, destination).await; + assert_eq!(archived.len(), 1); + assert_eq!(&archived[0].fields, fields); +} + +#[tokio::test] +async fn redis_dead_letter_transfer_concurrent_consumers_archive_exactly_once() { + for protocol in ["resp2", "resp3"] { + let Some((_server, runtime, mut admin)) = transfer_runtime(protocol).await else { + return; + }; + let source = "usage:{transfer}:concurrent"; + let destination = "usage:{transfer}:concurrent:dlq"; + let entry = seed_transfer_entry(&runtime, source).await; + let fields = transfer_archive_fields(&entry); + let barrier = Arc::new(tokio::sync::Barrier::new(16)); + let mut tasks = tokio::task::JoinSet::new(); + for _ in 0..16 { + let runtime = runtime.clone(); + let entry = entry.clone(); + let fields = fields.clone(); + let barrier = Arc::clone(&barrier); + tasks.spawn(async move { + barrier.wait().await; + transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect("concurrent transfer") + }); + } + let mut transferred_ids = Vec::new(); + let mut not_pending = 0; + while let Some(result) = tasks.join_next().await { + match result.expect("transfer task") { + RuntimeQueueTransferOutcome::Transferred { + destination_id, + acked, + deleted, + } => { + assert_eq!((acked, deleted), (1, 1)); + transferred_ids.push(destination_id); + } + RuntimeQueueTransferOutcome::NotPending => not_pending += 1, + } + } + assert_eq!(not_pending, 15); + assert_eq!(transferred_ids.len(), 1); + assert_transfer_completed(&mut admin, source, destination, &fields).await; + assert_eq!( + transfer_entries(&mut admin, destination).await[0].id, + transferred_ids[0] + ); + let diagnostics = runtime.redis_diagnostics().await.unwrap().unwrap(); + let lane = diagnostics + .lanes + .iter() + .find(|lane| lane.lane == "blocking_stream") + .expect("exclusive stream lane diagnostics"); + assert!(lane.command_count >= 16); + } +} + +#[tokio::test] +async fn redis_dead_letter_transfer_retry_after_ignored_success_does_not_archive_again() { + let Some((_server, runtime, mut admin)) = transfer_runtime("resp2").await else { + return; + }; + let source = "usage:{transfer}:retry"; + let destination = "usage:{transfer}:retry:dlq"; + let entry = seed_transfer_entry(&runtime, source).await; + let fields = transfer_archive_fields(&entry); + // Commit the operation, but discard its result as a caller missing the reply would. + let _ = transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect("first transfer commits"); + let changed_fields = BTreeMap::from([("payload".to_string(), "retry value".to_string())]); + assert_eq!( + transfer_entry(&runtime, source, &entry, destination, &changed_fields) + .await + .expect("retry is successful"), + RuntimeQueueTransferOutcome::NotPending + ); + assert_transfer_completed(&mut admin, source, destination, &fields).await; +} + +#[tokio::test] +async fn redis_dead_letter_transfer_acl_preflight_prevents_partial_writes() { + let Some((_server, runtime, mut admin)) = transfer_runtime("resp3").await else { + return; + }; + for forbidden in ["XADD", "XACK", "XDEL"] { + let source = format!("usage:{{transfer}}:acl-{forbidden}"); + let destination = format!("{source}:dlq"); + let entry = seed_transfer_entry(&runtime, &source).await; + let fields = transfer_archive_fields(&entry); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg(TRANSFER_USER) + .arg(format!("-{forbidden}")) + .query_async::<()>(&mut admin) + .await + .expect("deny one write command"); + let error = transfer_entry(&runtime, &source, &entry, &destination, &fields) + .await + .expect_err("denied write must fail before archiving"); + assert!( + error + .to_string() + .contains(&format!("requires {forbidden} permission")), + "expected the {forbidden} preflight error, got {error}" + ); + assert_transfer_source_unchanged(&mut admin, &source, &entry).await; + assert!(transfer_entries(&mut admin, &destination).await.is_empty()); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg(TRANSFER_USER) + .arg(format!("+{forbidden}")) + .query_async::<()>(&mut admin) + .await + .expect("restore one write command"); + assert!(matches!( + transfer_entry(&runtime, &source, &entry, &destination, &fields) + .await + .expect("retry after restoring permission"), + RuntimeQueueTransferOutcome::Transferred { + acked: 1, + deleted: 1, + .. + } + )); + assert_eq!( + transfer_entry(&runtime, &source, &entry, &destination, &fields) + .await + .expect("idempotent retry"), + RuntimeQueueTransferOutcome::NotPending + ); + assert_transfer_completed(&mut admin, &source, &destination, &fields).await; + } +} + +#[tokio::test] +async fn redis_dead_letter_transfer_invalid_state_preserves_source_until_repaired() { + let Some((_server, runtime, mut admin)) = transfer_runtime("resp2").await else { + return; + }; + let source = "usage:{transfer}:invalid"; + let destination = "usage:{transfer}:invalid:dlq"; + let entry = seed_transfer_entry(&runtime, source).await; + let fields = transfer_archive_fields(&entry); + ::redis::cmd("SET") + .arg(destination) + .arg("existing non-stream data") + .query_async::<()>(&mut admin) + .await + .expect("wrong-type destination"); + let error = transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect_err("wrong type must not acknowledge source"); + assert!(error.to_string().contains("WRONGTYPE")); + assert_transfer_source_unchanged(&mut admin, source, &entry).await; + assert_eq!( + ::redis::cmd("GET") + .arg(destination) + .query_async::(&mut admin) + .await + .unwrap(), + "existing non-stream data" + ); + ::redis::cmd("DEL") + .arg(destination) + .query_async::(&mut admin) + .await + .expect("repair destination type"); + + let error = RuntimeQueueStore::try_transfer_pending_to_stream( + &runtime, + source, + "missing-group", + &entry.id, + destination, + &fields, + ) + .await + .expect_err("missing group must fail before archiving"); + assert!(error.to_string().contains("NOGROUP")); + for invalid_id in [ + "", + "-", + "+", + "1", + "01-0", + "1-00", + "(1-0", + "1-+0", + "18446744073709551616-0", + ] { + assert!(matches!( + RuntimeQueueStore::try_transfer_pending_to_stream( + &runtime, + source, + TRANSFER_GROUP, + invalid_id, + destination, + &fields, + ) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + } + assert!(matches!( + transfer_entry(&runtime, source, &entry, source, &fields).await, + Err(DataLayerError::InvalidInput(_)) + )); + assert!(matches!( + transfer_entry(&runtime, source, &entry, destination, &BTreeMap::new()).await, + Err(DataLayerError::InvalidInput(_)) + )); + assert_transfer_source_unchanged(&mut admin, source, &entry).await; + assert!(transfer_entries(&mut admin, destination).await.is_empty()); + assert!(matches!( + transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect("valid transfer after failed attempts"), + RuntimeQueueTransferOutcome::Transferred { + acked: 1, + deleted: 1, + .. + } + )); + assert_transfer_completed(&mut admin, source, destination, &fields).await; +} + +#[tokio::test] +async fn redis_dead_letter_transfer_archives_retained_fields_after_source_trim() { + for protocol in ["resp2", "resp3"] { + let Some((_server, runtime, mut admin)) = transfer_runtime(protocol).await else { + return; + }; + let source = "usage:{transfer}:trimmed"; + let destination = "usage:{transfer}:trimmed:dlq"; + let entry = seed_transfer_entry(&runtime, source).await; + let fields = transfer_archive_fields(&entry); + let trimmed = ::redis::cmd("XTRIM") + .arg(source) + .arg("MAXLEN") + .arg(0) + .query_async::(&mut admin) + .await + .expect("trim source body while retaining PEL"); + assert_eq!(trimmed, 1); + assert!(transfer_entries(&mut admin, source).await.is_empty()); + assert_eq!(transfer_pending(&mut admin, source).await[0].0, entry.id); + assert!(matches!( + transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect("pending body remains recoverable from caller fields"), + RuntimeQueueTransferOutcome::Transferred { + acked: 1, + deleted: 0, + .. + } + )); + assert_eq!( + transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect("trimmed entry retry"), + RuntimeQueueTransferOutcome::NotPending + ); + assert_transfer_completed(&mut admin, source, destination, &fields).await; + } +} + +#[tokio::test] +async fn redis_dead_letter_transfer_preserves_ids_larger_than_lua_integer_precision() { + let Some((_server, runtime, mut admin)) = transfer_runtime("resp3").await else { + return; + }; + let source = "usage:{transfer}:large-id"; + let destination = "usage:{transfer}:large-id:dlq"; + let entry_id = "9007199254740993-18446744073709551614"; + RuntimeQueueStore::ensure_consumer_group(&runtime, source, TRANSFER_GROUP, "0-0") + .await + .expect("large-ID source group"); + ::redis::cmd("XADD") + .arg(source) + .arg(entry_id) + .arg("payload") + .arg("retained value") + .query_async::(&mut admin) + .await + .expect("large stream ID"); + let mut entries = + RuntimeQueueStore::read_group(&runtime, source, TRANSFER_GROUP, "owner", 1, None) + .await + .expect("read large-ID entry"); + assert_eq!(entries.len(), 1); + let entry = entries.pop().unwrap(); + assert_eq!(entry.id, entry_id); + let fields = transfer_archive_fields(&entry); + assert!(matches!( + transfer_entry(&runtime, source, &entry, destination, &fields) + .await + .expect("transfer exact large ID"), + RuntimeQueueTransferOutcome::Transferred { + acked: 1, + deleted: 1, + .. + } + )); + assert_transfer_completed(&mut admin, source, destination, &fields).await; +} diff --git a/crates/aether-runtime/state/src/redis/mod.rs b/crates/aether-runtime/state/src/redis/mod.rs index 69f5c1cd8..13bbca65b 100644 --- a/crates/aether-runtime/state/src/redis/mod.rs +++ b/crates/aether-runtime/state/src/redis/mod.rs @@ -4,6 +4,7 @@ mod lock; mod namespace; mod runtime; mod stream; +mod usage_cleanup; pub use client::{RedisClientConfig, RedisLaneDiagnostics}; pub use kv::{RedisKvRunner, RedisKvRunnerConfig}; diff --git a/crates/aether-runtime/state/src/redis/runtime.rs b/crates/aether-runtime/state/src/redis/runtime.rs index 4be7d84fb..e111ce87a 100644 --- a/crates/aether-runtime/state/src/redis/runtime.rs +++ b/crates/aether-runtime/state/src/redis/runtime.rs @@ -7,9 +7,13 @@ use crate::redis::{ }; use crate::{ DataLayerError, RateLimitCheck, RateLimitInput, RateLimitScope, RuntimeSemaphoreError, - UsageLimitCheck, UsageLimitInput, UsageLimitReleaseInput, + ScoreWindowU64Stats, UsageLimitCheck, UsageLimitInput, UsageLimitReleaseInput, + SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT, }; +const SCORE_WINDOW_STATS_SCRIPT: &str = include_str!("score_window.lua"); +const SCORE_WINDOW_STATS_PIPELINE_KEY_LIMIT: usize = 16; + const RATE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#" local user_key = KEYS[1] local key_key = KEYS[2] @@ -52,15 +56,119 @@ end return {1, 0, 0, remaining} "#; -const USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#" +pub(super) const USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#" local count = #KEYS local now = tonumber(ARGV[1]) local event_id = ARGV[2] +-- Large mixed windows are copied in bounded commands on an exclusive WATCH +-- connection. This read-only pass must finish before pruning any of the rules. +if ARGV[count * 3 + 3] ~= 'inline' and redis.acl_check_cmd then + local deferred = {2} + for i = 1, count do + local key = KEYS[i] + if redis.call('ZCARD', key) > 4096 then + local cutoff = now - tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 + local expired = redis.pcall('ZCOUNT', key, '-inf', cutoff) + if type(expired) == 'number' and expired > 4096 then + local live = redis.call('ZCARD', key) - expired + local temporary = key .. ':__usage_copy:acl' + if live > 256 + and redis.acl_check_cmd('WATCH', key) + and redis.acl_check_cmd('UNWATCH') + and redis.acl_check_cmd('MULTI') + and redis.acl_check_cmd('EXEC') + and redis.acl_check_cmd('EVAL', 'return 1', 0) + and redis.acl_check_cmd('EXISTS', temporary) + and redis.acl_check_cmd('ZRANGE', key, 0, 511, 'WITHSCORES') + and redis.acl_check_cmd('ZADD', temporary, 0, 'acl') + and redis.acl_check_cmd('PTTL', key) + and redis.acl_check_cmd('PEXPIRE', temporary, 60000) + and redis.acl_check_cmd('PERSIST', temporary) + and redis.acl_check_cmd('UNLINK', key, temporary) + and redis.acl_check_cmd('RENAME', temporary, key) then + deferred[#deferred + 1] = i + deferred[#deferred + 1] = expired + deferred[#deferred + 1] = live + end + end + end + end + if #deferred > 1 then return deferred end +end + +local function replace_with_survivors(key, live) + if not redis.acl_check_cmd then return false end + local temporary = key .. ':__usage_trim' + local exists = redis.pcall('EXISTS', temporary) + local ttl = redis.pcall('PTTL', key) + if exists ~= 0 or type(ttl) ~= 'number' or ttl == 0 or ttl < -1 then + return false + end + local rows = redis.call('ZRANGE', key, -live, -1, 'WITHSCORES') + local args = {} + for i = 1, #rows, 2 do + args[#args + 1] = rows[i + 1] + args[#args + 1] = rows[i] + end + -- Validate every write before detaching the original. The temporary key + -- is fully built first; restricted ACLs keep the original cleanup path. + if not redis.acl_check_cmd('ZADD', temporary, unpack(args)) + or not redis.acl_check_cmd('UNLINK', key) + or not redis.acl_check_cmd('UNLINK', temporary) + or not redis.acl_check_cmd('RENAME', temporary, key) + or (ttl > 0 and not redis.acl_check_cmd('PEXPIRE', temporary, ttl)) then + return false + end + if type(redis.pcall('ZADD', temporary, unpack(args))) ~= 'number' then + return false + end + if ttl > 0 and redis.pcall('PEXPIRE', temporary, ttl) ~= 1 then + redis.call('UNLINK', temporary) + return false + end + if type(redis.pcall('UNLINK', key)) ~= 'number' then + redis.call('UNLINK', temporary) + return false + end + redis.call('RENAME', temporary, key) + return true +end + +local function prune_window(key, cutoff) + local cardinality = redis.call('ZCARD', key) + if cardinality == 0 then return 0 end + if cardinality > 256 then + local earliest = redis.call('ZRANGE', key, 0, 0, 'WITHSCORES') + if #earliest >= 2 and tonumber(earliest[2]) > cutoff then + return cardinality + end + local latest = redis.call('ZRANGE', key, -1, -1, 'WITHSCORES') + if #latest >= 2 and tonumber(latest[2]) <= cutoff then + -- The entire window is expired. Redis can free the detached object + -- off its command thread while this key is reused. + local result = redis.pcall('UNLINK', key) + if type(result) == 'number' then return 0 end + elseif cardinality > 1024 then + local expired = redis.pcall('ZCOUNT', key, '-inf', cutoff) + if type(expired) == 'number' then + local live = cardinality - expired + if live > 0 and live <= 256 and expired > live * 4 + and replace_with_survivors(key, live) then + return live + end + end + end + end + -- Preserve the existing path for large live windows and restricted ACLs. + -- Partial deferred deletion could resurrect entries on out-of-order calls. + return cardinality - redis.call('ZREMRANGEBYSCORE', key, '-inf', cutoff) +end + +local current_counts = {} for i = 1, count do local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 - local cutoff = now - window_ms - redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', cutoff) + current_counts[i] = prune_window(KEYS[i], now - window_ms) end for i = 1, count do @@ -68,7 +176,7 @@ for i = 1, count do local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 local already_consumed = redis.call('ZSCORE', KEYS[i], event_id) if not already_consumed then - local current = redis.call('ZCARD', KEYS[i]) + local current = current_counts[i] if current >= limit then local earliest = redis.call('ZRANGE', KEYS[i], 0, 0, 'WITHSCORES') local retry_after = 1 @@ -362,30 +470,17 @@ impl RedisRuntimeRunner { &self, input: UsageLimitInput<'_>, ) -> Result { - let script = script(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT); - let mut invocation = script.prepare_invoke(); - for rule in input.rules { - invocation.key(self.keyspace.key(rule.key)); - } - invocation.arg(input.now_unix_ms as i64); - invocation.arg(input.event_id); - for rule in input.rules { - invocation.arg(rule.limit as i64); - invocation.arg(rule.window_seconds as i64); - invocation.arg(rule.retention_seconds as i64); - } + let keys = input + .rules + .iter() + .map(|rule| self.keyspace.key(rule.key)) + .collect::>(); let raw = run_lane_with_timeout( &self.connections, RedisConnectionLane::Fast, - self.command_timeout_ms, + Some(self.command_timeout_ms.unwrap_or(30_000)), "runtime usage limit check", - async { - let mut connection = self.connections.connection(RedisConnectionLane::Fast); - invocation - .invoke_async::>(&mut connection) - .await - .map_redis_err() - }, + super::usage_cleanup::check_and_consume(&self.connections, &keys, &input), ) .await?; match raw.first().copied() { @@ -539,6 +634,72 @@ impl RedisRuntimeRunner { .await } + pub(crate) async fn score_window_u64_stats_by_min( + &self, + keys: &[String], + min_score: f64, + ) -> Result>, DataLayerError> { + let script = script(SCORE_WINDOW_STATS_SCRIPT); + let mut output = Vec::with_capacity(keys.len()); + for batch in keys.chunks(SCORE_WINDOW_STATS_PIPELINE_KEY_LIMIT) { + let values: Vec<(u8, String, u64)> = run_lane_with_timeout( + &self.connections, + RedisConnectionLane::Admin, + self.command_timeout_ms, + "runtime score window stats", + async { + let mut connection = self.connections.connection(RedisConnectionLane::Admin); + let mut pipeline = redis::pipe(); + for key in batch { + pipeline + .cmd("EVALSHA") + .arg(script.get_hash()) + .arg(1) + .arg(self.keyspace.key(key)) + .arg(min_score) + .arg(SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT); + } + match pipeline.query_async(&mut connection).await { + Err(err) if err.kind() == redis::ErrorKind::NoScriptError => { + script + .prepare_invoke() + .load_async(&mut connection) + .await + .map_redis_err()?; + pipeline.query_async(&mut connection).await.map_redis_err() + } + result => result.map_redis_err(), + } + }, + ) + .await?; + if values.len() != batch.len() { + return Err(DataLayerError::UnexpectedValue( + "runtime score window stats result count mismatch".to_string(), + )); + } + for (aggregated, sum, positive_count) in values { + output.push(match aggregated { + 0 => None, + 1 => Some(ScoreWindowU64Stats { + sum: sum.parse().map_err(|_| { + DataLayerError::UnexpectedValue( + "runtime score window stats returned an invalid sum".to_string(), + ) + })?, + positive_count, + }), + _ => { + return Err(DataLayerError::UnexpectedValue( + "runtime score window stats returned an invalid status".to_string(), + )) + } + }); + } + } + Ok(output) + } + pub(crate) async fn score_remove_by_score( &self, key: &str, diff --git a/crates/aether-runtime/state/src/redis/score_window.lua b/crates/aether-runtime/state/src/redis/score_window.lua new file mode 100644 index 000000000..e330b50b6 --- /dev/null +++ b/crates/aether-runtime/state/src/redis/score_window.lua @@ -0,0 +1,36 @@ +-- Count before reading members so a large window cannot run an unbounded Lua loop. +local count = redis.call('ZCOUNT', KEYS[1], ARGV[1], '+inf') +if count > tonumber(ARGV[2]) then + return {0, '0', 0} +end + +local members = redis.call('ZRANGEBYSCORE', KEYS[1], ARGV[1], '+inf') +local high, low, positive = 0, 0, 0 +local max_high, max_low = 18446744073, 709551615 +for _, member in ipairs(members) do + local value = string.match(member, ':([^:]*)$') + if value and string.match(value, '^%+?%d+$') then + value = string.gsub(value, '^%+', '') + value = string.gsub(value, '^0+', '') + if #value > 0 and (#value < 20 or (#value == 20 and value <= '18446744073709551615')) then + positive = positive + 1 + -- Two base-1e9 limbs keep every integer operation exactly representable + -- in Redis Lua's doubles, including values above 2^53 and u64::MAX. + local split = math.max(0, #value - 9) + local value_high = tonumber(string.sub(value, 1, split)) or 0 + local value_low = tonumber(string.sub(value, split + 1)) + low = low + value_low + high = high + value_high + math.floor(low / 1000000000) + low = low % 1000000000 + if high > max_high or (high == max_high and low > max_low) then + high, low = max_high, max_low + end + end + end +end + +local total = string.format('%.0f', low) +if high > 0 then + total = string.format('%.0f', high) .. string.format('%09d', low) +end +return {1, total, positive} diff --git a/crates/aether-runtime/state/src/redis/stream.rs b/crates/aether-runtime/state/src/redis/stream.rs index 363a5f555..8ceb42878 100644 --- a/crates/aether-runtime/state/src/redis/stream.rs +++ b/crates/aether-runtime/state/src/redis/stream.rs @@ -1,16 +1,19 @@ -use std::collections::BTreeMap; +use std::collections::{BTreeMap, HashMap}; use std::future::Future; -use redis::from_redis_value; -use redis::streams::StreamReadReply; use redis::Value as RedisValue; +use redis::{from_owned_redis_value, from_redis_value}; use crate::error::{redis_error, RedisResultExt}; use crate::redis::{ run_lane_with_timeout, RedisClientConfig, RedisClientFactory, RedisConnectionLane, RedisConnectionRouter, RedisKeyspace, }; -use crate::{DataLayerError, RuntimeQueueStats}; +use crate::{ + validate_runtime_queue_transfer, DataLayerError, RuntimeQueueStats, RuntimeQueueTransferOutcome, +}; + +const DEAD_LETTER_TRANSFER_SCRIPT: &str = include_str!("dead_letter_transfer.lua"); #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct RedisStreamName(pub String); @@ -259,6 +262,41 @@ impl RedisStreamRunner { self.append_fields(stream, &fields).await } + /// Atomically archives a pending entry and removes it from its source group and stream. + /// Requires Redis 7+ for ACL preflight. Both stream keys must share a slot on Redis Cluster. + pub async fn try_transfer_pending_to_stream( + &self, + source: &str, + group: &str, + entry_id: &str, + destination: &str, + destination_fields: &BTreeMap, + ) -> Result { + validate_runtime_queue_transfer(source, group, entry_id, destination, destination_fields)?; + self.run_with_timeout( + RedisConnectionLane::BlockingStream, + "redis stream pending transfer", + async { + let mut lease = self.connections.blocking_stream_connection().await?; + let mut command = redis::cmd("EVAL"); + command + .arg(DEAD_LETTER_TRANSFER_SCRIPT) + .arg(2) + .arg(source) + .arg(destination) + .arg(group) + .arg(entry_id); + for (field, value) in destination_fields { + command.arg(field).arg(value); + } + let outcome = parse_transfer_result(lease.query(&command).await?)?; + lease.recycle(); + Ok(outcome) + }, + ) + .await + } + pub async fn read_group( &self, stream: &RedisStreamName, @@ -275,7 +313,6 @@ impl RedisStreamRunner { RedisConnectionLane::Stream }; self.run_with_timeout(lane, "redis stream read group", async { - let mut connection = self.connections.connection(lane); let mut command = redis::cmd("XREADGROUP"); command .arg("GROUP") @@ -288,6 +325,14 @@ impl RedisStreamRunner { } command.arg("STREAMS").arg(&stream.0).arg(">"); + if lane == RedisConnectionLane::BlockingStream { + let mut lease = self.connections.blocking_stream_connection().await?; + let entries = parse_stream_read_entries(lease.query(&command).await?)?; + lease.recycle(); + return Ok(entries); + } + + let mut connection = self.connections.connection(lane); let reply = command .query_async::(&mut connection) .await @@ -364,22 +409,25 @@ impl RedisStreamRunner { validate_stream_position(start_id)?; config.validate()?; - self.run_with_timeout(RedisConnectionLane::Stream, "redis stream reclaim", async { - let mut connection = self.connections.connection(RedisConnectionLane::Stream); - let reply = redis::cmd("XAUTOCLAIM") - .arg(&stream.0) - .arg(&group.0) - .arg(&consumer.0) - .arg(config.min_idle_ms) - .arg(start_id) - .arg("COUNT") - .arg(config.count) - .query_async::(&mut connection) - .await - .map_redis_err()?; - - parse_reclaim_result(reply) - }) + self.run_with_timeout( + RedisConnectionLane::BlockingStream, + "redis stream reclaim", + async { + let mut lease = self.connections.blocking_stream_connection().await?; + let mut command = redis::cmd("XAUTOCLAIM"); + command + .arg(&stream.0) + .arg(&group.0) + .arg(&consumer.0) + .arg(config.min_idle_ms) + .arg(start_id) + .arg("COUNT") + .arg(config.count); + let result = parse_reclaim_result(lease.query(&command).await?)?; + lease.recycle(); + Ok(result) + }, + ) .await } @@ -536,18 +584,21 @@ fn parse_stream_read_entries(value: RedisValue) -> Result, return Ok(Vec::new()); } - let reply = from_redis_value::(&value).map_err(redis_error)?; - Ok(reply - .keys + // StreamReadReply in redis 0.28 falls back to borrowed conversion even for an + // owned input. Its underlying containers support moving every payload buffer. + type StreamReadRows = Vec>>>>; + let rows = from_owned_redis_value::(value).map_err(redis_error)?; + Ok(rows .into_iter() - .flat_map(|key| key.ids.into_iter()) - .map(|id| RedisStreamEntry { - id: id.id, - fields: id - .map + .flat_map(HashMap::into_values) + .flatten() + .flat_map(HashMap::into_iter) + .map(|(id, fields)| RedisStreamEntry { + id, + fields: fields .into_iter() .filter_map(|(field, value)| { - redis::from_redis_value::(&value) + from_owned_redis_value::(value) .ok() .map(|text| (field, text)) }) @@ -556,6 +607,31 @@ fn parse_stream_read_entries(value: RedisValue) -> Result, .collect()) } +fn parse_transfer_result(value: RedisValue) -> Result { + if !matches!(&value, RedisValue::Array(parts) if parts.len() == 4) { + return Err(DataLayerError::UnexpectedValue( + "redis stream pending transfer returned an invalid result shape".to_string(), + )); + } + let (transferred, destination_id, acked, deleted) = + from_owned_redis_value::<(i64, String, usize, usize)>(value).map_err(redis_error)?; + match transferred { + 0 if destination_id.is_empty() && acked == 0 && deleted == 0 => { + Ok(RuntimeQueueTransferOutcome::NotPending) + } + 1 if !destination_id.is_empty() && acked == 1 && deleted <= 1 => { + Ok(RuntimeQueueTransferOutcome::Transferred { + destination_id, + acked, + deleted, + }) + } + _ => Err(DataLayerError::UnexpectedValue( + "redis stream pending transfer returned inconsistent outcome fields".to_string(), + )), + } +} + fn parse_reclaim_result(value: RedisValue) -> Result { let RedisValue::Array(parts) = value else { return Err(DataLayerError::UnexpectedValue( @@ -570,9 +646,13 @@ fn parse_reclaim_result(value: RedisValue) -> Result parse_string_array(value, "redis xautoclaim deleted_ids")?, None => Vec::new(), }; @@ -584,9 +664,9 @@ fn parse_reclaim_result(value: RedisValue) -> Result Result, DataLayerError> { +fn parse_reclaim_entries(value: RedisValue) -> Result, DataLayerError> { match value { - RedisValue::Array(entries) => entries.iter().map(parse_reclaim_entry).collect(), + RedisValue::Array(entries) => entries.into_iter().map(parse_reclaim_entry).collect(), RedisValue::Nil => Ok(Vec::new()), _ => Err(DataLayerError::UnexpectedValue( "redis xautoclaim entries payload was not an array".to_string(), @@ -594,7 +674,7 @@ fn parse_reclaim_entries(value: &RedisValue) -> Result, Da } } -fn parse_reclaim_entry(value: &RedisValue) -> Result { +fn parse_reclaim_entry(value: RedisValue) -> Result { let RedisValue::Array(parts) = value else { return Err(DataLayerError::UnexpectedValue( "redis xautoclaim entry was not an array".to_string(), @@ -607,13 +687,20 @@ fn parse_reclaim_entry(value: &RedisValue) -> Result Result, DataLayerError> { match value { @@ -625,19 +712,23 @@ fn parse_string_map( ))); } let mut fields = BTreeMap::new(); - for pair in values.chunks(2) { - let key = parse_string_value(&pair[0], context)?; - let value = parse_string_value(&pair[1], context)?; + let mut values = values.into_iter(); + while let Some(key) = values.next() { + let key = parse_owned_string_value(key, context)?; + let value = parse_owned_string_value( + values.next().expect("validated even number of fields"), + context, + )?; fields.insert(key, value); } Ok(fields) } RedisValue::Map(entries) => entries - .iter() + .into_iter() .map(|(key, value)| { Ok(( - parse_string_value(key, context)?, - parse_string_value(value, context)?, + parse_owned_string_value(key, context)?, + parse_owned_string_value(value, context)?, )) }) .collect(), @@ -648,11 +739,11 @@ fn parse_string_map( } } -fn parse_string_array(value: &RedisValue, context: &str) -> Result, DataLayerError> { +fn parse_string_array(value: RedisValue, context: &str) -> Result, DataLayerError> { match value { RedisValue::Array(values) => values - .iter() - .map(|value| parse_string_value(value, context)) + .into_iter() + .map(|value| parse_owned_string_value(value, context)) .collect(), RedisValue::Nil => Ok(Vec::new()), _ => Err(DataLayerError::UnexpectedValue(format!( @@ -669,6 +760,14 @@ fn parse_string_value(value: &RedisValue, context: &str) -> Result Result { + from_owned_redis_value::(value).map_err(|err| { + DataLayerError::UnexpectedValue(format!( + "{context} was not a string-compatible redis value: {err}" + )) + }) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] struct RedisXInfoGroupStats { pending: Option, @@ -788,19 +887,60 @@ fn parse_u64_value(value: &RedisValue, context: &str) -> Result RedisValue { + RedisValue::BulkString(text.as_bytes().to_vec()) +} + +fn fields_reply(fields: Vec<(RedisValue, RedisValue)>, resp3: bool) -> RedisValue { + if resp3 { + RedisValue::Map(fields) + } else { + RedisValue::Array( + fields + .into_iter() + .flat_map(|(key, value)| [key, value]) + .collect(), + ) + } +} + +fn read_reply(id: RedisValue, fields: RedisValue, resp3: bool) -> RedisValue { + let entries = RedisValue::Array(vec![RedisValue::Array(vec![id, fields])]); + if resp3 { + RedisValue::Map(vec![(bulk("usage:events"), entries)]) + } else { + RedisValue::Array(vec![RedisValue::Array(vec![bulk("usage:events"), entries])]) + } +} + +fn reclaim_reply(id: RedisValue, fields: RedisValue) -> RedisValue { + RedisValue::Array(vec![ + bulk("0-0"), + RedisValue::Array(vec![RedisValue::Array(vec![id, fields])]), + RedisValue::Nil, + ]) +} + +fn original_read_parser(value: &RedisValue) -> Result, DataLayerError> { + let reply = from_redis_value::(value).map_err(crate::error::redis_error)?; + Ok(reply + .keys + .into_iter() + .flat_map(|key| key.ids) + .map(|entry| RedisStreamEntry { + id: entry.id, + fields: entry + .map + .into_iter() + .filter_map(|(field, value)| { + from_redis_value::(&value) + .ok() + .map(|value| (field, value)) + }) + .collect(), + }) + .collect()) +} + +#[test] +fn owned_read_moves_large_payload_and_id_buffers_in_resp2_and_resp3() { + for resp3 in [false, true] { + let mut payload = Vec::with_capacity(512 * 1024); + payload.extend_from_slice(b"{\"message\":\""); + payload.resize(256 * 1024, b'x'); + payload.extend_from_slice(b"\",\"cache_read_input_tokens\":0}\n"); + let expected = payload.clone(); + let pointer = payload.as_ptr(); + let capacity = payload.capacity(); + let mut id = Vec::with_capacity(64); + id.extend_from_slice(b"1710000000000-0"); + let id_pointer = id.as_ptr(); + let id_capacity = id.capacity(); + let fields = fields_reply( + vec![(bulk("payload"), RedisValue::BulkString(payload))], + resp3, + ); + let parsed = + parse_stream_read_entries(read_reply(RedisValue::BulkString(id), fields, resp3)) + .expect("read reply"); + assert_eq!(parsed.len(), 1); + let body = &parsed[0].fields["payload"]; + assert_eq!(body.as_bytes(), expected); + assert_eq!( + body.as_ptr(), + pointer, + "the RESP payload allocation must be reused" + ); + assert_eq!(body.capacity(), capacity); + assert_eq!(parsed[0].id.as_ptr(), id_pointer); + assert_eq!(parsed[0].id.capacity(), id_capacity); + } +} + +#[test] +fn owned_reclaim_moves_payload_and_all_id_buffers() { + for resp3 in [false, true] { + let mut payload = Vec::with_capacity(512 * 1024); + payload.resize(256 * 1024, b'x'); + let pointer = payload.as_ptr(); + let capacity = payload.capacity(); + let mut id = Vec::with_capacity(64); + id.extend_from_slice(b"1710000000000-0"); + let id_pointer = id.as_ptr(); + let mut next_id = Vec::with_capacity(64); + next_id.extend_from_slice(b"1710000000001-0"); + let next_pointer = next_id.as_ptr(); + let mut deleted_id = Vec::with_capacity(64); + deleted_id.extend_from_slice(b"1709999999999-0"); + let deleted_pointer = deleted_id.as_ptr(); + let fields = fields_reply( + vec![(bulk("payload"), RedisValue::BulkString(payload))], + resp3, + ); + let parsed = parse_reclaim_result(RedisValue::Array(vec![ + RedisValue::BulkString(next_id), + RedisValue::Array(vec![RedisValue::Array(vec![ + RedisValue::BulkString(id), + fields, + ])]), + RedisValue::Array(vec![RedisValue::BulkString(deleted_id)]), + ])) + .expect("reclaim reply"); + let body = &parsed.entries[0].fields["payload"]; + assert_eq!(body.len(), 256 * 1024); + assert!(body.bytes().all(|byte| byte == b'x')); + assert_eq!(body.as_ptr(), pointer); + assert_eq!(body.capacity(), capacity); + assert_eq!(parsed.entries[0].id.as_ptr(), id_pointer); + assert_eq!(parsed.next_start_id.as_ptr(), next_pointer); + assert_eq!(parsed.deleted_ids[0].as_ptr(), deleted_pointer); + } +} + +#[test] +fn owned_read_matches_existing_redis_decoder_for_supported_reply_shapes() { + let mut replies = vec![ + RedisValue::Nil, + RedisValue::Array(vec![]), + RedisValue::Map(vec![]), + ]; + for resp3 in [false, true] { + for fields in [ + RedisValue::Nil, + fields_reply( + vec![( + bulk("payload"), + bulk("{ \"text\": \"caf\u{00e9}\", \"n\": 0 }\n"), + )], + resp3, + ), + fields_reply( + vec![ + (bulk("payload"), RedisValue::Nil), + (bulk("count"), RedisValue::Int(0)), + ], + resp3, + ), + fields_reply( + vec![(bulk("payload"), RedisValue::BulkString(vec![0xff]))], + resp3, + ), + fields_reply( + vec![ + (bulk("payload"), bulk("first")), + (bulk("payload"), bulk("last")), + (bulk("invalid"), RedisValue::Boolean(false)), + ], + resp3, + ), + fields_reply( + vec![( + bulk("payload"), + RedisValue::Attribute { + data: Box::new(bulk("annotated payload")), + attributes: vec![(bulk("encoding"), bulk("utf8"))], + }, + )], + resp3, + ), + ] { + replies.push(read_reply(bulk("1-0"), fields, resp3)); + } + replies.push(read_reply(RedisValue::Int(42), RedisValue::Nil, resp3)); + } + for reply in replies { + let expected = original_read_parser(&reply).expect("baseline reply"); + assert_eq!( + parse_stream_read_entries(reply).expect("owned reply"), + expected + ); + } +} + +#[test] +fn owned_read_preserves_duplicate_overwrite_before_value_filtering() { + for resp3 in [false, true] { + let reply = read_reply( + bulk("1-0"), + fields_reply( + vec![ + (bulk("payload"), bulk("valid earlier payload")), + (bulk("payload"), RedisValue::BulkString(vec![0xff])), + (bulk("retry"), bulk("first")), + (bulk("retry"), bulk("last")), + ], + resp3, + ), + resp3, + ); + let parsed = parse_stream_read_entries(reply) + .expect("invalid values are filtered after deduplication"); + assert_eq!( + parsed[0].fields, + BTreeMap::from([("retry".to_string(), "last".to_string())]) + ); + } +} + +#[test] +fn owned_read_keeps_invalid_id_key_and_shape_errors_in_redis_category() { + for reply in [ + RedisValue::Int(7), + RedisValue::Array(vec![RedisValue::Array(vec![bulk("stream")])]), + read_reply(RedisValue::BulkString(vec![0xff]), RedisValue::Nil, false), + read_reply(bulk("1-0"), RedisValue::Array(vec![bulk("orphan")]), false), + read_reply( + bulk("1-0"), + fields_reply( + vec![(RedisValue::BulkString(vec![0xff]), bulk("value"))], + true, + ), + true, + ), + ] { + assert!(matches!( + original_read_parser(&reply), + Err(DataLayerError::Redis(_)) + )); + assert!(matches!( + parse_stream_read_entries(reply), + Err(DataLayerError::Redis(_)) + )); + } +} + +#[test] +fn owned_reclaim_preserves_string_types_nil_and_duplicate_fields() { + for resp3 in [false, true] { + let fields = fields_reply( + vec![ + (bulk("payload"), bulk("first")), + (bulk("payload"), bulk("{\"text\":\"caf\u{00e9}\"}\n")), + (bulk("zero"), RedisValue::Int(0)), + (bulk("double"), RedisValue::Double(1.5)), + ( + bulk("simple"), + RedisValue::SimpleString("simple".to_string()), + ), + (bulk("okay"), RedisValue::Okay), + ( + bulk("verbatim"), + RedisValue::VerbatimString { + format: VerbatimFormat::Text, + text: "verbatim".to_string(), + }, + ), + ( + bulk("attribute"), + RedisValue::Attribute { + data: Box::new(bulk("annotated")), + attributes: vec![], + }, + ), + ], + resp3, + ); + let parsed = + parse_reclaim_result(reclaim_reply(bulk("1-0"), fields)).expect("reclaim reply"); + assert_eq!( + parsed.entries[0].fields, + BTreeMap::from([ + ( + "payload".to_string(), + "{\"text\":\"caf\u{00e9}\"}\n".to_string() + ), + ("zero".to_string(), "0".to_string()), + ("double".to_string(), "1.5".to_string()), + ("simple".to_string(), "simple".to_string()), + ("okay".to_string(), "OK".to_string()), + ("verbatim".to_string(), "verbatim".to_string()), + ("attribute".to_string(), "annotated".to_string()), + ]) + ); + assert!(parsed.deleted_ids.is_empty()); + } + let parsed = + parse_reclaim_result(reclaim_reply(bulk("1-0"), RedisValue::Nil)).expect("nil fields"); + assert!(parsed.entries[0].fields.is_empty()); + let parsed = parse_reclaim_result(RedisValue::Array(vec![bulk("0-0"), RedisValue::Nil])) + .expect("nil entries"); + assert!(parsed.entries.is_empty()); + assert!(parsed.deleted_ids.is_empty()); +} + +#[test] +fn owned_reclaim_preserves_strict_validation_and_error_context() { + for (reply, context) in [ + ( + RedisValue::Nil, + "redis xautoclaim returned non-array payload", + ), + ( + RedisValue::Array(vec![bulk("0-0")]), + "redis xautoclaim returned 1 top-level fields", + ), + ( + RedisValue::Array(vec![RedisValue::Nil, RedisValue::Nil]), + "redis xautoclaim next_start_id", + ), + ( + RedisValue::Array(vec![bulk("0-0"), RedisValue::Int(1)]), + "redis xautoclaim entries payload was not an array", + ), + ( + RedisValue::Array(vec![bulk("0-0"), RedisValue::Array(vec![RedisValue::Nil])]), + "redis xautoclaim entry was not an array", + ), + ( + RedisValue::Array(vec![ + bulk("0-0"), + RedisValue::Array(vec![RedisValue::Array(vec![])]), + ]), + "redis xautoclaim entry had 0 fields", + ), + ( + reclaim_reply(RedisValue::BulkString(vec![0xff]), RedisValue::Nil), + "redis xautoclaim entry id", + ), + ( + reclaim_reply(bulk("1-0"), RedisValue::Array(vec![bulk("orphan")])), + "redis xautoclaim entry fields expected an even number", + ), + ( + reclaim_reply(bulk("1-0"), RedisValue::Int(1)), + "redis xautoclaim entry fields expected a redis array/map payload", + ), + ( + reclaim_reply( + bulk("1-0"), + fields_reply( + vec![(bulk("payload"), RedisValue::BulkString(vec![0xff]))], + false, + ), + ), + "redis xautoclaim entry fields was not a string-compatible", + ), + ( + reclaim_reply( + bulk("1-0"), + fields_reply( + vec![(RedisValue::BulkString(vec![0xff]), bulk("value"))], + true, + ), + ), + "redis xautoclaim entry fields was not a string-compatible", + ), + ( + RedisValue::Array(vec![bulk("0-0"), RedisValue::Nil, RedisValue::Int(1)]), + "redis xautoclaim deleted_ids expected a redis array payload", + ), + ( + RedisValue::Array(vec![ + bulk("0-0"), + RedisValue::Nil, + RedisValue::Array(vec![RedisValue::Nil]), + ]), + "redis xautoclaim deleted_ids was not a string-compatible", + ), + ] { + let error = parse_reclaim_result(reply).expect_err("invalid reclaim reply"); + let DataLayerError::UnexpectedValue(message) = error else { + panic!("reclaim parse error must keep its classification: {error}"); + }; + assert!( + message.starts_with(context), + "expected {context}, got {message}" + ); + } + // Unlike read-group's filter, reclaim has always rejected an invalid value even if + // a later duplicate would overwrite it. Keep that validation order. + let reply = reclaim_reply( + bulk("1-0"), + fields_reply( + vec![ + (bulk("payload"), RedisValue::Nil), + (bulk("payload"), bulk("later valid payload")), + ], + false, + ), + ); + assert!(matches!( + parse_reclaim_result(reply), + Err(DataLayerError::UnexpectedValue(_)) + )); +} diff --git a/crates/aether-runtime/state/src/redis/stream_receive_tests.rs b/crates/aether-runtime/state/src/redis/stream_receive_tests.rs new file mode 100644 index 000000000..0d162dd6f --- /dev/null +++ b/crates/aether-runtime/state/src/redis/stream_receive_tests.rs @@ -0,0 +1,860 @@ +use super::*; + +use std::collections::BTreeSet; +use std::future::{poll_fn, Future}; +use std::pin::Pin; +use std::task::Poll; + +type TestConnection = ::redis::aio::MultiplexedConnection; +type ReadResult = Result, DataLayerError>; + +const TEST_GROUP: &str = "receive-workers"; +const OWNER_BLOCK_MS: u64 = 60_000; + +fn receive_lane_count() -> usize { + std::thread::available_parallelism() + .map(|value| value.get()) + .unwrap_or(4) + .clamp(4, 16) +} + +async fn receive_runtime( + protocol: &str, + command_timeout_ms: u64, +) -> Option<(TestRedisServer, RuntimeState, TestConnection)> { + let Some(server) = TestRedisServer::start().await else { + eprintln!( + "stream receive {protocol} skipped: isolated Redis fixture unavailable; check AETHER_REDIS_SERVER_BIN" + ); + return None; + }; + let mut admin = ::redis::Client::open(server.redis_url.clone()) + .expect("test admin client") + .get_multiplexed_async_connection() + .await + .expect("test admin connection"); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg("stream-reader") + .arg("on") + .arg(">stream-test-password") + .arg("~*") + .arg("+@all") + .query_async::<()>(&mut admin) + .await + .expect("test stream user"); + let runtime = RuntimeState::redis_with_blocking_stream_lanes( + RedisClientConfig { + url: format!( + "redis://stream-reader:stream-test-password@127.0.0.1:{}/7?protocol={protocol}", + server.port + ), + key_prefix: Some(format!("receive-{protocol}")), + }, + Some(command_timeout_ms), + Some(4), + ) + .await + .expect("authenticated receive runtime in database 7"); + ::redis::cmd("SELECT") + .arg(7) + .query_async::<()>(&mut admin) + .await + .expect("admin selects test database"); + eprintln!( + "stream receive fixture ready: protocol={protocol} db=7 port={} authenticated=true", + server.port + ); + Some((server, runtime, admin)) +} + +async fn receive_group(runtime: &RuntimeState, stream: &str) { + RuntimeQueueStore::ensure_consumer_group(runtime, stream, TEST_GROUP, "0-0") + .await + .expect("receive consumer group"); +} + +fn receive_fields(sequence: usize) -> BTreeMap { + BTreeMap::from([ + ( + "payload".to_string(), + format!("record-{sequence}\r\n\"quoted\"\\\u{4e2d}\u{6587}"), + ), + ("sequence".to_string(), sequence.to_string()), + ("legacy_field".to_string(), "preserve exactly".to_string()), + ]) +} + +async fn append_receive(runtime: &RuntimeState, stream: &str, sequence: usize) -> String { + RuntimeQueueStore::append_fields_with_maxlen(runtime, stream, &receive_fields(sequence), None) + .await + .expect("append receive entry") +} + +async fn client_rows(admin: &mut TestConnection) -> Vec> { + let value = ::redis::cmd("CLIENT") + .arg("LIST") + .query_async::(admin) + .await + .expect("Redis client list"); + value + .lines() + .map(|line| { + line.split_whitespace() + .filter_map(|field| field.split_once('=')) + .map(|(key, value)| (key.to_string(), value.to_string())) + .collect() + }) + .collect() +} + +fn blocked_rows(rows: &[BTreeMap]) -> Vec<&BTreeMap> { + rows.iter() + .filter(|row| row.get("flags").is_some_and(|flags| flags.contains('b'))) + .collect() +} + +async fn wait_for_blocked(admin: &mut TestConnection, expected: usize) { + tokio::time::timeout(Duration::from_secs(10), async { + loop { + if blocked_rows(&client_rows(admin).await).len() == expected { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("Redis blocked clients reach expected count"); +} + +fn spawn_owner( + owners: &mut tokio::task::JoinSet, + runtime: &RuntimeState, + stream: &'static str, + index: usize, +) { + let runtime = runtime.clone(); + owners.spawn(async move { + RuntimeQueueStore::read_group( + &runtime, + stream, + TEST_GROUP, + &format!("owner-{index}"), + 1, + Some(OWNER_BLOCK_MS), + ) + .await + }); +} + +async fn abort_owners(owners: &mut tokio::task::JoinSet) { + owners.abort_all(); + while let Some(result) = owners.join_next().await { + assert!(result.expect_err("owner must be cancelled").is_cancelled()); + } +} + +async fn assert_receive_pending(mut future: Pin<&mut F>) { + poll_fn(|context| { + assert!(future.as_mut().poll(context).is_pending()); + Poll::Ready(()) + }) + .await; +} + +async fn pending_consumers( + admin: &mut TestConnection, + stream: &str, +) -> Vec<(String, String, u64, u64)> { + ::redis::cmd("XPENDING") + .arg(stream) + .arg(TEST_GROUP) + .arg("-") + .arg("+") + .arg(100) + .query_async(admin) + .await + .expect("pending entry ownership") +} + +#[tokio::test] +async fn redis_stream_receive_full_pool_waits_without_sending_and_cancellation_preserves_pel() { + for protocol in ["resp2", "resp3"] { + let Some((_server, runtime, mut admin)) = receive_runtime(protocol, 5_000).await else { + return; + }; + let stream = "receive:blocked"; + let fast_stream = "receive:nonblocking"; + receive_group(&runtime, stream).await; + receive_group(&runtime, fast_stream).await; + let lanes = receive_lane_count(); + let mut owners = tokio::task::JoinSet::new(); + for index in 0..lanes { + spawn_owner(&mut owners, &runtime, stream, index); + } + wait_for_blocked(&mut admin, lanes).await; + let before = client_rows(&mut admin).await; + let owner_input_bytes = blocked_rows(&before) + .into_iter() + .map(|row| (row["id"].clone(), row.get("tot-net-in").cloned())) + .collect::>(); + + let mut waiter = Box::pin(RuntimeQueueStore::read_group( + &runtime, + stream, + TEST_GROUP, + "cancelled-waiter", + 1, + Some(OWNER_BLOCK_MS), + )); + assert_receive_pending(waiter.as_mut()).await; + + // Complete unrelated round trips while the waiter remains polled and the owners block. + tokio::time::timeout(Duration::from_secs(5), async { + runtime.kv_set("receive-fast", "ready", None).await.unwrap(); + assert_eq!( + runtime.kv_get("receive-fast").await.unwrap().as_deref(), + Some("ready") + ); + let expected_id = append_receive(&runtime, fast_stream, 7).await; + let entries = RuntimeQueueStore::read_group( + &runtime, + fast_stream, + TEST_GROUP, + "nonblocking-reader", + 1, + None, + ) + .await + .unwrap(); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].id, expected_id); + assert_eq!(entries[0].fields, receive_fields(7)); + }) + .await + .expect("full blocking pool must not delay fast or nonblocking stream lanes"); + + for _ in 0..3 { + tokio::task::yield_now().await; + let rows = client_rows(&mut admin).await; + let blocked = blocked_rows(&rows); + assert_eq!(blocked.len(), lanes); + for row in blocked { + assert_eq!( + row.get("tot-net-in"), + owner_input_bytes[&row["id"]].as_ref() + ); + assert_eq!(row["qbuf"], "0", "waiter must not be sent behind a BLOCK"); + } + assert_receive_pending(waiter.as_mut()).await; + } + drop(waiter); + assert_eq!(blocked_rows(&client_rows(&mut admin).await).len(), lanes); + assert!(pending_consumers(&mut admin, stream).await.is_empty()); + + abort_owners(&mut owners).await; + wait_for_blocked(&mut admin, 0).await; + let expected_id = append_receive(&runtime, stream, 8).await; + let entries = RuntimeQueueStore::read_group( + &runtime, + stream, + TEST_GROUP, + "replacement-reader", + 1, + Some(100), + ) + .await + .expect("replacement blocking connection"); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].id, expected_id); + assert_eq!(entries[0].fields, receive_fields(8)); + let pending = pending_consumers(&mut admin, stream).await; + assert_eq!(pending.len(), 1); + assert_eq!(pending[0].0, expected_id); + assert_eq!(pending[0].1, "replacement-reader"); + ::redis::cmd("SELECT") + .arg(0) + .query_async::<()>(&mut admin) + .await + .unwrap(); + let other_database_len = ::redis::cmd("XLEN") + .arg(stream) + .query_async::(&mut admin) + .await + .unwrap(); + assert_eq!( + other_database_len, 0, + "AUTH and SELECT must survive replacement connections" + ); + } +} + +#[tokio::test] +async fn redis_stream_receive_fast_consumer_reuses_free_lane_while_other_lanes_block() { + let Some((_server, runtime, mut admin)) = receive_runtime("resp2", 2_000).await else { + return; + }; + let slow_stream = "receive:slow"; + let fast_stream = "receive:ready"; + receive_group(&runtime, slow_stream).await; + receive_group(&runtime, fast_stream).await; + let lanes = receive_lane_count(); + let mut owners = tokio::task::JoinSet::new(); + for index in 0..lanes - 1 { + spawn_owner(&mut owners, &runtime, slow_stream, index); + } + wait_for_blocked(&mut admin, lanes - 1).await; + + let mut reader_connection_id = None; + for sequence in 0..lanes * 2 { + let expected_id = append_receive(&runtime, fast_stream, sequence).await; + let entries = RuntimeQueueStore::read_group( + &runtime, + fast_stream, + TEST_GROUP, + "fast-reader", + 1, + Some(100), + ) + .await + .expect("a free lane must remain reusable instead of rotating into a blocked lane"); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].id, expected_id); + assert_eq!(entries[0].fields, receive_fields(sequence)); + let rows = client_rows(&mut admin).await; + let idle_readers = rows + .iter() + .filter(|row| { + row.get("cmd").map(String::as_str) == Some("xreadgroup") + && row.get("flags").is_some_and(|flags| !flags.contains('b')) + }) + .collect::>(); + assert_eq!(idle_readers.len(), 1); + let current_id = &idle_readers[0]["id"]; + if let Some(previous_id) = reader_connection_id.as_ref() { + assert_eq!( + current_id, previous_id, + "successful reads must reuse the same free connection" + ); + } else { + reader_connection_id = Some(current_id.clone()); + } + } + assert_eq!( + blocked_rows(&client_rows(&mut admin).await).len(), + lanes - 1 + ); + abort_owners(&mut owners).await; + wait_for_blocked(&mut admin, 0).await; +} + +#[tokio::test] +async fn redis_stream_receive_timeout_discards_inflight_connection_before_reuse() { + let Some((_server, runtime, mut admin)) = receive_runtime("resp3", 1_000).await else { + return; + }; + let stream = "receive:timeout"; + receive_group(&runtime, stream).await; + let initial_connection_ids = client_rows(&mut admin) + .await + .into_iter() + .map(|row| row["id"].clone()) + .collect::>(); + // XREADGROUP is a write command; pausing writes makes its network response exceed the + // normal BLOCK-plus-grace timeout while read-only CLIENT diagnostics remain available. + ::redis::cmd("CLIENT") + .arg("PAUSE") + .arg(30_000) + .arg("WRITE") + .query_async::<()>(&mut admin) + .await + .expect("pause test Redis writes"); + let result = tokio::time::timeout( + Duration::from_secs(10), + RuntimeQueueStore::read_group( + &runtime, + stream, + TEST_GROUP, + "timed-out-reader", + 1, + Some(100), + ), + ) + .await + .expect("read reaches its configured command deadline"); + assert!(matches!(result, Err(DataLayerError::TimedOut(_)))); + tokio::time::timeout(Duration::from_secs(10), async { + loop { + let connection_ids = client_rows(&mut admin) + .await + .into_iter() + .map(|row| row["id"].clone()) + .collect::>(); + if connection_ids.len() + 1 == initial_connection_ids.len() + && connection_ids.is_subset(&initial_connection_ids) + { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("timed-out connection must disconnect while its old command is still paused"); + ::redis::cmd("CLIENT") + .arg("UNPAUSE") + .query_async::<()>(&mut admin) + .await + .unwrap(); + let expected_id = append_receive(&runtime, stream, 9).await; + let entries = + RuntimeQueueStore::read_group(&runtime, stream, TEST_GROUP, "after-timeout", 1, Some(100)) + .await + .expect("replacement read after timeout"); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].id, expected_id); + assert_eq!(entries[0].fields, receive_fields(9)); + let pending = pending_consumers(&mut admin, stream).await; + assert_eq!(pending.len(), 1); + assert_eq!(pending[0].0, expected_id); + assert_eq!(pending[0].1, "after-timeout"); +} + +async fn set_pending_idle(admin: &mut TestConnection, stream: &str, ids: &[String]) { + let mut command = ::redis::cmd("XCLAIM"); + command + .arg(stream) + .arg(TEST_GROUP) + .arg("initial-reader") + .arg(0); + for id in ids { + command.arg(id); + } + command.arg("IDLE").arg(120_000).arg("JUSTID"); + let changed = command + .query_async::>(admin) + .await + .expect("set pending idle"); + assert_eq!(changed, ids); +} + +#[tokio::test] +async fn redis_stream_receive_reclaim_pages_advance_past_fresh_prefix_and_deleted_entries() { + for protocol in ["resp2", "resp3"] { + let Some((_server, runtime, mut admin)) = receive_runtime(protocol, 5_000).await else { + return; + }; + let stream = "receive:reclaim-pages"; + receive_group(&runtime, stream).await; + let mut ids = Vec::new(); + for sequence in 0..28 { + ids.push(append_receive(&runtime, stream, sequence).await); + } + let entries = + RuntimeQueueStore::read_group(&runtime, stream, TEST_GROUP, "initial-reader", 28, None) + .await + .expect("seed pending entries"); + assert_eq!(entries.len(), ids.len()); + set_pending_idle(&mut admin, stream, &ids[25..]).await; + assert_eq!( + RuntimeQueueStore::delete(&runtime, stream, &ids[26..27]) + .await + .unwrap(), + 1 + ); + let config = RuntimeQueueReclaimConfig { + min_idle_ms: 60_000, + count: 2, + }; + let first = RuntimeQueueStore::claim_stale_page( + &runtime, + stream, + TEST_GROUP, + "reclaim-reader", + "0-0", + config, + ) + .await + .expect("first reclaim page"); + assert!( + first.entries.is_empty(), + "fresh prefix exceeds COUNT * 10 scan budget" + ); + assert!(first.deleted_ids.is_empty()); + assert_ne!( + first.next_start_id, "0-0", + "empty page must preserve continuation" + ); + let mut cursor = first.next_start_id; + let mut reclaimed = BTreeMap::new(); + let mut deleted = BTreeSet::new(); + for _ in 0..8 { + let page = RuntimeQueueStore::claim_stale_page( + &runtime, + stream, + TEST_GROUP, + "reclaim-reader", + &cursor, + config, + ) + .await + .expect("continued reclaim page"); + assert!(page.entries.len() <= config.count); + for entry in page.entries { + assert!(reclaimed.insert(entry.id, entry.fields).is_none()); + } + deleted.extend(page.deleted_ids); + cursor = page.next_start_id; + if cursor == "0-0" { + break; + } + } + assert_eq!(cursor, "0-0", "scan must eventually wrap"); + assert_eq!( + reclaimed, + BTreeMap::from([ + (ids[25].clone(), receive_fields(25)), + (ids[27].clone(), receive_fields(27)) + ]) + ); + assert_eq!(deleted, BTreeSet::from([ids[26].clone()])); + let pending = pending_consumers(&mut admin, stream).await; + assert_eq!(pending.len(), 27); + assert!(!pending.iter().any(|entry| entry.0 == ids[26])); + assert!(pending + .iter() + .filter(|entry| entry.1 == "reclaim-reader") + .all(|entry| entry.0 == ids[25] || entry.0 == ids[27])); + + set_pending_idle(&mut admin, stream, &ids[..1]).await; + let restarted = RuntimeQueueStore::claim_stale_page( + &runtime, + stream, + TEST_GROUP, + "rescan-reader", + "0-0", + config, + ) + .await + .expect("restart scan after cursor wraps"); + assert_eq!(restarted.entries.len(), 1); + assert_eq!(restarted.entries[0].id, ids[0]); + assert_eq!(restarted.entries[0].fields, receive_fields(0)); + assert!(restarted.deleted_ids.is_empty()); + } +} + +async fn interrupted_reclaim_preserves_pending(protocol: &str, cancel: bool) { + let command_timeout_ms = if cancel { 10_000 } else { 1_000 }; + let Some((_server, runtime, mut admin)) = receive_runtime(protocol, command_timeout_ms).await + else { + return; + }; + let stream = "receive:interrupted-reclaim"; + receive_group(&runtime, stream).await; + let mut ids = Vec::new(); + for sequence in 0..3 { + ids.push(append_receive(&runtime, stream, sequence).await); + } + assert_eq!( + RuntimeQueueStore::read_group(&runtime, stream, TEST_GROUP, "initial-reader", 3, None,) + .await + .unwrap() + .len(), + 3 + ); + set_pending_idle(&mut admin, stream, &ids).await; + RuntimeQueueStore::delete(&runtime, stream, &ids[1..2]) + .await + .unwrap(); + let before = client_rows(&mut admin).await; + let initial_ids = before + .iter() + .map(|row| row["id"].clone()) + .collect::>(); + let input_bytes = before + .iter() + .filter(|row| row.get("user").map(String::as_str) == Some("stream-reader")) + .map(|row| { + ( + row["id"].clone(), + row.get("tot-net-in") + .and_then(|value| value.parse::().ok()), + ) + }) + .collect::>(); + ::redis::cmd("CLIENT") + .arg("PAUSE") + .arg(30_000) + .arg("WRITE") + .query_async::<()>(&mut admin) + .await + .unwrap(); + let config = RuntimeQueueReclaimConfig { + min_idle_ms: 60_000, + count: 1, + }; + let claim_runtime = runtime.clone(); + let claim = tokio::spawn(async move { + RuntimeQueueStore::claim_stale_page( + &claim_runtime, + stream, + TEST_GROUP, + "interrupted-reader", + "0-0", + config, + ) + .await + }); + if cancel { + tokio::time::timeout(Duration::from_secs(5), async { + loop { + let sent = client_rows(&mut admin).await.iter().any(|row| { + let Some(previous) = input_bytes.get(&row["id"]) else { + return false; + }; + match ( + previous, + row.get("tot-net-in") + .and_then(|value| value.parse::().ok()), + ) { + (Some(previous), Some(current)) => current > *previous, + _ => row.get("cmd").map(String::as_str) == Some("xautoclaim"), + } + }); + if sent { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("claim reaches Redis before caller cancellation"); + claim.abort(); + assert!(claim.await.expect_err("cancelled reclaim").is_cancelled()); + } else { + let result = tokio::time::timeout(Duration::from_secs(10), claim) + .await + .expect("reclaim command deadline") + .expect("reclaim task"); + assert!(matches!(result, Err(DataLayerError::TimedOut(_)))); + } + tokio::time::timeout(Duration::from_secs(10), async { + loop { + let remaining = client_rows(&mut admin) + .await + .into_iter() + .map(|row| row["id"].clone()) + .collect::>(); + if remaining.len() + 1 == initial_ids.len() && remaining.is_subset(&initial_ids) { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("interrupted claim closes its socket before writes resume"); + let pending = pending_consumers(&mut admin, stream).await; + assert_eq!(pending.len(), 3); + assert!(pending.iter().all(|entry| entry.1 == "initial-reader")); + ::redis::cmd("CLIENT") + .arg("UNPAUSE") + .query_async::<()>(&mut admin) + .await + .unwrap(); + + let mut cursor = "0-0".to_string(); + let mut recovered = BTreeMap::new(); + let mut deleted = BTreeSet::new(); + let mut successful_claims = 0; + for _ in 0..5 { + let page = RuntimeQueueStore::claim_stale_page( + &runtime, + stream, + TEST_GROUP, + "recovery-reader", + &cursor, + config, + ) + .await + .expect("reclaim recovers after interrupted connection"); + successful_claims += 1; + for entry in page.entries { + assert!(recovered.insert(entry.id, entry.fields).is_none()); + } + deleted.extend(page.deleted_ids); + cursor = page.next_start_id; + if cursor == "0-0" { + break; + } + } + assert_eq!(cursor, "0-0"); + assert_eq!( + recovered, + BTreeMap::from([ + (ids[0].clone(), receive_fields(0)), + (ids[2].clone(), receive_fields(2)) + ]) + ); + assert_eq!(deleted, BTreeSet::from([ids[1].clone()])); + let pending = pending_consumers(&mut admin, stream).await; + assert_eq!(pending.len(), 2); + assert!(pending.iter().all(|entry| entry.1 == "recovery-reader")); + let diagnostics = runtime.redis_diagnostics().await.unwrap().unwrap(); + let lane = diagnostics + .lanes + .iter() + .find(|lane| lane.lane == "blocking_stream") + .unwrap(); + assert_eq!(lane.command_timeouts, u64::from(!cancel)); + assert!(lane.command_count >= successful_claims); +} + +#[tokio::test] +async fn redis_stream_receive_reclaim_cancellation_closes_connection_and_preserves_pending() { + interrupted_reclaim_preserves_pending("resp2", true).await; +} + +#[tokio::test] +async fn redis_stream_receive_reclaim_timeout_closes_connection_and_preserves_pending() { + interrupted_reclaim_preserves_pending("resp3", false).await; +} + +#[tokio::test] +async fn redis_stream_receive_reclaim_waits_for_read_lease_and_continues_after_release() { + let Some((_server, runtime, mut admin)) = receive_runtime("resp2", 1_000).await else { + return; + }; + let blocked_stream = "receive:reclaim-pool-blocked"; + let pending_stream = "receive:reclaim-pool-pending"; + receive_group(&runtime, blocked_stream).await; + receive_group(&runtime, pending_stream).await; + let ids = vec![ + append_receive(&runtime, pending_stream, 0).await, + append_receive(&runtime, pending_stream, 1).await, + ]; + assert_eq!( + RuntimeQueueStore::read_group( + &runtime, + pending_stream, + TEST_GROUP, + "initial-reader", + 2, + None + ) + .await + .unwrap() + .len(), + 2 + ); + set_pending_idle(&mut admin, pending_stream, &ids).await; + let lanes = receive_lane_count(); + let mut owners = tokio::task::JoinSet::new(); + for index in 0..lanes - 1 { + spawn_owner(&mut owners, &runtime, blocked_stream, index); + } + let owner_runtime = runtime.clone(); + let release_owner = owners.spawn(async move { + RuntimeQueueStore::read_group( + &owner_runtime, + blocked_stream, + TEST_GROUP, + "released-owner", + 1, + Some(OWNER_BLOCK_MS), + ) + .await + }); + wait_for_blocked(&mut admin, lanes).await; + let config = RuntimeQueueReclaimConfig { + min_idle_ms: 60_000, + count: 1, + }; + let result = tokio::time::timeout( + Duration::from_secs(10), + RuntimeQueueStore::claim_stale_page( + &runtime, + pending_stream, + TEST_GROUP, + "timed-out-waiter", + "0-0", + config, + ), + ) + .await + .expect("checkout uses the original command deadline"); + assert!(matches!(result, Err(DataLayerError::TimedOut(_)))); + assert_eq!(blocked_rows(&client_rows(&mut admin).await).len(), lanes); + assert!(pending_consumers(&mut admin, pending_stream) + .await + .iter() + .all(|entry| entry.1 == "initial-reader")); + let mut cancelled_claim = Box::pin(RuntimeQueueStore::claim_stale_page( + &runtime, + pending_stream, + TEST_GROUP, + "cancelled-waiter", + "0-0", + config, + )); + assert_receive_pending(cancelled_claim.as_mut()).await; + drop(cancelled_claim); + assert_eq!(blocked_rows(&client_rows(&mut admin).await).len(), lanes); + + let mut claim = Box::pin(RuntimeQueueStore::claim_stale_page( + &runtime, + pending_stream, + TEST_GROUP, + "after-release", + "0-0", + config, + )); + assert_receive_pending(claim.as_mut()).await; + release_owner.abort(); + assert!(owners + .join_next() + .await + .unwrap() + .expect_err("released owner cancelled") + .is_cancelled()); + let first = tokio::time::timeout(Duration::from_secs(5), claim) + .await + .expect("claim receives released pool capacity") + .expect("claim succeeds after read releases lease"); + assert_eq!(first.entries.len(), 1); + assert_eq!(first.entries[0].id, ids[0]); + assert_eq!(first.entries[0].fields, receive_fields(0)); + assert_ne!(first.next_start_id, "0-0"); + + let next_id = append_receive(&runtime, pending_stream, 2).await; + let next_read = RuntimeQueueStore::read_group( + &runtime, + pending_stream, + TEST_GROUP, + "read-after-claim", + 1, + Some(100), + ) + .await + .expect("completed claim returns its lease for the next read"); + assert_eq!(next_read.len(), 1); + assert_eq!(next_read[0].id, next_id); + let second = RuntimeQueueStore::claim_stale_page( + &runtime, + pending_stream, + TEST_GROUP, + "after-release", + &first.next_start_id, + config, + ) + .await + .expect("completed read returns its lease for the next claim"); + assert_eq!(second.entries.len(), 1); + assert_eq!(second.entries[0].id, ids[1]); + assert_eq!(second.entries[0].fields, receive_fields(1)); + assert_eq!( + blocked_rows(&client_rows(&mut admin).await).len(), + lanes - 1 + ); + abort_owners(&mut owners).await; + wait_for_blocked(&mut admin, 0).await; +} diff --git a/crates/aether-runtime/state/src/redis/usage_cleanup.rs b/crates/aether-runtime/state/src/redis/usage_cleanup.rs new file mode 100644 index 000000000..2dd8b68d7 --- /dev/null +++ b/crates/aether-runtime/state/src/redis/usage_cleanup.rs @@ -0,0 +1,175 @@ +use std::sync::OnceLock; + +use super::client::RedisBlockingStreamLease; +use super::runtime::USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT; +use super::{cmd, RedisConnectionLane, RedisConnectionRouter}; +use crate::error::RedisResultExt; +use crate::{DataLayerError, UsageLimitInput}; + +const COPY_CHUNK_SIZE: i64 = 512; +const MAX_COPY_ATTEMPTS: usize = 8; +const COPY_SCRIPT: &str = include_str!("usage_copy.lua"); +const COMMIT_PREFIX: &str = include_str!("usage_copy_commit.lua"); + +fn usage_args(command: &mut redis::Cmd, input: &UsageLimitInput<'_>) { + command.arg(input.now_unix_ms).arg(input.event_id); + for rule in input.rules { + command + .arg(rule.limit) + .arg(rule.window_seconds) + .arg(rule.retention_seconds); + } +} + +pub(super) async fn check_and_consume( + connections: &RedisConnectionRouter, + keys: &[String], + input: &UsageLimitInput<'_>, +) -> Result, DataLayerError> { + static SCRIPT: OnceLock = OnceLock::new(); + let script = SCRIPT.get_or_init(|| redis::Script::new(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT)); + let mut invocation = script.prepare_invoke(); + for key in keys { + invocation.key(key); + } + invocation.arg(input.now_unix_ms).arg(input.event_id); + for rule in input.rules { + invocation + .arg(rule.limit) + .arg(rule.window_seconds) + .arg(rule.retention_seconds); + } + let result: Vec = invocation + .invoke_async(&mut connections.connection(RedisConnectionLane::Fast)) + .await + .map_redis_err()?; + if result.first() != Some(&2) { + return Ok(result); + } + + let mut lease = connections.usage_cleanup_connection().await?; + for _ in 0..MAX_COPY_ATTEMPTS { + let mut temporary_keys = Vec::new(); + let result = copy_and_commit(&mut lease, keys, input, &mut temporary_keys).await; + // EXEC clears WATCH even on a conflict. Errors discard the lease, + // including any pending WATCH/MULTI state. + if let Ok(Some((result, committed))) = result.as_ref() { + // Never add another fallible round trip after admission. EXEC already + // cleared WATCH; the recheck fast path instead discards its watched lease. + if *committed { + lease.recycle(); + } + return Ok(result.clone()); + } else if result.is_ok() { + lease.query(&cmd("UNWATCH")).await?; + if !temporary_keys.is_empty() { + lease.query(cmd("UNLINK").arg(&temporary_keys)).await?; + } + } else if !temporary_keys.is_empty() { + // A fresh connection cannot accidentally queue cleanup inside a failed MULTI. + let mut connection = connections.connection(RedisConnectionLane::Admin); + let _ = cmd("UNLINK") + .arg(&temporary_keys) + .query_async::(&mut connection) + .await; + } + result?; + tokio::task::yield_now().await; + } + lease.recycle(); + Err(DataLayerError::Redis( + "usage window cleanup conflicted repeatedly; retry the request".to_string(), + )) +} + +async fn copy_and_commit( + lease: &mut RedisBlockingStreamLease, + keys: &[String], + input: &UsageLimitInput<'_>, + temporary_keys: &mut Vec, +) -> Result, bool)>, DataLayerError> { + lease.query(cmd("WATCH").arg(keys)).await?; + let mut check = cmd("EVAL"); + check + .arg(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT) + .arg(keys.len()) + .arg(keys); + usage_args(&mut check, input); + let plan: Vec = + redis::from_owned_redis_value(lease.query(&check).await?).map_redis_err()?; + if plan.first() != Some(&2) { + return Ok(Some((plan, false))); + } + if plan.len() < 4 || (plan.len() - 1) % 3 != 0 { + return Err(DataLayerError::UnexpectedValue( + "invalid usage window copy plan".to_string(), + )); + } + + let nonce = uuid::Uuid::new_v4(); + for window in plan[1..].chunks_exact(3) { + let [index, expired, live] = [window[0], window[1], window[2]]; + let key = usize::try_from(index - 1) + .ok() + .and_then(|index| keys.get(index)) + .filter(|_| expired > 0 && live > 0) + .ok_or_else(|| { + DataLayerError::UnexpectedValue("invalid usage window copy range".to_string()) + })?; + let temporary = format!("{key}:__usage_copy:{nonce}"); + let mut offset = 0; + while offset < live { + let take = COPY_CHUNK_SIZE.min(live - offset); + let value = lease + .query( + cmd("EVAL") + .arg(COPY_SCRIPT) + .arg(2) + .arg(key) + .arg(&temporary) + .arg(expired + offset) + .arg(expired + offset + take - 1) + .arg(offset), + ) + .await?; + let copied: i64 = redis::from_owned_redis_value(value).map_redis_err()?; + if offset == 0 && copied > 0 { + temporary_keys.push(temporary.clone()); + } + if copied != take { + return Ok(None); + } + offset += take; + } + } + + static COMMIT: OnceLock = OnceLock::new(); + let source = + COMMIT.get_or_init(|| format!("{COMMIT_PREFIX}\n{USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT}")); + let mut commit = cmd("EVAL"); + commit + .arg(source) + .arg(keys.len() + temporary_keys.len()) + .arg(keys) + .arg(&*temporary_keys) + .arg(keys.len()); + usage_args(&mut commit, input); + for window in plan[1..].chunks_exact(3) { + commit.arg(window[0]).arg(window[2]); + } + lease.query(&cmd("MULTI")).await?; + lease.query(&commit).await?; + let replies: Option> = + redis::from_owned_redis_value(lease.query(&cmd("EXEC")).await?).map_redis_err()?; + let Some(mut replies) = replies else { + return Ok(None); + }; + let reply = replies.pop().ok_or_else(|| { + DataLayerError::UnexpectedValue("empty usage window commit response".to_string()) + })?; + let result: Vec = redis::from_owned_redis_value(reply).map_redis_err()?; + if result.first() == Some(&2) { + return Ok(None); + } + Ok(Some((result, true))) +} diff --git a/crates/aether-runtime/state/src/redis/usage_copy.lua b/crates/aether-runtime/state/src/redis/usage_copy.lua new file mode 100644 index 000000000..f8fc5ede4 --- /dev/null +++ b/crates/aether-runtime/state/src/redis/usage_copy.lua @@ -0,0 +1,22 @@ +local source, target = KEYS[1], KEYS[2] +local offset = tonumber(ARGV[3]) +if offset == 0 and redis.call('EXISTS', target) ~= 0 then + return redis.error_reply('usage copy temporary key already exists') +end +if offset > 0 and redis.call('ZCARD', target) ~= offset then return -1 end +local rows = redis.call('ZRANGE', source, ARGV[1], ARGV[2], 'WITHSCORES') +if #rows == 0 then return 0 end +if not redis.acl_check_cmd('PEXPIRE', target, 60000) + or not redis.acl_check_cmd('UNLINK', target) then + return redis.error_reply('usage copy temporary key permission denied') +end +local args = {} +for i = 1, #rows, 2 do + args[#args + 1] = rows[i + 1] + args[#args + 1] = rows[i] +end +-- TTL is established in the same command as the first allocation. A cancelled +-- caller or a process crash cannot leave a permanent scratch key behind. +redis.call('ZADD', target, unpack(args)) +redis.call('PEXPIRE', target, 60000) +return #rows / 2 diff --git a/crates/aether-runtime/state/src/redis/usage_copy_commit.lua b/crates/aether-runtime/state/src/redis/usage_copy_commit.lua new file mode 100644 index 000000000..096fc68f7 --- /dev/null +++ b/crates/aether-runtime/state/src/redis/usage_copy_commit.lua @@ -0,0 +1,35 @@ +local rule_count = tonumber(table.remove(ARGV, 1)) +local swaps = #KEYS - rule_count +local ttls = {} +-- WATCH covers every source, including modifications made by older instances. +-- Check every temporary key and every permission before replacing any source. +for i = 1, swaps do + local position = rule_count * 3 + 2 + (i - 1) * 2 + local index = tonumber(ARGV[position + 1]) + local live = tonumber(ARGV[position + 2]) + local source, target = KEYS[index], KEYS[rule_count + i] + local ttl = redis.call('PTTL', source) + if ttl == 0 or ttl < -1 or redis.call('ZCARD', target) ~= live then return {2} end + if not redis.acl_check_cmd('UNLINK', source) + or not redis.acl_check_cmd('RENAME', target, source) + or not redis.acl_check_cmd('PEXPIRE', target, math.max(1, ttl)) + or not redis.acl_check_cmd('PERSIST', target) then + return redis.error_reply('usage copy commit permission denied') + end + ttls[i] = ttl +end +for i = 1, swaps do + local position = rule_count * 3 + 2 + (i - 1) * 2 + local index = tonumber(ARGV[position + 1]) + local source, target = KEYS[index], KEYS[rule_count + i] + if ttls[i] > 0 then + redis.call('PEXPIRE', target, ttls[i]) + else + redis.call('PERSIST', target) + end + redis.call('UNLINK', source) + redis.call('RENAME', target, source) +end +for i = #KEYS, rule_count + 1, -1 do KEYS[i] = nil end +ARGV[rule_count * 3 + 3] = 'inline' +-- The original prune/check/consume script follows in this same EXEC/EVAL. diff --git a/crates/aether-runtime/state/src/redis/usage_limit_cleanup_tests.rs b/crates/aether-runtime/state/src/redis/usage_limit_cleanup_tests.rs new file mode 100644 index 000000000..c769fb767 --- /dev/null +++ b/crates/aether-runtime/state/src/redis/usage_limit_cleanup_tests.rs @@ -0,0 +1,1125 @@ +use super::*; + +type TestConnection = ::redis::aio::MultiplexedConnection; + +// Frozen pre-optimization script, used as an oracle for out-of-order histories. +const LEGACY_USAGE_LIMIT_SCRIPT: &str = r#" +local count = #KEYS +local now = tonumber(ARGV[1]) +local event_id = ARGV[2] +for i = 1, count do + local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 + redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', now - window_ms) +end +for i = 1, count do + local limit = tonumber(ARGV[(i - 1) * 3 + 3]) + local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000 + if not redis.call('ZSCORE', KEYS[i], event_id) then + local current = redis.call('ZCARD', KEYS[i]) + if current >= limit then + local earliest = redis.call('ZRANGE', KEYS[i], 0, 0, 'WITHSCORES') + local retry_after = 1 + if #earliest >= 2 then + retry_after = math.ceil(math.max(1, tonumber(earliest[2]) + window_ms - now) / 1000) + end + return {0, i, limit, retry_after} + end + end +end +for i = 1, count do + local retention = tonumber(ARGV[(i - 1) * 3 + 5]) + redis.call('ZADD', KEYS[i], 'NX', now, event_id) + redis.call('EXPIRE', KEYS[i], retention + 1) +end +return {1, 0, 0, 0} +"#; + +struct Fixture { + _server: TestRedisServer, + runtime: RuntimeState, + admin: TestConnection, +} + +impl Fixture { + async fn start(protocol: &str) -> Option { + let Some(server) = TestRedisServer::start().await else { + eprintln!("usage cleanup {protocol} skipped: isolated Redis unavailable"); + return None; + }; + let mut admin = ::redis::Client::open(server.redis_url.clone()) + .unwrap() + .get_multiplexed_async_connection() + .await + .unwrap(); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg("usage-cleanup") + .arg("on") + .arg(">usage-cleanup-test-password") + .arg("~*") + .arg("+@all") + .query_async::<()>(&mut admin) + .await + .unwrap(); + let runtime = RuntimeState::redis( + RedisClientConfig { + url: format!( + "redis://usage-cleanup:usage-cleanup-test-password@127.0.0.1:{}/6?protocol={protocol}", + server.port + ), + key_prefix: Some("cleanup".to_string()), + }, + Some(5_000), + ) + .await + .unwrap(); + ::redis::cmd("SELECT") + .arg(6) + .query_async::<()>(&mut admin) + .await + .unwrap(); + eprintln!("usage cleanup fixture ready: protocol={protocol} db=6 authenticated=true"); + Some(Self { + _server: server, + runtime, + admin, + }) + } + + async fn consume( + &self, + rules: &[UsageLimitRule<'_>], + event_id: &str, + now_unix_ms: u64, + ) -> UsageLimitCheck { + self.runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules, + event_id, + now_unix_ms, + }) + .await + .unwrap() + } + + async fn seed(&mut self, prefix: &str, key: &str, count: usize, timestamp: u64) { + for start in (0..count).step_by(1_000) { + let mut command = ::redis::cmd("ZADD"); + command.arg(format!("{prefix}:{key}")); + for index in start..(start + 1_000).min(count) { + command.arg(timestamp).arg(format!("old-{index:08}")); + } + command.query_async::(&mut self.admin).await.unwrap(); + } + ::redis::cmd("PEXPIRE") + .arg(format!("{prefix}:{key}")) + .arg(600_000) + .query_async::(&mut self.admin) + .await + .unwrap(); + } + + async fn rows(&mut self, prefix: &str, key: &str) -> Vec<(String, f64)> { + ::redis::cmd("ZRANGE") + .arg(format!("{prefix}:{key}")) + .arg(0) + .arg(-1) + .arg("WITHSCORES") + .query_async(&mut self.admin) + .await + .unwrap() + } + + async fn seed_live(&mut self, prefix: &str, key: &str, count: usize) { + for start in (0..count).step_by(512) { + let mut command = ::redis::cmd("ZADD"); + command.arg(format!("{prefix}:{key}")); + for index in start..(start + 512).min(count) { + command.arg(100_001).arg(format!("live-{index:08}")); + } + command.query_async::(&mut self.admin).await.unwrap(); + } + } + + async fn wait_for_copy(&mut self) -> String { + tokio::time::timeout(Duration::from_secs(3), async { + loop { + let (_, keys): (u64, Vec) = ::redis::cmd("SCAN") + .arg(0) + .arg("MATCH") + .arg("cleanup:*:__usage_copy:*") + .arg("COUNT") + .arg(1000) + .query_async(&mut self.admin) + .await + .unwrap(); + if let Some(key) = keys.into_iter().next() { + return key; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("copy must become observable between its bounded chunks") + } + + async fn calls(&mut self, command: &str) -> u64 { + let info: ::redis::InfoDict = ::redis::cmd("INFO") + .arg("commandstats") + .query_async(&mut self.admin) + .await + .unwrap(); + info.get::(&format!("cmdstat_{command}")) + .and_then(|value| { + value.split(',').find_map(|field| { + field + .strip_prefix("calls=") + .and_then(|value| value.parse().ok()) + }) + }) + .unwrap_or(0) + } + + async fn reset_slowlog(&mut self) { + ::redis::cmd("CONFIG") + .arg("SET") + .arg("slowlog-log-slower-than") + .arg(0) + .arg("slowlog-max-len") + .arg(4096) + .query_async::<()>(&mut self.admin) + .await + .unwrap(); + ::redis::cmd("SLOWLOG") + .arg("RESET") + .query_async::<()>(&mut self.admin) + .await + .unwrap(); + } + + async fn max_command_us(&mut self) -> i64 { + let entries: Vec<(i64, i64, i64, Vec, String, String)> = ::redis::cmd("SLOWLOG") + .arg("GET") + .arg(4096) + .query_async(&mut self.admin) + .await + .unwrap(); + entries.into_iter().map(|entry| entry.2).max().unwrap_or(0) + } + + async fn legacy( + &mut self, + rules: &[UsageLimitRule<'_>], + event_id: &str, + now: u64, + ) -> UsageLimitCheck { + let script = ::redis::Script::new(LEGACY_USAGE_LIMIT_SCRIPT); + let mut invocation = script.prepare_invoke(); + for rule in rules { + invocation.key(format!("legacy:{}", rule.key)); + } + invocation.arg(now).arg(event_id); + for rule in rules { + invocation + .arg(rule.limit) + .arg(rule.window_seconds) + .arg(rule.retention_seconds); + } + let result: Vec = invocation.invoke_async(&mut self.admin).await.unwrap(); + match result.as_slice() { + [1, 0, 0, 0] => UsageLimitCheck::Allowed, + [0, index, limit, retry_after] => UsageLimitCheck::Rejected { + rule_index: *index as usize - 1, + limit: *limit as u64, + retry_after: *retry_after as u64, + }, + _ => panic!("unexpected legacy result: {result:?}"), + } + } +} + +fn rule(key: &str, limit: u64, window_seconds: u64) -> UsageLimitRule<'_> { + UsageLimitRule { + key, + limit, + window_seconds, + retention_seconds: 600, + } +} + +#[tokio::test] +async fn redis_usage_cleanup_whole_window_detaches_and_reuses_key_without_bulk_delete() { + for protocol in ["resp2", "resp3"] { + let Some(mut fixture) = Fixture::start(protocol).await else { + return; + }; + let rules = [UsageLimitRule { + retention_seconds: 10, + ..rule("usage:{user}:whole", 1, 10) + }]; + fixture + .seed("cleanup", rules[0].key, 100_000, 100_000) + .await; + let unlinks = fixture.calls("unlink").await; + let trims = fixture.calls("zremrangebyscore").await; + assert_eq!( + fixture.consume(&rules, "old-00000000", 110_000).await, + UsageLimitCheck::Allowed + ); + assert_eq!(fixture.calls("unlink").await, unlinks + 1); + assert_eq!(fixture.calls("zremrangebyscore").await, trims); + assert_eq!( + fixture.rows("cleanup", rules[0].key).await, + vec![("old-00000000".to_string(), 110_000.0)] + ); + assert_eq!( + fixture.consume(&rules, "old-00000000", 110_500).await, + UsageLimitCheck::Allowed + ); + assert_eq!(fixture.rows("cleanup", rules[0].key).await[0].1, 110_000.0); + let ttl: i64 = ::redis::cmd("PTTL") + .arg(format!("cleanup:{}", rules[0].key)) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert!( + (8_000..=11_000).contains(&ttl), + "new key must receive its current retention: {ttl}" + ); + fixture + .runtime + .release_usage_limits(UsageLimitReleaseInput { + rules: &rules, + event_id: "old-00000000", + }) + .await + .unwrap(); + assert!(fixture.rows("cleanup", rules[0].key).await.is_empty()); + assert_eq!( + fixture.consume(&rules, "replacement", 110_500).await, + UsageLimitCheck::Allowed + ); + } +} + +#[tokio::test] +async fn redis_usage_cleanup_preserves_live_boundary_and_out_of_order_rejection() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + let rules = [rule("usage:{user}:mixed", 1, 10)]; + fixture.seed("cleanup", rules[0].key, 4_096, 100_000).await; + ::redis::cmd("ZADD") + .arg(format!("cleanup:{}", rules[0].key)) + .arg(100_001) + .arg("live") + .query_async::(&mut fixture.admin) + .await + .unwrap(); + let unlinks = fixture.calls("unlink").await; + assert_eq!( + fixture.consume(&rules, "new", 110_000).await, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 1 + } + ); + assert_eq!(fixture.calls("unlink").await, unlinks + 1); + assert_eq!( + fixture.rows("cleanup", rules[0].key).await, + vec![("live".to_string(), 100_001.0)] + ); + assert_eq!( + fixture.consume(&rules, "new", 109_000).await, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 2 + } + ); + assert_eq!( + fixture.consume(&rules, "new", 110_001).await, + UsageLimitCheck::Allowed + ); +} + +#[tokio::test] +async fn redis_usage_cleanup_keeps_multi_window_denial_atomic_and_cleans_later_windows() { + let Some(mut fixture) = Fixture::start("resp2").await else { + return; + }; + let rules = [ + rule("usage:{user}:full", 1, 60), + rule("usage:{user}:expired", 1, 10), + ]; + for rule in &rules { + fixture + .seed( + "cleanup", + rule.key, + if rule.window_seconds == 10 { 1_024 } else { 1 }, + 100_000, + ) + .await; + } + assert_eq!( + fixture.consume(&rules, "denied", 110_000).await, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 50 + } + ); + assert!(fixture.rows("cleanup", rules[1].key).await.is_empty()); + assert_eq!( + fixture.rows("cleanup", rules[0].key).await, + vec![("old-00000000".to_string(), 100_000.0)] + ); + assert_eq!( + fixture.consume(&rules[1..], "denied", 101_000).await, + UsageLimitCheck::Allowed + ); + assert_eq!( + fixture.rows("cleanup", rules[1].key).await, + vec![("denied".to_string(), 101_000.0)] + ); +} + +#[tokio::test] +async fn redis_usage_cleanup_mixed_rebuild_preserves_rows_ttl_and_replay_boundaries() { + for protocol in ["resp2", "resp3"] { + let Some(mut fixture) = Fixture::start(protocol).await else { + return; + }; + for live in [1_usize, 16, 256, 257, 4_096, 16_384] { + let key = format!("usage:{{user}}:mixed-{live}"); + let rules = [rule(&key, live as u64, 10)]; + fixture.seed("cleanup", &key, 10_000, 100_000).await; + let mut command = ::redis::cmd("ZADD"); + command.arg(format!("cleanup:{key}")); + let rows = (0..live) + .map(|i| (format!("live-{i:03}"), 100_001.0 + i as f64)) + .collect::>(); + for (member, score) in &rows { + command.arg(score).arg(member); + } + command + .query_async::(&mut fixture.admin) + .await + .unwrap(); + let ttl_before: i64 = ::redis::cmd("PTTL") + .arg(format!("cleanup:{key}")) + .query_async(&mut fixture.admin) + .await + .unwrap(); + let unlinks = fixture.calls("unlink").await; + assert!(matches!( + fixture.consume(&rules, "new", 110_000).await, + UsageLimitCheck::Rejected { .. } + )); + assert_eq!(fixture.rows("cleanup", &key).await, rows); + assert_eq!(fixture.calls("unlink").await - unlinks, 1); + let ttl_after: i64 = ::redis::cmd("PTTL") + .arg(format!("cleanup:{key}")) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert!(ttl_after <= ttl_before && ttl_after >= ttl_before - 2_000); + assert_eq!( + fixture.consume(&rules, "live-000", 109_000).await, + UsageLimitCheck::Allowed + ); + assert_eq!(fixture.rows("cleanup", &key).await, rows); + assert!(matches!( + fixture.consume(&rules, "old-00000000", 109_000).await, + UsageLimitCheck::Rejected { .. } + )); + let temporary_exists: bool = ::redis::cmd("EXISTS") + .arg(format!("cleanup:{key}:__usage_trim")) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert!(!temporary_exists); + } + } +} + +#[tokio::test] +async fn redis_usage_cleanup_mixed_optional_acl_denials_keep_exact_state() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + for denied in ["zcount", "pttl", "exists", "unlink", "rename", "pexpire"] { + let key = format!("usage:{{user}}:acl-{denied}"); + fixture.seed("cleanup", &key, 2_048, 100_000).await; + ::redis::cmd("ZADD") + .arg(format!("cleanup:{key}")) + .arg(100_001) + .arg("live") + .query_async::(&mut fixture.admin) + .await + .unwrap(); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg("usage-cleanup") + .arg("+@all") + .arg(format!("-{denied}")) + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + assert!(matches!( + fixture.consume(&[rule(&key, 1, 10)], "new", 110_000).await, + UsageLimitCheck::Rejected { .. } + )); + assert_eq!( + fixture.rows("cleanup", &key).await, + vec![("live".into(), 100_001.0)] + ); + } +} + +#[tokio::test] +async fn redis_usage_cleanup_mixed_preserves_nonexpiring_keys_and_existing_temporary_keys() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + for collision in [false, true] { + let key = format!("usage:{{user}}:persistent-{collision}"); + fixture.seed("cleanup", &key, 2_048, 100_000).await; + ::redis::cmd("PERSIST") + .arg(format!("cleanup:{key}")) + .query_async::(&mut fixture.admin) + .await + .unwrap(); + ::redis::cmd("ZADD") + .arg(format!("cleanup:{key}")) + .arg(100_001) + .arg("live") + .query_async::(&mut fixture.admin) + .await + .unwrap(); + if collision { + ::redis::cmd("SET") + .arg(format!("cleanup:{key}:__usage_trim")) + .arg("unrelated") + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + } + assert!(matches!( + fixture.consume(&[rule(&key, 1, 10)], "new", 110_000).await, + UsageLimitCheck::Rejected { .. } + )); + let ttl: i64 = ::redis::cmd("PTTL") + .arg(format!("cleanup:{key}")) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert_eq!(ttl, -1); + assert_eq!( + fixture.rows("cleanup", &key).await, + vec![("live".into(), 100_001.0)] + ); + if collision { + let value: String = ::redis::cmd("GET") + .arg(format!("cleanup:{key}:__usage_trim")) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert_eq!(value, "unrelated"); + } + } +} + +#[tokio::test] +async fn redis_usage_cleanup_large_and_small_windows_keep_exact_cutoffs() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + for (index, count) in [0, 1, 255, 256, 257, 1_024].into_iter().enumerate() { + let key = format!("usage:{{user}}:boundary-{index}"); + let rules = [rule(&key, 1, 1)]; + let timestamp = MAX_REDIS_LUA_EXACT_INTEGER - 2_000_000; + fixture.seed("cleanup", &key, count, timestamp).await; + let unlinks = fixture.calls("unlink").await; + let result = fixture.consume(&rules, "new", timestamp + 999).await; + assert_eq!( + result, + if count == 0 { + UsageLimitCheck::Allowed + } else { + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 1, + } + } + ); + assert_eq!(fixture.calls("unlink").await, unlinks); + assert_eq!( + fixture.consume(&rules, "new", timestamp + 1_000).await, + UsageLimitCheck::Allowed + ); + assert_eq!( + fixture.calls("unlink").await, + unlinks + u64::from(count > 256) + ); + assert_eq!( + fixture.rows("cleanup", &key).await, + vec![( + "new".to_string(), + (timestamp + if count == 0 { 999 } else { 1_000 }) as f64 + )] + ); + } +} + +#[tokio::test] +async fn redis_usage_cleanup_unlink_acl_denial_uses_existing_atomic_trim() { + for protocol in ["resp2", "resp3"] { + let Some(mut fixture) = Fixture::start(protocol).await else { + return; + }; + let rules = [rule("usage:{user}:acl", 1, 10)]; + fixture.seed("cleanup", rules[0].key, 4_096, 100_000).await; + ::redis::cmd("ACL") + .arg("SETUSER") + .arg("usage-cleanup") + .arg("-unlink") + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + let trims = fixture.calls("zremrangebyscore").await; + assert_eq!( + fixture.consume(&rules, "new", 110_000).await, + UsageLimitCheck::Allowed + ); + assert_eq!(fixture.calls("zremrangebyscore").await, trims + 1); + assert_eq!( + fixture.rows("cleanup", rules[0].key).await, + vec![("new".to_string(), 110_000.0)] + ); + } +} + +#[tokio::test] +async fn redis_usage_cleanup_concurrent_expiration_preserves_the_admission_cap() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + let rules = [rule("usage:{user}:parallel-expired", 8, 10)]; + fixture.seed("cleanup", rules[0].key, 32_768, 100_000).await; + let barrier = Arc::new(tokio::sync::Barrier::new(64)); + let mut tasks = tokio::task::JoinSet::new(); + for index in 0..64 { + let runtime = fixture.runtime.clone(); + let barrier = Arc::clone(&barrier); + tasks.spawn(async move { + let event_id = format!("new-{index}"); + barrier.wait().await; + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: &event_id, + now_unix_ms: 110_000, + }) + .await + .unwrap() + }); + } + let mut allowed = 0; + while let Some(result) = tasks.join_next().await { + match result.unwrap() { + UsageLimitCheck::Allowed => allowed += 1, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 8, + retry_after: 10, + } => {} + other => panic!("unexpected admission: {other:?}"), + } + } + assert_eq!(allowed, 8); + assert_eq!(fixture.calls("unlink").await, 1); + assert_eq!(fixture.rows("cleanup", rules[0].key).await.len(), 8); +} + +#[tokio::test] +async fn redis_usage_cleanup_cached_script_recovers_after_flush_without_reconsuming() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + let rules = [rule("usage:{user}:script-reload", 1, 10)]; + assert_eq!( + fixture.consume(&rules, "original", 100_000).await, + UsageLimitCheck::Allowed + ); + ::redis::cmd("SCRIPT") + .arg("FLUSH") + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + assert_eq!( + fixture.consume(&rules, "original", 101_000).await, + UsageLimitCheck::Allowed + ); + assert_eq!( + fixture.consume(&rules, "extra", 101_000).await, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 9, + } + ); + assert_eq!( + fixture.rows("cleanup", rules[0].key).await, + vec![("original".to_string(), 100_000.0)] + ); +} + +#[tokio::test] +async fn redis_usage_cleanup_matches_legacy_for_replays_releases_and_regressing_time() { + for protocol in ["resp2", "resp3"] { + let Some(mut fixture) = Fixture::start(protocol).await else { + return; + }; + let mut rules = [ + rule("usage:{user}:one", 5, 1), + rule("usage:{user}:ten", 8, 10), + rule("usage:{user}:minute", 13, 60), + ]; + for rule in rules { + for prefix in ["cleanup", "legacy"] { + fixture.seed(prefix, rule.key, 512, 50_000).await; + } + } + let mut random = 0x1f28_e05b_a39d_c642_u64; + for index in 0..512 { + random ^= random << 13; + random ^= random >> 7; + random ^= random << 17; + let now = if index % 9 == 0 { + 110_000 + } else { + random % 130_001 + }; + let event = format!("event-{}", random % 19); + rules[0].limit = 1 + random % 7; + let actual = fixture.consume(&rules, &event, now).await; + let expected = fixture.legacy(&rules, &event, now).await; + assert_eq!( + actual, expected, + "protocol={protocol} step={index} now={now} event={event}" + ); + if index % 5 == 0 { + fixture + .runtime + .release_usage_limits(UsageLimitReleaseInput { + rules: &rules, + event_id: &event, + }) + .await + .unwrap(); + for rule in rules { + ::redis::cmd("ZREM") + .arg(format!("legacy:{}", rule.key)) + .arg(&event) + .query_async::(&mut fixture.admin) + .await + .unwrap(); + } + } + for rule in rules { + assert_eq!( + fixture.rows("cleanup", rule.key).await, + fixture.rows("legacy", rule.key).await, + "protocol={protocol} step={index} key={}", + rule.key + ); + } + } + } +} + +#[tokio::test] +#[ignore = "isolated large-window timing baseline"] +async fn redis_usage_cleanup_large_window_timing() { + let mut fixture = Fixture::start("resp3") + .await + .expect("timing baseline requires Redis"); + const MEMBERS: usize = 300_000; + let mut measurements = Vec::new(); + for round in 0..3 { + let key = format!("usage:{{user}}:timing-{round}"); + let rules = [rule(&key, 1, 10)]; + for prefix in ["legacy", "cleanup"] { + fixture.seed(prefix, &key, MEMBERS, 100_000).await; + } + let started = std::time::Instant::now(); + assert_eq!( + fixture.legacy(&rules, "new", 110_000).await, + UsageLimitCheck::Allowed + ); + let legacy_us = started.elapsed().as_micros(); + let started = std::time::Instant::now(); + assert_eq!( + fixture.consume(&rules, "new", 110_000).await, + UsageLimitCheck::Allowed + ); + let optimized_us = started.elapsed().as_micros(); + assert_eq!( + fixture.rows("cleanup", &key).await, + fixture.rows("legacy", &key).await + ); + measurements.push(serde_json::json!({"members":MEMBERS,"legacy_us":legacy_us,"optimized_us":optimized_us})); + } + eprintln!( + "usage cleanup timing: {}", + serde_json::to_string(&measurements).unwrap() + ); + + let mut mixed = Vec::new(); + for live in [1, 64, 256, 257, 4_096, 32_768, 150_000] { + let key = format!("usage:{{user}}:mixed-timing-{live}"); + let rules = [rule(&key, live, 10)]; + for prefix in ["legacy", "cleanup"] { + fixture.seed(prefix, &key, MEMBERS, 100_000).await; + fixture.seed_live(prefix, &key, live as usize).await; + } + fixture.reset_slowlog().await; + let started = std::time::Instant::now(); + let legacy = fixture.legacy(&rules, "new", 110_000).await; + let legacy_us = started.elapsed().as_micros(); + let legacy_max_command_us = fixture.max_command_us().await; + fixture.reset_slowlog().await; + let started = std::time::Instant::now(); + assert_eq!(fixture.consume(&rules, "new", 110_000).await, legacy); + let optimized_us = started.elapsed().as_micros(); + let optimized_max_command_us = fixture.max_command_us().await; + assert_eq!( + fixture.rows("cleanup", &key).await, + fixture.rows("legacy", &key).await + ); + mixed.push(serde_json::json!({"expired": MEMBERS, "live": live, "legacy_us": legacy_us, "optimized_us": optimized_us, + "legacy_max_command_us": legacy_max_command_us, "optimized_max_command_us": optimized_max_command_us})); + } + eprintln!( + "usage cleanup mixed timing: {}", + serde_json::to_string(&mixed).unwrap() + ); + + let rules = [rule("usage:{user}:steady", 10_000, 60)]; + for prefix in ["legacy", "cleanup"] { + fixture.seed(prefix, rules[0].key, 4_096, 100_000).await; + } + let mut legacy_us = Vec::new(); + let mut optimized_us = Vec::new(); + for index in 0..1_000 { + let event_id = format!("steady-{index}"); + let now = 101_000 + index; + let started = std::time::Instant::now(); + assert_eq!( + fixture.legacy(&rules, &event_id, now).await, + UsageLimitCheck::Allowed + ); + legacy_us.push(started.elapsed().as_micros()); + let started = std::time::Instant::now(); + assert_eq!( + fixture.consume(&rules, &event_id, now).await, + UsageLimitCheck::Allowed + ); + optimized_us.push(started.elapsed().as_micros()); + } + assert_eq!( + fixture.rows("cleanup", rules[0].key).await, + fixture.rows("legacy", rules[0].key).await + ); + legacy_us.sort_unstable(); + optimized_us.sort_unstable(); + eprintln!( + "usage cleanup steady timing: {}", + serde_json::json!({ + "requests_per_script": 1_000, + "seeded_live_members": 4_096, + "legacy_p50_us": legacy_us[499], + "legacy_p95_us": legacy_us[949], + "optimized_p50_us": optimized_us[499], + "optimized_p95_us": optimized_us[949], + }) + ); +} + +#[tokio::test] +async fn redis_usage_cleanup_chunked_copy_retries_when_another_rule_changes() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + let rules = [ + rule("usage:{copy}:short", 1, 10), + rule("usage:{copy}:large", 100_001, 10), + ]; + fixture.seed("cleanup", rules[0].key, 1, 100_000).await; + fixture.seed("cleanup", rules[1].key, 20_000, 100_000).await; + fixture.seed_live("cleanup", rules[1].key, 100_000).await; + let runtime = fixture.runtime.clone(); + let task = tokio::spawn(async move { + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "copier", + now_unix_ms: 110_000, + }) + .await + .unwrap() + }); + fixture.wait_for_copy().await; + // No earlier rule may have been pruned before the atomic commit. + assert_eq!(fixture.rows("cleanup", rules[0].key).await.len(), 1); + ::redis::cmd("ZADD") + .arg(format!("cleanup:{}", rules[0].key)) + .arg(109_000) + .arg("concurrent-consumer") + .query_async::(&mut fixture.admin) + .await + .unwrap(); + ::redis::cmd("ZREM") + .arg(format!("cleanup:{}", rules[1].key)) + .arg("live-00000000") + .query_async::(&mut fixture.admin) + .await + .unwrap(); + assert_eq!( + task.await.unwrap(), + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 1, + retry_after: 9, + } + ); + let rows = fixture.rows("cleanup", rules[1].key).await; + assert_eq!(rows.len(), 99_999); + assert!(!rows + .iter() + .any(|(id, _)| id == "copier" || id == "live-00000000")); + assert!(fixture.calls("watch").await >= 2); + assert_eq!(fixture.calls("zremrangebyscore").await, 1); +} + +#[tokio::test] +async fn redis_usage_cleanup_cancelled_copy_keeps_source_and_expires_scratch() { + let Some(mut fixture) = Fixture::start("resp2").await else { + return; + }; + let rules = [rule("usage:{copy}:cancelled", 100_001, 10)]; + fixture.seed("cleanup", rules[0].key, 20_000, 100_000).await; + fixture.seed_live("cleanup", rules[0].key, 100_000).await; + let runtime = fixture.runtime.clone(); + let task = tokio::spawn(async move { + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: "cancelled", + now_unix_ms: 110_000, + }) + .await + }); + let temporary = fixture.wait_for_copy().await; + task.abort(); + assert!(task.await.unwrap_err().is_cancelled()); + assert_eq!(fixture.rows("cleanup", rules[0].key).await.len(), 120_000); + let ttl: i64 = ::redis::cmd("PTTL") + .arg(&temporary) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert!(ttl > 0 && ttl <= 60_000); + assert_eq!( + fixture.consume(&rules, "replacement", 110_000).await, + UsageLimitCheck::Allowed + ); + let rows = fixture.rows("cleanup", rules[0].key).await; + assert_eq!(rows.len(), 100_001); + assert!(!rows.iter().any(|(id, _)| id == "cancelled")); +} + +#[tokio::test] +async fn redis_usage_cleanup_chunked_copy_concurrency_keeps_both_caps() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + let rules = [ + rule("usage:{copy}:parallel-one", 4104, 10), + rule("usage:{copy}:parallel-two", 4104, 10), + ]; + for rule in &rules { + fixture.seed("cleanup", rule.key, 20_000, 100_000).await; + fixture.seed_live("cleanup", rule.key, 4096).await; + } + let mut tasks = tokio::task::JoinSet::new(); + for index in 0..64 { + let runtime = fixture.runtime.clone(); + tasks.spawn(async move { + runtime + .check_and_consume_usage_limits(UsageLimitInput { + rules: &rules, + event_id: &format!("parallel-{index}"), + now_unix_ms: 110_000, + }) + .await + .unwrap() + }); + } + let mut allowed = 0; + while let Some(result) = tasks.join_next().await { + allowed += usize::from(result.unwrap() == UsageLimitCheck::Allowed); + } + assert_eq!(allowed, 8); + let first = fixture.rows("cleanup", rules[0].key).await; + assert_eq!(first.len(), 4104); + assert_eq!(first, fixture.rows("cleanup", rules[1].key).await); + assert_eq!(fixture.calls("zremrangebyscore").await, 0); +} + +#[tokio::test] +async fn redis_usage_cleanup_chunked_copy_optional_acl_denials_preserve_legacy_behavior() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + for denied in ["watch", "unwatch", "multi", "exec", "eval", "persist"] { + let key = format!("usage:{{copy}}:acl-{denied}"); + fixture.seed("cleanup", &key, 8192, 100_000).await; + fixture.seed_live("cleanup", &key, 512).await; + ::redis::cmd("ACL") + .arg("SETUSER") + .arg("usage-cleanup") + .arg("+@all") + .arg(format!("-{denied}")) + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + let trims = fixture.calls("zremrangebyscore").await; + assert_eq!( + fixture + .consume(&[rule(&key, 512, 10)], "new", 110_000) + .await, + UsageLimitCheck::Rejected { + rule_index: 0, + limit: 512, + retry_after: 1 + } + ); + assert_eq!(fixture.rows("cleanup", &key).await.len(), 512); + assert_eq!(fixture.calls("zremrangebyscore").await, trims + 1); + } +} + +#[tokio::test] +async fn redis_usage_cleanup_chunked_copy_matches_legacy_across_regressing_time() { + for protocol in ["resp2", "resp3"] { + let Some(mut fixture) = Fixture::start(protocol).await else { + return; + }; + let rules = [ + rule("usage:{history}:short", 512, 10), + rule("usage:{history}:long", 9000, 60), + ]; + for prefix in ["cleanup", "legacy"] { + for rule in &rules { + fixture.seed(prefix, rule.key, 8192, 100_000).await; + fixture.seed_live(prefix, rule.key, 512).await; + } + ::redis::cmd("PERSIST") + .arg(format!("{prefix}:{}", rules[0].key)) + .query_async::(&mut fixture.admin) + .await + .unwrap(); + } + for (event, now) in [ + ("denied", 110_000), + ("live-00000000", 109_000), + ("old-00000000", 1000), + ("expired", 110_001), + ("backdated", 102_000), + ("new", 160_001), + ] { + let expected = fixture.legacy(&rules, event, now).await; + assert_eq!(fixture.consume(&rules, event, now).await, expected); + for rule in &rules { + assert_eq!( + fixture.rows("cleanup", rule.key).await, + fixture.rows("legacy", rule.key).await + ); + } + } + assert!(fixture.calls("watch").await > 0); + } +} + +#[tokio::test] +async fn redis_usage_cleanup_copy_refuses_existing_temporary_and_denied_expiration() { + let Some(mut fixture) = Fixture::start("resp3").await else { + return; + }; + let source = "cleanup:usage:{copy}:source"; + let target = "cleanup:usage:{copy}:existing"; + fixture + .seed_live("cleanup", "usage:{copy}:source", 512) + .await; + ::redis::cmd("SET") + .arg(target) + .arg("unrelated") + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + let mut connection = ::redis::Client::open(format!( + "redis://usage-cleanup:usage-cleanup-test-password@127.0.0.1:{}/6", + fixture._server.port, + )) + .unwrap() + .get_multiplexed_async_connection() + .await + .unwrap(); + let script = ::redis::Script::new(include_str!("usage_copy.lua")); + let result = script + .key(source) + .key(target) + .arg(0) + .arg(511) + .arg(0) + .invoke_async::(&mut connection) + .await; + assert!(result.is_err()); + let value: String = ::redis::cmd("GET") + .arg(target) + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert_eq!(value, "unrelated"); + ::redis::cmd("ACL") + .arg("SETUSER") + .arg("usage-cleanup") + .arg("-pexpire") + .query_async::<()>(&mut fixture.admin) + .await + .unwrap(); + let result = script + .key(source) + .key("cleanup:usage:{copy}:denied") + .arg(0) + .arg(511) + .arg(0) + .invoke_async::(&mut connection) + .await; + assert!(result.is_err()); + let exists: bool = ::redis::cmd("EXISTS") + .arg("cleanup:usage:{copy}:denied") + .query_async(&mut fixture.admin) + .await + .unwrap(); + assert!(!exists); + assert_eq!( + fixture.rows("cleanup", "usage:{copy}:source").await.len(), + 512 + ); +} diff --git a/crates/aether-runtime/state/src/score_window.rs b/crates/aether-runtime/state/src/score_window.rs new file mode 100644 index 000000000..89bce45d5 --- /dev/null +++ b/crates/aether-runtime/state/src/score_window.rs @@ -0,0 +1,25 @@ +/// Maximum number of members examined by one server-side window aggregation. +pub const SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT: usize = 512; + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub struct ScoreWindowU64Stats { + pub sum: u64, + pub positive_count: u64, +} + +impl ScoreWindowU64Stats { + /// Members encode the value after their final colon. Invalid or zero values + /// do not contribute to the positive sample count. + pub fn from_members<'a>(members: impl IntoIterator) -> Self { + let mut stats = Self::default(); + for member in members { + let value = member + .rsplit_once(':') + .and_then(|(_, suffix)| suffix.parse::().ok()) + .unwrap_or(0); + stats.sum = stats.sum.saturating_add(value); + stats.positive_count += u64::from(value > 0); + } + stats + } +} diff --git a/crates/aether-scheduler-core/src/health.rs b/crates/aether-scheduler-core/src/health.rs index e6c507e5a..7358dfc10 100644 --- a/crates/aether-scheduler-core/src/health.rs +++ b/crates/aether-scheduler-core/src/health.rs @@ -720,6 +720,71 @@ mod tests { .expect("candidate should build") } + #[test] + fn runtime_projection_preserves_concurrency_rpm_and_failure_cooldown() { + let statuses = [ + RequestCandidateStatus::Available, + RequestCandidateStatus::Unused, + RequestCandidateStatus::Pending, + RequestCandidateStatus::Streaming, + RequestCandidateStatus::Success, + RequestCandidateStatus::Failed, + RequestCandidateStatus::Cancelled, + RequestCandidateStatus::Skipped, + ]; + let mut candidates = Vec::new(); + for (index, status) in statuses.into_iter().cycle().take(40).enumerate() { + let mut row = stored_candidate(&index.to_string(), status, 100 - index as i64); + row.api_key_id = Some("api-key".into()); + row.concurrent_requests = Some(index as u32); + row.started_at_unix_ms = (index % 2 == 0).then_some(99_000); + row.finished_at_unix_ms = (index % 3 == 0).then_some(101_000); + row.extra_data = Some(serde_json::json!({"stream_completed": true})); + candidates.push(row); + } + let projected = candidates + .iter() + .map(StoredRequestCandidate::runtime_snapshot) + .collect::>(); + for now in [90, 101, 160, 401] { + for rows in [&candidates[..], &candidates[5..], &candidates[39..]] { + let slim = &projected[candidates.len() - rows.len()..]; + assert_eq!( + count_recent_active_requests_for_api_key(rows, "api-key", now), + count_recent_active_requests_for_api_key(slim, "api-key", now) + ); + assert_eq!( + count_recent_active_requests_for_provider(rows, "provider-a", now), + count_recent_active_requests_for_provider(slim, "provider-a", now) + ); + assert_eq!( + count_recent_active_requests_for_provider_key(rows, "key-a", now), + count_recent_active_requests_for_provider_key(slim, "key-a", now) + ); + assert_eq!( + count_recent_rpm_requests_for_provider_key_since(rows, "key-a", now, Some(80)), + count_recent_rpm_requests_for_provider_key_since(slim, "key-a", now, Some(80)) + ); + assert_eq!( + is_candidate_in_recent_failure_cooldown( + rows, + "provider-a", + "endpoint-a", + "key-a", + now + ), + is_candidate_in_recent_failure_cooldown( + slim, + "provider-a", + "endpoint-a", + "key-a", + now + ) + ); + } + } + } + fn provider_catalog_key(id: &str) -> StoredProviderCatalogKey { StoredProviderCatalogKey::new( id.to_string(), diff --git a/crates/aether-testing/integration/Cargo.toml b/crates/aether-testing/integration/Cargo.toml index eafd3d999..1bc86bb8b 100644 --- a/crates/aether-testing/integration/Cargo.toml +++ b/crates/aether-testing/integration/Cargo.toml @@ -13,6 +13,7 @@ aether-crypto.workspace = true aether-data.workspace = true aether-data-contracts.workspace = true aether-gateway = { workspace = true, features = ["testkit"] } +aether-runtime.workspace = true aether-runtime-state.workspace = true aether-testkit = { workspace = true, features = ["gateway", "postgres"] } axum.workspace = true diff --git a/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs b/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs index b82892f51..8858ee33e 100644 --- a/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs +++ b/crates/aether-testing/integration/src/bin/capacity_curve_baseline.rs @@ -18,7 +18,7 @@ use axum::http::StatusCode; use axum::response::{IntoResponse, Response}; use axum::routing::any; use axum::{extract::Request, Json, Router}; -use futures_util::{SinkExt, StreamExt}; +use futures_util::{stream::FuturesUnordered, SinkExt, StreamExt}; use reqwest::Method; use serde::Serialize; use serde_json::json; @@ -83,6 +83,8 @@ struct CapacityCurvePointResult { successful_requests: usize, rejected_requests: usize, failed_requests: usize, + status_counts: BTreeMap, + non_success_status_samples: serde_json::Value, throughput_rps: u64, p50_ms: u64, p95_ms: u64, @@ -112,8 +114,16 @@ struct GateMetricSnapshot { rejected_total: u64, } -#[tokio::main] -async fn main() -> Result<(), Box> { +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + let runtime = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .thread_stack_size(8 * 1024 * 1024) + .build()?; + runtime.block_on(run()) +} + +async fn run() -> Result<(), Box> { init_test_runtime_for("capacity-curve-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; @@ -215,11 +225,11 @@ async fn run_gateway_curve( .await .map_err(std::io::Error::other)?; let duration_ms = started_at.elapsed().as_millis() as u64; - let metrics = capture_gate_metrics( - &format!("{}/_gateway/metrics", gateway.base_url()), - gate_name, - ) - .await?; + let samples = gateway + .metric_samples() + .await + .map_err(std::io::Error::other)?; + let metrics = gate_metrics(&samples, gate_name)?; points.push(capacity_point( *limit, total_requests, @@ -319,6 +329,21 @@ async fn run_tunnel_curve( let peer = connect_protocol_peer(tunnel.base_url(), config.tunnel_hold).await?; let total_requests = total_requests_for_limit(relay_concurrency, config.requests_per_point_multiplier); + let envelope = relay_envelope(); + let body_offset = + 4 + u32::from_be_bytes(envelope[..4].try_into().expect("metadata length")) as usize; + verify_tunnel_fixture(&tunnel, &envelope, body_offset, config.timeout).await?; + let header_sets = (0..total_requests) + .map(|_| { + let mut headers = + tunnel.relay_headers(&envelope[..body_offset], &envelope[body_offset..]); + headers.insert( + "content-type".to_string(), + "application/octet-stream".to_string(), + ); + headers + }) + .collect(); let probe = HttpLoadProbeConfig { url: format!( "{tunnel_base}{TUNNEL_RELAY_PATH_PREFIX}/node-baseline", @@ -329,7 +354,8 @@ async fn run_tunnel_curve( "content-type".to_string(), "application/octet-stream".to_string(), )]), - body: Some(relay_envelope()), + header_sets, + body: Some(envelope), total_requests, concurrency: relay_concurrency, timeout: config.timeout, @@ -350,7 +376,8 @@ async fn run_tunnel_curve( result, metrics, )); - drop(peer); + peer.abort(); + let _ = peer.await; } Ok(CapacityCurveScenarioReport { @@ -362,6 +389,58 @@ async fn run_tunnel_curve( }) } +async fn verify_tunnel_fixture( + tunnel: &TunnelHarness, + envelope: &[u8], + body_offset: usize, + timeout: Duration, +) -> Result<(), Box> { + let client = reqwest::Client::builder().timeout(timeout).build()?; + let url = format!( + "{}{TUNNEL_RELAY_PATH_PREFIX}/{TUNNEL_HARNESS_NODE_ID}", + tunnel.base_url() + ); + let unsigned = client.post(&url).body(envelope.to_vec()).send().await?; + if unsigned.status() != StatusCode::FORBIDDEN { + return Err(std::io::Error::other("unsigned tunnel preflight was not rejected").into()); + } + let mut signed = client.post(&url).body(envelope.to_vec()); + for (name, value) in tunnel.relay_headers(&envelope[..body_offset], &envelope[body_offset..]) { + signed = signed.header(name, value); + } + let signed = signed.build()?; + let mut tampered = signed + .try_clone() + .expect("buffered relay request should clone"); + let mut tampered_body = envelope.to_vec(); + *tampered_body + .last_mut() + .expect("relay body should be nonempty") ^= 1; + *tampered.body_mut() = Some(tampered_body.into()); + if client.execute(tampered).await?.status() != StatusCode::FORBIDDEN { + return Err(std::io::Error::other("tampered tunnel preflight was not rejected").into()); + } + let response = client + .execute( + signed + .try_clone() + .expect("buffered relay request should clone"), + ) + .await?; + let status = response.status(); + let body = response.text().await?; + if status != StatusCode::OK || body != "capacity-tunnel-stream" { + return Err(std::io::Error::other(format!( + "signed tunnel preflight failed: {status}: {body}" + )) + .into()); + } + if client.execute(signed).await?.status() != StatusCode::FORBIDDEN { + return Err(std::io::Error::other("replayed tunnel preflight was not rejected").into()); + } + Ok(()) +} + fn capacity_point( limit: usize, total_requests: usize, @@ -395,7 +474,10 @@ fn capacity_point( duration_ms, successful_requests, rejected_requests, - failed_requests: result.failed_requests, + failed_requests: total_requests.saturating_sub(successful_requests + rejected_requests), + status_counts: result.status_counts, + non_success_status_samples: serde_json::to_value(result.non_success_status_samples) + .expect("HTTP status samples should serialize"), throughput_rps, p50_ms: result.p50_ms, p95_ms: result.p95_ms, @@ -440,27 +522,29 @@ async fn capture_gate_metrics( let samples = fetch_prometheus_samples(metrics_url) .await .map_err(std::io::Error::other)?; + gate_metrics(&samples, gate_name) +} + +fn gate_metrics( + samples: &[aether_testkit::PrometheusSample], + gate_name: &str, +) -> Result> { + let required = |name| { + find_metric_value_u64(samples, name, &[("gate", gate_name)]) + .or_else(|| { + find_metric_value_u64( + samples, + &format!("aether_testkit_{name}"), + &[("gate", gate_name)], + ) + }) + .ok_or_else(|| std::io::Error::other(format!("missing {name} for gate {gate_name}"))) + }; Ok(GateMetricSnapshot { - in_flight: find_metric_value_u64(&samples, "concurrency_in_flight", &[("gate", gate_name)]) - .unwrap_or_default(), - available_permits: find_metric_value_u64( - &samples, - "concurrency_available_permits", - &[("gate", gate_name)], - ) - .unwrap_or_default(), - high_watermark: find_metric_value_u64( - &samples, - "concurrency_high_watermark", - &[("gate", gate_name)], - ) - .unwrap_or_default(), - rejected_total: find_metric_value_u64( - &samples, - "concurrency_rejected_total", - &[("gate", gate_name)], - ) - .unwrap_or_default(), + in_flight: required("concurrency_in_flight")?, + available_permits: required("concurrency_available_permits")?, + high_watermark: required("concurrency_high_watermark")?, + rejected_total: required("concurrency_rejected_total")?, }) } @@ -700,25 +784,35 @@ async fn connect_protocol_peer( )) .await?; Ok(tokio::spawn(async move { - while let Some(message) = stream.next().await { - let Ok(message) = message else { - break; - }; - match message { - Message::Binary(data) - if handle_binary_frame(&mut sink, data.to_vec(), hold) - .await - .is_err() => - { - break; + let mut responses = FuturesUnordered::new(); + loop { + tokio::select! { + message = stream.next() => { + match message { + Some(Ok(Message::Binary(data))) => { + match handle_binary_frame(&mut sink, data.to_vec()).await { + Ok(Some(stream_id)) => responses.push(async move { + tokio::time::sleep(hold).await; + stream_id + }), + Ok(None) => {}, + Err(_) => break, + } + } + Some(Ok(Message::Ping(payload))) => { + if sink.send(Message::Pong(payload)).await.is_err() { + break; + } + } + None | Some(Err(_)) | Some(Ok(Message::Close(_))) => break, + _ => {}, + } } - Message::Ping(payload) - if sink.send(Message::Pong(payload.clone())).await.is_err() => - { - break; + Some(stream_id) = responses.next(), if !responses.is_empty() => { + if send_protocol_response(&mut sink, stream_id).await.is_err() { + break; + } } - Message::Close(_) => break, - _ => {} } } let _ = sink.close().await; @@ -728,13 +822,12 @@ async fn connect_protocol_peer( async fn handle_binary_frame( sink: &mut S, data: Vec, - hold: Duration, -) -> Result<(), tokio_tungstenite::tungstenite::Error> +) -> Result, tokio_tungstenite::tungstenite::Error> where S: SinkExt + Unpin, { let Some(header) = protocol::FrameHeader::parse(&data) else { - return Ok(()); + return Ok(None); }; match header.msg_type { protocol::PING => { @@ -755,48 +848,57 @@ where .await?; } if header.flags & protocol::FLAG_END_STREAM == 0 { - return Ok(()); + return Ok(None); } - tokio::time::sleep(hold).await; - let response_meta = protocol::ResponseMeta { - status: 200, - headers: vec![( - "content-type".to_string(), - "text/plain; charset=utf-8".to_string(), - )], - }; - let response_meta_json = - serde_json::to_vec(&response_meta).expect("response metadata should serialize"); - sink.send(Message::Binary( - protocol::encode_frame( - header.stream_id, - protocol::RESPONSE_HEADERS, - 0, - &response_meta_json, - ) - .into(), - )) - .await?; - - for chunk in [ - b"capacity-".as_slice(), - b"tunnel-".as_slice(), - b"stream".as_slice(), - ] { - sink.send(Message::Binary( - protocol::encode_frame(header.stream_id, protocol::RESPONSE_BODY, 0, chunk) - .into(), - )) - .await?; - } - - sink.send(Message::Binary( - protocol::encode_frame(header.stream_id, protocol::STREAM_END, 0, &[]).into(), - )) - .await?; + return Ok(Some(header.stream_id)); } _ => {} } + Ok(None) +} + +async fn send_protocol_response( + sink: &mut S, + stream_id: u32, +) -> Result<(), tokio_tungstenite::tungstenite::Error> +where + S: SinkExt + Unpin, +{ + let response_meta = protocol::ResponseMeta { + status: 200, + headers: vec![( + "content-type".to_string(), + "text/plain; charset=utf-8".to_string(), + )], + }; + let response_meta_json = + serde_json::to_vec(&response_meta).expect("response metadata should serialize"); + sink.send(Message::Binary( + protocol::encode_frame( + stream_id, + protocol::RESPONSE_HEADERS, + 0, + &response_meta_json, + ) + .into(), + )) + .await?; + + for chunk in [ + b"capacity-".as_slice(), + b"tunnel-".as_slice(), + b"stream".as_slice(), + ] { + sink.send(Message::Binary( + protocol::encode_frame(stream_id, protocol::RESPONSE_BODY, 0, chunk).into(), + )) + .await?; + } + + sink.send(Message::Binary( + protocol::encode_frame(stream_id, protocol::STREAM_END, 0, &[]).into(), + )) + .await?; Ok(()) } diff --git a/crates/aether-testing/integration/src/bin/dependency_pressure_baseline.rs b/crates/aether-testing/integration/src/bin/dependency_pressure_baseline.rs index 37e8f66be..46f1a48a1 100644 --- a/crates/aether-testing/integration/src/bin/dependency_pressure_baseline.rs +++ b/crates/aether-testing/integration/src/bin/dependency_pressure_baseline.rs @@ -141,8 +141,13 @@ impl SummaryCollector { } } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("dependency-pressure-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs b/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs index bbc58bfae..160f12c78 100644 --- a/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs +++ b/crates/aether-testing/integration/src/bin/failure_recovery_baseline.rs @@ -172,8 +172,13 @@ impl RecoveryCollector { } } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("failure-recovery-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs b/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs index a5d293492..4240c6e34 100644 --- a/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs +++ b/crates/aether-testing/integration/src/bin/gateway_pressure_seed.rs @@ -7,6 +7,7 @@ use std::io; use std::io::Write; use std::path::{Path, PathBuf}; +use aether_crypto::PythonFernetCompat; use aether_data::repository::auth::CreateStandaloneApiKeyRecord; use aether_data::repository::wallet::WalletLookupKey; use aether_data::{ @@ -16,6 +17,7 @@ use aether_data_contracts::repository::global_models::{ CreateAdminGlobalModelRecord, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord, }; use aether_data_contracts::repository::provider_catalog::{ + ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, }; use serde_json::json; @@ -200,6 +202,12 @@ impl Config { #[tokio::main] async fn main() -> Result<(), Box> { let config = Config::from_env_and_args().map_err(|err| format!("invalid config: {err}"))?; + let encryption_key = env_value("AETHER_GATEWAY_DATA_ENCRYPTION_KEY") + .or_else(|| env_value("ENCRYPTION_KEY")) + .ok_or( + "set AETHER_GATEWAY_DATA_ENCRYPTION_KEY or ENCRYPTION_KEY to the gateway's encryption key before seeding", + )?; + let secret_cipher = PythonFernetCompat::from_secret(&encryption_key); let backends = DataBackends::from_config(DataLayerConfig::from_database(SqlDatabaseConfig { driver: DatabaseDriver::Postgres, @@ -215,10 +223,10 @@ async fn main() -> Result<(), Box> { }, }))?; - seed_provider_catalog(&backends, &config).await?; + seed_provider_catalog(&backends, &config, &secret_cipher).await?; seed_models(&backends, &config).await?; let operator_user_id = seed_operator_user(&backends, &config).await?; - seed_api_keys(&backends, &config, &operator_user_id).await?; + seed_api_keys(&backends, &config, &operator_user_id, &secret_cipher).await?; verify_candidate_selection(&backends, &config).await?; write_outputs(&config)?; @@ -243,6 +251,7 @@ async fn main() -> Result<(), Box> { async fn seed_provider_catalog( backends: &DataBackends, config: &Config, + secret_cipher: &PythonFernetCompat, ) -> Result<(), Box> { let reader = backends .read() @@ -326,7 +335,7 @@ async fn seed_provider_catalog( )? .with_transport_fields( Some(json!(["openai:chat"])), - Some(config.provider_api_key.clone()), + Some(secret_cipher.encrypt_plaintext(&config.provider_api_key)?), None, None, None, @@ -341,17 +350,40 @@ async fn seed_provider_catalog( Some(json!({"openai:chat": {"state": "closed"}})), ); - if reader - .list_keys_by_ids(std::slice::from_ref(&config.provider_key_id)) - .await? - .is_empty() - { - writer.create_key(&provider_key).await?; - } else { - writer.update_key(&provider_key).await?; + for _ in 0..8 { + let existing = reader + .list_keys_by_ids(std::slice::from_ref(&config.provider_key_id)) + .await? + .into_iter() + .next(); + let Some(existing) = existing else { + writer.create_key(&provider_key).await?; + return Ok(()); + }; + if existing.provider_id != config.provider_id { + return Err("existing pressure provider key belongs to a different provider".into()); + } + + // Randomized ciphertext changes on every seed. Fence against the observed + // credential and preserve runtime fields when rotating the configured key. + let update = ProviderCatalogKeyAdminCasUpdate { + expected_encrypted_auth_config: existing.encrypted_auth_config.clone(), + expected_credential: ProviderCatalogKeyOAuthCredentialFence { + encrypted_api_key: existing.encrypted_api_key, + auth_type: existing.auth_type, + provider_id: existing.provider_id, + provider_type: provider.provider_type.clone(), + }, + key: provider_key.clone(), + codex_rotation: None, + reset_oauth_runtime: true, + }; + if writer.compare_and_update_key_admin_state(&update).await? { + return Ok(()); + } } - Ok(()) + Err("pressure provider key changed repeatedly during seed; retry initialization".into()) } fn pressure_provider_transport_config(mock_upstream_h2c: bool) -> Option { @@ -466,9 +498,10 @@ async fn seed_api_keys( backends: &DataBackends, config: &Config, operator_user_id: &str, + secret_cipher: &PythonFernetCompat, ) -> Result<(), Box> { for index in 0..config.api_key_count { - seed_api_key(backends, config, operator_user_id, index).await?; + seed_api_key(backends, config, operator_user_id, index, secret_cipher).await?; } Ok(()) } @@ -478,6 +511,7 @@ async fn seed_api_key( config: &Config, operator_user_id: &str, key_index: usize, + secret_cipher: &PythonFernetCompat, ) -> Result<(), Box> { let auth_reader = backends .read() @@ -494,17 +528,28 @@ async fn seed_api_key( let api_key_id = pressure_api_key_id(config, key_index); let api_key_value = pressure_api_key_value(config, key_index); + let key_hash = sha256_hex(&api_key_value); + let key_encrypted = secret_cipher.encrypt_plaintext(&api_key_value)?; let existing = auth_reader .find_export_standalone_api_key_by_id(&api_key_id) .await?; + if existing + .as_ref() + .is_some_and(|record| record.key_hash != key_hash) + { + return Err(format!( + "existing pressure API key {api_key_id} has a different hash; use its original value or a new --api-key-id" + ) + .into()); + } if existing.is_none() { auth_writer .create_standalone_api_key(CreateStandaloneApiKeyRecord { user_id: operator_user_id.to_string(), api_key_id: api_key_id.clone(), - key_hash: sha256_hex(&api_key_value), - key_encrypted: Some(api_key_value), + key_hash, + key_encrypted: Some(key_encrypted), name: Some(format!("Local pressure API key {}", key_index + 1)), allowed_providers: Some(vec![config.provider_id.clone()]), allowed_api_formats: Some(vec!["openai:chat".to_string()]), @@ -526,8 +571,8 @@ async fn seed_api_key( .update_standalone_api_key_basic( aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord { api_key_id: api_key_id.clone(), - key_encrypted: None, - key_encrypted_present: false, + key_encrypted: Some(key_encrypted), + key_encrypted_present: true, name: Some(format!("Local pressure API key {}", key_index + 1)), name_present: true, force_capabilities: None, @@ -867,6 +912,8 @@ fn print_help() { println!( "Usage: cargo run -p aether-integration-tests --bin gateway_pressure_seed -- [options]\n\ \n\ +The seed and gateway must share AETHER_GATEWAY_DATA_ENCRYPTION_KEY (or ENCRYPTION_KEY).\n\ +\n\ Options:\n\ --database-url URL\n\ --output-env PATH\n\ diff --git a/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs b/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs index a29cf41c9..c7f40d512 100644 --- a/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs +++ b/crates/aether-testing/integration/src/bin/gateway_tunnel_stream_baseline.rs @@ -104,8 +104,13 @@ struct AcceptanceReport { reasons: Vec, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("gateway-tunnel-stream-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs b/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs index bbe8f2a74..6b930d46d 100644 --- a/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs +++ b/crates/aether-testing/integration/src/bin/llm_stream_stability_baseline.rs @@ -318,8 +318,13 @@ struct ProtocolPeer { stats: Arc, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("llm-stream-stability-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; config.validate().map_err(std::io::Error::other)?; diff --git a/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs b/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs index 85c96fb11..563e123cd 100644 --- a/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs +++ b/crates/aether-testing/integration/src/bin/mock_openai_upstream.rs @@ -441,6 +441,7 @@ async fn chat_completions(State(app): State, request: axum::extract::Reques completion .take() .expect("request completion guard should be present"), + false, ); } @@ -457,6 +458,14 @@ async fn chat_completions(State(app): State, request: axum::extract::Reques }; let stream = request_wants_stream(&body); if stream { + let include_usage = serde_json::from_slice::(&body) + .ok() + .and_then(|value| { + value + .pointer("/stream_options/include_usage") + .and_then(serde_json::Value::as_bool) + }) + .unwrap_or(false); record_response_header_created(&app, request_started.started_at.elapsed()); return build_chat_sse_response( app, @@ -464,6 +473,7 @@ async fn chat_completions(State(app): State, request: axum::extract::Reques completion .take() .expect("request completion guard should be present"), + include_usage, ); } // A stream truncation profile only applies after the request is known to be streaming. @@ -609,6 +619,7 @@ fn build_chat_sse_response( app: App, profile: RequestProfile, completion: RequestCompletionGuard, + include_usage: bool, ) -> Response { let response_created_at = Instant::now(); let config = app.config.clone(); @@ -659,6 +670,21 @@ fn build_chat_sse_response( yield Ok::(Bytes::from( "data: {\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n", )); + if include_usage { + let payload = json!({ + "id": "chatcmpl-mock", + "object": "chat.completion.chunk", + "created": current_unix_secs(), + "model": "mock-model", + "choices": [], + "usage": { + "prompt_tokens": 1, + "completion_tokens": config.chunks.max(1), + "total_tokens": config.chunks.max(1) + 1 + } + }); + yield Ok::(Bytes::from(format!("data: {payload}\n\n"))); + } yield Ok::(Bytes::from("data: [DONE]\n\n")); if let Some(completion) = completion.take() { completion.complete(); @@ -1465,7 +1491,11 @@ mod tests { ) { let response = client .post(url) - .json(&json!({"stream": true, "model": "mock-test"})) + .json(&json!({ + "stream": true, + "model": "mock-test", + "stream_options": {"include_usage": true} + })) .send() .await .expect("headers should arrive before the body error"); @@ -1505,6 +1535,78 @@ mod tests { !body.contains("[DONE]"), "truncated stream must not emit [DONE]" ); + assert!( + !body.contains("\"usage\""), + "truncated stream must not emit terminal usage" + ); + } + + #[tokio::test] + async fn chat_stream_usage_is_opt_in_and_precedes_done() { + for chunks in [0, 3] { + for (include_usage, assume_stream) in [ + (None, false), + (Some(false), false), + (Some(true), false), + (Some(true), true), + ] { + let config = Config { + chunks, + chunk_delay: Duration::ZERO, + assume_stream, + ..Default::default() + }; + let app = App { + metrics: Arc::new(Metrics::for_binds(&config.binds)), + bind_label: Arc::from(config.binds[0].to_string()), + config, + }; + let mut payload = json!({"stream": true, "model": "mock-test"}); + if let Some(include_usage) = include_usage { + payload["stream_options"] = json!({"include_usage": include_usage}); + } + let request = axum::http::Request::builder() + .body(Body::from(payload.to_string())) + .unwrap(); + let response = chat_completions(State(app), request).await; + assert_eq!(response.status(), StatusCode::OK); + let body = to_bytes(response.into_body(), 16 * 1024).await.unwrap(); + let body = std::str::from_utf8(&body).unwrap(); + let frames = body + .split("\n\n") + .filter_map(|frame| frame.strip_prefix("data: ")) + .collect::>(); + assert_eq!(frames.last(), Some(&"[DONE]")); + let payloads = frames[..frames.len() - 1] + .iter() + .map(|frame| serde_json::from_str::(frame).unwrap()) + .collect::>(); + let usage_chunks = payloads + .iter() + .filter(|payload| payload.get("usage").is_some()) + .collect::>(); + if include_usage == Some(true) && !assume_stream { + assert_eq!(usage_chunks.len(), 1); + assert_eq!(payloads.last(), Some(usage_chunks[0])); + assert_eq!(usage_chunks[0]["object"], "chat.completion.chunk"); + assert_eq!(usage_chunks[0]["choices"], json!([])); + assert_eq!( + usage_chunks[0]["usage"], + json!({ + "prompt_tokens": 1, + "completion_tokens": chunks.max(1), + "total_tokens": chunks.max(1) + 1 + }) + ); + assert_eq!( + payloads[payloads.len() - 2]["choices"][0]["finish_reason"], + "stop" + ); + } else { + assert!(usage_chunks.is_empty()); + } + } + } } async fn wait_for_completed(metrics: &Metrics, expected: u64) { diff --git a/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs b/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs index 29280165c..339309b32 100644 --- a/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs +++ b/crates/aether-testing/integration/src/bin/multi_instance_admission_baseline.rs @@ -94,8 +94,13 @@ struct WebSocketAdmissionProbeResult { runtime: BenchmarkRuntimeSnapshot, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("multi-instance-admission-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs b/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs index b4b8c5e6d..c82d02c3c 100644 --- a/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs +++ b/crates/aether-testing/integration/src/bin/multi_instance_owner_relay_baseline.rs @@ -66,8 +66,13 @@ struct RelayOverheadSnapshot { mean_delta_ms: i64, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("multi-instance-owner-relay-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/integration/src/bin/single_instance_baseline.rs b/crates/aether-testing/integration/src/bin/single_instance_baseline.rs index 72754efea..26e8779fc 100644 --- a/crates/aether-testing/integration/src/bin/single_instance_baseline.rs +++ b/crates/aether-testing/integration/src/bin/single_instance_baseline.rs @@ -54,8 +54,13 @@ struct SingleInstanceBaselineReport { scenarios: Vec, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("single-instance-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs b/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs index 14f3d9ea1..4194ddb1b 100644 --- a/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs +++ b/crates/aether-testing/integration/src/bin/usage_aux_counter_hotspot_baseline.rs @@ -123,8 +123,13 @@ struct LockSample { oldest_lock_wait_ms: i64, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("usage-aux-counter-hotspot-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; diff --git a/crates/aether-testing/integration/src/bin/usage_counter_hotspot_baseline.rs b/crates/aether-testing/integration/src/bin/usage_counter_hotspot_baseline.rs index 545a022a0..fb8562013 100644 --- a/crates/aether-testing/integration/src/bin/usage_counter_hotspot_baseline.rs +++ b/crates/aether-testing/integration/src/bin/usage_counter_hotspot_baseline.rs @@ -116,8 +116,13 @@ struct LockSample { oldest_lock_wait_ms: i64, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("usage-counter-hotspot-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; @@ -453,6 +458,7 @@ fn usage_record(index: usize) -> UpsertUsageRecord { let now_ms = now_unix_ms().saturating_add(index as u64); let now_secs = now_ms / 1_000; UpsertUsageRecord { + capture_retention: Default::default(), request_id: format!("usage-hotspot-{index:08}"), user_id: Some("user-hotspot".to_string()), api_key_id: Some("api-key-hotspot".to_string()), diff --git a/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs b/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs index 88d380721..894e14ef7 100644 --- a/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs +++ b/crates/aether-testing/integration/src/bin/usage_settlement_hotspot_baseline.rs @@ -18,6 +18,8 @@ use sqlx::{PgPool, Row}; use tokio::sync::Mutex; const PROVIDER_ID: &str = "provider-hotspot"; +const USER_ID: &str = "settlement-hotspot-user"; +const WALLET_ID: &str = "settlement-hotspot-wallet"; const REQUEST_PREFIX: &str = "settlement-hotspot"; const COST_PER_REQUEST_USD: f64 = 0.001; @@ -99,6 +101,7 @@ struct CounterReport { provider_monthly_outbox_rows: i64, provider_monthly_used_usd: f64, expected_provider_monthly_used_usd: f64, + wallet_consumed_usd: f64, } #[derive(Debug, Serialize, Clone, Copy, Default)] @@ -120,8 +123,13 @@ struct LockSample { oldest_lock_wait_ms: i64, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_test_runtime_for("usage-settlement-hotspot-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; @@ -239,6 +247,24 @@ async fn main() -> Result<(), Box> { } std::fs::write(path, format!("{raw}\n"))?; } + if report.failed_requests != 0 + || report.counters.settled_usage_rows != config.requests as i64 + || report.counters.settlement_snapshot_rows != config.requests as i64 + || report.counters.outbox_pending_rows != 0 + || (report.counters.provider_monthly_used_usd + - report.counters.expected_provider_monthly_used_usd) + .abs() + > 1e-8 + || (report.counters.wallet_consumed_usd + - report.counters.expected_provider_monthly_used_usd) + .abs() + > 1e-8 + { + return Err(std::io::Error::other( + "settlement baseline failed correctness checks; see report", + ) + .into()); + } Ok(()) } @@ -388,6 +414,24 @@ async fn wait_for_outbox_drain( } async fn seed_settlement_rows(pool: &PgPool, requests: usize) -> Result<(), sqlx::Error> { + sqlx::query( + "INSERT INTO users (id, username, email_verified) VALUES ($1, $1, true) ON CONFLICT (id) DO NOTHING", + ) + .bind(USER_ID) + .execute(pool) + .await?; + sqlx::query( + r#" +INSERT INTO wallets (id, user_id, balance, gift_balance, total_consumed, created_at, updated_at) +VALUES ($1, $2, $3, 0, 0, NOW(), NOW()) +ON CONFLICT (id) DO UPDATE SET balance = EXCLUDED.balance, gift_balance = 0, total_consumed = 0 +"#, + ) + .bind(WALLET_ID) + .bind(USER_ID) + .bind(requests as f64 * COST_PER_REQUEST_USD + 1.0) + .execute(pool) + .await?; sqlx::query( r#" INSERT INTO providers (id, name, provider_type, monthly_used_usd) @@ -432,6 +476,7 @@ WHERE request_id LIKE $1 INSERT INTO "usage" ( id, request_id, + user_id, provider_name, model, provider_id, @@ -445,6 +490,7 @@ INSERT INTO "usage" ( SELECT 'settlement-usage-' || LPAD(gs::TEXT, 8, '0'), $2 || '-' || LPAD(gs::TEXT, 8, '0'), + $7, 'Hotspot Provider', 'gpt-5', $3, @@ -463,6 +509,7 @@ FROM generate_series(0, $1::INTEGER - 1) AS gs .bind(COST_PER_REQUEST_USD) .bind(now_unix_ms() as i64) .bind(now_unix_secs() as i64) + .bind(USER_ID) .execute(pool) .await?; @@ -472,7 +519,7 @@ FROM generate_series(0, $1::INTEGER - 1) AS gs fn settlement_input(index: usize) -> UsageSettlementInput { UsageSettlementInput { request_id: format!("{REQUEST_PREFIX}-{index:08}"), - user_id: None, + user_id: Some(USER_ID.to_string()), api_key_id: None, api_key_is_standalone: false, provider_id: Some(PROVIDER_ID.to_string()), @@ -552,11 +599,13 @@ SELECT SELECT CAST(monthly_used_usd AS DOUBLE PRECISION) FROM providers WHERE id = $2 - ) AS provider_monthly_used_usd + ) AS provider_monthly_used_usd, + (SELECT CAST(total_consumed AS DOUBLE PRECISION) FROM wallets WHERE id = $3) AS wallet_consumed_usd "#, ) .bind(format!("{REQUEST_PREFIX}-%")) .bind(PROVIDER_ID) + .bind(WALLET_ID) .fetch_one(pool) .await?; Ok(CounterReport { @@ -568,6 +617,7 @@ SELECT provider_monthly_outbox_rows: row.try_get("provider_monthly_outbox_rows")?, provider_monthly_used_usd: row.try_get("provider_monthly_used_usd")?, expected_provider_monthly_used_usd: (requests as f64) * COST_PER_REQUEST_USD, + wallet_consumed_usd: row.try_get("wallet_consumed_usd")?, }) } diff --git a/crates/aether-testing/loadtools/src/bin/gateway_pressure_probe.rs b/crates/aether-testing/loadtools/src/bin/gateway_pressure_probe.rs index 086181516..68ef99fa5 100644 --- a/crates/aether-testing/loadtools/src/bin/gateway_pressure_probe.rs +++ b/crates/aether-testing/loadtools/src/bin/gateway_pressure_probe.rs @@ -83,6 +83,9 @@ struct GatewayPressureReport { settle_drain_elapsed_ms: u64, settle_required_metrics_available: bool, settle_missing_required_metrics: Vec, + settle_baseline: SettleDrainBaseline, + settle_final_metrics: BTreeMap, + settle_observations: Vec, load: HttpLoadProbeResult, metrics: GatewayPressureMetricsSummary, } @@ -573,9 +576,46 @@ struct SettleDrainResult { elapsed: Duration, required_metrics_available: bool, missing_required_metrics: Vec, + observations: Vec, } -#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +#[derive(Debug, Clone, Serialize)] +struct SettleDrainObservation { + elapsed_ms: u64, + drained: bool, + metrics: BTreeMap, +} + +fn observe_settle_drain( + observations: &mut Vec, + elapsed: Duration, + samples: &[PrometheusSample], + drained: bool, +) { + let observation = SettleDrainObservation { + elapsed_ms: elapsed.as_millis() as u64, + drained, + metrics: REQUIRED_SETTLE_DRAIN_METRICS + .iter() + .copied() + .chain([ + "gateway_http_connections_in_flight", + "gateway_process_open_fds", + "usage_runtime_producers_in_flight", + "usage_runtime_delayed_lifecycle_pending", + ]) + .filter(|name| metric_is_available(samples, name)) + .map(|name| (name.to_string(), metric_max(samples, name))) + .collect(), + }; + if observations.len() < 256 { + observations.push(observation); + } else { + observations[255] = observation; + } +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize)] struct SettleDrainBaseline { tokio_alive_tasks: u64, tokio_global_queue_depth: u64, @@ -2302,6 +2342,8 @@ impl GatewayDbPoolPressureWindow { fn samples_are_drained(samples: &[PrometheusSample], baseline: SettleDrainBaseline) -> bool { missing_required_settle_drain_metrics(samples).is_empty() + && metric_max(samples, "usage_runtime_producers_in_flight") == 0 + && metric_max(samples, "usage_runtime_delayed_lifecycle_pending") == 0 && metric_max(samples, "request_candidate_queue_depth") == 0 && metric_max(samples, "request_candidate_queue_pending_depth") == 0 && metric_max(samples, "request_candidate_active_queue_depth") == 0 @@ -2403,6 +2445,7 @@ async fn wait_for_settle_drain( ) -> SettleDrainResult { if settle_after.is_zero() { return SettleDrainResult { + observations: Vec::new(), completed: false, elapsed: Duration::ZERO, required_metrics_available: false, @@ -2422,12 +2465,14 @@ async fn wait_for_settle_drain( .min(Duration::from_millis(500)) .max(Duration::from_millis(50)); let mut stability = SettleDrainStability::default(); + let mut observations = Vec::new(); loop { match fetch_prometheus_samples(metrics_url).await { Ok(samples) => { missing_required_metrics = missing_required_settle_drain_metrics(&samples); let drained = samples_are_drained(&samples, baseline); + observe_settle_drain(&mut observations, started.elapsed(), &samples, drained); let terminal_loss_detected = { let mut snapshot = summary.lock().await; snapshot.observe(&samples); @@ -2435,6 +2480,7 @@ async fn wait_for_settle_drain( }; if terminal_loss_detected { return SettleDrainResult { + observations, completed: false, elapsed: started.elapsed(), required_metrics_available: missing_required_metrics.is_empty(), @@ -2443,6 +2489,7 @@ async fn wait_for_settle_drain( } if stability.observe(started.elapsed(), drained) { return SettleDrainResult { + observations, completed: true, elapsed: started.elapsed(), required_metrics_available: true, @@ -2465,6 +2512,12 @@ async fn wait_for_settle_drain( let mut completed = false; if let Ok(samples) = fetch_prometheus_samples(metrics_url).await { + observe_settle_drain( + &mut observations, + started.elapsed(), + &samples, + samples_are_drained(&samples, baseline), + ); missing_required_metrics = missing_required_settle_drain_metrics(&samples); completed = stability.observe(started.elapsed(), samples_are_drained(&samples, baseline)); let mut snapshot = summary.lock().await; @@ -2475,6 +2528,7 @@ async fn wait_for_settle_drain( } SettleDrainResult { + observations, completed, elapsed: started.elapsed(), required_metrics_available: missing_required_metrics.is_empty(), @@ -2540,13 +2594,25 @@ async fn main() -> Result<(), Box> { suite: "gateway_pressure_probe", acceptance_contract_version: ACCEPTANCE_CONTRACT_VERSION, target_url: config.load.url, - metrics_url: config.metrics_url, + metrics_url: config.metrics_url.clone(), sample_interval_ms: config.sample_interval.as_millis() as u64, settle_after_ms: config.settle_after.as_millis() as u64, settle_drain_completed: settle_drain.completed, settle_drain_elapsed_ms: settle_drain.elapsed.as_millis() as u64, settle_required_metrics_available: settle_drain.required_metrics_available, settle_missing_required_metrics: settle_drain.missing_required_metrics, + settle_baseline: settle_drain_baseline, + settle_observations: settle_drain.observations, + settle_final_metrics: fetch_prometheus_samples(&config.metrics_url) + .await + .map(|samples| { + REQUIRED_SETTLE_DRAIN_METRICS + .iter() + .filter(|name| metric_is_available(&samples, name)) + .map(|name| (name.to_string(), metric_max(&samples, name))) + .collect() + }) + .unwrap_or_default(), load, metrics: Arc::try_unwrap(summary) .unwrap_or_else(|_| panic!("metrics summary still referenced")) diff --git a/crates/aether-testing/loadtools/src/bin/redis_worker_baseline.rs b/crates/aether-testing/loadtools/src/bin/redis_worker_baseline.rs index 49b1e7663..cc547739a 100644 --- a/crates/aether-testing/loadtools/src/bin/redis_worker_baseline.rs +++ b/crates/aether-testing/loadtools/src/bin/redis_worker_baseline.rs @@ -52,8 +52,13 @@ struct RedisWorkerBaselineReport { ack: OperationSummary, } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_load_runtime_for("redis-worker-baseline"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/loadtools/src/bin/runtime_redis_pressure.rs b/crates/aether-testing/loadtools/src/bin/runtime_redis_pressure.rs index 8cf28bb75..0c23ac4b1 100644 --- a/crates/aether-testing/loadtools/src/bin/runtime_redis_pressure.rs +++ b/crates/aether-testing/loadtools/src/bin/runtime_redis_pressure.rs @@ -121,8 +121,13 @@ impl SummaryCollector { } } +fn main() -> Result<(), Box> { + let _log_shutdown = aether_runtime::LogShutdownGuard::new(); + run() +} + #[tokio::main] -async fn main() -> Result<(), Box> { +async fn run() -> Result<(), Box> { init_load_runtime_for("runtime-redis-pressure"); let config = parse_args(std::env::args().skip(1).collect())?; let report = run_suite(&config).await?; diff --git a/crates/aether-testing/testkit/src/gateway.rs b/crates/aether-testing/testkit/src/gateway.rs index 20e806193..dbe7bf37b 100644 --- a/crates/aether-testing/testkit/src/gateway.rs +++ b/crates/aether-testing/testkit/src/gateway.rs @@ -31,6 +31,7 @@ impl GatewayHarnessConfig { #[derive(Debug)] pub struct GatewayHarness { server: SpawnedServer, + state: AppState, } impl GatewayHarness { @@ -67,7 +68,7 @@ impl GatewayHarness { if let Some(gate) = config.distributed_request_gate { state = state.with_distributed_request_concurrency_gate(gate); } - let router = build_router_with_state(state); + let router = build_router_with_state(state.clone()); let server = match port { Some(port) => SpawnedServer::start_on_port(port, router) .await @@ -76,7 +77,7 @@ impl GatewayHarness { .await .map_err(|err| format!("failed to start gateway harness: {err}"))?, }; - Ok(Self { server }) + Ok(Self { server, state }) } pub fn base_url(&self) -> &str { @@ -86,4 +87,20 @@ impl GatewayHarness { pub fn port(&self) -> u16 { self.server.port() } + + pub async fn metric_samples(&self) -> Result, String> { + let samples = aether_gateway::testkit::gateway_metric_samples(&self.state).await?; + Ok(samples + .into_iter() + .map(|sample| crate::PrometheusSample { + name: sample.name.to_string(), + labels: sample + .labels + .into_iter() + .map(|label| (label.key.to_string(), label.value)) + .collect(), + value: sample.value.to_string(), + }) + .collect()) + } } diff --git a/crates/aether-testing/testkit/src/tunnel.rs b/crates/aether-testing/testkit/src/tunnel.rs index e3515b022..56b4c0975 100644 --- a/crates/aether-testing/testkit/src/tunnel.rs +++ b/crates/aether-testing/testkit/src/tunnel.rs @@ -1,9 +1,6 @@ use std::time::Duration; -use aether_gateway::{ - build_tunnel_runtime_router_with_state, TunnelConnConfig, TunnelControlPlaneClient, - TunnelRuntimeState, -}; +use aether_gateway::{TunnelConnConfig, TunnelControlPlaneClient, TunnelRuntimeState}; use aether_runtime_state::RuntimeSemaphore; use crate::server::SpawnedServer; @@ -11,6 +8,8 @@ use crate::server::SpawnedServer; pub const TUNNEL_HARNESS_NODE_ID: &str = "node-baseline"; pub const TUNNEL_HARNESS_GENERATION: &str = "tunnel-harness-generation-1"; pub const TUNNEL_HARNESS_MANAGEMENT_TOKEN: &str = "ae-tunnel-harness-management-token"; +const RELAY_INSTANCE: &str = "tunnel-harness"; +const RELAY_SECRET: &[u8] = b"tunnel-harness-relay-secret-32-bytes-minimum"; #[derive(Debug, Clone)] pub struct TunnelHarnessConfig { @@ -38,6 +37,7 @@ impl Default for TunnelHarnessConfig { #[derive(Debug)] pub struct TunnelHarness { server: SpawnedServer, + node_id: String, } impl TunnelHarness { @@ -74,7 +74,11 @@ impl TunnelHarness { TUNNEL_HARNESS_GENERATION, TUNNEL_HARNESS_MANAGEMENT_TOKEN, )?; - let router = build_tunnel_runtime_router_with_state(state); + let router = aether_gateway::testkit::build_tunnel_pressure_router( + state, + RELAY_INSTANCE, + RELAY_SECRET, + )?; let server = match port { Some(port) => SpawnedServer::start_on_port(port, router) .await @@ -83,7 +87,10 @@ impl TunnelHarness { .await .map_err(|err| format!("failed to start tunnel harness: {err}"))?, }; - Ok(Self { server }) + Ok(Self { + server, + node_id: config.node_id, + }) } pub fn base_url(&self) -> &str { @@ -93,6 +100,50 @@ impl TunnelHarness { pub fn port(&self) -> u16 { self.server.port() } + + pub fn relay_headers( + &self, + metadata_envelope: &[u8], + body: &[u8], + ) -> std::collections::BTreeMap { + use aether_contracts::tunnel::*; + use std::sync::atomic::{AtomicU64, Ordering}; + static NONCE: AtomicU64 = AtomicU64::new(0); + let timestamp = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs(); + let nonce = format!("harness-{}", NONCE.fetch_add(1, Ordering::Relaxed)); + let digest = tunnel_relay_payload_digest(metadata_envelope, body); + let signature = sign_tunnel_relay_request( + RELAY_SECRET, + "load-probe", + RELAY_INSTANCE, + &self.node_id, + "", + false, + timestamp, + &nonce, + &digest, + ); + [ + (TUNNEL_RELAY_AUTH_SENDER_HEADER, "load-probe".to_string()), + ( + TUNNEL_RELAY_OWNER_INSTANCE_HEADER, + RELAY_INSTANCE.to_string(), + ), + (TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp.to_string()), + (TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce), + ( + TUNNEL_RELAY_AUTH_PAYLOAD_HEADER, + digest.encode_header_value(), + ), + (TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature), + ] + .into_iter() + .map(|(key, value)| (key.to_string(), value)) + .collect() + } } pub fn insert_tunnel_harness_auth_headers( diff --git a/crates/aether-usage/runtime/src/body_capture.rs b/crates/aether-usage/runtime/src/body_capture.rs index 656e06d31..e0508dac3 100644 --- a/crates/aether-usage/runtime/src/body_capture.rs +++ b/crates/aether-usage/runtime/src/body_capture.rs @@ -623,6 +623,10 @@ fn append_body_capture_metadata_entry( ); } +pub(crate) fn mark_usage_event_capture_truncated(metadata: &mut Option, key: &str) { + aether_data_contracts::repository::usage::mark_usage_capture_memory_omitted(metadata, key); +} + fn upsert_body_capture_metadata_value_entry( metadata: &mut Option, key: &str, diff --git a/crates/aether-usage/runtime/src/config.rs b/crates/aether-usage/runtime/src/config.rs index cd9ca9487..d020c872c 100644 --- a/crates/aether-usage/runtime/src/config.rs +++ b/crates/aether-usage/runtime/src/config.rs @@ -15,6 +15,7 @@ pub struct UsageRuntimeConfig { pub consumer_group: String, pub dlq_stream_key: String, pub stream_maxlen: usize, + pub queue_payload_max_bytes: usize, pub consumer_batch_size: usize, pub consumer_block_ms: u64, pub reclaim_idle_ms: u64, @@ -47,6 +48,7 @@ impl Default for UsageRuntimeConfig { consumer_group: "usage_consumers".to_string(), dlq_stream_key: "usage:events:dlq".to_string(), stream_maxlen: 200_000, + queue_payload_max_bytes: 1024 * 1024, consumer_batch_size: 128, consumer_block_ms: 500, reclaim_idle_ms: 60_000, @@ -90,6 +92,11 @@ impl UsageRuntimeConfig { "usage runtime dlq_stream_key cannot be empty".to_string(), )); } + if self.stream_key == self.dlq_stream_key { + return Err(DataLayerError::InvalidConfiguration( + "usage runtime stream_key and dlq_stream_key must be different".to_string(), + )); + } if self.worker_count == 0 { return Err(DataLayerError::InvalidConfiguration( "usage runtime worker_count must be positive".to_string(), @@ -124,6 +131,11 @@ impl UsageRuntimeConfig { "usage runtime stream_maxlen must be positive".to_string(), )); } + if self.queue_payload_max_bytes == 0 { + return Err(DataLayerError::InvalidConfiguration( + "usage runtime queue_payload_max_bytes must be positive".to_string(), + )); + } if self.consumer_batch_size == 0 { return Err(DataLayerError::InvalidConfiguration( "usage runtime consumer_batch_size must be positive".to_string(), @@ -213,6 +225,19 @@ mod tests { assert!(config.validate().is_err()); } + #[test] + fn enabled_config_rejects_dead_letter_stream_equal_to_source() { + let mut config = UsageRuntimeConfig::default(); + config.dlq_stream_key = config.stream_key.clone(); + assert!(config.validate().is_ok()); + config.enabled = true; + assert!(matches!( + config.validate(), + Err(aether_data_contracts::DataLayerError::InvalidConfiguration(message)) + if message.contains("must be different") + )); + } + #[test] fn enabled_config_rejects_zero_terminal_submission_limit() { let config = UsageRuntimeConfig { @@ -222,4 +247,20 @@ mod tests { }; assert!(config.validate().is_err()); } + + #[test] + fn queue_payload_limit_defaults_to_one_mib_and_rejects_zero_when_enabled() { + let mut config = UsageRuntimeConfig::default(); + assert_eq!(config.queue_payload_max_bytes, 1024 * 1024); + config.queue_payload_max_bytes = 0; + assert!(config.validate().is_ok()); + config.enabled = true; + assert!(matches!( + config.validate(), + Err(aether_data_contracts::DataLayerError::InvalidConfiguration(message)) + if message.contains("queue_payload_max_bytes") + )); + config.queue_payload_max_bytes = 1; + assert!(config.validate().is_ok()); + } } diff --git a/crates/aether-usage/runtime/src/dead_letter_encoding.rs b/crates/aether-usage/runtime/src/dead_letter_encoding.rs new file mode 100644 index 000000000..7b95b56da --- /dev/null +++ b/crates/aether-usage/runtime/src/dead_letter_encoding.rs @@ -0,0 +1,541 @@ +use std::collections::BTreeMap; +use std::io::{self, Write}; +use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; +use std::sync::{Arc, LazyLock}; + +use aether_data_contracts::DataLayerError; +use aether_runtime_state::RuntimeQueueEntry; +use serde::Serialize; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; + +const DEFAULT_ENCODING_BUDGET_BYTES: usize = 64 * 1024 * 1024; +const DEFAULT_ENCODING_JOBS: usize = 4; +const MAX_ENCODING_JOBS: usize = 128; + +static ENCODING_BUDGET: LazyLock> = LazyLock::new(|| { + Arc::new(DeadLetterEncodingBudget::new( + configured_limit( + std::env::var("AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES") + .ok() + .as_deref(), + DEFAULT_ENCODING_BUDGET_BYTES, + maximum_budget_bytes(), + ), + configured_limit( + std::env::var("AETHER_USAGE_DLQ_ENCODING_MAX_JOBS") + .ok() + .as_deref(), + DEFAULT_ENCODING_JOBS, + MAX_ENCODING_JOBS, + ), + )) +}); + +pub(crate) fn shared_dead_letter_encoding_budget() -> Arc { + Arc::clone(&ENCODING_BUDGET) +} + +pub(crate) fn dead_letter_encoding_metrics() -> DeadLetterEncodingSnapshot { + ENCODING_BUDGET.snapshot() +} + +fn maximum_budget_bytes() -> usize { + Semaphore::MAX_PERMITS.min(u32::MAX as usize) +} + +fn configured_limit(raw: Option<&str>, fallback: usize, maximum: usize) -> usize { + raw.and_then(|raw| raw.trim().parse::().ok()) + .filter(|value| *value > 0) + .map(|value| value.min(maximum as u128) as usize) + .unwrap_or(fallback.min(maximum)) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct DeadLetterEncodingSnapshot { + pub(crate) limit_bytes: usize, + pub(crate) job_limit: usize, + pub(crate) reserved_bytes: usize, + pub(crate) active_jobs: usize, + pub(crate) capacity_rejected_total: u64, + pub(crate) oversized_rejected_total: u64, + pub(crate) encoded_total: u64, +} + +/// Reserves raw string lengths plus their worst-case JSON encoding before any +/// clone or encoding allocation. This excludes collection/allocation overhead and +/// Redis command, packed-command, and connection buffers; it is not an RSS limit. +pub(crate) struct DeadLetterEncodingBudget { + limit_bytes: usize, + job_limit: usize, + bytes: Arc, + jobs: Arc, + reserved_bytes: AtomicUsize, + active_jobs: AtomicUsize, + capacity_rejected_total: AtomicU64, + oversized_rejected_total: AtomicU64, + encoded_total: AtomicU64, + #[cfg(test)] + encode_hook: std::sync::Mutex>>, +} + +impl DeadLetterEncodingBudget { + pub(crate) fn new(limit_bytes: usize, job_limit: usize) -> Self { + let limit_bytes = limit_bytes.min(maximum_budget_bytes()); + let job_limit = job_limit.clamp(1, MAX_ENCODING_JOBS); + Self { + limit_bytes, + job_limit, + bytes: Arc::new(Semaphore::new(limit_bytes)), + jobs: Arc::new(Semaphore::new(job_limit)), + reserved_bytes: AtomicUsize::new(0), + active_jobs: AtomicUsize::new(0), + capacity_rejected_total: AtomicU64::new(0), + oversized_rejected_total: AtomicU64::new(0), + encoded_total: AtomicU64::new(0), + #[cfg(test)] + encode_hook: std::sync::Mutex::new(None), + } + } + + pub(crate) fn snapshot(&self) -> DeadLetterEncodingSnapshot { + DeadLetterEncodingSnapshot { + limit_bytes: self.limit_bytes, + job_limit: self.job_limit, + reserved_bytes: self.reserved_bytes.load(Ordering::Relaxed), + active_jobs: self.active_jobs.load(Ordering::Relaxed), + capacity_rejected_total: self.capacity_rejected_total.load(Ordering::Relaxed), + oversized_rejected_total: self.oversized_rejected_total.load(Ordering::Relaxed), + encoded_total: self.encoded_total.load(Ordering::Relaxed), + } + } + + pub(crate) fn try_reserve( + self: &Arc, + entry: &RuntimeQueueEntry, + error: &str, + ) -> Result { + let size = encoding_size(entry, error).filter(|size| size.total <= self.limit_bytes); + let Some(size) = size else { + self.oversized_rejected_total + .fetch_add(1, Ordering::Relaxed); + return Err(DataLayerError::InvalidInput(format!( + "dead-letter raw fields and worst-case JSON exceed the {}-byte encoding budget", + self.limit_bytes + ))); + }; + let job_permit = Arc::clone(&self.jobs) + .try_acquire_owned() + .map_err(|_| self.capacity_error())?; + let byte_permit = Arc::clone(&self.bytes) + .try_acquire_many_owned(size.total as u32) + .map_err(|_| self.capacity_error())?; + self.reserved_bytes.fetch_add(size.total, Ordering::Relaxed); + self.active_jobs.fetch_add(1, Ordering::Relaxed); + Ok(DeadLetterEncodingReservation { + budget: Arc::clone(self), + size, + _byte_permit: byte_permit, + _job_permit: job_permit, + }) + } + + fn capacity_error(&self) -> DataLayerError { + self.capacity_rejected_total.fetch_add(1, Ordering::Relaxed); + DataLayerError::TimedOut("dead-letter encoding capacity is exhausted".to_string()) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct EncodingSize { + json: usize, + total: usize, +} + +fn encoding_size(entry: &RuntimeQueueEntry, error: &str) -> Option { + let mut raw = entry.id.len().checked_add(error.len())?; + for (key, value) in &entry.fields { + raw = raw.checked_add(key.len())?.checked_add(value.len())?; + } + checked_encoding_size(raw, entry.fields.len()) +} + +fn checked_encoding_size(raw: usize, field_count: usize) -> Option { + // Every string byte needs at most six bytes (\u00XX). Each map entry adds + // four quotes, a colon and at most one comma to the empty envelope. + const EMPTY_ENVELOPE_BYTES: usize = br#"{"entry_id":"","fields":{},"error":""}"#.len(); + let json = raw + .checked_mul(6)? + .checked_add(field_count.checked_mul(6)?)? + .checked_add(EMPTY_ENVELOPE_BYTES)?; + Some(EncodingSize { + json, + total: raw.checked_add(json)?, + }) +} + +pub(crate) struct DeadLetterEncodingReservation { + budget: Arc, + size: EncodingSize, + _byte_permit: OwnedSemaphorePermit, + _job_permit: OwnedSemaphorePermit, +} + +impl DeadLetterEncodingReservation { + pub(crate) fn encode_owned( + self, + entry: RuntimeQueueEntry, + error: String, + ) -> impl std::future::Future> + Send { + let input = EncodingInput { + entry, + error, + reservation: self, + }; + async move { + tokio::task::spawn_blocking(move || input.encode()) + .await + .map_err(|error| { + DataLayerError::UnexpectedValue(format!( + "dead-letter encoding task failed: {error}" + )) + })? + } + } +} + +impl Drop for DeadLetterEncodingReservation { + fn drop(&mut self) { + self.budget + .reserved_bytes + .fetch_sub(self.size.total, Ordering::Relaxed); + self.budget.active_jobs.fetch_sub(1, Ordering::Relaxed); + } +} + +// Field order also protects cancellation/panic: raw data is dropped before its +// reservation, including when a queued blocking task never starts. +struct EncodingInput { + entry: RuntimeQueueEntry, + error: String, + reservation: DeadLetterEncodingReservation, +} + +#[derive(Serialize)] +struct DeadLetterPayload<'a> { + entry_id: &'a str, + fields: &'a BTreeMap, + error: &'a str, +} + +impl EncodingInput { + fn encode(self) -> Result { + #[cfg(test)] + { + let hook = self.reservation.budget.encode_hook.lock().unwrap().take(); + if let Some(hook) = hook { + hook(); + } + } + let mut writer = BoundedJsonWriter::new(self.reservation.size.json); + serde_json::to_writer( + &mut writer, + &DeadLetterPayload { + entry_id: &self.entry.id, + fields: &self.entry.fields, + error: &self.error, + }, + ) + .map_err(|error| { + DataLayerError::UnexpectedValue(format!( + "failed to encode complete dead-letter fields: {error}" + )) + })?; + let payload = String::from_utf8(writer.bytes).map_err(|error| { + DataLayerError::UnexpectedValue(format!("dead-letter JSON was not UTF-8: {error}")) + })?; + self.reservation + .budget + .encoded_total + .fetch_add(1, Ordering::Relaxed); + Ok(EncodedDeadLetter { + entry_id: self.entry.id, + fields: BTreeMap::from([("payload".to_string(), payload)]), + _reservation: self.reservation, + }) + } +} + +pub(crate) struct EncodedDeadLetter { + pub(crate) entry_id: String, + pub(crate) fields: BTreeMap, + // Remains owned by the result until transfer/append finishes. + _reservation: DeadLetterEncodingReservation, +} + +struct BoundedJsonWriter { + bytes: Vec, + max_bytes: usize, +} + +impl BoundedJsonWriter { + fn new(max_bytes: usize) -> Self { + Self { + bytes: Vec::new(), + max_bytes, + } + } +} + +impl Write for BoundedJsonWriter { + fn write(&mut self, bytes: &[u8]) -> io::Result { + if bytes.len() > self.max_bytes.saturating_sub(self.bytes.len()) { + return Err(io::Error::other("dead-letter JSON encoding bound exceeded")); + } + let required = self.bytes.len() + bytes.len(); + if required > self.bytes.capacity() { + let capacity = required + .max(self.bytes.capacity().saturating_mul(2)) + .min(self.max_bytes); + self.bytes + .try_reserve_exact(capacity - self.bytes.len()) + .map_err(|error| { + io::Error::other(format!("dead-letter JSON allocation failed: {error}")) + })?; + } + self.bytes.extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use super::*; + + fn test_entry() -> RuntimeQueueEntry { + RuntimeQueueEntry { + id: "10-3".to_string(), + fields: BTreeMap::from([ + ("payload".to_string(), "historical payload".to_string()), + ("extra".to_string(), "original metadata".to_string()), + ]), + } + } + + #[tokio::test] + async fn dead_letter_encoding_preserves_complete_wire_and_all_escape_forms() { + let control_bytes = (0u8..=31).map(char::from).collect::(); + let mut entry = test_entry(); + entry.id.push_str("\"\\\n"); + entry.fields.insert( + format!("{control_bytes}\"\\"), + format!("{control_bytes}\"\\\u{4e2d}\u{6587}\u{1f600}"), + ); + let error = format!("error:{control_bytes}\"\\\u{00e9}"); + let expected = serde_json::to_string(&DeadLetterPayload { + entry_id: &entry.id, + fields: &entry.fields, + error: &error, + }) + .unwrap(); + let size = encoding_size(&entry, &error).unwrap(); + assert!(expected.len() <= size.json); + let budget = Arc::new(DeadLetterEncodingBudget::new(size.total, 1)); + let encoded = budget + .try_reserve(&entry, &error) + .unwrap() + .encode_owned(entry.clone(), error.clone()) + .await + .unwrap(); + assert_eq!(encoded.entry_id, entry.id); + assert_eq!(encoded.fields.len(), 1); + assert_eq!(encoded.fields["payload"], expected); + let decoded: serde_json::Value = serde_json::from_str(&encoded.fields["payload"]).unwrap(); + assert_eq!( + decoded["fields"], + serde_json::to_value(entry.fields).unwrap() + ); + assert_eq!(decoded["error"], error); + assert_eq!(budget.snapshot().encoded_total, 1); + assert_eq!(budget.snapshot().reserved_bytes, size.total); + assert_eq!(budget.snapshot().active_jobs, 1); + drop(encoded); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().active_jobs, 0); + } + + #[test] + fn dead_letter_encoding_budget_rejects_oversize_and_saturation_without_waiters() { + let entry = test_entry(); + let size = encoding_size(&entry, "failure").unwrap(); + let too_small = Arc::new(DeadLetterEncodingBudget::new(size.total - 1, 1)); + assert!(matches!( + too_small.try_reserve(&entry, "failure"), + Err(DataLayerError::InvalidInput(_)) + )); + assert_eq!(too_small.snapshot().oversized_rejected_total, 1); + assert_eq!(too_small.snapshot().reserved_bytes, 0); + assert_eq!(too_small.snapshot().active_jobs, 0); + + for (limit, jobs) in [(size.total, 2), (size.total * 2, 1)] { + let budget = Arc::new(DeadLetterEncodingBudget::new(limit, jobs)); + let first = budget.try_reserve(&entry, "failure").unwrap(); + assert!(matches!( + budget.try_reserve(&entry, "failure"), + Err(DataLayerError::TimedOut(_)) + )); + assert_eq!(budget.snapshot().capacity_rejected_total, 1); + assert_eq!(budget.snapshot().reserved_bytes, size.total); + assert_eq!(budget.snapshot().active_jobs, 1); + assert_eq!(budget.jobs.available_permits(), jobs - 1); + drop(first); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.jobs.available_permits(), jobs); + assert_eq!(budget.bytes.available_permits(), limit); + drop(budget.try_reserve(&entry, "failure").unwrap()); + } + } + + #[test] + fn dead_letter_encoding_bounds_and_environment_cannot_overflow() { + assert!(checked_encoding_size(usize::MAX, 0).is_none()); + assert!(checked_encoding_size(usize::MAX / 6, 0).is_none()); + assert!(checked_encoding_size(0, usize::MAX).is_none()); + assert!(checked_encoding_size(usize::MAX / 7, 1).is_none()); + assert!(checked_encoding_size(0, 0).unwrap().total > 0); + assert_eq!(configured_limit(None, 4, 128), 4); + assert_eq!(configured_limit(Some("0"), 4, 128), 4); + assert_eq!(configured_limit(Some("invalid"), 4, 128), 4); + assert_eq!(configured_limit(Some(" 2 "), 4, 128), 2); + assert_eq!(configured_limit(Some("99999999999"), 4, 128), 128); + assert_eq!( + configured_limit(Some(&u128::MAX.to_string()), 4, maximum_budget_bytes()), + maximum_budget_bytes() + ); + let budget = DeadLetterEncodingBudget::new(usize::MAX, usize::MAX); + assert_eq!(budget.snapshot().limit_bytes, maximum_budget_bytes()); + assert_eq!(budget.snapshot().job_limit, MAX_ENCODING_JOBS); + } + + #[test] + fn dead_letter_encoding_bounded_writer_grows_geometrically_and_stops_at_limit() { + let mut writer = BoundedJsonWriter::new(4096); + let mut allocations = 0; + for _ in 0..4096 { + let previous_capacity = writer.bytes.capacity(); + writer.write_all(b"x").unwrap(); + allocations += usize::from(previous_capacity != writer.bytes.capacity()); + } + assert!(allocations <= 13, "allocations: {allocations}"); + assert!(writer.write_all(b"y").is_err()); + assert_eq!(writer.bytes.len(), 4096); + assert!(writer.bytes.iter().all(|byte| *byte == b'x')); + } + + #[tokio::test] + async fn dead_letter_encoding_unpolled_cancellation_releases_reservation() { + let entry = test_entry(); + let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1)); + let reservation = budget.try_reserve(&entry, "failure").unwrap(); + let future = reservation.encode_owned(entry, "failure".to_string()); + assert_eq!(budget.snapshot().active_jobs, 1); + drop(future); + assert_eq!(budget.snapshot().active_jobs, 0); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().encoded_total, 0); + } + + #[tokio::test] + async fn dead_letter_encoding_running_cancellation_holds_budget_until_closure_exits() { + let entry = test_entry(); + let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1)); + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = std::sync::mpsc::channel(); + *budget.encode_hook.lock().unwrap() = Some(Box::new(move || { + let _ = started_tx.send(()); + release_rx.recv_timeout(Duration::from_secs(2)).unwrap(); + })); + let reservation = budget.try_reserve(&entry, "failure").unwrap(); + let task = tokio::spawn(reservation.encode_owned(entry, "failure".to_string())); + tokio::time::timeout(Duration::from_secs(2), started_rx) + .await + .unwrap() + .unwrap(); + task.abort(); + assert!(matches!(task.await, Err(error) if error.is_cancelled())); + assert_eq!(budget.snapshot().active_jobs, 1); + assert!(budget.snapshot().reserved_bytes > 0); + assert!(budget.try_reserve(&test_entry(), "failure").is_err()); + release_tx.send(()).unwrap(); + tokio::time::timeout(Duration::from_secs(2), async { + while budget.snapshot().active_jobs != 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().encoded_total, 1); + } + + #[tokio::test] + async fn dead_letter_encoding_panic_releases_input_and_reservation() { + let entry = test_entry(); + let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1)); + *budget.encode_hook.lock().unwrap() = Some(Box::new(|| panic!("encoding test panic"))); + let result = budget + .try_reserve(&entry, "failure") + .unwrap() + .encode_owned(entry, "failure".to_string()) + .await; + assert!(matches!(result, Err(DataLayerError::UnexpectedValue(_)))); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().active_jobs, 0); + assert_eq!(budget.snapshot().encoded_total, 0); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn dead_letter_encoding_concurrent_reservations_have_no_hidden_waiting_queue() { + const TASKS: usize = 16; + let entry = Arc::new(test_entry()); + let size = encoding_size(&entry, "failure").unwrap(); + let budget = Arc::new(DeadLetterEncodingBudget::new(size.total * 2, 2)); + let ready = Arc::new(tokio::sync::Barrier::new(TASKS + 1)); + let release = Arc::new(tokio::sync::Barrier::new(TASKS + 1)); + let mut tasks = tokio::task::JoinSet::new(); + for _ in 0..TASKS { + let entry = Arc::clone(&entry); + let budget = Arc::clone(&budget); + let ready = Arc::clone(&ready); + let release = Arc::clone(&release); + tasks.spawn(async move { + let reservation = budget.try_reserve(&entry, "failure"); + ready.wait().await; + release.wait().await; + reservation.is_ok() + }); + } + tokio::time::timeout(Duration::from_secs(2), ready.wait()) + .await + .unwrap(); + assert_eq!(budget.snapshot().active_jobs, 2); + assert_eq!(budget.snapshot().reserved_bytes, size.total * 2); + assert_eq!( + budget.snapshot().capacity_rejected_total, + (TASKS - 2) as u64 + ); + release.wait().await; + let mut admitted = 0; + while let Some(result) = tasks.join_next().await { + admitted += usize::from(result.unwrap()); + } + assert_eq!(admitted, 2); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().active_jobs, 0); + } +} diff --git a/crates/aether-usage/runtime/src/event.rs b/crates/aether-usage/runtime/src/event.rs index 7c8f2d278..6c6e27e87 100644 --- a/crates/aether-usage/runtime/src/event.rs +++ b/crates/aether-usage/runtime/src/event.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; use aether_data_contracts::repository::usage::UsageBodyCaptureState; @@ -6,8 +7,18 @@ use aether_data_contracts::DataLayerError; use serde::{Deserialize, Serialize}; use serde_json::Value; +use crate::body_capture::mark_usage_event_capture_truncated; +pub use crate::event_capture_budget::UsageEventCaptureRetention; +use crate::event_capture_budget::{ + json_heap_estimate, shared_capture_memory_budget, EventCaptureMemoryBudget, +}; + pub const USAGE_EVENT_VERSION: u8 = 1; +#[path = "event_wire.rs"] +mod wire; +pub(crate) use wire::EncodedUsageEvent; + #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum UsageEventType { @@ -18,7 +29,7 @@ pub enum UsageEventType { Cancelled, } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)] +#[derive(Debug, PartialEq, Serialize, Deserialize, Default)] pub struct UsageEventData { #[serde(default, skip_serializing_if = "Option::is_none")] pub user_id: Option, @@ -144,6 +155,171 @@ pub struct UsageEventData { pub local_execution_runtime_miss_reason: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub request_metadata: Option, + #[doc(hidden)] + #[serde(skip)] + pub capture_retention: UsageEventCaptureRetention, +} + +impl UsageEventData { + fn capture_heap_estimate(&self) -> usize { + [ + &self.request_body, + &self.provider_request_body, + &self.response_body, + &self.client_response_body, + ] + .into_iter() + .flatten() + .fold(0usize, |bytes, body| { + bytes + .saturating_add(std::mem::size_of::()) + .saturating_add(json_heap_estimate(body)) + }) + } + + fn captured_fields(&self) -> [bool; 4] { + [ + self.request_body.is_some(), + self.provider_request_body.is_some(), + self.response_body.is_some(), + self.client_response_body.is_some(), + ] + } + + fn mark_capture_omitted(&mut self, captured: [bool; 4]) { + for (present, key, state) in [ + (captured[0], "request", &mut self.request_body_state), + ( + captured[1], + "provider_request", + &mut self.provider_request_body_state, + ), + (captured[2], "response", &mut self.response_body_state), + ( + captured[3], + "client_response", + &mut self.client_response_body_state, + ), + ] { + if present + && !matches!( + *state, + Some( + UsageBodyCaptureState::None + | UsageBodyCaptureState::Disabled + | UsageBodyCaptureState::Unavailable + ) + ) + { + *state = Some(UsageBodyCaptureState::Truncated); + mark_usage_event_capture_truncated(&mut self.request_metadata, key); + } + } + } + + pub(crate) fn apply_capture_memory_budget( + &mut self, + budget: std::sync::Arc, + ) { + let bytes = self.capture_heap_estimate(); + if self + .capture_retention + .reserve(std::sync::Arc::clone(&budget), bytes) + { + return; + } + let captured = self.captured_fields(); + self.request_body = None; + self.provider_request_body = None; + self.response_body = None; + self.client_response_body = None; + self.mark_capture_omitted(captured); + // The previous lease is released only after the owned JSON bodies are gone. + self.capture_retention.clear(budget); + } +} + +impl Clone for UsageEventData { + fn clone(&self) -> Self { + let (capture_retention, retain_bodies) = self + .capture_retention + .clone_for_bodies(|| self.capture_heap_estimate()); + // Enumerate every field so additions require an explicit ownership decision. + let mut cloned = Self { + user_id: self.user_id.clone(), + api_key_id: self.api_key_id.clone(), + username: self.username.clone(), + api_key_name: self.api_key_name.clone(), + provider_name: self.provider_name.clone(), + model: self.model.clone(), + target_model: self.target_model.clone(), + model_id: self.model_id.clone(), + global_model_id: self.global_model_id.clone(), + provider_id: self.provider_id.clone(), + provider_endpoint_id: self.provider_endpoint_id.clone(), + provider_api_key_id: self.provider_api_key_id.clone(), + request_type: self.request_type.clone(), + api_format: self.api_format.clone(), + api_family: self.api_family.clone(), + endpoint_kind: self.endpoint_kind.clone(), + endpoint_api_format: self.endpoint_api_format.clone(), + provider_api_family: self.provider_api_family.clone(), + provider_endpoint_kind: self.provider_endpoint_kind.clone(), + has_format_conversion: self.has_format_conversion, + is_stream: self.is_stream, + input_tokens: self.input_tokens, + output_tokens: self.output_tokens, + total_tokens: self.total_tokens, + cache_creation_input_tokens: self.cache_creation_input_tokens, + cache_creation_ephemeral_5m_input_tokens: self.cache_creation_ephemeral_5m_input_tokens, + cache_creation_ephemeral_1h_input_tokens: self.cache_creation_ephemeral_1h_input_tokens, + cache_read_input_tokens: self.cache_read_input_tokens, + cache_creation_cost_usd: self.cache_creation_cost_usd, + cache_read_cost_usd: self.cache_read_cost_usd, + output_price_per_1m: self.output_price_per_1m, + total_cost_usd: self.total_cost_usd, + actual_total_cost_usd: self.actual_total_cost_usd, + status_code: self.status_code, + error_message: self.error_message.clone(), + error_category: self.error_category.clone(), + response_time_ms: self.response_time_ms, + first_byte_time_ms: self.first_byte_time_ms, + request_headers: self.request_headers.clone(), + request_body: retain_bodies.then(|| self.request_body.clone()).flatten(), + request_body_ref: self.request_body_ref.clone(), + request_body_state: self.request_body_state, + provider_request_headers: self.provider_request_headers.clone(), + provider_request_body: retain_bodies + .then(|| self.provider_request_body.clone()) + .flatten(), + provider_request_body_ref: self.provider_request_body_ref.clone(), + provider_request_body_state: self.provider_request_body_state, + response_headers: self.response_headers.clone(), + response_body: retain_bodies.then(|| self.response_body.clone()).flatten(), + response_body_ref: self.response_body_ref.clone(), + response_body_state: self.response_body_state, + client_response_headers: self.client_response_headers.clone(), + client_response_body: retain_bodies + .then(|| self.client_response_body.clone()) + .flatten(), + client_response_body_ref: self.client_response_body_ref.clone(), + client_response_body_state: self.client_response_body_state, + candidate_id: self.candidate_id.clone(), + candidate_index: self.candidate_index, + key_name: self.key_name.clone(), + planner_kind: self.planner_kind.clone(), + route_family: self.route_family.clone(), + route_kind: self.route_kind.clone(), + execution_path: self.execution_path.clone(), + local_execution_runtime_miss_reason: self.local_execution_runtime_miss_reason.clone(), + request_metadata: self.request_metadata.clone(), + capture_retention, + }; + if !retain_bodies { + cloned.mark_capture_omitted(self.captured_fields()); + } + cloned + } } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] @@ -164,6 +340,16 @@ struct UsageEventEnvelope { data: UsageEventData, } +#[derive(Serialize)] +struct BorrowedUsageEventEnvelope<'a, T: ?Sized> { + v: u8, + #[serde(rename = "type")] + event_type: UsageEventType, + request_id: &'a str, + timestamp_ms: u64, + data: &'a T, +} + impl UsageEvent { pub fn new( event_type: UsageEventType, @@ -179,12 +365,12 @@ impl UsageEvent { } pub fn to_stream_fields(&self) -> Result, DataLayerError> { - let payload = UsageEventEnvelope { + let payload = BorrowedUsageEventEnvelope { v: USAGE_EVENT_VERSION, event_type: self.event_type, - request_id: self.request_id.clone(), + request_id: &self.request_id, timestamp_ms: self.timestamp_ms, - data: self.data.clone(), + data: &self.data, }; let payload = serde_json::to_string(&payload).map_err(|err| { DataLayerError::UnexpectedValue(format!( @@ -194,7 +380,21 @@ impl UsageEvent { Ok(BTreeMap::from([("payload".to_string(), payload)])) } + pub(crate) fn to_bounded_stream_fields( + &self, + max_bytes: usize, + ) -> Result { + wire::encode(self, max_bytes) + } + pub fn from_stream_fields(fields: &BTreeMap) -> Result { + Self::from_stream_fields_with_capture_budget(fields, shared_capture_memory_budget()) + } + + pub(crate) fn from_stream_fields_with_capture_budget( + fields: &BTreeMap, + budget: Arc, + ) -> Result { let payload = fields.get("payload").ok_or_else(|| { DataLayerError::UnexpectedValue( "usage event stream entry missing payload field".to_string(), @@ -212,12 +412,16 @@ impl UsageEvent { ))); } - Ok(Self { + let mut event = Self { event_type: envelope.event_type, request_id: envelope.request_id, timestamp_ms: envelope.timestamp_ms, data: envelope.data, - }) + }; + // The wire format has no ownership lease. Preserve billing facts before a decoded + // body can be omitted; the raw Redis response and serde allocation are not budgeted here. + crate::runtime::prepare_decoded_event_capture_memory(&mut event, budget); + Ok(event) } } @@ -230,8 +434,564 @@ pub fn now_ms() -> u64 { #[cfg(test)] mod tests { + use std::collections::BTreeMap; + use std::sync::Arc; + + use aether_data_contracts::repository::usage::UsageBodyCaptureState; + use aether_data_contracts::DataLayerError; + use serde_json::json; + + use crate::event_capture_budget::EventCaptureMemoryBudget; + use crate::{ + apply_usage_body_capture_policy_to_event, build_upsert_usage_record_from_event, + UsageBodyCapturePolicy, + }; + use super::{UsageEvent, UsageEventData, UsageEventType}; + fn captured_event() -> UsageEvent { + UsageEvent { + event_type: UsageEventType::Failed, + request_id: "capture-budget-request".to_string(), + timestamp_ms: 123_456, + data: UsageEventData { + provider_name: "provider".to_string(), + model: "model".to_string(), + input_tokens: Some(100), + output_tokens: Some(500), + total_tokens: Some(600), + cache_read_input_tokens: Some(0), + cache_creation_input_tokens: Some(25), + actual_total_cost_usd: Some(1.25), + status_code: Some(502), + error_category: Some("upstream_error".to_string()), + error_message: Some("upstream failed".to_string()), + request_body: Some(json!({"messages": [{"content": "request"}]})), + provider_request_body: Some(json!({"input": "upstream request"})), + response_body: Some(json!({"usage": {"input_tokens": 100, "output_tokens": 500}})), + client_response_body: Some(json!({"error": "client response"})), + request_body_state: Some(UsageBodyCaptureState::Inline), + provider_request_body_state: Some(UsageBodyCaptureState::Inline), + response_body_state: Some(UsageBodyCaptureState::Inline), + client_response_body_state: Some(UsageBodyCaptureState::Inline), + request_metadata: Some(json!({ + "requested_reasoning_effort": "high", + "provider_reasoning_effort": "medium", + "provider_service_tier": "priority", + "provider_actual_service_tier": "default", + "provider_cache_ttl_minutes": 60, + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000", + "body_capture": {"response": {"state": "inline", "source_bytes": 1000}} + })), + ..UsageEventData::default() + }, + } + } + + #[test] + fn event_capture_budget_zero_preserves_billing_refs_and_database_truncation() { + let budget = Arc::new(EventCaptureMemoryBudget::new(0)); + let mut event = captured_event(); + event.data.request_body_ref = + Some("usage://capture-budget-request/request_body".to_string()); + event.data.response_body_ref = + Some("usage://capture-budget-request/response_body".to_string()); + event.data.apply_capture_memory_budget(Arc::clone(&budget)); + assert!(event.data.request_body.is_none()); + assert!(event.data.provider_request_body.is_none()); + assert!(event.data.response_body.is_none()); + assert!(event.data.client_response_body.is_none()); + let capture_metadata = event + .data + .request_metadata + .as_ref() + .expect("capture metadata"); + assert_eq!( + capture_metadata["body_capture"]["response"]["source_bytes"], + 1000 + ); + assert_eq!( + capture_metadata["body_capture"]["response"]["stored_bytes"], + 0 + ); + assert_eq!( + capture_metadata["body_capture"]["response"]["reason"], + "usage_event_memory_budget_exceeded" + ); + let record = build_upsert_usage_record_from_event(&event).expect("record mapping"); + assert_eq!(record.status, "failed"); + assert_eq!(record.input_tokens, Some(100)); + assert_eq!(record.output_tokens, Some(500)); + assert_eq!(record.cache_read_input_tokens, Some(0)); + assert_eq!(record.cache_creation_input_tokens, Some(25)); + assert_eq!(record.actual_total_cost_usd, Some(1.25)); + assert_eq!(record.error_category.as_deref(), Some("upstream_error")); + assert_eq!(record.request_body_ref, event.data.request_body_ref); + assert_eq!(record.response_body_ref, event.data.response_body_ref); + for state in [ + record.request_body_state, + record.provider_request_body_state, + record.response_body_state, + record.client_response_body_state, + ] { + assert_eq!(state, Some(UsageBodyCaptureState::Truncated)); + } + let metadata = record.request_metadata.expect("preserved metadata"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 60); + assert_eq!( + metadata["plan_usage_reservation_token"], + "550e8400-e29b-41d4-a716-446655440000" + ); + // Persistence projects billing metadata; capture state remains in typed columns. + assert!(metadata.get("body_capture").is_none()); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 1); + } + + #[test] + fn event_capture_budget_clone_reserves_each_copy_before_cloning_bodies() { + let mut event = captured_event(); + let weight = event.data.capture_heap_estimate(); + let budget = Arc::new(EventCaptureMemoryBudget::new(weight * 2)); + event.data.apply_capture_memory_budget(Arc::clone(&budget)); + let copy = event.clone(); + assert_eq!(copy, event); + assert_eq!(budget.retained_bytes(), weight * 2); + let downgraded = event.clone(); + assert!(event.data.response_body.is_some()); + assert!(copy.data.response_body.is_some()); + assert!(downgraded.data.response_body.is_none()); + assert_eq!( + downgraded.data.response_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + assert_eq!(downgraded.data.total_tokens, event.data.total_tokens); + assert_eq!(downgraded.timestamp_ms, event.timestamp_ms); + assert_eq!(downgraded.event_type, event.event_type); + assert_eq!(budget.retained_bytes(), weight * 2); + drop(copy); + assert_eq!(budget.retained_bytes(), weight); + drop((event, downgraded)); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn event_capture_budget_serialization_borrows_bodies_without_charging_a_clone() { + let mut event = captured_event(); + let weight = event.data.capture_heap_estimate(); + let budget = Arc::new(EventCaptureMemoryBudget::new(weight)); + crate::runtime::prepare_event_capture_memory(&mut event, Arc::clone(&budget)); + let decoded_budget = Arc::new(EventCaptureMemoryBudget::new(usize::MAX)); + for _ in 0..3 { + let fields = event.to_stream_fields().expect("wire serialization"); + assert!(!fields["payload"].contains("capture_retention")); + let decoded = UsageEvent::from_stream_fields_with_capture_budget( + &fields, + Arc::clone(&decoded_budget), + ) + .expect("wire decode"); + assert_eq!(decoded, event); + assert_eq!(budget.retained_bytes(), weight); + assert!(decoded_budget.retained_bytes() > 0); + drop(decoded); + assert_eq!(decoded_budget.retained_bytes(), 0); + } + assert_eq!(budget.downgraded_total(), 0); + drop(event); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn from_stream_fields_legacy_body_budget_preserves_billing_and_request_facts() { + let mut event = captured_event(); + event.data.model = "gpt-5.6-sol".to_string(); + event.data.endpoint_api_format = Some("openai:responses".to_string()); + event.data.request_body = Some(json!({"reasoning": {"effort": "high"}})); + event.data.provider_request_body = Some(json!({ + "model": "gpt-5.6-sol", "reasoning": {"effort": "medium"}, + "service_tier": "priority" + })); + event.data.response_body = Some(json!({"service_tier": "Default"})); + event.data.request_body_state = None; + event.data.provider_request_body_state = None; + event.data.response_body_state = None; + event.data.client_response_body_state = None; + event.data.cache_creation_ephemeral_5m_input_tokens = Some(0); + event.data.cache_creation_ephemeral_1h_input_tokens = Some(25); + event.data.cache_read_cost_usd = Some(0.0); + event.data.request_body_ref = Some("usage://legacy/request".to_string()); + event.data.request_metadata = Some(json!({ + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000" + })); + let fields = event.to_stream_fields().expect("legacy wire serialization"); + assert!(!fields["payload"].contains("request_body_state")); + let budget = Arc::new(EventCaptureMemoryBudget::new(0)); + let decoded = + UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget)) + .expect("legacy wire decode"); + + assert_eq!(decoded.event_type, UsageEventType::Failed); + assert_eq!(decoded.request_id, event.request_id); + assert_eq!(decoded.timestamp_ms, event.timestamp_ms); + assert_eq!(decoded.data.input_tokens, Some(100)); + assert_eq!(decoded.data.output_tokens, Some(500)); + assert_eq!(decoded.data.total_tokens, Some(600)); + assert_eq!(decoded.data.cache_creation_input_tokens, Some(25)); + assert_eq!( + decoded.data.cache_creation_ephemeral_5m_input_tokens, + Some(0) + ); + assert_eq!( + decoded.data.cache_creation_ephemeral_1h_input_tokens, + Some(25) + ); + assert_eq!(decoded.data.cache_read_input_tokens, Some(0)); + assert_eq!(decoded.data.cache_read_cost_usd, Some(0.0)); + assert_eq!(decoded.data.actual_total_cost_usd, Some(1.25)); + assert_eq!(decoded.data.status_code, Some(502)); + assert_eq!( + decoded.data.error_category.as_deref(), + Some("upstream_error") + ); + assert_eq!(decoded.data.request_body_ref, event.data.request_body_ref); + assert!(decoded.data.request_body.is_none()); + assert!(decoded.data.provider_request_body.is_none()); + assert!(decoded.data.response_body.is_none()); + assert!(decoded.data.client_response_body.is_none()); + for state in [ + decoded.data.request_body_state, + decoded.data.provider_request_body_state, + decoded.data.response_body_state, + decoded.data.client_response_body_state, + ] { + assert_eq!(state, Some(UsageBodyCaptureState::Truncated)); + } + let metadata = decoded + .data + .request_metadata + .as_ref() + .expect("preserved facts"); + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_reasoning_effort"], "medium"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 30); + assert_eq!( + metadata["plan_usage_reservation_token"], + "550e8400-e29b-41d4-a716-446655440000" + ); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 1); + } + + #[test] + fn from_stream_fields_reconstructed_lease_also_bounds_recorder_clones() { + let fields = captured_event() + .to_stream_fields() + .expect("wire serialization"); + let probe_budget = Arc::new(EventCaptureMemoryBudget::new(usize::MAX)); + let probe = + UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&probe_budget)) + .expect("estimate decoded allocation"); + let weight = probe_budget.retained_bytes(); + assert!(weight > 0); + drop(probe); + assert_eq!(probe_budget.retained_bytes(), 0); + + let budget = Arc::new(EventCaptureMemoryBudget::new(weight * 2)); + let event = + UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget)) + .expect("wire decode"); + let recorder_copy = event.clone(); + assert!(recorder_copy.data.response_body.is_some()); + assert_eq!(budget.retained_bytes(), weight * 2); + let omitted_copy = event.clone(); + assert!(omitted_copy.data.response_body.is_none()); + assert_eq!(omitted_copy.data.total_tokens, Some(600)); + assert_eq!(omitted_copy.data.cache_read_input_tokens, Some(0)); + assert_eq!(budget.retained_bytes(), weight * 2); + drop(event); + assert_eq!(budget.retained_bytes(), weight); + drop((recorder_copy, omitted_copy)); + assert_eq!(budget.retained_bytes(), 0); + } + + fn assert_typed_body_clear_is_preserved(event: &UsageEvent, state: UsageBodyCaptureState) { + for (body, reference, actual_state) in [ + ( + &event.data.request_body, + &event.data.request_body_ref, + event.data.request_body_state, + ), + ( + &event.data.provider_request_body, + &event.data.provider_request_body_ref, + event.data.provider_request_body_state, + ), + ( + &event.data.response_body, + &event.data.response_body_ref, + event.data.response_body_state, + ), + ( + &event.data.client_response_body, + &event.data.client_response_body_ref, + event.data.client_response_body_state, + ), + ] { + assert!(body.is_none()); + assert_eq!(actual_state, Some(state)); + assert_eq!(reference.as_deref(), Some("usage://stale/reference")); + } + assert!(event + .data + .request_metadata + .as_ref() + .and_then(|value| value.get("body_capture")) + .is_none()); + } + + #[test] + fn from_stream_fields_and_clone_budget_preserve_typed_clear_with_residual_bodies() { + for state in [ + UsageBodyCaptureState::None, + UsageBodyCaptureState::Disabled, + UsageBodyCaptureState::Unavailable, + ] { + let mut source = captured_event(); + source.data.request_metadata = None; + source.data.request_body_state = Some(state); + source.data.provider_request_body_state = Some(state); + source.data.response_body_state = Some(state); + source.data.client_response_body_state = Some(state); + source.data.request_body_ref = Some("usage://stale/reference".to_string()); + source.data.provider_request_body_ref = Some("usage://stale/reference".to_string()); + source.data.response_body_ref = Some("usage://stale/reference".to_string()); + source.data.client_response_body_ref = Some("usage://stale/reference".to_string()); + let fields = source.to_stream_fields().expect("wire serialization"); + let decoded_budget = Arc::new(EventCaptureMemoryBudget::new(0)); + let decoded = UsageEvent::from_stream_fields_with_capture_budget( + &fields, + Arc::clone(&decoded_budget), + ) + .expect("wire decode"); + assert_typed_body_clear_is_preserved(&decoded, state); + assert_eq!(decoded_budget.retained_bytes(), 0); + + let clone_budget = Arc::new(EventCaptureMemoryBudget::new( + source.data.capture_heap_estimate(), + )); + crate::runtime::prepare_event_capture_memory(&mut source, Arc::clone(&clone_budget)); + let cloned = source.clone(); + assert_typed_body_clear_is_preserved(&cloned, state); + assert!(source.data.request_body.is_some()); + assert_eq!(clone_budget.downgraded_total(), 1); + drop(source); + assert_eq!(clone_budget.retained_bytes(), 0); + } + } + + #[test] + fn from_stream_fields_legacy_metadata_only_facts_survive_before_billing() { + for limit in [0, 8192] { + for include_response in [false, true] { + let mut source = legacy_metadata_only_event(); + if include_response { + source.data.response_body = Some(json!({"result": "response capture"})); + } + let fields = source + .to_stream_fields() + .expect("legacy wire serialization"); + let budget = Arc::new(EventCaptureMemoryBudget::new(limit)); + let decoded = UsageEvent::from_stream_fields_with_capture_budget( + &fields, + Arc::clone(&budget), + ) + .expect("legacy wire decode"); + // The worker enriches this clone before DTO conversion. Missing legacy bodies + // must not erase a previously derived TTL or turn an explicit zero into unknown. + let billing_event = decoded.clone(); + let metadata = billing_event + .data + .request_metadata + .as_ref() + .expect("legacy facts"); + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_reasoning_effort"], "medium"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 60); + assert_eq!(billing_event.data.input_tokens, Some(0)); + assert_eq!(billing_event.data.output_tokens, Some(0)); + assert_eq!(billing_event.data.total_tokens, Some(0)); + assert_eq!(billing_event.data.cache_read_input_tokens, Some(0)); + assert_eq!(billing_event.data.cache_creation_input_tokens, Some(0)); + assert_eq!(billing_event.data.actual_total_cost_usd, Some(0.0)); + assert_eq!(billing_event.data.request_body_state, None); + assert_eq!(billing_event.data.provider_request_body_state, None); + drop((billing_event, decoded)); + assert_eq!(budget.retained_bytes(), 0); + } + } + } + + #[test] + fn from_stream_fields_typed_none_still_clears_metadata_only_request_facts() { + for limit in [0, 8192] { + let mut source = legacy_metadata_only_event(); + source.data.request_body_state = Some(UsageBodyCaptureState::None); + source.data.provider_request_body_state = Some(UsageBodyCaptureState::None); + let fields = source.to_stream_fields().expect("wire serialization"); + let budget = Arc::new(EventCaptureMemoryBudget::new(limit)); + let decoded = + UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget)) + .expect("wire decode"); + let metadata = decoded + .data + .request_metadata + .as_ref() + .expect("response facts remain"); + for key in [ + "requested_reasoning_effort", + "provider_reasoning_effort", + "provider_service_tier", + "provider_cache_ttl_minutes", + ] { + assert!(metadata.get(key).is_none(), "typed none must clear {key}"); + } + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!( + decoded.data.request_body_state, + Some(UsageBodyCaptureState::None) + ); + assert_eq!( + decoded.data.provider_request_body_state, + Some(UsageBodyCaptureState::None) + ); + assert_eq!(decoded.data.cache_read_input_tokens, Some(0)); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 0); + } + } + + fn legacy_metadata_only_event() -> UsageEvent { + UsageEvent::new( + UsageEventType::Completed, + "legacy-metadata-only", + UsageEventData { + provider_name: "openai".to_string(), + model: "gpt-5.6-sol".to_string(), + endpoint_api_format: Some("openai:responses".to_string()), + input_tokens: Some(0), + output_tokens: Some(0), + total_tokens: Some(0), + cache_read_input_tokens: Some(0), + cache_creation_input_tokens: Some(0), + actual_total_cost_usd: Some(0.0), + request_metadata: Some(json!({ + "requested_reasoning_effort": "high", + "provider_reasoning_effort": "medium", + "provider_service_tier": "priority", + "provider_actual_service_tier": "default", + "provider_cache_ttl_minutes": 60 + })), + ..UsageEventData::default() + }, + ) + } + + #[test] + fn from_stream_fields_body_omission_keeps_unknown_usage_unknown() { + let budget = Arc::new(EventCaptureMemoryBudget::new(0)); + for event_type in [ + UsageEventType::Completed, + UsageEventType::Failed, + UsageEventType::Cancelled, + ] { + let event = UsageEvent::new( + event_type, + "usage-unavailable", + UsageEventData { + provider_name: "openai".to_string(), + model: "gpt-5".to_string(), + response_body: Some(json!({"error": "usage unavailable"})), + request_metadata: Some(json!({ + "usage_available": false, + "usage_pricing_available": false + })), + ..UsageEventData::default() + }, + ); + let fields = event.to_stream_fields().expect("wire serialization"); + let decoded = + UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget)) + .expect("wire decode"); + assert_eq!(decoded.event_type, event_type); + assert_eq!(decoded.data.input_tokens, None); + assert_eq!(decoded.data.output_tokens, None); + assert_eq!(decoded.data.total_tokens, None); + assert_eq!(decoded.data.cache_read_input_tokens, None); + assert_eq!(decoded.data.cache_creation_input_tokens, None); + assert_eq!(decoded.data.actual_total_cost_usd, None); + assert_eq!( + decoded.data.request_metadata.as_ref().expect("metadata")["usage_available"], + false + ); + assert_eq!( + decoded.data.request_metadata.as_ref().expect("metadata") + ["usage_pricing_available"], + false + ); + assert_eq!( + decoded.data.response_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + } + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 3); + } + + #[test] + fn from_stream_fields_invalid_envelopes_do_not_reserve_capture_memory() { + let budget = Arc::new(EventCaptureMemoryBudget::new(1024)); + let mut unsupported = captured_event() + .to_stream_fields() + .expect("wire serialization"); + let mut payload: serde_json::Value = + serde_json::from_str(&unsupported["payload"]).expect("json"); + payload["v"] = json!(99); + unsupported.insert("payload".to_string(), payload.to_string()); + for fields in [ + BTreeMap::new(), + BTreeMap::from([("payload".to_string(), "not json".to_string())]), + unsupported, + ] { + assert!(matches!( + UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget)), + Err(DataLayerError::UnexpectedValue(_)) + )); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 0); + } + } + + #[test] + fn event_capture_budget_basic_policy_needs_no_diagnostic_allocation() { + let mut event = captured_event(); + let budget = Arc::new(EventCaptureMemoryBudget::new(0)); + apply_usage_body_capture_policy_to_event(UsageBodyCapturePolicy::default(), &mut event); + event.data.apply_capture_memory_budget(Arc::clone(&budget)); + assert_eq!( + event.data.response_body_state, + Some(UsageBodyCaptureState::Disabled) + ); + assert_eq!(event.data.total_tokens, Some(600)); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 0); + } + #[test] fn usage_event_round_trips_through_stream_fields() { let event = UsageEvent::new( diff --git a/crates/aether-usage/runtime/src/event_capture_budget.rs b/crates/aether-usage/runtime/src/event_capture_budget.rs new file mode 100644 index 000000000..054168e7c --- /dev/null +++ b/crates/aether-usage/runtime/src/event_capture_budget.rs @@ -0,0 +1,102 @@ +use std::sync::{Arc, LazyLock}; + +#[cfg(test)] +use serde_json::Value; + +#[doc(hidden)] +pub use aether_data_contracts::repository::usage::UsageCaptureRetention as UsageEventCaptureRetention; +pub(crate) use aether_data_contracts::repository::usage::{ + usage_json_heap_estimate as json_heap_estimate, + UsageCaptureMemoryBudget as EventCaptureMemoryBudget, +}; + +const DEFAULT_CAPTURE_MEMORY_BUDGET_BYTES: usize = 128 * 1024 * 1024; + +static CAPTURE_MEMORY_BUDGET: LazyLock> = LazyLock::new(|| { + let limit = std::env::var("AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES") + .ok() + .and_then(|value| value.trim().parse().ok()) + .unwrap_or(DEFAULT_CAPTURE_MEMORY_BUDGET_BYTES); + Arc::new(EventCaptureMemoryBudget::new(limit)) +}); + +pub(crate) fn shared_capture_memory_budget() -> Arc { + Arc::clone(&CAPTURE_MEMORY_BUDGET) +} + +pub(crate) fn capture_memory_metrics() -> (usize, usize, u64) { + CAPTURE_MEMORY_BUDGET.snapshot() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn event_capture_budget_resize_and_drop_release_estimate() { + let budget = Arc::new(EventCaptureMemoryBudget::new(16)); + let mut retention = UsageEventCaptureRetention::default(); + assert!(retention.reserve(Arc::clone(&budget), 12)); + assert!(!retention.reserve(Arc::clone(&budget), 17)); + assert_eq!(budget.retained_bytes(), 12); + assert_eq!(budget.downgraded_total(), 1); + assert!(retention.reserve(Arc::clone(&budget), 4)); + assert_eq!(budget.retained_bytes(), 4); + drop(retention); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn event_capture_budget_unmanaged_clone_skips_estimation_and_empty_clone_is_free() { + let unmanaged = UsageEventCaptureRetention::default(); + let (_, retained) = + unmanaged.clone_for_bodies(|| panic!("unmanaged JSON must not be scanned")); + assert!(retained); + let budget = Arc::new(EventCaptureMemoryBudget::new(0)); + let mut managed = UsageEventCaptureRetention::default(); + assert!(managed.reserve(Arc::clone(&budget), 0)); + let (cloned, retained) = managed.clone_for_bodies(|| 0); + assert!(retained); + drop((managed, cloned)); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 0); + } + + #[test] + fn event_capture_budget_estimate_counts_string_and_array_spare_capacity() { + let mut text = String::with_capacity(1024); + text.push('x'); + let mut array = Vec::with_capacity(16); + let expected = text.capacity() + array.capacity() * std::mem::size_of::(); + array.push(Value::String(text)); + assert_eq!(json_heap_estimate(&Value::Array(array)), expected); + } + + #[test] + fn event_capture_budget_parallel_owners_never_exceed_shared_limit() { + let budget = Arc::new(EventCaptureMemoryBudget::new(1024)); + let barrier = Arc::new(std::sync::Barrier::new(8)); + std::thread::scope(|scope| { + for _ in 0..8 { + let budget = Arc::clone(&budget); + let barrier = Arc::clone(&barrier); + scope.spawn(move || { + for _ in 0..100 { + let mut retained = UsageEventCaptureRetention::default(); + barrier.wait(); + let _ = retained.reserve(Arc::clone(&budget), 400); + barrier.wait(); + assert_eq!(budget.retained_bytes(), 800); + barrier.wait(); + drop(retained); + barrier.wait(); + assert_eq!(budget.retained_bytes(), 0); + barrier.wait(); + } + }); + } + }); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(budget.downgraded_total(), 600); + } +} diff --git a/crates/aether-usage/runtime/src/event_wire.rs b/crates/aether-usage/runtime/src/event_wire.rs new file mode 100644 index 000000000..477603c0a --- /dev/null +++ b/crates/aether-usage/runtime/src/event_wire.rs @@ -0,0 +1,999 @@ +use std::collections::BTreeMap; +use std::io::{self, Write}; + +use aether_data_contracts::repository::usage::{ + resolve_provider_cache_ttl_minutes, UsageBodyCaptureState, + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY, + PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY, +}; +use aether_data_contracts::DataLayerError; +use serde::ser::{Impossible, SerializeMap, SerializeStruct}; +use serde::{Serialize, Serializer}; +use serde_json::Value; + +use super::{BorrowedUsageEventEnvelope, UsageEvent, UsageEventData, USAGE_EVENT_VERSION}; +use crate::body_capture::mark_usage_event_capture_truncated; +use crate::request_metadata::{ + attach_client_request_body_metadata, attach_provider_request_body_metadata, + attach_provider_response_body_metadata, clear_client_request_body_metadata, + clear_provider_request_body_metadata, request_body_derived_facts_action, + RequestBodyDerivedFactsAction, +}; + +const DIAGNOSTIC_FIELDS: [&str; 8] = [ + "request_body", + "provider_request_body", + "response_body", + "client_response_body", + "request_headers", + "provider_request_headers", + "response_headers", + "client_response_headers", +]; +const BODY_STATE_FIELDS: [&str; 4] = [ + "request_body_state", + "provider_request_body_state", + "response_body_state", + "client_response_body_state", +]; +const BODY_METADATA_KEYS: [&str; 4] = + ["request", "provider_request", "response", "client_response"]; + +#[derive(Debug)] +pub(crate) struct EncodedUsageEvent { + pub(crate) fields: BTreeMap, + pub(crate) diagnostics_omitted: bool, +} + +pub(super) fn encode( + event: &UsageEvent, + max_bytes: usize, +) -> Result { + let mut writer = BoundedJsonWriter::new(max_bytes); + if writer.serialize(&envelope(event, &event.data))? { + return writer.into_event(false); + } + + // Conservatively reject oversized original metadata before cloning it, even + // when later fact normalization could make that metadata smaller. + let core = ProjectedData { + data: &event.data, + overrides: None, + }; + if !writer.serialize(&envelope(event, &core))? { + return Err(wire_limit_error(max_bytes)); + } + + let overrides = WireOverrides::new(&event.data)?; + let projected = ProjectedData { + data: &event.data, + overrides: Some(&overrides), + }; + if !writer.serialize(&envelope(event, &projected))? { + return Err(wire_limit_error(max_bytes)); + } + writer.into_event(true) +} + +fn envelope<'a, T: Serialize + ?Sized>( + event: &'a UsageEvent, + data: &'a T, +) -> BorrowedUsageEventEnvelope<'a, T> { + BorrowedUsageEventEnvelope { + v: USAGE_EVENT_VERSION, + event_type: event.event_type, + request_id: &event.request_id, + timestamp_ms: event.timestamp_ms, + data, + } +} + +fn wire_limit_error(max_bytes: usize) -> DataLayerError { + DataLayerError::InvalidInput(format!( + "usage event exceeds the {max_bytes}-byte wire limit after omitting diagnostic bodies and headers" + )) +} + +struct BoundedJsonWriter { + bytes: Vec, + max_bytes: usize, + exceeded: bool, +} + +impl BoundedJsonWriter { + fn new(max_bytes: usize) -> Self { + Self { + bytes: Vec::new(), + max_bytes, + exceeded: false, + } + } + + fn serialize(&mut self, value: &T) -> Result { + self.bytes.clear(); + self.exceeded = false; + match serde_json::to_writer(&mut *self, value) { + Ok(()) => Ok(true), + Err(_) if self.exceeded => Ok(false), + Err(error) => Err(DataLayerError::UnexpectedValue(format!( + "failed to serialize usage event payload: {error}" + ))), + } + } + + fn into_event(self, diagnostics_omitted: bool) -> Result { + let payload = String::from_utf8(self.bytes).map_err(|error| { + DataLayerError::UnexpectedValue(format!("usage event JSON was not UTF-8: {error}")) + })?; + Ok(EncodedUsageEvent { + fields: BTreeMap::from([("payload".to_string(), payload)]), + diagnostics_omitted, + }) + } +} + +impl Write for BoundedJsonWriter { + fn write(&mut self, bytes: &[u8]) -> io::Result { + if bytes.len() > self.max_bytes.saturating_sub(self.bytes.len()) { + self.exceeded = true; + return Err(io::Error::other("usage event wire limit exceeded")); + } + let required = self.bytes.len() + bytes.len(); + if required > self.bytes.capacity() { + let capacity = required + .max(self.bytes.capacity().saturating_mul(2)) + .min(self.max_bytes); + self.bytes + .try_reserve_exact(capacity - self.bytes.len()) + .map_err(|error| { + io::Error::other(format!("usage event wire allocation failed: {error}")) + })?; + } + self.bytes.extend_from_slice(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +struct WireOverrides { + truncated: [bool; 4], + metadata: Option, +} + +impl WireOverrides { + fn new(data: &UsageEventData) -> Result { + // The full v1 consumer decodes a JSON null Option as no body. + let request_body = data.request_body.as_ref().filter(|body| !body.is_null()); + let provider_request_body = data + .provider_request_body + .as_ref() + .filter(|body| !body.is_null()); + let body_cache_ttl = resolve_provider_cache_ttl_minutes( + data.endpoint_api_format + .as_deref() + .or(data.api_format.as_deref()), + data.target_model.as_deref().or(Some(data.model.as_str())), + Some(data.model.as_str()), + provider_request_body, + ); + if body_cache_ttl.is_some() + && data.provider_request_body_state == Some(UsageBodyCaptureState::None) + { + return Err(DataLayerError::InvalidInput( + "usage event cannot omit a provider request body whose cache TTL would be cleared by its explicit none capture state" + .to_string(), + )); + } + let mut metadata = data.request_metadata.clone(); + match request_body_derived_facts_action(request_body, data.request_body_state) { + RequestBodyDerivedFactsAction::Refresh => { + if request_body.is_some_and(|body| !body.is_object()) { + if let Some(Value::Object(object)) = metadata.as_mut() { + object.remove(REQUESTED_REASONING_EFFORT_METADATA_KEY); + } + } else { + metadata = attach_client_request_body_metadata(metadata, request_body); + } + } + RequestBodyDerivedFactsAction::Clear + if request_body.is_some() || data.request_body_state.is_some() => + { + metadata = clear_client_request_body_metadata(metadata); + } + RequestBodyDerivedFactsAction::Clear | RequestBodyDerivedFactsAction::Preserve => {} + } + match request_body_derived_facts_action( + provider_request_body, + data.provider_request_body_state, + ) { + RequestBodyDerivedFactsAction::Refresh => { + if provider_request_body.is_some_and(|body| !body.is_object()) { + // An authoritative scalar/array has no tier or reasoning, + // but billing still falls back to the metadata's cache TTL. + if let Some(Value::Object(object)) = metadata.as_mut() { + object.remove(PROVIDER_REASONING_EFFORT_METADATA_KEY); + object.remove(PROVIDER_SERVICE_TIER_METADATA_KEY); + } + } else { + metadata = attach_provider_request_body_metadata( + metadata, + data.endpoint_api_format + .as_deref() + .or(data.api_format.as_deref()), + data.target_model.as_deref().or(Some(data.model.as_str())), + Some(data.model.as_str()), + provider_request_body, + ); + } + } + RequestBodyDerivedFactsAction::Clear + if provider_request_body.is_some() + || data.provider_request_body_state.is_some() => + { + metadata = clear_provider_request_body_metadata(metadata); + } + RequestBodyDerivedFactsAction::Clear | RequestBodyDerivedFactsAction::Preserve => {} + } + metadata = attach_provider_response_body_metadata(metadata, data.response_body.as_ref()); + // Billing reads raw-body TTL before metadata regardless of capture state. + // Preserve that precedence independently of reasoning and tier authority. + if let Some(cache_ttl) = body_cache_ttl { + let object = metadata + .get_or_insert_with(|| Value::Object(serde_json::Map::new())) + .as_object_mut() + .ok_or_else(|| { + DataLayerError::InvalidInput( + "usage event cannot preserve provider cache TTL in non-object metadata after omitting diagnostic bodies" + .to_string(), + ) + })?; + object.insert( + PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY.to_string(), + Value::Number(cache_ttl.into()), + ); + } + let truncated = [ + (data.request_body.as_ref(), data.request_body_state), + ( + data.provider_request_body.as_ref(), + data.provider_request_body_state, + ), + (data.response_body.as_ref(), data.response_body_state), + ( + data.client_response_body.as_ref(), + data.client_response_body_state, + ), + ] + .map(|(body, state)| { + body.is_some_and(|body| !body.is_null()) + && !matches!( + state, + Some( + UsageBodyCaptureState::None + | UsageBodyCaptureState::Disabled + | UsageBodyCaptureState::Unavailable + ) + ) + }); + for (truncated, key) in truncated.into_iter().zip(BODY_METADATA_KEYS) { + if !truncated { + continue; + } + mark_usage_event_capture_truncated(&mut metadata, key); + if let Some(entry) = metadata + .as_mut() + .and_then(|value| value.get_mut("body_capture")) + .and_then(|value| value.get_mut(key)) + .and_then(Value::as_object_mut) + { + entry.insert( + "reason".to_string(), + Value::String("wire_limit_exceeded".to_string()), + ); + } + } + Ok(Self { + truncated, + metadata, + }) + } +} + +struct ProjectedData<'a> { + data: &'a UsageEventData, + overrides: Option<&'a WireOverrides>, +} + +impl Serialize for ProjectedData<'_> { + fn serialize(&self, serializer: S) -> Result { + self.data.serialize(FieldProjectionSerializer { + map: serializer.serialize_map(None)?, + overrides: self.overrides, + }) + } +} + +// Reuse UsageEventData's derived field traversal, including future fields and +// skip_serializing_if rules. Only diagnostic fields and explicit overrides differ. +struct FieldProjectionSerializer<'a, M> { + map: M, + overrides: Option<&'a WireOverrides>, +} + +impl SerializeStruct for FieldProjectionSerializer<'_, M> { + type Ok = M::Ok; + type Error = M::Error; + + fn serialize_field( + &mut self, + key: &'static str, + value: &T, + ) -> Result<(), Self::Error> { + if DIAGNOSTIC_FIELDS.contains(&key) { + return Ok(()); + } + if let Some(overrides) = self.overrides { + if key == "request_metadata" + || BODY_STATE_FIELDS + .iter() + .zip(overrides.truncated) + .any(|(state, truncated)| *state == key && truncated) + { + return Ok(()); + } + } + self.map.serialize_entry(key, value) + } + + fn end(mut self) -> Result { + if let Some(overrides) = self.overrides { + for (key, truncated) in BODY_STATE_FIELDS.into_iter().zip(overrides.truncated) { + if truncated { + self.map + .serialize_entry(key, &UsageBodyCaptureState::Truncated)?; + } + } + if let Some(metadata) = overrides.metadata.as_ref() { + self.map.serialize_entry("request_metadata", metadata)?; + } + } + self.map.end() + } +} + +fn expected_struct() -> Result { + Err(E::custom("usage event data must serialize as a struct")) +} + +macro_rules! reject_scalar_serialization { + ($($name:ident($value:ident: $ty:ty)),* $(,)?) => { + $(fn $name(self, $value: $ty) -> Result { + let _ = $value; + expected_struct() + })* + }; +} + +impl Serializer for FieldProjectionSerializer<'_, M> { + type Ok = M::Ok; + type Error = M::Error; + type SerializeSeq = Impossible; + type SerializeTuple = Impossible; + type SerializeTupleStruct = Impossible; + type SerializeTupleVariant = Impossible; + type SerializeMap = Impossible; + type SerializeStruct = Self; + type SerializeStructVariant = Impossible; + + reject_scalar_serialization! { + serialize_bool(value: bool), serialize_i8(value: i8), serialize_i16(value: i16), + serialize_i32(value: i32), serialize_i64(value: i64), serialize_i128(value: i128), + serialize_u8(value: u8), serialize_u16(value: u16), serialize_u32(value: u32), + serialize_u64(value: u64), serialize_u128(value: u128), serialize_f32(value: f32), + serialize_f64(value: f64), serialize_char(value: char), serialize_str(value: &str), + serialize_bytes(value: &[u8]), + } + + fn serialize_none(self) -> Result { + expected_struct() + } + fn serialize_some(self, _: &T) -> Result { + expected_struct() + } + fn serialize_unit(self) -> Result { + expected_struct() + } + fn serialize_unit_struct(self, _: &'static str) -> Result { + expected_struct() + } + fn serialize_unit_variant( + self, + _: &'static str, + _: u32, + _: &'static str, + ) -> Result { + expected_struct() + } + fn serialize_newtype_struct( + self, + _: &'static str, + _: &T, + ) -> Result { + expected_struct() + } + fn serialize_newtype_variant( + self, + _: &'static str, + _: u32, + _: &'static str, + _: &T, + ) -> Result { + expected_struct() + } + fn serialize_seq(self, _: Option) -> Result { + expected_struct() + } + fn serialize_tuple(self, _: usize) -> Result { + expected_struct() + } + fn serialize_tuple_struct( + self, + _: &'static str, + _: usize, + ) -> Result { + expected_struct() + } + fn serialize_tuple_variant( + self, + _: &'static str, + _: u32, + _: &'static str, + _: usize, + ) -> Result { + expected_struct() + } + fn serialize_map(self, _: Option) -> Result { + expected_struct() + } + fn serialize_struct( + self, + _: &'static str, + _: usize, + ) -> Result { + Ok(self) + } + fn serialize_struct_variant( + self, + _: &'static str, + _: u32, + _: &'static str, + _: usize, + ) -> Result { + expected_struct() + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use serde_json::json; + + use super::*; + use crate::event::UsageEventType; + use crate::event_capture_budget::EventCaptureMemoryBudget; + + fn event() -> UsageEvent { + UsageEvent { + event_type: UsageEventType::Completed, + request_id: "wire-request".to_string(), + timestamp_ms: 1_234_567, + data: UsageEventData { + provider_name: "provider".to_string(), + model: "gpt-5.6-sol".to_string(), + endpoint_api_format: Some("openai:responses".to_string()), + user_id: Some("user-id".to_string()), + api_key_id: Some("key-id".to_string()), + provider_id: Some("provider-id".to_string()), + provider_endpoint_id: Some("endpoint-id".to_string()), + provider_api_key_id: Some("provider-key-id".to_string()), + input_tokens: Some(100), + output_tokens: Some(500), + total_tokens: Some(600), + cache_read_input_tokens: Some(0), + cache_creation_input_tokens: Some(25), + cache_creation_ephemeral_5m_input_tokens: Some(0), + cache_creation_ephemeral_1h_input_tokens: Some(25), + cache_read_cost_usd: Some(0.0), + total_cost_usd: Some(1.25), + actual_total_cost_usd: Some(1.25), + status_code: Some(200), + is_stream: Some(false), + candidate_index: Some(0), + first_byte_time_ms: Some(0), + error_message: Some(String::new()), + request_metadata: Some(json!({ + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000", + "dimensions": {"image_count": 2, "size": "1024x1024", "quality": "high"}, + "usage_available": true, + "usage_pricing_available": true + })), + ..UsageEventData::default() + }, + } + } + + fn wire_value(encoded: &EncodedUsageEvent) -> Value { + serde_json::from_str(&encoded.fields["payload"]).expect("complete JSON wire payload") + } + + #[test] + fn event_wire_exact_limit_accounts_for_json_escaping_and_utf8() { + let mut event = event(); + event.request_id = "escaped\0\n\r\t\"\\\u{03bb}\u{1f600}".to_string(); + let original = event.to_stream_fields().expect("original wire payload"); + let length = original["payload"].len(); + for limit in [length, length + 1] { + let encoded = event.to_bounded_stream_fields(limit).expect("exact fit"); + assert_eq!(encoded.fields, original); + assert!(!encoded.diagnostics_omitted); + } + for limit in [0, 1, length - 1] { + assert!(matches!( + event.to_bounded_stream_fields(limit), + Err(DataLayerError::InvalidInput(_)) + )); + } + } + + #[test] + fn event_wire_writer_does_not_append_past_limit_and_reuses_failed_buffer() { + let mut writer = BoundedJsonWriter::new(5); + writer.write_all(b"12345").expect("exact fit"); + assert!(writer.write_all(b"6").is_err()); + assert_eq!(writer.bytes, b"12345"); + assert!(writer.exceeded); + assert!(!writer.serialize(&"\u{0000}").expect("size rejection")); + assert!(writer.bytes.len() <= 5); + assert!(writer.serialize(&"\u{03bb}").expect("valid UTF-8 retry")); + assert_eq!(writer.bytes, "\"\u{03bb}\"".as_bytes()); + assert!(!writer.exceeded); + } + + #[test] + fn event_wire_projection_transparently_forwards_derived_fields() { + #[derive(Serialize)] + struct FutureFields<'a> { + request_body: &'a Value, + new_billing_field: &'a Value, + zero_count: u64, + enabled: bool, + #[serde(skip_serializing_if = "Option::is_none")] + missing_field: Option<&'a str>, + } + let diagnostic = json!({"large": "x".repeat(8_192)}); + let billing = json!({"nested": [0, false, "unchanged"]}); + let data = FutureFields { + request_body: &diagnostic, + new_billing_field: &billing, + zero_count: 0, + enabled: false, + missing_field: None, + }; + let mut bytes = Vec::new(); + let mut serializer = serde_json::Serializer::new(&mut bytes); + data.serialize(FieldProjectionSerializer { + map: (&mut serializer) + .serialize_map(None) + .expect("object serializer"), + overrides: None, + }) + .expect("project derived fields"); + assert_eq!( + serde_json::from_slice::(&bytes).expect("projected JSON"), + json!({"new_billing_field": billing, "zero_count": 0, "enabled": false}) + ); + } + + #[test] + fn event_wire_omission_preserves_billing_facts_refs_and_source_ownership() { + let mut event = event(); + let padding = "x".repeat(16_384); + event.data.request_body = Some(json!({"reasoning": {"effort": "high"}, "input": padding})); + event.data.provider_request_body = Some( + json!({"model": "gpt-5.6-sol", "reasoning": {"effort": "medium"}, "service_tier": "priority", "input": padding}), + ); + event.data.response_body = Some(json!({"service_tier": "Default", "output": padding})); + event.data.client_response_body = Some(json!({"output": padding})); + event.data.request_headers = Some(json!({"x-request": padding})); + event.data.provider_request_headers = Some(json!({"x-provider": padding})); + event.data.response_headers = Some(json!({"x-response": padding})); + event.data.client_response_headers = Some(json!({"x-client": padding})); + event.data.request_body_ref = Some("usage://wire-request/request_body".to_string()); + event.data.request_body_state = Some(UsageBodyCaptureState::Inline); + event.data.provider_request_body_state = Some(UsageBodyCaptureState::Inline); + event.data.response_body_state = Some(UsageBodyCaptureState::Inline); + event.data.request_metadata.as_mut().unwrap()["body_capture"] = json!({ + "response": {"state": "inline", "source_bytes": 123_456} + }); + let weight = event.data.capture_heap_estimate(); + let budget = Arc::new(EventCaptureMemoryBudget::new(weight)); + event.data.apply_capture_memory_budget(Arc::clone(&budget)); + let before = serde_json::to_value(&event).expect("source snapshot"); + let original_wire: Value = + serde_json::from_str(&event.to_stream_fields().unwrap()["payload"]).unwrap(); + let encoded = event + .to_bounded_stream_fields(8_192) + .expect("diagnostic omission"); + assert!(encoded.diagnostics_omitted); + assert!(encoded.fields["payload"].len() <= 8_192); + let value = wire_value(&encoded); + for field in DIAGNOSTIC_FIELDS { + assert!( + value["data"].get(field).is_none(), + "{field} must be omitted" + ); + } + for (field, original) in original_wire["data"].as_object().unwrap() { + if !DIAGNOSTIC_FIELDS.contains(&field.as_str()) + && !BODY_STATE_FIELDS.contains(&field.as_str()) + && field != "request_metadata" + { + assert_eq!(&value["data"][field], original, "core field {field}"); + } + } + for (field, key) in BODY_STATE_FIELDS.into_iter().zip(BODY_METADATA_KEYS) { + assert_eq!(value["data"][field], "truncated"); + assert_eq!( + value["data"]["request_metadata"]["body_capture"][key]["reason"], + "wire_limit_exceeded" + ); + assert_eq!( + value["data"]["request_metadata"]["body_capture"][key]["stored_bytes"], + 0 + ); + } + let metadata = &value["data"]["request_metadata"]; + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_reasoning_effort"], "medium"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 30); + assert_eq!( + metadata["body_capture"]["response"]["source_bytes"], + 123_456 + ); + assert_eq!( + metadata["dimensions"], + before["data"]["request_metadata"]["dimensions"] + ); + assert_eq!(serde_json::to_value(&event).unwrap(), before); + assert_eq!(budget.retained_bytes(), weight); + assert_eq!(budget.downgraded_total(), 0); + + let decoded = UsageEvent::from_stream_fields(&encoded.fields).expect("wire decode"); + let record = crate::build_upsert_usage_record_from_event(&decoded).expect("record mapping"); + assert_eq!(record.input_tokens, Some(100)); + assert_eq!(record.output_tokens, Some(500)); + assert_eq!(record.cache_read_input_tokens, Some(0)); + assert_eq!(record.cache_creation_ephemeral_5m_input_tokens, Some(0)); + assert_eq!(record.cache_creation_ephemeral_1h_input_tokens, Some(25)); + assert_eq!(record.error_message.as_deref(), Some("")); + assert_eq!(record.request_body_ref, event.data.request_body_ref); + assert_eq!( + record.request_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + let metadata = record.request_metadata.expect("record billing metadata"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 30); + assert_eq!(metadata["dimensions"]["image_count"], 2); + drop(event); + assert_eq!(budget.retained_bytes(), 0); + } + + #[test] + fn event_wire_preserves_explicit_capture_states_and_legacy_metadata() { + for state in [ + None, + Some(UsageBodyCaptureState::None), + Some(UsageBodyCaptureState::Disabled), + Some(UsageBodyCaptureState::Unavailable), + Some(UsageBodyCaptureState::Reference), + ] { + let mut event = event(); + event.data.request_body_state = state; + event.data.provider_request_body_state = state; + event.data.response_body_state = state; + event.data.client_response_body_state = state; + event.data.request_body_ref = Some("usage://wire-request/request_body".to_string()); + event.data.response_headers = Some(json!({"large-header": "x".repeat(16_384)})); + event.data.request_metadata = Some(json!({ + "requested_reasoning_effort": "high", + "provider_reasoning_effort": "medium", + "provider_service_tier": "priority", + "provider_cache_ttl_minutes": 60, + "provider_actual_service_tier": "flex" + })); + let before = serde_json::to_value(&event).unwrap(); + let encoded = event + .to_bounded_stream_fields(4_096) + .expect("headers omitted"); + let value = wire_value(&encoded); + for field in BODY_STATE_FIELDS { + assert_eq!(value["data"].get(field), before["data"].get(field)); + } + let metadata = &value["data"]["request_metadata"]; + if state == Some(UsageBodyCaptureState::None) { + assert!(metadata.get("requested_reasoning_effort").is_none()); + assert!(metadata.get("provider_cache_ttl_minutes").is_none()); + } else { + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 60); + } + assert_eq!(metadata["provider_actual_service_tier"], "flex"); + assert_eq!( + value["data"]["request_body_ref"], + before["data"]["request_body_ref"] + ); + assert_eq!(serde_json::to_value(&event).unwrap(), before); + if matches!( + state, + Some( + UsageBodyCaptureState::None + | UsageBodyCaptureState::Disabled + | UsageBodyCaptureState::Unavailable + ) + ) { + event.data.endpoint_api_format = Some("claude:messages".to_string()); + event.data.request_body = Some(json!({"reasoning": {"effort": "low"}})); + event.data.provider_request_body = + Some(json!({"reasoning": {"effort": "low"}, "service_tier": "default"})); + event.data.response_body = Some(json!({"service_tier": "priority"})); + event.data.client_response_body = Some(json!({"stale": true})); + let baseline = UsageEvent::from_stream_fields_with_capture_budget( + &event.to_stream_fields().expect("full v1 fields"), + Arc::new(EventCaptureMemoryBudget::new(usize::MAX)), + ) + .expect("full v1 consumer baseline"); + let with_stale_bodies = wire_value( + &event + .to_bounded_stream_fields(4_096) + .expect("explicit capture states override stale bodies"), + ); + for field in BODY_STATE_FIELDS { + assert_eq!( + with_stale_bodies["data"].get(field), + value["data"].get(field) + ); + } + assert_eq!( + with_stale_bodies["data"]["request_metadata"], + serde_json::to_value(&baseline.data.request_metadata).unwrap(), + "wire omission must preserve the full v1 consumer's billing facts" + ); + } + } + } + + #[test] + fn event_wire_preserves_raw_body_cache_ttl_across_capture_states() { + use aether_data_contracts::repository::usage::extract_provider_cache_ttl_minutes_from_metadata; + + for state in [ + None, + Some(UsageBodyCaptureState::Inline), + Some(UsageBodyCaptureState::Reference), + Some(UsageBodyCaptureState::Truncated), + Some(UsageBodyCaptureState::Disabled), + Some(UsageBodyCaptureState::Unavailable), + ] { + let mut event = event(); + event.data.provider_request_body_state = state; + event.data.provider_request_body = Some(json!({ + "prompt_cache_options": {"ttl": "30m"}, + "service_tier": "default", + "reasoning": {"effort": "low"} + })); + event.data.response_headers = Some(json!({"large": "x".repeat(16_384)})); + event.data.request_metadata = Some(json!({ + "provider_cache_ttl_minutes": 60, + "provider_service_tier": "priority", + "provider_reasoning_effort": "high" + })); + let before = serde_json::to_value(&event).unwrap(); + let baseline = UsageEvent::from_stream_fields_with_capture_budget( + &event.to_stream_fields().expect("full v1 fields"), + Arc::new(EventCaptureMemoryBudget::new(usize::MAX)), + ) + .expect("full v1 consumer baseline"); + let encoded = event.to_bounded_stream_fields(4_096).expect("body omitted"); + assert!(encoded.diagnostics_omitted); + let decoded = UsageEvent::from_stream_fields_with_capture_budget( + &encoded.fields, + Arc::new(EventCaptureMemoryBudget::new(usize::MAX)), + ) + .expect("projected consumer event"); + assert!(decoded.data.provider_request_body.is_none()); + assert_eq!( + extract_provider_cache_ttl_minutes_from_metadata( + decoded.data.request_metadata.as_ref() + ), + Some(30), + "raw body TTL must win over stale metadata for {state:?}" + ); + for field in ["provider_service_tier", "provider_reasoning_effort"] { + assert_eq!( + decoded.data.request_metadata.as_ref().unwrap().get(field), + baseline.data.request_metadata.as_ref().unwrap().get(field), + "TTL preservation must not change {field} authority for {state:?}" + ); + } + assert_eq!(serde_json::to_value(&event).unwrap(), before); + } + } + + #[test] + fn event_wire_non_object_bodies_match_full_consumer_facts() { + use aether_data_contracts::repository::usage::{ + extract_provider_cache_ttl_minutes_from_metadata, + resolve_provider_service_tier_from_request_capture, + }; + + for body in [ + json!(null), + json!("opaque"), + json!([]), + json!(42), + json!(true), + ] { + for state in [ + None, + Some(UsageBodyCaptureState::None), + Some(UsageBodyCaptureState::Inline), + Some(UsageBodyCaptureState::Reference), + Some(UsageBodyCaptureState::Truncated), + Some(UsageBodyCaptureState::Disabled), + Some(UsageBodyCaptureState::Unavailable), + ] { + let mut event = event(); + event.data.request_body = Some(body.clone()); + event.data.provider_request_body = Some(body.clone()); + event.data.request_body_state = state; + event.data.provider_request_body_state = state; + event.data.response_headers = Some(json!({"large": "x".repeat(16_384)})); + event.data.request_metadata = Some(json!({ + "requested_reasoning_effort": "medium", + "provider_reasoning_effort": "high", + "provider_service_tier": "priority", + "provider_cache_ttl_minutes": 60, + "unrelated": {"preserved": true} + })); + let original = event.to_stream_fields().expect("full v1 fields"); + let baseline = UsageEvent::from_stream_fields_with_capture_budget( + &original, + Arc::new(EventCaptureMemoryBudget::new(usize::MAX)), + ) + .expect("full v1 consumer baseline"); + let encoded = event + .to_bounded_stream_fields(4_096) + .expect("omit diagnostics"); + assert!(encoded.diagnostics_omitted); + let decoded = UsageEvent::from_stream_fields_with_capture_budget( + &encoded.fields, + Arc::new(EventCaptureMemoryBudget::new(usize::MAX)), + ) + .expect("projected consumer event"); + let tier = |data: &UsageEventData| { + resolve_provider_service_tier_from_request_capture( + data.provider_request_body.as_ref(), + data.provider_request_body_state, + data.request_metadata.as_ref(), + ) + }; + assert_eq!( + tier(&decoded.data), + tier(&baseline.data), + "provider tier for {body:?}, {state:?}" + ); + let metadata = decoded.data.request_metadata.as_ref().unwrap(); + let baseline_metadata = baseline.data.request_metadata.as_ref().unwrap(); + assert_eq!( + extract_provider_cache_ttl_minutes_from_metadata(Some(metadata)), + extract_provider_cache_ttl_minutes_from_metadata(Some(baseline_metadata)), + "metadata TTL fallback for {body:?}, {state:?}" + ); + let authoritative = !body.is_null() + && matches!( + state, + None | Some( + UsageBodyCaptureState::Inline | UsageBodyCaptureState::Reference + ) + ); + for field in [ + REQUESTED_REASONING_EFFORT_METADATA_KEY, + PROVIDER_REASONING_EFFORT_METADATA_KEY, + ] { + assert_eq!( + metadata.get(field), + if authoritative { + None + } else { + baseline_metadata.get(field) + }, + "reasoning authority for {body:?}, {state:?}, {field}" + ); + } + assert_eq!(metadata["unrelated"], baseline_metadata["unrelated"]); + if body.is_null() { + assert_eq!( + decoded.data.request_body_state, + baseline.data.request_body_state + ); + assert_eq!( + decoded.data.provider_request_body_state, + baseline.data.provider_request_body_state + ); + } + assert_eq!(event.to_stream_fields().unwrap(), original); + } + } + } + + #[test] + fn event_wire_rejects_cache_ttl_loss_from_explicit_none_capture_state() { + let mut event = event(); + event.data.provider_request_body_state = Some(UsageBodyCaptureState::None); + event.data.provider_request_body = Some(json!({ + "prompt_cache_options": {"ttl": "30m"}, + "large": "x".repeat(16_384) + })); + let original = event.to_stream_fields().expect("full v1 fields"); + let full = event + .to_bounded_stream_fields(original["payload"].len()) + .expect("complete diagnostics remain representable"); + assert_eq!(full.fields, original); + assert!(!full.diagnostics_omitted); + assert!(matches!( + event.to_bounded_stream_fields(4_096), + Err(DataLayerError::InvalidInput(_)) + )); + assert_eq!(event.to_stream_fields().unwrap(), original); + } + + #[test] + fn event_wire_rejects_oversized_core_and_post_projection_metadata() { + for field in ["request_metadata", "error_message"] { + let mut event = event(); + if field == "request_metadata" { + event.data.request_metadata = Some(json!({"large": "x".repeat(32_768)})); + } else { + event.data.error_message = Some("x".repeat(32_768)); + } + let error = event + .to_bounded_stream_fields(1_024) + .expect_err("oversized core must fail"); + assert!(matches!(error, DataLayerError::InvalidInput(_))); + assert!( + error.to_string().len() < 200, + "errors must not include payload contents" + ); + } + let mut event = event(); + event.data.response_body = Some(json!({"large": "x".repeat(8_192)})); + let core = ProjectedData { + data: &event.data, + overrides: None, + }; + let core_size = serde_json::to_vec(&envelope(&event, &core)).unwrap().len(); + assert!( + matches!( + event.to_bounded_stream_fields(core_size), + Err(DataLayerError::InvalidInput(_)) + ), + "added capture metadata must also fit the exact wire limit" + ); + } +} diff --git a/crates/aether-usage/runtime/src/executor.rs b/crates/aether-usage/runtime/src/executor.rs index 14246dc75..a1434e9d5 100644 --- a/crates/aether-usage/runtime/src/executor.rs +++ b/crates/aether-usage/runtime/src/executor.rs @@ -1,5 +1,13 @@ use std::future::Future; -use std::sync::OnceLock; +use std::sync::{Mutex, OnceLock}; +use std::time::Duration; + +struct UsageBackgroundRuntime { + owner: Mutex>, + handle: tokio::runtime::Handle, +} + +static RUNTIME: OnceLock = OnceLock::new(); const DEFAULT_USAGE_BACKGROUND_RUNTIME_THREADS: usize = 8; const MAX_USAGE_BACKGROUND_RUNTIME_THREADS: usize = 64; @@ -18,12 +26,10 @@ where F: Future + Send + 'static, F::Output: Send + 'static, { - usage_background_runtime().handle().spawn(task) + usage_background_runtime().handle.spawn(task) } -fn usage_background_runtime() -> &'static tokio::runtime::Runtime { - static RUNTIME: OnceLock<&'static tokio::runtime::Runtime> = OnceLock::new(); - +fn usage_background_runtime() -> &'static UsageBackgroundRuntime { RUNTIME.get_or_init(|| { let worker_threads = usage_background_runtime_threads(); let runtime = tokio::runtime::Builder::new_multi_thread() @@ -36,10 +42,28 @@ fn usage_background_runtime() -> &'static tokio::runtime::Runtime { .thread_stack_size(USAGE_BACKGROUND_RUNTIME_STACK_BYTES) .build() .expect("usage background runtime should build"); - Box::leak(Box::new(runtime)) + UsageBackgroundRuntime { + handle: runtime.handle().clone(), + owner: Mutex::new(Some(runtime)), + } }) } +/// Call outside Tokio after every UsageRuntime has drained. Does not start an unused runtime. +pub fn shutdown_usage_background_runtime(timeout: Duration) { + let Some(runtime) = RUNTIME.get() else { + return; + }; + let owner = runtime + .owner + .lock() + .unwrap_or_else(|p| p.into_inner()) + .take(); + if let Some(owner) = owner { + owner.shutdown_timeout(timeout); + } +} + fn usage_background_runtime_threads() -> usize { parse_usage_background_runtime_threads( std::env::var(GATEWAY_USAGE_BACKGROUND_RUNTIME_THREADS_ENV) diff --git a/crates/aether-usage/runtime/src/lib.rs b/crates/aether-usage/runtime/src/lib.rs index 71f6a5e57..027f575b9 100644 --- a/crates/aether-usage/runtime/src/lib.rs +++ b/crates/aether-usage/runtime/src/lib.rs @@ -1,15 +1,19 @@ mod body_capture; pub mod config; +mod dead_letter_encoding; pub mod event; +mod event_capture_budget; mod executor; mod keyed_lock; pub mod queue; +mod queue_read_budget; pub mod record; pub mod report; pub mod report_context; mod request_metadata; pub mod runtime; pub mod settlement; +mod shutdown; pub mod standardized_usage; pub mod usage_mapper; pub mod worker; @@ -21,6 +25,7 @@ pub use body_capture::{ }; pub use config::UsageRuntimeConfig; pub use event::{now_ms, UsageEvent, UsageEventData, UsageEventType, USAGE_EVENT_VERSION}; +pub use executor::shutdown_usage_background_runtime; pub use queue::UsageQueue; pub use record::build_upsert_usage_record_from_event; pub use report::{ @@ -50,6 +55,7 @@ pub use runtime::{ pub use settlement::{ reconcile_usage_policy_cost_for_event, settle_usage_if_needed, UsageSettlementWriter, }; +pub use shutdown::UsageProducerGuard; pub use standardized_usage::StandardizedUsage; pub use usage_mapper::{map_usage, map_usage_from_response, UsageMapper}; pub use worker::{ diff --git a/crates/aether-usage/runtime/src/queue.rs b/crates/aether-usage/runtime/src/queue.rs index 2dfd9c0fa..23b6dadb6 100644 --- a/crates/aether-usage/runtime/src/queue.rs +++ b/crates/aether-usage/runtime/src/queue.rs @@ -1,14 +1,30 @@ +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; -use serde_json::json; - use aether_data_contracts::DataLayerError; use aether_runtime_state::{ - RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats, RuntimeQueueStore, + RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueReclaimPage, RuntimeQueueStats, + RuntimeQueueStore, RuntimeQueueTransferOutcome, }; use super::config::UsageRuntimeConfig; -use super::event::UsageEvent; +use super::event::{EncodedUsageEvent, UsageEvent}; +use crate::dead_letter_encoding::{shared_dead_letter_encoding_budget, DeadLetterEncodingBudget}; +use crate::queue_read_budget::{shared_queue_read_budget, QueueReadBudget, QueueReadReservation}; + +static PAYLOAD_DOWNGRADED_TOTAL: AtomicU64 = AtomicU64::new(0); +static PAYLOAD_REJECTED_TOTAL: AtomicU64 = AtomicU64::new(0); + +pub(crate) fn payload_encoding_totals() -> (u64, u64) { + ( + PAYLOAD_DOWNGRADED_TOTAL.load(Ordering::Relaxed), + PAYLOAD_REJECTED_TOTAL.load(Ordering::Relaxed), + ) +} + +pub(crate) fn is_permanent_enqueue_error(error: &DataLayerError) -> bool { + matches!(error, DataLayerError::InvalidInput(_)) +} #[derive(Clone)] pub struct UsageQueue { @@ -17,6 +33,35 @@ pub struct UsageQueue { stream: String, group: String, dlq_stream: String, + read_budget: Arc, + dead_letter_encoding_budget: Arc, +} + +#[derive(Debug)] +pub(crate) enum UsageDeadLetterOutcome { + Transferred { + destination_id: String, + acked: usize, + }, + Appended { + destination_id: String, + }, + NotPending, + EncodingDeferred { + error: DataLayerError, + }, +} + +pub(crate) struct ReservedUsageReadBatch { + pub(crate) entries: Vec, + pub(crate) requested_count: usize, + // Keep last: batch data must be dropped before its reservation is returned. + pub(crate) reservation: QueueReadReservation, +} + +pub(crate) struct ReservedUsageReclaimPage { + pub(crate) page: RuntimeQueueReclaimPage, + pub(crate) reservation: QueueReadReservation, } impl UsageQueue { @@ -31,9 +76,26 @@ impl UsageQueue { group: config.consumer_group.clone(), dlq_stream: config.dlq_stream_key.clone(), config, + read_budget: shared_queue_read_budget(), + dead_letter_encoding_budget: shared_dead_letter_encoding_budget(), }) } + #[cfg(test)] + pub(crate) fn with_read_budget(mut self, read_budget: Arc) -> Self { + self.read_budget = read_budget; + self + } + + #[cfg(test)] + pub(crate) fn with_dead_letter_encoding_budget( + mut self, + budget: Arc, + ) -> Self { + self.dead_letter_encoding_budget = budget; + self + } + pub async fn ensure_consumer_group(&self) -> Result<(), DataLayerError> { self.runner .ensure_consumer_group(&self.stream, &self.group, "0-0") @@ -41,12 +103,36 @@ impl UsageQueue { } pub async fn enqueue(&self, event: &UsageEvent) -> Result { - let fields = event.to_stream_fields()?; + let encoded = self.encode_event(event)?; self.runner - .append_fields_with_maxlen(&self.stream, &fields, Some(self.config.stream_maxlen)) + .append_fields_with_maxlen( + &self.stream, + &encoded.fields, + Some(self.config.stream_maxlen), + ) .await } + pub(crate) fn validate_event(&self, event: &UsageEvent) -> Result<(), DataLayerError> { + self.encode_event(event).map(|_| ()) + } + + fn encode_event(&self, event: &UsageEvent) -> Result { + let encoded = match event.to_bounded_stream_fields(self.config.queue_payload_max_bytes) { + Ok(encoded) => encoded, + Err(error) => { + if is_permanent_enqueue_error(&error) { + PAYLOAD_REJECTED_TOTAL.fetch_add(1, Ordering::Relaxed); + } + return Err(error); + } + }; + if encoded.diagnostics_omitted { + PAYLOAD_DOWNGRADED_TOTAL.fetch_add(1, Ordering::Relaxed); + } + Ok(encoded) + } + pub async fn read_group( &self, consumer: &str, @@ -62,13 +148,52 @@ impl UsageQueue { .await } + /// Workers retain this lease through processing. The public Vec API remains + /// compatible, but cannot preserve a reservation after returning its entries. + pub(crate) async fn read_group_reserved( + &self, + consumer: &str, + ) -> Result { + let (requested_count, mut reservation) = self + .read_budget + .reserve( + self.config.consumer_batch_size, + self.config.queue_payload_max_bytes, + ) + .await?; + let entries = self + .runner + .read_group( + &self.stream, + &self.group, + consumer, + requested_count, + Some(self.config.consumer_block_ms.max(1)), + ) + .await?; + reservation.observe_entries(&entries, self.config.queue_payload_max_bytes); + Ok(ReservedUsageReadBatch { + entries, + requested_count, + reservation, + }) + } + pub async fn claim_stale( &self, consumer: &str, start_id: &str, ) -> Result, DataLayerError> { + Ok(self.claim_stale_page(consumer, start_id).await?.entries) + } + + pub async fn claim_stale_page( + &self, + consumer: &str, + start_id: &str, + ) -> Result { self.runner - .claim_stale( + .claim_stale_page( &self.stream, &self.group, consumer, @@ -81,10 +206,46 @@ impl UsageQueue { .await } + pub(crate) async fn claim_stale_page_reserved( + &self, + consumer: &str, + start_id: &str, + ) -> Result { + let (requested_count, mut reservation) = self + .read_budget + .reserve( + self.config.reclaim_count, + self.config.queue_payload_max_bytes, + ) + .await?; + let page = self + .runner + .claim_stale_page( + &self.stream, + &self.group, + consumer, + start_id, + RuntimeQueueReclaimConfig { + min_idle_ms: self.config.reclaim_idle_ms, + count: requested_count, + }, + ) + .await?; + reservation.observe_entries(&page.entries, self.config.queue_payload_max_bytes); + Ok(ReservedUsageReclaimPage { page, reservation }) + } + pub async fn ack_and_delete(&self, ids: &[String]) -> Result<(), DataLayerError> { - self.runner.ack(&self.stream, &self.group, ids).await?; + self.ack_and_delete_counted(ids).await.map(|_| ()) + } + + pub(crate) async fn ack_and_delete_counted( + &self, + ids: &[String], + ) -> Result { + let acked = self.runner.ack(&self.stream, &self.group, ids).await?; self.runner.delete(&self.stream, ids).await?; - Ok(()) + Ok(acked) } pub async fn push_dead_letter( @@ -92,20 +253,57 @@ impl UsageQueue { entry: &RuntimeQueueEntry, error: &str, ) -> Result { - let fields = std::collections::BTreeMap::from([( - "payload".to_string(), - serde_json::to_string(&json!({ - "entry_id": entry.id, - "fields": entry.fields, - "error": error, - })) - .map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?, - )]); + let reservation = self.dead_letter_encoding_budget.try_reserve(entry, error)?; + let encoded = reservation + .encode_owned(entry.clone(), error.to_string()) + .await?; self.runner - .append_fields_with_maxlen(&self.dlq_stream, &fields, None) + .append_fields_with_maxlen(&self.dlq_stream, &encoded.fields, None) .await } + pub(crate) async fn transfer_dead_letter_owned( + &self, + entry: RuntimeQueueEntry, + error: String, + ) -> Result { + let reservation = match self.dead_letter_encoding_budget.try_reserve(&entry, &error) { + Ok(reservation) => reservation, + Err(error) => return Ok(UsageDeadLetterOutcome::EncodingDeferred { error }), + }; + let encoded = match reservation.encode_owned(entry, error).await { + Ok(encoded) => encoded, + Err(error) => return Ok(UsageDeadLetterOutcome::EncodingDeferred { error }), + }; + match self + .runner + .try_transfer_pending_to_stream( + &self.stream, + &self.group, + &encoded.entry_id, + &self.dlq_stream, + &encoded.fields, + ) + .await? + { + Some(RuntimeQueueTransferOutcome::Transferred { + destination_id, + acked, + .. + }) => Ok(UsageDeadLetterOutcome::Transferred { + destination_id, + acked, + }), + Some(RuntimeQueueTransferOutcome::NotPending) => Ok(UsageDeadLetterOutcome::NotPending), + None => Ok(UsageDeadLetterOutcome::Appended { + destination_id: self + .runner + .append_fields_with_maxlen(&self.dlq_stream, &encoded.fields, None) + .await?, + }), + } + } + pub async fn stats(&self) -> Result { self.runner.stats(&self.stream, Some(&self.group)).await } @@ -136,10 +334,216 @@ fn usage_queue_runtime_settings(config: &UsageRuntimeConfig) -> UsageQueueRuntim #[cfg(test)] mod tests { - use super::{usage_queue_runtime_settings, UsageQueue, UsageQueueRuntimeSettings}; + use super::{ + usage_queue_runtime_settings, UsageDeadLetterOutcome, UsageQueue, UsageQueueRuntimeSettings, + }; + use crate::dead_letter_encoding::DeadLetterEncodingBudget; + use crate::queue_read_budget::QueueReadBudget; use crate::UsageRuntimeConfig; - use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState}; + use aether_data_contracts::DataLayerError; + use aether_runtime_state::{ + MemoryRuntimeStateConfig, RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats, + RuntimeQueueStore, RuntimeQueueTransferOutcome, RuntimeState, + }; + use async_trait::async_trait; + use std::collections::BTreeMap; + use std::future::Future; use std::sync::Arc; + use std::task::Poll; + use std::time::Duration; + + struct HeldDeadLetterStore { + started: tokio::sync::Notify, + } + + #[async_trait] + impl RuntimeQueueStore for HeldDeadLetterStore { + async fn ensure_consumer_group( + &self, + _stream: &str, + _group: &str, + _start_id: &str, + ) -> Result<(), DataLayerError> { + unreachable!("store only exercises dead-letter writes") + } + + async fn append_fields_with_maxlen( + &self, + _stream: &str, + fields: &BTreeMap, + _maxlen: Option, + ) -> Result { + assert!(fields.contains_key("payload")); + self.started.notify_one(); + std::future::pending().await + } + + async fn read_group( + &self, + _stream: &str, + _group: &str, + _consumer: &str, + _count: usize, + _block_ms: Option, + ) -> Result, DataLayerError> { + unreachable!("store only exercises dead-letter writes") + } + + async fn claim_stale( + &self, + _stream: &str, + _group: &str, + _consumer: &str, + _start_id: &str, + _config: RuntimeQueueReclaimConfig, + ) -> Result, DataLayerError> { + unreachable!("store only exercises dead-letter writes") + } + + async fn try_transfer_pending_to_stream( + &self, + _source: &str, + _group: &str, + _entry_id: &str, + _destination: &str, + fields: &BTreeMap, + ) -> Result, DataLayerError> { + assert!(fields.contains_key("payload")); + self.started.notify_one(); + std::future::pending().await + } + + async fn ack( + &self, + _stream: &str, + _group: &str, + _ids: &[String], + ) -> Result { + unreachable!("store only exercises dead-letter writes") + } + + async fn delete(&self, _stream: &str, _ids: &[String]) -> Result { + unreachable!("store only exercises dead-letter writes") + } + + async fn stats( + &self, + _stream: &str, + _group: Option<&str>, + ) -> Result { + unreachable!("store only exercises dead-letter writes") + } + } + + #[tokio::test] + async fn dead_letter_encoding_storage_wait_keeps_budget_until_public_or_owned_call_is_cancelled( + ) { + for owned_transfer in [false, true] { + let store = Arc::new(HeldDeadLetterStore { + started: tokio::sync::Notify::new(), + }); + let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1)); + let queue = UsageQueue::new(store.clone(), UsageRuntimeConfig::default()) + .unwrap() + .with_dead_letter_encoding_budget(Arc::clone(&budget)); + let entry = RuntimeQueueEntry { + id: "1-0".to_string(), + fields: BTreeMap::from([("payload".to_string(), "original fields".to_string())]), + }; + let task = tokio::spawn(async move { + if owned_transfer { + queue + .transfer_dead_letter_owned(entry, "failure".to_string()) + .await + .map(|_| ()) + } else { + queue.push_dead_letter(&entry, "failure").await.map(|_| ()) + } + }); + tokio::time::timeout(Duration::from_secs(2), store.started.notified()) + .await + .unwrap(); + assert_eq!(budget.snapshot().encoded_total, 1); + assert_eq!(budget.snapshot().active_jobs, 1); + assert!(budget.snapshot().reserved_bytes > 0); + task.abort(); + assert!(matches!(task.await, Err(error) if error.is_cancelled())); + assert_eq!(budget.snapshot().active_jobs, 0); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + } + + #[tokio::test] + async fn dead_letter_encoding_owned_defers_oversize_but_preserves_public_and_store_errors() { + let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let queue = UsageQueue::new(runner, UsageRuntimeConfig::default()) + .unwrap() + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(1, 1))); + let entry = RuntimeQueueEntry { + id: "1-0".to_string(), + fields: BTreeMap::from([("payload".to_string(), "original fields".to_string())]), + }; + assert!(matches!( + queue + .transfer_dead_letter_owned(entry.clone(), "failure".to_string()) + .await, + Ok(UsageDeadLetterOutcome::EncodingDeferred { + error: DataLayerError::InvalidInput(_), + }) + )); + assert!(matches!( + queue.push_dead_letter(&entry, "failure").await, + Err(DataLayerError::InvalidInput(_)) + )); + + let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1)); + let queue = queue.with_dead_letter_encoding_budget(Arc::clone(&budget)); + // The absent source causes a native store error after successful encoding. + assert!(matches!( + queue + .transfer_dead_letter_owned(entry, "failure".to_string()) + .await, + Err(DataLayerError::InvalidInput(_)) + )); + assert_eq!(budget.snapshot().encoded_total, 1); + assert_eq!(budget.snapshot().active_jobs, 0); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(queue.dlq_stats().await.unwrap().stream_length, 0); + } + + #[tokio::test] + async fn dead_letter_encoding_owned_defers_capacity_without_starting_store_work() { + let store = Arc::new(HeldDeadLetterStore { + started: tokio::sync::Notify::new(), + }); + let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1)); + let queue = UsageQueue::new(store, UsageRuntimeConfig::default()) + .unwrap() + .with_dead_letter_encoding_budget(Arc::clone(&budget)); + let entry = RuntimeQueueEntry { + id: "1-0".to_string(), + fields: BTreeMap::from([("payload".to_string(), "original fields".to_string())]), + }; + let held = budget.try_reserve(&entry, "failure").unwrap(); + let result = tokio::time::timeout( + Duration::from_secs(2), + queue.transfer_dead_letter_owned(entry, "failure".to_string()), + ) + .await + .unwrap(); + assert!(matches!( + result, + Ok(UsageDeadLetterOutcome::EncodingDeferred { + error: DataLayerError::TimedOut(_), + }) + )); + assert_eq!(budget.snapshot().encoded_total, 0); + assert_eq!(budget.snapshot().active_jobs, 1); + assert_eq!(budget.snapshot().capacity_rejected_total, 1); + drop(held); + assert_eq!(budget.snapshot().active_jobs, 0); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } #[test] fn usage_queue_applies_runtime_block_and_batch_settings() { @@ -164,4 +568,135 @@ mod tests { } ); } + + fn reserved_test_queue(runner: Arc, budget: Arc) -> UsageQueue { + UsageQueue::new( + runner, + UsageRuntimeConfig { + enabled: true, + queue_payload_max_bytes: 8, + consumer_batch_size: 128, + consumer_block_ms: 1, + reclaim_count: 128, + reclaim_idle_ms: 1, + ..UsageRuntimeConfig::default() + }, + ) + .expect("test queue") + .with_read_budget(budget) + } + + #[tokio::test] + async fn queue_read_budget_read_and_reclaim_share_a_reservation_across_clones() { + let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let budget = Arc::new(QueueReadBudget::new(16, 16)); + let queue = reserved_test_queue(Arc::clone(&runtime), Arc::clone(&budget)); + let other = queue.clone(); + queue.ensure_consumer_group().await.unwrap(); + for _ in 0..6 { + runtime + .append_fields_with_maxlen( + &queue.stream, + &BTreeMap::from([("payload".to_string(), "12345678".to_string())]), + None, + ) + .await + .unwrap(); + } + let first = queue.read_group_reserved("reader").await.unwrap(); + assert_eq!(first.requested_count, 2); + assert_eq!(first.entries.len(), 2); + let first_ids = first + .entries + .iter() + .map(|entry| entry.id.clone()) + .collect::>(); + let extra_pending = runtime + .read_group(&queue.stream, &queue.group, "previous-reader", 2, None) + .await + .unwrap(); + assert_eq!(extra_pending.len(), 2); + let next_cursor = extra_pending[0].id.clone(); + drop(extra_pending); + assert_eq!(budget.snapshot().reserved_bytes, 16); + + let mut reclaim = Box::pin(other.claim_stale_page_reserved("reclaimer", "0-0")); + std::future::poll_fn(|cx| { + assert!(reclaim.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + assert_eq!(budget.snapshot().waiters, 1); + tokio::time::sleep(Duration::from_millis(5)).await; + drop(first); + let claimed = reclaim.await.unwrap(); + assert_eq!(claimed.page.entries.len(), 2); + assert_eq!(claimed.page.next_start_id, next_cursor); + assert_eq!( + claimed + .page + .entries + .iter() + .map(|entry| entry.id.clone()) + .collect::>(), + first_ids + ); + assert_eq!(budget.snapshot().reserved_bytes, 16); + assert_eq!(budget.snapshot().waiters, 0); + + let mut next_read = Box::pin(queue.read_group_reserved("reader")); + std::future::poll_fn(|cx| { + assert!(next_read.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + drop(claimed); + let next = next_read.await.unwrap(); + assert_eq!(next.entries.len(), 2); + assert!(next + .entries + .iter() + .all(|entry| !first_ids.contains(&entry.id))); + drop(next); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().wait_total, 2); + } + + #[tokio::test] + async fn queue_read_budget_errors_and_empty_pages_release_reservations() { + let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let budget = Arc::new(QueueReadBudget::new(16, 16)); + let queue = reserved_test_queue(runtime, Arc::clone(&budget)); + assert!(queue.read_group_reserved("reader").await.is_err()); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert!(queue + .claim_stale_page_reserved("reader", "0-0") + .await + .is_err()); + assert_eq!(budget.snapshot().reserved_bytes, 0); + + queue.ensure_consumer_group().await.unwrap(); + let empty = queue.read_group_reserved("reader").await.unwrap(); + assert!(empty.entries.is_empty()); + assert_eq!(budget.snapshot().reserved_bytes, 0); + let page = queue + .claim_stale_page_reserved("reader", "0-0") + .await + .unwrap(); + assert!(page.page.entries.is_empty()); + assert_eq!(page.page.next_start_id, "0-0"); + assert_eq!(budget.snapshot().reserved_bytes, 0); + drop((empty, page)); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[test] + fn queue_read_budget_new_queues_and_clones_share_process_budget() { + let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let first = UsageQueue::new(runtime.clone(), UsageRuntimeConfig::default()).unwrap(); + let second = UsageQueue::new(runtime, UsageRuntimeConfig::default()).unwrap(); + let cloned = first.clone(); + assert!(Arc::ptr_eq(&first.read_budget, &second.read_budget)); + assert!(Arc::ptr_eq(&first.read_budget, &cloned.read_budget)); + } } diff --git a/crates/aether-usage/runtime/src/queue_read_budget.rs b/crates/aether-usage/runtime/src/queue_read_budget.rs new file mode 100644 index 000000000..a0f8a5d42 --- /dev/null +++ b/crates/aether-usage/runtime/src/queue_read_budget.rs @@ -0,0 +1,373 @@ +use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; +use std::sync::{Arc, LazyLock}; + +use aether_data_contracts::DataLayerError; +use aether_runtime_state::RuntimeQueueEntry; +use tokio::sync::{OwnedSemaphorePermit, Semaphore, TryAcquireError}; + +const DEFAULT_READ_PAYLOAD_BUDGET_BYTES: usize = 128 * 1024 * 1024; +const DEFAULT_READ_BATCH_PAYLOAD_BYTES: usize = 8 * 1024 * 1024; + +static READ_BUDGET: LazyLock> = LazyLock::new(|| { + let limit = std::env::var("AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES").ok(); + let batch = std::env::var("AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES").ok(); + Arc::new(QueueReadBudget::new( + configured_bytes(limit.as_deref(), DEFAULT_READ_PAYLOAD_BUDGET_BYTES), + configured_bytes(batch.as_deref(), DEFAULT_READ_BATCH_PAYLOAD_BYTES), + )) +}); + +pub(crate) fn shared_queue_read_budget() -> Arc { + Arc::clone(&READ_BUDGET) +} + +pub(crate) fn queue_read_budget_metrics() -> QueueReadBudgetSnapshot { + READ_BUDGET.snapshot() +} + +fn maximum_budget_bytes() -> usize { + Semaphore::MAX_PERMITS.min(u32::MAX as usize) +} + +fn configured_bytes(raw: Option<&str>, fallback: usize) -> usize { + raw.and_then(|raw| raw.trim().parse::().ok()) + .filter(|value| *value > 0) + .map(|value| value.min(maximum_budget_bytes() as u128) as usize) + .unwrap_or(fallback.min(maximum_budget_bytes())) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct QueueReadBudgetSnapshot { + pub(crate) limit_bytes: usize, + pub(crate) batch_limit_bytes: usize, + pub(crate) reserved_bytes: usize, + pub(crate) waiters: usize, + pub(crate) wait_total: u64, + /// Observed field names plus values, without cloning the returned strings. + pub(crate) actual_field_bytes_total: u64, + pub(crate) oversized_entries_total: u64, + pub(crate) oversized_batches_total: u64, +} + +/// A process-wide reservation based on the current producer payload limit. +/// Historical or externally written messages can exceed the estimate. Field names, +/// allocation capacity, RESP decoding, and decoded JSON are not an RSS bound here. +pub(crate) struct QueueReadBudget { + limit_bytes: usize, + batch_limit_bytes: usize, + permits: Arc, + reserved_bytes: AtomicUsize, + waiters: AtomicUsize, + wait_total: AtomicU64, + actual_field_bytes_total: AtomicU64, + oversized_entries_total: AtomicU64, + oversized_batches_total: AtomicU64, +} + +impl QueueReadBudget { + pub(crate) fn new(limit_bytes: usize, batch_limit_bytes: usize) -> Self { + let limit_bytes = limit_bytes.clamp(1, maximum_budget_bytes()); + let batch_limit_bytes = batch_limit_bytes.clamp(1, limit_bytes); + Self { + limit_bytes, + batch_limit_bytes, + permits: Arc::new(Semaphore::new(limit_bytes)), + reserved_bytes: AtomicUsize::new(0), + waiters: AtomicUsize::new(0), + wait_total: AtomicU64::new(0), + actual_field_bytes_total: AtomicU64::new(0), + oversized_entries_total: AtomicU64::new(0), + oversized_batches_total: AtomicU64::new(0), + } + } + + pub(crate) fn snapshot(&self) -> QueueReadBudgetSnapshot { + QueueReadBudgetSnapshot { + limit_bytes: self.limit_bytes, + batch_limit_bytes: self.batch_limit_bytes, + reserved_bytes: self.reserved_bytes.load(Ordering::Relaxed), + waiters: self.waiters.load(Ordering::Relaxed), + wait_total: self.wait_total.load(Ordering::Relaxed), + actual_field_bytes_total: self.actual_field_bytes_total.load(Ordering::Relaxed), + oversized_entries_total: self.oversized_entries_total.load(Ordering::Relaxed), + oversized_batches_total: self.oversized_batches_total.load(Ordering::Relaxed), + } + } + + pub(crate) async fn reserve( + self: &Arc, + requested_count: usize, + payload_limit: usize, + ) -> Result<(usize, QueueReadReservation), DataLayerError> { + if payload_limit == 0 || payload_limit > self.limit_bytes { + return Err(DataLayerError::InvalidConfiguration(format!( + "usage queue payload limit {payload_limit} must be positive and not exceed the {}-byte read payload budget", + self.limit_bytes + ))); + } + // A single valid payload may exceed the preferred batch target, but never + // the total budget. Clamp before multiplying or converting to u32 permits. + let count = requested_count + .max(1) + .min((self.batch_limit_bytes / payload_limit).max(1)); + let reserved_bytes = count * payload_limit; + let permits = reserved_bytes as u32; + let permit = match Arc::clone(&self.permits).try_acquire_many_owned(permits) { + Ok(permit) => permit, + Err(TryAcquireError::NoPermits) => { + self.wait_total.fetch_add(1, Ordering::Relaxed); + self.waiters.fetch_add(1, Ordering::Relaxed); + let _waiting = WaitingReservation { budget: self }; + Arc::clone(&self.permits) + .acquire_many_owned(permits) + .await + .map_err(|_| closed_budget_error())? + } + Err(TryAcquireError::Closed) => return Err(closed_budget_error()), + }; + self.reserved_bytes + .fetch_add(reserved_bytes, Ordering::Relaxed); + Ok(( + count, + QueueReadReservation { + budget: Arc::clone(self), + reserved_bytes, + permit: Some(permit), + }, + )) + } +} + +fn closed_budget_error() -> DataLayerError { + DataLayerError::InvalidConfiguration("usage queue read payload budget is closed".to_string()) +} + +struct WaitingReservation<'a> { + budget: &'a QueueReadBudget, +} + +impl Drop for WaitingReservation<'_> { + fn drop(&mut self) { + self.budget.waiters.fetch_sub(1, Ordering::Relaxed); + } +} + +// Deliberately not Clone: every concurrently retained batch needs its own lease. +pub(crate) struct QueueReadReservation { + budget: Arc, + reserved_bytes: usize, + permit: Option, +} + +impl QueueReadReservation { + pub(crate) fn observe_entries(&mut self, entries: &[RuntimeQueueEntry], payload_limit: usize) { + let mut value_bytes = 0usize; + let mut field_bytes = 0usize; + let mut oversized_entries = 0u64; + for entry in entries { + let mut entry_value_bytes = 0usize; + for (key, value) in &entry.fields { + entry_value_bytes = entry_value_bytes.saturating_add(value.len()); + field_bytes = field_bytes + .saturating_add(key.len()) + .saturating_add(value.len()); + } + value_bytes = value_bytes.saturating_add(entry_value_bytes); + oversized_entries += u64::from(entry_value_bytes > payload_limit); + } + self.budget.actual_field_bytes_total.fetch_add( + u64::try_from(field_bytes).unwrap_or(u64::MAX), + Ordering::Relaxed, + ); + self.budget + .oversized_entries_total + .fetch_add(oversized_entries, Ordering::Relaxed); + if value_bytes > self.reserved_bytes { + self.budget + .oversized_batches_total + .fetch_add(1, Ordering::Relaxed); + } + // Shrink unused payload estimates. Never wait for an upgrade after reading + // an oversized historical batch: other batches may hold all remaining bytes. + self.shrink_to(value_bytes.min(self.reserved_bytes)); + } + + fn shrink_to(&mut self, retained_bytes: usize) { + let released = self.reserved_bytes.saturating_sub(retained_bytes); + if released == 0 { + return; + } + let permit = self + .permit + .as_mut() + .expect("positive reservation must hold a permit") + .split(released) + .expect("released bytes must belong to this reservation"); + self.reserved_bytes -= released; + self.budget + .reserved_bytes + .fetch_sub(released, Ordering::Relaxed); + drop(permit); + } +} + +impl Drop for QueueReadReservation { + fn drop(&mut self) { + self.budget + .reserved_bytes + .fetch_sub(self.reserved_bytes, Ordering::Relaxed); + } +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + use std::future::Future; + use std::task::Poll; + + use super::*; + + #[tokio::test] + async fn queue_read_budget_cancelled_wait_preserves_current_reservations() { + let budget = Arc::new(QueueReadBudget::new(32, 16)); + let (_, first) = budget.reserve(8, 8).await.unwrap(); + let (_, second) = budget.reserve(8, 8).await.unwrap(); + let mut pending = Box::pin(budget.reserve(1, 8)); + std::future::poll_fn(|cx| { + assert!(pending.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + assert_eq!(budget.snapshot().reserved_bytes, 32); + assert_eq!(budget.snapshot().waiters, 1); + drop(pending); + assert_eq!(budget.snapshot().waiters, 0); + assert_eq!(budget.snapshot().wait_total, 1); + drop(first); + let (_, replacement) = budget.reserve(2, 8).await.unwrap(); + assert_eq!(budget.snapshot().reserved_bytes, 32); + drop((second, replacement)); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.permits.available_permits(), 32); + } + + #[tokio::test] + async fn queue_read_budget_counts_payload_values_and_observes_legacy_excess_without_waiting() { + let budget = Arc::new(QueueReadBudget::new(16, 16)); + let (count, mut reservation) = budget.reserve(2, 8).await.unwrap(); + assert_eq!(count, 2); + let entries = [RuntimeQueueEntry { + id: "1-0".to_string(), + fields: BTreeMap::from([("payload".to_string(), "x".repeat(8))]), + }]; + reservation.observe_entries(&entries, 8); + assert_eq!(budget.snapshot().reserved_bytes, 8); + assert_eq!(budget.snapshot().actual_field_bytes_total, 15); + assert_eq!(budget.snapshot().oversized_entries_total, 0); + let (_, mut legacy) = budget.reserve(1, 8).await.unwrap(); + let legacy_entries = [RuntimeQueueEntry { + id: "2-0".to_string(), + fields: BTreeMap::from([ + ("payload".to_string(), "x".repeat(8)), + ("extra".to_string(), "y".repeat(24)), + ]), + }]; + legacy.observe_entries(&legacy_entries, 8); + assert_eq!(budget.snapshot().reserved_bytes, 16); + assert_eq!(budget.snapshot().actual_field_bytes_total, 59); + assert_eq!(budget.snapshot().oversized_entries_total, 1); + assert_eq!(budget.snapshot().oversized_batches_total, 1); + assert_eq!(budget.snapshot().wait_total, 0); + drop((reservation, legacy)); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[tokio::test] + async fn queue_read_budget_empty_response_releases_all_reserved_bytes() { + let budget = Arc::new(QueueReadBudget::new(16, 16)); + let (_, mut reservation) = budget.reserve(2, 8).await.unwrap(); + reservation.observe_entries(&[], 8); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.permits.available_permits(), 16); + drop(reservation); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn queue_read_budget_concurrent_shrink_and_drop_preserve_shared_capacity() { + let budget = Arc::new(QueueReadBudget::new(64, 16)); + let barrier = Arc::new(tokio::sync::Barrier::new(16)); + let mut tasks = tokio::task::JoinSet::new(); + for _ in 0..16 { + let budget = Arc::clone(&budget); + let barrier = Arc::clone(&barrier); + tasks.spawn(async move { + barrier.wait().await; + for _ in 0..16 { + let (count, mut reservation) = budget.reserve(128, 8).await.unwrap(); + assert_eq!(count, 2); + assert!(budget.snapshot().reserved_bytes <= 64); + reservation.observe_entries( + &[RuntimeQueueEntry { + id: "1-0".to_string(), + fields: BTreeMap::from([( + "payload".to_string(), + "12345678".to_string(), + )]), + }], + 8, + ); + tokio::task::yield_now().await; + drop(reservation); + } + }); + } + tokio::time::timeout(std::time::Duration::from_secs(5), async { + while let Some(result) = tasks.join_next().await { + result.expect("reservation task"); + } + }) + .await + .expect("all shared reservations complete"); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().waiters, 0); + assert_eq!(budget.snapshot().actual_field_bytes_total, 16 * 16 * 15); + assert_eq!(budget.permits.available_permits(), 64); + } + + #[tokio::test] + async fn queue_read_budget_large_values_cannot_overflow_or_wait_for_impossible_permits() { + let budget = Arc::new(QueueReadBudget::new(usize::MAX, usize::MAX)); + let maximum = maximum_budget_bytes(); + assert_eq!(budget.snapshot().limit_bytes, maximum); + let (count, reservation) = budget.reserve(usize::MAX, 1).await.unwrap(); + assert_eq!(count, maximum); + assert_eq!(budget.snapshot().reserved_bytes, maximum); + drop(reservation); + assert!(matches!( + budget.reserve(1, maximum + 1).await, + Err(DataLayerError::InvalidConfiguration(_)) + )); + assert!(matches!( + budget.reserve(1, 0).await, + Err(DataLayerError::InvalidConfiguration(_)) + )); + let small_batch = Arc::new(QueueReadBudget::new(32, 4)); + let (count, reservation) = small_batch.reserve(usize::MAX, 16).await.unwrap(); + assert_eq!(count, 1); + assert_eq!(small_batch.snapshot().reserved_bytes, 16); + drop(reservation); + } + + #[test] + fn queue_read_budget_env_uses_positive_defaults_and_caps_extreme_values() { + for raw in [None, Some(""), Some("0"), Some("-1"), Some("bad")] { + assert_eq!(configured_bytes(raw, 128), 128); + } + assert_eq!(configured_bytes(Some(" 42 "), 128), 42); + assert_eq!( + configured_bytes(Some(&u128::MAX.to_string()), 128), + maximum_budget_bytes() + ); + } +} diff --git a/crates/aether-usage/runtime/src/record.rs b/crates/aether-usage/runtime/src/record.rs index 5b8ba7a20..ac46cc482 100644 --- a/crates/aether-usage/runtime/src/record.rs +++ b/crates/aether-usage/runtime/src/record.rs @@ -182,6 +182,7 @@ pub fn build_upsert_usage_record_from_event( finalized_at_unix_secs, created_at_unix_ms: Some(now_unix_secs), updated_at_unix_secs: now_unix_secs, + capture_retention: data.capture_retention, }) } @@ -262,6 +263,49 @@ mod tests { use super::build_upsert_usage_record_from_event; + #[test] + fn capture_retention_follows_event_bodies_into_record_and_its_clones() { + use aether_data_contracts::repository::usage::{ + usage_json_heap_estimate, UsageCaptureMemoryBudget, + }; + use std::sync::Arc; + + let body = serde_json::Value::String("retained diagnostic".repeat(8)); + let estimate = + 4 * (std::mem::size_of::() + usage_json_heap_estimate(&body)); + let budget = Arc::new(UsageCaptureMemoryBudget::new(3 * estimate)); + let mut event = UsageEvent::new( + UsageEventType::Completed, + "retained-record", + UsageEventData { + provider_name: "provider".to_owned(), + model: "model".to_owned(), + input_tokens: Some(5), + output_tokens: Some(7), + cache_read_input_tokens: Some(0), + request_body: Some(body.clone()), + provider_request_body: Some(body.clone()), + response_body: Some(body.clone()), + client_response_body: Some(body), + ..UsageEventData::default() + }, + ); + event.data.apply_capture_memory_budget(Arc::clone(&budget)); + assert_eq!(budget.retained_bytes(), estimate); + let record = build_upsert_usage_record_from_event(&event).unwrap(); + assert_eq!(budget.retained_bytes(), 2 * estimate); + drop(event); + assert_eq!(budget.retained_bytes(), estimate); + let cloned = record.clone(); + assert_eq!(budget.retained_bytes(), 2 * estimate); + assert_eq!(cloned.cache_read_input_tokens, Some(0)); + assert_eq!(cloned.response_body, record.response_body); + drop(record); + assert_eq!(budget.retained_bytes(), estimate); + drop(cloned); + assert_eq!(budget.retained_bytes(), 0); + } + #[test] fn builds_upsert_record_from_terminal_event() { let record = build_upsert_usage_record_from_event(&UsageEvent { diff --git a/crates/aether-usage/runtime/src/runtime.rs b/crates/aether-usage/runtime/src/runtime.rs index 5f26f893c..ffbc84733 100644 --- a/crates/aether-usage/runtime/src/runtime.rs +++ b/crates/aether-usage/runtime/src/runtime.rs @@ -14,21 +14,28 @@ use futures_util::{FutureExt, StreamExt}; use tokio::sync::mpsc; use tracing::{info, warn}; +use crate::event_capture_budget::{ + json_heap_estimate, EventCaptureMemoryBudget, UsageEventCaptureRetention, +}; use crate::executor::spawn_on_usage_background_runtime; +use crate::queue::is_permanent_enqueue_error; use crate::request_metadata::{ attach_client_request_body_metadata, attach_provider_request_body_metadata, attach_provider_response_body_metadata, clear_client_request_body_metadata, clear_provider_request_body_metadata, request_body_derived_facts_action, retain_first_byte_request_metadata, RequestBodyDerivedFactsAction, }; +use crate::settlement::{ + reconcile_usage_policy_cost_for_event_with_result, settle_usage_with_reconciled_cost, +}; +use crate::shutdown::{UsageBackgroundTasks, UsageShutdownState}; use crate::worker::{ build_usage_queue_worker_with_record_gate, UsageWorkerControl, UsageWorkerObservation, }; use crate::{ apply_usage_body_capture_policy_to_event, build_stream_terminal_usage_seed, build_sync_terminal_usage_seed, build_terminal_usage_event_from_seed, - build_upsert_usage_record_from_event, reconcile_usage_policy_cost_for_event, - settle_usage_if_needed, LifecycleUsageSeed, StreamTerminalUsagePayloadSeed, + build_upsert_usage_record_from_event, LifecycleUsageSeed, StreamTerminalUsagePayloadSeed, SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed, UsageEvent, UsageQueue, UsageRecordWriter, UsageRuntimeConfig, UsageSettlementWriter, }; @@ -87,6 +94,7 @@ pub trait UsageRuntimeAccess: #[derive(Debug, Clone)] pub struct UsageRuntime { config: UsageRuntimeConfig, + shutdown: Arc, body_policy_cache: Arc>>, enqueue_retry: Arc, worker_supervisor_state: Arc, @@ -451,6 +459,28 @@ impl UsageWorkerSupervisorState { #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub struct UsageRuntimeMetricsSnapshot { pub enabled: bool, + pub shutdown_started: bool, + pub producers_in_flight: usize, + pub delayed_lifecycle_pending: usize, + pub queue_payload_max_bytes: usize, + pub queue_payload_downgraded_total: u64, + pub queue_payload_rejected_total: u64, + pub queue_read_payload_budget_bytes: usize, + pub queue_read_batch_payload_bytes: usize, + pub queue_read_payload_reserved_bytes: usize, + pub queue_read_payload_waiters: usize, + pub queue_read_payload_wait_total: u64, + pub queue_read_actual_field_bytes_total: u64, + pub queue_read_oversized_entries_total: u64, + pub queue_read_oversized_batches_total: u64, + pub dlq_encoding_budget_bytes: usize, + pub dlq_encoding_max_jobs: usize, + pub dlq_encoding_reserved_bytes: usize, + pub dlq_encoding_active_jobs: usize, + pub dlq_encoding_capacity_rejected_total: u64, + pub dlq_encoding_oversized_rejected_total: u64, + pub dlq_encoding_encoded_total: u64, + pub enqueue_retry_permanent_failure_total: u64, pub queue_terminal_events: bool, pub queue_lifecycle_events: bool, pub worker_count: usize, @@ -536,6 +566,9 @@ pub struct UsageRuntimeMetricsSnapshot { pub enqueue_retry_pending: u64, pub enqueue_retry_failed_total: u64, pub enqueue_retry_closed_or_unavailable_total: u64, + pub event_capture_memory_budget_bytes: usize, + pub event_capture_memory_retained_bytes: usize, + pub event_capture_memory_downgraded_total: u64, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -1202,11 +1235,17 @@ enum LifecycleTerminalUsageSeed { Sync { context: TerminalUsageContextSeed, payload: SyncTerminalUsagePayloadSeed, + capture: Option, }, Stream { context: TerminalUsageContextSeed, payload: StreamTerminalUsagePayloadSeed, cancelled: bool, + capture: Option, + }, + Prepared { + kind: TerminalSeedKind, + result: Result, }, #[cfg(test)] BlockedBuild { @@ -1216,40 +1255,145 @@ enum LifecycleTerminalUsageSeed { }, } +#[derive(Clone, Copy)] +enum TerminalSeedKind { + Sync, + Stream, +} + +struct TerminalSeedCaptureRetention { + budget: Arc, + retention: UsageEventCaptureRetention, +} + +impl TerminalSeedCaptureRetention { + fn try_reserve( + bodies: [Option<&serde_json::Value>; 4], + budget: Arc, + ) -> Option { + let bytes = bodies.into_iter().flatten().fold(0usize, |bytes, body| { + bytes + .saturating_add(std::mem::size_of::()) + .saturating_add(json_heap_estimate(body)) + }); + let mut retention = UsageEventCaptureRetention::default(); + retention + .reserve(Arc::clone(&budget), bytes) + .then_some(Self { budget, retention }) + } + + fn attach(self, event: &mut UsageEvent) { + // The builder moves seed bodies into an unmanaged event. Transfer its existing + // reservation before resizing so concurrent seeds cannot claim the same bytes. + event.data.capture_retention = self.retention; + prepare_event_capture_memory(event, self.budget); + } +} + impl LifecycleTerminalUsageSeed { - async fn build(self, request_id: &str) -> Result { - match self { - Self::Sync { context, payload } => { - let result = build_sync_terminal_usage_event_offthread(context, payload).await; - if let Err(err) = &result { - warn!( - event_name = "usage_sync_terminal_build_failed", - log_type = "event", - request_id, - error = %err, - "usage runtime failed to build sync terminal usage event" - ); + fn prepare_capture_memory(self, budget: Arc) -> Self { + let (kind, result) = match self { + Self::Sync { + context, + payload, + capture: None, + } => { + let capture = TerminalSeedCaptureRetention::try_reserve( + [ + context.request_body.as_ref(), + context.provider_request.as_ref(), + payload.provider_response_full.as_ref(), + payload.client_response.as_ref(), + ], + Arc::clone(&budget), + ); + if capture.is_some() { + return Self::Sync { + context, + payload, + capture, + }; } - result + ( + TerminalSeedKind::Sync, + catch_unwind(AssertUnwindSafe(|| { + build_terminal_usage_event_from_seed(build_sync_terminal_usage_seed( + context, payload, + )) + })), + ) } Self::Stream { context, payload, cancelled, + capture: None, } => { - let result = - build_stream_terminal_usage_event_offthread(context, payload, cancelled).await; - if let Err(err) = &result { - warn!( - event_name = "usage_stream_terminal_build_failed", - log_type = "event", - request_id, - error = %err, - "usage runtime failed to build stream terminal usage event" - ); + let capture = TerminalSeedCaptureRetention::try_reserve( + [ + context.request_body.as_ref(), + context.provider_request.as_ref(), + payload.provider_response_full.as_ref(), + payload.client_response.as_ref(), + ], + Arc::clone(&budget), + ); + if capture.is_some() { + return Self::Stream { + context, + payload, + cancelled, + capture, + }; } - result + ( + TerminalSeedKind::Stream, + catch_unwind(AssertUnwindSafe(|| { + build_terminal_usage_event_from_seed(build_stream_terminal_usage_seed( + context, payload, cancelled, + )) + })), + ) } + prepared => return prepared, + }; + // These seeds still contain token fallbacks, image estimates and terminal + // evidence. Resolve the existing pure builder before discarding diagnostics. + // Keep failures queued so their ordering and error handling remain unchanged. + let result = result + .unwrap_or_else(|_| { + Err(DataLayerError::UnexpectedValue( + "usage builder panicked while preparing capture budget".to_string(), + )) + }) + .map(|mut event| { + prepare_event_capture_memory(&mut event, budget); + event + }); + Self::Prepared { kind, result } + } + + async fn build(self, request_id: &str) -> Result { + let (kind, result) = match self { + Self::Sync { + context, + payload, + capture, + } => ( + TerminalSeedKind::Sync, + build_sync_terminal_usage_event_offthread(context, payload, capture).await, + ), + Self::Stream { + context, + payload, + cancelled, + capture, + } => ( + TerminalSeedKind::Stream, + build_stream_terminal_usage_event_offthread(context, payload, cancelled, capture) + .await, + ), + Self::Prepared { kind, result } => (kind, result), #[cfg(test)] Self::BlockedBuild { event, @@ -1258,9 +1402,32 @@ impl LifecycleTerminalUsageSeed { } => { started.notify_one(); release.notified().await; - Ok(event) + return Ok(event); + } + }; + if let Err(err) = &result { + match kind { + TerminalSeedKind::Sync => { + warn!( + event_name = "usage_sync_terminal_build_failed", + log_type = "event", + request_id, + error = %err, + "usage runtime failed to build sync terminal usage event" + ); + } + TerminalSeedKind::Stream => { + warn!( + event_name = "usage_stream_terminal_build_failed", + log_type = "event", + request_id, + error = %err, + "usage runtime failed to build stream terminal usage event" + ); + } } } + result } } @@ -1673,7 +1840,7 @@ impl LifecycleSubmissionDispatcher { }) } - fn spawn(config: &UsageRuntimeConfig) -> Arc { + fn spawn(config: &UsageRuntimeConfig, tasks: &UsageBackgroundTasks) -> Arc { if !config.enabled { return Self::disabled(); } @@ -1690,7 +1857,7 @@ impl LifecycleSubmissionDispatcher { for _ in 0..workers { let (sender, receiver) = mpsc::unbounded_channel(); let slots = Arc::new(StdMutex::new(HashMap::new())); - spawn_on_usage_background_runtime(run_lifecycle_submission_worker( + tasks.spawn(run_lifecycle_submission_worker( receiver, Arc::clone(&slots), Arc::clone(&state), @@ -2187,10 +2354,11 @@ impl OrderedLifecycleDispatcher { fn spawn( config: &UsageRuntimeConfig, terminal_execution: Arc, + tasks: &UsageBackgroundTasks, ) -> Arc { let dispatcher = Self::disabled(terminal_execution); if config.enabled { - spawn_on_usage_background_runtime(run_ordered_lifecycle_dispatcher(Arc::downgrade( + tasks.spawn(run_ordered_lifecycle_dispatcher(Arc::downgrade( &dispatcher.core, ))); } @@ -2478,7 +2646,7 @@ impl PendingPersistenceDispatcher { }) } - fn spawn(config: &UsageRuntimeConfig) -> Arc { + fn spawn(config: &UsageRuntimeConfig, tasks: &UsageBackgroundTasks) -> Arc { if !config.enabled { return Self::disabled(); } @@ -2487,7 +2655,7 @@ impl PendingPersistenceDispatcher { .clamp(1_024, PENDING_PERSISTENCE_MAX_BUFFER); let state = Arc::new(PendingPersistenceState::new(capacity)); let (sender, receiver) = mpsc::channel(capacity); - spawn_on_usage_background_runtime(run_pending_persistence_dispatcher( + tasks.spawn(run_pending_persistence_dispatcher( receiver, Arc::clone(&state), )); @@ -3061,7 +3229,7 @@ impl FirstBytePersistenceDispatcher { }) } - fn spawn(config: &UsageRuntimeConfig) -> Arc { + fn spawn(config: &UsageRuntimeConfig, tasks: &UsageBackgroundTasks) -> Arc { if !config.enabled || !config.queue_lifecycle_events { return Self::disabled(); } @@ -3074,7 +3242,7 @@ impl FirstBytePersistenceDispatcher { .clamp(1, 256); let state = Arc::new(FirstBytePersistenceState::new(capacity)); let (sender, receiver) = mpsc::channel(capacity); - spawn_on_usage_background_runtime(run_first_byte_persistence_dispatcher( + tasks.spawn(run_first_byte_persistence_dispatcher( receiver, concurrency, Arc::clone(&state), @@ -3290,6 +3458,7 @@ impl UsageRuntime { let terminal_execution = TerminalExecutionDispatcher::disabled(); Self { config: UsageRuntimeConfig::disabled(), + shutdown: Arc::new(UsageShutdownState::default()), body_policy_cache: Arc::new(tokio::sync::Mutex::new(None)), enqueue_retry: UsageEnqueueRetryDispatcher::disabled(), worker_supervisor_state: Arc::new(UsageWorkerSupervisorState::default()), @@ -3312,7 +3481,8 @@ impl UsageRuntime { pub fn new(config: UsageRuntimeConfig) -> Result { config.validate()?; - let enqueue_retry = UsageEnqueueRetryDispatcher::spawn(config.clone()); + let shutdown = Arc::new(UsageShutdownState::default()); + let enqueue_retry = UsageEnqueueRetryDispatcher::spawn(config.clone(), &shutdown); let worker_record_gate = config .worker_record_concurrency_limit .map(UsageWorkerRecordConcurrencyGate::new) @@ -3322,17 +3492,22 @@ impl UsageRuntime { config.enqueue_retry_buffer_capacity, )); if config.enabled { - spawn_on_usage_background_runtime(run_lifecycle_coalescer_compactor(Arc::downgrade( - &lifecycle_coalescer, - ))); + shutdown + .tasks + .spawn(run_lifecycle_coalescer_compactor(Arc::downgrade( + &lifecycle_coalescer, + ))); } let terminal_submission_state = Arc::new(TerminalSubmissionState::new( terminal_submission_limit(&config), )); let terminal_execution = TerminalExecutionDispatcher::spawn(&config); - let ordered_lifecycle = - OrderedLifecycleDispatcher::spawn(&config, Arc::clone(&terminal_execution)); - let pending_persistence = PendingPersistenceDispatcher::spawn(&config); + let ordered_lifecycle = OrderedLifecycleDispatcher::spawn( + &config, + Arc::clone(&terminal_execution), + &shutdown.tasks, + ); + let pending_persistence = PendingPersistenceDispatcher::spawn(&config, &shutdown.tasks); let terminal_direct_fallback_state = Arc::new(TerminalDirectFallbackState::new( terminal_direct_fallback_limit(&config), )); @@ -3341,11 +3516,14 @@ impl UsageRuntime { Arc::clone(&lifecycle_coalescer), Arc::clone(&lifecycle_enqueue_state), Arc::clone(&enqueue_retry), + &shutdown, ); - let lifecycle_submission = LifecycleSubmissionDispatcher::spawn(&config); - let first_byte_persistence = FirstBytePersistenceDispatcher::spawn(&config); + let lifecycle_submission = LifecycleSubmissionDispatcher::spawn(&config, &shutdown.tasks); + let first_byte_persistence = + FirstBytePersistenceDispatcher::spawn(&config, &shutdown.tasks); Ok(Self { config, + shutdown, body_policy_cache: Arc::new(tokio::sync::Mutex::new(None)), enqueue_retry, worker_supervisor_state: Arc::new(UsageWorkerSupervisorState::default()), @@ -3368,10 +3546,155 @@ impl UsageRuntime { self.config.enabled } + /// Call before spawning a request finalizer that can outlive its HTTP body. + pub fn track_producer(&self) -> crate::UsageProducerGuard { + self.shutdown.producers.fetch_add(1, Ordering::AcqRel); + crate::UsageProducerGuard(Arc::clone(&self.shutdown.producers)) + } + + /// Stop request producers before calling. Queued Redis records remain durable + /// for the next consumer; local retry buffers must reach Redis or the database. + /// A timeout leaves the remaining work running so shutdown can be retried. + pub async fn shutdown(&self, timeout: Duration) -> Result<(), DataLayerError> { + self.shutdown_with_local_queue(timeout, None).await + } + + /// A process-local queue must also be consumed before stopping its workers. + pub async fn shutdown_with_local_queue( + &self, + timeout: Duration, + local_queue: Option>, + ) -> Result<(), DataLayerError> { + let drain = async { + let _lock = self.shutdown.lock.lock().await; + while self.shutdown.producers.load(Ordering::Acquire) != 0 { + tokio::time::sleep(Duration::from_millis(10)).await; + } + self.lifecycle_submission.state.admission.close(); + self.shutdown.drain.send_replace(true); + while self.local_work_pending() != 0 { + tokio::time::sleep(Duration::from_millis(10)).await; + } + if self.config.enabled { + if let Some(queue) = &local_queue { + loop { + let stats = queue + .stats(&self.config.stream_key, Some(&self.config.consumer_group)) + .await?; + if stats.stream_length == 0 + || (stats.group_pending == 0 && stats.group_lag == Some(0)) + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + if queue + .stats(&self.config.dlq_stream_key, None) + .await? + .stream_length + != 0 + { + return Err(DataLayerError::InvalidInput( + "local usage dead-letter queue must be recovered before shutdown" + .into(), + )); + } + } + } + self.terminal_submission_state.semaphore.close(); + self.shutdown.worker_control.request_shutdown(); + while self.shutdown.supervisors.load(Ordering::Acquire) != 0 { + tokio::time::sleep(Duration::from_millis(10)).await; + } + self.shutdown.tasks.stop_idle().await; + Ok(()) + }; + tokio::time::timeout(timeout, drain).await.map_err(|_| { + DataLayerError::TimedOut(format!( + "usage shutdown incomplete: producers={}, local_work={}, retry_pending={}, workers={}", + self.shutdown.producers.load(Ordering::Acquire), + self.local_work_pending(), + self.enqueue_retry.pending(), + self.shutdown.supervisors.load(Ordering::Acquire), + )) + })? + } + + fn local_work_pending(&self) -> u64 { + let snapshot = self.metrics_snapshot(); + // Admission lives through every ordered handoff, including gaps between + // per-stage gauges. Delayed events and retry buffers retain separate permits. + let admitted = self.lifecycle_submission.state.capacity.saturating_sub( + self.lifecycle_submission + .state + .admission + .available_permits(), + ); + [ + admitted as u64, + snapshot.delayed_lifecycle_pending as u64, + snapshot.lifecycle_submission_pending as u64, + snapshot.ordered_lifecycle_pending as u64, + snapshot.pending_persistence_pending as u64, + snapshot.first_byte_persistence_pending as u64, + snapshot.terminal_submission_pending as u64, + snapshot.terminal_submission_in_flight as u64, + snapshot.terminal_enqueue_in_flight, + snapshot.terminal_direct_fallback_in_flight as u64, + snapshot.lifecycle_enqueue_in_flight, + snapshot.enqueue_retry_pending, + ] + .into_iter() + .fold(0_u64, u64::saturating_add) + } + + fn track_worker(&self) -> crate::UsageProducerGuard { + self.shutdown.supervisors.fetch_add(1, Ordering::AcqRel); + crate::UsageProducerGuard(Arc::clone(&self.shutdown.supervisors)) + } + pub fn metrics_snapshot(&self) -> UsageRuntimeMetricsSnapshot { + let ( + event_capture_memory_budget_bytes, + event_capture_memory_retained_bytes, + event_capture_memory_downgraded_total, + ) = crate::event_capture_budget::capture_memory_metrics(); + let (queue_payload_downgraded_total, queue_payload_rejected_total) = + crate::queue::payload_encoding_totals(); + let queue_read = crate::queue_read_budget::queue_read_budget_metrics(); + let dlq_encoding = crate::dead_letter_encoding::dead_letter_encoding_metrics(); UsageRuntimeMetricsSnapshot { + queue_payload_max_bytes: self.config.queue_payload_max_bytes, + queue_payload_downgraded_total, + queue_payload_rejected_total, + queue_read_payload_budget_bytes: queue_read.limit_bytes, + queue_read_batch_payload_bytes: queue_read.batch_limit_bytes, + queue_read_payload_reserved_bytes: queue_read.reserved_bytes, + queue_read_payload_waiters: queue_read.waiters, + queue_read_payload_wait_total: queue_read.wait_total, + queue_read_actual_field_bytes_total: queue_read.actual_field_bytes_total, + queue_read_oversized_entries_total: queue_read.oversized_entries_total, + queue_read_oversized_batches_total: queue_read.oversized_batches_total, + dlq_encoding_budget_bytes: dlq_encoding.limit_bytes, + dlq_encoding_max_jobs: dlq_encoding.job_limit, + dlq_encoding_reserved_bytes: dlq_encoding.reserved_bytes, + dlq_encoding_active_jobs: dlq_encoding.active_jobs, + dlq_encoding_capacity_rejected_total: dlq_encoding.capacity_rejected_total, + dlq_encoding_oversized_rejected_total: dlq_encoding.oversized_rejected_total, + dlq_encoding_encoded_total: dlq_encoding.encoded_total, + enqueue_retry_permanent_failure_total: self.enqueue_retry.permanent_failure_total(), + event_capture_memory_budget_bytes, + event_capture_memory_retained_bytes, + event_capture_memory_downgraded_total, enabled: self.config.enabled, queue_terminal_events: self.config.queue_terminal_events, + shutdown_started: *self.shutdown.drain.borrow(), + producers_in_flight: self.shutdown.producers.load(Ordering::Acquire), + delayed_lifecycle_pending: self.lifecycle_delay.sender.as_ref().map_or(0, |sender| { + sender + .max_capacity() + .saturating_sub(self.lifecycle_delay.admission.available_permits()) + }), queue_lifecycle_events: self.config.queue_lifecycle_events, worker_count: self.config.worker_count, worker_autoscale_enabled: self.config.worker_autoscale_enabled, @@ -3716,7 +4039,12 @@ impl UsageRuntime { None, ) .ok()?; - Some(worker.spawn()) + let worker = worker.with_shutdown(self.shutdown.worker_control.clone()); + let guard = self.track_worker(); + Some(spawn_on_usage_background_runtime(async move { + let _guard = guard; + worker.run().await; + })) } pub fn spawn_workers(&self, data: Arc) -> Vec> @@ -3748,7 +4076,12 @@ impl UsageRuntime { ); continue; }; - handles.push(worker.spawn()); + let worker = worker.with_shutdown(self.shutdown.worker_control.clone()); + let guard = self.track_worker(); + handles.push(spawn_on_usage_background_runtime(async move { + let _guard = guard; + worker.run().await; + })); } handles } @@ -3761,15 +4094,20 @@ impl UsageRuntime { return None; } let runner = data.usage_worker_queue()?; - Some(spawn_on_usage_background_runtime( + let runtime = self.clone(); + let guard = self.track_worker(); + Some(spawn_on_usage_background_runtime(async move { + let _guard = guard; run_usage_worker_supervisor( runner, data, - self.config.clone(), - self.worker_record_gate.clone(), - Arc::clone(&self.worker_supervisor_state), - ), - )) + runtime.config.clone(), + runtime.worker_record_gate.clone(), + Arc::clone(&runtime.worker_supervisor_state), + runtime.shutdown.worker_control.clone(), + ) + .await; + })) } pub fn record_pending(&self, data: &T, seed: LifecycleUsageSeed) @@ -3836,12 +4174,16 @@ impl UsageRuntime { async fn dispatch_terminal( &self, data: &T, - event: UsageEvent, + mut event: UsageEvent, direct: bool, completion: Option>, ) where T: UsageRuntimeAccess + Clone + 'static, { + prepare_event_capture_memory( + &mut event, + crate::event_capture_budget::shared_capture_memory_budget(), + ); let request_id = event.request_id.clone(); self.lifecycle_submission .dispatch_terminal(Box::new(LifecycleSubmissionItemImpl { @@ -3865,6 +4207,9 @@ impl UsageRuntime { ) where T: UsageRuntimeAccess + Clone + 'static, { + let observed_at_unix_ms = now_unix_ms(); + let seed = seed + .prepare_capture_memory(crate::event_capture_budget::shared_capture_memory_budget()); self.lifecycle_submission .dispatch_terminal(Box::new(LifecycleSubmissionItemImpl { runtime: self.clone(), @@ -3872,7 +4217,7 @@ impl UsageRuntime { request_id, payload: LifecycleSubmissionPayload::TerminalSeed { seed, - observed_at_unix_ms: now_unix_ms(), + observed_at_unix_ms, }, })) .await; @@ -4018,6 +4363,7 @@ impl UsageRuntime { LifecycleTerminalUsageSeed::Sync { context: context_seed, payload: payload_seed, + capture: None, }, ) .await; @@ -4043,6 +4389,7 @@ impl UsageRuntime { context: context_seed, payload: payload_seed, cancelled, + capture: None, }, ) .await; @@ -4062,9 +4409,13 @@ impl UsageRuntime { where T: UsageRuntimeAccess, { - if !self.is_enabled() { + if !self.is_enabled() || self.lifecycle_submission.state.admission.is_closed() { return; } + prepare_event_capture_memory( + &mut event, + crate::event_capture_budget::shared_capture_memory_budget(), + ); let ordered_completion = self .await_lifecycle_submission_turn(&event.request_id) .await; @@ -4086,9 +4437,13 @@ impl UsageRuntime { where T: UsageRuntimeAccess, { - if !self.is_enabled() { + if !self.is_enabled() || self.lifecycle_submission.state.admission.is_closed() { return; } + prepare_event_capture_memory( + &mut event, + crate::event_capture_budget::shared_capture_memory_budget(), + ); let ordered_completion = self .await_lifecycle_submission_turn(&event.request_id) .await; @@ -4098,14 +4453,8 @@ impl UsageRuntime { }; self.apply_body_capture_policy_from_data(data, &mut event) .await; - if let Err(err) = data.enrich_usage_event(&mut event).await { - warn!( - event_name = "usage_terminal_billing_enrichment_failed", - log_type = "event", - request_id = %event.request_id, - error = %err, - "usage runtime failed to enrich terminal usage event with billing" - ); + if enrich_terminal_event(data, &mut event).await.is_err() { + return; } let request_id = event.request_id.clone(); if self.write_event_direct(data, &event).await { @@ -4122,8 +4471,25 @@ impl UsageRuntime { where T: UsageRuntimeAccess, { - preserve_request_facts(event); - preserve_provider_response_facts(event); + self.apply_body_capture_policy_with_budget( + data, + event, + crate::event_capture_budget::shared_capture_memory_budget(), + ) + .await; + } + + async fn apply_body_capture_policy_with_budget( + &self, + data: &T, + event: &mut UsageEvent, + budget: Arc, + ) where + T: UsageRuntimeAccess, + { + // A slow policy read must not retain unbudgeted JSON in mutex waiters. + // Denied captures count even when the eventual Basic policy disables them. + prepare_event_capture_memory(event, Arc::clone(&budget)); match self.cached_body_capture_policy(data).await { Ok(policy) => apply_usage_body_capture_policy_to_event(policy, event), Err(err) => { @@ -4138,6 +4504,7 @@ impl UsageRuntime { apply_usage_body_capture_policy_to_event(UsageBodyCapturePolicy::default(), event); } } + event.data.apply_capture_memory_budget(budget); } pub async fn body_capture_policy_for( @@ -4306,14 +4673,8 @@ impl UsageRuntime { return self.enqueue_or_write_terminal(data, event).await; } - if let Err(err) = data.enrich_usage_event(&mut event).await { - warn!( - event_name = "usage_terminal_billing_enrichment_failed", - log_type = "event", - request_id = %event.request_id, - error = %err, - "usage runtime failed to enrich ordered direct terminal usage event" - ); + if enrich_terminal_event(data, &mut event).await.is_err() { + return TerminalPersistenceOutcome::Failed; } let request_id = event.request_id.clone(); if self.write_event_direct(data, &event).await { @@ -4524,7 +4885,7 @@ impl UsageRuntime { let usage_event_type = event.event_type; let request_id = event.request_id.clone(); let direct_write_succeeded = self - .try_write_terminal_direct_fallback(data, &mut event) + .try_write_terminal_direct_fallback(data, &mut event, "bounded_local_enqueue_retry") .await; let deferred_fallback = if direct_write_succeeded { DeferredEnqueueFallback::DirectWrite @@ -4544,8 +4905,8 @@ impl UsageRuntime { TerminalPersistenceOutcome::Failed }; } - if event_phase == "terminal" { - enrich_terminal_event(data, &mut event).await; + if event_phase == "terminal" && enrich_terminal_event(data, &mut event).await.is_err() { + return TerminalPersistenceOutcome::Failed; } if self.write_event_direct(data, &event).await { TerminalPersistenceOutcome::PersistedDirectly @@ -4593,6 +4954,11 @@ impl UsageRuntime { if let Err(err) = queue.enqueue(&event).await { drop(_guard); + if is_permanent_enqueue_error(&err) { + return self + .defer_terminal_event(data, queue, event, "invalid_input", err) + .await; + } self.terminal_enqueue_state .open_circuit(now_unix_ms().saturating_add(LIFECYCLE_ENQUEUE_CIRCUIT_OPEN_MS)); let failures = self.terminal_enqueue_state.increment_failed_total(); @@ -4629,8 +4995,13 @@ impl UsageRuntime { { let usage_event_type = event.event_type; let request_id = event.request_id.clone(); + let fallback = if is_permanent_enqueue_error(&cause) { + "report_failure" + } else { + "bounded_local_enqueue_retry" + }; let direct_write_succeeded = self - .try_write_terminal_direct_fallback(data, &mut event) + .try_write_terminal_direct_fallback(data, &mut event, fallback) .await; let (deferred_fallback, outcome) = if direct_write_succeeded { ( @@ -4658,7 +5029,12 @@ impl UsageRuntime { outcome } - async fn try_write_terminal_direct_fallback(&self, data: &T, event: &mut UsageEvent) -> bool + async fn try_write_terminal_direct_fallback( + &self, + data: &T, + event: &mut UsageEvent, + fallback: &'static str, + ) -> bool where T: UsageRuntimeAccess, { @@ -4671,7 +5047,7 @@ impl UsageRuntime { usage_event_type = ?event.event_type, request_id = %event.request_id, rejected_total = rejected, - fallback = "bounded_local_enqueue_retry", + fallback, "usage runtime skipped terminal direct fallback because the writer is unavailable or under pressure" ); } @@ -4690,7 +5066,7 @@ impl UsageRuntime { request_id = %event.request_id, worker_record_limit = gate.limit(), rejected_total = rejected, - fallback = "bounded_local_enqueue_retry", + fallback, "usage runtime terminal direct fallback was rejected by the shared database concurrency gate" ); } @@ -4711,7 +5087,7 @@ impl UsageRuntime { request_id = %event.request_id, fallback_limit = self.terminal_direct_fallback_state.limit(), rejected_total = rejected, - fallback = "bounded_local_enqueue_retry", + fallback, "usage runtime terminal direct fallback is saturated" ); } @@ -4725,7 +5101,7 @@ impl UsageRuntime { usage_event_type = ?event.event_type, request_id = %event.request_id, error = %err, - fallback = "bounded_local_enqueue_retry", + fallback, "usage runtime could not enrich terminal event for direct fallback" ); false @@ -4754,7 +5130,7 @@ impl UsageRuntime { usage_event_type = ?event.event_type, request_id = %event.request_id, failed_total = failed, - fallback = "bounded_local_enqueue_retry", + fallback, "usage runtime terminal direct fallback failed" ); } @@ -4766,17 +5142,21 @@ impl UsageRuntime { where T: UsageRuntimeAccess, { - if let Err(err) = reconcile_usage_policy_cost_for_event(data, event).await { - warn!( - event_name = "usage_event_cost_reconciliation_failed", - log_type = "event", - usage_event_type = ?event.event_type, - request_id = %event.request_id, - error = %err, - "usage runtime failed to reconcile plan cost before direct usage upsert" - ); - return false; - } + let reconciled = match reconcile_usage_policy_cost_for_event_with_result(data, event).await + { + Ok(reconciled) => reconciled, + Err(err) => { + warn!( + event_name = "usage_event_cost_reconciliation_failed", + log_type = "event", + usage_event_type = ?event.event_type, + request_id = %event.request_id, + error = %err, + "usage runtime failed to reconcile plan cost before direct usage upsert" + ); + return false; + } + }; match build_upsert_usage_record_from_event(event) { Ok(record) => match catch_usage_writer_panic( "direct usage upsert", @@ -4785,7 +5165,9 @@ impl UsageRuntime { .await { Ok(Some(stored)) => { - if let Err(err) = settle_usage_if_needed(data, &stored).await { + if let Err(err) = + settle_usage_with_reconciled_cost(data, &stored, reconciled).await + { warn!( event_name = "usage_terminal_settlement_failed", log_type = "event", @@ -4825,7 +5207,32 @@ impl UsageRuntime { } } +pub(crate) fn prepare_event_capture_memory( + event: &mut UsageEvent, + budget: Arc, +) { + preserve_request_facts(event); + preserve_provider_response_facts(event); + event.data.apply_capture_memory_budget(budget); +} + +pub(crate) fn prepare_decoded_event_capture_memory( + event: &mut UsageEvent, + budget: Arc, +) { + preserve_request_facts_with_legacy_missing(event, true); + preserve_provider_response_facts(event); + event.data.apply_capture_memory_budget(budget); +} + fn preserve_request_facts(event: &mut UsageEvent) { + preserve_request_facts_with_legacy_missing(event, false); +} + +fn preserve_request_facts_with_legacy_missing( + event: &mut UsageEvent, + preserve_implicit_missing: bool, +) { let data = &mut event.data; match request_body_derived_facts_action(data.request_body.as_ref(), data.request_body_state) { RequestBodyDerivedFactsAction::Refresh => { @@ -4835,8 +5242,15 @@ fn preserve_request_facts(event: &mut UsageEvent) { ); } RequestBodyDerivedFactsAction::Clear => { - data.request_metadata = - clear_client_request_body_metadata(data.request_metadata.take()); + // Legacy wire events can carry derived facts without either capture field. + // An explicit typed `none` still authoritatively clears those facts. + if !preserve_implicit_missing + || data.request_body.is_some() + || data.request_body_state.is_some() + { + data.request_metadata = + clear_client_request_body_metadata(data.request_metadata.take()); + } } RequestBodyDerivedFactsAction::Preserve => {} } @@ -4856,8 +5270,13 @@ fn preserve_request_facts(event: &mut UsageEvent) { ); } RequestBodyDerivedFactsAction::Clear => { - data.request_metadata = - clear_provider_request_body_metadata(data.request_metadata.take()); + if !preserve_implicit_missing + || data.provider_request_body.is_some() + || data.provider_request_body_state.is_some() + { + data.request_metadata = + clear_provider_request_body_metadata(data.request_metadata.take()); + } } RequestBodyDerivedFactsAction::Preserve => {} } @@ -4878,7 +5297,7 @@ impl UsageQueueHealthSnapshot { } } -async fn enrich_terminal_event(data: &T, event: &mut UsageEvent) +async fn enrich_terminal_event(data: &T, event: &mut UsageEvent) -> Result<(), DataLayerError> where T: UsageBillingEventEnricher + Send + Sync, { @@ -4890,7 +5309,9 @@ where error = %err, "usage runtime failed to enrich terminal usage event with billing" ); + return Err(err); } + Ok(()) } struct ManagedUsageWorker { @@ -4919,6 +5340,7 @@ async fn run_usage_worker_supervisor( config: UsageRuntimeConfig, worker_record_gate: Option>, state: Arc, + control: UsageWorkerControl, ) where T: UsageRuntimeAccess + 'static, { @@ -4968,6 +5390,16 @@ async fn run_usage_worker_supervisor( loop { tokio::select! { + biased; + _ = control.wait_for_shutdown() => { + state.desired_count.store(0, Ordering::Release); + for worker in workers.values() { + worker.control.request_shutdown(); + } + while join_set.join_next().await.is_some() {} + state.active_count.store(0, Ordering::Release); + break; + } Some(observation) = telemetry_rx.recv() => { state.record_observation(observation); if observation.entries_read == 0 { @@ -5407,6 +5839,7 @@ impl LifecycleDelayDispatcher { coalescer: Arc, enqueue_state: Arc, enqueue_retry: Arc, + shutdown: &UsageShutdownState, ) -> Arc { if !config.enabled || !config.queue_lifecycle_events @@ -5419,12 +5852,13 @@ impl LifecycleDelayDispatcher { let delay = Duration::from_millis(config.lifecycle_enqueue_delay_ms.max(1)); let admission = Arc::new(tokio::sync::Semaphore::new(capacity)); let (sender, receiver) = mpsc::channel(capacity); - spawn_on_usage_background_runtime(run_lifecycle_delay_worker( + shutdown.tasks.spawn(run_lifecycle_delay_worker( config, coalescer, enqueue_state, enqueue_retry, receiver, + shutdown.drain.subscribe(), )); Arc::new(Self { delay, @@ -5470,11 +5904,31 @@ async fn run_lifecycle_delay_worker( enqueue_state: Arc, enqueue_retry: Arc, mut receiver: mpsc::Receiver, + mut drain: tokio::sync::watch::Receiver, ) { let mut pending = BTreeMap::>::new(); let mut receiver_open = true; loop { + if *drain.borrow_and_update() { + receiver.close(); + while let Ok(item) = receiver.try_recv() { + pending.entry(item.due_at).or_default().push(item); + } + for (_, items) in std::mem::take(&mut pending) { + for item in items { + item.item + .enqueue( + config.clone(), + Arc::clone(&coalescer), + Arc::clone(&enqueue_state), + Arc::clone(&enqueue_retry), + ) + .await; + } + } + break; + } if !pending.is_empty() { enqueue_due_lifecycle_items( &mut pending, @@ -5491,7 +5945,11 @@ async fn run_lifecycle_delay_worker( if !receiver_open { break; } - match receiver.recv().await { + let next = tokio::select! { + next = receiver.recv() => next, + _ = crate::shutdown::wait_for_drain(&mut drain) => continue, + }; + match next { Some(item) => { pending.entry(item.due_at).or_default().push(item); continue; @@ -5507,6 +5965,7 @@ async fn run_lifecycle_delay_worker( if receiver_open { tokio::select! { + _ = crate::shutdown::wait_for_drain(&mut drain) => continue, maybe_item = receiver.recv() => { match maybe_item { Some(item) => { @@ -5685,6 +6144,17 @@ where return true; }; + if is_permanent_enqueue_error(&err) { + enqueue_state.record_deferred( + "usage_lifecycle_event_enqueue_deferred", + "invalid_input", + event.event_type, + &event.request_id, + DeferredEnqueueFallback::Drop, + ); + return enqueue_retry.schedule(queue, event, "lifecycle", err); + } + enqueue_state.open_circuit(now_unix_ms().saturating_add(LIFECYCLE_ENQUEUE_CIRCUIT_OPEN_MS)); let failures = enqueue_state.increment_failed_total(); let retry_enabled = config.retry_deferred_lifecycle_events; @@ -5727,6 +6197,7 @@ struct UsageEnqueueRetryDispatcher { #[derive(Debug, Default)] struct UsageEnqueueDispatcherMetrics { scheduled_total: AtomicU64, + permanent_failure_total: AtomicU64, recovered_total: AtomicU64, pending: AtomicU64, retry_failed_total: AtomicU64, @@ -5749,7 +6220,7 @@ impl UsageEnqueueRetryDispatcher { }) } - fn spawn(config: UsageRuntimeConfig) -> Arc { + fn spawn(config: UsageRuntimeConfig, shutdown: &UsageShutdownState) -> Arc { let lifecycle_retry_enabled = config.queue_lifecycle_events && config.retry_deferred_lifecycle_events; if !config.enabled || !(config.queue_terminal_events || lifecycle_retry_enabled) { @@ -5769,12 +6240,14 @@ impl UsageEnqueueRetryDispatcher { senders.push(sender); let worker_config = config.clone(); let worker_metrics = Arc::clone(&metrics); - spawn_on_usage_background_runtime(async move { - run_usage_enqueue_retry_worker( + let drain = shutdown.drain.subscribe(); + shutdown.tasks.spawn(async move { + run_usage_enqueue_retry_worker_with_drain( worker_index, worker_config, receiver, worker_metrics, + drain, ) .await; }); @@ -5788,8 +6261,32 @@ impl UsageEnqueueRetryDispatcher { queue: UsageQueue, event: UsageEvent, event_phase: &'static str, - cause: DataLayerError, + mut cause: DataLayerError, ) -> bool { + // Circuit and admission failures can reach this path without an encoding attempt. + if !is_permanent_enqueue_error(&cause) { + if let Err(error) = queue.validate_event(&event) { + if is_permanent_enqueue_error(&error) { + cause = error; + } + } + } + if is_permanent_enqueue_error(&cause) { + let rejected = self.metrics.record_permanent_failure(); + if should_log_usage_retry_counter(rejected) { + warn!( + event_name = "usage_event_enqueue_invalid_input", + log_type = "ops", + event_phase, + usage_event_type = ?event.event_type, + request_id = %event.request_id, + rejected_total = rejected, + error = %cause, + "usage event could not be persisted; invalid queue input cannot be retried" + ); + } + return false; + } let event_type = event.event_type; let request_id = event.request_id.clone(); let cause_message = cause.to_string(); @@ -5887,6 +6384,10 @@ impl UsageEnqueueRetryDispatcher { self.metrics.scheduled_total.load(Ordering::Acquire) } + fn permanent_failure_total(&self) -> u64 { + self.metrics.permanent_failure_total.load(Ordering::Acquire) + } + fn recovered_total(&self) -> u64 { self.metrics.recovered_total.load(Ordering::Acquire) } @@ -5907,6 +6408,10 @@ impl UsageEnqueueRetryDispatcher { } impl UsageEnqueueDispatcherMetrics { + fn record_permanent_failure(&self) -> u64 { + self.permanent_failure_total.fetch_add(1, Ordering::AcqRel) + 1 + } + fn record_scheduled(&self) -> u64 { self.pending.fetch_add(1, Ordering::AcqRel); self.scheduled_total.fetch_add(1, Ordering::AcqRel) + 1 @@ -5929,17 +6434,32 @@ impl UsageEnqueueDispatcherMetrics { } } +#[cfg(test)] async fn run_usage_enqueue_retry_worker( + worker_index: usize, + config: UsageRuntimeConfig, + receiver: mpsc::Receiver, + metrics: Arc, +) { + let (_sender, drain) = tokio::sync::watch::channel(false); + run_usage_enqueue_retry_worker_with_drain(worker_index, config, receiver, metrics, drain).await; +} + +async fn run_usage_enqueue_retry_worker_with_drain( worker_index: usize, config: UsageRuntimeConfig, mut receiver: mpsc::Receiver, metrics: Arc, + mut drain: tokio::sync::watch::Receiver, ) { let mut initial_retry_delay_applied = false; while let Some(mut item) = receiver.recv().await { if item.delay_before_first_attempt && !initial_retry_delay_applied { initial_retry_delay_applied = true; - tokio::time::sleep(usage_enqueue_retry_delay(&config, 1)).await; + if !*drain.borrow() { + crate::shutdown::retry_delay(usage_enqueue_retry_delay(&config, 1), &mut drain) + .await; + } } loop { match item.queue.enqueue(&item.event).await { @@ -5960,6 +6480,24 @@ async fn run_usage_enqueue_retry_worker( } break; } + Err(err) if is_permanent_enqueue_error(&err) => { + let rejected = metrics.record_permanent_failure(); + metrics.pending.fetch_sub(1, Ordering::AcqRel); + if should_log_usage_retry_counter(rejected) { + warn!( + event_name = "usage_event_enqueue_retry_invalid_input", + log_type = "ops", + event_phase = item.event_phase, + usage_event_type = ?item.event.event_type, + request_id = %item.event.request_id, + worker_index, + rejected_total = rejected, + error = %err, + "usage enqueue retry ended for invalid input; advancing to the next event" + ); + } + break; + } Err(err) => { item.attempts = item.attempts.saturating_add(1); metrics.record_retry_failed(); @@ -5978,7 +6516,7 @@ async fn run_usage_enqueue_retry_worker( "usage runtime local enqueue retry failed; will retry" ); } - tokio::time::sleep(delay).await; + crate::shutdown::retry_delay(delay, &mut drain).await; } } } @@ -6043,23 +6581,24 @@ fn should_log_usage_retry_counter(value: u64) -> bool { async fn build_sync_terminal_usage_event_offthread( context_seed: TerminalUsageContextSeed, payload_seed: SyncTerminalUsagePayloadSeed, + capture: Option, ) -> Result { - tokio::task::spawn_blocking(move || { + build_terminal_usage_event_offthread(capture, move || { build_terminal_usage_event_from_seed(build_sync_terminal_usage_seed( context_seed, payload_seed, )) }) .await - .map_err(join_error_to_data_layer)? } async fn build_stream_terminal_usage_event_offthread( context_seed: TerminalUsageContextSeed, payload_seed: StreamTerminalUsagePayloadSeed, cancelled: bool, + capture: Option, ) -> Result { - tokio::task::spawn_blocking(move || { + build_terminal_usage_event_offthread(capture, move || { build_terminal_usage_event_from_seed(build_stream_terminal_usage_seed( context_seed, payload_seed, @@ -6067,6 +6606,20 @@ async fn build_stream_terminal_usage_event_offthread( )) }) .await +} + +async fn build_terminal_usage_event_offthread( + capture: Option, + build: impl FnOnce() -> Result + Send + 'static, +) -> Result { + tokio::task::spawn_blocking(move || { + let mut event = build()?; + if let Some(capture) = capture { + capture.attach(&mut event); + } + Ok(event) + }) + .await .map_err(join_error_to_data_layer)? } @@ -6087,6 +6640,14 @@ fn now_unix_ms() -> u64 { #[cfg(test)] mod tests { + mod shutdown { + include!("runtime_shutdown_tests.rs"); + } + + mod queue_payload { + include!("runtime_queue_payload_tests.rs"); + } + use std::collections::BTreeMap; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; @@ -6168,6 +6729,425 @@ mod tests { ) } + fn terminal_seed_capture_test_seeds( + request_id: &str, + padding_bytes: usize, + ) -> (TerminalUsageContextSeed, SyncTerminalUsagePayloadSeed) { + let (_, mut payload) = sync_terminal_test_seeds(request_id); + let provider_request = json!({ + "model": "gpt-5.6-sol", "reasoning": {"effort": "medium"}, + "service_tier": "priority", "prompt": "retained request facts" + }); + let mut plan = terminal_test_plan(request_id); + plan.model_name = Some("gpt-5.6-sol".to_string()); + plan.body = RequestBody::from_json(provider_request.clone()); + let report_context = json!({ + "original_request_body": {"reasoning": {"effort": "high"}}, + "provider_request_body": provider_request, + }); + let context = build_terminal_usage_context_seed(&plan, Some(&report_context)); + payload.provider_response_full = Some(json!({ + "id": request_id, + "output": "x".repeat(padding_bytes), + "service_tier": "default", + "usage": {"input_tokens": 100, "output_tokens": 500, "total_tokens": 600} + })); + (context, payload) + } + + #[tokio::test] + async fn terminal_seed_capture_budget_zero_preserves_sync_facts_and_explicit_zero() { + for tokens in [0, 100] { + let (context, mut payload) = terminal_seed_capture_test_seeds("seed-budget-sync", 4096); + payload.provider_response_full.as_mut().unwrap()["usage"] = json!({ + "input_tokens": tokens, "output_tokens": tokens, "total_tokens": 2 * tokens + }); + let mut expected = crate::build_terminal_usage_event_from_seed( + crate::build_sync_terminal_usage_seed(context.clone(), payload.clone()), + ) + .expect("original terminal event"); + let budget = Arc::new(super::EventCaptureMemoryBudget::new(0)); + let prepared = LifecycleTerminalUsageSeed::Sync { + context, + payload, + capture: None, + } + .prepare_capture_memory(Arc::clone(&budget)); + assert!(matches!( + prepared, + LifecycleTerminalUsageSeed::Prepared { .. } + )); + let event = prepared + .build("seed-budget-sync") + .await + .expect("prepared event"); + super::prepare_event_capture_memory(&mut expected, Arc::clone(&budget)); + assert_eq!(event.event_type, expected.event_type); + assert_eq!(event.data, expected.data); + assert_eq!(event.data.input_tokens, Some(tokens)); + assert_eq!(event.data.output_tokens, Some(tokens)); + assert_eq!(event.data.total_tokens, Some(2 * tokens)); + assert!(event.data.request_body.is_none()); + assert!(event.data.provider_request_body.is_none()); + assert!(event.data.response_body.is_none()); + assert!(event.data.client_response_body.is_none()); + assert_eq!( + event.data.response_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + let metadata = event.data.request_metadata.as_ref().expect("billing facts"); + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_reasoning_effort"], "medium"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 30); + assert_eq!(budget.retained_bytes(), 0); + } + } + + #[tokio::test] + async fn terminal_seed_capture_budget_zero_preserves_image_estimates_and_dimensions() { + let (mut context, mut payload) = terminal_seed_capture_test_seeds("seed-budget-image", 0); + context.client_contract = "openai:image".to_string(); + context.provider_contract = "openai:image".to_string(); + context.request_type = "image".to_string(); + context.provider_request = Some(json!({ + "model": "gpt-image-2", "prompt": "Draw a red kite", "n": 1, + "size": "1024x1024", "quality": "high", "output_format": "png" + })); + payload.provider_response_full = Some(json!({ + "data": [{"b64_json": "x".repeat(4096)}, {"b64_json": "y".repeat(4096)}] + })); + let expected = crate::build_terminal_usage_event_from_seed( + crate::build_sync_terminal_usage_seed(context.clone(), payload.clone()), + ) + .expect("original image event"); + let budget = Arc::new(super::EventCaptureMemoryBudget::new(0)); + let event = LifecycleTerminalUsageSeed::Sync { + context, + payload, + capture: None, + } + .prepare_capture_memory(Arc::clone(&budget)) + .build("seed-budget-image") + .await + .expect("image event"); + assert_eq!(event.event_type, UsageEventType::Completed); + assert_eq!(event.data.input_tokens, expected.data.input_tokens); + assert!(event.data.input_tokens.unwrap_or_default() > 0); + assert_eq!(event.data.total_tokens, expected.data.total_tokens); + let metadata = event + .data + .request_metadata + .as_ref() + .expect("image dimensions"); + assert_eq!(metadata["dimensions"]["image_count"], 2); + assert_eq!(metadata["dimensions"]["image_size"], "1024x1024"); + assert_eq!(metadata["dimensions"]["image_quality"], "high"); + assert_eq!(metadata["dimensions"]["image_output_format"], "png"); + assert!(event.data.response_body.is_none()); + assert!(event.data.provider_request_body.is_none()); + assert_eq!(budget.retained_bytes(), 0); + } + + #[tokio::test] + async fn terminal_seed_capture_budget_zero_preserves_stream_terminal_evidence() { + for (failed, cancelled) in [(false, false), (true, false), (false, true)] { + let (mut context, _) = terminal_seed_capture_test_seeds("seed-budget-stream", 0); + context.is_stream = true; + let provider_response_full = Some(json!({ + "chunks": [{ + "type": if failed { "response.failed" } else { "response.completed" }, + "response": { + "status": if failed { "failed" } else { "completed" }, + "service_tier": "default", + "usage": {"input_tokens": 100, "output_tokens": 500, "total_tokens": 600}, + "error": if failed { json!({"message": "provider refused"}) } else { json!(null) } + } + }] + })); + let payload = crate::StreamTerminalUsagePayloadSeed { + report_kind: "openai_responses_stream_success".to_string(), + status_code: if cancelled { 499 } else { 200 }, + response_time_ms: Some(12), + first_byte_time_ms: Some(3), + provider_response_headers: None, + client_response_headers: None, + provider_response_full, + provider_response_body_state: Some(UsageBodyCaptureState::Inline), + client_response: None, + client_response_body_state: Some(UsageBodyCaptureState::None), + standardized_usage: None, + provider_actual_service_tier: Some("default".to_string()), + observed_stream_finish: Some(true), + terminal_error_message: None, + capture_metadata: None, + }; + let mut expected = crate::build_terminal_usage_event_from_seed( + crate::build_stream_terminal_usage_seed( + context.clone(), + payload.clone(), + cancelled, + ), + ) + .expect("original stream event"); + let budget = Arc::new(super::EventCaptureMemoryBudget::new(0)); + let event = LifecycleTerminalUsageSeed::Stream { + context, + payload, + cancelled, + capture: None, + } + .prepare_capture_memory(Arc::clone(&budget)) + .build("seed-budget-stream") + .await + .expect("stream event"); + super::prepare_event_capture_memory(&mut expected, Arc::clone(&budget)); + assert_eq!(event.event_type, expected.event_type); + assert_eq!(event.data, expected.data); + assert_eq!(event.data.input_tokens, Some(100)); + assert_eq!(event.data.output_tokens, Some(500)); + assert_eq!(event.data.first_byte_time_ms, Some(3)); + assert_eq!( + event.event_type, + if cancelled { + UsageEventType::Cancelled + } else if failed { + UsageEventType::Failed + } else { + UsageEventType::Completed + } + ); + if failed { + assert_eq!( + event.data.error_message.as_deref(), + Some("provider refused") + ); + } + assert!(event.data.response_body.is_none()); + assert_eq!(budget.retained_bytes(), 0); + } + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn terminal_seed_capture_budget_bounds_concurrent_seeds_and_transfers_to_events() { + const COUNT: usize = 12; + const LIMIT: usize = 96 * 1024; + let budget = Arc::new(super::EventCaptureMemoryBudget::new(LIMIT)); + let barrier = Arc::new(tokio::sync::Barrier::new(COUNT)); + let mut tasks = Vec::new(); + for index in 0..COUNT { + let (context, payload) = + terminal_seed_capture_test_seeds(&format!("seed-budget-{index}"), 32 * 1024); + let budget = Arc::clone(&budget); + let barrier = Arc::clone(&barrier); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + LifecycleTerminalUsageSeed::Sync { + context, + payload, + capture: None, + } + .prepare_capture_memory(budget) + })); + } + let mut seeds = Vec::new(); + for task in tasks { + seeds.push(task.await.expect("concurrent seed preparation")); + } + let retained = budget.retained_bytes(); + assert!(retained > 0 && retained <= LIMIT); + assert!(seeds + .iter() + .any(|seed| matches!(seed, LifecycleTerminalUsageSeed::Prepared { .. }))); + let mut events = Vec::new(); + for seed in seeds { + let event = seed.build("seed-budget").await.expect("terminal event"); + assert_eq!(event.data.input_tokens, Some(100)); + assert_eq!(event.data.output_tokens, Some(500)); + events.push(event); + } + assert_eq!( + budget.retained_bytes(), + retained, + "ownership transfer must not drop or duplicate reservations" + ); + drop(events); + assert_eq!(budget.retained_bytes(), 0); + + let (context, payload) = terminal_seed_capture_test_seeds("seed-budget-dropped", 4096); + let queued = LifecycleTerminalUsageSeed::Sync { + context, + payload, + capture: None, + } + .prepare_capture_memory(Arc::clone(&budget)); + assert!(budget.retained_bytes() > 0); + drop(queued); + assert_eq!( + budget.retained_bytes(), + 0, + "dropping an unbuilt seed releases its bodies" + ); + } + + #[tokio::test] + async fn terminal_seed_capture_budget_cancellation_keeps_running_blocking_build_reserved() { + let budget = Arc::new(super::EventCaptureMemoryBudget::new(64 * 1024)); + let body = json!({"diagnostic": "x".repeat(4096)}); + let capture = super::TerminalSeedCaptureRetention::try_reserve( + [Some(&body), None, None, None], + Arc::clone(&budget), + ); + let retained = budget.retained_bytes(); + assert!(retained > 0); + let (started, started_rx) = tokio::sync::oneshot::channel(); + let (release, release_rx) = std::sync::mpsc::channel(); + let building = tokio::spawn(super::build_terminal_usage_event_offthread( + capture, + move || { + let _ = started.send(()); + release_rx + .recv_timeout(std::time::Duration::from_secs(2)) + .expect("release blocking builder"); + Ok(UsageEvent::new( + UsageEventType::Completed, + "seed-budget-cancel", + UsageEventData { + response_body: Some(body), + ..UsageEventData::default() + }, + )) + }, + )); + timeout(Duration::from_secs(1), started_rx) + .await + .expect("builder started") + .expect("start signal"); + building.abort(); + assert!(building + .await + .expect_err("cancelled wrapper") + .is_cancelled()); + assert_eq!( + budget.retained_bytes(), + retained, + "spawn_blocking survives cancellation of the caller" + ); + release.send(()).expect("release owned builder"); + timeout(Duration::from_secs(2), async { + while budget.retained_bytes() != 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("detached blocking result drops its body and reservation"); + + let body = json!({"diagnostic": "x".repeat(4096)}); + let capture = super::TerminalSeedCaptureRetention::try_reserve( + [Some(&body), None, None, None], + Arc::clone(&budget), + ); + let result = super::build_terminal_usage_event_offthread(capture, move || { + drop(body); + panic!("forced terminal seed builder panic") + }) + .await; + assert!(result.is_err()); + assert_eq!(budget.retained_bytes(), 0); + } + + #[tokio::test] + async fn terminal_seed_capture_budget_bounds_admission_backlog_and_preserves_failed_build_progress( + ) { + const COUNT: usize = 12; + for limit in [0, 96 * 1024] { + let config = UsageRuntimeConfig { + enabled: true, + terminal_submission_max_in_flight: 1, + ..UsageRuntimeConfig::default() + }; + let runtime = UsageRuntime::new(config).expect("usage runtime"); + let store = CloneQueueConfiguredUsageStore { + records: Arc::new(Mutex::new(Vec::new())), + queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), + }; + let budget = Arc::new(super::EventCaptureMemoryBudget::new(limit)); + let held = runtime + .terminal_submission_state + .acquire() + .await + .expect("hold terminal admission"); + runtime + .dispatch_terminal_seed( + &store, + "seed-budget-build-error".to_string(), + LifecycleTerminalUsageSeed::Prepared { + kind: super::TerminalSeedKind::Sync, + result: Err(DataLayerError::UnexpectedValue( + "forced prepared builder failure".to_string(), + )), + }, + ) + .await; + for index in 0..COUNT { + let request_id = format!("seed-budget-backlog-{index}"); + let (context, payload) = terminal_seed_capture_test_seeds(&request_id, 32 * 1024); + let seed = LifecycleTerminalUsageSeed::Sync { + context, + payload, + capture: None, + } + .prepare_capture_memory(Arc::clone(&budget)); + runtime + .dispatch_terminal_seed(&store, request_id, seed) + .await; + } + timeout(Duration::from_secs(2), async { + while runtime.metrics_snapshot().terminal_submission_pending < COUNT + 1 { + tokio::task::yield_now().await; + } + }) + .await + .expect("all seeds wait behind terminal admission"); + assert!(budget.retained_bytes() <= limit); + assert_eq!(budget.retained_bytes() > 0, limit > 0); + assert!(budget.downgraded_total() > 0); + assert!( + store.records.lock().unwrap().is_empty(), + "no persistence before admission" + ); + drop(held); + timeout(Duration::from_secs(5), async { + loop { + let snapshot = runtime.metrics_snapshot(); + if store.records.lock().unwrap().len() == COUNT + && snapshot.terminal_submission_pending == 0 + && snapshot.lifecycle_submission_pending == 0 + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("failed build releases its permit and every later terminal persists"); + let records = store.records.lock().unwrap(); + for record in records.iter() { + assert_eq!(record.status, "completed"); + assert_eq!(record.input_tokens, Some(100)); + assert_eq!(record.output_tokens, Some(500)); + assert_eq!(record.total_tokens, Some(600)); + assert!( + record.response_body.is_none(), + "default Basic policy still applies at its original position" + ); + } + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(runtime.metrics_snapshot().terminal_submission_in_flight, 0); + } + } + struct TestLifecycleSubmissionItem { request_id: String, priority: LifecycleSubmissionPriority, @@ -6292,6 +7272,9 @@ mod tests { #[derive(Default)] struct NoRedisUsageStore { records: Mutex>, + enrichment_failures: AtomicUsize, + enrichment_calls: AtomicUsize, + enriched_costs: Option<(f64, f64)>, } struct QueueConfiguredUsageStore { @@ -6473,7 +7456,23 @@ mod tests { #[async_trait] impl UsageBillingEventEnricher for NoRedisUsageStore { - async fn enrich_usage_event(&self, _event: &mut UsageEvent) -> Result<(), DataLayerError> { + async fn enrich_usage_event(&self, event: &mut UsageEvent) -> Result<(), DataLayerError> { + self.enrichment_calls.fetch_add(1, Ordering::AcqRel); + if self + .enrichment_failures + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |remaining| { + remaining.checked_sub(1) + }) + .is_ok() + { + // Real enrichers may update part of the event before a lookup fails. + event.data.total_cost_usd = Some(999.0); + return Err(DataLayerError::TimedOut("test billing lookup".to_string())); + } + if let Some((listed, actual)) = self.enriched_costs { + event.data.total_cost_usd = Some(listed); + event.data.actual_total_cost_usd = Some(actual); + } Ok(()) } } @@ -7574,6 +8573,211 @@ mod tests { } } + #[derive(Clone, Copy)] + enum DirectTerminalTestEntry { + PublicDirect, + OrderedDirect, + QueueDisabled, + } + + async fn invoke_direct_terminal_test_entry( + entry: DirectTerminalTestEntry, + runtime: &UsageRuntime, + store: &NoRedisUsageStore, + event: UsageEvent, + ) -> Option { + match entry { + DirectTerminalTestEntry::PublicDirect => { + runtime.record_terminal_event_direct(store, event).await; + None + } + DirectTerminalTestEntry::OrderedDirect => Some( + runtime + .persist_ordered_terminal_event(store, event, true) + .await, + ), + DirectTerminalTestEntry::QueueDisabled => { + Some(runtime.enqueue_or_write_terminal(store, event).await) + } + } + } + + async fn assert_direct_terminal_enrichment_failure_is_not_persisted( + entry: DirectTerminalTestEntry, + request_id: &str, + ) { + let runtime = UsageRuntime::new(UsageRuntimeConfig { + enabled: true, + queue_terminal_events: false, + ..UsageRuntimeConfig::default() + }) + .expect("usage runtime"); + let store = NoRedisUsageStore { + enrichment_failures: AtomicUsize::new(1), + enriched_costs: Some((0.456, 0.123)), + ..NoRedisUsageStore::default() + }; + let generation = runtime + .lifecycle_coalescer + .register(request_id.to_string()) + .await + .expect("delayed lifecycle generation"); + let event = UsageEvent::new( + UsageEventType::Completed, + request_id, + UsageEventData { + user_id: Some("user-direct-pricing-retry".to_string()), + provider_name: "openai".to_string(), + provider_id: Some("provider-direct-pricing-retry".to_string()), + model: "gpt-5".to_string(), + input_tokens: Some(4), + output_tokens: Some(8), + total_tokens: Some(12), + total_cost_usd: Some(0.9), + actual_total_cost_usd: Some(0.8), + status_code: Some(200), + ..UsageEventData::default() + }, + ); + + let outcome = timeout( + Duration::from_secs(2), + invoke_direct_terminal_test_entry(entry, &runtime, &store, event.clone()), + ) + .await + .expect("failed enrichment should release the terminal turn"); + if let Some(outcome) = outcome { + assert_eq!(outcome, super::TerminalPersistenceOutcome::Failed); + } + assert_eq!(store.enrichment_calls.load(Ordering::Acquire), 1); + assert!(store.records.lock().expect("records lock").is_empty()); + { + let coalescer = &runtime.lifecycle_coalescer; + let entries = coalescer.shards[coalescer.shard_index(request_id)] + .entries + .lock() + .await; + let marker = entries + .get(request_id) + .expect("delayed marker is preserved"); + assert_eq!(marker.generation, generation); + assert!(marker.terminal_seen_at.is_none()); + } + let snapshot = runtime.metrics_snapshot(); + assert_eq!(snapshot.ordered_lifecycle_pending, 0); + assert_eq!(snapshot.terminal_submission_in_flight, 0); + assert_eq!(snapshot.enqueue_retry_scheduled_total, 0); + + let outcome = timeout( + Duration::from_secs(2), + invoke_direct_terminal_test_entry(entry, &runtime, &store, event), + ) + .await + .expect("a later terminal attempt should be able to persist"); + if let Some(outcome) = outcome { + assert_eq!( + outcome, + super::TerminalPersistenceOutcome::PersistedDirectly + ); + } + assert_eq!(store.enrichment_calls.load(Ordering::Acquire), 2); + let records = store.records.lock().expect("records lock"); + assert_eq!(records.len(), 1); + assert_eq!(records[0].total_cost_usd, Some(0.456)); + assert_eq!(records[0].actual_total_cost_usd, Some(0.123)); + assert_eq!(records[0].total_tokens, Some(12)); + drop(records); + { + let coalescer = &runtime.lifecycle_coalescer; + let entries = coalescer.shards[coalescer.shard_index(request_id)] + .entries + .lock() + .await; + let marker = entries.get(request_id).expect("successful terminal marker"); + assert!(marker.terminal_seen_at.is_some()); + } + assert_eq!(runtime.metrics_snapshot().ordered_lifecycle_pending, 0); + assert_eq!(runtime.metrics_snapshot().terminal_submission_in_flight, 0); + } + + #[tokio::test] + async fn direct_terminal_pricing_failure_preserves_lifecycle_and_allows_correct_retry() { + assert_direct_terminal_enrichment_failure_is_not_persisted( + DirectTerminalTestEntry::PublicDirect, + "direct-terminal-pricing-retry", + ) + .await; + } + + #[tokio::test] + async fn ordered_direct_terminal_pricing_failure_returns_failed_before_correct_retry() { + assert_direct_terminal_enrichment_failure_is_not_persisted( + DirectTerminalTestEntry::OrderedDirect, + "ordered-direct-terminal-pricing-retry", + ) + .await; + } + + #[tokio::test] + async fn queue_disabled_terminal_pricing_failure_returns_failed_before_correct_retry() { + assert_direct_terminal_enrichment_failure_is_not_persisted( + DirectTerminalTestEntry::QueueDisabled, + "queue-disabled-terminal-pricing-retry", + ) + .await; + } + + #[tokio::test] + async fn bounded_direct_fallback_pricing_failure_does_not_write_or_report_success() { + let runtime = UsageRuntime::new(UsageRuntimeConfig { + enabled: true, + ..UsageRuntimeConfig::default() + }) + .expect("usage runtime"); + let store = NoRedisUsageStore { + enrichment_failures: AtomicUsize::new(1), + enriched_costs: Some((0.456, 0.123)), + ..NoRedisUsageStore::default() + }; + let mut event = UsageEvent::new( + UsageEventType::Completed, + "bounded-fallback-pricing-retry", + UsageEventData { + provider_name: "openai".to_string(), + model: "gpt-5".to_string(), + total_tokens: Some(12), + ..UsageEventData::default() + }, + ); + assert!( + !runtime + .try_write_terminal_direct_fallback(&store, &mut event, "test_retry") + .await + ); + assert!(store.records.lock().expect("records lock").is_empty()); + let snapshot = runtime.metrics_snapshot(); + assert_eq!(snapshot.terminal_direct_fallback_failed_total, 1); + assert_eq!(snapshot.terminal_direct_fallback_succeeded_total, 0); + assert_eq!(snapshot.terminal_direct_fallback_in_flight, 0); + + assert!( + runtime + .try_write_terminal_direct_fallback(&store, &mut event, "test_retry") + .await + ); + let records = store.records.lock().expect("records lock"); + assert_eq!(records.len(), 1); + assert_eq!(records[0].total_cost_usd, Some(0.456)); + assert_eq!(records[0].actual_total_cost_usd, Some(0.123)); + drop(records); + assert_eq!( + runtime + .metrics_snapshot() + .terminal_direct_fallback_succeeded_total, + 1 + ); + } + #[tokio::test] async fn terminal_usage_without_redis_writes_directly_to_usage_repository() { let runtime = UsageRuntime::new(UsageRuntimeConfig { @@ -8202,7 +9406,8 @@ mod tests { enqueue_retry_buffer_capacity: 1_024, ..UsageRuntimeConfig::default() }; - let dispatcher = super::PendingPersistenceDispatcher::spawn(&config); + let tasks = super::UsageBackgroundTasks::default(); + let dispatcher = super::PendingPersistenceDispatcher::spawn(&config, &tasks); let store = BlockingNonBatchPendingStore { release_writes: Arc::new(tokio::sync::Semaphore::new(0)), writes_in_flight: Arc::new(AtomicUsize::new(0)), @@ -8934,7 +10139,8 @@ mod tests { worker_record_concurrency_limit: Some(1), ..UsageRuntimeConfig::default() }; - let dispatcher = LifecycleSubmissionDispatcher::spawn(&config); + let tasks = super::UsageBackgroundTasks::default(); + let dispatcher = LifecycleSubmissionDispatcher::spawn(&config, &tasks); let request_id = "req-lifecycle-submission-panic"; dispatcher.dispatch(Box::new(PanickingLifecycleSubmissionItem { request_id: request_id.to_string(), @@ -8983,7 +10189,8 @@ mod tests { worker_record_concurrency_limit: Some(1), ..UsageRuntimeConfig::default() }; - let dispatcher = LifecycleSubmissionDispatcher::spawn(&config); + let tasks = super::UsageBackgroundTasks::default(); + let dispatcher = LifecycleSubmissionDispatcher::spawn(&config, &tasks); let started = Arc::new(tokio::sync::Notify::new()); let release = Arc::new(tokio::sync::Notify::new()); let seen = Arc::new(Mutex::new(Vec::new())); @@ -11082,6 +12289,280 @@ mod tests { assert_eq!(dispatcher.closed_or_unavailable_total(), 1); } + fn event_capture_budget_retry_event( + budget: &Arc, + ) -> UsageEvent { + let mut event = UsageEvent { + event_type: UsageEventType::Failed, + request_id: "event-capture-budget-retry".to_string(), + timestamp_ms: 123_000, + data: UsageEventData { + provider_name: "provider".to_string(), + model: "model".to_string(), + input_tokens: Some(100), + output_tokens: Some(500), + total_tokens: Some(600), + cache_read_input_tokens: Some(0), + status_code: Some(502), + error_category: Some("upstream_error".to_string()), + error_message: Some("upstream failed".to_string()), + response_body: Some(json!({"diagnostic": "x".repeat(1000)})), + response_body_state: Some(UsageBodyCaptureState::Inline), + request_metadata: Some(json!({ + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000", + "provider_service_tier": "priority", + "provider_cache_ttl_minutes": 60 + })), + ..UsageEventData::default() + }, + }; + event.data.apply_capture_memory_budget(Arc::clone(budget)); + event + } + + #[tokio::test] + async fn event_capture_budget_bounds_blocked_policy_waiters_and_releases_on_cancel_or_basic() { + for limit in [0, 64 * 1024] { + let runtime = UsageRuntime::new(UsageRuntimeConfig::default()).expect("runtime"); + let store = BlockingPolicyQueueConfiguredUsageStore { + queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), + policy_started: Arc::new(tokio::sync::Notify::new()), + release_policy: Arc::new(tokio::sync::Notify::new()), + policy_reads: Arc::new(AtomicUsize::new(0)), + }; + let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( + limit, + )); + let start = || { + let runtime = runtime.clone(); + let store = store.clone(); + let budget = Arc::clone(&budget); + tokio::spawn(async move { + let mut event = UsageEvent::new( + UsageEventType::Completed, + "event-capture-budget-policy", + UsageEventData { + input_tokens: Some(100), + output_tokens: Some(500), + response_body: Some(json!("x".repeat(40 * 1024))), + ..UsageEventData::default() + }, + ); + runtime + .apply_body_capture_policy_with_budget(&store, &mut event, budget) + .await; + event + }) + }; + let reading = start(); + timeout(Duration::from_secs(2), store.policy_started.notified()) + .await + .expect("policy read is blocked"); + let retained = budget.retained_bytes(); + assert_eq!(retained == 0, limit == 0); + assert!(retained <= limit); + let previous_denials = budget.downgraded_total(); + let waiting = start(); + timeout(Duration::from_secs(2), async { + while budget.downgraded_total() == previous_denials { + sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("waiting event must apply budget before acquiring policy mutex"); + assert_eq!(store.policy_reads.load(Ordering::Acquire), 1); + assert_eq!(budget.retained_bytes(), retained); + waiting.abort(); + assert!(waiting + .await + .expect_err("waiting task aborted") + .is_cancelled()); + assert_eq!(budget.retained_bytes(), retained); + reading.abort(); + assert!(reading + .await + .expect_err("reading task aborted") + .is_cancelled()); + assert_eq!(budget.retained_bytes(), 0); + + let completing = start(); + timeout(Duration::from_secs(2), store.policy_started.notified()) + .await + .expect("replacement policy read starts"); + assert_eq!(budget.retained_bytes(), retained); + store.release_policy.notify_one(); + let event = timeout(Duration::from_secs(2), completing) + .await + .expect("Basic policy completes") + .expect("policy task completion"); + assert!(event.data.response_body.is_none()); + assert_eq!( + event.data.response_body_state, + Some(UsageBodyCaptureState::Disabled) + ); + assert_eq!(event.data.input_tokens, Some(100)); + assert_eq!(event.data.output_tokens, Some(500)); + assert_eq!( + budget.retained_bytes(), + 0, + "Basic returns the lease while the event is still alive" + ); + } + } + + #[tokio::test] + async fn event_capture_budget_retry_holds_lease_and_preserves_event_after_recovery() { + for limit in [0, 64 * 1024] { + let config = UsageRuntimeConfig { + enabled: true, + queue_terminal_events: true, + consumer_block_ms: 1, + enqueue_retry_initial_backoff_ms: 1, + enqueue_retry_max_backoff_ms: 2, + ..UsageRuntimeConfig::default() + }; + let inner: Arc = + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let flaky = Arc::new(FlakyAppendQueueStore::new(inner, usize::MAX)); + let queue = UsageQueue::new(flaky.clone(), config.clone()).expect("retry queue"); + queue.ensure_consumer_group().await.expect("consumer group"); + let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( + limit, + )); + let event = event_capture_budget_retry_event(&budget); + let expected_fields = event.to_stream_fields().expect("expected wire event"); + let retained = budget.retained_bytes(); + assert_eq!(retained == 0, limit == 0); + let (sender, receiver) = mpsc::channel(1); + let metrics = Arc::new(super::UsageEnqueueDispatcherMetrics::default()); + let dispatcher = UsageEnqueueRetryDispatcher { + senders: vec![sender], + metrics: Arc::clone(&metrics), + }; + assert!(dispatcher.schedule( + queue.clone(), + event, + "terminal", + DataLayerError::Redis("initial failure".to_string()) + )); + let worker = tokio::spawn(super::run_usage_enqueue_retry_worker( + 0, + config, + receiver, + Arc::clone(&metrics), + )); + timeout(Duration::from_secs(2), async { + while metrics.retry_failed_total.load(Ordering::Acquire) < 2 { + sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("two retry failures"); + assert_eq!(budget.retained_bytes(), retained); + assert_eq!(dispatcher.pending(), 1); + flaky.remaining_failures.store(0, Ordering::Release); + drop(dispatcher); + timeout(Duration::from_secs(2), worker) + .await + .expect("retry recovery") + .expect("worker completion"); + assert_eq!(budget.retained_bytes(), 0); + assert_eq!(metrics.pending.load(Ordering::Acquire), 0); + let entries = queue + .read_group("capture-budget-consumer") + .await + .expect("persisted queue"); + assert_eq!(entries.len(), 1); + assert_eq!( + entries[0].fields, expected_fields, + "retry must preserve identity, terminal failure, billing, and truncation metadata" + ); + } + } + + #[tokio::test] + async fn event_capture_budget_retry_cancellation_releases_owned_body() { + let config = UsageRuntimeConfig { + enabled: true, + queue_terminal_events: true, + enqueue_retry_initial_backoff_ms: 1, + enqueue_retry_max_backoff_ms: 2, + ..UsageRuntimeConfig::default() + }; + let inner: Arc = + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let flaky = Arc::new(FlakyAppendQueueStore::new(inner, usize::MAX)); + let queue = UsageQueue::new(flaky, config.clone()).expect("retry queue"); + let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( + 64 * 1024, + )); + let event = event_capture_budget_retry_event(&budget); + let (sender, receiver) = mpsc::channel(1); + let metrics = Arc::new(super::UsageEnqueueDispatcherMetrics::default()); + let dispatcher = UsageEnqueueRetryDispatcher { + senders: vec![sender], + metrics: Arc::clone(&metrics), + }; + assert!(dispatcher.schedule( + queue, + event, + "terminal", + DataLayerError::Redis("failure".to_string()) + )); + let worker = tokio::spawn(super::run_usage_enqueue_retry_worker( + 0, + config, + receiver, + Arc::clone(&metrics), + )); + timeout(Duration::from_secs(2), async { + while metrics.retry_failed_total.load(Ordering::Acquire) == 0 { + sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("retry starts"); + assert!(budget.retained_bytes() > 0); + worker.abort(); + assert!(worker.await.expect_err("worker was aborted").is_cancelled()); + assert_eq!(budget.retained_bytes(), 0); + } + + #[tokio::test] + async fn event_capture_budget_retry_rejection_and_queued_drop_release_owned_bodies() { + let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( + 64 * 1024, + )); + let first = event_capture_budget_retry_event(&budget); + let retained_once = budget.retained_bytes(); + let second = event_capture_budget_retry_event(&budget); + assert_eq!(budget.retained_bytes(), retained_once * 2); + let config = UsageRuntimeConfig::default(); + let runner: Arc = + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let queue = UsageQueue::new(runner, config).expect("queue"); + let (sender, receiver) = mpsc::channel(1); + let dispatcher = UsageEnqueueRetryDispatcher { + senders: vec![sender], + metrics: Arc::new(super::UsageEnqueueDispatcherMetrics::default()), + }; + assert!(dispatcher.schedule( + queue.clone(), + first, + "terminal", + DataLayerError::Redis("failure".to_string()) + )); + assert!(!dispatcher.schedule( + queue, + second, + "terminal", + DataLayerError::Redis("failure".to_string()) + )); + assert_eq!(budget.retained_bytes(), retained_once); + drop(receiver); + assert_eq!(budget.retained_bytes(), 0); + } + #[tokio::test] async fn disabled_lifecycle_queue_writes_pending_directly() { let config = UsageRuntimeConfig { @@ -12709,6 +14190,69 @@ mod tests { assert_eq!(metadata["provider_actual_service_tier"], "priority"); } + #[test] + fn event_capture_budget_full_policy_preserves_facts_before_zero_budget_and_database_mapping() { + let mut event = UsageEvent::new( + UsageEventType::Completed, + "event-capture-budget-facts", + UsageEventData { + provider_name: "openai".to_string(), + model: "gpt-5.6-sol".to_string(), + endpoint_api_format: Some("openai:responses".to_string()), + input_tokens: Some(100), + output_tokens: Some(500), + total_tokens: Some(600), + cache_creation_input_tokens: Some(25), + cache_read_input_tokens: Some(0), + request_body: Some(json!({"reasoning": {"effort": "high"}})), + provider_request_body: Some(json!({ + "model": "gpt-5.6-sol", "reasoning": {"effort": "medium"}, + "service_tier": "priority" + })), + response_body: Some(json!({"service_tier": "Default"})), + request_metadata: Some(json!({ + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000" + })), + ..UsageEventData::default() + }, + ); + preserve_request_facts(&mut event); + preserve_provider_response_facts(&mut event); + apply_usage_body_capture_policy_to_event( + UsageBodyCapturePolicy { + record_level: UsageRequestRecordLevel::Full, + }, + &mut event, + ); + let budget = Arc::new(crate::event_capture_budget::EventCaptureMemoryBudget::new( + 0, + )); + event.data.apply_capture_memory_budget(Arc::clone(&budget)); + let record = crate::build_upsert_usage_record_from_event(&event).expect("database mapping"); + assert_eq!(record.input_tokens, Some(100)); + assert_eq!(record.output_tokens, Some(500)); + assert_eq!(record.cache_creation_input_tokens, Some(25)); + assert_eq!(record.cache_read_input_tokens, Some(0)); + assert!(record.request_body.is_none()); + assert!(record.provider_request_body.is_none()); + assert!(record.response_body.is_none()); + assert_eq!( + record.provider_request_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + let metadata = record.request_metadata.expect("preserved billing facts"); + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_reasoning_effort"], "medium"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + assert_eq!(metadata["provider_cache_ttl_minutes"], 30); + assert_eq!( + metadata["plan_usage_reservation_token"], + "550e8400-e29b-41d4-a716-446655440000" + ); + assert_eq!(budget.retained_bytes(), 0); + } + #[test] fn basic_request_record_level_strips_body_capture_but_preserves_derived_fields() { let mut event = UsageEvent::new( diff --git a/crates/aether-usage/runtime/src/runtime_queue_payload_tests.rs b/crates/aether-usage/runtime/src/runtime_queue_payload_tests.rs new file mode 100644 index 000000000..3bab6755e --- /dev/null +++ b/crates/aether-usage/runtime/src/runtime_queue_payload_tests.rs @@ -0,0 +1,432 @@ +use super::*; + +use super::super::{run_usage_enqueue_retry_worker, TerminalPersistenceOutcome}; + +const PAYLOAD_LIMIT: usize = 4 * 1024; + +fn payload_config(name: &str) -> UsageRuntimeConfig { + UsageRuntimeConfig { + enabled: true, + queue_terminal_events: true, + queue_lifecycle_events: true, + stream_key: format!("usage:events:test:payload:{name}"), + consumer_group: format!("usage_consumers_payload_{name}"), + queue_payload_max_bytes: PAYLOAD_LIMIT, + consumer_block_ms: 1, + enqueue_retry_buffer_capacity: 8, + enqueue_retry_workers: 1, + enqueue_retry_initial_backoff_ms: 1, + enqueue_retry_max_backoff_ms: 2, + ..UsageRuntimeConfig::default() + } +} + +fn payload_event(request_id: &str, oversized: bool) -> UsageEvent { + UsageEvent { + event_type: UsageEventType::Completed, + request_id: request_id.to_string(), + timestamp_ms: 123_000, + data: UsageEventData { + user_id: Some("user-payload".to_string()), + api_key_id: Some("key-payload".to_string()), + provider_name: "openai".to_string(), + provider_id: Some("provider-payload".to_string()), + provider_api_key_id: Some("provider-key-payload".to_string()), + // Model identity is a core field, so omitting diagnostic bodies cannot make this fit. + model: if oversized { + "m".repeat(PAYLOAD_LIMIT * 2) + } else { + "gpt-5".to_string() + }, + target_model: Some("gpt-5".to_string()), + api_format: Some("openai:responses".to_string()), + endpoint_api_format: Some("openai:responses".to_string()), + input_tokens: Some(100), + output_tokens: Some(25), + total_tokens: Some(125), + cache_creation_input_tokens: Some(30), + cache_creation_ephemeral_5m_input_tokens: Some(0), + cache_creation_ephemeral_1h_input_tokens: Some(30), + cache_read_input_tokens: Some(0), + total_cost_usd: Some(0.5), + actual_total_cost_usd: Some(0.25), + status_code: Some(200), + error_message: Some("preserve error presence and text".to_string()), + first_byte_time_ms: Some(12), + response_time_ms: Some(34), + request_headers: Some(json!({"x-request": "original"})), + provider_request_headers: Some(json!({"x-provider-request": "original"})), + response_headers: Some(json!({"x-provider-response": "original"})), + client_response_headers: Some(json!({"x-client-response": "original"})), + request_body: Some(json!({"reasoning": {"effort": "high"}, "input": "original"})), + request_body_state: Some(UsageBodyCaptureState::Inline), + provider_request_body: Some(json!({ + "model": "gpt-5", "service_tier": "priority", + "prompt_cache_retention": "24h", "input": "original provider input" + })), + provider_request_body_state: Some(UsageBodyCaptureState::Inline), + response_body: Some(json!({"service_tier": "default", "output": "original output"})), + response_body_state: Some(UsageBodyCaptureState::Inline), + client_response_body: Some(json!({"output": "original client output"})), + client_response_body_state: Some(UsageBodyCaptureState::Inline), + request_metadata: Some(json!({ + "api_key_is_standalone": true, + "plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000", + "plan_usage_reservation_deferred": true, + "usage_available": true, + "usage_pricing_available": true, + "provider_cache_ttl_minutes": 1440, + "provider_service_tier": "priority", + "provider_actual_service_tier": "default", + "dimensions": { + "image_count": 2, "image_size": "1024x1024", "image_quality": "high", + "image_output_format": "png", "reasoning_tokens": 0 + } + })), + ..UsageEventData::default() + }, + } +} + +async fn payload_queue(config: &UsageRuntimeConfig) -> (UsageQueue, Arc) { + let inner: Arc = + Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let runner = Arc::new(FlakyAppendQueueStore::new(inner, 0)); + let queue = UsageQueue::new(runner.clone(), config.clone()).expect("payload queue"); + queue.ensure_consumer_group().await.expect("payload group"); + (queue, runner) +} + +#[tokio::test] +async fn terminal_oversize_uses_original_event_for_direct_fallback_without_opening_circuit() { + let config = payload_config("terminal_direct"); + let (queue, runner) = payload_queue(&config).await; + let store = CloneQueueConfiguredUsageStore { + records: Arc::new(Mutex::new(Vec::new())), + queue: runner.clone(), + }; + let runtime = UsageRuntime::new(config).expect("usage runtime"); + let event = payload_event("payload-terminal-direct", true); + let expected = crate::build_upsert_usage_record_from_event(&event).expect("original record"); + + assert!(matches!( + queue.enqueue(&event).await, + Err(DataLayerError::InvalidInput(_)) + )); + assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0); + assert_eq!( + runtime.enqueue_or_write_terminal(&store, event).await, + TerminalPersistenceOutcome::PersistedDirectly + ); + { + let records = store.records.lock().expect("records lock"); + assert_eq!(records.as_slice(), &[expected]); + assert_eq!(records[0].cache_read_input_tokens, Some(0)); + assert_eq!(records[0].cache_creation_ephemeral_5m_input_tokens, Some(0)); + assert_eq!( + records[0].request_body_state, + Some(UsageBodyCaptureState::Inline) + ); + assert!(records[0].provider_request_body.is_some()); + assert!(records[0].request_headers.is_some()); + } + assert_eq!( + runtime + .terminal_enqueue_state + .circuit_open_until_unix_ms + .load(Ordering::Acquire), + 0 + ); + let snapshot = runtime.metrics_snapshot(); + assert_eq!(snapshot.terminal_direct_fallback_succeeded_total, 1); + assert_eq!(snapshot.enqueue_retry_scheduled_total, 0); + assert_eq!(snapshot.enqueue_retry_pending, 0); + + assert_eq!( + runtime + .enqueue_or_write_terminal(&store, payload_event("payload-after-direct", false)) + .await, + TerminalPersistenceOutcome::Queued + ); + assert_eq!(runner.successful_appends.load(Ordering::Acquire), 1); + assert_eq!(store.records.lock().expect("records lock").len(), 1); +} + +#[tokio::test] +async fn terminal_oversize_direct_failure_preserves_first_byte_and_does_not_retry_or_open_circuit() +{ + let config = payload_config("terminal_failed"); + let (_, runner) = payload_queue(&config).await; + let store = FailingWriteQueueConfiguredUsageStore { + queue: runner.clone(), + upsert_attempts: Arc::new(AtomicUsize::new(0)), + }; + let runtime = UsageRuntime::new(config).expect("usage runtime"); + let request_id = "payload-terminal-failed"; + let generation = runtime + .lifecycle_coalescer + .mark_first_byte(request_id) + .await + .expect("first byte"); + + assert_eq!( + runtime + .enqueue_or_write_terminal(&store, payload_event(request_id, true)) + .await, + TerminalPersistenceOutcome::Failed + ); + assert!( + runtime + .lifecycle_coalescer + .first_byte_is_current(request_id, generation) + .await + ); + assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0); + assert_eq!(store.upsert_attempts.load(Ordering::Acquire), 1); + assert_eq!( + runtime + .terminal_enqueue_state + .circuit_open_until_unix_ms + .load(Ordering::Acquire), + 0 + ); + let snapshot = runtime.metrics_snapshot(); + assert_eq!(snapshot.terminal_direct_fallback_failed_total, 1); + assert_eq!(snapshot.terminal_enqueue_deferred_dropped_total, 1); + assert_eq!(snapshot.enqueue_retry_scheduled_total, 0); + assert_eq!(snapshot.enqueue_retry_pending, 0); + + assert_eq!( + runtime + .enqueue_or_write_terminal(&store, payload_event("payload-after-failed", false)) + .await, + TerminalPersistenceOutcome::Queued + ); + assert_eq!(runner.successful_appends.load(Ordering::Acquire), 1); + assert_eq!(store.upsert_attempts.load(Ordering::Acquire), 1); +} + +#[tokio::test] +async fn terminal_oversize_direct_failure_stays_failed_when_primary_enqueue_is_deferred() { + for (name, circuit_open) in [("circuit_open", true), ("in_flight_limit", false)] { + let mut config = payload_config(name); + config.terminal_enqueue_max_in_flight = 1; + let (_, runner) = payload_queue(&config).await; + let store = FailingWriteQueueConfiguredUsageStore { + queue: runner.clone(), + upsert_attempts: Arc::new(AtomicUsize::new(0)), + }; + let runtime = UsageRuntime::new(config).expect("usage runtime"); + let request_id = format!("payload-terminal-{name}"); + let generation = runtime + .lifecycle_coalescer + .mark_first_byte(&request_id) + .await + .expect("first byte"); + let original_deadline = if circuit_open { + let deadline = super::super::now_unix_ms().saturating_add(60_000); + runtime.terminal_enqueue_state.open_circuit(deadline); + deadline + } else { + 0 + }; + let held_guard = if circuit_open { + None + } else { + Some( + runtime + .terminal_enqueue_state + .try_acquire_in_flight(1) + .expect("hold the only enqueue slot"), + ) + }; + + assert_eq!( + timeout( + Duration::from_secs(2), + runtime.enqueue_or_write_terminal(&store, payload_event(&request_id, true)), + ) + .await + .expect("bounded terminal fallback"), + TerminalPersistenceOutcome::Failed, + "{name} must not report an oversized event as buffered" + ); + assert!( + runtime + .lifecycle_coalescer + .first_byte_is_current(&request_id, generation) + .await + ); + assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0); + assert_eq!(store.upsert_attempts.load(Ordering::Acquire), 1); + assert_eq!( + runtime + .terminal_enqueue_state + .circuit_open_until_unix_ms + .load(Ordering::Acquire), + original_deadline + ); + let snapshot = runtime.metrics_snapshot(); + assert_eq!(snapshot.terminal_direct_fallback_failed_total, 1); + assert_eq!(snapshot.terminal_enqueue_deferred_dropped_total, 1); + assert_eq!(snapshot.terminal_enqueue_deferred_retry_total, 0); + assert_eq!(snapshot.enqueue_retry_permanent_failure_total, 1); + assert_eq!(snapshot.enqueue_retry_scheduled_total, 0); + assert_eq!(snapshot.enqueue_retry_pending, 0); + assert_eq!( + runtime.terminal_enqueue_state.in_flight(), + u64::from(!circuit_open) + ); + drop(held_guard); + assert_eq!(runtime.terminal_enqueue_state.in_flight(), 0); + } +} + +#[tokio::test] +async fn retry_worker_discards_oversize_and_drains_next_event_on_the_same_shard() { + let config = payload_config("retry_drain"); + let (queue, runner) = payload_queue(&config).await; + let (sender, receiver) = mpsc::channel(2); + let dispatcher = UsageEnqueueRetryDispatcher { + senders: vec![sender], + metrics: Arc::new(Default::default()), + }; + let metrics_view = UsageEnqueueRetryDispatcher { + senders: Vec::new(), + metrics: Arc::clone(&dispatcher.metrics), + }; + for (request_id, oversized) in [ + ("payload-retry-oversize", true), + ("payload-retry-small", false), + ] { + // Bypass admission to exercise the worker's defense for an already buffered event. + assert!(dispatcher + .schedule_item( + queue.clone(), + payload_event(request_id, oversized), + "terminal", + Some("prior transient failure"), + ) + .is_some()); + } + assert_eq!(dispatcher.pending(), 2); + drop(dispatcher); + + // Run the real worker as a cancellable future, so a regression cannot leave a detached retry. + timeout( + Duration::from_secs(2), + run_usage_enqueue_retry_worker(0, config, receiver, Arc::clone(&metrics_view.metrics)), + ) + .await + .expect("the permanent failure must not block the shard"); + + assert_eq!(metrics_view.permanent_failure_total(), 1); + assert_eq!(metrics_view.recovered_total(), 1); + assert_eq!(metrics_view.pending(), 0); + assert_eq!(runner.append_attempts.load(Ordering::Acquire), 1); + let entries = queue + .read_group("payload-retry-reader") + .await + .expect("queue read"); + assert_eq!(entries.len(), 1); + assert_eq!( + UsageEvent::from_stream_fields(&entries[0].fields) + .expect("queued event") + .request_id, + "payload-retry-small" + ); +} + +#[tokio::test] +async fn retry_dispatcher_rejects_oversize_for_permanent_and_transient_causes_without_consuming_slots( +) { + let config = payload_config("retry_reject"); + let (queue, runner) = payload_queue(&config).await; + let (sender, receiver) = mpsc::channel(1); + let dispatcher = UsageEnqueueRetryDispatcher { + senders: vec![sender], + metrics: Arc::new(Default::default()), + }; + let oversized = payload_event("payload-rejected", true); + let error = queue.enqueue(&oversized).await.expect_err("oversize input"); + assert!(matches!(error, DataLayerError::InvalidInput(_))); + assert!(!dispatcher.schedule(queue.clone(), oversized, "terminal", error)); + for cause in [ + DataLayerError::TimedOut("primary enqueue was deferred".to_string()), + DataLayerError::Redis("prior transient failure".to_string()), + ] { + assert!(!dispatcher.schedule( + queue.clone(), + payload_event("payload-rejected-transient-cause", true), + "terminal", + cause, + )); + } + assert_eq!(dispatcher.permanent_failure_total(), 3); + assert_eq!(dispatcher.pending(), 0); + assert_eq!(dispatcher.scheduled_total(), 0); + assert!(dispatcher.schedule( + queue, + payload_event("payload-after-reject", false), + "terminal", + DataLayerError::Redis("retryable failure".to_string()), + )); + let metrics_view = UsageEnqueueRetryDispatcher { + senders: Vec::new(), + metrics: Arc::clone(&dispatcher.metrics), + }; + drop(dispatcher); + + timeout( + Duration::from_secs(2), + run_usage_enqueue_retry_worker(0, config, receiver, Arc::clone(&metrics_view.metrics)), + ) + .await + .expect("retry worker drain"); + + assert_eq!(metrics_view.permanent_failure_total(), 3); + assert_eq!(metrics_view.recovered_total(), 1); + assert_eq!(metrics_view.pending(), 0); + assert_eq!(runner.append_attempts.load(Ordering::Acquire), 1); +} + +#[tokio::test] +async fn lifecycle_oversize_does_not_open_circuit_or_block_the_next_lifecycle_event() { + let config = payload_config("lifecycle"); + let (queue, runner) = payload_queue(&config).await; + let store = CloneQueueConfiguredUsageStore { + records: Arc::new(Mutex::new(Vec::new())), + queue: runner.clone(), + }; + let runtime = UsageRuntime::new(config).expect("usage runtime"); + let mut oversized = payload_event("payload-lifecycle-oversize", true); + oversized.event_type = UsageEventType::Streaming; + + assert!(!runtime.enqueue_lifecycle_event(&store, oversized).await); + assert_eq!( + runtime + .lifecycle_enqueue_state + .circuit_open_until_unix_ms + .load(Ordering::Acquire), + 0 + ); + assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0); + let snapshot = runtime.metrics_snapshot(); + assert_eq!(snapshot.enqueue_retry_scheduled_total, 0); + assert_eq!(snapshot.enqueue_retry_pending, 0); + assert!(store.records.lock().expect("records lock").is_empty()); + + let mut small = payload_event("payload-lifecycle-small", false); + small.event_type = UsageEventType::Streaming; + assert!(runtime.enqueue_lifecycle_event(&store, small).await); + assert_eq!(runner.successful_appends.load(Ordering::Acquire), 1); + let entries = queue + .read_group("payload-lifecycle-reader") + .await + .expect("queue read"); + assert_eq!(entries.len(), 1); + let queued = UsageEvent::from_stream_fields(&entries[0].fields).expect("queued lifecycle"); + assert_eq!(queued.request_id, "payload-lifecycle-small"); + assert_eq!(queued.event_type, UsageEventType::Streaming); + assert_eq!(queued.data.first_byte_time_ms, Some(12)); +} diff --git a/crates/aether-usage/runtime/src/runtime_shutdown_tests.rs b/crates/aether-usage/runtime/src/runtime_shutdown_tests.rs new file mode 100644 index 000000000..6a5ee8b34 --- /dev/null +++ b/crates/aether-usage/runtime/src/runtime_shutdown_tests.rs @@ -0,0 +1,366 @@ +use super::*; + +fn config() -> UsageRuntimeConfig { + UsageRuntimeConfig { + enabled: true, + queue_terminal_events: true, + queue_lifecycle_events: true, + worker_count: 2, + consumer_block_ms: 60_000, + enqueue_retry_buffer_capacity: 256, + enqueue_retry_initial_backoff_ms: 60_000, + enqueue_retry_max_backoff_ms: 60_000, + ..UsageRuntimeConfig::default() + } +} + +fn store() -> CloneQueueConfiguredUsageStore { + CloneQueueConfiguredUsageStore { + records: Arc::new(Mutex::new(Vec::new())), + queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())), + } +} + +fn terminal(request_id: &str) -> UsageEvent { + UsageEvent::new( + UsageEventType::Completed, + request_id, + UsageEventData { + provider_name: "openai".to_string(), + model: "test".to_string(), + status_code: Some(200), + input_tokens: Some(3), + output_tokens: Some(7), + total_tokens: Some(10), + ..UsageEventData::default() + }, + ) +} + +#[tokio::test] +async fn usage_shutdown_waits_for_a_producer_before_closing_admission() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = store(); + let producer = runtime.track_producer(); + let copy = runtime.clone(); + let task = tokio::spawn(async move { copy.shutdown(Duration::from_secs(3)).await }); + sleep(Duration::from_millis(30)).await; + assert!(!task.is_finished()); + runtime + .submit_terminal_event(&store, terminal("last-producer")) + .await; + drop(producer); + task.await.unwrap().unwrap(); + assert_eq!(runtime.local_work_pending(), 0); + assert_eq!( + store + .queue + .stats(&runtime.config.stream_key, None) + .await + .unwrap() + .stream_length, + 1 + ); + runtime.shutdown(Duration::from_secs(1)).await.unwrap(); + runtime + .submit_terminal_event(&store, terminal("closed")) + .await; + runtime + .record_terminal_event(&store, terminal("closed-direct-api")) + .await; + assert_eq!( + store + .queue + .stats(&runtime.config.stream_key, None) + .await + .unwrap() + .stream_length, + 1 + ); +} + +#[tokio::test] +async fn usage_shutdown_persists_concurrent_terminal_handoffs() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = store(); + let mut tasks = tokio::task::JoinSet::new(); + for index in 0..128 { + let runtime = runtime.clone(); + let store = store.clone(); + let producer = runtime.track_producer(); + tasks.spawn(async move { + let _producer = producer; + runtime + .submit_terminal_event(&store, terminal(&format!("drain-{index}"))) + .await; + }); + } + runtime.shutdown(Duration::from_secs(5)).await.unwrap(); + while let Some(result) = tasks.join_next().await { + result.unwrap(); + } + assert_eq!( + store + .queue + .stats(&runtime.config.stream_key, None) + .await + .unwrap() + .stream_length, + 128 + ); + assert_eq!(runtime.local_work_pending(), 0); +} + +#[tokio::test] +async fn usage_shutdown_wakes_retry_backoff_and_preserves_all_buffered_events() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = store(); + let queue = UsageQueue::new(Arc::clone(&store.queue), runtime.config.clone()).unwrap(); + for index in 0..32 { + assert!(runtime.enqueue_retry.schedule( + queue.clone(), + terminal(&format!("retry-{index}")), + "terminal", + DataLayerError::Redis("transient".into()) + )); + } + runtime.shutdown(Duration::from_secs(3)).await.unwrap(); + assert_eq!(runtime.enqueue_retry.pending(), 0); + assert_eq!(runtime.enqueue_retry.recovered_total(), 32); + assert_eq!( + store + .queue + .stats(&runtime.config.stream_key, None) + .await + .unwrap() + .stream_length, + 32 + ); +} + +#[tokio::test] +async fn usage_shutdown_failure_retains_retry_work_for_a_later_attempt() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = store(); + let flaky = Arc::new(FlakyAppendQueueStore::new( + Arc::clone(&store.queue), + usize::MAX, + )); + let queue = UsageQueue::new(flaky.clone(), runtime.config.clone()).unwrap(); + assert!(runtime.enqueue_retry.schedule( + queue, + terminal("recover-after-deadline"), + "terminal", + DataLayerError::Redis("unavailable".into()) + )); + let result = runtime.shutdown(Duration::from_millis(150)).await; + assert!(matches!(result, Err(DataLayerError::TimedOut(_)))); + assert_eq!(runtime.enqueue_retry.pending(), 1); + assert!(flaky.append_attempts.load(Ordering::Acquire) <= 4); + flaky.remaining_failures.store(0, Ordering::Release); + runtime.shutdown(Duration::from_secs(3)).await.unwrap(); + assert_eq!(runtime.enqueue_retry.recovered_total(), 1); + assert_eq!( + store + .queue + .stats(&runtime.config.stream_key, None) + .await + .unwrap() + .stream_length, + 1 + ); +} + +#[tokio::test] +async fn usage_shutdown_flushes_delayed_lifecycle_without_waiting_for_its_timer() { + let mut config = config(); + config.lifecycle_enqueue_delay_ms = 60_000; + let runtime = UsageRuntime::new(config).unwrap(); + let store = store(); + let event = UsageEvent::new( + UsageEventType::Pending, + "delayed", + UsageEventData::default(), + ); + runtime + .enqueue_lifecycle_event_with_config_delay(&store, event) + .await; + assert!(runtime.local_work_pending() > 0); + runtime.shutdown(Duration::from_secs(3)).await.unwrap(); + assert_eq!(runtime.local_work_pending(), 0); + assert_eq!( + store + .queue + .stats(&runtime.config.stream_key, None) + .await + .unwrap() + .stream_length, + 1 + ); +} + +#[tokio::test] +async fn usage_shutdown_stops_every_idle_worker_and_supervisor() { + for supervised in [false, true] { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = Arc::new(store()); + let handles = if supervised { + vec![runtime.spawn_worker_supervisor(store).unwrap()] + } else { + runtime.spawn_workers(store) + }; + sleep(Duration::from_millis(30)).await; + runtime.shutdown(Duration::from_secs(3)).await.unwrap(); + for handle in handles { + timeout(Duration::from_secs(1), handle) + .await + .unwrap() + .unwrap(); + } + assert_eq!(runtime.metrics_snapshot().worker_active_count, 0); + assert_eq!(runtime.shutdown.supervisors.load(Ordering::Acquire), 0); + } +} + +#[tokio::test] +async fn usage_shutdown_does_not_cancel_a_write_or_ack_its_unfinished_record() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = BlockingWriteQueueConfiguredUsageStore { + records: Arc::new(Mutex::new(Vec::new())), + queue: store().queue, + write_started: Arc::new(tokio::sync::Notify::new()), + release_writes: Arc::new(tokio::sync::Notify::new()), + writes_completed: Arc::new(AtomicUsize::new(0)), + }; + let queue = UsageQueue::new(Arc::clone(&store.queue), runtime.config.clone()).unwrap(); + queue.ensure_consumer_group().await.unwrap(); + queue.enqueue(&terminal("in-flight-worker")).await.unwrap(); + let worker = runtime.spawn_worker(Arc::new(store.clone())).unwrap(); + timeout(Duration::from_secs(3), store.write_started.notified()) + .await + .unwrap(); + assert!(runtime.shutdown(Duration::from_millis(50)).await.is_err()); + assert!(!worker.is_finished()); + assert_eq!( + store + .queue + .stats( + &runtime.config.stream_key, + Some(&runtime.config.consumer_group) + ) + .await + .unwrap() + .group_pending, + 1 + ); + store.release_writes.notify_one(); + runtime.shutdown(Duration::from_secs(3)).await.unwrap(); + worker.await.unwrap(); + assert_eq!(store.writes_completed.load(Ordering::Acquire), 1); + assert_eq!( + store + .queue + .stats( + &runtime.config.stream_key, + Some(&runtime.config.consumer_group) + ) + .await + .unwrap() + .group_pending, + 0 + ); +} + +#[tokio::test] +async fn usage_shutdown_disabled_runtime_is_immediate() { + UsageRuntime::disabled() + .shutdown(Duration::from_secs(1)) + .await + .unwrap(); +} + +#[tokio::test] +async fn usage_shutdown_consumes_process_local_queue_before_stopping_workers() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = store(); + let worker = runtime + .spawn_worker_supervisor(Arc::new(store.clone())) + .unwrap(); + for index in 0..64 { + runtime + .submit_terminal_event(&store, terminal(&format!("local-{index}"))) + .await; + } + runtime + .shutdown_with_local_queue(Duration::from_secs(5), Some(Arc::clone(&store.queue))) + .await + .unwrap(); + worker.await.unwrap(); + assert_eq!(store.records.lock().unwrap().len(), 64); + assert_eq!( + store + .queue + .stats( + &runtime.config.stream_key, + Some(&runtime.config.consumer_group) + ) + .await + .unwrap() + .group_pending, + 0 + ); +} + +#[tokio::test] +async fn usage_shutdown_rejects_unconsumed_memory_queue_as_success() { + let runtime = UsageRuntime::new(config()).unwrap(); + let store = store(); + runtime + .submit_terminal_event(&store, terminal("unconsumed")) + .await; + let result = runtime + .shutdown_with_local_queue(Duration::from_millis(50), Some(Arc::clone(&store.queue))) + .await; + assert!(result.is_err()); + let worker = runtime.spawn_worker(Arc::new(store.clone())).unwrap(); + runtime + .shutdown_with_local_queue(Duration::from_secs(3), Some(Arc::clone(&store.queue))) + .await + .unwrap(); + worker.await.unwrap(); + assert_eq!(store.records.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn usage_shutdown_does_not_miss_accepted_pending_to_terminal_handoffs() { + let mut config = config(); + config.queue_terminal_events = false; + let runtime = UsageRuntime::new(config).unwrap(); + let store = store(); + for index in 0..32 { + let id = format!("ordered-{index}"); + let plan = terminal_test_plan(&id); + runtime.record_pending(&store, build_lifecycle_usage_seed(&plan, None)); + runtime.record_stream_started( + &store, + &build_lifecycle_usage_seed(&plan, None), + 200, + Some(&ExecutionTelemetry { + ttfb_ms: Some(5), + elapsed_ms: None, + upstream_bytes: None, + }), + ); + runtime.submit_terminal_event(&store, terminal(&id)).await; + } + runtime.shutdown(Duration::from_secs(5)).await.unwrap(); + let records = store.records.lock().unwrap(); + for index in 0..32 { + let statuses: Vec<_> = records + .iter() + .filter(|r| r.request_id == format!("ordered-{index}")) + .map(|r| r.status.as_str()) + .collect(); + assert_eq!(statuses, ["pending", "streaming", "completed"]); + } +} diff --git a/crates/aether-usage/runtime/src/settlement.rs b/crates/aether-usage/runtime/src/settlement.rs index 89669a060..65d92b190 100644 --- a/crates/aether-usage/runtime/src/settlement.rs +++ b/crates/aether-usage/runtime/src/settlement.rs @@ -36,8 +36,19 @@ pub async fn reconcile_usage_policy_cost_for_event( writer: &dyn UsageSettlementWriter, event: &UsageEvent, ) -> Result<(), DataLayerError> { + reconcile_usage_policy_cost_for_event_with_result(writer, event) + .await + .map(|_| ()) +} + +pub(crate) struct ReconciledUsagePolicyCost(ReconcileUsagePolicyCostInput); + +pub(crate) async fn reconcile_usage_policy_cost_for_event_with_result( + writer: &dyn UsageSettlementWriter, + event: &UsageEvent, +) -> Result, DataLayerError> { if !writer.has_usage_settlement_writer() { - return Ok(()); + return Ok(None); } let terminal_state = match event.event_type { UsageEventType::Completed => UsagePolicyCostReservationState::Finalized, @@ -49,16 +60,16 @@ pub async fn reconcile_usage_policy_cost_for_event( UsageEventType::Failed | UsageEventType::Cancelled => { UsagePolicyCostReservationState::Released } - UsageEventType::Pending | UsageEventType::Streaming => return Ok(()), + UsageEventType::Pending | UsageEventType::Streaming => return Ok(None), }; if plan_usage_reservation_reconciliation_is_deferred(event.data.request_metadata.as_ref()) { - return Ok(()); + return Ok(None); } let Some(subject_id) = event.data.user_id.as_deref().and_then(non_empty_trimmed) else { - return Ok(()); + return Ok(None); }; let Some(reservation_token) = event_usage_policy_reservation_token(event) else { - return Ok(()); + return Ok(None); }; let actual_cost_units = if terminal_state == UsagePolicyCostReservationState::Finalized { let actual_cost_usd = event.data.actual_total_cost_usd.ok_or_else(|| { @@ -75,22 +86,39 @@ pub async fn reconcile_usage_policy_cost_for_event( 0 }; - let _ = writer - .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { - request_id: event.request_id.clone(), - subject_id: subject_id.to_string(), - reservation_token: reservation_token.to_string(), - actual_cost_units, - terminal_state, - finalized_at_unix_secs: event.timestamp_ms / 1_000, - }) - .await?; - Ok(()) + let input = ReconcileUsagePolicyCostInput { + request_id: event.request_id.clone(), + subject_id: subject_id.to_string(), + reservation_token: reservation_token.to_string(), + actual_cost_units, + terminal_state, + finalized_at_unix_secs: event.timestamp_ms / 1_000, + }; + let stored = writer.reconcile_usage_policy_cost(input.clone()).await?; + // A successful call alone is insufficient: None or a different existing terminal + // reservation must not suppress reconciliation of the subsequent stored usage row. + let matches = stored.is_some_and(|stored| { + stored.request_id == input.request_id + && stored.subject_id == input.subject_id + && stored.reservation_token == input.reservation_token + && stored.actual_cost_units == Some(input.actual_cost_units) + && stored.state == input.terminal_state + && stored.finalized_at_unix_secs == Some(input.finalized_at_unix_secs) + }); + Ok(matches.then_some(ReconciledUsagePolicyCost(input))) } pub async fn settle_usage_if_needed( writer: &dyn UsageSettlementWriter, usage: &StoredRequestUsageAudit, +) -> Result<(), DataLayerError> { + settle_usage_with_reconciled_cost(writer, usage, None).await +} + +pub(crate) async fn settle_usage_with_reconciled_cost( + writer: &dyn UsageSettlementWriter, + usage: &StoredRequestUsageAudit, + reconciled: Option, ) -> Result<(), DataLayerError> { if !writer.has_usage_settlement_writer() { return Ok(()); @@ -132,17 +160,21 @@ pub async fn settle_usage_if_needed( } else { (UsagePolicyCostReservationState::Released, 0) }; - let _ = writer - .reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput { - request_id: usage.request_id.clone(), - subject_id: subject_id.to_string(), - reservation_token: reservation_token.to_string(), - actual_cost_units, - terminal_state, - finalized_at_unix_secs: finalized_at_unix_secs - .unwrap_or(usage.updated_at_unix_secs), - }) - .await?; + let input = ReconcileUsagePolicyCostInput { + request_id: usage.request_id.clone(), + subject_id: subject_id.to_string(), + reservation_token: reservation_token.to_string(), + actual_cost_units, + terminal_state, + finalized_at_unix_secs: finalized_at_unix_secs + .unwrap_or(usage.updated_at_unix_secs), + }; + if !reconciled + .as_ref() + .is_some_and(|previous| previous.0 == input) + { + let _ = writer.reconcile_usage_policy_cost(input).await?; + } } } @@ -243,6 +275,10 @@ fn finite_cost(value: f64) -> Result { #[cfg(test)] mod tests { + mod reconciliation_reuse { + include!("settlement_reuse_tests.rs"); + } + use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Mutex; use std::time::Duration; diff --git a/crates/aether-usage/runtime/src/settlement_reuse_tests.rs b/crates/aether-usage/runtime/src/settlement_reuse_tests.rs new file mode 100644 index 000000000..34fc4a60b --- /dev/null +++ b/crates/aether-usage/runtime/src/settlement_reuse_tests.rs @@ -0,0 +1,517 @@ +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; + +use aether_data::repository::settlement::InMemorySettlementRepository; +use aether_data_contracts::repository::settlement::{ + ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement, + UsagePolicyCostReservationState, +}; +use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord}; +use aether_data_contracts::DataLayerError; +use aether_runtime_state::RuntimeQueueStore; +use async_trait::async_trait; +use serde_json::json; + +use super::sample_usage as base_usage; +use crate::settlement::{UsageSettlementInput, UsageSettlementWriter}; +use crate::worker::write_event_record; +use crate::{ + UsageBillingEventEnricher, UsageEvent, UsageEventData, UsageEventType, UsageRecordWriter, + UsageRuntime, UsageRuntimeAccess, UsageRuntimeConfig, +}; + +const RESERVATION_TOKEN: &str = "550e8400-e29b-41d4-a716-446655440000"; + +fn sample_usage() -> StoredRequestUsageAudit { + let mut usage = base_usage(); + usage.request_metadata = Some(json!({"plan_usage_reservation_token": RESERVATION_TOKEN})); + usage +} + +fn runtime() -> UsageRuntime { + UsageRuntime::new(UsageRuntimeConfig { + enabled: true, + ..Default::default() + }) + .unwrap() +} + +#[derive(Default)] +enum ReconcileResponse { + #[default] + Exact, + Missing, + Changed(fn(&mut StoredUsagePolicyCostReservation)), + Error, +} + +#[derive(Default)] +struct ReuseStore { + repository: Option, + response: ReconcileResponse, + stored_override: Option, + fail_next_upsert: AtomicBool, + upserts: AtomicUsize, + reconciliations: Mutex>, + settlements: Mutex>, +} + +#[async_trait] +impl UsageSettlementWriter for ReuseStore { + fn has_usage_settlement_writer(&self) -> bool { + true + } + + async fn reconcile_usage_policy_cost( + &self, + input: ReconcileUsagePolicyCostInput, + ) -> Result, DataLayerError> { + input.validate()?; + self.reconciliations.lock().unwrap().push(input.clone()); + tokio::task::yield_now().await; + if let Some(repository) = self.repository.as_ref() { + return aether_data_contracts::repository::settlement::SettlementWriteRepository::reconcile_usage_policy_cost(repository, input).await; + } + let mut stored = StoredUsagePolicyCostReservation { + request_id: input.request_id, + subject_id: input.subject_id, + reservation_token: input.reservation_token, + admitted_at_unix_secs: 100, + reserved_cost_units: 100_000_000, + actual_cost_units: Some(input.actual_cost_units), + state: input.terminal_state, + reservation_expires_at_unix_secs: 500, + retain_until_unix_secs: 1_000, + finalized_at_unix_secs: Some(input.finalized_at_unix_secs), + }; + match self.response { + ReconcileResponse::Exact => {} + ReconcileResponse::Missing => return Ok(None), + ReconcileResponse::Changed(change) => change(&mut stored), + ReconcileResponse::Error => { + return Err(DataLayerError::TimedOut("reconciliation".to_string())); + } + } + Ok(Some(stored)) + } + + async fn settle_usage( + &self, + input: UsageSettlementInput, + ) -> Result, DataLayerError> { + self.settlements.lock().unwrap().push(input.clone()); + if let Some(repository) = self.repository.as_ref() { + return aether_data_contracts::repository::settlement::SettlementWriteRepository::settle_usage(repository, input).await; + } + Ok(None) + } +} + +#[async_trait] +impl UsageRecordWriter for ReuseStore { + async fn upsert_usage_record( + &self, + record: UpsertUsageRecord, + ) -> Result, DataLayerError> { + self.upserts.fetch_add(1, Ordering::Relaxed); + if self.fail_next_upsert.swap(false, Ordering::Relaxed) { + return Err(DataLayerError::TimedOut("upsert".to_string())); + } + if let Some(stored) = self.stored_override.as_ref() { + return Ok(Some(stored.clone())); + } + let mut stored = sample_usage(); + stored.request_id = record.request_id; + stored.user_id = record.user_id; + stored.api_key_id = record.api_key_id; + stored.provider_id = record.provider_id; + stored.status = record.status; + stored.billing_status = record.billing_status; + stored.total_cost_usd = record.total_cost_usd.unwrap_or_default(); + stored.actual_total_cost_usd = record.actual_total_cost_usd.unwrap_or_default(); + stored.request_metadata = record.request_metadata; + stored.updated_at_unix_secs = record.updated_at_unix_secs; + stored.finalized_at_unix_secs = record.finalized_at_unix_secs; + Ok(Some(stored)) + } +} + +#[async_trait] +impl UsageBillingEventEnricher for ReuseStore { + async fn enrich_usage_event(&self, _event: &mut UsageEvent) -> Result<(), DataLayerError> { + Ok(()) + } +} + +impl UsageRuntimeAccess for ReuseStore { + fn has_usage_writer(&self) -> bool { + true + } + + fn has_usage_worker_queue(&self) -> bool { + false + } + + fn usage_worker_queue(&self) -> Option> { + None + } +} + +fn event() -> UsageEvent { + let mut event = UsageEvent::new( + UsageEventType::Completed, + "req-1", + UsageEventData { + user_id: Some("user-1".to_string()), + api_key_id: Some("key-1".to_string()), + provider_name: "openai".to_string(), + model: "gpt-5".to_string(), + total_cost_usd: Some(1.25), + actual_total_cost_usd: Some(0.75), + request_metadata: Some(json!({"plan_usage_reservation_token": RESERVATION_TOKEN})), + ..Default::default() + }, + ); + event.timestamp_ms = 200_999; + event +} + +async fn write(store: &ReuseStore, event: UsageEvent, direct: bool) { + if direct { + runtime().record_terminal_event_direct(store, event).await; + } else { + write_event_record(store, &event).await.unwrap(); + } +} + +#[tokio::test] +async fn worker_and_direct_writes_reuse_confirmed_reservation_and_still_settle_wallet() { + for direct in [false, true] { + let store = ReuseStore::default(); + write(&store, event(), direct).await; + assert_eq!(store.upserts.load(Ordering::Relaxed), 1); + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 1, "direct={direct}"); + assert_eq!(reconciliations[0].actual_cost_units, 75_000_000); + assert_eq!(reconciliations[0].finalized_at_unix_secs, 200); + let settlements = store.settlements.lock().unwrap(); + assert_eq!(settlements.len(), 1); + assert_eq!(settlements[0].request_id, "req-1"); + assert_eq!(settlements[0].actual_total_cost_usd, 0.75); + } +} + +#[tokio::test] +async fn missing_or_different_reconciliation_results_keep_stored_usage_reconciliation() { + let changes: [fn(&mut StoredUsagePolicyCostReservation); 9] = [ + |row| row.request_id = "other-request".to_string(), + |row| row.subject_id = "other-user".to_string(), + |row| row.reservation_token = "other-token".to_string(), + |row| row.actual_cost_units = Some(1), + |row| row.actual_cost_units = None, + |row| row.state = UsagePolicyCostReservationState::Reserved, + |row| row.state = UsagePolicyCostReservationState::Released, + |row| row.finalized_at_unix_secs = Some(199), + |row| row.finalized_at_unix_secs = None, + ]; + for direct in [false, true] { + for response in std::iter::once(ReconcileResponse::Missing) + .chain(changes.into_iter().map(ReconcileResponse::Changed)) + { + let store = ReuseStore { + response, + ..Default::default() + }; + write(&store, event(), direct).await; + assert_eq!(store.reconciliations.lock().unwrap().len(), 2); + assert_eq!(store.settlements.lock().unwrap().len(), 1); + } + } +} + +#[tokio::test] +async fn changed_stored_usage_is_reconciled_using_its_own_identity_cost_and_terminal_state() { + let changes: [fn(&mut StoredRequestUsageAudit); 6] = [ + |row| row.request_id = "other-request".to_string(), + |row| row.user_id = Some("other-user".to_string()), + |row| { + row.request_metadata.as_mut().unwrap()["plan_usage_reservation_token"] = + json!("other-token") + }, + |row| row.actual_total_cost_usd = 0.25, + |row| row.status = "failed".to_string(), + |row| row.finalized_at_unix_secs = Some(199), + ]; + for direct in [false, true] { + for change in changes { + let mut stored = sample_usage(); + change(&mut stored); + let store = ReuseStore { + stored_override: Some(stored.clone()), + ..Default::default() + }; + write(&store, event(), direct).await; + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 2); + assert_eq!(reconciliations[1].request_id, stored.request_id); + assert_eq!( + reconciliations[1].subject_id, + stored.user_id.as_ref().unwrap().as_str() + ); + assert_eq!( + reconciliations[1].reservation_token, + stored.request_metadata.as_ref().unwrap()["plan_usage_reservation_token"] + .as_str() + .unwrap() + ); + assert_ne!(reconciliations[0], reconciliations[1]); + let settlements = store.settlements.lock().unwrap(); + assert_eq!(settlements.len(), 1); + assert_eq!( + settlements[0].actual_total_cost_usd, + stored.actual_total_cost_usd + ); + assert_eq!(settlements[0].status, stored.status); + } + } +} + +#[tokio::test] +async fn cancellation_release_billable_cancellation_and_zero_cost_preserve_settlement_rules() { + for direct in [false, true] { + for (event_type, billable_cancel, cost, terminal_state, wallets) in [ + ( + UsageEventType::Failed, + false, + 0.0, + UsagePolicyCostReservationState::Released, + 0, + ), + ( + UsageEventType::Cancelled, + false, + 0.75, + UsagePolicyCostReservationState::Released, + 0, + ), + ( + UsageEventType::Cancelled, + true, + 0.75, + UsagePolicyCostReservationState::Finalized, + 1, + ), + ( + UsageEventType::Completed, + false, + 0.0, + UsagePolicyCostReservationState::Finalized, + 1, + ), + ] { + let store = ReuseStore::default(); + let mut event = event(); + event.event_type = event_type; + event.data.actual_total_cost_usd = Some(cost); + event.data.request_metadata.as_mut().unwrap()["cancelled_request_fee"] = + json!(billable_cancel); + write(&store, event, direct).await; + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), 1); + assert_eq!(reconciliations[0].terminal_state, terminal_state); + assert_eq!( + reconciliations[0].actual_cost_units, + if billable_cancel { 75_000_000 } else { 0 } + ); + assert_eq!(store.settlements.lock().unwrap().len(), wallets); + } + } +} + +#[tokio::test] +async fn reconciliation_failure_stops_both_writes_before_upsert_and_wallet_settlement() { + for direct in [false, true] { + let store = ReuseStore { + response: ReconcileResponse::Error, + ..Default::default() + }; + if direct { + write(&store, event(), true).await; + } else { + assert!(write_event_record(&store, &event()).await.is_err()); + } + assert_eq!(store.reconciliations.lock().unwrap().len(), 1); + assert_eq!(store.upserts.load(Ordering::Relaxed), 0); + assert!(store.settlements.lock().unwrap().is_empty()); + } +} + +#[tokio::test] +async fn retry_after_upsert_failure_reconciles_again_before_settling() { + for direct in [false, true] { + let store = ReuseStore { + fail_next_upsert: AtomicBool::new(true), + ..Default::default() + }; + if direct { + write(&store, event(), true).await; + } else { + assert!(write_event_record(&store, &event()).await.is_err()); + } + assert!(store.settlements.lock().unwrap().is_empty()); + write(&store, event(), direct).await; + assert_eq!(store.reconciliations.lock().unwrap().len(), 2); + assert_eq!(store.upserts.load(Ordering::Relaxed), 2); + assert_eq!(store.settlements.lock().unwrap().len(), 1); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn concurrent_same_user_worker_and_direct_writes_each_reconcile_once() { + const REQUESTS: usize = 256; + let store = Arc::new(ReuseStore::default()); + let runtime = runtime(); + let barrier = Arc::new(tokio::sync::Barrier::new(REQUESTS)); + let mut tasks = tokio::task::JoinSet::new(); + for index in 0..REQUESTS { + let store = store.clone(); + let runtime = runtime.clone(); + let barrier = barrier.clone(); + tasks.spawn(async move { + let mut event = event(); + event.request_id = format!("reuse-concurrent-{index}"); + event.data.request_metadata.as_mut().unwrap()["plan_usage_reservation_token"] = + json!(format!("550e8400-e29b-41d4-a716-{index:012x}")); + barrier.wait().await; + if index % 2 == 0 { + runtime + .record_terminal_event_direct(store.as_ref(), event) + .await; + } else { + write_event_record(store.as_ref(), &event).await.unwrap(); + } + }); + } + tokio::time::timeout(std::time::Duration::from_secs(10), async { + while let Some(result) = tasks.join_next().await { + result.unwrap(); + } + }) + .await + .unwrap(); + let reconciliations = store.reconciliations.lock().unwrap(); + assert_eq!(reconciliations.len(), REQUESTS); + let unique_tokens: std::collections::HashSet<_> = reconciliations + .iter() + .map(|input| &input.reservation_token) + .collect(); + assert_eq!(unique_tokens.len(), REQUESTS); + assert_eq!(store.upserts.load(Ordering::Relaxed), REQUESTS); + let settlements = store.settlements.lock().unwrap(); + assert_eq!(settlements.len(), REQUESTS); + let unique_requests: std::collections::HashSet<_> = + settlements.iter().map(|input| &input.request_id).collect(); + assert_eq!(unique_requests.len(), REQUESTS); + assert_eq!( + settlements + .iter() + .map(|input| input.actual_total_cost_usd) + .sum::(), + REQUESTS as f64 * 0.75 + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn concurrent_duplicate_delivery_debits_real_memory_wallet_only_once() { + use aether_data::repository::wallet::{ + InMemoryWalletRepository, StoredWalletSnapshot, WalletLookupKey, WalletReadRepository, + }; + use aether_data_contracts::repository::settlement::{ + ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, SettlementWriteRepository, + UsagePolicyCostWindow, + }; + + let wallet = StoredWalletSnapshot::new( + "wallet-1".to_string(), + Some("user-1".to_string()), + None, + 10.0, + 2.0, + "finite".to_string(), + "USD".to_string(), + "active".to_string(), + 0.0, + 0.0, + 0.0, + 0.0, + 100, + ) + .unwrap(); + let wallets = Arc::new(InMemoryWalletRepository::seed([wallet])); + let repository = InMemorySettlementRepository::from_wallet_repository(wallets.clone()); + let reservation = ReserveUsagePolicyCostInput { + request_id: "req-1".to_string(), + subject_id: "user-1".to_string(), + reservation_token: RESERVATION_TOKEN.to_string(), + admitted_at_unix_secs: 100, + reserved_cost_units: 100_000_000, + reservation_expires_at_unix_secs: 500, + retain_until_unix_secs: 1_000, + windows: vec![UsagePolicyCostWindow { + window_id: "window-1".to_string(), + starts_at_unix_secs: 0, + ends_at_unix_secs: 1_000, + limit_cost_units: 1_000_000_000, + }], + }; + assert!(matches!( + repository + .reserve_usage_policy_cost(reservation.clone()) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::Allowed { .. } + )); + let store = Arc::new(ReuseStore { + repository: Some(repository), + ..Default::default() + }); + let runtime = runtime(); + let mut tasks = tokio::task::JoinSet::new(); + for index in 0..32 { + let store = store.clone(); + let runtime = runtime.clone(); + tasks.spawn(async move { + if index % 2 == 0 { + write_event_record(store.as_ref(), &event()).await.unwrap(); + } else { + runtime + .record_terminal_event_direct(store.as_ref(), event()) + .await; + } + }); + } + while let Some(result) = tasks.join_next().await { + result.unwrap(); + } + assert_eq!(store.reconciliations.lock().unwrap().len(), 32); + assert_eq!(store.settlements.lock().unwrap().len(), 32); + let wallet = wallets + .find(WalletLookupKey::UserId("user-1")) + .await + .unwrap() + .unwrap(); + assert_eq!(wallet.balance + wallet.gift_balance, 11.25); + assert_eq!(wallet.total_consumed, 0.75); + assert!(matches!( + store + .repository + .as_ref() + .unwrap() + .reserve_usage_policy_cost(reservation) + .await + .unwrap(), + ReserveUsagePolicyCostOutcome::AlreadyTerminal { + state: UsagePolicyCostReservationState::Finalized + } + )); +} diff --git a/crates/aether-usage/runtime/src/shutdown.rs b/crates/aether-usage/runtime/src/shutdown.rs new file mode 100644 index 000000000..801bace2f --- /dev/null +++ b/crates/aether-usage/runtime/src/shutdown.rs @@ -0,0 +1,99 @@ +use std::future::Future; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use tokio::sync::watch; +use tokio::task::JoinHandle; + +#[derive(Debug, Default)] +pub(crate) struct UsageBackgroundTasks { + handles: Mutex>>, +} + +impl UsageBackgroundTasks { + pub(crate) fn spawn(&self, task: impl Future + Send + 'static) { + self.handles + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .push(crate::executor::spawn_on_usage_background_runtime(task)); + } + + pub(crate) async fn stop_idle(&self) { + let handles = std::mem::take( + &mut *self + .handles + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()), + ); + for handle in &handles { + handle.abort(); + } + for handle in handles { + let _ = handle.await; + } + } +} + +impl Drop for UsageBackgroundTasks { + fn drop(&mut self) { + for handle in self.handles.get_mut().unwrap_or_else(|p| p.into_inner()) { + handle.abort(); + } + } +} + +#[derive(Debug)] +pub(crate) struct UsageShutdownState { + pub(crate) producers: Arc, + pub(crate) drain: watch::Sender, + pub(crate) tasks: UsageBackgroundTasks, + pub(crate) worker_control: crate::worker::UsageWorkerControl, + pub(crate) supervisors: Arc, + pub(crate) lock: tokio::sync::Mutex<()>, +} + +impl Default for UsageShutdownState { + fn default() -> Self { + Self { + producers: Arc::new(AtomicUsize::new(0)), + drain: watch::channel(false).0, + tasks: UsageBackgroundTasks::default(), + worker_control: crate::worker::UsageWorkerControl::default(), + supervisors: Arc::new(AtomicUsize::new(0)), + lock: tokio::sync::Mutex::new(()), + } + } +} + +/// Retain across an owned request finalizer, including any detached handoff. +#[derive(Debug)] +pub struct UsageProducerGuard(pub(crate) Arc); + +impl Drop for UsageProducerGuard { + fn drop(&mut self) { + self.0.fetch_sub(1, Ordering::AcqRel); + } +} + +pub(crate) async fn wait_for_drain(signal: &mut watch::Receiver) { + loop { + if *signal.borrow_and_update() { + return; + } + if signal.changed().await.is_err() { + std::future::pending::<()>().await; + } + } +} + +pub(crate) async fn retry_delay(delay: Duration, signal: &mut watch::Receiver) { + if *signal.borrow_and_update() { + tokio::time::sleep(delay.min(Duration::from_millis(100))).await; + } else { + tokio::select! { + _ = tokio::time::sleep(delay) => {}, + _ = wait_for_drain(signal) => {}, + } + } +} diff --git a/crates/aether-usage/runtime/src/worker.rs b/crates/aether-usage/runtime/src/worker.rs index d894cefb9..64ae03f41 100644 --- a/crates/aether-usage/runtime/src/worker.rs +++ b/crates/aether-usage/runtime/src/worker.rs @@ -9,20 +9,30 @@ use async_trait::async_trait; use tokio::sync::{mpsc, Notify}; use tracing::warn; +use crate::event_capture_budget::{shared_capture_memory_budget, EventCaptureMemoryBudget}; use crate::executor::spawn_on_usage_background_runtime; use crate::keyed_lock::KeyedAsyncLockPool; +use crate::queue::UsageDeadLetterOutcome; use crate::runtime::{ UsageBillingEventEnricher, UsageRuntimeAccess, UsageWorkerRecordConcurrencyGate, }; +use crate::settlement::{ + reconcile_usage_policy_cost_for_event_with_result, settle_usage_with_reconciled_cost, +}; use crate::{ - build_upsert_usage_record_from_event, reconcile_usage_policy_cost_for_event, - settle_usage_if_needed, UsageEvent, UsageEventType, UsageQueue, UsageRuntimeConfig, - UsageSettlementWriter, + build_upsert_usage_record_from_event, UsageEvent, UsageEventType, UsageQueue, + UsageRuntimeConfig, UsageSettlementWriter, }; const USAGE_WORKER_DB_PRESSURE_DEFER_MS: u64 = 10; const USAGE_WORKER_ACK_CHUNK_SIZE: usize = 100; +enum EntryDisposition { + NeedsAck, + Complete, + Deferred(DataLayerError), +} + #[async_trait] pub trait UsageEventRecorder: Send + Sync { async fn record_usage_event(&self, event: &UsageEvent) -> Result<(), DataLayerError>; @@ -151,7 +161,7 @@ where let request_lock = usage_request_lock(&event.request_id); let _guard = request_lock.lock().await; let mut event = event.clone(); - enrich_terminal_event(self.data.as_ref(), &mut event).await; + enrich_terminal_event(self.data.as_ref(), &mut event).await?; write_event_record(self.data.as_ref(), &event).await } } @@ -171,9 +181,10 @@ pub struct UsageQueueWorker { control: Option, telemetry: Option>, config: UsageRuntimeConfig, + capture_memory_budget: Arc, } -#[derive(Clone, Default)] +#[derive(Debug, Clone, Default)] pub(crate) struct UsageWorkerControl { shutdown: Arc, shutdown_notify: Arc, @@ -182,6 +193,7 @@ pub(crate) struct UsageWorkerControl { impl UsageWorkerControl { pub(crate) fn request_shutdown(&self) { self.shutdown.store(true, Ordering::Release); + self.shutdown_notify.notify_waiters(); self.shutdown_notify.notify_one(); } @@ -189,9 +201,15 @@ impl UsageWorkerControl { self.shutdown.load(Ordering::Acquire) } - async fn wait_for_shutdown(&self) { - while !self.should_shutdown() { - self.shutdown_notify.notified().await; + pub(crate) async fn wait_for_shutdown(&self) { + loop { + let notified = self.shutdown_notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + if self.should_shutdown() { + return; + } + notified.await; } } } @@ -326,6 +344,7 @@ impl UsageQueueWorker { control: None, telemetry: None, config, + capture_memory_budget: shared_capture_memory_budget(), }) } @@ -343,6 +362,11 @@ impl UsageQueueWorker { spawn_on_usage_background_runtime(async move { self.run_forever().await }) } + pub(crate) fn with_shutdown(mut self, control: UsageWorkerControl) -> Self { + self.control = Some(control); + self + } + pub(crate) async fn run(self) { self.run_forever().await; } @@ -366,6 +390,7 @@ impl UsageQueueWorker { reclaim_interval.tick().await; let mut reclaim_due = false; + let mut reclaim_cursor = "0-0".to_string(); loop { if self.should_shutdown() { @@ -373,7 +398,7 @@ impl UsageQueueWorker { } let result = { - let mut read_future = Box::pin(self.queue.read_group(&self.consumer)); + let mut read_future = Box::pin(self.queue.read_group_reserved(&self.consumer)); loop { tokio::select! { biased; @@ -392,9 +417,14 @@ impl UsageQueueWorker { }; match result { - Ok(entries) => { - self.report_read(entries.len()); - if let Err(err) = self.process_entries(entries).await { + Ok(batch) => { + self.report_read(batch.entries.len(), batch.requested_count); + let reservation = batch.reservation; + let result = self.process_entries(batch.entries).await; + // The raw fields and decoded event must be gone, and ACK must finish, + // before another worker can reuse this batch's receive allowance. + drop(reservation); + if let Err(err) = result { self.report_process_failed(); warn!( event_name = "usage_worker_process_failed", @@ -427,7 +457,7 @@ impl UsageQueueWorker { if reclaim_due { reclaim_due = false; - self.reclaim_stale_entries().await; + self.reclaim_stale_entries(&mut reclaim_cursor).await; } } } @@ -439,11 +469,24 @@ impl UsageQueueWorker { } } - async fn reclaim_stale_entries(&self) { - match self.queue.claim_stale(&self.consumer, "0-0").await { - Ok(entries) => { - self.report_reclaimed(entries.len()); - if let Err(err) = self.process_entries(entries).await { + async fn reclaim_stale_entries(&self, cursor: &mut String) { + let result = tokio::select! { + biased; + _ = self.wait_for_shutdown() => return, + result = self.queue.claim_stale_page_reserved(&self.consumer, cursor) => result, + }; + match result { + Ok(batch) => { + let reservation = batch.reservation; + let page = batch.page; + // An empty page can still advance past a large active PEL prefix. + // Failed writes remain pending and are revisited after the scan wraps. + *cursor = page.next_start_id; + self.report_reclaimed(page.entries.len()); + let result = self.process_entries(page.entries).await; + drop(page.deleted_ids); + drop(reservation); + if let Err(err) = result { self.report_process_failed(); warn!( event_name = "usage_worker_reclaim_process_failed", @@ -475,11 +518,11 @@ impl UsageQueueWorker { .is_some_and(UsageWorkerControl::should_shutdown) } - fn report_read(&self, entries_read: usize) { + fn report_read(&self, entries_read: usize, requested_count: usize) { self.report(UsageWorkerObservation::read( self.worker_index, entries_read, - self.config.consumer_batch_size.max(1), + requested_count, )); } @@ -529,21 +572,27 @@ impl UsageQueueWorker { } let mut ack_ids = Vec::new(); + let mut deferred_error = None; for entry in entries { - match self.process_entry(&entry).await { - Ok(should_ack) => { - if should_ack { - ack_ids.push(entry.id.clone()); - if ack_ids.len() >= USAGE_WORKER_ACK_CHUNK_SIZE { - self.queue.ack_and_delete(&ack_ids).await?; - self.report_acked(ack_ids.len()); - ack_ids.clear(); - } + let id = entry.id.clone(); + let result = self.process_entry(entry).await; + match result { + Ok(EntryDisposition::NeedsAck) => { + ack_ids.push(id); + if ack_ids.len() >= USAGE_WORKER_ACK_CHUNK_SIZE { + self.acknowledge_entries(&ack_ids).await?; + ack_ids.clear(); } } + Ok(EntryDisposition::Complete) => {} + Ok(EntryDisposition::Deferred(err)) => { + // An entry that cannot fit the DLQ encoder must not indefinitely block + // the healthy entries returned with it on every reclaim pass. + deferred_error.get_or_insert(err); + } Err(err) => { if !ack_ids.is_empty() { - let _ = self.queue.ack_and_delete(&ack_ids).await; + let _ = self.acknowledge_entries(&ack_ids).await; } return Err(err); } @@ -551,37 +600,107 @@ impl UsageQueueWorker { } if !ack_ids.is_empty() { - self.queue.ack_and_delete(&ack_ids).await?; - self.report_acked(ack_ids.len()); + self.acknowledge_entries(&ack_ids).await?; } + match deferred_error { + Some(err) => Err(err), + None => Ok(()), + } + } + + async fn acknowledge_entries(&self, ids: &[String]) -> Result<(), DataLayerError> { + let acked = self.queue.ack_and_delete_counted(ids).await?; + if acked > 0 { + self.report_acked(acked); + } Ok(()) } - async fn process_entry(&self, entry: &RuntimeQueueEntry) -> Result { - let event = match UsageEvent::from_stream_fields(&entry.fields) { - Ok(event) => event, - Err(err) => { + async fn dead_letter_entry( + &self, + entry: RuntimeQueueEntry, + error: DataLayerError, + event_name: &'static str, + ) -> Result { + let id = entry.id.clone(); + let outcome = self + .queue + .transfer_dead_letter_owned(entry, error.to_string()) + .await?; + let (destination_id, disposition) = match outcome { + UsageDeadLetterOutcome::Transferred { + destination_id, + acked, + } => { + if acked > 0 { + self.report_acked(acked); + } + (destination_id, EntryDisposition::Complete) + } + UsageDeadLetterOutcome::Appended { destination_id } => { + (destination_id, EntryDisposition::NeedsAck) + } + UsageDeadLetterOutcome::NotPending => { warn!( - event_name = "usage_worker_entry_decode_dead_lettered", + event_name = "usage_worker_dead_letter_source_not_pending", log_type = "ops", worker_consumer = %self.consumer, worker_group = %self.config.consumer_group, - entry_id = %entry.id, - error = %err, - "usage worker moved malformed queue entry to dead letter" + entry_id = %id, + error = %error, + "usage worker skipped dead letter transfer because the source is no longer pending" ); - self.queue.push_dead_letter(entry, &err.to_string()).await?; - self.report_dead_lettered(1); - return Ok(true); + return Ok(EntryDisposition::Complete); + } + UsageDeadLetterOutcome::EncodingDeferred { error } => { + warn!( + event_name = "usage_worker_dead_letter_encoding_deferred", + log_type = "ops", + worker_consumer = %self.consumer, + worker_group = %self.config.consumer_group, + entry_id = %id, + error = %error, + "usage worker retained entry pending after dead letter encoding failed" + ); + return Ok(EntryDisposition::Deferred(error)); + } + }; + self.report_dead_lettered(1); + warn!( + event_name, + log_type = "ops", + worker_consumer = %self.consumer, + worker_group = %self.config.consumer_group, + entry_id = %id, + dead_letter_id = %destination_id, + error = %error, + "usage worker appended queue entry to dead letter" + ); + Ok(disposition) + } + + async fn process_entry( + &self, + entry: RuntimeQueueEntry, + ) -> Result { + let event = match UsageEvent::from_stream_fields_with_capture_budget( + &entry.fields, + Arc::clone(&self.capture_memory_budget), + ) { + Ok(event) => event, + Err(err) => { + return self + .dead_letter_entry(entry, err, "usage_worker_entry_decode_dead_lettered") + .await; } }; match self.recorder.record_usage_event(&event).await { - Ok(()) => Ok(true), + Ok(()) => Ok(EntryDisposition::NeedsAck), Err(err) if usage_event_record_error_is_permanent(&err) => { warn!( - event_name = "usage_worker_entry_record_dead_lettered", + event_name = "usage_worker_entry_record_permanent_failed", log_type = "ops", worker_consumer = %self.consumer, worker_group = %self.config.consumer_group, @@ -595,11 +714,11 @@ impl UsageQueueWorker { provider_endpoint_id = event.data.provider_endpoint_id.as_deref().unwrap_or(""), provider_api_key_id = event.data.provider_api_key_id.as_deref().unwrap_or(""), error = %err, - "usage worker moved non-retryable usage event to dead letter" + "usage worker encountered a non-retryable usage record failure" ); - self.queue.push_dead_letter(entry, &err.to_string()).await?; - self.report_dead_lettered(1); - Ok(true) + drop(event); + self.dead_letter_entry(entry, err, "usage_worker_entry_record_dead_lettered") + .await } Err(err) => { warn!( @@ -677,17 +796,17 @@ pub async fn write_event_record(data: &T, event: &UsageEvent) -> Result<(), D where T: UsageRecordWriter + UsageSettlementWriter + Send + Sync, { - reconcile_usage_policy_cost_for_event(data, event).await?; + let reconciled = reconcile_usage_policy_cost_for_event_with_result(data, event).await?; let record = build_upsert_usage_record_from_event(event)?; if let Some(stored) = data.upsert_usage_record(record).await? { - settle_usage_if_needed(data, &stored).await?; + settle_usage_with_reconciled_cost(data, &stored, reconciled).await?; } // Manual proxy traffic is counted at the actual transport-attempt boundary. Usage events are // replayable, so emitting that side effect here would count normal requests and reclaims twice. Ok(()) } -async fn enrich_terminal_event(data: &T, event: &mut UsageEvent) +async fn enrich_terminal_event(data: &T, event: &mut UsageEvent) -> Result<(), DataLayerError> where T: UsageBillingEventEnricher + Send + Sync, { @@ -695,7 +814,7 @@ where event.event_type, UsageEventType::Completed | UsageEventType::Failed | UsageEventType::Cancelled ) { - return; + return Ok(()); } if let Err(err) = data.enrich_usage_event(event).await { @@ -707,7 +826,9 @@ where error = %err, "usage worker failed to enrich terminal usage event with billing" ); + return Err(err); } + Ok(()) } fn consumer_name(worker_index: Option) -> String { @@ -724,7 +845,10 @@ fn consumer_name(worker_index: Option) -> String { #[cfg(test)] mod tests { - use std::collections::BTreeMap; + mod dead_letter_transfer { + include!("worker_dead_letter_tests.rs"); + } + use std::collections::{BTreeMap, VecDeque}; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -733,11 +857,13 @@ mod tests { ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement, UsageSettlementInput, }; - use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord}; + use aether_data_contracts::repository::usage::{ + StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState, + }; use aether_data_contracts::DataLayerError; use aether_runtime_state::{ - MemoryRuntimeStateConfig, RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats, - RuntimeQueueStore, RuntimeState, + MemoryRuntimeStateConfig, RuntimeQueueEntry, RuntimeQueueReclaimConfig, + RuntimeQueueReclaimPage, RuntimeQueueStats, RuntimeQueueStore, RuntimeState, }; use async_trait::async_trait; use tokio::sync::Notify; @@ -747,6 +873,9 @@ mod tests { write_event_record, ManualProxyNodeCounter, UsageEventRecorder, UsageQueueWorker, UsageRecordWriter, UsageWorkerControl, }; + use crate::dead_letter_encoding::DeadLetterEncodingBudget; + use crate::event_capture_budget::EventCaptureMemoryBudget; + use crate::queue_read_budget::QueueReadBudget; use crate::runtime::UsageWorkerRecordConcurrencyGate; use crate::UsageBillingEventEnricher; use crate::{ @@ -760,14 +889,153 @@ mod tests { settlements: Mutex>, reconciliations: Mutex>, enrich_calls: Mutex>, + enrich_outcomes: Mutex>, manual_proxy_counter_calls: AtomicUsize, } + enum TestEnrichmentOutcome { + TimedOut, + Unpriced, + Priced { listed: f64, actual: f64 }, + } + + #[derive(Default)] + struct ControlledRecorder { + entered: Notify, + release: Notify, + calls: AtomicUsize, + } + + #[async_trait] + impl UsageEventRecorder for ControlledRecorder { + async fn record_usage_event(&self, _event: &UsageEvent) -> Result<(), DataLayerError> { + self.calls.fetch_add(1, Ordering::AcqRel); + self.entered.notify_one(); + self.release.notified().await; + Ok(()) + } + } + #[derive(Default)] struct SelectiveFailingRecorder { calls: Mutex>, } + enum CaptureRecordOutcome { + RetryOnce, + PermanentFailure, + Wait, + } + + struct CaptureBudgetRecorder { + budget: Arc, + outcome: CaptureRecordOutcome, + calls: AtomicUsize, + entered: Notify, + } + + #[async_trait] + impl UsageEventRecorder for CaptureBudgetRecorder { + async fn record_usage_event(&self, event: &UsageEvent) -> Result<(), DataLayerError> { + let call = self.calls.fetch_add(1, Ordering::AcqRel); + assert_eq!(event.data.input_tokens, Some(4)); + assert_eq!(event.data.output_tokens, Some(6)); + assert_eq!(event.data.total_tokens, Some(10)); + assert_eq!(event.data.cache_read_input_tokens, Some(0)); + assert_eq!(event.data.actual_total_cost_usd, Some(0.123)); + let metadata = event.data.request_metadata.as_ref().expect("billing facts"); + assert_eq!(metadata["requested_reasoning_effort"], "high"); + assert_eq!(metadata["provider_service_tier"], "priority"); + assert_eq!(metadata["provider_actual_service_tier"], "default"); + if event.data.response_body.is_some() { + let retained = self.budget.retained_bytes(); + assert!(retained > 0); + let recorder_copy = event.clone(); + assert!(recorder_copy.data.response_body.is_some()); + assert!(self.budget.retained_bytes() > retained); + drop(recorder_copy); + assert_eq!(self.budget.retained_bytes(), retained); + } else { + assert_eq!(self.budget.retained_bytes(), 0); + assert_eq!( + event.data.response_body_state, + Some(UsageBodyCaptureState::Truncated) + ); + } + self.entered.notify_one(); + match self.outcome { + CaptureRecordOutcome::RetryOnce if call == 0 => Err(DataLayerError::TimedOut( + "retry test database write".to_string(), + )), + CaptureRecordOutcome::RetryOnce => Ok(()), + CaptureRecordOutcome::PermanentFailure => Err(DataLayerError::UnexpectedValue( + "permanent capture test error".to_string(), + )), + CaptureRecordOutcome::Wait => std::future::pending().await, + } + } + } + + async fn capture_budget_worker( + budget_bytes: usize, + outcome: CaptureRecordOutcome, + ) -> ( + Arc, + UsageQueueWorker, + Arc, + ) { + let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let queue_runner: Arc = runner.clone(); + let budget = Arc::new(EventCaptureMemoryBudget::new(budget_bytes)); + let recorder = Arc::new(CaptureBudgetRecorder { + budget: Arc::clone(&budget), + outcome, + calls: AtomicUsize::new(0), + entered: Notify::new(), + }); + let config = UsageRuntimeConfig { + enabled: true, + stream_key: "usage:test:worker:capture".to_string(), + consumer_group: "usage:test:worker:capture-group".to_string(), + dlq_stream_key: "usage:test:worker:capture-dlq".to_string(), + consumer_batch_size: 1, + consumer_block_ms: 1, + ..UsageRuntimeConfig::default() + }; + let mut worker = UsageQueueWorker::new(queue_runner, recorder.clone(), config, None) + .expect("worker should build"); + worker.capture_memory_budget = budget; + worker.queue = + worker + .queue + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new( + 64 * 1024 * 1024, + 4, + ))); + worker + .queue + .ensure_consumer_group() + .await + .expect("consumer group"); + (runner, worker, recorder) + } + + fn captured_worker_event() -> UsageEvent { + let mut event = sample_event(); + event.data.model = "gpt-5.6-sol".to_string(); + event.data.endpoint_api_format = Some("openai:responses".to_string()); + event.data.request_body = Some(serde_json::json!({"reasoning": {"effort": "high"}})); + event.data.provider_request_body = Some(serde_json::json!({ + "model": "gpt-5.6-sol", "service_tier": "priority" + })); + event.data.response_body = Some(serde_json::json!({ + "service_tier": "Default", "output": "x".repeat(1024) + })); + event.data.cache_read_input_tokens = Some(0); + event.data.actual_total_cost_usd = Some(0.123); + event + } + #[derive(Default)] struct SlowUsageStore { active: std::sync::atomic::AtomicUsize, @@ -783,7 +1051,15 @@ mod tests { read_completed: AtomicUsize, release_read: Notify, reclaim_calls: AtomicUsize, + reclaim_pages: Mutex>>>, + reclaim_cursors: Mutex>, acked: AtomicBool, + requested_counts: Mutex>, + block_ack: AtomicBool, + ack_entered: Notify, + release_ack: Notify, + block_reclaim: AtomicBool, + reclaim_cancelled: AtomicBool, } impl ReadReclaimRaceProbeQueue { @@ -795,7 +1071,15 @@ mod tests { read_completed: AtomicUsize::new(0), release_read: Notify::new(), reclaim_calls: AtomicUsize::new(0), + reclaim_pages: Mutex::new(None), + reclaim_cursors: Mutex::new(Vec::new()), acked: AtomicBool::new(false), + requested_counts: Mutex::new(Vec::new()), + block_ack: AtomicBool::new(false), + ack_entered: Notify::new(), + release_ack: Notify::new(), + block_reclaim: AtomicBool::new(false), + reclaim_cancelled: AtomicBool::new(false), } } } @@ -838,10 +1122,14 @@ mod tests { _stream: &str, _group: &str, _consumer: &str, - _count: usize, + count: usize, _block_ms: Option, ) -> Result, DataLayerError> { let call_index = self.read_calls.fetch_add(1, Ordering::AcqRel); + self.requested_counts + .lock() + .expect("requested counts lock") + .push(count); let mut first_read_guard = (call_index == 0).then(|| FirstReadDropGuard { cancelled: &self.first_read_cancelled, completed: false, @@ -872,12 +1160,54 @@ mod tests { .collect()) } + async fn claim_stale_page( + &self, + stream: &str, + group: &str, + consumer: &str, + start_id: &str, + config: RuntimeQueueReclaimConfig, + ) -> Result { + self.reclaim_cursors + .lock() + .expect("reclaim cursors lock") + .push(start_id.to_string()); + if self.block_reclaim.load(Ordering::Acquire) { + let _guard = FirstReadDropGuard { + cancelled: &self.reclaim_cancelled, + completed: false, + }; + return std::future::pending().await; + } + let scripted = self + .reclaim_pages + .lock() + .expect("reclaim pages lock") + .as_mut() + .map(|pages| pages.pop_front().expect("scripted reclaim page")); + if let Some(page) = scripted { + self.reclaim_calls.fetch_add(1, Ordering::AcqRel); + return page; + } + Ok(RuntimeQueueReclaimPage { + next_start_id: "0-0".to_string(), + entries: self + .claim_stale(stream, group, consumer, start_id, config) + .await?, + deleted_ids: Vec::new(), + }) + } + async fn ack( &self, _stream: &str, _group: &str, ids: &[String], ) -> Result { + if self.block_ack.load(Ordering::Acquire) { + self.ack_entered.notify_one(); + self.release_ack.notified().await; + } if ids.iter().any(|id| id == &self.entry.id) { self.acked.store(true, Ordering::Release); Ok(1) @@ -1007,6 +1337,23 @@ mod tests { .lock() .expect("enrich calls lock") .push(event.request_id.clone()); + match self + .enrich_outcomes + .lock() + .expect("enrich outcomes lock") + .pop_front() + { + Some(TestEnrichmentOutcome::TimedOut) => { + return Err(DataLayerError::TimedOut("test pricing lookup".to_string())); + } + Some(TestEnrichmentOutcome::Unpriced) => return Ok(()), + Some(TestEnrichmentOutcome::Priced { listed, actual }) => { + event.data.total_cost_usd = Some(listed); + event.data.actual_total_cost_usd = Some(actual); + return Ok(()); + } + None => {} + } event.data.total_cost_usd = Some(0.456); Ok(()) } @@ -1241,6 +1588,140 @@ mod tests { assert_eq!(records[0].total_cost_usd, Some(0.456)); } + #[tokio::test] + async fn data_event_recorder_pricing_timeout_stays_pending_until_successful_reclaim() { + let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let store = Arc::new(TestUsageStore::default()); + store + .enrich_outcomes + .lock() + .expect("enrich outcomes lock") + .extend([ + TestEnrichmentOutcome::TimedOut, + TestEnrichmentOutcome::Priced { + listed: 0.456, + actual: 0.123, + }, + ]); + let config = UsageRuntimeConfig { + consumer_batch_size: 1, + consumer_block_ms: 1, + reclaim_count: 1, + reclaim_idle_ms: 1, + queue_payload_max_bytes: 4096, + ..UsageRuntimeConfig::default() + }; + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let mut worker = + build_usage_queue_worker_with_record_gate(runner, store.clone(), config, None, None) + .expect("worker should build"); + worker.queue = worker.queue.with_read_budget(Arc::clone(&budget)); + worker.queue.ensure_consumer_group().await.expect("group"); + let mut event = sample_event(); + event.data.request_metadata = Some(serde_json::json!({ + "plan_usage_reservation_token": "pricing-retry-reservation" + })); + worker.queue.enqueue(&event).await.expect("enqueue"); + let batch = worker + .queue + .read_group_reserved(&worker.consumer) + .await + .expect("read"); + let error = worker + .process_entries(batch.entries) + .await + .expect_err("pricing timeout must reach the worker"); + drop(batch.reservation); + assert!(matches!(error, DataLayerError::TimedOut(_))); + assert!(store.records.lock().expect("records lock").is_empty()); + assert!(store + .reconciliations + .lock() + .expect("reconciliations lock") + .is_empty()); + assert!(store + .settlements + .lock() + .expect("settlements lock") + .is_empty()); + let stats = worker.queue.stats().await.expect("pending stats"); + assert_eq!((stats.stream_length, stats.group_pending), (1, 1)); + assert_eq!( + worker + .queue + .dlq_stats() + .await + .expect("dlq stats") + .stream_length, + 0 + ); + assert_eq!(budget.snapshot().reserved_bytes, 0); + + let mut cursor = "0-0".to_string(); + tokio::time::timeout(Duration::from_secs(1), async { + loop { + tokio::time::sleep(Duration::from_millis(1)).await; + worker.reclaim_stale_entries(&mut cursor).await; + if !store.records.lock().expect("records lock").is_empty() { + break; + } + } + }) + .await + .expect("pending entry should become reclaimable"); + + let records = store.records.lock().expect("records lock"); + assert_eq!(records.len(), 1); + assert_eq!(records[0].total_cost_usd, Some(0.456)); + assert_eq!(records[0].actual_total_cost_usd, Some(0.123)); + assert_eq!(records[0].total_tokens, Some(10)); + drop(records); + let reconciliations = store.reconciliations.lock().expect("reconciliations lock"); + assert_eq!(reconciliations.len(), 1); + assert_eq!(reconciliations[0].actual_cost_units, 12_300_000); + assert_eq!( + reconciliations[0].reservation_token, + "pricing-retry-reservation" + ); + drop(reconciliations); + assert_eq!(store.settlements.lock().expect("settlements lock").len(), 1); + assert_eq!( + store.enrich_calls.lock().expect("enrich calls lock").len(), + 2 + ); + let stats = worker.queue.stats().await.expect("acked stats"); + assert_eq!((stats.stream_length, stats.group_pending), (0, 0)); + assert_eq!( + worker + .queue + .dlq_stats() + .await + .expect("dlq stats") + .stream_length, + 0 + ); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[tokio::test] + async fn data_event_recorder_successful_unpriced_enrichment_keeps_existing_write_behavior() { + let store = Arc::new(TestUsageStore::default()); + store + .enrich_outcomes + .lock() + .expect("enrich outcomes lock") + .push_back(TestEnrichmentOutcome::Unpriced); + let recorder = super::UsageDataEventRecorder::new(Arc::clone(&store)); + recorder + .record_usage_event(&sample_event()) + .await + .expect("missing pricing is a successful enrichment outcome"); + let records = store.records.lock().expect("records lock"); + assert_eq!(records.len(), 1); + assert_eq!(records[0].total_cost_usd, None); + assert_eq!(records[0].actual_total_cost_usd, None); + } + #[tokio::test] async fn data_event_recorder_skips_enrichment_for_lifecycle_event() { let store = Arc::new(TestUsageStore::default()); @@ -1375,6 +1856,121 @@ mod tests { assert_eq!(store.records.lock().expect("records lock").len(), 4); } + #[tokio::test] + async fn usage_worker_reclaim_cursor_advances_on_empty_pages_and_retries_read_errors() { + let event = sample_event(); + let entry = RuntimeQueueEntry { + id: "43-0".to_string(), + fields: event.to_stream_fields().expect("event fields"), + }; + let runner = Arc::new(ReadReclaimRaceProbeQueue::new(entry.clone())); + let page = |next_start_id: &str, entries, deleted_ids| RuntimeQueueReclaimPage { + next_start_id: next_start_id.to_string(), + entries, + deleted_ids, + }; + *runner.reclaim_pages.lock().expect("reclaim pages lock") = Some(VecDeque::from([ + Ok(page("11-0", Vec::new(), vec!["9-0".to_string()])), + Err(DataLayerError::Redis("temporary reclaim error".to_string())), + Ok(page("42-0", Vec::new(), Vec::new())), + Ok(page("0-0", vec![entry], Vec::new())), + Ok(page("0-0", Vec::new(), Vec::new())), + ])); + let recorder = Arc::new(SelectiveFailingRecorder::default()); + let worker = UsageQueueWorker::new( + runner.clone(), + recorder.clone(), + UsageRuntimeConfig::default(), + None, + ) + .expect("cursor worker"); + let mut cursor = "0-0".to_string(); + for expected in ["11-0", "11-0", "42-0", "0-0", "0-0"] { + worker.reclaim_stale_entries(&mut cursor).await; + assert_eq!(cursor, expected); + } + assert_eq!( + runner + .reclaim_cursors + .lock() + .expect("reclaim cursors lock") + .as_slice(), + ["0-0", "11-0", "11-0", "42-0", "0-0"] + ); + assert_eq!( + recorder.calls.lock().expect("calls lock").as_slice(), + [event.request_id] + ); + assert!(runner.acked.load(Ordering::Acquire)); + assert_eq!(runner.reclaim_calls.load(Ordering::Acquire), 5); + } + + #[tokio::test] + async fn usage_worker_reclaim_cursor_advances_after_write_failure_and_revisits_on_wrap() { + struct FailOnceRecorder(AtomicUsize); + + #[async_trait] + impl UsageEventRecorder for FailOnceRecorder { + async fn record_usage_event(&self, _event: &UsageEvent) -> Result<(), DataLayerError> { + if self.0.fetch_add(1, Ordering::AcqRel) == 0 { + Err(DataLayerError::TimedOut( + "temporary write failure".to_string(), + )) + } else { + Ok(()) + } + } + } + + let entry = RuntimeQueueEntry { + id: "43-0".to_string(), + fields: sample_event().to_stream_fields().expect("event fields"), + }; + let runner = Arc::new(ReadReclaimRaceProbeQueue::new(entry.clone())); + *runner.reclaim_pages.lock().expect("reclaim pages lock") = Some(VecDeque::from([ + Ok(RuntimeQueueReclaimPage { + next_start_id: "50-0".to_string(), + entries: vec![entry.clone()], + deleted_ids: Vec::new(), + }), + Ok(RuntimeQueueReclaimPage { + next_start_id: "0-0".to_string(), + entries: Vec::new(), + deleted_ids: Vec::new(), + }), + Ok(RuntimeQueueReclaimPage { + next_start_id: "0-0".to_string(), + entries: vec![entry], + deleted_ids: Vec::new(), + }), + ])); + let recorder = Arc::new(FailOnceRecorder(AtomicUsize::new(0))); + let worker = UsageQueueWorker::new( + runner.clone(), + recorder.clone(), + UsageRuntimeConfig::default(), + None, + ) + .expect("cursor worker"); + let mut cursor = "0-0".to_string(); + worker.reclaim_stale_entries(&mut cursor).await; + assert_eq!(cursor, "50-0"); + assert!(!runner.acked.load(Ordering::Acquire)); + worker.reclaim_stale_entries(&mut cursor).await; + worker.reclaim_stale_entries(&mut cursor).await; + assert_eq!(cursor, "0-0"); + assert!(runner.acked.load(Ordering::Acquire)); + assert_eq!(recorder.0.load(Ordering::Acquire), 2); + assert_eq!( + runner + .reclaim_cursors + .lock() + .expect("reclaim cursors lock") + .as_slice(), + ["0-0", "50-0", "0-0"] + ); + } + #[tokio::test] async fn usage_worker_defers_reclaim_until_inflight_read_is_processed() { let event = sample_event(); @@ -1497,6 +2093,327 @@ mod tests { assert_eq!(queue.reclaim_calls.load(Ordering::Acquire), 0); } + fn receive_budget_worker_config() -> UsageRuntimeConfig { + UsageRuntimeConfig { + consumer_batch_size: 128, + consumer_block_ms: 60_000, + reclaim_interval_ms: 60_000, + queue_payload_max_bytes: 4096, + ..UsageRuntimeConfig::default() + } + } + + #[tokio::test] + async fn usage_worker_shared_read_budget_waits_for_slow_recorder_and_releases_on_cancel() { + let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let config = receive_budget_worker_config(); + let first_recorder = Arc::new(ControlledRecorder::default()); + let second_recorder = Arc::new(ControlledRecorder::default()); + let first_control = UsageWorkerControl::default(); + let (telemetry, _observations) = tokio::sync::mpsc::channel(16); + let mut first = UsageQueueWorker::new( + runner.clone(), + first_recorder.clone(), + config.clone(), + Some(0), + ) + .expect("first worker") + .with_supervisor(first_control.clone(), telemetry); + first.queue = first.queue.with_read_budget(Arc::clone(&budget)); + let queue = first.queue.clone(); + let mut second = UsageQueueWorker::new(runner, second_recorder.clone(), config, Some(1)) + .expect("second worker"); + second.queue = second.queue.with_read_budget(Arc::clone(&budget)); + for index in 0..2 { + let mut event = sample_event(); + event.request_id = format!("receive-budget-{index}"); + queue.enqueue(&event).await.expect("enqueue"); + } + + let first_handle = tokio::spawn(first.run()); + tokio::time::timeout(Duration::from_secs(1), first_recorder.entered.notified()) + .await + .expect("first recorder should hold a batch"); + assert!(budget.snapshot().reserved_bytes > 0); + let second_handle = tokio::spawn(second.run()); + tokio::time::timeout(Duration::from_secs(1), async { + while budget.snapshot().waiters == 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("second worker should wait before reading"); + assert_eq!(second_recorder.calls.load(Ordering::Acquire), 0); + let stats = queue.stats().await.expect("blocked stats"); + assert_eq!((stats.stream_length, stats.group_pending), (2, 1)); + + first_control.request_shutdown(); + first_recorder.release.notify_one(); + tokio::time::timeout(Duration::from_secs(1), first_handle) + .await + .expect("first worker should finish its acquired batch") + .expect("first worker task"); + tokio::time::timeout(Duration::from_secs(1), second_recorder.entered.notified()) + .await + .expect("second worker should acquire released allowance"); + assert!(budget.snapshot().reserved_bytes > 0); + second_handle.abort(); + assert!(second_handle + .await + .expect_err("cancel worker") + .is_cancelled()); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert_eq!(budget.snapshot().waiters, 0); + let stats = queue.stats().await.expect("cancelled stats"); + assert_eq!((stats.stream_length, stats.group_pending), (1, 1)); + } + + #[tokio::test] + async fn usage_worker_read_budget_survives_ack_and_reports_actual_requested_count() { + let runner = Arc::new(ReadReclaimRaceProbeQueue::new(RuntimeQueueEntry { + id: "7-0".to_string(), + fields: sample_event().to_stream_fields().expect("event fields"), + })); + runner.block_ack.store(true, Ordering::Release); + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let control = UsageWorkerControl::default(); + let (telemetry, mut observations) = tokio::sync::mpsc::channel(16); + let mut worker = UsageQueueWorker::new( + runner.clone(), + Arc::new(SelectiveFailingRecorder::default()), + receive_budget_worker_config(), + Some(3), + ) + .expect("worker") + .with_supervisor(control.clone(), telemetry); + worker.queue = worker.queue.with_read_budget(Arc::clone(&budget)); + runner.release_read.notify_one(); + let handle = tokio::spawn(worker.run()); + tokio::time::timeout(Duration::from_secs(1), runner.ack_entered.notified()) + .await + .expect("worker should reach ACK"); + assert!(budget.snapshot().reserved_bytes > 0); + assert!(!runner.acked.load(Ordering::Acquire)); + let observation = observations.recv().await.expect("read observation"); + assert_eq!(observation.worker_index, Some(3)); + assert_eq!(observation.entries_read, 1); + assert_eq!(observation.batch_size, 1); + assert_eq!( + runner + .requested_counts + .lock() + .expect("requested counts lock") + .as_slice(), + [1] + ); + + control.request_shutdown(); + runner.release_ack.notify_one(); + tokio::time::timeout(Duration::from_secs(1), handle) + .await + .expect("worker should finish ACK before shutdown") + .expect("worker task"); + assert!(runner.acked.load(Ordering::Acquire)); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[tokio::test] + async fn usage_worker_shutdown_cancels_reclaim_budget_wait_without_reading() { + use std::future::Future; + use std::task::Poll; + + let runner = Arc::new(ReadReclaimRaceProbeQueue::new(RuntimeQueueEntry { + id: "8-0".to_string(), + fields: sample_event().to_stream_fields().expect("event fields"), + })); + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let (_, occupied) = budget.reserve(1, 4096).await.expect("occupy budget"); + let control = UsageWorkerControl::default(); + let (telemetry, _observations) = tokio::sync::mpsc::channel(16); + let mut worker = UsageQueueWorker::new( + runner.clone(), + Arc::new(SelectiveFailingRecorder::default()), + receive_budget_worker_config(), + None, + ) + .expect("worker") + .with_supervisor(control.clone(), telemetry); + worker.queue = worker.queue.with_read_budget(Arc::clone(&budget)); + let mut cursor = "8-0".to_string(); + let mut reclaim = Box::pin(worker.reclaim_stale_entries(&mut cursor)); + std::future::poll_fn(|cx| { + assert!(reclaim.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + assert_eq!(budget.snapshot().waiters, 1); + assert!(runner + .reclaim_cursors + .lock() + .expect("cursors lock") + .is_empty()); + control.request_shutdown(); + tokio::time::timeout(Duration::from_secs(1), &mut reclaim) + .await + .expect("shutdown should cancel budget wait"); + drop(reclaim); + assert_eq!(cursor, "8-0"); + assert_eq!(budget.snapshot().waiters, 0); + assert_eq!(budget.snapshot().reserved_bytes, 4096); + drop(occupied); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[tokio::test] + async fn usage_worker_shutdown_cancels_inflight_reclaim_and_releases_budget() { + use std::future::Future; + use std::task::Poll; + + let runner = Arc::new(ReadReclaimRaceProbeQueue::new(RuntimeQueueEntry { + id: "9-0".to_string(), + fields: sample_event().to_stream_fields().expect("event fields"), + })); + runner.block_reclaim.store(true, Ordering::Release); + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let control = UsageWorkerControl::default(); + let (telemetry, _observations) = tokio::sync::mpsc::channel(16); + let mut worker = UsageQueueWorker::new( + runner.clone(), + Arc::new(SelectiveFailingRecorder::default()), + receive_budget_worker_config(), + None, + ) + .expect("worker") + .with_supervisor(control.clone(), telemetry); + worker.queue = worker.queue.with_read_budget(Arc::clone(&budget)); + let mut cursor = "9-0".to_string(); + let mut reclaim = Box::pin(worker.reclaim_stale_entries(&mut cursor)); + std::future::poll_fn(|cx| { + assert!(reclaim.as_mut().poll(cx).is_pending()); + Poll::Ready(()) + }) + .await; + assert_eq!(budget.snapshot().reserved_bytes, 4096); + assert_eq!( + runner + .reclaim_cursors + .lock() + .expect("cursors lock") + .as_slice(), + ["9-0"] + ); + control.request_shutdown(); + tokio::time::timeout(Duration::from_secs(1), &mut reclaim) + .await + .expect("shutdown should cancel reclaim I/O"); + drop(reclaim); + assert_eq!(cursor, "9-0"); + assert!(runner.reclaim_cancelled.load(Ordering::Acquire)); + assert_eq!(budget.snapshot().reserved_bytes, 0); + assert!(!runner.acked.load(Ordering::Acquire)); + } + + #[tokio::test] + async fn usage_worker_shutdown_finishes_acquired_reclaim_before_releasing_budget() { + let runner = Arc::new(ReadReclaimRaceProbeQueue::new(RuntimeQueueEntry { + id: "10-0".to_string(), + fields: sample_event().to_stream_fields().expect("event fields"), + })); + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let recorder = Arc::new(ControlledRecorder::default()); + let control = UsageWorkerControl::default(); + let (telemetry, _observations) = tokio::sync::mpsc::channel(16); + let mut worker = UsageQueueWorker::new( + runner.clone(), + recorder.clone(), + receive_budget_worker_config(), + None, + ) + .expect("worker") + .with_supervisor(control.clone(), telemetry); + worker.queue = worker.queue.with_read_budget(Arc::clone(&budget)); + let handle = tokio::spawn(async move { + let mut cursor = "10-0".to_string(); + worker.reclaim_stale_entries(&mut cursor).await; + cursor + }); + tokio::time::timeout(Duration::from_secs(1), recorder.entered.notified()) + .await + .expect("reclaimed entry should reach recorder"); + control.request_shutdown(); + tokio::task::yield_now().await; + assert!(!handle.is_finished()); + assert!(budget.snapshot().reserved_bytes > 0); + assert!(!runner.acked.load(Ordering::Acquire)); + recorder.release.notify_one(); + let cursor = tokio::time::timeout(Duration::from_secs(1), handle) + .await + .expect("acquired page should finish processing") + .expect("reclaim task"); + assert_eq!(cursor, "0-0"); + assert!(runner.acked.load(Ordering::Acquire)); + assert_eq!(budget.snapshot().reserved_bytes, 0); + } + + #[tokio::test] + async fn usage_worker_read_budget_processes_oversized_historical_payload() { + let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); + let store = Arc::new(TestUsageStore::default()); + let budget = Arc::new(QueueReadBudget::new(4096, 4096)); + let control = UsageWorkerControl::default(); + let (telemetry, mut observations) = tokio::sync::mpsc::channel(16); + let mut worker = build_usage_queue_worker_with_record_gate( + runner.clone(), + store.clone(), + receive_budget_worker_config(), + None, + None, + ) + .expect("worker") + .with_supervisor(control.clone(), telemetry); + worker.queue = worker.queue.with_read_budget(Arc::clone(&budget)); + worker.capture_memory_budget = Arc::new(EventCaptureMemoryBudget::new(128 * 1024)); + let queue = worker.queue.clone(); + let mut event = sample_event(); + event.data.response_body = Some(serde_json::json!({"legacy": "x".repeat(8192)})); + let fields = event.to_stream_fields().expect("historical envelope"); + assert!(fields.values().map(String::len).sum::() > 4096); + runner + .append_fields_with_maxlen(&worker.config.stream_key, &fields, None) + .await + .expect("append pre-limit message"); + let handle = tokio::spawn(worker.run()); + tokio::time::timeout(Duration::from_secs(1), async { + while observations + .recv() + .await + .expect("worker observation") + .acked_entries + == 0 + {} + }) + .await + .expect("oversized message should be recorded and ACKed"); + control.request_shutdown(); + tokio::time::timeout(Duration::from_secs(1), handle) + .await + .expect("worker should stop") + .expect("worker task"); + let records = store.records.lock().expect("records lock"); + assert_eq!(records.len(), 1); + assert_eq!(records[0].response_body, event.data.response_body); + assert_eq!(records[0].total_tokens, Some(10)); + drop(records); + let snapshot = budget.snapshot(); + assert_eq!(snapshot.reserved_bytes, 0); + assert_eq!(snapshot.oversized_entries_total, 1); + assert_eq!(snapshot.oversized_batches_total, 1); + let stats = queue.stats().await.expect("acked stats"); + assert_eq!((stats.stream_length, stats.group_pending), (0, 0)); + assert_eq!(queue.dlq_stats().await.expect("dlq stats").stream_length, 0); + } + #[test] fn usage_event_record_error_classifies_permanent_failures() { assert!(usage_event_record_error_is_permanent( @@ -1516,6 +2433,199 @@ mod tests { )); } + #[tokio::test] + async fn capture_budget_retry_releases_decoded_lease_and_preserves_pending_payload() { + let (_runner, worker, recorder) = + capture_budget_worker(64 * 1024, CaptureRecordOutcome::RetryOnce).await; + worker + .queue + .enqueue(&captured_worker_event()) + .await + .expect("enqueue"); + let entries = worker + .queue + .read_group(&worker.consumer) + .await + .expect("read event"); + assert_eq!(entries.len(), 1); + let retry_entries = entries.clone(); + assert!(matches!( + worker.process_entries(entries).await, + Err(DataLayerError::TimedOut(_)) + )); + assert_eq!(recorder.budget.retained_bytes(), 0); + let stats = worker.queue.stats().await.expect("pending stats"); + assert_eq!(stats.stream_length, 1); + assert_eq!(stats.group_pending, 1); + assert_eq!( + worker + .queue + .dlq_stats() + .await + .expect("dlq stats") + .stream_length, + 0 + ); + + // Replay the same pending entry, as reclamation does, without a timing-dependent idle wait. + worker + .process_entries(retry_entries) + .await + .expect("retry should succeed"); + assert_eq!(recorder.calls.load(Ordering::Acquire), 2); + assert_eq!(recorder.budget.retained_bytes(), 0); + let stats = worker.queue.stats().await.expect("ack stats"); + assert_eq!(stats.stream_length, 0); + assert_eq!(stats.group_pending, 0); + } + + #[tokio::test] + async fn capture_budget_downgrade_dead_letter_keeps_exact_original_fields() { + let (runner, worker, recorder) = + capture_budget_worker(0, CaptureRecordOutcome::PermanentFailure).await; + let mut original = captured_worker_event() + .to_stream_fields() + .expect("wire serialization"); + original.insert( + "legacy_marker".to_string(), + "preserve this field".to_string(), + ); + runner + .append_fields_with_maxlen(&worker.config.stream_key, &original, None) + .await + .expect("enqueue raw fields"); + let entries = worker + .queue + .read_group(&worker.consumer) + .await + .expect("read event"); + worker + .process_entries(entries) + .await + .expect("permanent failure should dead letter"); + assert_eq!(recorder.calls.load(Ordering::Acquire), 1); + assert_eq!(recorder.budget.retained_bytes(), 0); + assert_eq!(recorder.budget.downgraded_total(), 1); + let stats = worker.queue.stats().await.expect("ack stats"); + assert_eq!(stats.stream_length, 0); + assert_eq!(stats.group_pending, 0); + runner + .ensure_consumer_group( + &worker.config.dlq_stream_key, + "capture-dlq-inspection", + "0-0", + ) + .await + .expect("dlq group"); + let dlq = runner + .read_group( + &worker.config.dlq_stream_key, + "capture-dlq-inspection", + "capture-inspector", + 1, + Some(1), + ) + .await + .expect("read dlq"); + assert_eq!(dlq.len(), 1); + let payload: serde_json::Value = + serde_json::from_str(&dlq[0].fields["payload"]).expect("dlq json"); + assert_eq!( + payload["fields"], + serde_json::to_value(original).expect("original fields json") + ); + assert_eq!( + payload["error"].as_str(), + Some("unexpected database value: permanent capture test error") + ); + } + + #[tokio::test] + async fn capture_budget_malformed_entry_dead_letters_without_recording() { + let (runner, worker, recorder) = + capture_budget_worker(1024, CaptureRecordOutcome::PermanentFailure).await; + let original = BTreeMap::from([( + "payload".to_string(), + "malformed legacy payload".to_string(), + )]); + runner + .append_fields_with_maxlen(&worker.config.stream_key, &original, None) + .await + .expect("enqueue raw fields"); + let entries = worker + .queue + .read_group(&worker.consumer) + .await + .expect("read event"); + worker + .process_entries(entries) + .await + .expect("malformed event should dead letter"); + assert_eq!(recorder.calls.load(Ordering::Acquire), 0); + assert_eq!(recorder.budget.retained_bytes(), 0); + assert_eq!(recorder.budget.downgraded_total(), 0); + let stats = worker.queue.stats().await.expect("ack stats"); + assert_eq!(stats.stream_length, 0); + assert_eq!(stats.group_pending, 0); + runner + .ensure_consumer_group( + &worker.config.dlq_stream_key, + "capture-dlq-inspection", + "0-0", + ) + .await + .expect("dlq group"); + let dlq = runner + .read_group( + &worker.config.dlq_stream_key, + "capture-dlq-inspection", + "capture-inspector", + 1, + Some(1), + ) + .await + .expect("read dlq"); + assert_eq!(dlq.len(), 1); + let payload: serde_json::Value = + serde_json::from_str(&dlq[0].fields["payload"]).expect("dlq json"); + assert_eq!( + payload["fields"], + serde_json::to_value(original).expect("original fields json") + ); + } + + #[tokio::test] + async fn capture_budget_cancelled_record_releases_lease_without_acknowledging() { + let (_runner, worker, recorder) = + capture_budget_worker(64 * 1024, CaptureRecordOutcome::Wait).await; + worker + .queue + .enqueue(&captured_worker_event()) + .await + .expect("enqueue"); + let entries = worker + .queue + .read_group(&worker.consumer) + .await + .expect("read event"); + let worker = Arc::new(worker); + let task_worker = Arc::clone(&worker); + let task = tokio::spawn(async move { task_worker.process_entries(entries).await }); + tokio::time::timeout(Duration::from_secs(1), recorder.entered.notified()) + .await + .expect("recorder should start"); + assert!(recorder.budget.retained_bytes() > 0); + task.abort(); + assert!(task + .await + .expect_err("task should be cancelled") + .is_cancelled()); + assert_eq!(recorder.budget.retained_bytes(), 0); + let stats = worker.queue.stats().await.expect("pending stats"); + assert_eq!(stats.stream_length, 1); + assert_eq!(stats.group_pending, 1); + } + #[tokio::test] async fn process_entries_dead_letters_permanent_record_error_and_continues() { let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())); @@ -1530,8 +2640,15 @@ mod tests { consumer_block_ms: 1, ..UsageRuntimeConfig::default() }; - let worker = UsageQueueWorker::new(queue_runner, recorder.clone(), config, None) + let mut worker = UsageQueueWorker::new(queue_runner, recorder.clone(), config, None) .expect("worker should build"); + worker.queue = + worker + .queue + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new( + 64 * 1024 * 1024, + 4, + ))); worker .queue .ensure_consumer_group() diff --git a/crates/aether-usage/runtime/src/worker_dead_letter_tests.rs b/crates/aether-usage/runtime/src/worker_dead_letter_tests.rs new file mode 100644 index 000000000..c387c5873 --- /dev/null +++ b/crates/aether-usage/runtime/src/worker_dead_letter_tests.rs @@ -0,0 +1,449 @@ +use super::*; +use crate::dead_letter_encoding::DeadLetterEncodingBudget; +use crate::worker::UsageWorkerObservation; +use aether_runtime_state::RuntimeQueueTransferOutcome; + +struct TransferProbe { + inner: RuntimeState, + native: bool, + fail_write: AtomicBool, + lose_reply: AtomicBool, + ack_calls: Mutex>>, + append_calls: AtomicUsize, +} + +impl TransferProbe { + fn new(native: bool) -> Self { + Self { + inner: RuntimeState::memory(MemoryRuntimeStateConfig::default()), + native, + fail_write: AtomicBool::new(false), + lose_reply: AtomicBool::new(false), + ack_calls: Mutex::new(Vec::new()), + append_calls: AtomicUsize::new(0), + } + } +} + +#[async_trait] +impl RuntimeQueueStore for TransferProbe { + async fn ensure_consumer_group( + &self, + stream: &str, + group: &str, + start_id: &str, + ) -> Result<(), DataLayerError> { + self.inner + .ensure_consumer_group(stream, group, start_id) + .await + } + + async fn append_fields_with_maxlen( + &self, + stream: &str, + fields: &BTreeMap, + maxlen: Option, + ) -> Result { + self.append_calls.fetch_add(1, Ordering::Relaxed); + if self.fail_write.load(Ordering::Acquire) { + return Err(DataLayerError::TimedOut("test append failure".to_string())); + } + self.inner + .append_fields_with_maxlen(stream, fields, maxlen) + .await + } + + async fn try_transfer_pending_to_stream( + &self, + source: &str, + group: &str, + entry_id: &str, + destination: &str, + fields: &BTreeMap, + ) -> Result, DataLayerError> { + if !self.native { + return Ok(None); + } + if self.fail_write.load(Ordering::Acquire) { + return Err(DataLayerError::TimedOut( + "test atomic transfer failure".to_string(), + )); + } + let result = self + .inner + .try_transfer_pending_to_stream(source, group, entry_id, destination, fields) + .await?; + if self.lose_reply.swap(false, Ordering::AcqRel) { + return Err(DataLayerError::TimedOut( + "test committed transfer reply lost".to_string(), + )); + } + Ok(result) + } + + async fn read_group( + &self, + stream: &str, + group: &str, + consumer: &str, + count: usize, + block_ms: Option, + ) -> Result, DataLayerError> { + self.inner + .read_group(stream, group, consumer, count, block_ms) + .await + } + + async fn claim_stale( + &self, + stream: &str, + group: &str, + consumer: &str, + start_id: &str, + config: RuntimeQueueReclaimConfig, + ) -> Result, DataLayerError> { + self.inner + .claim_stale(stream, group, consumer, start_id, config) + .await + } + + async fn ack( + &self, + stream: &str, + group: &str, + ids: &[String], + ) -> Result { + self.ack_calls.lock().expect("ack calls").push(ids.to_vec()); + self.inner.ack(stream, group, ids).await + } + + async fn delete(&self, stream: &str, ids: &[String]) -> Result { + self.inner.delete(stream, ids).await + } + + async fn stats( + &self, + stream: &str, + group: Option<&str>, + ) -> Result { + self.inner.stats(stream, group).await + } +} + +async fn transfer_worker( + native: bool, +) -> ( + Arc, + UsageQueueWorker, + Arc, + tokio::sync::mpsc::Receiver, +) { + let runner = Arc::new(TransferProbe::new(native)); + let recorder = Arc::new(SelectiveFailingRecorder::default()); + let (telemetry, observations) = tokio::sync::mpsc::channel(32); + let mut worker = UsageQueueWorker::new( + runner.clone(), + recorder.clone(), + UsageRuntimeConfig { + enabled: true, + consumer_batch_size: 10, + consumer_block_ms: 1, + ..UsageRuntimeConfig::default() + }, + None, + ) + .expect("worker") + .with_supervisor(UsageWorkerControl::default(), telemetry); + worker.queue = + worker + .queue + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new( + 64 * 1024 * 1024, + 4, + ))); + worker.queue.ensure_consumer_group().await.expect("group"); + (runner, worker, recorder, observations) +} + +async fn malformed_entries( + runner: &TransferProbe, + worker: &UsageQueueWorker, +) -> Vec { + runner + .inner + .append_fields_with_maxlen( + &worker.config.stream_key, + &BTreeMap::from([ + ("payload".to_string(), "malformed\u{0000}\n\"\\".to_string()), + ( + "legacy".to_string(), + "preserve all original fields".to_string(), + ), + ]), + None, + ) + .await + .expect("raw append"); + worker + .queue + .read_group(&worker.consumer) + .await + .expect("read") +} + +fn observed_totals( + observations: &mut tokio::sync::mpsc::Receiver, +) -> (usize, usize) { + let mut totals = (0, 0); + while let Ok(observation) = observations.try_recv() { + totals.0 += observation.acked_entries; + totals.1 += observation.dead_lettered_entries; + } + totals +} + +#[tokio::test] +async fn native_transfer_replay_does_not_append_or_ack_twice() { + let (runner, worker, recorder, mut observations) = transfer_worker(true).await; + let entries = malformed_entries(&runner, &worker).await; + worker + .process_entries(entries.clone()) + .await + .expect("transfer"); + worker.process_entries(entries).await.expect("stale replay"); + assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 0); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 1 + ); + assert_eq!(observed_totals(&mut observations), (1, 1)); + assert!(runner.ack_calls.lock().expect("acks").is_empty()); + assert_eq!(runner.append_calls.load(Ordering::Relaxed), 0); + assert!(recorder.calls.lock().expect("record calls").is_empty()); +} + +#[tokio::test] +async fn committed_transfer_lost_reply_retries_without_duplicate_or_false_metrics() { + let (runner, worker, _, mut observations) = transfer_worker(true).await; + let entries = malformed_entries(&runner, &worker).await; + runner.lose_reply.store(true, Ordering::Release); + assert!(matches!( + worker.process_entries(entries.clone()).await, + Err(DataLayerError::TimedOut(_)) + )); + worker.process_entries(entries).await.expect("retry"); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 1 + ); + assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 0); + assert_eq!(observed_totals(&mut observations), (0, 0)); + assert!(runner.ack_calls.lock().expect("acks").is_empty()); + assert_eq!(runner.append_calls.load(Ordering::Relaxed), 0); +} + +#[tokio::test] +async fn source_no_longer_pending_is_not_reported_as_archived_or_deleted() { + let (runner, worker, _, mut observations) = transfer_worker(true).await; + let entries = malformed_entries(&runner, &worker).await; + runner + .inner + .ack( + &worker.config.stream_key, + &worker.config.consumer_group, + &[entries[0].id.clone()], + ) + .await + .expect("external ack"); + worker + .process_entries(entries) + .await + .expect("no longer pending"); + assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 1); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 0 + ); + assert_eq!(observed_totals(&mut observations), (0, 0)); + assert!(runner.ack_calls.lock().expect("acks").is_empty()); +} + +#[tokio::test] +async fn failed_transfer_acknowledges_successful_prefix_and_preserves_suffix_for_retry() { + let (runner, worker, recorder, mut observations) = transfer_worker(true).await; + for request_id in ["prefix", "req-worker-poison", "suffix"] { + let mut event = sample_event(); + event.request_id = request_id.to_string(); + worker.queue.enqueue(&event).await.expect("enqueue"); + } + let entries = worker + .queue + .read_group(&worker.consumer) + .await + .expect("read batch"); + let retry = entries[1..].to_vec(); + let prefix_id = entries[0].id.clone(); + runner.fail_write.store(true, Ordering::Release); + assert!(worker.process_entries(entries).await.is_err()); + assert_eq!( + recorder.calls.lock().expect("calls").as_slice(), + ["prefix", "req-worker-poison"] + ); + assert_eq!( + *runner.ack_calls.lock().expect("acks"), + vec![vec![prefix_id]] + ); + assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 2); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 0 + ); + assert_eq!(observed_totals(&mut observations), (1, 0)); + runner.fail_write.store(false, Ordering::Release); + worker.process_entries(retry).await.expect("retry suffix"); + assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 0); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 1 + ); + assert_eq!(observed_totals(&mut observations), (2, 1)); + assert_eq!(runner.append_calls.load(Ordering::Relaxed), 3); +} + +#[tokio::test] +async fn legacy_transfer_fallback_only_acknowledges_after_append_succeeds() { + let (runner, worker, _, mut observations) = transfer_worker(false).await; + let entries = malformed_entries(&runner, &worker).await; + runner.fail_write.store(true, Ordering::Release); + assert!(worker.process_entries(entries.clone()).await.is_err()); + assert!(runner.ack_calls.lock().expect("acks").is_empty()); + assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 1); + assert_eq!(observed_totals(&mut observations), (0, 0)); + runner.fail_write.store(false, Ordering::Release); + worker.process_entries(entries).await.expect("legacy retry"); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 1 + ); + assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 0); + assert_eq!(runner.ack_calls.lock().expect("acks").len(), 1); + assert_eq!(observed_totals(&mut observations), (1, 1)); +} + +#[tokio::test] +async fn encoding_rejection_preserves_pending_original_until_budget_allows_retry() { + let (runner, mut worker, _, mut observations) = transfer_worker(true).await; + let small = Arc::new(DeadLetterEncodingBudget::new(1, 1)); + worker.queue = worker.queue.with_dead_letter_encoding_budget(small.clone()); + let entries = malformed_entries(&runner, &worker).await; + let original = entries[0].fields.clone(); + assert!(worker.process_entries(entries.clone()).await.is_err()); + assert_eq!(small.snapshot().reserved_bytes, 0); + assert_eq!(small.snapshot().active_jobs, 0); + assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 1); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 0 + ); + assert_eq!(observed_totals(&mut observations), (0, 0)); + worker.queue = worker + .queue + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(64 * 1024, 1))); + worker + .process_entries(entries) + .await + .expect("retry with capacity"); + runner + .inner + .ensure_consumer_group(&worker.config.dlq_stream_key, "inspect", "0-0") + .await + .expect("dlq group"); + let dlq = runner + .inner + .read_group( + &worker.config.dlq_stream_key, + "inspect", + "inspector", + 1, + Some(1), + ) + .await + .expect("dlq read"); + let payload: serde_json::Value = + serde_json::from_str(&dlq[0].fields["payload"]).expect("wire JSON"); + assert_eq!( + payload["fields"], + serde_json::to_value(original).expect("original JSON") + ); + assert_eq!(observed_totals(&mut observations), (1, 1)); +} + +#[tokio::test] +async fn normal_record_replay_reports_actual_ack_count() { + let (_, worker, _, mut observations) = transfer_worker(true).await; + worker + .queue + .enqueue(&sample_event()) + .await + .expect("enqueue"); + let entries = worker + .queue + .read_group(&worker.consumer) + .await + .expect("read"); + worker + .process_entries(entries.clone()) + .await + .expect("first record"); + worker.process_entries(entries).await.expect("replay"); + assert_eq!(observed_totals(&mut observations), (1, 0)); +} + +#[tokio::test] +async fn oversized_dead_letter_does_not_block_healthy_entries_in_the_same_batch() { + let (runner, mut worker, recorder, mut observations) = transfer_worker(true).await; + let bad = malformed_entries(&runner, &worker).await; + for request_id in ["healthy-first", "healthy-second"] { + let mut event = sample_event(); + event.request_id = request_id.to_string(); + worker.queue.enqueue(&event).await.expect("healthy enqueue"); + } + let mut entries = bad.clone(); + entries.extend( + worker + .queue + .read_group(&worker.consumer) + .await + .expect("healthy read"), + ); + worker.queue = worker + .queue + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(1, 1))); + assert!(worker.process_entries(entries).await.is_err()); + assert_eq!( + recorder.calls.lock().expect("calls").as_slice(), + ["healthy-first", "healthy-second"] + ); + let stats = worker.queue.stats().await.expect("stats"); + assert_eq!((stats.group_pending, stats.stream_length), (1, 1)); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 0 + ); + assert_eq!(observed_totals(&mut observations), (2, 0)); + assert!(worker.process_entries(bad.clone()).await.is_err()); + assert_eq!(observed_totals(&mut observations), (0, 0)); + worker.queue = worker + .queue + .with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(64 * 1024, 1))); + worker + .process_entries(bad) + .await + .expect("archive after raising budget"); + assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 0); + assert_eq!( + worker.queue.dlq_stats().await.expect("dlq").stream_length, + 1 + ); + assert_eq!(observed_totals(&mut observations), (1, 1)); +} diff --git a/crates/aether-usage/runtime/src/write.rs b/crates/aether-usage/runtime/src/write.rs index 497553936..0ec9e393e 100644 --- a/crates/aether-usage/runtime/src/write.rs +++ b/crates/aether-usage/runtime/src/write.rs @@ -581,6 +581,7 @@ fn build_lifecycle_usage_event_from_record( execution_path: record.execution_path, local_execution_runtime_miss_reason: record.local_execution_runtime_miss_reason, request_metadata: record.request_metadata, + capture_retention: record.capture_retention, ..UsageEventData::default() }, } @@ -1534,6 +1535,7 @@ fn build_lifecycle_usage_record_owned( }; Ok(UpsertUsageRecord { + capture_retention: Default::default(), request_id, user_id, api_key_id, @@ -1631,6 +1633,7 @@ fn build_lifecycle_usage_record_impl( }; Ok(UpsertUsageRecord { + capture_retention: Default::default(), request_id: seed.request_id.clone(), user_id: seed.user_id.clone(), api_key_id: seed.api_key_id.clone(), diff --git a/docs/operations/concurrency-design-audit-2026-09-09.md b/docs/operations/concurrency-design-audit-2026-09-09.md new file mode 100644 index 000000000..a81109632 --- /dev/null +++ b/docs/operations/concurrency-design-audit-2026-09-09.md @@ -0,0 +1,526 @@ +# Aether 并发设计审查 + +- 日期:2026-09-09 +- 代码基线:`361952ada` +- 背景:RPM 增加后服务出现卡死现象。 +- 范围:当前代码、默认配置、隔离本地复现。尚未取得故障实例、实际 RPM、线上配置或卡死时指标;以下是确认的代码问题及条件性风险,不代表已经确认本次生产故障根因。 +- 初次审查仅新增报告;后续本地修复状态见下文。未部署、未修改生产配置、未向线上发起压测。 + +## 第一轮修复状态 + +以下问题描述及行号对应审查基线 `361952ada`,不是修复后的代码位置。 + +- **P1-1 已修复:** 压缩上传按声明的压缩大小预留,未知长度上传从一个额度单元开始,随缓冲容量增长申请预算;解压计入同时存活的输入、中间输出和最终输出。扩容额度不足立即返回 503,避免多个请求各持部分额度互相等待;取消和失败释放额度。请求体完整读取默认超时改为 120 秒,显式配置 0 仍可关闭。该预算覆盖读取和解压的显式缓冲,不等于整个请求生命周期或进程 RSS 上限。 +- **P1-3 已修复:** SQL 改为 `FOR UPDATE OF user_plan_entitlements`,不同用户不再争抢共享套餐行,同一 entitlement 的扣费仍串行。套餐 overage 配置按当前语句快照读取,后续语句可读取已提交的配置变更。 +- **P1-4 已修复:** reqwest、h2c、wreq 和 tunnel 的流统一接入空闲读取期限;执行配置 `read_ms` 优先,否则使用 `AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS`,默认 300 秒,显式 0 可关闭。首包期限独立保留,网关 keepalive 和空数据帧不重置空闲计时;已收到成功终止事件的请求不会因随后空闲被改记为失败。 +- **后续待处理:** P1-2 Redis 调度扫描、P1-5 响应捕获全局预算、P1-6 审计压缩及 P1-7 长期额度聚合和 SQL 期限,本轮未修改。 + +本地验证: + +- `cargo test -p aether-gateway-frontdoor`:34 项通过。 +- 网关请求体、解压、超时、流结束及模块边界的针对性回归:73 项通过;另跑 Anthropic 原生流及直通兼容性回归:24 项通过。 +- `cargo test -p aether-data-postgres --lib settlement::tests`:6 项通过,1 项需要数据库的测试默认忽略;该测试另在隔离 PostgreSQL 14.17 中显式执行通过,覆盖不同用户并行、同用户阻塞、扣费余额及配置并发更新,临时实例已停止。 +- 修改文件格式检查和 `git diff --check` 通过。未执行全工作区测试或阶梯吞吐压测,尚不能据此给出修复后的 RPM 容量。 + +## 第二轮修复状态 + +- **P1-2 已处理调度中的管理扫描和无关指标读取:** 增加调度专用运行态入口,跳过全池 sticky 会话扫描、会话计数及 cooldown TTL;sticky 直达只查询当前绑定。成本窗口仅在成本限额、`cost_first` 或 `quota_balanced` 启用时读取,延迟窗口仅在 `latency_first` 启用时读取。默认 64 个候选且未启用这些策略时,消除原先 128 次历史窗口查询;这是代码路径比较,不是压测吞吐结论。管理查询仍保留原统计,调度成本检查不再套用管理显示的 key 数量截断。 +- **P1-5 已增加流式诊断捕获共享预算:** `AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES` 默认 128 MiB,覆盖 provider/client 捕获的实际容量及扩容时的新旧分配,显式 0 关闭此类捕获。预算不足时保留连续前缀并标为截断,不阻塞客户端传输;取消和终态释放额度。计费观察独立消费完整数据,补充主解析器停用后的用量恢复,并保留同步 JSON 转流的终态摘要。完成语义恢复后及时释放预读副本。 +- **P1-6 已处理单条及 pending 批量写入的审计压缩:** 准备工作移到事务前的 blocking 任务,进程内最多 4 个工作任务、32 个含等待的准入任务,等待工作槽最多 1 秒、等待执行结果最多 30 秒。超时和容量不足返回可重试 `TimedOut`;取消后的运行任务继续持有许可,且只准备数据,不会自行写数据库。序列化直接写入带 8 KiB 缓冲的 gzip,避免完整 JSON 中间副本及普通路径一次额外 body 克隆。事务内仍保留依赖旧记录的生命周期、清空、恢复和幂等判断。 + +本地验证: + +- PostgreSQL usage 模块:124 项通过,13 项需要外部条件的测试默认忽略。 +- 隔离 PostgreSQL 14.17 的 5 项真实回归通过,覆盖单条/批量完整审计读写、重复事件计数、过期终态 no-op 和捕获清空;临时数据库已停止。日志:`/tmp/aether-usage-concurrency.68ZESf/test.log`。 +- 第一批网关回归 68 项通过,包含真实 Redis 命令计数、管理统计保留、策略读取矩阵、超过管理显示上限的成本检查及调度结果一致性。 +- 最终流式模块及相关超时回归 187 项全部通过,覆盖捕获预算耗尽、并发预算释放、用量回退、累计更新及显式归零、协议转换、同步 JSON 桥接和非对象字段兼容。日志:`/tmp/aether-concurrency-round2-stream-final.log`。 +- 修改过的 Rust 文件格式检查和 `git diff --check` 通过。 + +仍需后续处理的边界: + +- 启用成本/延迟策略时仍读取原始窗口,尚未改成增量聚合或合并刷新。 +- 128 MiB 不包含独立协议/计费解析缓冲、终态 base64 编码或 usage 队列副本。计费用量回退保留的单条协议记录仍有 Basic 5 MiB / Full 64 MiB 上限,记录完成即释放;大量并发超长单记录仍可能放大内存。审计准备许可约束任务数,也不是任意大 usage 记录的字节预算。 +- P1-7 长期额度聚合和 SQL/锁等待期限,以及独立实例阶梯压测,本轮未处理。 + +## 第三轮修复状态 + +- **P1-7 普通 SQL 等待期限:** PostgreSQL 连接默认设置 `statement_timeout=30000ms`、`lock_timeout=3000ms`,由 `AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS` 和 `AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS` 覆盖,显式 0 关闭,非法值在创建连接池时拒绝。覆盖直接连接池查询和事务;这是单条语句期限,不是整笔事务总期限。超时 SQLSTATE 仍按原可重试错误处理。 +- **P1-7 多窗口精确查询:** 长期请求额度和成本额度将多个窗口的独立扫描合并为一条带 `FILTER` 的聚合查询,保留用户锁、时间边界、有效预留过滤、当前事件排除、拒绝顺序和幂等语义。未引入近似额度或缓存;仍需扫描最大覆盖范围内的历史记录。 +- **维护任务隔离:** schema migration 和历史 backfill 使用关闭普通期限的专用连接,所有退出路径均关闭连接,避免配置泄漏回请求池。日/小时统计、钱包每日聚合、用量统计重建与 VACUUM 使用 5 分钟语句期限和 30 秒锁等待期限。 +- **P1-2 有界 Redis 聚合:** 单个窗口最多 512 条记录时在 Redis 内精确聚合,只返回总量和正值样本数;每批最多 16 个独立脚本。更大的窗口或聚合失败回到原完整查询;没有缓存金额,也没有无界 Lua 聚合。超过 512 条的成本窗口仍有原查询开销,并增加一次计数探测。 +- **P1-5/P1-6 事件正文保留:** 已构造的 usage 事件在同步保存计费相关字段后、进入终态队列或等待正文策略前,按四份诊断 JSON 正文的堆内存估算申请共享预算。`AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES` 默认 128 MiB,显式 0 关闭正文保留。预算申请不等待,额度不足舍弃正文并标记 `Truncated`,保留费用、用量、终态和引用;事件副本单独申请额度,重试随事件保留额度,取消及释放事件时归还。异步策略读取仍在原来的终态执行和准入位置,维持提交、排序与异常隔离语义。 +- **减少中间副本:** usage envelope 和死信队列序列化改为借用正文及字段,避免序列化前的完整深拷贝。新增 `usage_runtime_event_capture_memory_budget_bytes`、`usage_runtime_event_capture_memory_retained_bytes` 和 `usage_runtime_event_capture_memory_downgraded_total` 指标;降级次数也可能包含随后按 Basic 策略清除的正文。 + +本地验证: + +- Redis 运行态 3 项真实实例回归通过,覆盖 512/513 条边界、多批 Key、`u64` 精度和饱和、脚本重载、窗口变化及已完成的并发写入。日志:`/tmp/aether-round3-redis-tests.log`。 +- PostgreSQL 全量普通测试 229 项通过、24 项默认忽略;跨模块 `cargo check -p aether-data` 通过,维护 future 的 `Send` 回归通过。隔离 PostgreSQL 14.17 显式验证空库迁移和 usage 读写、超时回滚与期限隔离、精确多窗口准入和幂等、共享套餐锁,共 4 项通过。临时库已停止;日志:`/tmp/aether-postgres-deadlines.7sZ8w1/`。 +- 用量模块完整回归 264 项全部通过,含 13 项正文预算测试及终态队列容量、排序、异常隔离回归。日志:`/tmp/aether-round3-usage-agent-tests.log`。 +- 最终网关调度、真实 Redis 窗口回退及指标回归 61 项全部通过;最终网关、mock upstream 和 seed 工具构建通过。日志:`/tmp/aether-round3-gateway-tests.log`、`/tmp/aether-round3-verified-build.log`、`/tmp/aether-round3-seed-final-build.log`。 +- 修复压力 seed 工具按网关密钥加密 provider/client 凭据,并用已有凭据比较更新接口处理随机密文的重复初始化;隔离空库连续初始化两次已通过。mock chat 流按 `stream_options.include_usage=true` 输出终态用量,保留故障截断和未请求用量时的原行为,9 项测试全部通过。日志:`/tmp/aether-round3-mock-tests.log`;最终 mock 构建日志:`/tmp/aether-round3-mock-final-build.log`。 + +### 隔离端到端验收 + +使用本机 debug 构建、独立空 PostgreSQL 14.17 和 Redis、4 个 Tokio worker、入口并发上限 128、8 个客户端 API key。上游为单个零价 mock provider,每请求 20 个 256 字节内容块,首字节延迟 30 ms、块间隔 100 ms,完整流约 2 秒。请求启用流式 usage,客户端必须收到完整响应及 `[DONE]`。各档采样间隔 500 ms,停止流量后最多等待 45 秒排空;指标通过管理员会话采集,避免管理 Token 使用计数干扰排空判断。 + +| 并发 | 请求数 | 成功数 | 实测 RPS | 首字节 P95 (ms) | 总耗时 P95 (ms) | 网关峰值 RSS (MiB) | 排空 (s) | +| --- | --- | --- | --- | --- | --- | --- | --- | +| 8 | 64 | 64 | 4.01 | 117 | 2069 | 107.0 | 8.7 | +| 24 | 192 | 192 | 12.04 | 67 | 2015 | 122.3 | 9.9 | +| 48 | 384 | 384 | 23.91 | 87 | 2036 | 143.3 | 9.8 | + +- 三档共 640 个请求,全部 HTTP 200 且完整 SSE 结束,无请求错误。数据库恰有 640 条 `completed`,每条 input/output/total tokens 均为 `1/20/21`,合计 `640/12800/13440`,所有首字节时间有值。 +- 三档排空所需指标齐全,最终 usage 队列 lag、pending、计数 outbox、终态提交和有序生命周期待处理数均为 0;采样未观察到 PostgreSQL 锁等待,未出现 usage worker 处理失败。最终事件正文预算保留量为 0,未触发正文预算降级。 +- RSS 随并发档位上升,排空时未立即回落;不能据此证明长期无内存增长。最高档约 1435 RPM 是此短时、约 2 秒 mock 流的实测值,不代表生产容量,也没有修复前的同条件吞吐基线。真实长流、大请求、多候选账号池、长期额度和现金扣费负载未在本次阶梯压测覆盖;额度和结算语义由前述独立数据库回归验证。 +- 最终产物:`/tmp/aether-round3-pressure.T3sq7J/`,包含每档 JSON、原始日志、最终指标及 `usage-integrity.txt`;脚本:`/tmp/aether-concurrency-round3-pressure.sh`,总日志:`/tmp/aether-round3-pressure-final.log`。临时网关、mock、采样代理、Redis 和 PostgreSQL 均已停止。修改文件格式检查及 `git diff --check` 通过,未执行全工作区测试,未部署线上。 + +本轮仍未覆盖的内存与计算边界: + +- 事件正文预算是 JSON 堆内存估算,不是进程 RSS 上限。仍不包含计费提取前的 seed、Redis 消费批次、数据库写入 DTO、压缩和序列化字符串,以及协议观察缓冲。 +- Redis 超大成本窗口仍走原始历史查询;长期 SQL 额度仍需扫描历史,尚未改为带迁移和对账的增量账本聚合。 + +## 第四轮修复状态 + +- **受限等待期间减少重复候选准备:** 普通指定模型和无指定模型能力查询保留 150 ms 等待窗口,分页候选路径保留首次受限后的 100 ms 窗口。首次完整准备确认仅被客户端 API Key 并发限制阻塞后,等待期间只读取并发判断所需的近期候选;恢复或到达期限后重新执行完整动态校验,不直接复用旧候选放行。持续阻塞且首次准备未耗尽期限时,完整准备收敛为首次和最终两次;若短暂恢复后重新受限,仍可在原期限内重新准备。等待窗口不重置,也不是包含数据库 I/O 的端到端硬超时。 +- **探测任务创建前合并:** 同一 AppState 的请求补探测按 provider 在 `spawn` 前合并;运行期间的重复触发只保留一次后续执行信号,完成和取消时清理本地条目。AppState clone 共享协调器,切换到不同 RuntimeState 时隔离。最多同时保留 1024 个活跃 provider,超过容量跳过此次 best-effort 补探测,不创建等待任务;周期基础探测继续运行,但不等价于请求触发的 Burst 补探测。 +- **跨实例退出交接:** 保留 Redis pending 和带 token 的锁协议,释放锁后复查新 pending,避免另一实例在退出窗口提交的补探测无人接手。读取 pending 失败时释放锁并停止本轮处理,避免残留 pending 导致重新争锁的忙循环。 +- **可观测性:** 增加 `pool_quota_probe_replenish_provider_capacity`、`pool_quota_probe_replenish_active_providers`、`pool_quota_probe_replenish_started_total`、`pool_quota_probe_replenish_coalesced_total` 和 `pool_quota_probe_replenish_capacity_rejected_total`。 + +本地验证: + +- `cargo check --locked -p aether-gateway --lib` 通过。日志:`/tmp/aether-round4-check.log`。 +- 网关调度、分页候选、探测、模块边界、指标及真实 HTTP 并发等待回归共 376 项全部通过,包含本轮新增的 11 项调度等待和 8 项探测协调测试。日志:`/tmp/aether-round4-gateway-tests.log`。 +- 8 个同时受限的请求,通过真实候选筛选函数的依赖调用计数,确认完整候选读取共 16 次;恢复后重新加载被替换或移除的候选,轻量读取及完整重验的错误均正常传播。分页持续受限仅重启一次扫描,首次查询耗时计入原重试预算,取消后停止查询。 +- 单 provider 的探测运行中并发触发 64 次,仅创建一个任务并额外执行一次补探测。测试覆盖不同 provider 并行、100 次本地退出交接竞争、未首次执行即取消、运行中取消、panic、容量释放和 runtime 绑定隔离。两个本地协调器共享 Memory Runtime 模拟跨实例 pending/锁交接,精确覆盖旧任务最终检查后、解锁前的新触发;没有声称本轮运行了多进程 Redis 压测。 +- `cargo build --locked -p aether-gateway --bin aether-gateway` 通过,最终可执行文件已更新;日志:`/tmp/aether-round4-build.log`。修改文件格式和 `git diff --check` 通过,测试进程已退出。本轮未重新执行阶梯吞吐压测,第三轮的 640 请求结果仅代表其当时构建;未执行全工作区测试,未部署线上。 + +边界: + +- API Key 并发检查仍使用全局最近 128 条候选中的同 Key 活跃候选行,依赖状态持久化及原 300 秒活跃窗口;未引入原子分布式许可,也未改变候选行计数为请求去重计数。 +- 探测协调为 best-effort;取消后的 Redis 锁仍按原 30 秒 TTL 释放,未增加续期或 exactly-once 保证。同步日志、超大历史窗口增量聚合、未覆盖的内存副本以及目标环境长期压测仍待处理。 + +## 第五轮修复状态 + +- **运行日志 I/O 脱离请求线程:** stdout 和滚动文件各使用独立专用写线程,保留 Pretty/JSON、Stdout/File/Both、动态日志过滤及原有文件权限检查。文件写入、flush 和轮转重开均在写线程执行;定期文件清理移到 blocking 任务,单次完成后才安排下一次清理。启动时仍同步校验日志目录和目标文件,配置错误正常拒绝启动。 +- **日志过载保护:** 每个目标最多排队 4096 条事件,复制后的日志正文最多保留 8 MiB,单条事件上限 256 KiB。字节预算包含生产者已预留的复制、排队和正在写入的正文;队列满、预算不足或单条过大时整条舍弃,不等待设备、不在请求线程同步回退 stderr。Both 两个目标独立接收和降级,可能保留不同的事件集合。 +- **退出排空:** 网关、隧道、两个运行示例及 13 个基准工具的最外层入口持有日志 guard,先结束 Tokio runtime 再关闭队列。升级/回滚的显式进程退出也调用关闭接口。关闭先停止全部目标接收,再在共用的 2 秒预算内等待已接收事件及最终 flush;阻塞中的系统 I/O 无法强制取消,超时后不无限 join 写线程。 +- **信号处理:** 网关 ready 后收到 SIGTERM/SIGINT 会结束运行函数并经过日志 guard;现有连接未被统一追踪,此路径仍按进程终止处理业务,不声称请求或用量队列优雅排空。启动过程中尚未进入信号等待时,以及 SIGKILL/abort,不保证排空。 +- **日志健康指标:** 网关与隧道指标端点增加 `logging_stdout_*` / `logging_file_*`,记录队列和字节上限、当前正文保留量、接收数、按原因区分的丢弃数、写入错误、线程 panic、关闭超时及线程状态。沿用服务指标命名空间前缀。关闭 API 返回成功仅说明线程处理完队列且最终 flush 成功,之前的写入错误仍需查看指标,不代表 fsync 持久化成功。 + +本地验证: + +- 网关、隧道、loadtools 和 integration 的所有 binary/example 入口通过 `cargo check --locked`,未增加新的第三方依赖;integration 的依赖清单补充已有共享 runtime。日志:`/tmp/aether-round5-entrypoints-check.log`。 +- 共享 runtime 44 项单元测试及 2 项进程集成测试全部通过,包含 11 项新 writer 测试、2 项轮转回归、12 种初始化/输出格式/动态过滤场景,以及真实 stdout 堵塞测试。8 个生产者同时写入时,慢设备不阻塞生产者;4096 条和字节预算、整条拒绝、错误/panic 回收、退出竞争、stdout/file 双向隔离及阻塞 flush 均已覆盖。每种实际格式并发写入 512 条事件,记录完整且无重复。日志:`/tmp/aether-round5-runtime-final-tests.log`。 +- 真实 stdout 测试保持子进程管道完全不读,确认 stdout 队列触发丢弃后,文件仍收到完整 JSON 尾记录;guard 超时返回后,标准进程退出也成功完成。该测试耗时约 2.36 秒,包含日志的 2 秒退出等待;没有关闭管道来人为解除阻塞。 +- 网关入口 61 项、隧道 197 项回归全部通过,日志:`/tmp/aether-round5-service-tests.log`。以上合计 304 项测试通过,没有计入重复运行的首批测试。 +- Linux root/capabilities 专用日志 fixture 已适配异步排空,当前 macOS 环境未执行;Unix 文件权限、符号链接、多硬链接及安全轮转的普通测试已通过。 +- 网关和隧道最终二进制构建通过,日志:`/tmp/aether-round5-service-build.log`。隔离 PostgreSQL 和新网关实例的 Both/JSON 验收通过,stdout 和文件各保留 15 条完整日志,均包含唯一的 starting、ready 和 shutdown 事件;实际指标端点两路队列容量、接收及运行状态正常,丢弃、写入错误、panic 和关闭超时计数为 0。SIGTERM 后约 16 ms 正常退出,该耗时仅代表健康设备和无在途代理请求的此次烟测。产物:`/tmp/aether-logging-smoke-qvzN9U/result.json`,脚本:`/tmp/aether-nonblocking-logging-smoke.mjs`。 +- 临时网关和 PostgreSQL 已全部停止,测试与构建进程均已退出。修改文件格式检查及 `git diff --check` 通过;本轮未重跑业务阶梯吞吐压测,未执行全工作区测试,未部署线上。 + +本轮边界: + +- 日志格式化、字段 Debug 展开和 JSON 序列化仍发生在调用线程,订阅器线程局部字符串也可能保留历史容量;8 MiB 预算只约束交给后台写入的正文,不包含这些临时对象、队列元数据或进程 RSS。 +- 运行日志在过载、I/O 错误或关闭超时下可能丢失,不能作为可靠计费账本;账务持久化流程不使用这条日志队列。2 秒仅限制日志关闭等待,不是整个服务退出期限,既有业务或 blocking 任务清理仍可能更久。 +- 超大历史窗口增量聚合、尚未覆盖的内存副本、完整请求优雅排空及目标环境长期压测仍需后续处理。本轮不调整线上配置,也不据此给出生产 RPM 容量。 + +## 第六轮修复状态 + +- **补齐诊断正文预算的所有权传递:** 将纯预算令牌放到 data contracts,运行时仍使用原环境变量、128 MiB 默认值和指标。同步及流式终态 seed 在等待提交前申请预算;Redis 解码后的事件重新纳入本进程预算;事件生成的数据库写入 DTO 继续持有对应额度。事件和受管理 DTO 的正文副本分别申请额度,释放正文后才归还;序列化跳过令牌,不改变队列协议或数据库字段。 +- **保留完整计费与终态语义:** 正常终态构建仍在 blocking 任务执行;预算不足时同步执行既有纯构建逻辑,先解析 token、显式 cache=0、图像估算、错误及终态,再舍弃诊断正文,随后进入原有有序提交和准入路径。保留原始终态观测时间,构建异常仍按原终态失败路径隔离。已有 `None`、`Disabled`、`Unavailable` 状态不被预算降级改写为 `Truncated`,避免破坏清空指令。 +- **旧队列消息兼容:** 解码前的原始字段继续保留到记录或死信处理完成,重试不提前 ACK,DLQ 保存原始字段。旧消息缺正文且缺 typed state 时保留元数据内已有缓存 TTL、tier 和请求事实,显式 `None` 仍清空,避免预算接线改变后续计费输入。 +- **数据库准备减少正文副本:** 单条与 pending 批量准备先移走四份正文和 headers,再复制两个存储投影需要的少量元数据,避免原先为清洗而深拷贝整份 DTO。预算随输入进入 blocking 压缩闭包,调用方取消不会提前释放仍存活的正文额度;压缩结束后原始 JSON 释放,压缩结果继续走原事务、审计和正文存储流程。 + +本地验证: + +- data contracts 226 项、PostgreSQL 231 项、data runtime 356 项、usage runtime 282 项普通测试通过,合计 1095 项;其中本轮新增 27 项,覆盖队列积压、并发 seed、预算拒绝与复制、token/图像计费事实、typed clear、legacy metadata、重试/DLQ,以及取消和 panic 后额度释放。另有 25 项需要外部条件的测试默认忽略。日志:`/tmp/aether-round6-data-tests.log`、`/tmp/aether-round6-runtime-tests.log`。最终时间采样和反向 DTO 转事件的所有权修正后,usage runtime 282 项再次全部通过:`/tmp/aether-round6-usage-final-tests.log`,不重复计入总数。 +- 隔离 PostgreSQL 14.17 的 5 项真实回归通过。完整审计测试现覆盖受预算管理的单条、普通 pending 批量及同 request 重复批量写入,四份正文读取一致,准备结束后只保留调用方原对象的额度,最终释放为 0;同时验证过期终态 no-op、辅助计数幂等和 typed `None` 清空。日志:`/tmp/aether-round6-postgres.K8Aq5W/`;临时数据库已停止。 +- 网关、隧道、loadtools 和 integration 的 binary/example 入口通过 `cargo check --locked`,日志:`/tmp/aether-round6-entrypoints-check.log`;修改文件格式及 diff 检查通过。最终 `cargo build --locked -p aether-gateway --bin aether-gateway` 通过,网关可执行文件已更新,日志:`/tmp/aether-round6-gateway-build.log`。全部测试与构建进程已退出;尚未部署,未重跑业务阶梯吞吐压测,未执行全工作区测试。 + +本轮边界: + +- 预算覆盖上述运行时链路持有的四份诊断 JSON 正文及副本,不是进程 RSS 上限。原始 Redis RESP、批次字段字符串及反序列化临时分配仍不受此额度约束;直接通过契约自行构造或反序列化的 DTO 默认不启用运行时预算。 +- seed 进入本链路前的 JSON/base64 解析、构建过程的临时正文复制、序列化/压缩结果和 SQL bind 缓冲未纳入。预算不足时的纯构建会使用调用线程 CPU;它保留计费兼容性,没有消除解析和估算开销。 +- 协议观察器及用量恢复缓冲、超大历史窗口增量聚合、完整请求优雅排空和目标环境长期压测仍待后续处理。第三轮吞吐数据不能作为本轮构建或生产容量结论。 + +## 第七轮修复状态 + +- **Redis 回复转移正文缓冲:** `XREADGROUP` 采用当前 redis crate 支持 owned conversion 的底层容器结构,避开 `StreamReadReply` 的借用转换;字段正文直接由 RESP `BulkString` 转成 `String`。`XAUTOCLAIM` 的消息正文、游标和删除 ID 同样使用所有权转移。保留 RESP2/RESP3、nil、重复字段覆盖、无效 UTF-8 和错误分类等既有解析语义,没有改变队列协议。 +- **减少队列处理副本:** Memory 队列读取直接遍历待交付条目,去掉全局队列锁内的整批临时正文克隆;队列、PEL 和调用方继续独立持有自己的数据。usage worker 在处理完一条消息后先释放原始字段,再等待 ACK/DELETE,避免已处理的大正文跨确认 I/O 继续存活。失败和死信路径仍保留所需原文到处理结束。 +- **流式用量恢复收敛保留字段:** Claude 的跨记录状态只累计 mapper 使用的 token 和缓存字段,以及非空用量出现标记,未知大字段不再随流长度累积,也不再每次复制完整累计对象。候选和 chunks 数组省略没有用量或 tier 信息的空元素,保留 Gemini 的首候选位置和倒序查找规则。完整且未超限的 SSE 记录直接借用当前输入切片,跨分片与超限记录继续沿用原 carry 和恢复规则。 +- **投影兼容修复:** 显式 `usage: null` / `usageMetadata: null` 保留字段存在性,避免错误回退到旧快照;独立 null 不新增清零信号,带其他非空用量的嵌套结构遵循原 mapper 的优先级。图片回复仅提取 `data/result` 数量以恢复 `request_count/image_count`,不复制图片正文。 + +本地验证: + +- runtime state 76 项、usage runtime 282 项回归通过,合计 358 项。Redis parser 的 7 项新增测试用正文原指针和容量断言验证缓冲转移,另有 3 项 Memory 所有权和队列生命周期回归。日志:`/tmp/aether-round7-queue-tests.log`。 +- 新增真实 Redis 大消息回归,两种协议各 24 条消息、每条约 512 KiB 正文,3 个消费者同时分批读取,再以最多 5 条重领和 ACK/DELETE;48 条消息原文逐字节一致,无重复交付到不同读取结果,最终 pending、lag 和 stream length 都为 0。日志:`/tmp/aether-round7-large-redis-test.log`。测试自建 Redis 已关闭;此项已包含在前述 76 项中,不重复计数。 +- 网关流式、stream pump 和空闲读取期限 186 项回归全部通过,日志:`/tmp/aether-round7-stream-final-tests.log`。包含本轮新增的 9 项投影与借用测试,覆盖原始 Claude 累计对象对照、256 条带不同未知大字段的长流、10000 个空数组元素、正文复制计数、CR/LF/CRLF 分片及上限边界、显式 null 的实际 mapper 对照,以及仅有图片数量的回复。连同队列侧共 544 项通过,本轮新增 20 项;首批 parser、真实 Redis 和流式聚焦测试不重复计数。 +- 修改文件格式及 diff 检查通过,最终 `cargo build --locked -p aether-gateway --bin aether-gateway` 通过,网关可执行文件已更新,日志:`/tmp/aether-round7-gateway-build.log`。所有测试、构建进程和临时 Redis 均已退出;尚未部署,未执行新的业务吞吐压测或全工作区测试。 + +本轮边界: + +- 原始 Redis 消费批次仍按配置的条数读取,不是字节预算。默认每批最多 128 条、最多 32 个 worker;本轮减少重复分配,没有限制任意大消息或整体批次的最大驻留字节,也没有通过丢弃账务消息缩小批次。当前同 worker 的 read/reclaim 已串行,已处理条目原本就逐条释放。 +- 流式恢复仍保留必要的跨分片单记录,Basic 5 MiB / Full 64 MiB 的原限制未改变;多行 `data:` 拼接、JSON 解码临时值、非空用量数组和主协议观察器不在全局正文预算内。不能因减少诊断副本而直接停用计费恢复。 +- 主协议解析器的累计文本、原始队列批次字节控制、超大窗口增量聚合、完整请求优雅排空及目标环境长期压测仍待后续处理。 + +## 第八轮修复状态 + +- **主协议用量观察器取消正文累计:** `StreamingStandardTerminalObserver` 为 OpenAI Chat、Responses(包括 compact)及 Gemini 选用现有 provider parser 的终态观察模式。协议转换仍使用完整模式;观察器沿用同一事件分类和用量解析,只保留摘要需要的状态。Claude 当前没有正文累计,独立图片观察器也继续使用原实现。 +- **长流正文和工具参数:** Responses 不再保存文本、双份推理文本、工具参数、工具结果和图片项正文;保留工具索引、名称及 namespace 校验,它们影响未知事件计数和工具调用结束原因。Chat 不再等待迟到的工具名称或 ID 而持续缓存参数。Gemini 不再保存累计文本、推理、签名、媒体、工具参数和结果,只记录是否见过工具调用以保留结束原因。 +- **opaque 项去重:** Responses 观察模式将未知扩展项的去重键改为固定 32 字节 SHA-256 摘要,按原始键的相同字节增量计算,避免加密内容或整个序列化项成为常驻键。默认转换和客户端 emitter 的原键及完整输出保持不变。 + +本地验证: + +- `aether-ai-formats` 全部 922 项测试通过,包含本轮新增 16 项。逐个输入前缀及提前 EOF 对比完整解析器与观察模式的摘要,覆盖身份时点、用量、显式零、tier、错误、未知项去重、namespace、工具调用和 SSE / WebSocket 结构化入口。原有格式转换、会话历史、图片及同步转流回归均通过。日志:`/tmp/aether-round8-formats-final-tests.log`。 +- 长流回归在 2048 轮 1 KiB 文本、推理及工具参数输入后,直接断言 Responses 的正文缓冲容量仍为 0、Chat 工具缓冲为空;Gemini 在 32 轮多种 16 KiB 字段输入后所有正文状态 map 仍为空,并与完整模式核对终态帧。大 completed 项与 opaque 键的摘要字节、去重语义另有测试。此处验证内容保留行为,没有测量生产 RSS 或吞吐上限。 +- 网关流式、stream pump 和读取期限 186 项通过;Responses WebSocket 会话、上游和观察器 209 项通过。连同格式 crate 共 1317 项通过。日志:`/tmp/aether-round8-gateway-stream-tests.log`、`/tmp/aether-round8-gateway-websocket-final-tests.log`。WebSocket 回归直接运行同一份已编译测试程序;其间因现有 build script 监听 worktree 中不存在的 `.git/HEAD` 而触发的一次重复 Cargo 编译已主动停止,没有把中止当成测试通过。 +- 修改文件格式和 diff 检查通过,最终 `cargo build --locked -p aether-gateway --bin aether-gateway` 通过,日志:`/tmp/aether-round8-gateway-build.log`。本机内存压力较高,网关测试目标编译耗时 11 分 21 秒、最终构建 8 分 54 秒;这些是本地编译耗时,不是业务延迟。全部测试、构建进程已退出,未遗留本轮临时服务。尚未部署,未执行新的业务吞吐压测或全工作区测试。 + +本轮边界与后续: + +- Responses 的工具身份和 opaque 摘要集合仍按不同逻辑项数增长;单条 JSON 解码、未知事件错误载荷及部分 Gemini 分类 helper 仍有临时分配。完整格式转换和会话历史依赖的正文仍保留,本轮不宣称整个解析器或进程具有硬性 RSS 上限。 +- Redis 原始批次仍没有硬性字节上限。现有 redis 连接的超时或 future 取消不会保证后台立即停止接收已发命令的回复,仅读取后套预算或减少 COUNT 不能解决任意大消息。建议下一步先为新生产的完整队列 envelope 实施精确序列化字节上限,超限诊断降级须保留计费事实并沿现有失败路径重试;旧消息和 PEL 仍需兼容排空。真正的接收预算还需要读取前预留及受控连接/解码器,不能直接依赖当前 usage body blob 表作为入队旁路,该表依赖已存在的 usage 父记录。 +- 超大窗口增量聚合、完整请求优雅排空及目标环境长期压测继续待处理。尚未部署。 + +## 第九轮修复状态 + +- **新增队列消息的完整字节上限:** `AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES` 默认 1 MiB,启用 usage runtime 时显式 `0` 非法。`UsageQueue::enqueue` 在发送 Redis 命令前,以有界 writer 编码完整 v1 JSON envelope,包含 UTF-8、转义、metadata、正文、headers 及其他字段;不会先生成无限制的完整 JSON 字符串再检查长度。普通消息保持原格式,原公开编码接口及历史消息解码继续兼容。 +- **诊断降级保留计费语义:** 完整消息超限后,先借用检查去掉四份正文和四份 headers 的核心字段大小,核心可容纳才克隆 metadata 并生成诊断投影;复用同一字节缓冲。保留 token、费用、显式零、错误存在性、身份、时间、终态、正文引用、预留 token 和计费维度。按完整 v1 消费者规则保留请求档位、推理参数、响应实际档位及缓存 TTL,并标记被移除的正文为 `Truncated`。显式 `None`、`Disabled`、`Unavailable` 不改为可回退的状态;JSON null 与非对象正文分别按旧解码和权威规则处理。 +- **无法安全编码时的失败路径:** 核心仍超限,或去掉正文无法保留原缓存 TTL 计费语义时,返回 `InvalidInput`。终态沿现有有并发限制的数据库路径使用原事件回退;数据库不可用、受压或写入失败时明确返回 `Failed`,保留 first-byte 状态,不误报已入队或已缓冲。这类输入错误不打开 Redis 熔断,也不会无限重试。重试接收前再次校验,覆盖主路径已熔断或入队槽耗尽的旁路;重试 worker 也会终止单条永久失败并继续处理后续条目。 +- **指标与临时分配:** 导出队列 payload 上限、诊断降级、编码拒绝及永久重试失败计数。payload 计数是进程级编码尝试,包含入队与重试预校验,不是唯一事件数。超长 tier/reasoning 字符串先检查已有的 64 字节限制,再执行大小写规范化,避免明知非法仍复制整个字符串。 + +本地验证: + +- 用量模块全部 298 项通过,包含本轮新增的 9 项编码、6 项失败路径及 1 项配置回归。精确覆盖 UTF-8/转义字节边界、编码缓冲复用、字段透传、源事件不变、完整 v1 消费者对照,以及超限后的数据库回退、first-byte 保留、熔断/准入旁路拒绝和同重试分片继续排空。日志:`/tmp/aether-round9-usage-final-tests.log`。 +- 计费模块全部 87 项、数据契约全部 229 项通过。本轮新增 7 项真实队列读回计费对照与 3 项超长字段规范化回归;对照覆盖请求档位、缓存 TTL、显式零、未知价格、错误/取消、图片矩阵维度,以及 15 组非对象/null 正文组合。计费结果以原完整 v1 消息经旧解码路径后的行为为基线。连同用量模块共 614 项通过,不重复计入首批验证。日志:`/tmp/aether-round9-billing-contracts-final-tests.log`。 +- 网关指标、工作区模块边界、流式链路及 Responses WebSocket 共 438 项通过;网关入口配置全部 62 项通过,包括本轮新增的 payload 配置与指标回归。连同核心模块共 1114 项通过,本轮新增 28 项测试。直接复用本次构建生成的测试程序执行,日志:`/tmp/aether-round9-gateway-lib-tests.log`、`/tmp/aether-round9-gateway-main-tests.log`。 +- `cargo build --locked -j 1 -p aether-gateway --all-targets` 通过,最终网关可执行文件已更新;一次构建同时生成普通程序与测试目标,本地耗时 25 分 20 秒。日志:`/tmp/aether-round9-gateway-build.log`,产物清单:`/tmp/aether-round9-gateway-artifacts.jsonl`。修改文件格式及 diff 检查通过,所有本轮测试和构建进程均已退出;未创建临时服务。尚未部署,未重新执行业务吞吐压测或全工作区测试。 + +本轮边界: + +- 本次限制针对新生产的单条 JSON payload,不包含 RESP 外壳、完整读取批次、历史消息、PEL、DLQ 或进程总 RSS。原始 JSON 树、headers、metadata 的存活内存及解析临时值不因此获得统一字节预算;默认 128 条批次、32 个 worker 仍可能同时接收很多消息。真正的接收预算仍需要读取前预留和受控连接/解码器。 +- 对过大的原 metadata 采用保守拒绝,避免为判断最终能否缩小而先深拷贝整个对象;即使后续规范化理论上可使其变小,也使用原事件回退。模型分类 helper 的超长输入临时分配仍需后续审查,本轮不宣称所有计费提取都有硬性内存上限。 +- 没有新增持久化超限旁路。若消息无法编码且受限数据库回退也失败,用量落库会失败,并通过日志和指标暴露;不能把有限本地缓冲称为可靠落盘。超大窗口增量聚合、完整请求优雅排空及目标环境长期压测仍待处理。尚未部署。 + +## 第十轮修复状态 + +- **阻塞读取改为独占连接:** 原 blocking stream lane 通过轮询复用 `ConnectionManager`,快 worker 再次读取时可能命中其他 worker 正在执行 `BLOCK` 的连接,造成队头阻塞;取消调用 future 后,后台 driver 仍可能继续执行旧命令。现改为有固定容量和信号量的独占池,先取得空闲槽位再发命令,等待者取消不会影响现有读取。连接和未 spawn 的 driver 由同一查询持有,取消、超时及解析失败一起丢弃,下一次使用该槽位时重建;完整解析成功才回池复用。保留原 lane 数量、认证、数据库、RESP 配置及超时/延迟统计。 +- **积压重领继续扫描游标:** 原运行时丢弃 `XAUTOCLAIM` 返回的下一扫描位置,worker 每次从 `0-0` 开始。当大量近期活动的待确认消息占据前段,Redis 单次扫描可能返回空页,后段过期消息长期得不到重领。新增兼容的分页接口,保留下一位置和已删除消息 ID;worker 在成功响应后推进游标,空页和删除页也推进,读取错误保持原位置,扫描结束再从头开始。写入失败的消息不提前确认,仍留在 PEL,回绕或 worker 重启后可再次重领。原返回消息列表的公开接口和旧队列实现继续兼容。 +- **内存队列并发入队顺序:** 将序号分配移入队列插入的同一个锁范围。此前等待插入的生产者可能先取得较小 ID,其他生产者先插入并被消费后,较小 ID 会被读取游标永久跳过;现在 ID 顺序与实际插入顺序一致。内存后端同时实现重领分页,按数值序号推进。 + +本地验证: + +- `cargo check --locked -j 1 -p aether-runtime-state -p aether-usage-runtime` 通过。日志:`/tmp/aether-round10-core-check.log`。 +- 运行时状态全部 84 项、用量模块全部 300 项测试通过,共 384 项,包含本轮新增 10 项:连接池 2 项、内存队列 2 项、worker 游标 2 项、真实 Redis 4 项。覆盖连接容量、等待取消、多任务竞争、锁等待下的 ID 顺序、空页/删除页、读取失败重试、写入失败后回绕及原有记账流程。日志:`/tmp/aether-round10-queue-tests.log`。 +- 真实 Redis 测试另以 `--nocapture` 运行同一份测试程序,确认 6 次隔离实例就绪、没有跳过,4 项全部通过,不重复计入上述 384 项。使用本机 Redis、ACL 认证及数据库 7,覆盖 RESP2/RESP3、满池等待不发新命令、其他 lane 继续服务、快消费者复用空闲连接、取消/超时后服务端旧连接消失、重连保留认证/数据库,以及空扫描页后的 PEL 尾部和删除项恢复。日志:`/tmp/aether-round10-redis-receive-tests.log`。 +- 最终 `cargo build --locked -j 1 -p aether-gateway --bin aether-gateway` 在运行约 18 分 36 秒后因本机资源压力主动停止,退出码 143,不计为构建通过;停止前没有编译错误,日志:`/tmp/aether-round10-gateway-build.log`。期间清理了三份当前构建未使用的旧增量缓存,将可用磁盘从约 3 GiB 恢复至 7 GiB,但内存压力仍使编译持续缓慢。本轮没有修改网关入口、配置或指标,新增逻辑已在上述核心模块中验证;未重复编译网关测试目标,现有网关程序仍是上一轮产物。 +- 修改文件格式和 diff 检查通过。本轮测试、编译和临时 Redis 进程均已退出。尚未部署,未重新执行业务吞吐压测或全工作区测试;资源允许时仍需补跑最终网关构建。 + +本轮边界: + +- 本轮修复连接占用和积压恢复,不是完整读取批次的字节预算。历史超大消息、RESP 解码、默认 128 条批次和多个 worker 的总驻留内存仍需后续控制;上一轮新消息 1 MiB 上限保持有效。 +- 连接生命周期的取消控制目前仅用于阻塞 `XREADGROUP`。非阻塞 stream 命令及 `XAUTOCLAIM` 仍沿用原共享连接;关闭连接不能撤销 Redis 已执行的读取,已进入 PEL 的消息仍依赖重领。内存后端分页仍扫描并排序符合条件的待确认条目,本轮未建立该扫描的硬性内存上限。 +- 超大窗口增量聚合、完整请求优雅排空及目标环境长期压测继续待处理。尚未部署,测试结果不能用来推断生产 RPM 上限。 + +## 第十一轮修复状态 + +- **读取批次的共享预留:** 新增进程级 usage worker payload 预留,默认总量 128 MiB、单批目标 8 MiB。根据当前 `queue_payload_max_bytes` 推导实际 `COUNT`,默认由 128 条降到最多 8 条;读取与重领、所有 worker 和自动扩容后的新 worker 共用同一份额度。许可在发命令前取得,并保留到原始字段处理、记账及 ACK 完成;取消、失败和空响应释放,预算不足时在原任务内等待,不新增缓存任务。收到消息后按字段值实际长度缩减多余预留。 +- **兼容历史消息与配置:** `AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES`、`AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES` 分别控制总额与单批目标;0/非法值回退默认,极大值收敛到约 4 GiB 的有效额度,单批目标不超过总额。单条配置上限大于总额时明确返回配置错误,避免等待永远拿不到的许可。历史消息或其他生产者的大消息继续原记账流程,不在持有部分额度时等待追加,也不因超估算而删除或死信。 +- **重领连接与停止:** `XAUTOCLAIM` 改用上一轮的独占连接池,完整解析成功才回收连接。取消/超时会同时释放查询和 driver,避免归还读取预留后旧命令仍在后台接收。等待池容量也纳入原命令期限,指标归入实际使用的 `blocking_stream` lane;worker 取得重领响应前可被 shutdown 取消,取得响应后完成原处理与确认。读取与重领之间不嵌套持有池连接。 +- **扩缩容与观测:** worker 上报实际请求的 `COUNT`,防止读满 8 条却按 128 条误判为未满、压制扩容。新增 8 个 `usage_runtime_queue_read_*` 指标,涵盖总额、单批目标、当前预留、等待者、等待次数、累计字段字节及超估算条目/批次。累计字段字节包含字段名和值;预留及超估算判定使用值长度,合法 payload 恰好达到上限不会因字段名 `payload` 多出 7 字节而被误报。次数包含再次重领,不是唯一事件数。 +- **计费查询失败的恢复语义:** 原 worker 吞掉计费补全错误后继续写入/结算;暂时的数据库超时可能被后续“缺少实际费用”的永久错误覆盖,导致错误死信和 ACK。现在先传播原始错误,失败条目留在 PEL,后续重领再尝试;真正返回成功的无价格事件沿用原行为。三个直接落库入口也统一在补全失败时停止写入,不能误报已持久化或完成有序终态;已有受限数据库回退本来就正确停止,继续沿用。永久错误归档前先释放已解码事件,减少与原字段、死信编码的同时持有。 + +本地验证: + +- 运行时状态全部 87 项测试通过,包含本轮新增 3 项真实 Redis 重领回归。日志:`/tmp/aether-round11-state-tests.log`。 +- 用量模块最终全部 321 项通过,包含本轮新增 9 项预留、8 项 worker、4 项直接落库回归;连同运行时状态共 408 项,本轮新增 24 项。覆盖跨 Queue/worker 共享额度、16 个并发任务竞争、极大配置不溢出或永久等待、读/重领实际 COUNT、慢写/ACK 持有、错误和停止释放、历史超估算消息继续处理,以及价格查询超时不提前写库或确认、后续成功尝试使用准确费用。日志:`/tmp/aether-round11-usage-final-tests.log`。直接落库回归中的恢复是后续显式调用,不是新增自动重试。 +- 直接运行本轮已编译的真实 Redis 接收测试程序,7 项全部通过,确认 9 次隔离实例就绪、没有跳过;不重复计入上述测试数。RESP2/RESP3、ACL、数据库 7 下验证读/重领取消和超时、服务端旧连接释放、PEL/游标/删除项恢复,以及共池时等待期限、等待取消、释放单个槽位后继续读取和重领。日志:`/tmp/aether-round11-redis-receive-tests.log`。 +- 修改文件格式与 diff 检查通过;网关 `cargo check --locked -j 1 -p aether-gateway --bin aether-gateway` 通过,耗时 3 分 14 秒,日志:`/tmp/aether-round11-gateway-check.log`。本轮测试、检查和临时 Redis 进程均已退出。本轮没有重复执行上一轮因本机资源压力中止的完整代码生成与链接,最终网关程序仍需在资源允许时构建;未部署,未执行业务 RPM 压测或全工作区测试。 + +本轮边界: + +- 这是按当前生产配置估算的逻辑 payload 预留,不是网络接收字节或进程 RSS 的硬上限。滚动发布、其他实例使用更高上限、历史消息和直接写入额外字段都可超过估算。字段结构、字符串容量、Redis RESP 解码、连接缓冲高水位、解码 JSON 及死信序列化不由此获得硬上限;旧公开 Vec 读取接口保持兼容,不携带处理阶段许可。 +- 死信仍先追加再确认源消息,异常重试可能重复归档;历史大消息的死信 JSON 编码也没有独立硬字节预算。后续需要按源消息身份幂等的转移及受控超大消息恢复,不能以截掉原始账务字段或提前 ACK 代替。 +- 直接落库失败不会因此新增持久化重试渠道。超大窗口增量聚合、完整请求优雅排空及目标环境 RPM 压测仍待处理;尚未部署。 + +## 第十二轮修复状态 + +- **死信原子转移:** 内置 Redis 后端以一次 Lua 调用检查指定消费组的精确 PEL 身份,再追加完整死信、ACK 并删除源 ID;并发重领或提交成功但响应丢失后的重复调用不再次追加。Memory 后端在同一队列锁内完成转移,序号也在锁内分配。新增可选 trait 接口,默认返回 `None` 且无副作用;旧外部实现仍可使用原来的非原子追加后确认,内置转移报错不会降级为非原子写入。公开 `push_dead_letter` 保留原追加行为及 `{entry_id, fields, error}` JSON 格式。 +- **先检查,再写入:** 校验完整 canonical `u64-u64` ID、不同的源与目标键及非空字段;Redis 在首写前检查 PEL、目标类型和 XADD/XACK/XDEL 权限,避免可预见的脚本错误导致部分归档。Lua 不将 ID 转为浮点数,保留大整数精度。正文已被 trim 但 PEL 仍存在时,可使用调用者持有的完整字段归档。操作使用既有独占连接和 owned driver,等待池容量也计入超时,完整解析成功才回池;取消不能撤销已经完成的 Redis 命令,重试依靠 PEL 状态避免重复。 +- **独立的编码预留:** 新增 `AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES`,默认 64 MiB,及 `AETHER_USAGE_DLQ_ENCODING_MAX_JOBS`,默认 4。按原文字符串长度、JSON 最坏 6 倍转义及字段结构分隔符一次预留逻辑字节,使用 checked 运算和有界 writer;预留成功后才复制兼容接口输入或启动后台编码。worker 直接移动现有原始字段,永久记录错误时先释放解码事件。额度不足或单条过大立即失败,原消息保留在 PEL,不截断账务字段。后台任务取消等待后仍持有自己的许可,许可随编码结果保留到队列写入完成;与读取预留独立,避免持有部分读取额度再等待追加。 +- **worker 状态与指标:** 成功原子转移直接使用实际 ACK 数并报告一次死信,不再加入批次 ACK;返回 `NotPending` 不报告已归档,也不额外删除源消息。普通记录与旧后端追加仍批量确认,改为报告实际返回的 ACK 数。某条消息的存储转移失败时,只确认此前已成功处理的前缀,失败条目及尚未处理的后缀等待重领。编码预算拒绝及编码失败单独返回延期状态,保留坏消息并继续同批其他条目,只确认成功项,批次末尾仍报告失败;避免永远超过总预算的坏消息在每次重领时持续挡住正常账务。归档成功日志移到写入返回之后。新增 7 个 `usage_runtime_dlq_encoding_*` 指标:总额、最大/当前任务、预留字节、容量拒绝、超限拒绝及编码次数;次数包括重复尝试,编码成功不等于归档成功。 + +本地验证: + +- `cargo test --locked -j 1 -p aether-runtime-state` 全部 102 项通过;本轮新增 15 项,包括 8 项 Memory 转移、6 项真实 Redis 和 1 项返回解析验证。测试编译耗时 13.62 秒,执行 2.45 秒;日志:`/tmp/aether-round12-state-tests.log`。 +- `cargo test --locked -j 1 -p aether-usage-runtime` 全部 341 项通过;本轮新增 20 项,包括 8 项编码预算、3 项队列接线、8 项 worker、1 项配置验证。测试编译耗时 30.79 秒,执行 1.16 秒;日志:`/tmp/aether-round12-usage-tests.log`。连同状态模块共 443 项通过,本轮新增 35 项,无编译警告。覆盖完整原文及所有转义、精确上界和溢出、并发饱和、未 poll/后台运行/存储等待取消、panic 释放、原子失败不回退、ACK 实际计数、成功前缀确认、超大坏消息延期后正常后缀继续、后续调高预算恢复,以及旧后端和公开追加兼容。 +- 直接运行已编译的 `redis_dead_letter_transfer_ --nocapture`,6 项全部通过,确认 RESP2/RESP3、ACL 认证及数据库 5 下共 8 次隔离实例就绪,没有跳过;不重复计入上述 443 项。覆盖 16 个并发调用只追加一份、丢弃成功结果后重试、分别拒绝 XADD/XACK/XDEL 时无部分写且修复后可恢复、错误目标类型/消费组/ID/同键/空字段、正文 trim 后仍保留 PEL、超过 Lua 整数精度的 ID。日志:`/tmp/aether-round12-redis-transfer-tests.log`。取消/超时连接释放同时由状态模块中的前两轮 owned-driver 回归覆盖。 +- `cargo check --locked -j 1 -p aether-gateway --bin aether-gateway` 通过,耗时 2 分 23 秒,无警告;日志:`/tmp/aether-round12-gateway-check.log`。修改文件格式及 diff 检查通过;测试、检查和本轮临时 Redis 进程均已退出,现有服务未变更。本轮未重复执行前轮因资源压力中止的完整代码生成和链接,网关可执行产物仍需后续完整构建;未部署,未进行目标环境业务 RPM 压测或全工作区测试。 + +本轮边界: + +- 原子性和重复抑制按同一源 stream、消费组及源 ID 生效,不是永久去重索引。`NotPending` 只说明 PEL 不存在,外部 ACK、trim/delete、消费组销毁重建及其他消费者都可能改变这个状态,不能据此推断曾经归档。Redis XDEL 后其他消费组仍可能保留 PEL,Memory 沿用删除全部组 PEL 的既有语义;本轮不提供跨消费组的全局归档去重。 +- Redis 转移要求 7+ 的 `redis.acl_check_cmd` 和 `EVAL/TYPE/XPENDING/XADD/XACK/XDEL` 权限。Lua 错误本身没有回滚能力;脚本预检当前已知可失败的前置条件,不替代 Redis 持久化及故障恢复保证。Cluster 两键必须同 slot,当前默认 stream 键没有自动改名或迁移;缺能力、权限或跨 slot 失败会保留源消息,不能自动降级。 +- 编码预算覆盖任务所持原文长度与 JSON 上界,包含存储等待阶段,仍不是进程 RSS 或网络缓冲硬上限。字段容器、字符串多余容量、内存后端持久副本、Redis Cmd/packed-command/连接缓冲副本不计入;公开追加及旧后端可能沿用共享 driver。0/非法环境值回退默认,bytes 最大约 4 GiB,jobs 最大 128。保守 6 倍估算会拒绝实际编码较小的超大原文;这类存量需调高额度后重试或受控恢复,本轮未新增自动大消息通道。 +- 已提交但响应丢失时可以避免重复归档,当前进程的成功计数可能少计;普通 ACK 成功后 XDEL 失败也可能少计 ACK,指标不是持久化账本。死信总量及保留期限仍需运营管理,本轮没有静默修剪死信。完整请求优雅排空、超大窗口增量聚合与目标环境业务 RPM 压测仍待处理;尚未部署。 + +## 第十三轮修复状态 + +- **TCP 层先限量:** 原入口每次 accept 后都创建独立连接任务,HTTP 请求限流尚未生效的握手及空闲 keep-alive 不受请求 gate 保护。新增 frontdoor `HttpConnectionBudget`,二进制入口的全部监听分片和指标使用同一个 Arc;accept 后立即尝试取得许可,满额直接关闭刚接入的 socket 并让出执行,不创建 HTTP 任务或额度等待队列。先 accept 再准入,避免无流量的 reuseport 分片预占许可、饿住有流量的分片。 +- **许可跟随 socket:** 将许可放入底层 AsyncRead/AsyncWrite 包装,完整透传读、写、flush、shutdown 及 vectored write。socket 先释放、许可随后归还;HTTP/1 upgrade 会携带整个 IO 包装,所以连接 future 返回后 WebSocket 仍占一个许可,直到升级连接释放。HTTP/2 同一 TCP 内多路流共用一个许可,原 HTTP 请求和 WebSocket 会话 gate 保持独立;取消、解析失败及首个请求头超时也会释放底层 IO。 +- **accept 错误恢复:** 原 `listener.accept().await?` 会退出任意出错的监听任务,主入口随后停止其他监听分片。改用共同的 accept helper:ConnectionRefused/Aborted/Reset 重试,其他错误记录后退避一秒再试,资源耗尽不会触发无等待的重复 accept。兼容 `serve_tcp` 入口也使用相同预算和恢复逻辑。 +- **容量与观测:** 新增 `AETHER_GATEWAY_MAX_HTTP_CONNECTIONS` / `--max-http-connections`。未配置或 0 时按请求上限与 WebSocket 上限之和自动推导;自动和显式配置都限制在 1 至 65536,已知 FD soft limit 时再限制为 `max(1, (FD - 256) / 2)`。启动日志记录实际值,新增 `gateway_http_connections_limit/in_flight/high_watermark/rejected_total/accept_errors_total` 五项指标。 + +本地验证: + +- `cargo test --locked -j 1 -p aether-gateway-frontdoor` 全部 43 项通过,无警告,包含本轮新增 9 项连接回归。覆盖默认/显式/0/FD 极值、不溢出的容量推导、IO 先于许可释放、拒绝立即消费 socket、读写及 vectored write/半关闭透传、两个真实 listener 共享 limit=1 后恢复、首次 poll 前取消、真实 HTTP/1 首头超时和解析失败、101 upgrade 后 HTTP future 已结束但升级 echo 仍持许可、真实 HTTP/2 两条并行请求共用一个连接许可,以及注入 accept 错误后继续服务、资源错误每次精确退避一秒。日志:`/tmp/aether-round13-frontdoor-tests.log`。暂停时钟的退避测试未真实耗尽进程 FD。 +- `cargo check --locked -j 1 -p aether-gateway --all-targets` 最终通过,耗时 3 分 42 秒,无警告,覆盖主程序、库及测试/示例等目标;日志:`/tmp/aether-round13-gateway-final-check.log`。本轮另新增 2 项 AppState 共享预算和指标接线测试,已随全部目标完成编译检查,但没有重新生成及执行大型网关测试程序,不计入上述 43 项执行通过数。初次检查发现新增测试误用了不存在的 `for_tests` 方法,已按现有测试模式修正为 `AppState::new()`;初次退出码 101,不计为通过。 +- 修改文件格式及 diff 检查通过。frontdoor 测试新增三个已锁定版本的开发依赖引用,Cargo.lock 仅增加对应依赖名,没有升级依赖版本。本轮测试、检查及临时 TCP 连接/任务均已退出,现有服务未变更;没有重复进行网关完整代码生成和链接,也没有部署或执行业务 RPM 压测。 + +本轮边界与后续发现: + +- 这是已准入的入站 socket 数量上限,不是 HTTP/2 请求任务、全部 FD 或进程 RSS 上限。每个监听分片在 accept 与同步准入之间可短暂持有一个待判定 socket,kernel backlog、上游、Redis、数据库及其他入口不计入。容量规划仍需实测;健康检查使用同一入口,也可能在连接满额时被拒绝。 +- 满额发生在 HTTP 解析前,客户端看到连接关闭而不是 429/503;客户端和反向代理的重试应有退避。成功请求结束后 keep-alive 连接继续占额度,直到连接释放;本轮没有新增空闲连接回收、强制流期限或完整请求优雅排空。现有超时及流式传输语义保持原样。 +- 公共兼容 `serve_tcp` 每次调用使用独立预算,默认 4096,可用同名环境变量覆盖,最大 65536;它没有二进制入口的请求/WS 容量配置和 FD 探测,也未将预算绑定到其默认 AppState 的五项指标。二进制入口拥有 FD 约束、跨分片共享及指标接线。 +- 另确认 usage 窗口 Lua 每次准入全量清理过期成员,集中到期可能阻塞 Redis。不能直接换成分批删除加当前窗口 ZCOUNT:合法并发请求的时间参数可能乱序到达,后来的较早 cutoff 会重新计入此前逻辑上已过期但未物理删除的成员。保持精确额度需配套单调清理水位等状态设计;本轮没有修改这段逻辑。相同终态的事件级及 stored-record 级成本预留协调可能重复取得 PostgreSQL 用户行锁,后续可研究携带匹配身份的首次结果,不能无条件删除任一调用。 +- usage 独立 runtime 的本地重试渠道仍缺少完整停止与排空生命周期;当前数量有界,但进程退出前未入 Redis 的缓冲终态仍需统一处理。本轮没有修改停机流程、迁移架构或部署,业务 RPM 与长期排空压测仍待执行。 + +## 第十四轮修复状态 + +- **消除同次写入的重复成本协调:** worker 和直接落库原来先按事件协调成本预留,upsert 返回后又按存储记录协调一次,正常终态重复获取 PostgreSQL 用户行及预留行锁。现在首次调用返回的持久化结果同时匹配 request ID、用户、服务端预留 token、实际费用单位、终态和完成秒数时,携带一个仅限本次写入的内部结果;存储记录再次匹配全部字段才省去第二次协调。钱包结算、原用户级锁、费用检查及幂等性继续执行。 +- **保留冲突和重试语义:** 首次返回 `None`、返回其他终态或身份、upsert 后记录变化,都继续按存储记录重新协调。首次协调出错仍在 upsert 前停止;upsert 失败后的新尝试重新协调,不跨请求或重试缓存结果。公开协调与结算 API 签名不变;没有更改数据库表或放宽账务条件。 +- **修正验证工具:** PostgreSQL 结算基线补齐有余额的钱包及所属用户,并核对结算记录、快照、outbox、钱包消费和供应商累计值,检查失败返回非零退出码。此前无钱包的全量拒绝不能算吞吐成功。网关夹具补齐启动流程会创建的默认路由,通过 testkit 对象读取受保护的指标,缺失必要指标直接报错;容量报告保留 HTTP 状态和错误样本。探针新增排空基线与最终必要指标字段,判定阈值不变。网关阶梯夹具使用与主程序相同的 8 MiB Tokio worker 栈,避免 debug 路由调用链在默认小栈上溢出。 +- **隧道夹具按当前入口运行:** 仅在 `testkit` 功能下增加适配入口,将测试请求交给现有网关签名验证、正文完整性校验和临时 spool 流程,再进入内部 relay;没有开放跳过鉴权的 HTTP 路由。每条压力请求使用独立签名和 nonce。每个档位先检查合法签名成功并返回完整正文,无签名、篡改正文及重放均被拒绝;模拟节点通过并发计时等待响应,避免旧夹具在单个 WebSocket 读循环逐条 sleep 造成串行瓶颈。每个档位结束后显式取消并等待模拟节点任务。 + +本地验证: + +- `cargo test --locked -j1 -p aether-usage-runtime`:349 项全部通过,本轮新增 8 项,包含多字段不匹配矩阵、失败/取消/零费用、失败重试、256 个同用户并发写入和 32 次重复投递。重复投递使用真实 Memory 结算仓库,余额只扣一次。日志:`/tmp/aether-round14-usage-tests.log`。`gateway_pressure_probe` 既有 21 项测试全部通过,日志:`/tmp/aether-round14-probe-tests.log`。 +- 本轮已成功完整构建并链接网关、容量基线、结算基线和压力探针,补上第十至十三轮只做编译检查而未更新可执行文件的验证。主要日志:`/tmp/aether-round14-full-build.log`、`/tmp/aether-round14-final-harness-build.log`;无编译警告。 +- 主网关真实 TCP 冒烟:2 个监听分片共享连接上限 4,保持 4 条空闲连接后,另外 32 条连接在共 7 ms 内被关闭;释放后健康检查恢复 200,认证指标确认 high-watermark=4、rejected=32。网关正常退出,日志:`/tmp/aether-round14-connection-smoke.log`。 +- 隔离 PostgreSQL 14 结算热点:同一用户、同一钱包,100 并发完成 2000 笔真实结算,0 失败;2000 条 usage、结算快照及已处理 outbox 完整,待处理为 0,钱包消费与供应商费用均为 2.00 USD。耗时 2332 ms,857 次/秒,P95=251 ms;采样锁等待最高 62 个、最长 549 ms,同一钱包仍需串行扣费。报告:`/tmp/aether-round14-postgres-settlement-final.json`。 +- 网关同步/流式、执行运行时同步/流式和隧道共 5 个阶梯场景,在 8/32/128/256 的 gate 上限各执行 8 倍请求,共 16928 条全部 HTTP 200,无拒绝或读取失败。前四场景各 3392 条,隧道共 3360 条,隧道并发为上限减 1,因为 WebSocket 会话占一个许可。最终前四场景在途归零,隧道采样时只剩该会话,high-watermark 均达到对应上限。每个档位的签名、正文和重放预检通过。报告:`/tmp/aether-round14-capacity-final.json`;最终构建日志:`/tmp/aether-round14-capacity-auth-final-build.log`。这些场景是 Memory 仓库/模拟节点功能和容量验证,与下面真实数据库业务压测分开解读。 +- 主网关 + 隔离 PostgreSQL/Redis + 本地模拟上游,8 个 API key、4 个网关 Tokio worker、2 个监听分片、HTTP 请求上限 256、TCP 上限 512、数据库池上限 24,执行以下约 2 秒的流式请求,要求完整读取且出现 SSE `[DONE]`: + +| 并发 | 请求数 | 成功 / 失败 | 实测 RPS | 首字节 P95 | 完整响应 P95 | 排空判定耗时 | +| --- | --- | --- | --- | --- | --- | --- | +| 16 | 128 | 128 / 0 | 8.01 | 123 ms | 2062 ms | 9.93 s | +| 64 | 512 | 512 / 0 | 32.09 | 69 ms | 2046 ms | 9.86 s | +| 128 | 1024 | 1024 / 0 | 57.34 | 106 ms | 2074 ms | 119.21 s | + +业务压测报告目录:`/tmp/aether-round14-pressure.UYCeSq`。全部 1664 个请求为 HTTP 200;SQL 验证全部 completed,input/output/total tokens 分别为 1664/33280/34944,没有缺失首字节时间或单条 token 不匹配。所有档位必要排空指标齐全并通过连续安静期检查,最终用量队列、PEL、DLQ、outbox、请求在途和数据库锁等待均为 0。128 并发档网关 RSS 峰值约 190 MiB、末值约 147 MiB,FD 从峰值 338 降至 62,Tokio 活跃任务降至 85;没有要求 RSS 或常驻后台任务归零。 + +本轮边界: + +- 本地 debug 构建、模拟上游和短时压测不能推导生产 RPM 上限,也没有同环境修复前后的性能对照。模拟上游费用为 0,非零钱包金额另由上述 PostgreSQL 热点及 Memory 幂等测试验证。临时 PostgreSQL 夹具关闭 fsync、同步提交及 full-page writes,热点吞吐数不代表开启生产持久化后的性能。 +- 阶梯夹具使用 75 ms 模拟同步/隧道响应或 3 次 25 ms 流式间隔。256 并发时网关同步/流式 P95 分别为 770/808 ms,超过该夹具的 300 ms 延迟预算,吞吐从 128 并发的 596/754 RPS 降至 537/524 RPS;全部请求成功不等于此档位容量合格。执行运行时 P95 为 104/96 ms,隧道为 115 ms。需在目标环境定位网关层 CPU、调度及后台写入开销后决定并发上限,不应直接按本地峰值放大配置。 +- 128 并发后的完整排空约需两分钟;此前 45 秒检查及一次 120 秒检查未通过,不能作为快速回收的证据。当前复测确认最终回收,但仍需目标环境下长期压力、连接空闲回收和停机排空验证,不能归因为已确认的泄漏或宣称完全消除积压风险。 +- 第十三轮发现的 Redis 过期窗口集中清理和 usage 本地缓冲优雅停止仍待后续处理。尚未部署、未调整生产配置,未执行全工作区测试。 + +本轮构建、测试、压测及临时服务均已退出,原有开发服务未改动;格式及 diff 检查通过。验证过程中出现的缺失默认路由、旧指标入口、未签名 relay、夹具小栈溢出和流 ID 类型编译错误均已修正;此前失败或全量拒绝的报告不计入上述通过结果。 + +## 第十五轮修复状态 + +- **整窗过期时异步释放:** Redis 用量准入脚本原来逐条同步删除整个过期 sorted set,集中到期的大集合会长时间占用命令线程。现在对超过 256 条的集合检查首尾时间:首条未过期时跳过清理,末条也已过期时在同一次 Lua 调用内 `UNLINK` 整个键,让 Redis 在后台释放旧对象,随后仍按原规则准入并写入新事件。小集合及混合新旧记录的窗口继续精确删除,空集合直接使用计数 0。 +- **保留时间和事务语义:** 所有规则先清理,再检查,最后全部消费;任一规则拒绝时仍不写入新事件。后续规则的过期清理不会被前面规则的拒绝跳过。没有新增清理水位、辅助键或改变事件时间戳;乱序时间、重放、释放补偿、`Retry-After` 和键 TTL 沿用旧语义,无需迁移 Redis 数据。`UNLINK` 不可用或被 ACL 拒绝时退回原同步删除,不把未清理的旧记录当作已清理。 +- **减去请求内重复工作:** 清理后计数只在当前 Lua 调用内复用,不再重新 `ZCARD`;Rust 端用 `OnceLock` 复用不可变脚本对象,避免每个请求复制脚本文本和计算 SHA1。每次调用仍单独构造键和参数;Redis 脚本缓存丢失时仍由现有客户端处理 `NOSCRIPT` 并重新加载。 + +本地验证: + +- `cargo test --locked -j1 -p aether-runtime-state -p aether-usage-runtime`:状态模块 110 项、用量模块 349 项全部通过,共 459 项,0 失败;1 项较大数据计时基线默认忽略,已另行显式执行通过。日志:`/tmp/aether-round15-regression-tests.log`。 +- 新增 8 项真实 Redis 回归及 1 项计时基线。单独执行新增回归时确认 11 次隔离实例就绪,无 fixture 跳过,覆盖 Redis 8.0.3、RESP2/RESP3、ACL 认证及数据库 6。验证 10 万条记录一次 detach 后键复用、截止边界及 Lua 精确整数上界、混合窗口保留有效项、首规则拒绝时后续窗口仍清理、拒绝 UNLINK 后兼容、64 个并发争抢只放行 8 个、`SCRIPT FLUSH` 后幂等重载,以及 2 个协议各 512 步新旧脚本差分对照。对照逐步核对判定及完整剩余记录,包含时间回退、重复事件、额度变化和释放补偿。日志:`/tmp/aether-round15-redis-verified.log`。 +- 隔离 Redis 的 3 次对照各装入 30 万条同时过期记录。旧准入调用耗时 51.968 / 49.395 / 51.431 ms,新路径为 0.576 / 0.222 / 0.320 ms;另外装入 4096 条有效记录,各执行 1000 次持续准入,旧路径 P50/P95=86/123 us,新路径为 74/111 us,最终记录一致。日志:`/tmp/aether-round15-cleanup-timing-final.log`。测试验证逻辑和数据一致性,不以容易受本机负载影响的耗时阈值作为断言。 +- `cargo build --locked -j1 -p aether-gateway --bin aether-gateway` 完整构建与链接通过,耗时 4 分 15 秒,无警告;日志:`/tmp/aether-round15-gateway-build.log`。全工作区格式检查、新增 include 测试文件的独立格式检查及 diff 检查通过;本轮构建、测试和临时 Redis 实例均已退出,原有开发服务未改动。 + +本轮边界: + +- 这是整窗过期的优化,不是所有清理操作的硬耗时上限。混合窗口仍会同步删除其全部过期项;若大量旧记录与有效项同时存在,阻塞风险仍在。分批逻辑清理需要额外状态与完整的乱序、重试和滚动升级设计,不能直接以当前 cutoff 计数代替物理删除。 +- Redis 自动 TTL 过期的释放由服务端过期策略决定,可能先于本脚本发生,不由此改为异步。UNLINK 加速也依赖服务端支持及 ACL 许可;兼容回退仍有原同步释放成本。异步释放期间旧对象仍可能占内存,本轮没有提供 Redis 内存或 RSS 硬上限。 +- 数字来自本机 debug 构建、隔离 Redis 和单窗口局部对照,不代表网关端到端吞吐或生产 RPM 容量;未修改生产 Redis 配置、未部署,未重复上一轮业务压测。usage 停机排空及前轮观察到的高并发延迟仍待处理。 + +## 第十六轮修复状态 + +- **混合窗口的稀疏存活项重建:** 窗口超过 1024 条、有效项为 1–256 条且过期项超过有效项四倍时,先复制有效成员及原始 score 到临时键,再在同一次 Lua 内 `UNLINK` 旧集合并替换。先验证权限并完整构建临时集合,保留原 TTL;临时键已存在、Redis 缺少 ACL 检查接口或可选命令受限时继续精确同步清理。无清理水位或持久辅助数据,乱序、重放、多规则拒绝和释放语义不变。有效项超过 256 条的混合窗口仍走同步路径,不能据此宣称所有窗口清理已有硬耗时上限。 +- **请求与用量统一停机:** 收到退出信号后停止所有监听分片接收请求,HTTP/1 和 HTTP/2 等待已接收响应;默认 30 秒到期后取消剩余连接任务,并使升级后的连接读写返回关闭错误。对可能脱离 HTTP 生命周期的流式终态、取消回调及 WebSocket 审计,在创建后台任务前登记 producer。另将整条代理请求及响应体纳入登记,覆盖响应头前断开、内联响应体及断开后的后台消费,避免先关闭用量入口再提交终态的竞态。路由的 `cancel_on_client_disconnect=false` 策略保持继续完成上游请求,后台处理受后续用量排空期限约束;配置为 true 时按原规则取消并持久化取消终态。 +- **排空本地用量缓冲:** 等待 producer 结束后关闭生命周期入口,立即冲刷延迟事件并唤醒重试退避,等待各阶段与候选记录写入完成。已确认写入 Redis 的事件可留待下次消费;Memory 队列必须实际消费,存在未恢复的本地 DLQ 不判定为成功。worker 完成当前写入及 ACK 后停止,随后结束空闲分发任务和专用 Tokio runtime。超时明确报错;运行库层面保留未完成工作供再次调用停机,主进程不会把超时打印成排空成功。 +- **缩短上游空闲连接滞留:** reqwest 和浏览器传输的每客户端、每 origin 空闲连接默认上限由 1024 降至 32,默认空闲期限为 15 秒;H2C 池默认上限由 512 降至 32,并补上主动清理所需 timer。活动请求及响应流不受空闲期限限制。配置项为 `AETHER_GATEWAY_UPSTREAM_POOL_MAX_IDLE_PER_HOST` 和 `AETHER_GATEWAY_UPSTREAM_POOL_IDLE_TIMEOUT_MS`;HTTP、用量停止期限分别用 `AETHER_GATEWAY_HTTP_SHUTDOWN_TIMEOUT_MS`、`AETHER_GATEWAY_USAGE_SHUTDOWN_TIMEOUT_MS` 配置,均默认 30000 ms。 +- **保留恢复过程证据:** 压测探针记录有界的逐次排空观测,包含用量 producer、延迟事件、业务队列、Tokio 任务、连接和 FD。必要指标、任务基线容差以及连续 7 秒安静期保持原判定,新增本地待处理指标有值时也必须归零。 + +已完成的本地验证: + +- frontdoor 45 项、用量运行库 360 项、负载工具库 17 项、压力探针 21 项、状态模块 113 项、网关主程序 66 项、请求生命周期 8 项、传输模块 115 项及流式断开策略 1 项,共 746 项通过。状态模块另有 1 项默认忽略的计时基线,已显式执行通过。新增停机验证覆盖 producer 竞态、并发终态、重试失败后再次停止、延迟事件、多个 worker、写入途中停止、Memory 队列、响应头前断开、后台响应体、HTTP/1 和 HTTP/2 挂起处理器的实际释放,以及升级连接。日志:`/tmp/aether-round16-final-regressions.log`、`/tmp/aether-round16-state-regression.log`、`/tmp/aether-round16-main-final-tests.log`、`/tmp/aether-round16-request-lifecycle-tests.log`、`/tmp/aether-round16-transport-tests.log`。流式策略测试首次使用默认小栈时发生栈溢出;按网关入口相同的 8 MiB 栈设置 `RUST_MIN_STACK=8388608` 后通过,日志:`/tmp/aether-round16-disconnect-policy-test.log`。 +- Redis 新增混合重建测试覆盖 RESP2/RESP3、1/16/256/257 个有效项、TTL、非过期键、重复提交、已有临时键,以及 ZCOUNT/PTTL/EXISTS/UNLINK/RENAME/PEXPIRE 权限受限后的精确回退。定向 11 项测试确认 15 次隔离 Redis 实例就绪,无 fixture 跳过;原有 1024 步新旧脚本差分测试继续通过。日志:`/tmp/aether-round16-redis-tests.log`。 +- 同机 Redis 8.0.3 对照,30 万条过期记录分别混合 1/64/256 个有效项,旧脚本为 48.009/54.660/52.924 ms,新脚本为 0.320/0.377/0.496 ms,最终成员完全一致。整窗过期 3 次为旧脚本 56.802/51.804/50.886 ms、新脚本 0.744/0.258/0.318 ms。4096 条有效记录的持续调用 P50/P95 为旧脚本 87/131 us、新脚本 76/123 us;计时不作为断言阈值。日志:`/tmp/aether-round16-redis-timing.log`。 +- 网关主程序和压力探针完整构建、链接通过,无警告,日志:`/tmp/aether-round16-final-build.log`。补齐请求生命周期登记后再次完整构建主程序通过,日志:`/tmp/aether-round16-final-gateway-build.log`。 + +真实入口压力和停机验证: + +- 使用隔离 PostgreSQL 14、Redis 8.0.3 和约 2 秒的流式模拟上游,8 个 API key、4 个网关 worker、2 个监听分片、请求上限 384、TCP 上限 768、数据库池上限 24。要求完整读取且出现 SSE `[DONE]`。本轮 PostgreSQL 未关闭持久化,启动对照实查 `fsync`、`synchronous_commit`、`full_page_writes` 均为 on;Redis 夹具关闭 AOF/RDB,因此只验证保留 Redis 进程时的网关重启恢复。 +- 连接池单变量对照使用同一构建,仅设置旧空闲上限/期限 1024/90000 ms 或新默认 32/15000 ms。旧配置的 128、256 并发均全部 HTTP 200,业务待处理指标分别约 2.8/5.1 秒归零,但两档都未通过 150 秒恢复检查,最后 Tokio 活跃任务为 229/350、FD 为 206/327;脚本返回失败,不记为恢复通过。新配置分别在 22.217/26.542 秒通过原有恢复判定,FD 降至 65/63。这确认本地长时间恢复的主要滞留来自空闲连接及相关任务,而不是业务队列持续积压。对照目录分别为 `/tmp/aether-round16-pressure.3xsT7P`、`/tmp/aether-round16-pressure.iysPcj`;不把一次对照的吞吐或首字节波动视为可靠容量提升。 +- 最后补齐整个代理请求的 producer 登记后,使用最终构建再次执行新配置,结果如下。必要排空指标齐全,最终队列 lag、PEL、DLQ、outbox、本地用量待处理及数据库锁等待均归零,无 worker 处理失败;恢复判定包含连续 7 秒安静期。 + +| 并发 | 请求数 | 成功 / 失败 | RPS | 首字节 P95 | 完整响应 P95 | 恢复判定耗时 | 最终 FD / Tokio 任务 | +| --- | --- | --- | --- | --- | --- | --- | --- | +| 128 | 1024 | 1024 / 0 | 56 | 150 ms | 2141 ms | 22.325 s | 65 / 87 | +| 256 | 2048 | 2048 / 0 | 115 | 155 ms | 2192 ms | 26.243 s | 63 / 85 | + +最终报告目录:`/tmp/aether-round16-pressure.JI8vPF`,执行日志:`/tmp/aether-round16-pressure-final-new.log`,验证脚本:`/tmp/aether-round16-pressure.sh`。该轮 256 并发 FD 峰值 595、最终 63;RSS 峰值约 283 MiB,检查结束仍约 282 MiB,没有把短时间 RSS 不下降判为泄漏或声称内存已回到启动值。 + +- 正常 SIGTERM:32 条已开始输出的流全部完整结束,网关成功退出;重启同一数据库与 Redis 后,核对 3104 条 completed,逐条 input/output/total tokens 为 1/20/21,首字节时间无缺失。 +- 将 HTTP 停止期限缩短为 50 ms:默认继续完成策略下,32 个客户端都观察到流中断,网关继续完成其后台请求后成功退出;恢复后 3136 条全部 completed,token 合计 3136/62720/65856,逐条仍为 1/20/21。随后仅在测试数据库启用断开取消策略,另外 32 条流中断后全部落为 cancelled、HTTP 499、billing_status=void、token 为 0,首字节时间保留。最终数据库为 3136 条 completed 加 32 条 cancelled,未把默认继续完成误判为取消。 +- 最终脚本对 SQL 完整性失败显式退出。早期验证脚本误将默认断开策略预期为取消,并遇到旧 Bash 的条件失败未中止问题;已修正断言和预期,上述停机终态与 token 结论只采用最终目录的复测结果。 + +全工作区格式检查、新增 include 测试文件的独立格式检查及 diff 检查通过。本轮构建、测试、压测、临时 PostgreSQL/Redis 和模拟服务均已退出,原有开发服务未改动。 + +边界: + +- 新连接池设置只限制空闲缓存,不能推导整个进程的 FD 或 RSS 上限,也不能替代活动请求并发预算。减少空闲连接会增加突发流量之间重新建立连接的次数;需要在目标环境观察握手成本及连接复用。 +- 有界停机无法保证 SIGKILL、依赖持续故障或超过停止期限时的本地未持久数据不丢失;本轮未引入磁盘日志或更换消息架构。进程管理器的强杀期限应覆盖两阶段停止和额外收尾时间。 +- 第十四轮短响应夹具在 256 并发的延迟超预算尚不能判定已解决;本地 debug、模拟上游和短时压力结果不代表生产容量。未部署或修改生产配置,未执行全工作区测试。 + +## 第十七轮修复状态 + +本轮继续处理混合过期 Redis 大窗口和短响应网关的历史记录读取开销,未部署生产。 + +- **Redis 混合大窗口:** 当过期项超过 4096、有效项超过 256 时,先执行只读规划,随后在独占连接上 WATCH 全部规则,每条命令最多复制 512 个有效成员;通过 MULTI/EXEC 校验期间没有源数据变化,最后在一个 Lua 调用中替换窗口并完成原有多规则检查与消费。旧进程的 ZADD/ZREM、TTL 变化也会使 WATCH 失效。复制期间不提前清理其他规则,避免乱序请求、重复事件和跨规则拒绝改变原有额度语义。 +- **维护资源限制:** 每个 Redis runtime router 最多两条按需维护连接,最多尝试八次,包含连接池等待的总耗时受原有命令超时约束;未配置超时时使用 30 秒上限。取消或失败会丢弃仍有 WATCH/MULTI 状态的独占连接,未提交副本有 60 秒 TTL。成功提交后没有额外的可失败网络收尾操作。缺少所需 ACL 能力或 Redis 不支持 ACL 预检时保留原有精确清理路径。 +- **候选内存仓库索引:** 请求更新不再全表查找逻辑主键,按 request_id 查候选;最近记录查询改用创建时间索引,在复制之前取 limit。三个索引在同一写锁中维护,覆盖 ID 替换、相同时间排序和删除空索引。所有写入入口保留清洗,读取不再重复重建已经清洗的 JSON。 +- **公共候选诊断清洗:** 提取借用 JSON 对象的内部函数,移除持久化清洗前的整份输入克隆、清洗结果克隆及诊断对象的再次克隆;字段白名单、诊断大小限制、管理员与公开投影语义不变。PostgreSQL 与 Memory 都调用这一公共函数。 +- **调度读取精简投影:** 新增 `list_recent_runtime` 读取身份、状态、计数及时间,Memory 从时间索引直接构建不带诊断的记录,PostgreSQL 查询不读取诊断和能力 JSON。调度与自适应 RPM 观察使用此路径,管理员和请求可观测性继续读取完整记录。精简前后并发计数、RPM 和失败冷却判断保持一致。 + +已完成 241 项检查:状态模块 119 项、候选仓库与数据契约 28 项、调度核心 92 项、隔离 PostgreSQL 14 精简投影检查 1 项,以及单独执行的 Redis 大窗口计时检查 1 项。PostgreSQL 检查使用独立临时实例和连接级临时表,包含 32 KiB 诊断字段,验证精简投影与完整记录的运行时字段一致,且完整读取仍保留诊断。最终网关压测程序编译通过;临时数据库均已停止。日志:`/tmp/aether-round17-state-complete-tests.log`、`/tmp/aether-round17-projection-tests.log`、`/tmp/aether-round17-scheduler-tests.log`、`/tmp/aether-round17-postgres-live.log`。 + +首次短响应采样确认调度等待 `list_recent` 的全局读锁,读路径原本复制全部历史记录再逐条重做 JSON 清洗;仅修索引后最近 128 条完整 JSON 的复制与清洗仍是热点,因此继续将调度读取改为精简投影。最终采样仍有候选仓库读写锁等待,但没有原先全量历史复制的放大行为,其他开销分布在请求处理、凭据解密和 JSON 构建。本机其他进程负载较高,不能直接与第十四轮数字比较。 + +短响应对照使用本轮修改前保留的二进制与最终二进制,同为本机 debug 构建、Memory 网关夹具、模拟上游约 75 ms、每点请求量为并发数的八倍。最终版本及随后复跑的对照版各执行 15344 个请求,全部成功,无失败或拒绝。网关 P95 单位 ms: + +| 场景 | 并发 | 修改前复测 P95 | 修改后 P95 | 修改后 RPS | +| --- | --- | --- | --- | --- | +| 同步 | 128 | 1025 | 222 | 857 | +| 同步 | 256 | 2049 | 408 | 800 | +| 流式 | 128 | 569 | 263 | 678 | +| 流式 | 256 | 1564 | 451 | 794 | + +128 并发达到该夹具的 300 ms 预算;256 并发两种网关路径仍超过预算。最终版本中独立执行器和隧道曲线的 P95 均低于预算。结果位于 `/tmp/aether-round17-capacity-final.json`、`/tmp/aether-round17-capacity-control-recheck.json`。本机对照时系统 CPU 约 92%–99%,同一进程包含压测客户端、网关与模拟上游,不能将这里的 RPS 视为生产容量。用于 CPU 采样的额外长测有采样器干扰,未用于上述性能对照。 + +Redis 8.0.3 同机前后对照,均含 300000 条过期记录,单位 ms: + +| 有效项 | 旧整次调用 | 新整次调用 | 旧最长命令 | 新最长命令 | +| --- | --- | --- | --- | --- | +| 257 | 64.750 | 2.199 | 64.351 | 0.218 | +| 4096 | 77.099 | 7.331 | 76.301 | 0.681 | +| 32768 | 62.999 | 45.464 | 62.583 | 0.914 | +| 150000 | 77.200 | 242.083 | 76.840 | 1.585 | + +最终成员与旧脚本一致;最长命令来自隔离 Redis 的 SLOWLOG,包含 EVAL/EXEC。有效项多时整次清理更慢,但复制批次之间允许其他命令执行,降低对同一 Redis 其他请求的连续阻塞。原有整窗过期和不超过 256 个有效项的快速路径继续通过。4096 条有效项持续调用 1000 次,旧 P50/P95 为 155/286 us,新为 152/310 us;未将计时作为测试断言。记录:`/tmp/aether-round17-redis-timing.log`。 + +边界:复制需要临时保存有效成员,额外内存与有效项数量相关;频繁并发修改可能触发重试或达到命令期限,未提交的复制不会替换原始窗口。提交成功但回复丢失仍具有现有 Redis 命令的结果不确定性,应使用同一事件 ID 重试。ACL 不足或不支持预检的 Redis 仍走同步回退,不应把本优化描述为所有配置下的 Redis 总阻塞硬上限。代码层面的已复现放大路径已修复,但不能据此宣称所有容量问题均已解决;256 并发的最终验收仍需在代表性环境运行 release 构建并结合持续压测确认。未部署生产。 + +## 主要判断 + +系统已经有入口并发门、请求体内存预算、认证缓存合并、前后台数据库池隔离及多种有界队列。问题主要在于部分请求路径仍执行全量统计、锁范围过大,以及资源预算和超时没有覆盖完整生命周期。 + +这些问题会相互放大:每请求成本增加,使请求停留更久;在途请求增加后继续竞争 Redis、数据库和内存,客户端重试又增加负载。只提高入口并发数或数据库连接数不能消除这些瓶颈。 + +平均在途请求约为 `RPM / 60 × 平均请求持续秒数`。例如 600 RPM、平均持续 60 秒,约有 600 个在途请求;该例是容量计算,不是当前实例测量值。 + +## P1-1:压缩或未知长度请求一次占满全局请求体预算 + +**代码:** + +- `crates/aether-gateway/frontdoor/src/body.rs:392`:带非 identity Content-Encoding,或缺少 Content-Length 时,预留 `min(单请求上限, 全局预算)`。 +- `apps/aether-gateway/src/state/app.rs:40`:默认全局预算 256 MiB;`apps/aether-gateway/src/headers.rs:20` 的单请求默认上限同为 256 MiB。 +- `apps/aether-gateway/src/state/app.rs:195`、`crates/aether-gateway/frontdoor/src/body.rs:201`:默认没有请求体完整读取超时。 +- `crates/aether-gateway/frontdoor/src/body.rs:154`:等不到预算则返回过载;默认等待预算 250 ms。 + +**结果:** 即使压缩后的请求只有 1 KiB,也会预留全部 256 MiB。压缩上传或未知长度上传的读取阶段因此被串行化;一条一直未传完的请求可以长期占住预算,其他需要读取请求体的请求陆续收到 503。这里的 256 MiB 是预留额度,并非立即分配的实际内存。 + +**验证:** 使用实际 `BodyBufferPolicy` 的独立程序复现:1 KiB gzip 请求预留 268435456 字节,剩余许可为 0;第二个普通 1 KiB 请求等待约 252 ms 后被拒绝,释放第一个预留后立即恢复。程序显式采用默认参数,验证的是预算准入,不是完整 HTTP 慢上传压测。源码位于 `/tmp/aether-frontdoor-budget-repro.rs`,可执行程序位于 `/tmp/aether-frontdoor-budget-repro`。`cargo test -p aether-gateway-frontdoor body::tests -- --nocapture` 的 13 项既有请求体测试通过。 + +**建议:** 设置适合实际上传大小和速率的读取期限;根据业务明确单请求大小上限,使一个普通上传不能耗尽全局预算。进一步改为分段、有界读取及解压预算,或为大型上传分配独立额度。增量预留必须避免多个请求各持部分预算、同时等待扩容导致死锁,不能简单取消当前保护。 + +## P1-2:调度请求执行 Redis 全库扫描和历史样本聚合 + +**代码:** + +- `apps/aether-gateway/src/dispatch/pool_scheduler.rs:142`、`:985`:候选分页和 sticky 路径调用管理侧 runtime 统计函数。 +- `apps/aether-gateway/src/handlers/admin/provider/pool/runtime/reads.rs:126`:cache affinity 且 sticky TTL 非零时,扫描所有匹配会话,再 MGET 全部结果。 +- `crates/aether-runtime/state/src/redis/runtime.rs:289`:SCAN 循环直到游标归零,COUNT 200 是每轮提示,不是总量上限。 +- `apps/aether-gateway/src/dispatch/pool_scheduler.rs:1984`:扫描生成的会话总数和按 key 统计并不进入实际调度状态。 +- `apps/aether-gateway/src/handlers/admin/provider/pool/runtime/reads.rs:208`、`:225`:无条件对每个候选 key 拉取成本和延迟窗口的原始成员,即使相应排序或限额未启用。 + +**结果:** 一页 64 个不同候选 key 就产生 128 次窗口查询,另加会话扫描等操作。成本窗口默认 5 小时,历史样本按时间修剪;随着 RPM 增加,每个请求需要读取和聚合的历史数据也增加。`join_all` 并发等待不会消除 Redis 执行量及返回数据量。SCAN 和窗口读取都使用 Admin lane,管理请求也可能受到影响。 + +**建议:** 为调度建立独立读取接口,只读取当前 sticky 绑定及调度需要的状态;会话总数留给管理统计。按启用策略读取成本或延迟,改用增量计数、时间桶或后台维护的短时快照;保留严格额度检查的原子性。对相同池的刷新合并,避免每个请求重复拉取历史。 + +## P1-3:结算错误地锁住公共套餐行,跨用户串行 + +**代码:** `crates/aether-data/adapters/postgres/src/settlement.rs:482` 的 `user_plan_entitlements JOIN billing_plans ... FOR UPDATE` 没有限定锁的表。 + +**触发:** 普通用户的非零费用结算,且用户有有效套餐。即使套餐没有 daily_quota,查询也先锁行,之后才判断 `grants.is_empty`。 + +**结果:** 不同用户只要使用同一个套餐,就会争抢同一条 `billing_plans` 行。锁持续到事务结束,其间还可能汇总每日账本、写额度流水、更新钱包和结算快照。首先影响后台结算和队列消化速度,不能据此断言前台数据库池必然同时耗尽。 + +**验证:** 在隔离 PostgreSQL 14.17 中,两用户拥有不同 entitlement、共享一个 plan。事务 A 持原 SQL 锁时,事务 B 因 500 ms lock_timeout 失败,错误明确指向 `relation "billing_plans"`;对照使用 `FOR UPDATE OF user_plan_entitlements` 后事务 B 成功。临时 PostgreSQL 已停止。该验证证明锁冲突,不是完整结算吞吐压测;部署 Compose 使用 PostgreSQL 15。 + +复现脚本:`/tmp/aether-plan-lock-repro-20260909.sh`;本次日志:`/tmp/aether-plan-lock-repro.XoYxfO/`。 + +**建议:** 将锁范围限定为需要更新的用户 entitlement;明确套餐配置并发变更的一致性规则。添加真实双连接回归:同用户额度不能重复扣,不同用户同套餐不能互相阻塞。 + +## P1-4:首包之后的流缺少空闲期限 + +**代码:** + +- `apps/aether-gateway/src/execution_runtime/stream/execution.rs:2699`:默认 inline 直通路径仅在首包前使用 timeout,首包后直接 `upstream.next().await`。 +- `apps/aether-gateway/src/execution_runtime/transport.rs:3375`:stream 不使用总请求 timeout。 +- `apps/aether-gateway/src/execution_runtime/transport.rs:4212`:该 reqwest 客户端配置连接超时,没有设置读取超时。 + +**结果:** 上游已经发送首包、随后不再发送数据也不关闭连接时,只要下游继续保持连接,该流就可能长期占据请求名额、provider 并发守卫及缓冲。坏流逐渐积累会压缩可用容量。target permit 在首次向客户端 yield 时释放,不能把它算作整条流一直占用的资源。 + +**建议:** 增加可按 provider 配置的上游空闲期限;计时以真实上游活动为依据,网关自己的 keepalive 不应重置它。超时后关闭上游并走一次终态结算,确保断开、取消和超时都释放请求及 provider 守卫。对合法长思考模型采用匹配其行为的阈值。 + +## P1-5:并发上限与实际常驻内存不匹配 + +**代码:** + +- `apps/aether-gateway/src/main.rs:433`、`:467`:自动入口上限为每 CPU 1024,结合 FD 下调,但没有内存预算。 +- `apps/aether-gateway/src/execution_runtime/stream/execution.rs:167`、`:400`:Basic 模式每份流分析缓冲上限 5 MiB;Full 为 64 MiB。 +- 同文件 `:1993`、`:2031`:分别累积 provider 和 client 两份 body,包括直通路径。 + +**结果:** 两份捕获缓冲达到上限时,Basic 单流约 10 MiB,Full 约 128 MiB,尚未计入请求 JSON、转换状态、队列及其他内存。1000 条都达到 Basic 捕获上限的流,仅这两份内容就约 9.77 GiB。这是达到上限时的预算估算,不是普通小回复的固定内存或已测 RSS。 + +**建议:** 用实测每请求内存和 cgroup/物理内存确定入口容量;为响应捕获建立全局字节预算和截断策略。Basic 优先使用增量 usage/error 解析,直通时避免重复捕获同一内容。请求体读取预算不能充当整个流生命周期的内存保护。 + +## P1-6:审计压缩在 Tokio 线程和数据库事务内同步执行 + +**代码:** `crates/aether-data/adapters/postgres/src/usage/mod.rs:8457` 开始事务并锁 request;`:8522` 同步准备审计内容;`:12362` 序列化 JSON,`:12376` 执行 gzip level 6。`:12478` 附近最多处理四份请求/响应 body;inline 阈值为 0。 + +**结果:** 有 body 捕获,尤其大上下文、高并发时,CPU 压缩同时占据 Tokio 工作线程、数据库连接和事务锁。前后台连接池隔离不能隔离同一进程内的 CPU 和内存争用。 + +**建议:** 在开启事务前完成可独立准备的序列化和压缩,放到有并发和字节预算的 blocking worker;事务内只保留必须原子执行的读写。检查相同 body 的去重,避免单纯增加 worker 数量导致 CPU 和内存进一步饱和。 + +## P1-7:长期额度准入按用户串行扫描,且默认无锁等待期限 + +**代码:** `crates/aether-data/adapters/postgres/src/settlement.rs:276` 锁 `users` 行;`:611`、`:850` 在准入中取得该锁后,分别逐窗口 COUNT 请求预留或 SUM 成本预留。释放及对账还会争同一用户锁。`crates/aether-data/adapters/postgres/src/tx.rs:15`、`:116` 的默认读写事务不设置 lock_timeout 或 statement_timeout。 + +**触发:** 配置长期请求额度或成本上限;不能把普通短期 RPM 规则一概算入这条路径。 + +**结果:** 同一用户多个 key 的准入串行;窗口内历史越多,锁内工作越多。等待锁的事务还占用连接。池 acquire_timeout 只管取得连接之前的等待;Compose 的 idle_in_transaction_session_timeout 也不能终止正在执行的锁等待 SQL。 + +**建议:** 按用户和额度窗口维护原子聚合及预留,避免准入反复扫描明细。为前台和后台事务分别设定锁等待和 SQL 期限,失败时回滚并限制重试;精确额度和幂等结算约束必须保留。 + +## 次要放大器 + +- **P2,过载后重复计算:** `apps/aether-gateway/src/ai_serving/planner/state/scheduler.rs:100` 在 API key 并发受限时短间隔重做候选读取和排序。建议在昂贵规划前做准入,使用有界等待或通知。 +- **P2,探测前置任务未合并:** `apps/aether-gateway/src/maintenance/runtime/pool_quota_probe.rs:1675` 先 spawn,再去 Redis 去重及争锁。池恶化时仍随请求量创建任务。建议在 spawn 前按 provider 合并触发信号。 +- **P2,同步日志输出:** `crates/aether-runtime/base/src/tracing.rs:403` 直接写 stdout,`:695` 的文件写持全局 Mutex。磁盘或日志收集变慢时可能阻塞 Tokio worker。改为有界日志队列,并明确队列满时策略;当前没有证据证明这是本次故障主因。 + +## 修改和验收顺序 + +1. 先修公共套餐锁范围、读取期限及压缩上传预算问题;这些都有明确且局部的触发条件。 +2. 从请求调度移走管理扫描,按需读取运行态,避免每请求重算历史统计。 +3. 修流空闲期限,建立响应捕获全局预算,将压缩移出事务及异步工作线程。 +4. 优化长期额度聚合,并为数据库锁等待、SQL 执行和过载重试设置预算。 +5. 用独立测试实例及可控制延迟的 mock upstream 做阶梯压测;每档记录吞吐、P95/P99 首包及总耗时、RSS、CPU、队列积压,停止流量后检查资源能否回落。 + +生产定位至少需要:故障实例和版本、CPU/内存/FD 配额、实际 RPM 和平均流时长、是否使用压缩上传/账号池/套餐/完整 body 捕获。采集 `pool_runtime_state` 阶段延迟、Redis Admin lane 延迟、数据库 checked-out/lock waiting、usage 队列 lag 和 Tokio 任务数。CPU 低且锁等待高、Redis 延迟高、RSS 持续涨、请求体大量 503 分别对应不同路径,不能仅凭“卡死”选择一个原因。 + +现有 `crates/aether-testing/loadtools/src/bin/gateway_pressure_probe.rs` 可用于受控环境的压力和排空观测;不要用健康检查接口的吞吐代替真实 AI 调度及结算链路的容量。