summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--.env.control.example26
-rw-r--r--.github/workflows/ci.yml95
-rw-r--r--Dockerfile2
-rw-r--r--Makefile12
-rw-r--r--README.md28
-rw-r--r--cmd/aigw/main.go149
-rw-r--r--cmd/aigw/main_test.go14
-rw-r--r--cmd/migrate/main.go10
-rw-r--r--cmd/reconcile-billing/main.go79
-rw-r--r--cmd/rotate-credentials/main.go41
-rw-r--r--config.control.example.json23
-rw-r--r--docker-compose.yml20
-rw-r--r--docs/runbook.md82
-rw-r--r--internal/adminapi/api.go208
-rw-r--r--internal/adminui/assets/app.js23
-rw-r--r--internal/adminui/assets/index.html18
-rw-r--r--internal/billing/ledger.go41
-rw-r--r--internal/billing/operations.go881
-rw-r--r--internal/billing/operations_test.go36
-rw-r--r--internal/billing/service.go319
-rw-r--r--internal/billing/service_test.go174
-rw-r--r--internal/billing/stripe.go207
-rw-r--r--internal/billing/types.go162
-rw-r--r--internal/catalog/catalog.go38
-rw-r--r--internal/config/config.go289
-rw-r--r--internal/config/config_test.go8
-rw-r--r--internal/controlplane/mail_operations.go213
-rw-r--r--internal/controlplane/mail_operations_test.go29
-rw-r--r--internal/controlplane/manager.go44
-rw-r--r--internal/controlplane/mutations.go143
-rw-r--r--internal/controlplane/outbox.go19
-rw-r--r--internal/controlplane/queries.go73
-rw-r--r--internal/controlplane/retention.go57
-rw-r--r--internal/controlplane/rotation.go69
-rw-r--r--internal/controlplane/schema.sql265
-rw-r--r--internal/controlplane/snapshot.go108
-rw-r--r--internal/controlplane/store.go82
-rw-r--r--internal/controlplane/store_integration_test.go66
-rw-r--r--internal/controlplane/types.go72
-rw-r--r--internal/controlplane/usage.go16
-rw-r--r--internal/domain/types.go35
-rw-r--r--internal/httpapi/api.go202
-rw-r--r--internal/httpapi/api_test.go2
-rw-r--r--internal/httpapi/proxy.go44
-rw-r--r--internal/mailer/mailer.go5
-rw-r--r--internal/operations/operations.go143
-rw-r--r--internal/provider/forwarder.go19
-rw-r--r--internal/provider/forwarder_test.go39
-rw-r--r--internal/security/credentials.go106
-rw-r--r--internal/security/credentials_test.go28
-rw-r--r--internal/telemetry/metrics.go48
-rw-r--r--internal/usage/observer.go11
-rw-r--r--internal/usage/observer_test.go19
-rwxr-xr-xscripts/backup-postgres.sh14
-rwxr-xr-xscripts/load-smoke.sh22
-rwxr-xr-xscripts/redis-fault-drill.sh18
-rwxr-xr-xscripts/restore-drill.sh14
-rwxr-xr-xscripts/start-debug.sh44
58 files changed, 4710 insertions, 344 deletions
diff --git a/.env.control.example b/.env.control.example
index 8122e07..5020f32 100644
--- a/.env.control.example
+++ b/.env.control.example
@@ -1,4 +1,12 @@
AIGW_SERVER_ADDRESS=:8080
+AIGW_PUBLIC_ADDRESS=:8080
+AIGW_ADMIN_ADDRESS=:8081
+AIGW_WEBHOOK_ADDRESS=:8082
+AIGW_OPERATIONS_ADDRESS=:9090
+# Empty means direct connections only; add only your load balancer CIDRs.
+AIGW_TRUSTED_PROXY_CIDRS=
+AIGW_REQUIRE_HTTPS=false
+AIGW_DEPLOYMENT_REGION=
AIGW_POSTGRES_USER=aigw-local
AIGW_POSTGRES_PASSWORD=replace-with-a-long-random-password
AIGW_POSTGRES_DB=aigw-local
@@ -8,20 +16,30 @@ AIGW_REDIS_URL=redis://127.0.0.1:6379/0
AIGW_REDIS_URL_DOCKER=redis://redis:6379/0
# base64 of exactly 32 random bytes; generate with: openssl rand -base64 32
AIGW_CREDENTIAL_KEY=replace-with-base64-32-byte-key
+# Comma-separated prior keys during a rotation window; remove after re-encryption.
+AIGW_CREDENTIAL_PREVIOUS_KEYS=
AIGW_ADMIN_TOKEN=replace-with-a-long-random-admin-token
-AIGW_PUBLIC_URL=http://localhost:8080/admin/
+AIGW_PUBLIC_URL=http://localhost:8081/admin/
AIGW_WEBAUTHN_RP_ID=localhost
-AIGW_WEBAUTHN_ORIGINS=http://localhost:8080
+AIGW_WEBAUTHN_ORIGINS=http://localhost:8081
AIGW_SMTP_FROM_ADDRESS=no-reply@aigw.local
AIGW_SMTP_ADDRESS=127.0.0.1:1025
AIGW_SMTP_ADDRESS_DOCKER=mailpit:1025
# Leave both empty for a local Mailpit server; set both through a secret manager in production.
AIGW_SMTP_USERNAME=
AIGW_SMTP_PASSWORD=
+# HMAC-SHA256 secret for normalized bounce/complaint callbacks. Inject from a secret manager.
+AIGW_MAIL_FEEDBACK_SECRET=
# Prefer a restricted test key (rk_test_) with Checkout Session write access.
AIGW_STRIPE_API_KEY=replace-with-a-stripe-restricted-key
# Development only: use a separate key with Debugging Tools Write for Stripe CLI.
AIGW_STRIPE_CLI_API_KEY=replace-with-a-separate-cli-restricted-key
AIGW_STRIPE_WEBHOOK_SECRET=replace-with-a-webhook-signing-secret
-AIGW_STRIPE_SUCCESS_URL=http://localhost:8080/admin/?topup=success
-AIGW_STRIPE_CANCEL_URL=http://localhost:8080/admin/?topup=cancel
+AIGW_STRIPE_SUCCESS_URL=http://localhost:8081/admin/?topup=success
+AIGW_STRIPE_CANCEL_URL=http://localhost:8081/admin/?topup=cancel
+AIGW_STRIPE_PORTAL_RETURN_URL=http://localhost:8081/admin/?billing=portal
+AIGW_STRIPE_AUTOMATIC_TAX_ENABLED=false
+AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED=false
+# Set only after a tax adviser confirms the canonical Stripe Tax code for the product.
+AIGW_STRIPE_PRODUCT_TAX_CODE=
+AIGW_SETTLEMENT_SPOOL_PATH=/var/lib/aigw/settlements.jsonl
diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml
new file mode 100644
index 0000000..ebfe3cb
--- /dev/null
+++ b/.github/workflows/ci.yml
@@ -0,0 +1,95 @@
+name: ci
+
+on:
+ push:
+ branches: [main]
+ tags: ['v*']
+ pull_request:
+
+permissions:
+ contents: read
+
+jobs:
+ test:
+ runs-on: ubuntu-latest
+ services:
+ postgres:
+ image: postgres:16-alpine
+ env:
+ POSTGRES_USER: aigw
+ POSTGRES_PASSWORD: integration-only
+ POSTGRES_DB: aigw_test
+ ports: ['5432:5432']
+ options: >-
+ --health-cmd "pg_isready -U aigw -d aigw_test"
+ --health-interval 5s --health-timeout 3s --health-retries 20
+ redis:
+ image: redis:7-alpine
+ ports: ['6379:6379']
+ options: >-
+ --health-cmd "redis-cli ping" --health-interval 5s --health-timeout 3s --health-retries 20
+ env:
+ AIGW_TEST_DATABASE_URL: postgres://aigw:integration-only@127.0.0.1:5432/aigw_test?sslmode=disable
+ steps:
+ - uses: actions/checkout@v4
+ - uses: actions/setup-go@v5
+ with:
+ go-version-file: go.mod
+ cache: true
+ - run: go test -count=1 -p=1 ./...
+ - run: go test -race -count=1 -p=1 ./...
+ - run: go vet ./...
+ - run: CGO_ENABLED=0 go build -buildvcs=false -trimpath ./cmd/...
+
+ image:
+ needs: test
+ runs-on: ubuntu-latest
+ permissions:
+ contents: read
+ packages: write
+ id-token: write
+ steps:
+ - uses: actions/checkout@v4
+ - uses: docker/setup-buildx-action@v3
+ - uses: docker/build-push-action@v6
+ with:
+ context: .
+ load: true
+ tags: aigw:${{ github.sha }}
+ - uses: anchore/sbom-action@v0
+ with:
+ image: aigw:${{ github.sha }}
+ format: spdx-json
+ output-file: sbom.spdx.json
+ - uses: actions/upload-artifact@v4
+ with:
+ name: sbom-spdx
+ path: sbom.spdx.json
+ - uses: aquasecurity/trivy-action@0.28.0
+ with:
+ image-ref: aigw:${{ github.sha }}
+ format: table
+ exit-code: '1'
+ ignore-unfixed: true
+ severity: HIGH,CRITICAL
+ - uses: docker/login-action@v3
+ if: startsWith(github.ref, 'refs/tags/v')
+ with:
+ registry: ghcr.io
+ username: ${{ github.actor }}
+ password: ${{ secrets.GITHUB_TOKEN }}
+ - id: publish
+ uses: docker/build-push-action@v6
+ if: startsWith(github.ref, 'refs/tags/v')
+ with:
+ context: .
+ push: true
+ tags: ghcr.io/${{ github.repository }}:${{ github.ref_name }}
+ - uses: sigstore/cosign-installer@v3
+ if: startsWith(github.ref, 'refs/tags/v')
+ - name: Sign image by digest
+ if: startsWith(github.ref, 'refs/tags/v')
+ env:
+ IMAGE: ghcr.io/${{ github.repository }}
+ DIGEST: ${{ steps.publish.outputs.digest }}
+ run: cosign sign --yes "$IMAGE@$DIGEST"
diff --git a/Dockerfile b/Dockerfile
index 4f31386..f7d2c28 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -12,6 +12,6 @@ FROM alpine:3.22
RUN apk add --no-cache ca-certificates && adduser -D -H -u 10001 aigw
USER aigw
COPY --from=build /out/aigw /usr/local/bin/aigw
-EXPOSE 8080
+EXPOSE 8080 8081 8082 9090
ENTRYPOINT ["/usr/local/bin/aigw"]
CMD ["-config", "/etc/aigw/config.json"]
diff --git a/Makefile b/Makefile
index b8e638f..a51ea8d 100644
--- a/Makefile
+++ b/Makefile
@@ -1,4 +1,4 @@
-.PHONY: build test vet run mock migrate
+.PHONY: build test test-race test-integration vet run mock migrate reconcile
build:
CGO_ENABLED=0 go build -buildvcs=false -trimpath -o aigw ./cmd/aigw
@@ -6,6 +6,13 @@ build:
test:
go test ./...
+test-race:
+ go test -race ./...
+
+test-integration:
+ test -n "$$AIGW_TEST_DATABASE_URL"
+ go test -count=1 -p=1 ./...
+
vet:
go vet ./...
@@ -17,3 +24,6 @@ mock:
migrate:
go run ./cmd/migrate
+
+reconcile:
+ go run ./cmd/reconcile-billing
diff --git a/README.md b/README.md
index dfec494..fbf22a4 100644
--- a/README.md
+++ b/README.md
@@ -95,14 +95,14 @@ curl http://127.0.0.1:8080/anthropic/v1/messages \
## 生产边界
-当前版本可以作为带预付计费的数据面,但上线前仍需完成业务侧对账:
+当前版本可以作为带预付计费的数据面,并已把控制面、账务和运营入口拆开:
-- 控制面计费模式会把 UsageEvent、冻结记录和扣费流水同步、幂等写入 PostgreSQL;异步结构化日志只是可观测副本,丢弃不会影响账本。
+- 控制面计费模式会把 UsageEvent、冻结记录和扣费流水同步、幂等写入 PostgreSQL;响应结束只负责把结算事件投递到持久化队列,worker 负责重试、过期冻结恢复和本地 JSONL spool 补偿。
- 客户密钥和模型目录在控制面模式下从 PostgreSQL 载入到原子内存快照;Redis 只是可选的变更广播加速层,故障时通过 PostgreSQL generation 轮询收敛。
- 当前只把 OpenAI 入口发给 OpenAI 兼容上游、Anthropic 入口发给 Anthropic 兼容上游,不做跨协议转换。
- 自动故障转移可能在极少数网络错误下造成上游重复执行。正式计费时需要上游幂等能力、请求去重策略和重复成本对账。
-- 上游必须返回 usage 才能按 token 结算;未返回 usage 的成功响应当前记为零费用,应在上生产前为每个供应商做账单对账或补充 token 计算器。
-- `/metrics` 应仅在内网暴露;公网 TLS、WAF 和连接层限速应放在负载均衡器或边缘代理。
+- 计费上游必须返回 usage。OpenAI 流请求会强制请求 `stream_options.include_usage=true`;成功响应缺少 usage 时不会按零费用放行,也不会猜测 token,而是把授权保持为 `metering_failed`、触发 readiness/Prometheus 告警,直到运营切断或修复该上游。显式的零 token usage 仍可正常结算。
+- 默认生产配置使用独立的推理、管理、支付/邮件 Webhook 和 operations listener。`/healthz` 只表示进程存活,`/readyz` 会检查 PG、快照、Stripe 对账、Webhook、退款、未收款、缺失用量、结算队列和邮件积压;Redis 是可降级传播层。`/metrics` 应仅在内网暴露;公网 TLS、WAF 和连接层限速应放在负载均衡器或边缘代理。
## PostgreSQL + Redis 控制面
@@ -126,9 +126,9 @@ curl http://127.0.0.1:8080/anthropic/v1/messages \
源码未变化时可跳过镜像构建以快速重启:`AIGW_DEBUG_SKIP_BUILD=1 ./scripts/start-debug.sh`。默认构建使用 Docker host network;特殊环境可以通过 `AIGW_DOCKER_BUILD_NETWORK=default` 覆盖。
-然后打开 `http://127.0.0.1:8080/admin/`,本地邮件在 `http://127.0.0.1:8025/` 查看。启用注册时,新账号必须通过一次性邮件链接验证;团队成员由管理员邀请并自行设置密码。平台管理员也可以从权限为 `0600` 的环境文件读取 `AIGW_ADMIN_TOKEN`,将其作为 bootstrap/break-glass 凭证。日常操作使用邮箱/密码、TOTP 或 Passkey,服务端创建可逐设备撤销的数据库会话,所有写请求需要 CSRF token。第一套资源的创建顺序是:Tenant → Project → API key → Provider → Model route。客户 API Key 明文只在创建成功时返回一次;团队成员使用自己的账号,不共享管理员令牌。
+然后打开 `http://127.0.0.1:8081/admin/`,本地邮件在 `http://127.0.0.1:8025/` 查看;健康检查在 `http://127.0.0.1:9090/readyz`。Stripe CLI Webhook 转发到 `http://127.0.0.1:8082/billing/stripe/webhook`。启用注册时,新账号必须通过一次性邮件链接验证;团队成员由管理员邀请并自行设置密码。平台管理员也可以从权限为 `0600` 的环境文件读取 `AIGW_ADMIN_TOKEN`,将其作为 bootstrap/break-glass 凭证。日常操作使用邮箱/密码、TOTP 或 Passkey,服务端创建可逐设备撤销的数据库会话,所有写请求需要 CSRF token。第一套资源的创建顺序是:Tenant → Project → API key → Provider → Model route。客户 API Key 明文只在创建成功时返回一次;团队成员使用自己的账号,不共享管理员令牌。
-账号邮件先在 PostgreSQL outbox 中加密持久化,再由后台 worker 发送;SMTP 临时不可用不会回滚注册、邀请或重置请求,worker 会重试并记录失败。生产环境把 `AIGW_PUBLIC_URL` 设置为 HTTPS 控制台 URL,把 `AIGW_SMTP_ADDRESS`、`AIGW_SMTP_FROM_ADDRESS`、`AIGW_SMTP_USERNAME`、`AIGW_SMTP_PASSWORD` 通过密钥管理服务注入,并将 `admin.mail.tls_mode` 改为 `starttls` 或 `tls`。WebAuthn 的 `AIGW_WEBAUTHN_RP_ID` 必须是控制台有效域名,`AIGW_WEBAUTHN_ORIGINS` 是逗号分隔的 HTTPS origin。
+账号邮件先在 PostgreSQL outbox 中加密持久化,再由后台 worker 发送;SMTP 临时不可用不会回滚注册、邀请或重置请求。普通失败指数退避,10 次后进入 dead-letter。生产环境把 `AIGW_PUBLIC_URL` 设置为 HTTPS 控制台 URL,把 `AIGW_SMTP_ADDRESS`、`AIGW_SMTP_FROM_ADDRESS`、`AIGW_SMTP_USERNAME`、`AIGW_SMTP_PASSWORD` 和 `AIGW_MAIL_FEEDBACK_SECRET` 通过密钥管理服务注入,并将 `admin.mail.tls_mode` 改为 `starttls` 或 `tls`。邮件供应商的 bounce/complaint 事件应由边缘适配器规范化后签名发送到 `/mail/feedback`;永久退信和投诉地址会进入抑制表。worker 会按日幂等发送低余额和异常消费通知。域名 DNS 仍必须在邮件供应商处配置 SPF、DKIM 和 DMARC,这不是应用代码可以代替的步骤。WebAuthn 的 `AIGW_WEBAUTHN_RP_ID` 必须是控制台有效域名,`AIGW_WEBAUTHN_ORIGINS` 是逗号分隔的 HTTPS origin。
控制台角色分为:`platform_admin`、`platform_viewer`、`tenant_admin`、`tenant_billing`、`tenant_developer`、`tenant_viewer`。租户角色的查询条件在服务端下推到 PostgreSQL,不能读取其他租户的项目、密钥、余额、Usage 或审计事件;供应商凭证和路由管理只对平台角色开放。
@@ -147,7 +147,7 @@ Redis 可用时,RPM/TPM/并发通过 Lua 原子执行并在多实例间共享ï
## 余额与 Stripe 充值
-模型价格在后台按“币种单位 / 100 万 token”配置,数据库使用 `amount_micros` 固定精度整数保存金额。推理请求会先按请求体字节数和 `max_tokens`/`max_completion_tokens` 保守冻结额度;成功响应按上游返回的输入、输出和缓存 token 结算,失败请求释放冻结。`request_id` 是用量与扣费幂等键。
+模型价格在后台按“币种单位 / 100 万 token”配置,数据库使用 `amount_micros` 固定精度整数保存金额。Stripe 不直接为推理请求结账,只向 PostgreSQL 预付钱包充值;推理请求先按请求体字节数和 `max_tokens`/`max_completion_tokens` 保守冻结余额,成功响应按可信 usage 扣款,失败请求释放冻结。`request_id` 是用量与扣费幂等键,余额、冻结和不可变 ledger 都在同一个 PG 事务中更新。
Stripe 使用托管 Checkout,服务端不会接触卡号,也没有硬编码支付方式;支付方式由 Stripe Dashboard 动态配置。充值只在签名校验通过的 Webhook 确认 `payment_status=paid` 后入账,成功跳转页不会直接修改余额。
@@ -158,7 +158,7 @@ cp .env.control.example .env
# 将 AIGW_STRIPE_API_KEY 设为最小权限的 rk_test_ restricted key
# AIGW_STRIPE_CLI_API_KEY 使用另一把仅有 Debugging Tools Write 的测试 key
# 本地转发会显示 whsec_...,填入 AIGW_STRIPE_WEBHOOK_SECRET
-stripe listen --api-key "$AIGW_STRIPE_CLI_API_KEY" --forward-to http://127.0.0.1:8080/billing/stripe/webhook
+stripe listen --api-key "$AIGW_STRIPE_CLI_API_KEY" --forward-to http://127.0.0.1:8082/billing/stripe/webhook
docker compose up --build
```
@@ -169,7 +169,11 @@ Webhook 至少订阅:
- `checkout.session.async_payment_failed`
- `checkout.session.expired`
-当前没有默认启用 Stripe Tax,因为是否有有效税务注册不能由代码推断。确认注册和税务处理方案后再显式加入 `automatic_tax`。生产环境应把 Stripe restricted key 和 Webhook signing secret 放入云平台的密钥管理服务,并限制密钥权限和来源 IP,不要放进镜像或仓库。
+网关 restricted key 的最小权限按实际启用功能配置:Checkout Sessions Write(充值与对账读取)、Customer Portal Write、Customers Write、Charges and Refunds Write、Payment Intents Read 和 Invoices Read。不要给主网关 Debugging Tools 权限;Stripe CLI 使用独立的测试 key。Dashboard 保存权限时可能要求账户持有人完成二次验证。
+
+后台已覆盖 Checkout 重试、客户门户、退款队列、争议/发票/收据记录、失败重试、周期性 Stripe 对账和 CSV 财务导出。Stripe 已支付但本地 pending 的订单会自动补账且传播错误不会被忽略;Stripe 中不存在的未入账订单只能通过 RBAC/审计保护的 resolution 作废,已入账孤立充值只能用等额负向账本冲销,原记录不会删除或改写。退款与争议会先冻结/扣除本地余额,余额不足进入 `uncollected_micros`,不会静默丢账。当前没有默认启用 Stripe Tax,因为是否有有效税务注册不能由代码推断;确认注册和 canonical product tax code 后再显式打开。生产环境应把 Stripe restricted key 和 Webhook signing secret 放入云平台的密钥管理服务,并限制密钥权限和来源 IP,不要放进镜像或仓库。
+
+模型目录支持输入/输出模态、上下文窗口、最大输出、能力集合、生命周期、弃用替代模型、区域和租户/API key allowlist;价格在 `model_price_versions` 中按生效时间版本化。推理请求会在路由前执行区域、生命周期、能力和 allowlist 校验,旧 alias 只映射到 canonical model ID。
静态模式的上游 Base URL 与 API Key 分别使用 `base_url_env` 和 `api_key_env`;控制面、Redis、Stripe、SMTP、WebAuthn 和监听地址同样只通过环境变量或密钥管理服务注入。严格 JSON 解析会拒绝旧的 `address`、`base_url`、`success_url` 和 `cancel_url` 字面量字段,避免环境隔离被配置文件绕过。版本化配置只保存 `*_env` 名称,不保存外部服务密钥或部署域名。
@@ -177,6 +181,8 @@ Webhook 至少订阅:
```bash
AIGW_DATABASE_URL="postgres://..." go run ./cmd/migrate
+# 查看当前迁移版本和 checksum
+AIGW_DATABASE_URL="postgres://..." go run ./cmd/migrate -status
```
管理 API 支持账号密码、TOTP、Passkey 会话和仅用于初始化/故障恢复的 bootstrap token。即使已经启用 RBAC,生产环境仍应设置 HTTPS、把推理端口和管理端口分离、按需关闭公开注册,并把 bootstrap token 存入密钥管理服务;企业部署可再接 OIDC/SAML 与强制 MFA 策略。
@@ -187,5 +193,9 @@ AIGW_DATABASE_URL="postgres://..." go run ./cmd/migrate
```bash
go test ./...
+go test -race ./...
go vet ./...
+CGO_ENABLED=0 go build -buildvcs=false ./cmd/...
```
+
+设置 `AIGW_TEST_DATABASE_URL` 后,测试会强制执行 PostgreSQL Webhook/结算/迁移升级用例。CI 还构建容器镜像、生成 SPDX SBOM、以 Trivy 阻断 HIGH/CRITICAL 漏洞,并在 `v*` 标签发布时使用 OIDC keyless Cosign 签名。负载与 Redis 降级演练分别使用 `scripts/load-smoke.sh` 和 `scripts/redis-fault-drill.sh`;备份恢复演练使用 `scripts/backup-postgres.sh` 与 `scripts/restore-drill.sh`。
diff --git a/cmd/aigw/main.go b/cmd/aigw/main.go
index c7d9ed3..108a849 100644
--- a/cmd/aigw/main.go
+++ b/cmd/aigw/main.go
@@ -21,6 +21,7 @@ import (
"aigw/internal/httpapi"
"aigw/internal/limits"
"aigw/internal/mailer"
+ "aigw/internal/operations"
"aigw/internal/provider"
"aigw/internal/routing"
"aigw/internal/telemetry"
@@ -69,7 +70,8 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error {
store, err = controlplane.NewStore(ctx, controlplane.Options{
DatabaseURL: cfg.ControlPlane.DatabaseURL, RedisURL: cfg.ControlPlane.RedisURL,
CredentialKey: cfg.ControlPlane.CredentialKey, RedisChannel: cfg.ControlPlane.RedisChannel,
- VersionCacheKey: cfg.ControlPlane.SnapshotCacheKey,
+ PreviousCredentialKeys: cfg.ControlPlane.PreviousCredentialKeys,
+ VersionCacheKey: cfg.ControlPlane.SnapshotCacheKey,
})
if err != nil {
return err
@@ -138,6 +140,15 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error {
return fmt.Errorf("configure SMTP: %w", err)
}
go mailer.NewWorker(store, sender, logger).Run(ctx)
+ go store.RunMailNotificationWorker(ctx, controlplane.MailNotificationConfig{
+ LowBalanceMicros: cfg.Admin.Mail.LowBalanceMicros,
+ SpendAnomalyMultiplier: cfg.Admin.Mail.SpendAnomalyMultiplier,
+ SpendAnomalyMinMicros: cfg.Admin.Mail.SpendAnomalyMinMicros,
+ Interval: time.Duration(cfg.Admin.Mail.NotificationIntervalSeconds) * time.Second,
+ }, logger)
+ }
+ if cfg.Admin.Enabled {
+ go store.RunRetentionWorker(ctx, cfg.Admin.AuditRetentionDays, cfg.Admin.SecurityRetentionDays, logger)
}
var billingService *billing.Service
@@ -150,11 +161,18 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error {
StripeEnabled: cfg.Billing.Stripe.Enabled, StripeAPIKey: cfg.Billing.Stripe.APIKey,
StripeWebhookSecret: cfg.Billing.Stripe.WebhookSecret,
StripeSuccessURL: cfg.Billing.Stripe.SuccessURL, StripeCancelURL: cfg.Billing.Stripe.CancelURL,
+ StripePortalReturnURL: cfg.Billing.Stripe.PortalReturnURL,
+ StripeAutomaticTax: cfg.Billing.Stripe.AutomaticTaxEnabled,
+ StripeProductTaxCode: cfg.Billing.Stripe.ProductTaxCode,
+ SettlementSpoolPath: cfg.Billing.SettlementSpoolPath,
+ Metrics: metrics,
})
if err != nil {
return err
}
defer billingService.Close()
+ go billingService.RunSettlementWorker(ctx)
+ go billingService.RunStripeOperations(ctx)
}
var billingMeter billing.Meter
if billingService != nil {
@@ -162,58 +180,109 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error {
}
inferenceAPI := httpapi.New(httpapi.Options{
- Authenticator: authenticator,
- Catalog: modelCatalog,
- Router: routing.New(modelCatalog),
- Forwarder: provider.New(cfg.UpstreamHTTP, metrics),
- UsageSink: usageSink,
- BillingMeter: billingMeter,
- Limiter: requestLimiter,
- UsageRecorder: store,
- Metrics: metrics,
- Logger: logger,
- MaxBodyBytes: cfg.Server.MaxBodyBytes,
- ExposeMetrics: cfg.Observability.ExposeMetrics,
+ Authenticator: authenticator,
+ Catalog: modelCatalog,
+ Router: routing.New(modelCatalog),
+ Forwarder: provider.New(cfg.UpstreamHTTP, metrics),
+ UsageSink: usageSink,
+ BillingMeter: billingMeter,
+ Limiter: requestLimiter,
+ UsageRecorder: optionalUsageRecorder(store),
+ Metrics: metrics,
+ Logger: logger,
+ MaxBodyBytes: cfg.Server.MaxBodyBytes,
+ ExposeMetrics: cfg.Observability.ExposeMetrics,
+ DeploymentRegion: cfg.Server.DeploymentRegion,
})
- root := http.NewServeMux()
- root.Handle("/", inferenceAPI.Handler())
+ adminHandler := http.Handler(nil)
if cfg.Admin.Enabled {
- adminHandler := adminapi.New(adminapi.Options{
+ adminHandler = adminapi.New(adminapi.Options{
Store: store, Manager: manager, Billing: billingService, Token: cfg.Admin.Token,
Logger: logger, Prefix: cfg.Admin.BasePath, RegistrationEnabled: cfg.Admin.RegistrationEnabled,
SessionTTL: time.Duration(cfg.Admin.SessionTTLHours) * time.Hour, Currency: cfg.Billing.Currency,
PublicURL: cfg.Admin.PublicURL, WebAuthn: webAuthn, MailEnabled: cfg.Admin.Mail.Enabled,
}).Handler()
- root.Handle(cfg.Admin.BasePath, adminHandler)
- root.Handle(cfg.Admin.BasePath+"/", adminHandler)
}
- if billingService != nil && billingService.StripeEnabled() {
- root.Handle("/billing/stripe/webhook", billingService.WebhookHandler())
+ operationHandler := operations.Handler{Store: store, Manager: manager, Billing: billingService, Metrics: metrics,
+ MaxSnapshotAge: 3 * time.Duration(cfg.ControlPlane.ReloadIntervalSeconds) * time.Second}
+ return serve(ctx, cfg, logger, store, inferenceAPI, adminHandler, billingService, operationHandler)
+}
+
+func optionalUsageRecorder(store *controlplane.Store) httpapi.UsageRecorder {
+ if store == nil {
+ return nil
}
+ return store
+}
- server := &http.Server{
- Addr: cfg.Server.Address, Handler: root,
- ReadHeaderTimeout: cfg.Server.ReadHeaderTimeout(),
- IdleTimeout: cfg.Server.IdleTimeout(),
- }
- errCh := make(chan error, 1)
- go func() {
- logger.Info("gateway_listening", "address", cfg.Server.Address, "control_plane", cfg.ControlPlane.Enabled,
- "billing", cfg.Billing.Enabled, "stripe", cfg.Billing.Stripe.Enabled)
- errCh <- server.ListenAndServe()
- }()
+func serve(ctx context.Context, cfg config.Config, logger *slog.Logger, store *controlplane.Store, inferenceAPI *httpapi.API, adminHandler http.Handler, billingService *billing.Service, operationHandler http.Handler) error {
+ type listener struct {
+ name, address string
+ handler http.Handler
+ }
+ var listeners []listener
+ webhookMux := http.NewServeMux()
+ hasWebhooks := false
+ if billingService != nil && billingService.StripeEnabled() {
+ webhookMux.Handle("/billing/stripe/webhook", billingService.WebhookHandler())
+ hasWebhooks = true
+ }
+ if cfg.Admin.Enabled && cfg.Admin.Mail.Enabled && cfg.Admin.Mail.FeedbackSecret != "" {
+ webhookMux.Handle("/mail/feedback", store.MailFeedbackHandler(cfg.Admin.Mail.FeedbackSecret))
+ hasWebhooks = true
+ }
+ if cfg.Server.SplitListeners {
+ listeners = append(listeners, listener{"inference", cfg.Server.PublicAddress, inferenceAPI.InferenceHandler()}, listener{"operations", cfg.Server.OperationsAddress, operationHandler})
+ if adminHandler != nil {
+ listeners = append(listeners, listener{"admin", cfg.Server.AdminAddress, adminHandler})
+ }
+ if hasWebhooks {
+ listeners = append(listeners, listener{"webhooks", cfg.Server.WebhookAddress, webhookMux})
+ }
+ } else {
+ root := http.NewServeMux()
+ root.Handle("/healthz", operationHandler)
+ root.Handle("/readyz", operationHandler)
+ root.Handle("/metrics", operationHandler)
+ root.Handle("/", inferenceAPI.Handler())
+ if adminHandler != nil {
+ root.Handle(cfg.Admin.BasePath, adminHandler)
+ root.Handle(cfg.Admin.BasePath+"/", adminHandler)
+ }
+ if hasWebhooks {
+ root.Handle("/billing/stripe/webhook", webhookMux)
+ root.Handle("/mail/feedback", webhookMux)
+ }
+ listeners = []listener{{"combined", cfg.Server.Address, root}}
+ }
+ servers := make([]*http.Server, 0, len(listeners))
+ errCh := make(chan error, len(listeners))
+ for _, item := range listeners {
+ handler, err := httpapi.TrustProxyHeaders(item.handler, cfg.Server.TrustedProxyCIDRs, cfg.Server.RequireHTTPS && item.name != "operations")
+ if err != nil {
+ return err
+ }
+ server := &http.Server{Addr: item.address, Handler: handler, ReadHeaderTimeout: cfg.Server.ReadHeaderTimeout(), IdleTimeout: cfg.Server.IdleTimeout()}
+ servers = append(servers, server)
+ go func(item listener, server *http.Server) {
+ logger.Info("http_listener_started", "name", item.name, "address", item.address)
+ errCh <- server.ListenAndServe()
+ }(item, server)
+ }
select {
case <-ctx.Done():
- shutdownContext, cancel := context.WithTimeout(context.Background(), cfg.Server.ShutdownTimeout())
- defer cancel()
- if err := server.Shutdown(shutdownContext); err != nil {
- return fmt.Errorf("shutdown HTTP server: %w", err)
- }
- return nil
case err := <-errCh:
- if errors.Is(err, http.ErrServerClosed) {
- return nil
+ if !errors.Is(err, http.ErrServerClosed) {
+ return fmt.Errorf("serve HTTP: %w", err)
+ }
+ }
+ shutdownCtx, cancel := context.WithTimeout(context.Background(), cfg.Server.ShutdownTimeout())
+ defer cancel()
+ var shutdownErr error
+ for _, server := range servers {
+ if err := server.Shutdown(shutdownCtx); err != nil && shutdownErr == nil {
+ shutdownErr = err
}
- return fmt.Errorf("serve HTTP: %w", err)
}
+ return shutdownErr
}
diff --git a/cmd/aigw/main_test.go b/cmd/aigw/main_test.go
new file mode 100644
index 0000000..449c35d
--- /dev/null
+++ b/cmd/aigw/main_test.go
@@ -0,0 +1,14 @@
+package main
+
+import (
+ "testing"
+
+ "aigw/internal/controlplane"
+)
+
+func TestOptionalUsageRecorderRejectsTypedNilStore(t *testing.T) {
+ var store *controlplane.Store
+ if recorder := optionalUsageRecorder(store); recorder != nil {
+ t.Fatal("nil control-plane store became a non-nil usage recorder")
+ }
+}
diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go
index 20452f9..a24b7bc 100644
--- a/cmd/migrate/main.go
+++ b/cmd/migrate/main.go
@@ -12,6 +12,7 @@ import (
func main() {
environment := flag.String("database-url-env", "AIGW_DATABASE_URL", "environment variable containing the PostgreSQL URL")
+ status := flag.Bool("status", false, "print the latest applied migration instead of changing the schema")
flag.Parse()
databaseURL := os.Getenv(*environment)
if databaseURL == "" {
@@ -20,6 +21,15 @@ func main() {
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
+ if *status {
+ migration, err := controlplane.MigrationStatusDatabase(ctx, databaseURL)
+ if err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ fmt.Printf("version=%d name=%s checksum=%s applied_at=%s\n", migration.Version, migration.Name, migration.Checksum, migration.AppliedAt.Format(time.RFC3339))
+ return
+ }
if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil {
fmt.Fprintln(os.Stderr, err)
os.Exit(1)
diff --git a/cmd/reconcile-billing/main.go b/cmd/reconcile-billing/main.go
new file mode 100644
index 0000000..79329ad
--- /dev/null
+++ b/cmd/reconcile-billing/main.go
@@ -0,0 +1,79 @@
+package main
+
+import (
+ "context"
+ "encoding/json"
+ "flag"
+ "fmt"
+ "os"
+ "time"
+
+ "aigw/internal/billing"
+)
+
+func main() {
+ databaseURLEnv := flag.String("database-url-env", "AIGW_DATABASE_URL", "environment variable containing the PostgreSQL URL")
+ stripeKeyEnv := flag.String("stripe-key-env", "AIGW_STRIPE_API_KEY", "environment variable containing the Stripe restricted key")
+ currency := flag.String("currency", "usd", "wallet currency")
+ limit := flag.Int("limit", 200, "maximum recent top-up orders to reconcile")
+ resolveMissing := flag.Bool("resolve-confirmed-missing", false, "close uncredited pending orders whose Stripe Sessions are confirmed missing")
+ resolutionReason := flag.String("resolution-reason", "maintenance reconciliation confirmed the uncredited Checkout Session is absent from the configured Stripe account", "auditable reason for resolving missing orders")
+ flag.Parse()
+ databaseURL, stripeKey := os.Getenv(*databaseURLEnv), os.Getenv(*stripeKeyEnv)
+ if databaseURL == "" || stripeKey == "" {
+ fmt.Fprintln(os.Stderr, "database and Stripe key environment variables are required")
+ os.Exit(1)
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute)
+ defer cancel()
+ service, err := billing.New(ctx, billing.Options{DatabaseURL: databaseURL, Currency: *currency, StripeEnabled: true, StripeAPIKey: stripeKey})
+ if err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ defer service.Close()
+ result, err := service.Reconcile(ctx, *limit)
+ if err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ if *resolveMissing {
+ for _, mismatch := range result.Mismatches {
+ if mismatch["type"] != "stripe_session_missing" {
+ continue
+ }
+ orderID, _ := mismatch["order_id"].(string)
+ order, getErr := service.GetTopUpOrder(ctx, "", orderID)
+ if getErr != nil {
+ fmt.Fprintln(os.Stderr, getErr)
+ os.Exit(1)
+ }
+ var resolveErr error
+ if order.Status == "paid" {
+ _, resolveErr = service.ReverseMissingTopUpCredit(ctx, order.TenantID, order.ID,
+ billing.ResolveMissingTopUpInput{Reason: *resolutionReason}, billing.ResolutionActor{ID: "reconcile-billing", Type: "maintenance"})
+ } else {
+ _, resolveErr = service.ResolveMissingTopUp(ctx, order.TenantID, order.ID,
+ billing.ResolveMissingTopUpInput{Reason: *resolutionReason}, billing.ResolutionActor{ID: "reconcile-billing", Type: "maintenance"})
+ }
+ if resolveErr != nil {
+ fmt.Fprintln(os.Stderr, resolveErr)
+ os.Exit(1)
+ }
+ }
+ result, err = service.Reconcile(ctx, *limit)
+ if err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ }
+ encoder := json.NewEncoder(os.Stdout)
+ encoder.SetIndent("", " ")
+ if err := encoder.Encode(result); err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ if result.MismatchCount > 0 {
+ os.Exit(2)
+ }
+}
diff --git a/cmd/rotate-credentials/main.go b/cmd/rotate-credentials/main.go
new file mode 100644
index 0000000..598035d
--- /dev/null
+++ b/cmd/rotate-credentials/main.go
@@ -0,0 +1,41 @@
+package main
+
+import (
+ "context"
+ "fmt"
+ "os"
+ "strings"
+ "time"
+
+ "aigw/internal/controlplane"
+)
+
+func main() {
+ databaseURL := strings.TrimSpace(os.Getenv("AIGW_DATABASE_URL"))
+ current := strings.TrimSpace(os.Getenv("AIGW_CREDENTIAL_KEY"))
+ previousRaw := os.Getenv("AIGW_CREDENTIAL_PREVIOUS_KEYS")
+ if databaseURL == "" || current == "" || strings.TrimSpace(previousRaw) == "" {
+ fmt.Fprintln(os.Stderr, "AIGW_DATABASE_URL, AIGW_CREDENTIAL_KEY, and AIGW_CREDENTIAL_PREVIOUS_KEYS are required")
+ os.Exit(2)
+ }
+ previous := []string{}
+ for _, value := range strings.Split(previousRaw, ",") {
+ if value = strings.TrimSpace(value); value != "" {
+ previous = append(previous, value)
+ }
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute)
+ defer cancel()
+ store, err := controlplane.NewStore(ctx, controlplane.Options{DatabaseURL: databaseURL, CredentialKey: current, PreviousCredentialKeys: previous})
+ if err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ defer store.Close()
+ count, err := store.RotateCredentials(ctx)
+ if err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ fmt.Printf("re-encrypted %d credential records\n", count)
+}
diff --git a/config.control.example.json b/config.control.example.json
index 2b05636..58ef2e7 100644
--- a/config.control.example.json
+++ b/config.control.example.json
@@ -1,6 +1,14 @@
{
"server": {
"address_env": "AIGW_SERVER_ADDRESS",
+ "split_listeners": true,
+ "public_address_env": "AIGW_PUBLIC_ADDRESS",
+ "admin_address_env": "AIGW_ADMIN_ADDRESS",
+ "webhook_address_env": "AIGW_WEBHOOK_ADDRESS",
+ "operations_address_env": "AIGW_OPERATIONS_ADDRESS",
+ "trusted_proxy_cidrs_env": "AIGW_TRUSTED_PROXY_CIDRS",
+ "require_https_env": "AIGW_REQUIRE_HTTPS",
+ "deployment_region_env": "AIGW_DEPLOYMENT_REGION",
"max_body_bytes": 16777216,
"read_header_timeout_seconds": 10,
"idle_timeout_seconds": 120,
@@ -14,6 +22,7 @@
"database_url_env": "AIGW_DATABASE_URL",
"redis_url_env": "AIGW_REDIS_URL",
"credential_key_env": "AIGW_CREDENTIAL_KEY",
+ "previous_credential_keys_env": "AIGW_CREDENTIAL_PREVIOUS_KEYS",
"redis_channel": "aigw:control:changed",
"snapshot_cache_key": "aigw:control:generation",
"reload_interval_seconds": 30,
@@ -25,6 +34,8 @@
"base_path": "/admin",
"registration_enabled": true,
"session_ttl_hours": 12,
+ "audit_retention_days": 2555,
+ "security_retention_days": 30,
"public_url_env": "AIGW_PUBLIC_URL",
"mail": {
"enabled": true,
@@ -33,6 +44,11 @@
"smtp_address_env": "AIGW_SMTP_ADDRESS",
"smtp_username_env": "AIGW_SMTP_USERNAME",
"smtp_password_env": "AIGW_SMTP_PASSWORD",
+ "feedback_secret_env": "AIGW_MAIL_FEEDBACK_SECRET",
+ "low_balance_micros": 5000000,
+ "spend_anomaly_multiplier": 3,
+ "spend_anomaly_min_micros": 10000000,
+ "notification_interval_seconds": 300,
"tls_mode": "none"
},
"webauthn": {
@@ -48,12 +64,17 @@
"default_max_output_tokens": 4096,
"min_top_up_minor": 500,
"max_top_up_minor": 1000000,
+ "settlement_spool_path_env": "AIGW_SETTLEMENT_SPOOL_PATH",
"stripe": {
"enabled": true,
"api_key_env": "AIGW_STRIPE_API_KEY",
"webhook_secret_env": "AIGW_STRIPE_WEBHOOK_SECRET",
"success_url_env": "AIGW_STRIPE_SUCCESS_URL",
- "cancel_url_env": "AIGW_STRIPE_CANCEL_URL"
+ "cancel_url_env": "AIGW_STRIPE_CANCEL_URL",
+ "portal_return_url_env": "AIGW_STRIPE_PORTAL_RETURN_URL",
+ "automatic_tax_enabled_env": "AIGW_STRIPE_AUTOMATIC_TAX_ENABLED",
+ "tax_registration_confirmed_env": "AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED",
+ "product_tax_code_env": "AIGW_STRIPE_PRODUCT_TAX_CODE"
}
},
"upstream_http": {
diff --git a/docker-compose.yml b/docker-compose.yml
index 0fc531b..8b6b2f1 100644
--- a/docker-compose.yml
+++ b/docker-compose.yml
@@ -48,11 +48,22 @@ services:
condition: service_healthy
ports:
- "127.0.0.1:8080:8080"
+ - "127.0.0.1:8081:8081"
+ - "127.0.0.1:8082:8082"
+ - "127.0.0.1:9090:9090"
environment:
AIGW_SERVER_ADDRESS: ${AIGW_SERVER_ADDRESS:?set AIGW_SERVER_ADDRESS}
+ AIGW_PUBLIC_ADDRESS: ${AIGW_PUBLIC_ADDRESS:-:8080}
+ AIGW_ADMIN_ADDRESS: ${AIGW_ADMIN_ADDRESS:-:8081}
+ AIGW_WEBHOOK_ADDRESS: ${AIGW_WEBHOOK_ADDRESS:-:8082}
+ AIGW_OPERATIONS_ADDRESS: ${AIGW_OPERATIONS_ADDRESS:-:9090}
+ AIGW_TRUSTED_PROXY_CIDRS: ${AIGW_TRUSTED_PROXY_CIDRS:-}
+ AIGW_REQUIRE_HTTPS: ${AIGW_REQUIRE_HTTPS:-false}
+ AIGW_DEPLOYMENT_REGION: ${AIGW_DEPLOYMENT_REGION:-}
AIGW_DATABASE_URL: ${AIGW_DATABASE_URL_DOCKER:?set AIGW_DATABASE_URL_DOCKER}
AIGW_REDIS_URL: ${AIGW_REDIS_URL_DOCKER:-}
AIGW_CREDENTIAL_KEY: ${AIGW_CREDENTIAL_KEY:?set AIGW_CREDENTIAL_KEY}
+ AIGW_CREDENTIAL_PREVIOUS_KEYS: ${AIGW_CREDENTIAL_PREVIOUS_KEYS:-}
AIGW_ADMIN_TOKEN: ${AIGW_ADMIN_TOKEN:?set AIGW_ADMIN_TOKEN}
AIGW_PUBLIC_URL: ${AIGW_PUBLIC_URL:?set AIGW_PUBLIC_URL}
AIGW_WEBAUTHN_RP_ID: ${AIGW_WEBAUTHN_RP_ID:?set AIGW_WEBAUTHN_RP_ID}
@@ -65,10 +76,16 @@ services:
AIGW_STRIPE_WEBHOOK_SECRET: ${AIGW_STRIPE_WEBHOOK_SECRET:?set AIGW_STRIPE_WEBHOOK_SECRET}
AIGW_STRIPE_SUCCESS_URL: ${AIGW_STRIPE_SUCCESS_URL:?set AIGW_STRIPE_SUCCESS_URL}
AIGW_STRIPE_CANCEL_URL: ${AIGW_STRIPE_CANCEL_URL:?set AIGW_STRIPE_CANCEL_URL}
+ AIGW_STRIPE_PORTAL_RETURN_URL: ${AIGW_STRIPE_PORTAL_RETURN_URL:-http://localhost:8081/admin/?billing=portal}
+ AIGW_STRIPE_AUTOMATIC_TAX_ENABLED: ${AIGW_STRIPE_AUTOMATIC_TAX_ENABLED:-false}
+ AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED: ${AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED:-false}
+ AIGW_STRIPE_PRODUCT_TAX_CODE: ${AIGW_STRIPE_PRODUCT_TAX_CODE:-}
+ AIGW_SETTLEMENT_SPOOL_PATH: /var/lib/aigw/settlements.jsonl
volumes:
- ./config.control.example.json:/etc/aigw/config.json:ro
+ - aigw-settlements:/var/lib/aigw
healthcheck:
- test: ["CMD", "wget", "-q", "-O", "-", "http://127.0.0.1:8080/readyz"]
+ test: ["CMD", "wget", "-q", "-O", "-", "http://127.0.0.1:9090/readyz"]
interval: 3s
timeout: 2s
retries: 20
@@ -77,3 +94,4 @@ services:
volumes:
aigw-postgres:
aigw-redis:
+ aigw-settlements:
diff --git a/docs/runbook.md b/docs/runbook.md
new file mode 100644
index 0000000..ddd1a8c
--- /dev/null
+++ b/docs/runbook.md
@@ -0,0 +1,82 @@
+# Production operations runbook
+
+## HTTP boundary
+
+Production uses four listeners. Expose only `AIGW_PUBLIC_ADDRESS` and the Stripe path on
+`AIGW_WEBHOOK_ADDRESS` to the internet. Put `AIGW_ADMIN_ADDRESS` behind SSO/VPN or an access
+proxy. Keep `AIGW_OPERATIONS_ADDRESS` private to the orchestrator and monitoring network.
+Set `AIGW_REQUIRE_HTTPS=true` behind TLS termination and list only the load balancer CIDRs in
+`AIGW_TRUSTED_PROXY_CIDRS`; forwarding headers from other peers are discarded.
+
+## Alerts
+
+Page when `/readyz` is non-200 for five minutes, `aigw_ready == 0`,
+`aigw_billing_settlement_spool_records > 0`, or the settlement backlog has an item older than
+15 minutes. Warn when Redis is degraded, webhook events are unprocessed,
+`aigw_stripe_refund_backlog` or `aigw_stripe_reconciliation_mismatches` is non-zero, disputes
+need response, dead-letter mail exists, `aigw_billing_unmetered_successes` is non-zero, or
+`aigw_billing_uncollected_micros` increases.
+
+## Stripe reconciliation
+
+Stripe only funds the prepaid wallet; inference never creates a Stripe charge. Run
+`go run ./cmd/reconcile-billing` after deploys and investigate any non-zero exit. A paid Stripe
+session with a local pending order is repaired through the same idempotent webhook transaction.
+Never delete an orphaned local order. Resolve an uncredited missing session through the Admin UI,
+or reverse an already credited missing session with the dedicated equal negative ledger action.
+Both require `billing.adjust` and create an immutable resolution row. The maintenance CLI can do
+the same only with the explicit `-resolve-confirmed-missing` flag and a written reason.
+
+The gateway restricted key should grant only Checkout Sessions Write, Customer Portal Write,
+Customers Write, Charges and Refunds Write, Payment Intents Read, and Invoices Read. Use a separate
+test key for Stripe CLI. Configure the production Webhook signing secret independently and alert on
+signature failures or a five-minute unprocessed backlog.
+
+## Mail deliverability
+
+Use a production SMTP provider with STARTTLS/TLS and inject all endpoints, credentials, sender, and
+`AIGW_MAIL_FEEDBACK_SECRET` through the deployment secret manager. Configure SPF, DKIM, and a DMARC
+policy for the From domain in DNS. Map the provider's delivered/bounce/complaint event into the
+normalized `/mail/feedback` payload and HMAC-sign the exact raw body; bounce/complaint events suppress
+future delivery. Monitor dead-letter outbox rows. Low-balance and anomalous-spend notifications are
+deduplicated per recipient and UTC day.
+
+## Database migrations
+
+Run `cmd/migrate` before deploying application instances and keep `auto_migrate=false` in
+production. Migrations use a PostgreSQL advisory lock and a recorded checksum. Schema rollback
+is always a reviewed forward migration; restore a database backup only for whole-release
+rollback after stopping writers. Never edit an already applied migration body.
+
+## Backup and restore
+
+Run `scripts/backup-postgres.sh` from a host with `pg_dump`, encrypted storage, and a scoped
+database credential. Test `scripts/restore-drill.sh` into a disposable isolated database at
+least monthly. Record row-count evidence and application smoke tests before deleting the drill.
+
+Before each release, run `scripts/load-smoke.sh` against a non-production upstream, then
+`scripts/redis-fault-drill.sh` to prove Redis is optional and PostgreSQL polling keeps readiness.
+Exercise PostgreSQL failover separately and verify pending settlement jobs resume without duplicate
+ledger entries. Archive the command output with the release evidence.
+
+## Credential rotation
+
+1. Put the new key in `AIGW_CREDENTIAL_KEY` and the old key in
+ `AIGW_CREDENTIAL_PREVIOUS_KEYS` on every instance.
+2. Deploy and verify snapshot/MFA/mail decryption.
+3. Run `go run ./cmd/rotate-credentials` once with both variables configured.
+4. Restart with only the new key, then revoke the old key from the secret manager.
+
+Stripe credentials should be separate restricted keys per environment and service. Rotate the
+Webhook endpoint secret with overlapping endpoints, then remove the previous endpoint after all
+instances use the new value.
+
+## Retention
+
+Keep immutable billing ledger and Stripe financial event evidence for the statutory period set
+with finance/legal. Partition and archive Usage by month. Audit events should be exported to
+append-only object storage before database deletion. `admin.audit_retention_days` defaults to
+2555 days and `admin.security_retention_days` defaults to 30 days; a daily worker enforces both.
+Account action tokens, expired sessions, WebAuthn challenges, sent mail and login throttles are
+security data covered by the shorter window. Set the audit period only after finance/legal and
+incident-response owners approve the archive and retrieval procedure.
diff --git a/internal/adminapi/api.go b/internal/adminapi/api.go
index 4e66d7f..9f460e3 100644
--- a/internal/adminapi/api.go
+++ b/internal/adminapi/api.go
@@ -144,6 +144,7 @@ func (a *API) Handler() http.Handler {
mux.HandleFunc("GET "+apiPrefix+"/models", a.withAuth("platform.read", a.listModels))
mux.HandleFunc("POST "+apiPrefix+"/models", a.withAuth("platform.write", a.createModel))
mux.HandleFunc("POST "+apiPrefix+"/models/{id}/toggle", a.withAuth("platform.write", a.toggleModel))
+ mux.HandleFunc("POST "+apiPrefix+"/models/{id}/prices", a.withAuth("platform.write", a.createModelPriceVersion))
mux.HandleFunc("POST "+apiPrefix+"/reload", a.withAuth("platform.write", a.reload))
if a.billing != nil {
mux.HandleFunc("GET "+apiPrefix+"/billing/accounts", a.withAuth("billing.read", a.listBillingAccounts))
@@ -152,6 +153,16 @@ func (a *API) Handler() http.Handler {
mux.HandleFunc("GET "+apiPrefix+"/billing/orders/{id}", a.withAuth("billing.read", a.getTopUpOrder))
mux.HandleFunc("POST "+apiPrefix+"/billing/adjustments", a.withAuth("billing.adjust", a.adjustBalance))
mux.HandleFunc("POST "+apiPrefix+"/billing/checkout-sessions", a.withAuth("billing.topup", a.createCheckoutSession))
+ mux.HandleFunc("POST "+apiPrefix+"/billing/portal-sessions", a.withAuth("billing.topup", a.createPortalSession))
+ mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/retry", a.withAuth("billing.topup", a.retryCheckoutSession))
+ mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/resolve-missing", a.withAuth("billing.adjust", a.resolveMissingTopUp))
+ mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/reverse-missing-credit", a.withAuth("billing.adjust", a.reverseMissingTopUpCredit))
+ mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/refund", a.withAuth("billing.adjust", a.createRefund))
+ mux.HandleFunc("GET "+apiPrefix+"/billing/refunds", a.withAuth("billing.read", a.listRefunds))
+ mux.HandleFunc("GET "+apiPrefix+"/billing/disputes", a.withAuth("billing.read", a.listDisputes))
+ mux.HandleFunc("GET "+apiPrefix+"/billing/invoices", a.withAuth("billing.read", a.listInvoices))
+ mux.HandleFunc("POST "+apiPrefix+"/billing/reconcile", a.withAuth("billing.adjust", a.reconcileBilling))
+ mux.HandleFunc("GET "+apiPrefix+"/billing/export.csv", a.withAuth("billing.read", a.exportBillingCSV))
}
mux.HandleFunc("GET "+apiPrefix+"/usage", a.withAuth("usage.read", a.listUsage))
mux.HandleFunc("GET "+apiPrefix+"/usage/summary", a.withAuth("usage.read", a.usageSummary))
@@ -945,11 +956,11 @@ func (a *API) listBillingLedger(w http.ResponseWriter, r *http.Request) {
func (a *API) listTopUpOrders(w http.ResponseWriter, r *http.Request) {
actor := a.actor(r)
- if actor.TenantID == "" {
- apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r))
- return
+ tenantID := actor.TenantID
+ if tenantID == "" {
+ tenantID = r.URL.Query().Get("tenant_id")
}
- result, err := a.billing.ListTopUpOrders(r.Context(), actor.TenantID, 50)
+ result, err := a.billing.ListTopUpOrders(r.Context(), tenantID, 100)
if err != nil {
a.billingError(w, r, err)
return
@@ -959,10 +970,6 @@ func (a *API) listTopUpOrders(w http.ResponseWriter, r *http.Request) {
func (a *API) getTopUpOrder(w http.ResponseWriter, r *http.Request) {
actor := a.actor(r)
- if actor.TenantID == "" {
- apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r))
- return
- }
result, err := a.billing.GetTopUpOrder(r.Context(), actor.TenantID, r.PathValue("id"))
if err != nil {
a.billingError(w, r, err)
@@ -992,6 +999,7 @@ func (a *API) createCheckoutSession(w http.ResponseWriter, r *http.Request) {
if tenantID := a.actor(r).TenantID; tenantID != "" {
input.TenantID = tenantID
}
+ input.CustomerEmail = a.actor(r).Email
result, err := a.billing.CreateCheckout(r.Context(), input)
if err != nil {
a.billingError(w, r, err)
@@ -1000,6 +1008,169 @@ func (a *API) createCheckoutSession(w http.ResponseWriter, r *http.Request) {
writeStatusJSON(w, http.StatusCreated, result)
}
+func (a *API) createPortalSession(w http.ResponseWriter, r *http.Request) {
+ tenantID := a.actor(r).TenantID
+ if tenantID == "" {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r))
+ return
+ }
+ result, err := a.billing.CreatePortalSession(r.Context(), tenantID)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) retryCheckoutSession(w http.ResponseWriter, r *http.Request) {
+ actor := a.actor(r)
+ if actor.TenantID == "" {
+ a.scopeError(w, r)
+ return
+ }
+ result, err := a.billing.RetryCheckout(r.Context(), actor.TenantID, r.PathValue("id"), actor.Email)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) resolveMissingTopUp(w http.ResponseWriter, r *http.Request) {
+ var input billing.ResolveMissingTopUpInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ tenantID := r.URL.Query().Get("tenant_id")
+ if actorTenant := a.actor(r).TenantID; actorTenant != "" {
+ tenantID = actorTenant
+ }
+ if tenantID == "" {
+ a.scopeError(w, r)
+ return
+ }
+ actor := a.actor(r)
+ actorType := "console_user"
+ if actor.Bootstrap {
+ actorType = "bootstrap"
+ }
+ result, err := a.billing.ResolveMissingTopUp(r.Context(), tenantID, r.PathValue("id"), input,
+ billing.ResolutionActor{ID: actor.ID, Type: actorType})
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) reverseMissingTopUpCredit(w http.ResponseWriter, r *http.Request) {
+ var input billing.ResolveMissingTopUpInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ tenantID := r.URL.Query().Get("tenant_id")
+ actor := a.actor(r)
+ if actor.TenantID != "" {
+ tenantID = actor.TenantID
+ }
+ if tenantID == "" {
+ a.scopeError(w, r)
+ return
+ }
+ actorType := "console_user"
+ if actor.Bootstrap {
+ actorType = "bootstrap"
+ }
+ result, err := a.billing.ReverseMissingTopUpCredit(r.Context(), tenantID, r.PathValue("id"), input,
+ billing.ResolutionActor{ID: actor.ID, Type: actorType})
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) createRefund(w http.ResponseWriter, r *http.Request) {
+ var input billing.RefundInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ tenantID := r.URL.Query().Get("tenant_id")
+ if actorTenant := a.actor(r).TenantID; actorTenant != "" {
+ tenantID = actorTenant
+ }
+ if tenantID == "" {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "tenant_id is required"}, requestID(r))
+ return
+ }
+ result, err := a.billing.CreateRefund(r.Context(), tenantID, r.PathValue("id"), input)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeStatusJSON(w, http.StatusAccepted, result)
+}
+
+func (a *API) listRefunds(w http.ResponseWriter, r *http.Request) {
+ tenantID := a.actor(r).TenantID
+ if tenantID == "" {
+ tenantID = r.URL.Query().Get("tenant_id")
+ }
+ result, err := a.billing.ListRefunds(r.Context(), tenantID, 200)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) listDisputes(w http.ResponseWriter, r *http.Request) {
+ tenantID := a.actor(r).TenantID
+ if tenantID == "" {
+ tenantID = r.URL.Query().Get("tenant_id")
+ }
+ result, err := a.billing.ListDisputes(r.Context(), tenantID, 200)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) listInvoices(w http.ResponseWriter, r *http.Request) {
+ tenantID := a.actor(r).TenantID
+ if tenantID == "" {
+ tenantID = r.URL.Query().Get("tenant_id")
+ }
+ result, err := a.billing.ListInvoices(r.Context(), tenantID, 200)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) reconcileBilling(w http.ResponseWriter, r *http.Request) {
+ result, err := a.billing.Reconcile(r.Context(), 200)
+ if err != nil {
+ a.billingError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) exportBillingCSV(w http.ResponseWriter, r *http.Request) {
+ tenantID := a.actor(r).TenantID
+ if tenantID == "" {
+ tenantID = r.URL.Query().Get("tenant_id")
+ }
+ w.Header().Set("Content-Type", "text/csv; charset=utf-8")
+ w.Header().Set("Content-Disposition", `attachment; filename="aigw-financial-ledger.csv"`)
+ if err := a.billing.WriteFinancialCSV(r.Context(), tenantID, w); err != nil {
+ a.logger.Error("billing_export_failed", "error", err)
+ }
+}
+
func (a *API) listTenants(w http.ResponseWriter, r *http.Request) {
result, err := a.store.ListTenantsFor(r.Context(), a.actor(r).TenantID)
if err != nil {
@@ -1185,6 +1356,23 @@ func (a *API) toggleModel(w http.ResponseWriter, r *http.Request) {
writeJSON(w, map[string]any{"id": id, "enabled": input.Enabled})
}
+func (a *API) createModelPriceVersion(w http.ResponseWriter, r *http.Request) {
+ var input controlplane.CreatePriceVersionInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ id := r.PathValue("id")
+ generation, err := a.store.CreateModelPriceVersion(r.Context(), id, input)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "model_price", id) {
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, map[string]any{"model_id": id, "generation": generation})
+}
+
func (a *API) reload(w http.ResponseWriter, r *http.Request) {
generation, err := a.manager.Reload(r.Context())
if err != nil {
@@ -1375,6 +1563,10 @@ func (a *API) billingError(w http.ResponseWriter, r *http.Request, err error) {
status = http.StatusNotFound
typeName = "topup_order_not_found"
message = "Top-up order was not found"
+ case errors.Is(err, billing.ErrCannotResolveTopUp):
+ status = http.StatusConflict
+ typeName = "topup_order_not_resolvable"
+ message = err.Error()
default:
a.logger.Error("admin_billing_error", "error", err)
}
diff --git a/internal/adminui/assets/app.js b/internal/adminui/assets/app.js
index 1a965fc..8d29e23 100644
--- a/internal/adminui/assets/app.js
+++ b/internal/adminui/assets/app.js
@@ -1,7 +1,7 @@
const state = {
token: '', csrf: '', actor: {}, permissions: new Set(), overview: {},
tenants: [], projects: [], keys: [], providers: [], models: [], billingAccounts: [], ledger: [],
- usage: [], usageSummary: [], limits: [], users: [], audit: [], orders: [], sessions: [],
+ usage: [], usageSummary: [], limits: [], users: [], audit: [], orders: [], refunds: [], disputes: [], invoices: [], sessions: [],
mfa: {totp_enabled:false,passkeys:[]}, pendingMFA: null, authConfig: {}
};
const $ = (selector) => document.querySelector(selector);
@@ -102,9 +102,12 @@ async function loadAll(knownSession = null) {
state.overview.billing_enabled ? permitted('billing.read','/billing/ledger') : [],
state.actor.id ? api('/auth/mfa') : {totp_enabled:false,passkeys:[]},
state.actor.id ? api('/auth/sessions') : [],
- state.actor.tenant_id && state.overview.billing_enabled && can('billing.read') ? api('/billing/orders') : []
+ state.overview.billing_enabled && can('billing.read') ? api('/billing/orders') : [],
+ state.overview.billing_enabled ? permitted('billing.read','/billing/refunds') : [],
+ state.overview.billing_enabled ? permitted('billing.read','/billing/disputes') : [],
+ state.overview.billing_enabled ? permitted('billing.read','/billing/invoices') : []
]);
- [state.tenants,state.projects,state.keys,state.providers,state.models,state.usage,state.usageSummary,state.limits,state.users,state.audit,state.billingAccounts,state.ledger,state.mfa,state.sessions,state.orders] = results;
+ [state.tenants,state.projects,state.keys,state.providers,state.models,state.usage,state.usageSummary,state.limits,state.users,state.audit,state.billingAccounts,state.ledger,state.mfa,state.sessions,state.orders,state.refunds,state.disputes,state.invoices] = results;
renderAll(); setConnected(true); return true;
} catch (error) { setConnected(false); if (error.status !== 401) toast(error.message, true); return false; }
}
@@ -137,11 +140,14 @@ function renderProjects() { $('#projects-body').innerHTML = state.projects.map(i
function renderKeyProjects() { const tenant = $('#key-tenant').value; const projects = state.projects.filter(item => !tenant || item.tenant_id === tenant); $('#key-project').innerHTML = selectOptions(projects,'id','name'); }
function renderKeys() { $('#keys-body').innerHTML = state.keys.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><code>${esc(item.key_prefix)}</code></td><td><code>${shortID(item.project_id)}</code></td><td>${(item.scopes||[]).map(scope=>`<span class="tag">${esc(scope)}</span>`).join('')}</td><td><span class="badge ${item.status}">${esc(item.status)}</span></td><td>${item.status==='active'&&can('keys.write')?`<button class="text-button danger" data-revoke-key="${esc(item.id)}">Revoke</button>`:''}</td></tr>`).join('') || emptyRow(6); }
function renderProviders() { $('#providers-body').innerHTML = state.providers.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><span class="tag">${esc(item.protocol)}</span></td><td class="truncate">${esc(item.base_url)}</td><td>${integer(item.route_count)}</td><td><span class="badge ${item.enabled?'active':'suspended'}">${item.enabled?'enabled':'disabled'}</span></td><td>${can('platform.write')?`<button class="text-button" data-toggle-provider="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled?'Disable':'Enable'}</button>`:''}</td></tr>`).join('') || emptyRow(6); }
-function renderModels() { $('#models-body').innerHTML = state.models.map(item => `<tr><td><strong>${esc(item.public_id)}</strong></td><td>${esc(item.owned_by||'—')}<small class="price-line">in ${money(item.input_price_micros_per_million)}/1M · out ${money(item.output_price_micros_per_million)}/1M</small></td><td><div class="route-list">${(item.routes||[]).map(route=>`<span>${esc(route.provider_name||route.provider_id).slice(0,24)} → ${esc(route.upstream_model)} <em>p${route.priority} / w${route.weight}</em></span>`).join('')}</div></td><td><span class="badge ${item.enabled?'active':'suspended'}">${item.enabled?'enabled':'disabled'}</span></td><td>${can('platform.write')?`<button class="text-button" data-toggle-model="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled?'Disable':'Enable'}</button>`:''}</td></tr>`).join('') || emptyRow(5); }
+function renderModels() { $('#models-body').innerHTML = state.models.map(item => `<tr><td><strong>${esc(item.public_id)}</strong><small class="price-line">${esc(item.display_name||'')} · ${integer(item.context_window)} ctx · ${esc((item.input_modalities||[]).join('+'))} → ${esc((item.output_modalities||[]).join('+'))}</small></td><td>${esc(item.owned_by||'—')}<small class="price-line">v${item.price_version||1} ${esc(item.price_currency||'usd')} · in ${money(item.input_price_micros_per_million,item.price_currency)}/1M · out ${money(item.output_price_micros_per_million,item.price_currency)}/1M</small></td><td><div class="route-list">${(item.routes||[]).map(route=>`<span>${esc(route.provider_name||route.provider_id).slice(0,24)} → ${esc(route.upstream_model)} <em>p${route.priority} / w${route.weight}</em></span>`).join('')}</div></td><td><span class="badge ${item.enabled&&item.lifecycle!=='retired'?'active':'suspended'}">${esc(item.lifecycle||'active')}</span></td><td>${can('platform.write')?`<button class="text-button" data-toggle-model="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled?'Disable':'Enable'}</button>`:''}</td></tr>`).join('') || emptyRow(5); }
function renderBilling() {
$('#billing-currency').textContent=(state.overview.billing_currency||'').toUpperCase();
$('#billing-accounts-body').innerHTML=state.billingAccounts.map(item=>`<tr><td><strong>${esc(item.tenant_name)}</strong><br><code>${shortID(item.tenant_id)}</code></td><td>${money(item.balance_micros,item.currency)}</td><td>${money(item.reserved_micros,item.currency)}</td><td><strong>${money(item.available_micros,item.currency)}</strong></td><td>${date(item.updated_at)}</td></tr>`).join('')||emptyRow(5);
$('#billing-ledger-body').innerHTML=state.ledger.map(item=>`<tr><td>${date(item.created_at)}</td><td><code>${shortID(item.tenant_id)}</code></td><td><span class="tag">${esc(item.kind)}</span></td><td class="${item.amount_micros>=0?'money-positive':'money-negative'}">${money(item.amount_micros,item.currency)}</td><td>${money(item.balance_after_micros,item.currency)}</td><td title="${esc(item.description)}"><code>${shortID(item.source_id)}</code></td></tr>`).join('')||emptyRow(6);
+ const orders=state.orders||[];$('#billing-orders-body').innerHTML=orders.map(item=>`<tr><td>${date(item.created_at)}</td><td>${money(item.amount_micros,item.currency)}</td><td><span class="badge ${item.status==='paid'?'active':'suspended'}">${esc(item.status)}</span></td><td title="${esc(item.reconciliation_error||'')}"><span class="badge ${['ok','repaired','resolved'].includes(item.reconciliation_status)?'active':'suspended'}">${esc(item.reconciliation_status||'unknown')}</span></td><td>${item.invoice_url?`<a href="${esc(item.invoice_url)}" target="_blank" rel="noopener">Invoice</a>`:''} ${item.invoice_pdf_url?`<a href="${esc(item.invoice_pdf_url)}" target="_blank" rel="noopener">PDF</a>`:''} ${item.receipt_url?`<a href="${esc(item.receipt_url)}" target="_blank" rel="noopener">Receipt</a>`:''}</td><td>${['failed','expired'].includes(item.status)?`<button class="text-button" data-retry-order="${esc(item.id)}">Retry</button>`:''}${can('billing.adjust')&&item.status==='pending'&&item.reconciliation_status==='missing'?` <button class="text-button danger" data-resolve-order="${esc(item.id)}" data-order-tenant="${esc(item.tenant_id)}">Resolve</button>`:''}${can('billing.adjust')&&item.status==='paid'&&item.reconciliation_status==='missing'&&!item.stripe_payment_intent_id?` <button class="text-button danger" data-reverse-order="${esc(item.id)}" data-order-tenant="${esc(item.tenant_id)}">Reverse</button>`:''}${can('billing.adjust')&&['paid','partially_refunded'].includes(item.status)?` <button class="text-button danger" data-refund-order="${esc(item.id)}" data-order-amount="${item.amount_minor}">Refund</button>`:''}</td></tr>`).join('')||emptyRow(6);
+ $('#refunds-body').innerHTML=(state.refunds||[]).map(item=>`<tr><td>${date(item.created_at)}</td><td><code>${shortID(item.topup_order_id)}</code></td><td>${money(item.amount_micros,item.currency)}</td><td><span class="badge ${item.status==='succeeded'?'active':'suspended'}">${esc(item.status)}</span></td><td>${esc(item.last_error||'—')}</td></tr>`).join('')||emptyRow(5);
+ $('#disputes-body').innerHTML=(state.disputes||[]).map(item=>`<tr><td>${date(item.updated_at)}</td><td>${money(item.amount_micros,item.currency)}</td><td><span class="badge ${item.status==='won'?'active':'suspended'}">${esc(item.status)}</span></td><td>${esc(item.reason)}</td><td>${date(item.due_by)}</td></tr>`).join('')||emptyRow(5);
}
function renderSecurity() {
const totp=Boolean(state.mfa?.totp_enabled);
@@ -182,6 +188,13 @@ document.addEventListener('click',async(event)=>{
const session=event.target.closest('[data-revoke-session]');if(session&&confirm('Sign out this device?')){try{const result=await api(`/auth/sessions/${session.dataset.revokeSession}/revoke`,{method:'POST',body:'{}'});if(result.current){state.csrf='';state.actor={};setConnected(false);}else{await loadAll();toast('Device signed out');}}catch(error){toast(error.message,true);}}
const passkey=event.target.closest('[data-delete-passkey]');if(passkey&&confirm('Delete this passkey?')){const password=$('#passkey-form [name=current_password]').value;if(!password){toast('Enter your current password in the Passkey panel',true);}else{try{await api(`/auth/passkeys/${passkey.dataset.deletePasskey}/delete`,{method:'POST',body:JSON.stringify({current_password:password})});await loadAll();toast('Passkey deleted');}catch(error){toast(error.message,true);}}}
const save=event.target.closest('.save-limit');if(save){const row=save.closest('[data-limit-project]');try{await api(`/limits/${row.dataset.limitProject}`,{method:'POST',body:JSON.stringify({requests_per_minute:Number(row.querySelector('.limit-rpm').value),tokens_per_minute:Number(row.querySelector('.limit-tpm').value),concurrent_requests:Number(row.querySelector('.limit-concurrency').value),monthly_spend_micros:decimalToScaled(row.querySelector('.limit-spend').value,6)})});await loadAll();toast('Project limits updated');}catch(error){toast(error.message,true);}}
+ if(event.target.id==='billing-portal'){try{const result=await api('/billing/portal-sessions',{method:'POST',body:'{}'});window.location.assign(result.url);}catch(error){toast(error.message,true);}}
+ if(event.target.id==='billing-export'){try{const headers={};if(state.token)headers.Authorization=`Bearer ${state.token}`;const response=await fetch('./api/billing/export.csv',{credentials:'same-origin',headers});if(!response.ok){const payload=await response.json().catch(()=>({}));throw new Error(payload?.error?.message||`Export failed (${response.status})`);}const blob=await response.blob();const url=URL.createObjectURL(blob);const link=document.createElement('a');link.href=url;link.download=`aigw-ledger-${new Date().toISOString().slice(0,10)}.csv`;document.body.appendChild(link);link.click();link.remove();URL.revokeObjectURL(url);toast('Financial export downloaded');}catch(error){toast(error.message,true);}}
+ if(event.target.id==='billing-reconcile'){try{const result=await api('/billing/reconcile',{method:'POST',body:'{}'});toast(`Reconciliation ${result.status}: ${result.mismatch_count} mismatches`,result.mismatch_count>0);await loadAll();}catch(error){toast(error.message,true);}}
+ const retryOrder=event.target.closest('[data-retry-order]');if(retryOrder){try{const result=await api(`/billing/orders/${retryOrder.dataset.retryOrder}/retry`,{method:'POST',body:'{}'});window.location.assign(result.url);}catch(error){toast(error.message,true);}}
+ const resolveOrder=event.target.closest('[data-resolve-order]');if(resolveOrder){const reason=prompt('Resolution reason for this missing Stripe Session');if(reason){try{await api(`/billing/orders/${resolveOrder.dataset.resolveOrder}/resolve-missing?tenant_id=${encodeURIComponent(resolveOrder.dataset.orderTenant)}`,{method:'POST',body:JSON.stringify({reason})});await loadAll();toast('Missing top-up resolved');}catch(error){toast(error.message,true);}}}
+ const reverseOrder=event.target.closest('[data-reverse-order]');if(reverseOrder){const reason=prompt('Reason for reversing this unverified local credit');if(reason&&confirm('Post an equal negative ledger entry for this credit?')){try{await api(`/billing/orders/${reverseOrder.dataset.reverseOrder}/reverse-missing-credit?tenant_id=${encodeURIComponent(reverseOrder.dataset.orderTenant)}`,{method:'POST',body:JSON.stringify({reason})});await loadAll();toast('Unverified credit reversed');}catch(error){toast(error.message,true);}}}
+ const refundOrder=event.target.closest('[data-refund-order]');if(refundOrder){const amount=prompt('Refund amount in account currency');if(amount){try{const digits=currencyDigits(state.overview.billing_currency||'usd');await api(`/billing/orders/${refundOrder.dataset.refundOrder}/refund?tenant_id=${encodeURIComponent(state.orders.find(item=>item.id===refundOrder.dataset.refundOrder)?.tenant_id||'')}`,{method:'POST',body:JSON.stringify({amount_minor:decimalToScaled(amount,digits),reason:'requested_by_customer'})});await loadAll();toast('Refund queued');}catch(error){toast(error.message,true);}}}
});
$('#key-tenant').addEventListener('change',renderKeyProjects);
@@ -200,7 +213,7 @@ $('#tenant-form').addEventListener('submit',async(event)=>{event.preventDefault(
$('#project-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/projects',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Project created');}catch(error){toast(error.message,true);}});
$('#key-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.scopes=data.scopes.split(',').map(value=>value.trim()).filter(Boolean);const result=await api('/keys',{method:'POST',body:JSON.stringify(data)});event.target.reset();showSecret('API key created',result.key);await loadAll();}catch(error){toast(error.message,true);}});
$('#provider-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/providers',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Provider added');}catch(error){toast(error.message,true);}});
-$('#model-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.input_price_micros_per_million=decimalToScaled(data.input_price,6);data.output_price_micros_per_million=decimalToScaled(data.output_price,6);data.cache_read_price_micros_per_million=decimalToScaled(data.cache_read_price,6);data.cache_write_price_micros_per_million=decimalToScaled(data.cache_write_price,6);delete data.input_price;delete data.output_price;delete data.cache_read_price;delete data.cache_write_price;data.routes=$$('.route-row').map(row=>({provider_id:row.querySelector('.route-provider').value,upstream_model:row.querySelector('.route-upstream').value,priority:Number(row.querySelector('.route-priority').value),weight:Number(row.querySelector('.route-weight').value)}));await api('/models',{method:'POST',body:JSON.stringify(data)});event.target.reset();$('#route-editor').innerHTML='';renderRouteEditor();await loadAll();toast('Model created');}catch(error){toast(error.message,true);}});
+$('#model-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.input_price_micros_per_million=decimalToScaled(data.input_price,6);data.output_price_micros_per_million=decimalToScaled(data.output_price,6);data.cache_read_price_micros_per_million=decimalToScaled(data.cache_read_price,6);data.cache_write_price_micros_per_million=decimalToScaled(data.cache_write_price,6);delete data.input_price;delete data.output_price;delete data.cache_read_price;delete data.cache_write_price;for(const field of ['capabilities','input_modalities','output_modalities','regions','aliases','allowed_tenant_ids','allowed_key_ids'])data[field]=String(data[field]||'').split(',').map(value=>value.trim()).filter(Boolean);data.context_window=Number(data.context_window||0);data.max_output_tokens=Number(data.max_output_tokens||0);data.price_currency=state.overview.billing_currency||'usd';data.routes=$$('.route-row').map(row=>({provider_id:row.querySelector('.route-provider').value,upstream_model:row.querySelector('.route-upstream').value,priority:Number(row.querySelector('.route-priority').value),weight:Number(row.querySelector('.route-weight').value)}));await api('/models',{method:'POST',body:JSON.stringify(data)});event.target.reset();$('#route-editor').innerHTML='';renderRouteEditor();await loadAll();toast('Model created');}catch(error){toast(error.message,true);}});
$('#topup-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);const digits=currencyDigits(state.overview.billing_currency||'usd');const result=await api('/billing/checkout-sessions',{method:'POST',body:JSON.stringify({tenant_id:data.tenant_id,amount_minor:decimalToScaled(data.amount,digits)})});window.location.assign(result.url);}catch(error){toast(error.message,true);}});
$('#adjustment-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/billing/adjustments',{method:'POST',body:JSON.stringify({tenant_id:data.tenant_id,amount_micros:decimalToScaled(data.amount,6),description:data.description})});event.target.reset();await loadAll();toast('Balance adjusted');}catch(error){toast(error.message,true);}});
$('#user-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/users',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Invitation sent');}catch(error){toast(error.message,true);}});
diff --git a/internal/adminui/assets/index.html b/internal/adminui/assets/index.html
index 95ab153..04b4e87 100644
--- a/internal/adminui/assets/index.html
+++ b/internal/adminui/assets/index.html
@@ -46,14 +46,14 @@
<form class="auth-pane" id="reset-complete-pane">
<div><span class="eyebrow">ACCOUNT RECOVERY</span><h1>Choose a new password</h1></div>
<input name="token" id="reset-token" type="hidden">
- <input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1">
+ <input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1">
<label>New password<input name="new_password" type="password" minlength="12" maxlength="128" required autocomplete="new-password"></label>
<button class="button primary" type="submit">Update password</button>
</form>
<form class="auth-pane" id="invite-pane">
<div><span class="eyebrow">TEAM ACCESS</span><h1>Accept invitation</h1></div>
<input name="token" id="invite-token" type="hidden">
- <input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1">
+ <input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1">
<label>Name<input name="display_name" required autocomplete="name"></label>
<label>Password<input name="password" type="password" minlength="12" maxlength="128" required autocomplete="new-password"></label>
<button class="button primary" type="submit">Join workspace</button>
@@ -124,18 +124,18 @@
<section id="providers" class="section">
<div class="section-heading"><div><span class="eyebrow">UPSTREAMS</span><h1>Providers</h1></div></div>
- <form class="panel form-grid" id="provider-form" data-permission="platform.write"><input class="visually-hidden" autocomplete="username" value="aigw-provider" aria-hidden="true" tabindex="-1"><label>Name<input name="name" required placeholder="openai-primary"></label><label>Protocol<select name="protocol"><option value="openai">OpenAI</option><option value="anthropic">Anthropic</option></select></label><label>Base URL<input name="base_url" type="url" required placeholder="https://api.example.com/v1"></label><label>API key<input name="api_key" type="password" required autocomplete="new-password" placeholder="Stored encrypted"></label><button class="button primary" type="submit">Add provider</button></form>
+ <form class="panel form-grid" id="provider-form" data-permission="platform.write"><input class="visually-hidden" name="username" autocomplete="username" value="aigw-provider" aria-hidden="true" tabindex="-1"><label>Name<input name="name" autocomplete="off" required placeholder="openai-primary"></label><label>Protocol<select name="protocol"><option value="openai">OpenAI</option><option value="anthropic">Anthropic</option></select></label><label>Base URL<input name="base_url" type="url" autocomplete="off" required placeholder="https://api.example.com/v1"></label><label>API key<input name="api_key" type="password" required autocomplete="new-password" placeholder="Stored encrypted"></label><button class="button primary" type="submit">Add provider</button></form>
<div class="panel table-wrap"><table><thead><tr><th>Name</th><th>Protocol</th><th>Base URL</th><th>Routes</th><th>Status</th><th></th></tr></thead><tbody id="providers-body"></tbody></table></div>
</section>
<section id="models" class="section">
<div class="section-heading"><div><span class="eyebrow">ROUTING</span><h1>Models & routes</h1></div></div>
- <form class="panel form-grid" id="model-form" data-permission="platform.write"><label>Public model ID<input name="public_id" required placeholder="openai/gpt-4.1-mini"></label><label>Owned by<input name="owned_by" placeholder="openai"></label><label>Input price / 1M<input name="input_price" inputmode="decimal" value="0" required></label><label>Output price / 1M<input name="output_price" inputmode="decimal" value="0" required></label><label>Cache read / 1M<input name="cache_read_price" inputmode="decimal" value="0" required></label><label>Cache write / 1M<input name="cache_write_price" inputmode="decimal" value="0" required></label><div class="route-editor" id="route-editor"></div><button class="button subtle" type="button" id="add-route">Add route</button><button class="button primary" type="submit">Create model</button></form>
+ <form class="panel form-grid" id="model-form" data-permission="platform.write"><label>Public model ID<input name="public_id" required placeholder="openai/gpt-4.1-mini"></label><label>Display name<input name="display_name" placeholder="GPT 4.1 mini"></label><label>Owned by<input name="owned_by" placeholder="openai"></label><label>Lifecycle<select name="lifecycle"><option value="preview">Preview</option><option value="active" selected>Active</option><option value="deprecated">Deprecated</option><option value="retired">Retired</option></select></label><label>Description<input name="description" placeholder="Fast general-purpose text model"></label><label>Context window<input name="context_window" type="number" min="0" value="0"></label><label>Max output tokens<input name="max_output_tokens" type="number" min="0" value="0"></label><label>Capabilities<input name="capabilities" value="chat,streaming" placeholder="chat,streaming,tools,json"></label><label>Input modalities<input name="input_modalities" value="text" placeholder="text,image"></label><label>Output modalities<input name="output_modalities" value="text" placeholder="text,image"></label><label>Regions<input name="regions" placeholder="us-east,nz"></label><label>Deprecated aliases<input name="aliases" placeholder="old/model-id"></label><label>Tenant allowlist IDs<input name="allowed_tenant_ids" placeholder="UUIDs, blank is public"></label><label>API key allowlist IDs<input name="allowed_key_ids" placeholder="UUIDs, blank is unrestricted"></label><label>Input price / 1M<input name="input_price" inputmode="decimal" value="0" required></label><label>Output price / 1M<input name="output_price" inputmode="decimal" value="0" required></label><label>Cache read / 1M<input name="cache_read_price" inputmode="decimal" value="0" required></label><label>Cache write / 1M<input name="cache_write_price" inputmode="decimal" value="0" required></label><div class="route-editor" id="route-editor"></div><button class="button subtle" type="button" id="add-route">Add route</button><button class="button primary" type="submit">Create model</button></form>
<div class="panel table-wrap"><table><thead><tr><th>Public ID</th><th>Owner</th><th>Routes</th><th>Status</th><th></th></tr></thead><tbody id="models-body"></tbody></table></div>
</section>
<section id="billing" class="section">
- <div class="section-heading"><div><span class="eyebrow">REVENUE</span><h1>Balances & ledger</h1></div><span class="currency-label" id="billing-currency"></span></div>
+ <div class="section-heading"><div><span class="eyebrow">REVENUE</span><h1>Balances & ledger</h1></div><div class="form-actions"><button class="button secondary" id="billing-portal" data-permission="billing.topup">Customer portal</button><button class="button subtle" id="billing-export" data-permission="billing.read">Export CSV</button><button class="button subtle" id="billing-reconcile" data-permission="billing.adjust">Reconcile Stripe</button><span class="currency-label" id="billing-currency"></span></div></div>
<div class="billing-actions">
<form class="panel form-grid compact-form" id="topup-form" data-permission="billing.topup"><label>Tenant<select name="tenant_id" id="topup-tenant" required></select></label><label>Amount<input name="amount" inputmode="decimal" min="0" required placeholder="25.00"></label><button class="button primary" type="submit">Open Stripe Checkout</button></form>
<form class="panel form-grid compact-form" id="adjustment-form" data-permission="billing.adjust"><label>Tenant<select name="tenant_id" id="adjustment-tenant" required></select></label><label>Signed amount<input name="amount" inputmode="decimal" required placeholder="10.00 or -5.00"></label><label>Reference<input name="description" maxlength="240" placeholder="Support credit"></label><button class="button secondary" type="submit">Post adjustment</button></form>
@@ -143,6 +143,10 @@
<div class="panel table-wrap"><table><thead><tr><th>Tenant</th><th>Balance</th><th>Reserved</th><th>Available</th><th>Updated</th></tr></thead><tbody id="billing-accounts-body"></tbody></table></div>
<div class="section-heading ledger-heading"><div><span class="eyebrow">AUDIT</span><h2>Recent ledger entries</h2></div></div>
<div class="panel table-wrap"><table><thead><tr><th>Time</th><th>Tenant</th><th>Kind</th><th>Amount</th><th>Balance after</th><th>Reference</th></tr></thead><tbody id="billing-ledger-body"></tbody></table></div>
+ <div class="section-heading ledger-heading"><div><span class="eyebrow">PAYMENTS</span><h2>Top-ups, invoices & refunds</h2></div></div>
+ <div class="panel table-wrap"><table><thead><tr><th>Created</th><th>Amount</th><th>Status</th><th>Reconciliation</th><th>Documents</th><th>Action</th></tr></thead><tbody id="billing-orders-body"></tbody></table></div>
+ <div class="panel table-wrap"><table><thead><tr><th>Created</th><th>Order</th><th>Amount</th><th>Status</th><th>Failure</th></tr></thead><tbody id="refunds-body"></tbody></table></div>
+ <div class="panel table-wrap"><table><thead><tr><th>Updated</th><th>Amount</th><th>Status</th><th>Reason</th><th>Due</th></tr></thead><tbody id="disputes-body"></tbody></table></div>
</section>
<section id="limits" class="section">
@@ -160,9 +164,9 @@
<div class="section-heading"><div><span class="eyebrow">SECURITY</span><h1>Account</h1></div></div>
<form class="panel form-grid compact-form" id="password-form"><input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label>New password<input name="new_password" type="password" required minlength="12" maxlength="128" autocomplete="new-password"></label><button class="button primary" type="submit">Change password</button></form>
<div class="account-grid">
- <form class="panel form-grid compact-form" id="totp-begin-form"><h2>Authenticator app</h2><input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label class="hidden" id="totp-disable-code">Authenticator code<input name="code" inputmode="numeric" autocomplete="one-time-code"></label><button class="button secondary" id="totp-action" type="submit">Set up TOTP</button></form>
+ <form class="panel form-grid compact-form" id="totp-begin-form"><h2>Authenticator app</h2><input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label class="hidden" id="totp-disable-code">Authenticator code<input name="code" inputmode="numeric" autocomplete="one-time-code"></label><button class="button secondary" id="totp-action" type="submit">Set up TOTP</button></form>
<form class="panel form-grid compact-form hidden" id="totp-confirm-form"><h2>Confirm authenticator</h2><img id="totp-qr" alt="TOTP QR code"><code id="totp-secret"></code><label>Code<input name="code" inputmode="numeric" autocomplete="one-time-code" required></label><button class="button primary" type="submit">Enable TOTP</button></form>
- <form class="panel form-grid compact-form" id="passkey-form"><h2>Passkey</h2><input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label>Passkey name<input name="name" value="This device" required></label><button class="button secondary" type="submit">Add passkey</button></form>
+ <form class="panel form-grid compact-form" id="passkey-form"><h2>Passkey</h2><input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label>Passkey name<input name="name" autocomplete="off" value="This device" required></label><button class="button secondary" type="submit">Add passkey</button></form>
</div>
<div class="panel table-wrap"><div class="section-heading"><h2>Devices</h2><button class="button subtle" id="revoke-other-sessions" type="button">Sign out other devices</button></div><table><thead><tr><th>Device</th><th>IP</th><th>Method</th><th>Last seen</th><th></th></tr></thead><tbody id="sessions-body"></tbody></table></div>
<div class="panel table-wrap"><div class="section-heading"><h2>Passkeys</h2></div><table><thead><tr><th>Name</th><th>Created</th><th>Last used</th><th></th></tr></thead><tbody id="passkeys-body"></tbody></table></div>
diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go
index c408cfe..2eb3d87 100644
--- a/internal/billing/ledger.go
+++ b/internal/billing/ledger.go
@@ -77,9 +77,20 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
if limit < 1 || limit > 200 {
limit = 50
}
- rows, err := s.db.Query(ctx, `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
- COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at FROM topup_orders
- WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`, tenantID, limit)
+ query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
+ COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at,
+ COALESCE(stripe_customer_id,''),COALESCE(stripe_payment_intent_id,''),COALESCE(stripe_charge_id,''),
+ COALESCE(stripe_invoice_id,''),COALESCE(invoice_url,''),COALESCE(invoice_pdf_url,''),COALESCE(receipt_url,''),
+ refunded_micros,disputed_micros,reconciliation_status,reconciled_at,reconciliation_error FROM topup_orders`
+ args := []any{}
+ if strings.TrimSpace(tenantID) != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY created_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("query top-up orders: %w", err)
}
@@ -88,7 +99,10 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
for rows.Next() {
var item TopUpOrder
if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status,
- &item.StripeSessionID, &item.CheckoutURL, &item.CreatedAt, &item.PaidAt); err != nil {
+ &item.StripeSessionID, &item.CheckoutURL, &item.CreatedAt, &item.PaidAt, &item.StripeCustomerID,
+ &item.StripePaymentIntentID, &item.StripeChargeID, &item.StripeInvoiceID, &item.InvoiceURL,
+ &item.InvoicePDFURL, &item.ReceiptURL, &item.RefundedMicros, &item.DisputedMicros,
+ &item.ReconciliationStatus, &item.ReconciledAt, &item.ReconciliationError); err != nil {
return nil, err
}
result = append(result, item)
@@ -98,10 +112,21 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in
func (s *Service) GetTopUpOrder(ctx context.Context, tenantID, orderID string) (TopUpOrder, error) {
var result TopUpOrder
- err := s.db.QueryRow(ctx, `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
- COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at FROM topup_orders
- WHERE id=$1 AND tenant_id=$2`, orderID, tenantID).Scan(&result.ID, &result.TenantID, &result.AmountMinor, &result.AmountMicros,
- &result.Currency, &result.Status, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt)
+ query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,
+ COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at,
+ COALESCE(stripe_customer_id,''),COALESCE(stripe_payment_intent_id,''),COALESCE(stripe_charge_id,''),
+ COALESCE(stripe_invoice_id,''),COALESCE(invoice_url,''),COALESCE(invoice_pdf_url,''),COALESCE(receipt_url,''),
+ refunded_micros,disputed_micros,reconciliation_status,reconciled_at,reconciliation_error FROM topup_orders WHERE id=$1`
+ args := []any{orderID}
+ if strings.TrimSpace(tenantID) != "" {
+ query += ` AND tenant_id=$2`
+ args = append(args, tenantID)
+ }
+ err := s.db.QueryRow(ctx, query, args...).Scan(&result.ID, &result.TenantID, &result.AmountMinor, &result.AmountMicros,
+ &result.Currency, &result.Status, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt,
+ &result.StripeCustomerID, &result.StripePaymentIntentID, &result.StripeChargeID, &result.StripeInvoiceID,
+ &result.InvoiceURL, &result.InvoicePDFURL, &result.ReceiptURL, &result.RefundedMicros, &result.DisputedMicros,
+ &result.ReconciliationStatus, &result.ReconciledAt, &result.ReconciliationError)
if errors.Is(err, pgx.ErrNoRows) {
return TopUpOrder{}, ErrTopUpOrderNotFound
}
diff --git a/internal/billing/operations.go b/internal/billing/operations.go
new file mode 100644
index 0000000..a6dc653
--- /dev/null
+++ b/internal/billing/operations.go
@@ -0,0 +1,881 @@
+package billing
+
+import (
+ "context"
+ "encoding/csv"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+ "github.com/stripe/stripe-go/v86"
+)
+
+func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (PortalResult, error) {
+ if !s.stripeEnabled || s.stripeClient == nil {
+ return PortalResult{}, ErrStripeDisabled
+ }
+ var customerID string
+ if err := s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return PortalResult{}, errors.New("no Stripe customer exists for this account")
+ }
+ return PortalResult{}, err
+ }
+ session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{
+ Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL),
+ })
+ if err != nil {
+ return PortalResult{}, fmt.Errorf("create Stripe customer portal session: %w", err)
+ }
+ if session.URL == "" {
+ return PortalResult{}, errors.New("Stripe returned an incomplete portal session")
+ }
+ return PortalResult{URL: session.URL}, nil
+}
+
+func (s *Service) RetryCheckout(ctx context.Context, tenantID, orderID, email string) (CheckoutResult, error) {
+ var amount int64
+ var status string
+ if err := s.db.QueryRow(ctx, `SELECT amount_minor,status FROM topup_orders WHERE id=$1 AND tenant_id=$2`, orderID, tenantID).Scan(&amount, &status); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return CheckoutResult{}, ErrTopUpOrderNotFound
+ }
+ return CheckoutResult{}, err
+ }
+ if status != "failed" && status != "expired" {
+ return CheckoutResult{}, errors.New("only failed or expired top-ups can be retried")
+ }
+ return s.CreateCheckout(ctx, CheckoutInput{TenantID: tenantID, AmountMinor: amount, CustomerEmail: email})
+}
+
+// ResolveMissingTopUp closes an uncredited local order only after reconciliation
+// proved that its Checkout Session does not exist in the configured Stripe account.
+func (s *Service) ResolveMissingTopUp(ctx context.Context, tenantID, orderID string, input ResolveMissingTopUpInput, actor ResolutionActor) (TopUpOrder, error) {
+ reason := normalizeDescription(input.Reason)
+ if reason == "" {
+ return TopUpOrder{}, fmt.Errorf("%w: a resolution reason is required", ErrCannotResolveTopUp)
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return TopUpOrder{}, err
+ }
+ defer tx.Rollback(ctx)
+ var status, reconciliationStatus, paymentIntentID string
+ if err := tx.QueryRow(ctx, `SELECT status,reconciliation_status,COALESCE(stripe_payment_intent_id,'')
+ FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).
+ Scan(&status, &reconciliationStatus, &paymentIntentID); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TopUpOrder{}, ErrTopUpOrderNotFound
+ }
+ return TopUpOrder{}, err
+ }
+ if status != "pending" || reconciliationStatus != "missing" || paymentIntentID != "" {
+ return TopUpOrder{}, fmt.Errorf("%w: only an uncredited pending order confirmed missing by reconciliation can be resolved", ErrCannotResolveTopUp)
+ }
+ var credits int64
+ if err := tx.QueryRow(ctx, `SELECT count(*) FROM billing_ledger
+ WHERE source_type='stripe_checkout' AND source_id=(SELECT stripe_session_id FROM topup_orders WHERE id=$1)`, orderID).Scan(&credits); err != nil {
+ return TopUpOrder{}, err
+ }
+ if credits != 0 {
+ return TopUpOrder{}, fmt.Errorf("%w: a credited top-up cannot be resolved as missing", ErrCannotResolveTopUp)
+ }
+ if actor.Type != "console_user" && actor.Type != "bootstrap" && actor.Type != "maintenance" {
+ return TopUpOrder{}, fmt.Errorf("%w: a valid resolution actor is required", ErrCannotResolveTopUp)
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_reconciliation_resolutions
+ (topup_order_id,tenant_id,actor_id,actor_type,reason) VALUES ($1,$2,$3,$4,$5)`,
+ orderID, tenantID, actor.ID, actor.Type, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_status='resolved',
+ reconciled_at=now(),reconciliation_error=$2 WHERE id=$1`, orderID, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return TopUpOrder{}, err
+ }
+ return s.GetTopUpOrder(ctx, tenantID, orderID)
+}
+
+// ReverseMissingTopUpCredit preserves the original credit and adds an equal
+// negative ledger entry when the configured Stripe account cannot prove the
+// payment. It refuses to consume funds reserved for in-flight requests.
+func (s *Service) ReverseMissingTopUpCredit(ctx context.Context, tenantID, orderID string, input ResolveMissingTopUpInput, actor ResolutionActor) (TopUpOrder, error) {
+ reason := normalizeDescription(input.Reason)
+ if reason == "" || (actor.Type != "console_user" && actor.Type != "bootstrap" && actor.Type != "maintenance") {
+ return TopUpOrder{}, fmt.Errorf("%w: a reason and valid resolution actor are required", ErrCannotResolveTopUp)
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return TopUpOrder{}, err
+ }
+ defer tx.Rollback(ctx)
+ var status, reconciliationStatus, paymentIntentID, currency, sessionID string
+ var amount int64
+ if err := tx.QueryRow(ctx, `SELECT status,reconciliation_status,COALESCE(stripe_payment_intent_id,''),
+ currency,amount_micros,COALESCE(stripe_session_id,'') FROM topup_orders
+ WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).
+ Scan(&status, &reconciliationStatus, &paymentIntentID, &currency, &amount, &sessionID); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TopUpOrder{}, ErrTopUpOrderNotFound
+ }
+ return TopUpOrder{}, err
+ }
+ if status != "paid" || reconciliationStatus != "missing" || paymentIntentID != "" || sessionID == "" {
+ return TopUpOrder{}, fmt.Errorf("%w: only a paid credit confirmed missing with no PaymentIntent can be reversed", ErrCannotResolveTopUp)
+ }
+ var originalCredits int64
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(sum(amount_micros),0) FROM billing_ledger
+ WHERE tenant_id=$1 AND source_type='stripe_checkout' AND source_id=$2`, tenantID, sessionID).Scan(&originalCredits); err != nil {
+ return TopUpOrder{}, err
+ }
+ if originalCredits != amount {
+ return TopUpOrder{}, fmt.Errorf("%w: original Stripe credit does not match the order", ErrCannotResolveTopUp)
+ }
+ var balance, reserved int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil {
+ return TopUpOrder{}, err
+ }
+ if balance-reserved < amount {
+ return TopUpOrder{}, fmt.Errorf("%w: available balance is insufficient to reverse the orphaned credit", ErrCannotResolveTopUp)
+ }
+ newBalance := balance - amount
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger
+ (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description)
+ VALUES ($1,$2,$3,$4,'adjustment','stripe_reconciliation',$5,$6)`,
+ tenantID, currency, -amount, newBalance, orderID, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_reconciliation_resolutions
+ (topup_order_id,tenant_id,actor_id,actor_type,reason) VALUES ($1,$2,$3,$4,$5)`,
+ orderID, tenantID, actor.ID, actor.Type, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='reversed',reconciliation_status='resolved',
+ reconciled_at=now(),reconciliation_error=$2 WHERE id=$1`, orderID, reason); err != nil {
+ return TopUpOrder{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return TopUpOrder{}, err
+ }
+ return s.GetTopUpOrder(ctx, tenantID, orderID)
+}
+
+func (s *Service) CreateRefund(ctx context.Context, tenantID, orderID string, input RefundInput) (Refund, error) {
+ if !s.stripeEnabled {
+ return Refund{}, ErrStripeDisabled
+ }
+ input.Reason = strings.TrimSpace(input.Reason)
+ if input.Reason == "" {
+ input.Reason = "requested_by_customer"
+ }
+ if input.Reason != "requested_by_customer" && input.Reason != "duplicate" && input.Reason != "fraudulent" {
+ return Refund{}, errors.New("refund reason must be requested_by_customer, duplicate, or fraudulent")
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return Refund{}, err
+ }
+ defer tx.Rollback(ctx)
+ var orderTenant, currency, status, paymentIntentID string
+ var orderAmount, refunded int64
+ if err := tx.QueryRow(ctx, `SELECT tenant_id::text,currency,status,COALESCE(stripe_payment_intent_id,''),amount_micros,refunded_micros
+ FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).Scan(&orderTenant, &currency, &status, &paymentIntentID, &orderAmount, &refunded); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return Refund{}, ErrTopUpOrderNotFound
+ }
+ return Refund{}, err
+ }
+ if status != "paid" && status != "partially_refunded" {
+ return Refund{}, errors.New("only a paid top-up can be refunded")
+ }
+ if paymentIntentID == "" {
+ return Refund{}, errors.New("top-up has no Stripe PaymentIntent")
+ }
+ amountMicros, err := minorToMicros(currency, input.AmountMinor)
+ if err != nil {
+ return Refund{}, ErrInvalidAmount
+ }
+ var pendingRefunds int64
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(sum(amount_micros),0) FROM stripe_refunds
+ WHERE topup_order_id=$1 AND status IN ('queued','submitting','pending','requires_action','succeeded')`, orderID).Scan(&pendingRefunds); err != nil {
+ return Refund{}, err
+ }
+ // Legacy successful refunds are already reflected in refunded_micros. New
+ // rows are included in the sum, so use the larger value without double-counting.
+ committedRefunds := refunded
+ if pendingRefunds > committedRefunds {
+ committedRefunds = pendingRefunds
+ }
+ if amountMicros > orderAmount-committedRefunds {
+ return Refund{}, ErrInvalidAmount
+ }
+ var balance, held int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &held); err != nil {
+ return Refund{}, err
+ }
+ available := balance - held
+ if available < 0 {
+ available = 0
+ }
+ refundHold := amountMicros
+ if refundHold > available {
+ refundHold = available
+ }
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+$2,updated_at=now() WHERE tenant_id=$1`, tenantID, refundHold); err != nil {
+ return Refund{}, err
+ }
+ var result Refund
+ err = tx.QueryRow(ctx, `INSERT INTO stripe_refunds (tenant_id,topup_order_id,amount_minor,amount_micros,held_micros,currency,reason)
+ VALUES ($1,$2,$3,$4,$5,$6,$7) RETURNING id::text,tenant_id::text,topup_order_id::text,amount_minor,amount_micros,currency,reason,status,last_error,created_at,completed_at`,
+ tenantID, orderID, input.AmountMinor, amountMicros, refundHold, currency, input.Reason).Scan(&result.ID, &result.TenantID, &result.TopUpOrderID,
+ &result.AmountMinor, &result.AmountMicros, &result.Currency, &result.Reason, &result.Status, &result.LastError, &result.CreatedAt, &result.CompletedAt)
+ if err != nil {
+ return Refund{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Refund{}, err
+ }
+ return result, nil
+}
+
+func (s *Service) ListRefunds(ctx context.Context, tenantID string, limit int) ([]Refund, error) {
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ query := `SELECT id::text,tenant_id::text,topup_order_id::text,COALESCE(stripe_refund_id,''),amount_minor,amount_micros,
+ currency,reason,status,last_error,created_at,completed_at FROM stripe_refunds`
+ args := []any{}
+ if tenantID != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY created_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ result := make([]Refund, 0)
+ for rows.Next() {
+ var item Refund
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.StripeRefundID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Reason, &item.Status, &item.LastError, &item.CreatedAt, &item.CompletedAt); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Service) ListDisputes(ctx context.Context, tenantID string, limit int) ([]PaymentDispute, error) {
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ query := `SELECT stripe_dispute_id,COALESCE(tenant_id::text,''),COALESCE(topup_order_id::text,''),
+ amount_minor,amount_micros,currency,status,reason,debited_micros,uncollected_micros,due_by,updated_at FROM stripe_disputes`
+ args := []any{}
+ if tenantID != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY updated_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY updated_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ result := []PaymentDispute{}
+ for rows.Next() {
+ var item PaymentDispute
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, &item.Reason, &item.DebitedMicros, &item.UncollectedMicros, &item.DueBy, &item.UpdatedAt); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Service) ListInvoices(ctx context.Context, tenantID string, limit int) ([]Invoice, error) {
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ query := `SELECT stripe_invoice_id,COALESCE(tenant_id::text,''),COALESCE(topup_order_id::text,''),status,currency,
+ amount_due_minor,amount_paid_minor,attempt_count,next_payment_attempt,hosted_invoice_url,invoice_pdf_url,last_failure,updated_at FROM stripe_invoices`
+ args := []any{}
+ if tenantID != "" {
+ query += ` WHERE tenant_id=$1 ORDER BY updated_at DESC LIMIT $2`
+ args = []any{tenantID, limit}
+ } else {
+ query += ` ORDER BY updated_at DESC LIMIT $1`
+ args = []any{limit}
+ }
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+ result := []Invoice{}
+ for rows.Next() {
+ var item Invoice
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.Status, &item.Currency, &item.AmountDueMinor, &item.AmountPaidMinor, &item.AttemptCount, &item.NextPaymentAttempt, &item.HostedInvoiceURL, &item.InvoicePDFURL, &item.LastFailure, &item.UpdatedAt); err != nil {
+ return nil, err
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Service) RunStripeOperations(ctx context.Context) {
+ if !s.stripeEnabled || s.stripeClient == nil {
+ return
+ }
+ ticker := time.NewTicker(2 * time.Second)
+ reconcile := time.NewTicker(30 * time.Minute)
+ metrics := time.NewTicker(15 * time.Second)
+ defer ticker.Stop()
+ defer reconcile.Stop()
+ defer metrics.Stop()
+ initialCtx, cancel := context.WithTimeout(ctx, 2*time.Minute)
+ _, _ = s.Reconcile(initialCtx, 100)
+ cancel()
+ s.refreshOperationalMetrics(ctx)
+ for {
+ for i := 0; i < 8; i++ {
+ ok, _ := s.processRefundOperation(ctx)
+ if !ok {
+ break
+ }
+ }
+ _, _ = s.pollPendingRefund(ctx)
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ case <-reconcile.C:
+ _, _ = s.Reconcile(ctx, 100)
+ case <-metrics.C:
+ s.refreshOperationalMetrics(ctx)
+ }
+ }
+}
+
+func (s *Service) refreshOperationalMetrics(ctx context.Context) {
+ if s.metrics == nil {
+ return
+ }
+ status, err := s.OperationalStatus(ctx)
+ if err == nil {
+ s.metrics.SetStripeOperations(status)
+ }
+}
+
+func (s *Service) OperationalStatus(ctx context.Context) (OperationalStatus, error) {
+ status := OperationalStatus{StripeEnabled: s.stripeEnabled, ReconciliationStatus: "disabled"}
+ if err := s.db.QueryRow(ctx, `SELECT
+ count(*) FILTER (WHERE status IN ('queued','submitting','pending','requires_action')),
+ min(created_at) FILTER (WHERE status IN ('queued','submitting','pending','requires_action')),
+ COALESCE((SELECT sum(uncollected_micros) FROM stripe_refunds),0)+
+ COALESCE((SELECT sum(uncollected_micros) FROM stripe_disputes),0)+
+ COALESCE((SELECT sum(uncollected_micros) FROM usage_events),0)
+ FROM stripe_refunds`).Scan(&status.RefundBacklog, &status.OldestRefund, &status.UncollectedMicros); err != nil {
+ return status, err
+ }
+ if err := s.db.QueryRow(ctx, `SELECT count(*),min(created_at) FROM stripe_webhook_events
+ WHERE processed_at IS NULL`).Scan(&status.UnprocessedWebhooks, &status.OldestUnprocessedWebhook); err != nil {
+ return status, err
+ }
+ if err := s.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE metering_status='missing'`).Scan(&status.UnmeteredSuccesses); err != nil {
+ return status, err
+ }
+ if !s.stripeEnabled {
+ return status, nil
+ }
+ err := s.db.QueryRow(ctx, `SELECT status,mismatch_count,completed_at,error FROM billing_reconciliation_runs
+ WHERE status<>'running' ORDER BY started_at DESC LIMIT 1`).Scan(&status.ReconciliationStatus,
+ &status.ReconciliationMismatches, &status.ReconciliationCompletedAt, &status.ReconciliationError)
+ if errors.Is(err, pgx.ErrNoRows) {
+ status.ReconciliationStatus = "never_run"
+ return status, nil
+ }
+ return status, err
+}
+
+func (s *Service) processRefundOperation(ctx context.Context) (bool, error) {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return false, err
+ }
+ defer tx.Rollback(ctx)
+ var id, orderID, paymentIntentID, reason string
+ var amount int64
+ err = tx.QueryRow(ctx, `WITH selected AS (
+ SELECT r.id FROM stripe_refunds r WHERE r.available_at<=now() AND
+ (r.status='queued' OR (r.status='submitting' AND r.updated_at<now()-interval '5 minutes'))
+ ORDER BY r.available_at,r.created_at FOR UPDATE SKIP LOCKED LIMIT 1)
+ UPDATE stripe_refunds r SET status='submitting',attempts=attempts+1,updated_at=now()
+ FROM selected,topup_orders o WHERE r.id=selected.id AND o.id=r.topup_order_id
+ RETURNING r.id::text,r.topup_order_id::text,o.stripe_payment_intent_id,r.amount_minor,r.reason`).Scan(&id, &orderID, &paymentIntentID, &amount, &reason)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, tx.Commit(ctx)
+ }
+ if err != nil {
+ return false, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return false, err
+ }
+ params := &stripe.RefundCreateParams{Amount: stripe.Int64(amount), PaymentIntent: stripe.String(paymentIntentID), Reason: stripe.String(reason),
+ Metadata: map[string]string{"aigw_refund_id": id, "aigw_topup_order_id": orderID}}
+ params.SetIdempotencyKey("aigw_refund_" + id)
+ refund, err := s.stripeClient.V1Refunds.Create(ctx, params)
+ if err != nil {
+ message := err.Error()
+ if len(message) > 1000 {
+ message = message[:1000]
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE stripe_refunds SET status='queued',last_error=$2,
+ available_at=now()+make_interval(secs=>LEAST(1800,power(2,LEAST(attempts,10))::int)),updated_at=now() WHERE id=$1`, id, message)
+ return true, err
+ }
+ return true, s.applyRefund(ctx, refund)
+}
+
+func (s *Service) pollPendingRefund(ctx context.Context) (bool, error) {
+ var id string
+ err := s.db.QueryRow(ctx, `UPDATE stripe_refunds SET available_at=now()+interval '1 minute',updated_at=now()
+ WHERE id=(SELECT id FROM stripe_refunds WHERE status IN ('pending','requires_action')
+ AND stripe_refund_id IS NOT NULL AND available_at<=now() ORDER BY available_at LIMIT 1 FOR UPDATE SKIP LOCKED)
+ RETURNING stripe_refund_id`).Scan(&id)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, nil
+ }
+ if err != nil {
+ return false, err
+ }
+ refund, err := s.stripeClient.V1Refunds.Retrieve(ctx, id, &stripe.RefundRetrieveParams{})
+ if err != nil {
+ return true, err
+ }
+ return true, s.applyRefund(ctx, refund)
+}
+
+func (s *Service) applyRefund(ctx context.Context, refund *stripe.Refund) error {
+ if refund == nil || refund.ID == "" {
+ return ErrInvalidAmount
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ if err := s.applyRefundTx(ctx, tx, refund); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Service) applyRefundTx(ctx context.Context, tx pgx.Tx, refund *stripe.Refund) error {
+ if refund == nil || refund.ID == "" {
+ return ErrInvalidAmount
+ }
+ localID := refund.Metadata["aigw_refund_id"]
+ var id, tenantID, orderID, currentStatus string
+ var amountMicros, held int64
+ query := `SELECT id::text,tenant_id::text,topup_order_id::text,amount_micros,held_micros,status FROM stripe_refunds WHERE `
+ arg := refund.ID
+ if localID != "" {
+ query += `id=$1 FOR UPDATE`
+ arg = localID
+ } else {
+ query += `stripe_refund_id=$1 FOR UPDATE`
+ }
+ err := tx.QueryRow(ctx, query, arg).Scan(&id, &tenantID, &orderID, &amountMicros, &held, &currentStatus)
+ if errors.Is(err, pgx.ErrNoRows) {
+ paymentIntentID := ""
+ if refund.PaymentIntent != nil {
+ paymentIntentID = refund.PaymentIntent.ID
+ }
+ if paymentIntentID == "" {
+ return ErrInvalidAmount
+ }
+ var currency string
+ if err := tx.QueryRow(ctx, `SELECT id::text,tenant_id::text,currency FROM topup_orders WHERE stripe_payment_intent_id=$1 FOR UPDATE`, paymentIntentID).Scan(&orderID, &tenantID, &currency); err != nil {
+ return err
+ }
+ amountMicros, err = minorToMicros(currency, refund.Amount)
+ if err != nil {
+ return err
+ }
+ err = tx.QueryRow(ctx, `INSERT INTO stripe_refunds (tenant_id,topup_order_id,stripe_refund_id,amount_minor,amount_micros,currency,reason,status)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8) RETURNING id::text`, tenantID, orderID, refund.ID, refund.Amount, amountMicros, currency, string(refund.Reason), string(refund.Status)).Scan(&id)
+ if err != nil {
+ return err
+ }
+ currentStatus = ""
+ } else if err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET stripe_refund_id=$2,status=$3,last_error=$4,updated_at=now(),
+ completed_at=CASE WHEN $3 IN ('succeeded','failed','canceled') THEN now() ELSE completed_at END WHERE id=$1`, id, refund.ID, string(refund.Status), string(refund.FailureReason)); err != nil {
+ return err
+ }
+ if refund.Status == stripe.RefundStatusSucceeded && currentStatus != "succeeded" {
+ var balance, reserved int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil {
+ return err
+ }
+ if held > 0 {
+ if held > reserved {
+ return errors.New("refund hold invariant violated")
+ }
+ reserved -= held
+ }
+ debit := amountMicros
+ if debit > balance-reserved {
+ debit = balance - reserved
+ }
+ if debit < 0 {
+ debit = 0
+ }
+ newBalance := balance - debit
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,reserved_micros=$3,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance, reserved); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET held_micros=0,uncollected_micros=$2 WHERE id=$1`, id, amountMicros-debit); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description)
+ SELECT $1,currency,$2,$3,'refund','stripe_refund',$4,'Stripe top-up refund' FROM topup_orders WHERE id=$5
+ ON CONFLICT (source_type,source_id) DO NOTHING`, tenantID, -debit, newBalance, refund.ID, orderID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET refunded_micros=LEAST(amount_micros,refunded_micros+$2),
+ status=CASE WHEN refunded_micros+$2>=amount_micros THEN 'refunded' ELSE 'partially_refunded' END WHERE id=$1`, orderID, amountMicros); err != nil {
+ return err
+ }
+ } else if (refund.Status == stripe.RefundStatusFailed || refund.Status == stripe.RefundStatusCanceled) && held > 0 {
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros-$2,updated_at=now() WHERE tenant_id=$1`, tenantID, held); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET held_micros=0 WHERE id=$1`, id); err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+func (s *Service) applyDisputeTx(ctx context.Context, tx pgx.Tx, dispute *stripe.Dispute, eventType stripe.EventType) error {
+ paymentIntentID := ""
+ if dispute.PaymentIntent != nil {
+ paymentIntentID = dispute.PaymentIntent.ID
+ }
+ if paymentIntentID == "" && dispute.Charge != nil && dispute.Charge.PaymentIntent != nil {
+ paymentIntentID = dispute.Charge.PaymentIntent.ID
+ }
+ if paymentIntentID == "" {
+ return ErrInvalidAmount
+ }
+ var orderID, tenantID, currency string
+ if err := tx.QueryRow(ctx, `SELECT id::text,tenant_id::text,currency FROM topup_orders WHERE stripe_payment_intent_id=$1 FOR UPDATE`, paymentIntentID).Scan(&orderID, &tenantID, &currency); err != nil {
+ return err
+ }
+ amountMicros, err := minorToMicros(currency, dispute.Amount)
+ if err != nil {
+ return err
+ }
+ dueBy := (*time.Time)(nil)
+ if dispute.EvidenceDetails != nil && dispute.EvidenceDetails.DueBy > 0 {
+ value := time.Unix(dispute.EvidenceDetails.DueBy, 0).UTC()
+ dueBy = &value
+ }
+ var previousDebited int64
+ queryErr := tx.QueryRow(ctx, `SELECT debited_micros FROM stripe_disputes WHERE stripe_dispute_id=$1 FOR UPDATE`, dispute.ID).Scan(&previousDebited)
+ if queryErr != nil && !errors.Is(queryErr, pgx.ErrNoRows) {
+ return queryErr
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_disputes (stripe_dispute_id,tenant_id,topup_order_id,stripe_payment_intent_id,
+ amount_minor,amount_micros,currency,status,reason,due_by) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10)
+ ON CONFLICT (stripe_dispute_id) DO UPDATE SET status=EXCLUDED.status,reason=EXCLUDED.reason,
+ due_by=EXCLUDED.due_by,updated_at=now(),closed_at=CASE WHEN EXCLUDED.status IN ('won','lost') THEN now() ELSE stripe_disputes.closed_at END`,
+ dispute.ID, tenantID, orderID, paymentIntentID, dispute.Amount, amountMicros, currency, string(dispute.Status), string(dispute.Reason), dueBy); err != nil {
+ return err
+ }
+ shouldDebit := eventType == stripe.EventTypeChargeDisputeCreated || eventType == stripe.EventTypeChargeDisputeFundsWithdrawn
+ shouldReverse := eventType == stripe.EventTypeChargeDisputeFundsReinstated || dispute.Status == stripe.DisputeStatusWon
+ if shouldDebit && previousDebited == 0 {
+ var balance, reserved int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil {
+ return err
+ }
+ debit := amountMicros
+ if debit > balance-reserved {
+ debit = balance - reserved
+ }
+ if debit < 0 {
+ debit = 0
+ }
+ newBalance := balance - debit
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_disputes SET debited_micros=$2,uncollected_micros=$3 WHERE stripe_dispute_id=$1`, dispute.ID, debit, amountMicros-debit); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description)
+ VALUES ($1,$2,$3,$4,'dispute','stripe_dispute',$5,'Stripe payment dispute') ON CONFLICT DO NOTHING`, tenantID, currency, -debit, newBalance, dispute.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET disputed_micros=GREATEST(disputed_micros,$2),status='disputed' WHERE id=$1`, orderID, amountMicros); err != nil {
+ return err
+ }
+ } else if shouldReverse && previousDebited > 0 {
+ var balance int64
+ if err := tx.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance); err != nil {
+ return err
+ }
+ newBalance := balance + previousDebited
+ if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_disputes SET debited_micros=0,uncollected_micros=0 WHERE stripe_dispute_id=$1`, dispute.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description)
+ VALUES ($1,$2,$3,$4,'dispute_reversal','stripe_dispute_reversal',$5,'Stripe dispute funds reinstated') ON CONFLICT DO NOTHING`, tenantID, currency, previousDebited, newBalance, dispute.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET disputed_micros=0,status=CASE WHEN refunded_micros=0 THEN 'paid' WHEN refunded_micros<amount_micros THEN 'partially_refunded' ELSE 'refunded' END WHERE id=$1`, orderID); err != nil {
+ return err
+ }
+ }
+ return nil
+}
+
+func (s *Service) applyInvoiceTx(ctx context.Context, tx pgx.Tx, invoice *stripe.Invoice) error {
+ orderID := invoice.Metadata["aigw_topup_order_id"]
+ tenantID := invoice.Metadata["aigw_tenant_id"]
+ customerID := ""
+ if invoice.Customer != nil {
+ customerID = invoice.Customer.ID
+ }
+ if tenantID == "" && customerID != "" {
+ _ = tx.QueryRow(ctx, `SELECT tenant_id::text FROM stripe_customers WHERE stripe_customer_id=$1`, customerID).Scan(&tenantID)
+ }
+ if orderID == "" {
+ _ = tx.QueryRow(ctx, `SELECT id::text FROM topup_orders WHERE stripe_invoice_id=$1`, invoice.ID).Scan(&orderID)
+ }
+ failure := ""
+ if invoice.LastFinalizationError != nil {
+ failure = invoice.LastFinalizationError.Msg
+ }
+ var nextAttempt *time.Time
+ if invoice.NextPaymentAttempt > 0 {
+ value := time.Unix(invoice.NextPaymentAttempt, 0).UTC()
+ nextAttempt = &value
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_invoices (stripe_invoice_id,tenant_id,topup_order_id,stripe_customer_id,status,currency,
+ amount_due_minor,amount_paid_minor,attempt_count,next_payment_attempt,hosted_invoice_url,invoice_pdf_url,last_failure)
+ VALUES ($1,NULLIF($2,'')::uuid,NULLIF($3,'')::uuid,NULLIF($4,''),$5,$6,$7,$8,$9,$10,$11,$12,$13)
+ ON CONFLICT (stripe_invoice_id) DO UPDATE SET status=EXCLUDED.status,amount_due_minor=EXCLUDED.amount_due_minor,
+ amount_paid_minor=EXCLUDED.amount_paid_minor,attempt_count=EXCLUDED.attempt_count,next_payment_attempt=EXCLUDED.next_payment_attempt,
+ hosted_invoice_url=EXCLUDED.hosted_invoice_url,invoice_pdf_url=EXCLUDED.invoice_pdf_url,last_failure=EXCLUDED.last_failure,updated_at=now()`,
+ invoice.ID, tenantID, orderID, customerID, string(invoice.Status), string(invoice.Currency), invoice.AmountDue, invoice.AmountPaid,
+ invoice.AttemptCount, nextAttempt, invoice.HostedInvoiceURL, invoice.InvoicePDF, failure); err != nil {
+ return err
+ }
+ if orderID != "" {
+ _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_invoice_id=$2,invoice_url=COALESCE(NULLIF($3,''),invoice_url),
+ invoice_pdf_url=COALESCE(NULLIF($4,''),invoice_pdf_url) WHERE id=$1`, orderID, invoice.ID, invoice.HostedInvoiceURL, invoice.InvoicePDF)
+ return err
+ }
+ return nil
+}
+
+func (s *Service) Reconcile(ctx context.Context, limit int) (ReconciliationResult, error) {
+ if !s.stripeEnabled || s.stripeClient == nil {
+ return ReconciliationResult{}, ErrStripeDisabled
+ }
+ if limit < 1 || limit > 500 {
+ limit = 100
+ }
+ var result ReconciliationResult
+ if err := s.db.QueryRow(ctx, `INSERT INTO billing_reconciliation_runs (status) VALUES ('running') RETURNING id::text`).Scan(&result.ID); err != nil {
+ return result, err
+ }
+ rows, err := s.db.Query(ctx, `SELECT id::text,stripe_session_id,status,amount_minor,currency,reconciliation_status FROM topup_orders
+ WHERE stripe_session_id IS NOT NULL ORDER BY created_at DESC LIMIT $1`, limit)
+ if err != nil {
+ return s.failReconciliation(ctx, result, err)
+ }
+ type order struct {
+ id, session, status, currency, reconciliationStatus string
+ amount int64
+ }
+ var orders []order
+ for rows.Next() {
+ var item order
+ if err := rows.Scan(&item.id, &item.session, &item.status, &item.amount, &item.currency, &item.reconciliationStatus); err != nil {
+ rows.Close()
+ return s.failReconciliation(ctx, result, err)
+ }
+ orders = append(orders, item)
+ }
+ rows.Close()
+ for _, item := range orders {
+ session, retrieveErr := s.stripeClient.V1CheckoutSessions.Retrieve(ctx, item.session, &stripe.CheckoutSessionRetrieveParams{})
+ result.CheckedOrders++
+ if retrieveErr != nil {
+ if stripeResourceMissing(retrieveErr) && item.reconciliationStatus == "resolved" && (item.status == "failed" || item.status == "reversed") {
+ continue
+ }
+ message := truncateError(retrieveErr)
+ typeName := "stripe_retrieve_failed"
+ state := "mismatch"
+ if stripeResourceMissing(retrieveErr) {
+ typeName = "stripe_session_missing"
+ state = "missing"
+ }
+ if err := s.updateOrderReconciliation(ctx, item.id, state, message); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record reconciliation retrieval failure for order %s: %w", item.id, err))
+ }
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": typeName, "error": message})
+ continue
+ }
+ expected := item.status == "paid" || item.status == "partially_refunded" || item.status == "refunded" || item.status == "disputed"
+ stripePaid := session.PaymentStatus == stripe.CheckoutSessionPaymentStatusPaid
+ if session.AmountTotal != item.amount || string(session.Currency) != item.currency {
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_mismatch", "local_status": item.status, "stripe_payment_status": session.PaymentStatus, "local_amount": item.amount, "stripe_amount": session.AmountTotal})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "amount or currency mismatch"); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record checkout mismatch for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ if stripePaid && !expected {
+ raw, marshalErr := json.Marshal(session)
+ if marshalErr != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("encode Stripe Checkout Session for order %s: %w", item.id, marshalErr))
+ }
+ if repairErr := s.processStripeEvent(ctx, stripe.Event{ID: "reconcile_" + session.ID + "_paid", Type: stripe.EventTypeCheckoutSessionCompleted, Data: &stripe.EventData{Raw: raw}}); repairErr != nil {
+ message := truncateError(repairErr)
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_repair_failed", "error": message})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record checkout repair failure for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ result.Repairs = append(result.Repairs, map[string]any{"order_id": item.id, "type": "credited_paid_checkout"})
+ if err := s.updateOrderReconciliation(ctx, item.id, "repaired", ""); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record checkout repair for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ if expected != stripePaid {
+ result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_payment_state_mismatch", "local_status": item.status, "stripe_payment_status": session.PaymentStatus})
+ if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "payment state mismatch"); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record payment-state mismatch for order %s: %w", item.id, err))
+ }
+ continue
+ }
+ if err := s.updateOrderReconciliation(ctx, item.id, "ok", ""); err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("record clean reconciliation for order %s: %w", item.id, err))
+ }
+ }
+ result.MismatchCount = int64(len(result.Mismatches))
+ result.Status = "clean"
+ if result.MismatchCount > 0 {
+ result.Status = "mismatch"
+ }
+ report, marshalErr := json.Marshal(result.Mismatches)
+ if marshalErr != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("encode reconciliation report: %w", marshalErr))
+ }
+ _, err = s.db.Exec(ctx, `UPDATE billing_reconciliation_runs SET status=$2,checked_orders=$3,mismatch_count=$4,report=$5,completed_at=now() WHERE id=$1`, result.ID, result.Status, result.CheckedOrders, result.MismatchCount, report)
+ if err != nil {
+ return s.failReconciliation(ctx, result, fmt.Errorf("complete reconciliation run: %w", err))
+ }
+ return result, nil
+}
+
+func (s *Service) updateOrderReconciliation(ctx context.Context, orderID, status, message string) error {
+ command, err := s.db.Exec(ctx, `UPDATE topup_orders SET reconciliation_status=$2,reconciled_at=now(),reconciliation_error=$3 WHERE id=$1`, orderID, status, message)
+ if err != nil {
+ return err
+ }
+ if command.RowsAffected() != 1 {
+ return fmt.Errorf("expected one top-up order, updated %d", command.RowsAffected())
+ }
+ return nil
+}
+
+func stripeResourceMissing(err error) bool {
+ var stripeErr *stripe.Error
+ return errors.As(err, &stripeErr) && (stripeErr.Code == stripe.ErrorCodeResourceMissing || stripeErr.HTTPStatusCode == 404)
+}
+
+func truncateError(err error) string {
+ if err == nil {
+ return ""
+ }
+ message := err.Error()
+ if len(message) > 1000 {
+ return message[:1000]
+ }
+ return message
+}
+
+func (s *Service) failReconciliation(ctx context.Context, result ReconciliationResult, cause error) (ReconciliationResult, error) {
+ if _, err := s.db.Exec(ctx, `UPDATE billing_reconciliation_runs SET status='failed',error=$2,completed_at=now() WHERE id=$1`, result.ID, truncateError(cause)); err != nil {
+ return result, errors.Join(cause, fmt.Errorf("record failed reconciliation run: %w", err))
+ }
+ return result, cause
+}
+
+func (s *Service) WriteFinancialCSV(ctx context.Context, tenantID string, output io.Writer) error {
+ query := `
+ SELECT id::text, tenant_id::text, COALESCE(project_id::text, ''), currency, amount_micros,
+ balance_after_micros, kind, source_type, source_id, description, created_at
+ FROM billing_ledger`
+ args := []any{}
+ if strings.TrimSpace(tenantID) != "" {
+ query += ` WHERE tenant_id=$1`
+ args = append(args, tenantID)
+ }
+ query += ` ORDER BY created_at, id`
+ rows, err := s.db.Query(ctx, query, args...)
+ if err != nil {
+ return err
+ }
+ defer rows.Close()
+ w := csv.NewWriter(output)
+ if err := w.Write([]string{"timestamp", "tenant_id", "project_id", "currency", "kind", "amount_micros", "balance_after_micros", "source_type", "source_id", "description"}); err != nil {
+ return err
+ }
+ for rows.Next() {
+ var e LedgerEntry
+ if err := rows.Scan(&e.ID, &e.TenantID, &e.ProjectID, &e.Currency, &e.AmountMicros,
+ &e.BalanceAfterMicros, &e.Kind, &e.SourceType, &e.SourceID, &e.Description, &e.CreatedAt); err != nil {
+ return err
+ }
+ if err := w.Write([]string{e.CreatedAt.UTC().Format(time.RFC3339Nano), e.TenantID, e.ProjectID, e.Currency, e.Kind, strconv.FormatInt(e.AmountMicros, 10), strconv.FormatInt(e.BalanceAfterMicros, 10), e.SourceType, e.SourceID, e.Description}); err != nil {
+ return err
+ }
+ }
+ if err := rows.Err(); err != nil {
+ return err
+ }
+ w.Flush()
+ return w.Error()
+}
diff --git a/internal/billing/operations_test.go b/internal/billing/operations_test.go
new file mode 100644
index 0000000..46bedeb
--- /dev/null
+++ b/internal/billing/operations_test.go
@@ -0,0 +1,36 @@
+package billing
+
+import (
+ "testing"
+ "time"
+)
+
+func TestOperationalStatusReadiness(t *testing.T) {
+ now := time.Now().UTC()
+ recent := now.Add(-time.Minute)
+ clean := OperationalStatus{
+ StripeEnabled: true, ReconciliationStatus: "clean", ReconciliationCompletedAt: &recent,
+ }
+ if !clean.Ready(now) {
+ t.Fatal("clean operational state should be ready")
+ }
+ oldWebhook := now.Add(-6 * time.Minute)
+ stuckWebhook := clean
+ stuckWebhook.UnprocessedWebhooks = 1
+ stuckWebhook.OldestUnprocessedWebhook = &oldWebhook
+ if stuckWebhook.Ready(now) {
+ t.Fatal("stuck webhook should fail readiness")
+ }
+ freshRefund := now.Add(-time.Minute)
+ processingRefund := clean
+ processingRefund.RefundBacklog = 1
+ processingRefund.OldestRefund = &freshRefund
+ if !processingRefund.Ready(now) {
+ t.Fatal("fresh refund operation should remain ready during its processing window")
+ }
+ unmetered := clean
+ unmetered.UnmeteredSuccesses = 1
+ if unmetered.Ready(now) {
+ t.Fatal("unmetered success must fail readiness")
+ }
+}
diff --git a/internal/billing/service.go b/internal/billing/service.go
index b87a0bd..30f0e32 100644
--- a/internal/billing/service.go
+++ b/internal/billing/service.go
@@ -1,6 +1,7 @@
package billing
import (
+ "bufio"
"context"
"crypto/rand"
"encoding/hex"
@@ -8,13 +9,17 @@ import (
"errors"
"fmt"
"math/big"
+ "os"
+ "path/filepath"
"strings"
+ "sync"
"time"
"aigw/internal/domain"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
+ "github.com/stripe/stripe-go/v86"
)
const microsPerUnit = int64(1_000_000)
@@ -29,8 +34,15 @@ type Service struct {
stripeWebhookSecret string
stripeSuccessURL string
stripeCancelURL string
+ stripePortalReturnURL string
+ stripeAutomaticTax bool
+ stripeProductTaxCode string
integrationIdentifier string
createStripeCheckout stripeCheckoutCreator
+ stripeClient *stripe.Client
+ settlementSpoolPath string
+ metrics OperationalMetrics
+ spoolMu sync.Mutex
}
func New(ctx context.Context, options Options) (*Service, error) {
@@ -47,10 +59,15 @@ func New(ctx context.Context, options Options) (*Service, error) {
minTopUpMinor: options.MinTopUpMinor, maxTopUpMinor: options.MaxTopUpMinor,
stripeEnabled: options.StripeEnabled, stripeWebhookSecret: options.StripeWebhookSecret,
stripeSuccessURL: options.StripeSuccessURL, stripeCancelURL: options.StripeCancelURL,
+ stripePortalReturnURL: options.StripePortalReturnURL, stripeAutomaticTax: options.StripeAutomaticTax,
+ stripeProductTaxCode: options.StripeProductTaxCode,
integrationIdentifier: "aigw_balance_" + randomLetters(8),
+ settlementSpoolPath: strings.TrimSpace(options.SettlementSpoolPath),
+ metrics: options.Metrics,
}
if options.StripeEnabled {
- service.createStripeCheckout = newStripeCheckoutCreator(options.StripeAPIKey)
+ service.stripeClient = stripe.NewClient(options.StripeAPIKey)
+ service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create
}
return service, nil
}
@@ -67,10 +84,15 @@ func (s *Service) Currency() string {
return s.currency
}
+func (s *Service) Ping(ctx context.Context) error { return s.db.Ping(ctx) }
+
func (s *Service) Authorize(ctx context.Context, input Authorization) error {
if input.RequestID == "" || input.Principal.TenantID == "" || input.Principal.ProjectID == "" || input.Principal.KeyID == "" {
return errors.New("billing authorization identity is incomplete")
}
+ if input.Model.PriceCurrency != "" && input.Model.PriceCurrency != s.currency {
+ return fmt.Errorf("model price currency %s does not match wallet currency %s", input.Model.PriceCurrency, s.currency)
+ }
reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens)
if err != nil {
return err
@@ -101,7 +123,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
var used, pending int64
if err := tx.QueryRow(ctx, `SELECT
COALESCE((SELECT cost_micros FROM usage_monthly_rollups WHERE project_id=$1 AND period_start=$2),0),
- COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status='pending' AND created_at >= $2 AND created_at < $3),0)`,
+ COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`,
input.Principal.ProjectID, period, nextPeriod).Scan(&used, &pending); err != nil {
return fmt.Errorf("read monthly spend quota: %w", err)
}
@@ -115,15 +137,19 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
if _, err := tx.Exec(ctx, `
INSERT INTO billing_reservations (
request_id, tenant_id, project_id, key_id, public_model, currency, reserved_micros,
+ price_version_id,
input_price_micros_per_million, output_price_micros_per_million,
cache_read_price_micros_per_million, cache_write_price_micros_per_million)
- VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)`,
+ VALUES ($1,$2,$3,$4,$5,$6,$7,NULLIF($8,'')::uuid,$9,$10,$11,$12)`,
input.RequestID, input.Principal.TenantID, input.Principal.ProjectID, input.Principal.KeyID,
- input.Model.ID, s.currency, reserved, input.Model.InputPriceMicrosPerMillion,
+ input.Model.ID, s.currency, reserved, input.Model.PriceVersionID, input.Model.InputPriceMicrosPerMillion,
input.Model.OutputPriceMicrosPerMillion, input.Model.CacheReadPriceMicrosPerMillion,
input.Model.CacheWritePriceMicrosPerMillion); err != nil {
return fmt.Errorf("create billing reservation: %w", err)
}
+ if _, err := tx.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, input.RequestID); err != nil {
+ return fmt.Errorf("create settlement job: %w", err)
+ }
if _, err := tx.Exec(ctx, `
UPDATE tenant_wallets SET reserved_micros = reserved_micros + $2, updated_at = now()
WHERE tenant_id = $1`, input.Principal.TenantID, reserved); err != nil {
@@ -135,7 +161,27 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
return nil
}
-func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
+func (s *Service) EnqueueSettlement(ctx context.Context, event domain.UsageEvent) error {
+ payload, err := json.Marshal(event)
+ if err != nil {
+ return fmt.Errorf("encode settlement event: %w", err)
+ }
+ _, err = s.db.Exec(ctx, `
+ INSERT INTO billing_settlement_jobs (request_id, event, status, available_at, updated_at)
+ VALUES ($1,$2,'pending',now(),now())
+ ON CONFLICT (request_id) DO UPDATE SET event=EXCLUDED.event,
+ status=CASE WHEN billing_settlement_jobs.status='done' THEN 'done' ELSE 'pending' END,
+ available_at=now(), locked_at=NULL, last_error='', updated_at=now()`, event.RequestID, payload)
+ if err == nil {
+ return nil
+ }
+ if spoolErr := s.appendSettlementSpool(payload); spoolErr != nil {
+ return fmt.Errorf("enqueue settlement in PostgreSQL: %v; append durable spool: %w", err, spoolErr)
+ }
+ return nil
+}
+
+func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error {
tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
if err != nil {
return fmt.Errorf("begin usage settlement: %w", err)
@@ -159,7 +205,51 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
}
actualCost := int64(0)
- if event.StatusCode >= 200 && event.StatusCode < 300 {
+ billableSuccess := event.StatusCode >= 200 && event.StatusCode < 300 && event.Success
+ if billableSuccess && !event.UsageReported {
+ // Fail closed: keep the authorization hold in place and make the request
+ // visible to reconciliation. Releasing it would turn an unmetered success
+ // into a free request; guessing tokens here could overcharge the customer.
+ if _, err := tx.Exec(ctx, `UPDATE billing_reservations SET status='metering_failed', settled_at=now()
+ WHERE request_id=$1`, event.RequestID); err != nil {
+ return fmt.Errorf("mark unmetered reservation: %w", err)
+ }
+ var usageAlreadyRecorded bool
+ if err := tx.QueryRow(ctx, `SELECT true FROM usage_events WHERE request_id=$1 FOR UPDATE`, event.RequestID).Scan(&usageAlreadyRecorded); err != nil && !errors.Is(err, pgx.ErrNoRows) {
+ return fmt.Errorf("lock unmetered usage event: %w", err)
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO usage_events (
+ request_id,tenant_id,project_id,key_id,public_model,provider_id,upstream_model,protocol,stream,
+ status_code,success,error_type,attempts,started_at,duration_ms,input_tokens,output_tokens,total_tokens,
+ cache_creation_input_tokens,cache_read_input_tokens,cost_micros,charged_micros,uncollected_micros,
+ usage_reported,metering_status)
+ VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,'usage_not_reported',$12,$13,$14,
+ $15,$16,$17,$18,$19,0,0,0,false,'missing')
+ ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',usage_reported=false,metering_status='missing'`,
+ event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel,
+ string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.Attempts, event.StartedAt,
+ event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens,
+ event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil {
+ return fmt.Errorf("persist unmetered usage event: %w", err)
+ }
+ if !usageAlreadyRecorded {
+ period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC)
+ if _, err := tx.Exec(ctx, `INSERT INTO usage_monthly_rollups
+ (period_start,tenant_id,project_id,request_count,successful_requests,input_tokens,output_tokens,total_tokens)
+ VALUES ($1,$2,$3,1,1,$4,$5,$6)
+ ON CONFLICT (project_id,period_start) DO UPDATE SET
+ request_count=usage_monthly_rollups.request_count+1,
+ successful_requests=usage_monthly_rollups.successful_requests+1,
+ input_tokens=usage_monthly_rollups.input_tokens+EXCLUDED.input_tokens,
+ output_tokens=usage_monthly_rollups.output_tokens+EXCLUDED.output_tokens,
+ total_tokens=usage_monthly_rollups.total_tokens+EXCLUDED.total_tokens,updated_at=now()`,
+ period, tenantID, projectID, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens); err != nil {
+ return fmt.Errorf("roll up unmetered usage event: %w", err)
+ }
+ }
+ return tx.Commit(ctx)
+ }
+ if billableSuccess {
actualCost, err = usageCost(event.Usage, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice)
if err != nil {
return err
@@ -195,8 +285,9 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
return fmt.Errorf("lock existing usage event: %w", err)
}
if usageAlreadyRecorded {
- if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4 WHERE request_id=$1`,
- event.RequestID, actualCost, charged, uncollected); err != nil {
+ if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4,
+ usage_reported=$5,metering_status=$6 WHERE request_id=$1`,
+ event.RequestID, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil {
return fmt.Errorf("apply usage charge: %w", err)
}
} else if _, err := tx.Exec(ctx, `
@@ -204,14 +295,14 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model,
protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms,
input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens,
- cost_micros, charged_micros, uncollected_micros)
- VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23)
+ cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status)
+ VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25)
ON CONFLICT (request_id) DO NOTHING`,
event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel,
string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts,
event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens,
event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens,
- actualCost, charged, uncollected); err != nil {
+ actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil {
return fmt.Errorf("persist usage event: %w", err)
}
if charged > 0 {
@@ -251,6 +342,199 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error {
return nil
}
+// RunSettlementWorker processes jobs using PostgreSQL row locks so any gateway
+// instance can resume work left by another instance after a crash.
+func (s *Service) RunSettlementWorker(ctx context.Context) {
+ ticker := time.NewTicker(time.Second)
+ defer ticker.Stop()
+ for {
+ s.drainSettlementSpool(ctx)
+ s.recoverStaleSettlements(ctx)
+ for i := 0; i < 32; i++ {
+ processed, err := s.processSettlementJob(ctx)
+ if err != nil || !processed {
+ break
+ }
+ }
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ }
+ }
+}
+
+func (s *Service) recoverStaleSettlements(ctx context.Context) {
+ // A process may die after authorization but before it can attach a response
+ // event. Releasing after an hour prevents permanent holds; such synthetic
+ // events stay visible in Usage for reconciliation.
+ rows, err := s.db.Query(ctx, `SELECT request_id,tenant_id::text,project_id::text,key_id::text,public_model,created_at
+ FROM billing_reservations WHERE status='pending' AND created_at<now()-interval '1 hour'
+ AND EXISTS(SELECT 1 FROM billing_settlement_jobs j WHERE j.request_id=billing_reservations.request_id AND j.status='awaiting_event') LIMIT 100`)
+ if err != nil {
+ return
+ }
+ defer rows.Close()
+ for rows.Next() {
+ var event domain.UsageEvent
+ if rows.Scan(&event.RequestID, &event.TenantID, &event.ProjectID, &event.KeyID, &event.PublicModel, &event.StartedAt) != nil {
+ continue
+ }
+ event.StatusCode = 500
+ event.Success = false
+ event.ErrorType = "gateway_interrupted_before_settlement"
+ event.DurationMS = time.Since(event.StartedAt).Milliseconds()
+ _ = s.EnqueueSettlement(ctx, event)
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry',locked_at=NULL,available_at=now(),updated_at=now(),last_error='recovered stale processing lease'
+ WHERE status='processing' AND locked_at<now()-interval '5 minutes'`)
+}
+
+func (s *Service) processSettlementJob(ctx context.Context) (bool, error) {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return false, err
+ }
+ defer tx.Rollback(ctx)
+ var requestID string
+ var payload []byte
+ err = tx.QueryRow(ctx, `
+ WITH selected AS (
+ SELECT request_id FROM billing_settlement_jobs
+ WHERE status IN ('pending','retry') AND available_at <= now()
+ ORDER BY available_at, created_at FOR UPDATE SKIP LOCKED LIMIT 1
+ )
+ UPDATE billing_settlement_jobs j SET status='processing', attempts=attempts+1,
+ locked_at=now(), updated_at=now()
+ FROM selected WHERE j.request_id=selected.request_id
+ RETURNING j.request_id, j.event`,
+ ).Scan(&requestID, &payload)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return false, tx.Commit(ctx)
+ }
+ if err != nil {
+ return false, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return false, err
+ }
+ var event domain.UsageEvent
+ if err := json.Unmarshal(payload, &event); err != nil {
+ s.retrySettlement(ctx, requestID, fmt.Errorf("decode settlement event: %w", err))
+ return true, err
+ }
+ jobCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
+ err = s.settle(jobCtx, event)
+ cancel()
+ if err != nil {
+ s.retrySettlement(ctx, requestID, err)
+ return true, err
+ }
+ _, err = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='done', locked_at=NULL,
+ last_error='', updated_at=now(), completed_at=now() WHERE request_id=$1`, requestID)
+ return true, err
+}
+
+func (s *Service) retrySettlement(ctx context.Context, requestID string, cause error) {
+ message := cause.Error()
+ if len(message) > 1000 {
+ message = message[:1000]
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry', locked_at=NULL,
+ last_error=$2, available_at=now() + make_interval(secs => LEAST(300, power(2, LEAST(attempts, 8))::int)),
+ updated_at=now() WHERE request_id=$1`, requestID, message)
+}
+
+func (s *Service) appendSettlementSpool(payload []byte) error {
+ if s.settlementSpoolPath == "" {
+ return errors.New("settlement spool path is not configured")
+ }
+ s.spoolMu.Lock()
+ defer s.spoolMu.Unlock()
+ if err := os.MkdirAll(filepath.Dir(s.settlementSpoolPath), 0o700); err != nil {
+ return err
+ }
+ file, err := os.OpenFile(s.settlementSpoolPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600)
+ if err != nil {
+ return err
+ }
+ defer file.Close()
+ if _, err := file.Write(append(payload, '\n')); err != nil {
+ return err
+ }
+ return file.Sync()
+}
+
+func (s *Service) drainSettlementSpool(ctx context.Context) {
+ if s.settlementSpoolPath == "" {
+ return
+ }
+ s.spoolMu.Lock()
+ defer s.spoolMu.Unlock()
+ file, err := os.Open(s.settlementSpoolPath)
+ if errors.Is(err, os.ErrNotExist) {
+ return
+ }
+ if err != nil {
+ return
+ }
+ var pending [][]byte
+ scanner := bufio.NewScanner(file)
+ scanner.Buffer(make([]byte, 64*1024), 2<<20)
+ for scanner.Scan() {
+ line := append([]byte(nil), scanner.Bytes()...)
+ var event domain.UsageEvent
+ if json.Unmarshal(line, &event) != nil || event.RequestID == "" {
+ pending = append(pending, line)
+ continue
+ }
+ if _, err := s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET event=$2, status=CASE WHEN status='done' THEN 'done' ELSE 'pending' END,
+ available_at=now(), locked_at=NULL, updated_at=now() WHERE request_id=$1`, event.RequestID, line); err != nil {
+ pending = append(pending, line)
+ }
+ }
+ _ = file.Close()
+ if scanner.Err() != nil {
+ return
+ }
+ temporary := s.settlementSpoolPath + ".tmp"
+ out, err := os.OpenFile(temporary, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600)
+ if err != nil {
+ return
+ }
+ for _, line := range pending {
+ _, _ = out.Write(append(line, '\n'))
+ }
+ _ = out.Sync()
+ _ = out.Close()
+ _ = os.Rename(temporary, s.settlementSpoolPath)
+}
+
+func (s *Service) SettlementQueueStatus(ctx context.Context) (SettlementQueueStatus, error) {
+ var result SettlementQueueStatus
+ err := s.db.QueryRow(ctx, `SELECT
+ count(*) FILTER (WHERE status='awaiting_event'), count(*) FILTER (WHERE status='pending'),
+ count(*) FILTER (WHERE status='processing'), count(*) FILTER (WHERE status='retry'),
+ min(created_at) FILTER (WHERE status IN ('awaiting_event','pending','processing','retry'))
+ FROM billing_settlement_jobs`).Scan(&result.AwaitingEvent, &result.Pending, &result.Processing, &result.Retrying, &result.OldestPending)
+ if err != nil {
+ return result, err
+ }
+ if s.settlementSpoolPath != "" {
+ s.spoolMu.Lock()
+ file, openErr := os.Open(s.settlementSpoolPath)
+ if openErr == nil {
+ scanner := bufio.NewScanner(file)
+ for scanner.Scan() {
+ result.SpoolRecords++
+ }
+ _ = file.Close()
+ }
+ s.spoolMu.Unlock()
+ }
+ return result, nil
+}
+
func boolToInt(value bool) int {
if value {
return 1
@@ -258,8 +542,21 @@ func boolToInt(value bool) int {
return 0
}
+func meteringStatus(event domain.UsageEvent) string {
+ if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success {
+ return "upstream_failed"
+ }
+ if event.UsageReported {
+ return "reported"
+ }
+ return "missing"
+}
+
func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64) (int64, error) {
maxOutput := defaultMaxOutput
+ if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) {
+ maxOutput = model.MaxOutputTokens
+ }
var limits struct {
MaxTokens int64 `json:"max_tokens"`
MaxCompletionTokens int64 `json:"max_completion_tokens"`
diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go
index f675da7..6a1bb49 100644
--- a/internal/billing/service_test.go
+++ b/internal/billing/service_test.go
@@ -2,6 +2,7 @@ package billing
import (
"context"
+ "crypto/sha256"
"encoding/json"
"fmt"
"net/http"
@@ -97,6 +98,26 @@ func TestCheckoutReturnURLPreservesCallbackAndSessionPlaceholder(t *testing.T) {
}
}
+func TestSettlementSpoolWritesOneDurableRecord(t *testing.T) {
+ path := t.TempDir() + "/settlements.jsonl"
+ service := &Service{settlementSpoolPath: path}
+ event := domain.UsageEvent{RequestID: "req_spool_test", TenantID: "tenant", StartedAt: time.Now().UTC()}
+ payload, err := json.Marshal(event)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := service.appendSettlementSpool(payload); err != nil {
+ t.Fatal(err)
+ }
+ contents, err := os.ReadFile(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if strings.Count(string(contents), "req_spool_test") != 1 || !strings.HasSuffix(string(contents), "\n") {
+ t.Fatalf("unexpected spool contents %q", contents)
+ }
+}
+
func TestWebhookRejectsInvalidSignatureBeforeProcessing(t *testing.T) {
service := &Service{stripeWebhookSecret: "whsec_test"}
request := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(`{"id":"evt_fake"}`))
@@ -133,6 +154,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
t.Fatal(err)
}
eventID := fmt.Sprintf("evt_aigw_%d", time.Now().UnixNano())
+ followupEventID := eventID + "_async"
t.Cleanup(func() {
cleanupCtx := context.Background()
for _, statement := range []struct {
@@ -140,6 +162,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
arg string
}{
{`DELETE FROM stripe_webhook_events WHERE event_id=$1`, eventID},
+ {`DELETE FROM stripe_webhook_events WHERE event_id=$1`, followupEventID},
{`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID},
{`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID},
{`DELETE FROM topup_orders WHERE tenant_id=$1`, tenantID},
@@ -177,6 +200,28 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
t.Fatalf("delivery %d status = %d, body = %s", delivery+1, response.Code, response.Body.String())
}
}
+ if _, err := service.db.Exec(ctx, `UPDATE topup_orders SET status='refunded',refunded_micros=amount_micros WHERE id=$1`, orderID); err != nil {
+ t.Fatal(err)
+ }
+ followupPayload, err := json.Marshal(map[string]any{
+ "id": followupEventID, "object": "event", "api_version": stripe.APIVersion,
+ "type": string(stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded),
+ "data": map[string]any{"object": map[string]any{
+ "id": sessionID, "object": "checkout.session", "client_reference_id": orderID,
+ "amount_total": amountMinor, "currency": "usd", "payment_status": "paid",
+ }},
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ followupSigned := webhook.GenerateTestSignedPayload(&webhook.UnsignedPayload{Payload: followupPayload, Secret: service.stripeWebhookSecret})
+ followupRequest := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(string(followupPayload)))
+ followupRequest.Header.Set("Stripe-Signature", followupSigned.Header)
+ followupResponse := httptest.NewRecorder()
+ service.WebhookHandler().ServeHTTP(followupResponse, followupRequest)
+ if followupResponse.Code != http.StatusOK {
+ t.Fatalf("follow-up status = %d, body = %s", followupResponse.Code, followupResponse.Body.String())
+ }
var balance int64
if err := service.db.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance); err != nil {
@@ -190,13 +235,136 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) {
if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1`, sessionID).Scan(&ledgerCount); err != nil {
t.Fatal(err)
}
- if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id=$1`, eventID).Scan(&webhookCount); err != nil {
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id IN ($1,$2)`, eventID, followupEventID).Scan(&webhookCount); err != nil {
t.Fatal(err)
}
if err := service.db.QueryRow(ctx, `SELECT status FROM topup_orders WHERE id=$1`, orderID).Scan(&orderStatus); err != nil {
t.Fatal(err)
}
- if ledgerCount != 1 || webhookCount != 1 || orderStatus != "paid" {
- t.Fatalf("ledger=%d webhook=%d order=%s, want 1/1/paid", ledgerCount, webhookCount, orderStatus)
+ if ledgerCount != 1 || webhookCount != 2 || orderStatus != "refunded" {
+ t.Fatalf("ledger=%d webhook=%d order=%s, want 1/2/refunded", ledgerCount, webhookCount, orderStatus)
+ }
+}
+
+func TestSettlementWorkerPersistsUsageAndReleasesReservationPostgres(t *testing.T) {
+ databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL")
+ if databaseURL == "" {
+ t.Skip("AIGW_TEST_DATABASE_URL is not set")
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil {
+ t.Fatal(err)
+ }
+ service, err := New(ctx, Options{DatabaseURL: databaseURL, Currency: "usd", SettlementSpoolPath: t.TempDir() + "/settlements.jsonl"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(service.Close)
+ slug := fmt.Sprintf("settle-%d", time.Now().UnixNano())
+ var tenantID, projectID, keyID string
+ if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Settlement integration') RETURNING id::text`, slug).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Settlement') RETURNING id::text`, tenantID).Scan(&projectID); err != nil {
+ t.Fatal(err)
+ }
+ keyHash := sha256.Sum256([]byte(slug))
+ if err := service.db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'integration','sk-test',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil {
+ t.Fatal(err)
+ }
+ requestID := fmt.Sprintf("req_settle_%d", time.Now().UnixNano())
+ missingRequestID := requestID + "_missing_usage"
+ t.Cleanup(func() {
+ for _, statement := range []struct {
+ query string
+ arg string
+ }{
+ {`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, requestID},
+ {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID},
+ {`DELETE FROM billing_reservations WHERE request_id=$1`, requestID},
+ {`DELETE FROM billing_reservations WHERE request_id=$1`, missingRequestID},
+ {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM api_keys WHERE id=$1`, keyID},
+ {`DELETE FROM projects WHERE id=$1`, projectID},
+ {`DELETE FROM tenants WHERE id=$1`, tenantID},
+ } {
+ if _, cleanupErr := service.db.Exec(context.Background(), statement.query, statement.arg); cleanupErr != nil {
+ t.Errorf("cleanup settlement integration data: %v", cleanupErr)
+ }
+ }
+ })
+ if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros,reserved_micros) VALUES ($1,'usd',1000,20)`, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ($1,$2,$3,$4,'demo/model','usd',20,1000000,1000000,0,0)`, requestID, tenantID, projectID, keyID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, requestID); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: requestID, TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: time.Now().UTC(), Usage: domain.Usage{InputTokens: 10}}); err != nil {
+ t.Fatal(err)
+ }
+ processed, err := service.processSettlementJob(ctx)
+ if err != nil || !processed {
+ t.Fatalf("process settlement: processed=%v err=%v", processed, err)
+ }
+ var reservationStatus, jobStatus string
+ var balance, reserved, charged, usageCount int64
+ if err := service.db.QueryRow(ctx, `SELECT status,charged_micros FROM billing_reservations WHERE request_id=$1`, requestID).Scan(&reservationStatus, &charged); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, requestID).Scan(&jobStatus); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE request_id=$1`, requestID).Scan(&usageCount); err != nil {
+ t.Fatal(err)
+ }
+ if reservationStatus != "settled" || jobStatus != "done" || balance != 990 || reserved != 0 || charged != 10 || usageCount != 1 {
+ t.Fatalf("reservation=%s job=%s balance=%d reserved=%d charged=%d usage=%d", reservationStatus, jobStatus, balance, reserved, charged, usageCount)
+ }
+
+ if _, err := service.db.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+30 WHERE tenant_id=$1`, tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations
+ (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million)
+ VALUES ($1,$2,$3,$4,'demo/model','usd',30,1000000,1000000,0,0)`, missingRequestID, tenantID, projectID, keyID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, missingRequestID); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: missingRequestID, TenantID: tenantID, ProjectID: projectID,
+ KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true,
+ UsageReported: false, StartedAt: time.Now().UTC()}); err != nil {
+ t.Fatal(err)
+ }
+ processed, err = service.processSettlementJob(ctx)
+ if err != nil || !processed {
+ t.Fatalf("process missing-usage settlement: processed=%v err=%v", processed, err)
+ }
+ var missingReservationStatus, missingJobStatus, meteringStatus string
+ if err := service.db.QueryRow(ctx, `SELECT status FROM billing_reservations WHERE request_id=$1`, missingRequestID).Scan(&missingReservationStatus); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID).Scan(&missingJobStatus); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil {
+ t.Fatal(err)
+ }
+ if err := service.db.QueryRow(ctx, `SELECT metering_status FROM usage_events WHERE request_id=$1`, missingRequestID).Scan(&meteringStatus); err != nil {
+ t.Fatal(err)
+ }
+ if missingReservationStatus != "metering_failed" || missingJobStatus != "done" || balance != 990 || reserved != 30 || meteringStatus != "missing" {
+ t.Fatalf("missing usage reservation=%s job=%s balance=%d reserved=%d metering=%s", missingReservationStatus, missingJobStatus, balance, reserved, meteringStatus)
}
}
diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go
index b887035..59ee9bd 100644
--- a/internal/billing/stripe.go
+++ b/internal/billing/stripe.go
@@ -39,6 +39,13 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec
"aigw_topup_order_id": orderID,
"aigw_tenant_id": strings.TrimSpace(input.TenantID),
},
+ InvoiceCreation: &stripe.CheckoutSessionCreateInvoiceCreationParams{
+ Enabled: stripe.Bool(true),
+ InvoiceData: &stripe.CheckoutSessionCreateInvoiceCreationInvoiceDataParams{
+ Description: stripe.String("AIGW prepaid API usage credit"),
+ Metadata: map[string]string{"aigw_topup_order_id": orderID, "aigw_tenant_id": strings.TrimSpace(input.TenantID)},
+ },
+ },
LineItems: []*stripe.CheckoutSessionCreateLineItemParams{{
Quantity: stripe.Int64(1),
PriceData: &stripe.CheckoutSessionCreateLineItemPriceDataParams{
@@ -51,10 +58,27 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec
},
}},
}
+ var customerID string
+ _ = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, input.TenantID).Scan(&customerID)
+ if customerID != "" {
+ params.Customer = stripe.String(customerID)
+ } else {
+ params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways))
+ if strings.TrimSpace(input.CustomerEmail) != "" {
+ params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail))
+ }
+ }
+ if s.stripeAutomaticTax {
+ params.AutomaticTax = &stripe.CheckoutSessionCreateAutomaticTaxParams{Enabled: stripe.Bool(true)}
+ params.TaxIDCollection = &stripe.CheckoutSessionCreateTaxIDCollectionParams{Enabled: stripe.Bool(true)}
+ params.LineItems[0].PriceData.ProductData.TaxCode = stripe.String(s.stripeProductTaxCode)
+ }
params.SetIdempotencyKey("aigw_topup_" + orderID)
session, err := s.createStripeCheckout(ctx, params)
if err != nil {
- _, _ = s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID)
+ if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil {
+ return CheckoutResult{}, errors.Join(fmt.Errorf("create Stripe Checkout Session: %w", err), fmt.Errorf("mark top-up order failed: %w", updateErr))
+ }
return CheckoutResult{}, fmt.Errorf("create Stripe Checkout Session: %w", err)
}
if session.ID == "" || session.URL == "" {
@@ -103,7 +127,12 @@ func (s *Service) WebhookHandler() http.Handler {
http.Error(w, "invalid webhook signature", http.StatusBadRequest)
return
}
+ if err := s.recordWebhookAttempt(r.Context(), event); err != nil {
+ http.Error(w, "webhook persistence failed", http.StatusInternalServerError)
+ return
+ }
if err := s.processStripeEvent(r.Context(), event); err != nil {
+ s.recordWebhookFailure(r.Context(), event.ID, err)
if errors.Is(err, ErrInvalidAmount) || isNotFound(err) {
http.Error(w, "invalid checkout event", http.StatusBadRequest)
return
@@ -117,15 +146,29 @@ func (s *Service) WebhookHandler() http.Handler {
}
func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) error {
- typeName := string(event.Type)
switch event.Type {
case stripe.EventTypeCheckoutSessionCompleted,
stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded,
stripe.EventTypeCheckoutSessionAsyncPaymentFailed,
stripe.EventTypeCheckoutSessionExpired:
+ return s.processCheckoutEvent(ctx, event)
+ case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded,
+ stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed,
+ stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated,
+ stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn,
+ stripe.EventTypeChargeDisputeFundsReinstated,
+ stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized,
+ stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed:
+ return s.processOperationalStripeEvent(ctx, event)
default:
- return nil
+ _, err := s.db.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='ignored event type'
+ WHERE event_id=$1 AND processed_at IS NULL`, event.ID)
+ return err
}
+}
+
+func (s *Service) processCheckoutEvent(ctx context.Context, event stripe.Event) error {
+ typeName := string(event.Type)
if event.Data == nil {
return ErrInvalidAmount
}
@@ -141,6 +184,9 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
return err
}
defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, event.ID); err != nil {
+ return err
+ }
tag, err := tx.Exec(ctx, `
INSERT INTO stripe_webhook_events (event_id, event_type) VALUES ($1,$2)
ON CONFLICT (event_id) DO NOTHING`, event.ID, typeName)
@@ -148,7 +194,13 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
return fmt.Errorf("record Stripe event: %w", err)
}
if tag.RowsAffected() == 0 {
- return tx.Commit(ctx)
+ var processed bool
+ if err := tx.QueryRow(ctx, `SELECT processed_at IS NOT NULL FROM stripe_webhook_events WHERE event_id=$1`, event.ID).Scan(&processed); err != nil {
+ return err
+ }
+ if processed {
+ return tx.Commit(ctx)
+ }
}
var tenantID, currency, status string
@@ -164,6 +216,32 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
if (storedSessionID != nil && *storedSessionID != session.ID) || amountMinor != session.AmountTotal || currency != string(session.Currency) {
return ErrInvalidAmount
}
+ customerID, paymentIntentID, invoiceID := "", "", ""
+ if session.Customer != nil {
+ customerID = session.Customer.ID
+ }
+ if session.PaymentIntent != nil {
+ paymentIntentID = session.PaymentIntent.ID
+ }
+ if session.Invoice != nil {
+ invoiceID = session.Invoice.ID
+ }
+ if customerID != "" {
+ email := session.CustomerEmail
+ if session.CustomerDetails != nil && session.CustomerDetails.Email != "" {
+ email = session.CustomerDetails.Email
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,$3)
+ ON CONFLICT (tenant_id) DO UPDATE SET stripe_customer_id=EXCLUDED.stripe_customer_id,
+ email=CASE WHEN EXCLUDED.email='' THEN stripe_customers.email ELSE EXCLUDED.email END,updated_at=now()`, tenantID, customerID, email); err != nil {
+ return err
+ }
+ }
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_customer_id=COALESCE(NULLIF($2,''),stripe_customer_id),
+ stripe_payment_intent_id=COALESCE(NULLIF($3,''),stripe_payment_intent_id),
+ stripe_invoice_id=COALESCE(NULLIF($4,''),stripe_invoice_id) WHERE id=$1`, session.ClientReferenceID, customerID, paymentIntentID, invoiceID); err != nil {
+ return err
+ }
if event.Type == stripe.EventTypeCheckoutSessionAsyncPaymentFailed || event.Type == stripe.EventTypeCheckoutSessionExpired {
orderStatus := "failed"
if event.Type == stripe.EventTypeCheckoutSessionExpired {
@@ -172,20 +250,24 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status = $2, stripe_session_id = COALESCE(stripe_session_id, $3) WHERE id = $1 AND status = 'pending'`, session.ClientReferenceID, orderStatus, session.ID); err != nil {
return err
}
- _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID)
+ _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID)
if err != nil {
return err
}
return tx.Commit(ctx)
}
if session.PaymentStatus != stripe.CheckoutSessionPaymentStatusPaid {
- _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID)
+ _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID)
if err != nil {
return err
}
return tx.Commit(ctx)
}
- if status != "paid" {
+ var alreadyCredited bool
+ if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1)`, session.ID).Scan(&alreadyCredited); err != nil {
+ return err
+ }
+ if !alreadyCredited {
if _, err := tx.Exec(ctx, `
INSERT INTO tenant_wallets (tenant_id, currency) VALUES ($1,$2)
ON CONFLICT (tenant_id) DO NOTHING`, tenantID, currency); err != nil {
@@ -209,14 +291,117 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er
ON CONFLICT (source_type, source_id) DO NOTHING`, tenantID, currency, amountMicros, newBalance, session.ID); err != nil {
return err
}
- if _, err := tx.Exec(ctx, `
- UPDATE topup_orders SET status = 'paid', stripe_session_id = COALESCE(stripe_session_id, $2), paid_at = now()
- WHERE id = $1`, session.ClientReferenceID, session.ID); err != nil {
+ }
+ if _, err := tx.Exec(ctx, `
+ UPDATE topup_orders SET status=CASE WHEN status IN ('pending','failed','expired') THEN 'paid' ELSE status END,
+ stripe_session_id=COALESCE(stripe_session_id,$2), paid_at=COALESCE(paid_at,now())
+ WHERE id=$1`, session.ClientReferenceID, session.ID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Service) processOperationalStripeEvent(ctx context.Context, event stripe.Event) error {
+ if event.ID == "" || event.Data == nil {
+ return ErrInvalidAmount
+ }
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, event.ID); err != nil {
+ return err
+ }
+ tag, err := tx.Exec(ctx, `INSERT INTO stripe_webhook_events (event_id,event_type) VALUES ($1,$2) ON CONFLICT DO NOTHING`, event.ID, string(event.Type))
+ if err != nil {
+ return err
+ }
+ if tag.RowsAffected() == 0 {
+ var processed bool
+ if err := tx.QueryRow(ctx, `SELECT processed_at IS NOT NULL FROM stripe_webhook_events WHERE event_id=$1`, event.ID).Scan(&processed); err != nil {
return err
}
+ if processed {
+ return tx.Commit(ctx)
+ }
}
- if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID); err != nil {
+ switch event.Type {
+ case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded:
+ var charge stripe.Charge
+ if json.Unmarshal(event.Data.Raw, &charge) != nil || charge.ID == "" {
+ return ErrInvalidAmount
+ }
+ paymentIntentID := ""
+ if charge.PaymentIntent != nil {
+ paymentIntentID = charge.PaymentIntent.ID
+ }
+ if paymentIntentID != "" {
+ if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_charge_id=$2,receipt_url=COALESCE(NULLIF($3,''),receipt_url) WHERE stripe_payment_intent_id=$1`, paymentIntentID, charge.ID, charge.ReceiptURL); err != nil {
+ return err
+ }
+ }
+ if event.Type == stripe.EventTypeChargeRefunded && charge.Refunds != nil {
+ for _, refund := range charge.Refunds.Data {
+ if err := s.applyRefundTx(ctx, tx, refund); err != nil {
+ return err
+ }
+ }
+ }
+ case stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed:
+ var refund stripe.Refund
+ if json.Unmarshal(event.Data.Raw, &refund) != nil || refund.ID == "" {
+ return ErrInvalidAmount
+ }
+ if err := s.applyRefundTx(ctx, tx, &refund); err != nil {
+ return err
+ }
+ case stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated,
+ stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn, stripe.EventTypeChargeDisputeFundsReinstated:
+ var dispute stripe.Dispute
+ if json.Unmarshal(event.Data.Raw, &dispute) != nil || dispute.ID == "" {
+ return ErrInvalidAmount
+ }
+ if err := s.applyDisputeTx(ctx, tx, &dispute, event.Type); err != nil {
+ return err
+ }
+ case stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized, stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed:
+ var invoice stripe.Invoice
+ if json.Unmarshal(event.Data.Raw, &invoice) != nil || invoice.ID == "" {
+ return ErrInvalidAmount
+ }
+ if err := s.applyInvoiceTx(ctx, tx, &invoice); err != nil {
+ return err
+ }
+ }
+ if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID); err != nil {
return err
}
return tx.Commit(ctx)
}
+
+func (s *Service) recordWebhookAttempt(ctx context.Context, event stripe.Event) error {
+ if event.ID == "" {
+ return ErrInvalidAmount
+ }
+ _, err := s.db.Exec(ctx, `INSERT INTO stripe_webhook_events
+ (event_id,event_type,attempts,last_attempt_at) VALUES ($1,$2,1,now())
+ ON CONFLICT (event_id) DO UPDATE SET attempts=stripe_webhook_events.attempts+1,last_attempt_at=now()`,
+ event.ID, string(event.Type))
+ return err
+}
+
+func (s *Service) recordWebhookFailure(ctx context.Context, eventID string, cause error) {
+ message := "webhook processing failed"
+ if cause != nil {
+ message = cause.Error()
+ }
+ if len(message) > 1000 {
+ message = message[:1000]
+ }
+ _, _ = s.db.Exec(ctx, `UPDATE stripe_webhook_events SET processing_error=$2,last_attempt_at=now()
+ WHERE event_id=$1 AND processed_at IS NULL`, eventID, message)
+}
diff --git a/internal/billing/types.go b/internal/billing/types.go
index 3d0461e..fe5df1e 100644
--- a/internal/billing/types.go
+++ b/internal/billing/types.go
@@ -14,11 +14,13 @@ var (
ErrInvalidAmount = errors.New("invalid amount")
ErrQuotaExceeded = errors.New("monthly spend quota exceeded")
ErrTopUpOrderNotFound = errors.New("top-up order not found")
+ ErrUsageNotReported = errors.New("billable successful response did not report usage")
+ ErrCannotResolveTopUp = errors.New("top-up order cannot be resolved as missing")
)
type Meter interface {
Authorize(context.Context, Authorization) error
- Settle(context.Context, domain.UsageEvent) error
+ EnqueueSettlement(context.Context, domain.UsageEvent) error
}
type Authorization struct {
@@ -40,6 +42,54 @@ type Options struct {
StripeWebhookSecret string
StripeSuccessURL string
StripeCancelURL string
+ StripePortalReturnURL string
+ StripeAutomaticTax bool
+ StripeProductTaxCode string
+ SettlementSpoolPath string
+ Metrics OperationalMetrics
+}
+
+type OperationalMetrics interface {
+ SetStripeOperations(OperationalStatus)
+}
+
+type OperationalStatus struct {
+ StripeEnabled bool `json:"stripe_enabled"`
+ ReconciliationStatus string `json:"reconciliation_status"`
+ ReconciliationMismatches int64 `json:"reconciliation_mismatches"`
+ ReconciliationCompletedAt *time.Time `json:"reconciliation_completed_at,omitempty"`
+ ReconciliationError string `json:"reconciliation_error,omitempty"`
+ UnprocessedWebhooks int64 `json:"unprocessed_webhooks"`
+ OldestUnprocessedWebhook *time.Time `json:"oldest_unprocessed_webhook,omitempty"`
+ RefundBacklog int64 `json:"refund_backlog"`
+ OldestRefund *time.Time `json:"oldest_refund,omitempty"`
+ UncollectedMicros int64 `json:"uncollected_micros"`
+ UnmeteredSuccesses int64 `json:"unmetered_successes"`
+}
+
+func (s OperationalStatus) Ready(now time.Time) bool {
+ if !s.StripeEnabled {
+ return s.UnmeteredSuccesses == 0
+ }
+ if s.ReconciliationCompletedAt == nil || now.Sub(*s.ReconciliationCompletedAt) > 2*time.Hour {
+ return false
+ }
+ webhooksStuck := s.UnprocessedWebhooks > 0 &&
+ (s.OldestUnprocessedWebhook == nil || now.Sub(*s.OldestUnprocessedWebhook) > 5*time.Minute)
+ refundsStuck := s.RefundBacklog > 0 &&
+ (s.OldestRefund == nil || now.Sub(*s.OldestRefund) > 15*time.Minute)
+ return s.ReconciliationStatus == "clean" && s.ReconciliationMismatches == 0 &&
+ s.ReconciliationError == "" && !webhooksStuck && !refundsStuck &&
+ s.UncollectedMicros == 0 && s.UnmeteredSuccesses == 0
+}
+
+type SettlementQueueStatus struct {
+ AwaitingEvent int64 `json:"awaiting_event"`
+ Pending int64 `json:"pending"`
+ Processing int64 `json:"processing"`
+ Retrying int64 `json:"retrying"`
+ OldestPending *time.Time `json:"oldest_pending,omitempty"`
+ SpoolRecords int `json:"spool_records"`
}
type Account struct {
@@ -73,8 +123,9 @@ type AdjustmentInput struct {
}
type CheckoutInput struct {
- TenantID string `json:"tenant_id"`
- AmountMinor int64 `json:"amount_minor"`
+ TenantID string `json:"tenant_id"`
+ AmountMinor int64 `json:"amount_minor"`
+ CustomerEmail string `json:"-"`
}
type CheckoutResult struct {
@@ -84,14 +135,99 @@ type CheckoutResult struct {
}
type TopUpOrder struct {
- ID string `json:"id"`
- TenantID string `json:"tenant_id"`
- AmountMinor int64 `json:"amount_minor"`
- AmountMicros int64 `json:"amount_micros"`
- Currency string `json:"currency"`
- Status string `json:"status"`
- StripeSessionID string `json:"stripe_session_id,omitempty"`
- CheckoutURL string `json:"checkout_url,omitempty"`
- CreatedAt time.Time `json:"created_at"`
- PaidAt *time.Time `json:"paid_at,omitempty"`
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ AmountMinor int64 `json:"amount_minor"`
+ AmountMicros int64 `json:"amount_micros"`
+ Currency string `json:"currency"`
+ Status string `json:"status"`
+ StripeSessionID string `json:"stripe_session_id,omitempty"`
+ CheckoutURL string `json:"checkout_url,omitempty"`
+ CreatedAt time.Time `json:"created_at"`
+ PaidAt *time.Time `json:"paid_at,omitempty"`
+ StripeCustomerID string `json:"stripe_customer_id,omitempty"`
+ StripePaymentIntentID string `json:"stripe_payment_intent_id,omitempty"`
+ StripeChargeID string `json:"stripe_charge_id,omitempty"`
+ StripeInvoiceID string `json:"stripe_invoice_id,omitempty"`
+ InvoiceURL string `json:"invoice_url,omitempty"`
+ InvoicePDFURL string `json:"invoice_pdf_url,omitempty"`
+ ReceiptURL string `json:"receipt_url,omitempty"`
+ RefundedMicros int64 `json:"refunded_micros"`
+ DisputedMicros int64 `json:"disputed_micros"`
+ ReconciliationStatus string `json:"reconciliation_status"`
+ ReconciledAt *time.Time `json:"reconciled_at,omitempty"`
+ ReconciliationError string `json:"reconciliation_error,omitempty"`
+}
+
+type ResolveMissingTopUpInput struct {
+ Reason string `json:"reason"`
+}
+
+type ResolutionActor struct {
+ ID string
+ Type string
+}
+
+type RefundInput struct {
+ AmountMinor int64 `json:"amount_minor"`
+ Reason string `json:"reason"`
+}
+
+type Refund struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ TopUpOrderID string `json:"topup_order_id"`
+ StripeRefundID string `json:"stripe_refund_id,omitempty"`
+ AmountMinor int64 `json:"amount_minor"`
+ AmountMicros int64 `json:"amount_micros"`
+ Currency string `json:"currency"`
+ Reason string `json:"reason"`
+ Status string `json:"status"`
+ LastError string `json:"last_error,omitempty"`
+ CreatedAt time.Time `json:"created_at"`
+ CompletedAt *time.Time `json:"completed_at,omitempty"`
+}
+
+type PortalResult struct {
+ URL string `json:"url"`
+}
+
+type ReconciliationResult struct {
+ ID string `json:"id"`
+ Status string `json:"status"`
+ CheckedOrders int64 `json:"checked_orders"`
+ MismatchCount int64 `json:"mismatch_count"`
+ Mismatches []map[string]any `json:"mismatches"`
+ Repairs []map[string]any `json:"repairs,omitempty"`
+}
+
+type PaymentDispute struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id,omitempty"`
+ TopUpOrderID string `json:"topup_order_id,omitempty"`
+ AmountMinor int64 `json:"amount_minor"`
+ AmountMicros int64 `json:"amount_micros"`
+ Currency string `json:"currency"`
+ Status string `json:"status"`
+ Reason string `json:"reason"`
+ DebitedMicros int64 `json:"debited_micros"`
+ UncollectedMicros int64 `json:"uncollected_micros"`
+ DueBy *time.Time `json:"due_by,omitempty"`
+ UpdatedAt time.Time `json:"updated_at"`
+}
+
+type Invoice struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id,omitempty"`
+ TopUpOrderID string `json:"topup_order_id,omitempty"`
+ Status string `json:"status"`
+ Currency string `json:"currency"`
+ AmountDueMinor int64 `json:"amount_due_minor"`
+ AmountPaidMinor int64 `json:"amount_paid_minor"`
+ AttemptCount int `json:"attempt_count"`
+ NextPaymentAttempt *time.Time `json:"next_payment_attempt,omitempty"`
+ HostedInvoiceURL string `json:"hosted_invoice_url,omitempty"`
+ InvoicePDFURL string `json:"invoice_pdf_url,omitempty"`
+ LastFailure string `json:"last_failure,omitempty"`
+ UpdatedAt time.Time `json:"updated_at"`
}
diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go
index 7966d9a..09a6fba 100644
--- a/internal/catalog/catalog.go
+++ b/internal/catalog/catalog.go
@@ -15,8 +15,9 @@ type Catalog struct {
}
type snapshot struct {
- models map[string]domain.Model
- list []domain.Model
+ models map[string]domain.Model
+ aliases map[string]string
+ list []domain.Model
}
func New(cfg config.Config) *Catalog {
@@ -61,15 +62,19 @@ func NewModels(models []domain.Model) *Catalog {
func (c *Catalog) Replace(source []domain.Model) {
models := make(map[string]domain.Model, len(source))
+ aliases := make(map[string]string)
list := make([]domain.Model, 0, len(source))
for _, sourceModel := range source {
model := sourceModel
model.Routes = append([]domain.Route(nil), sourceModel.Routes...)
models[model.ID] = model
+ for _, alias := range model.Aliases {
+ aliases[alias] = model.ID
+ }
list = append(list, model)
}
sort.Slice(list, func(i, j int) bool { return list[i].ID < list[j].ID })
- c.state.Store(&snapshot{models: models, list: list})
+ c.state.Store(&snapshot{models: models, aliases: aliases, list: list})
}
func (c *Catalog) Model(id string) (domain.Model, error) {
@@ -79,11 +84,38 @@ func (c *Catalog) Model(id string) (domain.Model, error) {
}
model, ok := current.models[id]
if !ok {
+ if canonical, aliasOK := current.aliases[id]; aliasOK {
+ model, ok = current.models[canonical]
+ }
+ }
+ if !ok {
return domain.Model{}, fmt.Errorf("model %q not found", id)
}
return model, nil
}
+func (c *Catalog) ModelForPrincipal(id string, principal domain.Principal) (domain.Model, error) {
+ model, err := c.Model(id)
+ if err != nil {
+ return domain.Model{}, err
+ }
+ if !model.Allows(principal) {
+ return domain.Model{}, fmt.Errorf("model %q not allowed", id)
+ }
+ return model, nil
+}
+
+func (c *Catalog) ModelsFor(protocol domain.Protocol, principal domain.Principal) []domain.Model {
+ models := c.Models(protocol)
+ result := models[:0]
+ for _, model := range models {
+ if model.Allows(principal) {
+ result = append(result, model)
+ }
+ }
+ return result
+}
+
func (c *Catalog) Models(protocol domain.Protocol) []domain.Model {
current := c.state.Load()
if current == nil {
diff --git a/internal/config/config.go b/internal/config/config.go
index 376cf77..a67f221 100644
--- a/internal/config/config.go
+++ b/internal/config/config.go
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"io"
+ "net"
"net/url"
"os"
"strings"
@@ -26,12 +27,27 @@ type Config struct {
}
type ServerConfig struct {
- Address string `json:"-"`
- AddressEnv string `json:"address_env"`
- MaxBodyBytes int64 `json:"max_body_bytes"`
- ReadHeaderTimeoutSecs int `json:"read_header_timeout_seconds"`
- IdleTimeoutSecs int `json:"idle_timeout_seconds"`
- ShutdownTimeoutSecs int `json:"shutdown_timeout_seconds"`
+ Address string `json:"-"`
+ AddressEnv string `json:"address_env"`
+ SplitListeners bool `json:"split_listeners"`
+ PublicAddressEnv string `json:"public_address_env"`
+ AdminAddressEnv string `json:"admin_address_env"`
+ WebhookAddressEnv string `json:"webhook_address_env"`
+ OperationsAddressEnv string `json:"operations_address_env"`
+ TrustedProxyCIDRsEnv string `json:"trusted_proxy_cidrs_env"`
+ RequireHTTPSEnv string `json:"require_https_env"`
+ DeploymentRegionEnv string `json:"deployment_region_env"`
+ PublicAddress string `json:"-"`
+ AdminAddress string `json:"-"`
+ WebhookAddress string `json:"-"`
+ OperationsAddress string `json:"-"`
+ TrustedProxyCIDRs []string `json:"-"`
+ RequireHTTPS bool `json:"-"`
+ DeploymentRegion string `json:"-"`
+ MaxBodyBytes int64 `json:"max_body_bytes"`
+ ReadHeaderTimeoutSecs int `json:"read_header_timeout_seconds"`
+ IdleTimeoutSecs int `json:"idle_timeout_seconds"`
+ ShutdownTimeoutSecs int `json:"shutdown_timeout_seconds"`
}
type AuthConfig struct {
@@ -40,45 +56,55 @@ type AuthConfig struct {
}
type ControlPlaneConfig struct {
- Enabled bool `json:"enabled"`
- DatabaseURLEnv string `json:"database_url_env"`
- RedisURLEnv string `json:"redis_url_env"`
- CredentialKeyEnv string `json:"credential_key_env"`
- RedisChannel string `json:"redis_channel"`
- SnapshotCacheKey string `json:"snapshot_cache_key"`
- ReloadIntervalSeconds int `json:"reload_interval_seconds"`
- AutoMigrate bool `json:"auto_migrate"`
- DatabaseURL string `json:"-"`
- RedisURL string `json:"-"`
- CredentialKey string `json:"-"`
+ Enabled bool `json:"enabled"`
+ DatabaseURLEnv string `json:"database_url_env"`
+ RedisURLEnv string `json:"redis_url_env"`
+ CredentialKeyEnv string `json:"credential_key_env"`
+ PreviousCredentialKeysEnv string `json:"previous_credential_keys_env"`
+ RedisChannel string `json:"redis_channel"`
+ SnapshotCacheKey string `json:"snapshot_cache_key"`
+ ReloadIntervalSeconds int `json:"reload_interval_seconds"`
+ AutoMigrate bool `json:"auto_migrate"`
+ DatabaseURL string `json:"-"`
+ RedisURL string `json:"-"`
+ CredentialKey string `json:"-"`
+ PreviousCredentialKeys []string `json:"-"`
}
type AdminConfig struct {
- Enabled bool `json:"enabled"`
- TokenEnv string `json:"token_env"`
- BasePath string `json:"base_path"`
- RegistrationEnabled bool `json:"registration_enabled"`
- SessionTTLHours int `json:"session_ttl_hours"`
- PublicURL string `json:"-"`
- PublicURLEnv string `json:"public_url_env"`
- Mail MailConfig `json:"mail"`
- WebAuthn WebAuthnConfig `json:"webauthn"`
- Token string `json:"-"`
+ Enabled bool `json:"enabled"`
+ TokenEnv string `json:"token_env"`
+ BasePath string `json:"base_path"`
+ RegistrationEnabled bool `json:"registration_enabled"`
+ SessionTTLHours int `json:"session_ttl_hours"`
+ AuditRetentionDays int `json:"audit_retention_days"`
+ SecurityRetentionDays int `json:"security_retention_days"`
+ PublicURL string `json:"-"`
+ PublicURLEnv string `json:"public_url_env"`
+ Mail MailConfig `json:"mail"`
+ WebAuthn WebAuthnConfig `json:"webauthn"`
+ Token string `json:"-"`
}
type MailConfig struct {
- Enabled bool `json:"enabled"`
- FromName string `json:"from_name"`
- TLSMode string `json:"tls_mode"`
- FromAddressEnv string `json:"from_address_env"`
- SMTPAddressEnv string `json:"smtp_address_env"`
- SMTPUsernameEnv string `json:"smtp_username_env"`
- SMTPPasswordEnv string `json:"smtp_password_env"`
- SMTPImplicitTLS bool `json:"smtp_implicit_tls"`
- FromAddress string `json:"-"`
- SMTPAddress string `json:"-"`
- SMTPUsername string `json:"-"`
- SMTPPassword string `json:"-"`
+ Enabled bool `json:"enabled"`
+ FromName string `json:"from_name"`
+ TLSMode string `json:"tls_mode"`
+ FromAddressEnv string `json:"from_address_env"`
+ SMTPAddressEnv string `json:"smtp_address_env"`
+ SMTPUsernameEnv string `json:"smtp_username_env"`
+ SMTPPasswordEnv string `json:"smtp_password_env"`
+ FeedbackSecretEnv string `json:"feedback_secret_env"`
+ LowBalanceMicros int64 `json:"low_balance_micros"`
+ SpendAnomalyMultiplier int64 `json:"spend_anomaly_multiplier"`
+ SpendAnomalyMinMicros int64 `json:"spend_anomaly_min_micros"`
+ NotificationIntervalSeconds int `json:"notification_interval_seconds"`
+ SMTPImplicitTLS bool `json:"smtp_implicit_tls"`
+ FromAddress string `json:"-"`
+ SMTPAddress string `json:"-"`
+ SMTPUsername string `json:"-"`
+ SMTPPassword string `json:"-"`
+ FeedbackSecret string `json:"-"`
}
type WebAuthnConfig struct {
@@ -134,19 +160,29 @@ type BillingConfig struct {
DefaultMaxOutputTokens int64 `json:"default_max_output_tokens"`
MinTopUpMinor int64 `json:"min_top_up_minor"`
MaxTopUpMinor int64 `json:"max_top_up_minor"`
+ SettlementSpoolPathEnv string `json:"settlement_spool_path_env"`
+ SettlementSpoolPath string `json:"-"`
Stripe StripeConfig `json:"stripe"`
}
type StripeConfig struct {
- Enabled bool `json:"enabled"`
- APIKeyEnv string `json:"api_key_env"`
- WebhookSecretEnv string `json:"webhook_secret_env"`
- SuccessURL string `json:"-"`
- CancelURL string `json:"-"`
- SuccessURLEnv string `json:"success_url_env"`
- CancelURLEnv string `json:"cancel_url_env"`
- APIKey string `json:"-"`
- WebhookSecret string `json:"-"`
+ Enabled bool `json:"enabled"`
+ APIKeyEnv string `json:"api_key_env"`
+ WebhookSecretEnv string `json:"webhook_secret_env"`
+ SuccessURLEnv string `json:"success_url_env"`
+ CancelURLEnv string `json:"cancel_url_env"`
+ PortalReturnURLEnv string `json:"portal_return_url_env"`
+ AutomaticTaxEnabledEnv string `json:"automatic_tax_enabled_env"`
+ TaxRegistrationConfirmedEnv string `json:"tax_registration_confirmed_env"`
+ ProductTaxCodeEnv string `json:"product_tax_code_env"`
+ SuccessURL string `json:"-"`
+ CancelURL string `json:"-"`
+ PortalReturnURL string `json:"-"`
+ APIKey string `json:"-"`
+ WebhookSecret string `json:"-"`
+ AutomaticTaxEnabled bool `json:"-"`
+ TaxRegistrationConfirmed bool `json:"-"`
+ ProductTaxCode string `json:"-"`
}
func Load(path string) (Config, error) {
@@ -186,6 +222,27 @@ func applyDefaults(cfg *Config) {
if cfg.Server.Address == "" {
cfg.Server.Address = ":8080"
}
+ if cfg.Server.PublicAddressEnv == "" {
+ cfg.Server.PublicAddressEnv = "AIGW_PUBLIC_ADDRESS"
+ }
+ if cfg.Server.AdminAddressEnv == "" {
+ cfg.Server.AdminAddressEnv = "AIGW_ADMIN_ADDRESS"
+ }
+ if cfg.Server.WebhookAddressEnv == "" {
+ cfg.Server.WebhookAddressEnv = "AIGW_WEBHOOK_ADDRESS"
+ }
+ if cfg.Server.OperationsAddressEnv == "" {
+ cfg.Server.OperationsAddressEnv = "AIGW_OPERATIONS_ADDRESS"
+ }
+ if cfg.Server.TrustedProxyCIDRsEnv == "" {
+ cfg.Server.TrustedProxyCIDRsEnv = "AIGW_TRUSTED_PROXY_CIDRS"
+ }
+ if cfg.Server.RequireHTTPSEnv == "" {
+ cfg.Server.RequireHTTPSEnv = "AIGW_REQUIRE_HTTPS"
+ }
+ if cfg.Server.DeploymentRegionEnv == "" {
+ cfg.Server.DeploymentRegionEnv = "AIGW_DEPLOYMENT_REGION"
+ }
if cfg.Server.MaxBodyBytes == 0 {
cfg.Server.MaxBodyBytes = 16 << 20
}
@@ -210,6 +267,9 @@ func applyDefaults(cfg *Config) {
if cfg.ControlPlane.CredentialKeyEnv == "" {
cfg.ControlPlane.CredentialKeyEnv = "AIGW_CREDENTIAL_KEY"
}
+ if cfg.ControlPlane.PreviousCredentialKeysEnv == "" {
+ cfg.ControlPlane.PreviousCredentialKeysEnv = "AIGW_CREDENTIAL_PREVIOUS_KEYS"
+ }
if cfg.ControlPlane.RedisChannel == "" {
cfg.ControlPlane.RedisChannel = "aigw:control:changed"
}
@@ -228,6 +288,12 @@ func applyDefaults(cfg *Config) {
if cfg.Admin.SessionTTLHours == 0 {
cfg.Admin.SessionTTLHours = 12
}
+ if cfg.Admin.AuditRetentionDays == 0 {
+ cfg.Admin.AuditRetentionDays = 2555
+ }
+ if cfg.Admin.SecurityRetentionDays == 0 {
+ cfg.Admin.SecurityRetentionDays = 30
+ }
if cfg.Admin.PublicURLEnv == "" {
cfg.Admin.PublicURLEnv = "AIGW_PUBLIC_URL"
}
@@ -249,6 +315,21 @@ func applyDefaults(cfg *Config) {
if cfg.Admin.Mail.SMTPPasswordEnv == "" {
cfg.Admin.Mail.SMTPPasswordEnv = "AIGW_SMTP_PASSWORD"
}
+ if cfg.Admin.Mail.FeedbackSecretEnv == "" {
+ cfg.Admin.Mail.FeedbackSecretEnv = "AIGW_MAIL_FEEDBACK_SECRET"
+ }
+ if cfg.Admin.Mail.LowBalanceMicros == 0 {
+ cfg.Admin.Mail.LowBalanceMicros = 5_000_000
+ }
+ if cfg.Admin.Mail.SpendAnomalyMultiplier == 0 {
+ cfg.Admin.Mail.SpendAnomalyMultiplier = 3
+ }
+ if cfg.Admin.Mail.SpendAnomalyMinMicros == 0 {
+ cfg.Admin.Mail.SpendAnomalyMinMicros = 10_000_000
+ }
+ if cfg.Admin.Mail.NotificationIntervalSeconds == 0 {
+ cfg.Admin.Mail.NotificationIntervalSeconds = 300
+ }
if cfg.Admin.WebAuthn.RPDisplayName == "" {
cfg.Admin.WebAuthn.RPDisplayName = "AIGW Console"
}
@@ -285,6 +366,9 @@ func applyDefaults(cfg *Config) {
if cfg.Billing.MinTopUpMinor == 0 {
cfg.Billing.MinTopUpMinor = 500
}
+ if cfg.Billing.SettlementSpoolPathEnv == "" {
+ cfg.Billing.SettlementSpoolPathEnv = "AIGW_SETTLEMENT_SPOOL_PATH"
+ }
if cfg.Billing.Stripe.APIKeyEnv == "" {
cfg.Billing.Stripe.APIKeyEnv = "AIGW_STRIPE_API_KEY"
}
@@ -297,6 +381,18 @@ func applyDefaults(cfg *Config) {
if cfg.Billing.Stripe.CancelURLEnv == "" {
cfg.Billing.Stripe.CancelURLEnv = "AIGW_STRIPE_CANCEL_URL"
}
+ if cfg.Billing.Stripe.PortalReturnURLEnv == "" {
+ cfg.Billing.Stripe.PortalReturnURLEnv = "AIGW_STRIPE_PORTAL_RETURN_URL"
+ }
+ if cfg.Billing.Stripe.AutomaticTaxEnabledEnv == "" {
+ cfg.Billing.Stripe.AutomaticTaxEnabledEnv = "AIGW_STRIPE_AUTOMATIC_TAX_ENABLED"
+ }
+ if cfg.Billing.Stripe.TaxRegistrationConfirmedEnv == "" {
+ cfg.Billing.Stripe.TaxRegistrationConfirmedEnv = "AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED"
+ }
+ if cfg.Billing.Stripe.ProductTaxCodeEnv == "" {
+ cfg.Billing.Stripe.ProductTaxCodeEnv = "AIGW_STRIPE_PRODUCT_TAX_CODE"
+ }
for i := range cfg.Models {
for j := range cfg.Models[i].Routes {
if cfg.Models[i].Routes[j].Weight == 0 {
@@ -310,10 +406,40 @@ func resolveSecrets(cfg *Config) error {
if value := strings.TrimSpace(os.Getenv(cfg.Server.AddressEnv)); value != "" {
cfg.Server.Address = value
}
+ if cfg.Server.SplitListeners {
+ if err := resolveRequiredEnv(&cfg.Server.PublicAddress, cfg.Server.PublicAddressEnv, "server.public_address"); err != nil {
+ return err
+ }
+ if err := resolveRequiredEnv(&cfg.Server.AdminAddress, cfg.Server.AdminAddressEnv, "server.admin_address"); err != nil {
+ return err
+ }
+ if err := resolveRequiredEnv(&cfg.Server.WebhookAddress, cfg.Server.WebhookAddressEnv, "server.webhook_address"); err != nil {
+ return err
+ }
+ if err := resolveRequiredEnv(&cfg.Server.OperationsAddress, cfg.Server.OperationsAddressEnv, "server.operations_address"); err != nil {
+ return err
+ }
+ }
+ for _, cidr := range strings.Split(os.Getenv(cfg.Server.TrustedProxyCIDRsEnv), ",") {
+ if cidr = strings.TrimSpace(cidr); cidr != "" {
+ cfg.Server.TrustedProxyCIDRs = append(cfg.Server.TrustedProxyCIDRs, cidr)
+ }
+ }
+ var err error
+ cfg.Server.RequireHTTPS, err = envBool(cfg.Server.RequireHTTPSEnv)
+ if err != nil {
+ return err
+ }
+ cfg.Server.DeploymentRegion = strings.ToLower(strings.TrimSpace(os.Getenv(cfg.Server.DeploymentRegionEnv)))
if cfg.ControlPlane.Enabled {
cfg.ControlPlane.DatabaseURL = os.Getenv(cfg.ControlPlane.DatabaseURLEnv)
cfg.ControlPlane.RedisURL = os.Getenv(cfg.ControlPlane.RedisURLEnv)
cfg.ControlPlane.CredentialKey = os.Getenv(cfg.ControlPlane.CredentialKeyEnv)
+ for _, value := range strings.Split(os.Getenv(cfg.ControlPlane.PreviousCredentialKeysEnv), ",") {
+ if value = strings.TrimSpace(value); value != "" {
+ cfg.ControlPlane.PreviousCredentialKeys = append(cfg.ControlPlane.PreviousCredentialKeys, value)
+ }
+ }
}
if cfg.Admin.Enabled {
cfg.Admin.Token = os.Getenv(cfg.Admin.TokenEnv)
@@ -325,6 +451,7 @@ func resolveSecrets(cfg *Config) error {
cfg.Admin.Mail.SMTPAddress = strings.TrimSpace(os.Getenv(cfg.Admin.Mail.SMTPAddressEnv))
cfg.Admin.Mail.SMTPUsername = os.Getenv(cfg.Admin.Mail.SMTPUsernameEnv)
cfg.Admin.Mail.SMTPPassword = os.Getenv(cfg.Admin.Mail.SMTPPasswordEnv)
+ cfg.Admin.Mail.FeedbackSecret = os.Getenv(cfg.Admin.Mail.FeedbackSecretEnv)
}
if cfg.Admin.WebAuthn.Enabled {
cfg.Admin.WebAuthn.RPID = strings.TrimSpace(os.Getenv(cfg.Admin.WebAuthn.RPIDEnv))
@@ -344,6 +471,22 @@ func resolveSecrets(cfg *Config) error {
if err := resolveRequiredEnv(&cfg.Billing.Stripe.CancelURL, cfg.Billing.Stripe.CancelURLEnv, "billing.stripe.cancel_url"); err != nil {
return err
}
+ if err := resolveRequiredEnv(&cfg.Billing.Stripe.PortalReturnURL, cfg.Billing.Stripe.PortalReturnURLEnv, "billing.stripe.portal_return_url"); err != nil {
+ return err
+ }
+ var err error
+ cfg.Billing.Stripe.AutomaticTaxEnabled, err = envBool(cfg.Billing.Stripe.AutomaticTaxEnabledEnv)
+ if err != nil {
+ return err
+ }
+ cfg.Billing.Stripe.TaxRegistrationConfirmed, err = envBool(cfg.Billing.Stripe.TaxRegistrationConfirmedEnv)
+ if err != nil {
+ return err
+ }
+ cfg.Billing.Stripe.ProductTaxCode = strings.TrimSpace(os.Getenv(cfg.Billing.Stripe.ProductTaxCodeEnv))
+ }
+ if cfg.Billing.Enabled {
+ cfg.Billing.SettlementSpoolPath = strings.TrimSpace(os.Getenv(cfg.Billing.SettlementSpoolPathEnv))
}
for i := range cfg.Providers {
provider := &cfg.Providers[i]
@@ -364,6 +507,17 @@ func resolveSecrets(cfg *Config) error {
return nil
}
+func envBool(name string) (bool, error) {
+ value := strings.TrimSpace(os.Getenv(name))
+ if value == "" || strings.EqualFold(value, "false") || value == "0" {
+ return false, nil
+ }
+ if strings.EqualFold(value, "true") || value == "1" {
+ return true, nil
+ }
+ return false, fmt.Errorf("environment variable %s must be true/false or 1/0", name)
+}
+
func resolveRequiredEnv(target *string, environment, field string) error {
if environment == "" {
return nil
@@ -383,6 +537,23 @@ func Validate(cfg Config) error {
if cfg.Observability.UsageBuffer < 1 {
return errors.New("observability.usage_buffer must be positive")
}
+ if cfg.Server.SplitListeners {
+ seen := map[string]string{}
+ for name, address := range map[string]string{"public": cfg.Server.PublicAddress, "admin": cfg.Server.AdminAddress, "webhook": cfg.Server.WebhookAddress, "operations": cfg.Server.OperationsAddress} {
+ if strings.TrimSpace(address) == "" {
+ return fmt.Errorf("server %s listener address is empty", name)
+ }
+ if previous, ok := seen[address]; ok {
+ return fmt.Errorf("server %s and %s listeners must use different addresses", previous, name)
+ }
+ seen[address] = name
+ }
+ }
+ for _, cidr := range cfg.Server.TrustedProxyCIDRs {
+ if _, _, err := net.ParseCIDR(cidr); err != nil {
+ return fmt.Errorf("invalid trusted proxy CIDR %q", cidr)
+ }
+ }
if cfg.ControlPlane.Enabled {
if cfg.ControlPlane.DatabaseURL == "" {
@@ -408,6 +579,9 @@ func Validate(cfg Config) error {
if cfg.Admin.SessionTTLHours < 1 || cfg.Admin.SessionTTLHours > 720 {
return errors.New("admin.session_ttl_hours must be between 1 and 720")
}
+ if cfg.Admin.AuditRetentionDays < 30 || cfg.Admin.SecurityRetentionDays < 1 {
+ return errors.New("admin audit retention must be at least 30 days and security retention at least 1 day")
+ }
publicURL, err := url.Parse(cfg.Admin.PublicURL)
if err != nil || publicURL.Host == "" || (publicURL.Scheme != "http" && publicURL.Scheme != "https") {
return errors.New("admin.public_url must resolve from an environment variable to an absolute http(s) URL")
@@ -419,6 +593,13 @@ func Validate(cfg Config) error {
if (cfg.Admin.Mail.SMTPUsername == "") != (cfg.Admin.Mail.SMTPPassword == "") {
return errors.New("admin.mail SMTP username and password must both be set or both be empty")
}
+ if cfg.Admin.Mail.FeedbackSecret != "" && len(cfg.Admin.Mail.FeedbackSecret) < 32 {
+ return errors.New("admin.mail feedback secret must contain at least 32 characters")
+ }
+ if cfg.Admin.Mail.LowBalanceMicros < 0 || cfg.Admin.Mail.SpendAnomalyMultiplier < 2 ||
+ cfg.Admin.Mail.SpendAnomalyMinMicros < 0 || cfg.Admin.Mail.NotificationIntervalSeconds < 30 {
+ return errors.New("admin.mail notification thresholds are invalid")
+ }
switch cfg.Admin.Mail.TLSMode {
case "starttls", "tls", "none":
default:
@@ -453,6 +634,9 @@ func Validate(cfg Config) error {
if cfg.Billing.MinTopUpMinor < 1 || cfg.Billing.MaxTopUpMinor < cfg.Billing.MinTopUpMinor {
return errors.New("billing top-up bounds are invalid")
}
+ if strings.TrimSpace(cfg.Billing.SettlementSpoolPath) == "" {
+ return fmt.Errorf("billing: environment variable %s is empty; durable settlement fallback is required", cfg.Billing.SettlementSpoolPathEnv)
+ }
if cfg.Billing.Stripe.Enabled {
if cfg.Billing.Stripe.APIKey == "" {
return fmt.Errorf("billing.stripe: environment variable %s is empty", cfg.Billing.Stripe.APIKeyEnv)
@@ -460,14 +644,17 @@ func Validate(cfg Config) error {
if cfg.Billing.Stripe.WebhookSecret == "" {
return fmt.Errorf("billing.stripe: environment variable %s is empty", cfg.Billing.Stripe.WebhookSecretEnv)
}
- if strings.TrimSpace(cfg.Billing.Stripe.SuccessURL) == "" || strings.TrimSpace(cfg.Billing.Stripe.CancelURL) == "" {
- return errors.New("billing.stripe.success_url and cancel_url are required")
+ if strings.TrimSpace(cfg.Billing.Stripe.SuccessURL) == "" || strings.TrimSpace(cfg.Billing.Stripe.CancelURL) == "" || strings.TrimSpace(cfg.Billing.Stripe.PortalReturnURL) == "" {
+ return errors.New("billing.stripe success, cancel, and portal return URLs are required")
}
- for name, value := range map[string]string{"success_url": cfg.Billing.Stripe.SuccessURL, "cancel_url": cfg.Billing.Stripe.CancelURL} {
+ for name, value := range map[string]string{"success_url": cfg.Billing.Stripe.SuccessURL, "cancel_url": cfg.Billing.Stripe.CancelURL, "portal_return_url": cfg.Billing.Stripe.PortalReturnURL} {
parsed, err := url.Parse(value)
if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") {
return fmt.Errorf("billing.stripe.%s must be an absolute http(s) URL", name)
}
+ if cfg.Billing.Stripe.AutomaticTaxEnabled && (!cfg.Billing.Stripe.TaxRegistrationConfirmed || cfg.Billing.Stripe.ProductTaxCode == "") {
+ return errors.New("Stripe automatic tax requires an explicit confirmed registration and product tax code")
+ }
}
}
}
diff --git a/internal/config/config_test.go b/internal/config/config_test.go
index 1ecf70f..8e2e9bb 100644
--- a/internal/config/config_test.go
+++ b/internal/config/config_test.go
@@ -110,6 +110,8 @@ func TestLoadResolvesStripeSecrets(t *testing.T) {
t.Setenv("TEST_STRIPE_WEBHOOK", "whsec_example")
t.Setenv("TEST_STRIPE_SUCCESS", "https://console.example.test/admin/?topup=success")
t.Setenv("TEST_STRIPE_CANCEL", "https://console.example.test/admin/?topup=cancel")
+ t.Setenv("AIGW_STRIPE_PORTAL_RETURN_URL", "https://console.example.test/admin/?billing=portal")
+ t.Setenv("AIGW_SETTLEMENT_SPOOL_PATH", filepath.Join(t.TempDir(), "settlements.jsonl"))
path := writeConfig(t, `{
"control_plane":{"enabled":true},
"billing":{"enabled":true,"stripe":{"enabled":true,"api_key_env":"TEST_STRIPE_KEY","webhook_secret_env":"TEST_STRIPE_WEBHOOK","success_url_env":"TEST_STRIPE_SUCCESS","cancel_url_env":"TEST_STRIPE_CANCEL"}}
@@ -147,6 +149,10 @@ func TestLoadRejectsBillingWithoutControlPlane(t *testing.T) {
func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) {
t.Setenv("AIGW_SERVER_ADDRESS", "127.0.0.1:18081")
+ t.Setenv("AIGW_PUBLIC_ADDRESS", "127.0.0.1:18081")
+ t.Setenv("AIGW_ADMIN_ADDRESS", "127.0.0.1:18082")
+ t.Setenv("AIGW_WEBHOOK_ADDRESS", "127.0.0.1:18083")
+ t.Setenv("AIGW_OPERATIONS_ADDRESS", "127.0.0.1:19090")
t.Setenv("AIGW_API_KEYS", `[{"key":"test","key_id":"key","tenant_id":"tenant","project_id":"project","scopes":["inference"]}]`)
t.Setenv("OPENAI_BASE_URL", "https://openai.example.test/v1")
t.Setenv("OPENAI_API_KEY", "openai-secret")
@@ -175,6 +181,8 @@ func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) {
t.Setenv("AIGW_STRIPE_WEBHOOK_SECRET", "whsec_example")
t.Setenv("AIGW_STRIPE_SUCCESS_URL", "https://console.example.test/admin/?topup=success")
t.Setenv("AIGW_STRIPE_CANCEL_URL", "https://console.example.test/admin/?topup=cancel")
+ t.Setenv("AIGW_STRIPE_PORTAL_RETURN_URL", "https://console.example.test/admin/?billing=portal")
+ t.Setenv("AIGW_SETTLEMENT_SPOOL_PATH", filepath.Join(t.TempDir(), "settlements.jsonl"))
controlConfig, err := Load(filepath.Join("..", "..", "config.control.example.json"))
if err != nil {
t.Fatalf("load control-plane example: %v", err)
diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go
new file mode 100644
index 0000000..aaaf2c5
--- /dev/null
+++ b/internal/controlplane/mail_operations.go
@@ -0,0 +1,213 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/hex"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "log/slog"
+ "net/http"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+)
+
+const maxMailFeedbackBytes = 64 << 10
+
+type MailNotificationConfig struct {
+ LowBalanceMicros int64
+ SpendAnomalyMultiplier int64
+ SpendAnomalyMinMicros int64
+ Interval time.Duration
+}
+
+type mailFeedback struct {
+ EventID string `json:"event_id"`
+ EventType string `json:"event_type"`
+ Recipient string `json:"recipient"`
+ Provider string `json:"provider"`
+ Detail string `json:"detail"`
+}
+
+// MailFeedbackHandler accepts a provider-neutral normalized callback. An edge
+// adapter maps the provider's native event into this payload and signs
+// "<unix timestamp>.<raw body>" with HMAC-SHA256.
+func (s *Store) MailFeedbackHandler(secret string) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ w.Header().Set("Allow", http.MethodPost)
+ http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
+ return
+ }
+ body, err := io.ReadAll(io.LimitReader(r.Body, maxMailFeedbackBytes+1))
+ if err != nil || len(body) > maxMailFeedbackBytes || !validMailFeedbackSignature(secret, r.Header.Get("X-AIGW-Mail-Timestamp"), r.Header.Get("X-AIGW-Mail-Signature"), body, time.Now().UTC()) {
+ http.Error(w, "invalid feedback signature", http.StatusBadRequest)
+ return
+ }
+ var event mailFeedback
+ decoder := json.NewDecoder(strings.NewReader(string(body)))
+ decoder.DisallowUnknownFields()
+ if decoder.Decode(&event) != nil || strings.TrimSpace(event.EventID) == "" ||
+ !strings.Contains(event.Recipient, "@") || (event.EventType != "delivered" && event.EventType != "bounce" && event.EventType != "complaint") {
+ http.Error(w, "invalid feedback event", http.StatusBadRequest)
+ return
+ }
+ if err := s.applyMailFeedback(r.Context(), event); err != nil {
+ http.Error(w, "feedback persistence failed", http.StatusInternalServerError)
+ return
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = io.WriteString(w, "{\"received\":true}\n")
+ })
+}
+
+func validMailFeedbackSignature(secret, timestamp, signature string, body []byte, now time.Time) bool {
+ if secret == "" || timestamp == "" || !strings.HasPrefix(signature, "sha256=") {
+ return false
+ }
+ seconds, err := strconv.ParseInt(timestamp, 10, 64)
+ if err != nil || now.Sub(time.Unix(seconds, 0)).Abs() > 5*time.Minute {
+ return false
+ }
+ provided, err := hex.DecodeString(strings.TrimPrefix(signature, "sha256="))
+ if err != nil {
+ return false
+ }
+ mac := hmac.New(sha256.New, []byte(secret))
+ _, _ = mac.Write([]byte(timestamp))
+ _, _ = mac.Write([]byte("."))
+ _, _ = mac.Write(body)
+ return hmac.Equal(provided, mac.Sum(nil))
+}
+
+func (s *Store) applyMailFeedback(ctx context.Context, event mailFeedback) error {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ tag, err := tx.Exec(ctx, `INSERT INTO mail_feedback_events (event_id,event_type,recipient,provider)
+ VALUES ($1,$2,lower($3),$4) ON CONFLICT DO NOTHING`, event.EventID, event.EventType, event.Recipient, event.Provider)
+ if err != nil || tag.RowsAffected() == 0 {
+ if err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+ }
+ if event.EventType == "bounce" || event.EventType == "complaint" {
+ detail := event.Detail
+ if len(detail) > 1000 {
+ detail = detail[:1000]
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO mail_suppressions
+ (recipient,reason,provider,provider_event_id,detail) VALUES (lower($1),$2,$3,$4,$5)
+ ON CONFLICT (recipient) DO UPDATE SET reason=EXCLUDED.reason,provider=EXCLUDED.provider,
+ provider_event_id=EXCLUDED.provider_event_id,detail=EXCLUDED.detail,updated_at=now()`,
+ event.Recipient, event.EventType, event.Provider, event.EventID, detail); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_mail_outbox SET status='suppressed',last_error=$2,claimed_at=NULL
+ WHERE lower(recipient)=lower($1) AND status IN ('pending','retry','sending')`, event.Recipient, "recipient suppressed after "+event.EventType); err != nil {
+ return err
+ }
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) RunMailNotificationWorker(ctx context.Context, config MailNotificationConfig, logger *slog.Logger) {
+ if config.Interval <= 0 {
+ config.Interval = 5 * time.Minute
+ }
+ if logger == nil {
+ logger = slog.Default()
+ }
+ ticker := time.NewTicker(config.Interval)
+ defer ticker.Stop()
+ for {
+ if err := s.queueBillingNotifications(ctx, config); err != nil && !errors.Is(err, context.Canceled) {
+ logger.Warn("billing_notification_scan_failed", "error", err)
+ }
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ }
+ }
+}
+
+func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotificationConfig) error {
+ rows, err := s.db.Query(ctx, `WITH spend AS (
+ SELECT tenant_id,
+ COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())),0)::bigint today,
+ (COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())-interval '7 days' AND created_at<date_trunc('day',now())),0)/7)::bigint baseline
+ FROM billing_ledger GROUP BY tenant_id)
+ SELECT w.tenant_id::text,w.currency,w.balance_micros-w.reserved_micros,u.email,u.display_name,
+ COALESCE(spend.today,0),COALESCE(spend.baseline,0)
+ FROM tenant_wallets w JOIN console_users u ON u.tenant_id=w.tenant_id
+ LEFT JOIN spend ON spend.tenant_id=w.tenant_id
+ WHERE u.status='active' AND u.email_verified_at IS NOT NULL AND u.role IN ('tenant_admin','tenant_billing')
+ AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`)
+ if err != nil {
+ return err
+ }
+ defer rows.Close()
+ for rows.Next() {
+ var tenantID, currency, email, name string
+ var available, today, baseline int64
+ if err := rows.Scan(&tenantID, &currency, &available, &email, &name, &today, &baseline); err != nil {
+ return err
+ }
+ day := time.Now().UTC().Format("2006-01-02")
+ if available <= config.LowBalanceMicros {
+ body := fmt.Sprintf("Hi %s,\n\nYour AIGW prepaid balance is low: %.6f %s remains available. Add funds to avoid interrupted API access.\n", displayName(name), float64(available)/1_000_000, strings.ToUpper(currency))
+ if err := s.queueNotification(ctx, tenantID, email, "low_balance", day, "AIGW balance is low", body); err != nil {
+ return err
+ }
+ }
+ if baseline > 0 && today >= config.SpendAnomalyMinMicros && today >= baseline*config.SpendAnomalyMultiplier {
+ body := fmt.Sprintf("Hi %s,\n\nAIGW detected unusual API spend today: %.6f %s versus a seven-day daily baseline of %.6f %s. Review API keys and usage in the console.\n", displayName(name), float64(today)/1_000_000, strings.ToUpper(currency), float64(baseline)/1_000_000, strings.ToUpper(currency))
+ if err := s.queueNotification(ctx, tenantID, email, "spend_anomaly", day, "Unusual AIGW API spend detected", body); err != nil {
+ return err
+ }
+ }
+ }
+ return rows.Err()
+}
+
+func (s *Store) queueNotification(ctx context.Context, tenantID, recipient, kind, dedupe, subject, body string) error {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ tag, err := tx.Exec(ctx, `INSERT INTO mail_notification_events (tenant_id,recipient,notification_type,dedupe_key)
+ VALUES ($1,lower($2),$3,$4) ON CONFLICT DO NOTHING`, tenantID, recipient, kind, dedupe)
+ if err != nil || tag.RowsAffected() == 0 {
+ if err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+ }
+ ciphertext, err := s.cipher.Encrypt(body)
+ if err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO console_mail_outbox (recipient,template,subject,body_ciphertext)
+ VALUES (lower($1),$2,$3,$4)`, recipient, kind, subject, ciphertext); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func displayName(value string) string {
+ if value = strings.TrimSpace(value); value != "" {
+ return value
+ }
+ return "there"
+}
diff --git a/internal/controlplane/mail_operations_test.go b/internal/controlplane/mail_operations_test.go
new file mode 100644
index 0000000..0b47cdb
--- /dev/null
+++ b/internal/controlplane/mail_operations_test.go
@@ -0,0 +1,29 @@
+package controlplane
+
+import (
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/hex"
+ "strconv"
+ "testing"
+ "time"
+)
+
+func TestMailFeedbackSignature(t *testing.T) {
+ now := time.Unix(1_800_000_000, 0).UTC()
+ timestamp := strconv.FormatInt(now.Unix(), 10)
+ body := []byte(`{"event_id":"evt_1","event_type":"bounce","recipient":"test@example.com","provider":"test"}`)
+ mac := hmac.New(sha256.New, []byte("a-production-length-feedback-secret"))
+ _, _ = mac.Write([]byte(timestamp + "."))
+ _, _ = mac.Write(body)
+ signature := "sha256=" + hex.EncodeToString(mac.Sum(nil))
+ if !validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, body, now) {
+ t.Fatal("valid signature was rejected")
+ }
+ if validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, []byte(`{}`), now) {
+ t.Fatal("signature must bind the raw body")
+ }
+ if validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, body, now.Add(6*time.Minute)) {
+ t.Fatal("stale signature was accepted")
+ }
+}
diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go
index b6be748..212963b 100644
--- a/internal/controlplane/manager.go
+++ b/internal/controlplane/manager.go
@@ -28,16 +28,19 @@ type policyReplacer interface {
}
type Manager struct {
- store managerStore
- catalog *catalog.Catalog
- authenticator *auth.StaticAuthenticator
- logger *slog.Logger
- pollInterval time.Duration
- generation atomic.Int64
- redisConnected atomic.Bool
- reloadMu sync.Mutex
- broadcasts chan ChangeEvent
- policyTarget policyReplacer
+ store managerStore
+ catalog *catalog.Catalog
+ authenticator *auth.StaticAuthenticator
+ logger *slog.Logger
+ pollInterval time.Duration
+ generation atomic.Int64
+ redisConnected atomic.Bool
+ loaded atomic.Bool
+ lastReloadUnix atomic.Int64
+ lastHealthyUnix atomic.Int64
+ reloadMu sync.Mutex
+ broadcasts chan ChangeEvent
+ policyTarget policyReplacer
}
func NewManager(store managerStore, modelCatalog *catalog.Catalog, authenticator *auth.StaticAuthenticator, logger *slog.Logger, pollInterval time.Duration, policyTargets ...policyReplacer) *Manager {
@@ -67,10 +70,30 @@ func (m *Manager) Reload(ctx context.Context) (int64, error) {
m.policyTarget.ReplacePolicies(snapshot.Limits)
}
m.generation.Store(snapshot.Generation)
+ m.loaded.Store(true)
+ m.lastReloadUnix.Store(time.Now().UTC().Unix())
+ m.lastHealthyUnix.Store(time.Now().UTC().Unix())
m.logger.Info("control_plane_reloaded", "generation", snapshot.Generation, "models", len(snapshot.Models), "api_keys", len(snapshot.APIKeys))
return snapshot.Generation, nil
}
+func (m *Manager) Loaded() bool { return m.loaded.Load() }
+func (m *Manager) LastReloadAt() time.Time {
+ value := m.lastReloadUnix.Load()
+ if value == 0 {
+ return time.Time{}
+ }
+ return time.Unix(value, 0).UTC()
+}
+
+func (m *Manager) LastHealthyAt() time.Time {
+ value := m.lastHealthyUnix.Load()
+ if value == 0 {
+ return time.Time{}
+ }
+ return time.Unix(value, 0).UTC()
+}
+
func (m *Manager) AfterMutation(ctx context.Context, generation int64, resource, id string) error {
loadedGeneration, err := m.Reload(ctx)
if err != nil {
@@ -135,6 +158,7 @@ func (m *Manager) runPolling(ctx context.Context) {
m.logger.Warn("control_plane_generation_check_failed", "error", err)
continue
}
+ m.lastHealthyUnix.Store(time.Now().UTC().Unix())
if generation > m.generation.Load() {
if _, err := m.Reload(ctx); err != nil {
m.logger.Error("control_plane_reload_failed", "source", "postgres", "error", err)
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
index 3cedf70..c2cb7d8 100644
--- a/internal/controlplane/mutations.go
+++ b/internal/controlplane/mutations.go
@@ -11,6 +11,7 @@ import (
"net/url"
"regexp"
"strings"
+ "time"
"github.com/jackc/pgx/v5"
)
@@ -172,10 +173,45 @@ func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool)
func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, int64, error) {
input.PublicID = strings.TrimSpace(input.PublicID)
+ input.DisplayName = strings.TrimSpace(input.DisplayName)
+ input.Description = strings.TrimSpace(input.Description)
input.OwnedBy = strings.TrimSpace(input.OwnedBy)
+ input.Lifecycle = strings.ToLower(strings.TrimSpace(input.Lifecycle))
+ input.PriceCurrency = strings.ToLower(strings.TrimSpace(input.PriceCurrency))
+ if input.DisplayName == "" {
+ input.DisplayName = input.PublicID
+ }
+ if input.Lifecycle == "" {
+ input.Lifecycle = "active"
+ }
+ if input.PriceCurrency == "" {
+ input.PriceCurrency = "usd"
+ }
+ if len(input.InputModalities) == 0 {
+ input.InputModalities = []string{"text"}
+ }
+ if len(input.OutputModalities) == 0 {
+ input.OutputModalities = []string{"text"}
+ }
+ if len(input.Capabilities) == 0 {
+ input.Capabilities = []string{"chat", "streaming"}
+ }
+ input.InputModalities = uniqueStrings(input.InputModalities)
+ input.OutputModalities = uniqueStrings(input.OutputModalities)
+ input.Capabilities = uniqueStrings(input.Capabilities)
+ input.Regions = uniqueStrings(input.Regions)
+ input.Aliases = uniqueStrings(input.Aliases)
+ input.AllowedTenantIDs = uniqueStrings(input.AllowedTenantIDs)
+ input.AllowedKeyIDs = uniqueStrings(input.AllowedKeyIDs)
if input.PublicID == "" || len(input.Routes) == 0 {
return Model{}, 0, errors.New("model requires public_id and at least one route")
}
+ if input.ContextWindow < 0 || input.MaxOutputTokens < 0 || len(input.PriceCurrency) != 3 {
+ return Model{}, 0, errors.New("model context, output limit, or price currency is invalid")
+ }
+ if input.Lifecycle != "preview" && input.Lifecycle != "active" && input.Lifecycle != "deprecated" && input.Lifecycle != "retired" {
+ return Model{}, 0, errors.New("model lifecycle must be preview, active, deprecated, or retired")
+ }
if input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 || input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 {
return Model{}, 0, errors.New("model prices cannot be negative")
}
@@ -194,21 +230,63 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model,
return Model{}, 0, err
}
defer tx.Rollback(ctx)
+ inputModalitiesJSON, _ := json.Marshal(input.InputModalities)
+ outputModalitiesJSON, _ := json.Marshal(input.OutputModalities)
+ capabilitiesJSON, _ := json.Marshal(input.Capabilities)
+ regionsJSON, _ := json.Marshal(input.Regions)
var result Model
err = tx.QueryRow(ctx, `
- INSERT INTO models (public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million,
- cache_read_price_micros_per_million, cache_write_price_micros_per_million)
- VALUES ($1, $2, $3, $4, $5, $6)
- RETURNING id::text, public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million,
+ INSERT INTO models (public_id, display_name, description, owned_by, input_modalities, output_modalities,
+ context_window, max_output_tokens, capabilities, regions, lifecycle, released_at,
+ deprecated_at, retired_at, replacement_model, input_price_micros_per_million,
+ output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,NULLIF($15,''),$16,$17,$18,$19)
+ RETURNING id::text, public_id, display_name, description, owned_by,
+ input_price_micros_per_million, output_price_micros_per_million,
cache_read_price_micros_per_million, cache_write_price_micros_per_million, enabled, created_at`,
- input.PublicID, input.OwnedBy, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion,
+ input.PublicID, input.DisplayName, input.Description, input.OwnedBy, inputModalitiesJSON, outputModalitiesJSON,
+ input.ContextWindow, input.MaxOutputTokens, capabilitiesJSON, regionsJSON, input.Lifecycle,
+ input.ReleasedAt, input.DeprecatedAt, input.RetiredAt, input.ReplacementModel,
+ input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion,
input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion,
- ).Scan(&result.ID, &result.PublicID, &result.OwnedBy, &result.InputPriceMicrosPerMillion,
+ ).Scan(&result.ID, &result.PublicID, &result.DisplayName, &result.Description, &result.OwnedBy, &result.InputPriceMicrosPerMillion,
&result.OutputPriceMicrosPerMillion, &result.CacheReadPriceMicrosPerMillion,
&result.CacheWritePriceMicrosPerMillion, &result.Enabled, &result.CreatedAt)
if err != nil {
return Model{}, 0, fmt.Errorf("create model: %w", err)
}
+ result.InputModalities, result.OutputModalities = input.InputModalities, input.OutputModalities
+ result.ContextWindow, result.MaxOutputTokens = input.ContextWindow, input.MaxOutputTokens
+ result.Capabilities, result.Regions, result.Lifecycle = input.Capabilities, input.Regions, input.Lifecycle
+ result.ReleasedAt, result.DeprecatedAt, result.RetiredAt = input.ReleasedAt, input.DeprecatedAt, input.RetiredAt
+ result.ReplacementModel, result.Aliases = input.ReplacementModel, input.Aliases
+ result.AllowedTenantIDs, result.AllowedKeyIDs = input.AllowedTenantIDs, input.AllowedKeyIDs
+ result.PriceCurrency, result.PriceVersion = input.PriceCurrency, 1
+ err = tx.QueryRow(ctx, `INSERT INTO model_price_versions (model_id,version,currency,
+ input_price_micros_per_million,output_price_micros_per_million,
+ cache_read_price_micros_per_million,cache_write_price_micros_per_million)
+ VALUES ($1,1,$2,$3,$4,$5,$6) RETURNING id::text,effective_from`, result.ID, input.PriceCurrency,
+ input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion,
+ input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion,
+ ).Scan(&result.PriceVersionID, &result.PriceEffectiveFrom)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model price version: %w", err)
+ }
+ for _, alias := range input.Aliases {
+ if _, err := tx.Exec(ctx, `INSERT INTO model_aliases (alias,model_id,deprecated) VALUES ($1,$2,true)`, alias, result.ID); err != nil {
+ return Model{}, 0, fmt.Errorf("create model alias: %w", err)
+ }
+ }
+ for _, tenantID := range input.AllowedTenantIDs {
+ if _, err := tx.Exec(ctx, `INSERT INTO model_tenant_allowlist (model_id,tenant_id) VALUES ($1,$2)`, result.ID, tenantID); err != nil {
+ return Model{}, 0, fmt.Errorf("create tenant model allowlist: %w", err)
+ }
+ }
+ for _, keyID := range input.AllowedKeyIDs {
+ if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_allowlist (api_key_id,model_id) VALUES ($1,$2)`, keyID, result.ID); err != nil {
+ return Model{}, 0, fmt.Errorf("create key model allowlist: %w", err)
+ }
+ }
result.Routes = make([]Route, 0, len(input.Routes))
for _, route := range input.Routes {
var created Route
@@ -233,6 +311,59 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model,
return result, generation, nil
}
+func (s *Store) CreateModelPriceVersion(ctx context.Context, modelID string, input CreatePriceVersionInput) (int64, error) {
+ input.Currency = strings.ToLower(strings.TrimSpace(input.Currency))
+ if len(input.Currency) != 3 || input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 ||
+ input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 {
+ return 0, errors.New("price version has invalid currency or negative price")
+ }
+ if input.EffectiveFrom.IsZero() {
+ input.EffectiveFrom = time.Now().UTC()
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ // Lock the model row before allocating the next version. PostgreSQL does
+ // not allow FOR UPDATE on an aggregate query.
+ var modelExists bool
+ if err := tx.QueryRow(ctx, `SELECT true FROM models WHERE id=$1 FOR UPDATE`, modelID).Scan(&modelExists); errors.Is(err, pgx.ErrNoRows) {
+ return 0, ErrNotFound
+ } else if err != nil {
+ return 0, err
+ }
+ var currentEffectiveFrom time.Time
+ err = tx.QueryRow(ctx, `SELECT effective_from FROM model_price_versions WHERE model_id=$1 AND effective_to IS NULL`, modelID).Scan(&currentEffectiveFrom)
+ if err != nil && !errors.Is(err, pgx.ErrNoRows) {
+ return 0, err
+ }
+ if err == nil && !input.EffectiveFrom.After(currentEffectiveFrom) {
+ return 0, errors.New("new price version must become effective after the current open version")
+ }
+ var version int
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(max(version),0)+1 FROM model_price_versions WHERE model_id=$1`, modelID).Scan(&version); err != nil {
+ return 0, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE model_price_versions SET effective_to=$2 WHERE model_id=$1 AND effective_to IS NULL AND effective_from < $2`, modelID, input.EffectiveFrom); err != nil {
+ return 0, err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency,input_price_micros_per_million,
+ output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million,effective_from)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8)`, modelID, version, input.Currency, input.InputPriceMicrosPerMillion,
+ input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, input.EffectiveFrom); err != nil {
+ return 0, err
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return 0, err
+ }
+ return generation, nil
+}
+
func (s *Store) SetModelEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
return s.toggle(ctx, `UPDATE models SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
}
diff --git a/internal/controlplane/outbox.go b/internal/controlplane/outbox.go
index b54a064..b708a90 100644
--- a/internal/controlplane/outbox.go
+++ b/internal/controlplane/outbox.go
@@ -109,8 +109,9 @@ func (s *Store) ClaimMail(ctx context.Context) (mailer.Message, bool, error) {
var ciphertext []byte
err = tx.QueryRow(ctx, `WITH candidate AS (
SELECT id FROM console_mail_outbox
- WHERE ((status IN ('pending','failed') AND available_at <= now())
+ WHERE ((status IN ('pending','retry') AND available_at <= now())
OR (status='sending' AND claimed_at < now()-interval '5 minutes'))
+ AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(console_mail_outbox.recipient))
ORDER BY available_at,created_at FOR UPDATE SKIP LOCKED LIMIT 1
) UPDATE console_mail_outbox o SET status='sending',claimed_at=now(),attempts=attempts+1,last_error=''
FROM candidate WHERE o.id=candidate.id
@@ -146,7 +147,19 @@ func (s *Store) MarkMailFailed(ctx context.Context, id string, deliveryErr error
if len(message) > 1000 {
message = message[:1000]
}
- _, err := s.db.Exec(ctx, `UPDATE console_mail_outbox SET status='failed',claimed_at=NULL,last_error=$2,
- available_at=now()+make_interval(secs => LEAST(300, 5 * attempts)) WHERE id=$1 AND status='sending'`, id, message)
+ _, err := s.db.Exec(ctx, `UPDATE console_mail_outbox SET
+ status=CASE WHEN attempts>=10 THEN 'dead' ELSE 'retry' END,claimed_at=NULL,last_error=$2,
+ available_at=now()+make_interval(secs => LEAST(3600, 5 * power(2,LEAST(attempts,9))::int))
+ WHERE id=$1 AND status='sending'`, id, message)
return err
}
+
+func (s *Store) MailQueueStatus(ctx context.Context) (MailQueueStatus, error) {
+ var result MailQueueStatus
+ err := s.db.QueryRow(ctx, `SELECT
+ count(*) FILTER (WHERE status IN ('pending','sending','retry')),
+ count(*) FILTER (WHERE status='dead'),
+ min(created_at) FILTER (WHERE status IN ('pending','sending','retry'))
+ FROM console_mail_outbox`).Scan(&result.Backlog, &result.Failed, &result.OldestPending)
+ return result, err
+}
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go
index 9b76fad..73d869f 100644
--- a/internal/controlplane/queries.go
+++ b/internal/controlplane/queries.go
@@ -173,10 +173,17 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) {
func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
rows, err := s.db.Query(ctx, `
- SELECT id::text, public_id, owned_by, input_price_micros_per_million,
- output_price_micros_per_million, cache_read_price_micros_per_million,
- cache_write_price_micros_per_million, enabled, created_at
- FROM models ORDER BY public_id`)
+ SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by,
+ m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens,
+ m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at, m.retired_at,
+ COALESCE(m.replacement_model,''), pv.id::text, pv.version, pv.currency, pv.effective_from,
+ pv.input_price_micros_per_million, pv.output_price_micros_per_million,
+ pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million,
+ m.enabled, m.created_at
+ FROM models m JOIN LATERAL (
+ SELECT * FROM model_price_versions v WHERE v.model_id=m.id
+ ORDER BY v.effective_from DESC LIMIT 1
+ ) pv ON TRUE ORDER BY m.public_id`)
if err != nil {
return nil, fmt.Errorf("query models: %w", err)
}
@@ -184,13 +191,23 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
positions := make(map[string]int)
for rows.Next() {
var item Model
- if err := rows.Scan(&item.ID, &item.PublicID, &item.OwnedBy, &item.InputPriceMicrosPerMillion,
+ var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte
+ if err := rows.Scan(&item.ID, &item.PublicID, &item.DisplayName, &item.Description, &item.OwnedBy,
+ &inputModalitiesJSON, &outputModalitiesJSON, &item.ContextWindow, &item.MaxOutputTokens,
+ &capabilitiesJSON, &regionsJSON, &item.Lifecycle, &item.ReleasedAt, &item.DeprecatedAt,
+ &item.RetiredAt, &item.ReplacementModel, &item.PriceVersionID, &item.PriceVersion,
+ &item.PriceCurrency, &item.PriceEffectiveFrom, &item.InputPriceMicrosPerMillion,
&item.OutputPriceMicrosPerMillion, &item.CacheReadPriceMicrosPerMillion,
&item.CacheWritePriceMicrosPerMillion, &item.Enabled, &item.CreatedAt); err != nil {
rows.Close()
return nil, fmt.Errorf("scan model: %w", err)
}
+ _ = json.Unmarshal(inputModalitiesJSON, &item.InputModalities)
+ _ = json.Unmarshal(outputModalitiesJSON, &item.OutputModalities)
+ _ = json.Unmarshal(capabilitiesJSON, &item.Capabilities)
+ _ = json.Unmarshal(regionsJSON, &item.Regions)
item.Routes = []Route{}
+ item.Aliases, item.AllowedTenantIDs, item.AllowedKeyIDs = []string{}, []string{}, []string{}
positions[item.ID] = len(models)
models = append(models, item)
}
@@ -200,6 +217,52 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
}
rows.Close()
+ aliasRows, err := s.db.Query(ctx, `SELECT model_id::text, alias FROM model_aliases ORDER BY alias`)
+ if err != nil {
+ return nil, fmt.Errorf("query model aliases: %w", err)
+ }
+ for aliasRows.Next() {
+ var modelID, alias string
+ if err := aliasRows.Scan(&modelID, &alias); err != nil {
+ aliasRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].Aliases = append(models[p].Aliases, alias)
+ }
+ }
+ aliasRows.Close()
+ tenantRows, err := s.db.Query(ctx, `SELECT model_id::text, tenant_id::text FROM model_tenant_allowlist`)
+ if err != nil {
+ return nil, fmt.Errorf("query tenant model allowlist: %w", err)
+ }
+ for tenantRows.Next() {
+ var modelID, tenantID string
+ if err := tenantRows.Scan(&modelID, &tenantID); err != nil {
+ tenantRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].AllowedTenantIDs = append(models[p].AllowedTenantIDs, tenantID)
+ }
+ }
+ tenantRows.Close()
+ keyRows, err := s.db.Query(ctx, `SELECT model_id::text, api_key_id::text FROM api_key_model_allowlist`)
+ if err != nil {
+ return nil, fmt.Errorf("query key model allowlist: %w", err)
+ }
+ for keyRows.Next() {
+ var modelID, keyID string
+ if err := keyRows.Scan(&modelID, &keyID); err != nil {
+ keyRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].AllowedKeyIDs = append(models[p].AllowedKeyIDs, keyID)
+ }
+ }
+ keyRows.Close()
+
routeRows, err := s.db.Query(ctx, `
SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol,
r.upstream_model, r.priority, r.weight, r.enabled
diff --git a/internal/controlplane/retention.go b/internal/controlplane/retention.go
new file mode 100644
index 0000000..4b20611
--- /dev/null
+++ b/internal/controlplane/retention.go
@@ -0,0 +1,57 @@
+package controlplane
+
+import (
+ "context"
+ "log/slog"
+ "time"
+)
+
+func (s *Store) RunRetentionWorker(ctx context.Context, auditDays, securityDays int, logger *slog.Logger) {
+ run := func() {
+ cleanupCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
+ defer cancel()
+ if err := s.pruneExpiredSecurityData(cleanupCtx, auditDays, securityDays); err != nil {
+ logger.Error("retention_cleanup_failed", "error", err)
+ }
+ }
+ run()
+ ticker := time.NewTicker(24 * time.Hour)
+ defer ticker.Stop()
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ run()
+ }
+ }
+}
+
+func (s *Store) pruneExpiredSecurityData(ctx context.Context, auditDays, securityDays int) error {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ auditCutoff := time.Now().UTC().AddDate(0, 0, -auditDays)
+ securityCutoff := time.Now().UTC().AddDate(0, 0, -securityDays)
+ queries := []struct {
+ query string
+ cutoff time.Time
+ }{
+ {`DELETE FROM audit_logs WHERE created_at < $1`, auditCutoff},
+ {`DELETE FROM console_sessions WHERE expires_at < $1 OR (revoked_at IS NOT NULL AND revoked_at < $1)`, securityCutoff},
+ {`DELETE FROM console_action_tokens WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff},
+ {`DELETE FROM console_auth_challenges WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff},
+ {`DELETE FROM console_webauthn_challenges WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff},
+ {`DELETE FROM console_login_throttles WHERE updated_at < $1`, securityCutoff},
+ {`DELETE FROM console_rate_limits WHERE updated_at < $1`, securityCutoff},
+ {`DELETE FROM console_mail_outbox WHERE status='sent' AND sent_at < $1`, securityCutoff},
+ }
+ for _, item := range queries {
+ if _, err := tx.Exec(ctx, item.query, item.cutoff); err != nil {
+ return err
+ }
+ }
+ return tx.Commit(ctx)
+}
diff --git a/internal/controlplane/rotation.go b/internal/controlplane/rotation.go
new file mode 100644
index 0000000..1fc8084
--- /dev/null
+++ b/internal/controlplane/rotation.go
@@ -0,0 +1,69 @@
+package controlplane
+
+import (
+ "context"
+ "fmt"
+)
+
+type encryptedColumn struct{ table, key, column string }
+
+// RotateCredentials re-encrypts every control-plane secret with the primary
+// key in the configured keyring. Run all gateway instances with both new and
+// previous keys before invoking this operation.
+func (s *Store) RotateCredentials(ctx context.Context) (int, error) {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ columns := []encryptedColumn{
+ {"providers", "id", "api_key_ciphertext"},
+ {"console_mail_outbox", "id", "body_ciphertext"},
+ {"console_totp_credentials", "user_id", "secret_ciphertext"},
+ {"console_passkeys", "id", "credential_ciphertext"},
+ {"console_webauthn_challenges", "id", "session_ciphertext"},
+ }
+ total := 0
+ for _, item := range columns {
+ rows, queryErr := tx.Query(ctx, fmt.Sprintf(`SELECT %s::text,%s FROM %s`, item.key, item.column, item.table))
+ if queryErr != nil {
+ return total, queryErr
+ }
+ type record struct {
+ id string
+ ciphertext []byte
+ }
+ records := []record{}
+ for rows.Next() {
+ var value record
+ if err := rows.Scan(&value.id, &value.ciphertext); err != nil {
+ rows.Close()
+ return total, err
+ }
+ records = append(records, value)
+ }
+ if err := rows.Err(); err != nil {
+ rows.Close()
+ return total, err
+ }
+ rows.Close()
+ for _, value := range records {
+ plaintext, err := s.cipher.Decrypt(value.ciphertext)
+ if err != nil {
+ return total, fmt.Errorf("decrypt %s %s: %w", item.table, value.id, err)
+ }
+ ciphertext, err := s.cipher.Encrypt(plaintext)
+ if err != nil {
+ return total, err
+ }
+ if _, err := tx.Exec(ctx, fmt.Sprintf(`UPDATE %s SET %s=$2 WHERE %s=$1`, item.table, item.column, item.key), value.id, ciphertext); err != nil {
+ return total, err
+ }
+ total++
+ }
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return total, err
+ }
+ return total, nil
+}
diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql
index 8a5a605..88db49b 100644
--- a/internal/controlplane/schema.sql
+++ b/internal/controlplane/schema.sql
@@ -1,5 +1,12 @@
CREATE EXTENSION IF NOT EXISTS pgcrypto;
+CREATE TABLE IF NOT EXISTS schema_migrations (
+ version BIGINT PRIMARY KEY,
+ name TEXT NOT NULL,
+ checksum TEXT NOT NULL,
+ applied_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
CREATE TABLE IF NOT EXISTS control_state (
singleton BOOLEAN PRIMARY KEY DEFAULT TRUE CHECK (singleton),
generation BIGINT NOT NULL DEFAULT 0,
@@ -71,6 +78,72 @@ ALTER TABLE models ADD COLUMN IF NOT EXISTS input_price_micros_per_million BIGIN
ALTER TABLE models ADD COLUMN IF NOT EXISTS output_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (output_price_micros_per_million >= 0);
ALTER TABLE models ADD COLUMN IF NOT EXISTS cache_read_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_read_price_micros_per_million >= 0);
ALTER TABLE models ADD COLUMN IF NOT EXISTS cache_write_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_write_price_micros_per_million >= 0);
+ALTER TABLE models ADD COLUMN IF NOT EXISTS display_name TEXT NOT NULL DEFAULT '';
+ALTER TABLE models ADD COLUMN IF NOT EXISTS description TEXT NOT NULL DEFAULT '';
+ALTER TABLE models ADD COLUMN IF NOT EXISTS input_modalities JSONB NOT NULL DEFAULT '["text"]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS output_modalities JSONB NOT NULL DEFAULT '["text"]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS context_window BIGINT NOT NULL DEFAULT 0 CHECK (context_window >= 0);
+ALTER TABLE models ADD COLUMN IF NOT EXISTS max_output_tokens BIGINT NOT NULL DEFAULT 0 CHECK (max_output_tokens >= 0);
+ALTER TABLE models ADD COLUMN IF NOT EXISTS capabilities JSONB NOT NULL DEFAULT '["chat","streaming"]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS regions JSONB NOT NULL DEFAULT '[]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS lifecycle TEXT NOT NULL DEFAULT 'active';
+ALTER TABLE models ADD COLUMN IF NOT EXISTS released_at TIMESTAMPTZ;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS deprecated_at TIMESTAMPTZ;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS retired_at TIMESTAMPTZ;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS replacement_model TEXT;
+ALTER TABLE models DROP CONSTRAINT IF EXISTS models_lifecycle_check;
+ALTER TABLE models ADD CONSTRAINT models_lifecycle_check CHECK (lifecycle IN ('preview','active','deprecated','retired'));
+
+CREATE TABLE IF NOT EXISTS model_price_versions (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ version INTEGER NOT NULL CHECK (version > 0),
+ currency TEXT NOT NULL CHECK (currency = lower(currency) AND length(currency) = 3),
+ input_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (input_price_micros_per_million >= 0),
+ output_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (output_price_micros_per_million >= 0),
+ cache_read_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_read_price_micros_per_million >= 0),
+ cache_write_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_write_price_micros_per_million >= 0),
+ effective_from TIMESTAMPTZ NOT NULL DEFAULT now(),
+ effective_to TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ UNIQUE (model_id, version),
+ CHECK (effective_to IS NULL OR effective_to > effective_from)
+);
+CREATE UNIQUE INDEX IF NOT EXISTS model_price_versions_one_open_idx
+ ON model_price_versions (model_id) WHERE effective_to IS NULL;
+CREATE INDEX IF NOT EXISTS model_price_versions_effective_idx
+ ON model_price_versions (model_id, effective_from DESC);
+
+INSERT INTO model_price_versions (
+ model_id, version, currency, input_price_micros_per_million,
+ output_price_micros_per_million, cache_read_price_micros_per_million,
+ cache_write_price_micros_per_million, effective_from)
+SELECT id, 1, 'usd', input_price_micros_per_million, output_price_micros_per_million,
+ cache_read_price_micros_per_million, cache_write_price_micros_per_million, created_at
+FROM models
+ON CONFLICT (model_id, version) DO NOTHING;
+
+CREATE TABLE IF NOT EXISTS model_aliases (
+ alias TEXT PRIMARY KEY,
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ deprecated BOOLEAN NOT NULL DEFAULT FALSE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+CREATE INDEX IF NOT EXISTS model_aliases_model_idx ON model_aliases (model_id);
+
+CREATE TABLE IF NOT EXISTS model_tenant_allowlist (
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (model_id, tenant_id)
+);
+
+CREATE TABLE IF NOT EXISTS api_key_model_allowlist (
+ api_key_id UUID NOT NULL REFERENCES api_keys(id) ON DELETE CASCADE,
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (api_key_id, model_id)
+);
CREATE TABLE IF NOT EXISTS model_routes (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -106,7 +179,7 @@ CREATE TABLE IF NOT EXISTS billing_reservations (
public_model TEXT NOT NULL,
currency TEXT NOT NULL,
reserved_micros BIGINT NOT NULL CHECK (reserved_micros >= 0),
- status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'settled', 'released')),
+ status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'settled', 'released', 'metering_failed')),
input_price_micros_per_million BIGINT NOT NULL,
output_price_micros_per_million BIGINT NOT NULL,
cache_read_price_micros_per_million BIGINT NOT NULL,
@@ -118,6 +191,30 @@ CREATE TABLE IF NOT EXISTS billing_reservations (
settled_at TIMESTAMPTZ,
FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE
);
+ALTER TABLE billing_reservations ADD COLUMN IF NOT EXISTS price_version_id UUID REFERENCES model_price_versions(id) ON DELETE SET NULL;
+ALTER TABLE billing_reservations DROP CONSTRAINT IF EXISTS billing_reservations_status_check;
+ALTER TABLE billing_reservations ADD CONSTRAINT billing_reservations_status_check
+ CHECK (status IN ('pending','settled','released','metering_failed'));
+
+CREATE TABLE IF NOT EXISTS billing_settlement_jobs (
+ request_id TEXT PRIMARY KEY REFERENCES billing_reservations(request_id) ON DELETE CASCADE,
+ event JSONB,
+ status TEXT NOT NULL DEFAULT 'awaiting_event'
+ CHECK (status IN ('awaiting_event', 'pending', 'processing', 'retry', 'done')),
+ attempts INTEGER NOT NULL DEFAULT 0 CHECK (attempts >= 0),
+ available_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ locked_at TIMESTAMPTZ,
+ last_error TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
+CREATE INDEX IF NOT EXISTS billing_settlement_jobs_ready_idx
+ ON billing_settlement_jobs (available_at, created_at)
+ WHERE status IN ('pending', 'retry');
+CREATE INDEX IF NOT EXISTS billing_settlement_jobs_stale_idx
+ ON billing_settlement_jobs (created_at)
+ WHERE status IN ('awaiting_event', 'processing', 'retry');
CREATE TABLE IF NOT EXISTS usage_events (
request_id TEXT PRIMARY KEY,
@@ -143,9 +240,17 @@ CREATE TABLE IF NOT EXISTS usage_events (
cost_micros BIGINT NOT NULL DEFAULT 0,
charged_micros BIGINT NOT NULL DEFAULT 0,
uncollected_micros BIGINT NOT NULL DEFAULT 0,
+ usage_reported BOOLEAN NOT NULL DEFAULT FALSE,
+ metering_status TEXT NOT NULL DEFAULT 'not_billable'
+ CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')),
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE RESTRICT
);
+ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS usage_reported BOOLEAN NOT NULL DEFAULT FALSE;
+ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS metering_status TEXT NOT NULL DEFAULT 'not_billable';
+ALTER TABLE usage_events DROP CONSTRAINT IF EXISTS usage_events_metering_status_check;
+ALTER TABLE usage_events ADD CONSTRAINT usage_events_metering_status_check
+ CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed'));
-- Usage persistence is independent from billing. Older installations created this
-- foreign key, which prevented recording requests when prepaid billing was disabled.
@@ -165,6 +270,9 @@ CREATE TABLE IF NOT EXISTS billing_ledger (
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
UNIQUE (source_type, source_id)
);
+ALTER TABLE billing_ledger DROP CONSTRAINT IF EXISTS billing_ledger_kind_check;
+ALTER TABLE billing_ledger ADD CONSTRAINT billing_ledger_kind_check
+ CHECK (kind IN ('topup','usage','adjustment','refund','release','dispute','dispute_reversal'));
CREATE TABLE IF NOT EXISTS topup_orders (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -178,13 +286,127 @@ CREATE TABLE IF NOT EXISTS topup_orders (
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
paid_at TIMESTAMPTZ
);
+ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_status_check;
+ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_status_check
+ CHECK (status IN ('pending','paid','failed','expired','partially_refunded','refunded','disputed','reversed'));
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_customer_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_payment_intent_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_charge_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_invoice_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS invoice_url TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS invoice_pdf_url TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS receipt_url TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS refunded_micros BIGINT NOT NULL DEFAULT 0 CHECK (refunded_micros >= 0);
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS disputed_micros BIGINT NOT NULL DEFAULT 0 CHECK (disputed_micros >= 0);
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciliation_status TEXT NOT NULL DEFAULT 'unknown';
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciled_at TIMESTAMPTZ;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciliation_error TEXT NOT NULL DEFAULT '';
+ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_reconciliation_status_check;
+ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_reconciliation_status_check
+ CHECK (reconciliation_status IN ('unknown','ok','repaired','missing','mismatch','resolved'));
+CREATE INDEX IF NOT EXISTS topup_orders_payment_intent_idx ON topup_orders (stripe_payment_intent_id) WHERE stripe_payment_intent_id IS NOT NULL;
+CREATE INDEX IF NOT EXISTS topup_orders_customer_idx ON topup_orders (stripe_customer_id) WHERE stripe_customer_id IS NOT NULL;
+
+CREATE TABLE IF NOT EXISTS billing_reconciliation_resolutions (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ topup_order_id UUID NOT NULL UNIQUE REFERENCES topup_orders(id) ON DELETE RESTRICT,
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT,
+ actor_id TEXT NOT NULL DEFAULT '',
+ actor_type TEXT NOT NULL CHECK (actor_type IN ('console_user','bootstrap','maintenance')),
+ reason TEXT NOT NULL,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS stripe_customers (
+ tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE,
+ stripe_customer_id TEXT NOT NULL UNIQUE,
+ email TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS stripe_refunds (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT,
+ topup_order_id UUID NOT NULL REFERENCES topup_orders(id) ON DELETE RESTRICT,
+ stripe_refund_id TEXT UNIQUE,
+ amount_minor BIGINT NOT NULL CHECK (amount_minor > 0),
+ amount_micros BIGINT NOT NULL CHECK (amount_micros > 0),
+ held_micros BIGINT NOT NULL DEFAULT 0 CHECK (held_micros >= 0),
+ uncollected_micros BIGINT NOT NULL DEFAULT 0 CHECK (uncollected_micros >= 0),
+ currency TEXT NOT NULL,
+ reason TEXT NOT NULL DEFAULT 'requested_by_customer',
+ status TEXT NOT NULL DEFAULT 'queued'
+ CHECK (status IN ('queued','submitting','pending','requires_action','succeeded','failed','canceled')),
+ attempts INTEGER NOT NULL DEFAULT 0,
+ available_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ last_error TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
+CREATE INDEX IF NOT EXISTS stripe_refunds_ready_idx ON stripe_refunds (available_at, created_at)
+ WHERE status IN ('queued','submitting');
+
+CREATE TABLE IF NOT EXISTS stripe_disputes (
+ stripe_dispute_id TEXT PRIMARY KEY,
+ tenant_id UUID REFERENCES tenants(id) ON DELETE SET NULL,
+ topup_order_id UUID REFERENCES topup_orders(id) ON DELETE SET NULL,
+ stripe_payment_intent_id TEXT,
+ amount_minor BIGINT NOT NULL CHECK (amount_minor >= 0),
+ amount_micros BIGINT NOT NULL CHECK (amount_micros >= 0),
+ currency TEXT NOT NULL,
+ status TEXT NOT NULL,
+ reason TEXT NOT NULL DEFAULT '',
+ debited_micros BIGINT NOT NULL DEFAULT 0 CHECK (debited_micros >= 0),
+ uncollected_micros BIGINT NOT NULL DEFAULT 0 CHECK (uncollected_micros >= 0),
+ due_by TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ closed_at TIMESTAMPTZ
+);
+
+CREATE TABLE IF NOT EXISTS stripe_invoices (
+ stripe_invoice_id TEXT PRIMARY KEY,
+ tenant_id UUID REFERENCES tenants(id) ON DELETE SET NULL,
+ topup_order_id UUID REFERENCES topup_orders(id) ON DELETE SET NULL,
+ stripe_customer_id TEXT,
+ status TEXT NOT NULL DEFAULT '',
+ currency TEXT NOT NULL DEFAULT '',
+ amount_due_minor BIGINT NOT NULL DEFAULT 0,
+ amount_paid_minor BIGINT NOT NULL DEFAULT 0,
+ attempt_count INTEGER NOT NULL DEFAULT 0,
+ next_payment_attempt TIMESTAMPTZ,
+ hosted_invoice_url TEXT NOT NULL DEFAULT '',
+ invoice_pdf_url TEXT NOT NULL DEFAULT '',
+ last_failure TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS billing_reconciliation_runs (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ status TEXT NOT NULL CHECK (status IN ('running','clean','mismatch','failed')),
+ checked_orders BIGINT NOT NULL DEFAULT 0,
+ mismatch_count BIGINT NOT NULL DEFAULT 0,
+ report JSONB NOT NULL DEFAULT '[]'::jsonb,
+ error TEXT NOT NULL DEFAULT '',
+ started_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
CREATE TABLE IF NOT EXISTS stripe_webhook_events (
event_id TEXT PRIMARY KEY,
event_type TEXT NOT NULL,
processed_at TIMESTAMPTZ,
+ attempts INTEGER NOT NULL DEFAULT 0,
+ last_attempt_at TIMESTAMPTZ,
+ processing_error TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
+ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS attempts INTEGER NOT NULL DEFAULT 0;
+ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS last_attempt_at TIMESTAMPTZ;
+ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS processing_error TEXT NOT NULL DEFAULT '';
CREATE TABLE IF NOT EXISTS console_users (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -279,7 +501,7 @@ CREATE INDEX IF NOT EXISTS console_action_tokens_active_idx
CREATE TABLE IF NOT EXISTS console_mail_outbox (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
recipient TEXT NOT NULL,
- template TEXT NOT NULL CHECK (template IN ('verify_email', 'password_reset', 'invite')),
+ template TEXT NOT NULL CHECK (template IN ('verify_email', 'password_reset', 'invite', 'low_balance', 'spend_anomaly')),
subject TEXT NOT NULL,
body_ciphertext BYTEA NOT NULL,
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'sending', 'sent', 'failed')),
@@ -290,8 +512,43 @@ CREATE TABLE IF NOT EXISTS console_mail_outbox (
last_error TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
-CREATE INDEX IF NOT EXISTS console_mail_outbox_pending_idx
- ON console_mail_outbox (available_at, created_at) WHERE status IN ('pending', 'failed');
+ALTER TABLE console_mail_outbox DROP CONSTRAINT IF EXISTS console_mail_outbox_template_check;
+ALTER TABLE console_mail_outbox ADD CONSTRAINT console_mail_outbox_template_check
+ CHECK (template IN ('verify_email','password_reset','invite','low_balance','spend_anomaly'));
+ALTER TABLE console_mail_outbox DROP CONSTRAINT IF EXISTS console_mail_outbox_status_check;
+UPDATE console_mail_outbox SET status='retry' WHERE status='failed';
+ALTER TABLE console_mail_outbox ADD CONSTRAINT console_mail_outbox_status_check
+ CHECK (status IN ('pending','sending','sent','retry','dead','suppressed'));
+DROP INDEX IF EXISTS console_mail_outbox_pending_idx;
+CREATE INDEX console_mail_outbox_pending_idx
+ ON console_mail_outbox (available_at, created_at) WHERE status IN ('pending', 'retry');
+
+CREATE TABLE IF NOT EXISTS mail_suppressions (
+ recipient TEXT PRIMARY KEY,
+ reason TEXT NOT NULL CHECK (reason IN ('bounce','complaint','manual')),
+ provider TEXT NOT NULL DEFAULT '',
+ provider_event_id TEXT NOT NULL DEFAULT '',
+ detail TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS mail_feedback_events (
+ event_id TEXT PRIMARY KEY,
+ event_type TEXT NOT NULL CHECK (event_type IN ('delivered','bounce','complaint')),
+ recipient TEXT NOT NULL,
+ provider TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS mail_notification_events (
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ recipient TEXT NOT NULL,
+ notification_type TEXT NOT NULL CHECK (notification_type IN ('low_balance','spend_anomaly')),
+ dedupe_key TEXT NOT NULL,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (tenant_id,recipient,notification_type,dedupe_key)
+);
CREATE TABLE IF NOT EXISTS console_auth_challenges (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go
index bc04e7c..b9bddaf 100644
--- a/internal/controlplane/snapshot.go
+++ b/internal/controlplane/snapshot.go
@@ -6,6 +6,7 @@ import (
"encoding/json"
"fmt"
"strings"
+ "time"
"aigw/internal/auth"
"aigw/internal/domain"
@@ -102,13 +103,23 @@ func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain
func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) {
rows, err := tx.Query(ctx, `
- SELECT m.public_id, m.owned_by, m.input_price_micros_per_million, m.output_price_micros_per_million,
- m.cache_read_price_micros_per_million, m.cache_write_price_micros_per_million,
+ SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by,
+ m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens,
+ m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at,
+ m.retired_at, COALESCE(m.replacement_model,''),
+ pv.id::text, pv.version, pv.currency, pv.effective_from,
+ pv.input_price_micros_per_million, pv.output_price_micros_per_million,
+ pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million,
r.provider_id::text, r.upstream_model, r.priority, r.weight
FROM models m
+ JOIN LATERAL (
+ SELECT * FROM model_price_versions v WHERE v.model_id=m.id
+ AND v.effective_from <= now() AND (v.effective_to IS NULL OR v.effective_to > now())
+ ORDER BY v.effective_from DESC LIMIT 1
+ ) pv ON TRUE
JOIN model_routes r ON r.model_id = m.id AND r.enabled = TRUE
JOIN providers p ON p.id = r.provider_id AND p.enabled = TRUE
- WHERE m.enabled = TRUE
+ WHERE m.enabled = TRUE AND m.lifecycle <> 'retired'
ORDER BY m.public_id, r.priority, r.created_at`)
if err != nil {
return nil, fmt.Errorf("query model routes: %w", err)
@@ -117,10 +128,21 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov
models := make([]domain.Model, 0)
index := make(map[string]int)
for rows.Next() {
- var publicID, ownedBy, providerID, upstreamModel string
+ var modelID, publicID, displayName, description, ownedBy, replacement, providerID, upstreamModel string
+ var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte
+ var contextWindow, maxOutput int64
+ var lifecycle, priceVersionID, priceCurrency string
+ var releasedAt, deprecatedAt, retiredAt *time.Time
+ var priceVersion int
+ var priceEffectiveFrom time.Time
var inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64
var priority, weight int
- if err := rows.Scan(&publicID, &ownedBy, &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice, &providerID, &upstreamModel, &priority, &weight); err != nil {
+ if err := rows.Scan(&modelID, &publicID, &displayName, &description, &ownedBy,
+ &inputModalitiesJSON, &outputModalitiesJSON, &contextWindow, &maxOutput,
+ &capabilitiesJSON, &regionsJSON, &lifecycle, &releasedAt, &deprecatedAt, &retiredAt, &replacement,
+ &priceVersionID, &priceVersion, &priceCurrency, &priceEffectiveFrom,
+ &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice,
+ &providerID, &upstreamModel, &priority, &weight); err != nil {
return nil, fmt.Errorf("scan model route: %w", err)
}
provider, ok := providers[providerID]
@@ -129,13 +151,32 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov
}
position, exists := index[publicID]
if !exists {
+ var inputModalities, outputModalities, capabilities, regions []string
+ if err := json.Unmarshal(inputModalitiesJSON, &inputModalities); err != nil {
+ return nil, fmt.Errorf("decode model input modalities: %w", err)
+ }
+ if err := json.Unmarshal(outputModalitiesJSON, &outputModalities); err != nil {
+ return nil, fmt.Errorf("decode model output modalities: %w", err)
+ }
+ if err := json.Unmarshal(capabilitiesJSON, &capabilities); err != nil {
+ return nil, fmt.Errorf("decode model capabilities: %w", err)
+ }
+ if err := json.Unmarshal(regionsJSON, &regions); err != nil {
+ return nil, fmt.Errorf("decode model regions: %w", err)
+ }
position = len(models)
index[publicID] = position
models = append(models, domain.Model{
- ID: publicID, OwnedBy: ownedBy,
+ ID: publicID, DisplayName: displayName, Description: description, OwnedBy: ownedBy,
+ InputModalities: inputModalities, OutputModalities: outputModalities,
+ ContextWindow: contextWindow, MaxOutputTokens: maxOutput, Capabilities: capabilities,
+ Regions: regions, Lifecycle: lifecycle, ReleasedAt: releasedAt, DeprecatedAt: deprecatedAt,
+ RetiredAt: retiredAt, ReplacementModel: replacement, PriceVersionID: priceVersionID,
+ PriceVersion: priceVersion, PriceCurrency: priceCurrency, PriceEffectiveFrom: priceEffectiveFrom,
InputPriceMicrosPerMillion: inputPrice, OutputPriceMicrosPerMillion: outputPrice,
CacheReadPriceMicrosPerMillion: cacheReadPrice, CacheWritePriceMicrosPerMillion: cacheWritePrice,
})
+ _ = modelID
}
models[position].Routes = append(models[position].Routes, domain.Route{
Provider: provider, UpstreamModel: upstreamModel, Priority: priority, Weight: weight,
@@ -144,6 +185,61 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("read model routes: %w", err)
}
+ byID := make(map[string]int, len(models))
+ for i := range models {
+ byID[models[i].ID] = i
+ }
+ aliasRows, err := tx.Query(ctx, `SELECT a.alias, m.public_id FROM model_aliases a JOIN models m ON m.id=a.model_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query model aliases: %w", err)
+ }
+ for aliasRows.Next() {
+ var alias, publicID string
+ if err := aliasRows.Scan(&alias, &publicID); err != nil {
+ aliasRows.Close()
+ return nil, err
+ }
+ if position, ok := byID[publicID]; ok {
+ models[position].Aliases = append(models[position].Aliases, alias)
+ }
+ }
+ aliasRows.Close()
+ tenantRows, err := tx.Query(ctx, `SELECT m.public_id, a.tenant_id::text FROM model_tenant_allowlist a JOIN models m ON m.id=a.model_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query model tenant allowlist: %w", err)
+ }
+ for tenantRows.Next() {
+ var publicID, tenantID string
+ if err := tenantRows.Scan(&publicID, &tenantID); err != nil {
+ tenantRows.Close()
+ return nil, err
+ }
+ if position, ok := byID[publicID]; ok {
+ if models[position].AllowedTenantIDs == nil {
+ models[position].AllowedTenantIDs = map[string]struct{}{}
+ }
+ models[position].AllowedTenantIDs[tenantID] = struct{}{}
+ }
+ }
+ tenantRows.Close()
+ keyRows, err := tx.Query(ctx, `SELECT m.public_id, a.api_key_id::text FROM api_key_model_allowlist a JOIN models m ON m.id=a.model_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query model key allowlist: %w", err)
+ }
+ for keyRows.Next() {
+ var publicID, keyID string
+ if err := keyRows.Scan(&publicID, &keyID); err != nil {
+ keyRows.Close()
+ return nil, err
+ }
+ if position, ok := byID[publicID]; ok {
+ if models[position].AllowedKeyIDs == nil {
+ models[position].AllowedKeyIDs = map[string]struct{}{}
+ }
+ models[position].AllowedKeyIDs[keyID] = struct{}{}
+ }
+ }
+ keyRows.Close()
return models, nil
}
diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go
index 833f846..fb48d78 100644
--- a/internal/controlplane/store.go
+++ b/internal/controlplane/store.go
@@ -2,7 +2,9 @@ package controlplane
import (
"context"
+ "crypto/sha256"
_ "embed"
+ "encoding/hex"
"encoding/json"
"errors"
"fmt"
@@ -11,6 +13,7 @@ import (
"aigw/internal/security"
+ "github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/redis/go-redis/v9"
)
@@ -20,12 +23,15 @@ var schemaSQL string
var ErrRedisDisabled = errors.New("Redis propagation is disabled")
+const migrationVersion int64 = 2026080504
+
type Options struct {
- DatabaseURL string
- RedisURL string
- CredentialKey string
- RedisChannel string
- VersionCacheKey string
+ DatabaseURL string
+ RedisURL string
+ CredentialKey string
+ PreviousCredentialKeys []string
+ RedisChannel string
+ VersionCacheKey string
}
type Store struct {
@@ -37,7 +43,7 @@ type Store struct {
}
func NewStore(ctx context.Context, options Options) (*Store, error) {
- cipher, err := security.NewCredentialCipher(options.CredentialKey)
+ cipher, err := security.NewCredentialKeyring(options.CredentialKey, options.PreviousCredentialKeys)
if err != nil {
return nil, err
}
@@ -76,11 +82,10 @@ func (s *Store) RedisEnabled() bool {
return s.redis != nil
}
+func (s *Store) Ping(ctx context.Context) error { return s.db.Ping(ctx) }
+
func (s *Store) Migrate(ctx context.Context) error {
- if _, err := s.db.Exec(ctx, schemaSQL); err != nil {
- return fmt.Errorf("apply control-plane schema: %w", err)
- }
- return nil
+ return applySchema(ctx, s.db)
}
func MigrateDatabase(ctx context.Context, databaseURL string) error {
@@ -89,12 +94,67 @@ func MigrateDatabase(ctx context.Context, databaseURL string) error {
return fmt.Errorf("configure PostgreSQL: %w", err)
}
defer db.Close()
- if _, err := db.Exec(ctx, schemaSQL); err != nil {
+ return applySchema(ctx, db)
+}
+
+func MigrationStatusDatabase(ctx context.Context, databaseURL string) (MigrationStatus, error) {
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ return MigrationStatus{}, fmt.Errorf("configure PostgreSQL: %w", err)
+ }
+ defer db.Close()
+ var result MigrationStatus
+ err = db.QueryRow(ctx, `SELECT version,name,checksum,applied_at FROM schema_migrations ORDER BY version DESC LIMIT 1`).Scan(&result.Version, &result.Name, &result.Checksum, &result.AppliedAt)
+ if err != nil {
+ return MigrationStatus{}, fmt.Errorf("read migration status: %w", err)
+ }
+ return result, nil
+}
+
+func applySchema(ctx context.Context, db *pgxpool.Pool) error {
+ tx, err := db.Begin(ctx)
+ if err != nil {
+ return fmt.Errorf("begin migration: %w", err)
+ }
+ defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, migrationVersion); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, schemaSQL); err != nil {
return fmt.Errorf("apply control-plane schema: %w", err)
}
+ hash := sha256.Sum256([]byte(schemaSQL))
+ checksum := hex.EncodeToString(hash[:])
+ var existing string
+ err = tx.QueryRow(ctx, `SELECT checksum FROM schema_migrations WHERE version=$1`, migrationVersion).Scan(&existing)
+ if err == nil && existing != checksum {
+ return fmt.Errorf("migration %d checksum changed; deploy an explicit new migration version", migrationVersion)
+ }
+ if !errors.Is(err, pgx.ErrNoRows) && err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "commercial-control-plane", checksum); err != nil {
+ return err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return fmt.Errorf("commit migration: %w", err)
+ }
return nil
}
+type MigrationStatus struct {
+ Version int64 `json:"version"`
+ Name string `json:"name"`
+ Checksum string `json:"checksum"`
+ AppliedAt time.Time `json:"applied_at"`
+}
+
+func (s *Store) MigrationStatus(ctx context.Context) (MigrationStatus, error) {
+ var result MigrationStatus
+ err := s.db.QueryRow(ctx, `SELECT version,name,checksum,applied_at FROM schema_migrations ORDER BY version DESC LIMIT 1`).Scan(&result.Version, &result.Name, &result.Checksum, &result.AppliedAt)
+ return result, err
+}
+
func (s *Store) DatabaseGeneration(ctx context.Context) (int64, error) {
var generation int64
err := s.db.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&generation)
diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go
new file mode 100644
index 0000000..4f55e10
--- /dev/null
+++ b/internal/controlplane/store_integration_test.go
@@ -0,0 +1,66 @@
+package controlplane
+
+import (
+ "context"
+ "fmt"
+ "net/url"
+ "os"
+ "testing"
+ "time"
+
+ "github.com/jackc/pgx/v5/pgxpool"
+)
+
+func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) {
+ databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL")
+ if databaseURL == "" {
+ t.Skip("AIGW_TEST_DATABASE_URL is not set")
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer db.Close()
+ schema := fmt.Sprintf("migration_drill_%d", time.Now().UnixNano())
+ if _, err := db.Exec(ctx, "CREATE SCHEMA "+schema); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() { _, _ = db.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") })
+ if _, err := db.Exec(ctx, `CREATE TABLE `+schema+`.schema_migrations (
+ version BIGINT PRIMARY KEY,name TEXT NOT NULL,checksum TEXT NOT NULL,applied_at TIMESTAMPTZ NOT NULL DEFAULT now())`); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := db.Exec(ctx, `INSERT INTO `+schema+`.schema_migrations(version,name,checksum) VALUES ($1,'previous-release','immutable-previous-checksum')`, migrationVersion-1); err != nil {
+ t.Fatal(err)
+ }
+ parsed, err := url.Parse(databaseURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ query := parsed.Query()
+ query.Set("search_path", schema)
+ parsed.RawQuery = query.Encode()
+ isolatedURL := parsed.String()
+ if err := MigrateDatabase(ctx, isolatedURL); err != nil {
+ t.Fatal(err)
+ }
+ if err := MigrateDatabase(ctx, isolatedURL); err != nil {
+ t.Fatalf("second migration must be idempotent: %v", err)
+ }
+ status, err := MigrationStatusDatabase(ctx, isolatedURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if status.Version != migrationVersion {
+ t.Fatalf("migration version = %d, want %d", status.Version, migrationVersion)
+ }
+ var count int
+ if err := db.QueryRow(ctx, `SELECT count(*) FROM `+schema+`.schema_migrations`).Scan(&count); err != nil {
+ t.Fatal(err)
+ }
+ if count != 2 {
+ t.Fatalf("migration history contains %d rows, want previous and current", count)
+ }
+}
diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go
index f81e63b..c24dde9 100644
--- a/internal/controlplane/types.go
+++ b/internal/controlplane/types.go
@@ -7,6 +7,12 @@ import (
"aigw/internal/domain"
)
+type MailQueueStatus struct {
+ Backlog int64 `json:"backlog"`
+ Failed int64 `json:"failed"`
+ OldestPending *time.Time `json:"oldest_pending,omitempty"`
+}
+
type Tenant struct {
ID string `json:"id"`
Slug string `json:"slug"`
@@ -62,16 +68,36 @@ type Route struct {
}
type Model struct {
- ID string `json:"id"`
- PublicID string `json:"public_id"`
- OwnedBy string `json:"owned_by"`
- InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
- OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
- CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
- CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
- Enabled bool `json:"enabled"`
- Routes []Route `json:"routes"`
- CreatedAt time.Time `json:"created_at"`
+ ID string `json:"id"`
+ PublicID string `json:"public_id"`
+ DisplayName string `json:"display_name"`
+ Description string `json:"description"`
+ OwnedBy string `json:"owned_by"`
+ InputModalities []string `json:"input_modalities"`
+ OutputModalities []string `json:"output_modalities"`
+ ContextWindow int64 `json:"context_window"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Capabilities []string `json:"capabilities"`
+ Regions []string `json:"regions"`
+ Lifecycle string `json:"lifecycle"`
+ ReleasedAt *time.Time `json:"released_at,omitempty"`
+ DeprecatedAt *time.Time `json:"deprecated_at,omitempty"`
+ RetiredAt *time.Time `json:"retired_at,omitempty"`
+ ReplacementModel string `json:"replacement_model,omitempty"`
+ Aliases []string `json:"aliases"`
+ AllowedTenantIDs []string `json:"allowed_tenant_ids"`
+ AllowedKeyIDs []string `json:"allowed_key_ids"`
+ PriceVersionID string `json:"price_version_id"`
+ PriceVersion int `json:"price_version"`
+ PriceCurrency string `json:"price_currency"`
+ PriceEffectiveFrom time.Time `json:"price_effective_from"`
+ InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
+ OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
+ CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
+ CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
+ Enabled bool `json:"enabled"`
+ Routes []Route `json:"routes"`
+ CreatedAt time.Time `json:"created_at"`
}
type Overview struct {
@@ -141,7 +167,24 @@ type RouteInput struct {
type CreateModelInput struct {
PublicID string `json:"public_id"`
+ DisplayName string `json:"display_name"`
+ Description string `json:"description"`
OwnedBy string `json:"owned_by"`
+ InputModalities []string `json:"input_modalities"`
+ OutputModalities []string `json:"output_modalities"`
+ ContextWindow int64 `json:"context_window"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Capabilities []string `json:"capabilities"`
+ Regions []string `json:"regions"`
+ Lifecycle string `json:"lifecycle"`
+ ReleasedAt *time.Time `json:"released_at"`
+ DeprecatedAt *time.Time `json:"deprecated_at"`
+ RetiredAt *time.Time `json:"retired_at"`
+ ReplacementModel string `json:"replacement_model"`
+ Aliases []string `json:"aliases"`
+ AllowedTenantIDs []string `json:"allowed_tenant_ids"`
+ AllowedKeyIDs []string `json:"allowed_key_ids"`
+ PriceCurrency string `json:"price_currency"`
InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
@@ -149,6 +192,15 @@ type CreateModelInput struct {
Routes []RouteInput `json:"routes"`
}
+type CreatePriceVersionInput struct {
+ Currency string `json:"currency"`
+ InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
+ OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
+ CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
+ CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
+ EffectiveFrom time.Time `json:"effective_from"`
+}
+
type ConsoleActor struct {
ID string `json:"id,omitempty"`
TenantID string `json:"tenant_id,omitempty"`
diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go
index 6c69a9f..b436d8c 100644
--- a/internal/controlplane/usage.go
+++ b/internal/controlplane/usage.go
@@ -29,13 +29,13 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error
request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model,
protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms,
input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens,
- cost_micros, charged_micros, uncollected_micros)
- VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0)
+ cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status)
+ VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0,$21,$22)
ON CONFLICT (request_id) DO NOTHING`, event.RequestID, event.TenantID, event.ProjectID, event.KeyID,
event.PublicModel, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream,
event.StatusCode, event.Success, event.ErrorType, event.Attempts, event.StartedAt, event.DurationMS,
event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens,
- event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens)
+ event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, event.UsageReported, usageMeteringStatus(event))
if err != nil {
return fmt.Errorf("persist usage event: %w", err)
}
@@ -51,6 +51,16 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error
return nil
}
+func usageMeteringStatus(event domain.UsageEvent) string {
+ if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success {
+ return "upstream_failed"
+ }
+ if event.UsageReported {
+ return "reported"
+ }
+ return "missing"
+}
+
func upsertUsageRollup(ctx context.Context, tx pgx.Tx, event domain.UsageEvent, cost, charged, uncollected int64) error {
period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC)
_, err := tx.Exec(ctx, `
diff --git a/internal/domain/types.go b/internal/domain/types.go
index b1decc9..94e56c1 100644
--- a/internal/domain/types.go
+++ b/internal/domain/types.go
@@ -32,7 +32,27 @@ type Route struct {
type Model struct {
ID string
+ DisplayName string
+ Description string
OwnedBy string
+ InputModalities []string
+ OutputModalities []string
+ ContextWindow int64
+ MaxOutputTokens int64
+ Capabilities []string
+ Regions []string
+ Lifecycle string
+ ReleasedAt *time.Time
+ DeprecatedAt *time.Time
+ RetiredAt *time.Time
+ ReplacementModel string
+ Aliases []string
+ AllowedTenantIDs map[string]struct{}
+ AllowedKeyIDs map[string]struct{}
+ PriceVersionID string
+ PriceVersion int
+ PriceCurrency string
+ PriceEffectiveFrom time.Time
InputPriceMicrosPerMillion int64
OutputPriceMicrosPerMillion int64
CacheReadPriceMicrosPerMillion int64
@@ -40,6 +60,20 @@ type Model struct {
Routes []Route
}
+func (m Model) Allows(principal Principal) bool {
+ if len(m.AllowedTenantIDs) > 0 {
+ if _, ok := m.AllowedTenantIDs[principal.TenantID]; !ok {
+ return false
+ }
+ }
+ if len(m.AllowedKeyIDs) > 0 {
+ if _, ok := m.AllowedKeyIDs[principal.KeyID]; !ok {
+ return false
+ }
+ }
+ return true
+}
+
type Usage struct {
InputTokens int64 `json:"input_tokens,omitempty"`
OutputTokens int64 `json:"output_tokens,omitempty"`
@@ -65,6 +99,7 @@ type UsageEvent struct {
StartedAt time.Time `json:"started_at"`
DurationMS int64 `json:"duration_ms"`
Usage Usage `json:"usage"`
+ UsageReported bool `json:"usage_reported"`
}
// LimitPolicy is the immutable runtime view of a project's commercial limits.
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go
index becfbf3..6ab437c 100644
--- a/internal/httpapi/api.go
+++ b/internal/httpapi/api.go
@@ -28,54 +28,57 @@ import (
type requestIDKey struct{}
+type UsageRecorder interface {
+ RecordUsage(context.Context, domain.UsageEvent) error
+}
+
type API struct {
- authenticator auth.Authenticator
- catalog *catalog.Catalog
- router *routing.Router
- forwarder *provider.Forwarder
- usageSink telemetry.UsageSink
- billingMeter billing.Meter
- limiter *limits.Limiter
- usageRecorder interface {
- RecordUsage(context.Context, domain.UsageEvent) error
- }
- metrics *telemetry.Metrics
- logger *slog.Logger
- maxBodyBytes int64
- exposeMetrics bool
+ authenticator auth.Authenticator
+ catalog *catalog.Catalog
+ router *routing.Router
+ forwarder *provider.Forwarder
+ usageSink telemetry.UsageSink
+ billingMeter billing.Meter
+ limiter *limits.Limiter
+ usageRecorder UsageRecorder
+ metrics *telemetry.Metrics
+ logger *slog.Logger
+ maxBodyBytes int64
+ exposeMetrics bool
+ deploymentRegion string
}
type Options struct {
- Authenticator auth.Authenticator
- Catalog *catalog.Catalog
- Router *routing.Router
- Forwarder *provider.Forwarder
- UsageSink telemetry.UsageSink
- BillingMeter billing.Meter
- Limiter *limits.Limiter
- UsageRecorder interface {
- RecordUsage(context.Context, domain.UsageEvent) error
- }
- Metrics *telemetry.Metrics
- Logger *slog.Logger
- MaxBodyBytes int64
- ExposeMetrics bool
+ Authenticator auth.Authenticator
+ Catalog *catalog.Catalog
+ Router *routing.Router
+ Forwarder *provider.Forwarder
+ UsageSink telemetry.UsageSink
+ BillingMeter billing.Meter
+ Limiter *limits.Limiter
+ UsageRecorder UsageRecorder
+ Metrics *telemetry.Metrics
+ Logger *slog.Logger
+ MaxBodyBytes int64
+ ExposeMetrics bool
+ DeploymentRegion string
}
func New(options Options) *API {
return &API{
- authenticator: options.Authenticator,
- catalog: options.Catalog,
- router: options.Router,
- forwarder: options.Forwarder,
- usageSink: options.UsageSink,
- billingMeter: options.BillingMeter,
- limiter: options.Limiter,
- usageRecorder: options.UsageRecorder,
- metrics: options.Metrics,
- logger: options.Logger,
- maxBodyBytes: options.MaxBodyBytes,
- exposeMetrics: options.ExposeMetrics,
+ authenticator: options.Authenticator,
+ catalog: options.Catalog,
+ router: options.Router,
+ forwarder: options.Forwarder,
+ usageSink: options.UsageSink,
+ billingMeter: options.BillingMeter,
+ limiter: options.Limiter,
+ usageRecorder: options.UsageRecorder,
+ metrics: options.Metrics,
+ logger: options.Logger,
+ maxBodyBytes: options.MaxBodyBytes,
+ exposeMetrics: options.ExposeMetrics,
+ deploymentRegion: strings.ToLower(strings.TrimSpace(options.DeploymentRegion)),
}
}
@@ -87,6 +90,17 @@ func (a *API) Handler() http.Handler {
mux.Handle("GET /metrics", a.metrics)
}
+ a.registerInference(mux)
+ return a.withRequestID(a.recoverPanics(mux))
+}
+
+func (a *API) InferenceHandler() http.Handler {
+ mux := http.NewServeMux()
+ a.registerInference(mux)
+ return a.withRequestID(a.recoverPanics(mux))
+}
+
+func (a *API) registerInference(mux *http.ServeMux) {
mux.HandleFunc("GET /v1/models", a.openAIModels)
mux.HandleFunc("GET /api/v1/models", a.openAIModels)
mux.HandleFunc("POST /v1/chat/completions", a.openAIChat)
@@ -97,7 +111,6 @@ func (a *API) Handler() http.Handler {
mux.HandleFunc("POST /anthropic/v1/messages", a.anthropicMessages)
mux.HandleFunc("POST /api/anthropic/v1/messages", a.anthropicMessages)
- return a.withRequestID(a.recoverPanics(mux))
}
func (a *API) health(w http.ResponseWriter, _ *http.Request) {
@@ -143,8 +156,13 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
}
var envelope struct {
- Model string `json:"model"`
- Stream bool `json:"stream"`
+ Model string `json:"model"`
+ Stream bool `json:"stream"`
+ MaxTokens int64 `json:"max_tokens"`
+ MaxCompletionTokens int64 `json:"max_completion_tokens"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Tools json.RawMessage `json:"tools"`
+ ResponseFormat json.RawMessage `json:"response_format"`
}
if err := json.Unmarshal(body, &envelope); err != nil || strings.TrimSpace(envelope.Model) == "" {
apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Parameter model is required and the body must be valid JSON"}, requestID)
@@ -170,7 +188,30 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
defer lease.Release()
}
- routes, err := a.router.Plan(envelope.Model, protocol)
+ model, modelErr := a.catalog.ModelForPrincipal(envelope.Model, principal)
+ if modelErr != nil || !a.modelAvailableInRegion(model) {
+ apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist or is not allowed for this API key"}, requestID)
+ return
+ }
+ requestedOutput := max(envelope.MaxTokens, envelope.MaxCompletionTokens, envelope.MaxOutputTokens)
+ if model.MaxOutputTokens > 0 && requestedOutput > model.MaxOutputTokens {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Requested output exceeds the model maximum"}, requestID)
+ return
+ }
+ if envelope.Stream && !modelHasCapability(model, "streaming") {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support streaming"}, requestID)
+ return
+ }
+ if len(envelope.Tools) > 0 && string(envelope.Tools) != "null" && !modelHasCapability(model, "tools") {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support tools"}, requestID)
+ return
+ }
+ if len(envelope.ResponseFormat) > 0 && string(envelope.ResponseFormat) != "null" && !modelHasCapability(model, "json") {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support structured JSON output"}, requestID)
+ return
+ }
+
+ routes, err := a.router.Plan(model.ID, protocol)
if err != nil {
if errors.Is(err, routing.ErrNoRoute) {
apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "model_not_supported", Message: "Model does not support this API protocol"}, requestID)
@@ -180,11 +221,6 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
return
}
if a.billingMeter != nil {
- model, modelErr := a.catalog.Model(envelope.Model)
- if modelErr != nil {
- apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist"}, requestID)
- return
- }
policy := domain.LimitPolicy{}
if a.limiter != nil {
policy, _ = a.limiter.Policy(principal.ProjectID)
@@ -257,6 +293,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
observer := usage.NewObserver(protocol, stream)
copyErr := copyResponse(w, result.Response.Body, observer, stream)
usageResult := observer.Usage()
+ usageReported := observer.Reported()
success = copyErr == nil
errorType := ""
if copyErr != nil {
@@ -267,6 +304,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
PublicModel: envelope.Model, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel,
Protocol: protocol, Stream: stream, StatusCode: result.Response.StatusCode, Success: success,
ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), Usage: usageResult,
+ UsageReported: usageReported,
})
a.logger.Info("inference_request",
"request_id", requestID,
@@ -281,22 +319,27 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do
}
func (a *API) openAIModels(w http.ResponseWriter, r *http.Request) {
- if !a.authorize(w, r) {
+ principal, ok := a.authorize(w, r)
+ if !ok {
return
}
- models := a.catalog.Models(domain.ProtocolOpenAI)
+ models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal))
data := make([]map[string]any, 0, len(models))
for _, model := range models {
- data = append(data, map[string]any{"id": model.ID, "object": "model", "created": 0, "owned_by": model.OwnedBy})
+ data = append(data, map[string]any{"id": model.ID, "object": "model", "created": modelCreated(model), "owned_by": model.OwnedBy,
+ "display_name": model.DisplayName, "context_window": model.ContextWindow, "max_output_tokens": model.MaxOutputTokens,
+ "input_modalities": model.InputModalities, "output_modalities": model.OutputModalities, "capabilities": model.Capabilities,
+ "lifecycle": model.Lifecycle, "regions": model.Regions, "replacement_model": model.ReplacementModel})
}
writeJSON(w, map[string]any{"object": "list", "data": data})
}
func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) {
- if !a.authorize(w, r) {
+ principal, ok := a.authorize(w, r)
+ if !ok {
return
}
- models := a.catalog.Models(domain.ProtocolAnthropic)
+ models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolAnthropic, principal))
data := make([]map[string]any, 0, len(models))
for _, model := range models {
data = append(data, map[string]any{"id": model.ID, "display_name": model.ID, "created_at": "1970-01-01T00:00:00Z", "type": "model"})
@@ -309,13 +352,52 @@ func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) {
writeJSON(w, response)
}
-func (a *API) authorize(w http.ResponseWriter, r *http.Request) bool {
+func (a *API) modelAvailableInRegion(model domain.Model) bool {
+ if a.deploymentRegion == "" || len(model.Regions) == 0 {
+ return true
+ }
+ for _, region := range model.Regions {
+ if strings.EqualFold(region, a.deploymentRegion) || region == "*" {
+ return true
+ }
+ }
+ return false
+}
+
+func (a *API) availableModels(models []domain.Model) []domain.Model {
+ result := models[:0]
+ for _, model := range models {
+ if a.modelAvailableInRegion(model) {
+ result = append(result, model)
+ }
+ }
+ return result
+}
+func modelHasCapability(model domain.Model, wanted string) bool {
+ if len(model.Capabilities) == 0 {
+ return true
+ }
+ for _, value := range model.Capabilities {
+ if value == wanted || value == "*" {
+ return true
+ }
+ }
+ return false
+}
+func modelCreated(model domain.Model) int64 {
+ if model.ReleasedAt != nil {
+ return model.ReleasedAt.Unix()
+ }
+ return 0
+}
+
+func (a *API) authorize(w http.ResponseWriter, r *http.Request) (domain.Principal, bool) {
principal, err := a.authenticator.Authenticate(r)
if err != nil || !hasScope(principal, "inference") {
apierror.Write(w, apierror.Error{Status: http.StatusForbidden, Type: "access_denied", Message: "Invalid API key or insufficient permission"}, requestIDFrom(r.Context()))
- return false
+ return domain.Principal{}, false
}
- return true
+ return principal, true
}
func (a *API) publishUsage(event domain.UsageEvent) {
@@ -327,11 +409,11 @@ func (a *API) publishUsage(event domain.UsageEvent) {
func (a *API) finishUsage(r *http.Request, event domain.UsageEvent) {
settled := false
if a.billingMeter != nil {
- ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 5*time.Second)
- err := a.billingMeter.Settle(ctx, event)
+ ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 500*time.Millisecond)
+ err := a.billingMeter.EnqueueSettlement(ctx, event)
cancel()
if err != nil {
- a.logger.Error("billing_settlement_failed", "request_id", event.RequestID, "tenant_id", event.TenantID, "error", err)
+ a.logger.Error("billing_settlement_enqueue_failed", "request_id", event.RequestID, "tenant_id", event.TenantID, "error", err)
} else {
settled = true
}
diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go
index 014b99b..9407530 100644
--- a/internal/httpapi/api_test.go
+++ b/internal/httpapi/api_test.go
@@ -37,7 +37,7 @@ func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) err
return m.authorizeErr
}
-func (m *fakeBillingMeter) Settle(_ context.Context, event domain.UsageEvent) error {
+func (m *fakeBillingMeter) EnqueueSettlement(_ context.Context, event domain.UsageEvent) error {
m.settled <- event
return nil
}
diff --git a/internal/httpapi/proxy.go b/internal/httpapi/proxy.go
new file mode 100644
index 0000000..3cd2797
--- /dev/null
+++ b/internal/httpapi/proxy.go
@@ -0,0 +1,44 @@
+package httpapi
+
+import (
+ "net"
+ "net/http"
+ "strings"
+)
+
+// TrustProxyHeaders accepts forwarding metadata only from explicitly trusted
+// CIDRs. This prevents a direct client from forging HTTPS or audit IP state.
+func TrustProxyHeaders(next http.Handler, trustedCIDRs []string, requireHTTPS bool) (http.Handler, error) {
+ trusted := make([]*net.IPNet, 0, len(trustedCIDRs))
+ for _, value := range trustedCIDRs {
+ _, network, err := net.ParseCIDR(value)
+ if err != nil {
+ return nil, err
+ }
+ trusted = append(trusted, network)
+ }
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ host, _, _ := net.SplitHostPort(r.RemoteAddr)
+ remote := net.ParseIP(host)
+ trustedPeer := false
+ for _, network := range trusted {
+ if remote != nil && network.Contains(remote) {
+ trustedPeer = true
+ break
+ }
+ }
+ if !trustedPeer {
+ for _, header := range []string{"Forwarded", "X-Forwarded-For", "X-Forwarded-Host", "X-Forwarded-Port", "X-Forwarded-Proto", "X-Real-IP"} {
+ r.Header.Del(header)
+ }
+ } else if forwarded := strings.TrimSpace(strings.Split(r.Header.Get("X-Forwarded-For"), ",")[0]); net.ParseIP(forwarded) != nil {
+ r.RemoteAddr = net.JoinHostPort(forwarded, "0")
+ }
+ secure := r.TLS != nil || (trustedPeer && strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https"))
+ if requireHTTPS && !secure {
+ http.Error(w, "HTTPS is required", http.StatusUpgradeRequired)
+ return
+ }
+ next.ServeHTTP(w, r)
+ }), nil
+}
diff --git a/internal/mailer/mailer.go b/internal/mailer/mailer.go
index 4ec11ac..ed6b124 100644
--- a/internal/mailer/mailer.go
+++ b/internal/mailer/mailer.go
@@ -102,8 +102,9 @@ func (s *Sender) Send(ctx context.Context, message Message) error {
}
body := strings.ReplaceAll(message.Body, "\r\n", "\n")
body = strings.ReplaceAll(body, "\n", "\r\n")
- _, writeErr := fmt.Fprintf(w, "From: %s\r\nTo: %s\r\nSubject: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: 8bit\r\n\r\n%s\r\n",
- from, cleanHeader(message.Recipient), cleanHeader(message.Subject), body)
+ messageID := cleanHeader(message.ID) + "@aigw.local"
+ _, writeErr := fmt.Fprintf(w, "From: %s\r\nTo: %s\r\nSubject: %s\r\nMessage-ID: <%s>\r\nX-AIGW-Message-ID: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: 8bit\r\n\r\n%s\r\n",
+ from, cleanHeader(message.Recipient), cleanHeader(message.Subject), messageID, cleanHeader(message.ID), body)
closeErr := w.Close()
if writeErr != nil {
return fmt.Errorf("write SMTP body: %w", writeErr)
diff --git a/internal/operations/operations.go b/internal/operations/operations.go
new file mode 100644
index 0000000..e566488
--- /dev/null
+++ b/internal/operations/operations.go
@@ -0,0 +1,143 @@
+package operations
+
+import (
+ "context"
+ "encoding/json"
+ "net/http"
+ "time"
+
+ "aigw/internal/billing"
+ "aigw/internal/controlplane"
+ "aigw/internal/telemetry"
+)
+
+type Handler struct {
+ Store *controlplane.Store
+ Manager *controlplane.Manager
+ Billing *billing.Service
+ Metrics *telemetry.Metrics
+ MaxSnapshotAge time.Duration
+}
+
+func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path == "/healthz" {
+ write(w, http.StatusOK, map[string]any{"status": "alive"})
+ return
+ }
+ if r.URL.Path == "/metrics" && h.Metrics != nil {
+ h.Metrics.ServeHTTP(w, r)
+ return
+ }
+ if r.URL.Path != "/readyz" {
+ http.NotFound(w, r)
+ return
+ }
+ ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second)
+ defer cancel()
+ checks := map[string]any{}
+ ready := true
+ if h.Store != nil {
+ if err := h.Store.Ping(ctx); err != nil {
+ checks["postgres"] = map[string]any{"status": "failed", "error": err.Error()}
+ ready = false
+ } else {
+ checks["postgres"] = map[string]any{"status": "ok"}
+ }
+ }
+ if h.Manager != nil {
+ healthyAt := h.Manager.LastHealthyAt()
+ if healthyAt.IsZero() {
+ healthyAt = h.Manager.LastReloadAt()
+ }
+ age := time.Since(healthyAt)
+ snapshotOK := h.Manager.Loaded() && age <= h.MaxSnapshotAge
+ checks["snapshot"] = map[string]any{"status": status(snapshotOK), "generation": h.Manager.Generation(), "age_seconds": int64(age.Seconds())}
+ if !snapshotOK {
+ ready = false
+ }
+ redisStatus := "disabled"
+ if h.Manager.RedisConfigured() {
+ redisStatus = "degraded"
+ if h.Manager.RedisConnected() {
+ redisStatus = "ok"
+ }
+ }
+ checks["redis"] = map[string]any{"status": redisStatus, "required": false}
+ }
+ if h.Billing != nil {
+ if err := h.Billing.Ping(ctx); err != nil {
+ checks["billing_postgres"] = map[string]any{"status": "failed", "error": err.Error()}
+ ready = false
+ } else {
+ checks["billing_postgres"] = map[string]any{"status": "ok"}
+ }
+ queue, err := h.Billing.SettlementQueueStatus(ctx)
+ if err != nil {
+ checks["settlement_queue"] = map[string]any{"status": "failed", "error": err.Error()}
+ ready = false
+ } else {
+ backlog := queue.AwaitingEvent + queue.Pending + queue.Processing + queue.Retrying
+ queueOK := queue.SpoolRecords == 0
+ if queue.OldestPending != nil && time.Since(*queue.OldestPending) > 15*time.Minute {
+ queueOK = false
+ }
+ if !queueOK {
+ ready = false
+ }
+ checks["settlement_queue"] = map[string]any{"status": status(queueOK), "backlog": backlog, "spool_records": queue.SpoolRecords, "oldest_pending": queue.OldestPending}
+ if h.Metrics != nil {
+ h.Metrics.SetSettlementQueue(backlog, queue.SpoolRecords)
+ }
+ billingHealth, err := h.Billing.OperationalStatus(ctx)
+ if err != nil {
+ checks["billing_operations"] = map[string]any{"status": "failed", "error": err.Error()}
+ ready = false
+ } else {
+ billingOK := billingHealth.Ready(time.Now().UTC())
+ checks["billing_operations"] = map[string]any{"status": status(billingOK), "details": billingHealth}
+ if !billingOK {
+ ready = false
+ }
+ if h.Metrics != nil {
+ h.Metrics.SetStripeOperations(billingHealth)
+ }
+ }
+ }
+ if h.Store != nil {
+ mail, err := h.Store.MailQueueStatus(ctx)
+ if err != nil {
+ checks["mail_queue"] = map[string]any{"status": "failed", "error": err.Error()}
+ ready = false
+ } else {
+ mailOK := mail.Failed == 0
+ if mail.OldestPending != nil && time.Since(*mail.OldestPending) > 10*time.Minute {
+ mailOK = false
+ }
+ checks["mail_queue"] = map[string]any{"status": status(mailOK), "details": mail}
+ if !mailOK {
+ ready = false
+ }
+ }
+ }
+ }
+ if h.Metrics != nil {
+ h.Metrics.SetReady(ready)
+ }
+ code := http.StatusOK
+ if !ready {
+ code = http.StatusServiceUnavailable
+ }
+ write(w, code, map[string]any{"status": status(ready), "checks": checks})
+}
+
+func status(ok bool) string {
+ if ok {
+ return "ok"
+ }
+ return "failed"
+}
+func write(w http.ResponseWriter, code int, value any) {
+ w.Header().Set("Content-Type", "application/json")
+ w.WriteHeader(code)
+ _ = json.NewEncoder(w).Encode(value)
+}
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go
index a9d5734..0d40bb1 100644
--- a/internal/provider/forwarder.go
+++ b/internal/provider/forwarder.go
@@ -49,7 +49,7 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque
if err := ctx.Err(); err != nil {
return Result{Attempts: i}, err
}
- body, err := rewriteModel(originalBody, route.UpstreamModel)
+ body, err := rewriteRequest(originalBody, route.UpstreamModel, protocol)
if err != nil {
return Result{Attempts: i}, err
}
@@ -80,12 +80,29 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque
}
func rewriteModel(body []byte, upstreamModel string) ([]byte, error) {
+ return rewriteRequest(body, upstreamModel, "")
+}
+
+func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) ([]byte, error) {
var object map[string]json.RawMessage
if err := json.Unmarshal(body, &object); err != nil {
return nil, fmt.Errorf("decode request body: %w", err)
}
encoded, _ := json.Marshal(upstreamModel)
object["model"] = encoded
+ if protocol == domain.ProtocolOpenAI {
+ var stream bool
+ _ = json.Unmarshal(object["stream"], &stream)
+ if stream {
+ var options map[string]json.RawMessage
+ _ = json.Unmarshal(object["stream_options"], &options)
+ if options == nil {
+ options = map[string]json.RawMessage{}
+ }
+ options["include_usage"] = json.RawMessage("true")
+ object["stream_options"], _ = json.Marshal(options)
+ }
+ }
result, err := json.Marshal(object)
if err != nil {
return nil, fmt.Errorf("encode upstream request: %w", err)
diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go
new file mode 100644
index 0000000..2e9b4a1
--- /dev/null
+++ b/internal/provider/forwarder_test.go
@@ -0,0 +1,39 @@
+package provider
+
+import (
+ "encoding/json"
+ "testing"
+
+ "aigw/internal/domain"
+)
+
+func TestRewriteRequestForcesOpenAIStreamUsage(t *testing.T) {
+ result, err := rewriteRequest([]byte(`{"model":"public/model","stream":true,"stream_options":{"other":true}}`), "upstream/model", domain.ProtocolOpenAI)
+ if err != nil {
+ t.Fatal(err)
+ }
+ var body struct {
+ Model string `json:"model"`
+ StreamOptions map[string]any `json:"stream_options"`
+ }
+ if err := json.Unmarshal(result, &body); err != nil {
+ t.Fatal(err)
+ }
+ if body.Model != "upstream/model" || body.StreamOptions["include_usage"] != true || body.StreamOptions["other"] != true {
+ t.Fatalf("unexpected rewritten body: %s", result)
+ }
+}
+
+func TestRewriteRequestDoesNotAddStreamOptionsToAnthropic(t *testing.T) {
+ result, err := rewriteRequest([]byte(`{"model":"public/model","stream":true}`), "upstream/model", domain.ProtocolAnthropic)
+ if err != nil {
+ t.Fatal(err)
+ }
+ var body map[string]json.RawMessage
+ if err := json.Unmarshal(result, &body); err != nil {
+ t.Fatal(err)
+ }
+ if _, exists := body["stream_options"]; exists {
+ t.Fatalf("unexpected OpenAI stream options in Anthropic request: %s", result)
+ }
+}
diff --git a/internal/security/credentials.go b/internal/security/credentials.go
index b55fb6b..900261e 100644
--- a/internal/security/credentials.go
+++ b/internal/security/credentials.go
@@ -4,54 +4,114 @@ import (
"crypto/aes"
"crypto/cipher"
"crypto/rand"
+ "crypto/sha256"
"encoding/base64"
"errors"
"fmt"
"io"
+ "strings"
)
type CredentialCipher struct {
+ current credentialKey
+ keys []credentialKey
+}
+
+type credentialKey struct {
+ id []byte
aead cipher.AEAD
}
+var credentialEnvelopeMagic = []byte("AGK1")
+
func NewCredentialCipher(encodedKey string) (*CredentialCipher, error) {
- key, err := base64.StdEncoding.DecodeString(encodedKey)
- if err != nil {
- return nil, fmt.Errorf("decode credential key: %w", err)
- }
- if len(key) != 32 {
- return nil, errors.New("credential key must be a base64-encoded 32-byte key")
- }
- block, err := aes.NewCipher(key)
- if err != nil {
- return nil, fmt.Errorf("create credential cipher: %w", err)
+ return NewCredentialKeyring(encodedKey, nil)
+}
+
+func NewCredentialKeyring(encodedKey string, previous []string) (*CredentialCipher, error) {
+ values := append([]string{encodedKey}, previous...)
+ keys := make([]credentialKey, 0, len(values))
+ seen := map[string]struct{}{}
+ for _, value := range values {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ continue
+ }
+ key, err := base64.StdEncoding.DecodeString(value)
+ if err != nil {
+ return nil, fmt.Errorf("decode credential key: %w", err)
+ }
+ if len(key) != 32 {
+ return nil, errors.New("credential key must be a base64-encoded 32-byte key")
+ }
+ block, err := aes.NewCipher(key)
+ if err != nil {
+ return nil, fmt.Errorf("create credential cipher: %w", err)
+ }
+ aead, err := cipher.NewGCM(block)
+ if err != nil {
+ return nil, fmt.Errorf("create credential AEAD: %w", err)
+ }
+ fingerprint := sha256.Sum256(key)
+ id := base64.RawURLEncoding.EncodeToString(fingerprint[:6])
+ if _, ok := seen[id]; ok {
+ continue
+ }
+ seen[id] = struct{}{}
+ keys = append(keys, credentialKey{id: []byte(id), aead: aead})
}
- aead, err := cipher.NewGCM(block)
- if err != nil {
- return nil, fmt.Errorf("create credential AEAD: %w", err)
+ if len(keys) == 0 {
+ return nil, errors.New("at least one credential key is required")
}
- return &CredentialCipher{aead: aead}, nil
+ return &CredentialCipher{current: keys[0], keys: keys}, nil
}
func (c *CredentialCipher) Encrypt(plaintext string) ([]byte, error) {
if plaintext == "" {
return nil, errors.New("credential cannot be empty")
}
- nonce := make([]byte, c.aead.NonceSize())
+ nonce := make([]byte, c.current.aead.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return nil, fmt.Errorf("generate credential nonce: %w", err)
}
- return c.aead.Seal(nonce, nonce, []byte(plaintext), nil), nil
+ prefix := append(append([]byte{}, credentialEnvelopeMagic...), c.current.id...)
+ sealed := c.current.aead.Seal(nil, nonce, []byte(plaintext), prefix)
+ return append(append(prefix, nonce...), sealed...), nil
}
func (c *CredentialCipher) Decrypt(ciphertext []byte) (string, error) {
- if len(ciphertext) < c.aead.NonceSize() {
- return "", errors.New("credential ciphertext is truncated")
+ if len(ciphertext) >= len(credentialEnvelopeMagic) && string(ciphertext[:len(credentialEnvelopeMagic)]) == string(credentialEnvelopeMagic) {
+ idStart := len(credentialEnvelopeMagic)
+ idEnd := idStart + 8
+ if len(ciphertext) < idEnd {
+ return "", errors.New("credential ciphertext is truncated")
+ }
+ prefix := ciphertext[:idEnd]
+ for _, key := range c.keys {
+ if string(key.id) != string(ciphertext[idStart:idEnd]) {
+ continue
+ }
+ nonceEnd := idEnd + key.aead.NonceSize()
+ if len(ciphertext) < nonceEnd {
+ return "", errors.New("credential ciphertext is truncated")
+ }
+ plaintext, err := key.aead.Open(nil, ciphertext[idEnd:nonceEnd], ciphertext[nonceEnd:], prefix)
+ if err != nil {
+ return "", errors.New("decrypt credential: authentication failed")
+ }
+ return string(plaintext), nil
+ }
+ return "", errors.New("credential key is not present in the configured keyring")
}
- nonce := ciphertext[:c.aead.NonceSize()]
- plaintext, err := c.aead.Open(nil, nonce, ciphertext[c.aead.NonceSize():], nil)
- if err != nil {
- return "", errors.New("decrypt credential: authentication failed")
+ for _, key := range c.keys {
+ if len(ciphertext) < key.aead.NonceSize() {
+ continue
+ }
+ nonce := ciphertext[:key.aead.NonceSize()]
+ plaintext, err := key.aead.Open(nil, nonce, ciphertext[key.aead.NonceSize():], nil)
+ if err == nil {
+ return string(plaintext), nil
+ }
}
- return string(plaintext), nil
+ return "", errors.New("decrypt credential: authentication failed")
}
diff --git a/internal/security/credentials_test.go b/internal/security/credentials_test.go
index 07fa59e..570ec50 100644
--- a/internal/security/credentials_test.go
+++ b/internal/security/credentials_test.go
@@ -2,29 +2,35 @@ package security
import (
"encoding/base64"
- "strings"
"testing"
)
-func TestCredentialCipherRoundTrip(t *testing.T) {
- key := base64.StdEncoding.EncodeToString([]byte(strings.Repeat("k", 32)))
- cipher, err := NewCredentialCipher(key)
+func TestCredentialKeyringDecryptsPreviousAndReencryptsWithPrimary(t *testing.T) {
+ primary := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901"))
+ previous := base64.StdEncoding.EncodeToString([]byte("abcdefghijklmnopqrstuvwxyzabcdef"))
+ oldCipher, err := NewCredentialCipher(previous)
if err != nil {
t.Fatal(err)
}
- ciphertext, err := cipher.Encrypt("upstream-secret")
+ legacy, err := oldCipher.Encrypt("provider-secret")
if err != nil {
t.Fatal(err)
}
- plaintext, err := cipher.Decrypt(ciphertext)
+ keyring, err := NewCredentialKeyring(primary, []string{previous})
if err != nil {
t.Fatal(err)
}
- if plaintext != "upstream-secret" {
- t.Fatalf("unexpected plaintext: %q", plaintext)
+ if got, err := keyring.Decrypt(legacy); err != nil || got != "provider-secret" {
+ t.Fatalf("decrypt previous key: got %q, err %v", got, err)
}
- ciphertext[len(ciphertext)-1] ^= 1
- if _, err := cipher.Decrypt(ciphertext); err == nil {
- t.Fatal("expected authentication failure for modified ciphertext")
+ rotated, err := keyring.Encrypt("provider-secret")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if string(rotated[:4]) != "AGK1" {
+ t.Fatalf("expected versioned credential envelope, got %q", rotated[:4])
+ }
+ if got, err := keyring.Decrypt(rotated); err != nil || got != "provider-secret" {
+ t.Fatalf("decrypt primary key: got %q, err %v", got, err)
}
}
diff --git a/internal/telemetry/metrics.go b/internal/telemetry/metrics.go
index 4942d8d..04cc494 100644
--- a/internal/telemetry/metrics.go
+++ b/internal/telemetry/metrics.go
@@ -4,14 +4,44 @@ import (
"fmt"
"net/http"
"sync/atomic"
+
+ "aigw/internal/billing"
)
type Metrics struct {
- requests atomic.Uint64
- failed atomic.Uint64
- inFlight atomic.Int64
- attempts atomic.Uint64
- droppedUsage atomic.Uint64
+ requests atomic.Uint64
+ failed atomic.Uint64
+ inFlight atomic.Int64
+ attempts atomic.Uint64
+ droppedUsage atomic.Uint64
+ settlementBacklog atomic.Int64
+ settlementSpool atomic.Int64
+ stripeRefundBacklog atomic.Int64
+ stripeUncollected atomic.Int64
+ stripeMismatches atomic.Int64
+ stripeWebhooks atomic.Int64
+ unmeteredSuccesses atomic.Int64
+ ready atomic.Int64
+}
+
+func (m *Metrics) SetSettlementQueue(backlog int64, spool int) {
+ m.settlementBacklog.Store(backlog)
+ m.settlementSpool.Store(int64(spool))
+}
+func (m *Metrics) SetReady(ready bool) {
+ if ready {
+ m.ready.Store(1)
+ } else {
+ m.ready.Store(0)
+ }
+}
+
+func (m *Metrics) SetStripeOperations(status billing.OperationalStatus) {
+ m.stripeRefundBacklog.Store(status.RefundBacklog)
+ m.stripeUncollected.Store(status.UncollectedMicros)
+ m.stripeMismatches.Store(status.ReconciliationMismatches)
+ m.stripeWebhooks.Store(status.UnprocessedWebhooks)
+ m.unmeteredSuccesses.Store(status.UnmeteredSuccesses)
}
func (m *Metrics) RequestStarted() {
@@ -41,4 +71,12 @@ func (m *Metrics) ServeHTTP(w http.ResponseWriter, _ *http.Request) {
fmt.Fprintf(w, "# TYPE aigw_requests_in_flight gauge\naigw_requests_in_flight %d\n", m.inFlight.Load())
fmt.Fprintf(w, "# TYPE aigw_upstream_attempts_total counter\naigw_upstream_attempts_total %d\n", m.attempts.Load())
fmt.Fprintf(w, "# TYPE aigw_usage_events_dropped_total counter\naigw_usage_events_dropped_total %d\n", m.droppedUsage.Load())
+ fmt.Fprintf(w, "# TYPE aigw_billing_settlement_backlog gauge\naigw_billing_settlement_backlog %d\n", m.settlementBacklog.Load())
+ fmt.Fprintf(w, "# TYPE aigw_billing_settlement_spool_records gauge\naigw_billing_settlement_spool_records %d\n", m.settlementSpool.Load())
+ fmt.Fprintf(w, "# TYPE aigw_stripe_refund_backlog gauge\naigw_stripe_refund_backlog %d\n", m.stripeRefundBacklog.Load())
+ fmt.Fprintf(w, "# TYPE aigw_billing_uncollected_micros gauge\naigw_billing_uncollected_micros %d\n", m.stripeUncollected.Load())
+ fmt.Fprintf(w, "# TYPE aigw_stripe_reconciliation_mismatches gauge\naigw_stripe_reconciliation_mismatches %d\n", m.stripeMismatches.Load())
+ fmt.Fprintf(w, "# TYPE aigw_stripe_webhook_backlog gauge\naigw_stripe_webhook_backlog %d\n", m.stripeWebhooks.Load())
+ fmt.Fprintf(w, "# TYPE aigw_billing_unmetered_successes gauge\naigw_billing_unmetered_successes %d\n", m.unmeteredSuccesses.Load())
+ fmt.Fprintf(w, "# TYPE aigw_ready gauge\naigw_ready %d\n", m.ready.Load())
}
diff --git a/internal/usage/observer.go b/internal/usage/observer.go
index cc8b52f..808781f 100644
--- a/internal/usage/observer.go
+++ b/internal/usage/observer.go
@@ -44,6 +44,17 @@ func (o *Observer) Usage() domain.Usage {
return o.usage
}
+// Reported distinguishes a real zero-token usage object from a response that
+// omitted usage entirely. Billing must never infer this from token totals.
+func (o *Observer) Reported() bool {
+ if !o.stream {
+ o.parseJSON(o.buffer)
+ } else if len(o.line) > 0 {
+ o.parseSSELine(o.line)
+ }
+ return o.found
+}
+
func (o *Observer) captureTail(p []byte) {
if len(p) >= maxCaptureBytes {
o.buffer = append(o.buffer[:0], p[len(p)-maxCaptureBytes:]...)
diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go
index 8fcf408..3104acb 100644
--- a/internal/usage/observer_test.go
+++ b/internal/usage/observer_test.go
@@ -10,6 +10,9 @@ func TestObserverReadsOpenAIJSONUsage(t *testing.T) {
observer := NewObserver(domain.ProtocolOpenAI, false)
_, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`))
got := observer.Usage()
+ if !observer.Reported() {
+ t.Fatal("expected usage to be marked as reported")
+ }
if got.InputTokens != 11 || got.OutputTokens != 7 || got.TotalTokens != 18 {
t.Fatalf("unexpected usage: %+v", got)
}
@@ -20,7 +23,23 @@ func TestObserverCombinesAnthropicSSEUsage(t *testing.T) {
_, _ = observer.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":1}}}\n\n"))
_, _ = observer.Write([]byte("event: message_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":8}}\n\n"))
got := observer.Usage()
+ if !observer.Reported() {
+ t.Fatal("expected streaming usage to be marked as reported")
+ }
if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 {
t.Fatalf("unexpected usage: %+v", got)
}
}
+
+func TestObserverDistinguishesMissingUsageFromReportedZero(t *testing.T) {
+ missing := NewObserver(domain.ProtocolOpenAI, false)
+ _, _ = missing.Write([]byte(`{"choices":[]}`))
+ if missing.Reported() {
+ t.Fatal("response without usage must not be reported")
+ }
+ reported := NewObserver(domain.ProtocolOpenAI, false)
+ _, _ = reported.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":0,"completion_tokens":0,"total_tokens":0}}`))
+ if !reported.Reported() {
+ t.Fatal("explicit zero usage must be distinguished from a missing usage object")
+ }
+}
diff --git a/scripts/backup-postgres.sh b/scripts/backup-postgres.sh
new file mode 100755
index 0000000..f6444c9
--- /dev/null
+++ b/scripts/backup-postgres.sh
@@ -0,0 +1,14 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+: "${AIGW_DATABASE_URL:?AIGW_DATABASE_URL is required}"
+backup_dir="${AIGW_BACKUP_DIR:-./backups}"
+retention_days="${AIGW_BACKUP_RETENTION_DAYS:-30}"
+mkdir -p -- "$backup_dir"
+chmod 700 "$backup_dir"
+timestamp="$(date -u +%Y%m%dT%H%M%SZ)"
+target="$backup_dir/aigw-$timestamp.dump"
+pg_dump --dbname="$AIGW_DATABASE_URL" --format=custom --compress=9 --file="$target"
+sha256sum "$target" >"$target.sha256"
+find "$backup_dir" -maxdepth 1 -type f -name 'aigw-*.dump*' -mtime "+$retention_days" -delete
+printf 'backup=%s\nchecksum=%s.sha256\n' "$target" "$target"
diff --git a/scripts/load-smoke.sh b/scripts/load-smoke.sh
new file mode 100755
index 0000000..536edc1
--- /dev/null
+++ b/scripts/load-smoke.sh
@@ -0,0 +1,22 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+: "${AIGW_LOAD_BASE_URL:?AIGW_LOAD_BASE_URL is required}"
+: "${AIGW_LOAD_API_KEY:?AIGW_LOAD_API_KEY is required}"
+: "${AIGW_LOAD_MODEL:?AIGW_LOAD_MODEL is required}"
+requests="${AIGW_LOAD_REQUESTS:-1000}"
+concurrency="${AIGW_LOAD_CONCURRENCY:-50}"
+payload="$(mktemp)"
+results="$(mktemp)"
+trap 'rm -f -- "$payload" "$results"' EXIT
+chmod 600 "$payload" "$results"
+printf '{"model":"%s","messages":[{"role":"user","content":"health probe"}],"max_tokens":1}\n' "$AIGW_LOAD_MODEL" >"$payload"
+export AIGW_LOAD_BASE_URL AIGW_LOAD_API_KEY payload results
+seq "$requests" | xargs -P "$concurrency" -n 1 sh -c '
+ curl --silent --show-error --output /dev/null --write-out "%{http_code}\n" \
+ --header "Authorization: Bearer $AIGW_LOAD_API_KEY" --header "Content-Type: application/json" \
+ --data-binary "@$payload" "$AIGW_LOAD_BASE_URL/v1/chat/completions" >>"$results"
+' _
+failures="$(awk '$1 < 200 || $1 >= 300 { count++ } END { print count+0 }' "$results")"
+printf 'requests=%s concurrency=%s failures=%s\n' "$requests" "$concurrency" "$failures"
+test "$failures" -eq 0
diff --git a/scripts/redis-fault-drill.sh b/scripts/redis-fault-drill.sh
new file mode 100755
index 0000000..8be92ff
--- /dev/null
+++ b/scripts/redis-fault-drill.sh
@@ -0,0 +1,18 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+: "${AIGW_READY_URL:?AIGW_READY_URL is required}"
+: "${AIGW_REDIS_CONTAINER:?AIGW_REDIS_CONTAINER is required}"
+cleanup() { docker start "$AIGW_REDIS_CONTAINER" >/dev/null 2>&1 || true; }
+trap cleanup EXIT
+docker stop "$AIGW_REDIS_CONTAINER" >/dev/null
+for _ in $(seq 1 20); do
+ body="$(curl --silent --show-error --fail "$AIGW_READY_URL" || true)"
+ if printf '%s' "$body" | grep -q '"redis"' && printf '%s' "$body" | grep -q '"required":false'; then
+ printf 'redis degradation confirmed; gateway remained ready\n'
+ exit 0
+ fi
+ sleep 1
+done
+printf 'gateway did not report recoverable Redis degradation\n' >&2
+exit 1
diff --git a/scripts/restore-drill.sh b/scripts/restore-drill.sh
new file mode 100755
index 0000000..53b8cd2
--- /dev/null
+++ b/scripts/restore-drill.sh
@@ -0,0 +1,14 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+: "${AIGW_RESTORE_DATABASE_URL:?AIGW_RESTORE_DATABASE_URL must point to an isolated drill database}"
+backup="${1:?usage: restore-drill.sh BACKUP.dump}"
+[[ -f "$backup" && -f "$backup.sha256" ]] || { printf 'backup or checksum is missing\n' >&2; exit 2; }
+sha256sum -c "$backup.sha256"
+case "$AIGW_RESTORE_DATABASE_URL" in
+ *localhost*|*127.0.0.1*|*restore*|*drill*) ;;
+ *) printf 'refusing restore: target URL must visibly identify localhost, restore, or drill\n' >&2; exit 2 ;;
+esac
+pg_restore --dbname="$AIGW_RESTORE_DATABASE_URL" --clean --if-exists --no-owner "$backup"
+psql "$AIGW_RESTORE_DATABASE_URL" -v ON_ERROR_STOP=1 -c "SELECT count(*) AS tenants FROM tenants; SELECT count(*) AS usage_events FROM usage_events; SELECT count(*) AS ledger_entries FROM billing_ledger;"
+printf 'restore drill completed for %s\n' "$backup"
diff --git a/scripts/start-debug.sh b/scripts/start-debug.sh
index 7a22a89..bf18de5 100755
--- a/scripts/start-debug.sh
+++ b/scripts/start-debug.sh
@@ -26,6 +26,13 @@ if [[ ! -f "$env_file" ]]; then
umask 077
{
printf 'AIGW_SERVER_ADDRESS=:8080\n'
+ printf 'AIGW_PUBLIC_ADDRESS=:8080\n'
+ printf 'AIGW_ADMIN_ADDRESS=:8081\n'
+ printf 'AIGW_WEBHOOK_ADDRESS=:8082\n'
+ printf 'AIGW_OPERATIONS_ADDRESS=:9090\n'
+ printf 'AIGW_TRUSTED_PROXY_CIDRS=\n'
+ printf 'AIGW_REQUIRE_HTTPS=false\n'
+ printf 'AIGW_DEPLOYMENT_REGION=\n'
printf 'AIGW_POSTGRES_USER=aigw\n'
printf 'AIGW_POSTGRES_PASSWORD=%s\n' "$postgres_password"
printf 'AIGW_POSTGRES_DB=aigw\n'
@@ -34,10 +41,11 @@ if [[ ! -f "$env_file" ]]; then
printf 'AIGW_REDIS_URL=redis://127.0.0.1:6379/0\n'
printf 'AIGW_REDIS_URL_DOCKER=redis://redis:6379/0\n'
printf 'AIGW_CREDENTIAL_KEY=%s\n' "$credential_key"
+ printf 'AIGW_CREDENTIAL_PREVIOUS_KEYS=\n'
printf 'AIGW_ADMIN_TOKEN=%s\n' "$admin_token"
- printf 'AIGW_PUBLIC_URL=http://localhost:8080/admin/\n'
+ printf 'AIGW_PUBLIC_URL=http://localhost:8081/admin/\n'
printf 'AIGW_WEBAUTHN_RP_ID=localhost\n'
- printf 'AIGW_WEBAUTHN_ORIGINS=http://localhost:8080\n'
+ printf 'AIGW_WEBAUTHN_ORIGINS=http://localhost:8081\n'
printf 'AIGW_SMTP_FROM_ADDRESS=no-reply@aigw.local\n'
printf 'AIGW_SMTP_ADDRESS=127.0.0.1:1025\n'
printf 'AIGW_SMTP_ADDRESS_DOCKER=mailpit:1025\n'
@@ -46,13 +54,37 @@ if [[ ! -f "$env_file" ]]; then
printf 'AIGW_STRIPE_API_KEY=rk_test_replace_me\n'
printf 'AIGW_STRIPE_CLI_API_KEY=rk_test_replace_me\n'
printf 'AIGW_STRIPE_WEBHOOK_SECRET=whsec_replace_me\n'
- printf 'AIGW_STRIPE_SUCCESS_URL=http://localhost:8080/admin/?topup=success\n'
- printf 'AIGW_STRIPE_CANCEL_URL=http://localhost:8080/admin/?topup=cancel\n'
+ printf 'AIGW_STRIPE_SUCCESS_URL=http://localhost:8081/admin/?topup=success\n'
+ printf 'AIGW_STRIPE_CANCEL_URL=http://localhost:8081/admin/?topup=cancel\n'
+ printf 'AIGW_STRIPE_PORTAL_RETURN_URL=http://localhost:8081/admin/?billing=portal\n'
+ printf 'AIGW_STRIPE_AUTOMATIC_TAX_ENABLED=false\n'
+ printf 'AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED=false\n'
+ printf 'AIGW_STRIPE_PRODUCT_TAX_CODE=\n'
+ printf 'AIGW_SETTLEMENT_SPOOL_PATH=/var/lib/aigw/settlements.jsonl\n'
} >"$env_file"
chmod 600 "$env_file"
log "created $env_file"
fi
+# Backfill non-secret deployment settings when an older local environment file
+# is reused. Exact legacy localhost values are moved to the split admin port.
+sed -i \
+ -e 's|^AIGW_PUBLIC_URL=http://localhost:8080/admin/$|AIGW_PUBLIC_URL=http://localhost:8081/admin/|' \
+ -e 's|^AIGW_WEBAUTHN_ORIGINS=http://localhost:8080$|AIGW_WEBAUTHN_ORIGINS=http://localhost:8081|' \
+ -e 's|^AIGW_STRIPE_SUCCESS_URL=http://localhost:8080/admin/?topup=success$|AIGW_STRIPE_SUCCESS_URL=http://localhost:8081/admin/?topup=success|' \
+ -e 's|^AIGW_STRIPE_CANCEL_URL=http://localhost:8080/admin/?topup=cancel$|AIGW_STRIPE_CANCEL_URL=http://localhost:8081/admin/?topup=cancel|' \
+ "$env_file"
+for setting in \
+ 'AIGW_PUBLIC_ADDRESS=:8080' \
+ 'AIGW_ADMIN_ADDRESS=:8081' \
+ 'AIGW_WEBHOOK_ADDRESS=:8082' \
+ 'AIGW_OPERATIONS_ADDRESS=:9090' \
+ 'AIGW_STRIPE_PORTAL_RETURN_URL=http://localhost:8081/admin/?billing=portal' \
+ 'AIGW_SETTLEMENT_SPOOL_PATH=/var/lib/aigw/settlements.jsonl'; do
+ name="${setting%%=*}"
+ grep -q "^${name}=" "$env_file" || printf '%s\n' "$setting" >>"$env_file"
+done
+
admin_token="$(sed -n 's/^AIGW_ADMIN_TOKEN=//p' "$env_file" | head -n 1)"
[[ -n "$admin_token" ]] || fail "AIGW_ADMIN_TOKEN is missing from $env_file"
for required_name in AIGW_SERVER_ADDRESS AIGW_POSTGRES_USER AIGW_POSTGRES_PASSWORD AIGW_POSTGRES_DB AIGW_DATABASE_URL_DOCKER AIGW_CREDENTIAL_KEY AIGW_PUBLIC_URL AIGW_WEBAUTHN_RP_ID AIGW_WEBAUTHN_ORIGINS AIGW_SMTP_FROM_ADDRESS AIGW_SMTP_ADDRESS_DOCKER AIGW_STRIPE_API_KEY AIGW_STRIPE_WEBHOOK_SECRET AIGW_STRIPE_SUCCESS_URL AIGW_STRIPE_CANCEL_URL; do
@@ -81,9 +113,9 @@ if ! docker compose --env-file "$env_file" up -d --no-build --wait --wait-timeou
fi
log "services are ready"
-printf '\nAdmin UI: http://localhost:8080/admin/\n'
+printf '\nAdmin UI: http://localhost:8081/admin/\n'
printf 'Mail inbox: http://127.0.0.1:8025/\n'
-printf 'Health: http://127.0.0.1:8080/readyz\n'
+printf 'Health: http://127.0.0.1:9090/readyz\n'
printf 'Secrets: %s (mode 0600)\n' "$env_file"
printf '\nLogs: docker compose --env-file %q logs -f aigw\n' "$env_file"
printf 'Stop: ./scripts/stop-debug.sh\n'