diff options
| author | Chia <Chia@93.nz> | 2026-08-06 09:29:41 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 09:32:46 +1200 |
| commit | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch) | |
| tree | c730526150e55e39b822d5197e4a20318ecaa449 | |
| parent | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (diff) | |
feat: complete commercial control plane, billing, auth, and model catalog
- add PostgreSQL control-plane persistence with Redis-degraded hot reload
- implement prepaid balance, usage ledger, Stripe top-up and reconciliation
- add registration, email verification, password reset, invitations and RBAC
- support TOTP, Passkey MFA, device sessions, quotas and rate limits
- add tenant billing profiles, audit logs and operational readiness checks
- build authenticated admin console, Quickstart, Playground and usage analytics
- add public model catalog with pricing, filtering and cost estimation
- support OpenAI Responses providers and provider health failover
- validate real upstream usage reporting and balance settlement
Diffstat (limited to '')
67 files changed, 5973 insertions, 423 deletions
diff --git a/.env.control.example b/.env.control.example index 5020f32..fd29a99 100644 --- a/.env.control.example +++ b/.env.control.example @@ -20,6 +20,7 @@ AIGW_CREDENTIAL_KEY=replace-with-base64-32-byte-key AIGW_CREDENTIAL_PREVIOUS_KEYS= AIGW_ADMIN_TOKEN=replace-with-a-long-random-admin-token AIGW_PUBLIC_URL=http://localhost:8081/admin/ +AIGW_INFERENCE_PUBLIC_URL=http://localhost:8080 AIGW_WEBAUTHN_RP_ID=localhost AIGW_WEBAUTHN_ORIGINS=http://localhost:8081 AIGW_SMTP_FROM_ADDRESS=no-reply@aigw.local @@ -30,7 +31,9 @@ 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. +# Prefer a restricted test key (rk_test_) with only the permissions documented in README. +# Automatic top-up additionally requires Setup Intents Read and Payment Intents Write. +AIGW_STRIPE_ENABLED=false 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 @@ -8,6 +8,9 @@ COPY cmd ./cmd COPY internal ./internal RUN CGO_ENABLED=0 go build -buildvcs=false -trimpath -ldflags="-s -w" -o /out/aigw ./cmd/aigw +FROM build AS test +CMD ["go", "test", "./..."] + FROM alpine:3.22 RUN apk add --no-cache ca-certificates && adduser -D -H -u 10001 aigw USER aigw @@ -1,27 +1,32 @@ # AIGW -AIGW 是一个轻量、无状态的 AI API 中转后端。当前阶段聚焦上游接入与 API 分发:提供 OpenAI Chat Completions 和 Anthropic Messages 兼容入口,支持公开模型名映射、多上游路由、加权分流、故障转移、SSE 直通、客户密钥鉴权以及异步用量事件。 +AIGW 是一个轻量、无状态的 AI API 中转后端。当前阶段聚焦上游接入与 API 分发:提供 OpenAI Chat Completions、OpenAI Responses 和 Anthropic Messages 兼容入口,支持公开模型名映射、多上游路由、加权分流、故障转移、SSE 直通、客户密钥鉴权以及异步用量事件。 它借鉴了 ZenMux 的双协议、`provider/model` 模型命名、统一错误、请求 ID、路由和可观测性边界,但没有复制其业务实现。 ## 当前能力 -- OpenAI:`POST /v1/chat/completions`、`GET /v1/models` +- OpenAI:`POST /v1/chat/completions`、`POST /v1/responses`、`GET /v1/models` - Anthropic:`POST /anthropic/v1/messages`、`GET /anthropic/v1/models` - ZenMux 风格别名:所有入口同时提供 `/api/...` 路径 - 一个公开模型可配置多个同协议上游,低 `priority` 优先,同级按 `weight` 分流 +- 请求可用 `model-id:provider-slug` 固定到模型目录公开的某个供应商;固定后不会回退到其他供应商 - 在 429、502、503、504 或连接失败时,于响应开始前自动尝试下一条路由 +- 每条模型/供应商 route 记录真实请求的近期可用率和响应头延迟;连续 3 次可重试失败后熔断 30 秒,冷却期间路由自动绕行 - SSE 增量直通、主动断连传播、共享 HTTP/2 连接池 - `Authorization: Bearer` 和 `x-api-key` 客户鉴权 - 统一 JSON 错误、`X-AIGW-Request-ID`、Prometheus 文本指标 - 非阻塞用量事件,包含租户、项目、模型、上游、尝试次数、耗时和 token 用量 - PostgreSQL 预付余额、请求额度冻结、实际 token 结算和不可变账本 - Stripe 托管 Checkout 充值、签名 Webhook 与事件/订单双重幂等 +- Stripe 托管支付方式保存与低余额自动充值,off-session PaymentIntent 失败会暂停自动充值并提示客户处理 - PostgreSQL Usage Ledger、月度项目汇总和单请求成本追溯 - 邮箱验证、邀请注册、密码重置、登录限流、可撤销设备会话和管理 API 审计日志 - TOTP(含一次性恢复码)与 WebAuthn Passkey 注册、二次验证和无密码登录 - 六种 RBAC 角色、租户数据隔离和 CSRF 防护 - 项目级 RPM、估算 TPM、并发限制和月度消费配额 +- API Key 级模型白名单、月度消费上限、过期时间、标签和最后使用时间 +- 无需登录的 `/admin/models` 模型与价格目录,支持搜索、协议/输入/开发者筛选、详情和 token 成本估算 ## 快速运行 @@ -40,6 +45,11 @@ go run ./cmd/aigw -config config.json 如果当前只有一种上游,从 `config.json` 删除未使用的 provider 和对应 model,避免启动时要求该密钥。 +控制面模式下,公开目录位于 `AIGW_PUBLIC_URL` 同源的 `/models`(本地默认 +`http://localhost:8081/admin/models`),匿名数据接口为 +`GET /admin/api/public/models`。公开响应只包含全局可用模型的公开 ID、能力、 +协议、价格与聚合可用性;租户/Key 限定模型和内部路由字段不会返回。 + 不使用真实密钥的本地体验方式: ```bash @@ -65,6 +75,21 @@ curl http://127.0.0.1:8080/v1/chat/completions \ -d '{"model":"openai/gpt-4.1-mini","messages":[{"role":"user","content":"hello"}],"stream":true}' ``` +OpenAI Responses 供应商使用独立的 wire API 配置,不会与 Chat Completions 隐式互转: + +```bash +export OPENAI_BASE_URL='https://your-provider.example/v1' +export OPENAI_API_KEY='your-upstream-key' +export AIGW_SERVER_ADDRESS='127.0.0.1:18081' +export AIGW_API_KEYS='[{"key":"sk-local-change-me","key_id":"local-key","tenant_id":"tenant-demo","project_id":"project-default","scopes":["inference"]}]' +go run ./cmd/aigw -config config.responses.example.json + +curl http://127.0.0.1:18081/v1/responses \ + -H 'Authorization: Bearer sk-local-change-me' \ + -H 'Content-Type: application/json' \ + -d '{"model":"openai/gpt-5.5","input":"Reply with OK.","max_output_tokens":32}' +``` + Anthropic 调用示例: ```bash @@ -77,7 +102,7 @@ curl http://127.0.0.1:8080/anthropic/v1/messages \ ## 配置路由 -每条 route 把一个对外模型映射到一个上游模型: +每条 route 把一个对外模型映射到一个上游模型。供应商的 `protocol` 表示 OpenAI/Anthropic 协议族,`wire_api` 表示实际调用 `chat_completions`、`responses` 或 `messages`。静态配置可为供应商设置唯一的公开 `slug`;省略时使用符合相同格式的 `id`: ```json { @@ -93,14 +118,26 @@ curl http://127.0.0.1:8080/anthropic/v1/messages \ 这里 A/B 承担约 80/20 的首选流量,C 只作为更低优先级的后备。所有上游密钥仅通过 `api_key_env` 指向的环境变量读取。客户密钥从 `AIGW_API_KEYS` JSON 数组读取,进程内只保存 SHA-256 摘要。 +默认请求只传基础模型 ID,由网关自动进行权重分流、熔断绕行和故障转移。需要复现特定供应商行为时,可以先从 `GET /v1/models` 的 `providers` 字段读取公开 slug,再把它附加到模型名: + +```bash +curl http://127.0.0.1:8080/v1/chat/completions \ + -H 'Authorization: Bearer sk-local-change-me' \ + -H 'Content-Type: application/json' \ + -d '{"model":"vendor/model-public:provider-b","messages":[{"role":"user","content":"hello"}]}' +``` + +指定供应商时,路由器只会尝试该 slug 下与当前 API 协议兼容的 route,不会静默切换到其他供应商。该供应商不存在时返回 `404 provider_not_found`,正在熔断冷却时返回 `503 provider_unavailable`。API Key 模型白名单仍按基础模型 ID 校验,用量与扣费也归集到基础模型,避免供应商后缀拆分账单。 + ## 生产边界 当前版本可以作为带预付计费的数据面,并已把控制面、账务和运营入口拆开: - 控制面计费模式会把 UsageEvent、冻结记录和扣费流水同步、幂等写入 PostgreSQL;响应结束只负责把结算事件投递到持久化队列,worker 负责重试、过期冻结恢复和本地 JSONL spool 补偿。 - 客户密钥和模型目录在控制面模式下从 PostgreSQL 载入到原子内存快照;Redis 只是可选的变更广播加速层,故障时通过 PostgreSQL generation 轮询收敛。 -- 当前只把 OpenAI 入口发给 OpenAI 兼容上游、Anthropic 入口发给 Anthropic 兼容上游,不做跨协议转换。 +- 当前只把 Chat Completions、Responses、Anthropic Messages 入口发给相同 wire API 的兼容上游,不做隐式跨协议转换。 - 自动故障转移可能在极少数网络错误下造成上游重复执行。正式计费时需要上游幂等能力、请求去重策略和重复成本对账。 +- 供应商健康状态是每个网关实例基于真实流量维护的 100 次滑动窗口,不是主动探测、全局 SLA 或首 token 延迟;新 route 在首个请求前显示为未采样。熔断状态不写入 PG/Redis,实例重启后重新学习。 - 计费上游必须返回 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 和连接层限速应放在负载均衡器或边缘代理。 @@ -122,19 +159,23 @@ curl http://127.0.0.1:8080/anthropic/v1/messages \ ./scripts/stop-debug.sh ``` -关闭脚本保留 PostgreSQL/Redis 数据卷,下一次启动仍可继续使用已有控制面数据。需要测试 Stripe Checkout 时,把 `.env.debug` 中的 restricted key、Webhook signing secret、成功 URL 和取消 URL 设置为对应环境的值。JSON 配置只保存环境变量名称,不保存外部服务 URL 或密钥。 +关闭脚本保留 PostgreSQL/Redis 数据卷,下一次启动仍可继续使用已有控制面数据。默认 `AIGW_STRIPE_ENABLED=false`,预付余额、扣费账本和人工调账仍可使用,但 Checkout、Customer Portal、退款和 Stripe 对账关闭。需要测试 Stripe 时,把 `.env.debug` 中的 restricted key、Webhook signing secret、成功 URL 和取消 URL 设置为对应环境的值,再显式设置 `AIGW_STRIPE_ENABLED=true`。JSON 配置只保存环境变量名称,不保存外部服务 URL 或密钥。 源码未变化时可跳过镜像构建以快速重启:`AIGW_DEBUG_SKIP_BUILD=1 ./scripts/start-debug.sh`。默认构建使用 Docker host network;特殊环境可以通过 `AIGW_DOCKER_BUILD_NETWORK=default` 覆盖。 -然后打开 `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 明文只在创建成功时返回一次;团队成员使用自己的账号,不共享管理员令牌。 +然后打开 `http://127.0.0.1:8081/admin/`,本地邮件在 `http://127.0.0.1:8025/` 查看;健康检查在 `http://127.0.0.1:9090/readyz`。开发者控制台的调用示例使用 `AIGW_INFERENCE_PUBLIC_URL` 作为网关地址,分离 listener 时不要把它误配成管理地址。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 临时不可用不会回滚注册、邀请或重置请求。普通失败指数退避,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 会按租户保存的阈值、按日幂等发送低余额通知,并发送异常消费通知;未配置租户继续使用 `admin.mail.low_balance_micros` 的全局默认值。域名 DNS 仍必须在邮件供应商处配置 SPF、DKIM 和 DMARC,这不是应用代码可以代替的步骤。WebAuthn 的 `AIGW_WEBAUTHN_RP_ID` 必须是控制台有效域名,`AIGW_WEBAUTHN_ORIGINS` 是逗号分隔的 HTTPS origin。 + +租户登录后的 Quickstart 首屏会显示充值、密钥、可用模型和首次成功请求四步状态;模型目录按协议、输入模态和开发者筛选,并可按发布时间、名称、输入/输出价格和上下文长度排序,再生成使用 `AIGW_INFERENCE_PUBLIC_URL` 的 cURL、Python 和 Node.js 示例。新工作区可以在 Quickstart 直接为默认项目和当前模型创建 starter key,明文仍只显示一次,同时自动放入当前页面内存中的 Playground。连接面板集中列出 OpenAI/Anthropic SDK Base URL、Chat Completions、Responses、Messages、Models 端点,并可复制不含真实密钥的环境变量模板。每个模型都有客户安全的详情视图,展示协议、模态、上下文、最大输出、能力、生命周期和别名,并按当前价格版本实时估算输入、输出、缓存读取和缓存写入成本;逐供应商运行状态只暴露名称、协议、近期可用率、响应头延迟、样本量和熔断恢复时间,不暴露地址、凭证或上游模型。完全熔断的模型不能从快捷入口发起请求。API Playground 会直接从浏览器调用该推理地址,使用当前客户 API Key 经过完整鉴权、路由、余额冻结/结算和 Usage 链路;它只把 Key 保留在当前页面内存,刷新或退出立即清除。失败时页面保留结构化响应和请求 ID,并把权限、余额、模型、限流或上游错误引导到对应控制台页面。分离 listener 时推理服务只允许 `AIGW_PUBLIC_URL` 的精确 Origin,并且不带管理 Cookie。租户可把目录内可用模型保存为默认模型和不同的 fallback 模型,Quickstart 会立即采用该默认值;billing 成员可以单独启停低余额邮件并设置阈值。 -账号邮件先在 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。 +API keys 页面会展示每个 Key 当月已结算费用、待结算冻结、请求数、月度上限、剩余额度、过期时间和最后使用时间;月度统计直接读取 Usage Ledger 与 billing reservation,不在浏览器侧计算。 -控制台角色分为:`platform_admin`、`platform_viewer`、`tenant_admin`、`tenant_billing`、`tenant_developer`、`tenant_viewer`。租户角色的查询条件在服务端下推到 PostgreSQL,不能读取其他租户的项目、密钥、余额、Usage 或审计事件;供应商凭证和路由管理只对平台角色开放。 +控制台角色分为:`platform_admin`、`platform_viewer`、`tenant_admin`、`tenant_billing`、`tenant_developer`、`tenant_viewer`。租户角色的查询条件在服务端下推到 PostgreSQL,不能读取其他租户的项目、密钥、余额、Usage 或审计事件;供应商凭证和路由管理只对平台角色开放。默认模型与财务告警使用独立写权限:developer 不能修改低余额阈值,billing 成员不能修改 API 默认模型。 ## Usage、配额与限流 -每个完成上游尝试的请求都会按 `request_id` 幂等写入 `usage_events`,并更新 `usage_monthly_rollups`。Usage 持久化独立于预付费冻结记录,因此关闭计费也不会关闭用量账本。Admin WebUI 提供本月汇总、单请求状态、模型、token、成本、未收金额和延迟查询。 +每个完成上游尝试的请求都会按 `request_id` 幂等写入 `usage_events`,并更新 `usage_monthly_rollups`。Usage 持久化独立于预付费冻结记录,因此关闭计费也不会关闭用量账本。Admin WebUI 提供本月汇总、单请求状态、模型、token、成本、未收金额和延迟查询;同一时间、项目、API Key、模型、供应商 slug、协议、流式状态、错误类型与成功状态过滤会下推到模型成本排行和供应商性能聚合,展示本期/上期费用变化、成功率、缓存命中、P95 延迟和缺失 usage 请求。每条请求可以打开详情,查看完整 request ID、项目与 Key、路由上游、协议、流式状态、重试、吞吐量、缓存 token 和结算状态,并复制不含 prompt、响应正文或客户密钥的诊断 JSON。 Limits 页面按项目配置: @@ -147,9 +188,9 @@ Redis 可用时,RPM/TPM/并发通过 Lua 原子执行并在多实例间共享ï ## 余额与 Stripe 充值 -模型价格在后台按“币种单位 / 100 万 token”配置,数据库使用 `amount_micros` 固定精度整数保存金额。Stripe 不直接为推理请求结账,只向 PostgreSQL 预付钱包充值;推理请求先按请求体字节数和 `max_tokens`/`max_completion_tokens` 保守冻结余额,成功响应按可信 usage 扣款,失败请求释放冻结。`request_id` 是用量与扣费幂等键,余额、冻结和不可变 ledger 都在同一个 PG 事务中更新。 +模型价格在后台按“币种单位 / 100 万 token”配置,数据库使用 `amount_micros` 固定精度整数保存金额。Stripe 不直接为推理请求结账,只向 PostgreSQL 预付钱包充值;推理请求先按请求体字节数和 `max_tokens`/`max_completion_tokens` 保守冻结余额,成功响应按可信 usage 扣款,失败请求释放冻结。计量层把 OpenAI 总 input 中的 `cached_tokens`/cache-write 子集归一化为互斥的非缓存输入、缓存读取和缓存写入桶,Anthropic 已独立报告的缓存字段则保持不变,避免按输入价和缓存价重复扣费。`request_id` 是用量与扣费幂等键,余额、冻结和不可变 ledger 都在同一个 PG 事务中更新。 -Stripe 使用托管 Checkout,服务端不会接触卡号,也没有硬编码支付方式;支付方式由 Stripe Dashboard 动态配置。充值只在签名校验通过的 Webhook 确认 `payment_status=paid` 后入账,成功跳转页不会直接修改余额。 +Stripe 使用托管 Checkout,服务端不会接触卡号,也没有硬编码支付方式;支付方式由 Stripe Dashboard 动态配置。租户可在 Billing 页面保存开票名称、邮箱和地址,资料用稳定幂等键创建或更新 Stripe Customer;Stripe 暂时不可用时本地资料标记为待修复,下一次充值、保存支付方式或打开客户门户会再次同步。本地不保存税号,启用且确认 Stripe Tax 注册后由 Checkout/Customer Portal 托管税号。手动充值只在签名校验通过的 Webhook 确认 `payment_status=paid` 后入账,成功跳转页不会直接修改余额。自动充值先通过 Checkout Setup Session 保存支付方式;余额低于客户阈值时,后台 worker 用订单 ID 作为幂等键创建并确认 off-session PaymentIntent。同步结果、Webhook 重放和周期对账都只能生成一笔钱包入账;需要客户认证或支付方式失效时会暂停自动充值,不会循环扣款。 配置 Stripe 测试环境: @@ -168,12 +209,28 @@ Webhook 至少订阅: - `checkout.session.async_payment_succeeded` - `checkout.session.async_payment_failed` - `checkout.session.expired` - -网关 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 保存权限时可能要求账户持有人完成二次验证。 +- `payment_intent.succeeded` +- `payment_intent.payment_failed` +- `payment_intent.canceled` +- `charge.succeeded` +- `charge.updated` +- `charge.refunded` +- `refund.created` +- `refund.updated` +- `refund.failed` +- `charge.dispute.created` +- `charge.dispute.updated` +- `charge.dispute.closed` +- `invoice.created` +- `invoice.finalized` +- `invoice.paid` +- `invoice.payment_failed` + +网关 restricted key 的最小权限按实际启用功能配置:Checkout Sessions Write、Customer Portal Sessions Write、Setup Intents Read、Payment Intents Write、Refunds Write,以及 Customers、Charges、Disputes、Invoices 的 Read;如果 Dashboard 将 Checkout 自动创建 Customer 归入 Customers 写权限,再增加 Customers Write。不要给主网关 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。 +模型目录支持输入/输出模态、上下文窗口、最大输出、能力集合、生命周期、弃用替代模型、区域和租户/API key allowlist;价格在 `model_price_versions` 中按生效时间版本化。每个客户 API Key 还可以独立设置模型白名单、月度消费上限和过期时间:限制随 PostgreSQL generation 热加载到鉴权快照,模型检查发生在路由前,金额上限在钱包行锁事务内和待结算冻结一起检查。旧 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` 名称,不保存外部服务密钥或部署域名。 @@ -196,6 +253,9 @@ go test ./... go test -race ./... go vet ./... CGO_ENABLED=0 go build -buildvcs=false ./cmd/... + +# 在容器网络内运行真实 PostgreSQL 集成测试 +docker compose --env-file .env.debug --profile test run --rm integration-test ``` 设置 `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 108a849..3ac388a 100644 --- a/cmd/aigw/main.go +++ b/cmd/aigw/main.go @@ -23,6 +23,7 @@ import ( "aigw/internal/mailer" "aigw/internal/operations" "aigw/internal/provider" + "aigw/internal/providerhealth" "aigw/internal/routing" "aigw/internal/telemetry" @@ -178,12 +179,13 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { if billingService != nil { billingMeter = billingService } + routeHealth := providerhealth.New(providerhealth.Options{}) inferenceAPI := httpapi.New(httpapi.Options{ Authenticator: authenticator, Catalog: modelCatalog, - Router: routing.New(modelCatalog), - Forwarder: provider.New(cfg.UpstreamHTTP, metrics), + Router: routing.New(modelCatalog, routeHealth), + Forwarder: provider.New(cfg.UpstreamHTTP, metrics, routeHealth), UsageSink: usageSink, BillingMeter: billingMeter, Limiter: requestLimiter, @@ -193,6 +195,7 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { MaxBodyBytes: cfg.Server.MaxBodyBytes, ExposeMetrics: cfg.Observability.ExposeMetrics, DeploymentRegion: cfg.Server.DeploymentRegion, + BrowserOrigin: cfg.Admin.PublicURL, }) adminHandler := http.Handler(nil) if cfg.Admin.Enabled { @@ -201,6 +204,10 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { 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, + InferencePublicURL: cfg.Admin.InferencePublicURL, + DefaultLowBalanceMicros: cfg.Admin.Mail.LowBalanceMicros, + Catalog: modelCatalog, + ProviderHealth: routeHealth, }).Handler() } operationHandler := operations.Handler{Store: store, Manager: manager, Billing: billingService, Metrics: metrics, diff --git a/config.control.example.json b/config.control.example.json index 58ef2e7..26de205 100644 --- a/config.control.example.json +++ b/config.control.example.json @@ -37,6 +37,7 @@ "audit_retention_days": 2555, "security_retention_days": 30, "public_url_env": "AIGW_PUBLIC_URL", + "inference_public_url_env": "AIGW_INFERENCE_PUBLIC_URL", "mail": { "enabled": true, "from_name": "AIGW", @@ -66,7 +67,7 @@ "max_top_up_minor": 1000000, "settlement_spool_path_env": "AIGW_SETTLEMENT_SPOOL_PATH", "stripe": { - "enabled": true, + "enabled_env": "AIGW_STRIPE_ENABLED", "api_key_env": "AIGW_STRIPE_API_KEY", "webhook_secret_env": "AIGW_STRIPE_WEBHOOK_SECRET", "success_url_env": "AIGW_STRIPE_SUCCESS_URL", diff --git a/config.example.json b/config.example.json index 8ac2d0c..a6a5dc1 100644 --- a/config.example.json +++ b/config.example.json @@ -23,13 +23,17 @@ "providers": [ { "id": "openai-primary", + "slug": "openai-primary", "protocol": "openai", + "wire_api": "chat_completions", "base_url_env": "OPENAI_BASE_URL", "api_key_env": "OPENAI_API_KEY" }, { "id": "anthropic-primary", + "slug": "anthropic-primary", "protocol": "anthropic", + "wire_api": "messages", "base_url_env": "ANTHROPIC_BASE_URL", "api_key_env": "ANTHROPIC_API_KEY" } diff --git a/config.local.json b/config.local.json index 74a324f..f798b2c 100644 --- a/config.local.json +++ b/config.local.json @@ -19,13 +19,17 @@ "providers": [ { "id": "local-openai", + "slug": "local-openai", "protocol": "openai", + "wire_api": "chat_completions", "base_url_env": "MOCK_UPSTREAM_BASE_URL", "api_key_env": "MOCK_UPSTREAM_KEY" }, { "id": "local-anthropic", + "slug": "local-anthropic", "protocol": "anthropic", + "wire_api": "messages", "base_url_env": "MOCK_UPSTREAM_BASE_URL", "api_key_env": "MOCK_UPSTREAM_KEY" } diff --git a/config.responses.example.json b/config.responses.example.json new file mode 100644 index 0000000..af1611b --- /dev/null +++ b/config.responses.example.json @@ -0,0 +1,47 @@ +{ + "server": { + "address_env": "AIGW_SERVER_ADDRESS", + "max_body_bytes": 16777216, + "read_header_timeout_seconds": 10, + "idle_timeout_seconds": 120, + "shutdown_timeout_seconds": 20 + }, + "auth": { + "keys_env": "AIGW_API_KEYS", + "allow_anonymous": false + }, + "upstream_http": { + "max_idle_connections": 4096, + "max_idle_connections_per_host": 1024, + "idle_connection_timeout_seconds": 90, + "response_header_timeout_seconds": 60 + }, + "observability": { + "usage_buffer": 8192, + "expose_metrics": true + }, + "providers": [ + { + "id": "openai-responses", + "slug": "openai-responses", + "protocol": "openai", + "wire_api": "responses", + "base_url_env": "OPENAI_BASE_URL", + "api_key_env": "OPENAI_API_KEY" + } + ], + "models": [ + { + "id": "openai/gpt-5.5", + "owned_by": "openai", + "routes": [ + { + "provider": "openai-responses", + "upstream_model": "gpt-5.5", + "priority": 0, + "weight": 100 + } + ] + } + ] +} diff --git a/docker-compose.yml b/docker-compose.yml index 8b6b2f1..cc7ecec 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -66,16 +66,18 @@ services: 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_INFERENCE_PUBLIC_URL: ${AIGW_INFERENCE_PUBLIC_URL:?set AIGW_INFERENCE_PUBLIC_URL} AIGW_WEBAUTHN_RP_ID: ${AIGW_WEBAUTHN_RP_ID:?set AIGW_WEBAUTHN_RP_ID} AIGW_WEBAUTHN_ORIGINS: ${AIGW_WEBAUTHN_ORIGINS:?set AIGW_WEBAUTHN_ORIGINS} AIGW_SMTP_FROM_ADDRESS: ${AIGW_SMTP_FROM_ADDRESS:?set AIGW_SMTP_FROM_ADDRESS} AIGW_SMTP_ADDRESS: ${AIGW_SMTP_ADDRESS_DOCKER:?set AIGW_SMTP_ADDRESS_DOCKER} AIGW_SMTP_USERNAME: ${AIGW_SMTP_USERNAME:-} AIGW_SMTP_PASSWORD: ${AIGW_SMTP_PASSWORD:-} - AIGW_STRIPE_API_KEY: ${AIGW_STRIPE_API_KEY:?set AIGW_STRIPE_API_KEY} - 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_ENABLED: ${AIGW_STRIPE_ENABLED:-false} + AIGW_STRIPE_API_KEY: ${AIGW_STRIPE_API_KEY:-} + AIGW_STRIPE_WEBHOOK_SECRET: ${AIGW_STRIPE_WEBHOOK_SECRET:-} + AIGW_STRIPE_SUCCESS_URL: ${AIGW_STRIPE_SUCCESS_URL:-} + AIGW_STRIPE_CANCEL_URL: ${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} @@ -91,7 +93,26 @@ services: retries: 20 start_period: 5s + integration-test: + image: aigw-test:local + build: + context: . + target: test + profiles: ["test"] + working_dir: /src + depends_on: + postgres: + condition: service_healthy + environment: + AIGW_TEST_DATABASE_URL: ${AIGW_DATABASE_URL_DOCKER:?set AIGW_DATABASE_URL_DOCKER} + GOCACHE: /go-build-cache + volumes: + - .:/src:ro + - aigw-go-build:/go-build-cache + command: ["go", "test", "-p=1", "./internal/controlplane", "./internal/billing", "-count=1"] + volumes: aigw-postgres: aigw-redis: aigw-settlements: + aigw-go-build: diff --git a/docs/architecture.md b/docs/architecture.md index c17c65e..f8e617f 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -29,35 +29,35 @@ flowchart LR ## 热路径 1. 入口生成不可预测的请求 ID,并通过 Bearer 或 `x-api-key` 解析客户身份。 -2. 身份包含 `key_id`、`tenant_id`、`project_id` 和 scopes;控制面模式使用 PostgreSQL 摘要快照,代理只依赖 `Authenticator` 接口。 +2. 身份包含 `key_id`、`tenant_id`、`project_id`、scopes、密钥模型白名单、月度上限和过期时间;控制面模式使用 PostgreSQL 摘要快照,代理只依赖 `Authenticator` 接口。 3. 请求体在配置上限内读取一次,以提取公开模型名并支持故障转移时重放。 4. 项目策略从 PG 热更新快照读取。Redis Lua 原子占用 RPM、估算 TPM 和并发额度;Redis 不可用时退回本机窗口。 -5. Router 按协议过滤路由,先按优先级分组,再在同级内按权重选择首选上游。 -6. 计费开启时,请求进入上游前在 PostgreSQL 原子检查月度消费、冻结保守估算额度;余额不足返回 402,月度额度耗尽返回 429。 -7. Provider Adapter 重写上游模型名和凭证,使用进程级共享 Transport 发送请求。上游返回成功头后即锁定路由;SSE 逐块 flush。 +5. Router 按协议过滤路由并跳过仍处于冷却期的 route,再按优先级分组、在同级内按权重选择首选上游。请求使用 `model:provider-slug` 时,基础模型先经过租户/API Key 白名单校验,然后只保留该公开 slug 的兼容 route,绝不跨供应商回退;全部匹配 route 熔断时直接返回可重试的 503。 +6. 计费开启时,请求进入上游前在锁定钱包的 PostgreSQL 事务中同时检查项目月度额度、API Key 月度额度和未结冻结,再冻结保守估算额度;余额不足返回 402,任一月度额度耗尽返回 429。 +7. Provider Adapter 重写上游模型名和凭证,使用进程级共享 Transport 发送请求。每次上游尝试记录响应头延迟与可重试失败;连续 3 次连接失败或 429/502/503/504 后熔断该模型/供应商 route 30 秒。上游返回成功头后即锁定路由;SSE 逐块 flush。 8. 请求结束后释放并发 lease,并用 `request_id` 幂等持久化 UsageEvent、更新月度汇总、按实际 token 结算;结构化日志仅是异步副本。 ## 已固定的扩展边界 | 边界 | 当前实现 | 下一阶段替换 | | --- | --- | --- | -| 客户身份 | PostgreSQL 快照、内存 SHA-256 索引 | SSO/OIDC、SCIM、模型 allowlist | +| 客户身份 | PostgreSQL 快照、内存 SHA-256 索引、密钥过期/模型白名单/月度上限 | IP/CIDR 策略、短期服务身份 | | 控制台权限 | 邮箱验证/邀请/重置、登录限流、设备会话、TOTP/恢复码、Passkey、CSRF、六角色 RBAC、租户 SQL scope、审计日志 | SSO/OIDC、SCIM、组织级 MFA 策略、自定义角色、审批流 | -| 权限 | `Principal.Scopes` 中的 `inference` + 控制台 RBAC | ABAC、IP 与模型策略 | -| 模型目录 | PostgreSQL 快照 + 可选 Redis generation 广播 + PG 轮询兜底 | 版本化控制面、热更新、灰度发布 | -| 路由 | priority + weighted selection + failover | 健康评分、延迟 EWMA、成本/质量策略、熔断 | +| 权限 | `Principal.Scopes` 中的 `inference`、API Key 模型限制 + 控制台 RBAC | ABAC、IP 与条件策略 | +| 模型目录 | PostgreSQL 快照 + 可选 Redis generation 广播 + PG 轮询兜底;租户安全目录 API、供应商运行状态与 Quickstart | 公开 SEO 目录、状态历史、灰度发布 | +| 路由 | priority + weighted selection + failover + `model:provider-slug` 固定供应商 + 真实流量滑动窗口 + 响应头延迟 EWMA + 熔断 | 主动探测、TTFT/吞吐量、成本/质量策略、跨实例健康聚合 | | 计量 | PostgreSQL UsageEvent + 月度汇总 + 异步日志副本 | 分区表、持久消息流、供应商账单对账 | -| 计费 | 预付余额、冻结/结算、不可变流水、Stripe Checkout | 价格版本、退款/争议、信用额度、Metronome 企业合同 | +| 计费 | 版本价格、预付余额、冻结/结算、不可变流水、Stripe 手动/自动充值、退款/争议/对账 | 信用额度、合同价、Metronome 企业合同 | | 限流 | Redis Lua 全局 RPM/估算 TPM/并发,故障时本机降级;PG 月度消费配额 | 滑动窗口、层级策略、边缘 token bucket | -| 协议 | 同协议透传 | 规范化 IR + OpenAI/Anthropic/Google 双向转换 | +| 协议 | Chat Completions、Responses、Anthropic Messages 同 wire API 透传 | 规范化 IR + OpenAI/Anthropic/Google 双向转换 | ## 计费数据原则 真实计费不依赖请求日志。当前以 `request_id` 作为 reservation、usage 和扣费幂等键,价格快照随冻结记录保存,账本流水不可变。余额使用百万分之一币种单位,所有变更在锁定 tenant wallet 的 PostgreSQL 事务中完成。 -Stripe 充值使用 Checkout Session:本地先创建 top-up order,Stripe 请求使用 order ID 作为幂等键;Webhook 验证签名后再次核对 event ID、order ID、session ID、金额和币种。浏览器成功跳转不具有入账权威性。 +Stripe 手动充值使用 Checkout Session:本地先创建 top-up order,Stripe 请求使用 order ID 作为幂等键;Webhook 验证签名后再次核对 event ID、order ID、session ID、金额和币种。自动充值使用 Checkout Setup Session 保存支付方式,低余额 worker 再创建 off-session PaymentIntent;同一订单的同步成功、Webhook 与对账共享幂等账本来源。浏览器成功跳转不具有入账权威性。 -流式请求的 usage 可能只在最后事件出现。当前 observer 会在线解析 OpenAI/Anthropic SSE 中的 usage;如果上游不返回 usage,事件中的 token 为零。接入商业扣费前,应为每种上游建立有测试的 usage normalizer,并使用上游账单进行日对账。 +流式请求的 usage 可能只在最后事件出现。当前 observer 会在线解析 Chat Completions、Responses 和 Anthropic SSE 中的 usage,并把 OpenAI details 中属于总 input 子集的缓存 token 拆成互斥计费桶;Anthropic 独立缓存字段不做减法。如果计费上游不返回 usage,请求会进入 `metering_failed` 并保持余额冻结,而不是按零费用结算。每种上游仍需使用供应商账单做日对账。 ## 扩容方式 @@ -67,6 +67,60 @@ Stripe 充值使用 Checkout Session:本地先创建 top-up order,Stripe 请 - 请求体默认最多 16 MiB,响应仅保留最多 64 KiB 用于非流式 usage 提取;正文直接传输。 - 余额冻结和结算当前各需要一次 PostgreSQL 事务,以正确性优先。高吞吐阶段可把计费拆成独立服务并做账户分片,但不能用最终一致缓存替代权威账本事务。 +## 开发者上手路径 + +租户控制台的 Quickstart 首屏把充值、API key、可用模型和首次成功请求显示为 +可追踪的四步状态。模型目录通过 `GET /admin/api/developer/models` 返回公开元数据和 +当前生效价格,只保留该租户可用的模型与 wire API;它不会返回供应商地址、上游模型名、 +路由权重或 allowlist。控制台可按开发者、协议和输入模态过滤,并按发布时间、价格或上下文 +排序。`GET /admin/api/developer/config` 返回由 +`AIGW_INFERENCE_PUBLIC_URL` 注入的推理地址和协议端点,控制台据此生成 cURL、Python +和 Node.js 示例。新建客户密钥仍只在创建响应和一次性对话框中出现,代码示例默认使用 +`$AIGW_API_KEY`,不会把明文密钥写入本地存储或 HTML。 + +同一 Quickstart 页面提供 API Playground。它使用所选客户 Key 直接调用配置中的公开推理 +地址,因此请求仍经过正式鉴权、模型限制、路由、余额冻结、结算与 Usage Ledger。Key 只 +存在于当前页面内存。分离 listener 时,推理服务仅允许从 `AIGW_PUBLIC_URL` 派生出的精确 +Origin,且不接受浏览器 credentials,只向页面暴露 `X-AIGW-Request-ID`。 + +模型详情由同一份租户安全目录数据渲染,不暴露供应商凭证、内部路由或上游模型名。成本 +估算按当前版本化单价分别计算输入、输出、缓存读取和缓存写入 token;它是请求前预算工具, +实际扣款仍只认 Usage Ledger。Playground 对 `401/403`、`402`、`404`、`429` 和 `5xx` +提供不同的恢复入口,同时原样保留结构化错误与 request ID,便于开发者和支持人员定位。 +目录还把当前内存 route 状态按模型投影为 `online`、`degraded` 或 `unavailable`,详情只返回 +供应商公开 slug、显示名、协议、近期可用率、响应头延迟、样本量和熔断恢复时间。Quickstart +和 Playground 默认使用自动路由,也可以把所选 slug 编入 `model:provider-slug` 来固定供应商。 +`GET /v1/models` 和 Anthropic 模型列表同样只公开 slug 与 wire API,不返回内部 UUID、URL、 +凭证、上游模型名或权重。统计窗口是当前实例 +最近 100 次真实尝试;它不冒充主动健康检查、首 token 延迟、持久状态历史或全局 SLA。 + +Usage 页不保存 prompt 或响应正文。单请求详情仅把已持久化的身份边界、模型路由、协议、 +重试、延迟、token、缓存和结算字段组成可复制诊断 JSON,因此既能支持工单排障,也不会 +把客户输入扩大为新的控制面敏感数据面。 + +`GET /admin/api/usage/analytics` 直接聚合 PostgreSQL Usage Ledger,并复用 Usage 页的租户、 +项目、API Key、模型、状态和时间过滤。结果按模型与供应商返回请求量、成功率、token、 +缓存命中、费用、未收金额、缺失 usage 和 P95 延迟,同时用等长前一周期计算费用变化。 +它不在推理热路径执行,也不从浏览器当前加载的有限请求列表推算财务数据。 + +Quickstart 的 starter key 表单复用正式 `POST /admin/api/keys` 写入链路,为当前租户、所选 +项目和当前模型生成仅含 `inference` scope 的 Key。`api_keys(project_id, tenant_id)` 到 +`projects(id, tenant_id)` 的复合外键保证项目不能跨租户绑定;明文 Key 仍只在创建响应中 +返回一次,随后仅保存在页面内存并填入 Playground。连接面板直接消费 +`GET /admin/api/developer/config`,集中输出两个 SDK Base URL、四个推理/模型端点和不含 +真实凭证的环境变量模板。 + +密钥表单可设置模型白名单、月度金额上限、到期时间和标签。白名单与密钥摘要在同一个 +PostgreSQL 事务内创建,任何未知模型都会让事务整体回滚。`last_used_at` 在 Usage/结算 +事务内单调更新;过期时间既用于快照过滤,也在每次鉴权时检查,避免长轮询间隔延迟失效。 +密钥列表还通过 key/time 索引从 PostgreSQL 返回本月已结算费用、待结算冻结和请求数。 + +`tenant_preferences` 保存租户默认模型、不同的 fallback 模型和低余额提醒阈值。 +`GET /admin/api/developer/preferences` 返回当前租户值;两个独立的写接口分别要求 +`developer.preferences.write` 与 `billing.preferences.write`,因此 developer 与 billing 角色 +不能越权修改对方的设置。保存默认/fallback 时服务端会重新验证该模型当前对租户可见、 +未退役且至少存在一条启用路由。邮件扫描直接读取 PostgreSQL 中的租户阈值,不依赖 Redis。 + ## 管理面安全 bootstrap token 只映射为 `platform_admin`,用于首次建号和故障恢复,不是日常用户凭证。租户注册在同一 PostgreSQL 事务内创建租户、默认项目、钱包和待验证的 `tenant_admin` 账号;邮件动作使用仅保存摘要的一次性 token,邮件正文在 outbox 中加密。密码使用 PBKDF2-HMAC-SHA-256 哈希,TOTP secret、Passkey credential 和 WebAuthn challenge 使用 AES-256-GCM 加密。登录创建 HttpOnly、SameSite 会话 Cookie,并为所有写请求校验独立 CSRF Cookie/header;改密、密码重置和撤销成员会立即失效旧会话。平台角色没有 `tenant_id`,租户角色必须绑定一个 tenant。所有管理 API 在 handler 执行前校验 permission,租户过滤在 SQL 查询或资源所有权检查中完成,前端隐藏菜单不承担安全职责。 @@ -76,7 +130,7 @@ bootstrap token 只映射为 `platform_admin`,用于首次建号和故障恢å¤ ## 建议的后续顺序 1. 将管理监听端口与公网推理端口分离,并为企业客户接入 OIDC/SAML、SCIM 与组织级强制 MFA 策略。 -2. 增加价格版本、退款/冲正、Stripe dispute 处理和供应商日账单对账。 +2. 增加供应商日账单对账和合同价/信用额度。 3. 为不返回 usage 的上游增加可靠 token 计算器,并监控 `uncollected_micros`。 -4. 把 UsageEvent 做时间分区和归档,增加 CSV 导出与对账作业。 -5. 主动健康检查、熔断、延迟 EWMA 和按成本路由。 +4. 把 UsageEvent 做时间分区和归档,增加定时导出与报告。 +5. 增加主动健康检查、TTFT/吞吐量采样、跨实例状态聚合和按成本/性能路由。 diff --git a/docs/commercial-readiness.md b/docs/commercial-readiness.md index b78ae2d..6dc9064 100644 --- a/docs/commercial-readiness.md +++ b/docs/commercial-readiness.md @@ -7,55 +7,95 @@ commercial feature. ## What works now -- OpenAI Chat Completions and Anthropic Messages proxying, streaming, routing, +- OpenAI Chat Completions, OpenAI Responses, and Anthropic Messages proxying, streaming, routing, retry, authentication, persistent usage, prepaid billing, quotas, rate limits, concurrent request limits, RBAC, audit logs, and a PostgreSQL-backed console. -- Stripe-hosted Checkout with signed, idempotent Webhook crediting. The gateway - never accepts card details and never credits a success redirect. +- Stripe-hosted manual top-up and payment-method setup, off-session automatic + top-up, signed/idempotent Webhook crediting, refund/dispute handling, and + reconciliation. The gateway never accepts card details and never credits a + success redirect. +- Tenant billing profiles persist invoice name, email, and postal address and + synchronize them to Stripe Customer with a stable idempotency key. Stripe + failures preserve the local profile for retry; tax IDs stay in Stripe-hosted + Checkout or Customer Portal rather than this database. - PostgreSQL is the source of truth. Redis accelerates invalidation and shared counters but is not required for startup, control-plane writes, billing, or balance correctness. - Verified self-service registration, invitation acceptance, password reset, persistent login throttles, per-device session revocation, encrypted email outbox delivery, TOTP with recovery codes, and WebAuthn Passkeys. +- Tenant-scoped default/fallback model preferences and RBAC-separated low-balance + notification thresholds, persisted in PostgreSQL and applied by Quickstart and + the notification worker. +- Per-key model restrictions, monthly spend caps, expiration, tags, and last-use + tracking. Restrictions are enforced by the runtime snapshot and billing + transaction, not only rendered by the console. The developer console also has + a page-memory API Playground and displays current-month settled spend, pending + reservations, request count, remaining cap, and last use for each key. +- The authenticated model catalog includes a customer-safe detail view, current + price-version cost estimates for input/output/cache tokens, copyable model IDs, + developer filtering, release/price/context sorting, one-click Playground + selection, cURL/Python/Node examples, and actionable diagnostics for + authentication, balance, model, rate-limit, and provider errors. +- The unauthenticated `/admin/models` catalog exposes only globally available + models and supports search, protocol/input/developer filters, release/price/context + sorting, model details, versioned token prices, aggregate route availability, + and a pre-registration cost estimate. Tenant/key allowlists, upstream model IDs, + provider IDs, URLs, and routing weights are excluded by a dedicated public type. +- Quickstart colocates balance state, direct starter-key creation, OpenAI and + Anthropic SDK base URLs, copyable REST endpoints, environment configuration, + code examples, and the live Playground. A newly created key is scoped to the + selected project/model and is kept only in page memory after its one-time reveal. +- Usage events open into a privacy-safe request diagnostic with the complete + request ID, route, retry, protocol, latency, throughput, cache-token, and + settlement fields; copied JSON excludes prompts, responses, and secrets. +- Ledger-backed model cost ranking and provider performance views share the + Usage filters and report period-over-period charge change, success rate, + cache hit, P95 latency, and missing-usage exposure without sampling browser data. +- Runtime routes use a per-instance 100-attempt availability window, response-header + latency EWMA, and a 3-failure/30-second circuit breaker. The customer model detail + compares safe provider runtime fields and prevents launching a model while every + route is cooling down. +- Developers can keep automatic failover or pin a request with + `model-id:provider-slug`. Provider slugs are stable public identifiers returned by + the safe model catalog and selectable in Quickstart and Playground; pinned + requests never fail over to another provider, while billing and key allowlists + remain keyed by the canonical base model. ## Customer product gaps ### P0 before a public commercial launch -- Production transactional-email provider selection, domain authentication, - bounce/complaint handling, low-balance notifications, and monitoring for the - durable mail outbox. Local development currently uses Mailpit. -- A public model catalog and detail page containing provider/developer, release - and retirement dates, input/output modalities, context and maximum output, - supported parameters and protocols, regional availability, and versioned price - dimensions. -- Tenant and API-key model allowlists, budget alerts, low-balance notifications, - downloadable invoices/receipts, payment history, refunds/disputes operations, - and explicit tax handling after registrations are confirmed. +- Production mail provider DNS authentication (SPF/DKIM/DMARC) and provider-side + bounce/complaint wiring remain deployment tasks; the signed feedback endpoint, + suppression table, low-balance notifications, retries, and dead-letter mail + outbox are implemented. Local development uses Mailpit. +- Explicit tax treatment after registrations are confirmed still needs legal + and product sign-off. Invoice details, hosted invoice/PDF/receipt links, + payment history, refunds, disputes, reconciliation, and CSV ledger export are + implemented. - Operational separation of the public inference listener from the management listener, HTTPS-only cookies behind a trusted proxy, backup/restore drills, migration rollback policy, secret rotation, and alerting for usage settlement or Webhook backlogs. -- A durable settlement retry/outbox and reconciliation worker. Persistence and - balance settlement currently use bounded synchronous database calls after the - response; a database timeout is logged but is not queued for retry, so a long - outage can leave reservations pending or usage unbilled. +- Enterprise identity integrations (OIDC/SAML/SCIM), custom roles, and approval + workflows are not included in the current console; password, invite, session, + TOTP, Passkey, RBAC, and audit flows are implemented. ### P1 for ZenMux-like breadth -- OpenAI Responses, Embeddings, Images, Speech and Transcriptions; Gemini native +- OpenAI Embeddings, Images, Speech and Transcriptions; Gemini native APIs; rerank and other media endpoints. The existing protocol field does not make these APIs implemented. -- Provider health measurements per model and route: availability, first-token - latency, throughput, error history, health-aware routing, and customer-visible - status history. -- Provider comparison and price ranges, cache/search/image/audio pricing units, - lifecycle aliases and deprecation notices, searchable filters, release sorting, - SDK examples, and copyable endpoint snippets. -- Usage exports, cost attribution, budgets, scheduled reports, organization - invites, custom roles, OIDC/SAML SSO, SCIM, and support impersonation with - approval and full audit evidence. +- Active provider probes, first-token latency, throughput-aware selection, + cross-instance health aggregation, and customer-visible status history. Runtime + circuit breaking and request-derived route health are implemented; historical + success, total latency, cache hit, missing usage, and cost come from the Usage Ledger. +- Provider price ranges, non-token search/image/audio pricing units, and richer + deprecation notices. Provider runtime comparison and release sorting are implemented. +- Usage exports, scheduled reports, organization invites, custom roles, + OIDC/SAML SSO, SCIM, and support impersonation with approval and full audit + evidence. Model/provider cost attribution is implemented in the console. ## Environment boundary @@ -70,11 +110,13 @@ with mode `0600`. | Redis URLs | `AIGW_REDIS_URL`, `AIGW_REDIS_URL_DOCKER` | | Provider credential encryption | `AIGW_CREDENTIAL_KEY` | | Bootstrap administrator | `AIGW_ADMIN_TOKEN` | +| Stripe integration switch | `AIGW_STRIPE_ENABLED` | | Stripe application key | `AIGW_STRIPE_API_KEY` | | Stripe CLI development key | `AIGW_STRIPE_CLI_API_KEY` | | Stripe Webhook signing secret | `AIGW_STRIPE_WEBHOOK_SECRET` | -| Stripe result URLs | `AIGW_STRIPE_SUCCESS_URL`, `AIGW_STRIPE_CANCEL_URL` | +| Stripe result URLs | `AIGW_STRIPE_SUCCESS_URL`, `AIGW_STRIPE_CANCEL_URL`, `AIGW_STRIPE_PORTAL_RETURN_URL` | | Console public URL | `AIGW_PUBLIC_URL` | +| Public inference/API URL used by customer examples | `AIGW_INFERENCE_PUBLIC_URL` | | SMTP endpoint/sender | `AIGW_SMTP_ADDRESS`, `AIGW_SMTP_FROM_ADDRESS` | | SMTP credentials | `AIGW_SMTP_USERNAME`, `AIGW_SMTP_PASSWORD` | | WebAuthn RP/origins | `AIGW_WEBAUTHN_RP_ID`, `AIGW_WEBAUTHN_ORIGINS` | diff --git a/internal/adminapi/api.go b/internal/adminapi/api.go index 9f460e3..ff87ac9 100644 --- a/internal/adminapi/api.go +++ b/internal/adminapi/api.go @@ -13,13 +13,18 @@ import ( "log/slog" "net" "net/http" + "sort" + "strconv" "strings" "time" "aigw/internal/adminui" "aigw/internal/apierror" "aigw/internal/billing" + "aigw/internal/catalog" "aigw/internal/controlplane" + "aigw/internal/domain" + "aigw/internal/providerhealth" "github.com/go-webauthn/webauthn/webauthn" "github.com/jackc/pgx/v5/pgconn" @@ -38,6 +43,10 @@ type API struct { publicURL string webauthn *webauthn.WebAuthn mailEnabled bool + inferencePublicURL string + defaultLowBalance int64 + catalog *catalog.Catalog + health *providerhealth.Tracker } type actorKey struct{} @@ -58,18 +67,22 @@ func (w *auditWriter) Write(body []byte) (int, error) { } type Options struct { - Store *controlplane.Store - Manager *controlplane.Manager - Billing *billing.Service - Token string - Logger *slog.Logger - Prefix string - RegistrationEnabled bool - SessionTTL time.Duration - Currency string - PublicURL string - WebAuthn *webauthn.WebAuthn - MailEnabled bool + Store *controlplane.Store + Manager *controlplane.Manager + Billing *billing.Service + Token string + Logger *slog.Logger + Prefix string + RegistrationEnabled bool + SessionTTL time.Duration + Currency string + PublicURL string + WebAuthn *webauthn.WebAuthn + MailEnabled bool + InferencePublicURL string + DefaultLowBalanceMicros int64 + Catalog *catalog.Catalog + ProviderHealth *providerhealth.Tracker } func New(options Options) *API { @@ -89,7 +102,8 @@ func New(options Options) *API { return &API{store: options.Store, manager: options.Manager, billing: options.Billing, token: []byte(options.Token), logger: options.Logger, prefix: prefix, registrationEnabled: options.RegistrationEnabled, sessionTTL: options.SessionTTL, currency: options.Currency, publicURL: strings.TrimRight(options.PublicURL, "/") + "/", - webauthn: options.WebAuthn, mailEnabled: options.MailEnabled} + webauthn: options.WebAuthn, mailEnabled: options.MailEnabled, inferencePublicURL: strings.TrimRight(options.InferencePublicURL, "/"), + defaultLowBalance: options.DefaultLowBalanceMicros, catalog: options.Catalog, health: options.ProviderHealth} } func (a *API) Handler() http.Handler { @@ -103,6 +117,8 @@ func (a *API) Handler() http.Handler { http.Redirect(w, r, target, http.StatusTemporaryRedirect) }) mux.Handle(a.prefix+"/", http.StripPrefix(a.prefix, adminui.Handler())) + mux.HandleFunc("GET "+apiPrefix+"/public/models", a.public(a.publicModels)) + mux.HandleFunc("GET "+apiPrefix+"/public/models/{id...}", a.public(a.publicModel)) mux.HandleFunc("GET "+apiPrefix+"/auth/config", a.public(a.authConfig)) mux.HandleFunc("GET "+apiPrefix+"/auth/session", a.public(a.authSession)) mux.HandleFunc("POST "+apiPrefix+"/auth/register", a.public(a.register)) @@ -131,6 +147,11 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("POST "+apiPrefix+"/auth/passkeys/{id}/delete", a.withAuth("overview.read", a.deletePasskey)) mux.HandleFunc("GET "+apiPrefix+"/overview", a.withAuth("overview.read", a.overview)) + mux.HandleFunc("GET "+apiPrefix+"/developer/config", a.withAuth("overview.read", a.developerConfig)) + mux.HandleFunc("GET "+apiPrefix+"/developer/models", a.withAuth("overview.read", a.developerModels)) + mux.HandleFunc("GET "+apiPrefix+"/developer/preferences", a.withAuth("preferences.read", a.developerPreferences)) + mux.HandleFunc("PUT "+apiPrefix+"/developer/preferences", a.withAuth("developer.preferences.write", a.updateDeveloperPreferences)) + mux.HandleFunc("PUT "+apiPrefix+"/developer/preferences/billing", a.withAuth("billing.preferences.write", a.updateBillingPreferences)) mux.HandleFunc("GET "+apiPrefix+"/tenants", a.withAuth("tenants.read", a.listTenants)) mux.HandleFunc("POST "+apiPrefix+"/tenants", a.withAuth("tenants.write", a.createTenant)) mux.HandleFunc("GET "+apiPrefix+"/projects", a.withAuth("projects.read", a.listProjects)) @@ -148,12 +169,17 @@ func (a *API) Handler() http.Handler { 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)) + mux.HandleFunc("GET "+apiPrefix+"/billing/profile", a.withAuth("billing.read", a.getBillingProfile)) + mux.HandleFunc("PUT "+apiPrefix+"/billing/profile", a.withAuth("billing.topup", a.updateBillingProfile)) mux.HandleFunc("GET "+apiPrefix+"/billing/ledger", a.withAuth("billing.read", a.listBillingLedger)) mux.HandleFunc("GET "+apiPrefix+"/billing/orders", a.withAuth("billing.read", a.listTopUpOrders)) 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("GET "+apiPrefix+"/billing/auto-topup", a.withAuth("billing.read", a.getAutoTopUp)) + mux.HandleFunc("PUT "+apiPrefix+"/billing/auto-topup", a.withAuth("billing.topup", a.updateAutoTopUp)) + mux.HandleFunc("POST "+apiPrefix+"/billing/auto-topup/setup-sessions", a.withAuth("billing.topup", a.createAutoTopUpSetupSession)) 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)) @@ -166,6 +192,8 @@ func (a *API) Handler() http.Handler { } mux.HandleFunc("GET "+apiPrefix+"/usage", a.withAuth("usage.read", a.listUsage)) mux.HandleFunc("GET "+apiPrefix+"/usage/summary", a.withAuth("usage.read", a.usageSummary)) + mux.HandleFunc("GET "+apiPrefix+"/usage/daily", a.withAuth("usage.read", a.usageDaily)) + mux.HandleFunc("GET "+apiPrefix+"/usage/analytics", a.withAuth("usage.read", a.usageAnalytics)) mux.HandleFunc("GET "+apiPrefix+"/limits", a.withAuth("limits.read", a.listLimits)) mux.HandleFunc("POST "+apiPrefix+"/limits/{project_id}", a.withAuth("limits.write", a.setLimit)) mux.HandleFunc("GET "+apiPrefix+"/users", a.withAuth("users.read", a.listUsers)) @@ -932,6 +960,228 @@ func (a *API) overview(w http.ResponseWriter, r *http.Request) { writeJSON(w, result) } +func (a *API) developerConfig(w http.ResponseWriter, _ *http.Request) { + writeJSON(w, map[string]any{ + "base_url": a.inferencePublicURL, + "endpoints": map[string]string{ + "chat_completions": "/v1/chat/completions", + "responses": "/v1/responses", + "messages": "/anthropic/v1/messages", + "models": "/v1/models", + }, + }) +} + +func (a *API) developerModels(w http.ResponseWriter, r *http.Request) { + result, err := a.store.ListDeveloperModels(r.Context(), a.actor(r).TenantID) + if err != nil { + a.databaseError(w, r, err) + return + } + if err := a.addDeveloperModelHealth(r.Context(), result); err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) publicModels(w http.ResponseWriter, r *http.Request) { + result, err := a.store.ListPublicModels(r.Context()) + if err != nil { + a.databaseError(w, r, err) + return + } + a.addPublicModelHealth(result) + w.Header().Set("Cache-Control", "public, max-age=30, stale-while-revalidate=120") + writeJSON(w, map[string]any{ + "data": result, + "inference_base_url": a.inferencePublicURL, + "registration_enabled": a.registrationEnabled, + }) +} + +func (a *API) publicModel(w http.ResponseWriter, r *http.Request) { + wanted := strings.Trim(strings.TrimSpace(r.PathValue("id")), "/") + result, err := a.store.ListPublicModels(r.Context()) + if err != nil { + a.databaseError(w, r, err) + return + } + a.addPublicModelHealth(result) + for _, item := range result { + if item.PublicID == wanted { + w.Header().Set("Cache-Control", "public, max-age=30, stale-while-revalidate=120") + writeJSON(w, item) + return + } + } + apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "model_not_found", Message: "Model not found"}, requestID(r)) +} + +func (a *API) addPublicModelHealth(models []controlplane.PublicModel) { + if a.catalog == nil { + return + } + statusByRoute := make(map[providerhealth.RouteKey]providerhealth.Status) + if a.health != nil { + for _, item := range a.health.Snapshot() { + statusByRoute[providerhealth.RouteKey{ModelID: item.ModelID, ProviderID: item.ProviderID, WireAPI: item.WireAPI}] = item + } + } + for index := range models { + model, err := a.catalog.Model(models[index].PublicID) + if err != nil { + continue + } + seen := make(map[string]struct{}) + available := 0 + for _, route := range model.Routes { + wireAPI := route.Provider.EffectiveWireAPI() + dedupe := route.Provider.ID + "\x00" + wireAPI + if _, exists := seen[dedupe]; exists { + continue + } + seen[dedupe] = struct{}{} + status, measured := statusByRoute[providerhealth.RouteKey{ModelID: model.ID, ProviderID: route.Provider.ID, WireAPI: wireAPI}] + if !measured || status.State != "open" { + available++ + } + } + models[index].ProviderCount = len(seen) + models[index].AvailableProviderCount = available + switch { + case len(seen) == 0 || available == 0: + models[index].HealthStatus = "unavailable" + case available < len(seen): + models[index].HealthStatus = "degraded" + default: + models[index].HealthStatus = "available" + } + } +} + +func (a *API) addDeveloperModelHealth(ctx context.Context, models []controlplane.DeveloperModel) error { + if a.catalog == nil { + return nil + } + providers, err := a.store.ListProviders(ctx) + if err != nil { + return err + } + providersByID := make(map[string]controlplane.Provider, len(providers)) + for _, item := range providers { + providersByID[item.ID] = item + } + statusByRoute := make(map[providerhealth.RouteKey]providerhealth.Status) + if a.health != nil { + for _, item := range a.health.Snapshot() { + statusByRoute[providerhealth.RouteKey{ModelID: item.ModelID, ProviderID: item.ProviderID, WireAPI: item.WireAPI}] = item + } + } + for index := range models { + model, err := a.catalog.Model(models[index].PublicID) + if err != nil { + continue + } + seen := make(map[string]struct{}) + items := make([]controlplane.DeveloperProviderHealth, 0, len(model.Routes)) + available := 0 + for _, route := range model.Routes { + key := providerhealth.RouteKey{ModelID: model.ID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()} + dedupe := key.ProviderID + "\x00" + key.WireAPI + if _, exists := seen[dedupe]; exists { + continue + } + seen[dedupe] = struct{}{} + status, measured := statusByRoute[key] + state := "unknown" + if measured { + state = status.State + } + if state != "open" { + available++ + } + provider := providersByID[route.Provider.ID] + items = append(items, controlplane.DeveloperProviderHealth{Slug: route.Provider.EffectiveSlug(), Name: provider.Name, + Protocol: string(route.Provider.Protocol), WireAPI: key.WireAPI, State: state, Attempts: status.Attempts, + RecentSamples: status.RecentSamples, AvailabilityPercent: status.AvailabilityPercent, + HeaderLatencyEWMA: status.HeaderLatencyEWMA, ConsecutiveFailures: status.ConsecutiveFailures, + LastObservedAt: status.LastObservedAt, CircuitOpenUntil: status.CircuitOpenUntil}) + } + sort.Slice(items, func(i, j int) bool { + iOpen := items[i].State == "open" + jOpen := items[j].State == "open" + if iOpen != jOpen { + return !iOpen + } + if items[i].Name != items[j].Name { + return items[i].Name < items[j].Name + } + return items[i].WireAPI < items[j].WireAPI + }) + models[index].Providers = items + models[index].ProviderCount = len(items) + models[index].AvailableProviderCount = available + switch { + case len(items) == 0 || available == 0: + models[index].HealthStatus = "unavailable" + case available < len(items): + models[index].HealthStatus = "degraded" + default: + models[index].HealthStatus = "online" + } + } + return nil +} + +func (a *API) developerPreferences(w http.ResponseWriter, r *http.Request) { + result, err := a.store.GetTenantPreferences(r.Context(), a.preferenceTenantID(r, ""), a.defaultLowBalance) + if err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) updateDeveloperPreferences(w http.ResponseWriter, r *http.Request) { + var input controlplane.SetDeveloperPreferencesInput + if !decodeBody(w, r, &input) { + return + } + input.TenantID = a.preferenceTenantID(r, input.TenantID) + result, err := a.store.SetDeveloperPreferences(r.Context(), input) + if err != nil { + a.mutationError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) updateBillingPreferences(w http.ResponseWriter, r *http.Request) { + var input controlplane.SetBillingPreferencesInput + if !decodeBody(w, r, &input) { + return + } + input.TenantID = a.preferenceTenantID(r, input.TenantID) + result, err := a.store.SetBillingPreferences(r.Context(), input, a.defaultLowBalance) + if err != nil { + a.mutationError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) preferenceTenantID(r *http.Request, requested string) string { + actor := a.actor(r) + if actor.TenantID != "" { + return actor.TenantID + } + if requested = strings.TrimSpace(requested); requested != "" { + return requested + } + return strings.TrimSpace(r.URL.Query().Get("tenant_id")) +} + func (a *API) listBillingAccounts(w http.ResponseWriter, r *http.Request) { result, err := a.billing.ListAccounts(r.Context(), a.actor(r).TenantID) if err != nil { @@ -941,6 +1191,38 @@ func (a *API) listBillingAccounts(w http.ResponseWriter, r *http.Request) { writeJSON(w, result) } +func (a *API) getBillingProfile(w http.ResponseWriter, r *http.Request) { + tenantID := a.preferenceTenantID(r, "") + if tenantID == "" { + a.scopeError(w, r) + return + } + result, err := a.billing.GetBillingProfile(r.Context(), tenantID) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) updateBillingProfile(w http.ResponseWriter, r *http.Request) { + var input billing.UpdateBillingProfileInput + if !decodeBody(w, r, &input) { + return + } + input.TenantID = a.preferenceTenantID(r, input.TenantID) + if input.TenantID == "" { + a.scopeError(w, r) + return + } + result, err := a.billing.UpdateBillingProfile(r.Context(), input) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + func (a *API) listBillingLedger(w http.ResponseWriter, r *http.Request) { tenantID := a.actor(r).TenantID if tenantID == "" { @@ -1022,6 +1304,46 @@ func (a *API) createPortalSession(w http.ResponseWriter, r *http.Request) { writeStatusJSON(w, http.StatusCreated, result) } +func (a *API) getAutoTopUp(w http.ResponseWriter, r *http.Request) { + tenantID := a.preferenceTenantID(r, "") + result, err := a.billing.GetAutoTopUpSettings(r.Context(), tenantID) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) updateAutoTopUp(w http.ResponseWriter, r *http.Request) { + var input billing.UpdateAutoTopUpInput + if !decodeBody(w, r, &input) { + return + } + input.TenantID = a.preferenceTenantID(r, input.TenantID) + result, err := a.billing.UpdateAutoTopUp(r.Context(), input) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) createAutoTopUpSetupSession(w http.ResponseWriter, r *http.Request) { + var input billing.AutoTopUpSetupInput + if !decodeBody(w, r, &input) { + return + } + actor := a.actor(r) + input.TenantID = a.preferenceTenantID(r, input.TenantID) + input.CustomerEmail = actor.Email + result, err := a.billing.CreateAutoTopUpSetupSession(r.Context(), input) + 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 == "" { @@ -1383,7 +1705,11 @@ func (a *API) reload(w http.ResponseWriter, r *http.Request) { } func (a *API) listUsage(w http.ResponseWriter, r *http.Request) { - query := controlplane.UsageQuery{TenantID: a.actor(r).TenantID, ProjectID: r.URL.Query().Get("project_id"), Model: r.URL.Query().Get("model"), Limit: 200} + query, err := a.usageQuery(r) + if err != nil { + a.mutationError(w, r, err) + return + } result, err := a.store.ListUsage(r.Context(), query) if err != nil { a.databaseError(w, r, err) @@ -1392,6 +1718,101 @@ func (a *API) listUsage(w http.ResponseWriter, r *http.Request) { writeJSON(w, result) } +func (a *API) usageDaily(w http.ResponseWriter, r *http.Request) { + query, err := a.usageQuery(r) + if err != nil { + a.mutationError(w, r, err) + return + } + result, err := a.store.UsageDaily(r.Context(), query) + if err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) usageAnalytics(w http.ResponseWriter, r *http.Request) { + query, err := a.usageQuery(r) + if err != nil { + a.mutationError(w, r, err) + return + } + result, err := a.store.UsageAnalytics(r.Context(), query) + if err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) usageQuery(r *http.Request) (controlplane.UsageQuery, error) { + actor := a.actor(r) + tenantID := actor.TenantID + if tenantID == "" { + tenantID = strings.TrimSpace(r.URL.Query().Get("tenant_id")) + } + query := controlplane.UsageQuery{ + TenantID: tenantID, ProjectID: strings.TrimSpace(r.URL.Query().Get("project_id")), + KeyID: strings.TrimSpace(r.URL.Query().Get("key_id")), Model: strings.TrimSpace(r.URL.Query().Get("model")), + Provider: strings.ToLower(strings.TrimSpace(r.URL.Query().Get("provider"))), Protocol: strings.TrimSpace(r.URL.Query().Get("protocol")), + ErrorType: strings.TrimSpace(r.URL.Query().Get("error_type")), RequestID: strings.TrimSpace(r.URL.Query().Get("request_id")), + Status: strings.TrimSpace(r.URL.Query().Get("status")), Limit: 200, + } + if raw := strings.TrimSpace(r.URL.Query().Get("limit")); raw != "" { + limit, err := strconv.Atoi(raw) + if err != nil || limit < 1 || limit > 1000 { + return controlplane.UsageQuery{}, errors.New("usage limit must be between 1 and 1000") + } + query.Limit = limit + } + if query.Status != "" && query.Status != "success" && query.Status != "error" { + return controlplane.UsageQuery{}, errors.New("usage status must be success or error") + } + if query.Protocol != "" && query.Protocol != string(domain.ProtocolOpenAI) && query.Protocol != string(domain.ProtocolOpenAIResponses) && query.Protocol != string(domain.ProtocolAnthropic) { + return controlplane.UsageQuery{}, errors.New("usage protocol is invalid") + } + if len(query.Provider) > 64 || len(query.ErrorType) > 128 { + return controlplane.UsageQuery{}, errors.New("usage provider or error type is too long") + } + if raw := strings.TrimSpace(r.URL.Query().Get("stream")); raw != "" { + value, err := strconv.ParseBool(raw) + if err != nil { + return controlplane.UsageQuery{}, errors.New("usage stream must be true or false") + } + query.Stream = &value + } + var err error + if query.From, err = parseUsageTime(r.URL.Query().Get("from"), false); err != nil { + return controlplane.UsageQuery{}, err + } + if query.To, err = parseUsageTime(r.URL.Query().Get("to"), true); err != nil { + return controlplane.UsageQuery{}, err + } + if !query.From.IsZero() && !query.To.IsZero() && !query.To.After(query.From) { + return controlplane.UsageQuery{}, errors.New("usage to must be after from") + } + return query, nil +} + +func parseUsageTime(raw string, endOfDay bool) (time.Time, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return time.Time{}, nil + } + if parsed, err := time.Parse(time.RFC3339, raw); err == nil { + return parsed.UTC(), nil + } + parsed, err := time.Parse("2006-01-02", raw) + if err != nil { + return time.Time{}, errors.New("usage dates must be RFC3339 or YYYY-MM-DD") + } + if endOfDay { + parsed = parsed.AddDate(0, 0, 1) + } + return parsed.UTC(), nil +} + func (a *API) usageSummary(w http.ResponseWriter, r *http.Request) { result, err := a.store.UsageSummary(r.Context(), a.actor(r).TenantID, r.URL.Query().Get("project_id")) if err != nil { @@ -1567,6 +1988,27 @@ func (a *API) billingError(w http.ResponseWriter, r *http.Request, err error) { status = http.StatusConflict typeName = "topup_order_not_resolvable" message = err.Error() + case errors.Is(err, billing.ErrBillingAccountNotFound): + status = http.StatusNotFound + typeName = "billing_account_not_found" + message = "Billing account was not found" + case errors.Is(err, billing.ErrPaymentMethodRequired): + status = http.StatusConflict + typeName = "payment_method_required" + message = "Save a payment method before enabling automatic top-up" + case errors.Is(err, billing.ErrAutoTopUpNeedsAttention): + status = http.StatusConflict + typeName = "payment_method_attention_required" + message = "Replace or re-authorize the saved payment method before enabling automatic top-up" + case errors.Is(err, billing.ErrInvalidBillingProfile): + status = http.StatusBadRequest + typeName = "invalid_billing_profile" + message = strings.TrimPrefix(err.Error(), billing.ErrInvalidBillingProfile.Error()+": ") + case errors.Is(err, billing.ErrBillingProfileSync): + status = http.StatusBadGateway + typeName = "billing_profile_sync_failed" + message = "Billing details were saved, but Stripe synchronization failed" + a.logger.Error("billing_profile_sync_failed", "error", err) default: a.logger.Error("admin_billing_error", "error", err) } diff --git a/internal/adminui/assets/app.js b/internal/adminui/assets/app.js index 8d29e23..e30a0c8 100644 --- a/internal/adminui/assets/app.js +++ b/internal/adminui/assets/app.js @@ -1,8 +1,10 @@ const state = { token: '', csrf: '', actor: {}, permissions: new Set(), overview: {}, tenants: [], projects: [], keys: [], providers: [], models: [], billingAccounts: [], ledger: [], - usage: [], usageSummary: [], limits: [], users: [], audit: [], orders: [], refunds: [], disputes: [], invoices: [], sessions: [], - mfa: {totp_enabled:false,passkeys:[]}, pendingMFA: null, authConfig: {} + usage: [], usageSummary: [], usageDaily: [], usageAnalytics: {models:[],providers:[]}, limits: [], users: [], audit: [], orders: [], refunds: [], disputes: [], invoices: [], sessions: [], + developerConfig: {base_url:'',endpoints:{}}, developerModels: [], preferences: {}, + autoTopUp: {}, billingProfile: {}, + mfa: {totp_enabled:false,passkeys:[]}, pendingMFA: null, authConfig: {}, playgroundKey: '', playgroundController: null, detailModel: null, detailUsage: null }; const $ = (selector) => document.querySelector(selector); const $$ = (selector) => [...document.querySelectorAll(selector)]; @@ -26,7 +28,7 @@ async function api(path, options = {}) { function setConnected(connected) { $('#auth-screen').classList.toggle('hidden', connected); $('#console-app').classList.toggle('hidden', !connected); - if (!connected) return; + if (!connected) { state.playgroundController?.abort();state.playgroundController=null;state.playgroundKey=''; const key=$('#playground-key');if(key)key.value=''; return; } $('#connection-state').textContent = state.actor.role?.replaceAll('_', ' ') || 'connected'; $('#actor-label').textContent = state.actor.display_name || state.actor.email || 'Operator'; } @@ -42,6 +44,7 @@ function scaledToDecimal(value, digits) { const number = BigInt(value || 0); con function currencyDigits(currency) { return ['bif','clp','djf','gnf','jpy','kmf','krw','mga','pyg','rwf','ugx','vnd','vuv','xaf','xof','xpf'].includes(currency) ? 0 : ['bhd','jod','kwd','omr','tnd'].includes(currency) ? 3 : 2; } function money(micros, currency = state.overview.billing_currency || 'usd') { return new Intl.NumberFormat(undefined, { style:'currency', currency:currency.toUpperCase(), minimumFractionDigits:2, maximumFractionDigits:6 }).format(Number(micros || 0) / 1_000_000); } function integer(value) { return new Intl.NumberFormat().format(Number(value || 0)); } +function chartHeightClass(value, maximum) { return `chart-height-${Math.max(1, Math.min(20, Math.ceil(Number(value || 0) / Math.max(1, Number(maximum || 0)) * 20)))}`; } function emptyRow(span) { return `<tr><td colspan="${span}" class="empty">No records yet</td></tr>`; } function showSecret(title, value) { $('#secret-title').textContent = title; $('#created-secret').textContent = value; $('#secret-dialog').showModal(); } function cookie(name) { const prefix=`${encodeURIComponent(name)}=`; const value=document.cookie.split('; ').find(item=>item.startsWith(prefix)); return value ? decodeURIComponent(value.slice(prefix.length)) : ''; } @@ -90,13 +93,16 @@ async function getPasskey(options) { } async function permitted(permission, path) { if (!can(permission)) return []; return api(path); } +function defaultUsageQuery() { const to=new Date();const from=new Date(to.getTime()-29*86400000);return new URLSearchParams({from:from.toISOString().slice(0,10),to:to.toISOString().slice(0,10)}).toString(); } async function loadAll(knownSession = null) { try { const session = knownSession || await api('/me'); state.actor = session.actor || {}; state.permissions = new Set(session.permissions || []); state.overview = await api('/overview'); + const usageQuery=defaultUsageQuery(); const results = await Promise.all([ permitted('tenants.read','/tenants'), permitted('projects.read','/projects'), permitted('keys.read','/keys'), - permitted('platform.read','/providers'), permitted('platform.read','/models'), permitted('usage.read','/usage'), + permitted('platform.read','/providers'), permitted('platform.read','/models'), permitted('overview.read','/developer/config'), + permitted('overview.read','/developer/models'), can('preferences.read') ? api('/developer/preferences') : {}, permitted('usage.read',`/usage?${usageQuery}`), permitted('usage.read','/usage/summary'), permitted('limits.read','/limits'), permitted('users.read','/users'), permitted('audit.read','/audit'), state.overview.billing_enabled ? permitted('billing.read','/billing/accounts') : [], state.overview.billing_enabled ? permitted('billing.read','/billing/ledger') : [], @@ -105,9 +111,12 @@ async function loadAll(knownSession = null) { 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.overview.billing_enabled ? permitted('billing.read','/billing/invoices') : [], + permitted('usage.read',`/usage/daily?${usageQuery}`), permitted('usage.read',`/usage/analytics?${usageQuery}`), + state.overview.billing_enabled && state.actor.tenant_id && can('billing.read') ? api('/billing/auto-topup') : {}, + state.overview.billing_enabled && state.actor.tenant_id && can('billing.read') ? api('/billing/profile') : {} ]); - [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; + [state.tenants,state.projects,state.keys,state.providers,state.models,state.developerConfig,state.developerModels,state.preferences,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,state.usageDaily,state.usageAnalytics,state.autoTopUp,state.billingProfile] = results; renderAll(); setConnected(true); return true; } catch (error) { setConnected(false); if (error.status !== 401) toast(error.message, true); return false; } } @@ -116,11 +125,16 @@ function applyPermissions() { $$('[data-permission]').forEach(node => node.classList.toggle('hidden', !can(node.dataset.permission))); $('#billing-tab').classList.toggle('hidden', !state.overview.billing_enabled || !can('billing.read')); $('#topup-form').classList.toggle('hidden', !state.overview.stripe_enabled || !can('billing.topup')); + $('#billing-portal').classList.toggle('hidden', !state.overview.stripe_enabled || !can('billing.topup')); $('#account-tab').classList.toggle('hidden', !state.actor.id); - const active = $('.tab.active'); if (active?.classList.contains('hidden')) $('.tab[data-section="overview"]').click(); + $('#developer-preferences-form').classList.toggle('hidden', !state.actor.tenant_id || !can('preferences.read')); + $('#billing-preferences-form').classList.toggle('hidden', !state.actor.tenant_id || !state.overview.billing_enabled || !can('billing.read')); + $('#auto-topup-panel').classList.toggle('hidden', !state.actor.tenant_id || !state.overview.billing_enabled || !can('billing.read')); + $('#billing-profile-panel').classList.toggle('hidden', !state.actor.tenant_id || !state.overview.billing_enabled || !can('billing.read')); + const active = $('.tab.active'); if (active?.classList.contains('hidden')) $('.tab[data-section="quickstart"]').click(); } function renderAll() { - applyPermissions(); renderOverview(); renderTenants(); renderProjects(); renderKeys(); renderProviders(); renderModels(); renderBilling(); + applyPermissions(); renderOverview(); renderQuickstart(); renderPreferences(); renderCatalog(); renderTenants(); renderProjects(); renderKeys(); renderProviders(); renderModels(); renderBilling(); renderUsage(); renderLimits(); renderUsers(); renderAudit(); renderRouteEditor(); renderSecurity(); } function renderOverview() { @@ -131,23 +145,228 @@ function renderOverview() { $('#metrics').innerHTML = items.map(([label,value,sub]) => `<article class="metric"><span>${label}</span><strong>${esc(value)}</strong><small>${esc(sub)}</small></article>`).join(''); $('#overview-usage-body').innerHTML = current.map(item => `<tr><td><strong>${esc(item.project_name)}</strong></td><td>${integer(item.request_count)}</td><td>${percent(item.successful_requests,item.request_count)}</td><td>${integer(item.total_tokens)}</td><td>${money(item.cost_micros)}</td><td class="${item.uncollected_micros ? 'money-negative':''}">${money(item.uncollected_micros)}</td></tr>`).join('') || emptyRow(6); } +function goTo(section) { const node=$(`.tab[data-section="${section}"]`); if(node){ node.click(); window.scrollTo({top:0,behavior:'smooth'}); } } +function latestAvailableBalance() { return state.billingAccounts.find(item => !state.actor.tenant_id || item.tenant_id===state.actor.tenant_id)?.available_micros || 0; } +function inferenceBaseURL() { return String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,''); } +function developerEndpoint(name) { const defaults={chat_completions:'/v1/chat/completions',responses:'/v1/responses',messages:'/anthropic/v1/messages',models:'/v1/models'};const path=state.developerConfig.endpoints?.[name]||defaults[name]||'';return `${inferenceBaseURL()}${path}`; } +function modelIsUnavailable(model) { return model?.health_status==='unavailable'; } +function developerModelOption(item) { const unavailable=modelIsUnavailable(item);return `<option value="${esc(item.public_id)}" ${unavailable?'disabled':''}>${esc(item.display_name||item.public_id)}${unavailable?' (unavailable)':''}</option>`; } +function providerHealthSummary(item) { + const total=Number(item.provider_count||0);const available=Number(item.available_provider_count||0);const status=item.health_status||'online'; + if(status==='unavailable')return {status,label:'Unavailable',detail:'No provider is currently accepting traffic'}; + if(status==='degraded')return {status,label:'Degraded',detail:`${available} of ${total} providers accepting traffic`}; + return {status:'online',label:'Online',detail:total?`${available} of ${total} providers accepting traffic`:'Routing is available'}; +} +function renderQuickstart() { + const balance=latestAvailableBalance(); + const balanceText=state.overview.billing_enabled && can('billing.read') ? `Available balance ${money(balance)}` : state.overview.billing_enabled ? 'Balance managed by billing admin' : 'Prepaid billing disabled'; + $('#quickstart-balance').textContent=balanceText; + const hasKey=state.keys.some(item=>item.status==='active'); + const hasUsage=state.usage.some(item=>item.success); + const funded=!state.overview.billing_enabled || !can('billing.read') || balance>0; + const steps=[ + {title:'Workspace ready',body:'Your account and default project are ready.',done:Boolean(state.actor.tenant_id||state.actor.bootstrap),action:'projects'}, + {title:'Add balance',body:state.overview.billing_enabled?(!can('billing.read')?'Billing is managed by workspace billing members.':funded?'Funds are available for inference.':'Add funds before the first billable request.'):'Billing is not enabled.',done:funded,action:can('billing.read')?'billing':'team'}, + {title:'Create an API key',body:hasKey?'An active key can call the gateway.':'Create a key and copy its secret once.',done:hasKey,action:hasKey?'keys':'',target:hasKey?'':'starter-key-form'}, + {title:'Make a request',body:hasUsage?'Your first successful request is recorded.':'Run the example below and watch it appear here.',done:hasUsage,action:hasUsage?'usage':'',target:hasUsage?'':'playground-form'} + ]; + $('#onboarding').innerHTML=steps.map(step=>`<button type="button" class="onboarding-step ${step.done?'done':''}" ${step.action?`data-goto="${esc(step.action)}"`:''} ${step.target?`data-scroll-to="${esc(step.target)}"`:''}><strong>${step.done?'✓ ':''}${esc(step.title)}</strong><small>${esc(step.body)}</small><span class="step-state">${step.done?'Complete':'Open step →'}</span></button>`).join(''); + const signals=[['Models available',integer(state.developerModels.length)],['Active API keys',integer(state.keys.filter(item=>item.status==='active').length)],['This month',money(state.usageSummary.reduce((sum,item)=>sum+Number(item.cost_micros||0),0))],['Auto top-up',state.autoTopUp.enabled?'Enabled':state.autoTopUp.payment_method_configured?'Ready':'Not configured'],['Fallback model',state.preferences.fallback_model||'Not configured'],['Runtime',state.overview.redis_connected?'Live propagation':'PG fallback']]; + $('#quickstart-signals').innerHTML=signals.map(([label,value])=>`<div class="signal"><span>${esc(label)}</span><strong>${esc(value)}</strong></div>`).join(''); + const recent=state.usage.slice(0,5); + $('#quickstart-usage-body').innerHTML=recent.map(item=>`<tr><td>${date(item.started_at)}</td><td><strong>${esc(item.public_model)}</strong></td><td><span class="badge ${item.success?'active':'suspended'}">${item.status_code}</span></td><td>${integer(item.total_tokens)}</td><td>${money(item.charged_micros||item.cost_micros)}</td><td>${integer(item.duration_ms)} ms</td></tr>`).join('')||emptyRow(6); + const modelSelect=$('#quickstart-model'); const previous=modelSelect.value||state.preferences.default_model; + modelSelect.innerHTML=state.developerModels.map(developerModelOption).join('')||'<option value="">No available model</option>'; + if(state.developerModels.some(item=>item.public_id===previous&&!modelIsUnavailable(item)))modelSelect.value=previous; + else{const firstAvailable=state.developerModels.find(item=>!modelIsUnavailable(item));if(firstAvailable)modelSelect.value=firstAvailable.public_id;} + const playgroundModel=$('#playground-model'); const playgroundPrevious=playgroundModel.value||modelSelect.value; + playgroundModel.innerHTML=modelSelect.innerHTML; + if(state.developerModels.some(item=>item.public_id===playgroundPrevious&&!modelIsUnavailable(item)))playgroundModel.value=playgroundPrevious; + if(state.playgroundKey)$('#playground-key').value=state.playgroundKey; + $('#playground-endpoint').textContent=inferenceBaseURL(); + syncQuickstartProtocols(); syncPlaygroundProtocols(); renderQuickstartCode(); renderDeveloperAccess(); +} +function renderDeveloperAccess() { + const form=$('#starter-key-form');const projects=state.projects.filter(item=>item.status==='active'&&(!state.actor.tenant_id||item.tenant_id===state.actor.tenant_id));const select=$('#starter-project');const previous=select.value; + select.innerHTML=projects.map(item=>`<option value="${esc(item.id)}">${esc(item.name)}</option>`).join('')||'<option value="">No active project</option>'; + if(projects.some(item=>item.id===previous))select.value=previous;else{const preferred=projects.find(item=>item.slug==='default')||projects[0];if(preferred)select.value=preferred.id;} + const model=selectedDeveloperModel();$('#starter-model-label').textContent=model.public_id||'No model'; + form.classList.toggle('hidden',!state.actor.tenant_id||!can('keys.write')); + $('#starter-key-submit').disabled=!select.value||!model.public_id||modelIsUnavailable(model); + const rows=[['OpenAI SDK base',`${inferenceBaseURL()}/v1`],['Anthropic SDK base',`${inferenceBaseURL()}/anthropic`],['Chat Completions',developerEndpoint('chat_completions')],['Responses',developerEndpoint('responses')],['Anthropic Messages',developerEndpoint('messages')],['Models',developerEndpoint('models')]]; + $('#endpoint-list').innerHTML=rows.map(([label,value],index)=>`<div class="endpoint-row"><span>${esc(label)}</span><code>${esc(value)}</code><button class="text-button" type="button" data-copy-endpoint="${index}">Copy</button></div>`).join(''); + $('#endpoint-list').dataset.values=JSON.stringify(rows.map(([,value])=>value)); + $('#endpoint-env').textContent=`export AIGW_API_KEY="your-key"\nexport OPENAI_BASE_URL="${inferenceBaseURL()}/v1"\nexport ANTHROPIC_BASE_URL="${inferenceBaseURL()}/anthropic"`; +} +function renderPreferences() { + const preferences=state.preferences||{}; const modelOptions=state.developerModels.map(item=>`<option value="${esc(item.public_id)}">${esc(item.display_name||item.public_id)}</option>`).join(''); + const defaultSelect=$('#preference-default-model'); const fallbackSelect=$('#preference-fallback-model'); + defaultSelect.innerHTML=`<option value="">First available model</option>${modelOptions}`; + fallbackSelect.innerHTML=`<option value="">No workspace fallback</option>${modelOptions}`; + if(state.developerModels.some(item=>item.public_id===preferences.default_model))defaultSelect.value=preferences.default_model; + if(state.developerModels.some(item=>item.public_id===preferences.fallback_model))fallbackSelect.value=preferences.fallback_model; + const canWriteDeveloper=can('developer.preferences.write'); defaultSelect.disabled=!canWriteDeveloper; fallbackSelect.disabled=!canWriteDeveloper; + $('#low-balance-enabled').checked=preferences.low_balance_enabled!==false; + $('#low-balance-threshold').value=scaledToDecimal(preferences.low_balance_threshold_micros||0,6); + const canWriteBilling=can('billing.preferences.write'); $('#low-balance-enabled').disabled=!canWriteBilling; $('#low-balance-threshold').disabled=!canWriteBilling; + $('#balance-alert-status').textContent=state.authConfig.email_delivery_enabled?'Verified billing members receive at most one low-balance alert per day.':'The preference is saved now and activates when SMTP delivery is configured.'; +} +function selectedDeveloperModel() { return state.developerModels.find(item=>item.public_id===$('#quickstart-model')?.value) || state.developerModels.find(item=>!modelIsUnavailable(item)) || state.developerModels[0] || {}; } +function syncQuickstartProtocols() { + const model=selectedDeveloperModel(); const select=$('#quickstart-protocol'); const previous=select.value; + select.innerHTML=(model.supported_wire_apis||[]).map(apiName=>`<option value="${esc(apiName)}">${esc(apiName==='chat_completions'?'OpenAI Chat Completions':apiName==='responses'?'OpenAI Responses':apiName==='messages'?'Anthropic Messages':apiName)}</option>`).join('')||'<option value="chat_completions">OpenAI Chat Completions</option>'; + if((model.supported_wire_apis||[]).includes(previous))select.value=previous; + syncProviderSelect('#quickstart-provider',model,select.value); +} +function selectedPlaygroundModel() { return state.developerModels.find(item=>item.public_id===$('#playground-model')?.value) || state.developerModels.find(item=>!modelIsUnavailable(item)) || state.developerModels[0] || {}; } +function syncPlaygroundProtocols() { + const model=selectedPlaygroundModel(); const select=$('#playground-protocol'); const previous=select.value; + select.innerHTML=(model.supported_wire_apis||[]).map(apiName=>`<option value="${esc(apiName)}">${esc(apiName==='chat_completions'?'OpenAI Chat Completions':apiName==='responses'?'OpenAI Responses':apiName==='messages'?'Anthropic Messages':apiName)}</option>`).join('')||'<option value="chat_completions">OpenAI Chat Completions</option>'; + if((model.supported_wire_apis||[]).includes(previous))select.value=previous; + syncProviderSelect('#playground-provider',model,select.value); + const max=Number(model.max_output_tokens||4096);$('#playground-max-output').max=String(max>0?max:4096); + if(Number($('#playground-max-output').value)>max&&max>0)$('#playground-max-output').value=String(max); +} +function syncProviderSelect(selector,model,wire) { + const select=$(selector);const previous=select.value;const providers=(model.providers||[]).filter(item=>item.wire_api===wire); + select.innerHTML=`<option value="">Automatic routing</option>${providers.map(item=>`<option value="${esc(item.slug)}" ${item.state==='open'?'disabled':''}>${esc(item.name||item.slug)} (${esc(item.slug)})${item.state==='open'?' — unavailable':''}</option>`).join('')}`; + if(providers.some(item=>item.slug===previous&&item.state!=='open'))select.value=previous; +} +function modelSelector(model,providerSelector) { const provider=$(providerSelector)?.value||'';const publicID=model.public_id||'model-id';return provider?`${publicID}:${provider}`:publicID; } +function quickstartEndpoint(wire) { return wire==='messages'?developerEndpoint('messages'):wire==='responses'?developerEndpoint('responses'):developerEndpoint('chat_completions'); } +function renderQuickstartCode() { + const model=selectedDeveloperModel(); const selectedModel=modelSelector(model,'#quickstart-provider'); const wire=$('#quickstart-protocol').value||'chat_completions'; const language=$('#quickstart-language').value||'curl'; const endpoint=quickstartEndpoint(wire); const key='$AIGW_API_KEY'; + let code=''; + if(language==='curl') { + const headers=wire==='messages'?`-H "x-api-key: ${key}"\n -H "anthropic-version: 2023-06-01"`:`-H "Authorization: Bearer ${key}"`; + const body=wire==='messages'?`{"model":"${selectedModel}","max_tokens":256,"messages":[{"role":"user","content":"Say hello in one sentence."}]}`:wire==='responses'?`{"model":"${selectedModel}","input":"Say hello in one sentence."}`:`{"model":"${selectedModel}","messages":[{"role":"user","content":"Say hello in one sentence."}],"stream":false}`; + code=`export AIGW_API_KEY="your-key"\ncurl ${endpoint} \\\n ${headers} \\\n -H "Content-Type: application/json" \\\n -d '${body}'`; + } else if(language==='python') { + code=wire==='messages'?`import os\nfrom anthropic import Anthropic\n\nclient = Anthropic(api_key=os.environ["AIGW_API_KEY"], base_url="${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/anthropic")\nmessage = client.messages.create(model="${selectedModel}", max_tokens=256, messages=[{"role":"user", "content":"Say hello in one sentence."}])\nprint(message.content[0].text)`:wire==='responses'?`import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/v1")\nresponse = client.responses.create(model="${selectedModel}", input="Say hello in one sentence.")\nprint(response.output_text)`: `import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/v1")\nresponse = client.chat.completions.create(model="${selectedModel}", messages=[{"role":"user", "content":"Say hello in one sentence."}])\nprint(response.choices[0].message.content)`; + } else { + code=wire==='messages'?`import Anthropic from "@anthropic-ai/sdk";\n\nconst client = new Anthropic({ apiKey: process.env.AIGW_API_KEY, baseURL: "${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/anthropic" });\nconst message = await client.messages.create({ model: "${selectedModel}", max_tokens: 256, messages: [{ role: "user", content: "Say hello in one sentence." }] });\nconsole.log(message.content[0].text);`: `import OpenAI from "openai";\n\nconst client = new OpenAI({ apiKey: process.env.AIGW_API_KEY, baseURL: "${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/v1" });\nconst response = await client.${wire==='responses'?'responses.create({ model: "'+selectedModel+'", input: "Say hello in one sentence." })':'chat.completions.create({ model: "'+selectedModel+'", messages: [{ role: "user", content: "Say hello in one sentence." }] })'};\nconsole.log(${wire==='responses'?'response.output_text':'response.choices[0].message.content'});`; + } + $('#quickstart-code-block code').textContent=code; +} +function playgroundBody(wire,model,prompt,maxOutput) { + if(wire==='messages')return {model,max_tokens:maxOutput,messages:[{role:'user',content:prompt}]}; + if(wire==='responses')return {model,input:prompt,max_output_tokens:maxOutput,stream:false}; + return {model,messages:[{role:'user',content:prompt}],max_tokens:maxOutput,stream:false}; +} +function playgroundTokenCount(payload) { + const usage=payload?.usage;if(!usage)return 0; + return Number(usage.total_tokens ?? (Number(usage.input_tokens||usage.prompt_tokens||0)+Number(usage.output_tokens||usage.completion_tokens||0))); +} +function playgroundDiagnosis(status,network=false) { + if(network)return {title:'Gateway unreachable',message:'Verify the inference public URL, TLS certificate, and exact admin Origin allowed by the gateway.',action:'Open API keys',target:'keys'}; + if(status===400)return {title:'Request rejected',message:'Check the selected protocol, prompt, and output limit against the model details.',action:'Review model',target:'catalog'}; + if(status===401||status===403)return {title:'API key not authorized',message:'Use an active inference key and check its expiry, scopes, model allowlist, and monthly spend cap.',action:'Open API keys',target:'keys'}; + if(status===402)return {title:'Balance required',message:'Add funds to the prepaid wallet, then retry this request. Automatic top-up can prevent future interruptions.',action:'Add funds',target:'billing'}; + if(status===404)return {title:'Model route unavailable',message:'The model may be unavailable to this key or may not support the selected protocol.',action:'Choose a model',target:'catalog'}; + if(status===409)return {title:'Request conflict',message:'Retry with a new request after the in-flight billing reservation or account change completes.',action:'View usage',target:'usage'}; + if(status===429)return {title:'Limit reached',message:'Review request, token, concurrency, project, and API key limits before retrying.',action:'Review limits',target:'limits'}; + if(status>=500)return {title:'Provider unavailable',message:'The gateway could not complete the upstream request after routing and failover. Keep the request ID for support.',action:'Choose another model',target:'catalog'}; + return {title:'Request failed',message:'Review the response body and request ID. The request was not accepted as successful.',action:'View usage',target:'usage'}; +} +function hidePlaygroundDiagnostic(){const node=$('#playground-diagnostic');node.classList.add('hidden');$('#playground-diagnostic-action').classList.add('hidden');} +function showPlaygroundDiagnostic(status,network=false) { + const diagnosis=playgroundDiagnosis(status,network);const action=$('#playground-diagnostic-action'); + $('#playground-diagnostic-title').textContent=diagnosis.title;$('#playground-diagnostic-message').textContent=diagnosis.message; + action.textContent=diagnosis.action;action.dataset.target=diagnosis.target;action.classList.toggle('hidden',!$(`.tab[data-section="${diagnosis.target}"]`)||$(`.tab[data-section="${diagnosis.target}"]`).classList.contains('hidden')); + $('#playground-diagnostic').classList.remove('hidden'); +} +async function runPlayground(event) { + event.preventDefault(); + const key=$('#playground-key').value.trim();const model=modelSelector(selectedPlaygroundModel(),'#playground-provider');const wire=$('#playground-protocol').value;const prompt=$('#playground-prompt').value.trim();const maxOutput=Number($('#playground-max-output').value); + if(!key||!model||!prompt||!Number.isSafeInteger(maxOutput)||maxOutput<1)return toast('Complete the API request fields',true); + state.playgroundKey=key; + const endpoint=quickstartEndpoint(wire);const headers={'Content-Type':'application/json'}; + if(wire==='messages'){headers['X-API-Key']=key;headers['Anthropic-Version']='2023-06-01';}else headers.Authorization=`Bearer ${key}`; + const result=$('#playground-result');const status=$('#playground-status');const send=$('#playground-send');const stop=$('#playground-stop'); + result.classList.remove('hidden');hidePlaygroundDiagnostic();status.className='badge';status.textContent='Sending';$('#playground-request-id').textContent='';$('#playground-duration').textContent='';$('#playground-tokens').textContent='';$('#playground-response').textContent=''; + send.disabled=true;stop.classList.remove('hidden');const controller=new AbortController();state.playgroundController=controller;const started=performance.now(); + try { + const response=await fetch(endpoint,{method:'POST',mode:'cors',credentials:'omit',cache:'no-store',headers,body:JSON.stringify(playgroundBody(wire,model,prompt,maxOutput)),signal:controller.signal}); + const raw=await response.text();let payload;try{payload=raw?JSON.parse(raw):{};}catch{payload=raw;} + const requestID=response.headers.get('X-AIGW-Request-ID')||payload?.error?.request_id||'';const duration=Math.round(performance.now()-started);const tokens=playgroundTokenCount(payload); + status.className=`badge ${response.ok?'active':'suspended'}`;status.textContent=`${response.status} ${response.ok?'OK':'Error'}`; + $('#playground-request-id').textContent=requestID?`Request ${requestID}`:'';$('#playground-duration').textContent=`${integer(duration)} ms`;$('#playground-tokens').textContent=tokens?`${integer(tokens)} tokens`:''; + $('#playground-response').textContent=typeof payload==='string'?payload:JSON.stringify(payload,null,2); + if(!response.ok){showPlaygroundDiagnostic(response.status);toast(payload?.error?.message||`Request failed (${response.status})`,true);}else{hidePlaygroundDiagnostic();toast('Request completed');setTimeout(()=>loadAll(),900);} + } catch(error) { + const stopped=error.name==='AbortError';status.className='badge suspended';status.textContent=stopped?'Stopped':'Network error';$('#playground-duration').textContent=`${integer(Math.round(performance.now()-started))} ms`;$('#playground-response').textContent=stopped?'Request cancelled.':error.message;if(stopped)hidePlaygroundDiagnostic();else showPlaygroundDiagnostic(0,true);toast(stopped?'Request stopped':error.message,true); + } finally { + if(state.playgroundController===controller)state.playgroundController=null;send.disabled=false;stop.classList.add('hidden'); + } +} +function renderCatalog() { + const ownerSelect=$('#catalog-owner');const ownerValue=ownerSelect?.value||'';const owners=[...new Set(state.developerModels.map(item=>String(item.owned_by||'').trim()).filter(Boolean))].sort((a,b)=>a.localeCompare(b)); + if(ownerSelect){ownerSelect.innerHTML=`<option value="">All developers</option>${owners.map(owner=>`<option value="${esc(owner)}">${esc(owner)}</option>`).join('')}`;if(owners.includes(ownerValue))ownerSelect.value=ownerValue;} + const search=String($('#catalog-search')?.value||'').trim().toLowerCase(); const protocol=$('#catalog-protocol')?.value||''; const input=$('#catalog-input')?.value||'';const owner=ownerSelect?.value||'';const sort=$('#catalog-sort')?.value||'newest'; + const models=state.developerModels.filter(item=>{const hay=[item.public_id,item.display_name,item.owned_by,item.description,...(item.capabilities||[])].join(' ').toLowerCase();return (!search||hay.includes(search))&&(!protocol||(item.supported_wire_apis||[]).includes(protocol))&&(!input||(item.input_modalities||[]).includes(input))&&(!owner||item.owned_by===owner);}); + const compareText=(a,b)=>String(a.display_name||a.public_id).localeCompare(String(b.display_name||b.public_id)); + models.sort((a,b)=>sort==='name'?compareText(a,b):sort==='input_price'?Number(a.input_price_micros_per_million||0)-Number(b.input_price_micros_per_million||0)||compareText(a,b):sort==='output_price'?Number(a.output_price_micros_per_million||0)-Number(b.output_price_micros_per_million||0)||compareText(a,b):sort==='context'?Number(b.context_window||0)-Number(a.context_window||0)||compareText(a,b):new Date(b.released_at||0)-new Date(a.released_at||0)||compareText(a,b)); + $('#catalog-count').textContent=`${models.length} of ${state.developerModels.length} models`; + $('#catalog-grid').innerHTML=models.map(item=>{const health=providerHealthSummary(item);return `<article class="catalog-card"><div><span class="eyebrow">${esc(item.owned_by||'MODEL')}</span><h2>${esc(item.display_name||item.public_id)}</h2><div class="model-id">${esc(item.public_id)}</div></div><div class="provider-summary"><span class="badge ${health.status==='online'?'active':health.status}">${esc(health.label)}</span><small>${esc(health.detail)}</small></div><p>${esc(item.description||'No description provided.')}</p><div class="catalog-meta">${(item.supported_wire_apis||[]).map(apiName=>`<span class="tag">${esc(apiName)}</span>`).join('')}${(item.input_modalities||[]).map(modality=>`<span class="tag">${esc(modality)} input</span>`).join('')}</div><div class="catalog-price">${money(item.input_price_micros_per_million,item.price_currency)} in · ${money(item.output_price_micros_per_million,item.price_currency)} out / 1M tokens</div><small class="muted">${integer(item.context_window)} context · ${integer(item.max_output_tokens)} max output</small><div class="catalog-actions"><button type="button" class="button subtle" data-model-details="${esc(item.public_id)}">Details</button><button type="button" class="button secondary" data-use-model="${esc(item.public_id)}" ${modelIsUnavailable(item)?'disabled':''}>Use this model</button></div></article>`;}).join('')||'<div class="panel empty">No models match these filters.</div>'; +} +function modelProtocolLabel(value) { return value==='chat_completions'?'OpenAI Chat Completions':value==='responses'?'OpenAI Responses':value==='messages'?'Anthropic Messages':value; } +function showModelDetails(publicID) { + const model=state.developerModels.find(item=>item.public_id===publicID);if(!model)return; + state.detailModel=model;$('#model-dialog-title').textContent=model.display_name||model.public_id;$('#model-dialog-id').textContent=model.public_id; + const lifecycle=model.lifecycle||'active';const status=lifecycle==='retired'?'retired':lifecycle==='deprecated'?'deprecated':lifecycle==='preview'?'preview':'available'; + const health=providerHealthSummary(model);const rows=[['Status',status],['Runtime',health.label],['Providers',health.detail],['Owner',model.owned_by||'—'],['Protocols',(model.supported_wire_apis||[]).map(modelProtocolLabel).join(', ')||'—'],['Input',(model.input_modalities||[]).join(', ')||'—'],['Output',(model.output_modalities||[]).join(', ')||'—'],['Context',`${integer(model.context_window)} tokens`],['Max output',`${integer(model.max_output_tokens)} tokens`],['Released',date(model.released_at)],['Capabilities',(model.capabilities||[]).join(', ')||'—'],['Aliases',(model.aliases||[]).join(', ')||'—']]; + $('#model-detail-grid').innerHTML=rows.map(([label,value])=>`<div class="model-detail-row"><span>${esc(label)}</span><strong>${esc(value)}</strong></div>`).join(''); + const providers=model.providers||[];$('#model-provider-health').classList.toggle('hidden',providers.length===0);$('#model-provider-health-body').innerHTML=providers.map(item=>{const measured=Number(item.recent_samples||0)>0;const availability=measured?`${Number(item.availability_percent||0).toFixed(1)}%`:'Not measured';const latency=Number(item.header_latency_ewma_ms||0)>0?`${integer(item.header_latency_ewma_ms)} ms`:'Not measured';const retry=item.circuit_open_until?`Retry ${date(item.circuit_open_until)}`:'';return `<tr><td><strong>${esc(item.name||item.slug)}</strong><small class="price-line">${esc(item.slug)} · ${esc(modelProtocolLabel(item.wire_api)||item.protocol||'')}</small></td><td><span class="badge ${item.state==='healthy'?'active':item.state==='open'?'suspended':item.state==='degraded'?'degraded':''}">${esc(item.state==='open'?'Circuit open':item.state)}</span>${retry?`<small class="provider-retry">${esc(retry)}</small>`:''}</td><td>${esc(availability)}</td><td>${esc(latency)}</td><td>${integer(item.attempts||0)}</td></tr>`;}).join(''); + $('#model-dialog-playground').disabled=modelIsUnavailable(model);$('#estimate-input').value='1000';$('#estimate-output').value='500';$('#estimate-cache-read').value='0';$('#estimate-cache-write').value='0';renderModelEstimate();$('#model-dialog').showModal(); +} +function renderModelEstimate() { + const model=state.detailModel;if(!model)return;const input=Math.max(0,Number($('#estimate-input').value||0));const output=Math.max(0,Number($('#estimate-output').value||0));const cacheRead=Math.max(0,Number($('#estimate-cache-read').value||0));const cacheWrite=Math.max(0,Number($('#estimate-cache-write').value||0));const micros=(input*Number(model.input_price_micros_per_million||0)+output*Number(model.output_price_micros_per_million||0)+cacheRead*Number(model.cache_read_price_micros_per_million||0)+cacheWrite*Number(model.cache_write_price_micros_per_million||0))/1_000_000;$('#model-estimate').textContent=`Estimated cost ${money(Math.round(micros),model.price_currency||state.overview.billing_currency||'usd')}`; +} +function useModelInPlayground(publicID) { const model=state.developerModels.find(item=>item.public_id===publicID);if(!model||modelIsUnavailable(model)){toast('This model has no provider currently accepting traffic',true);return;}$('#quickstart-model').value=publicID;$('#playground-model').value=publicID;syncQuickstartProtocols();syncPlaygroundProtocols();$('#quickstart-provider').value='';$('#playground-provider').value='';renderQuickstartCode();if($('#model-dialog').open)$('#model-dialog').close();goTo('quickstart'); } function renderTenants() { $('#tenants-body').innerHTML = state.tenants.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><code>${esc(item.slug)}</code></td><td><span class="badge ${item.status}">${esc(item.status)}</span></td><td>${date(item.created_at)}</td></tr>`).join('') || emptyRow(4); ['project-tenant','key-tenant','topup-tenant','adjustment-tenant','user-tenant'].forEach(id => { const node=$(`#${id}`); if (node) node.innerHTML=selectOptions(state.tenants,'id','name', state.actor.tenant_id ? 'Current tenant' : 'Select tenant…'); }); if (state.actor.tenant_id) ['project-tenant','key-tenant','topup-tenant','adjustment-tenant','user-tenant'].forEach(id => { const node=$(`#${id}`); if (node) node.value=state.actor.tenant_id; }); } function renderProjects() { $('#projects-body').innerHTML = state.projects.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><code>${shortID(item.tenant_id)}</code></td><td>${esc(item.slug)}</td><td><span class="badge ${item.status}">${esc(item.status)}</span></td></tr>`).join('') || emptyRow(4); renderKeyProjects(); } -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 renderKeyProjects() { const tenant=$('#key-tenant').value;const node=$('#key-project');const selected=node.value;const projects=state.projects.filter(item=>!tenant||item.tenant_id===tenant);node.innerHTML=selectOptions(projects,'id','name');if(projects.some(item=>item.id===selected))node.value=selected;else if(projects.length===1)node.value=projects[0].id; } +function renderKeys() { + const picker=$('#key-models');const selected=new Set([...picker.selectedOptions].map(option=>option.value));picker.innerHTML=state.developerModels.map(item=>`<option value="${esc(item.public_id)}">${esc(item.display_name||item.public_id)} (${esc(item.public_id)})</option>`).join('');[...picker.options].forEach(option=>{option.selected=selected.has(option.value);}); + $('#keys-body').innerHTML = state.keys.map(item => { + const models=item.allowed_models||[];const expires=item.expires_at?date(item.expires_at):'Never';const tags=(item.tags||[]).map(tag=>`<span class="tag">${esc(tag)}</span>`).join(''); + const spent=Number(item.current_month_spend_micros||0);const reserved=Number(item.current_month_reserved_micros||0);const cap=Number(item.monthly_spend_micros||0);const remaining=Math.max(0,cap-spent-reserved); + const effectiveStatus=item.status==='active'&&item.expires_at&&new Date(item.expires_at)<=new Date()?'expired':item.status; + return `<tr><td><strong>${esc(item.name)}</strong>${tags?`<small class="key-tags">${tags}</small>`:''}</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('')}<small class="price-line">${models.length?`${integer(models.length)} selected model${models.length===1?'':'s'}`:'All visible models'}</small></td><td><strong>${money(spent)}</strong><small class="price-line">${integer(item.current_month_requests)} requests · ${money(reserved)} reserved</small></td><td>${cap?money(cap):'Unlimited'}<small class="price-line">${cap?`${money(remaining)} remaining`:'No key-level cap'} · Expires: ${esc(expires)}</small></td><td><span class="badge ${effectiveStatus}">${esc(effectiveStatus)}</span><small class="price-line">Last used: ${esc(date(item.last_used_at))}</small></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(8); +} +function renderProviders() { $('#providers-body').innerHTML = state.providers.map(item => `<tr><td><strong>${esc(item.name)}</strong><small class="price-line"><code>${esc(item.slug)}</code></small></td><td><span class="tag">${esc(item.protocol)}</span><small class="price-line">${esc(item.wire_api)}</small></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><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); + const orders=state.orders||[];$('#billing-orders-body').innerHTML=orders.map(item=>`<tr><td>${date(item.created_at)}${item.trigger_type==='auto'?'<small class="price-line">automatic</small>':''}</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>${item.trigger_type!=='auto'&&['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); + renderBillingProfile(); + renderAutoTopUp(); +} +function renderBillingProfile(){ + if(!state.actor.tenant_id)return;const item=state.billingProfile||{};const form=$('#billing-profile-form'); + for(const name of ['legal_name','billing_email','address_line1','address_line2','city','region','postal_code','country'])form.elements[name].value=item[name]||''; + const labels={not_configured:'Not configured',checkout_managed:'Stripe managed',disabled:'Saved locally',pending:'Syncing',synced:'Stripe synced',failed:'Sync needed'};const status=$('#billing-profile-status');const label=labels[item.stripe_sync_status]||'Not configured';status.textContent=label;status.className=`badge ${item.stripe_sync_status==='synced'?'active':item.stripe_sync_status==='failed'?'suspended':''}`; + $('#billing-profile-error').textContent=item.stripe_sync_error||'';form.querySelector('button[type=submit]').disabled=!can('billing.topup'); + $('#billing-portal').disabled=!item.stripe_customer_configured; +} +function renderAutoTopUp(){ + const item=state.autoTopUp||{};if(!state.actor.tenant_id)return; + const currency=item.currency||state.overview.billing_currency||'usd';const digits=currencyDigits(currency); + $('#auto-topup-enabled').checked=Boolean(item.enabled);$('#auto-topup-threshold').value=scaledToDecimal(item.threshold_micros||0,6);$('#auto-topup-amount').value=scaledToDecimal(item.topup_amount_minor||0,digits); + const status=$('#auto-topup-status');status.textContent=String(item.status||'not_configured').replaceAll('_',' ');status.className=`badge ${item.enabled&&item.status==='ready'?'active':item.status==='action_required'||item.status==='failed'?'suspended':''}`; + $('#auto-topup-payment-method').textContent=item.payment_method_configured?`${item.payment_method_brand||item.payment_method_type||'payment method'} •••• ${item.payment_method_last4||''}${item.payment_method_exp_month?` · ${String(item.payment_method_exp_month).padStart(2,'0')}/${item.payment_method_exp_year}`:''}`:'Not saved'; + $('#auto-topup-error').textContent=item.last_error||''; + const writable=can('billing.topup');$$('#auto-topup-form input').forEach(node=>node.disabled=!writable);$('#auto-topup-form button[type=submit]').disabled=!writable;$('#auto-topup-payment-setup').disabled=!writable||!item.stripe_enabled; + $('#auto-topup-payment-setup').textContent=item.payment_method_configured?'Replace payment method':'Save payment method'; } function renderSecurity() { const totp=Boolean(state.mfa?.totp_enabled); @@ -160,8 +379,30 @@ function renderSecurity() { $('#orders-body').innerHTML=(state.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><code>${shortID(item.id)}</code></td></tr>`).join('')||emptyRow(4); } function renderUsage() { + const project=$('#usage-project'),key=$('#usage-key'),model=$('#usage-model'),provider=$('#usage-provider');const projectValue=project.value,keyValue=key.value,modelValue=model.value,providerValue=provider.value; + const providers=new Map();state.developerModels.forEach(item=>(item.providers||[]).forEach(route=>providers.set(route.slug,route.name||route.slug))); + project.innerHTML=`<option value="">All projects</option>${state.projects.map(item=>`<option value="${esc(item.id)}">${esc(item.name)}</option>`).join('')}`;key.innerHTML=`<option value="">All API keys</option>${state.keys.map(item=>`<option value="${esc(item.id)}">${esc(item.name)} (${esc(item.key_prefix)})</option>`).join('')}`;model.innerHTML=`<option value="">All models</option>${state.developerModels.map(item=>`<option value="${esc(item.public_id)}">${esc(item.display_name||item.public_id)}</option>`).join('')}`;provider.innerHTML=`<option value="">All providers</option>${[...providers].sort((a,b)=>a[1].localeCompare(b[1])).map(([slug,name])=>`<option value="${esc(slug)}">${esc(name)} (${esc(slug)})</option>`).join('')}`; + if(projectValue)project.value=projectValue;if(keyValue)key.value=keyValue;if(modelValue)model.value=modelValue;if(providerValue)provider.value=providerValue; + if(!$('#usage-from').value){const query=new URLSearchParams(defaultUsageQuery());$('#usage-from').value=query.get('from');$('#usage-to').value=query.get('to');} + const points=state.usageDaily||[];const totals=points.reduce((acc,item)=>{acc.requests+=Number(item.request_count||0);acc.success+=Number(item.successful_requests||0);acc.tokens+=Number(item.total_tokens||0);acc.charged+=Number(item.charged_micros||0);acc.duration+=Number(item.average_duration_ms||0)*Number(item.request_count||0);acc.p95=Math.max(acc.p95,Number(item.p95_duration_ms||0));return acc;},{requests:0,success:0,tokens:0,charged:0,duration:0,p95:0}); + const metrics=[['Requests',integer(totals.requests),'selected range'],['Success rate',percent(totals.success,totals.requests),'completed requests'],['Tokens',integer(totals.tokens),'input and output'],['Charged',money(totals.charged),'wallet debit'],['Average latency',totals.requests?`${integer(Math.round(totals.duration/totals.requests))} ms`:'—','request weighted'],['P95 latency',totals.p95?`${integer(totals.p95)} ms`:'—','highest daily P95']];$('#usage-metrics').innerHTML=metrics.map(([label,value,sub])=>`<article class="metric"><span>${label}</span><strong>${esc(value)}</strong><small>${esc(sub)}</small></article>`).join(''); + const maxRequests=Math.max(1,...points.map(item=>Number(item.request_count||0)));$('#usage-chart').innerHTML=points.length?`<div class="chart-bars">${points.map(item=>`<div class="chart-day" title="${esc(new Date(item.day).toLocaleDateString())}: ${integer(item.request_count)} requests, ${percent(item.successful_requests,item.request_count)} success, ${money(item.charged_micros)} charged"><div class="chart-bar"><span class="${chartHeightClass(item.request_count,maxRequests)}"></span></div><small>${new Date(item.day).toLocaleDateString(undefined,{month:'short',day:'numeric'})}</small></div>`).join('')}</div>`:'<div class="empty">No usage in this range</div>'; $('#usage-summary-body').innerHTML=state.usageSummary.map(item=>`<tr><td>${new Date(item.period_start).toLocaleDateString(undefined,{year:'numeric',month:'short'})}</td><td><strong>${esc(item.project_name)}</strong></td><td>${integer(item.request_count)}</td><td>${percent(item.successful_requests,item.request_count)}</td><td>${integer(item.input_tokens)}</td><td>${integer(item.output_tokens)}</td><td>${money(item.cost_micros)}</td></tr>`).join('')||emptyRow(7); - $('#usage-events-body').innerHTML=state.usage.map(item=>`<tr><td>${date(item.started_at)}</td><td><code title="${esc(item.request_id)}">${shortID(item.request_id)}</code></td><td>${esc(item.public_model)}</td><td><span class="badge ${item.success?'active':'suspended'}">${item.status_code}</span>${item.error_type?`<small class="error-label">${esc(item.error_type)}</small>`:''}</td><td>${integer(item.total_tokens)}</td><td>${money(item.cost_micros)}</td><td>${integer(item.duration_ms)} ms</td></tr>`).join('')||emptyRow(7); + const analytics=Array.isArray(state.usageAnalytics)?{models:[],providers:[]}:state.usageAnalytics||{models:[],providers:[]}; + const changeLabel=item=>item.charge_change_percent==null?(Number(item.previous_charged_micros||0)===0&&Number(item.charged_micros||0)>0?'New':'—'):`${Number(item.charge_change_percent)>=0?'+':''}${Number(item.charge_change_percent).toFixed(1)}%`; + const changeClass=item=>item.charge_change_percent==null?(Number(item.charged_micros||0)>0?'positive':''):Number(item.charge_change_percent)>0?'money-negative':'positive'; + const cacheRate=item=>{const denominator=Number(item.input_tokens||0)+Number(item.cache_read_input_tokens||0)+Number(item.cache_creation_input_tokens||0);return denominator?percent(item.cache_read_input_tokens,denominator):'—';}; + $('#usage-model-analytics-body').innerHTML=(analytics.models||[]).map(item=>`<tr><td><strong>${esc(item.public_model)}</strong><small class="price-line">${integer(item.provider_count)} provider${Number(item.provider_count)===1?'':'s'}</small></td><td>${integer(item.request_count)}</td><td>${percent(item.successful_requests,item.request_count)}</td><td>${integer(item.total_tokens)}</td><td>${money(item.charged_micros)}</td><td class="${changeClass(item)}">${changeLabel(item)}</td><td>${integer(item.p95_duration_ms)} ms</td><td>${integer(item.missing_usage_requests)}</td></tr>`).join('')||emptyRow(8); + $('#usage-provider-analytics-body').innerHTML=(analytics.providers||[]).map(item=>`<tr><td><strong>${esc(item.provider_name)}</strong><small class="price-line">${esc(item.wire_api||'unknown')}</small></td><td>${integer(item.request_count)}</td><td>${percent(item.successful_requests,item.request_count)}</td><td>${integer(item.model_count)}</td><td>${cacheRate(item)}</td><td>${money(item.charged_micros)}</td><td class="${changeClass(item)}">${changeLabel(item)}</td><td>${integer(item.p95_duration_ms)} ms</td></tr>`).join('')||emptyRow(8); + $('#usage-events-body').innerHTML=state.usage.map((item,index)=>`<tr><td>${date(item.started_at)}</td><td><code title="${esc(item.request_id)}">${shortID(item.request_id)}</code><small class="price-line">${esc(item.protocol)}${item.attempts>1?` · ${item.attempts} attempts`:''}</small></td><td>${esc(item.project_name||shortID(item.project_id))}<small class="price-line">${esc(item.key_name||shortID(item.key_id))}</small></td><td>${esc(item.public_model)}<small class="price-line">${esc(item.provider_name||'—')}</small></td><td><span class="badge ${item.success?'active':'suspended'}">${item.status_code}</span>${item.error_type?`<small class="error-label">${esc(item.error_type)}</small>`:''}<small class="price-line">${esc(item.metering_status||'')}</small></td><td>${integer(item.total_tokens)}</td><td>${money(item.charged_micros)}${item.uncollected_micros?`<small class="error-label">${money(item.uncollected_micros)} uncollected</small>`:''}</td><td>${integer(item.duration_ms)} ms</td><td><button class="text-button" type="button" data-usage-details="${index}">Details</button></td></tr>`).join('')||emptyRow(9); +} +function usageDiagnostic(item) { + return {request_id:item.request_id,started_at:item.started_at,status_code:item.status_code,success:Boolean(item.success),error_type:item.error_type||'',project:{id:item.project_id,name:item.project_name||''},api_key:{id:item.key_id,name:item.key_name||''},model:{public_id:item.public_model,provider:item.provider_name||item.provider_id||'',upstream_id:item.upstream_model||''},transport:{protocol:item.protocol,stream:Boolean(item.stream),attempts:Number(item.attempts||0),duration_ms:Number(item.duration_ms||0)},usage:{input_tokens:Number(item.input_tokens||0),output_tokens:Number(item.output_tokens||0),total_tokens:Number(item.total_tokens||0),cache_read_input_tokens:Number(item.cache_read_input_tokens||0),cache_creation_input_tokens:Number(item.cache_creation_input_tokens||0),reported:Boolean(item.usage_reported)},billing:{cost_micros:Number(item.cost_micros||0),charged_micros:Number(item.charged_micros||0),uncollected_micros:Number(item.uncollected_micros||0),metering_status:item.metering_status||''}}; +} +function showUsageDetails(index) { + const item=state.usage[Number(index)];if(!item)return;state.detailUsage=item;$('#usage-dialog-title').textContent=item.public_model||'Request details';$('#usage-dialog-id').textContent=item.request_id; + const throughput=item.duration_ms>0&&item.output_tokens>0?`${(Number(item.output_tokens)*1000/Number(item.duration_ms)).toFixed(1)} output tokens/s`:'—';const rows=[['Time',date(item.started_at)],['Status',`${item.status_code} · ${item.success?'success':'error'}`],['Project',item.project_name||shortID(item.project_id)],['API key',item.key_name||shortID(item.key_id)],['Protocol',`${item.protocol}${item.stream?' · stream':''}`],['Public model',item.public_model],['Provider',item.provider_name||item.provider_id||'—'],['Upstream model',item.upstream_model||'—'],['Attempts',integer(item.attempts)],['Latency',`${integer(item.duration_ms)} ms`],['Throughput',throughput],['Tokens',`${integer(item.input_tokens)} in · ${integer(item.output_tokens)} out`],['Cache',`${integer(item.cache_read_input_tokens)} read · ${integer(item.cache_creation_input_tokens)} write`],['Charged',money(item.charged_micros)],['Metering',item.metering_status||'—'],['Error',item.error_type||'—']]; + $('#usage-detail-grid').innerHTML=rows.map(([label,value])=>`<div class="model-detail-row"><span>${esc(label)}</span><strong>${esc(value)}</strong></div>`).join('');$('#usage-diagnostic').textContent=JSON.stringify(usageDiagnostic(item),null,2);$('#usage-dialog').showModal(); } function renderLimits() { const existing=new Map(state.limits.map(item=>[item.project_id,item])); @@ -179,6 +420,12 @@ function addRoute() { const wrapper=document.createElement('div');wrapper.classN document.addEventListener('click',async(event)=>{ const authTab=event.target.closest('.auth-tab');if(authTab){$$('.auth-tab').forEach(node=>node.classList.toggle('active',node===authTab));$$('.auth-pane').forEach(node=>node.classList.toggle('active',node.id===authTab.dataset.authPane));authError();return;} const tab=event.target.closest('.tab');if(tab){$$('.tab').forEach(node=>node.classList.toggle('active',node===tab));$$('.section').forEach(node=>node.classList.toggle('active',node.id===tab.dataset.section));return;} + const goto=event.target.closest('[data-goto]');if(goto){goTo(goto.dataset.goto);return;} + const scroll=event.target.closest('[data-scroll-to]');if(scroll){document.getElementById(scroll.dataset.scrollTo)?.scrollIntoView({behavior:'smooth',block:'start'});return;} + const endpointCopy=event.target.closest('[data-copy-endpoint]');if(endpointCopy){try{const values=JSON.parse($('#endpoint-list').dataset.values||'[]');await navigator.clipboard.writeText(values[Number(endpointCopy.dataset.copyEndpoint)]||'');toast('Endpoint copied');}catch(error){toast('Copy failed',true);}return;} + const details=event.target.closest('[data-model-details]');if(details){showModelDetails(details.dataset.modelDetails);return;} + const useModel=event.target.closest('[data-use-model]');if(useModel){useModelInPlayground(useModel.dataset.useModel);return;} + const usageDetails=event.target.closest('[data-usage-details]');if(usageDetails){showUsageDetails(usageDetails.dataset.usageDetails);return;} if(event.target.id==='reload'){try{await api('/reload',{method:'POST',body:'{}'});await loadAll();toast('Snapshot reloaded');}catch(error){toast(error.message,true);}} if(event.target.id==='add-route')addRoute();if(event.target.closest('.remove-route'))event.target.closest('.route-row').remove(); const revokeKey=event.target.closest('[data-revoke-key]');if(revokeKey&&confirm('Revoke this API key?')){try{await api(`/keys/${revokeKey.dataset.revokeKey}/revoke`,{method:'POST',body:'{}'});await loadAll();toast('API key revoked');}catch(error){toast(error.message,true);}} @@ -211,10 +458,18 @@ $('#bootstrap-pane').addEventListener('submit',async(event)=>{event.preventDefau $('#sign-out').addEventListener('click',async()=>{try{if(!state.token)await api('/auth/logout',{method:'POST',body:'{}'});}catch(error){if(error.status!==401)toast(error.message,true);}state.token='';state.csrf='';state.actor={};state.permissions=new Set();setConnected(false);}); $('#tenant-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/tenants',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Tenant created');}catch(error){toast(error.message,true);}}); $('#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);}}); +$('#key-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const form=event.target;const data=formJSON(form);data.scopes=data.scopes.split(',').map(value=>value.trim()).filter(Boolean);data.tags=data.tags.split(',').map(value=>value.trim()).filter(Boolean);data.allowed_models=[...$('#key-models').selectedOptions].map(option=>option.value);data.monthly_spend_micros=data.monthly_spend.trim()?decimalToScaled(data.monthly_spend,6):0;delete data.monthly_spend;data.expires_at=data.expires_at?new Date(data.expires_at).toISOString():null;const result=await api('/keys',{method:'POST',body:JSON.stringify(data)});state.playgroundKey=result.key;form.reset();showSecret('API key created',result.key);await loadAll();goTo('quickstart');}catch(error){toast(error.message,true);}}); +$('#starter-key-form').addEventListener('submit',async(event)=>{event.preventDefault();const button=$('#starter-key-submit');try{const projectID=$('#starter-project').value;const name=$('#starter-key-name').value.trim();const model=selectedDeveloperModel();if(!projectID||!name||!model.public_id)throw new Error('An active project and model are required');button.disabled=true;const result=await api('/keys',{method:'POST',body:JSON.stringify({tenant_id:state.actor.tenant_id,project_id:projectID,name,scopes:['inference'],tags:['quickstart'],allowed_models:[model.public_id],monthly_spend_micros:0,expires_at:null})});state.playgroundKey=result.key;await loadAll();$('#playground-key').value=result.key;showSecret('Starter API key created',result.key);toast('Starter key is ready in the Playground');}catch(error){toast(error.message,true);}finally{button.disabled=false;}}); +function syncProviderWireAPI(){const form=$('#provider-form');const wire=form.elements.wire_api;const protocol=form.elements.protocol.value;wire.innerHTML=protocol==='anthropic'?'<option value="messages">Messages</option>':'<option value="chat_completions">Chat Completions</option><option value="responses">Responses</option>';} +$('#provider-form [name=protocol]').addEventListener('change',syncProviderWireAPI); +$('#provider-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/providers',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();syncProviderWireAPI();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;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);}}); +$('#billing-profile-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.country=String(data.country||'').toUpperCase();state.billingProfile=await api('/billing/profile',{method:'PUT',body:JSON.stringify(data)});renderBillingProfile();toast(state.billingProfile.stripe_sync_status==='synced'?'Invoice details saved and synced':'Invoice details saved');}catch(error){try{state.billingProfile=await api('/billing/profile');renderBillingProfile();}catch{}toast(error.message,true);}}); +$('#auto-topup-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const currency=state.autoTopUp.currency||state.overview.billing_currency||'usd';const result=await api('/billing/auto-topup',{method:'PUT',body:JSON.stringify({enabled:$('#auto-topup-enabled').checked,threshold_micros:decimalToScaled($('#auto-topup-threshold').value,6),topup_amount_minor:decimalToScaled($('#auto-topup-amount').value,currencyDigits(currency))})});state.autoTopUp=result;renderAutoTopUp();renderQuickstart();toast(result.enabled?'Automatic top-up enabled':'Automatic top-up settings saved');}catch(error){toast(error.message,true);}}); +$('#auto-topup-payment-setup').addEventListener('click',async()=>{try{const result=await api('/billing/auto-topup/setup-sessions',{method:'POST',body:'{}'});window.location.assign(result.url);}catch(error){toast(error.message,true);}}); +$('#developer-preferences-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/developer/preferences',{method:'PUT',body:JSON.stringify({default_model:data.default_model||'',fallback_model:data.fallback_model||''})});await loadAll();toast('API defaults saved');}catch(error){toast(error.message,true);}}); +$('#billing-preferences-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/developer/preferences/billing',{method:'PUT',body:JSON.stringify({low_balance_enabled:$('#low-balance-enabled').checked,low_balance_threshold_micros:decimalToScaled($('#low-balance-threshold').value,6)})});await loadAll();toast('Balance alert saved');}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);}}); $('#password-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/auth/password',{method:'POST',body:JSON.stringify({current_password:data.current_password,new_password:data.new_password})});event.target.reset();state.csrf='';state.actor={};state.permissions=new Set();setConnected(false);authError('Password changed. Sign in again.',true);}catch(error){toast(error.message,true);}}); @@ -223,6 +478,26 @@ $('#totp-confirm-form').addEventListener('submit',async(event)=>{event.preventDe $('#passkey-form').addEventListener('submit',async(event)=>{event.preventDefault();try{if(!state.authConfig.passkeys_enabled)throw new Error('Passkeys are not configured');const data=formJSON(event.target);const begin=await api('/auth/mfa/passkey/options',{method:'POST',body:JSON.stringify({current_password:data.current_password})});const credential=await createPasskey(begin.options);await api('/auth/mfa/passkey',{method:'POST',body:JSON.stringify({challenge_token:begin.challenge_token,name:data.name,credential})});event.target.reset();await loadAll();toast('Passkey added');}catch(error){toast(error.message,true);}}); $('#revoke-other-sessions').addEventListener('click',async()=>{if(!confirm('Sign out every other device?'))return;try{await api('/auth/sessions/revoke-others',{method:'POST',body:'{}'});await loadAll();toast('Other devices signed out');}catch(error){toast(error.message,true);}}); $('#close-dialog').addEventListener('click',()=>$('#secret-dialog').close());$('#copy-secret').addEventListener('click',async()=>{await navigator.clipboard.writeText($('#created-secret').textContent);toast('Credential copied');}); +$('#close-model-dialog').addEventListener('click',()=>$('#model-dialog').close());$('#model-dialog-close').addEventListener('click',()=>$('#model-dialog').close());$('#model-dialog-playground').addEventListener('click',()=>{if(state.detailModel)useModelInPlayground(state.detailModel.public_id);});$('#copy-model-id').addEventListener('click',async()=>{await navigator.clipboard.writeText($('#model-dialog-id').textContent);toast('Model ID copied');}); +$('#close-usage-dialog').addEventListener('click',()=>$('#usage-dialog').close());$('#usage-dialog-close').addEventListener('click',()=>$('#usage-dialog').close());$('#copy-request-id').addEventListener('click',async()=>{await navigator.clipboard.writeText(state.detailUsage?.request_id||'');toast('Request ID copied');});$('#copy-request-diagnostic').addEventListener('click',async()=>{await navigator.clipboard.writeText($('#usage-diagnostic').textContent);toast('Diagnostic copied');}); +$('#model-estimate-form').addEventListener('submit',event=>event.preventDefault()); +['estimate-input','estimate-output','estimate-cache-read','estimate-cache-write'].forEach(id=>$('#'+id).addEventListener('input',renderModelEstimate)); +$('#quickstart-model').addEventListener('change',()=>{syncQuickstartProtocols();renderQuickstartCode();renderDeveloperAccess();}); +$('#quickstart-protocol').addEventListener('change',()=>{syncProviderSelect('#quickstart-provider',selectedDeveloperModel(),$('#quickstart-protocol').value);renderQuickstartCode();}); +$('#quickstart-provider').addEventListener('change',renderQuickstartCode); +$('#quickstart-language').addEventListener('change',renderQuickstartCode); +$('#playground-model').addEventListener('change',syncPlaygroundProtocols); +$('#playground-protocol').addEventListener('change',()=>syncProviderSelect('#playground-provider',selectedPlaygroundModel(),$('#playground-protocol').value)); +$('#playground-form').addEventListener('submit',runPlayground); +$('#playground-stop').addEventListener('click',()=>state.playgroundController?.abort()); +$('#playground-diagnostic-action').addEventListener('click',event=>goTo(event.currentTarget.dataset.target)); +$('#catalog-search').addEventListener('input',renderCatalog);$('#catalog-protocol').addEventListener('change',renderCatalog);$('#catalog-input').addEventListener('change',renderCatalog);$('#catalog-owner').addEventListener('change',renderCatalog);$('#catalog-sort').addEventListener('change',renderCatalog); +$('#copy-quickstart').addEventListener('click',async()=>{try{await navigator.clipboard.writeText($('#quickstart-code-block code').textContent);toast('Example copied');}catch(error){toast('Copy failed; select the example manually',true);}}); +$('#copy-endpoint-env').addEventListener('click',async()=>{try{await navigator.clipboard.writeText($('#endpoint-env').textContent);toast('Environment copied');}catch(error){toast('Copy failed; select the environment manually',true);}}); +async function loadUsageFilters(){const params=new URLSearchParams(new FormData($('#usage-filter-form')));for(const [key,value] of [...params.entries()])if(!String(value).trim())params.delete(key);const query=params.toString();try{[state.usage,state.usageDaily,state.usageAnalytics]=await Promise.all([api(`/usage${query?`?${query}`:''}`),api(`/usage/daily${query?`?${query}`:''}`),api(`/usage/analytics${query?`?${query}`:''}`)]);renderUsage();}catch(error){toast(error.message,true);}} +$('#usage-filter-form').addEventListener('submit',async(event)=>{event.preventDefault();await loadUsageFilters();}); +$('#usage-filter-reset').addEventListener('click',async()=>{const form=$('#usage-filter-form');form.reset();const query=new URLSearchParams(defaultUsageQuery());$('#usage-from').value=query.get('from');$('#usage-to').value=query.get('to');await loadUsageFilters();}); async function pollTopUp(orderID){for(let attempt=0;attempt<20;attempt++){const order=await api(`/billing/orders/${encodeURIComponent(orderID)}`);if(order.status==='paid'){await loadAll();toast('Balance credited');return;}if(['failed','expired'].includes(order.status)){toast(`Top-up ${order.status}`,true);return;}await new Promise(resolve=>setTimeout(resolve,1500));}toast('Payment is still processing');} -async function start(){try{state.authConfig=await api('/auth/config');$('#register-tab').classList.toggle('hidden',!state.authConfig.registration_enabled);if(!state.authConfig.registration_enabled&&$('#register-tab').classList.contains('active'))showAuthPane('login-pane');const params=new URLSearchParams(location.search);const action=params.get('action');const token=params.get('token');if(action==='reset-password'&&token){$('#reset-token').value=token;showAuthPane('reset-complete-pane');setConnected(false);return;}if(action==='accept-invite'&&token){$('#invite-token').value=token;showAuthPane('invite-pane');setConnected(false);return;}state.csrf=cookie('aigw_csrf');if(action==='verify-email'&&token){const result=await api('/auth/email/verify',{method:'POST',body:JSON.stringify({token})});history.replaceState({},'',location.pathname);await completeBrowserLogin(result);return;}const session=await api('/auth/session');if(session.authenticated){await loadAll(session);const orderID=params.get('order_id');if(params.get('topup')==='success'&&orderID){history.replaceState({},'',location.pathname);pollTopUp(orderID).catch(error=>toast(error.message,true));}else if(params.get('topup')==='cancel'){history.replaceState({},'',location.pathname);toast('Top-up cancelled');}}else setConnected(false);}catch(error){setConnected(false);authError(error.message);}} +async function pollAutoTopUpSetup(){for(let attempt=0;attempt<20;attempt++){const settings=await api('/billing/auto-topup');if(settings.payment_method_configured){state.autoTopUp=settings;renderBilling();renderQuickstart();toast('Payment method saved; review the threshold and enable automatic top-up');return;}await new Promise(resolve=>setTimeout(resolve,1500));}toast('Payment method setup is still processing',true);} +async function start(){try{state.authConfig=await api('/auth/config');$('#register-tab').classList.toggle('hidden',!state.authConfig.registration_enabled);if(!state.authConfig.registration_enabled&&$('#register-tab').classList.contains('active'))showAuthPane('login-pane');const params=new URLSearchParams(location.search);const action=params.get('action');const token=params.get('token');if(params.get('auth')==='register'&&state.authConfig.registration_enabled)showAuthPane('register-pane');if(action==='reset-password'&&token){$('#reset-token').value=token;showAuthPane('reset-complete-pane');setConnected(false);return;}if(action==='accept-invite'&&token){$('#invite-token').value=token;showAuthPane('invite-pane');setConnected(false);return;}state.csrf=cookie('aigw_csrf');if(action==='verify-email'&&token){const result=await api('/auth/email/verify',{method:'POST',body:JSON.stringify({token})});history.replaceState({},'',location.pathname);await completeBrowserLogin(result);return;}const session=await api('/auth/session');if(session.authenticated){await loadAll(session);const orderID=params.get('order_id');if(params.get('topup')==='success'&&orderID){history.replaceState({},'',location.pathname);pollTopUp(orderID).catch(error=>toast(error.message,true));}else if(params.get('topup')==='cancel'){history.replaceState({},'',location.pathname);toast('Top-up cancelled');}else if(params.get('autotopup')==='setup'){history.replaceState({},'',location.pathname);pollAutoTopUpSetup().catch(error=>toast(error.message,true));}else if(params.get('autotopup')==='cancel'){history.replaceState({},'',location.pathname);toast('Payment method setup cancelled');}}else setConnected(false);}catch(error){setConnected(false);authError(error.message);}} start(); diff --git a/internal/adminui/assets/index.html b/internal/adminui/assets/index.html index 04b4e87..8f31e1c 100644 --- a/internal/adminui/assets/index.html +++ b/internal/adminui/assets/index.html @@ -75,7 +75,9 @@ </header> <main class="shell"> <nav class="tabs" aria-label="Admin sections"> - <button class="tab active" data-section="overview">Overview</button> + <button class="tab active" data-section="quickstart">Quickstart</button> + <button class="tab" data-section="overview">Overview</button> + <button class="tab" data-section="catalog">Model catalog</button> <button class="tab" data-section="usage" data-permission="usage.read">Usage</button> <button class="tab" data-section="billing" data-permission="billing.read" id="billing-tab">Billing</button> <button class="tab" data-section="projects" data-permission="projects.read">Projects</button> @@ -86,10 +88,77 @@ <button class="tab" data-section="audit" data-permission="audit.read">Audit</button> <button class="tab" data-section="tenants" data-permission="tenants.read">Tenants</button> <button class="tab" data-section="providers" data-permission="platform.read">Providers</button> - <button class="tab" data-section="models" data-permission="platform.read">Models & routes</button> + <button class="tab" data-section="models" data-permission="platform.read">Routing</button> </nav> - <section id="overview" class="section active"> + <section id="quickstart" class="section active"> + <div class="section-heading"><div><span class="eyebrow">DEVELOPER WORKSPACE</span><h1>Start building</h1></div><div class="form-actions"><span class="currency-label" id="quickstart-balance"></span><button class="button secondary" type="button" id="quickstart-topup" data-goto="billing" data-permission="billing.topup">Add funds</button></div></div> + <div class="onboarding panel" id="onboarding"></div> + <div class="quick-access-grid"> + <form class="panel starter-key-panel" id="starter-key-form" data-permission="keys.write"> + <div class="section-heading"><div><span class="eyebrow">API ACCESS</span><h2>Create a starter key</h2></div><span class="badge" id="starter-model-label"></span></div> + <div class="starter-key-fields"><label>Project<select id="starter-project" required></select></label><label>Key name<input id="starter-key-name" maxlength="120" value="Quickstart key" required></label><button class="button primary" id="starter-key-submit" type="submit">Create & use key</button></div> + <p class="muted">The key is limited to the selected Quickstart model and placed in the Playground for this page only.</p> + </form> + <section class="panel endpoint-panel"> + <div class="section-heading"><div><span class="eyebrow">CONNECTION</span><h2>API endpoints</h2></div><button class="button subtle" type="button" id="copy-endpoint-env">Copy environment</button></div> + <div class="endpoint-list" id="endpoint-list"></div> + <pre class="endpoint-env"><code id="endpoint-env"></code></pre> + </section> + </div> + <div class="quickstart-grid"> + <div class="panel quickstart-code"> + <div class="section-heading"><div><span class="eyebrow">FIRST REQUEST</span><h2>Copy a working example</h2></div><button class="button subtle" type="button" id="copy-quickstart">Copy</button></div> + <div class="form-grid compact-form"><label>Model<select id="quickstart-model"></select></label><label>Protocol<select id="quickstart-protocol"></select></label><label>Provider<select id="quickstart-provider"></select></label><label>Language<select id="quickstart-language"><option value="curl">cURL</option><option value="python">Python</option><option value="node">Node.js</option></select></label></div> + <pre id="quickstart-code-block"><code></code></pre> + <p class="muted">Use an active API key in <code>AIGW_API_KEY</code>. The example never stores your key in the browser.</p> + </div> + <div class="panel quickstart-next"><span class="eyebrow">WORKSPACE SIGNALS</span><h2>What to do next</h2><div id="quickstart-signals" class="signal-list"></div></div> + </div> + <form class="panel playground" id="playground-form"> + <div class="section-heading"><div><span class="eyebrow">LIVE API</span><h2>API playground</h2></div><span class="currency-label" id="playground-endpoint"></span></div> + <div class="playground-fields"> + <label class="playground-key">API key<input id="playground-key" type="password" autocomplete="off" spellcheck="false" required placeholder="sk-aigw-..."></label> + <label>Model<select id="playground-model" required></select></label> + <label>Protocol<select id="playground-protocol" required></select></label> + <label>Provider<select id="playground-provider"></select></label> + <label class="playground-prompt">Prompt<textarea id="playground-prompt" required rows="4">Reply with exactly: AIGW ready</textarea></label> + <div class="playground-actions"><label>Max output<input id="playground-max-output" type="number" min="1" max="4096" value="128" required></label><button class="button primary" id="playground-send" type="submit">Send request</button><button class="button secondary hidden" id="playground-stop" type="button">Stop</button></div> + </div> + <div class="playground-result hidden" id="playground-result" aria-live="polite"> + <div class="playground-meta"><strong class="badge" id="playground-status"></strong><span id="playground-request-id"></span><span id="playground-duration"></span><span id="playground-tokens"></span></div> + <pre><code id="playground-response"></code></pre> + <div class="playground-diagnostic hidden" id="playground-diagnostic"><div><strong id="playground-diagnostic-title"></strong><p id="playground-diagnostic-message"></p></div><button class="button secondary hidden" id="playground-diagnostic-action" type="button"></button></div> + </div> + </form> + <div class="section-heading ledger-heading"><div><span class="eyebrow">DEFAULTS & ALERTS</span><h2>Workspace preferences</h2></div></div> + <div class="preferences-grid"> + <form class="panel form-grid compact-form" id="developer-preferences-form" data-permission="preferences.read"> + <h2>API defaults</h2> + <label>Default model<select id="preference-default-model" name="default_model"></select></label> + <label>Fallback model<select id="preference-fallback-model" name="fallback_model"></select></label> + <p class="muted form-note">These defaults drive workspace examples. Your API request can still select any model available to this tenant.</p> + <button class="button primary" type="submit" data-permission="developer.preferences.write">Save API defaults</button> + </form> + <form class="panel form-grid compact-form" id="billing-preferences-form" data-permission="billing.read"> + <h2>Balance alert</h2> + <label class="toggle-row"><input id="low-balance-enabled" name="low_balance_enabled" type="checkbox"><span>Email billing members when available balance is low</span></label> + <label>Alert below<input id="low-balance-threshold" name="low_balance_threshold" inputmode="decimal" required></label> + <p class="muted form-note" id="balance-alert-status"></p> + <button class="button primary" type="submit" data-permission="billing.preferences.write">Save balance alert</button> + </form> + </div> + <div class="section-heading ledger-heading"><div><span class="eyebrow">RECENT ACTIVITY</span><h2>Latest requests</h2></div><button class="button subtle" type="button" data-goto="usage">View all usage</button></div> + <div class="panel table-wrap"><table><thead><tr><th>Time</th><th>Model</th><th>Status</th><th>Tokens</th><th>Charged</th><th>Latency</th></tr></thead><tbody id="quickstart-usage-body"></tbody></table></div> + </section> + + <section id="catalog" class="section"> + <div class="section-heading"><div><span class="eyebrow">DISCOVER</span><h1>Model catalog</h1></div><span class="currency-label" id="catalog-count"></span></div> + <div class="panel catalog-toolbar"><label>Search models<input id="catalog-search" type="search" placeholder="Search by model, developer, capability"></label><label>Protocol<select id="catalog-protocol"><option value="">All protocols</option><option value="chat_completions">OpenAI Chat Completions</option><option value="responses">OpenAI Responses</option><option value="messages">Anthropic Messages</option></select></label><label>Input<select id="catalog-input"><option value="">Any input</option><option value="text">Text</option><option value="image">Image</option><option value="audio">Audio</option><option value="video">Video</option></select></label><label>Developer<select id="catalog-owner"><option value="">All developers</option></select></label><label>Sort<select id="catalog-sort"><option value="newest">Newest</option><option value="name">Name</option><option value="input_price">Lowest input price</option><option value="output_price">Lowest output price</option><option value="context">Largest context</option></select></label></div> + <div id="catalog-grid" class="catalog-grid"></div> + </section> + + <section id="overview" class="section"> <div class="section-heading"><div><span class="eyebrow">OPERATIONS</span><h1>Account overview</h1></div><button class="button secondary" id="reload" data-permission="platform.write">Reload snapshot</button></div> <div class="metric-grid" id="metrics"></div> <div class="section-heading ledger-heading"><div><span class="eyebrow">CURRENT PERIOD</span><h2>Project usage</h2></div></div> @@ -98,9 +167,29 @@ <section id="usage" class="section"> <div class="section-heading"><div><span class="eyebrow">METERING</span><h1>Usage ledger</h1></div></div> + <form class="panel usage-filters" id="usage-filter-form"> + <label>From<input id="usage-from" name="from" type="date"></label> + <label>To<input id="usage-to" name="to" type="date"></label> + <label>Project<select id="usage-project" name="project_id"><option value="">All projects</option></select></label> + <label>API key<select id="usage-key" name="key_id"><option value="">All API keys</option></select></label> + <label>Model<select id="usage-model" name="model"><option value="">All models</option></select></label> + <label>Provider<select id="usage-provider" name="provider"><option value="">All providers</option></select></label> + <label>Protocol<select id="usage-protocol" name="protocol"><option value="">All protocols</option><option value="openai">OpenAI Chat Completions</option><option value="openai_responses">OpenAI Responses</option><option value="anthropic">Anthropic Messages</option></select></label> + <label>Transport<select id="usage-stream" name="stream"><option value="">Streaming and non-streaming</option><option value="true">Streaming</option><option value="false">Non-streaming</option></select></label> + <label>Status<select id="usage-status" name="status"><option value="">All statuses</option><option value="success">Successful</option><option value="error">Errors</option></select></label> + <label>Error type<input id="usage-error-type" name="error_type" placeholder="provider_error"></label> + <label class="usage-request-filter">Request ID<input id="usage-request-id" name="request_id" placeholder="req_..."></label> + <div class="form-actions"><button class="button primary" type="submit">Apply filters</button><button class="button subtle" id="usage-filter-reset" type="button">Reset</button></div> + </form> + <div class="metric-grid usage-metrics" id="usage-metrics"></div> + <div class="panel usage-chart" id="usage-chart"></div> + <div class="analytics-grid"> + <div class="panel table-wrap analytics-panel"><div class="section-heading ledger-heading"><div><span class="eyebrow">COST RANKING</span><h2>Model cost</h2></div><span class="currency-label">Current range vs previous</span></div><table class="analytics-table"><thead><tr><th>Model</th><th>Requests</th><th>Success</th><th>Tokens</th><th>Charged</th><th>Change</th><th>P95</th><th>Missing usage</th></tr></thead><tbody id="usage-model-analytics-body"></tbody></table></div> + <div class="panel table-wrap analytics-panel"><div class="section-heading ledger-heading"><div><span class="eyebrow">ROUTE HEALTH</span><h2>Provider performance</h2></div><span class="currency-label">Cache hit and latency</span></div><table class="analytics-table"><thead><tr><th>Provider</th><th>Requests</th><th>Success</th><th>Models</th><th>Cache hit</th><th>Charged</th><th>Change</th><th>P95</th></tr></thead><tbody id="usage-provider-analytics-body"></tbody></table></div> + </div> <div class="panel table-wrap"><table><thead><tr><th>Period</th><th>Project</th><th>Requests</th><th>Success</th><th>Input</th><th>Output</th><th>Cost</th></tr></thead><tbody id="usage-summary-body"></tbody></table></div> <div class="section-heading ledger-heading"><div><span class="eyebrow">REQUESTS</span><h2>Recent events</h2></div></div> - <div class="panel table-wrap"><table><thead><tr><th>Time</th><th>Request</th><th>Model</th><th>Status</th><th>Tokens</th><th>Cost</th><th>Latency</th></tr></thead><tbody id="usage-events-body"></tbody></table></div> + <div class="panel table-wrap"><table><thead><tr><th>Time</th><th>Request</th><th>Project / key</th><th>Model / provider</th><th>Status</th><th>Tokens</th><th>Charged</th><th>Latency</th><th></th></tr></thead><tbody id="usage-events-body"></tbody></table></div> </section> <section id="tenants" class="section"> @@ -117,14 +206,24 @@ <section id="keys" class="section"> <div class="section-heading"><div><span class="eyebrow">ACCESS</span><h1>API keys</h1></div></div> - <form class="panel form-grid" id="key-form" data-permission="keys.write"><label>Tenant<select name="tenant_id" id="key-tenant" required></select></label><label>Project<select name="project_id" id="key-project" required></select></label><label>Name<input name="name" required placeholder="CLI production key"></label><label>Scopes<input name="scopes" value="inference" placeholder="inference"></label><button class="button primary" type="submit">Create key</button></form> + <form class="panel form-grid key-form" id="key-form" data-permission="keys.write"> + <label>Tenant<select name="tenant_id" id="key-tenant" required></select></label> + <label>Project<select name="project_id" id="key-project" required></select></label> + <label>Name<input name="name" required maxlength="120" placeholder="CLI production key"></label> + <label>Scopes<input name="scopes" value="inference" placeholder="inference"></label> + <label>Tags<input name="tags" placeholder="production, backend"></label> + <label>Monthly spend cap<input name="monthly_spend" inputmode="decimal" value="0" placeholder="0 = unlimited"></label> + <label>Expires at<input name="expires_at" type="datetime-local"></label> + <label class="key-model-picker">Allowed models<select name="allowed_models" id="key-models" multiple size="5" aria-describedby="key-model-help"></select><small id="key-model-help">No selection allows every model visible to this workspace.</small></label> + <button class="button primary" type="submit">Create key</button> + </form> <div class="panel warning"><strong>Key visibility</strong><span>The secret is shown only once after creation.</span></div> - <div class="panel table-wrap"><table><thead><tr><th>Name</th><th>Prefix</th><th>Project</th><th>Scopes</th><th>Status</th><th></th></tr></thead><tbody id="keys-body"></tbody></table></div> + <div class="panel table-wrap"><table class="keys-table"><thead><tr><th>Name</th><th>Prefix</th><th>Project</th><th>Restrictions</th><th>Month usage</th><th>Monthly cap</th><th>Activity</th><th></th></tr></thead><tbody id="keys-body"></tbody></table></div> </section> <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" 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> + <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>Public slug<input name="slug" autocomplete="off" required pattern="[a-z0-9][a-z0-9-]{1,62}[a-z0-9]" maxlength="64" placeholder="openai-primary"></label><label>Protocol<select name="protocol"><option value="openai">OpenAI</option><option value="anthropic">Anthropic</option></select></label><label>Wire API<select name="wire_api"><option value="chat_completions">Chat Completions</option><option value="responses">Responses</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> @@ -140,6 +239,30 @@ <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> </div> + <div class="panel billing-profile-panel" id="billing-profile-panel" data-permission="billing.read"> + <div class="section-heading"><div><span class="eyebrow">INVOICE DETAILS</span><h2>Billing profile</h2></div><span class="badge" id="billing-profile-status">Not configured</span></div> + <form class="form-grid" id="billing-profile-form"> + <label>Legal or billing name<input name="legal_name" maxlength="150" autocomplete="organization" required></label> + <label>Billing email<input name="billing_email" type="email" maxlength="254" autocomplete="email" required></label> + <label class="billing-address-wide">Address line 1<input name="address_line1" maxlength="200" autocomplete="address-line1" required></label> + <label class="billing-address-wide">Address line 2<input name="address_line2" maxlength="200" autocomplete="address-line2"></label> + <label>City<input name="city" maxlength="100" autocomplete="address-level2" required></label> + <label>State or region<input name="region" maxlength="100" autocomplete="address-level1"></label> + <label>Postal code<input name="postal_code" maxlength="32" autocomplete="postal-code" required></label> + <label>Country code<input name="country" minlength="2" maxlength="2" pattern="[A-Za-z]{2}" autocomplete="country" autocapitalize="characters" placeholder="NZ" required></label> + <div class="form-actions billing-profile-actions"><small class="error-label" id="billing-profile-error"></small><button class="button primary" type="submit" data-permission="billing.topup">Save invoice details</button></div> + </form> + </div> + <div class="panel auto-topup-panel" id="auto-topup-panel" data-permission="billing.read"> + <div class="section-heading"><div><span class="eyebrow">BALANCE PROTECTION</span><h2>Automatic top-up</h2></div><span class="badge" id="auto-topup-status">Not configured</span></div> + <form class="form-grid" id="auto-topup-form"> + <label class="toggle-row"><input id="auto-topup-enabled" name="enabled" type="checkbox"><span>Automatically add funds when available balance reaches the threshold</span></label> + <label>Balance threshold<input id="auto-topup-threshold" name="threshold" inputmode="decimal" required></label> + <label>Top-up amount<input id="auto-topup-amount" name="amount" inputmode="decimal" required></label> + <div><span class="muted">Payment method</span><strong id="auto-topup-payment-method">Not saved</strong><small class="error-label" id="auto-topup-error"></small></div> + <div class="form-actions auto-topup-actions"><button class="button secondary" id="auto-topup-payment-setup" type="button">Save payment method</button><button class="button primary" type="submit">Save automatic top-up</button></div> + </form> + </div> <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> @@ -180,6 +303,8 @@ </main> </div> <div id="toast" class="toast" role="status"></div> + <dialog id="model-dialog"><div class="dialog-content model-dialog-content"><div class="section-heading"><div><span class="eyebrow">MODEL DETAIL</span><h2 id="model-dialog-title">Model</h2></div><button class="icon-button" id="close-model-dialog" aria-label="Close">×</button></div><div class="model-dialog-id"><code id="model-dialog-id"></code><button class="text-button" id="copy-model-id" type="button">Copy ID</button></div><div class="model-detail-grid" id="model-detail-grid"></div><section class="model-provider-health" id="model-provider-health"><div class="section-heading"><div><h3>Provider health</h3><small>Measured from requests handled by this gateway instance.</small></div></div><div class="table-wrap"><table class="provider-health-table"><thead><tr><th>Provider</th><th>State</th><th>Recent availability</th><th>Header latency</th><th>Samples</th></tr></thead><tbody id="model-provider-health-body"></tbody></table></div></section><form class="model-estimator" id="model-estimate-form"><h3>Cost estimate</h3><div class="form-grid compact-form"><label>Input tokens<input id="estimate-input" type="number" min="0" step="1" value="1000" required></label><label>Output tokens<input id="estimate-output" type="number" min="0" step="1" value="500" required></label><label>Cache read tokens<input id="estimate-cache-read" type="number" min="0" step="1" value="0" required></label><label>Cache write tokens<input id="estimate-cache-write" type="number" min="0" step="1" value="0" required></label><output id="model-estimate" class="model-estimate" aria-live="polite"></output></div></form><div class="form-actions model-dialog-actions"><button class="button secondary" id="model-dialog-playground" type="button">Use in Playground</button><button class="button subtle" id="model-dialog-close" type="button">Close</button></div></div></dialog> + <dialog id="usage-dialog"><div class="dialog-content request-dialog-content"><div class="section-heading"><div><span class="eyebrow">API REQUEST</span><h2 id="usage-dialog-title">Request details</h2></div><button class="icon-button" id="close-usage-dialog" aria-label="Close">×</button></div><div class="model-dialog-id"><code id="usage-dialog-id"></code><button class="text-button" id="copy-request-id" type="button">Copy request ID</button></div><div class="request-detail-grid" id="usage-detail-grid"></div><pre class="request-diagnostic"><code id="usage-diagnostic"></code></pre><div class="form-actions model-dialog-actions"><button class="button secondary" id="copy-request-diagnostic" type="button">Copy diagnostic</button><button class="button subtle" id="usage-dialog-close" type="button">Close</button></div></div></dialog> <dialog id="secret-dialog"><div class="dialog-content"><div class="section-heading"><div><span class="eyebrow">ONE-TIME SECRET</span><h2 id="secret-title">Credential created</h2></div><button class="icon-button" id="close-dialog" aria-label="Close">×</button></div><p>Copy this credential now. It will not be shown again.</p><code id="created-secret"></code><button class="button primary" id="copy-secret">Copy credential</button></div></dialog> <script src="./app.js" defer></script> </body> diff --git a/internal/adminui/assets/models.css b/internal/adminui/assets/models.css new file mode 100644 index 0000000..dc7c9c9 --- /dev/null +++ b/internal/adminui/assets/models.css @@ -0,0 +1,64 @@ +:root { --bg:#f4f6f7; --panel:#fff; --ink:#17212b; --muted:#657583; --line:#d9e1e5; --nav:#102a3a; --accent:#146c94; --accent-soft:#e6f2f6; --good:#187151; --warn:#7b5e14; --bad:#a23f45; font-family:Inter,ui-sans-serif,system-ui,-apple-system,BlinkMacSystemFont,"Segoe UI",sans-serif; } +* { box-sizing:border-box; } +body { margin:0; background:var(--bg); color:var(--ink); font-size:14px; } +button,input,select { font:inherit; } +button { cursor:pointer; } +.catalog-header { min-height:66px; padding:12px max(24px,calc((100% - 1240px)/2)); display:flex; align-items:center; justify-content:space-between; gap:20px; background:var(--nav); } +.catalog-brand { display:grid; grid-template-columns:38px auto; grid-template-rows:20px 14px; column-gap:10px; color:#fff; text-decoration:none; } +.catalog-brand > span { grid-row:1/-1; width:38px; height:38px; border:1px solid #8fd0df; display:grid; place-items:center; color:#b8eef7; font-weight:800; } +.catalog-brand strong { font-size:15px; align-self:end; } +.catalog-brand small { color:#8ba9b9; font-size:9px; } +.catalog-header nav { display:flex; gap:8px; } +.button { min-height:40px; padding:0 15px; border:1px solid transparent; display:inline-flex; align-items:center; justify-content:center; font-weight:700; text-decoration:none; } +.button.primary { color:#fff; background:var(--accent); } +.button.secondary { color:var(--accent); background:#fff; border-color:var(--line); } +.catalog-header .button.secondary { color:#b8eef7; background:transparent; border-color:#537486; } +main { width:min(1240px,calc(100% - 48px)); margin:30px auto 64px; } +.catalog-intro { display:flex; justify-content:space-between; align-items:end; gap:24px; margin-bottom:22px; } +.eyebrow { color:var(--accent); font-size:10px; font-weight:800; } +h1 { margin:7px 0 0; font-size:32px; line-height:1.1; } +.catalog-stats { display:flex; gap:22px; color:var(--muted); font-size:12px; } +.catalog-stats strong { display:block; color:var(--ink); font-size:20px; } +.catalog-filters { display:grid; grid-template-columns:2fr repeat(4,minmax(0,1fr)); gap:12px; padding:18px; margin-bottom:18px; background:var(--panel); border:1px solid var(--line); } +label { display:flex; flex-direction:column; gap:7px; color:var(--muted); font-size:12px; font-weight:650; } +input,select { width:100%; min-height:40px; padding:9px 11px; border:1px solid var(--line); color:var(--ink); background:#fff; outline:none; } +input:focus,select:focus { border-color:#69a9bf; box-shadow:0 0 0 3px var(--accent-soft); } +.catalog-grid { display:grid; grid-template-columns:repeat(3,minmax(0,1fr)); gap:14px; } +.model-card { min-width:0; display:flex; flex-direction:column; gap:12px; padding:18px; background:var(--panel); border:1px solid var(--line); box-shadow:0 8px 24px rgba(29,47,61,.05); } +.model-card header { display:flex; justify-content:space-between; align-items:start; gap:12px; } +.model-card header small { color:var(--muted); } +.model-card h2 { margin:4px 0 0; font-size:17px; overflow-wrap:anywhere; } +.model-card > code { color:#486071; font-size:12px; overflow-wrap:anywhere; } +.model-card > p { min-height:63px; margin:0; color:#526673; line-height:1.5; display:-webkit-box; -webkit-line-clamp:3; -webkit-box-orient:vertical; overflow:hidden; } +.health { flex:0 0 auto; padding:4px 7px; border:1px solid var(--line); font-size:11px; white-space:nowrap; } +.health.available { color:var(--good); background:#eaf7f0; border-color:#c7e9d9; } +.health.degraded { color:var(--warn); background:#fff8dc; border-color:#e9d990; } +.health.unavailable { color:var(--bad); background:#fff0f0; border-color:#f0cccc; } +.model-tags,.detail-tags { display:flex; flex-wrap:wrap; gap:5px; } +.model-tags span,.detail-tags span { padding:4px 7px; color:#4c6572; background:#eef3f5; font-size:11px; } +.model-card dl { display:grid; grid-template-columns:1fr 1fr 1fr; gap:8px; margin:0; padding-top:12px; border-top:1px solid var(--line); } +.model-card dl div { min-width:0; } +dt { color:var(--muted); font-size:10px; } +dd { margin:5px 0 0; font-size:12px; font-weight:700; overflow-wrap:anywhere; } +.model-card > .button { margin-top:auto; align-self:flex-start; } +.catalog-status { padding:28px; color:var(--muted); text-align:center; background:#fff; border:1px solid var(--line); } +.catalog-status[role="alert"] { color:var(--bad); background:#fff6f6; } +.hidden { display:none !important; } +dialog { width:min(720px,calc(100% - 32px)); max-height:calc(100vh - 32px); padding:0; border:0; box-shadow:0 18px 70px rgba(0,0,0,.25); } +dialog::backdrop { background:rgba(16,42,58,.5); } +.model-dialog { padding:24px; overflow:auto; } +.model-dialog > header { display:flex; justify-content:space-between; align-items:start; gap:16px; padding-bottom:16px; border-bottom:1px solid var(--line); } +.model-dialog h2 { margin:5px 0 5px; font-size:24px; } +.model-dialog code { color:#486071; overflow-wrap:anywhere; } +.icon-button { width:36px; height:36px; border:1px solid var(--line); background:#fff; color:var(--muted); font-size:23px; line-height:1; } +.model-dialog > p { margin:18px 0; color:#526673; line-height:1.55; } +.detail-grid { display:grid; grid-template-columns:1fr 1fr; gap:0 22px; margin:18px 0; } +.detail-grid > div { display:flex; justify-content:space-between; gap:14px; padding:11px 0; border-bottom:1px solid var(--line); } +.detail-grid dd { text-align:right; } +.price-estimator { padding-top:18px; border-top:1px solid var(--line); } +.price-estimator h3 { margin:5px 0 14px; font-size:16px; } +.estimate-inputs { display:grid; grid-template-columns:1fr 1fr 1fr; align-items:end; gap:12px; } +.estimate-inputs output { min-height:40px; display:flex; align-items:center; justify-content:center; color:var(--good); background:#f0fbf5; border:1px solid #c7e9d9; font-weight:750; } +.model-dialog footer { display:flex; justify-content:flex-end; gap:8px; margin-top:22px; } +@media (max-width:900px) { .catalog-filters { grid-template-columns:repeat(2,minmax(0,1fr)); } .search-field { grid-column:1/-1; } .catalog-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } } +@media (max-width:620px) { .catalog-header { padding:12px; } .catalog-brand small { display:none; } .catalog-header .button { padding:0 11px; } main { width:calc(100% - 24px); margin-top:20px; } .catalog-intro { align-items:start; flex-direction:column; } .catalog-stats { width:100%; justify-content:space-between; } .catalog-filters,.catalog-grid,.detail-grid,.estimate-inputs { grid-template-columns:1fr; } .search-field { grid-column:auto; } .model-card > p { min-height:0; } .model-dialog { padding:20px; } .model-dialog footer { flex-direction:column; } } diff --git a/internal/adminui/assets/models.html b/internal/adminui/assets/models.html new file mode 100644 index 0000000..0af0f50 --- /dev/null +++ b/internal/adminui/assets/models.html @@ -0,0 +1,50 @@ +<!doctype html> +<html lang="en"> +<head> + <meta charset="utf-8"> + <meta name="viewport" content="width=device-width, initial-scale=1"> + <meta name="description" content="Browse AIGW models, API protocols, context limits, and current token prices."> + <title>Models | AIGW</title> + <link rel="icon" href="data:image/svg+xml,%3Csvg xmlns='http://www.w3.org/2000/svg' viewBox='0 0 32 32'%3E%3Crect width='32' height='32' fill='%23102a3a'/%3E%3Ctext x='16' y='22' text-anchor='middle' font-family='Arial' font-weight='700' font-size='18' fill='%23b8eef7'%3EA%3C/text%3E%3C/svg%3E"> + <link rel="stylesheet" href="./models.css"> + <script src="./models.js" defer></script> +</head> +<body> + <header class="catalog-header"> + <a class="catalog-brand" href="./models" aria-label="AIGW model catalog"><span>A</span><strong>AIGW</strong><small>MODEL CATALOG</small></a> + <nav aria-label="Account access"><a class="button secondary" href="./">Sign in</a><a class="button primary" id="create-account" href="./?auth=register">Create account</a></nav> + </header> + <main> + <section class="catalog-intro"> + <div><span class="eyebrow">UNIFIED API</span><h1>Models</h1></div> + <div class="catalog-stats" aria-live="polite"><span><strong id="model-count">0</strong> models</span><span><strong id="provider-count">0</strong> routes</span><span><strong id="protocol-count">0</strong> APIs</span></div> + </section> + + <form class="catalog-filters" id="catalog-filters"> + <label class="search-field">Search<input id="search" type="search" placeholder="Model, developer, or capability" autocomplete="off"></label> + <label>Protocol<select id="protocol"><option value="">All protocols</option></select></label> + <label>Input<select id="input"><option value="">Any input</option><option value="text">Text</option><option value="image">Image</option><option value="audio">Audio</option><option value="video">Video</option></select></label> + <label>Developer<select id="owner"><option value="">All developers</option></select></label> + <label>Sort<select id="sort"><option value="newest">Newest</option><option value="name">Name</option><option value="input_price">Lowest input price</option><option value="output_price">Lowest output price</option><option value="context">Largest context</option></select></label> + </form> + + <div class="catalog-status hidden" id="catalog-error" role="alert"></div> + <div class="catalog-grid" id="catalog-grid" aria-live="polite"></div> + <div class="catalog-status hidden" id="catalog-empty">No models match these filters.</div> + </main> + + <dialog id="model-dialog"> + <article class="model-dialog"> + <header><div><span class="eyebrow" id="detail-owner"></span><h2 id="detail-name"></h2><code id="detail-id"></code></div><button class="icon-button" id="close-dialog" type="button" aria-label="Close">×</button></header> + <p id="detail-description"></p> + <div class="detail-tags" id="detail-tags"></div> + <dl class="detail-grid" id="detail-grid"></dl> + <section class="price-estimator"> + <div><span class="eyebrow">COST ESTIMATE</span><h3>Token estimate</h3></div> + <div class="estimate-inputs"><label>Input tokens<input id="estimate-input" type="number" min="0" step="100" value="1000"></label><label>Output tokens<input id="estimate-output" type="number" min="0" step="100" value="500"></label><output id="estimate-total"></output></div> + </section> + <footer><a class="button secondary" href="./">Sign in</a><a class="button primary" id="detail-register" href="./?auth=register">Start with this model</a></footer> + </article> + </dialog> +</body> +</html> diff --git a/internal/adminui/assets/models.js b/internal/adminui/assets/models.js new file mode 100644 index 0000000..4057ac5 --- /dev/null +++ b/internal/adminui/assets/models.js @@ -0,0 +1,126 @@ +'use strict'; + +const state = { models: [], selected: null, registrationEnabled: false }; +const $ = selector => document.querySelector(selector); +const esc = value => String(value ?? '').replace(/[&<>'"]/g, char => ({'&':'&','<':'<','>':'>',"'":''','"':'"'}[char])); +const integer = value => new Intl.NumberFormat().format(Number(value || 0)); +const protocolName = value => ({chat_completions:'Chat Completions',responses:'Responses',messages:'Anthropic Messages'}[value] || value); +const price = (micros, currency='usd') => new Intl.NumberFormat(undefined, {style:'currency', currency:String(currency).toUpperCase(), minimumFractionDigits:2, maximumFractionDigits:6}).format(Number(micros || 0) / 1_000_000); +const date = value => value ? new Intl.DateTimeFormat(undefined, {year:'numeric',month:'short',day:'numeric'}).format(new Date(value)) : 'Not published'; + +function healthLabel(model) { + if (model.health_status === 'unavailable') return ['Unavailable', 'unavailable']; + if (model.health_status === 'degraded') return [`${model.available_provider_count}/${model.provider_count} routes`, 'degraded']; + return ['Available', 'available']; +} + +function renderStats(models) { + const protocols = new Set(models.flatMap(model => model.supported_wire_apis || [])); + $('#model-count').textContent = integer(models.length); + $('#provider-count').textContent = integer(models.reduce((sum, model) => sum + Number(model.provider_count || 0), 0)); + $('#protocol-count').textContent = integer(protocols.size); +} + +function renderFilters() { + const owners = [...new Set(state.models.map(model => model.owned_by).filter(Boolean))].sort(); + const protocols = [...new Set(state.models.flatMap(model => model.supported_wire_apis || []))].sort(); + $('#owner').innerHTML = '<option value="">All developers</option>' + owners.map(owner => `<option value="${esc(owner)}">${esc(owner)}</option>`).join(''); + $('#protocol').innerHTML = '<option value="">All protocols</option>' + protocols.map(item => `<option value="${esc(item)}">${esc(protocolName(item))}</option>`).join(''); +} + +function filteredModels() { + const query = $('#search').value.trim().toLowerCase(); + const protocol = $('#protocol').value; + const input = $('#input').value; + const owner = $('#owner').value; + const result = state.models.filter(model => { + const haystack = [model.public_id, model.display_name, model.description, model.owned_by, ...(model.capabilities || []), ...(model.aliases || [])].join(' ').toLowerCase(); + return (!query || haystack.includes(query)) && (!protocol || (model.supported_wire_apis || []).includes(protocol)) && (!input || (model.input_modalities || []).includes(input)) && (!owner || model.owned_by === owner); + }); + const sort = $('#sort').value; + result.sort((a,b) => { + if (sort === 'name') return String(a.display_name || a.public_id).localeCompare(String(b.display_name || b.public_id)); + if (sort === 'input_price') return Number(a.input_price_micros_per_million) - Number(b.input_price_micros_per_million); + if (sort === 'output_price') return Number(a.output_price_micros_per_million) - Number(b.output_price_micros_per_million); + if (sort === 'context') return Number(b.context_window) - Number(a.context_window); + return new Date(b.released_at || 0) - new Date(a.released_at || 0) || String(a.public_id).localeCompare(String(b.public_id)); + }); + return result; +} + +function renderCatalog() { + const models = filteredModels(); + $('#catalog-empty').classList.toggle('hidden', models.length !== 0); + $('#catalog-grid').innerHTML = models.map(model => { + const [health, healthClass] = healthLabel(model); + return `<article class="model-card"> + <header><div><small>${esc(model.owned_by || 'Independent')}</small><h2>${esc(model.display_name || model.public_id)}</h2></div><span class="health ${healthClass}">${esc(health)}</span></header> + <code>${esc(model.public_id)}</code> + <p>${esc(model.description || 'No description published.')}</p> + <div class="model-tags">${(model.supported_wire_apis || []).map(item => `<span>${esc(protocolName(item))}</span>`).join('')}${(model.input_modalities || []).map(item => `<span>${esc(item)}</span>`).join('')}</div> + <dl><div><dt>Input</dt><dd>${price(model.input_price_micros_per_million, model.price_currency)} / 1M</dd></div><div><dt>Output</dt><dd>${price(model.output_price_micros_per_million, model.price_currency)} / 1M</dd></div><div><dt>Context</dt><dd>${integer(model.context_window)}</dd></div></dl> + <button class="button secondary" type="button" data-model="${esc(model.public_id)}">View model</button> + </article>`; + }).join(''); +} + +function renderEstimate() { + if (!state.selected) return; + const input = Math.max(0, Number($('#estimate-input').value || 0)); + const output = Math.max(0, Number($('#estimate-output').value || 0)); + const totalMicros = (Number(state.selected.input_price_micros_per_million || 0) * input + Number(state.selected.output_price_micros_per_million || 0) * output) / 1_000_000; + $('#estimate-total').textContent = `${price(totalMicros, state.selected.price_currency)} estimated`; +} + +function openModel(publicID) { + const model = state.models.find(item => item.public_id === publicID); + if (!model) return; + state.selected = model; + const [health, healthClass] = healthLabel(model); + $('#detail-owner').textContent = model.owned_by || 'Independent'; + $('#detail-name').textContent = model.display_name || model.public_id; + $('#detail-id').textContent = model.public_id; + $('#detail-description').textContent = model.description || 'No description published.'; + $('#detail-tags').innerHTML = [...(model.supported_wire_apis || []).map(protocolName), ...(model.capabilities || []), ...(model.input_modalities || []).map(item => `${item} input`), ...(model.output_modalities || []).map(item => `${item} output`)].map(item => `<span>${esc(item)}</span>`).join(''); + const rows = [ + ['Status', `<span class="health ${healthClass}">${esc(health)}</span>`], + ['Input price', `${price(model.input_price_micros_per_million, model.price_currency)} / 1M tokens`], + ['Output price', `${price(model.output_price_micros_per_million, model.price_currency)} / 1M tokens`], + ['Cached input', `${price(model.cache_read_price_micros_per_million, model.price_currency)} / 1M tokens`], + ['Context window', `${integer(model.context_window)} tokens`], + ['Max output', `${integer(model.max_output_tokens)} tokens`], + ['Regions', esc((model.regions || []).join(', ') || 'Global')], + ['Released', esc(date(model.released_at))] + ]; + $('#detail-grid').innerHTML = rows.map(([label,value]) => `<div><dt>${esc(label)}</dt><dd>${value}</dd></div>`).join(''); + const register = $('#detail-register'); + register.classList.toggle('hidden', !state.registrationEnabled); + register.href = `./?auth=register&model=${encodeURIComponent(model.public_id)}`; + renderEstimate(); + $('#model-dialog').showModal(); +} + +async function start() { + try { + const response = await fetch('./api/public/models', {headers:{Accept:'application/json'}}); + if (!response.ok) throw new Error(`Catalog request failed (${response.status})`); + const payload = await response.json(); + state.models = Array.isArray(payload.data) ? payload.data : []; + state.registrationEnabled = Boolean(payload.registration_enabled); + $('#create-account').classList.toggle('hidden', !state.registrationEnabled); + renderStats(state.models); + renderFilters(); + renderCatalog(); + } catch (error) { + $('#catalog-error').textContent = error.message || 'The model catalog is temporarily unavailable.'; + $('#catalog-error').classList.remove('hidden'); + } +} + +$('#catalog-filters').addEventListener('input', renderCatalog); +$('#catalog-grid').addEventListener('click', event => { const button = event.target.closest('[data-model]'); if (button) openModel(button.dataset.model); }); +$('#close-dialog').addEventListener('click', () => $('#model-dialog').close()); +$('#model-dialog').addEventListener('click', event => { if (event.target === $('#model-dialog')) $('#model-dialog').close(); }); +$('#estimate-input').addEventListener('input', renderEstimate); +$('#estimate-output').addEventListener('input', renderEstimate); +start(); diff --git a/internal/adminui/assets/style.css b/internal/adminui/assets/style.css index e712da3..b71dc20 100644 --- a/internal/adminui/assets/style.css +++ b/internal/adminui/assets/style.css @@ -1,18 +1,140 @@ :root { --bg:#f3f5f7; --panel:#fff; --ink:#18212b; --muted:#71808e; --line:#dce3e8; --accent:#146c94; --accent-soft:#e5f2f7; --danger:#b4494d; --shadow:0 8px 24px rgba(29,47,61,.06); font-family:Inter,ui-sans-serif,system-ui,-apple-system,BlinkMacSystemFont,"Segoe UI",sans-serif; } -* { box-sizing:border-box; } body { margin:0; color:var(--ink); background:var(--bg); font-size:14px; } button,input,select { font:inherit; } button { cursor:pointer; } +* { box-sizing:border-box; } body { margin:0; color:var(--ink); background:var(--bg); font-size:14px; } button,input,select,textarea { font:inherit; } button { cursor:pointer; } .auth-screen { min-height:100vh; display:grid; grid-template-columns:minmax(260px,1fr) minmax(360px,560px); background:#102a3a; } .auth-brand { color:#fff; display:flex; align-items:flex-start; gap:12px; padding:38px; } .auth-brand strong { display:block; font-size:18px; } .auth-brand small { display:block; color:#8ba9b9; font-size:10px; margin-top:3px; } .auth-panel { background:#fff; padding:clamp(30px,7vh,72px) clamp(28px,5vw,64px); overflow:auto; } .auth-tabs { display:flex; overflow:auto; border-bottom:1px solid var(--line); margin-bottom:34px; } .auth-tab { border:0; border-bottom:2px solid transparent; background:transparent; color:var(--muted); padding:11px 12px; white-space:nowrap; } .auth-tab.active { color:var(--accent); border-bottom-color:var(--accent); font-weight:700; } .auth-pane { display:none; gap:18px; } .auth-pane.active { display:grid; } .auth-pane h1 { margin-bottom:8px; } .auth-pane .button { margin-top:4px; } .auth-error { color:var(--danger); min-height:20px; margin:18px 0 0; font-size:12px; } .auth-error.success { color:#187151; } .form-actions { display:flex; gap:8px; } .form-actions .button { flex:1; } .topbar { height:72px; background:#102a3a; color:#fff; padding:0 32px; display:flex; align-items:center; justify-content:space-between; gap:24px; } .brand { display:flex; align-items:center; gap:11px; letter-spacing:0; } .brand-mark { width:32px; height:32px; display:grid; place-items:center; border:1px solid #8fd0df; color:#b8eef7; font-weight:800; } .brand strong { display:block; font-size:15px; } .brand small { color:#8ba9b9; font-size:9px; letter-spacing:0; } .session { display:flex; align-items:center; gap:12px; } .session div { display:flex; flex-direction:column; align-items:flex-end; gap:2px; } .session button { min-height:36px; border:1px solid #8fd0df; background:transparent; color:#b8eef7; padding:0 13px; font-weight:700; } .session button:hover { background:#18384b; } .state { color:#86d5ad; font-size:11px; text-transform:capitalize; } .actor-label { max-width:220px; overflow:hidden; text-overflow:ellipsis; white-space:nowrap; color:#fff; font-size:12px; } .shell { width:min(1240px,calc(100% - 48px)); margin:28px auto 60px; } .tabs { display:flex; flex-wrap:wrap; gap:4px; border-bottom:1px solid var(--line); margin-bottom:26px; } .tab { white-space:nowrap; border:0; background:transparent; color:var(--muted); padding:12px 15px; border-bottom:2px solid transparent; } .tab.active { color:var(--accent); border-bottom-color:var(--accent); font-weight:700; } .section { display:none; } .section.active { display:block; } .section-heading { display:flex; justify-content:space-between; align-items:flex-end; gap:20px; margin-bottom:19px; } .eyebrow { color:var(--accent); font-size:10px; letter-spacing:0; font-weight:800; } h1 { font-size:28px; line-height:1.1; margin:7px 0 0; letter-spacing:0; } h2 { margin:4px 0 0; font-size:20px; } .metric-grid { display:grid; grid-template-columns:repeat(6,1fr); gap:12px; } .metric { background:var(--panel); border:1px solid var(--line); padding:18px; box-shadow:var(--shadow); } .metric span,.metric small { display:block; color:var(--muted); } .metric strong { display:block; font-size:28px; margin:12px 0 3px; font-weight:750; } .metric small { font-size:11px; } .panel { background:var(--panel); border:1px solid var(--line); box-shadow:var(--shadow); padding:20px; margin-bottom:16px; } .note,.warning { display:flex; align-items:flex-start; gap:12px; } .note-icon { flex:0 0 22px; height:22px; border:1px solid var(--accent); color:var(--accent); display:grid; place-items:center; font-weight:700; } .note p { margin:5px 0 0; color:var(--muted); } .warning { color:#6d5523; background:#fff9e9; border-color:#ead9a9; box-shadow:none; } .warning span { margin-left:8px; color:#887650; } -.form-grid { display:grid; grid-template-columns:repeat(3,minmax(0,1fr)); align-items:end; gap:13px; } label { display:flex; flex-direction:column; gap:7px; color:var(--muted); font-size:12px; font-weight:650; } input,select { width:100%; border:1px solid var(--line); background:#fff; color:var(--ink); padding:10px 11px; min-height:40px; outline:none; } input:focus,select:focus { border-color:#69a9bf; box-shadow:0 0 0 3px var(--accent-soft); } .button { border:1px solid transparent; min-height:40px; padding:0 15px; font-weight:700; } .button.primary { color:#fff; background:var(--accent); } .button.primary:hover { background:#0d5879; } .button.secondary { color:var(--accent); background:var(--accent-soft); border-color:#c5e1ea; } .button.subtle { color:var(--accent); background:#fff; border-color:var(--line); grid-column:1; } -.table-wrap { overflow:auto; padding:0; } table { width:100%; border-collapse:collapse; min-width:700px; } th,td { padding:14px 18px; text-align:left; border-bottom:1px solid var(--line); vertical-align:middle; } th { color:var(--muted); font-size:11px; font-weight:700; text-transform:uppercase; letter-spacing:0; background:#fbfcfd; } tbody tr:last-child td { border-bottom:0; } td { font-size:13px; } code { font-family:"SFMono-Regular",Consolas,monospace; font-size:12px; color:#486071; } .badge,.tag { display:inline-flex; align-items:center; padding:4px 7px; font-size:11px; line-height:1; } .badge { border:1px solid #d9e0e4; color:var(--muted); } .badge.active { color:#187151; background:#e9f7f0; border-color:#c7e9d9; } .badge.revoked,.badge.suspended { color:var(--danger); background:#fff0f0; border-color:#f0cccc; } .tag { color:#4c6572; background:#eef3f5; margin:2px 3px 2px 0; } .text-button { border:0; background:transparent; color:var(--accent); padding:5px 0; } .text-button.danger { color:var(--danger); } .empty { color:var(--muted); text-align:center; padding:32px; } .truncate { max-width:280px; overflow:hidden; text-overflow:ellipsis; white-space:nowrap; } +.form-grid { display:grid; grid-template-columns:repeat(3,minmax(0,1fr)); align-items:end; gap:13px; } label { display:flex; flex-direction:column; gap:7px; color:var(--muted); font-size:12px; font-weight:650; } input,select,textarea { width:100%; border:1px solid var(--line); background:#fff; color:var(--ink); padding:10px 11px; min-height:40px; outline:none; } textarea { resize:vertical; line-height:1.5; } input:focus,select:focus,textarea:focus { border-color:#69a9bf; box-shadow:0 0 0 3px var(--accent-soft); } .button { border:1px solid transparent; min-height:40px; padding:0 15px; font-weight:700; } .button.primary { color:#fff; background:var(--accent); } .button.primary:hover { background:#0d5879; } .button.secondary { color:var(--accent); background:var(--accent-soft); border-color:#c5e1ea; } .button.subtle { color:var(--accent); background:#fff; border-color:var(--line); grid-column:1; } +.table-wrap { overflow:auto; padding:0; } table { width:100%; border-collapse:collapse; min-width:700px; } th,td { padding:14px 18px; text-align:left; border-bottom:1px solid var(--line); vertical-align:middle; } th { color:var(--muted); font-size:11px; font-weight:700; text-transform:uppercase; letter-spacing:0; background:#fbfcfd; } tbody tr:last-child td { border-bottom:0; } td { font-size:13px; } code { font-family:"SFMono-Regular",Consolas,monospace; font-size:12px; color:#486071; } .badge,.tag { display:inline-flex; align-items:center; padding:4px 7px; font-size:11px; line-height:1; } .badge { border:1px solid #d9e0e4; color:var(--muted); } .badge.active { color:#187151; background:#e9f7f0; border-color:#c7e9d9; } .badge.revoked,.badge.suspended,.badge.expired,.badge.unavailable { color:var(--danger); background:#fff0f0; border-color:#f0cccc; } .badge.degraded { color:#765b16; background:#fff8dc; border-color:#e9d990; } .tag { color:#4c6572; background:#eef3f5; margin:2px 3px 2px 0; } .text-button { border:0; background:transparent; color:var(--accent); padding:5px 0; } .text-button.danger { color:var(--danger); } .empty { color:var(--muted); text-align:center; padding:32px; } .truncate { max-width:280px; overflow:hidden; text-overflow:ellipsis; white-space:nowrap; } .route-editor { grid-column:1/-1; display:flex; flex-direction:column; gap:8px; } .route-row { display:grid; grid-template-columns:1.4fr 1.4fr .6fr .6fr 34px; gap:8px; } .icon-button { border:1px solid var(--line); background:#fff; color:var(--muted); width:34px; height:34px; font-size:19px; } .route-list { display:flex; flex-direction:column; gap:3px; color:#486071; font-size:12px; } .route-list em { color:var(--muted); font-style:normal; margin-left:4px; } .hidden { display:none !important; } .visually-hidden { position:absolute !important; width:1px !important; height:1px !important; padding:0 !important; margin:-1px !important; overflow:hidden !important; clip:rect(0,0,0,0) !important; white-space:nowrap !important; border:0 !important; } .billing-actions { display:grid; grid-template-columns:1fr 1fr; gap:16px; } .compact-form { grid-template-columns:1fr 1fr; } .compact-form .button { grid-column:1/-1; } .currency-label { color:var(--muted); font-size:12px; text-transform:uppercase; } .ledger-heading { margin-top:28px; } .money-positive { color:#187151; } .money-negative { color:var(--danger); } .account-grid { display:grid; grid-template-columns:1fr 1fr; gap:16px; } .account-grid .panel { align-content:start; } #totp-qr { width:180px; height:180px; border:1px solid var(--line); } #totp-secret { overflow-wrap:anywhere; } .table-wrap > .section-heading { padding:18px; margin:0; align-items:center; border-bottom:1px solid var(--line); } .price-line { display:block; color:var(--muted); font-size:10px; margin-top:5px; white-space:nowrap; } +.onboarding { display:grid; grid-template-columns:repeat(4,minmax(0,1fr)); gap:10px; padding:12px; } +.onboarding-step { display:flex; flex-direction:column; gap:7px; min-height:112px; padding:14px; border:1px solid var(--line); background:#fbfcfd; text-align:left; color:var(--ink); } +.onboarding-step strong { font-size:13px; } +.onboarding-step small { color:var(--muted); line-height:1.45; flex:1; } +.onboarding-step.done { border-color:#b9e1cf; background:#f0fbf5; } +.onboarding-step.done .step-state { color:#187151; } +.step-state { color:var(--accent); font-size:11px; font-weight:750; } +.quick-access-grid { display:grid; grid-template-columns:minmax(320px,.85fr) minmax(0,1.4fr); gap:16px; } +.starter-key-panel,.endpoint-panel { align-content:start; } +.starter-key-fields { display:grid; grid-template-columns:1fr 1fr auto; gap:12px; align-items:end; } +.starter-key-panel .muted { margin-bottom:0; } +.endpoint-list { display:grid; gap:0; } +.endpoint-row { display:grid; grid-template-columns:150px minmax(0,1fr) auto; align-items:center; gap:12px; padding:9px 0; border-bottom:1px solid var(--line); } +.endpoint-row span { color:var(--muted); font-size:12px; } +.endpoint-row code { overflow:hidden; text-overflow:ellipsis; white-space:nowrap; } +.endpoint-env { margin:14px 0 0; padding:13px; background:#102a3a; overflow:auto; white-space:pre-wrap; word-break:break-word; } +.endpoint-env code { color:#e9f7fb; line-height:1.55; } +.quickstart-grid { display:grid; grid-template-columns:minmax(0,1.5fr) minmax(260px,1fr); gap:16px; } +.quickstart-code pre { margin:17px 0 0; min-height:220px; background:#102a3a; color:#e9f7fb; padding:17px; overflow:auto; white-space:pre-wrap; word-break:break-word; } +.quickstart-code pre code { color:inherit; font-size:12px; line-height:1.6; } +.quickstart-next { align-content:start; } +.playground { margin-top:0; } +.playground .section-heading { align-items:center; } +.playground-fields { display:grid; grid-template-columns:1.4fr 1fr 1fr 1fr .72fr; align-items:end; gap:13px; } +.playground-key { grid-column:span 2; } +.playground-prompt { grid-column:span 4; } +.playground-prompt textarea { min-height:112px; } +.playground-actions { display:grid; gap:13px; } +.playground-result { margin-top:20px; padding-top:18px; border-top:1px solid var(--line); } +.playground-meta { display:flex; flex-wrap:wrap; align-items:center; gap:8px 16px; color:var(--muted); font-size:11px; } +.playground-result pre { margin:12px 0 0; min-height:160px; max-height:440px; background:#102a3a; color:#e9f7fb; padding:17px; overflow:auto; white-space:pre-wrap; word-break:break-word; } +.playground-result pre code { color:inherit; font-size:12px; line-height:1.6; } +.playground-diagnostic { display:flex; justify-content:space-between; align-items:center; gap:16px; margin-top:12px; padding:13px 15px; border:1px solid #e5c4c4; background:#fff6f6; color:#713b3b; } +.playground-diagnostic strong { font-size:12px; } +.playground-diagnostic p { margin:5px 0 0; color:#7b5a5a; font-size:12px; line-height:1.45; } +.preferences-grid { display:grid; grid-template-columns:1fr 1fr; gap:16px; } +.preferences-grid .panel { align-content:start; } +.preferences-grid h2,.preferences-grid .form-note { grid-column:1/-1; } +.toggle-row { grid-column:1/-1; flex-direction:row; align-items:center; min-height:40px; } +.toggle-row input { width:18px; min-height:18px; flex:0 0 18px; } +.form-note { margin:0; line-height:1.5; } +.signal-list { display:grid; gap:10px; margin-top:16px; } +.signal { display:flex; justify-content:space-between; align-items:center; gap:12px; border-bottom:1px solid var(--line); padding-bottom:10px; } +.signal:last-child { border-bottom:0; } +.signal span { color:var(--muted); font-size:12px; } +.catalog-toolbar { display:grid; grid-template-columns:2fr repeat(4,minmax(0,1fr)); gap:13px; } +.catalog-grid { display:grid; grid-template-columns:repeat(3,minmax(0,1fr)); gap:14px; } +.catalog-card { display:flex; flex-direction:column; gap:12px; background:var(--panel); border:1px solid var(--line); padding:18px; box-shadow:var(--shadow); min-width:0; } +.catalog-card h2 { font-size:17px; overflow-wrap:anywhere; } +.catalog-card .model-id { color:var(--muted); font-family:"SFMono-Regular",Consolas,monospace; font-size:12px; overflow-wrap:anywhere; } +.catalog-card p { color:#526673; line-height:1.5; margin:0; display:-webkit-box; -webkit-line-clamp:3; -webkit-box-orient:vertical; overflow:hidden; } +.provider-summary { display:flex; align-items:center; gap:8px; min-height:24px; } +.provider-summary small { color:var(--muted); line-height:1.35; } +.catalog-meta { display:flex; flex-wrap:wrap; gap:5px; } +.catalog-price { color:#187151; font-weight:750; font-size:12px; } +.catalog-card .button { margin-top:auto; align-self:flex-start; } +.catalog-actions { display:flex; flex-wrap:wrap; gap:8px; margin-top:auto; } +.catalog-actions .button { margin-top:0; } +.usage-filters { display:grid; grid-template-columns:repeat(6,minmax(0,1fr)); align-items:end; gap:12px; } +.usage-request-filter { grid-column:span 2; } +.usage-filters .form-actions { grid-column:span 2; } +.usage-metrics { margin-bottom:16px; } +.usage-chart { height:220px; overflow:auto; } +.chart-bars { min-width:680px; height:176px; display:flex; align-items:stretch; gap:7px; } +.chart-day { flex:1 0 22px; min-width:22px; display:grid; grid-template-rows:1fr 28px; gap:7px; text-align:center; } +.chart-bar { height:140px; display:flex; align-items:flex-end; background:#f3f6f7; border-bottom:1px solid var(--line); } +.chart-bar span { display:block; width:100%; min-height:4px; background:var(--accent); } +.chart-height-1 { height:5%; } .chart-height-2 { height:10%; } .chart-height-3 { height:15%; } .chart-height-4 { height:20%; } +.chart-height-5 { height:25%; } .chart-height-6 { height:30%; } .chart-height-7 { height:35%; } .chart-height-8 { height:40%; } +.chart-height-9 { height:45%; } .chart-height-10 { height:50%; } .chart-height-11 { height:55%; } .chart-height-12 { height:60%; } +.chart-height-13 { height:65%; } .chart-height-14 { height:70%; } .chart-height-15 { height:75%; } .chart-height-16 { height:80%; } +.chart-height-17 { height:85%; } .chart-height-18 { height:90%; } .chart-height-19 { height:95%; } .chart-height-20 { height:100%; } +.chart-day small { color:var(--muted); font-size:9px; white-space:nowrap; overflow:hidden; } +.analytics-grid { display:grid; grid-template-columns:1fr; gap:16px; margin:16px 0; } +.analytics-panel { padding-top:0; } +.analytics-panel .section-heading { min-width:760px; margin:0; padding:18px 18px 14px; } +.analytics-table { min-width:960px; } +.positive { color:#187151; } +.auto-topup-panel .section-heading { align-items:center; } +.billing-profile-panel { margin-top:16px; } +.billing-profile-panel .section-heading { align-items:center; } +.billing-profile-panel .form-grid { grid-template-columns:repeat(4,minmax(0,1fr)); } +.billing-profile-panel .billing-address-wide { grid-column:span 2; } +.billing-profile-actions { grid-column:1/-1; justify-content:space-between; align-items:center; } +.billing-profile-actions .error-label { margin:0; } +.billing-profile-panel input[name="country"] { text-transform:uppercase; } +.auto-topup-panel .form-grid { grid-template-columns:1.2fr 1fr 1fr; } +.auto-topup-panel .toggle-row { grid-column:1/-1; } +.auto-topup-actions { align-items:end; } +.auto-topup-actions .button { flex:1; } +#auto-topup-payment-method { display:block; margin-top:8px; font-size:13px; overflow-wrap:anywhere; } +.key-form .key-model-picker { grid-column:span 2; } +.key-model-picker select { min-height:126px; } +.key-model-picker small { color:var(--muted); font-weight:400; line-height:1.4; } +.key-form .button.primary { align-self:end; } +.keys-table { min-width:1220px; } +.key-tags { display:flex; flex-wrap:wrap; margin-top:5px; } .muted,.error-label { display:block; color:var(--muted); font-size:11px; margin-top:4px; } .error-label { color:var(--danger); } .limits-table input { min-width:118px; padding:8px 9px; } .limits-table .button { min-height:36px; } input:disabled,select:disabled { background:#f5f7f8; color:#697985; cursor:not-allowed; } .toast { position:fixed; bottom:24px; right:24px; background:#102a3a; color:#fff; padding:12px 16px; opacity:0; transform:translateY(8px); pointer-events:none; transition:.2s; } .toast.visible { opacity:1; transform:none; } .toast.error { background:#8f3d42; } dialog { border:0; padding:0; width:min(460px,calc(100% - 32px)); box-shadow:0 18px 70px rgba(0,0,0,.22); } dialog::backdrop { background:rgba(16,42,58,.45); } .dialog-content { padding:24px; } .dialog-content p { color:var(--muted); } .dialog-content code { display:block; background:#f3f5f7; padding:15px; overflow:auto; color:var(--ink); margin:18px 0; white-space:pre-wrap; } -@media (max-width:900px) { .metric-grid { grid-template-columns:repeat(3,1fr); } .form-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } .form-grid .button.primary { grid-column:1/-1; } .billing-actions,.account-grid { grid-template-columns:1fr; } } -@media (max-width:620px) { .auth-screen { grid-template-columns:1fr; background:#fff; } .auth-brand { background:#102a3a; padding:22px; } .auth-panel { padding:28px 22px 50px; } .auth-tabs { overflow:auto; } .topbar { height:auto; padding:16px; align-items:flex-start; } .session { margin-left:auto; } .session div { align-items:flex-end; max-width:150px; } .session .actor-label { max-width:150px; } .shell { width:calc(100% - 24px); margin-top:18px; } .tabs { margin-bottom:20px; flex-wrap:nowrap; overflow:auto; } .metric-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } .metric { padding:14px; } .metric strong { font-size:22px; overflow-wrap:anywhere; } .form-grid { grid-template-columns:1fr; } .route-row { grid-template-columns:minmax(0,1fr) minmax(0,1fr) 34px; } .route-provider,.route-upstream { grid-column:1/-1; } h1 { font-size:24px; } .section-heading { align-items:flex-start; } } +#model-dialog { width:min(680px,calc(100vw - 32px)); } +#usage-dialog { width:min(760px,calc(100vw - 32px)); } +.model-dialog-content { width:100%; max-height:calc(100vh - 32px); overflow:auto; } +.model-dialog-id { display:flex; justify-content:space-between; align-items:center; gap:12px; padding:10px 0 16px; border-bottom:1px solid var(--line); } +.model-dialog-id code { margin:0; padding:0; background:none; overflow-wrap:anywhere; } +.model-detail-grid { display:grid; grid-template-columns:1fr 1fr; gap:0 20px; margin:16px 0; } +.model-detail-row { display:flex; justify-content:space-between; align-items:flex-start; gap:12px; padding:10px 0; border-bottom:1px solid var(--line); min-width:0; } +.model-detail-row span { color:var(--muted); font-size:12px; flex:0 0 auto; } +.model-detail-row strong { text-align:right; overflow-wrap:anywhere; font-size:12px; } +.model-provider-health { border-top:1px solid var(--line); padding-top:16px; margin-bottom:18px; } +.model-provider-health .section-heading { margin-bottom:10px; } +.model-provider-health h3 { margin:0; font-size:15px; } +.model-provider-health small { display:block; margin-top:4px; color:var(--muted); } +.provider-health-table { min-width:620px; border:1px solid var(--line); } +.provider-health-table th,.provider-health-table td { padding:10px 12px; } +.provider-retry { max-width:155px; line-height:1.3; } +.model-estimator { border-top:1px solid var(--line); padding-top:16px; } +.model-estimator h3 { margin:0 0 12px; font-size:15px; } +.model-estimator .form-grid { grid-template-columns:repeat(4,minmax(0,1fr)); } +.model-estimate { grid-column:1/-1; align-self:end; color:#187151; font-weight:750; min-height:40px; display:flex; align-items:center; } +.model-dialog-actions { margin-top:18px; justify-content:flex-end; } +.request-dialog-content { width:100%; max-height:calc(100vh - 32px); overflow:auto; } +.request-detail-grid { display:grid; grid-template-columns:1fr 1fr; gap:0 20px; margin:16px 0; } +.request-diagnostic { margin:16px 0 0; max-height:260px; background:#102a3a; padding:15px; overflow:auto; white-space:pre-wrap; word-break:break-word; } +.request-diagnostic code { color:#e9f7fb; margin:0; padding:0; background:transparent; } +@media (max-width:900px) { .metric-grid { grid-template-columns:repeat(3,1fr); } .form-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } .form-grid .button.primary { grid-column:1/-1; } .billing-actions,.account-grid,.quickstart-grid,.quick-access-grid,.preferences-grid { grid-template-columns:1fr; } .catalog-toolbar { grid-template-columns:repeat(2,minmax(0,1fr)); } .catalog-toolbar label:first-child { grid-column:1/-1; } .catalog-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } .onboarding { grid-template-columns:repeat(2,minmax(0,1fr)); } .usage-filters { grid-template-columns:repeat(3,minmax(0,1fr)); } .auto-topup-panel .form-grid,.billing-profile-panel .form-grid { grid-template-columns:1fr 1fr; } .playground-fields { grid-template-columns:1fr 1fr; } .playground-key,.playground-prompt { grid-column:1/-1; } } +@media (max-width:620px) { .auth-screen { grid-template-columns:1fr; background:#fff; } .auth-brand { background:#102a3a; padding:22px; } .auth-panel { padding:28px 22px 50px; } .auth-tabs { overflow:auto; } .topbar { height:auto; padding:16px; align-items:flex-start; } .session { margin-left:auto; } .session div { align-items:flex-end; max-width:150px; } .session .actor-label { max-width:150px; } .shell { width:calc(100% - 24px); margin-top:18px; } .tabs { margin-bottom:20px; flex-wrap:nowrap; overflow:auto; } .metric-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } .metric { padding:14px; } .metric strong { font-size:22px; overflow-wrap:anywhere; } .form-grid,.catalog-toolbar,.usage-filters,.auto-topup-panel .form-grid,.billing-profile-panel .form-grid,.playground-fields,.starter-key-fields { grid-template-columns:1fr; } .catalog-toolbar label:first-child,.usage-request-filter,.usage-filters .form-actions,.key-form .key-model-picker,.billing-profile-panel .billing-address-wide,.playground-key,.playground-prompt { grid-column:auto; } .endpoint-row { grid-template-columns:1fr auto; } .endpoint-row code { grid-column:1/-1; grid-row:2; white-space:normal; overflow-wrap:anywhere; } .route-row { grid-template-columns:minmax(0,1fr) minmax(0,1fr) 34px; } .route-provider,.route-upstream { grid-column:1/-1; } h1 { font-size:24px; } .section-heading { align-items:flex-start; } .onboarding { grid-template-columns:1fr; } .catalog-grid { grid-template-columns:1fr; } .quickstart-code pre { min-height:250px; } .auto-topup-actions,.billing-profile-actions { flex-direction:column; align-items:stretch; } .playground-meta { align-items:flex-start; flex-direction:column; } .playground-diagnostic { align-items:flex-start; flex-direction:column; } .model-detail-grid,.request-detail-grid,.model-estimator .form-grid { grid-template-columns:1fr; } .model-dialog-actions { flex-direction:column; align-items:stretch; } } diff --git a/internal/adminui/ui.go b/internal/adminui/ui.go index 7f58030..49899ea 100644 --- a/internal/adminui/ui.go +++ b/internal/adminui/ui.go @@ -14,5 +14,14 @@ func Handler() http.Handler { if err != nil { panic(err) } - return http.FileServer(http.FS(content)) + files := http.FileServer(http.FS(content)) + mux := http.NewServeMux() + mux.HandleFunc("GET /models", func(w http.ResponseWriter, r *http.Request) { + http.ServeFileFS(w, r, content, "models.html") + }) + mux.HandleFunc("GET /models/", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "../models", http.StatusPermanentRedirect) + }) + mux.Handle("/", files) + return mux } diff --git a/internal/auth/static.go b/internal/auth/static.go index 495f90c..2b44158 100644 --- a/internal/auth/static.go +++ b/internal/auth/static.go @@ -8,6 +8,7 @@ import ( "net/http" "strings" "sync/atomic" + "time" "aigw/internal/domain" ) @@ -19,11 +20,14 @@ type Authenticator interface { } type KeyRecord struct { - Key string `json:"key"` - KeyID string `json:"key_id"` - TenantID string `json:"tenant_id"` - ProjectID string `json:"project_id"` - Scopes []string `json:"scopes"` + Key string `json:"key"` + KeyID string `json:"key_id"` + TenantID string `json:"tenant_id"` + ProjectID string `json:"project_id"` + Scopes []string `json:"scopes"` + AllowedModels []string `json:"allowed_models,omitempty"` + MonthlySpendMicros int64 `json:"monthly_spend_micros,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` } type StaticAuthenticator struct { @@ -65,9 +69,16 @@ func NewStatic(raw string, allowAnonymous bool) (*StaticAuthenticator, error) { return nil, fmt.Errorf("duplicate client API key at record %d", i) } seen[hash] = struct{}{} + allowedModels := make(map[string]struct{}, len(record.AllowedModels)) + for _, model := range record.AllowedModels { + if model = strings.TrimSpace(model); model != "" { + allowedModels[model] = struct{}{} + } + } hashed = append(hashed, HashedKeyRecord{Hash: hash, Principal: domain.Principal{ KeyID: record.KeyID, TenantID: record.TenantID, ProjectID: record.ProjectID, - Scopes: append([]string(nil), record.Scopes...), + Scopes: append([]string(nil), record.Scopes...), AllowedModels: allowedModels, + MonthlySpendMicros: record.MonthlySpendMicros, ExpiresAt: record.ExpiresAt, }}) } if len(hashed) == 0 && !allowAnonymous { @@ -88,6 +99,7 @@ func (a *StaticAuthenticator) ReplaceHashed(records []HashedKeyRecord) { for _, record := range records { principal := record.Principal principal.Scopes = append([]string(nil), principal.Scopes...) + principal.AllowedModels = cloneSet(principal.AllowedModels) keys[record.Hash] = principal } a.state.Store(&keySnapshot{keys: keys}) @@ -109,12 +121,23 @@ func (a *StaticAuthenticator) Authenticate(r *http.Request) (domain.Principal, e return domain.Principal{}, ErrUnauthorized } principal, ok := snapshot.keys[sha256.Sum256([]byte(key))] - if !ok { + if !ok || (principal.ExpiresAt != nil && !principal.ExpiresAt.After(time.Now())) { return domain.Principal{}, ErrUnauthorized } return principal, nil } +func cloneSet(source map[string]struct{}) map[string]struct{} { + if len(source) == 0 { + return nil + } + result := make(map[string]struct{}, len(source)) + for key := range source { + result[key] = struct{}{} + } + return result +} + func bearerToken(header string) string { parts := strings.Fields(header) if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") { diff --git a/internal/auth/static_test.go b/internal/auth/static_test.go index cf54ba6..a037ce1 100644 --- a/internal/auth/static_test.go +++ b/internal/auth/static_test.go @@ -4,6 +4,7 @@ import ( "crypto/sha256" "net/http" "testing" + "time" "aigw/internal/domain" ) @@ -74,3 +75,33 @@ func TestStaticAuthenticatorAcceptsAnthropicHeader(t *testing.T) { t.Fatal(err) } } + +func TestStaticAuthenticatorLoadsRestrictionsAndRejectsExpiredKey(t *testing.T) { + future := time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano) + authenticator, err := NewStatic(`[{"key":"sk-limited","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","allowed_models":["model/allowed"],"monthly_spend_micros":1250000,"expires_at":"`+future+`"}]`, false) + if err != nil { + t.Fatal(err) + } + request, _ := http.NewRequest(http.MethodGet, "http://gateway.test/v1/models", nil) + request.Header.Set("Authorization", "Bearer sk-limited") + principal, err := authenticator.Authenticate(request) + if err != nil { + t.Fatal(err) + } + if principal.MonthlySpendMicros != 1_250_000 { + t.Fatalf("monthly spend limit = %d", principal.MonthlySpendMicros) + } + if _, ok := principal.AllowedModels["model/allowed"]; !ok { + t.Fatalf("allowed model was not loaded: %+v", principal.AllowedModels) + } + + past := time.Now().Add(-time.Minute) + hash := sha256.Sum256([]byte("sk-expired")) + authenticator.ReplaceHashed([]HashedKeyRecord{{Hash: hash, Principal: domain.Principal{ + KeyID: "expired", TenantID: "tenant-1", ProjectID: "project-1", ExpiresAt: &past, + }}}) + request.Header.Set("Authorization", "Bearer sk-expired") + if _, err := authenticator.Authenticate(request); err != ErrUnauthorized { + t.Fatalf("expired key error = %v, want ErrUnauthorized", err) + } +} diff --git a/internal/billing/auto_topup.go b/internal/billing/auto_topup.go new file mode 100644 index 0000000..a90405b --- /dev/null +++ b/internal/billing/auto_topup.go @@ -0,0 +1,541 @@ +package billing + +import ( + "context" + "errors" + "fmt" + "net/url" + "strings" + "time" + + "github.com/jackc/pgx/v5" + "github.com/stripe/stripe-go/v86" +) + +type stripeSetupIntentRetriever func(context.Context, string, *stripe.SetupIntentRetrieveParams) (*stripe.SetupIntent, error) +type stripePaymentIntentCreator func(context.Context, *stripe.PaymentIntentCreateParams) (*stripe.PaymentIntent, error) +type stripePaymentIntentRetriever func(context.Context, string, *stripe.PaymentIntentRetrieveParams) (*stripe.PaymentIntent, error) + +const autoTopUpAction = "auto_topup_setup" + +func (s *Service) defaultAutoTopUpAmountMinor() int64 { + amount := int64(2000) + if amount < s.minTopUpMinor { + amount = s.minTopUpMinor + } + if amount > s.maxTopUpMinor { + amount = s.maxTopUpMinor + } + return amount +} + +func (s *Service) defaultAutoTopUpThresholdMicros() int64 { + amount, err := minorToMicros(s.currency, s.defaultAutoTopUpAmountMinor()) + if err != nil || amount <= 0 { + return 0 + } + if amount/4 > 5*microsPerUnit { + return 5 * microsPerUnit + } + return amount / 4 +} + +func (s *Service) GetAutoTopUpSettings(ctx context.Context, tenantID string) (AutoTopUpSettings, error) { + tenantID = strings.TrimSpace(tenantID) + if tenantID == "" { + return AutoTopUpSettings{}, ErrBillingAccountNotFound + } + var result AutoTopUpSettings + var paymentMethodID string + err := s.db.QueryRow(ctx, ` + SELECT t.id::text, COALESCE(w.currency,$2), $3::boolean, + COALESCE(a.enabled,FALSE), COALESCE(a.threshold_micros,$4), + COALESCE(a.topup_amount_minor,$5), COALESCE(a.stripe_payment_method_id,''), + COALESCE(a.payment_method_type,''), COALESCE(a.payment_method_brand,''), + COALESCE(a.payment_method_last4,''), COALESCE(a.payment_method_exp_month,0), + COALESCE(a.payment_method_exp_year,0), COALESCE(a.status,'not_configured'), + COALESCE(a.last_error,''), a.last_attempt_at, a.last_succeeded_at, + a.next_attempt_at, COALESCE(a.updated_at,t.created_at) + FROM tenants t + LEFT JOIN tenant_wallets w ON w.tenant_id=t.id + LEFT JOIN tenant_auto_topup_settings a ON a.tenant_id=t.id + WHERE t.id=$1`, tenantID, s.currency, s.stripeEnabled, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor()). + Scan(&result.TenantID, &result.Currency, &result.StripeEnabled, &result.Enabled, + &result.ThresholdMicros, &result.TopUpAmountMinor, &paymentMethodID, + &result.PaymentMethodType, &result.PaymentMethodBrand, &result.PaymentMethodLast4, + &result.PaymentMethodExpMonth, &result.PaymentMethodExpYear, &result.Status, + &result.LastError, &result.LastAttemptAt, &result.LastSucceededAt, + &result.NextAttemptAt, &result.UpdatedAt) + if errors.Is(err, pgx.ErrNoRows) { + return AutoTopUpSettings{}, ErrBillingAccountNotFound + } + if err != nil { + return AutoTopUpSettings{}, fmt.Errorf("query automatic top-up settings: %w", err) + } + result.PaymentMethodConfigured = paymentMethodID != "" + return result, nil +} + +func (s *Service) UpdateAutoTopUp(ctx context.Context, input UpdateAutoTopUpInput) (AutoTopUpSettings, error) { + input.TenantID = strings.TrimSpace(input.TenantID) + if input.TenantID == "" || input.ThresholdMicros < 0 || input.TopUpAmountMinor < s.minTopUpMinor || input.TopUpAmountMinor > s.maxTopUpMinor { + return AutoTopUpSettings{}, ErrInvalidAmount + } + topUpMicros, err := minorToMicros(s.currency, input.TopUpAmountMinor) + if err != nil || topUpMicros <= input.ThresholdMicros { + return AutoTopUpSettings{}, fmt.Errorf("%w: automatic top-up amount must exceed the balance threshold", ErrInvalidAmount) + } + if input.Enabled && !s.stripeEnabled { + return AutoTopUpSettings{}, ErrStripeDisabled + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return AutoTopUpSettings{}, err + } + defer tx.Rollback(ctx) + if _, err := tx.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor) + VALUES ($1,$2,$3) ON CONFLICT (tenant_id) DO NOTHING`, input.TenantID, input.ThresholdMicros, input.TopUpAmountMinor); err != nil { + return AutoTopUpSettings{}, fmt.Errorf("initialize automatic top-up settings: %w", err) + } + var paymentMethodID, status string + if err := tx.QueryRow(ctx, `SELECT COALESCE(stripe_payment_method_id,''),status FROM tenant_auto_topup_settings WHERE tenant_id=$1 FOR UPDATE`, input.TenantID).Scan(&paymentMethodID, &status); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return AutoTopUpSettings{}, ErrBillingAccountNotFound + } + return AutoTopUpSettings{}, err + } + if input.Enabled && paymentMethodID == "" { + return AutoTopUpSettings{}, ErrPaymentMethodRequired + } + if input.Enabled && status == "action_required" { + return AutoTopUpSettings{}, ErrAutoTopUpNeedsAttention + } + next := any(nil) + newStatus := "not_configured" + if paymentMethodID != "" { + newStatus = "ready" + } + if input.Enabled { + newStatus = "ready" + next = time.Now().UTC() + } + if _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=$2,threshold_micros=$3,topup_amount_minor=$4, + status=$5,last_error=CASE WHEN $2 THEN '' ELSE last_error END, + next_attempt_at=$6,updated_at=now() WHERE tenant_id=$1`, input.TenantID, input.Enabled, + input.ThresholdMicros, input.TopUpAmountMinor, newStatus, next); err != nil { + return AutoTopUpSettings{}, fmt.Errorf("update automatic top-up settings: %w", err) + } + if err := tx.Commit(ctx); err != nil { + return AutoTopUpSettings{}, err + } + return s.GetAutoTopUpSettings(ctx, input.TenantID) +} + +func (s *Service) DisableAutoTopUp(ctx context.Context, tenantID string) (AutoTopUpSettings, error) { + tenantID = strings.TrimSpace(tenantID) + if tenantID == "" { + return AutoTopUpSettings{}, ErrBillingAccountNotFound + } + if _, err := s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=FALSE,status=CASE WHEN stripe_payment_method_id IS NULL THEN 'not_configured' ELSE 'ready' END,next_attempt_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil { + return AutoTopUpSettings{}, err + } + return s.GetAutoTopUpSettings(ctx, tenantID) +} + +// CreateAutoTopUpSetupSession opens a Stripe-hosted SetupIntent flow. Stripe +// owns card collection; this service only receives a PaymentMethod ID after a +// signed webhook confirms that the setup succeeded. +func (s *Service) CreateAutoTopUpSetupSession(ctx context.Context, input AutoTopUpSetupInput) (AutoTopUpSetupResult, error) { + if !s.stripeEnabled || s.createStripeCheckout == nil { + return AutoTopUpSetupResult{}, ErrStripeDisabled + } + input.TenantID = strings.TrimSpace(input.TenantID) + if input.TenantID == "" { + return AutoTopUpSetupResult{}, ErrBillingAccountNotFound + } + if _, err := s.GetAutoTopUpSettings(ctx, input.TenantID); err != nil { + return AutoTopUpSetupResult{}, err + } + customerID, err := s.ensureStripeCustomer(ctx, input.TenantID) + if err != nil { + return AutoTopUpSetupResult{}, err + } + params := &stripe.CheckoutSessionCreateParams{ + Mode: stripe.String(string(stripe.CheckoutSessionModeSetup)), + Currency: stripe.String(s.currency), + ClientReferenceID: stripe.String(input.TenantID), + IntegrationIdentifier: stripe.String(s.integrationIdentifier), + SuccessURL: stripe.String(autoTopUpReturnURL(s.stripeSuccessURL, true)), + CancelURL: stripe.String(autoTopUpReturnURL(s.stripeCancelURL, false)), + Metadata: map[string]string{ + "aigw_action": autoTopUpAction, + "aigw_tenant_id": input.TenantID, + }, + } + 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)) + } + } + params.SetIdempotencyKey("aigw_autotopup_setup_" + randomHex(16)) + session, err := s.createStripeCheckout(ctx, params) + if err != nil { + return AutoTopUpSetupResult{}, fmt.Errorf("create automatic top-up setup session: %w", err) + } + if session.ID == "" || session.URL == "" { + return AutoTopUpSetupResult{}, errors.New("Stripe returned an incomplete setup session") + } + if _, err := s.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id) + VALUES ($1,$2,$3,$4) ON CONFLICT (tenant_id) DO UPDATE SET stripe_setup_session_id=EXCLUDED.stripe_setup_session_id,updated_at=now()`, + input.TenantID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), session.ID); err != nil { + return AutoTopUpSetupResult{}, fmt.Errorf("persist automatic top-up setup session: %w", err) + } + return AutoTopUpSetupResult{SessionID: session.ID, URL: session.URL}, nil +} + +func autoTopUpReturnURL(raw string, success bool) string { + parsed, err := url.Parse(raw) + if err != nil { + return raw + } + query := parsed.Query() + query.Set("autotopup", "setup") + if success { + query.Set("session_id", "{CHECKOUT_SESSION_ID}") + } else { + query.Set("autotopup", "cancel") + query.Del("session_id") + } + parsed.RawQuery = strings.ReplaceAll(query.Encode(), url.QueryEscape("{CHECKOUT_SESSION_ID}"), "{CHECKOUT_SESSION_ID}") + return parsed.String() +} + +func (s *Service) processAutoTopUpSetupEvent(ctx context.Context, event stripe.Event, session *stripe.CheckoutSession) error { + if s.retrieveStripeSetupIntent == nil || session == nil || event.ID == "" || session.ID == "" || session.ClientReferenceID == "" { + return ErrInvalidAmount + } + if event.Type != stripe.EventTypeCheckoutSessionCompleted { + return nil + } + setupIntentID := "" + if session.SetupIntent != nil { + setupIntentID = session.SetupIntent.ID + } + if setupIntentID == "" { + return ErrPaymentMethodRequired + } + intent, err := s.retrieveStripeSetupIntent(ctx, setupIntentID, &stripe.SetupIntentRetrieveParams{}) + if err != nil { + return fmt.Errorf("retrieve automatic top-up setup intent: %w", err) + } + if intent == nil || intent.Status != stripe.SetupIntentStatusSucceeded || intent.PaymentMethod == nil || intent.PaymentMethod.ID == "" { + return ErrPaymentMethodRequired + } + 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 + } + if _, 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)); err != nil { + return err + } + var storedSetupSession string + if err := tx.QueryRow(ctx, `SELECT COALESCE(stripe_setup_session_id,'') FROM tenant_auto_topup_settings WHERE tenant_id=$1 FOR UPDATE`, session.ClientReferenceID).Scan(&storedSetupSession); err != nil || storedSetupSession != session.ID { + return ErrInvalidAmount + } + customerID := "" + if session.Customer != nil { + customerID = session.Customer.ID + } + if customerID == "" && intent.Customer != nil { + customerID = intent.Customer.ID + } + if customerID == "" { + return ErrPaymentMethodRequired + } + 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()`, session.ClientReferenceID, customerID, email); err != nil { + return err + } + methodType, brand, last4 := intent.PaymentMethod.Type, "", "" + var expMonth, expYear int64 + if intent.PaymentMethod.Card != nil { + brand, last4 = string(intent.PaymentMethod.Card.Brand), intent.PaymentMethod.Card.Last4 + expMonth, expYear = intent.PaymentMethod.Card.ExpMonth, intent.PaymentMethod.Card.ExpYear + } + if _, err := tx.Exec(ctx, `INSERT INTO tenant_auto_topup_settings + (tenant_id,threshold_micros,topup_amount_minor,stripe_payment_method_id,payment_method_type,payment_method_brand,payment_method_last4,payment_method_exp_month,payment_method_exp_year,stripe_setup_session_id,status,last_error,failure_count,updated_at) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,'ready','',0,now()) + ON CONFLICT (tenant_id) DO UPDATE SET stripe_payment_method_id=EXCLUDED.stripe_payment_method_id, + payment_method_type=EXCLUDED.payment_method_type,payment_method_brand=EXCLUDED.payment_method_brand, + payment_method_last4=EXCLUDED.payment_method_last4,payment_method_exp_month=EXCLUDED.payment_method_exp_month, + payment_method_exp_year=EXCLUDED.payment_method_exp_year,stripe_setup_session_id=EXCLUDED.stripe_setup_session_id, + status='ready',last_error='',failure_count=0,next_attempt_at=NULL,updated_at=now()`, + session.ClientReferenceID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), intent.PaymentMethod.ID, + methodType, brand, last4, expMonth, expYear, 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) +} + +// processAutoTopUpOnce claims one eligible tenant before making a Stripe call. +// The row lock and unique pending-order index make this safe across gateways. +func (s *Service) processAutoTopUpOnce(ctx context.Context) (bool, error) { + if !s.stripeEnabled || s.createStripePaymentIntent == nil { + return false, nil + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return false, err + } + defer tx.Rollback(ctx) + var tenantID, customerID, customerEmail, paymentMethodID, currency, orderID string + var amountMinor, threshold, balance, reserved int64 + row := tx.QueryRow(ctx, ` + SELECT a.tenant_id::text,c.stripe_customer_id,c.email,a.stripe_payment_method_id,w.currency, + a.topup_amount_minor,a.threshold_micros,w.balance_micros,w.reserved_micros + ,COALESCE(o.id::text,'') + FROM tenant_auto_topup_settings a + JOIN stripe_customers c ON c.tenant_id=a.tenant_id + JOIN tenant_wallets w ON w.tenant_id=a.tenant_id + LEFT JOIN LATERAL (SELECT id FROM topup_orders WHERE tenant_id=a.tenant_id AND trigger_type='auto' AND status='pending' ORDER BY created_at DESC LIMIT 1) o ON TRUE + WHERE a.enabled AND a.stripe_payment_method_id IS NOT NULL + AND a.status IN ('ready','failed','charging') + AND (a.next_attempt_at IS NULL OR a.next_attempt_at <= now()) + AND w.balance_micros-w.reserved_micros <= a.threshold_micros + ORDER BY w.balance_micros-w.reserved_micros,a.updated_at + FOR UPDATE OF a SKIP LOCKED LIMIT 1`) + if err := row.Scan(&tenantID, &customerID, &customerEmail, &paymentMethodID, ¤cy, &amountMinor, &threshold, &balance, &reserved, &orderID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return false, tx.Commit(ctx) + } + return false, err + } + if balance-reserved > threshold { + return false, tx.Commit(ctx) + } + amountMicros, err := minorToMicros(currency, amountMinor) + if err != nil { + return false, err + } + if orderID == "" { + if err := tx.QueryRow(ctx, `INSERT INTO topup_orders (tenant_id,amount_minor,amount_micros,currency,trigger_type) + VALUES ($1,$2,$3,$4,'auto') RETURNING id::text`, tenantID, amountMinor, amountMicros, currency).Scan(&orderID); err != nil { + return false, err + } + } + if _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='charging',last_attempt_at=now(),next_attempt_at=now()+interval '15 minutes',updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil { + return false, err + } + if err := tx.Commit(ctx); err != nil { + return false, err + } + params := &stripe.PaymentIntentCreateParams{ + Amount: stripe.Int64(amountMinor), Currency: stripe.String(currency), Customer: stripe.String(customerID), + PaymentMethod: stripe.String(paymentMethodID), Confirm: stripe.Bool(true), OffSession: stripe.Bool(true), + ErrorOnRequiresAction: stripe.Bool(true), Description: stripe.String("AIGW automatic prepaid balance top-up"), + Metadata: map[string]string{"aigw_action": "auto_topup", "aigw_tenant_id": tenantID, "aigw_topup_order_id": orderID}, + } + if strings.TrimSpace(customerEmail) != "" { + params.ReceiptEmail = stripe.String(strings.TrimSpace(customerEmail)) + } + params.SetIdempotencyKey("aigw_autotopup_" + orderID) + intent, callErr := s.createStripePaymentIntent(ctx, params) + if callErr != nil { + var stripeErr *stripe.Error + if errors.As(callErr, &stripeErr) && stripeErr.PaymentIntent != nil { + intent = stripeErr.PaymentIntent + if intent.ID != "" { + _, _ = s.db.Exec(ctx, `UPDATE topup_orders SET stripe_payment_intent_id=$2,stripe_customer_id=$3 WHERE id=$1`, orderID, intent.ID, customerID) + } + if intent.Status == stripe.PaymentIntentStatusSucceeded { + return true, s.creditAutoTopUpPaymentIntent(ctx, intent) + } + return true, s.applyAutoTopUpPaymentIntentFailure(ctx, intent) + } + return true, s.scheduleAutoTopUpRetry(ctx, tenantID, orderID, callErr) + } + if intent == nil || intent.ID == "" { + return true, s.failAutoTopUp(ctx, tenantID, orderID, errors.New("Stripe returned an incomplete automatic top-up PaymentIntent")) + } + if _, err := s.db.Exec(ctx, `UPDATE topup_orders SET stripe_payment_intent_id=$2,stripe_customer_id=$3 WHERE id=$1`, orderID, intent.ID, customerID); err != nil { + return true, err + } + if intent.Status == stripe.PaymentIntentStatusSucceeded { + return true, s.creditAutoTopUpPaymentIntent(ctx, intent) + } + if intent.Status == stripe.PaymentIntentStatusProcessing { + return true, nil + } + if intent.Status == stripe.PaymentIntentStatusRequiresAction || intent.Status == stripe.PaymentIntentStatusRequiresPaymentMethod || intent.Status == stripe.PaymentIntentStatusCanceled { + return true, s.markAutoTopUpAttention(ctx, tenantID, orderID, ErrAutoTopUpNeedsAttention) + } + return true, s.failAutoTopUp(ctx, tenantID, orderID, fmt.Errorf("automatic top-up PaymentIntent ended in status %s", intent.Status)) +} + +func (s *Service) scheduleAutoTopUpRetry(ctx context.Context, tenantID, orderID string, cause error) error { + message := truncateError(cause) + _, err := s.db.Exec(ctx, `UPDATE topup_orders SET reconciliation_error=$2 WHERE id=$1 AND status='pending'`, orderID, message) + if err != nil { + return errors.Join(cause, err) + } + _, settingsErr := s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='failed',failure_count=failure_count+1,last_error=$2, + next_attempt_at=now()+interval '1 hour',updated_at=now() WHERE tenant_id=$1`, tenantID, message) + if settingsErr != nil { + return errors.Join(cause, settingsErr) + } + return cause +} + +func (s *Service) failAutoTopUp(ctx context.Context, tenantID, orderID string, cause error) error { + message := truncateError(cause) + _, err := s.db.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_error=$2 WHERE id=$1`, orderID, message) + if err != nil { + return errors.Join(cause, err) + } + _, settingsErr := s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='failed',failure_count=failure_count+1,last_error=$2, + next_attempt_at=now()+interval '1 hour',updated_at=now() WHERE tenant_id=$1`, tenantID, message) + if settingsErr != nil { + return errors.Join(cause, settingsErr) + } + return cause +} + +func (s *Service) markAutoTopUpAttention(ctx context.Context, tenantID, orderID string, cause error) error { + message := truncateError(cause) + _, err := s.db.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_error=$2 WHERE id=$1`, orderID, message) + if err != nil { + return err + } + _, err = s.db.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=FALSE,status='action_required',last_error=$2,next_attempt_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID, message) + return err +} + +func (s *Service) creditAutoTopUpPaymentIntent(ctx context.Context, intent *stripe.PaymentIntent) error { + if intent == nil || intent.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 := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, intent.ID); err != nil { + return err + } + if err := s.applyAutoTopUpPaymentIntentTx(ctx, tx, intent); err != nil { + return err + } + return tx.Commit(ctx) +} + +func (s *Service) applyAutoTopUpPaymentIntentTx(ctx context.Context, tx pgx.Tx, intent *stripe.PaymentIntent) error { + if intent.Status != stripe.PaymentIntentStatusSucceeded || intent.Metadata["aigw_action"] != "auto_topup" { + return nil + } + orderID, tenantID := intent.Metadata["aigw_topup_order_id"], intent.Metadata["aigw_tenant_id"] + if orderID == "" || tenantID == "" { + return ErrInvalidAmount + } + var amountMinor, amountMicros int64 + var currency, status, storedPI string + if err := tx.QueryRow(ctx, `SELECT amount_minor,amount_micros,currency,status,COALESCE(stripe_payment_intent_id,'') FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).Scan(&amountMinor, &amountMicros, ¤cy, &status, &storedPI); err != nil { + return err + } + if intent.Amount != amountMinor || string(intent.Currency) != currency || (storedPI != "" && storedPI != intent.ID) { + return ErrInvalidAmount + } + if intent.AmountReceived != 0 && intent.AmountReceived != amountMinor { + return ErrInvalidAmount + } + if intent.Customer != nil && intent.Customer.ID != "" { + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_customer_id=$2 WHERE id=$1`, orderID, intent.Customer.ID); err != nil { + return err + } + } + if status != "paid" { + if _, err := tx.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency) VALUES ($1,$2) ON CONFLICT DO NOTHING`, tenantID, currency); err != nil { + return err + } + 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 balance > int64(^uint64(0)>>1)-amountMicros { + return ErrInvalidAmount + } + newBalance := balance + amountMicros + if newBalance < reserved { + return ErrInsufficientBalance + } + var inserted int64 + if err := tx.QueryRow(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description) + VALUES ($1,$2,$3,$4,'topup','stripe_payment_intent',$5,'Automatic prepaid balance top-up') + ON CONFLICT (source_type,source_id) DO NOTHING RETURNING amount_micros`, tenantID, currency, amountMicros, newBalance, intent.ID).Scan(&inserted); err != nil && !errors.Is(err, pgx.ErrNoRows) { + return err + } + if inserted != 0 { + 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 topup_orders SET status='paid',paid_at=COALESCE(paid_at,now()),stripe_payment_intent_id=$2,reconciliation_status='ok',reconciled_at=now(),reconciliation_error='' WHERE id=$1`, orderID, intent.ID); err != nil { + return err + } + _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET status='ready',last_error='',failure_count=0,last_succeeded_at=now(),next_attempt_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID) + return err +} + +func (s *Service) applyAutoTopUpPaymentIntentFailure(ctx context.Context, intent *stripe.PaymentIntent) error { + if intent == nil || intent.Metadata["aigw_action"] != "auto_topup" { + return nil + } + tenantID, orderID := intent.Metadata["aigw_tenant_id"], intent.Metadata["aigw_topup_order_id"] + message := "automatic top-up payment failed" + if intent.LastPaymentError != nil && intent.LastPaymentError.Msg != "" { + message = intent.LastPaymentError.Msg + } + if intent.Status == stripe.PaymentIntentStatusRequiresAction || intent.Status == stripe.PaymentIntentStatusRequiresPaymentMethod || intent.Status == stripe.PaymentIntentStatusCanceled { + return s.markAutoTopUpAttention(ctx, tenantID, orderID, errors.New(message)) + } + return s.failAutoTopUp(ctx, tenantID, orderID, errors.New(message)) +} + +func (s *Service) applyAutoTopUpPaymentIntentFailureTx(ctx context.Context, tx pgx.Tx, intent *stripe.PaymentIntent) error { + if intent == nil || intent.Metadata["aigw_action"] != "auto_topup" { + return nil + } + tenantID, orderID := intent.Metadata["aigw_tenant_id"], intent.Metadata["aigw_topup_order_id"] + if tenantID == "" || orderID == "" { + return ErrInvalidAmount + } + message := "automatic top-up payment failed" + if intent.LastPaymentError != nil && intent.LastPaymentError.Msg != "" { + message = truncateError(errors.New(intent.LastPaymentError.Msg)) + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_error=$2,stripe_payment_intent_id=COALESCE(NULLIF($3,''),stripe_payment_intent_id) WHERE id=$1 AND status='pending'`, orderID, message, intent.ID); err != nil { + return err + } + status, enabled := "failed", true + if intent.Status == stripe.PaymentIntentStatusRequiresAction || intent.Status == stripe.PaymentIntentStatusRequiresPaymentMethod || intent.Status == stripe.PaymentIntentStatusCanceled { + status, enabled = "action_required", false + } + _, err := tx.Exec(ctx, `UPDATE tenant_auto_topup_settings SET enabled=$2,status=$3,last_error=$4,failure_count=failure_count+1, + next_attempt_at=CASE WHEN $2 THEN now()+interval '1 hour' ELSE NULL END,updated_at=now() WHERE tenant_id=$1`, tenantID, enabled, status, message) + return err +} diff --git a/internal/billing/auto_topup_test.go b/internal/billing/auto_topup_test.go new file mode 100644 index 0000000..85edb3d --- /dev/null +++ b/internal/billing/auto_topup_test.go @@ -0,0 +1,155 @@ +package billing + +import ( + "context" + "encoding/json" + "fmt" + "os" + "strings" + "testing" + "time" + + "aigw/internal/controlplane" + + "github.com/stripe/stripe-go/v86" +) + +func TestAutoTopUpReturnURL(t *testing.T) { + success := autoTopUpReturnURL("https://console.example.test/admin/?topup=success", true) + if !strings.Contains(success, "autotopup=setup") || !strings.Contains(success, "session_id={CHECKOUT_SESSION_ID}") { + t.Fatalf("unexpected setup return URL %q", success) + } + cancel := autoTopUpReturnURL("https://console.example.test/admin/?topup=cancel&session_id=old", false) + if !strings.Contains(cancel, "autotopup=cancel") || strings.Contains(cancel, "session_id=") { + t.Fatalf("unexpected setup cancel URL %q", cancel) + } +} + +func TestAutomaticTopUpSetupAndCreditAreIdempotentPostgres(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", MinTopUpMinor: 500, MaxTopUpMinor: 1_000_000, + StripeEnabled: true, StripeAPIKey: "rk_test_placeholder", StripeWebhookSecret: "whsec_integration_test", + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + + suffix := time.Now().UnixNano() + eventID := fmt.Sprintf("evt_auto_topup_%d", suffix) + var tenantID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Auto top-up integration') RETURNING id::text`, fmt.Sprintf("auto-topup-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency) VALUES ($1,'usd')`, tenantID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + cleanupCtx := context.Background() + if _, cleanupErr := service.db.Exec(cleanupCtx, `DELETE FROM stripe_webhook_events WHERE event_id=$1`, eventID); cleanupErr != nil { + t.Errorf("cleanup automatic top-up webhook: %v", cleanupErr) + } + for _, query := range []string{ + `DELETE FROM billing_ledger WHERE tenant_id=$1`, + `DELETE FROM topup_orders WHERE tenant_id=$1`, + `DELETE FROM tenant_auto_topup_settings WHERE tenant_id=$1`, + `DELETE FROM stripe_customers WHERE tenant_id=$1`, + `DELETE FROM tenant_wallets WHERE tenant_id=$1`, + `DELETE FROM tenants WHERE id=$1`, + } { + if _, cleanupErr := service.db.Exec(cleanupCtx, query, tenantID); cleanupErr != nil { + t.Errorf("cleanup automatic top-up integration data: %v", cleanupErr) + } + } + }) + + customerID := fmt.Sprintf("cus_auto_%d", suffix) + paymentMethodID := fmt.Sprintf("pm_auto_%d", suffix) + setupIntentID := fmt.Sprintf("seti_auto_%d", suffix) + setupSessionID := fmt.Sprintf("cs_auto_%d", suffix) + service.retrieveStripeSetupIntent = func(context.Context, string, *stripe.SetupIntentRetrieveParams) (*stripe.SetupIntent, error) { + return &stripe.SetupIntent{ + ID: setupIntentID, Status: stripe.SetupIntentStatusSucceeded, + Customer: &stripe.Customer{ID: customerID}, + PaymentMethod: &stripe.PaymentMethod{ID: paymentMethodID, Type: stripe.PaymentMethodTypeCard, + Card: &stripe.PaymentMethodCard{Brand: stripe.PaymentMethodCardBrandVisa, Last4: "4242", ExpMonth: 12, ExpYear: 2035}}, + }, nil + } + if _, err := service.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id) VALUES ($1,5000000,2000,$2)`, tenantID, setupSessionID); err != nil { + t.Fatal(err) + } + raw, err := json.Marshal(map[string]any{ + "id": setupSessionID, "object": "checkout.session", "client_reference_id": tenantID, + "customer": customerID, "customer_email": "developer@example.test", "setup_intent": setupIntentID, + "metadata": map[string]string{"aigw_action": autoTopUpAction, "aigw_tenant_id": tenantID}, + }) + if err != nil { + t.Fatal(err) + } + event := stripe.Event{ID: eventID, Type: stripe.EventTypeCheckoutSessionCompleted, Data: &stripe.EventData{Raw: raw}} + if err := service.processStripeEvent(ctx, event); err != nil { + t.Fatal(err) + } + if err := service.processStripeEvent(ctx, event); err != nil { + t.Fatalf("replayed setup event: %v", err) + } + settings, err := service.UpdateAutoTopUp(ctx, UpdateAutoTopUpInput{ + TenantID: tenantID, Enabled: true, ThresholdMicros: 1_000_000, TopUpAmountMinor: 2000, + }) + if err != nil { + t.Fatal(err) + } + if !settings.Enabled || !settings.PaymentMethodConfigured || settings.PaymentMethodLast4 != "4242" { + t.Fatalf("unexpected settings: %+v", settings) + } + + var createdIntent *stripe.PaymentIntent + stripeCalls := 0 + service.createStripePaymentIntent = func(_ context.Context, params *stripe.PaymentIntentCreateParams) (*stripe.PaymentIntent, error) { + stripeCalls++ + createdIntent = &stripe.PaymentIntent{ + ID: fmt.Sprintf("pi_auto_%d", suffix), Status: stripe.PaymentIntentStatusSucceeded, + Amount: *params.Amount, AmountReceived: *params.Amount, Currency: stripe.Currency(*params.Currency), + Customer: &stripe.Customer{ID: *params.Customer}, PaymentMethod: &stripe.PaymentMethod{ID: *params.PaymentMethod}, + Metadata: params.Metadata, + } + return createdIntent, nil + } + processed, err := service.processAutoTopUpOnce(ctx) + if err != nil || !processed { + t.Fatalf("process automatic top-up: processed=%v err=%v", processed, err) + } + processed, err = service.processAutoTopUpOnce(ctx) + if err != nil || processed { + t.Fatalf("second automatic top-up: processed=%v err=%v", processed, err) + } + if stripeCalls != 1 { + t.Fatalf("Stripe calls = %d, want 1", stripeCalls) + } + if err := service.creditAutoTopUpPaymentIntent(ctx, createdIntent); err != nil { + t.Fatalf("replayed successful PaymentIntent: %v", err) + } + + var balance, ledgerCount, paidOrders int64 + if err := service.db.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE tenant_id=$1 AND source_type='stripe_payment_intent'`, tenantID).Scan(&ledgerCount); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT count(*) FROM topup_orders WHERE tenant_id=$1 AND trigger_type='auto' AND status='paid'`, tenantID).Scan(&paidOrders); err != nil { + t.Fatal(err) + } + if balance != 20_000_000 || ledgerCount != 1 || paidOrders != 1 { + t.Fatalf("balance=%d ledger=%d paid_orders=%d", balance, ledgerCount, paidOrders) + } +} diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go index 2eb3d87..7a00081 100644 --- a/internal/billing/ledger.go +++ b/internal/billing/ledger.go @@ -77,7 +77,7 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in if limit < 1 || limit > 200 { limit = 50 } - query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status, + query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,trigger_type, 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,''), @@ -98,7 +98,7 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in result := make([]TopUpOrder, 0) for rows.Next() { var item TopUpOrder - if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, + if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, &item.TriggerType, &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, @@ -112,7 +112,7 @@ 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 - query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status, + query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status,trigger_type, 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,''), @@ -123,7 +123,7 @@ func (s *Service) GetTopUpOrder(ctx context.Context, tenantID, orderID string) ( 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.Currency, &result.Status, &result.TriggerType, &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) diff --git a/internal/billing/operations.go b/internal/billing/operations.go index a6dc653..461c59a 100644 --- a/internal/billing/operations.go +++ b/internal/billing/operations.go @@ -19,13 +19,13 @@ func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (Por 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") - } + customerID, err := s.ensureStripeCustomer(ctx, tenantID) + if err != nil { return PortalResult{}, err } + if customerID == "" { + return PortalResult{}, errors.New("no Stripe customer exists for this account") + } session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{ Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL), }) @@ -353,6 +353,12 @@ func (s *Service) RunStripeOperations(ctx context.Context) { cancel() s.refreshOperationalMetrics(ctx) for { + for i := 0; i < 4; i++ { + ok, _ := s.processAutoTopUpOnce(ctx) + if !ok { + break + } + } for i := 0; i < 8; i++ { ok, _ := s.processRefundOperation(ctx) if !ok { @@ -793,6 +799,74 @@ func (s *Service) Reconcile(ctx context.Context, limit int) (ReconciliationResul return s.failReconciliation(ctx, result, fmt.Errorf("record clean reconciliation for order %s: %w", item.id, err)) } } + if s.retrieveStripePaymentIntent == nil { + return s.failReconciliation(ctx, result, errors.New("Stripe PaymentIntent retrieval is unavailable")) + } + piRows, err := s.db.Query(ctx, `SELECT id::text,stripe_payment_intent_id,status,amount_minor,currency FROM topup_orders + WHERE trigger_type='auto' AND stripe_payment_intent_id IS NOT NULL ORDER BY created_at DESC LIMIT $1`, limit) + if err != nil { + return s.failReconciliation(ctx, result, err) + } + type paymentIntentOrder struct { + id, paymentIntent, status, currency string + amount int64 + } + var paymentIntentOrders []paymentIntentOrder + for piRows.Next() { + var item paymentIntentOrder + if err := piRows.Scan(&item.id, &item.paymentIntent, &item.status, &item.amount, &item.currency); err != nil { + piRows.Close() + return s.failReconciliation(ctx, result, err) + } + paymentIntentOrders = append(paymentIntentOrders, item) + } + piRows.Close() + for _, item := range paymentIntentOrders { + intent, retrieveErr := s.retrieveStripePaymentIntent(ctx, item.paymentIntent, &stripe.PaymentIntentRetrieveParams{}) + result.CheckedOrders++ + if retrieveErr != nil { + message := truncateError(retrieveErr) + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record automatic top-up retrieval failure for order %s: %w", item.id, err)) + } + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_retrieve_failed", "error": message}) + continue + } + if intent == nil || intent.Amount != item.amount || string(intent.Currency) != item.currency { + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "amount or currency mismatch"); err != nil { + return s.failReconciliation(ctx, result, err) + } + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_mismatch"}) + continue + } + expectedPaid := item.status == "paid" || item.status == "partially_refunded" || item.status == "refunded" || item.status == "disputed" + stripePaid := intent.Status == stripe.PaymentIntentStatusSucceeded + if stripePaid && !expectedPaid { + if repairErr := s.creditAutoTopUpPaymentIntent(ctx, intent); repairErr != nil { + message := truncateError(repairErr) + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_repair_failed", "error": message}) + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil { + return s.failReconciliation(ctx, result, err) + } + continue + } + result.Repairs = append(result.Repairs, map[string]any{"order_id": item.id, "type": "credited_paid_payment_intent"}) + if err := s.updateOrderReconciliation(ctx, item.id, "repaired", ""); err != nil { + return s.failReconciliation(ctx, result, err) + } + continue + } + if expectedPaid != stripePaid { + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "payment_intent_state_mismatch", "local_status": item.status, "stripe_status": intent.Status}) + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "payment state mismatch"); err != nil { + return s.failReconciliation(ctx, result, err) + } + continue + } + if err := s.updateOrderReconciliation(ctx, item.id, "ok", ""); err != nil { + return s.failReconciliation(ctx, result, err) + } + } result.MismatchCount = int64(len(result.Mismatches)) result.Status = "clean" if result.MismatchCount > 0 { diff --git a/internal/billing/profile.go b/internal/billing/profile.go new file mode 100644 index 0000000..3b08a3c --- /dev/null +++ b/internal/billing/profile.go @@ -0,0 +1,197 @@ +package billing + +import ( + "context" + "errors" + "fmt" + "net/mail" + "strings" + "unicode" + + "github.com/jackc/pgx/v5" + "github.com/stripe/stripe-go/v86" +) + +type stripeCustomerCreator func(context.Context, *stripe.CustomerCreateParams) (*stripe.Customer, error) +type stripeCustomerUpdater func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error) + +func (s *Service) GetBillingProfile(ctx context.Context, tenantID string) (BillingProfile, error) { + tenantID = strings.TrimSpace(tenantID) + if tenantID == "" { + return BillingProfile{}, ErrBillingAccountNotFound + } + var result BillingProfile + err := s.db.QueryRow(ctx, `SELECT t.id::text, + COALESCE(p.legal_name,t.name),COALESCE(p.billing_email,sc.email,''), + COALESCE(p.address_line1,''),COALESCE(p.address_line2,''),COALESCE(p.city,''), + COALESCE(p.region,''),COALESCE(p.postal_code,''),COALESCE(p.country,''), + p.tenant_id IS NOT NULL,sc.stripe_customer_id IS NOT NULL, + COALESCE(p.stripe_sync_status,CASE WHEN sc.stripe_customer_id IS NOT NULL THEN 'checkout_managed' ELSE 'not_configured' END), + p.stripe_synced_at,COALESCE(p.stripe_sync_error,''),p.updated_at + FROM tenants t + LEFT JOIN tenant_billing_profiles p ON p.tenant_id=t.id + LEFT JOIN stripe_customers sc ON sc.tenant_id=t.id + WHERE t.id=$1`, tenantID).Scan(&result.TenantID, &result.LegalName, &result.BillingEmail, + &result.AddressLine1, &result.AddressLine2, &result.City, &result.Region, &result.PostalCode, + &result.Country, &result.Configured, &result.StripeCustomerConfigured, &result.StripeSyncStatus, + &result.StripeSyncedAt, &result.StripeSyncError, &result.UpdatedAt) + if errors.Is(err, pgx.ErrNoRows) { + return BillingProfile{}, ErrBillingAccountNotFound + } + if err != nil { + return BillingProfile{}, fmt.Errorf("get billing profile: %w", err) + } + return result, nil +} + +func (s *Service) UpdateBillingProfile(ctx context.Context, input UpdateBillingProfileInput) (BillingProfile, error) { + normalized, err := normalizeBillingProfile(input) + if err != nil { + return BillingProfile{}, err + } + status := "disabled" + if s.stripeEnabled { + status = "pending" + } + _, err = s.db.Exec(ctx, `INSERT INTO tenant_billing_profiles + (tenant_id,legal_name,billing_email,address_line1,address_line2,city,region,postal_code,country,stripe_sync_status,stripe_synced_at,stripe_sync_error) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,NULL,'') + ON CONFLICT (tenant_id) DO UPDATE SET legal_name=EXCLUDED.legal_name, + billing_email=EXCLUDED.billing_email,address_line1=EXCLUDED.address_line1,address_line2=EXCLUDED.address_line2, + city=EXCLUDED.city,region=EXCLUDED.region,postal_code=EXCLUDED.postal_code,country=EXCLUDED.country, + stripe_sync_status=EXCLUDED.stripe_sync_status,stripe_synced_at=NULL,stripe_sync_error='',updated_at=now()`, + normalized.TenantID, normalized.LegalName, normalized.BillingEmail, normalized.AddressLine1, + normalized.AddressLine2, normalized.City, normalized.Region, normalized.PostalCode, normalized.Country, status) + if err != nil { + return BillingProfile{}, fmt.Errorf("save billing profile: %w", err) + } + if !s.stripeEnabled { + return s.GetBillingProfile(ctx, normalized.TenantID) + } + if _, err := s.ensureStripeCustomer(ctx, normalized.TenantID); err != nil { + return BillingProfile{}, err + } + return s.GetBillingProfile(ctx, normalized.TenantID) +} + +func (s *Service) ensureStripeCustomer(ctx context.Context, tenantID string) (string, error) { + var profile BillingProfile + err := s.db.QueryRow(ctx, `SELECT tenant_id::text,legal_name,billing_email,address_line1,address_line2, + city,region,postal_code,country,TRUE FROM tenant_billing_profiles WHERE tenant_id=$1`, tenantID). + Scan(&profile.TenantID, &profile.LegalName, &profile.BillingEmail, &profile.AddressLine1, + &profile.AddressLine2, &profile.City, &profile.Region, &profile.PostalCode, &profile.Country, &profile.Configured) + if errors.Is(err, pgx.ErrNoRows) { + var customerID string + err = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID) + if errors.Is(err, pgx.ErrNoRows) { + return "", nil + } + return customerID, err + } + if err != nil { + return "", fmt.Errorf("load billing profile for Stripe: %w", err) + } + if !s.stripeEnabled || s.createStripeCustomer == nil || s.updateStripeCustomer == nil { + return "", ErrStripeDisabled + } + var customerID string + err = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return "", fmt.Errorf("load Stripe customer: %w", err) + } + if customerID == "" { + params := billingProfileCustomerCreateParams(profile) + params.SetIdempotencyKey("aigw_customer_" + tenantID) + customer, createErr := s.createStripeCustomer(ctx, params) + if createErr != nil || customer == nil || strings.TrimSpace(customer.ID) == "" { + if createErr == nil { + createErr = errors.New("Stripe returned an incomplete Customer") + } + return "", s.failBillingProfileSync(ctx, tenantID, createErr) + } + customerID = customer.ID + } else { + customer, updateErr := s.updateStripeCustomer(ctx, customerID, billingProfileCustomerUpdateParams(profile)) + if updateErr != nil || customer == nil || strings.TrimSpace(customer.ID) == "" { + if updateErr == nil { + updateErr = errors.New("Stripe returned an incomplete Customer") + } + return "", s.failBillingProfileSync(ctx, tenantID, updateErr) + } + } + if _, err := s.db.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=EXCLUDED.email,updated_at=now()`, + tenantID, customerID, profile.BillingEmail); err != nil { + return "", fmt.Errorf("persist Stripe customer: %w", err) + } + if _, err := s.db.Exec(ctx, `UPDATE tenant_billing_profiles SET stripe_sync_status='synced',stripe_synced_at=now(), + stripe_sync_error='',updated_at=now() WHERE tenant_id=$1`, tenantID); err != nil { + return "", fmt.Errorf("record billing profile Stripe synchronization: %w", err) + } + return customerID, nil +} + +func (s *Service) failBillingProfileSync(ctx context.Context, tenantID string, cause error) error { + _, _ = s.db.Exec(ctx, `UPDATE tenant_billing_profiles SET stripe_sync_status='failed', + stripe_sync_error='Stripe customer synchronization failed',stripe_synced_at=NULL,updated_at=now() WHERE tenant_id=$1`, tenantID) + return fmt.Errorf("%w: %v", ErrBillingProfileSync, cause) +} + +func billingProfileCustomerCreateParams(profile BillingProfile) *stripe.CustomerCreateParams { + return &stripe.CustomerCreateParams{ + Name: stripe.String(profile.LegalName), BusinessName: stripe.String(profile.LegalName), + Email: stripe.String(profile.BillingEmail), Address: billingProfileAddress(profile), + Metadata: map[string]string{"aigw_tenant_id": profile.TenantID}, + } +} + +func billingProfileCustomerUpdateParams(profile BillingProfile) *stripe.CustomerUpdateParams { + return &stripe.CustomerUpdateParams{ + Name: stripe.String(profile.LegalName), BusinessName: stripe.String(profile.LegalName), + Email: stripe.String(profile.BillingEmail), Address: billingProfileAddress(profile), + Metadata: map[string]string{"aigw_tenant_id": profile.TenantID}, + } +} + +func billingProfileAddress(profile BillingProfile) *stripe.AddressParams { + return &stripe.AddressParams{Line1: stripe.String(profile.AddressLine1), Line2: stripe.String(profile.AddressLine2), + City: stripe.String(profile.City), State: stripe.String(profile.Region), PostalCode: stripe.String(profile.PostalCode), + Country: stripe.String(profile.Country)} +} + +func normalizeBillingProfile(input UpdateBillingProfileInput) (UpdateBillingProfileInput, error) { + input.TenantID = strings.TrimSpace(input.TenantID) + if input.TenantID == "" { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: tenant is required", ErrInvalidBillingProfile) + } + var err error + for _, field := range []struct { + value *string + name string + max int + required bool + }{ + {&input.LegalName, "legal name", 150, true}, {&input.BillingEmail, "billing email", 254, true}, + {&input.AddressLine1, "address line 1", 200, true}, {&input.AddressLine2, "address line 2", 200, false}, + {&input.City, "city", 100, true}, {&input.Region, "state or region", 100, false}, + {&input.PostalCode, "postal code", 32, true}, + } { + *field.value = strings.TrimSpace(*field.value) + if field.required && *field.value == "" { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: %s is required", ErrInvalidBillingProfile, field.name) + } + if len([]rune(*field.value)) > field.max || strings.IndexFunc(*field.value, unicode.IsControl) >= 0 { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: %s is invalid", ErrInvalidBillingProfile, field.name) + } + } + parsedEmail, err := mail.ParseAddress(input.BillingEmail) + if err != nil || !strings.EqualFold(parsedEmail.Address, input.BillingEmail) { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: billing email is invalid", ErrInvalidBillingProfile) + } + input.BillingEmail = strings.ToLower(parsedEmail.Address) + input.Country = strings.ToUpper(strings.TrimSpace(input.Country)) + if len(input.Country) != 2 || input.Country[0] < 'A' || input.Country[0] > 'Z' || input.Country[1] < 'A' || input.Country[1] > 'Z' { + return UpdateBillingProfileInput{}, fmt.Errorf("%w: country must be a two-letter code", ErrInvalidBillingProfile) + } + return input, nil +} diff --git a/internal/billing/profile_test.go b/internal/billing/profile_test.go new file mode 100644 index 0000000..e263bb8 --- /dev/null +++ b/internal/billing/profile_test.go @@ -0,0 +1,135 @@ +package billing + +import ( + "context" + "errors" + "fmt" + "os" + "testing" + "time" + + "aigw/internal/controlplane" + + "github.com/stripe/stripe-go/v86" +) + +func TestNormalizeBillingProfile(t *testing.T) { + input := UpdateBillingProfileInput{TenantID: " tenant ", LegalName: " Example Limited ", BillingEmail: "BILLING@EXAMPLE.TEST", + AddressLine1: " 1 Queen Street ", City: " Auckland ", PostalCode: "1010", Country: "nz"} + result, err := normalizeBillingProfile(input) + if err != nil { + t.Fatal(err) + } + if result.TenantID != "tenant" || result.LegalName != "Example Limited" || result.BillingEmail != "billing@example.test" || result.Country != "NZ" { + t.Fatalf("unexpected normalized profile: %+v", result) + } + for _, invalid := range []UpdateBillingProfileInput{ + {TenantID: "tenant", LegalName: "Example", BillingEmail: "not-an-email", AddressLine1: "1 Street", City: "Auckland", PostalCode: "1010", Country: "NZ"}, + {TenantID: "tenant", LegalName: "Example", BillingEmail: "billing@example.test", AddressLine1: "1 Street", City: "Auckland", PostalCode: "1010", Country: "New Zealand"}, + {TenantID: "tenant", BillingEmail: "billing@example.test", AddressLine1: "1 Street", City: "Auckland", PostalCode: "1010", Country: "NZ"}, + } { + if _, err := normalizeBillingProfile(invalid); !errors.Is(err, ErrInvalidBillingProfile) { + t.Fatalf("error = %v, want invalid billing profile", err) + } + } +} + +func TestBillingProfileStripeSynchronizationPostgres(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", StripeEnabled: true, StripeAPIKey: "rk_test_placeholder"}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + + suffix := time.Now().UnixNano() + var tenantID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Billing profile integration') RETURNING id::text`, fmt.Sprintf("billing-profile-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + cleanupCtx := context.Background() + for _, query := range []string{ + `DELETE FROM tenant_billing_profiles WHERE tenant_id=$1`, + `DELETE FROM stripe_customers WHERE tenant_id=$1`, + `DELETE FROM tenants WHERE id=$1`, + } { + if _, cleanupErr := service.db.Exec(cleanupCtx, query, tenantID); cleanupErr != nil { + t.Errorf("cleanup billing profile: %v", cleanupErr) + } + } + }) + + customerID := fmt.Sprintf("cus_profile_%d", suffix) + createCalls, updateCalls := 0, 0 + service.createStripeCustomer = func(_ context.Context, params *stripe.CustomerCreateParams) (*stripe.Customer, error) { + createCalls++ + if params.IdempotencyKey == nil || *params.IdempotencyKey != "aigw_customer_"+tenantID || params.Address == nil || *params.Address.Country != "NZ" || *params.Name != "Example Limited" || params.Metadata["aigw_tenant_id"] != tenantID { + t.Fatalf("unexpected Stripe create params: %+v", params) + } + return &stripe.Customer{ID: customerID}, nil + } + service.updateStripeCustomer = func(_ context.Context, id string, params *stripe.CustomerUpdateParams) (*stripe.Customer, error) { + updateCalls++ + if id != customerID || params.Address == nil || *params.Address.PostalCode != "1010" { + t.Fatalf("unexpected Stripe update: id=%s params=%+v", id, params) + } + return &stripe.Customer{ID: customerID}, nil + } + + input := UpdateBillingProfileInput{TenantID: tenantID, LegalName: "Example Limited", BillingEmail: "billing@example.test", + AddressLine1: "1 Queen Street", City: "Auckland", Region: "Auckland", PostalCode: "1010", Country: "nz"} + profile, err := service.UpdateBillingProfile(ctx, input) + if err != nil { + t.Fatal(err) + } + if createCalls != 1 || updateCalls != 0 || !profile.Configured || !profile.StripeCustomerConfigured || profile.StripeSyncStatus != "synced" || profile.StripeSyncedAt == nil { + t.Fatalf("unexpected first synchronization: create=%d update=%d profile=%+v", createCalls, updateCalls, profile) + } + + input.LegalName = "Example API Limited" + profile, err = service.UpdateBillingProfile(ctx, input) + if err != nil { + t.Fatal(err) + } + if createCalls != 1 || updateCalls != 1 || profile.LegalName != "Example API Limited" || profile.StripeSyncStatus != "synced" { + t.Fatalf("unexpected update synchronization: create=%d update=%d profile=%+v", createCalls, updateCalls, profile) + } + + service.updateStripeCustomer = func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error) { + return nil, errors.New("temporary Stripe outage") + } + input.City = "Wellington" + if _, err := service.UpdateBillingProfile(ctx, input); !errors.Is(err, ErrBillingProfileSync) { + t.Fatalf("error = %v, want Stripe sync error", err) + } + profile, err = service.GetBillingProfile(ctx, tenantID) + if err != nil { + t.Fatal(err) + } + if profile.City != "Wellington" || profile.StripeSyncStatus != "failed" || profile.StripeSyncError == "" { + t.Fatalf("failed synchronization did not preserve local profile: %+v", profile) + } + + service.updateStripeCustomer = func(context.Context, string, *stripe.CustomerUpdateParams) (*stripe.Customer, error) { + return &stripe.Customer{ID: customerID}, nil + } + if _, err := service.ensureStripeCustomer(ctx, tenantID); err != nil { + t.Fatal(err) + } + profile, err = service.GetBillingProfile(ctx, tenantID) + if err != nil { + t.Fatal(err) + } + if profile.StripeSyncStatus != "synced" || profile.StripeSyncError != "" { + t.Fatalf("profile did not recover after retry: %+v", profile) + } +} diff --git a/internal/billing/service.go b/internal/billing/service.go index 30f0e32..8a91839 100644 --- a/internal/billing/service.go +++ b/internal/billing/service.go @@ -25,24 +25,29 @@ import ( const microsPerUnit = int64(1_000_000) type Service struct { - db *pgxpool.Pool - currency string - defaultMaxOutputTokens int64 - minTopUpMinor int64 - maxTopUpMinor int64 - stripeEnabled bool - 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 + db *pgxpool.Pool + currency string + defaultMaxOutputTokens int64 + minTopUpMinor int64 + maxTopUpMinor int64 + stripeEnabled bool + stripeWebhookSecret string + stripeSuccessURL string + stripeCancelURL string + stripePortalReturnURL string + stripeAutomaticTax bool + stripeProductTaxCode string + integrationIdentifier string + createStripeCheckout stripeCheckoutCreator + createStripeCustomer stripeCustomerCreator + updateStripeCustomer stripeCustomerUpdater + retrieveStripeSetupIntent stripeSetupIntentRetriever + createStripePaymentIntent stripePaymentIntentCreator + retrieveStripePaymentIntent stripePaymentIntentRetriever + stripeClient *stripe.Client + settlementSpoolPath string + metrics OperationalMetrics + spoolMu sync.Mutex } func New(ctx context.Context, options Options) (*Service, error) { @@ -68,6 +73,11 @@ func New(ctx context.Context, options Options) (*Service, error) { if options.StripeEnabled { service.stripeClient = stripe.NewClient(options.StripeAPIKey) service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create + service.createStripeCustomer = service.stripeClient.V1Customers.Create + service.updateStripeCustomer = service.stripeClient.V1Customers.Update + service.retrieveStripeSetupIntent = service.stripeClient.V1SetupIntents.Retrieve + service.createStripePaymentIntent = service.stripeClient.V1PaymentIntents.Create + service.retrieveStripePaymentIntent = service.stripeClient.V1PaymentIntents.Retrieve } return service, nil } @@ -131,6 +141,21 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { return ErrQuotaExceeded } } + if input.Principal.MonthlySpendMicros > 0 { + period := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + nextPeriod := period.AddDate(0, 1, 0) + var used, pending int64 + if err := tx.QueryRow(ctx, `SELECT + COALESCE((SELECT sum(cost_micros) FROM usage_events WHERE key_id=$1 AND started_at >= $2 AND started_at < $3),0), + COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE key_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`, + input.Principal.KeyID, period, nextPeriod).Scan(&used, &pending); err != nil { + return fmt.Errorf("read API key monthly spend quota: %w", err) + } + limit := input.Principal.MonthlySpendMicros + if reserved > limit || used > limit-reserved || pending > limit-used-reserved { + return ErrQuotaExceeded + } + } if balance-held < reserved { return ErrInsufficientBalance } @@ -200,6 +225,11 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { if err != nil { return fmt.Errorf("load billing reservation: %w", err) } + if _, err := tx.Exec(ctx, `UPDATE api_keys + SET last_used_at = GREATEST(COALESCE(last_used_at, $2), $2) + WHERE id = $1`, keyID, event.StartedAt); err != nil { + return fmt.Errorf("update API key last used time: %w", err) + } if status != "pending" { return tx.Commit(ctx) } diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go index 6a1bb49..5e21b20 100644 --- a/internal/billing/service_test.go +++ b/internal/billing/service_test.go @@ -4,6 +4,7 @@ import ( "context" "crypto/sha256" "encoding/json" + "errors" "fmt" "net/http" "net/http/httptest" @@ -19,6 +20,72 @@ import ( "github.com/stripe/stripe-go/v86/webhook" ) +func TestAuthorizeEnforcesAPIKeyMonthlySpendCapPostgres(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", DefaultMaxOutputTokens: 10}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + slug := fmt.Sprintf("key-budget-%d", time.Now().UnixNano()) + var tenantID, projectID, keyID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Key budget') 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','Default') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { + t.Fatal(err) + } + hash := sha256.Sum256([]byte(slug)) + if err := service.db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash,monthly_spend_micros) VALUES ($1,$2,'limited','sk-test',$3,9) RETURNING id::text`, tenantID, projectID, hash[:]).Scan(&keyID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + for _, statement := range []struct { + query string + args []any + }{ + {`DELETE FROM billing_settlement_jobs WHERE request_id LIKE 'req_key_budget_%'`, nil}, + {`DELETE FROM billing_reservations WHERE tenant_id=$1`, []any{tenantID}}, + {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, []any{tenantID}}, + {`DELETE FROM api_keys WHERE id=$1`, []any{keyID}}, + {`DELETE FROM projects WHERE id=$1`, []any{projectID}}, + {`DELETE FROM tenants WHERE id=$1`, []any{tenantID}}, + } { + if _, cleanupErr := service.db.Exec(context.Background(), statement.query, statement.args...); cleanupErr != nil { + t.Errorf("cleanup API key budget test data: %v", cleanupErr) + } + } + }) + if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros) VALUES ($1,'usd',1000000)`, tenantID); err != nil { + t.Fatal(err) + } + model := domain.Model{ID: "model/budget", PriceCurrency: "usd", OutputPriceMicrosPerMillion: 1_000_000} + principal := domain.Principal{TenantID: tenantID, ProjectID: projectID, KeyID: keyID, MonthlySpendMicros: 9} + err = service.Authorize(ctx, Authorization{RequestID: "req_key_budget_rejected", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}) + if !errors.Is(err, ErrQuotaExceeded) { + t.Fatalf("Authorize error = %v, want ErrQuotaExceeded", err) + } + var rejectedReservations int + if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_reservations WHERE request_id='req_key_budget_rejected'`).Scan(&rejectedReservations); err != nil { + t.Fatal(err) + } + if rejectedReservations != 0 { + t.Fatalf("quota rejection left %d reservation rows", rejectedReservations) + } + principal.MonthlySpendMicros = 10 + if err := service.Authorize(ctx, Authorization{RequestID: "req_key_budget_allowed", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}); err != nil { + t.Fatalf("Authorize at exact cap: %v", err) + } +} + func TestUsageCostUsesFixedPointAndRoundsOnce(t *testing.T) { usage := domain.Usage{InputTokens: 3, OutputTokens: 2, CacheReadInputTokens: 5} cost, err := usageCost(usage, 150_000, 600_000, 30_000, 0) diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go index 59ee9bd..032eb02 100644 --- a/internal/billing/stripe.go +++ b/internal/billing/stripe.go @@ -58,8 +58,13 @@ 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) + customerID, err := s.ensureStripeCustomer(ctx, input.TenantID) + if err != nil { + 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(err, fmt.Errorf("mark top-up order failed: %w", updateErr)) + } + return CheckoutResult{}, err + } if customerID != "" { params.Customer = stripe.String(customerID) } else { @@ -154,6 +159,7 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er return s.processCheckoutEvent(ctx, event) case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded, stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed, + stripe.EventTypePaymentIntentSucceeded, stripe.EventTypePaymentIntentPaymentFailed, stripe.EventTypePaymentIntentCanceled, stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated, stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn, stripe.EventTypeChargeDisputeFundsReinstated, @@ -176,6 +182,9 @@ func (s *Service) processCheckoutEvent(ctx context.Context, event stripe.Event) if err := json.Unmarshal(event.Data.Raw, &session); err != nil { return ErrInvalidAmount } + if session.Metadata["aigw_action"] == autoTopUpAction { + return s.processAutoTopUpSetupEvent(ctx, event, &session) + } if event.ID == "" || session.ID == "" || session.ClientReferenceID == "" { return ErrInvalidAmount } @@ -330,6 +339,18 @@ func (s *Service) processOperationalStripeEvent(ctx context.Context, event strip } } switch event.Type { + case stripe.EventTypePaymentIntentSucceeded, stripe.EventTypePaymentIntentPaymentFailed, stripe.EventTypePaymentIntentCanceled: + var intent stripe.PaymentIntent + if json.Unmarshal(event.Data.Raw, &intent) != nil || intent.ID == "" { + return ErrInvalidAmount + } + if event.Type == stripe.EventTypePaymentIntentSucceeded { + if err := s.applyAutoTopUpPaymentIntentTx(ctx, tx, &intent); err != nil { + return err + } + } else if err := s.applyAutoTopUpPaymentIntentFailureTx(ctx, tx, &intent); err != nil { + return err + } case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded: var charge stripe.Charge if json.Unmarshal(event.Data.Raw, &charge) != nil || charge.ID == "" { diff --git a/internal/billing/types.go b/internal/billing/types.go index fe5df1e..1633ecc 100644 --- a/internal/billing/types.go +++ b/internal/billing/types.go @@ -9,13 +9,18 @@ import ( ) var ( - ErrInsufficientBalance = errors.New("insufficient balance") - ErrStripeDisabled = errors.New("Stripe top-ups are disabled") - 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") + ErrInsufficientBalance = errors.New("insufficient balance") + ErrStripeDisabled = errors.New("Stripe top-ups are disabled") + 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") + ErrBillingAccountNotFound = errors.New("billing account not found") + ErrPaymentMethodRequired = errors.New("a saved payment method is required") + ErrAutoTopUpNeedsAttention = errors.New("automatic top-up payment method requires attention") + ErrInvalidBillingProfile = errors.New("invalid billing profile") + ErrBillingProfileSync = errors.New("billing profile Stripe synchronization failed") ) type Meter interface { @@ -134,6 +139,36 @@ type CheckoutResult struct { URL string `json:"url"` } +type BillingProfile struct { + TenantID string `json:"tenant_id"` + LegalName string `json:"legal_name"` + BillingEmail string `json:"billing_email"` + AddressLine1 string `json:"address_line1"` + AddressLine2 string `json:"address_line2"` + City string `json:"city"` + Region string `json:"region"` + PostalCode string `json:"postal_code"` + Country string `json:"country"` + Configured bool `json:"configured"` + StripeCustomerConfigured bool `json:"stripe_customer_configured"` + StripeSyncStatus string `json:"stripe_sync_status"` + StripeSyncedAt *time.Time `json:"stripe_synced_at,omitempty"` + StripeSyncError string `json:"stripe_sync_error,omitempty"` + UpdatedAt *time.Time `json:"updated_at,omitempty"` +} + +type UpdateBillingProfileInput struct { + TenantID string `json:"tenant_id"` + LegalName string `json:"legal_name"` + BillingEmail string `json:"billing_email"` + AddressLine1 string `json:"address_line1"` + AddressLine2 string `json:"address_line2"` + City string `json:"city"` + Region string `json:"region"` + PostalCode string `json:"postal_code"` + Country string `json:"country"` +} + type TopUpOrder struct { ID string `json:"id"` TenantID string `json:"tenant_id"` @@ -141,6 +176,7 @@ type TopUpOrder struct { AmountMicros int64 `json:"amount_micros"` Currency string `json:"currency"` Status string `json:"status"` + TriggerType string `json:"trigger_type"` StripeSessionID string `json:"stripe_session_id,omitempty"` CheckoutURL string `json:"checkout_url,omitempty"` CreatedAt time.Time `json:"created_at"` @@ -159,6 +195,44 @@ type TopUpOrder struct { ReconciliationError string `json:"reconciliation_error,omitempty"` } +type AutoTopUpSettings struct { + TenantID string `json:"tenant_id"` + Currency string `json:"currency"` + StripeEnabled bool `json:"stripe_enabled"` + Enabled bool `json:"enabled"` + ThresholdMicros int64 `json:"threshold_micros"` + TopUpAmountMinor int64 `json:"topup_amount_minor"` + PaymentMethodConfigured bool `json:"payment_method_configured"` + PaymentMethodType string `json:"payment_method_type,omitempty"` + PaymentMethodBrand string `json:"payment_method_brand,omitempty"` + PaymentMethodLast4 string `json:"payment_method_last4,omitempty"` + PaymentMethodExpMonth int64 `json:"payment_method_exp_month,omitempty"` + PaymentMethodExpYear int64 `json:"payment_method_exp_year,omitempty"` + Status string `json:"status"` + LastError string `json:"last_error,omitempty"` + LastAttemptAt *time.Time `json:"last_attempt_at,omitempty"` + LastSucceededAt *time.Time `json:"last_succeeded_at,omitempty"` + NextAttemptAt *time.Time `json:"next_attempt_at,omitempty"` + UpdatedAt time.Time `json:"updated_at"` +} + +type UpdateAutoTopUpInput struct { + TenantID string `json:"tenant_id"` + Enabled bool `json:"enabled"` + ThresholdMicros int64 `json:"threshold_micros"` + TopUpAmountMinor int64 `json:"topup_amount_minor"` +} + +type AutoTopUpSetupInput struct { + TenantID string `json:"tenant_id"` + CustomerEmail string `json:"-"` +} + +type AutoTopUpSetupResult struct { + SessionID string `json:"session_id"` + URL string `json:"url"` +} + type ResolveMissingTopUpInput struct { Reason string `json:"reason"` } diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go index 09a6fba..362d5b0 100644 --- a/internal/catalog/catalog.go +++ b/internal/catalog/catalog.go @@ -23,9 +23,15 @@ type snapshot struct { func New(cfg config.Config) *Catalog { providers := make(map[string]domain.Provider, len(cfg.Providers)) for _, provider := range cfg.Providers { + slug := provider.Slug + if slug == "" { + slug = provider.ID + } providers[provider.ID] = domain.Provider{ ID: provider.ID, + Slug: slug, Protocol: provider.Protocol, + WireAPI: provider.WireAPI, BaseURL: strings.TrimRight(provider.BaseURL, "/"), APIKey: provider.APIKey, } @@ -124,7 +130,7 @@ func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { result := make([]domain.Model, 0, len(current.list)) for _, model := range current.list { for _, route := range model.Routes { - if route.Provider.Protocol == protocol { + if protocolCompatible(route.Provider, protocol) { result = append(result, model) break } @@ -133,6 +139,19 @@ func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { return result } +func protocolCompatible(provider domain.Provider, requestProtocol domain.Protocol) bool { + switch requestProtocol { + case domain.ProtocolOpenAI: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions" + case domain.ProtocolOpenAIResponses: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses" + case domain.ProtocolAnthropic: + return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages" + default: + return false + } +} + func (c *Catalog) Count() int { current := c.state.Load() if current == nil { diff --git a/internal/catalog/catalog_test.go b/internal/catalog/catalog_test.go index 07fb71f..2ce8779 100644 --- a/internal/catalog/catalog_test.go +++ b/internal/catalog/catalog_test.go @@ -41,3 +41,13 @@ func TestCatalogReplaceCopiesRouteSlices(t *testing.T) { t.Fatal("catalog snapshot aliases the caller's route slice") } } + +func TestModelRestrictionIsAppliedBeforeCatalogAccess(t *testing.T) { + principal := domain.Principal{AllowedModels: map[string]struct{}{"model/allowed": {}}} + if !(domain.Model{ID: "model/allowed"}).Allows(principal) { + t.Fatal("allowed model was rejected") + } + if (domain.Model{ID: "model/other"}).Allows(principal) { + t.Fatal("model outside the API key restriction was allowed") + } +} diff --git a/internal/config/config.go b/internal/config/config.go index a67f221..047f372 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -81,6 +81,8 @@ type AdminConfig struct { SecurityRetentionDays int `json:"security_retention_days"` PublicURL string `json:"-"` PublicURLEnv string `json:"public_url_env"` + InferencePublicURL string `json:"-"` + InferencePublicURLEnv string `json:"inference_public_url_env"` Mail MailConfig `json:"mail"` WebAuthn WebAuthnConfig `json:"webauthn"` Token string `json:"-"` @@ -125,7 +127,9 @@ type UpstreamHTTPConfig struct { type ProviderConfig struct { ID string `json:"id"` + Slug string `json:"slug"` Protocol domain.Protocol `json:"protocol"` + WireAPI string `json:"wire_api"` BaseURL string `json:"-"` BaseURLEnv string `json:"base_url_env"` APIKeyEnv string `json:"api_key_env"` @@ -167,6 +171,7 @@ type BillingConfig struct { type StripeConfig struct { Enabled bool `json:"enabled"` + EnabledEnv string `json:"enabled_env"` APIKeyEnv string `json:"api_key_env"` WebhookSecretEnv string `json:"webhook_secret_env"` SuccessURLEnv string `json:"success_url_env"` @@ -297,6 +302,9 @@ func applyDefaults(cfg *Config) { if cfg.Admin.PublicURLEnv == "" { cfg.Admin.PublicURLEnv = "AIGW_PUBLIC_URL" } + if cfg.Admin.InferencePublicURLEnv == "" { + cfg.Admin.InferencePublicURLEnv = "AIGW_INFERENCE_PUBLIC_URL" + } if cfg.Admin.Mail.FromName == "" { cfg.Admin.Mail.FromName = "AIGW" } @@ -400,6 +408,18 @@ func applyDefaults(cfg *Config) { } } } + for i := range cfg.Providers { + if cfg.Providers[i].Slug == "" { + cfg.Providers[i].Slug = cfg.Providers[i].ID + } + if cfg.Providers[i].WireAPI == "" { + if cfg.Providers[i].Protocol == domain.ProtocolAnthropic { + cfg.Providers[i].WireAPI = "messages" + } else { + cfg.Providers[i].WireAPI = "chat_completions" + } + } + } } func resolveSecrets(cfg *Config) error { @@ -446,6 +466,13 @@ func resolveSecrets(cfg *Config) error { if err := resolveRequiredEnv(&cfg.Admin.PublicURL, cfg.Admin.PublicURLEnv, "admin.public_url"); err != nil { return err } + cfg.Admin.InferencePublicURL = strings.TrimRight(strings.TrimSpace(os.Getenv(cfg.Admin.InferencePublicURLEnv)), "/") + if cfg.Admin.InferencePublicURL == "" { + publicURL, err := url.Parse(cfg.Admin.PublicURL) + if err == nil && publicURL.Scheme != "" && publicURL.Host != "" { + cfg.Admin.InferencePublicURL = publicURL.Scheme + "://" + publicURL.Host + } + } if cfg.Admin.Mail.Enabled { cfg.Admin.Mail.FromAddress = strings.TrimSpace(os.Getenv(cfg.Admin.Mail.FromAddressEnv)) cfg.Admin.Mail.SMTPAddress = strings.TrimSpace(os.Getenv(cfg.Admin.Mail.SMTPAddressEnv)) @@ -462,6 +489,13 @@ func resolveSecrets(cfg *Config) error { } } } + if cfg.Billing.Enabled && cfg.Billing.Stripe.EnabledEnv != "" { + var err error + cfg.Billing.Stripe.Enabled, err = envBool(cfg.Billing.Stripe.EnabledEnv) + if err != nil { + return err + } + } if cfg.Billing.Enabled && cfg.Billing.Stripe.Enabled { cfg.Billing.Stripe.APIKey = os.Getenv(cfg.Billing.Stripe.APIKeyEnv) cfg.Billing.Stripe.WebhookSecret = os.Getenv(cfg.Billing.Stripe.WebhookSecretEnv) @@ -586,6 +620,10 @@ func Validate(cfg Config) error { 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") } + inferenceURL, err := url.Parse(cfg.Admin.InferencePublicURL) + if err != nil || inferenceURL.Host == "" || (inferenceURL.Scheme != "http" && inferenceURL.Scheme != "https") { + return errors.New("admin.inference_public_url must resolve from an environment variable to an absolute http(s) URL") + } if cfg.Admin.Mail.Enabled { if cfg.Admin.Mail.FromAddress == "" || cfg.Admin.Mail.SMTPAddress == "" { return errors.New("admin.mail requires SMTP address and from address environment variables") @@ -660,16 +698,30 @@ func Validate(cfg Config) error { } providers := make(map[string]ProviderConfig, len(cfg.Providers)) + providerSlugs := make(map[string]string, len(cfg.Providers)) for _, provider := range cfg.Providers { if provider.ID == "" { return errors.New("provider id is required") } + if !validProviderSlug(provider.Slug) { + return fmt.Errorf("provider %q: slug must be 3-64 lowercase letters, numbers, or hyphens", provider.ID) + } + if existingID, exists := providerSlugs[provider.Slug]; exists { + return fmt.Errorf("provider %q: duplicate public slug already used by provider %q", provider.ID, existingID) + } + providerSlugs[provider.Slug] = provider.ID if _, exists := providers[provider.ID]; exists { return fmt.Errorf("duplicate provider id %q", provider.ID) } if provider.Protocol != domain.ProtocolOpenAI && provider.Protocol != domain.ProtocolAnthropic { return fmt.Errorf("provider %q: unsupported protocol %q", provider.ID, provider.Protocol) } + if provider.Protocol == domain.ProtocolOpenAI && provider.WireAPI != "chat_completions" && provider.WireAPI != "responses" { + return fmt.Errorf("provider %q: wire_api must be chat_completions or responses", provider.ID) + } + if provider.Protocol == domain.ProtocolAnthropic && provider.WireAPI != "messages" { + return fmt.Errorf("provider %q: wire_api must be messages", provider.ID) + } parsed, err := url.Parse(provider.BaseURL) if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { return fmt.Errorf("provider %q: base_url must be an absolute http(s) URL", provider.ID) @@ -719,6 +771,18 @@ func Validate(cfg Config) error { return nil } +func validProviderSlug(value string) bool { + if len(value) < 3 || len(value) > 64 || value[0] == '-' || value[len(value)-1] == '-' { + return false + } + for _, character := range value { + if (character < 'a' || character > 'z') && (character < '0' || character > '9') && character != '-' { + return false + } + } + return true +} + func (c ServerConfig) ReadHeaderTimeout() time.Duration { return time.Duration(c.ReadHeaderTimeoutSecs) * time.Second } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 8e2e9bb..bb3e814 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -27,11 +27,59 @@ func TestLoadAppliesDefaultsAndResolvesSecrets(t *testing.T) { if cfg.Providers[0].BaseURL != "https://example.com/v1" { t.Fatal("provider URL was not resolved") } + if cfg.Providers[0].WireAPI != "chat_completions" { + t.Fatalf("default wire API = %q", cfg.Providers[0].WireAPI) + } + if cfg.Providers[0].Slug != "primary" { + t.Fatalf("default provider slug = %q", cfg.Providers[0].Slug) + } if cfg.Models[0].Routes[0].Weight != 1 { t.Fatalf("expected default route weight 1, got %d", cfg.Models[0].Routes[0].Weight) } } +func TestLoadValidatesPublicProviderSlugs(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com") + for _, providers := range []string{ + `[{"id":"primary","slug":"Not Valid","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}]`, + `[{"id":"one","slug":"shared-provider","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"},{"id":"two","slug":"shared-provider","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}]`, + } { + path := writeConfig(t, `{"providers":`+providers+`,"models":[{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}]}`) + if _, err := Load(path); err == nil { + t.Fatalf("expected invalid provider slugs to be rejected: %s", providers) + } + } +} + +func TestLoadAcceptsOpenAIResponsesWireAPI(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com") + path := writeConfig(t, `{ + "providers": [{"id":"responses","protocol":"openai","wire_api":"responses","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], + "models": [{"id":"example/model","routes":[{"provider":"responses","upstream_model":"gpt-example"}]}] +}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if cfg.Providers[0].WireAPI != "responses" { + t.Fatalf("wire API = %q", cfg.Providers[0].WireAPI) + } +} + +func TestLoadRejectsIncompatibleWireAPI(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com") + path := writeConfig(t, `{ + "providers": [{"id":"bad","protocol":"anthropic","wire_api":"responses","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], + "models": [{"id":"example/model","routes":[{"provider":"bad","upstream_model":"model"}]}] +}`) + if _, err := Load(path); err == nil { + t.Fatal("expected incompatible wire API to be rejected") + } +} + func TestLoadRejectsUnknownFieldsAndTrailingData(t *testing.T) { t.Setenv("TEST_UPSTREAM_KEY", "secret") unknown := writeConfig(t, `{"unknown":true}`) @@ -63,6 +111,7 @@ func TestLoadControlPlaneModeWithoutStaticRoutes(t *testing.T) { t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret") t.Setenv("AIGW_PUBLIC_URL", "http://localhost:8080/admin/") + t.Setenv("AIGW_INFERENCE_PUBLIC_URL", "https://api.example.test") path := writeConfig(t, `{ "control_plane": {"enabled":true}, "admin": {"enabled":true} @@ -78,6 +127,24 @@ func TestLoadControlPlaneModeWithoutStaticRoutes(t *testing.T) { if len(cfg.Providers) != 0 || len(cfg.Models) != 0 { t.Fatal("control-plane mode unexpectedly requires static providers or models") } + if cfg.Admin.InferencePublicURL != "https://api.example.test" { + t.Fatalf("inference public URL = %q", cfg.Admin.InferencePublicURL) + } +} + +func TestLoadDerivesInferenceURLFromConsoleOrigin(t *testing.T) { + t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw") + t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") + t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret") + t.Setenv("AIGW_PUBLIC_URL", "https://console.example.test/admin/") + path := writeConfig(t, `{"control_plane":{"enabled":true},"admin":{"enabled":true}}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if cfg.Admin.InferencePublicURL != "https://console.example.test" { + t.Fatalf("derived inference URL = %q", cfg.Admin.InferencePublicURL) + } } func TestLoadControlPlaneModeWithoutRedis(t *testing.T) { @@ -128,6 +195,25 @@ func TestLoadResolvesStripeSecrets(t *testing.T) { } } +func TestLoadCanDisableStripeFromEnvironmentWithoutCredentials(t *testing.T) { + t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw") + t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") + t.Setenv("AIGW_SETTLEMENT_SPOOL_PATH", filepath.Join(t.TempDir(), "settlements.jsonl")) + t.Setenv("TEST_STRIPE_ENABLED", "false") + path := writeConfig(t, `{ + "control_plane":{"enabled":true}, + "billing":{"enabled":true,"currency":"usd","default_max_output_tokens":1024,"min_top_up_minor":500,"max_top_up_minor":1000000, + "stripe":{"enabled_env":"TEST_STRIPE_ENABLED"}} +}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if cfg.Billing.Stripe.Enabled || cfg.Billing.Stripe.APIKey != "" { + t.Fatalf("Stripe should be disabled without credentials: %+v", cfg.Billing.Stripe) + } +} + func TestLoadRejectsEmptyExternalServiceEnvironment(t *testing.T) { t.Setenv("TEST_UPSTREAM_KEY", "secret") t.Setenv("TEST_UPSTREAM_URL", "") @@ -165,6 +251,13 @@ func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) { if staticConfig.Providers[0].BaseURL != "https://openai.example.test/v1" { t.Fatalf("static provider URL = %q", staticConfig.Providers[0].BaseURL) } + responsesConfig, err := Load(filepath.Join("..", "..", "config.responses.example.json")) + if err != nil { + t.Fatalf("load Responses example: %v", err) + } + if len(responsesConfig.Providers) != 1 || responsesConfig.Providers[0].WireAPI != "responses" { + t.Fatalf("Responses example provider = %+v", responsesConfig.Providers) + } t.Setenv("AIGW_DATABASE_URL", "postgres://example:secret@postgres.example.test/aigw") t.Setenv("AIGW_REDIS_URL", "redis://redis.example.test:6379/0") @@ -178,6 +271,7 @@ func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) { t.Setenv("AIGW_WEBAUTHN_RP_ID", "console.example.test") t.Setenv("AIGW_WEBAUTHN_ORIGINS", "https://console.example.test") t.Setenv("AIGW_STRIPE_API_KEY", "rk_test_example") + t.Setenv("AIGW_STRIPE_ENABLED", "true") 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") diff --git a/internal/controlplane/access.go b/internal/controlplane/access.go index 47e9a8f..1b04db6 100644 --- a/internal/controlplane/access.go +++ b/internal/controlplane/access.go @@ -43,23 +43,24 @@ func (a ConsoleActor) Can(permission string) bool { case RoleTenantAdmin: switch permission { case "overview.read", "tenants.read", "projects.read", "projects.write", "keys.read", "keys.write", - "billing.read", "billing.topup", "usage.read", "audit.read", "limits.read", "limits.write", "users.read", "users.write": + "billing.read", "billing.topup", "usage.read", "audit.read", "limits.read", "limits.write", "users.read", "users.write", + "preferences.read", "developer.preferences.write", "billing.preferences.write": return true } return false case RoleTenantBilling: - return permission == "overview.read" || permission == "billing.read" || permission == "billing.topup" || permission == "usage.read" || permission == "audit.read" + return permission == "overview.read" || permission == "billing.read" || permission == "billing.topup" || permission == "usage.read" || permission == "audit.read" || permission == "preferences.read" || permission == "billing.preferences.write" case RoleTenantDeveloper: - return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "keys.write" || permission == "usage.read" || permission == "limits.read" + return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "keys.write" || permission == "usage.read" || permission == "limits.read" || permission == "preferences.read" || permission == "developer.preferences.write" case RoleTenantViewer: - return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "billing.read" || permission == "usage.read" || permission == "limits.read" || permission == "audit.read" + return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "billing.read" || permission == "usage.read" || permission == "limits.read" || permission == "audit.read" || permission == "preferences.read" default: return false } } func (a ConsoleActor) Permissions() []string { - all := []string{"overview.read", "tenants.read", "tenants.write", "projects.read", "projects.write", "keys.read", "keys.write", "platform.read", "platform.write", "billing.read", "billing.topup", "billing.adjust", "usage.read", "limits.read", "limits.write", "users.read", "users.write", "audit.read"} + all := []string{"overview.read", "preferences.read", "developer.preferences.write", "billing.preferences.write", "tenants.read", "tenants.write", "projects.read", "projects.write", "keys.read", "keys.write", "platform.read", "platform.write", "billing.read", "billing.topup", "billing.adjust", "usage.read", "limits.read", "limits.write", "users.read", "users.write", "audit.read"} result := make([]string, 0, len(all)) for _, permission := range all { if a.Can(permission) { diff --git a/internal/controlplane/access_test.go b/internal/controlplane/access_test.go index 7767e0d..d871024 100644 --- a/internal/controlplane/access_test.go +++ b/internal/controlplane/access_test.go @@ -13,8 +13,12 @@ func TestConsoleRolePermissions(t *testing.T) { {RoleTenantAdmin, "keys.write", true}, {RoleTenantAdmin, "limits.write", true}, {RoleTenantBilling, "billing.topup", true}, + {RoleTenantBilling, "billing.preferences.write", true}, + {RoleTenantBilling, "developer.preferences.write", false}, {RoleTenantBilling, "keys.read", false}, {RoleTenantDeveloper, "keys.write", true}, + {RoleTenantDeveloper, "developer.preferences.write", true}, + {RoleTenantDeveloper, "billing.preferences.write", false}, {RoleTenantDeveloper, "billing.read", false}, {RoleTenantViewer, "usage.read", true}, {RoleTenantViewer, "users.read", false}, diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go index aaaf2c5..d42a328 100644 --- a/internal/controlplane/mail_operations.go +++ b/internal/controlplane/mail_operations.go @@ -148,23 +148,26 @@ func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotifi (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) + COALESCE(spend.today,0),COALESCE(spend.baseline,0), + COALESCE(pref.low_balance_enabled,TRUE),COALESCE(pref.low_balance_threshold_micros,$1) 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 + LEFT JOIN tenant_preferences pref ON pref.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))`) + AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`, config.LowBalanceMicros) 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, ¤cy, &available, &email, &name, &today, &baseline); err != nil { + var available, today, baseline, lowBalanceThreshold int64 + var lowBalanceEnabled bool + if err := rows.Scan(&tenantID, ¤cy, &available, &email, &name, &today, &baseline, &lowBalanceEnabled, &lowBalanceThreshold); err != nil { return err } day := time.Now().UTC().Format("2006-01-02") - if available <= config.LowBalanceMicros { + if lowBalanceEnabled && available <= lowBalanceThreshold { 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 diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go index c2cb7d8..9f81d6e 100644 --- a/internal/controlplane/mutations.go +++ b/internal/controlplane/mutations.go @@ -17,8 +17,9 @@ import ( ) var ( - ErrNotFound = errors.New("control-plane resource not found") - slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`) + ErrNotFound = errors.New("control-plane resource not found") + slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`) + nonSlugCharacters = regexp.MustCompile(`[^a-z0-9]+`) ) func (s *Store) CreateTenant(ctx context.Context, input CreateTenantInput) (Tenant, int64, error) { @@ -84,11 +85,28 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea if input.TenantID == "" || input.ProjectID == "" || input.Name == "" { return CreatedAPIKey{}, 0, errors.New("API key requires tenant_id, project_id, and name") } + if len(input.Name) > 120 || input.MonthlySpendMicros < 0 { + return CreatedAPIKey{}, 0, errors.New("API key name or monthly spend limit is invalid") + } + if input.ExpiresAt != nil && !input.ExpiresAt.After(time.Now()) { + return CreatedAPIKey{}, 0, errors.New("API key expiry must be in the future") + } if len(input.Scopes) == 0 { input.Scopes = []string{"inference"} } scopes := uniqueStrings(input.Scopes) + tags := uniqueStrings(input.Tags) + allowedModels := uniqueStrings(input.AllowedModels) + if len(scopes) > 20 || len(tags) > 20 || len(allowedModels) > 200 { + return CreatedAPIKey{}, 0, errors.New("API key has too many scopes, tags, or model restrictions") + } + for _, value := range append(append(append([]string{}, scopes...), tags...), allowedModels...) { + if len(value) > 160 { + return CreatedAPIKey{}, 0, errors.New("API key scope, tag, or model ID is too long") + } + } scopesJSON, _ := json.Marshal(scopes) + tagsJSON, _ := json.Marshal(tags) random := make([]byte, 32) if _, err := rand.Read(random); err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("generate API key: %w", err) @@ -104,15 +122,32 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea defer tx.Rollback(ctx) var result CreatedAPIKey err = tx.QueryRow(ctx, ` - INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_hash, scopes) - VALUES ($1, $2, $3, $4, $5, $6) - RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at`, - input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON, - ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &scopesJSON, &result.Status, &result.CreatedAt) + INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_hash, scopes, tags, monthly_spend_micros, expires_at) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) + RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, tags, + monthly_spend_micros, status, expires_at, last_used_at, created_at`, + input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON, tagsJSON, + input.MonthlySpendMicros, input.ExpiresAt, + ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, + &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.Status, &result.ExpiresAt, + &result.LastUsedAt, &result.CreatedAt) if err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err) } + if len(allowedModels) > 0 { + command, err := tx.Exec(ctx, ` + INSERT INTO api_key_model_restrictions (api_key_id, model_id) + SELECT $1, id FROM models WHERE public_id = ANY($2::text[])`, result.ID, allowedModels) + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("restrict API key models: %w", err) + } + if command.RowsAffected() != int64(len(allowedModels)) { + return CreatedAPIKey{}, 0, errors.New("one or more allowed model IDs do not exist") + } + } result.Scopes = scopes + result.Tags = tags + result.AllowedModels = allowedModels result.Key = rawKey generation, err := bumpGeneration(ctx, tx) if err != nil { @@ -129,10 +164,29 @@ func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) { } func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) { + input.Slug = strings.ToLower(strings.TrimSpace(input.Slug)) input.Name = strings.TrimSpace(input.Name) + if input.Slug == "" { + input.Slug = strings.Trim(nonSlugCharacters.ReplaceAllString(strings.ToLower(input.Name), "-"), "-") + if len(input.Slug) > 64 { + input.Slug = strings.TrimRight(input.Slug[:64], "-") + } + } input.BaseURL = strings.TrimRight(strings.TrimSpace(input.BaseURL), "/") - if input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") { - return Provider{}, 0, errors.New("provider requires name, protocol openai|anthropic, base_url, and api_key") + input.WireAPI = strings.TrimSpace(input.WireAPI) + if input.WireAPI == "" { + if input.Protocol == "anthropic" { + input.WireAPI = "messages" + } else { + input.WireAPI = "chat_completions" + } + } + if !slugPattern.MatchString(input.Slug) || input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") { + return Provider{}, 0, errors.New("provider requires a unique 3-64 character lowercase slug, name, protocol openai|anthropic, base_url, and api_key") + } + if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses") || + (input.Protocol == "anthropic" && input.WireAPI != "messages") { + return Provider{}, 0, errors.New("provider wire_api is incompatible with protocol") } parsed, err := url.Parse(input.BaseURL) if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { @@ -149,11 +203,11 @@ func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) ( defer tx.Rollback(ctx) var result Provider err = tx.QueryRow(ctx, ` - INSERT INTO providers (name, protocol, base_url, api_key_ciphertext) - VALUES ($1, $2, $3, $4) - RETURNING id::text, name, protocol, base_url, enabled, created_at`, - input.Name, input.Protocol, input.BaseURL, ciphertext, - ).Scan(&result.ID, &result.Name, &result.Protocol, &result.BaseURL, &result.Enabled, &result.CreatedAt) + INSERT INTO providers (slug, name, protocol, wire_api, base_url, api_key_ciphertext) + VALUES ($1, $2, $3, $4, $5, $6) + RETURNING id::text, slug, name, protocol, wire_api, base_url, enabled, created_at`, + input.Slug, input.Name, input.Protocol, input.WireAPI, input.BaseURL, ciphertext, + ).Scan(&result.ID, &result.Slug, &result.Name, &result.Protocol, &result.WireAPI, &result.BaseURL, &result.Enabled, &result.CreatedAt) if err != nil { return Provider{}, 0, fmt.Errorf("create provider: %w", err) } diff --git a/internal/controlplane/preferences.go b/internal/controlplane/preferences.go new file mode 100644 index 0000000..73bcce6 --- /dev/null +++ b/internal/controlplane/preferences.go @@ -0,0 +1,156 @@ +package controlplane + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/jackc/pgx/v5" +) + +const maxLowBalanceThresholdMicros int64 = 1_000_000_000_000_000 + +func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defaultThresholdMicros int64) (TenantPreferences, error) { + if defaultThresholdMicros <= 0 { + defaultThresholdMicros = 5_000_000 + } + result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, LowBalanceThresholdMicros: defaultThresholdMicros} + if result.TenantID == "" { + return result, nil + } + var defaultModel, fallbackModel *string + var updatedAt time.Time + err := s.db.QueryRow(ctx, ` + SELECT default_model, fallback_model, low_balance_enabled, low_balance_threshold_micros, updated_at + FROM tenant_preferences WHERE tenant_id=$1`, result.TenantID).Scan( + &defaultModel, &fallbackModel, &result.LowBalanceEnabled, &result.LowBalanceThresholdMicros, &updatedAt) + if errors.Is(err, pgx.ErrNoRows) { + return result, nil + } + if err != nil { + return TenantPreferences{}, fmt.Errorf("query tenant preferences: %w", err) + } + if defaultModel != nil { + result.DefaultModel = *defaultModel + } + if fallbackModel != nil { + result.FallbackModel = *fallbackModel + } + result.UpdatedAt = &updatedAt + return result, nil +} + +func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput) (TenantPreferences, error) { + input.TenantID = strings.TrimSpace(input.TenantID) + input.DefaultModel = strings.TrimSpace(input.DefaultModel) + input.FallbackModel = strings.TrimSpace(input.FallbackModel) + if input.TenantID == "" { + return TenantPreferences{}, errors.New("tenant_id is required") + } + if input.DefaultModel != "" && input.DefaultModel == input.FallbackModel { + return TenantPreferences{}, errors.New("default_model and fallback_model must be different") + } + available, err := s.ListDeveloperModels(ctx, input.TenantID) + if err != nil { + return TenantPreferences{}, err + } + allowed := make(map[string]struct{}, len(available)) + for _, model := range available { + allowed[model.PublicID] = struct{}{} + } + for field, model := range map[string]string{"default_model": input.DefaultModel, "fallback_model": input.FallbackModel} { + if model != "" { + if _, ok := allowed[model]; !ok { + return TenantPreferences{}, fmt.Errorf("%s is not available to this tenant", field) + } + } + } + tx, err := s.db.Begin(ctx) + if err != nil { + return TenantPreferences{}, err + } + defer tx.Rollback(ctx) + if err := tx.QueryRow(ctx, `SELECT id FROM tenants WHERE id=$1`, input.TenantID).Scan(new(string)); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return TenantPreferences{}, ErrNotFound + } + return TenantPreferences{}, err + } + var result TenantPreferences + var defaultModel, fallbackModel *string + var updatedAt time.Time + if err := tx.QueryRow(ctx, ` + INSERT INTO tenant_preferences (tenant_id, default_model, fallback_model) + VALUES ($1, NULLIF($2,''), NULLIF($3,'')) + ON CONFLICT (tenant_id) DO UPDATE SET default_model=EXCLUDED.default_model, + fallback_model=EXCLUDED.fallback_model, updated_at=now() + RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled, + low_balance_threshold_micros, updated_at`, input.TenantID, input.DefaultModel, input.FallbackModel).Scan( + &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled, + &result.LowBalanceThresholdMicros, &updatedAt); err != nil { + return TenantPreferences{}, fmt.Errorf("save developer preferences: %w", err) + } + if defaultModel != nil { + result.DefaultModel = *defaultModel + } + if fallbackModel != nil { + result.FallbackModel = *fallbackModel + } + result.UpdatedAt = &updatedAt + if err := tx.Commit(ctx); err != nil { + return TenantPreferences{}, err + } + return result, nil +} + +func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaultThresholdMicros int64) (TenantPreferences, error) { + input.TenantID = strings.TrimSpace(input.TenantID) + if input.TenantID == "" { + return TenantPreferences{}, errors.New("tenant_id is required") + } + if defaultThresholdMicros <= 0 { + defaultThresholdMicros = 5_000_000 + } + if input.LowBalanceThresholdMicros != nil && (*input.LowBalanceThresholdMicros < 0 || *input.LowBalanceThresholdMicros > maxLowBalanceThresholdMicros) { + return TenantPreferences{}, errors.New("low balance threshold is outside the supported range") + } + tx, err := s.db.Begin(ctx) + if err != nil { + return TenantPreferences{}, err + } + defer tx.Rollback(ctx) + if err := tx.QueryRow(ctx, `SELECT id FROM tenants WHERE id=$1`, input.TenantID).Scan(new(string)); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return TenantPreferences{}, ErrNotFound + } + return TenantPreferences{}, err + } + var result TenantPreferences + var defaultModel, fallbackModel *string + var updatedAt time.Time + if err := tx.QueryRow(ctx, ` + INSERT INTO tenant_preferences (tenant_id, low_balance_enabled, low_balance_threshold_micros) + VALUES ($1, COALESCE($2::boolean, TRUE), COALESCE($3::bigint, $4::bigint)) + ON CONFLICT (tenant_id) DO UPDATE SET + low_balance_enabled=COALESCE($2::boolean, tenant_preferences.low_balance_enabled), + low_balance_threshold_micros=COALESCE($3::bigint, tenant_preferences.low_balance_threshold_micros), updated_at=now() + RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled, + low_balance_threshold_micros, updated_at`, input.TenantID, input.LowBalanceEnabled, input.LowBalanceThresholdMicros, defaultThresholdMicros).Scan( + &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled, + &result.LowBalanceThresholdMicros, &updatedAt); err != nil { + return TenantPreferences{}, fmt.Errorf("save billing preferences: %w", err) + } + if defaultModel != nil { + result.DefaultModel = *defaultModel + } + if fallbackModel != nil { + result.FallbackModel = *fallbackModel + } + result.UpdatedAt = &updatedAt + if err := tx.Commit(ctx); err != nil { + return TenantPreferences{}, err + } + return result, nil +} diff --git a/internal/controlplane/preferences_test.go b/internal/controlplane/preferences_test.go new file mode 100644 index 0000000..88a3d50 --- /dev/null +++ b/internal/controlplane/preferences_test.go @@ -0,0 +1,31 @@ +package controlplane + +import ( + "context" + "testing" +) + +func TestGetTenantPreferencesWithoutTenantUsesConfiguredDefault(t *testing.T) { + result, err := (&Store{}).GetTenantPreferences(context.Background(), "", 12_500_000) + if err != nil { + t.Fatal(err) + } + if !result.LowBalanceEnabled || result.LowBalanceThresholdMicros != 12_500_000 { + t.Fatalf("unexpected defaults: %+v", result) + } +} + +func TestPreferenceValidationRejectsUnsafeValuesBeforeDatabaseAccess(t *testing.T) { + store := &Store{} + if _, err := store.SetDeveloperPreferences(context.Background(), SetDeveloperPreferencesInput{ + TenantID: "tenant", DefaultModel: "same", FallbackModel: "same", + }); err == nil { + t.Fatal("expected identical default and fallback models to fail") + } + threshold := maxLowBalanceThresholdMicros + 1 + if _, err := store.SetBillingPreferences(context.Background(), SetBillingPreferencesInput{ + TenantID: "tenant", LowBalanceThresholdMicros: &threshold, + }, 5_000_000); err == nil { + t.Fatal("expected excessive low balance threshold to fail") + } +} diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go index 73d869f..48610c7 100644 --- a/internal/controlplane/queries.go +++ b/internal/controlplane/queries.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "time" ) func (s *Store) Overview(ctx context.Context) (Overview, error) { @@ -84,15 +85,35 @@ func (s *Store) ListAPIKeys(ctx context.Context) ([]APIKey, error) { } func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, error) { + periodStart := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + periodEnd := periodStart.AddDate(0, 1, 0) query := ` - SELECT id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at - FROM api_keys` - args := []any{} + SELECT k.id::text, k.tenant_id::text, k.project_id::text, k.name, k.key_prefix, k.scopes, + k.tags, k.monthly_spend_micros, k.status, k.expires_at, k.last_used_at, k.created_at, + usage.month_spend, usage.month_requests, pending.month_reserved, + COALESCE(( + SELECT jsonb_agg(m.public_id ORDER BY m.public_id) + FROM api_key_model_restrictions r + JOIN models m ON m.id = r.model_id + WHERE r.api_key_id = k.id + ), '[]'::jsonb) + FROM api_keys k + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(u.cost_micros), 0)::bigint AS month_spend, COUNT(*)::bigint AS month_requests + FROM usage_events u WHERE u.key_id = k.id AND u.started_at >= $1 AND u.started_at < $2 + ) usage + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(b.reserved_micros), 0)::bigint AS month_reserved + FROM billing_reservations b + WHERE b.key_id = k.id AND b.status IN ('pending', 'metering_failed') + AND b.created_at >= $1 AND b.created_at < $2 + ) pending` + args := []any{periodStart, periodEnd} if tenantID != "" { - query += ` WHERE tenant_id=$1` + query += ` WHERE k.tenant_id=$3` args = append(args, tenantID) } - query += ` ORDER BY created_at DESC` + query += ` ORDER BY k.created_at DESC` rows, err := s.db.Query(ctx, query, args...) if err != nil { return nil, fmt.Errorf("query API keys: %w", err) @@ -101,13 +122,22 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, result := make([]APIKey, 0) for rows.Next() { var item APIKey - var scopesJSON []byte - if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, &scopesJSON, &item.Status, &item.CreatedAt); err != nil { + var scopesJSON, tagsJSON, allowedModelsJSON []byte + if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, + &scopesJSON, &tagsJSON, &item.MonthlySpendMicros, &item.Status, &item.ExpiresAt, + &item.LastUsedAt, &item.CreatedAt, &item.CurrentMonthSpendMicros, &item.CurrentMonthRequests, + &item.CurrentMonthReservedMicros, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if err := json.Unmarshal(scopesJSON, &item.Scopes); err != nil { return nil, fmt.Errorf("decode API key scopes: %w", err) } + if err := json.Unmarshal(tagsJSON, &item.Tags); err != nil { + return nil, fmt.Errorf("decode API key tags: %w", err) + } + if err := json.Unmarshal(allowedModelsJSON, &item.AllowedModels); err != nil { + return nil, fmt.Errorf("decode API key model restrictions: %w", err) + } result = append(result, item) } return result, rows.Err() @@ -153,7 +183,7 @@ func (s *Store) ResourceTenantID(ctx context.Context, resource, id string) (stri func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) { rows, err := s.db.Query(ctx, ` - SELECT p.id::text, p.name, p.protocol, p.base_url, p.enabled, count(r.id), p.created_at + SELECT p.id::text, p.slug, p.name, p.protocol, p.wire_api, p.base_url, p.enabled, count(r.id), p.created_at FROM providers p LEFT JOIN model_routes r ON r.provider_id = p.id GROUP BY p.id ORDER BY p.created_at DESC`) if err != nil { @@ -163,7 +193,7 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) { result := make([]Provider, 0) for rows.Next() { var item Provider - if err := rows.Scan(&item.ID, &item.Name, &item.Protocol, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil { + if err := rows.Scan(&item.ID, &item.Slug, &item.Name, &item.Protocol, &item.WireAPI, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil { return nil, fmt.Errorf("scan provider: %w", err) } result = append(result, item) @@ -182,6 +212,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { m.enabled, m.created_at 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 ORDER BY m.public_id`) if err != nil { @@ -264,7 +295,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { keyRows.Close() routeRows, err := s.db.Query(ctx, ` - SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol, + SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol, p.wire_api, p.enabled, r.upstream_model, r.priority, r.weight, r.enabled FROM model_routes r JOIN providers p ON p.id = r.provider_id ORDER BY r.priority, r.created_at`) @@ -275,7 +306,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { for routeRows.Next() { var route Route var modelID string - if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil { + if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.WireAPI, &route.ProviderEnabled, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil { return nil, fmt.Errorf("scan model route: %w", err) } if position, ok := positions[modelID]; ok { @@ -284,3 +315,141 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { } return models, routeRows.Err() } + +func (s *Store) ListDeveloperModels(ctx context.Context, tenantID string) ([]DeveloperModel, error) { + models, err := s.ListModels(ctx) + if err != nil { + return nil, err + } + keyIDs := map[string]struct{}{} + if tenantID != "" { + rows, err := s.db.Query(ctx, `SELECT id::text FROM api_keys WHERE tenant_id=$1 AND status='active'`, tenantID) + if err != nil { + return nil, fmt.Errorf("query developer API keys: %w", err) + } + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + rows.Close() + return nil, err + } + keyIDs[id] = struct{}{} + } + if err := rows.Err(); err != nil { + rows.Close() + return nil, err + } + rows.Close() + } + return developerModelsFor(models, tenantID, keyIDs), nil +} + +func (s *Store) ListPublicModels(ctx context.Context) ([]PublicModel, error) { + models, err := s.ListModels(ctx) + if err != nil { + return nil, err + } + return publicModelsFor(models), nil +} + +func publicModelsFor(models []Model) []PublicModel { + result := make([]PublicModel, 0, len(models)) + for _, model := range models { + if !model.Enabled || model.Lifecycle == "retired" || len(model.AllowedTenantIDs) != 0 || len(model.AllowedKeyIDs) != 0 { + continue + } + wireSet := make(map[string]struct{}) + providerSet := make(map[string]struct{}) + wireAPIs := make([]string, 0, len(model.Routes)) + for _, route := range model.Routes { + if !route.Enabled || !route.ProviderEnabled { + continue + } + wireAPI := route.WireAPI + if wireAPI == "" && route.Protocol == "anthropic" { + wireAPI = "messages" + } else if wireAPI == "" { + wireAPI = "chat_completions" + } + if _, exists := wireSet[wireAPI]; !exists { + wireSet[wireAPI] = struct{}{} + wireAPIs = append(wireAPIs, wireAPI) + } + providerSet[route.ProviderID] = struct{}{} + } + if len(wireAPIs) == 0 { + continue + } + result = append(result, PublicModel{PublicID: model.PublicID, DisplayName: model.DisplayName, + Description: model.Description, OwnedBy: model.OwnedBy, InputModalities: model.InputModalities, + OutputModalities: model.OutputModalities, ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens, + Capabilities: model.Capabilities, Regions: model.Regions, Lifecycle: model.Lifecycle, ReleasedAt: model.ReleasedAt, + ReplacementModel: model.ReplacementModel, Aliases: model.Aliases, PriceCurrency: model.PriceCurrency, + InputPriceMicrosPerMillion: model.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: model.OutputPriceMicrosPerMillion, + CacheReadPriceMicrosPerMillion: model.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: model.CacheWritePriceMicrosPerMillion, + SupportedWireAPIs: wireAPIs, ProviderCount: len(providerSet), AvailableProviderCount: len(providerSet), HealthStatus: "available"}) + } + return result +} + +func developerModelsFor(models []Model, tenantID string, keyIDs map[string]struct{}) []DeveloperModel { + result := make([]DeveloperModel, 0, len(models)) + for _, model := range models { + if !model.Enabled || model.Lifecycle == "retired" || !stringAllowed(model.AllowedTenantIDs, tenantID) || !keyAllowed(model.AllowedKeyIDs, keyIDs) { + continue + } + wireSet := map[string]struct{}{} + wireAPIs := make([]string, 0, len(model.Routes)) + for _, route := range model.Routes { + if !route.Enabled || !route.ProviderEnabled { + continue + } + wireAPI := route.WireAPI + if wireAPI == "" && route.Protocol == "anthropic" { + wireAPI = "messages" + } else if wireAPI == "" { + wireAPI = "chat_completions" + } + if _, exists := wireSet[wireAPI]; !exists { + wireSet[wireAPI] = struct{}{} + wireAPIs = append(wireAPIs, wireAPI) + } + } + if len(wireAPIs) == 0 { + continue + } + result = append(result, DeveloperModel{ID: model.ID, PublicID: model.PublicID, DisplayName: model.DisplayName, + Description: model.Description, OwnedBy: model.OwnedBy, InputModalities: model.InputModalities, + OutputModalities: model.OutputModalities, ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens, + Capabilities: model.Capabilities, Regions: model.Regions, Lifecycle: model.Lifecycle, ReleasedAt: model.ReleasedAt, + ReplacementModel: model.ReplacementModel, Aliases: model.Aliases, PriceCurrency: model.PriceCurrency, + InputPriceMicrosPerMillion: model.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: model.OutputPriceMicrosPerMillion, + CacheReadPriceMicrosPerMillion: model.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: model.CacheWritePriceMicrosPerMillion, + SupportedWireAPIs: wireAPIs}) + } + return result +} + +func stringAllowed(allowed []string, value string) bool { + if len(allowed) == 0 || value == "" { + return true + } + for _, item := range allowed { + if item == value { + return true + } + } + return false +} + +func keyAllowed(allowed []string, keys map[string]struct{}) bool { + if len(allowed) == 0 { + return true + } + for _, id := range allowed { + if _, ok := keys[id]; ok { + return true + } + } + return false +} diff --git a/internal/controlplane/queries_test.go b/internal/controlplane/queries_test.go new file mode 100644 index 0000000..4cd44b9 --- /dev/null +++ b/internal/controlplane/queries_test.go @@ -0,0 +1,80 @@ +package controlplane + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestDeveloperModelsForFiltersScopeAndRedactsRouting(t *testing.T) { + models := []Model{ + {ID: "public", PublicID: "acme/public", DisplayName: "Public", Enabled: true, Lifecycle: "active", + AllowedTenantIDs: nil, Routes: []Route{{Protocol: "openai", WireAPI: "responses", Enabled: true, ProviderEnabled: true, ProviderName: "secret-provider", UpstreamModel: "secret-model", Priority: 1, Weight: 100}}}, + {ID: "tenant", PublicID: "acme/private", Enabled: true, Lifecycle: "active", AllowedTenantIDs: []string{"tenant-a"}, + Routes: []Route{{Protocol: "anthropic", WireAPI: "messages", Enabled: true, ProviderEnabled: true}}}, + {ID: "key", PublicID: "acme/key", Enabled: true, Lifecycle: "active", AllowedKeyIDs: []string{"key-a"}, + Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}}, + {ID: "retired", PublicID: "acme/retired", Enabled: true, Lifecycle: "retired", Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}}, + {ID: "disabled", PublicID: "acme/disabled", Enabled: false, Lifecycle: "active", Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}}, + {ID: "noroute", PublicID: "acme/noroute", Enabled: true, Lifecycle: "active", Routes: []Route{{Protocol: "openai", Enabled: false, ProviderEnabled: true}}}, + {ID: "provider-off", PublicID: "acme/provider-off", Enabled: true, Lifecycle: "active", Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: false}}}, + } + got := developerModelsFor(models, "tenant-a", map[string]struct{}{"key-a": {}}) + if len(got) != 3 { + t.Fatalf("developer model count = %d, want 3", len(got)) + } + if got[0].PublicID != "acme/key" && got[1].PublicID != "acme/key" && got[2].PublicID != "acme/key" { + t.Fatal("key-allowlisted model was not included for an active tenant key") + } + for _, item := range got { + if len(item.SupportedWireAPIs) == 0 { + t.Fatalf("unexpected developer model: %+v", item) + } + if item.PublicID == "acme/public" && item.SupportedWireAPIs[0] != "responses" { + t.Fatalf("responses wire API was not preserved: %+v", item) + } + } + for _, item := range got { + if item.PublicID == "acme/private" && item.SupportedWireAPIs[0] != "messages" { + t.Fatalf("anthropic wire API was not preserved: %+v", item) + } + } + encoded, err := json.Marshal(got) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(encoded), "secret-provider") || strings.Contains(string(encoded), "secret-model") || strings.Contains(string(encoded), "provider_id") { + t.Fatalf("developer model response leaked routing metadata: %s", encoded) + } +} + +func TestDeveloperModelsForDoesNotExposeKeyRestrictedModelsWithoutTenantKey(t *testing.T) { + models := []Model{{ID: "key", PublicID: "key-only", Enabled: true, Lifecycle: "active", AllowedKeyIDs: []string{"key-a"}, Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}}} + if got := developerModelsFor(models, "tenant-a", map[string]struct{}{}); len(got) != 0 { + t.Fatalf("key-restricted models visible without a matching key: %+v", got) + } +} + +func TestPublicModelsForOnlyExposesUnrestrictedCatalogData(t *testing.T) { + models := []Model{ + {ID: "public", PublicID: "acme/public", DisplayName: "Public", Enabled: true, Lifecycle: "active", + Routes: []Route{{ProviderID: "provider-a", ProviderName: "internal provider", Protocol: "openai", WireAPI: "responses", UpstreamModel: "secret-model", Enabled: true, ProviderEnabled: true}}}, + {ID: "tenant", PublicID: "acme/tenant", Enabled: true, Lifecycle: "active", AllowedTenantIDs: []string{"tenant-a"}, + Routes: []Route{{ProviderID: "provider-a", Protocol: "openai", Enabled: true, ProviderEnabled: true}}}, + {ID: "key", PublicID: "acme/key", Enabled: true, Lifecycle: "active", AllowedKeyIDs: []string{"key-a"}, + Routes: []Route{{ProviderID: "provider-a", Protocol: "openai", Enabled: true, ProviderEnabled: true}}}, + } + got := publicModelsFor(models) + if len(got) != 1 || got[0].PublicID != "acme/public" || got[0].ProviderCount != 1 || got[0].SupportedWireAPIs[0] != "responses" { + t.Fatalf("unexpected public catalog: %+v", got) + } + encoded, err := json.Marshal(got) + if err != nil { + t.Fatal(err) + } + for _, forbidden := range []string{"secret-model", "internal provider", "provider_id", "upstream_model", "allowed_tenant"} { + if strings.Contains(string(encoded), forbidden) { + t.Fatalf("public catalog leaked %q: %s", forbidden, encoded) + } + } +} diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql index 88db49b..ef1ccdd 100644 --- a/internal/controlplane/schema.sql +++ b/internal/controlplane/schema.sql @@ -49,17 +49,53 @@ CREATE TABLE IF NOT EXISTS api_keys ( revoked_at TIMESTAMPTZ, FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE ); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS monthly_spend_micros BIGINT NOT NULL DEFAULT 0 CHECK (monthly_spend_micros >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS expires_at TIMESTAMPTZ; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS tags JSONB NOT NULL DEFAULT '[]'::jsonb; CREATE TABLE IF NOT EXISTS providers ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + slug TEXT NOT NULL UNIQUE CHECK (slug ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$'), name TEXT NOT NULL UNIQUE, protocol TEXT NOT NULL CHECK (protocol IN ('openai', 'anthropic')), + wire_api TEXT NOT NULL DEFAULT 'chat_completions' CHECK (wire_api IN ('chat_completions', 'responses', 'messages')), base_url TEXT NOT NULL, api_key_ciphertext BYTEA NOT NULL, enabled BOOLEAN NOT NULL DEFAULT TRUE, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), updated_at TIMESTAMPTZ NOT NULL DEFAULT now() ); +ALTER TABLE providers ADD COLUMN IF NOT EXISTS slug TEXT; +WITH normalized AS ( + SELECT id, + left(trim(both '-' from regexp_replace(lower(name), '[^a-z0-9]+', '-', 'g')), 64) AS base + FROM providers +), ranked AS ( + SELECT id, base, count(*) OVER (PARTITION BY base) AS base_count + FROM normalized +) +UPDATE providers p +SET slug = CASE + WHEN length(r.base) BETWEEN 3 AND 64 + AND r.base ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$' + AND r.base_count = 1 THEN r.base + ELSE 'provider-' || left(replace(p.id::text, '-', ''), 12) +END +FROM ranked r +WHERE p.id = r.id AND (p.slug IS NULL OR p.slug = ''); +ALTER TABLE providers ALTER COLUMN slug SET NOT NULL; +CREATE UNIQUE INDEX IF NOT EXISTS providers_slug_unique_idx ON providers (slug); +ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_slug_check; +ALTER TABLE providers ADD CONSTRAINT providers_slug_check CHECK (slug ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$'); +ALTER TABLE providers ADD COLUMN IF NOT EXISTS wire_api TEXT NOT NULL DEFAULT 'chat_completions'; +UPDATE providers SET wire_api = 'messages' WHERE protocol = 'anthropic' AND wire_api = 'chat_completions'; +ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_wire_api_check; +ALTER TABLE providers ADD CONSTRAINT providers_wire_api_check CHECK (wire_api IN ('chat_completions', 'responses', 'messages')); +ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_protocol_wire_api_check; +ALTER TABLE providers ADD CONSTRAINT providers_protocol_wire_api_check CHECK ( + (protocol = 'openai' AND wire_api IN ('chat_completions', 'responses')) OR + (protocol = 'anthropic' AND wire_api = 'messages') +); CREATE TABLE IF NOT EXISTS models ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -145,6 +181,15 @@ CREATE TABLE IF NOT EXISTS api_key_model_allowlist ( PRIMARY KEY (api_key_id, model_id) ); +-- Per-key restrictions are separate from the platform model allowlist above: +-- an empty set means that the key may use every model visible to its tenant. +CREATE TABLE IF NOT EXISTS api_key_model_restrictions ( + 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(), model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE, @@ -171,6 +216,18 @@ CREATE TABLE IF NOT EXISTS tenant_wallets ( updated_at TIMESTAMPTZ NOT NULL DEFAULT now() ); +-- Tenant-scoped developer and billing preferences. These values are control +-- plane data, but are intentionally not loaded into the inference snapshot. +CREATE TABLE IF NOT EXISTS tenant_preferences ( + tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE, + default_model TEXT REFERENCES models(public_id) ON DELETE SET NULL, + fallback_model TEXT REFERENCES models(public_id) ON DELETE SET NULL, + low_balance_enabled BOOLEAN NOT NULL DEFAULT TRUE, + low_balance_threshold_micros BIGINT NOT NULL DEFAULT 5000000 CHECK (low_balance_threshold_micros >= 0), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + CHECK (default_model IS NULL OR fallback_model IS NULL OR default_model <> fallback_model) +); + CREATE TABLE IF NOT EXISTS billing_reservations ( request_id TEXT PRIMARY KEY, tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE, @@ -286,6 +343,10 @@ CREATE TABLE IF NOT EXISTS topup_orders ( created_at TIMESTAMPTZ NOT NULL DEFAULT now(), paid_at TIMESTAMPTZ ); +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS trigger_type TEXT NOT NULL DEFAULT 'manual'; +ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_trigger_type_check; +ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_trigger_type_check + CHECK (trigger_type IN ('manual','auto')); 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')); @@ -306,6 +367,8 @@ 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 UNIQUE INDEX IF NOT EXISTS topup_orders_auto_pending_idx ON topup_orders (tenant_id) + WHERE trigger_type='auto' AND status='pending'; CREATE TABLE IF NOT EXISTS billing_reconciliation_resolutions ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -325,6 +388,51 @@ CREATE TABLE IF NOT EXISTS stripe_customers ( updated_at TIMESTAMPTZ NOT NULL DEFAULT now() ); +CREATE TABLE IF NOT EXISTS tenant_billing_profiles ( + tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE, + legal_name TEXT NOT NULL, + billing_email TEXT NOT NULL, + address_line1 TEXT NOT NULL, + address_line2 TEXT NOT NULL DEFAULT '', + city TEXT NOT NULL, + region TEXT NOT NULL DEFAULT '', + postal_code TEXT NOT NULL, + country TEXT NOT NULL CHECK (country ~ '^[A-Z]{2}$'), + stripe_sync_status TEXT NOT NULL DEFAULT 'pending' + CHECK (stripe_sync_status IN ('pending','synced','failed','disabled')), + stripe_synced_at TIMESTAMPTZ, + stripe_sync_error TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS tenant_auto_topup_settings ( + tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE, + enabled BOOLEAN NOT NULL DEFAULT FALSE, + threshold_micros BIGINT NOT NULL CHECK (threshold_micros >= 0), + topup_amount_minor BIGINT NOT NULL CHECK (topup_amount_minor > 0), + stripe_payment_method_id TEXT, + payment_method_type TEXT NOT NULL DEFAULT '', + payment_method_brand TEXT NOT NULL DEFAULT '', + payment_method_last4 TEXT NOT NULL DEFAULT '', + payment_method_exp_month INTEGER NOT NULL DEFAULT 0 CHECK (payment_method_exp_month BETWEEN 0 AND 12), + payment_method_exp_year INTEGER NOT NULL DEFAULT 0 CHECK (payment_method_exp_year >= 0), + stripe_setup_session_id TEXT UNIQUE, + status TEXT NOT NULL DEFAULT 'not_configured' + CHECK (status IN ('not_configured','ready','charging','action_required','failed')), + last_error TEXT NOT NULL DEFAULT '', + failure_count INTEGER NOT NULL DEFAULT 0 CHECK (failure_count >= 0), + last_attempt_at TIMESTAMPTZ, + last_succeeded_at TIMESTAMPTZ, + next_attempt_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + CHECK (stripe_payment_method_id IS NOT NULL OR enabled = FALSE) +); +CREATE INDEX IF NOT EXISTS tenant_auto_topup_ready_idx + ON tenant_auto_topup_settings (next_attempt_at, tenant_id) + WHERE enabled AND stripe_payment_method_id IS NOT NULL; + 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, @@ -658,8 +766,10 @@ CREATE INDEX IF NOT EXISTS billing_ledger_tenant_idx ON billing_ledger (tenant_i CREATE INDEX IF NOT EXISTS usage_events_tenant_idx ON usage_events (tenant_id, created_at DESC); CREATE INDEX IF NOT EXISTS usage_events_project_idx ON usage_events (project_id, created_at DESC); CREATE INDEX IF NOT EXISTS usage_events_model_idx ON usage_events (public_model, created_at DESC); +CREATE INDEX IF NOT EXISTS usage_events_key_idx ON usage_events (key_id, started_at DESC); CREATE INDEX IF NOT EXISTS billing_reservations_pending_idx ON billing_reservations (status, created_at) WHERE status = 'pending'; CREATE INDEX IF NOT EXISTS billing_reservations_project_pending_idx ON billing_reservations (project_id, created_at) WHERE status = 'pending'; +CREATE INDEX IF NOT EXISTS billing_reservations_key_period_idx ON billing_reservations (key_id, created_at DESC) WHERE status IN ('pending', 'metering_failed'); CREATE INDEX IF NOT EXISTS console_users_tenant_idx ON console_users (tenant_id, created_at DESC); CREATE INDEX IF NOT EXISTS console_sessions_user_idx ON console_sessions (user_id, created_at DESC); CREATE INDEX IF NOT EXISTS console_sessions_expiry_idx ON console_sessions (expires_at) WHERE revoked_at IS NULL; diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go index b9bddaf..75a3a80 100644 --- a/internal/controlplane/snapshot.go +++ b/internal/controlplane/snapshot.go @@ -74,7 +74,7 @@ func loadLimitPolicies(ctx context.Context, tx pgx.Tx) ([]domain.LimitPolicy, er func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain.Provider, error) { rows, err := tx.Query(ctx, ` - SELECT id::text, name, protocol, base_url, api_key_ciphertext + SELECT id::text, slug, name, protocol, wire_api, base_url, api_key_ciphertext FROM providers WHERE enabled = TRUE ORDER BY name`) @@ -84,16 +84,16 @@ func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain defer rows.Close() providers := make(map[string]domain.Provider) for rows.Next() { - var id, name, protocol, baseURL string + var id, slug, name, protocol, wireAPI, baseURL string var ciphertext []byte - if err := rows.Scan(&id, &name, &protocol, &baseURL, &ciphertext); err != nil { + if err := rows.Scan(&id, &slug, &name, &protocol, &wireAPI, &baseURL, &ciphertext); err != nil { return nil, fmt.Errorf("scan provider: %w", err) } apiKey, err := s.cipher.Decrypt(ciphertext) if err != nil { return nil, fmt.Errorf("decrypt provider %q credential: %w", name, err) } - providers[id] = domain.Provider{ID: id, Protocol: domain.Protocol(protocol), BaseURL: strings.TrimRight(baseURL, "/"), APIKey: apiKey} + providers[id] = domain.Provider{ID: id, Slug: slug, Protocol: domain.Protocol(protocol), WireAPI: wireAPI, BaseURL: strings.TrimRight(baseURL, "/"), APIKey: apiKey} } if err := rows.Err(); err != nil { return nil, fmt.Errorf("read providers: %w", err) @@ -245,11 +245,18 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) { rows, err := tx.Query(ctx, ` - SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes + SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes, + k.monthly_spend_micros, k.expires_at, + COALESCE(( + SELECT jsonb_agg(m.public_id ORDER BY m.public_id) + FROM api_key_model_restrictions r + JOIN models m ON m.id = r.model_id + WHERE r.api_key_id = k.id + ), '[]'::jsonb) FROM api_keys k JOIN tenants t ON t.id = k.tenant_id AND t.status = 'active' JOIN projects p ON p.id = k.project_id AND p.status = 'active' - WHERE k.status = 'active'`) + WHERE k.status = 'active' AND (k.expires_at IS NULL OR k.expires_at > now())`) if err != nil { return nil, fmt.Errorf("query API keys: %w", err) } @@ -257,8 +264,11 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) records := make([]auth.HashedKeyRecord, 0) for rows.Next() { var keyID, tenantID, projectID string - var hashBytes, scopesJSON []byte - if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON); err != nil { + var hashBytes, scopesJSON, allowedModelsJSON []byte + var monthlySpendMicros int64 + var expiresAt *time.Time + if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON, + &monthlySpendMicros, &expiresAt, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if len(hashBytes) != sha256.Size { @@ -270,8 +280,17 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) if err := json.Unmarshal(scopesJSON, &scopes); err != nil { return nil, fmt.Errorf("decode API key %s scopes: %w", keyID, err) } + var modelIDs []string + if err := json.Unmarshal(allowedModelsJSON, &modelIDs); err != nil { + return nil, fmt.Errorf("decode API key %s model restrictions: %w", keyID, err) + } + allowedModels := make(map[string]struct{}, len(modelIDs)) + for _, modelID := range modelIDs { + allowedModels[modelID] = struct{}{} + } records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{ KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes, + AllowedModels: allowedModels, MonthlySpendMicros: monthlySpendMicros, ExpiresAt: expiresAt, }}) } if err := rows.Err(); err != nil { diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go index fb48d78..c4d7016 100644 --- a/internal/controlplane/store.go +++ b/internal/controlplane/store.go @@ -23,7 +23,7 @@ var schemaSQL string var ErrRedisDisabled = errors.New("Redis propagation is disabled") -const migrationVersion int64 = 2026080504 +const migrationVersion int64 = 2026080605 type Options struct { DatabaseURL string @@ -133,7 +133,7 @@ func applySchema(ctx context.Context, db *pgxpool.Pool) error { 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 { + if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "tenant-billing-profiles", checksum); err != nil { return err } if err := tx.Commit(ctx); err != nil { diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go index 4f55e10..9bf093c 100644 --- a/internal/controlplane/store_integration_test.go +++ b/internal/controlplane/store_integration_test.go @@ -2,6 +2,7 @@ package controlplane import ( "context" + "encoding/base64" "fmt" "net/url" "os" @@ -11,6 +12,128 @@ import ( "github.com/jackc/pgx/v5/pgxpool" ) +func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(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() + rootDB, err := pgxpool.New(ctx, databaseURL) + if err != nil { + t.Fatal(err) + } + defer rootDB.Close() + schema := fmt.Sprintf("key_controls_%d", time.Now().UnixNano()) + if _, err := rootDB.Exec(ctx, "CREATE SCHEMA "+schema); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _, _ = rootDB.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") }) + 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) + } + credentialKey := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) + store, err := NewStore(ctx, Options{DatabaseURL: isolatedURL, CredentialKey: credentialKey}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + tenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "key-control-test", Name: "Key Control Test"}) + if err != nil { + t.Fatal(err) + } + project, _, err := store.CreateProject(ctx, CreateProjectInput{TenantID: tenant.ID, Slug: "production", Name: "Production"}) + if err != nil { + t.Fatal(err) + } + provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "responses", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}) + if err != nil { + t.Fatal(err) + } + if provider.Slug != "key-control-provider" { + t.Fatalf("derived provider slug = %q, want key-control-provider", provider.Slug) + } + model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}}) + if err != nil { + t.Fatal(err) + } + expiresAt := time.Now().Add(24 * time.Hour).UTC().Truncate(time.Microsecond) + created, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID, + Name: "production backend", Scopes: []string{"inference"}, Tags: []string{"production", "backend"}, + AllowedModels: []string{model.PublicID}, MonthlySpendMicros: 25_000_000, ExpiresAt: &expiresAt}) + if err != nil { + t.Fatal(err) + } + if created.Key == "" || created.MonthlySpendMicros != 25_000_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 { + t.Fatalf("unexpected created key: %+v", created.APIKey) + } + if _, err := store.db.Exec(ctx, `INSERT INTO usage_events + (request_id,tenant_id,project_id,key_id,public_model,protocol,status_code,success,started_at,cost_micros) + VALUES ('req_key_current_month',$1,$2,$3,$4,'responses',200,TRUE,now(),42000), + ('req_key_previous_month',$1,$2,$3,$4,'responses',200,TRUE,now()-interval '2 months',99000)`, + tenant.ID, project.ID, created.ID, model.PublicID); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO billing_reservations + (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,status, + input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) + VALUES ('req_key_pending',$1,$2,$3,$4,'usd',9000,'pending',100000,200000,0,0)`, + tenant.ID, project.ID, created.ID, model.PublicID); err != nil { + t.Fatal(err) + } + keys, err := store.ListAPIKeysFor(ctx, tenant.ID) + if err != nil { + t.Fatal(err) + } + if len(keys) != 1 || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) { + t.Fatalf("key restrictions did not round trip: %+v", keys) + } + if keys[0].CurrentMonthSpendMicros != 42_000 || keys[0].CurrentMonthReservedMicros != 9_000 || keys[0].CurrentMonthRequests != 1 { + t.Fatalf("key month activity is incorrect: %+v", keys[0]) + } + snapshot, err := store.LoadSnapshot(ctx) + if err != nil { + t.Fatal(err) + } + if len(snapshot.APIKeys) != 1 { + t.Fatalf("snapshot API keys = %d, want 1", len(snapshot.APIKeys)) + } + principal := snapshot.APIKeys[0].Principal + if principal.MonthlySpendMicros != 25_000_000 || principal.ExpiresAt == nil { + t.Fatalf("snapshot lost API key controls: %+v", principal) + } + if _, ok := principal.AllowedModels[model.PublicID]; !ok { + t.Fatalf("snapshot lost allowed model: %+v", principal.AllowedModels) + } + otherTenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "other-key-control-test", Name: "Other Key Control Test"}) + if err != nil { + t.Fatal(err) + } + otherProject, _, err := store.CreateProject(ctx, CreateProjectInput{TenantID: otherTenant.ID, Slug: "default", Name: "Other Default"}) + if err != nil { + t.Fatal(err) + } + if _, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: otherProject.ID, Name: "cross tenant"}); err == nil { + t.Fatal("cross-tenant project binding should be rejected by the composite foreign key") + } + if _, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID, + Name: "invalid model", AllowedModels: []string{"model/does-not-exist"}}); err == nil { + t.Fatal("unknown allowed model should reject the whole API key transaction") + } + keys, err = store.ListAPIKeysFor(ctx, tenant.ID) + if err != nil || len(keys) != 1 { + t.Fatalf("failed key transaction leaked a row: keys=%d err=%v", len(keys), err) + } +} + func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") if databaseURL == "" { @@ -63,4 +186,48 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { if count != 2 { t.Fatalf("migration history contains %d rows, want previous and current", count) } + + scopedDB, err := pgxpool.New(ctx, isolatedURL) + if err != nil { + t.Fatal(err) + } + defer scopedDB.Close() + var tenantID, providerID, modelID string + if err := scopedDB.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ('preference-test','Preference Test') RETURNING id::text`).Scan(&tenantID); err != nil { + t.Fatal(err) + } + if err := scopedDB.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext) + VALUES ('preference-provider','Preference provider','openai','responses','https://example.invalid',decode('00','hex')) RETURNING id::text`).Scan(&providerID); err != nil { + t.Fatal(err) + } + if err := scopedDB.QueryRow(ctx, `INSERT INTO models (public_id,display_name) VALUES ('model/preference-test','Preference Model') RETURNING id::text`).Scan(&modelID); err != nil { + t.Fatal(err) + } + if _, err := scopedDB.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency) VALUES ($1,1,'usd')`, modelID); err != nil { + t.Fatal(err) + } + if _, err := scopedDB.Exec(ctx, `INSERT INTO model_routes (model_id,provider_id,upstream_model) VALUES ($1,$2,'upstream-test')`, modelID, providerID); err != nil { + t.Fatal(err) + } + store := &Store{db: scopedDB} + prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"}) + if err != nil { + t.Fatal(err) + } + if prefs.DefaultModel != "model/preference-test" || prefs.FallbackModel != "" { + t.Fatalf("unexpected developer preferences: %+v", prefs) + } + enabled := false + threshold := int64(9_750_000) + if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenantID, + LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold}, 5_000_000); err != nil { + t.Fatal(err) + } + prefs, err = store.GetTenantPreferences(ctx, tenantID, 5_000_000) + if err != nil { + t.Fatal(err) + } + if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.DefaultModel != "model/preference-test" { + t.Fatalf("preferences did not round trip: %+v", prefs) + } } diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go index c24dde9..f6402d8 100644 --- a/internal/controlplane/types.go +++ b/internal/controlplane/types.go @@ -31,14 +31,22 @@ type Project struct { } type APIKey struct { - ID string `json:"id"` - TenantID string `json:"tenant_id"` - ProjectID string `json:"project_id"` - Name string `json:"name"` - KeyPrefix string `json:"key_prefix"` - Scopes []string `json:"scopes"` - Status string `json:"status"` - CreatedAt time.Time `json:"created_at"` + ID string `json:"id"` + TenantID string `json:"tenant_id"` + ProjectID string `json:"project_id"` + Name string `json:"name"` + KeyPrefix string `json:"key_prefix"` + Scopes []string `json:"scopes"` + Tags []string `json:"tags"` + AllowedModels []string `json:"allowed_models"` + MonthlySpendMicros int64 `json:"monthly_spend_micros"` + CurrentMonthSpendMicros int64 `json:"current_month_spend_micros"` + CurrentMonthReservedMicros int64 `json:"current_month_reserved_micros"` + CurrentMonthRequests int64 `json:"current_month_requests"` + Status string `json:"status"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + LastUsedAt *time.Time `json:"last_used_at,omitempty"` + CreatedAt time.Time `json:"created_at"` } type CreatedAPIKey struct { @@ -48,8 +56,10 @@ type CreatedAPIKey struct { type Provider struct { ID string `json:"id"` + Slug string `json:"slug"` Name string `json:"name"` Protocol string `json:"protocol"` + WireAPI string `json:"wire_api"` BaseURL string `json:"base_url"` Enabled bool `json:"enabled"` RouteCount int `json:"route_count"` @@ -57,14 +67,112 @@ type Provider struct { } type Route struct { - ID string `json:"id"` - ProviderID string `json:"provider_id"` - ProviderName string `json:"provider_name"` - Protocol string `json:"protocol"` - UpstreamModel string `json:"upstream_model"` - Priority int `json:"priority"` - Weight int `json:"weight"` - Enabled bool `json:"enabled"` + ID string `json:"id"` + ProviderID string `json:"provider_id"` + ProviderName string `json:"provider_name"` + Protocol string `json:"protocol"` + WireAPI string `json:"wire_api"` + UpstreamModel string `json:"upstream_model"` + Priority int `json:"priority"` + Weight int `json:"weight"` + Enabled bool `json:"enabled"` + ProviderEnabled bool `json:"provider_enabled"` +} + +// DeveloperModel is the customer-safe model catalog view. It intentionally +// omits internal provider URLs, upstream model names, routing weights, and +// allowlist membership. +type DeveloperModel struct { + ID string `json:"-"` + 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"` + ReplacementModel string `json:"replacement_model,omitempty"` + Aliases []string `json:"aliases"` + 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"` + CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"` + SupportedWireAPIs []string `json:"supported_wire_apis"` + ProviderCount int `json:"provider_count"` + AvailableProviderCount int `json:"available_provider_count"` + HealthStatus string `json:"health_status"` + Providers []DeveloperProviderHealth `json:"providers"` +} + +type DeveloperProviderHealth struct { + Slug string `json:"slug"` + Name string `json:"name"` + Protocol string `json:"protocol"` + WireAPI string `json:"wire_api"` + State string `json:"state"` + Attempts uint64 `json:"attempts"` + RecentSamples int `json:"recent_samples"` + AvailabilityPercent float64 `json:"availability_percent"` + HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"` + ConsecutiveFailures uint64 `json:"consecutive_failures"` + LastObservedAt *time.Time `json:"last_observed_at,omitempty"` + CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"` +} + +// PublicModel is the unauthenticated catalog view. It is deliberately smaller +// than DeveloperModel so authenticated-only fields cannot become public by +// accident when the tenant catalog evolves. +type PublicModel 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,omitempty"` + ReplacementModel string `json:"replacement_model,omitempty"` + Aliases []string `json:"aliases"` + 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"` + CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"` + SupportedWireAPIs []string `json:"supported_wire_apis"` + ProviderCount int `json:"provider_count"` + AvailableProviderCount int `json:"available_provider_count"` + HealthStatus string `json:"health_status"` +} + +type TenantPreferences struct { + TenantID string `json:"tenant_id,omitempty"` + DefaultModel string `json:"default_model,omitempty"` + FallbackModel string `json:"fallback_model,omitempty"` + LowBalanceEnabled bool `json:"low_balance_enabled"` + LowBalanceThresholdMicros int64 `json:"low_balance_threshold_micros"` + UpdatedAt *time.Time `json:"updated_at,omitempty"` +} + +type SetDeveloperPreferencesInput struct { + TenantID string `json:"tenant_id"` + DefaultModel string `json:"default_model"` + FallbackModel string `json:"fallback_model"` +} + +type SetBillingPreferencesInput struct { + TenantID string `json:"tenant_id"` + LowBalanceEnabled *bool `json:"low_balance_enabled"` + LowBalanceThresholdMicros *int64 `json:"low_balance_threshold_micros"` } type Model struct { @@ -145,15 +253,21 @@ type CreateProjectInput struct { } type CreateAPIKeyInput struct { - TenantID string `json:"tenant_id"` - ProjectID string `json:"project_id"` - Name string `json:"name"` - Scopes []string `json:"scopes"` + TenantID string `json:"tenant_id"` + ProjectID string `json:"project_id"` + Name string `json:"name"` + Scopes []string `json:"scopes"` + Tags []string `json:"tags"` + AllowedModels []string `json:"allowed_models"` + MonthlySpendMicros int64 `json:"monthly_spend_micros"` + ExpiresAt *time.Time `json:"expires_at"` } type CreateProviderInput struct { + Slug string `json:"slug"` Name string `json:"name"` Protocol string `json:"protocol"` + WireAPI string `json:"wire_api"` BaseURL string `json:"base_url"` APIKey string `json:"api_key"` } @@ -340,9 +454,12 @@ type UsageRecord struct { RequestID string `json:"request_id"` TenantID string `json:"tenant_id"` ProjectID string `json:"project_id"` + ProjectName string `json:"project_name"` KeyID string `json:"key_id"` + KeyName string `json:"key_name"` PublicModel string `json:"public_model"` ProviderID string `json:"provider_id,omitempty"` + ProviderName string `json:"provider_name,omitempty"` UpstreamModel string `json:"upstream_model,omitempty"` Protocol string `json:"protocol"` Stream bool `json:"stream"` @@ -360,6 +477,72 @@ type UsageRecord struct { CostMicros int64 `json:"cost_micros"` ChargedMicros int64 `json:"charged_micros"` UncollectedMicros int64 `json:"uncollected_micros"` + UsageReported bool `json:"usage_reported"` + MeteringStatus string `json:"metering_status"` +} + +type UsageDailyPoint struct { + Day time.Time `json:"day"` + RequestCount int64 `json:"request_count"` + SuccessfulRequests int64 `json:"successful_requests"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + ChargedMicros int64 `json:"charged_micros"` + UncollectedMicros int64 `json:"uncollected_micros"` + AverageDurationMS int64 `json:"average_duration_ms"` + P95DurationMS int64 `json:"p95_duration_ms"` +} + +// UsageAnalytics is a persisted-ledger aggregation used by the developer +// console. It deliberately contains no prompt or response content. +type UsageAnalytics struct { + RangeStart time.Time `json:"range_start"` + RangeEnd time.Time `json:"range_end"` + Models []UsageModelAnalytics `json:"models"` + Providers []UsageProviderAnalytics `json:"providers"` +} + +type UsageModelAnalytics struct { + PublicModel string `json:"public_model"` + RequestCount int64 `json:"request_count"` + SuccessfulRequests int64 `json:"successful_requests"` + ErrorCount int64 `json:"error_count"` + ProviderCount int64 `json:"provider_count"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + CacheReadInputTokens int64 `json:"cache_read_input_tokens"` + CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"` + ChargedMicros int64 `json:"charged_micros"` + UncollectedMicros int64 `json:"uncollected_micros"` + MissingUsageRequests int64 `json:"missing_usage_requests"` + AverageDurationMS int64 `json:"average_duration_ms"` + P95DurationMS int64 `json:"p95_duration_ms"` + PreviousChargedMicros int64 `json:"previous_charged_micros"` + ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"` +} + +type UsageProviderAnalytics struct { + ProviderID string `json:"provider_id"` + ProviderName string `json:"provider_name"` + WireAPI string `json:"wire_api"` + RequestCount int64 `json:"request_count"` + SuccessfulRequests int64 `json:"successful_requests"` + ErrorCount int64 `json:"error_count"` + ModelCount int64 `json:"model_count"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + CacheReadInputTokens int64 `json:"cache_read_input_tokens"` + CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"` + ChargedMicros int64 `json:"charged_micros"` + UncollectedMicros int64 `json:"uncollected_micros"` + MissingUsageRequests int64 `json:"missing_usage_requests"` + AverageDurationMS int64 `json:"average_duration_ms"` + P95DurationMS int64 `json:"p95_duration_ms"` + PreviousChargedMicros int64 `json:"previous_charged_micros"` + ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"` } type UsageSummary struct { diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go index b436d8c..ebabaf6 100644 --- a/internal/controlplane/usage.go +++ b/internal/controlplane/usage.go @@ -14,7 +14,16 @@ import ( type UsageQuery struct { TenantID string ProjectID string + KeyID string Model string + Provider string + Protocol string + ErrorType string + Stream *bool + RequestID string + Status string + From time.Time + To time.Time Limit int } @@ -39,6 +48,11 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error if err != nil { return fmt.Errorf("persist usage event: %w", err) } + if _, err := tx.Exec(ctx, `UPDATE api_keys + SET last_used_at = GREATEST(COALESCE(last_used_at, $2), $2) + WHERE id = $1`, event.KeyID, event.StartedAt); err != nil { + return fmt.Errorf("update API key last used time: %w", err) + } if command.RowsAffected() == 0 { return tx.Commit(ctx) } @@ -96,20 +110,55 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, limit = 200 } where := []string{"1=1"} - args := make([]any, 0, 5) + args := make([]any, 0, 13) index := 1 - for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.Model, "public_model=$"}} { + for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.KeyID, "key_id=$"}, {query.Model, "public_model=$"}, {query.RequestID, "request_id=$"}} { if strings.TrimSpace(item.value) != "" { where = append(where, item.clause+fmt.Sprint(index)) args = append(args, item.value) index++ } } + for _, item := range []struct{ value, clause string }{{query.Protocol, "protocol=$"}, {query.ErrorType, "error_type=$"}} { + if strings.TrimSpace(item.value) != "" { + where = append(where, item.clause+fmt.Sprint(index)) + args = append(args, item.value) + index++ + } + } + if strings.TrimSpace(query.Provider) != "" { + where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=usage_events.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")") + args = append(args, query.Provider) + index++ + } + if query.Stream != nil { + where = append(where, "stream=$"+fmt.Sprint(index)) + args = append(args, *query.Stream) + index++ + } + if query.Status == "success" { + where = append(where, "success=TRUE") + } else if query.Status == "error" { + where = append(where, "success=FALSE") + } + if !query.From.IsZero() { + where = append(where, "started_at >= $"+fmt.Sprint(index)) + args = append(args, query.From) + index++ + } + if !query.To.IsZero() { + where = append(where, "started_at < $"+fmt.Sprint(index)) + args = append(args, query.To) + index++ + } args = append(args, limit) - rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, key_id::text, public_model, - COALESCE(provider_id,''), COALESCE(upstream_model,''), protocol, stream, status_code, success, error_type, + rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, + COALESCE((SELECT name FROM projects p WHERE p.id=usage_events.project_id),''), key_id::text, + COALESCE((SELECT name FROM api_keys k WHERE k.id=usage_events.key_id),''), public_model, + COALESCE(provider_id,''), COALESCE((SELECT name FROM providers p WHERE p.id::text=usage_events.provider_id),''), + COALESCE(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 FROM usage_events WHERE `+ + cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status FROM usage_events WHERE `+ strings.Join(where, " AND ")+` ORDER BY created_at DESC LIMIT $`+fmt.Sprint(index), args...) if err != nil { return nil, fmt.Errorf("query usage events: %w", err) @@ -118,10 +167,10 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, result := make([]UsageRecord, 0) for rows.Next() { var item UsageRecord - if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.KeyID, &item.PublicModel, &item.ProviderID, &item.UpstreamModel, + if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.ProjectName, &item.KeyID, &item.KeyName, &item.PublicModel, &item.ProviderID, &item.ProviderName, &item.UpstreamModel, &item.Protocol, &item.Stream, &item.StatusCode, &item.Success, &item.ErrorType, &item.Attempts, &item.StartedAt, &item.DurationMS, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, - &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros); err != nil { + &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros, &item.UsageReported, &item.MeteringStatus); err != nil { return nil, fmt.Errorf("scan usage event: %w", err) } result = append(result, item) @@ -129,6 +178,71 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, return result, rows.Err() } +func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyPoint, error) { + where := []string{"1=1"} + args := make([]any, 0, 12) + index := 1 + for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.KeyID, "key_id=$"}, {query.Model, "public_model=$"}, {query.RequestID, "request_id=$"}} { + if strings.TrimSpace(item.value) != "" { + where = append(where, item.clause+fmt.Sprint(index)) + args = append(args, item.value) + index++ + } + } + for _, item := range []struct{ value, clause string }{{query.Protocol, "protocol=$"}, {query.ErrorType, "error_type=$"}} { + if strings.TrimSpace(item.value) != "" { + where = append(where, item.clause+fmt.Sprint(index)) + args = append(args, item.value) + index++ + } + } + if strings.TrimSpace(query.Provider) != "" { + where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=usage_events.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")") + args = append(args, query.Provider) + index++ + } + if query.Stream != nil { + where = append(where, "stream=$"+fmt.Sprint(index)) + args = append(args, *query.Stream) + index++ + } + if query.Status == "success" { + where = append(where, "success=TRUE") + } else if query.Status == "error" { + where = append(where, "success=FALSE") + } + if !query.From.IsZero() { + where = append(where, "started_at >= $"+fmt.Sprint(index)) + args = append(args, query.From) + index++ + } + if !query.To.IsZero() { + where = append(where, "started_at < $"+fmt.Sprint(index)) + args = append(args, query.To) + index++ + } + rows, err := s.db.Query(ctx, `SELECT date_trunc('day', started_at AT TIME ZONE 'UTC'), count(*), + count(*) FILTER (WHERE success), COALESCE(sum(input_tokens),0), COALESCE(sum(output_tokens),0), + COALESCE(sum(total_tokens),0), COALESCE(sum(charged_micros),0), COALESCE(sum(uncollected_micros),0), + COALESCE(round(avg(duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint + FROM usage_events WHERE `+strings.Join(where, " AND ")+` GROUP BY 1 ORDER BY 1`, args...) + if err != nil { + return nil, fmt.Errorf("query daily usage: %w", err) + } + defer rows.Close() + result := make([]UsageDailyPoint, 0) + for rows.Next() { + var item UsageDailyPoint + if err := rows.Scan(&item.Day, &item.RequestCount, &item.SuccessfulRequests, &item.InputTokens, &item.OutputTokens, + &item.TotalTokens, &item.ChargedMicros, &item.UncollectedMicros, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + return nil, fmt.Errorf("scan daily usage: %w", err) + } + result = append(result, item) + } + return result, rows.Err() +} + func (s *Store) UsageSummary(ctx context.Context, tenantID, projectID string) ([]UsageSummary, error) { where := []string{"1=1"} args := make([]any, 0, 2) diff --git a/internal/controlplane/usage_analytics.go b/internal/controlplane/usage_analytics.go new file mode 100644 index 0000000..7cc042f --- /dev/null +++ b/internal/controlplane/usage_analytics.go @@ -0,0 +1,215 @@ +package controlplane + +import ( + "context" + "fmt" + "strings" + "time" +) + +// UsageAnalytics aggregates the immutable usage ledger for the developer +// console. Queries run outside the inference path and are scoped by the +// caller's tenant before reaching this store. +func (s *Store) UsageAnalytics(ctx context.Context, query UsageQuery) (UsageAnalytics, error) { + to := query.To + if to.IsZero() { + to = time.Now().UTC() + } + from := query.From + if from.IsZero() { + from = to.Add(-30 * 24 * time.Hour) + } + if !to.After(from) { + return UsageAnalytics{}, fmt.Errorf("usage analytics range must be positive") + } + result := UsageAnalytics{RangeStart: from, RangeEnd: to, Models: make([]UsageModelAnalytics, 0), Providers: make([]UsageProviderAnalytics, 0)} + + modelPrevious, err := s.usageModelCharges(ctx, query, from.Add(-to.Sub(from)), from) + if err != nil { + return UsageAnalytics{}, err + } + models, err := s.usageModelAnalytics(ctx, query, from, to) + if err != nil { + return UsageAnalytics{}, err + } + for index := range models { + models[index].PreviousChargedMicros = modelPrevious[models[index].PublicModel] + models[index].ChargeChangePercent = chargeChange(models[index].ChargedMicros, models[index].PreviousChargedMicros) + } + + providerPrevious, err := s.usageProviderCharges(ctx, query, from.Add(-to.Sub(from)), from) + if err != nil { + return UsageAnalytics{}, err + } + providers, err := s.usageProviderAnalytics(ctx, query, from, to) + if err != nil { + return UsageAnalytics{}, err + } + for index := range providers { + providers[index].PreviousChargedMicros = providerPrevious[providers[index].ProviderID] + providers[index].ChargeChangePercent = chargeChange(providers[index].ChargedMicros, providers[index].PreviousChargedMicros) + } + result.Models = models + result.Providers = providers + return result, nil +} + +func chargeChange(current, previous int64) *float64 { + if previous == 0 { + return nil + } + value := (float64(current) - float64(previous)) / float64(previous) * 100 + return &value +} + +func analyticsUsageWhere(query UsageQuery, from, to time.Time) (string, []any) { + where := []string{"1=1"} + args := make([]any, 0, 12) + index := 1 + for _, item := range []struct { + value string + clause string + }{ + {query.TenantID, "e.tenant_id=$"}, + {query.ProjectID, "e.project_id=$"}, + {query.KeyID, "e.key_id=$"}, + {query.Model, "e.public_model=$"}, + {query.RequestID, "e.request_id=$"}, + } { + if strings.TrimSpace(item.value) != "" { + where = append(where, item.clause+fmt.Sprint(index)) + args = append(args, item.value) + index++ + } + } + for _, item := range []struct { + value string + clause string + }{ + {query.Protocol, "e.protocol=$"}, + {query.ErrorType, "e.error_type=$"}, + } { + if strings.TrimSpace(item.value) != "" { + where = append(where, item.clause+fmt.Sprint(index)) + args = append(args, item.value) + index++ + } + } + if strings.TrimSpace(query.Provider) != "" { + where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=e.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")") + args = append(args, query.Provider) + index++ + } + if query.Stream != nil { + where = append(where, "e.stream=$"+fmt.Sprint(index)) + args = append(args, *query.Stream) + index++ + } + if query.Status == "success" { + where = append(where, "e.success=TRUE") + } else if query.Status == "error" { + where = append(where, "e.success=FALSE") + } + if !from.IsZero() { + where = append(where, "e.started_at >= $"+fmt.Sprint(index)) + args = append(args, from) + index++ + } + if !to.IsZero() { + where = append(where, "e.started_at < $"+fmt.Sprint(index)) + args = append(args, to) + } + return strings.Join(where, " AND "), args +} + +func (s *Store) usageModelAnalytics(ctx context.Context, query UsageQuery, from, to time.Time) ([]UsageModelAnalytics, error) { + where, args := analyticsUsageWhere(query, from, to) + rows, err := s.db.Query(ctx, `SELECT e.public_model, count(*), count(*) FILTER (WHERE e.success), count(*) FILTER (WHERE NOT e.success), + count(DISTINCT NULLIF(e.provider_id,'')), COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), + COALESCE(sum(e.total_tokens),0), COALESCE(sum(e.cache_read_input_tokens),0), COALESCE(sum(e.cache_creation_input_tokens),0), + COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), + COALESCE(round(avg(e.duration_ms)),0)::bigint, COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint + FROM usage_events e WHERE `+where+` GROUP BY e.public_model ORDER BY sum(e.charged_micros) DESC, e.public_model`, args...) + if err != nil { + return nil, fmt.Errorf("query usage model analytics: %w", err) + } + defer rows.Close() + result := make([]UsageModelAnalytics, 0) + for rows.Next() { + var item UsageModelAnalytics + if err := rows.Scan(&item.PublicModel, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, &item.ProviderCount, + &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheReadInputTokens, &item.CacheCreationInputTokens, + &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + return nil, fmt.Errorf("scan usage model analytics: %w", err) + } + result = append(result, item) + } + return result, rows.Err() +} + +func (s *Store) usageModelCharges(ctx context.Context, query UsageQuery, from, to time.Time) (map[string]int64, error) { + where, args := analyticsUsageWhere(query, from, to) + rows, err := s.db.Query(ctx, `SELECT e.public_model, COALESCE(sum(e.charged_micros),0) + FROM usage_events e WHERE `+where+` GROUP BY e.public_model`, args...) + if err != nil { + return nil, fmt.Errorf("query previous model charges: %w", err) + } + defer rows.Close() + result := make(map[string]int64) + for rows.Next() { + var model string + var charged int64 + if err := rows.Scan(&model, &charged); err != nil { + return nil, fmt.Errorf("scan previous model charges: %w", err) + } + result[model] = charged + } + return result, rows.Err() +} + +func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, from, to time.Time) ([]UsageProviderAnalytics, error) { + where, args := analyticsUsageWhere(query, from, to) + rows, err := s.db.Query(ctx, `SELECT COALESCE(e.provider_id,''), COALESCE(NULLIF(p.name,''),'Unassigned'), COALESCE(p.wire_api,''), + count(*), count(*) FILTER (WHERE e.success), count(*) FILTER (WHERE NOT e.success), count(DISTINCT e.public_model), + COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), + COALESCE(sum(e.cache_read_input_tokens),0), COALESCE(sum(e.cache_creation_input_tokens),0), COALESCE(sum(e.charged_micros),0), + COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), COALESCE(round(avg(e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint + FROM usage_events e LEFT JOIN providers p ON p.id::text=e.provider_id WHERE `+where+` + GROUP BY e.provider_id, p.name, p.wire_api ORDER BY sum(e.charged_micros) DESC, COALESCE(NULLIF(p.name,''),'Unassigned')`, args...) + if err != nil { + return nil, fmt.Errorf("query usage provider analytics: %w", err) + } + defer rows.Close() + result := make([]UsageProviderAnalytics, 0) + for rows.Next() { + var item UsageProviderAnalytics + if err := rows.Scan(&item.ProviderID, &item.ProviderName, &item.WireAPI, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, + &item.ModelCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheReadInputTokens, &item.CacheCreationInputTokens, + &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + return nil, fmt.Errorf("scan usage provider analytics: %w", err) + } + result = append(result, item) + } + return result, rows.Err() +} + +func (s *Store) usageProviderCharges(ctx context.Context, query UsageQuery, from, to time.Time) (map[string]int64, error) { + where, args := analyticsUsageWhere(query, from, to) + rows, err := s.db.Query(ctx, `SELECT COALESCE(e.provider_id,''), COALESCE(sum(e.charged_micros),0) + FROM usage_events e WHERE `+where+` GROUP BY e.provider_id`, args...) + if err != nil { + return nil, fmt.Errorf("query previous provider charges: %w", err) + } + defer rows.Close() + result := make(map[string]int64) + for rows.Next() { + var providerID string + var charged int64 + if err := rows.Scan(&providerID, &charged); err != nil { + return nil, fmt.Errorf("scan previous provider charges: %w", err) + } + result[providerID] = charged + } + return result, rows.Err() +} diff --git a/internal/controlplane/usage_analytics_test.go b/internal/controlplane/usage_analytics_test.go new file mode 100644 index 0000000..2372be6 --- /dev/null +++ b/internal/controlplane/usage_analytics_test.go @@ -0,0 +1,34 @@ +package controlplane + +import "testing" + +func TestChargeChange(t *testing.T) { + tests := []struct { + name string + current, previous int64 + want *float64 + }{ + {name: "no activity", current: 0, previous: 0, want: nil}, + {name: "new spend has no finite percentage", current: 25, previous: 0, want: nil}, + {name: "increase", current: 125, previous: 25, want: floatPointer(400)}, + {name: "decrease", current: 25, previous: 100, want: floatPointer(-75)}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got := chargeChange(test.current, test.previous) + if test.want == nil { + if got != nil { + t.Fatalf("chargeChange() = %v, want nil", *got) + } + return + } + if got == nil || *got != *test.want { + t.Fatalf("chargeChange() = %v, want %v", got, *test.want) + } + }) + } +} + +func floatPointer(value float64) *float64 { + return &value +} diff --git a/internal/controlplane/usage_integration_test.go b/internal/controlplane/usage_integration_test.go new file mode 100644 index 0000000..2545c1d --- /dev/null +++ b/internal/controlplane/usage_integration_test.go @@ -0,0 +1,152 @@ +package controlplane + +import ( + "context" + "crypto/sha256" + "fmt" + "os" + "testing" + "time" + + "aigw/internal/domain" + + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestUsageFiltersAndDailyAggregationPostgres(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 := MigrateDatabase(ctx, databaseURL); err != nil { + t.Fatal(err) + } + db, err := pgxpool.New(ctx, databaseURL) + if err != nil { + t.Fatal(err) + } + t.Cleanup(db.Close) + store := &Store{db: db} + suffix := time.Now().UnixNano() + var tenantID, projectID, keyID, providerID, otherTenantID, otherProjectID, otherKeyID string + if err := db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Usage integration') RETURNING id::text`, fmt.Sprintf("usage-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'production','Production') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { + t.Fatal(err) + } + keyHash := sha256.Sum256([]byte(fmt.Sprint(suffix))) + if err := db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Production key','sk-usage',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext) + VALUES ($1,$2,'openai','responses','https://usage.test',$3) RETURNING id::text`, fmt.Sprintf("usage-provider-%d", suffix), fmt.Sprintf("Usage provider %d", suffix), []byte("encrypted-test-value")).Scan(&providerID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Other usage tenant') RETURNING id::text`, fmt.Sprintf("usage-other-%d", suffix)).Scan(&otherTenantID); err != nil { + t.Fatal(err) + } + if err := db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'production','Other production') RETURNING id::text`, otherTenantID).Scan(&otherProjectID); err != nil { + t.Fatal(err) + } + otherKeyHash := sha256.Sum256([]byte(fmt.Sprintf("other-%d", suffix))) + if err := db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Other key','sk-other',$3) RETURNING id::text`, otherTenantID, otherProjectID, otherKeyHash[:]).Scan(&otherKeyID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + cleanupCtx := context.Background() + for _, statement := range []struct{ query, arg string }{ + {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, otherTenantID}, + {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM usage_events WHERE tenant_id=$1`, otherTenantID}, + {`DELETE FROM providers WHERE id=$1`, providerID}, + {`DELETE FROM api_keys WHERE id=$1`, keyID}, + {`DELETE FROM api_keys WHERE id=$1`, otherKeyID}, + {`DELETE FROM projects WHERE id=$1`, projectID}, + {`DELETE FROM projects WHERE id=$1`, otherProjectID}, + {`DELETE FROM tenants WHERE id=$1`, tenantID}, + {`DELETE FROM tenants WHERE id=$1`, otherTenantID}, + } { + if _, cleanupErr := db.Exec(cleanupCtx, statement.query, statement.arg); cleanupErr != nil { + t.Errorf("cleanup usage integration data: %v", cleanupErr) + } + } + }) + + started := time.Now().UTC().Add(-time.Hour).Truncate(time.Second) + events := []domain.UsageEvent{ + {RequestID: fmt.Sprintf("req_usage_ok_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started, DurationMS: 120, Usage: domain.Usage{InputTokens: 10, OutputTokens: 4, TotalTokens: 14, CacheReadInputTokens: 5}}, + {RequestID: fmt.Sprintf("req_usage_error_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAIResponses, Stream: true, StatusCode: 502, Success: false, ErrorType: "provider_error", StartedAt: started.Add(time.Minute), DurationMS: 350}, + } + for _, event := range events { + if err := store.RecordUsage(ctx, event); err != nil { + t.Fatal(err) + } + } + if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=125,cost_micros=125 WHERE request_id=$1`, events[0].RequestID); err != nil { + t.Fatal(err) + } + + records, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, KeyID: keyID, Status: "success", From: started.Add(-time.Minute), To: started.Add(time.Hour), Limit: 20}) + if err != nil { + t.Fatal(err) + } + if len(records) != 1 || records[0].RequestID != events[0].RequestID || records[0].ProjectName != "Production" || records[0].KeyName != "Production key" || records[0].MeteringStatus != "reported" { + t.Fatalf("unexpected filtered usage: %+v", records) + } + streaming := true + providerSlug := fmt.Sprintf("usage-provider-%d", suffix) + records, err = store.ListUsage(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(time.Hour), Limit: 20}) + if err != nil { + t.Fatal(err) + } + if len(records) != 1 || records[0].RequestID != events[1].RequestID { + t.Fatalf("unexpected provider/protocol/transport usage filter: %+v", records) + } + filteredPoints, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(time.Hour)}) + if err != nil { + t.Fatal(err) + } + if len(filteredPoints) != 1 || filteredPoints[0].RequestCount != 1 || filteredPoints[0].SuccessfulRequests != 0 { + t.Fatalf("unexpected filtered daily usage: %+v", filteredPoints) + } + points, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(2 * time.Hour)}) + if err != nil { + t.Fatal(err) + } + if len(points) != 1 || points[0].RequestCount != 2 || points[0].SuccessfulRequests != 1 || points[0].TotalTokens != 14 || points[0].ChargedMicros != 125 || points[0].P95DurationMS < 120 { + t.Fatalf("unexpected daily usage: %+v", points) + } + + previous := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_previous_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(-5 * time.Minute), DurationMS: 80, Usage: domain.Usage{InputTokens: 3, OutputTokens: 2, TotalTokens: 5}} + missing := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_missing_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, StartedAt: started.Add(2 * time.Minute), DurationMS: 200} + otherTenantEvent := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_other_%d", suffix), TenantID: otherTenantID, ProjectID: otherProjectID, KeyID: otherKeyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(3 * time.Minute), DurationMS: 900, Usage: domain.Usage{InputTokens: 1000, OutputTokens: 1000, TotalTokens: 2000}} + for _, event := range []domain.UsageEvent{previous, missing, otherTenantEvent} { + if err := store.RecordUsage(ctx, event); err != nil { + t.Fatal(err) + } + } + if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=25,cost_micros=25 WHERE request_id=$1`, previous.RequestID); err != nil { + t.Fatal(err) + } + analytics, err := store.UsageAnalytics(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Minute), To: started.Add(10 * time.Minute)}) + if err != nil { + t.Fatal(err) + } + if len(analytics.Models) != 1 || analytics.Models[0].RequestCount != 3 || analytics.Models[0].SuccessfulRequests != 2 || analytics.Models[0].ProviderCount != 1 || analytics.Models[0].CacheReadInputTokens != 5 || analytics.Models[0].MissingUsageRequests != 1 || analytics.Models[0].ChargedMicros != 125 || analytics.Models[0].PreviousChargedMicros != 25 || analytics.Models[0].ChargeChangePercent == nil || *analytics.Models[0].ChargeChangePercent != 400 { + t.Fatalf("unexpected model analytics: %+v", analytics.Models) + } + if len(analytics.Providers) != 1 || analytics.Providers[0].ProviderID != providerID || analytics.Providers[0].WireAPI != "responses" || analytics.Providers[0].RequestCount != 3 || analytics.Providers[0].MissingUsageRequests != 1 || analytics.Providers[0].P95DurationMS < 200 { + t.Fatalf("unexpected provider analytics: %+v", analytics.Providers) + } + filteredAnalytics, err := store.UsageAnalytics(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(10 * time.Minute)}) + if err != nil { + t.Fatal(err) + } + if len(filteredAnalytics.Models) != 1 || filteredAnalytics.Models[0].RequestCount != 1 || filteredAnalytics.Models[0].ErrorCount != 1 || len(filteredAnalytics.Providers) != 1 || filteredAnalytics.Providers[0].RequestCount != 1 { + t.Fatalf("unexpected filtered analytics: %+v", filteredAnalytics) + } +} diff --git a/internal/domain/types.go b/internal/domain/types.go index 94e56c1..81da8f2 100644 --- a/internal/domain/types.go +++ b/internal/domain/types.go @@ -5,24 +5,47 @@ import "time" type Protocol string const ( - ProtocolOpenAI Protocol = "openai" - ProtocolAnthropic Protocol = "anthropic" + ProtocolOpenAI Protocol = "openai" + ProtocolOpenAIResponses Protocol = "openai_responses" + ProtocolAnthropic Protocol = "anthropic" ) type Principal struct { - KeyID string - TenantID string - ProjectID string - Scopes []string + KeyID string + TenantID string + ProjectID string + Scopes []string + AllowedModels map[string]struct{} + MonthlySpendMicros int64 + ExpiresAt *time.Time } type Provider struct { ID string + Slug string Protocol Protocol + WireAPI string BaseURL string APIKey string } +func (p Provider) EffectiveSlug() string { + if p.Slug != "" { + return p.Slug + } + return p.ID +} + +func (p Provider) EffectiveWireAPI() string { + if p.WireAPI != "" { + return p.WireAPI + } + if p.Protocol == ProtocolAnthropic { + return "messages" + } + return "chat_completions" +} + type Route struct { Provider Provider UpstreamModel string @@ -61,6 +84,11 @@ type Model struct { } func (m Model) Allows(principal Principal) bool { + if len(principal.AllowedModels) > 0 { + if _, ok := principal.AllowedModels[m.ID]; !ok { + return false + } + } if len(m.AllowedTenantIDs) > 0 { if _, ok := m.AllowedTenantIDs[principal.TenantID]; !ok { return false diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go index 6ab437c..e66b170 100644 --- a/internal/httpapi/api.go +++ b/internal/httpapi/api.go @@ -10,6 +10,7 @@ import ( "io" "log/slog" "net/http" + "net/url" "runtime/debug" "strings" "time" @@ -46,6 +47,7 @@ type API struct { maxBodyBytes int64 exposeMetrics bool deploymentRegion string + browserOrigin string } type Options struct { @@ -62,6 +64,7 @@ type Options struct { MaxBodyBytes int64 ExposeMetrics bool DeploymentRegion string + BrowserOrigin string } func New(options Options) *API { @@ -79,6 +82,7 @@ func New(options Options) *API { maxBodyBytes: options.MaxBodyBytes, exposeMetrics: options.ExposeMetrics, deploymentRegion: strings.ToLower(strings.TrimSpace(options.DeploymentRegion)), + browserOrigin: browserOrigin(options.BrowserOrigin), } } @@ -91,13 +95,52 @@ func (a *API) Handler() http.Handler { } a.registerInference(mux) - return a.withRequestID(a.recoverPanics(mux)) + return a.withRequestID(a.withBrowserOrigin(a.recoverPanics(mux))) } func (a *API) InferenceHandler() http.Handler { mux := http.NewServeMux() a.registerInference(mux) - return a.withRequestID(a.recoverPanics(mux)) + return a.withRequestID(a.withBrowserOrigin(a.recoverPanics(mux))) +} + +// withBrowserOrigin permits the authenticated console to call the split +// inference listener. It deliberately does not allow credentials or wildcard +// origins: the API key remains an explicit bearer credential in the request. +func (a *API) withBrowserOrigin(next http.Handler) http.Handler { + if a.browserOrigin == "" { + return next + } + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + origin := strings.TrimSpace(r.Header.Get("Origin")) + if origin == "" { + next.ServeHTTP(w, r) + return + } + if origin != a.browserOrigin { + w.WriteHeader(http.StatusForbidden) + return + } + w.Header().Set("Vary", "Origin") + w.Header().Set("Access-Control-Allow-Origin", a.browserOrigin) + w.Header().Set("Access-Control-Expose-Headers", "X-AIGW-Request-ID") + if r.Method == http.MethodOptions { + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, X-API-Key, Anthropic-Version, X-Request-ID") + w.Header().Set("Access-Control-Max-Age", "600") + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} + +func browserOrigin(raw string) string { + u, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { + return "" + } + return strings.ToLower(u.Scheme) + "://" + u.Host } func (a *API) registerInference(mux *http.ServeMux) { @@ -105,6 +148,8 @@ func (a *API) registerInference(mux *http.ServeMux) { mux.HandleFunc("GET /api/v1/models", a.openAIModels) mux.HandleFunc("POST /v1/chat/completions", a.openAIChat) mux.HandleFunc("POST /api/v1/chat/completions", a.openAIChat) + mux.HandleFunc("POST /v1/responses", a.openAIResponses) + mux.HandleFunc("POST /api/v1/responses", a.openAIResponses) mux.HandleFunc("GET /anthropic/v1/models", a.anthropicModels) mux.HandleFunc("GET /api/anthropic/v1/models", a.anthropicModels) @@ -122,6 +167,10 @@ func (a *API) openAIChat(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolOpenAI) } +func (a *API) openAIResponses(w http.ResponseWriter, r *http.Request) { + a.serveInference(w, r, domain.ProtocolOpenAIResponses) +} + func (a *API) anthropicMessages(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolAnthropic) } @@ -188,11 +237,12 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do defer lease.Release() } - model, modelErr := a.catalog.ModelForPrincipal(envelope.Model, principal) + model, providerSlug, modelErr := a.resolveModelSelector(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 } + publicModel := model.ID 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) @@ -211,9 +261,17 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do return } - routes, err := a.router.Plan(model.ID, protocol) + routes, err := a.router.PlanProvider(model.ID, protocol, providerSlug) if err != nil { - if errors.Is(err, routing.ErrNoRoute) { + if errors.Is(err, routing.ErrProviderNotFound) { + apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "provider_not_found", Message: "Requested provider is not configured for this model and API protocol"}, requestID) + } else if errors.Is(err, routing.ErrNoHealthyRoute) { + message := "All compatible providers are cooling down after retryable failures" + if providerSlug != "" { + message = "Requested provider is cooling down after retryable failures" + } + apierror.Write(w, apierror.Error{Status: http.StatusServiceUnavailable, Type: "provider_unavailable", Message: message}, requestID) + } else 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) } else { apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist"}, requestID) @@ -231,28 +289,30 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do if errors.Is(err, billing.ErrInsufficientBalance) { apierror.Write(w, apierror.Error{Status: http.StatusPaymentRequired, Type: "insufficient_balance", Message: "Account balance is insufficient"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusPaymentRequired, Success: false, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusPaymentRequired, Success: false, ErrorType: "insufficient_balance", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } if errors.Is(err, billing.ErrQuotaExceeded) { apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "monthly_quota_exceeded", Message: "Project monthly spend quota exceeded"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false, ErrorType: "monthly_quota_exceeded", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } a.logger.Error("billing_authorization_failed", "request_id", requestID, "tenant_id", principal.TenantID, "error", err) apierror.Write(w, apierror.Error{Status: http.StatusServiceUnavailable, Type: "billing_unavailable", Message: "Billing service is temporarily unavailable"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusServiceUnavailable, Success: false, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusServiceUnavailable, Success: false, ErrorType: "billing_unavailable", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } } - result, err := a.forwarder.Forward(r.Context(), protocol, requestID, body, r.Header, routes) + result, err := a.forwarder.Forward(r.Context(), protocol, requestID, model.ID, body, r.Header, routes) if err != nil { + a.logger.Warn("inference_upstream_failed", "request_id", requestID, "model", publicModel, + "protocol", protocol, "attempts", result.Attempts, "error", err) errorType := "no_provider_available" status := http.StatusBadGateway if errors.Is(err, context.Canceled) { @@ -263,7 +323,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } a.finishUsage(r, domain.UsageEvent{ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: status, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: status, Success: false, ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), }) return @@ -275,7 +335,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do apierror.Write(w, gatewayError, requestID) a.finishUsage(r, domain.UsageEvent{ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, + PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: gatewayError.Status, Success: false, ErrorType: gatewayError.Type, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), }) @@ -301,7 +361,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } a.finishUsage(r, domain.UsageEvent{ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, + PublicModel: publicModel, 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, @@ -310,7 +370,8 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do "request_id", requestID, "tenant_id", principal.TenantID, "project_id", principal.ProjectID, - "model", envelope.Model, + "model", publicModel, + "provider_slug", result.Route.Provider.EffectiveSlug(), "provider", result.Route.Provider.ID, "status", result.Response.StatusCode, "attempts", result.Attempts, @@ -323,17 +384,62 @@ func (a *API) openAIModels(w http.ResponseWriter, r *http.Request) { if !ok { return } - models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal)) + models := a.availableOpenAIModels(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": 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}) + "lifecycle": model.Lifecycle, "regions": model.Regions, "replacement_model": model.ReplacementModel, + "supported_wire_apis": modelWireAPIs(model), "providers": modelProviderDescriptors(model)}) } writeJSON(w, map[string]any{"object": "list", "data": data}) } +func modelWireAPIs(model domain.Model) []string { + seen := map[string]struct{}{} + result := make([]string, 0, len(model.Routes)) + for _, route := range model.Routes { + wireAPI := route.Provider.EffectiveWireAPI() + if _, exists := seen[wireAPI]; exists { + continue + } + seen[wireAPI] = struct{}{} + result = append(result, wireAPI) + } + return result +} + +func modelProviderDescriptors(model domain.Model) []map[string]string { + seen := map[string]struct{}{} + result := make([]map[string]string, 0, len(model.Routes)) + for _, route := range model.Routes { + slug := route.Provider.EffectiveSlug() + wireAPI := route.Provider.EffectiveWireAPI() + key := slug + "\x00" + wireAPI + if _, exists := seen[key]; exists { + continue + } + seen[key] = struct{}{} + result = append(result, map[string]string{"slug": slug, "wire_api": wireAPI}) + } + return result +} + +func (a *API) availableOpenAIModels(principal domain.Principal) []domain.Model { + combined := append(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal), a.catalog.ModelsFor(domain.ProtocolOpenAIResponses, principal)...) + seen := make(map[string]struct{}, len(combined)) + result := make([]domain.Model, 0, len(combined)) + for _, model := range combined { + if _, exists := seen[model.ID]; exists || !a.modelAvailableInRegion(model) { + continue + } + seen[model.ID] = struct{}{} + result = append(result, model) + } + return result +} + func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { principal, ok := a.authorize(w, r) if !ok { @@ -342,7 +448,7 @@ func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { 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"}) + data = append(data, map[string]any{"id": model.ID, "display_name": model.ID, "created_at": "1970-01-01T00:00:00Z", "type": "model", "providers": modelProviderDescriptors(model)}) } response := map[string]any{"data": data, "has_more": false, "first_id": nil, "last_id": nil} if len(models) > 0 { @@ -373,6 +479,43 @@ func (a *API) availableModels(models []domain.Model) []domain.Model { } return result } + +func (a *API) resolveModelSelector(selector string, principal domain.Principal) (domain.Model, string, error) { + selector = strings.TrimSpace(selector) + if model, err := a.catalog.Model(selector); err == nil { + if !model.Allows(principal) { + return domain.Model{}, "", fmt.Errorf("model %q not allowed", selector) + } + return model, "", nil + } + separator := strings.LastIndexByte(selector, ':') + if separator <= 0 || separator == len(selector)-1 { + return domain.Model{}, "", fmt.Errorf("model %q not found or not allowed", selector) + } + modelID := strings.TrimSpace(selector[:separator]) + providerSlug := strings.TrimSpace(selector[separator+1:]) + if modelID == "" || !validProviderSlug(providerSlug) { + return domain.Model{}, "", fmt.Errorf("invalid model provider selector %q", selector) + } + model, err := a.catalog.ModelForPrincipal(modelID, principal) + if err != nil { + return domain.Model{}, "", err + } + return model, providerSlug, nil +} + +func validProviderSlug(value string) bool { + if len(value) < 3 || len(value) > 64 || value[0] == '-' || value[len(value)-1] == '-' { + return false + } + for _, character := range value { + if (character < 'a' || character > 'z') && (character < '0' || character > '9') && character != '-' { + return false + } + } + return true +} + func modelHasCapability(model domain.Model, wanted string) bool { if len(model.Capabilities) == 0 { return true diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go index 9407530..cbb3967 100644 --- a/internal/httpapi/api_test.go +++ b/internal/httpapi/api_test.go @@ -20,6 +20,7 @@ import ( "aigw/internal/config" "aigw/internal/domain" "aigw/internal/provider" + "aigw/internal/providerhealth" "aigw/internal/routing" "aigw/internal/telemetry" ) @@ -28,6 +29,49 @@ type captureUsageSink struct { events chan domain.UsageEvent } +func TestInferenceBrowserOriginCORS(t *testing.T) { + api := New(Options{ + BrowserOrigin: "https://console.example.test/admin/", + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + Metrics: &telemetry.Metrics{}, + }) + handler := api.InferenceHandler() + + request := httptest.NewRequest(http.MethodOptions, "/v1/responses", nil) + request.Header.Set("Origin", "https://console.example.test") + request.Header.Set("Access-Control-Request-Method", http.MethodPost) + request.Header.Set("Access-Control-Request-Headers", "authorization,content-type") + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != http.StatusNoContent { + t.Fatalf("preflight status = %d, want %d", response.Code, http.StatusNoContent) + } + if got := response.Header().Get("Access-Control-Allow-Origin"); got != "https://console.example.test" { + t.Fatalf("allow origin = %q", got) + } + if response.Header().Get("Access-Control-Allow-Credentials") != "" { + t.Fatal("inference CORS must not allow browser credentials") + } + if !strings.Contains(response.Header().Get("Access-Control-Expose-Headers"), "X-AIGW-Request-ID") { + t.Fatal("request ID is not exposed to the developer console") + } + + blocked := httptest.NewRequest(http.MethodOptions, "/v1/responses", nil) + blocked.Header.Set("Origin", "https://attacker.example") + blockedResponse := httptest.NewRecorder() + handler.ServeHTTP(blockedResponse, blocked) + if blockedResponse.Code != http.StatusForbidden { + t.Fatalf("untrusted preflight status = %d, want %d", blockedResponse.Code, http.StatusForbidden) + } + blockedRequest := httptest.NewRequest(http.MethodGet, "/v1/models", nil) + blockedRequest.Header.Set("Origin", "https://attacker.example") + blockedActual := httptest.NewRecorder() + handler.ServeHTTP(blockedActual, blockedRequest) + if blockedActual.Code != http.StatusForbidden { + t.Fatalf("untrusted actual status = %d, want %d", blockedActual.Code, http.StatusForbidden) + } +} + type fakeBillingMeter struct { authorizeErr error settled chan domain.UsageEvent @@ -134,6 +178,199 @@ func TestProxyFailsOverBeforeWritingResponse(t *testing.T) { } } +func TestProviderSelectorPinsRouteAndKeepsCanonicalUsageModel(t *testing.T) { + var primaryCalls atomic.Int64 + primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + primaryCalls.Add(1) + w.WriteHeader(http.StatusServiceUnavailable) + })) + defer primary.Close() + var backupCalls atomic.Int64 + backup := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + backupCalls.Add(1) + var request map[string]any + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + if request["model"] != "backup-model" { + t.Errorf("upstream model = %v, want backup-model", request["model"]) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer backup.Close() + + gateway, sink := newTestGateway(t, + []config.ProviderConfig{ + {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: primary.URL + "/v1", APIKey: "one"}, + {ID: "backup", Protocol: domain.ProtocolOpenAI, BaseURL: backup.URL + "/v1", APIKey: "two"}, + }, + []config.RouteConfig{ + {Provider: "primary", UpstreamModel: "primary-model", Priority: 0, Weight: 1}, + {Provider: "backup", UpstreamModel: "backup-model", Priority: 10, Weight: 1}, + }, + ) + defer gateway.Close() + + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model:backup","messages":[{"role":"user","content":"hello"}]}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + payload, _ := io.ReadAll(response.Body) + t.Fatalf("status = %d: %s", response.StatusCode, payload) + } + if primaryCalls.Load() != 0 || backupCalls.Load() != 1 { + t.Fatalf("pinned routing calls: primary=%d backup=%d", primaryCalls.Load(), backupCalls.Load()) + } + event := <-sink.events + if event.PublicModel != "public/model" || event.ProviderID != "backup" || event.Attempts != 1 { + t.Fatalf("unexpected pinned usage event: %+v", event) + } +} + +func TestProviderSelectorRejectsUnknownProviderWithoutCallingUpstream(t *testing.T) { + var calls atomic.Int64 + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + calls.Add(1) + w.WriteHeader(http.StatusOK) + })) + defer upstream.Close() + gateway, _ := newTestGateway(t, + []config.ProviderConfig{{ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "one"}}, + []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}, + ) + defer gateway.Close() + + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model:missing","messages":[{"role":"user","content":"hello"}]}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + raw, err := io.ReadAll(response.Body) + if err != nil { + t.Fatal(err) + } + var payload struct { + Error struct { + Type string `json:"type"` + } `json:"error"` + } + if err := json.Unmarshal(raw, &payload); err != nil { + t.Fatal(err) + } + if response.StatusCode != http.StatusNotFound || payload.Error.Type != "provider_not_found" || calls.Load() != 0 { + t.Fatalf("status=%d type=%q calls=%d", response.StatusCode, payload.Error.Type, calls.Load()) + } +} + +func TestResolveModelSelectorChecksBaseModelAllowlistAndPreservesExactColonID(t *testing.T) { + modelCatalog := catalog.NewModels([]domain.Model{ + {ID: "public/model"}, + {ID: "exact:model", AllowedKeyIDs: map[string]struct{}{"other-key": {}}}, + {ID: "exact"}, + }) + api := &API{catalog: modelCatalog} + principal := domain.Principal{KeyID: "key-1", AllowedModels: map[string]struct{}{"public/model": {}, "exact": {}}} + model, providerSlug, err := api.resolveModelSelector("public/model:backup", principal) + if err != nil || model.ID != "public/model" || providerSlug != "backup" { + t.Fatalf("base allowlist selector: model=%+v provider=%q err=%v", model, providerSlug, err) + } + if _, _, err := api.resolveModelSelector("exact:model", principal); err == nil { + t.Fatal("an unauthorized exact colon model ID must not be reinterpreted as a provider selector") + } +} + +func TestModelsListPublishesProviderSlugsWithoutUpstreamDetails(t *testing.T) { + gateway, _ := newTestGateway(t, + []config.ProviderConfig{{ID: "openai-primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://secret-upstream.example/v1", APIKey: "secret"}}, + []config.RouteConfig{{Provider: "openai-primary", UpstreamModel: "secret-upstream-model", Weight: 1}}, + ) + defer gateway.Close() + request, _ := http.NewRequest(http.MethodGet, gateway.URL+"/v1/models", nil) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + raw, err := io.ReadAll(response.Body) + if err != nil { + t.Fatal(err) + } + var payload struct { + Data []struct { + ID string `json:"id"` + Providers []struct { + Slug string `json:"slug"` + WireAPI string `json:"wire_api"` + } `json:"providers"` + } `json:"data"` + } + if err := json.Unmarshal(raw, &payload); err != nil { + t.Fatal(err) + } + if response.StatusCode != http.StatusOK || len(payload.Data) != 1 || len(payload.Data[0].Providers) != 1 || payload.Data[0].Providers[0].Slug != "openai-primary" { + t.Fatalf("unexpected models payload: status=%d payload=%+v", response.StatusCode, payload) + } + if strings.Contains(string(raw), "secret-upstream") { + t.Fatalf("models payload leaked upstream detail: %s", raw) + } +} + +func TestCircuitBreakerSkipsFailingProviderOnSubsequentRequests(t *testing.T) { + var primaryCalls atomic.Int64 + primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + primaryCalls.Add(1) + w.WriteHeader(http.StatusServiceUnavailable) + })) + defer primary.Close() + var fallbackCalls atomic.Int64 + fallback := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fallbackCalls.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer fallback.Close() + + gateway, sink := newTestGateway(t, + []config.ProviderConfig{ + {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: primary.URL + "/v1", APIKey: "one"}, + {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: fallback.URL + "/v1", APIKey: "two"}, + }, + []config.RouteConfig{ + {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1}, + {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1}, + }, + ) + defer gateway.Close() + + for requestNumber := range 4 { + response := postOpenAI(t, gateway.URL, false) + _, _ = io.Copy(io.Discard, response.Body) + _ = response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("request %d status = %d, want %d", requestNumber+1, response.StatusCode, http.StatusOK) + } + select { + case <-sink.events: + case <-time.After(time.Second): + t.Fatalf("request %d did not emit usage", requestNumber+1) + } + } + if primaryCalls.Load() != 3 { + t.Fatalf("primary calls = %d, want 3 before circuit opens", primaryCalls.Load()) + } + if fallbackCalls.Load() != 4 { + t.Fatalf("fallback calls = %d, want 4", fallbackCalls.Load()) + } +} + func TestSSEIsFlushedBeforeUpstreamCompletes(t *testing.T) { release := make(chan struct{}) upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -194,6 +431,45 @@ func TestAnthropicHeadersAndPath(t *testing.T) { } } +func TestOpenAIResponsesProxyRewritesModelAndEmitsUsage(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/responses" { + t.Errorf("unexpected path: %s", r.URL.Path) + } + var request map[string]any + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + if request["model"] != "gpt-upstream" || request["input"] != "hello" { + t.Errorf("unexpected Responses request: %+v", request) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"id":"resp_1","object":"response","status":"completed","model":"gpt-upstream","output":[],"usage":{"input_tokens":9,"output_tokens":4,"total_tokens":13}}`) + })) + defer upstream.Close() + + gateway, sink := newTestGateway(t, []config.ProviderConfig{{ + ID: "responses", Protocol: domain.ProtocolOpenAI, WireAPI: "responses", BaseURL: upstream.URL, APIKey: "upstream-secret", + }}, []config.RouteConfig{{Provider: "responses", UpstreamModel: "gpt-upstream", Weight: 1}}) + defer gateway.Close() + + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/responses", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + payload, _ := io.ReadAll(response.Body) + t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) + } + event := <-sink.events + if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || !event.UsageReported { + t.Fatalf("unexpected Responses usage event: %+v", event) + } +} + func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -276,12 +552,13 @@ func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig, } metrics := &telemetry.Metrics{} modelCatalog := catalog.New(cfg) + routeHealth := providerhealth.New(providerhealth.Options{}) sink := &captureUsageSink{events: make(chan domain.UsageEvent, 10)} api := New(Options{ Authenticator: authenticator, Catalog: modelCatalog, - Router: routing.New(modelCatalog), - Forwarder: provider.New(cfg.UpstreamHTTP, metrics), + Router: routing.New(modelCatalog, routeHealth), + Forwarder: provider.New(cfg.UpstreamHTTP, metrics, routeHealth), UsageSink: sink, BillingMeter: meter, Metrics: metrics, diff --git a/internal/operations/operations.go b/internal/operations/operations.go index e566488..502c7f8 100644 --- a/internal/operations/operations.go +++ b/internal/operations/operations.go @@ -12,11 +12,29 @@ import ( ) type Handler struct { - Store *controlplane.Store - Manager *controlplane.Manager - Billing *billing.Service - Metrics *telemetry.Metrics - MaxSnapshotAge time.Duration + Store readinessStore + Manager *controlplane.Manager + Billing readinessBilling + Metrics *telemetry.Metrics + MaxSnapshotAge time.Duration + ReadinessTimeout time.Duration +} + +type readinessStore interface { + Ping(context.Context) error + MailQueueStatus(context.Context) (controlplane.MailQueueStatus, error) +} + +type readinessBilling interface { + Ping(context.Context) error + SettlementQueueStatus(context.Context) (billing.SettlementQueueStatus, error) + OperationalStatus(context.Context) (billing.OperationalStatus, error) +} + +type readinessResult struct { + name string + value map[string]any + ready bool } func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { @@ -32,17 +50,47 @@ func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { http.NotFound(w, r) return } - ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second) + + timeout := h.ReadinessTimeout + if timeout <= 0 { + timeout = 1500 * time.Millisecond + } + // Readiness probes are independently bounded below. Some container health + // clients close their request side aggressively after sending the GET; do + // not let that client lifecycle make every backend look unavailable. + ctx, cancel := context.WithTimeout(context.Background(), timeout) defer cancel() checks := map[string]any{} ready := true + results := make(chan readinessResult, 5) + expected := map[string]struct{}{} + launch := func(name string, check func(context.Context) readinessResult) { + expected[name] = struct{}{} + go func() { + result := check(ctx) + result.name = name + results <- result + }() + } + 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"} - } + launch("postgres", func(ctx context.Context) readinessResult { + if err := h.Store.Ping(ctx); err != nil { + return failedResult(err) + } + return okResult() + }) + launch("mail_queue", func(ctx context.Context) readinessResult { + mail, err := h.Store.MailQueueStatus(ctx) + if err != nil { + return failedResult(err) + } + mailOK := mail.Failed == 0 + if mail.OldestPending != nil && time.Since(*mail.OldestPending) > 10*time.Minute { + mailOK = false + } + return readinessResult{ready: mailOK, value: map[string]any{"status": status(mailOK), "details": mail}} + }) } if h.Manager != nil { healthyAt := h.Manager.LastHealthyAt() @@ -65,61 +113,60 @@ func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { 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 { + launch("billing_postgres", func(ctx context.Context) readinessResult { + if err := h.Billing.Ping(ctx); err != nil { + return failedResult(err) + } + return okResult() + }) + launch("settlement_queue", func(ctx context.Context) readinessResult { + queue, err := h.Billing.SettlementQueueStatus(ctx) + if err != nil { + return failedResult(err) + } 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) } + return readinessResult{ready: queueOK, value: map[string]any{"status": status(queueOK), "backlog": backlog, "spool_records": queue.SpoolRecords, "oldest_pending": queue.OldestPending}} + }) + launch("billing_operations", func(ctx context.Context) readinessResult { 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) - } + return failedResult(err) } - } - if h.Store != nil { - mail, err := h.Store.MailQueueStatus(ctx) - if err != nil { - checks["mail_queue"] = map[string]any{"status": "failed", "error": err.Error()} + billingOK := billingHealth.Ready(time.Now().UTC()) + if h.Metrics != nil { + h.Metrics.SetStripeOperations(billingHealth) + } + return readinessResult{ready: billingOK, value: map[string]any{"status": status(billingOK), "details": billingHealth}} + }) + } + + for len(expected) > 0 { + select { + case result := <-results: + if _, ok := expected[result.name]; !ok { + continue + } + delete(expected, result.name) + checks[result.name] = result.value + if !result.ready { 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 - } } + case <-ctx.Done(): + for name := range expected { + checks[name] = map[string]any{"status": "failed", "error": "check timed out"} + } + ready = false + expected = map[string]struct{}{} } } + if h.Metrics != nil { h.Metrics.SetReady(ready) } @@ -130,12 +177,25 @@ func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { write(w, code, map[string]any{"status": status(ready), "checks": checks}) } +func okResult() readinessResult { + return readinessResult{ready: true, value: map[string]any{"status": "ok"}} +} + +func failedResult(err error) readinessResult { + message := err.Error() + if err == context.Canceled || err == context.DeadlineExceeded { + message = "check timed out" + } + return readinessResult{value: map[string]any{"status": "failed", "error": message}} +} + 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) diff --git a/internal/operations/operations_test.go b/internal/operations/operations_test.go new file mode 100644 index 0000000..6936b1a --- /dev/null +++ b/internal/operations/operations_test.go @@ -0,0 +1,89 @@ +package operations + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "aigw/internal/billing" + "aigw/internal/controlplane" +) + +type healthyStore struct{} + +func (healthyStore) Ping(context.Context) error { return nil } +func (healthyStore) MailQueueStatus(context.Context) (controlplane.MailQueueStatus, error) { + return controlplane.MailQueueStatus{}, nil +} + +type slowBilling struct{} + +func (slowBilling) Ping(context.Context) error { return nil } +func (slowBilling) SettlementQueueStatus(context.Context) (billing.SettlementQueueStatus, error) { + return billing.SettlementQueueStatus{}, nil +} +func (slowBilling) OperationalStatus(ctx context.Context) (billing.OperationalStatus, error) { + <-ctx.Done() + return billing.OperationalStatus{}, ctx.Err() +} + +func TestReadinessTimeoutDoesNotMislabelCompletedChecks(t *testing.T) { + handler := Handler{Store: healthyStore{}, Billing: slowBilling{}, ReadinessTimeout: 20 * time.Millisecond} + request := httptest.NewRequest(http.MethodGet, "/readyz", nil) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, request) + if recorder.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want 503", recorder.Code) + } + var payload struct { + Checks map[string]struct { + Status string `json:"status"` + Error string `json:"error"` + } `json:"checks"` + } + if err := json.NewDecoder(recorder.Body).Decode(&payload); err != nil { + t.Fatal(err) + } + for _, name := range []string{"postgres", "mail_queue", "billing_postgres", "settlement_queue"} { + if payload.Checks[name].Status != "ok" { + t.Fatalf("%s status = %q, want ok; payload=%s", name, payload.Checks[name].Status, recorder.Body.String()) + } + } + if payload.Checks["billing_operations"].Status != "failed" || payload.Checks["billing_operations"].Error != "check timed out" { + t.Fatalf("billing operations = %+v", payload.Checks["billing_operations"]) + } +} + +func TestMailQueueIsCheckedWithoutBilling(t *testing.T) { + handler := Handler{Store: healthyStore{}, ReadinessTimeout: time.Second} + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/readyz", nil)) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, body=%s", recorder.Code, recorder.Body.String()) + } + var payload struct { + Checks map[string]any `json:"checks"` + } + if err := json.NewDecoder(recorder.Body).Decode(&payload); err != nil { + t.Fatal(err) + } + if _, ok := payload.Checks["mail_queue"]; !ok { + t.Fatal("mail_queue check missing when billing is disabled") + } +} + +func TestReadinessChecksDoNotInheritCanceledClientContext(t *testing.T) { + handler := Handler{Store: healthyStore{}, ReadinessTimeout: time.Second} + request := httptest.NewRequest(http.MethodGet, "/readyz", nil) + ctx, cancel := context.WithCancel(request.Context()) + cancel() + request = request.WithContext(ctx) + recorder := httptest.NewRecorder() + handler.ServeHTTP(recorder, request) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, body=%s", recorder.Code, recorder.Body.String()) + } +} diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index 0d40bb1..d7d850d 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -14,6 +14,7 @@ import ( "aigw/internal/config" "aigw/internal/domain" + "aigw/internal/providerhealth" "aigw/internal/telemetry" ) @@ -26,9 +27,10 @@ type Result struct { type Forwarder struct { client *http.Client metrics *telemetry.Metrics + health *providerhealth.Tracker } -func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder { +func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics, trackers ...*providerhealth.Tracker) *Forwarder { transport := &http.Transport{ Proxy: http.ProxyFromEnvironment, DialContext: (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext, @@ -40,31 +42,41 @@ func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder { ResponseHeaderTimeout: time.Duration(cfg.ResponseHeaderTimeoutSecs) * time.Second, ExpectContinueTimeout: time.Second, } - return &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} + forwarder := &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} + if len(trackers) > 0 { + forwarder.health = trackers[0] + } + return forwarder } -func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID string, originalBody []byte, sourceHeaders http.Header, routes []domain.Route) (Result, error) { +func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID, modelID string, originalBody []byte, sourceHeaders http.Header, routes []domain.Route) (Result, error) { var lastErr error for i, route := range routes { if err := ctx.Err(); err != nil { return Result{Attempts: i}, err } - body, err := rewriteRequest(originalBody, route.UpstreamModel, protocol) + body, err := rewriteRequestWithWireAPI(originalBody, route.UpstreamModel, protocol, route.Provider.EffectiveWireAPI()) if err != nil { return Result{Attempts: i}, err } - request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider.BaseURL, protocol), bytes.NewReader(body)) + request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider, protocol), bytes.NewReader(body)) if err != nil { return Result{Attempts: i}, fmt.Errorf("build upstream request: %w", err) } setHeaders(request.Header, sourceHeaders, route.Provider, protocol, requestID) f.metrics.UpstreamAttempt() + attemptStarted := time.Now() response, err := f.client.Do(request) if err != nil { + if ctx.Err() != nil { + return Result{Attempts: i + 1}, ctx.Err() + } + f.observe(modelID, route, 0, time.Since(attemptStarted), true) lastErr = err continue } attempts := i + 1 + f.observe(modelID, route, response.StatusCode, time.Since(attemptStarted), retryableStatus(response.StatusCode)) if retryableStatus(response.StatusCode) && attempts < len(routes) { _, _ = io.CopyN(io.Discard, response.Body, 8<<10) _ = response.Body.Close() @@ -79,18 +91,34 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque return Result{Attempts: len(routes)}, lastErr } +func (f *Forwarder) observe(modelID string, route domain.Route, statusCode int, latency time.Duration, failed bool) { + if f.health == nil { + return + } + f.health.Observe(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}, + providerhealth.Observation{StatusCode: statusCode, Latency: latency, Failed: failed}) +} + func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { - return rewriteRequest(body, upstreamModel, "") + return rewriteRequest(body, upstreamModel, domain.ProtocolOpenAI) } func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) ([]byte, error) { + wireAPI := "chat_completions" + if protocol == domain.ProtocolAnthropic { + wireAPI = "messages" + } + return rewriteRequestWithWireAPI(body, upstreamModel, protocol, wireAPI) +} + +func rewriteRequestWithWireAPI(body []byte, upstreamModel string, protocol domain.Protocol, wireAPI string) ([]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 { + if protocol == domain.ProtocolOpenAI && wireAPI == "chat_completions" { var stream bool _ = json.Unmarshal(object["stream"], &stream) if stream { @@ -110,12 +138,18 @@ func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) return result, nil } -func endpointURL(baseURL string, protocol domain.Protocol) string { - baseURL = strings.TrimRight(baseURL, "/") - if protocol == domain.ProtocolAnthropic { +func endpointURL(provider domain.Provider, _ domain.Protocol) string { + baseURL := strings.TrimRight(provider.BaseURL, "/") + switch provider.EffectiveWireAPI() { + case "responses": + return baseURL + "/responses" + case "messages": return baseURL + "/messages" + case "chat_completions": + fallthrough + default: + return baseURL + "/chat/completions" } - return baseURL + "/chat/completions" } func setHeaders(target, source http.Header, provider domain.Provider, protocol domain.Protocol, requestID string) { diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go index 2e9b4a1..e9751ce 100644 --- a/internal/provider/forwarder_test.go +++ b/internal/provider/forwarder_test.go @@ -37,3 +37,24 @@ func TestRewriteRequestDoesNotAddStreamOptionsToAnthropic(t *testing.T) { t.Fatalf("unexpected OpenAI stream options in Anthropic request: %s", result) } } + +func TestResponsesWireAPIUsesResponsesEndpointWithoutChatStreamOptions(t *testing.T) { + result, err := rewriteRequestWithWireAPI([]byte(`{"model":"public/model","input":"hello","stream":true}`), "gpt-upstream", domain.ProtocolOpenAIResponses, "responses") + if err != nil { + t.Fatal(err) + } + var body map[string]json.RawMessage + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if string(body["model"]) != `"gpt-upstream"` { + t.Fatalf("model was not rewritten: %s", result) + } + if _, exists := body["stream_options"]; exists { + t.Fatalf("Responses request contains Chat Completions stream options: %s", result) + } + provider := domain.Provider{BaseURL: "https://example.test", Protocol: domain.ProtocolOpenAI, WireAPI: "responses"} + if got := endpointURL(provider, domain.ProtocolOpenAIResponses); got != "https://example.test/responses" { + t.Fatalf("endpoint URL = %q", got) + } +} diff --git a/internal/providerhealth/tracker.go b/internal/providerhealth/tracker.go new file mode 100644 index 0000000..a5af2b2 --- /dev/null +++ b/internal/providerhealth/tracker.go @@ -0,0 +1,200 @@ +package providerhealth + +import ( + "sort" + "sync" + "time" +) + +const recentWindow = 100 + +type RouteKey struct { + ModelID string + ProviderID string + WireAPI string +} + +type Observation struct { + StatusCode int + Latency time.Duration + Failed bool + ObservedAt time.Time +} + +type Status struct { + ModelID string `json:"model_id"` + ProviderID string `json:"provider_id"` + WireAPI string `json:"wire_api"` + State string `json:"state"` + Attempts uint64 `json:"attempts"` + RecentSamples int `json:"recent_samples"` + AvailabilityPercent float64 `json:"availability_percent"` + HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"` + ConsecutiveFailures uint64 `json:"consecutive_failures"` + LastStatusCode int `json:"last_status_code,omitempty"` + LastObservedAt *time.Time `json:"last_observed_at,omitempty"` + LastHealthyAt *time.Time `json:"last_healthy_at,omitempty"` + CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"` +} + +type Options struct { + FailureThreshold uint64 + OpenDuration time.Duration + Now func() time.Time +} + +type Tracker struct { + states sync.Map + failureThreshold uint64 + openDuration time.Duration + now func() time.Time +} + +type routeState struct { + mu sync.RWMutex + attempts uint64 + consecutiveFailures uint64 + lastStatusCode int + lastObservedAt time.Time + lastHealthyAt time.Time + openUntil time.Time + headerLatencyEWMA float64 + recent [recentWindow]bool + recentCount int + recentPosition int + recentHealthy int +} + +func New(options Options) *Tracker { + if options.FailureThreshold == 0 { + options.FailureThreshold = 3 + } + if options.OpenDuration <= 0 { + options.OpenDuration = 30 * time.Second + } + if options.Now == nil { + options.Now = time.Now + } + return &Tracker{failureThreshold: options.FailureThreshold, openDuration: options.OpenDuration, now: options.Now} +} + +func (t *Tracker) Observe(key RouteKey, observation Observation) { + if t == nil || key.ProviderID == "" { + return + } + if observation.ObservedAt.IsZero() { + observation.ObservedAt = t.now() + } + value, _ := t.states.LoadOrStore(key, &routeState{}) + state := value.(*routeState) + state.mu.Lock() + defer state.mu.Unlock() + + state.attempts++ + state.lastStatusCode = observation.StatusCode + state.lastObservedAt = observation.ObservedAt + if observation.Latency > 0 { + latency := float64(observation.Latency.Milliseconds()) + if latency < 1 { + latency = 1 + } + if state.headerLatencyEWMA == 0 { + state.headerLatencyEWMA = latency + } else { + state.headerLatencyEWMA = state.headerLatencyEWMA*0.8 + latency*0.2 + } + } + state.addRecent(!observation.Failed) + if observation.Failed { + state.consecutiveFailures++ + if state.consecutiveFailures >= t.failureThreshold { + state.openUntil = observation.ObservedAt.Add(t.openDuration) + } + return + } + state.consecutiveFailures = 0 + state.openUntil = time.Time{} + state.lastHealthyAt = observation.ObservedAt +} + +func (s *routeState) addRecent(healthy bool) { + if s.recentCount == recentWindow { + if s.recent[s.recentPosition] { + s.recentHealthy-- + } + } else { + s.recentCount++ + } + s.recent[s.recentPosition] = healthy + if healthy { + s.recentHealthy++ + } + s.recentPosition = (s.recentPosition + 1) % recentWindow +} + +func (t *Tracker) CircuitOpen(key RouteKey) bool { + if t == nil { + return false + } + value, ok := t.states.Load(key) + if !ok { + return false + } + state := value.(*routeState) + state.mu.RLock() + defer state.mu.RUnlock() + return state.openUntil.After(t.now()) +} + +func (t *Tracker) Snapshot() []Status { + if t == nil { + return []Status{} + } + now := t.now() + result := make([]Status, 0) + t.states.Range(func(rawKey, rawState any) bool { + key := rawKey.(RouteKey) + state := rawState.(*routeState) + state.mu.RLock() + item := statusFromState(key, state, now) + state.mu.RUnlock() + result = append(result, item) + return true + }) + sort.Slice(result, func(i, j int) bool { + if result[i].ModelID != result[j].ModelID { + return result[i].ModelID < result[j].ModelID + } + return result[i].ProviderID < result[j].ProviderID + }) + return result +} + +func statusFromState(key RouteKey, state *routeState, now time.Time) Status { + item := Status{ModelID: key.ModelID, ProviderID: key.ProviderID, WireAPI: key.WireAPI, + Attempts: state.attempts, RecentSamples: state.recentCount, HeaderLatencyEWMA: int64(state.headerLatencyEWMA + 0.5), + ConsecutiveFailures: state.consecutiveFailures, LastStatusCode: state.lastStatusCode} + if state.recentCount > 0 { + item.AvailabilityPercent = float64(state.recentHealthy) / float64(state.recentCount) * 100 + } + if !state.lastObservedAt.IsZero() { + value := state.lastObservedAt + item.LastObservedAt = &value + } + if !state.lastHealthyAt.IsZero() { + value := state.lastHealthyAt + item.LastHealthyAt = &value + } + if state.openUntil.After(now) { + value := state.openUntil + item.CircuitOpenUntil = &value + item.State = "open" + } else if state.attempts == 0 { + item.State = "unknown" + } else if state.consecutiveFailures > 0 || (state.recentCount >= 5 && item.AvailabilityPercent < 95) { + item.State = "degraded" + } else { + item.State = "healthy" + } + return item +} diff --git a/internal/providerhealth/tracker_test.go b/internal/providerhealth/tracker_test.go new file mode 100644 index 0000000..73bd6d5 --- /dev/null +++ b/internal/providerhealth/tracker_test.go @@ -0,0 +1,50 @@ +package providerhealth + +import ( + "sync" + "testing" + "time" +) + +func TestTrackerOpensAndRecoversCircuit(t *testing.T) { + now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) + tracker := New(Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }}) + key := RouteKey{ModelID: "openai/test", ProviderID: "primary", WireAPI: "responses"} + for _, status := range []int{503, 429, 504} { + tracker.Observe(key, Observation{StatusCode: status, Latency: 100 * time.Millisecond, Failed: true}) + } + if !tracker.CircuitOpen(key) { + t.Fatal("circuit did not open after consecutive retryable failures") + } + snapshot := tracker.Snapshot() + if len(snapshot) != 1 || snapshot[0].State != "open" || snapshot[0].AvailabilityPercent != 0 || snapshot[0].CircuitOpenUntil == nil { + t.Fatalf("unexpected open snapshot: %+v", snapshot) + } + now = now.Add(31 * time.Second) + if tracker.CircuitOpen(key) { + t.Fatal("circuit did not permit a recovery attempt after cooldown") + } + tracker.Observe(key, Observation{StatusCode: 200, Latency: 50 * time.Millisecond}) + snapshot = tracker.Snapshot() + if snapshot[0].State != "healthy" || snapshot[0].ConsecutiveFailures != 0 || snapshot[0].LastHealthyAt == nil || snapshot[0].HeaderLatencyEWMA < 50 { + t.Fatalf("unexpected recovered snapshot: %+v", snapshot[0]) + } +} + +func TestTrackerConcurrentObservations(t *testing.T) { + tracker := New(Options{FailureThreshold: 1000}) + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "chat_completions"} + var group sync.WaitGroup + for index := range 200 { + group.Add(1) + go func(failed bool) { + defer group.Done() + tracker.Observe(key, Observation{StatusCode: 200, Latency: time.Millisecond, Failed: failed}) + }(index%2 == 0) + } + group.Wait() + snapshot := tracker.Snapshot() + if len(snapshot) != 1 || snapshot[0].Attempts != 200 || snapshot[0].RecentSamples != recentWindow || snapshot[0].AvailabilityPercent < 0 || snapshot[0].AvailabilityPercent > 100 { + t.Fatalf("unexpected concurrent snapshot: %+v", snapshot) + } +} diff --git a/internal/routing/router.go b/internal/routing/router.go index 53e5261..37de633 100644 --- a/internal/routing/router.go +++ b/internal/routing/router.go @@ -9,31 +9,65 @@ import ( "aigw/internal/catalog" "aigw/internal/domain" + "aigw/internal/providerhealth" ) -var ErrNoRoute = errors.New("no compatible upstream route") +var ( + ErrNoRoute = errors.New("no compatible upstream route") + ErrNoHealthyRoute = errors.New("all compatible upstream routes have open circuits") + ErrProviderNotFound = errors.New("requested provider is not configured for this model and protocol") +) type Router struct { catalog *catalog.Catalog + health *providerhealth.Tracker counters sync.Map } -func New(catalog *catalog.Catalog) *Router { - return &Router{catalog: catalog} +func New(catalog *catalog.Catalog, trackers ...*providerhealth.Tracker) *Router { + router := &Router{catalog: catalog} + if len(trackers) > 0 { + router.health = trackers[0] + } + return router } func (r *Router) Plan(modelID string, protocol domain.Protocol) ([]domain.Route, error) { + return r.plan(modelID, protocol, "") +} + +func (r *Router) PlanProvider(modelID string, protocol domain.Protocol, providerSlug string) ([]domain.Route, error) { + return r.plan(modelID, protocol, providerSlug) +} + +func (r *Router) plan(modelID string, protocol domain.Protocol, providerSlug string) ([]domain.Route, error) { model, err := r.catalog.Model(modelID) if err != nil { return nil, err } routes := make([]domain.Route, 0, len(model.Routes)) + compatible := 0 + matched := 0 for _, route := range model.Routes { - if route.Provider.Protocol == protocol { + if protocolCompatible(route.Provider, protocol) { + compatible++ + if providerSlug != "" && route.Provider.EffectiveSlug() != providerSlug { + continue + } + matched++ + if r.health != nil && r.health.CircuitOpen(providerhealth.RouteKey{ModelID: model.ID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}) { + continue + } routes = append(routes, route) } } if len(routes) == 0 { + if providerSlug != "" && matched == 0 { + return nil, ErrProviderNotFound + } + if compatible > 0 { + return nil, ErrNoHealthyRoute + } return nil, ErrNoRoute } @@ -50,6 +84,19 @@ func (r *Router) Plan(modelID string, protocol domain.Protocol) ([]domain.Route, return result, nil } +func protocolCompatible(provider domain.Provider, requestProtocol domain.Protocol) bool { + switch requestProtocol { + case domain.ProtocolOpenAI: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions" + case domain.ProtocolOpenAIResponses: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses" + case domain.ProtocolAnthropic: + return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages" + default: + return false + } +} + func (r *Router) rotate(modelID string, protocol domain.Protocol, routes []domain.Route) []domain.Route { if len(routes) < 2 { return append([]domain.Route(nil), routes...) diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go index 62dc656..2ad4685 100644 --- a/internal/routing/router_test.go +++ b/internal/routing/router_test.go @@ -1,11 +1,14 @@ package routing import ( + "errors" "testing" + "time" "aigw/internal/catalog" "aigw/internal/config" "aigw/internal/domain" + "aigw/internal/providerhealth" ) func TestPlanHonorsPriorityAndProtocol(t *testing.T) { @@ -34,6 +37,40 @@ func TestPlanHonorsPriorityAndProtocol(t *testing.T) { } } +func TestPlanSkipsOpenCircuitAndRecoversAfterCooldown(t *testing.T) { + now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) + health := providerhealth.New(providerhealth.Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }}) + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}, + {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: "https://fallback.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1}, + {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1}, + }}}, + } + for range 3 { + health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "primary", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true}) + } + router := New(catalog.New(cfg), health) + plan, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil || len(plan) != 1 || plan[0].Provider.ID != "fallback" { + t.Fatalf("open primary was not skipped: plan=%+v err=%v", plan, err) + } + for range 3 { + health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "fallback", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true}) + } + if _, err := router.Plan("public/model", domain.ProtocolOpenAI); !errors.Is(err, ErrNoHealthyRoute) { + t.Fatalf("Plan() error = %v, want ErrNoHealthyRoute", err) + } + now = now.Add(31 * time.Second) + plan, err = router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil || len(plan) != 2 || plan[0].Provider.ID != "primary" { + t.Fatalf("routes did not recover after cooldown: plan=%+v err=%v", plan, err) + } +} + func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { cfg := config.Config{ Providers: []config.ProviderConfig{ @@ -58,3 +95,60 @@ func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { t.Fatalf("unexpected weighted distribution: %+v", counts) } } + +func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) { + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "chat", Protocol: domain.ProtocolOpenAI, WireAPI: "chat_completions", BaseURL: "https://chat.test", APIKey: "one"}, + {ID: "responses", Protocol: domain.ProtocolOpenAI, WireAPI: "responses", BaseURL: "https://responses.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "chat", UpstreamModel: "chat-model", Weight: 1}, + {Provider: "responses", UpstreamModel: "responses-model", Weight: 1}, + }}}, + } + router := New(catalog.New(cfg)) + chat, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil || len(chat) != 1 || chat[0].Provider.ID != "chat" { + t.Fatalf("unexpected Chat plan: %+v err=%v", chat, err) + } + responses, err := router.Plan("public/model", domain.ProtocolOpenAIResponses) + if err != nil || len(responses) != 1 || responses[0].Provider.ID != "responses" { + t.Fatalf("unexpected Responses plan: %+v err=%v", responses, err) + } +} + +func TestPlanProviderPinsWithoutFallbackToOtherProviders(t *testing.T) { + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}, + {ID: "backup", Protocol: domain.ProtocolOpenAI, BaseURL: "https://backup.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "primary", UpstreamModel: "primary-model", Priority: 0, Weight: 1}, + {Provider: "backup", UpstreamModel: "backup-model", Priority: 10, Weight: 1}, + }}}, + } + router := New(catalog.New(cfg)) + plan, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "backup") + if err != nil || len(plan) != 1 || plan[0].Provider.ID != "backup" { + t.Fatalf("unexpected pinned plan: %+v err=%v", plan, err) + } + if _, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "missing"); !errors.Is(err, ErrProviderNotFound) { + t.Fatalf("missing provider error = %v, want ErrProviderNotFound", err) + } +} + +func TestPlanProviderHonorsCircuitBreaker(t *testing.T) { + now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) + health := providerhealth.New(providerhealth.Options{FailureThreshold: 1, OpenDuration: time.Minute, Now: func() time.Time { return now }}) + cfg := config.Config{ + Providers: []config.ProviderConfig{{ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}}, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}}}, + } + health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "primary", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true}) + router := New(catalog.New(cfg), health) + if _, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "primary"); !errors.Is(err, ErrNoHealthyRoute) { + t.Fatalf("open pinned provider error = %v, want ErrNoHealthyRoute", err) + } +} diff --git a/internal/usage/observer.go b/internal/usage/observer.go index 808781f..71c1379 100644 --- a/internal/usage/observer.go +++ b/internal/usage/observer.go @@ -95,18 +95,28 @@ func (o *Observer) parseSSELine(line []byte) { o.parseJSON(payload) } +type tokenDetails struct { + CachedTokens *int64 `json:"cached_tokens"` + CacheWriteTokens *int64 `json:"cache_write_tokens"` +} + type usageFields struct { - PromptTokens *int64 `json:"prompt_tokens"` - CompletionTokens *int64 `json:"completion_tokens"` - TotalTokens *int64 `json:"total_tokens"` - InputTokens *int64 `json:"input_tokens"` - OutputTokens *int64 `json:"output_tokens"` - CacheCreationInputTokens *int64 `json:"cache_creation_input_tokens"` - CacheReadInputTokens *int64 `json:"cache_read_input_tokens"` + PromptTokens *int64 `json:"prompt_tokens"` + CompletionTokens *int64 `json:"completion_tokens"` + TotalTokens *int64 `json:"total_tokens"` + InputTokens *int64 `json:"input_tokens"` + OutputTokens *int64 `json:"output_tokens"` + CacheCreationInputTokens *int64 `json:"cache_creation_input_tokens"` + CacheReadInputTokens *int64 `json:"cache_read_input_tokens"` + PromptTokensDetails *tokenDetails `json:"prompt_tokens_details"` + InputTokensDetails *tokenDetails `json:"input_tokens_details"` } type responseEnvelope struct { - Usage *usageFields `json:"usage"` + Usage *usageFields `json:"usage"` + Response *struct { + Usage *usageFields `json:"usage"` + } `json:"response"` Message *struct { Usage *usageFields `json:"usage"` } `json:"message"` @@ -131,16 +141,22 @@ func (o *Observer) parseJSON(payload []byte) { if envelope.Message != nil && envelope.Message.Usage != nil { o.apply(envelope.Message.Usage) } + if envelope.Response != nil && envelope.Response.Usage != nil { + o.apply(envelope.Response.Usage) + } } func (o *Observer) apply(fields *usageFields) { + inputReported := false if fields.PromptTokens != nil { o.usage.InputTokens = *fields.PromptTokens o.found = true + inputReported = true } if fields.InputTokens != nil { o.usage.InputTokens = *fields.InputTokens o.found = true + inputReported = true } if fields.CompletionTokens != nil { o.usage.OutputTokens = *fields.CompletionTokens @@ -163,6 +179,26 @@ func (o *Observer) apply(fields *usageFields) { o.usage.CacheReadInputTokens = *fields.CacheReadInputTokens o.found = true } + details := fields.InputTokensDetails + if details == nil { + details = fields.PromptTokensDetails + } + if details != nil { + if details.CachedTokens != nil { + o.usage.CacheReadInputTokens = *details.CachedTokens + o.found = true + } + if details.CacheWriteTokens != nil { + o.usage.CacheCreationInputTokens = *details.CacheWriteTokens + o.found = true + } + // OpenAI reports cached token details as subsets of prompt/input_tokens. + // Normalize them into mutually exclusive buckets before billing. Anthropic + // reports its top-level cache fields separately, so they are not adjusted. + if inputReported && (o.protocol == domain.ProtocolOpenAI || o.protocol == domain.ProtocolOpenAIResponses) { + o.usage.InputTokens = max(0, o.usage.InputTokens-o.usage.CacheReadInputTokens-o.usage.CacheCreationInputTokens) + } + } if !o.explicitTotal && o.found { o.usage.TotalTokens = o.usage.InputTokens + o.usage.OutputTokens } diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go index 3104acb..4a4eba6 100644 --- a/internal/usage/observer_test.go +++ b/internal/usage/observer_test.go @@ -8,25 +8,25 @@ import ( func TestObserverReadsOpenAIJSONUsage(t *testing.T) { observer := NewObserver(domain.ProtocolOpenAI, false) - _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`)) + _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"prompt_tokens_details":{"cached_tokens":3,"cache_write_tokens":2},"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 { + if got.InputTokens != 6 || got.CacheReadInputTokens != 3 || got.CacheCreationInputTokens != 2 || got.OutputTokens != 7 || got.TotalTokens != 18 { t.Fatalf("unexpected usage: %+v", got) } } func TestObserverCombinesAnthropicSSEUsage(t *testing.T) { observer := NewObserver(domain.ProtocolAnthropic, true) - _, _ = 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_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":1,\"cache_read_input_tokens\":5,\"cache_creation_input_tokens\":2}}}\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 { + if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 || got.CacheReadInputTokens != 5 || got.CacheCreationInputTokens != 2 { t.Fatalf("unexpected usage: %+v", got) } } @@ -43,3 +43,17 @@ func TestObserverDistinguishesMissingUsageFromReportedZero(t *testing.T) { t.Fatal("explicit zero usage must be distinguished from a missing usage object") } } + +func TestObserverReadsResponsesUsage(t *testing.T) { + nonStream := NewObserver(domain.ProtocolOpenAIResponses, false) + _, _ = nonStream.Write([]byte(`{"object":"response","usage":{"input_tokens":11,"input_tokens_details":{"cached_tokens":3,"cache_write_tokens":2},"output_tokens":7,"total_tokens":18}}`)) + if got := nonStream.Usage(); got.InputTokens != 6 || got.OutputTokens != 7 || got.TotalTokens != 18 || got.CacheReadInputTokens != 3 || got.CacheCreationInputTokens != 2 || !nonStream.Reported() { + t.Fatalf("unexpected non-stream Responses usage: %+v reported=%v", got, nonStream.Reported()) + } + + stream := NewObserver(domain.ProtocolOpenAIResponses, true) + _, _ = stream.Write([]byte("event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":13,\"output_tokens\":5,\"total_tokens\":18}}}\n\n")) + if got := stream.Usage(); got.InputTokens != 13 || got.OutputTokens != 5 || got.TotalTokens != 18 || !stream.Reported() { + t.Fatalf("unexpected streaming Responses usage: %+v reported=%v", got, stream.Reported()) + } +} diff --git a/scripts/start-debug.sh b/scripts/start-debug.sh index bf18de5..80db5de 100755 --- a/scripts/start-debug.sh +++ b/scripts/start-debug.sh @@ -22,104 +22,43 @@ if [[ ! -f "$env_file" ]]; then command -v openssl >/dev/null 2>&1 || fail "openssl is required to generate local credentials" credential_key="$(openssl rand -base64 32 | tr -d '\n')" admin_token="aigw-admin-$(openssl rand -hex 24)" - postgres_password="$(openssl rand -hex 24)" 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' - printf 'AIGW_DATABASE_URL=postgres://aigw:%s@127.0.0.1:5432/aigw?sslmode=disable\n' "$postgres_password" - printf 'AIGW_DATABASE_URL_DOCKER=postgres://aigw:%s@postgres:5432/aigw?sslmode=disable\n' "$postgres_password" - 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:8081/admin/\n' - printf 'AIGW_WEBAUTHN_RP_ID=localhost\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' - printf 'AIGW_SMTP_USERNAME=\n' - printf 'AIGW_SMTP_PASSWORD=\n' - 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: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" +for required_name in AIGW_CREDENTIAL_KEY AIGW_ADMIN_TOKEN; do + grep -q "^${required_name}=." "$env_file" || fail "$required_name is missing from $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 - grep -q "^${required_name}=." "$env_file" || fail "$required_name is missing from $env_file; compare it with .env.control.example" -done +compose=(docker compose --project-directory "$repo_dir") +if [[ -n "${AIGW_COMPOSE_PROJECT_NAME:-}" ]]; then + compose+=(--project-name "$AIGW_COMPOSE_PROJECT_NAME") +fi +compose+=(--env-file "$env_file") cd "$repo_dir" if [[ "${AIGW_DEBUG_SKIP_BUILD:-0}" == "1" ]]; then - log "skipping image build (AIGW_DEBUG_SKIP_BUILD=1)" + log "starting all services without rebuilding the gateway image" + up_args=(up -d --no-build --remove-orphans --wait --wait-timeout 120) else - build_network="${AIGW_DOCKER_BUILD_NETWORK:-host}" - log "building gateway image (network: $build_network)" - docker build --network="$build_network" -t aigw-debug:local . + log "building and starting all services" + up_args=(up -d --build --remove-orphans --wait --wait-timeout 120) fi -log "starting PostgreSQL, Redis, Mailpit, and gateway" -if ! docker compose --env-file "$env_file" up -d --no-build --wait --wait-timeout 120; then - docker compose --env-file "$env_file" ps >&2 || true - gateway_logs="$(docker compose --env-file "$env_file" logs --no-color --tail=120 aigw 2>&1 || true)" - printf '%s\n' "$gateway_logs" >&2 - if [[ "$gateway_logs" == *"decrypt credential"* ]]; then - printf '[aigw-debug] the AIGW_CREDENTIAL_KEY in %s does not match credentials already stored in the PostgreSQL volume\n' "$env_file" >&2 - printf '[aigw-debug] restore the previous key, or explicitly remove the debug volumes if the stored control-plane data is disposable\n' >&2 - fi +if ! "${compose[@]}" "${up_args[@]}"; then + "${compose[@]}" ps >&2 || true + "${compose[@]}" logs --no-color --tail=120 aigw >&2 || true fail "services did not become healthy" fi -log "services are ready" -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:9090/readyz\n' +log "all services are ready" +printf '\nAdmin UI: http://127.0.0.1:8080/admin/\n' +printf 'Health: http://127.0.0.1:8080/readyz\n' printf 'Secrets: %s (mode 0600)\n' "$env_file" -printf '\nLogs: docker compose --env-file %q logs -f aigw\n' "$env_file" +printf '\nLogs: docker compose --project-directory %q --env-file %q logs -f\n' "$repo_dir" "$env_file" printf 'Stop: ./scripts/stop-debug.sh\n' - -if grep -q '^AIGW_STRIPE_API_KEY=rk_test_replace_me$' "$env_file" 2>/dev/null; then - printf '\nStripe uses placeholders. Replace the two Stripe values in %s before testing Checkout.\n' "$env_file" -fi diff --git a/scripts/stop-debug.sh b/scripts/stop-debug.sh index 43de07d..661d7c1 100755 --- a/scripts/stop-debug.sh +++ b/scripts/stop-debug.sh @@ -5,31 +5,63 @@ set -euo pipefail script_dir="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" repo_dir="$(cd -- "$script_dir/.." && pwd)" env_file="${AIGW_DEBUG_ENV_FILE:-$repo_dir/.env.debug}" +stop_timeout="${AIGW_DEBUG_STOP_TIMEOUT:-30}" -command -v docker >/dev/null 2>&1 || { - printf '[aigw-debug] error: docker is not installed\n' >&2 +log() { + printf '[aigw-debug] %s\n' "$*" +} + +fail() { + printf '[aigw-debug] error: %s\n' "$*" >&2 exit 1 } +command -v docker >/dev/null 2>&1 || fail "docker is not installed" +docker info >/dev/null 2>&1 || fail "cannot access Docker; check that the daemon is running and your user belongs to the docker group" + +if [[ ! "$stop_timeout" =~ ^[0-9]+$ ]]; then + fail "AIGW_DEBUG_STOP_TIMEOUT must be a non-negative integer" +fi + +compose=(docker compose --project-directory "$repo_dir") +if [[ -n "${AIGW_COMPOSE_PROJECT_NAME:-}" ]]; then + compose+=(--project-name "$AIGW_COMPOSE_PROJECT_NAME") +fi + if [[ -f "$env_file" ]]; then - compose=(docker compose --env-file "$env_file") + compose+=(--env-file "$env_file") else - export AIGW_SERVER_ADDRESS=:8080 - export AIGW_POSTGRES_USER=debug-stop-placeholder - export AIGW_POSTGRES_PASSWORD=debug-stop-placeholder - export AIGW_POSTGRES_DB=debug-stop-placeholder - export AIGW_DATABASE_URL_DOCKER=postgres://debug-stop-placeholder:debug-stop-placeholder@postgres:5432/debug-stop-placeholder - export AIGW_REDIS_URL_DOCKER= + # Compose still interpolates required variables for `down`; these values are + # only used to resolve the file and are never sent to a running service. export AIGW_CREDENTIAL_KEY=debug-stop-placeholder export AIGW_ADMIN_TOKEN=debug-stop-placeholder - export AIGW_STRIPE_API_KEY=debug-stop-placeholder - export AIGW_STRIPE_WEBHOOK_SECRET=debug-stop-placeholder - export AIGW_STRIPE_SUCCESS_URL=http://127.0.0.1:8080/admin/?topup=success - export AIGW_STRIPE_CANCEL_URL=http://127.0.0.1:8080/admin/?topup=cancel - compose=(docker compose) fi cd "$repo_dir" -printf '[aigw-debug] stopping services\n' -"${compose[@]}" down --remove-orphans -printf '[aigw-debug] stopped; PostgreSQL and Redis volumes were preserved\n' +log "stopping all Compose services" +down_status=0 +if "${compose[@]}" down --remove-orphans --timeout "$stop_timeout"; then + : +else + down_status=$? +fi + +# A forced cleanup handles containers left behind by an interrupted `down` or +# by a previous compose configuration, while leaving named data volumes intact. +mapfile -t remaining_containers < <("${compose[@]}" ps -aq 2>/dev/null || true) +if ((${#remaining_containers[@]} > 0)); then + log "removing ${#remaining_containers[@]} remaining service container(s)" + if ! docker rm -f -- "${remaining_containers[@]}"; then + fail "could not remove all remaining service containers" + fi +fi + +mapfile -t remaining_containers < <("${compose[@]}" ps -aq 2>/dev/null || true) +if ((${#remaining_containers[@]} > 0)); then + fail "${#remaining_containers[@]} service container(s) are still running or present" +fi +if ((down_status != 0)); then + fail "docker compose down failed with exit code $down_status" +fi + +log "all services stopped; named PostgreSQL and Redis volumes were preserved" |
