summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--.env.control.example6
-rw-r--r--.env.example5
-rw-r--r--.gitignore5
-rw-r--r--Dockerfile17
-rw-r--r--Makefile19
-rw-r--r--README.md125
-rw-r--r--cmd/migrate/main.go28
-rw-r--r--cmd/mockupstream/main.go65
-rw-r--r--config.control.example.json40
-rw-r--r--config.example.json64
-rw-r--r--config.local.json50
-rw-r--r--docker-compose.yml47
-rw-r--r--docs/architecture.md68
-rw-r--r--go.mod18
-rw-r--r--go.sum25
-rw-r--r--internal/adminapi/api.go352
-rw-r--r--internal/adminui/assets/app.js70
-rw-r--r--internal/adminui/assets/index.html66
-rw-r--r--internal/adminui/assets/style.css13
-rw-r--r--internal/adminui/ui.go18
-rw-r--r--internal/apierror/apierror.go26
-rw-r--r--internal/auth/static.go124
-rw-r--r--internal/auth/static_test.go76
-rw-r--r--internal/catalog/catalog.go103
-rw-r--r--internal/catalog/catalog_test.go43
-rw-r--r--internal/config/config.go311
-rw-r--r--internal/config/config_test.go95
-rw-r--r--internal/controlplane/manager.go213
-rw-r--r--internal/controlplane/manager_test.go212
-rw-r--r--internal/controlplane/mutations.go281
-rw-r--r--internal/controlplane/queries.go146
-rw-r--r--internal/controlplane/schema.sql82
-rw-r--r--internal/controlplane/snapshot.go150
-rw-r--r--internal/controlplane/store.go170
-rw-r--r--internal/controlplane/types.go138
-rw-r--r--internal/domain/types.go64
-rw-r--r--internal/httpapi/api.go369
-rw-r--r--internal/httpapi/api_test.go224
-rw-r--r--internal/provider/forwarder.go134
-rw-r--r--internal/routing/router.go81
-rw-r--r--internal/routing/router_test.go60
-rw-r--r--internal/security/credentials.go57
-rw-r--r--internal/security/credentials_test.go30
-rw-r--r--internal/telemetry/metrics.go44
-rw-r--r--internal/telemetry/usage_sink.go57
-rw-r--r--internal/usage/observer.go200
-rw-r--r--internal/usage/observer_test.go26
47 files changed, 4617 insertions, 0 deletions
diff --git a/.env.control.example b/.env.control.example
new file mode 100644
index 0000000..4654d0c
--- /dev/null
+++ b/.env.control.example
@@ -0,0 +1,6 @@
+AIGW_DATABASE_URL=postgres://aigw:aigw@127.0.0.1:5432/aigw?sslmode=disable
+AIGW_REDIS_URL=redis://127.0.0.1:6379/0
+# base64 of exactly 32 random bytes; generate with: openssl rand -base64 32
+AIGW_CREDENTIAL_KEY=replace-with-base64-32-byte-key
+AIGW_ADMIN_TOKEN=replace-with-a-long-random-admin-token
+
diff --git a/.env.example b/.env.example
new file mode 100644
index 0000000..4024793
--- /dev/null
+++ b/.env.example
@@ -0,0 +1,5 @@
+# Use shell quotes when exporting this JSON value.
+AIGW_API_KEYS=[{"key":"sk-local-change-me","key_id":"local-key","tenant_id":"tenant-demo","project_id":"project-default","scopes":["inference"]}]
+OPENAI_API_KEY=replace-me
+ANTHROPIC_API_KEY=replace-me
+
diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000..ff4471e
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1,5 @@
+.env
+aigw
+coverage.out
+.playwright-cli/
+output/
diff --git a/Dockerfile b/Dockerfile
new file mode 100644
index 0000000..4f31386
--- /dev/null
+++ b/Dockerfile
@@ -0,0 +1,17 @@
+FROM golang:1.26-alpine AS build
+
+WORKDIR /src
+COPY go.mod ./
+COPY go.sum ./
+RUN go mod download
+COPY cmd ./cmd
+COPY internal ./internal
+RUN CGO_ENABLED=0 go build -buildvcs=false -trimpath -ldflags="-s -w" -o /out/aigw ./cmd/aigw
+
+FROM alpine:3.22
+RUN apk add --no-cache ca-certificates && adduser -D -H -u 10001 aigw
+USER aigw
+COPY --from=build /out/aigw /usr/local/bin/aigw
+EXPOSE 8080
+ENTRYPOINT ["/usr/local/bin/aigw"]
+CMD ["-config", "/etc/aigw/config.json"]
diff --git a/Makefile b/Makefile
new file mode 100644
index 0000000..b8e638f
--- /dev/null
+++ b/Makefile
@@ -0,0 +1,19 @@
+.PHONY: build test vet run mock migrate
+
+build:
+ CGO_ENABLED=0 go build -buildvcs=false -trimpath -o aigw ./cmd/aigw
+
+test:
+ go test ./...
+
+vet:
+ go vet ./...
+
+run:
+ go run ./cmd/aigw -config config.json
+
+mock:
+ go run ./cmd/mockupstream
+
+migrate:
+ go run ./cmd/migrate
diff --git a/README.md b/README.md
new file mode 100644
index 0000000..96d2540
--- /dev/null
+++ b/README.md
@@ -0,0 +1,125 @@
+# AIGW
+
+AIGW 是一个轻量、无状态的 AI API 中转后端。当前阶段聚焦上游接入与 API 分发:提供 OpenAI Chat Completions 和 Anthropic Messages 兼容入口,支持公开模型名映射、多上游路由、加权分流、故障转移、SSE 直通、客户密钥鉴权以及异步用量事件。
+
+它借鉴了 ZenMux 的双协议、`provider/model` 模型命名、统一错误、请求 ID、路由和可观测性边界,但没有复制其业务实现。
+
+## 当前能力
+
+- OpenAI:`POST /v1/chat/completions`、`GET /v1/models`
+- Anthropic:`POST /anthropic/v1/messages`、`GET /anthropic/v1/models`
+- ZenMux 风格别名:所有入口同时提供 `/api/...` 路径
+- 一个公开模型可配置多个同协议上游,低 `priority` 优先,同级按 `weight` 分流
+- 在 429、502、503、504 或连接失败时,于响应开始前自动尝试下一条路由
+- SSE 增量直通、主动断连传播、共享 HTTP/2 连接池
+- `Authorization: Bearer` 和 `x-api-key` 客户鉴权
+- 统一 JSON 错误、`X-AIGW-Request-ID`、Prometheus 文本指标
+- 非阻塞用量事件,包含租户、项目、模型、上游、尝试次数、耗时和 token 用量
+
+## 快速运行
+
+需要 Go 1.26。先准备配置和环境变量:
+
+```bash
+cp config.example.json config.json
+export AIGW_API_KEYS='[{"key":"sk-local-change-me","key_id":"local-key","tenant_id":"tenant-demo","project_id":"project-default","scopes":["inference"]}]'
+export OPENAI_API_KEY='your-upstream-key'
+export ANTHROPIC_API_KEY='your-upstream-key'
+go run ./cmd/aigw -config config.json
+```
+
+如果当前只有一种上游,从 `config.json` 删除未使用的 provider 和对应 model,避免启动时要求该密钥。
+
+不使用真实密钥的本地体验方式:
+
+```bash
+# terminal 1
+go run ./cmd/mockupstream
+
+# terminal 2
+export MOCK_UPSTREAM_KEY='local-only'
+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.local.json
+```
+
+本地网关会监听 `http://127.0.0.1:18081`,公开模型为 `demo/openai` 和 `demo/anthropic`。`cmd/mockupstream` 只用于开发验证,不应部署到生产环境。
+
+OpenAI 调用示例:
+
+```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":"openai/gpt-4.1-mini","messages":[{"role":"user","content":"hello"}],"stream":true}'
+```
+
+Anthropic 调用示例:
+
+```bash
+curl http://127.0.0.1:8080/anthropic/v1/messages \
+ -H 'x-api-key: sk-local-change-me' \
+ -H 'anthropic-version: 2023-06-01' \
+ -H 'Content-Type: application/json' \
+ -d '{"model":"anthropic/claude-sonnet","max_tokens":256,"messages":[{"role":"user","content":"hello"}]}'
+```
+
+## 配置路由
+
+每条 route 把一个对外模型映射到一个上游模型:
+
+```json
+{
+ "id": "vendor/model-public",
+ "owned_by": "vendor",
+ "routes": [
+ {"provider": "provider-a", "upstream_model": "model-v3", "priority": 0, "weight": 80},
+ {"provider": "provider-b", "upstream_model": "model-v3", "priority": 0, "weight": 20},
+ {"provider": "provider-c", "upstream_model": "model-v2", "priority": 10, "weight": 100}
+ ]
+}
+```
+
+这里 A/B 承担约 80/20 的首选流量,C 只作为更低优先级的后备。所有上游密钥仅通过 `api_key_env` 指向的环境变量读取。客户密钥从 `AIGW_API_KEYS` JSON 数组读取,进程内只保存 SHA-256 摘要。
+
+## 生产边界
+
+当前版本可以作为第一阶段的数据面,但还不是完整商业计费平台:
+
+- 用量事件写到结构化日志,队列满时会丢弃并增加 `aigw_usage_events_dropped_total`。正式扣费前必须替换为持久、幂等的账本写入或消息队列。
+- 客户密钥和模型目录在控制面模式下从 PostgreSQL 载入到原子内存快照;Redis 只是可选的变更广播加速层,故障时通过 PostgreSQL generation 轮询收敛。
+- 当前只把 OpenAI 入口发给 OpenAI 兼容上游、Anthropic 入口发给 Anthropic 兼容上游,不做跨协议转换。
+- 自动故障转移可能在极少数网络错误下造成上游重复执行。正式计费时需要上游幂等能力、请求去重策略和重复成本对账。
+- `/metrics` 应仅在内网暴露;公网 TLS、WAF 和连接层限速应放在负载均衡器或边缘代理。
+
+## PostgreSQL + Redis 控制面
+
+控制面模式使用 [config.control.example.json](config.control.example.json)。PostgreSQL 保存事实数据:租户、项目、API Key 摘要、加密后的供应商密钥、模型和路由。Redis 保存当前 generation 并广播变更事件;它不是客户账本,PG 才是控制面事实来源。Redis 短暂不可用或未配置时,网关仍从 PostgreSQL 启动并通过周期轮询收敛,恢复后会在后台自动重连订阅。管理写入先提交 PG 并替换本机快照,Redis 广播失败只记录告警,不会把已提交的写入报告为失败。
+
+凭证密钥必须是 base64 编码的 32 字节随机值。供应商 API Key 使用 AES-256-GCM 加密后写入 PostgreSQL,运行时解密到当前内存快照,管理接口永远不会返回它。
+
+开发环境可以直接启动:
+
+```bash
+export AIGW_CREDENTIAL_KEY="$(openssl rand -base64 32)"
+export AIGW_ADMIN_TOKEN="$(openssl rand -hex 32)"
+docker compose up --build
+```
+
+然后打开 `http://127.0.0.1:8080/admin/`,输入 `AIGW_ADMIN_TOKEN`。第一套资源的创建顺序是:Tenant → Project → API key → Provider → Model route。创建客户 API Key 时,明文只在成功响应中出现一次。
+
+生产环境建议把 `auto_migrate` 改为 `false`,先执行:
+
+```bash
+AIGW_DATABASE_URL="postgres://..." go run ./cmd/migrate
+```
+
+管理 API 仅使用一个 bootstrap token,应该放在内网、VPN 或反向代理后。下一阶段需要把它替换为管理员用户、角色和审计日志。
+
+完整的扩展边界见 [架构说明](docs/architecture.md)。
+
+## 验证
+
+```bash
+go test ./...
+go vet ./...
+```
diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go
new file mode 100644
index 0000000..20452f9
--- /dev/null
+++ b/cmd/migrate/main.go
@@ -0,0 +1,28 @@
+package main
+
+import (
+ "context"
+ "flag"
+ "fmt"
+ "os"
+ "time"
+
+ "aigw/internal/controlplane"
+)
+
+func main() {
+ environment := flag.String("database-url-env", "AIGW_DATABASE_URL", "environment variable containing the PostgreSQL URL")
+ flag.Parse()
+ databaseURL := os.Getenv(*environment)
+ if databaseURL == "" {
+ fmt.Fprintf(os.Stderr, "%s is empty\n", *environment)
+ os.Exit(1)
+ }
+ ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
+ defer cancel()
+ if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil {
+ fmt.Fprintln(os.Stderr, err)
+ os.Exit(1)
+ }
+ fmt.Println("control-plane schema is up to date")
+}
diff --git a/cmd/mockupstream/main.go b/cmd/mockupstream/main.go
new file mode 100644
index 0000000..cf695bb
--- /dev/null
+++ b/cmd/mockupstream/main.go
@@ -0,0 +1,65 @@
+// Command mockupstream is a local-only upstream used for manual gateway checks.
+package main
+
+import (
+ "encoding/json"
+ "flag"
+ "fmt"
+ "log"
+ "net/http"
+ "time"
+)
+
+func main() {
+ address := flag.String("address", "127.0.0.1:18080", "listen address")
+ flag.Parse()
+ mux := http.NewServeMux()
+ mux.HandleFunc("POST /v1/chat/completions", openAI)
+ mux.HandleFunc("POST /v1/messages", anthropic)
+ server := &http.Server{Addr: *address, Handler: mux, ReadHeaderTimeout: 5 * time.Second}
+ log.Printf("mock upstream listening on http://%s", *address)
+ log.Fatal(server.ListenAndServe())
+}
+
+func openAI(w http.ResponseWriter, r *http.Request) {
+ var request struct {
+ Model string `json:"model"`
+ Stream bool `json:"stream"`
+ }
+ if json.NewDecoder(r.Body).Decode(&request) != nil {
+ http.Error(w, "invalid JSON", http.StatusBadRequest)
+ return
+ }
+ if request.Stream {
+ w.Header().Set("Content-Type", "text/event-stream")
+ w.Header().Set("Cache-Control", "no-cache")
+ fmt.Fprintf(w, "data: {\"id\":\"chatcmpl-local\",\"object\":\"chat.completion.chunk\",\"model\":%q,\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"hello from mock upstream\"},\"finish_reason\":null}]}\n\n", request.Model)
+ w.(http.Flusher).Flush()
+ time.Sleep(50 * time.Millisecond)
+ fmt.Fprint(w, "data: {\"id\":\"chatcmpl-local\",\"object\":\"chat.completion.chunk\",\"choices\":[],\"usage\":{\"prompt_tokens\":4,\"completion_tokens\":4,\"total_tokens\":8}}\n\ndata: [DONE]\n\n")
+ return
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": "chatcmpl-local", "object": "chat.completion", "model": request.Model,
+ "choices": []any{map[string]any{"index": 0, "message": map[string]string{"role": "assistant", "content": "hello from mock upstream"}, "finish_reason": "stop"}},
+ "usage": map[string]int{"prompt_tokens": 4, "completion_tokens": 4, "total_tokens": 8},
+ })
+}
+
+func anthropic(w http.ResponseWriter, r *http.Request) {
+ var request struct {
+ Model string `json:"model"`
+ }
+ if json.NewDecoder(r.Body).Decode(&request) != nil {
+ http.Error(w, "invalid JSON", http.StatusBadRequest)
+ return
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "id": "msg_local", "type": "message", "role": "assistant", "model": request.Model,
+ "content": []any{map[string]string{"type": "text", "text": "hello from mock upstream"}},
+ "stop_reason": "end_turn", "stop_sequence": nil,
+ "usage": map[string]int{"input_tokens": 4, "output_tokens": 4},
+ })
+}
diff --git a/config.control.example.json b/config.control.example.json
new file mode 100644
index 0000000..d2affaf
--- /dev/null
+++ b/config.control.example.json
@@ -0,0 +1,40 @@
+{
+ "server": {
+ "address": ":8080",
+ "max_body_bytes": 16777216,
+ "read_header_timeout_seconds": 10,
+ "idle_timeout_seconds": 120,
+ "shutdown_timeout_seconds": 20
+ },
+ "auth": {
+ "allow_anonymous": false
+ },
+ "control_plane": {
+ "enabled": true,
+ "database_url_env": "AIGW_DATABASE_URL",
+ "redis_url_env": "AIGW_REDIS_URL",
+ "credential_key_env": "AIGW_CREDENTIAL_KEY",
+ "redis_channel": "aigw:control:changed",
+ "snapshot_cache_key": "aigw:control:generation",
+ "reload_interval_seconds": 30,
+ "auto_migrate": true
+ },
+ "admin": {
+ "enabled": true,
+ "token_env": "AIGW_ADMIN_TOKEN",
+ "base_path": "/admin"
+ },
+ "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": [],
+ "models": []
+}
+
diff --git a/config.example.json b/config.example.json
new file mode 100644
index 0000000..fbf677d
--- /dev/null
+++ b/config.example.json
@@ -0,0 +1,64 @@
+{
+ "server": {
+ "address": ":8080",
+ "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-primary",
+ "protocol": "openai",
+ "base_url": "https://api.openai.com/v1",
+ "api_key_env": "OPENAI_API_KEY"
+ },
+ {
+ "id": "anthropic-primary",
+ "protocol": "anthropic",
+ "base_url": "https://api.anthropic.com/v1",
+ "api_key_env": "ANTHROPIC_API_KEY"
+ }
+ ],
+ "models": [
+ {
+ "id": "openai/gpt-4.1-mini",
+ "owned_by": "openai",
+ "routes": [
+ {
+ "provider": "openai-primary",
+ "upstream_model": "gpt-4.1-mini",
+ "priority": 0,
+ "weight": 100
+ }
+ ]
+ },
+ {
+ "id": "anthropic/claude-sonnet",
+ "owned_by": "anthropic",
+ "routes": [
+ {
+ "provider": "anthropic-primary",
+ "upstream_model": "claude-sonnet-4-5",
+ "priority": 0,
+ "weight": 100
+ }
+ ]
+ }
+ ]
+}
+
diff --git a/config.local.json b/config.local.json
new file mode 100644
index 0000000..d4273b1
--- /dev/null
+++ b/config.local.json
@@ -0,0 +1,50 @@
+{
+ "server": {
+ "address": "127.0.0.1:18081",
+ "max_body_bytes": 16777216
+ },
+ "auth": {
+ "keys_env": "AIGW_API_KEYS",
+ "allow_anonymous": false
+ },
+ "upstream_http": {
+ "max_idle_connections": 128,
+ "max_idle_connections_per_host": 64,
+ "response_header_timeout_seconds": 10
+ },
+ "observability": {
+ "usage_buffer": 128,
+ "expose_metrics": true
+ },
+ "providers": [
+ {
+ "id": "local-openai",
+ "protocol": "openai",
+ "base_url": "http://127.0.0.1:18080/v1",
+ "api_key_env": "MOCK_UPSTREAM_KEY"
+ },
+ {
+ "id": "local-anthropic",
+ "protocol": "anthropic",
+ "base_url": "http://127.0.0.1:18080/v1",
+ "api_key_env": "MOCK_UPSTREAM_KEY"
+ }
+ ],
+ "models": [
+ {
+ "id": "demo/openai",
+ "owned_by": "local",
+ "routes": [
+ {"provider": "local-openai", "upstream_model": "mock-openai", "priority": 0, "weight": 1}
+ ]
+ },
+ {
+ "id": "demo/anthropic",
+ "owned_by": "local",
+ "routes": [
+ {"provider": "local-anthropic", "upstream_model": "mock-anthropic", "priority": 0, "weight": 1}
+ ]
+ }
+ ]
+}
+
diff --git a/docker-compose.yml b/docker-compose.yml
new file mode 100644
index 0000000..9e74612
--- /dev/null
+++ b/docker-compose.yml
@@ -0,0 +1,47 @@
+services:
+ postgres:
+ image: postgres:16-alpine
+ environment:
+ POSTGRES_USER: aigw
+ POSTGRES_PASSWORD: aigw
+ POSTGRES_DB: aigw
+ ports:
+ - "127.0.0.1:5432:5432"
+ volumes:
+ - aigw-postgres:/var/lib/postgresql/data
+ healthcheck:
+ test: ["CMD-SHELL", "pg_isready -U aigw -d aigw"]
+ interval: 5s
+ timeout: 3s
+ retries: 20
+
+ redis:
+ image: redis:7-alpine
+ ports:
+ - "127.0.0.1:6379:6379"
+ volumes:
+ - aigw-redis:/data
+ healthcheck:
+ test: ["CMD", "redis-cli", "ping"]
+ interval: 5s
+ timeout: 3s
+ retries: 20
+
+ aigw:
+ build: .
+ depends_on:
+ postgres:
+ condition: service_healthy
+ ports:
+ - "127.0.0.1:8080:8080"
+ environment:
+ AIGW_DATABASE_URL: postgres://aigw:aigw@postgres:5432/aigw?sslmode=disable
+ AIGW_REDIS_URL: redis://redis:6379/0
+ AIGW_CREDENTIAL_KEY: ${AIGW_CREDENTIAL_KEY:?set AIGW_CREDENTIAL_KEY}
+ AIGW_ADMIN_TOKEN: ${AIGW_ADMIN_TOKEN:?set AIGW_ADMIN_TOKEN}
+ volumes:
+ - ./config.control.example.json:/etc/aigw/config.json:ro
+
+volumes:
+ aigw-postgres:
+ aigw-redis:
diff --git a/docs/architecture.md b/docs/architecture.md
new file mode 100644
index 0000000..640a37d
--- /dev/null
+++ b/docs/architecture.md
@@ -0,0 +1,68 @@
+# 架构说明
+
+## 设计目标
+
+第一阶段只做高性能、可横向扩容的数据面,并提前固定未来计费、权限和控制面的接口位置。代理核心不在请求热路径同步访问数据库,也不把供应商差异扩散到对外 API 层。
+
+```mermaid
+flowchart LR
+ C["客户 SDK"] --> E["OpenAI / Anthropic 入口"]
+ E --> A["Authenticator + Scope"]
+ A --> R["Model Catalog + Router"]
+ R --> P["共享连接池 + Provider Adapter"]
+ P --> U1["上游 A"]
+ P --> U2["上游 B"]
+ P --> U3["后备上游"]
+ E -. "非阻塞" .-> M["Usage Event Sink"]
+ E -.-> O["Metrics + Structured Logs"]
+ M --> L["未来:持久计费账本"]
+ D["PostgreSQL 控制面"] -. "事实数据" .-> S["Redis generation / PubSub"]
+ S -. "变更广播" .-> A
+ D -. "启动与轮询快照" .-> A
+ D -. "模型和路由快照" .-> R
+```
+
+## 热路径
+
+1. 入口生成不可预测的请求 ID,并通过 Bearer 或 `x-api-key` 解析客户身份。
+2. 身份包含 `key_id`、`tenant_id`、`project_id` 和 scopes;当前是静态实现,代理只依赖 `Authenticator` 接口。
+3. 请求体在配置上限内读取一次,以提取公开模型名并支持故障转移时重放。
+4. Router 按协议过滤路由,先按优先级分组,再在同级内按权重选择首选上游。
+5. Provider Adapter 重写上游模型名和凭证,使用进程级共享 Transport 发送请求。
+6. 上游返回成功头后即锁定路由。SSE 使用固定缓冲区复制并逐块 flush,不等待完整响应。
+7. 请求结束后发布 UsageEvent;发布操作有界、非阻塞,当前消费者输出结构化日志。
+
+## 已固定的扩展边界
+
+| 边界 | 当前实现 | 下一阶段替换 |
+| --- | --- | --- |
+| 客户身份 | PostgreSQL 快照、内存 SHA-256 索引 | 管理员 RBAC、撤销审计、Redis/Lua 配额 |
+| 权限 | `Principal.Scopes` 中的 `inference` | RBAC/ABAC、模型 allowlist、IP 与预算策略 |
+| 模型目录 | PostgreSQL 快照 + 可选 Redis generation 广播 + PG 轮询兜底 | 版本化控制面、热更新、灰度发布 |
+| 路由 | priority + weighted selection + failover | 健康评分、延迟 EWMA、成本/质量策略、熔断 |
+| 计量 | UsageEvent 异步日志 | 持久消息流、幂等消费、不可变用量账本 |
+| 计费 | 未执行扣费 | 价格版本、预授权/额度、结算、退款与对账 |
+| 限流 | 仅上游 429 故障转移 | Redis/Lua 或边缘 token bucket、租户并发额度 |
+| 协议 | 同协议透传 | 规范化 IR + OpenAI/Anthropic/Google 双向转换 |
+
+## 计费数据原则
+
+真实计费不能直接依赖请求日志。建议下一阶段建立不可变 `usage_ledger`,以 `request_id + attempt` 做幂等键,并同时记录:公开模型、实际上游、价格表版本、输入/输出/cache token、币种、上游成本、客户价格、状态和冲正关系。
+
+流式请求的 usage 可能只在最后事件出现。当前 observer 会在线解析 OpenAI/Anthropic SSE 中的 usage;如果上游不返回 usage,事件中的 token 为零。接入商业扣费前,应为每种上游建立有测试的 usage normalizer,并使用上游账单进行日对账。
+
+## 扩容方式
+
+- 进程无状态,可以在 L4/L7 负载均衡器后横向扩容。
+- HTTP server 不设置 WriteTimeout,避免杀死长时间流;上游只限制响应头等待时间,客户端取消会终止上游请求。
+- 连接池按进程共享,默认保留 4096 个空闲连接、单 host 1024 个,可按实例并发和上游限制调整。
+- 请求体默认最多 16 MiB,响应仅保留最多 64 KiB 用于非流式 usage 提取;正文直接传输。
+- 计量队列有界,慢消费者不会拖住推理请求。正式计费时不能仅靠“丢弃并打指标”,需要本地 WAL 或高可用消息系统。
+
+## 建议的后续顺序
+
+1. 把管理 bootstrap token 替换为管理员账户、RBAC、审计事件,并与公网推理监听端口隔离。
+2. 增加价格表和账本:tenant、project、api_key、provider_credential、model、route、price_book。
+3. 持久 UsageEvent 和幂等账本,然后再实现余额预授权、扣费和退款。
+4. 主动健康检查、熔断、延迟 EWMA 和按成本路由。
+5. 规范化中间表示,实现可靠的跨协议转换与 Responses/Embeddings 等更多端点。
diff --git a/go.mod b/go.mod
new file mode 100644
index 0000000..840da94
--- /dev/null
+++ b/go.mod
@@ -0,0 +1,18 @@
+module aigw
+
+go 1.26
+
+require (
+ github.com/jackc/pgx/v5 v5.10.0
+ github.com/redis/go-redis/v9 v9.21.0
+)
+
+require (
+ github.com/cespare/xxhash/v2 v2.3.0 // indirect
+ github.com/jackc/pgpassfile v1.0.0 // indirect
+ github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
+ github.com/jackc/puddle/v2 v2.2.2 // indirect
+ go.uber.org/atomic v1.11.0 // indirect
+ golang.org/x/sync v0.17.0 // indirect
+ golang.org/x/text v0.29.0 // indirect
+)
diff --git a/go.sum b/go.sum
new file mode 100644
index 0000000..cd11f69
--- /dev/null
+++ b/go.sum
@@ -0,0 +1,25 @@
+github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
+github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
+github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
+github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
+github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
+github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
+github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
+github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0=
+github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
+github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
+github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
+github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
+github.com/redis/go-redis/v9 v9.21.0 h1:FPBE4hhbAke+TLmcY3WkpbDffJEomdqPn3HYiqAtL9E=
+github.com/redis/go-redis/v9 v9.21.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA=
+github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
+github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
+github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
+go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
+go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
+golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
+golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
+golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
+golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
+gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
+gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
diff --git a/internal/adminapi/api.go b/internal/adminapi/api.go
new file mode 100644
index 0000000..1eb2a5d
--- /dev/null
+++ b/internal/adminapi/api.go
@@ -0,0 +1,352 @@
+package adminapi
+
+import (
+ "crypto/subtle"
+ "encoding/json"
+ "errors"
+ "log/slog"
+ "net/http"
+ "strings"
+
+ "aigw/internal/adminui"
+ "aigw/internal/apierror"
+ "aigw/internal/controlplane"
+
+ "github.com/jackc/pgx/v5/pgconn"
+)
+
+type API struct {
+ store *controlplane.Store
+ manager *controlplane.Manager
+ token []byte
+ logger *slog.Logger
+ prefix string
+}
+
+type Options struct {
+ Store *controlplane.Store
+ Manager *controlplane.Manager
+ Token string
+ Logger *slog.Logger
+ Prefix string
+}
+
+func New(options Options) *API {
+ prefix := strings.TrimRight(options.Prefix, "/")
+ if prefix == "" {
+ prefix = "/admin"
+ }
+ return &API{store: options.Store, manager: options.Manager, token: []byte(options.Token), logger: options.Logger, prefix: prefix}
+}
+
+func (a *API) Handler() http.Handler {
+ mux := http.NewServeMux()
+ apiPrefix := a.prefix + "/api"
+ mux.HandleFunc("GET "+a.prefix, func(w http.ResponseWriter, r *http.Request) {
+ http.Redirect(w, r, a.prefix+"/", http.StatusTemporaryRedirect)
+ })
+ mux.Handle(a.prefix+"/", http.StripPrefix(a.prefix, adminui.Handler()))
+
+ mux.HandleFunc("GET "+apiPrefix+"/overview", a.withAuth(a.overview))
+ mux.HandleFunc("GET "+apiPrefix+"/tenants", a.withAuth(a.listTenants))
+ mux.HandleFunc("POST "+apiPrefix+"/tenants", a.withAuth(a.createTenant))
+ mux.HandleFunc("GET "+apiPrefix+"/projects", a.withAuth(a.listProjects))
+ mux.HandleFunc("POST "+apiPrefix+"/projects", a.withAuth(a.createProject))
+ mux.HandleFunc("GET "+apiPrefix+"/keys", a.withAuth(a.listKeys))
+ mux.HandleFunc("POST "+apiPrefix+"/keys", a.withAuth(a.createKey))
+ mux.HandleFunc("POST "+apiPrefix+"/keys/{id}/revoke", a.withAuth(a.revokeKey))
+ mux.HandleFunc("GET "+apiPrefix+"/providers", a.withAuth(a.listProviders))
+ mux.HandleFunc("POST "+apiPrefix+"/providers", a.withAuth(a.createProvider))
+ mux.HandleFunc("POST "+apiPrefix+"/providers/{id}/toggle", a.withAuth(a.toggleProvider))
+ mux.HandleFunc("GET "+apiPrefix+"/models", a.withAuth(a.listModels))
+ mux.HandleFunc("POST "+apiPrefix+"/models", a.withAuth(a.createModel))
+ mux.HandleFunc("POST "+apiPrefix+"/models/{id}/toggle", a.withAuth(a.toggleModel))
+ mux.HandleFunc("POST "+apiPrefix+"/reload", a.withAuth(a.reload))
+ return mux
+}
+
+func (a *API) withAuth(next http.HandlerFunc) http.HandlerFunc {
+ return func(w http.ResponseWriter, r *http.Request) {
+ provided := strings.TrimSpace(r.Header.Get("X-Admin-Token"))
+ if provided == "" {
+ provided = bearerToken(r.Header.Get("Authorization"))
+ }
+ if len(provided) == 0 || subtle.ConstantTimeCompare([]byte(provided), a.token) != 1 {
+ apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "admin_unauthorized", Message: "Administrator authentication required"}, requestID(r))
+ return
+ }
+ next(w, r)
+ }
+}
+
+func (a *API) overview(w http.ResponseWriter, r *http.Request) {
+ result, err := a.store.Overview(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ result.RuntimeGeneration = a.manager.Generation()
+ result.RedisConfigured = a.manager.RedisConfigured()
+ result.RedisConnected = a.manager.RedisConnected()
+ writeJSON(w, result)
+}
+
+func (a *API) listTenants(w http.ResponseWriter, r *http.Request) {
+ result, err := a.store.ListTenants(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) createTenant(w http.ResponseWriter, r *http.Request) {
+ var input controlplane.CreateTenantInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ result, generation, err := a.store.CreateTenant(r.Context(), input)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "tenant", result.ID) {
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) listProjects(w http.ResponseWriter, r *http.Request) {
+ result, err := a.store.ListProjects(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) createProject(w http.ResponseWriter, r *http.Request) {
+ var input controlplane.CreateProjectInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ result, generation, err := a.store.CreateProject(r.Context(), input)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "project", result.ID) {
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) listKeys(w http.ResponseWriter, r *http.Request) {
+ result, err := a.store.ListAPIKeys(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) createKey(w http.ResponseWriter, r *http.Request) {
+ var input controlplane.CreateAPIKeyInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ result, generation, err := a.store.CreateAPIKey(r.Context(), input)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "api_key", result.ID) {
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) revokeKey(w http.ResponseWriter, r *http.Request) {
+ generation, err := a.store.RevokeAPIKey(r.Context(), r.PathValue("id"))
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "api_key", r.PathValue("id")) {
+ return
+ }
+ writeJSON(w, map[string]any{"status": "revoked"})
+}
+
+func (a *API) listProviders(w http.ResponseWriter, r *http.Request) {
+ result, err := a.store.ListProviders(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) createProvider(w http.ResponseWriter, r *http.Request) {
+ var input controlplane.CreateProviderInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ result, generation, err := a.store.CreateProvider(r.Context(), input)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "provider", result.ID) {
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) toggleProvider(w http.ResponseWriter, r *http.Request) {
+ var input struct {
+ Enabled bool `json:"enabled"`
+ }
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ id := r.PathValue("id")
+ generation, err := a.store.SetProviderEnabled(r.Context(), id, input.Enabled)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "provider", id) {
+ return
+ }
+ writeJSON(w, map[string]any{"id": id, "enabled": input.Enabled})
+}
+
+func (a *API) listModels(w http.ResponseWriter, r *http.Request) {
+ result, err := a.store.ListModels(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ writeJSON(w, result)
+}
+
+func (a *API) createModel(w http.ResponseWriter, r *http.Request) {
+ var input controlplane.CreateModelInput
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ result, generation, err := a.store.CreateModel(r.Context(), input)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "model", result.ID) {
+ return
+ }
+ writeStatusJSON(w, http.StatusCreated, result)
+}
+
+func (a *API) toggleModel(w http.ResponseWriter, r *http.Request) {
+ var input struct {
+ Enabled bool `json:"enabled"`
+ }
+ if !decodeBody(w, r, &input) {
+ return
+ }
+ id := r.PathValue("id")
+ generation, err := a.store.SetModelEnabled(r.Context(), id, input.Enabled)
+ if err != nil {
+ a.mutationError(w, r, err)
+ return
+ }
+ if !a.changed(w, r, generation, "model", id) {
+ return
+ }
+ writeJSON(w, map[string]any{"id": id, "enabled": input.Enabled})
+}
+
+func (a *API) reload(w http.ResponseWriter, r *http.Request) {
+ generation, err := a.manager.Reload(r.Context())
+ if err != nil {
+ a.databaseError(w, r, err)
+ return
+ }
+ writeJSON(w, map[string]any{"generation": generation, "status": "reloaded"})
+}
+
+func (a *API) changed(w http.ResponseWriter, r *http.Request, generation int64, resource, id string) bool {
+ if err := a.manager.AfterMutation(r.Context(), generation, resource, id); err != nil {
+ a.logger.Error("admin_control_plane_sync_failed", "resource", resource, "id", id, "error", err)
+ apierror.Write(w, apierror.Error{Status: http.StatusServiceUnavailable, Type: "control_plane_sync_failed", Message: "Change was persisted but runtime reload failed; retry reload"}, requestID(r))
+ return false
+ }
+ return true
+}
+
+func (a *API) databaseError(w http.ResponseWriter, r *http.Request, err error) {
+ a.logger.Error("admin_database_error", "error", err)
+ apierror.Write(w, apierror.Error{Status: http.StatusInternalServerError, Type: "control_plane_error", Message: "Control plane unavailable"}, requestID(r))
+}
+
+func (a *API) mutationError(w http.ResponseWriter, r *http.Request, err error) {
+ status := http.StatusBadRequest
+ typeName := "invalid_params"
+ message := err.Error()
+ if errors.Is(err, controlplane.ErrNotFound) {
+ status = http.StatusNotFound
+ typeName = "not_found"
+ message = "Resource not found"
+ }
+ var pgError *pgconn.PgError
+ if errors.As(err, &pgError) {
+ switch pgError.Code {
+ case "23505":
+ status = http.StatusConflict
+ typeName = "already_exists"
+ message = "Resource already exists"
+ case "23503":
+ status = http.StatusBadRequest
+ typeName = "invalid_reference"
+ message = "Referenced resource does not exist"
+ }
+ }
+ apierror.Write(w, apierror.Error{Status: status, Type: typeName, Message: message}, requestID(r))
+}
+
+func decodeBody(w http.ResponseWriter, r *http.Request, destination any) bool {
+ r.Body = http.MaxBytesReader(w, r.Body, 1<<20)
+ defer r.Body.Close()
+ decoder := json.NewDecoder(r.Body)
+ decoder.DisallowUnknownFields()
+ if err := decoder.Decode(destination); err != nil {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Invalid JSON request body"}, requestID(r))
+ return false
+ }
+ return true
+}
+
+func writeJSON(w http.ResponseWriter, value any) {
+ writeStatusJSON(w, http.StatusOK, value)
+}
+
+func writeStatusJSON(w http.ResponseWriter, status int, value any) {
+ w.Header().Set("Content-Type", "application/json")
+ w.WriteHeader(status)
+ _ = json.NewEncoder(w).Encode(value)
+}
+
+func bearerToken(header string) string {
+ parts := strings.Fields(header)
+ if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
+ return parts[1]
+ }
+ return ""
+}
+
+func requestID(r *http.Request) string {
+ if value := r.Header.Get("X-AIGW-Request-ID"); value != "" {
+ return value
+ }
+ return "admin"
+}
diff --git a/internal/adminui/assets/app.js b/internal/adminui/assets/app.js
new file mode 100644
index 0000000..f5ba048
--- /dev/null
+++ b/internal/adminui/assets/app.js
@@ -0,0 +1,70 @@
+const state = { token: sessionStorage.getItem('aigw_admin_token') || '', tenants: [], projects: [], providers: [], models: [], keys: [] };
+const $ = (selector) => document.querySelector(selector);
+const $$ = (selector) => [...document.querySelectorAll(selector)];
+
+function esc(value) {
+ return String(value ?? '').replace(/[&<>'"]/g, (char) => ({ '&':'&amp;', '<':'&lt;', '>':'&gt;', "'":'&#39;', '"':'&quot;' }[char]));
+}
+function date(value) { return value ? new Date(value).toLocaleString() : '—'; }
+function toast(message, error = false) {
+ const node = $('#toast'); node.textContent = message; node.className = `toast visible ${error ? 'error' : ''}`;
+ setTimeout(() => { node.className = 'toast'; }, 3200);
+}
+async function api(path, options = {}) {
+ const response = await fetch(`./api${path}`, { ...options, headers: { 'Content-Type': 'application/json', 'Authorization': `Bearer ${state.token}`, ...(options.headers || {}) } });
+ const payload = await response.json().catch(() => ({}));
+ if (!response.ok) throw new Error(payload?.error?.message || `Request failed (${response.status})`);
+ return payload;
+}
+function setConnected(connected) {
+ $('#connection-state').textContent = connected ? 'Connected' : 'Offline';
+ $('#connection-state').className = `state ${connected ? 'online' : ''}`;
+}
+function formJSON(form) { return Object.fromEntries(new FormData(form).entries()); }
+function selectOptions(items, valueKey, labelKey, empty = 'Select…') { return `<option value="">${empty}</option>${items.map(item => `<option value="${esc(item[valueKey])}">${esc(item[labelKey])}</option>`).join('')}`; }
+
+async function loadAll() {
+ if (!state.token) { setConnected(false); return; }
+ try {
+ [state.tenants, state.projects, state.providers, state.models, state.keys] = await Promise.all([
+ api('/tenants'), api('/projects'), api('/providers'), api('/models'), api('/keys')
+ ]);
+ const overview = await api('/overview');
+ renderOverview(overview); renderTenants(); renderProjects(); renderKeys(); renderProviders(); renderModels(); renderRouteEditor(); setConnected(true);
+ } catch (error) { setConnected(false); toast(error.message, true); }
+}
+function renderOverview(data) {
+ const propagation = !data.redis_configured ? 'PG polling' : data.redis_connected ? 'Redis live' : 'PG fallback';
+ const items = [['Tenants', data.tenants, 'active accounts'], ['Projects', data.projects, 'workspaces'], ['API keys', data.api_keys, 'active credentials'], ['Providers', data.providers, 'enabled upstreams'], ['Models', data.models, 'public routes'], ['Runtime', data.runtime_generation, `PG generation ${data.generation}`], ['Propagation', propagation, data.redis_configured ? 'automatic recovery' : 'Redis not configured']];
+ $('#metrics').innerHTML = items.map(([label, value, sub]) => `<article class="metric"><span>${label}</span><strong>${esc(value)}</strong><small>${sub}</small></article>`).join('');
+}
+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').innerHTML = selectOptions(state.tenants, 'id', 'name'); $('#key-tenant').innerHTML = selectOptions(state.tenants, 'id', 'name'); }
+function renderProjects() { $('#projects-body').innerHTML = state.projects.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><code>${esc(item.tenant_id).slice(0, 8)}…</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>${esc(item.project_id).slice(0, 8)}…</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' ? `<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>${esc(item.route_count)}</td><td><span class="badge ${item.enabled ? 'active' : 'suspended'}">${item.enabled ? 'enabled' : 'disabled'}</span></td><td><button class="text-button" data-toggle-provider="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled ? 'Disable' : 'Enable'}</button></td></tr>`).join('') || emptyRow(6); }
+function renderModels() { $('#models-body').innerHTML = state.models.map(item => `<tr><td><strong>${esc(item.public_id)}</strong></td><td>${esc(item.owned_by || '—')}</td><td><div class="route-list">${(item.routes || []).map(route => `<span>${esc(route.provider_name || route.provider_id).slice(0, 24)} → ${esc(route.upstream_model)} <em>p${route.priority} / w${route.weight}</em></span>`).join('')}</div></td><td><span class="badge ${item.enabled ? 'active' : 'suspended'}">${item.enabled ? 'enabled' : 'disabled'}</span></td><td><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 renderRouteEditor() { const current = $('#route-editor'); if (!current.children.length) addRoute(); $$('.route-provider').forEach(select => { const selected = select.value; select.innerHTML = selectOptions(state.providers.filter(item => item.enabled), 'id', 'name', 'Provider…'); select.value = selected; }); }
+function addRoute() { const wrapper = document.createElement('div'); wrapper.className = 'route-row'; wrapper.innerHTML = `<select class="route-provider" required></select><input class="route-upstream" required placeholder="Upstream model"><input class="route-priority" type="number" min="0" value="0" title="Priority"><input class="route-weight" type="number" min="1" max="100" value="100" title="Weight"><button class="icon-button remove-route" type="button" aria-label="Remove route">×</button>`; $('#route-editor').appendChild(wrapper); renderRouteEditor(); }
+function emptyRow(span) { return `<tr><td colspan="${span}" class="empty">No records yet</td></tr>`; }
+
+document.addEventListener('click', async (event) => {
+ 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; }
+ 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 revoke = event.target.closest('[data-revoke-key]'); if (revoke && confirm('Revoke this API key?')) { try { await api(`/keys/${revoke.dataset.revokeKey}/revoke`, { method: 'POST', body: '{}' }); await loadAll(); toast('API key revoked'); } catch (error) { toast(error.message, true); } }
+ const provider = event.target.closest('[data-toggle-provider]'); if (provider) { try { await api(`/providers/${provider.dataset.toggleProvider}/toggle`, { method: 'POST', body: JSON.stringify({ enabled: provider.dataset.enabled === 'true' }) }); await loadAll(); toast('Provider updated'); } catch (error) { toast(error.message, true); } }
+ const model = event.target.closest('[data-toggle-model]'); if (model) { try { await api(`/models/${model.dataset.toggleModel}/toggle`, { method: 'POST', body: JSON.stringify({ enabled: model.dataset.enabled === 'true' }) }); await loadAll(); toast('Model updated'); } catch (error) { toast(error.message, true); } }
+});
+$('#key-tenant').addEventListener('change', renderKeyProjects);
+$('#session-form').addEventListener('submit', async (event) => { event.preventDefault(); state.token = $('#admin-token').value.trim(); sessionStorage.setItem('aigw_admin_token', state.token); await loadAll(); });
+$('#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(); $('#created-secret').textContent = result.key; $('#secret-dialog').showModal(); await loadAll(); } catch (error) { toast(error.message, true); } });
+$('#provider-form').addEventListener('submit', async (event) => { event.preventDefault(); try { await api('/providers', { method: 'POST', body: JSON.stringify(formJSON(event.target)) }); event.target.reset(); await loadAll(); toast('Provider added'); } catch (error) { toast(error.message, true); } });
+$('#model-form').addEventListener('submit', async (event) => { event.preventDefault(); try { const data = formJSON(event.target); data.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); } });
+$('#close-dialog').addEventListener('click', () => $('#secret-dialog').close());
+$('#copy-secret').addEventListener('click', async () => { await navigator.clipboard.writeText($('#created-secret').textContent); toast('Key copied'); });
+$('#admin-token').value = state.token;
+if (state.token) loadAll();
diff --git a/internal/adminui/assets/index.html b/internal/adminui/assets/index.html
new file mode 100644
index 0000000..7869855
--- /dev/null
+++ b/internal/adminui/assets/index.html
@@ -0,0 +1,66 @@
+<!doctype html>
+<html lang="zh-CN">
+<head>
+ <meta charset="utf-8">
+ <meta name="viewport" content="width=device-width, initial-scale=1">
+ <title>AIGW Control Plane</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="./style.css">
+</head>
+<body>
+ <header class="topbar">
+ <div class="brand"><span class="brand-mark">A</span><div><strong>AIGW</strong><small>CONTROL PLANE</small></div></div>
+ <form class="session" id="session-form"><input id="admin-token" name="admin-token" type="password" placeholder="Admin token" autocomplete="current-password" aria-label="Admin token"><button type="submit">Connect</button><span id="connection-state" class="state">Offline</span></form>
+ </header>
+ <main class="shell">
+ <nav class="tabs" aria-label="Admin sections">
+ <button class="tab active" data-section="overview">Overview</button>
+ <button class="tab" data-section="tenants">Tenants</button>
+ <button class="tab" data-section="projects">Projects</button>
+ <button class="tab" data-section="keys">API keys</button>
+ <button class="tab" data-section="providers">Providers</button>
+ <button class="tab" data-section="models">Models & routes</button>
+ </nav>
+
+ <section id="overview" class="section active">
+ <div class="section-heading"><div><span class="eyebrow">OPERATIONS</span><h1>Control plane overview</h1></div><button class="button secondary" id="reload">Reload snapshot</button></div>
+ <div class="metric-grid" id="metrics"></div>
+ <div class="panel note"><div class="note-icon">i</div><div><strong>Runtime snapshot</strong><p>Writes commit to PostgreSQL and apply locally first. Redis accelerates propagation when available; PostgreSQL polling keeps every gateway convergent.</p></div></div>
+ </section>
+
+ <section id="tenants" class="section">
+ <div class="section-heading"><div><span class="eyebrow">IDENTITY</span><h1>Tenants</h1></div></div>
+ <form class="panel form-grid" id="tenant-form"><label>Slug<input name="slug" required pattern="[a-z0-9][a-z0-9-]{1,62}[a-z0-9]" placeholder="acme"></label><label>Name<input name="name" required placeholder="Acme Inc."></label><button class="button primary" type="submit">Create tenant</button></form>
+ <div class="panel table-wrap"><table><thead><tr><th>Name</th><th>Slug</th><th>Status</th><th>Created</th></tr></thead><tbody id="tenants-body"></tbody></table></div>
+ </section>
+
+ <section id="projects" class="section">
+ <div class="section-heading"><div><span class="eyebrow">IDENTITY</span><h1>Projects</h1></div></div>
+ <form class="panel form-grid" id="project-form"><label>Tenant<select name="tenant_id" id="project-tenant" required></select></label><label>Slug<input name="slug" required placeholder="production"></label><label>Name<input name="name" required placeholder="Production API"></label><button class="button primary" type="submit">Create project</button></form>
+ <div class="panel table-wrap"><table><thead><tr><th>Name</th><th>Tenant</th><th>Slug</th><th>Status</th></tr></thead><tbody id="projects-body"></tbody></table></div>
+ </section>
+
+ <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"><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,admin"></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>
+ </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"><label>Name<input name="name" required placeholder="openai-primary"></label><label>Protocol<select name="protocol"><option value="openai">OpenAI</option><option value="anthropic">Anthropic</option></select></label><label>Base URL<input name="base_url" type="url" required placeholder="https://api.example.com/v1"></label><label>API key<input name="api_key" type="password" required autocomplete="new-password" placeholder="Stored encrypted"></label><button class="button primary" type="submit">Add provider</button></form>
+ <div class="panel table-wrap"><table><thead><tr><th>Name</th><th>Protocol</th><th>Base URL</th><th>Routes</th><th>Status</th><th></th></tr></thead><tbody id="providers-body"></tbody></table></div>
+ </section>
+
+ <section id="models" class="section">
+ <div class="section-heading"><div><span class="eyebrow">ROUTING</span><h1>Models & routes</h1></div></div>
+ <form class="panel form-grid" id="model-form"><label>Public model ID<input name="public_id" required placeholder="openai/gpt-4.1-mini"></label><label>Owned by<input name="owned_by" placeholder="openai"></label><div class="route-editor" id="route-editor"></div><button class="button subtle" type="button" id="add-route">Add route</button><button class="button primary" type="submit">Create model</button></form>
+ <div class="panel table-wrap"><table><thead><tr><th>Public ID</th><th>Owner</th><th>Routes</th><th>Status</th><th></th></tr></thead><tbody id="models-body"></tbody></table></div>
+ </section>
+ </main>
+ <div id="toast" class="toast" role="status"></div>
+ <dialog id="secret-dialog"><div class="dialog-content"><div class="section-heading"><div><span class="eyebrow">ONE-TIME SECRET</span><h2>API key created</h2></div><button class="icon-button" id="close-dialog" aria-label="Close">×</button></div><p>Copy this key now. It will not be shown again.</p><code id="created-secret"></code><button class="button primary" id="copy-secret">Copy key</button></div></dialog>
+ <script src="./app.js" defer></script>
+</body>
+</html>
diff --git a/internal/adminui/assets/style.css b/internal/adminui/assets/style.css
new file mode 100644
index 0000000..dadb019
--- /dev/null
+++ b/internal/adminui/assets/style.css
@@ -0,0 +1,13 @@
+: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; }
+.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:8px; } .session input { width:220px; border:1px solid #3b5b6c; background:#18384b; color:#fff; padding:9px 11px; outline:none; } .session input::placeholder { color:#91acb9; } .session button { min-height:40px; border:1px solid #8fd0df; background:#b8eef7; color:#102a3a; padding:0 14px; font-weight:750; } .session button:hover { background:#d4f6fb; } .state { color:#9db0bb; font-size:12px; } .state.online { color:#86d5ad; }
+.shell { width:min(1240px,calc(100% - 48px)); margin:28px auto 60px; } .tabs { display:flex; gap:4px; border-bottom:1px solid var(--line); margin-bottom:26px; overflow:auto; } .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; }
+.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; }
+.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; }
+@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; } }
+@media (max-width:620px) { .topbar { height:auto; padding:16px; align-items:flex-start; flex-direction:column; } .session { width:100%; } .session input { flex:1; width:auto; min-width:0; } .shell { width:calc(100% - 24px); margin-top:18px; } .metric-grid { grid-template-columns:repeat(2,1fr); } .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; } }
diff --git a/internal/adminui/ui.go b/internal/adminui/ui.go
new file mode 100644
index 0000000..7f58030
--- /dev/null
+++ b/internal/adminui/ui.go
@@ -0,0 +1,18 @@
+package adminui
+
+import (
+ "embed"
+ "io/fs"
+ "net/http"
+)
+
+//go:embed assets/*
+var assets embed.FS
+
+func Handler() http.Handler {
+ content, err := fs.Sub(assets, "assets")
+ if err != nil {
+ panic(err)
+ }
+ return http.FileServer(http.FS(content))
+}
diff --git a/internal/apierror/apierror.go b/internal/apierror/apierror.go
new file mode 100644
index 0000000..e88781a
--- /dev/null
+++ b/internal/apierror/apierror.go
@@ -0,0 +1,26 @@
+package apierror
+
+import (
+ "encoding/json"
+ "net/http"
+ "strconv"
+)
+
+type Error struct {
+ Status int
+ Type string
+ Message string
+}
+
+func Write(w http.ResponseWriter, err Error, requestID string) {
+ w.Header().Set("Content-Type", "application/json")
+ w.Header().Set("X-AIGW-Request-ID", requestID)
+ w.WriteHeader(err.Status)
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "error": map[string]string{
+ "code": strconv.Itoa(err.Status),
+ "type": err.Type,
+ "message": err.Message,
+ },
+ })
+}
diff --git a/internal/auth/static.go b/internal/auth/static.go
new file mode 100644
index 0000000..495f90c
--- /dev/null
+++ b/internal/auth/static.go
@@ -0,0 +1,124 @@
+package auth
+
+import (
+ "crypto/sha256"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "net/http"
+ "strings"
+ "sync/atomic"
+
+ "aigw/internal/domain"
+)
+
+var ErrUnauthorized = errors.New("invalid or missing API key")
+
+type Authenticator interface {
+ Authenticate(*http.Request) (domain.Principal, error)
+}
+
+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"`
+}
+
+type StaticAuthenticator struct {
+ state atomic.Pointer[keySnapshot]
+ allowAnonymous bool
+}
+
+type keySnapshot struct {
+ keys map[[sha256.Size]byte]domain.Principal
+}
+
+type HashedKeyRecord struct {
+ Hash [sha256.Size]byte
+ Principal domain.Principal
+}
+
+func NewStatic(raw string, allowAnonymous bool) (*StaticAuthenticator, error) {
+ result := &StaticAuthenticator{allowAnonymous: allowAnonymous}
+ if strings.TrimSpace(raw) == "" {
+ if allowAnonymous {
+ result.ReplaceHashed(nil)
+ return result, nil
+ }
+ return nil, errors.New("client API key environment variable is empty")
+ }
+
+ var records []KeyRecord
+ if err := json.Unmarshal([]byte(raw), &records); err != nil {
+ return nil, fmt.Errorf("parse client API keys JSON: %w", err)
+ }
+ hashed := make([]HashedKeyRecord, 0, len(records))
+ seen := make(map[[sha256.Size]byte]struct{}, len(records))
+ for i, record := range records {
+ if record.Key == "" || record.KeyID == "" || record.TenantID == "" || record.ProjectID == "" {
+ return nil, fmt.Errorf("client API key record %d requires key, key_id, tenant_id, and project_id", i)
+ }
+ hash := sha256.Sum256([]byte(record.Key))
+ if _, exists := seen[hash]; exists {
+ return nil, fmt.Errorf("duplicate client API key at record %d", i)
+ }
+ seen[hash] = 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...),
+ }})
+ }
+ if len(hashed) == 0 && !allowAnonymous {
+ return nil, errors.New("at least one client API key is required")
+ }
+ result.ReplaceHashed(hashed)
+ return result, nil
+}
+
+func NewDynamic(records []HashedKeyRecord, allowAnonymous bool) *StaticAuthenticator {
+ result := &StaticAuthenticator{allowAnonymous: allowAnonymous}
+ result.ReplaceHashed(records)
+ return result
+}
+
+func (a *StaticAuthenticator) ReplaceHashed(records []HashedKeyRecord) {
+ keys := make(map[[sha256.Size]byte]domain.Principal, len(records))
+ for _, record := range records {
+ principal := record.Principal
+ principal.Scopes = append([]string(nil), principal.Scopes...)
+ keys[record.Hash] = principal
+ }
+ a.state.Store(&keySnapshot{keys: keys})
+}
+
+func (a *StaticAuthenticator) Authenticate(r *http.Request) (domain.Principal, error) {
+ key := bearerToken(r.Header.Get("Authorization"))
+ if key == "" {
+ key = strings.TrimSpace(r.Header.Get("x-api-key"))
+ }
+ if key == "" && a.allowAnonymous {
+ return domain.Principal{KeyID: "anonymous", TenantID: "anonymous", ProjectID: "anonymous"}, nil
+ }
+ if key == "" {
+ return domain.Principal{}, ErrUnauthorized
+ }
+ snapshot := a.state.Load()
+ if snapshot == nil {
+ return domain.Principal{}, ErrUnauthorized
+ }
+ principal, ok := snapshot.keys[sha256.Sum256([]byte(key))]
+ if !ok {
+ return domain.Principal{}, ErrUnauthorized
+ }
+ return principal, nil
+}
+
+func bearerToken(header string) string {
+ parts := strings.Fields(header)
+ if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
+ return ""
+ }
+ return parts[1]
+}
diff --git a/internal/auth/static_test.go b/internal/auth/static_test.go
new file mode 100644
index 0000000..cf54ba6
--- /dev/null
+++ b/internal/auth/static_test.go
@@ -0,0 +1,76 @@
+package auth
+
+import (
+ "crypto/sha256"
+ "net/http"
+ "testing"
+
+ "aigw/internal/domain"
+)
+
+func TestStaticAuthenticator(t *testing.T) {
+ raw := `[{"key":"sk-client","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`
+ authenticator, err := NewStatic(raw, false)
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ request, _ := http.NewRequest(http.MethodGet, "http://gateway.test/v1/models", nil)
+ request.Header.Set("Authorization", "Bearer sk-client")
+ principal, err := authenticator.Authenticate(request)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if principal.TenantID != "tenant-1" || principal.ProjectID != "project-1" || principal.KeyID != "key-1" {
+ t.Fatalf("unexpected principal: %+v", principal)
+ }
+
+ request.Header.Set("Authorization", "Bearer wrong")
+ if _, err := authenticator.Authenticate(request); err != ErrUnauthorized {
+ t.Fatalf("expected ErrUnauthorized, got %v", err)
+ }
+}
+
+func TestDynamicAuthenticatorReplacesSnapshot(t *testing.T) {
+ oldHash := sha256.Sum256([]byte("sk-old"))
+ newHash := sha256.Sum256([]byte("sk-new"))
+ authenticator := NewDynamic([]HashedKeyRecord{{
+ Hash: oldHash,
+ Principal: domain.Principal{
+ KeyID: "old", TenantID: "tenant-1", ProjectID: "project-1", Scopes: []string{"inference"},
+ },
+ }}, false)
+
+ request, _ := http.NewRequest(http.MethodGet, "http://gateway.test/v1/models", nil)
+ request.Header.Set("Authorization", "Bearer sk-old")
+ if _, err := authenticator.Authenticate(request); err != nil {
+ t.Fatalf("old key should initially authenticate: %v", err)
+ }
+
+ authenticator.ReplaceHashed([]HashedKeyRecord{{
+ Hash: newHash,
+ Principal: domain.Principal{
+ KeyID: "new", TenantID: "tenant-1", ProjectID: "project-1", Scopes: []string{"inference"},
+ },
+ }})
+ if _, err := authenticator.Authenticate(request); err != ErrUnauthorized {
+ t.Fatalf("old key remained valid after snapshot replacement: %v", err)
+ }
+ request.Header.Set("Authorization", "Bearer sk-new")
+ principal, err := authenticator.Authenticate(request)
+ if err != nil || principal.KeyID != "new" {
+ t.Fatalf("new key was not loaded: principal=%+v err=%v", principal, err)
+ }
+}
+
+func TestStaticAuthenticatorAcceptsAnthropicHeader(t *testing.T) {
+ authenticator, err := NewStatic(`[{"key":"sk-client","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1"}]`, false)
+ if err != nil {
+ t.Fatal(err)
+ }
+ request, _ := http.NewRequest(http.MethodGet, "http://gateway.test/anthropic/v1/models", nil)
+ request.Header.Set("x-api-key", "sk-client")
+ if _, err := authenticator.Authenticate(request); err != nil {
+ t.Fatal(err)
+ }
+}
diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go
new file mode 100644
index 0000000..76e33b1
--- /dev/null
+++ b/internal/catalog/catalog.go
@@ -0,0 +1,103 @@
+package catalog
+
+import (
+ "fmt"
+ "sort"
+ "strings"
+ "sync/atomic"
+
+ "aigw/internal/config"
+ "aigw/internal/domain"
+)
+
+type Catalog struct {
+ state atomic.Pointer[snapshot]
+}
+
+type snapshot struct {
+ models map[string]domain.Model
+ list []domain.Model
+}
+
+func New(cfg config.Config) *Catalog {
+ providers := make(map[string]domain.Provider, len(cfg.Providers))
+ for _, provider := range cfg.Providers {
+ providers[provider.ID] = domain.Provider{
+ ID: provider.ID,
+ Protocol: provider.Protocol,
+ BaseURL: strings.TrimRight(provider.BaseURL, "/"),
+ APIKey: provider.APIKey,
+ }
+ }
+
+ models := make([]domain.Model, 0, len(cfg.Models))
+ for _, modelCfg := range cfg.Models {
+ model := domain.Model{ID: modelCfg.ID, OwnedBy: modelCfg.OwnedBy}
+ for _, route := range modelCfg.Routes {
+ model.Routes = append(model.Routes, domain.Route{
+ Provider: providers[route.Provider],
+ UpstreamModel: route.UpstreamModel,
+ Priority: route.Priority,
+ Weight: route.Weight,
+ })
+ }
+ models = append(models, model)
+ }
+ return NewModels(models)
+}
+
+func NewModels(models []domain.Model) *Catalog {
+ catalog := &Catalog{}
+ catalog.Replace(models)
+ return catalog
+}
+
+func (c *Catalog) Replace(source []domain.Model) {
+ models := make(map[string]domain.Model, len(source))
+ list := make([]domain.Model, 0, len(source))
+ for _, sourceModel := range source {
+ model := sourceModel
+ model.Routes = append([]domain.Route(nil), sourceModel.Routes...)
+ models[model.ID] = model
+ list = append(list, model)
+ }
+ sort.Slice(list, func(i, j int) bool { return list[i].ID < list[j].ID })
+ c.state.Store(&snapshot{models: models, list: list})
+}
+
+func (c *Catalog) Model(id string) (domain.Model, error) {
+ current := c.state.Load()
+ if current == nil {
+ return domain.Model{}, fmt.Errorf("model %q not found", id)
+ }
+ model, ok := current.models[id]
+ if !ok {
+ return domain.Model{}, fmt.Errorf("model %q not found", id)
+ }
+ return model, nil
+}
+
+func (c *Catalog) Models(protocol domain.Protocol) []domain.Model {
+ current := c.state.Load()
+ if current == nil {
+ return nil
+ }
+ result := make([]domain.Model, 0, len(current.list))
+ for _, model := range current.list {
+ for _, route := range model.Routes {
+ if route.Provider.Protocol == protocol {
+ result = append(result, model)
+ break
+ }
+ }
+ }
+ return result
+}
+
+func (c *Catalog) Count() int {
+ current := c.state.Load()
+ if current == nil {
+ return 0
+ }
+ return len(current.list)
+}
diff --git a/internal/catalog/catalog_test.go b/internal/catalog/catalog_test.go
new file mode 100644
index 0000000..07fb71f
--- /dev/null
+++ b/internal/catalog/catalog_test.go
@@ -0,0 +1,43 @@
+package catalog
+
+import (
+ "testing"
+
+ "aigw/internal/domain"
+)
+
+func TestCatalogReplaceSwapsModelSnapshot(t *testing.T) {
+ openAI := domain.Provider{ID: "openai", Protocol: domain.ProtocolOpenAI}
+ anthropic := domain.Provider{ID: "anthropic", Protocol: domain.ProtocolAnthropic}
+ catalog := NewModels([]domain.Model{{
+ ID: "old", Routes: []domain.Route{{Provider: openAI, UpstreamModel: "old-upstream"}},
+ }})
+
+ catalog.Replace([]domain.Model{{
+ ID: "new", Routes: []domain.Route{{Provider: anthropic, UpstreamModel: "new-upstream"}},
+ }})
+ if _, err := catalog.Model("old"); err == nil {
+ t.Fatal("old model remained after snapshot replacement")
+ }
+ model, err := catalog.Model("new")
+ if err != nil || len(model.Routes) != 1 || model.Routes[0].Provider.ID != "anthropic" {
+ t.Fatalf("new model was not loaded: model=%+v err=%v", model, err)
+ }
+ if got := catalog.Models(domain.ProtocolOpenAI); len(got) != 0 {
+ t.Fatalf("unexpected OpenAI models after replacement: %+v", got)
+ }
+}
+
+func TestCatalogReplaceCopiesRouteSlices(t *testing.T) {
+ models := []domain.Model{{ID: "model", Routes: []domain.Route{{UpstreamModel: "before"}}}}
+ catalog := NewModels(models)
+ models[0].Routes[0].UpstreamModel = "after"
+
+ model, err := catalog.Model("model")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if model.Routes[0].UpstreamModel != "before" {
+ t.Fatal("catalog snapshot aliases the caller's route slice")
+ }
+}
diff --git a/internal/config/config.go b/internal/config/config.go
new file mode 100644
index 0000000..5e1e518
--- /dev/null
+++ b/internal/config/config.go
@@ -0,0 +1,311 @@
+package config
+
+import (
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/url"
+ "os"
+ "strings"
+ "time"
+
+ "aigw/internal/domain"
+)
+
+type Config struct {
+ Server ServerConfig `json:"server"`
+ Auth AuthConfig `json:"auth"`
+ ControlPlane ControlPlaneConfig `json:"control_plane"`
+ Admin AdminConfig `json:"admin"`
+ UpstreamHTTP UpstreamHTTPConfig `json:"upstream_http"`
+ Providers []ProviderConfig `json:"providers"`
+ Models []ModelConfig `json:"models"`
+ Observability ObservabilityConfig `json:"observability"`
+}
+
+type ServerConfig struct {
+ Address string `json:"address"`
+ MaxBodyBytes int64 `json:"max_body_bytes"`
+ ReadHeaderTimeoutSecs int `json:"read_header_timeout_seconds"`
+ IdleTimeoutSecs int `json:"idle_timeout_seconds"`
+ ShutdownTimeoutSecs int `json:"shutdown_timeout_seconds"`
+}
+
+type AuthConfig struct {
+ KeysEnv string `json:"keys_env"`
+ AllowAnonymous bool `json:"allow_anonymous"`
+}
+
+type ControlPlaneConfig struct {
+ Enabled bool `json:"enabled"`
+ DatabaseURLEnv string `json:"database_url_env"`
+ RedisURLEnv string `json:"redis_url_env"`
+ CredentialKeyEnv string `json:"credential_key_env"`
+ RedisChannel string `json:"redis_channel"`
+ SnapshotCacheKey string `json:"snapshot_cache_key"`
+ ReloadIntervalSeconds int `json:"reload_interval_seconds"`
+ AutoMigrate bool `json:"auto_migrate"`
+ DatabaseURL string `json:"-"`
+ RedisURL string `json:"-"`
+ CredentialKey string `json:"-"`
+}
+
+type AdminConfig struct {
+ Enabled bool `json:"enabled"`
+ TokenEnv string `json:"token_env"`
+ BasePath string `json:"base_path"`
+ Token string `json:"-"`
+}
+
+type UpstreamHTTPConfig struct {
+ MaxIdleConnections int `json:"max_idle_connections"`
+ MaxIdleConnectionsPerHost int `json:"max_idle_connections_per_host"`
+ IdleConnectionTimeoutSecs int `json:"idle_connection_timeout_seconds"`
+ ResponseHeaderTimeoutSecs int `json:"response_header_timeout_seconds"`
+}
+
+type ProviderConfig struct {
+ ID string `json:"id"`
+ Protocol domain.Protocol `json:"protocol"`
+ BaseURL string `json:"base_url"`
+ APIKeyEnv string `json:"api_key_env"`
+ APIKey string `json:"-"`
+}
+
+type ModelConfig struct {
+ ID string `json:"id"`
+ OwnedBy string `json:"owned_by"`
+ Routes []RouteConfig `json:"routes"`
+}
+
+type RouteConfig struct {
+ Provider string `json:"provider"`
+ UpstreamModel string `json:"upstream_model"`
+ Priority int `json:"priority"`
+ Weight int `json:"weight"`
+}
+
+type ObservabilityConfig struct {
+ UsageBuffer int `json:"usage_buffer"`
+ ExposeMetrics bool `json:"expose_metrics"`
+}
+
+func Load(path string) (Config, error) {
+ f, err := os.Open(path)
+ if err != nil {
+ return Config{}, fmt.Errorf("open config: %w", err)
+ }
+ defer f.Close()
+
+ var cfg Config
+ decoder := json.NewDecoder(f)
+ decoder.DisallowUnknownFields()
+ if err := decoder.Decode(&cfg); err != nil {
+ return Config{}, fmt.Errorf("decode config: %w", err)
+ }
+ if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) {
+ if err == nil {
+ return Config{}, errors.New("decode config: multiple JSON values")
+ }
+ return Config{}, fmt.Errorf("decode config trailing data: %w", err)
+ }
+
+ applyDefaults(&cfg)
+ if err := resolveSecrets(&cfg); err != nil {
+ return Config{}, err
+ }
+ if err := Validate(cfg); err != nil {
+ return Config{}, err
+ }
+ return cfg, nil
+}
+
+func applyDefaults(cfg *Config) {
+ if cfg.Server.Address == "" {
+ cfg.Server.Address = ":8080"
+ }
+ if cfg.Server.MaxBodyBytes == 0 {
+ cfg.Server.MaxBodyBytes = 16 << 20
+ }
+ if cfg.Server.ReadHeaderTimeoutSecs == 0 {
+ cfg.Server.ReadHeaderTimeoutSecs = 10
+ }
+ if cfg.Server.IdleTimeoutSecs == 0 {
+ cfg.Server.IdleTimeoutSecs = 120
+ }
+ if cfg.Server.ShutdownTimeoutSecs == 0 {
+ cfg.Server.ShutdownTimeoutSecs = 20
+ }
+ if cfg.Auth.KeysEnv == "" {
+ cfg.Auth.KeysEnv = "AIGW_API_KEYS"
+ }
+ if cfg.ControlPlane.DatabaseURLEnv == "" {
+ cfg.ControlPlane.DatabaseURLEnv = "AIGW_DATABASE_URL"
+ }
+ if cfg.ControlPlane.RedisURLEnv == "" {
+ cfg.ControlPlane.RedisURLEnv = "AIGW_REDIS_URL"
+ }
+ if cfg.ControlPlane.CredentialKeyEnv == "" {
+ cfg.ControlPlane.CredentialKeyEnv = "AIGW_CREDENTIAL_KEY"
+ }
+ if cfg.ControlPlane.RedisChannel == "" {
+ cfg.ControlPlane.RedisChannel = "aigw:control:changed"
+ }
+ if cfg.ControlPlane.SnapshotCacheKey == "" {
+ cfg.ControlPlane.SnapshotCacheKey = "aigw:control:snapshot:v1"
+ }
+ if cfg.ControlPlane.ReloadIntervalSeconds == 0 {
+ cfg.ControlPlane.ReloadIntervalSeconds = 30
+ }
+ if cfg.Admin.TokenEnv == "" {
+ cfg.Admin.TokenEnv = "AIGW_ADMIN_TOKEN"
+ }
+ if cfg.Admin.BasePath == "" {
+ cfg.Admin.BasePath = "/admin"
+ }
+ if cfg.UpstreamHTTP.MaxIdleConnections == 0 {
+ cfg.UpstreamHTTP.MaxIdleConnections = 4096
+ }
+ if cfg.UpstreamHTTP.MaxIdleConnectionsPerHost == 0 {
+ cfg.UpstreamHTTP.MaxIdleConnectionsPerHost = 1024
+ }
+ if cfg.UpstreamHTTP.IdleConnectionTimeoutSecs == 0 {
+ cfg.UpstreamHTTP.IdleConnectionTimeoutSecs = 90
+ }
+ if cfg.UpstreamHTTP.ResponseHeaderTimeoutSecs == 0 {
+ cfg.UpstreamHTTP.ResponseHeaderTimeoutSecs = 60
+ }
+ if cfg.Observability.UsageBuffer == 0 {
+ cfg.Observability.UsageBuffer = 8192
+ }
+ for i := range cfg.Models {
+ for j := range cfg.Models[i].Routes {
+ if cfg.Models[i].Routes[j].Weight == 0 {
+ cfg.Models[i].Routes[j].Weight = 1
+ }
+ }
+ }
+}
+
+func resolveSecrets(cfg *Config) error {
+ if cfg.ControlPlane.Enabled {
+ cfg.ControlPlane.DatabaseURL = os.Getenv(cfg.ControlPlane.DatabaseURLEnv)
+ cfg.ControlPlane.RedisURL = os.Getenv(cfg.ControlPlane.RedisURLEnv)
+ cfg.ControlPlane.CredentialKey = os.Getenv(cfg.ControlPlane.CredentialKeyEnv)
+ }
+ if cfg.Admin.Enabled {
+ cfg.Admin.Token = os.Getenv(cfg.Admin.TokenEnv)
+ }
+ for i := range cfg.Providers {
+ provider := &cfg.Providers[i]
+ if provider.APIKeyEnv == "" {
+ continue
+ }
+ provider.APIKey = os.Getenv(provider.APIKeyEnv)
+ if provider.APIKey == "" {
+ return fmt.Errorf("provider %q: environment variable %s is empty", provider.ID, provider.APIKeyEnv)
+ }
+ }
+ return nil
+}
+
+func Validate(cfg Config) error {
+ if cfg.Server.MaxBodyBytes < 1024 {
+ return errors.New("server.max_body_bytes must be at least 1024")
+ }
+ if cfg.Observability.UsageBuffer < 1 {
+ return errors.New("observability.usage_buffer must be positive")
+ }
+
+ if cfg.ControlPlane.Enabled {
+ if cfg.ControlPlane.DatabaseURL == "" {
+ return fmt.Errorf("control_plane: environment variable %s is empty", cfg.ControlPlane.DatabaseURLEnv)
+ }
+ if cfg.ControlPlane.CredentialKey == "" {
+ return fmt.Errorf("control_plane: environment variable %s is empty", cfg.ControlPlane.CredentialKeyEnv)
+ }
+ if cfg.ControlPlane.ReloadIntervalSeconds < 1 {
+ return errors.New("control_plane.reload_interval_seconds must be positive")
+ }
+ }
+ if cfg.Admin.Enabled {
+ if !cfg.ControlPlane.Enabled {
+ return errors.New("admin requires control_plane.enabled")
+ }
+ if cfg.Admin.Token == "" {
+ return fmt.Errorf("admin: environment variable %s is empty", cfg.Admin.TokenEnv)
+ }
+ if !strings.HasPrefix(cfg.Admin.BasePath, "/") || cfg.Admin.BasePath == "/" {
+ return errors.New("admin.base_path must start with / and cannot be /")
+ }
+ }
+
+ providers := make(map[string]ProviderConfig, len(cfg.Providers))
+ for _, provider := range cfg.Providers {
+ if provider.ID == "" {
+ return errors.New("provider id is required")
+ }
+ 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)
+ }
+ 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)
+ }
+ if provider.APIKeyEnv == "" {
+ return fmt.Errorf("provider %q: api_key_env is required", provider.ID)
+ }
+ providers[provider.ID] = provider
+ }
+ if len(providers) == 0 && !cfg.ControlPlane.Enabled {
+ return errors.New("at least one provider is required")
+ }
+
+ models := make(map[string]struct{}, len(cfg.Models))
+ for _, model := range cfg.Models {
+ if strings.TrimSpace(model.ID) == "" {
+ return errors.New("model id is required")
+ }
+ if _, exists := models[model.ID]; exists {
+ return fmt.Errorf("duplicate model id %q", model.ID)
+ }
+ models[model.ID] = struct{}{}
+ if len(model.Routes) == 0 {
+ return fmt.Errorf("model %q: at least one route is required", model.ID)
+ }
+ for _, route := range model.Routes {
+ if _, exists := providers[route.Provider]; !exists {
+ return fmt.Errorf("model %q: unknown provider %q", model.ID, route.Provider)
+ }
+ if route.UpstreamModel == "" {
+ return fmt.Errorf("model %q: upstream_model is required", model.ID)
+ }
+ if route.Priority < 0 {
+ return fmt.Errorf("model %q: route priority cannot be negative", model.ID)
+ }
+ if route.Weight < 1 || route.Weight > 100 {
+ return fmt.Errorf("model %q: route weight must be between 1 and 100", model.ID)
+ }
+ }
+ }
+ if len(models) == 0 && !cfg.ControlPlane.Enabled {
+ return errors.New("at least one model is required")
+ }
+ return nil
+}
+
+func (c ServerConfig) ReadHeaderTimeout() time.Duration {
+ return time.Duration(c.ReadHeaderTimeoutSecs) * time.Second
+}
+
+func (c ServerConfig) IdleTimeout() time.Duration {
+ return time.Duration(c.IdleTimeoutSecs) * time.Second
+}
+
+func (c ServerConfig) ShutdownTimeout() time.Duration {
+ return time.Duration(c.ShutdownTimeoutSecs) * time.Second
+}
diff --git a/internal/config/config_test.go b/internal/config/config_test.go
new file mode 100644
index 0000000..e681508
--- /dev/null
+++ b/internal/config/config_test.go
@@ -0,0 +1,95 @@
+package config
+
+import (
+ "os"
+ "path/filepath"
+ "testing"
+)
+
+func TestLoadAppliesDefaultsAndResolvesSecrets(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ path := writeConfig(t, `{
+ "providers": [{"id":"primary","protocol":"openai","base_url":"https://example.com/v1","api_key_env":"TEST_UPSTREAM_KEY"}],
+ "models": [{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}]
+}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.Server.Address != ":8080" || cfg.Server.MaxBodyBytes == 0 {
+ t.Fatalf("defaults not applied: %+v", cfg.Server)
+ }
+ if cfg.Providers[0].APIKey != "secret" {
+ t.Fatal("provider secret was not resolved")
+ }
+ if cfg.Models[0].Routes[0].Weight != 1 {
+ t.Fatalf("expected default route weight 1, got %d", cfg.Models[0].Routes[0].Weight)
+ }
+}
+
+func TestLoadRejectsUnknownFieldsAndTrailingData(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ unknown := writeConfig(t, `{"unknown":true}`)
+ if _, err := Load(unknown); err == nil {
+ t.Fatal("expected unknown field error")
+ }
+
+ trailing := writeConfig(t, `{} {}`)
+ if _, err := Load(trailing); err == nil {
+ t.Fatal("expected trailing JSON error")
+ }
+}
+
+func TestLoadControlPlaneModeWithoutStaticRoutes(t *testing.T) {
+ t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw")
+ t.Setenv("AIGW_REDIS_URL", "redis://redis:6379/0")
+ t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=")
+ t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret")
+ path := writeConfig(t, `{
+ "control_plane": {"enabled":true},
+ "admin": {"enabled":true}
+}`)
+
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.ControlPlane.DatabaseURL == "" || cfg.ControlPlane.RedisURL == "" || cfg.Admin.Token != "admin-secret" {
+ t.Fatalf("control-plane secrets were not resolved: %+v", cfg)
+ }
+ if len(cfg.Providers) != 0 || len(cfg.Models) != 0 {
+ t.Fatal("control-plane mode unexpectedly requires static providers or models")
+ }
+}
+
+func TestLoadControlPlaneModeWithoutRedis(t *testing.T) {
+ t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw")
+ t.Setenv("AIGW_REDIS_URL", "")
+ t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=")
+ path := writeConfig(t, `{"control_plane":{"enabled":true}}`)
+
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.ControlPlane.RedisURL != "" {
+ t.Fatalf("Redis URL = %q, want empty", cfg.ControlPlane.RedisURL)
+ }
+}
+
+func TestLoadRejectsAdminWithoutControlPlane(t *testing.T) {
+ t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret")
+ path := writeConfig(t, `{"admin":{"enabled":true}}`)
+ if _, err := Load(path); err == nil {
+ t.Fatal("expected admin without control plane to be rejected")
+ }
+}
+
+func writeConfig(t *testing.T, content string) string {
+ t.Helper()
+ path := filepath.Join(t.TempDir(), "config.json")
+ if err := os.WriteFile(path, []byte(content), 0o600); err != nil {
+ t.Fatal(err)
+ }
+ return path
+}
diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go
new file mode 100644
index 0000000..42f5efe
--- /dev/null
+++ b/internal/controlplane/manager.go
@@ -0,0 +1,213 @@
+package controlplane
+
+import (
+ "context"
+ "encoding/json"
+ "log/slog"
+ "sync"
+ "sync/atomic"
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/catalog"
+)
+
+const broadcastQueueSize = 128
+
+type managerStore interface {
+ LoadSnapshot(context.Context) (Snapshot, error)
+ DatabaseGeneration(context.Context) (int64, error)
+ PublishChange(context.Context, ChangeEvent) error
+ Subscribe(context.Context) (<-chan ChangeMessage, func() error, error)
+ RedisEnabled() bool
+}
+
+type Manager struct {
+ store managerStore
+ catalog *catalog.Catalog
+ authenticator *auth.StaticAuthenticator
+ logger *slog.Logger
+ pollInterval time.Duration
+ generation atomic.Int64
+ redisConnected atomic.Bool
+ reloadMu sync.Mutex
+ broadcasts chan ChangeEvent
+}
+
+func NewManager(store managerStore, modelCatalog *catalog.Catalog, authenticator *auth.StaticAuthenticator, logger *slog.Logger, pollInterval time.Duration) *Manager {
+ if pollInterval <= 0 {
+ pollInterval = 30 * time.Second
+ }
+ return &Manager{
+ store: store, catalog: modelCatalog, authenticator: authenticator,
+ logger: logger, pollInterval: pollInterval, broadcasts: make(chan ChangeEvent, broadcastQueueSize),
+ }
+}
+
+func (m *Manager) Reload(ctx context.Context) (int64, error) {
+ m.reloadMu.Lock()
+ defer m.reloadMu.Unlock()
+ snapshot, err := m.store.LoadSnapshot(ctx)
+ if err != nil {
+ return m.generation.Load(), err
+ }
+ m.catalog.Replace(snapshot.Models)
+ m.authenticator.ReplaceHashed(snapshot.APIKeys)
+ m.generation.Store(snapshot.Generation)
+ m.logger.Info("control_plane_reloaded", "generation", snapshot.Generation, "models", len(snapshot.Models), "api_keys", len(snapshot.APIKeys))
+ return snapshot.Generation, nil
+}
+
+func (m *Manager) AfterMutation(ctx context.Context, generation int64, resource, id string) error {
+ loadedGeneration, err := m.Reload(ctx)
+ if err != nil {
+ return err
+ }
+ if !m.store.RedisEnabled() {
+ return nil
+ }
+ if loadedGeneration > generation {
+ generation = loadedGeneration
+ }
+ event := newChange(generation, resource, id)
+ select {
+ case m.broadcasts <- event:
+ default:
+ m.logger.Warn("control_plane_publish_dropped", "generation", generation, "resource", resource, "id", id, "error", "broadcast queue is full")
+ }
+ return nil
+}
+
+func (m *Manager) Generation() int64 {
+ return m.generation.Load()
+}
+
+func (m *Manager) RedisConfigured() bool {
+ return m.store.RedisEnabled()
+}
+
+func (m *Manager) RedisConnected() bool {
+ return m.redisConnected.Load()
+}
+
+func (m *Manager) Run(ctx context.Context) {
+ var workers sync.WaitGroup
+ if m.store.RedisEnabled() {
+ workers.Add(2)
+ go func() {
+ defer workers.Done()
+ m.runSubscriptions(ctx)
+ }()
+ go func() {
+ defer workers.Done()
+ m.runBroadcasts(ctx)
+ }()
+ } else {
+ m.logger.Info("control_plane_redis_disabled", "fallback", "postgres_polling")
+ }
+ m.runPolling(ctx)
+ workers.Wait()
+}
+
+func (m *Manager) runPolling(ctx context.Context) {
+ ticker := time.NewTicker(m.pollInterval)
+ defer ticker.Stop()
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ generation, err := m.store.DatabaseGeneration(ctx)
+ if err != nil {
+ m.logger.Warn("control_plane_generation_check_failed", "error", err)
+ continue
+ }
+ if generation > m.generation.Load() {
+ if _, err := m.Reload(ctx); err != nil {
+ m.logger.Error("control_plane_reload_failed", "source", "postgres", "error", err)
+ }
+ }
+ }
+ }
+}
+
+func (m *Manager) runSubscriptions(ctx context.Context) {
+ retryDelay := m.pollInterval
+ if retryDelay > time.Second {
+ retryDelay = time.Second
+ }
+ if retryDelay <= 0 {
+ retryDelay = time.Second
+ }
+
+ for ctx.Err() == nil {
+ messages, closeSubscription, err := m.store.Subscribe(ctx)
+ if err != nil {
+ m.redisConnected.Store(false)
+ m.logger.Warn("control_plane_subscription_failed", "error", err)
+ if !waitForRetry(ctx, retryDelay) {
+ return
+ }
+ continue
+ }
+ m.redisConnected.Store(true)
+ m.logger.Info("control_plane_subscription_connected")
+
+ closed := false
+ for !closed {
+ select {
+ case <-ctx.Done():
+ closed = true
+ case message, ok := <-messages:
+ if !ok {
+ closed = true
+ continue
+ }
+ var event ChangeEvent
+ if json.Unmarshal([]byte(message.Payload), &event) != nil || event.Generation <= m.generation.Load() {
+ continue
+ }
+ if _, err := m.Reload(ctx); err != nil {
+ m.logger.Error("control_plane_reload_failed", "source", "redis", "error", err)
+ }
+ }
+ }
+ m.redisConnected.Store(false)
+ if closeSubscription != nil {
+ _ = closeSubscription()
+ }
+ if ctx.Err() == nil {
+ m.logger.Warn("control_plane_subscription_disconnected")
+ if !waitForRetry(ctx, retryDelay) {
+ return
+ }
+ }
+ }
+}
+
+func (m *Manager) runBroadcasts(ctx context.Context) {
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case event := <-m.broadcasts:
+ publishContext, cancel := context.WithTimeout(ctx, 2*time.Second)
+ err := m.store.PublishChange(publishContext, event)
+ cancel()
+ if err != nil {
+ m.logger.Warn("control_plane_publish_failed", "generation", event.Generation, "resource", event.Resource, "id", event.ID, "error", err)
+ }
+ }
+ }
+}
+
+func waitForRetry(ctx context.Context, delay time.Duration) bool {
+ timer := time.NewTimer(delay)
+ defer timer.Stop()
+ select {
+ case <-ctx.Done():
+ return false
+ case <-timer.C:
+ return true
+ }
+}
diff --git a/internal/controlplane/manager_test.go b/internal/controlplane/manager_test.go
new file mode 100644
index 0000000..8dd4012
--- /dev/null
+++ b/internal/controlplane/manager_test.go
@@ -0,0 +1,212 @@
+package controlplane
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "log/slog"
+ "strings"
+ "sync"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/catalog"
+)
+
+type fakeManagerStore struct {
+ redisEnabled bool
+ snapshot atomic.Value
+ databaseGeneration atomic.Int64
+ publishCalls atomic.Int64
+ subscribeCalls atomic.Int64
+ publishErr error
+ published chan ChangeEvent
+ subscribe func(context.Context, int64) (<-chan ChangeMessage, func() error, error)
+}
+
+type safeLogBuffer struct {
+ mu sync.Mutex
+ buf bytes.Buffer
+}
+
+func (b *safeLogBuffer) Write(data []byte) (int, error) {
+ b.mu.Lock()
+ defer b.mu.Unlock()
+ return b.buf.Write(data)
+}
+
+func (b *safeLogBuffer) String() string {
+ b.mu.Lock()
+ defer b.mu.Unlock()
+ return b.buf.String()
+}
+
+func newFakeManagerStore(generation int64) *fakeManagerStore {
+ store := &fakeManagerStore{published: make(chan ChangeEvent, 8)}
+ store.snapshot.Store(Snapshot{Generation: generation})
+ store.databaseGeneration.Store(generation)
+ return store
+}
+
+func (s *fakeManagerStore) LoadSnapshot(context.Context) (Snapshot, error) {
+ return s.snapshot.Load().(Snapshot), nil
+}
+
+func (s *fakeManagerStore) DatabaseGeneration(context.Context) (int64, error) {
+ return s.databaseGeneration.Load(), nil
+}
+
+func (s *fakeManagerStore) PublishChange(_ context.Context, event ChangeEvent) error {
+ s.publishCalls.Add(1)
+ s.published <- event
+ return s.publishErr
+}
+
+func (s *fakeManagerStore) Subscribe(ctx context.Context) (<-chan ChangeMessage, func() error, error) {
+ call := s.subscribeCalls.Add(1)
+ if s.subscribe != nil {
+ return s.subscribe(ctx, call)
+ }
+ channel := make(chan ChangeMessage)
+ return channel, func() error { return nil }, nil
+}
+
+func (s *fakeManagerStore) RedisEnabled() bool {
+ return s.redisEnabled
+}
+
+func newTestManager(store managerStore, logger *slog.Logger, interval time.Duration) *Manager {
+ return NewManager(store, catalog.NewModels(nil), auth.NewDynamic(nil, false), logger, interval)
+}
+
+func TestAfterMutationReloadsLocallyWhenRedisPublishFails(t *testing.T) {
+ store := newFakeManagerStore(7)
+ store.redisEnabled = true
+ store.publishErr = errors.New("redis unavailable")
+ var logs safeLogBuffer
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&logs, nil)), 10*time.Millisecond)
+ ctx, cancel := context.WithCancel(context.Background())
+ done := make(chan struct{})
+ go func() {
+ manager.Run(ctx)
+ close(done)
+ }()
+
+ if err := manager.AfterMutation(context.Background(), 7, "model", "model-1"); err != nil {
+ t.Fatalf("mutation unexpectedly failed: %v", err)
+ }
+ if manager.Generation() != 7 {
+ t.Fatalf("local generation = %d, want 7", manager.Generation())
+ }
+ select {
+ case event := <-store.published:
+ if event.Generation != 7 || event.Resource != "model" {
+ t.Fatalf("unexpected event: %+v", event)
+ }
+ case <-time.After(time.Second):
+ t.Fatal("broadcast was not attempted")
+ }
+ waitUntil(t, time.Second, func() bool { return strings.Contains(logs.String(), "control_plane_publish_failed") })
+
+ cancel()
+ select {
+ case <-done:
+ case <-time.After(time.Second):
+ t.Fatal("manager did not stop")
+ }
+}
+
+func TestPollingContinuesWhileRedisSubscribeIsBlocked(t *testing.T) {
+ store := newFakeManagerStore(1)
+ store.redisEnabled = true
+ store.subscribe = func(ctx context.Context, _ int64) (<-chan ChangeMessage, func() error, error) {
+ <-ctx.Done()
+ return nil, nil, ctx.Err()
+ }
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&safeLogBuffer{}, nil)), 10*time.Millisecond)
+ if _, err := manager.Reload(context.Background()); err != nil {
+ t.Fatal(err)
+ }
+
+ ctx, cancel := context.WithCancel(context.Background())
+ done := make(chan struct{})
+ go func() {
+ manager.Run(ctx)
+ close(done)
+ }()
+ store.snapshot.Store(Snapshot{Generation: 2})
+ store.databaseGeneration.Store(2)
+ waitUntil(t, time.Second, func() bool { return manager.Generation() == 2 })
+
+ cancel()
+ select {
+ case <-done:
+ case <-time.After(time.Second):
+ t.Fatal("manager did not stop")
+ }
+}
+
+func TestSubscriptionReconnectsAfterChannelCloses(t *testing.T) {
+ store := newFakeManagerStore(1)
+ store.redisEnabled = true
+ first := make(chan ChangeMessage)
+ second := make(chan ChangeMessage, 1)
+ store.subscribe = func(_ context.Context, call int64) (<-chan ChangeMessage, func() error, error) {
+ if call == 1 {
+ return first, func() error { return nil }, nil
+ }
+ return second, func() error { return nil }, nil
+ }
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&safeLogBuffer{}, nil)), 10*time.Millisecond)
+ if _, err := manager.Reload(context.Background()); err != nil {
+ t.Fatal(err)
+ }
+
+ ctx, cancel := context.WithCancel(context.Background())
+ done := make(chan struct{})
+ go func() {
+ manager.Run(ctx)
+ close(done)
+ }()
+ waitUntil(t, time.Second, func() bool { return store.subscribeCalls.Load() == 1 })
+ close(first)
+ waitUntil(t, time.Second, func() bool { return store.subscribeCalls.Load() >= 2 && manager.RedisConnected() })
+ store.snapshot.Store(Snapshot{Generation: 2})
+ second <- ChangeMessage{Payload: `{"generation":2,"resource":"model"}`}
+ waitUntil(t, time.Second, func() bool { return manager.Generation() == 2 })
+
+ cancel()
+ select {
+ case <-done:
+ case <-time.After(time.Second):
+ t.Fatal("manager did not stop")
+ }
+}
+
+func TestRedisCanBeDisabled(t *testing.T) {
+ store := newFakeManagerStore(3)
+ manager := newTestManager(store, slog.New(slog.NewTextHandler(&bytes.Buffer{}, nil)), 10*time.Millisecond)
+ if err := manager.AfterMutation(context.Background(), 3, "tenant", "tenant-1"); err != nil {
+ t.Fatal(err)
+ }
+ if manager.RedisConfigured() || manager.RedisConnected() {
+ t.Fatal("Redis unexpectedly reported as available")
+ }
+ if store.publishCalls.Load() != 0 || store.subscribeCalls.Load() != 0 {
+ t.Fatal("Redis operations were attempted while disabled")
+ }
+}
+
+func waitUntil(t *testing.T, timeout time.Duration, condition func() bool) {
+ t.Helper()
+ deadline := time.Now().Add(timeout)
+ for time.Now().Before(deadline) {
+ if condition() {
+ return
+ }
+ time.Sleep(5 * time.Millisecond)
+ }
+ t.Fatal("condition was not met before timeout")
+}
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
new file mode 100644
index 0000000..a781994
--- /dev/null
+++ b/internal/controlplane/mutations.go
@@ -0,0 +1,281 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/rand"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "net/url"
+ "regexp"
+ "strings"
+
+ "github.com/jackc/pgx/v5"
+)
+
+var (
+ ErrNotFound = errors.New("control-plane resource not found")
+ slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`)
+)
+
+func (s *Store) CreateTenant(ctx context.Context, input CreateTenantInput) (Tenant, int64, error) {
+ input.Slug = strings.ToLower(strings.TrimSpace(input.Slug))
+ input.Name = strings.TrimSpace(input.Name)
+ if !slugPattern.MatchString(input.Slug) || input.Name == "" {
+ return Tenant{}, 0, errors.New("tenant requires a 3-64 character lowercase slug and a name")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Tenant{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Tenant
+ err = tx.QueryRow(ctx, `
+ INSERT INTO tenants (slug, name) VALUES ($1, $2)
+ RETURNING id::text, slug, name, status, created_at`, input.Slug, input.Name,
+ ).Scan(&result.ID, &result.Slug, &result.Name, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return Tenant{}, 0, fmt.Errorf("create tenant: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Tenant{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Tenant{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) CreateProject(ctx context.Context, input CreateProjectInput) (Project, int64, error) {
+ input.Slug = strings.ToLower(strings.TrimSpace(input.Slug))
+ input.Name = strings.TrimSpace(input.Name)
+ if input.TenantID == "" || !slugPattern.MatchString(input.Slug) || input.Name == "" {
+ return Project{}, 0, errors.New("project requires tenant_id, a 3-64 character lowercase slug, and a name")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Project{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Project
+ err = tx.QueryRow(ctx, `
+ INSERT INTO projects (tenant_id, slug, name) VALUES ($1, $2, $3)
+ RETURNING id::text, tenant_id::text, slug, name, status, created_at`, input.TenantID, input.Slug, input.Name,
+ ).Scan(&result.ID, &result.TenantID, &result.Slug, &result.Name, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return Project{}, 0, fmt.Errorf("create project: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Project{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Project{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (CreatedAPIKey, int64, error) {
+ input.Name = strings.TrimSpace(input.Name)
+ if input.TenantID == "" || input.ProjectID == "" || input.Name == "" {
+ return CreatedAPIKey{}, 0, errors.New("API key requires tenant_id, project_id, and name")
+ }
+ if len(input.Scopes) == 0 {
+ input.Scopes = []string{"inference"}
+ }
+ scopes := uniqueStrings(input.Scopes)
+ scopesJSON, _ := json.Marshal(scopes)
+ random := make([]byte, 32)
+ if _, err := rand.Read(random); err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("generate API key: %w", err)
+ }
+ rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random)
+ hash := sha256.Sum256([]byte(rawKey))
+ prefix := rawKey[:min(18, len(rawKey))] + "..."
+
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ 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)
+ if err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err)
+ }
+ result.Scopes = scopes
+ result.Key = rawKey
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) {
+ return s.toggle(ctx, `UPDATE api_keys SET status = 'revoked', revoked_at = now() WHERE id = $1 AND status <> 'revoked'`, id)
+}
+
+func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) {
+ input.Name = strings.TrimSpace(input.Name)
+ 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")
+ }
+ parsed, err := url.Parse(input.BaseURL)
+ if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") {
+ return Provider{}, 0, errors.New("provider base_url must be an absolute http(s) URL")
+ }
+ ciphertext, err := s.cipher.Encrypt(input.APIKey)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ 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)
+ if err != nil {
+ return Provider{}, 0, fmt.Errorf("create provider: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Provider{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
+ return s.toggle(ctx, `UPDATE providers SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
+}
+
+func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, int64, error) {
+ input.PublicID = strings.TrimSpace(input.PublicID)
+ input.OwnedBy = strings.TrimSpace(input.OwnedBy)
+ if input.PublicID == "" || len(input.Routes) == 0 {
+ return Model{}, 0, errors.New("model requires public_id and at least one route")
+ }
+ for i := range input.Routes {
+ input.Routes[i].ProviderID = strings.TrimSpace(input.Routes[i].ProviderID)
+ input.Routes[i].UpstreamModel = strings.TrimSpace(input.Routes[i].UpstreamModel)
+ if input.Routes[i].Weight == 0 {
+ input.Routes[i].Weight = 1
+ }
+ if input.Routes[i].ProviderID == "" || input.Routes[i].UpstreamModel == "" || input.Routes[i].Priority < 0 || input.Routes[i].Weight < 1 || input.Routes[i].Weight > 100 {
+ return Model{}, 0, fmt.Errorf("route %d has invalid provider, upstream model, priority, or weight", i+1)
+ }
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Model{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Model
+ err = tx.QueryRow(ctx, `
+ INSERT INTO models (public_id, owned_by) VALUES ($1, $2)
+ RETURNING id::text, public_id, owned_by, enabled, created_at`, input.PublicID, input.OwnedBy,
+ ).Scan(&result.ID, &result.PublicID, &result.OwnedBy, &result.Enabled, &result.CreatedAt)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model: %w", err)
+ }
+ result.Routes = make([]Route, 0, len(input.Routes))
+ for _, route := range input.Routes {
+ var created Route
+ err := tx.QueryRow(ctx, `
+ INSERT INTO model_routes (model_id, provider_id, upstream_model, priority, weight)
+ VALUES ($1, $2, $3, $4, $5)
+ RETURNING id::text, provider_id::text, upstream_model, priority, weight, enabled`,
+ result.ID, route.ProviderID, route.UpstreamModel, route.Priority, route.Weight,
+ ).Scan(&created.ID, &created.ProviderID, &created.UpstreamModel, &created.Priority, &created.Weight, &created.Enabled)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model route: %w", err)
+ }
+ result.Routes = append(result.Routes, created)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Model{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Model{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) SetModelEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
+ return s.toggle(ctx, `UPDATE models SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
+}
+
+func (s *Store) toggle(ctx context.Context, query, id string, args ...any) (int64, error) {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ parameters := append([]any{id}, args...)
+ command, err := tx.Exec(ctx, query, parameters...)
+ if err != nil {
+ return 0, err
+ }
+ if command.RowsAffected() == 0 {
+ return 0, ErrNotFound
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return 0, err
+ }
+ return generation, nil
+}
+
+func bumpGeneration(ctx context.Context, tx pgx.Tx) (int64, error) {
+ var generation int64
+ err := tx.QueryRow(ctx, `
+ UPDATE control_state SET generation = generation + 1, updated_at = now()
+ WHERE singleton = TRUE RETURNING generation`,
+ ).Scan(&generation)
+ if err != nil {
+ return 0, fmt.Errorf("advance control-plane generation: %w", err)
+ }
+ return generation, nil
+}
+
+func uniqueStrings(values []string) []string {
+ seen := make(map[string]struct{}, len(values))
+ result := make([]string, 0, len(values))
+ for _, value := range values {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ continue
+ }
+ if _, exists := seen[value]; exists {
+ continue
+ }
+ seen[value] = struct{}{}
+ result = append(result, value)
+ }
+ return result
+}
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go
new file mode 100644
index 0000000..732cb46
--- /dev/null
+++ b/internal/controlplane/queries.go
@@ -0,0 +1,146 @@
+package controlplane
+
+import (
+ "context"
+ "encoding/json"
+ "fmt"
+)
+
+func (s *Store) Overview(ctx context.Context) (Overview, error) {
+ var result Overview
+ err := s.db.QueryRow(ctx, `
+ SELECT
+ (SELECT generation FROM control_state WHERE singleton = TRUE),
+ (SELECT count(*) FROM tenants WHERE status = 'active'),
+ (SELECT count(*) FROM projects WHERE status = 'active'),
+ (SELECT count(*) FROM api_keys WHERE status = 'active'),
+ (SELECT count(*) FROM providers WHERE enabled = TRUE),
+ (SELECT count(*) FROM models WHERE enabled = TRUE)`,
+ ).Scan(&result.Generation, &result.Tenants, &result.Projects, &result.APIKeys, &result.Providers, &result.Models)
+ if err != nil {
+ return Overview{}, fmt.Errorf("query control-plane overview: %w", err)
+ }
+ return result, nil
+}
+
+func (s *Store) ListTenants(ctx context.Context) ([]Tenant, error) {
+ rows, err := s.db.Query(ctx, `SELECT id::text, slug, name, status, created_at FROM tenants ORDER BY created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query tenants: %w", err)
+ }
+ defer rows.Close()
+ result := make([]Tenant, 0)
+ for rows.Next() {
+ var item Tenant
+ if err := rows.Scan(&item.ID, &item.Slug, &item.Name, &item.Status, &item.CreatedAt); err != nil {
+ return nil, fmt.Errorf("scan tenant: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListProjects(ctx context.Context) ([]Project, error) {
+ rows, err := s.db.Query(ctx, `SELECT id::text, tenant_id::text, slug, name, status, created_at FROM projects ORDER BY created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query projects: %w", err)
+ }
+ defer rows.Close()
+ result := make([]Project, 0)
+ for rows.Next() {
+ var item Project
+ if err := rows.Scan(&item.ID, &item.TenantID, &item.Slug, &item.Name, &item.Status, &item.CreatedAt); err != nil {
+ return nil, fmt.Errorf("scan project: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListAPIKeys(ctx context.Context) ([]APIKey, error) {
+ rows, err := s.db.Query(ctx, `
+ SELECT id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at
+ FROM api_keys ORDER BY created_at DESC`)
+ if err != nil {
+ return nil, fmt.Errorf("query API keys: %w", err)
+ }
+ defer rows.Close()
+ 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 {
+ 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)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+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
+ 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 {
+ return nil, fmt.Errorf("query providers: %w", err)
+ }
+ defer rows.Close()
+ 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 {
+ return nil, fmt.Errorf("scan provider: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
+ rows, err := s.db.Query(ctx, `SELECT id::text, public_id, owned_by, enabled, created_at FROM models ORDER BY public_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query models: %w", err)
+ }
+ models := make([]Model, 0)
+ positions := make(map[string]int)
+ for rows.Next() {
+ var item Model
+ if err := rows.Scan(&item.ID, &item.PublicID, &item.OwnedBy, &item.Enabled, &item.CreatedAt); err != nil {
+ rows.Close()
+ return nil, fmt.Errorf("scan model: %w", err)
+ }
+ item.Routes = []Route{}
+ positions[item.ID] = len(models)
+ models = append(models, item)
+ }
+ if err := rows.Err(); err != nil {
+ rows.Close()
+ return nil, err
+ }
+ rows.Close()
+
+ routeRows, err := s.db.Query(ctx, `
+ SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol,
+ r.upstream_model, r.priority, r.weight, r.enabled
+ FROM model_routes r JOIN providers p ON p.id = r.provider_id
+ ORDER BY r.priority, r.created_at`)
+ if err != nil {
+ return nil, fmt.Errorf("query model routes: %w", err)
+ }
+ defer routeRows.Close()
+ 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 {
+ return nil, fmt.Errorf("scan model route: %w", err)
+ }
+ if position, ok := positions[modelID]; ok {
+ models[position].Routes = append(models[position].Routes, route)
+ }
+ }
+ return models, routeRows.Err()
+}
diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql
new file mode 100644
index 0000000..25af7c9
--- /dev/null
+++ b/internal/controlplane/schema.sql
@@ -0,0 +1,82 @@
+CREATE EXTENSION IF NOT EXISTS pgcrypto;
+
+CREATE TABLE IF NOT EXISTS control_state (
+ singleton BOOLEAN PRIMARY KEY DEFAULT TRUE CHECK (singleton),
+ generation BIGINT NOT NULL DEFAULT 0,
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+INSERT INTO control_state (singleton) VALUES (TRUE) ON CONFLICT DO NOTHING;
+
+CREATE TABLE IF NOT EXISTS tenants (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ slug TEXT NOT NULL UNIQUE,
+ name TEXT NOT NULL,
+ status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'suspended')),
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS projects (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ slug TEXT NOT NULL,
+ name TEXT NOT NULL,
+ status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'suspended')),
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ UNIQUE (tenant_id, slug),
+ UNIQUE (id, tenant_id)
+);
+
+CREATE TABLE IF NOT EXISTS api_keys (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ project_id UUID NOT NULL,
+ name TEXT NOT NULL,
+ key_prefix TEXT NOT NULL,
+ key_hash BYTEA NOT NULL UNIQUE CHECK (octet_length(key_hash) = 32),
+ scopes JSONB NOT NULL DEFAULT '["inference"]'::jsonb,
+ status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'revoked')),
+ last_used_at TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ revoked_at TIMESTAMPTZ,
+ FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE
+);
+
+CREATE TABLE IF NOT EXISTS providers (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ name TEXT NOT NULL UNIQUE,
+ protocol TEXT NOT NULL CHECK (protocol IN ('openai', 'anthropic')),
+ 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()
+);
+
+CREATE TABLE IF NOT EXISTS models (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ public_id TEXT NOT NULL UNIQUE,
+ owned_by TEXT NOT NULL DEFAULT '',
+ enabled BOOLEAN NOT NULL DEFAULT TRUE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+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,
+ provider_id UUID NOT NULL REFERENCES providers(id) ON DELETE RESTRICT,
+ upstream_model TEXT NOT NULL,
+ priority INTEGER NOT NULL DEFAULT 0 CHECK (priority >= 0),
+ weight INTEGER NOT NULL DEFAULT 1 CHECK (weight BETWEEN 1 AND 100),
+ enabled BOOLEAN NOT NULL DEFAULT TRUE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ UNIQUE (model_id, provider_id, upstream_model)
+);
+
+CREATE INDEX IF NOT EXISTS api_keys_active_hash_idx ON api_keys (key_hash) WHERE status = 'active';
+CREATE INDEX IF NOT EXISTS projects_tenant_idx ON projects (tenant_id);
+CREATE INDEX IF NOT EXISTS model_routes_model_idx ON model_routes (model_id) WHERE enabled;
+CREATE INDEX IF NOT EXISTS model_routes_provider_idx ON model_routes (provider_id) WHERE enabled;
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go
new file mode 100644
index 0000000..c8931ab
--- /dev/null
+++ b/internal/controlplane/snapshot.go
@@ -0,0 +1,150 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/sha256"
+ "encoding/json"
+ "fmt"
+ "strings"
+
+ "aigw/internal/auth"
+ "aigw/internal/domain"
+
+ "github.com/jackc/pgx/v5"
+)
+
+func (s *Store) LoadSnapshot(ctx context.Context) (Snapshot, error) {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly})
+ if err != nil {
+ return Snapshot{}, fmt.Errorf("begin snapshot transaction: %w", err)
+ }
+ defer tx.Rollback(ctx)
+
+ var result Snapshot
+ if err := tx.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&result.Generation); err != nil {
+ return Snapshot{}, fmt.Errorf("read snapshot generation: %w", err)
+ }
+
+ providers, err := s.loadProviders(ctx, tx)
+ if err != nil {
+ return Snapshot{}, err
+ }
+ result.Models, err = loadModels(ctx, tx, providers)
+ if err != nil {
+ return Snapshot{}, err
+ }
+ result.APIKeys, err = loadAPIKeys(ctx, tx)
+ if err != nil {
+ return Snapshot{}, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Snapshot{}, fmt.Errorf("commit snapshot transaction: %w", err)
+ }
+ return result, nil
+}
+
+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
+ FROM providers
+ WHERE enabled = TRUE
+ ORDER BY name`)
+ if err != nil {
+ return nil, fmt.Errorf("query providers: %w", err)
+ }
+ defer rows.Close()
+ providers := make(map[string]domain.Provider)
+ for rows.Next() {
+ var id, name, protocol, baseURL string
+ var ciphertext []byte
+ if err := rows.Scan(&id, &name, &protocol, &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}
+ }
+ if err := rows.Err(); err != nil {
+ return nil, fmt.Errorf("read providers: %w", err)
+ }
+ return providers, nil
+}
+
+func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) {
+ rows, err := tx.Query(ctx, `
+ SELECT m.public_id, m.owned_by, r.provider_id::text, r.upstream_model, r.priority, r.weight
+ FROM models m
+ JOIN model_routes r ON r.model_id = m.id AND r.enabled = TRUE
+ JOIN providers p ON p.id = r.provider_id AND p.enabled = TRUE
+ WHERE m.enabled = TRUE
+ ORDER BY m.public_id, r.priority, r.created_at`)
+ if err != nil {
+ return nil, fmt.Errorf("query model routes: %w", err)
+ }
+ defer rows.Close()
+ models := make([]domain.Model, 0)
+ index := make(map[string]int)
+ for rows.Next() {
+ var publicID, ownedBy, providerID, upstreamModel string
+ var priority, weight int
+ if err := rows.Scan(&publicID, &ownedBy, &providerID, &upstreamModel, &priority, &weight); err != nil {
+ return nil, fmt.Errorf("scan model route: %w", err)
+ }
+ provider, ok := providers[providerID]
+ if !ok {
+ continue
+ }
+ position, exists := index[publicID]
+ if !exists {
+ position = len(models)
+ index[publicID] = position
+ models = append(models, domain.Model{ID: publicID, OwnedBy: ownedBy})
+ }
+ models[position].Routes = append(models[position].Routes, domain.Route{
+ Provider: provider, UpstreamModel: upstreamModel, Priority: priority, Weight: weight,
+ })
+ }
+ if err := rows.Err(); err != nil {
+ return nil, fmt.Errorf("read model routes: %w", err)
+ }
+ return models, nil
+}
+
+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
+ 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'`)
+ if err != nil {
+ return nil, fmt.Errorf("query API keys: %w", err)
+ }
+ defer rows.Close()
+ 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 {
+ return nil, fmt.Errorf("scan API key: %w", err)
+ }
+ if len(hashBytes) != sha256.Size {
+ return nil, fmt.Errorf("API key %s has invalid hash length", keyID)
+ }
+ var hash [sha256.Size]byte
+ copy(hash[:], hashBytes)
+ var scopes []string
+ if err := json.Unmarshal(scopesJSON, &scopes); err != nil {
+ return nil, fmt.Errorf("decode API key %s scopes: %w", keyID, err)
+ }
+ records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{
+ KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes,
+ }})
+ }
+ if err := rows.Err(); err != nil {
+ return nil, fmt.Errorf("read API keys: %w", err)
+ }
+ return records, nil
+}
diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go
new file mode 100644
index 0000000..833f846
--- /dev/null
+++ b/internal/controlplane/store.go
@@ -0,0 +1,170 @@
+package controlplane
+
+import (
+ "context"
+ _ "embed"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "aigw/internal/security"
+
+ "github.com/jackc/pgx/v5/pgxpool"
+ "github.com/redis/go-redis/v9"
+)
+
+//go:embed schema.sql
+var schemaSQL string
+
+var ErrRedisDisabled = errors.New("Redis propagation is disabled")
+
+type Options struct {
+ DatabaseURL string
+ RedisURL string
+ CredentialKey string
+ RedisChannel string
+ VersionCacheKey string
+}
+
+type Store struct {
+ db *pgxpool.Pool
+ redis *redis.Client
+ cipher *security.CredentialCipher
+ redisChannel string
+ versionCacheKey string
+}
+
+func NewStore(ctx context.Context, options Options) (*Store, error) {
+ cipher, err := security.NewCredentialCipher(options.CredentialKey)
+ if err != nil {
+ return nil, err
+ }
+ db, err := pgxpool.New(ctx, options.DatabaseURL)
+ if err != nil {
+ return nil, fmt.Errorf("configure PostgreSQL: %w", err)
+ }
+ if err := db.Ping(ctx); err != nil {
+ db.Close()
+ return nil, fmt.Errorf("connect PostgreSQL: %w", err)
+ }
+ var redisClient *redis.Client
+ if strings.TrimSpace(options.RedisURL) != "" {
+ redisOptions, err := redis.ParseURL(options.RedisURL)
+ if err != nil {
+ db.Close()
+ return nil, fmt.Errorf("parse Redis URL: %w", err)
+ }
+ redisClient = redis.NewClient(redisOptions)
+ }
+ return &Store{
+ db: db, redis: redisClient, cipher: cipher,
+ redisChannel: options.RedisChannel, versionCacheKey: options.VersionCacheKey,
+ }, nil
+}
+
+func (s *Store) Close() error {
+ s.db.Close()
+ if s.redis == nil {
+ return nil
+ }
+ return s.redis.Close()
+}
+
+func (s *Store) RedisEnabled() bool {
+ return s.redis != nil
+}
+
+func (s *Store) Migrate(ctx context.Context) error {
+ if _, err := s.db.Exec(ctx, schemaSQL); err != nil {
+ return fmt.Errorf("apply control-plane schema: %w", err)
+ }
+ return nil
+}
+
+func MigrateDatabase(ctx context.Context, databaseURL string) error {
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ return fmt.Errorf("configure PostgreSQL: %w", err)
+ }
+ defer db.Close()
+ if _, err := db.Exec(ctx, schemaSQL); err != nil {
+ return fmt.Errorf("apply control-plane schema: %w", err)
+ }
+ return nil
+}
+
+func (s *Store) DatabaseGeneration(ctx context.Context) (int64, error) {
+ var generation int64
+ err := s.db.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&generation)
+ if err != nil {
+ return 0, fmt.Errorf("read control-plane generation: %w", err)
+ }
+ return generation, nil
+}
+
+func (s *Store) RedisGeneration(ctx context.Context) (int64, error) {
+ if s.redis == nil {
+ return 0, ErrRedisDisabled
+ }
+ generation, err := s.redis.Get(ctx, s.versionCacheKey).Int64()
+ if errors.Is(err, redis.Nil) {
+ return 0, nil
+ }
+ return generation, err
+}
+
+func (s *Store) PublishChange(ctx context.Context, event ChangeEvent) error {
+ if s.redis == nil {
+ return ErrRedisDisabled
+ }
+ payload, err := json.Marshal(event)
+ if err != nil {
+ return err
+ }
+ pipeline := s.redis.TxPipeline()
+ pipeline.Set(ctx, s.versionCacheKey, event.Generation, 0)
+ pipeline.Publish(ctx, s.redisChannel, payload)
+ _, err = pipeline.Exec(ctx)
+ if err != nil {
+ return fmt.Errorf("publish control-plane change: %w", err)
+ }
+ return nil
+}
+
+func (s *Store) Subscribe(ctx context.Context) (<-chan ChangeMessage, func() error, error) {
+ if s.redis == nil {
+ return nil, nil, ErrRedisDisabled
+ }
+ pubsub := s.redis.Subscribe(ctx, s.redisChannel)
+ if _, err := pubsub.Receive(ctx); err != nil {
+ _ = pubsub.Close()
+ return nil, nil, fmt.Errorf("subscribe control-plane changes: %w", err)
+ }
+ messages := make(chan ChangeMessage)
+ redisMessages := pubsub.Channel()
+ go func() {
+ defer close(messages)
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case message, ok := <-redisMessages:
+ if !ok {
+ return
+ }
+ select {
+ case messages <- ChangeMessage{Payload: message.Payload}:
+ case <-ctx.Done():
+ return
+ }
+ }
+ }
+ }()
+ return messages, pubsub.Close, nil
+}
+
+func newChange(generation int64, resource, id string) ChangeEvent {
+ return ChangeEvent{Generation: generation, Resource: resource, ID: id, ChangedAt: time.Now().UTC().Format(time.RFC3339Nano)}
+}
diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go
new file mode 100644
index 0000000..24f2843
--- /dev/null
+++ b/internal/controlplane/types.go
@@ -0,0 +1,138 @@
+package controlplane
+
+import (
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/domain"
+)
+
+type Tenant struct {
+ ID string `json:"id"`
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+ Status string `json:"status"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type Project struct {
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+ Status string `json:"status"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+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"`
+}
+
+type CreatedAPIKey struct {
+ APIKey
+ Key string `json:"key"`
+}
+
+type Provider struct {
+ ID string `json:"id"`
+ Name string `json:"name"`
+ Protocol string `json:"protocol"`
+ BaseURL string `json:"base_url"`
+ Enabled bool `json:"enabled"`
+ RouteCount int `json:"route_count"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+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"`
+}
+
+type Model struct {
+ ID string `json:"id"`
+ PublicID string `json:"public_id"`
+ OwnedBy string `json:"owned_by"`
+ Enabled bool `json:"enabled"`
+ Routes []Route `json:"routes"`
+ CreatedAt time.Time `json:"created_at"`
+}
+
+type Overview struct {
+ Generation int64 `json:"generation"`
+ RuntimeGeneration int64 `json:"runtime_generation"`
+ RedisConfigured bool `json:"redis_configured"`
+ RedisConnected bool `json:"redis_connected"`
+ Tenants int64 `json:"tenants"`
+ Projects int64 `json:"projects"`
+ APIKeys int64 `json:"api_keys"`
+ Providers int64 `json:"providers"`
+ Models int64 `json:"models"`
+}
+
+type Snapshot struct {
+ Generation int64
+ Models []domain.Model
+ APIKeys []auth.HashedKeyRecord
+}
+
+type ChangeEvent struct {
+ Generation int64 `json:"generation"`
+ Resource string `json:"resource"`
+ ID string `json:"id,omitempty"`
+ ChangedAt string `json:"changed_at"`
+}
+
+type ChangeMessage struct {
+ Payload string
+}
+
+type CreateTenantInput struct {
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+}
+
+type CreateProjectInput struct {
+ TenantID string `json:"tenant_id"`
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+}
+
+type CreateAPIKeyInput struct {
+ TenantID string `json:"tenant_id"`
+ ProjectID string `json:"project_id"`
+ Name string `json:"name"`
+ Scopes []string `json:"scopes"`
+}
+
+type CreateProviderInput struct {
+ Name string `json:"name"`
+ Protocol string `json:"protocol"`
+ BaseURL string `json:"base_url"`
+ APIKey string `json:"api_key"`
+}
+
+type RouteInput struct {
+ ProviderID string `json:"provider_id"`
+ UpstreamModel string `json:"upstream_model"`
+ Priority int `json:"priority"`
+ Weight int `json:"weight"`
+}
+
+type CreateModelInput struct {
+ PublicID string `json:"public_id"`
+ OwnedBy string `json:"owned_by"`
+ Routes []RouteInput `json:"routes"`
+}
diff --git a/internal/domain/types.go b/internal/domain/types.go
new file mode 100644
index 0000000..e2586ea
--- /dev/null
+++ b/internal/domain/types.go
@@ -0,0 +1,64 @@
+package domain
+
+import "time"
+
+type Protocol string
+
+const (
+ ProtocolOpenAI Protocol = "openai"
+ ProtocolAnthropic Protocol = "anthropic"
+)
+
+type Principal struct {
+ KeyID string
+ TenantID string
+ ProjectID string
+ Scopes []string
+}
+
+type Provider struct {
+ ID string
+ Protocol Protocol
+ BaseURL string
+ APIKey string
+}
+
+type Route struct {
+ Provider Provider
+ UpstreamModel string
+ Priority int
+ Weight int
+}
+
+type Model struct {
+ ID string
+ OwnedBy string
+ Routes []Route
+}
+
+type Usage struct {
+ InputTokens int64 `json:"input_tokens,omitempty"`
+ OutputTokens int64 `json:"output_tokens,omitempty"`
+ TotalTokens int64 `json:"total_tokens,omitempty"`
+ CacheCreationInputTokens int64 `json:"cache_creation_input_tokens,omitempty"`
+ CacheReadInputTokens int64 `json:"cache_read_input_tokens,omitempty"`
+}
+
+type UsageEvent struct {
+ RequestID string `json:"request_id"`
+ KeyID string `json:"key_id"`
+ TenantID string `json:"tenant_id"`
+ ProjectID string `json:"project_id"`
+ PublicModel string `json:"public_model"`
+ ProviderID string `json:"provider_id,omitempty"`
+ UpstreamModel string `json:"upstream_model,omitempty"`
+ Protocol Protocol `json:"protocol"`
+ Stream bool `json:"stream"`
+ StatusCode int `json:"status_code"`
+ Success bool `json:"success"`
+ ErrorType string `json:"error_type,omitempty"`
+ Attempts int `json:"attempts"`
+ StartedAt time.Time `json:"started_at"`
+ DurationMS int64 `json:"duration_ms"`
+ Usage Usage `json:"usage"`
+}
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go
new file mode 100644
index 0000000..ed7939f
--- /dev/null
+++ b/internal/httpapi/api.go
@@ -0,0 +1,369 @@
+package httpapi
+
+import (
+ "context"
+ "crypto/rand"
+ "encoding/hex"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "log/slog"
+ "net/http"
+ "strings"
+ "time"
+
+ "aigw/internal/apierror"
+ "aigw/internal/auth"
+ "aigw/internal/catalog"
+ "aigw/internal/domain"
+ "aigw/internal/provider"
+ "aigw/internal/routing"
+ "aigw/internal/telemetry"
+ "aigw/internal/usage"
+)
+
+type requestIDKey struct{}
+
+type API struct {
+ authenticator auth.Authenticator
+ catalog *catalog.Catalog
+ router *routing.Router
+ forwarder *provider.Forwarder
+ usageSink telemetry.UsageSink
+ metrics *telemetry.Metrics
+ logger *slog.Logger
+ maxBodyBytes int64
+ exposeMetrics bool
+}
+
+type Options struct {
+ Authenticator auth.Authenticator
+ Catalog *catalog.Catalog
+ Router *routing.Router
+ Forwarder *provider.Forwarder
+ UsageSink telemetry.UsageSink
+ Metrics *telemetry.Metrics
+ Logger *slog.Logger
+ MaxBodyBytes int64
+ ExposeMetrics bool
+}
+
+func New(options Options) *API {
+ return &API{
+ authenticator: options.Authenticator,
+ catalog: options.Catalog,
+ router: options.Router,
+ forwarder: options.Forwarder,
+ usageSink: options.UsageSink,
+ metrics: options.Metrics,
+ logger: options.Logger,
+ maxBodyBytes: options.MaxBodyBytes,
+ exposeMetrics: options.ExposeMetrics,
+ }
+}
+
+func (a *API) Handler() http.Handler {
+ mux := http.NewServeMux()
+ mux.HandleFunc("GET /healthz", a.health)
+ mux.HandleFunc("GET /readyz", a.health)
+ if a.exposeMetrics {
+ mux.Handle("GET /metrics", a.metrics)
+ }
+
+ mux.HandleFunc("GET /v1/models", a.openAIModels)
+ 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("GET /anthropic/v1/models", a.anthropicModels)
+ mux.HandleFunc("GET /api/anthropic/v1/models", a.anthropicModels)
+ mux.HandleFunc("POST /anthropic/v1/messages", a.anthropicMessages)
+ mux.HandleFunc("POST /api/anthropic/v1/messages", a.anthropicMessages)
+
+ return a.withRequestID(a.recoverPanics(mux))
+}
+
+func (a *API) health(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = io.WriteString(w, `{"status":"ok"}`+"\n")
+}
+
+func (a *API) openAIChat(w http.ResponseWriter, r *http.Request) {
+ a.serveInference(w, r, domain.ProtocolOpenAI)
+}
+
+func (a *API) anthropicMessages(w http.ResponseWriter, r *http.Request) {
+ a.serveInference(w, r, domain.ProtocolAnthropic)
+}
+
+func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol domain.Protocol) {
+ requestID := requestIDFrom(r.Context())
+ startedAt := time.Now().UTC()
+ a.metrics.RequestStarted()
+ success := false
+ defer func() { a.metrics.RequestFinished(success) }()
+
+ principal, authErr := a.authenticator.Authenticate(r)
+ if authErr != nil {
+ apierror.Write(w, apierror.Error{Status: http.StatusForbidden, Type: "access_denied", Message: "Invalid or missing API key"}, requestID)
+ return
+ }
+ if !hasScope(principal, "inference") {
+ apierror.Write(w, apierror.Error{Status: http.StatusForbidden, Type: "access_denied", Message: "API key does not have inference permission"}, requestID)
+ return
+ }
+
+ body, err := readBody(w, r, a.maxBodyBytes)
+ if err != nil {
+ status := http.StatusBadRequest
+ message := "Invalid request body"
+ if errors.As(err, new(*http.MaxBytesError)) {
+ status = http.StatusRequestEntityTooLarge
+ message = "Request body is too large"
+ }
+ apierror.Write(w, apierror.Error{Status: status, Type: "invalid_params", Message: message}, requestID)
+ return
+ }
+
+ var envelope struct {
+ Model string `json:"model"`
+ Stream bool `json:"stream"`
+ }
+ if err := json.Unmarshal(body, &envelope); err != nil || strings.TrimSpace(envelope.Model) == "" {
+ apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Parameter model is required and the body must be valid JSON"}, requestID)
+ return
+ }
+
+ routes, err := a.router.Plan(envelope.Model, protocol)
+ if err != nil {
+ if errors.Is(err, routing.ErrNoRoute) {
+ apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "model_not_supported", Message: "Model does not support this API protocol"}, requestID)
+ } else {
+ apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist"}, requestID)
+ }
+ return
+ }
+
+ result, err := a.forwarder.Forward(r.Context(), protocol, requestID, body, r.Header, routes)
+ if err != nil {
+ errorType := "no_provider_available"
+ status := http.StatusBadGateway
+ if errors.Is(err, context.Canceled) {
+ errorType = "client_disconnected"
+ status = 499
+ } else {
+ apierror.Write(w, apierror.Error{Status: status, Type: errorType, Message: "No upstream provider is currently available"}, requestID)
+ }
+ a.publishUsage(domain.UsageEvent{
+ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID,
+ PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: status,
+ Success: false, ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(),
+ })
+ return
+ }
+ defer result.Response.Body.Close()
+
+ if result.Response.StatusCode >= 400 {
+ gatewayError := normalizeProviderError(result.Response)
+ apierror.Write(w, gatewayError, requestID)
+ a.publishUsage(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,
+ Protocol: protocol, Stream: envelope.Stream, StatusCode: gatewayError.Status, Success: false,
+ ErrorType: gatewayError.Type, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(),
+ })
+ return
+ }
+
+ stream := envelope.Stream || strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "text/event-stream")
+ copyResponseHeaders(w.Header(), result.Response.Header, stream)
+ w.Header().Set("X-AIGW-Request-ID", requestID)
+ if stream {
+ w.Header().Set("X-Accel-Buffering", "no")
+ }
+ w.WriteHeader(result.Response.StatusCode)
+
+ observer := usage.NewObserver(protocol, stream)
+ copyErr := copyResponse(w, result.Response.Body, observer, stream)
+ usageResult := observer.Usage()
+ success = copyErr == nil
+ errorType := ""
+ if copyErr != nil {
+ errorType = "stream_interrupted"
+ }
+ a.publishUsage(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,
+ Protocol: protocol, Stream: stream, StatusCode: result.Response.StatusCode, Success: success,
+ ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), Usage: usageResult,
+ })
+ a.logger.Info("inference_request",
+ "request_id", requestID,
+ "tenant_id", principal.TenantID,
+ "project_id", principal.ProjectID,
+ "model", envelope.Model,
+ "provider", result.Route.Provider.ID,
+ "status", result.Response.StatusCode,
+ "attempts", result.Attempts,
+ "duration_ms", time.Since(startedAt).Milliseconds(),
+ )
+}
+
+func (a *API) openAIModels(w http.ResponseWriter, r *http.Request) {
+ if !a.authorize(w, r) {
+ return
+ }
+ models := a.catalog.Models(domain.ProtocolOpenAI)
+ data := make([]map[string]any, 0, len(models))
+ for _, model := range models {
+ data = append(data, map[string]any{"id": model.ID, "object": "model", "created": 0, "owned_by": model.OwnedBy})
+ }
+ writeJSON(w, map[string]any{"object": "list", "data": data})
+}
+
+func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) {
+ if !a.authorize(w, r) {
+ return
+ }
+ models := a.catalog.Models(domain.ProtocolAnthropic)
+ 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"})
+ }
+ response := map[string]any{"data": data, "has_more": false, "first_id": nil, "last_id": nil}
+ if len(models) > 0 {
+ response["first_id"] = models[0].ID
+ response["last_id"] = models[len(models)-1].ID
+ }
+ writeJSON(w, response)
+}
+
+func (a *API) authorize(w http.ResponseWriter, r *http.Request) bool {
+ principal, err := a.authenticator.Authenticate(r)
+ if err != nil || !hasScope(principal, "inference") {
+ apierror.Write(w, apierror.Error{Status: http.StatusForbidden, Type: "access_denied", Message: "Invalid API key or insufficient permission"}, requestIDFrom(r.Context()))
+ return false
+ }
+ return true
+}
+
+func (a *API) publishUsage(event domain.UsageEvent) {
+ if a.usageSink != nil {
+ a.usageSink.Publish(event)
+ }
+}
+
+func (a *API) withRequestID(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ id := newRequestID()
+ w.Header().Set("X-AIGW-Request-ID", id)
+ next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), requestIDKey{}, id)))
+ })
+}
+
+func (a *API) recoverPanics(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ defer func() {
+ if recovered := recover(); recovered != nil {
+ a.logger.Error("request_panic", "request_id", requestIDFrom(r.Context()), "error", recovered)
+ apierror.Write(w, apierror.Error{Status: http.StatusInternalServerError, Type: "internal_server_error", Message: "Internal server error"}, requestIDFrom(r.Context()))
+ }
+ }()
+ next.ServeHTTP(w, r)
+ })
+}
+
+func newRequestID() string {
+ var value [16]byte
+ if _, err := rand.Read(value[:]); err != nil {
+ return fmt.Sprintf("req_%d", time.Now().UnixNano())
+ }
+ return "req_" + hex.EncodeToString(value[:])
+}
+
+func requestIDFrom(ctx context.Context) string {
+ id, _ := ctx.Value(requestIDKey{}).(string)
+ return id
+}
+
+func hasScope(principal domain.Principal, wanted string) bool {
+ if len(principal.Scopes) == 0 {
+ return true
+ }
+ for _, scope := range principal.Scopes {
+ if scope == wanted || scope == "*" {
+ return true
+ }
+ }
+ return false
+}
+
+func readBody(w http.ResponseWriter, r *http.Request, limit int64) ([]byte, error) {
+ r.Body = http.MaxBytesReader(w, r.Body, limit)
+ defer r.Body.Close()
+ return io.ReadAll(r.Body)
+}
+
+func copyResponseHeaders(target, source http.Header, stream bool) {
+ for _, name := range []string{"Content-Type", "Cache-Control", "Retry-After"} {
+ if value := source.Get(name); value != "" {
+ target.Set(name, value)
+ }
+ }
+ if !stream {
+ if value := source.Get("Content-Length"); value != "" {
+ target.Set("Content-Length", value)
+ }
+ }
+}
+
+func copyResponse(w http.ResponseWriter, body io.Reader, observer io.Writer, stream bool) error {
+ var destination io.Writer = w
+ if stream {
+ destination = &flushingWriter{writer: w, controller: http.NewResponseController(w)}
+ }
+ _, err := io.CopyBuffer(io.MultiWriter(destination, observer), body, make([]byte, 32<<10))
+ return err
+}
+
+type flushingWriter struct {
+ writer io.Writer
+ controller *http.ResponseController
+}
+
+func (w *flushingWriter) Write(p []byte) (int, error) {
+ n, err := w.writer.Write(p)
+ if err == nil {
+ _ = w.controller.Flush()
+ }
+ return n, err
+}
+
+func normalizeProviderError(response *http.Response) apierror.Error {
+ payload, _ := io.ReadAll(io.LimitReader(response.Body, 64<<10))
+ message := "Upstream provider rejected the request"
+ var common struct {
+ Error struct {
+ Message string `json:"message"`
+ } `json:"error"`
+ }
+ if json.Unmarshal(payload, &common) == nil && common.Error.Message != "" {
+ message = common.Error.Message
+ }
+ switch response.StatusCode {
+ case http.StatusBadRequest:
+ return apierror.Error{Status: http.StatusUnprocessableEntity, Type: "provider_unprocessable_entity_error", Message: message}
+ case http.StatusRequestEntityTooLarge:
+ return apierror.Error{Status: http.StatusRequestEntityTooLarge, Type: "invalid_params", Message: message}
+ case http.StatusTooManyRequests:
+ return apierror.Error{Status: http.StatusTooManyRequests, Type: "rate_limit", Message: message}
+ default:
+ return apierror.Error{Status: http.StatusBadGateway, Type: "provider_error", Message: message}
+ }
+}
+
+func writeJSON(w http.ResponseWriter, value any) {
+ w.Header().Set("Content-Type", "application/json")
+ _ = json.NewEncoder(w).Encode(value)
+}
diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go
new file mode 100644
index 0000000..66a344b
--- /dev/null
+++ b/internal/httpapi/api_test.go
@@ -0,0 +1,224 @@
+package httpapi
+
+import (
+ "bufio"
+ "bytes"
+ "encoding/json"
+ "io"
+ "log/slog"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "sync/atomic"
+ "testing"
+ "time"
+
+ "aigw/internal/auth"
+ "aigw/internal/catalog"
+ "aigw/internal/config"
+ "aigw/internal/domain"
+ "aigw/internal/provider"
+ "aigw/internal/routing"
+ "aigw/internal/telemetry"
+)
+
+type captureUsageSink struct {
+ events chan domain.UsageEvent
+}
+
+func (s *captureUsageSink) Publish(event domain.UsageEvent) {
+ s.events <- event
+}
+
+func TestOpenAIProxyRewritesModelAndEmitsUsage(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path != "/v1/chat/completions" {
+ t.Errorf("unexpected path: %s", r.URL.Path)
+ }
+ if r.Header.Get("Authorization") != "Bearer upstream-secret" {
+ t.Errorf("upstream authorization leaked or missing: %q", r.Header.Get("Authorization"))
+ }
+ var request map[string]any
+ if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
+ t.Error(err)
+ }
+ if request["model"] != "upstream-model" {
+ t.Errorf("model was not rewritten: %+v", request)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = io.WriteString(w, `{"id":"chat-1","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`)
+ }))
+ defer upstream.Close()
+
+ gateway, sink := newTestGateway(t, []config.ProviderConfig{{
+ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret",
+ }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "upstream-model", Weight: 1}})
+ defer gateway.Close()
+
+ request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model","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("unexpected status %d: %s", response.StatusCode, payload)
+ }
+ if response.Header.Get("X-AIGW-Request-ID") == "" {
+ t.Fatal("missing gateway request id")
+ }
+
+ select {
+ case event := <-sink.events:
+ if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || !event.Success {
+ t.Fatalf("unexpected usage event: %+v", event)
+ }
+ case <-time.After(time.Second):
+ t.Fatal("usage event was not emitted")
+ }
+}
+
+func TestProxyFailsOverBeforeWritingResponse(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()
+
+ response := postOpenAI(t, gateway.URL, false)
+ defer response.Body.Close()
+ if response.StatusCode != http.StatusOK || primaryCalls.Load() != 1 || fallbackCalls.Load() != 1 {
+ t.Fatalf("failover did not complete: status=%d primary=%d fallback=%d", response.StatusCode, primaryCalls.Load(), fallbackCalls.Load())
+ }
+ event := <-sink.events
+ if event.Attempts != 2 || event.ProviderID != "fallback" {
+ t.Fatalf("unexpected failover event: %+v", event)
+ }
+}
+
+func TestSSEIsFlushedBeforeUpstreamCompletes(t *testing.T) {
+ release := make(chan struct{})
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "text/event-stream")
+ _, _ = io.WriteString(w, "data: {\"id\":\"chunk-1\",\"choices\":[]}\n\n")
+ w.(http.Flusher).Flush()
+ <-release
+ _, _ = io.WriteString(w, "data: [DONE]\n\n")
+ }))
+ defer upstream.Close()
+ gateway, _ := newTestGateway(t, []config.ProviderConfig{{
+ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret",
+ }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}})
+ defer gateway.Close()
+
+ response := postOpenAI(t, gateway.URL, true)
+ defer response.Body.Close()
+ reader := bufio.NewReader(response.Body)
+ line, err := reader.ReadString('\n')
+ if err != nil {
+ close(release)
+ t.Fatal(err)
+ }
+ if !strings.Contains(line, "chunk-1") {
+ close(release)
+ t.Fatalf("unexpected first SSE line: %q", line)
+ }
+ close(release)
+}
+
+func TestAnthropicHeadersAndPath(t *testing.T) {
+ upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Path != "/v1/messages" || r.Header.Get("x-api-key") != "anthropic-upstream" || r.Header.Get("anthropic-version") != "2023-06-01" {
+ t.Errorf("unexpected anthropic request: path=%s key=%q version=%q", r.URL.Path, r.Header.Get("x-api-key"), r.Header.Get("anthropic-version"))
+ }
+ _, _ = io.WriteString(w, `{"type":"message","content":[],"usage":{"input_tokens":4,"output_tokens":6}}`)
+ }))
+ defer upstream.Close()
+ gateway, sink := newTestGateway(t, []config.ProviderConfig{{
+ ID: "anthropic", Protocol: domain.ProtocolAnthropic, BaseURL: upstream.URL + "/v1", APIKey: "anthropic-upstream",
+ }}, []config.RouteConfig{{Provider: "anthropic", UpstreamModel: "claude-upstream", Weight: 1}})
+ defer gateway.Close()
+
+ request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/anthropic/v1/messages", strings.NewReader(`{"model":"public/model","max_tokens":10,"messages":[{"role":"user","content":"hello"}]}`))
+ request.Header.Set("x-api-key", "client-secret")
+ request.Header.Set("anthropic-version", "2023-06-01")
+ response, err := http.DefaultClient.Do(request)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer response.Body.Close()
+ if response.StatusCode != http.StatusOK {
+ t.Fatalf("unexpected status: %d", response.StatusCode)
+ }
+ event := <-sink.events
+ if event.Usage.InputTokens != 4 || event.Usage.OutputTokens != 6 {
+ t.Fatalf("unexpected anthropic usage: %+v", event.Usage)
+ }
+}
+
+func newTestGateway(t *testing.T, providers []config.ProviderConfig, routes []config.RouteConfig) (*httptest.Server, *captureUsageSink) {
+ t.Helper()
+ authenticator, err := auth.NewStatic(`[{"key":"client-secret","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`, false)
+ if err != nil {
+ t.Fatal(err)
+ }
+ cfg := config.Config{
+ Providers: providers,
+ Models: []config.ModelConfig{{ID: "public/model", OwnedBy: "test", Routes: routes}},
+ UpstreamHTTP: config.UpstreamHTTPConfig{
+ MaxIdleConnections: 100, MaxIdleConnectionsPerHost: 20,
+ IdleConnectionTimeoutSecs: 10, ResponseHeaderTimeoutSecs: 2,
+ },
+ }
+ metrics := &telemetry.Metrics{}
+ modelCatalog := catalog.New(cfg)
+ 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),
+ UsageSink: sink,
+ Metrics: metrics,
+ Logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
+ MaxBodyBytes: 1 << 20,
+ })
+ return httptest.NewServer(api.Handler()), sink
+}
+
+func postOpenAI(t *testing.T, gatewayURL string, stream bool) *http.Response {
+ t.Helper()
+ payload := []byte(`{"model":"public/model","messages":[{"role":"user","content":"hello"}],"stream":false}`)
+ if stream {
+ payload = []byte(`{"model":"public/model","messages":[{"role":"user","content":"hello"}],"stream":true}`)
+ }
+ request, _ := http.NewRequest(http.MethodPost, gatewayURL+"/v1/chat/completions", bytes.NewReader(payload))
+ request.Header.Set("Authorization", "Bearer client-secret")
+ response, err := http.DefaultClient.Do(request)
+ if err != nil {
+ t.Fatal(err)
+ }
+ return response
+}
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go
new file mode 100644
index 0000000..a9d5734
--- /dev/null
+++ b/internal/provider/forwarder.go
@@ -0,0 +1,134 @@
+package provider
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net"
+ "net/http"
+ "strings"
+ "time"
+
+ "aigw/internal/config"
+ "aigw/internal/domain"
+ "aigw/internal/telemetry"
+)
+
+type Result struct {
+ Response *http.Response
+ Route domain.Route
+ Attempts int
+}
+
+type Forwarder struct {
+ client *http.Client
+ metrics *telemetry.Metrics
+}
+
+func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder {
+ transport := &http.Transport{
+ Proxy: http.ProxyFromEnvironment,
+ DialContext: (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext,
+ ForceAttemptHTTP2: true,
+ MaxIdleConns: cfg.MaxIdleConnections,
+ MaxIdleConnsPerHost: cfg.MaxIdleConnectionsPerHost,
+ IdleConnTimeout: time.Duration(cfg.IdleConnectionTimeoutSecs) * time.Second,
+ TLSHandshakeTimeout: 10 * time.Second,
+ ResponseHeaderTimeout: time.Duration(cfg.ResponseHeaderTimeoutSecs) * time.Second,
+ ExpectContinueTimeout: time.Second,
+ }
+ return &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics}
+}
+
+func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID 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 := rewriteModel(originalBody, route.UpstreamModel)
+ if err != nil {
+ return Result{Attempts: i}, err
+ }
+ request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider.BaseURL, 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()
+ response, err := f.client.Do(request)
+ if err != nil {
+ lastErr = err
+ continue
+ }
+ attempts := i + 1
+ if retryableStatus(response.StatusCode) && attempts < len(routes) {
+ _, _ = io.CopyN(io.Discard, response.Body, 8<<10)
+ _ = response.Body.Close()
+ lastErr = fmt.Errorf("upstream %s returned %d", route.Provider.ID, response.StatusCode)
+ continue
+ }
+ return Result{Response: response, Route: route, Attempts: attempts}, nil
+ }
+ if lastErr == nil {
+ lastErr = errors.New("all upstream routes failed")
+ }
+ return Result{Attempts: len(routes)}, lastErr
+}
+
+func rewriteModel(body []byte, upstreamModel 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
+ result, err := json.Marshal(object)
+ if err != nil {
+ return nil, fmt.Errorf("encode upstream request: %w", err)
+ }
+ return result, nil
+}
+
+func endpointURL(baseURL string, protocol domain.Protocol) string {
+ baseURL = strings.TrimRight(baseURL, "/")
+ if protocol == domain.ProtocolAnthropic {
+ return baseURL + "/messages"
+ }
+ return baseURL + "/chat/completions"
+}
+
+func setHeaders(target, source http.Header, provider domain.Provider, protocol domain.Protocol, requestID string) {
+ target.Set("Content-Type", "application/json")
+ target.Set("Accept", source.Get("Accept"))
+ if target.Get("Accept") == "" {
+ target.Set("Accept", "application/json")
+ }
+ target.Set("User-Agent", "aigw/0.1")
+ target.Set("X-Request-ID", requestID)
+ if protocol == domain.ProtocolAnthropic {
+ target.Set("x-api-key", provider.APIKey)
+ version := source.Get("anthropic-version")
+ if version == "" {
+ version = "2023-06-01"
+ }
+ target.Set("anthropic-version", version)
+ if beta := source.Get("anthropic-beta"); beta != "" {
+ target.Set("anthropic-beta", beta)
+ }
+ return
+ }
+ target.Set("Authorization", "Bearer "+provider.APIKey)
+}
+
+func retryableStatus(status int) bool {
+ switch status {
+ case http.StatusTooManyRequests, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout:
+ return true
+ default:
+ return false
+ }
+}
diff --git a/internal/routing/router.go b/internal/routing/router.go
new file mode 100644
index 0000000..53e5261
--- /dev/null
+++ b/internal/routing/router.go
@@ -0,0 +1,81 @@
+package routing
+
+import (
+ "errors"
+ "sort"
+ "strconv"
+ "sync"
+ "sync/atomic"
+
+ "aigw/internal/catalog"
+ "aigw/internal/domain"
+)
+
+var ErrNoRoute = errors.New("no compatible upstream route")
+
+type Router struct {
+ catalog *catalog.Catalog
+ counters sync.Map
+}
+
+func New(catalog *catalog.Catalog) *Router {
+ return &Router{catalog: catalog}
+}
+
+func (r *Router) Plan(modelID string, protocol domain.Protocol) ([]domain.Route, error) {
+ model, err := r.catalog.Model(modelID)
+ if err != nil {
+ return nil, err
+ }
+ routes := make([]domain.Route, 0, len(model.Routes))
+ for _, route := range model.Routes {
+ if route.Provider.Protocol == protocol {
+ routes = append(routes, route)
+ }
+ }
+ if len(routes) == 0 {
+ return nil, ErrNoRoute
+ }
+
+ sort.SliceStable(routes, func(i, j int) bool { return routes[i].Priority < routes[j].Priority })
+ result := make([]domain.Route, 0, len(routes))
+ for start := 0; start < len(routes); {
+ end := start + 1
+ for end < len(routes) && routes[end].Priority == routes[start].Priority {
+ end++
+ }
+ result = append(result, r.rotate(modelID, protocol, routes[start:end])...)
+ start = end
+ }
+ return result, nil
+}
+
+func (r *Router) rotate(modelID string, protocol domain.Protocol, routes []domain.Route) []domain.Route {
+ if len(routes) < 2 {
+ return append([]domain.Route(nil), routes...)
+ }
+ key := modelID + "\x00" + string(protocol) + "\x00" + strconv.Itoa(routes[0].Priority)
+ counterValue, _ := r.counters.LoadOrStore(key, &atomic.Uint64{})
+ counter := counterValue.(*atomic.Uint64).Add(1) - 1
+
+ totalWeight := 0
+ for _, route := range routes {
+ totalWeight += route.Weight
+ }
+ position := int(counter % uint64(totalWeight))
+ selected := 0
+ for i, route := range routes {
+ if position < route.Weight {
+ selected = i
+ break
+ }
+ position -= route.Weight
+ }
+
+ result := make([]domain.Route, 0, len(routes))
+ result = append(result, routes[selected])
+ for offset := 1; offset < len(routes); offset++ {
+ result = append(result, routes[(selected+offset)%len(routes)])
+ }
+ return result
+}
diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go
new file mode 100644
index 0000000..62dc656
--- /dev/null
+++ b/internal/routing/router_test.go
@@ -0,0 +1,60 @@
+package routing
+
+import (
+ "testing"
+
+ "aigw/internal/catalog"
+ "aigw/internal/config"
+ "aigw/internal/domain"
+)
+
+func TestPlanHonorsPriorityAndProtocol(t *testing.T) {
+ cfg := config.Config{
+ Providers: []config.ProviderConfig{
+ {ID: "openai-primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://one.test", APIKey: "one"},
+ {ID: "openai-fallback", Protocol: domain.ProtocolOpenAI, BaseURL: "https://two.test", APIKey: "two"},
+ {ID: "anthropic", Protocol: domain.ProtocolAnthropic, BaseURL: "https://three.test", APIKey: "three"},
+ },
+ Models: []config.ModelConfig{{
+ ID: "public/model",
+ Routes: []config.RouteConfig{
+ {Provider: "openai-fallback", UpstreamModel: "fallback", Priority: 10, Weight: 1},
+ {Provider: "anthropic", UpstreamModel: "claude", Priority: 0, Weight: 1},
+ {Provider: "openai-primary", UpstreamModel: "primary", Priority: 0, Weight: 1},
+ },
+ }},
+ }
+ router := New(catalog.New(cfg))
+ plan, err := router.Plan("public/model", domain.ProtocolOpenAI)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(plan) != 2 || plan[0].Provider.ID != "openai-primary" || plan[1].Provider.ID != "openai-fallback" {
+ t.Fatalf("unexpected plan: %+v", plan)
+ }
+}
+
+func TestPlanUsesWeightsForPrimarySelection(t *testing.T) {
+ cfg := config.Config{
+ Providers: []config.ProviderConfig{
+ {ID: "one", Protocol: domain.ProtocolOpenAI, BaseURL: "https://one.test", APIKey: "one"},
+ {ID: "two", Protocol: domain.ProtocolOpenAI, BaseURL: "https://two.test", APIKey: "two"},
+ },
+ Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{
+ {Provider: "one", UpstreamModel: "one", Weight: 3},
+ {Provider: "two", UpstreamModel: "two", Weight: 1},
+ }}},
+ }
+ router := New(catalog.New(cfg))
+ counts := map[string]int{}
+ for range 8 {
+ plan, err := router.Plan("public/model", domain.ProtocolOpenAI)
+ if err != nil {
+ t.Fatal(err)
+ }
+ counts[plan[0].Provider.ID]++
+ }
+ if counts["one"] != 6 || counts["two"] != 2 {
+ t.Fatalf("unexpected weighted distribution: %+v", counts)
+ }
+}
diff --git a/internal/security/credentials.go b/internal/security/credentials.go
new file mode 100644
index 0000000..b55fb6b
--- /dev/null
+++ b/internal/security/credentials.go
@@ -0,0 +1,57 @@
+package security
+
+import (
+ "crypto/aes"
+ "crypto/cipher"
+ "crypto/rand"
+ "encoding/base64"
+ "errors"
+ "fmt"
+ "io"
+)
+
+type CredentialCipher struct {
+ aead cipher.AEAD
+}
+
+func NewCredentialCipher(encodedKey string) (*CredentialCipher, error) {
+ key, err := base64.StdEncoding.DecodeString(encodedKey)
+ if err != nil {
+ return nil, fmt.Errorf("decode credential key: %w", err)
+ }
+ if len(key) != 32 {
+ return nil, errors.New("credential key must be a base64-encoded 32-byte key")
+ }
+ block, err := aes.NewCipher(key)
+ if err != nil {
+ return nil, fmt.Errorf("create credential cipher: %w", err)
+ }
+ aead, err := cipher.NewGCM(block)
+ if err != nil {
+ return nil, fmt.Errorf("create credential AEAD: %w", err)
+ }
+ return &CredentialCipher{aead: aead}, nil
+}
+
+func (c *CredentialCipher) Encrypt(plaintext string) ([]byte, error) {
+ if plaintext == "" {
+ return nil, errors.New("credential cannot be empty")
+ }
+ nonce := make([]byte, c.aead.NonceSize())
+ if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
+ return nil, fmt.Errorf("generate credential nonce: %w", err)
+ }
+ return c.aead.Seal(nonce, nonce, []byte(plaintext), nil), nil
+}
+
+func (c *CredentialCipher) Decrypt(ciphertext []byte) (string, error) {
+ if len(ciphertext) < c.aead.NonceSize() {
+ return "", errors.New("credential ciphertext is truncated")
+ }
+ nonce := ciphertext[:c.aead.NonceSize()]
+ plaintext, err := c.aead.Open(nil, nonce, ciphertext[c.aead.NonceSize():], nil)
+ if err != nil {
+ return "", errors.New("decrypt credential: authentication failed")
+ }
+ return string(plaintext), nil
+}
diff --git a/internal/security/credentials_test.go b/internal/security/credentials_test.go
new file mode 100644
index 0000000..07fa59e
--- /dev/null
+++ b/internal/security/credentials_test.go
@@ -0,0 +1,30 @@
+package security
+
+import (
+ "encoding/base64"
+ "strings"
+ "testing"
+)
+
+func TestCredentialCipherRoundTrip(t *testing.T) {
+ key := base64.StdEncoding.EncodeToString([]byte(strings.Repeat("k", 32)))
+ cipher, err := NewCredentialCipher(key)
+ if err != nil {
+ t.Fatal(err)
+ }
+ ciphertext, err := cipher.Encrypt("upstream-secret")
+ if err != nil {
+ t.Fatal(err)
+ }
+ plaintext, err := cipher.Decrypt(ciphertext)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if plaintext != "upstream-secret" {
+ t.Fatalf("unexpected plaintext: %q", plaintext)
+ }
+ ciphertext[len(ciphertext)-1] ^= 1
+ if _, err := cipher.Decrypt(ciphertext); err == nil {
+ t.Fatal("expected authentication failure for modified ciphertext")
+ }
+}
diff --git a/internal/telemetry/metrics.go b/internal/telemetry/metrics.go
new file mode 100644
index 0000000..4942d8d
--- /dev/null
+++ b/internal/telemetry/metrics.go
@@ -0,0 +1,44 @@
+package telemetry
+
+import (
+ "fmt"
+ "net/http"
+ "sync/atomic"
+)
+
+type Metrics struct {
+ requests atomic.Uint64
+ failed atomic.Uint64
+ inFlight atomic.Int64
+ attempts atomic.Uint64
+ droppedUsage atomic.Uint64
+}
+
+func (m *Metrics) RequestStarted() {
+ m.requests.Add(1)
+ m.inFlight.Add(1)
+}
+
+func (m *Metrics) RequestFinished(success bool) {
+ m.inFlight.Add(-1)
+ if !success {
+ m.failed.Add(1)
+ }
+}
+
+func (m *Metrics) UpstreamAttempt() {
+ m.attempts.Add(1)
+}
+
+func (m *Metrics) UsageDropped() {
+ m.droppedUsage.Add(1)
+}
+
+func (m *Metrics) ServeHTTP(w http.ResponseWriter, _ *http.Request) {
+ w.Header().Set("Content-Type", "text/plain; version=0.0.4")
+ fmt.Fprintf(w, "# TYPE aigw_requests_total counter\naigw_requests_total %d\n", m.requests.Load())
+ fmt.Fprintf(w, "# TYPE aigw_requests_failed_total counter\naigw_requests_failed_total %d\n", m.failed.Load())
+ fmt.Fprintf(w, "# TYPE aigw_requests_in_flight gauge\naigw_requests_in_flight %d\n", m.inFlight.Load())
+ fmt.Fprintf(w, "# TYPE aigw_upstream_attempts_total counter\naigw_upstream_attempts_total %d\n", m.attempts.Load())
+ fmt.Fprintf(w, "# TYPE aigw_usage_events_dropped_total counter\naigw_usage_events_dropped_total %d\n", m.droppedUsage.Load())
+}
diff --git a/internal/telemetry/usage_sink.go b/internal/telemetry/usage_sink.go
new file mode 100644
index 0000000..aefe6f5
--- /dev/null
+++ b/internal/telemetry/usage_sink.go
@@ -0,0 +1,57 @@
+package telemetry
+
+import (
+ "context"
+ "log/slog"
+ "sync"
+
+ "aigw/internal/domain"
+)
+
+type UsageSink interface {
+ Publish(domain.UsageEvent)
+}
+
+type AsyncUsageLogger struct {
+ logger *slog.Logger
+ metrics *Metrics
+ events chan domain.UsageEvent
+ done chan struct{}
+ once sync.Once
+}
+
+func NewAsyncUsageLogger(logger *slog.Logger, metrics *Metrics, buffer int) *AsyncUsageLogger {
+ sink := &AsyncUsageLogger{
+ logger: logger,
+ metrics: metrics,
+ events: make(chan domain.UsageEvent, buffer),
+ done: make(chan struct{}),
+ }
+ go sink.run()
+ return sink
+}
+
+func (s *AsyncUsageLogger) Publish(event domain.UsageEvent) {
+ select {
+ case s.events <- event:
+ default:
+ s.metrics.UsageDropped()
+ }
+}
+
+func (s *AsyncUsageLogger) Close(ctx context.Context) error {
+ s.once.Do(func() { close(s.events) })
+ select {
+ case <-s.done:
+ return nil
+ case <-ctx.Done():
+ return ctx.Err()
+ }
+}
+
+func (s *AsyncUsageLogger) run() {
+ defer close(s.done)
+ for event := range s.events {
+ s.logger.Info("usage_event", "event", event)
+ }
+}
diff --git a/internal/usage/observer.go b/internal/usage/observer.go
new file mode 100644
index 0000000..cc8b52f
--- /dev/null
+++ b/internal/usage/observer.go
@@ -0,0 +1,200 @@
+package usage
+
+import (
+ "bytes"
+ "encoding/json"
+ "strings"
+
+ "aigw/internal/domain"
+)
+
+const maxCaptureBytes = 64 << 10
+
+type Observer struct {
+ protocol domain.Protocol
+ stream bool
+ buffer []byte
+ line []byte
+ usage domain.Usage
+ found bool
+ explicitTotal bool
+}
+
+func NewObserver(protocol domain.Protocol, stream bool) *Observer {
+ return &Observer{protocol: protocol, stream: stream}
+}
+
+func (o *Observer) Write(p []byte) (int, error) {
+ if o.stream {
+ o.observeSSE(p)
+ } else {
+ o.captureTail(p)
+ }
+ return len(p), nil
+}
+
+func (o *Observer) Usage() domain.Usage {
+ if o.stream {
+ if len(o.line) > 0 {
+ o.parseSSELine(o.line)
+ }
+ return o.usage
+ }
+ o.parseJSON(o.buffer)
+ return o.usage
+}
+
+func (o *Observer) captureTail(p []byte) {
+ if len(p) >= maxCaptureBytes {
+ o.buffer = append(o.buffer[:0], p[len(p)-maxCaptureBytes:]...)
+ return
+ }
+ if len(o.buffer)+len(p) > maxCaptureBytes {
+ drop := len(o.buffer) + len(p) - maxCaptureBytes
+ copy(o.buffer, o.buffer[drop:])
+ o.buffer = o.buffer[:len(o.buffer)-drop]
+ }
+ o.buffer = append(o.buffer, p...)
+}
+
+func (o *Observer) observeSSE(p []byte) {
+ o.line = append(o.line, p...)
+ for {
+ index := bytes.IndexByte(o.line, '\n')
+ if index < 0 {
+ if len(o.line) > maxCaptureBytes {
+ o.line = append(o.line[:0], o.line[len(o.line)-maxCaptureBytes:]...)
+ }
+ return
+ }
+ line := bytes.TrimSpace(o.line[:index])
+ o.parseSSELine(line)
+ o.line = o.line[index+1:]
+ }
+}
+
+func (o *Observer) parseSSELine(line []byte) {
+ if !bytes.HasPrefix(line, []byte("data:")) || !bytes.Contains(line, []byte("\"usage\"")) {
+ return
+ }
+ payload := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:")))
+ if bytes.Equal(payload, []byte("[DONE]")) {
+ return
+ }
+ o.parseJSON(payload)
+}
+
+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"`
+}
+
+type responseEnvelope struct {
+ Usage *usageFields `json:"usage"`
+ Message *struct {
+ Usage *usageFields `json:"usage"`
+ } `json:"message"`
+}
+
+func (o *Observer) parseJSON(payload []byte) {
+ var envelope responseEnvelope
+ if err := json.Unmarshal(payload, &envelope); err != nil {
+ payload = extractUsageObject(payload)
+ if len(payload) == 0 {
+ return
+ }
+ var fields usageFields
+ if json.Unmarshal(payload, &fields) == nil {
+ o.apply(&fields)
+ }
+ return
+ }
+ if envelope.Usage != nil {
+ o.apply(envelope.Usage)
+ }
+ if envelope.Message != nil && envelope.Message.Usage != nil {
+ o.apply(envelope.Message.Usage)
+ }
+}
+
+func (o *Observer) apply(fields *usageFields) {
+ if fields.PromptTokens != nil {
+ o.usage.InputTokens = *fields.PromptTokens
+ o.found = true
+ }
+ if fields.InputTokens != nil {
+ o.usage.InputTokens = *fields.InputTokens
+ o.found = true
+ }
+ if fields.CompletionTokens != nil {
+ o.usage.OutputTokens = *fields.CompletionTokens
+ o.found = true
+ }
+ if fields.OutputTokens != nil {
+ o.usage.OutputTokens = *fields.OutputTokens
+ o.found = true
+ }
+ if fields.TotalTokens != nil {
+ o.usage.TotalTokens = *fields.TotalTokens
+ o.found = true
+ o.explicitTotal = true
+ }
+ if fields.CacheCreationInputTokens != nil {
+ o.usage.CacheCreationInputTokens = *fields.CacheCreationInputTokens
+ o.found = true
+ }
+ if fields.CacheReadInputTokens != nil {
+ o.usage.CacheReadInputTokens = *fields.CacheReadInputTokens
+ o.found = true
+ }
+ if !o.explicitTotal && o.found {
+ o.usage.TotalTokens = o.usage.InputTokens + o.usage.OutputTokens
+ }
+}
+
+func extractUsageObject(payload []byte) []byte {
+ index := strings.LastIndex(string(payload), `"usage"`)
+ if index < 0 {
+ return nil
+ }
+ rest := payload[index+len(`"usage"`):]
+ start := bytes.IndexByte(rest, '{')
+ if start < 0 {
+ return nil
+ }
+ rest = rest[start:]
+ depth := 0
+ inString := false
+ escaped := false
+ for i, b := range rest {
+ if inString {
+ if escaped {
+ escaped = false
+ continue
+ }
+ if b == '\\' {
+ escaped = true
+ } else if b == '"' {
+ inString = false
+ }
+ continue
+ }
+ switch b {
+ case '"':
+ inString = true
+ case '{':
+ depth++
+ case '}':
+ depth--
+ if depth == 0 {
+ return rest[:i+1]
+ }
+ }
+ }
+ return nil
+}
diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go
new file mode 100644
index 0000000..8fcf408
--- /dev/null
+++ b/internal/usage/observer_test.go
@@ -0,0 +1,26 @@
+package usage
+
+import (
+ "testing"
+
+ "aigw/internal/domain"
+)
+
+func TestObserverReadsOpenAIJSONUsage(t *testing.T) {
+ observer := NewObserver(domain.ProtocolOpenAI, false)
+ _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`))
+ got := observer.Usage()
+ if got.InputTokens != 11 || 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_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":8}}\n\n"))
+ got := observer.Usage()
+ if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 {
+ t.Fatalf("unexpected usage: %+v", got)
+ }
+}