From 5b651488b081b65fda8a323f228e139adb79a35d Mon Sep 17 00:00:00 2001 From: Chia Date: Tue, 4 Aug 2026 19:58:52 +1200 Subject: Build AI gateway control plane and admin UI --- .env.control.example | 6 + .env.example | 5 + .gitignore | 5 + Dockerfile | 17 ++ Makefile | 19 ++ README.md | 125 ++++++++++++ cmd/migrate/main.go | 28 +++ cmd/mockupstream/main.go | 65 ++++++ config.control.example.json | 40 ++++ config.example.json | 64 ++++++ config.local.json | 50 +++++ docker-compose.yml | 47 +++++ docs/architecture.md | 68 +++++++ go.mod | 18 ++ go.sum | 25 +++ internal/adminapi/api.go | 352 ++++++++++++++++++++++++++++++++ internal/adminui/assets/app.js | 70 +++++++ internal/adminui/assets/index.html | 66 ++++++ internal/adminui/assets/style.css | 13 ++ internal/adminui/ui.go | 18 ++ internal/apierror/apierror.go | 26 +++ internal/auth/static.go | 124 ++++++++++++ internal/auth/static_test.go | 76 +++++++ internal/catalog/catalog.go | 103 ++++++++++ internal/catalog/catalog_test.go | 43 ++++ internal/config/config.go | 311 ++++++++++++++++++++++++++++ internal/config/config_test.go | 95 +++++++++ internal/controlplane/manager.go | 213 ++++++++++++++++++++ internal/controlplane/manager_test.go | 212 +++++++++++++++++++ internal/controlplane/mutations.go | 281 ++++++++++++++++++++++++++ internal/controlplane/queries.go | 146 ++++++++++++++ internal/controlplane/schema.sql | 82 ++++++++ internal/controlplane/snapshot.go | 150 ++++++++++++++ internal/controlplane/store.go | 170 ++++++++++++++++ internal/controlplane/types.go | 138 +++++++++++++ internal/domain/types.go | 64 ++++++ internal/httpapi/api.go | 369 ++++++++++++++++++++++++++++++++++ internal/httpapi/api_test.go | 224 +++++++++++++++++++++ internal/provider/forwarder.go | 134 ++++++++++++ internal/routing/router.go | 81 ++++++++ internal/routing/router_test.go | 60 ++++++ internal/security/credentials.go | 57 ++++++ internal/security/credentials_test.go | 30 +++ internal/telemetry/metrics.go | 44 ++++ internal/telemetry/usage_sink.go | 57 ++++++ internal/usage/observer.go | 200 ++++++++++++++++++ internal/usage/observer_test.go | 26 +++ 47 files changed, 4617 insertions(+) create mode 100644 .env.control.example create mode 100644 .env.example create mode 100644 .gitignore create mode 100644 Dockerfile create mode 100644 Makefile create mode 100644 README.md create mode 100644 cmd/migrate/main.go create mode 100644 cmd/mockupstream/main.go create mode 100644 config.control.example.json create mode 100644 config.example.json create mode 100644 config.local.json create mode 100644 docker-compose.yml create mode 100644 docs/architecture.md create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal/adminapi/api.go create mode 100644 internal/adminui/assets/app.js create mode 100644 internal/adminui/assets/index.html create mode 100644 internal/adminui/assets/style.css create mode 100644 internal/adminui/ui.go create mode 100644 internal/apierror/apierror.go create mode 100644 internal/auth/static.go create mode 100644 internal/auth/static_test.go create mode 100644 internal/catalog/catalog.go create mode 100644 internal/catalog/catalog_test.go create mode 100644 internal/config/config.go create mode 100644 internal/config/config_test.go create mode 100644 internal/controlplane/manager.go create mode 100644 internal/controlplane/manager_test.go create mode 100644 internal/controlplane/mutations.go create mode 100644 internal/controlplane/queries.go create mode 100644 internal/controlplane/schema.sql create mode 100644 internal/controlplane/snapshot.go create mode 100644 internal/controlplane/store.go create mode 100644 internal/controlplane/types.go create mode 100644 internal/domain/types.go create mode 100644 internal/httpapi/api.go create mode 100644 internal/httpapi/api_test.go create mode 100644 internal/provider/forwarder.go create mode 100644 internal/routing/router.go create mode 100644 internal/routing/router_test.go create mode 100644 internal/security/credentials.go create mode 100644 internal/security/credentials_test.go create mode 100644 internal/telemetry/metrics.go create mode 100644 internal/telemetry/usage_sink.go create mode 100644 internal/usage/observer.go create mode 100644 internal/usage/observer_test.go 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) => ({ '&':'&', '<':'<', '>':'>', "'":''', '"':'"' }[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 `${items.map(item => ``).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]) => `
${label}${esc(value)}${sub}
`).join(''); +} +function renderTenants() { $('#tenants-body').innerHTML = state.tenants.map(item => `${esc(item.name)}${esc(item.slug)}${esc(item.status)}${date(item.created_at)}`).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 => `${esc(item.name)}${esc(item.tenant_id).slice(0, 8)}…${esc(item.slug)}${esc(item.status)}`).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 => `${esc(item.name)}${esc(item.key_prefix)}${esc(item.project_id).slice(0, 8)}…${(item.scopes || []).map(scope => `${esc(scope)}`).join('')}${esc(item.status)}${item.status === 'active' ? `` : ''}`).join('') || emptyRow(6); } +function renderProviders() { $('#providers-body').innerHTML = state.providers.map(item => `${esc(item.name)}${esc(item.protocol)}${esc(item.base_url)}${esc(item.route_count)}${item.enabled ? 'enabled' : 'disabled'}`).join('') || emptyRow(6); } +function renderModels() { $('#models-body').innerHTML = state.models.map(item => `${esc(item.public_id)}${esc(item.owned_by || '—')}
${(item.routes || []).map(route => `${esc(route.provider_name || route.provider_id).slice(0, 24)} → ${esc(route.upstream_model)} p${route.priority} / w${route.weight}`).join('')}
${item.enabled ? 'enabled' : 'disabled'}`).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 = ``; $('#route-editor').appendChild(wrapper); renderRouteEditor(); } +function emptyRow(span) { return `No records yet`; } + +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 @@ + + + + + + AIGW Control Plane + + + + +
+
A
AIGWCONTROL PLANE
+
Offline
+
+
+ + +
+
OPERATIONS

Control plane overview

+
+
i
Runtime snapshot

Writes commit to PostgreSQL and apply locally first. Redis accelerates propagation when available; PostgreSQL polling keeps every gateway convergent.

+
+ +
+
IDENTITY

Tenants

+
+
NameSlugStatusCreated
+
+ +
+
IDENTITY

Projects

+
+
NameTenantSlugStatus
+
+ +
+
ACCESS

API keys

+
+
Key visibilityThe secret is shown only once after creation.
+
NamePrefixProjectScopesStatus
+
+ +
+
UPSTREAMS

Providers

+
+
NameProtocolBase URLRoutesStatus
+
+ +
+
ROUTING

Models & routes

+
+
Public IDOwnerRoutesStatus
+
+
+
+
ONE-TIME SECRET

API key created

Copy this key now. It will not be shown again.

+ + + 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) + } +} -- cgit v1.2.3