fix: restore reviewable migration evidence

This commit is contained in:
MiMoCode
2026-07-10 18:26:48 +08:00
commit d1bbb5370c
42 changed files with 6419 additions and 0 deletions
+8
View File
@@ -0,0 +1,8 @@
tests/
.env
.env.*
*.log
*.exe
models.json
tmp_*.go
dbcheck.exe
+57
View File
@@ -0,0 +1,57 @@
# 监听地址
LISTEN_ADDR=:8080
# newapi 后端;本地运行和 Compose 均必须显式设置
UPSTREAM_URL=
# 是否跳过上游 HTTPS 证书校验;仅开发/可信内网自签证书场景使用
UPSTREAM_TLS_INSECURE_SKIP_VERIFY=false
# ClickHouse 原生协议地址;必须显式设置,支持 host:port、clickhouse:// 或 clickhouses://
CLICKHOUSE_URL=
# 单条请求/响应 body 最大记录字节数
MAX_BODY_BYTES=1048576
# 异步日志队列容量
LOG_QUEUE_SIZE=256
# 队列中所有待写日志的总字节预算
LOG_QUEUE_BYTES=67108864
# 批量写入条数
LOG_BATCH_SIZE=50
# 批量写入间隔(Go duration: 500ms, 2s, 1m...
LOG_BATCH_INTERVAL=2s
# 后台日志 worker 数
LOG_WORKERS=2
# 黑白名单文件
FILTER_FILE=./filter.yaml
# DB 故障后重连尝试间隔
DB_RECONNECT_INTERVAL=10s
# HTTP 超时(Go duration: 30s, 10m...
READ_TIMEOUT=30s
WRITE_TIMEOUT=10m
IDLE_TIMEOUT=5m
UPSTREAM_TIMEOUT=30s
UPSTREAM_RESPONSE_TIMEOUT=30s
UPSTREAM_STREAM_IDLE_TIMEOUT=2m
# 可信反向代理 IP/CIDR,逗号分隔;为空时忽略所有客户端转发头
TRUSTED_PROXIES=
# ===== Docker Compose 可选配置 =====
# 对外暴露的代理端口(容器内固定 :8080)
LISTEN_PORT=8080
# Compose 内置 ClickHouse 配置
CLICKHOUSE_USER=tokenthief
CLICKHOUSE_PASSWORD=
CLICKHOUSE_DB=tokenthief
# 容器时区
TZ=Asia/Shanghai
+54
View File
@@ -0,0 +1,54 @@
# Environment / secrets
.env
.env.*
!.env.example
# Local binaries
TokenThief
TokenThief.exe
tokenthief
tokenthief.exe
*.exe
*.dll
*.so
*.dylib
# Go build/test outputs
bin/
dist/
build/
*.test
*.out
coverage.out
coverage.html
*.coverprofile
# Runtime logs and temporary outputs
*.log
proxy.log
proxy.err.log
models.json
tmp_*.go
dbcheck.exe
# Local database / volumes
clickhouse_data/
data/
*.db
*.sqlite
*.sqlite3
# IDE / editor
.idea/
.vscode/
*.swp
*.swo
*~
# OS files
.DS_Store
Thumbs.db
# Docker compose overrides
compose.override.yml
docker-compose.override.yml
+18
View File
@@ -0,0 +1,18 @@
# syntax=docker/dockerfile:1
FROM golang:1.25-alpine AS build
WORKDIR /src
COPY go.mod go.sum ./
RUN go mod download
COPY . .
RUN CGO_ENABLED=0 go build -ldflags="-s -w" -o /out/TokenThief ./
FROM alpine:3.20
RUN apk add --no-cache ca-certificates tzdata && \
addgroup -S app && adduser -S -G app app
COPY --from=build /out/TokenThief /usr/local/bin/TokenThief
COPY filter.yaml /etc/tokenthief/filter.yaml
ENV FILTER_FILE=/etc/tokenthief/filter.yaml \
TZ=Asia/Shanghai
USER app
EXPOSE 8080
ENTRYPOINT ["/usr/local/bin/TokenThief"]
+228
View File
@@ -0,0 +1,228 @@
# TokenThief
newapi 反向代理,采集请求/响应(包括 SSE/chunked 流式响应)并异步批量写入 ClickHouse。
## 特性
- 透明反代任意 HTTP 后端(默认目标为 newapi)。
- 流式响应(SSE、chunked、OpenAI 兼容 chat completions)边转发边缓冲,结束后整体入库。
- 异步队列 + 批量写入,主路径零阻塞。
- **DB 故障不影响代理服务**:连接失败时丢弃日志,后台持续重连。
- 通过 yaml 配置 glob 风格的黑白名单。
- 原样存储 headers/bodyString 字段保存 JSON 文本与 body 内容)。
## 快速开始
### 本地运行
```bash
cp .env.example .env
# 编辑 .env 设置 UPSTREAM_URL / CLICKHOUSE_URL
set -a; source .env; set +a
go run .
```
### Docker
```bash
docker build -t tokenthief .
docker run --rm -p 8080:8080 \
-e UPSTREAM_URL=http://newapi:3000 \
-e CLICKHOUSE_URL='clickhouses://tokenthief:<URL_ENCODED_PASSWORD>@clickhouse:9440/tokenthief' \
-v $(pwd)/filter.yaml:/etc/tokenthief/filter.yaml \
tokenthief
```
镜像内二进制路径为 `/usr/local/bin/TokenThief`
### Docker Compose
`compose.yml` 已包含 TokenThief 与 ClickHouse
```bash
cp .env.example .env
# 编辑 .env,设置 UPSTREAM_URL、强密码 CLICKHOUSE_PASSWORD 和对应的 CLICKHOUSE_URL
docker compose -f compose.yml up -d --build
```
默认端口:
| 服务 | 地址 |
|---|---|
| TokenThief | `http://localhost:8080` |
| ClickHouse | 仅 Compose 内部网络,不默认发布宿主机端口 |
Compose 在渲染配置时要求显式提供 `UPSTREAM_URL`、应用 DSN 和非空 ClickHouse 密码。密码中的 URL 特殊字符必须编码:
```text
clickhouse://tokenthief:<URL_ENCODED_PASSWORD>@thief_clickhouse:9000/tokenthief
```
### 运行测试
测试代码主要集中在 `tests/` 目录,数据库包还包含同包测试。
```bash
gofmt -w .
go test ./...
go vet ./...
```
构建二进制时不会引入测试内容:`_test.go` 不参与 `go build``tests/scripts/` 也不被主程序 importDocker 构建时 `.dockerignore` 会把整个 `tests/` 目录排除在 build context 之外。
## 环境变量
| 变量 | 说明 | 默认 |
|---|---|---|
| `LISTEN_ADDR` | 监听地址 | `:8080` |
| `UPSTREAM_URL` | 后端地址(必填) | - |
| `UPSTREAM_TLS_INSECURE_SKIP_VERIFY` | 跳过上游 HTTPS 证书校验;仅开发/可信内网自签证书场景使用 | `false` |
| `CLICKHOUSE_URL` | ClickHouse 原生协议地址(必填);支持严格 `host:port``clickhouse://` 或 TLS `clickhouses://` | - |
| `MAX_BODY_BYTES` | 单个请求体和响应体的记录上限 | `1048576` |
| `LOG_QUEUE_SIZE` | 异步队列条数上限 | `256` |
| `LOG_QUEUE_BYTES` | 异步队列总字节预算 | `67108864` |
| `LOG_BATCH_SIZE` | 批量写入条数 | `50` |
| `LOG_BATCH_INTERVAL` | 批量刷新间隔 | `2s` |
| `LOG_WORKERS` | worker 数 | `2` |
| `FILTER_FILE` | 黑白名单文件 | `./filter.yaml` |
| `DB_RECONNECT_INTERVAL` | DB 重连间隔 | `10s` |
| `READ_TIMEOUT` | 请求读取总超时 | `30s` |
| `WRITE_TIMEOUT` | 响应写入总超时 | `10m` |
| `IDLE_TIMEOUT` | HTTP keep-alive 空闲超时 | `5m` |
| `UPSTREAM_TIMEOUT` | 上游连接、TLS 握手、响应头等待超时 | `30s` |
| `UPSTREAM_RESPONSE_TIMEOUT` | 普通上游响应体总超时 | `30s` |
| `UPSTREAM_STREAM_IDLE_TIMEOUT` | SSE 上游响应体空闲超时,每次成功读取后重置 | `2m` |
| `TRUSTED_PROXIES` | 可信反向代理 IP/CIDR,逗号分隔;为空时忽略转发头 | 空 |
## 黑白名单(filter.yaml
`CLICKHOUSE_URL` 中的账号、密码和数据库名会传给 ClickHouse 原生协议连接。密码如果包含 `@``:``/``#` 等 URL 特殊字符,需要先做 URL encode。仅填写 `host:port` 时会使用 ClickHouse 默认用户、空密码和默认数据库。
默认 `filter.yaml` 已按 [newapi 官方文档](https://docs.newapi.pro/zh/docs/api) 列出全部 AI 模型接口(chat、completions、embeddings、moderations、rerank、realtime、audio、images、videos、Claude、Gemini 等)。
```yaml
mode: whitelist # whitelist | blacklist | disabled
patterns:
- /v1/chat/completions
- /v1/audio/* # 单段通配
- /v1/videos/** # 跨段通配
- /v1beta/models/*:generateContent # 段内通配
```
Pattern 语法:
| 通配符 | 含义 |
|---|---|
| `*` | 匹配单个路径段内除 `/` 之外的任意字符 |
| `**` | 跨段匹配任意字符(包括 `/` |
| `?` | 匹配单个非 `/` 字符 |
## 端点
- `/healthz` — 始终返回 200,不参与反代与日志。
- 其余路径 — 全部反代到 `UPSTREAM_URL`
## 数据库表
启动时自动创建单张 ClickHouse MergeTree 表 `proxy_logs`(详见 `db/migrate.go`),按 `started_at` 月份分区。DDL 如下:
```sql
CREATE TABLE IF NOT EXISTS proxy_logs (
request_id String,
method String,
path String,
query String,
client_ip String,
request_headers String,
request_body String,
request_truncated Bool DEFAULT false,
status_code Int32,
response_headers String,
response_body String,
response_truncated Bool DEFAULT false,
is_stream Bool DEFAULT false,
latency_ms Int64,
started_at DateTime64(3),
finished_at DateTime64(3),
error String
)
ENGINE = MergeTree
PARTITION BY toYYYYMM(started_at)
ORDER BY (started_at, request_id);
```
完整字段说明如下。
### 字段一览
| 字段 | 类型 | 可空 | 说明 |
|---|---|---|---|
| `request_id` | `String` | 否 | 每次请求由代理生成的 16 字节随机 hex(32 个字符),同时写入响应头 `X-Request-Id`,便于客户端日志关联 |
| `method` | `String` | 否 | HTTP 方法(`GET` / `POST` / …),与客户端实际发送一致 |
| `path` | `String` | 否 | 请求路径,不含 query string,例如 `/v1/chat/completions` |
| `query` | `String` | 否 | 原始 query string(不含 `?`),例如 `model=gpt-4&stream=true`;无 query 时为空串 |
| `client_ip` | `String` | 否 | 默认记录 TCP 对端;仅直接对端命中 `TRUSTED_PROXIES` 时,从右向左解析 `X-Forwarded-For`,无 XFF 时使用有效的 `X-Real-IP` |
| `request_headers` | `String` | 否 | 完整请求头,序列化为 `{"Header-Name": ["value1", "value2"], ...}` 的 JSON 对象;**注意 Authorization、Cookie、API-Key 等敏感头未脱敏**,按设计原样存储 |
| `request_body` | `String` | 否 | 请求体内容。最多保留 `MAX_BODY_BYTES`(默认 1 MiB)字节,超出部分丢弃;读取失败时请求不会转发 |
| `request_truncated` | `Bool` | 否 | 请求体是否被 `MAX_BODY_BYTES` 截断。即使截断,下游 newapi 仍会通过 `MultiReader` 接收到完整 body,不影响功能 |
| `status_code` | `Int32` | 否 | 上游返回的 HTTP 状态码。`502` 通常意味着上游连接失败;WebSocket 成功升级记录为 `101` |
| `response_headers` | `String` | 否 | 响应头,结构同 `request_headers`。对于 SSE,会包含 `Content-Type: text/event-stream` 等 |
| `response_body` | `String` | 否 | 响应体内容。对于流式响应(SSE / chunked),这里保存的是**所有 chunk 拼接后的完整字节流**(包含 `data:` 前缀、`\n\n` 分隔符以及最后的 `[DONE]`),方便事后离线解析。最多保留 `MAX_BODY_BYTES` 字节 |
| `response_truncated` | `Bool` | 否 | 响应体是否被截断。截断只影响数据库存储,客户端始终收到完整数据 |
| `is_stream` | `Bool` | 否 | 是否 SSE 响应,仅接受媒体类型 `text/event-stream` |
| `latency_ms` | `Int64` | 否 | 端到端耗时(毫秒),从代理接收到请求到响应完成。对流式响应 = 从首请求到最后一个 chunk 发出 |
| `started_at` | `DateTime64(3)` | 否 | 代理接收到请求的时刻;ClickHouse 按 `toYYYYMM(started_at)` 月度分区 |
| `finished_at` | `DateTime64(3)` | 否 | 响应完全写回客户端(包括所有 chunk)的时刻 |
| `error` | `String` | 否 | 仅在反代过程中出现错误时填充。常见值:上游不可达、超时、读取请求体失败等 |
### 常用查询示例
```sql
-- 查看最近 20 次失败请求
SELECT started_at, path, status_code, error
FROM proxy_logs
WHERE status_code >= 400 OR error != ''
ORDER BY started_at DESC
LIMIT 20;
-- 查看某次请求的完整内容
SELECT
request_id,
method, path,
request_body AS req_text,
response_body AS resp_text,
latency_ms, is_stream
FROM proxy_logs
WHERE request_id = '0123456789abcdef0123456789abcdef';
-- 按模型统计调用量(从请求体里提取 JSON 字段)
SELECT
JSONExtractString(request_body, 'model') AS model,
count(*) AS calls,
toInt32(avg(latency_ms)) AS avg_ms
FROM proxy_logs
WHERE path = '/v1/chat/completions'
AND started_at > now() - INTERVAL 1 DAY
GROUP BY 1
ORDER BY calls DESC;
-- 查 Authorization(注意:敏感信息)
SELECT JSONExtractRaw(request_headers, 'Authorization') FROM proxy_logs LIMIT 5;
```
### 注意事项
- **敏感信息**:请求头中的 `Authorization``Cookie``X-Api-Key` 等**未脱敏**。如需脱敏请在 `proxy/capture.go``headersJSON` 中改造,或对数据库做列级权限控制。
- **body 编码**ClickHouse 以 `String` 保存 body 内容;文本接口可直接查询,二进制或压缩内容需按业务格式离线解析。
- **WebSocket**:只记录 101 握手元数据,`response_body` 为空;WS 帧内容不采集,进程关闭时会关闭受管升级连接。
- **截断**`request_truncated` / `response_truncated``true` 时,对应 `*_body` 仅包含前 `MAX_BODY_BYTES` 字节。需保留完整内容请调高 `MAX_BODY_BYTES`,但要警惕数据库膨胀。
## 设计要点
- 请求体读取使用 `LimitReader`,超长仅记录前 `MAX_BODY_BYTES` 字节,下游仍能拿到完整 body。
- `httputil.ReverseProxy` + `FlushInterval = -1`,自定义 `ResponseWriter` 同时实现 `Flusher`/`Hijacker`,写入时先转发再缓冲,保证流式实时性。
- 日志通过非阻塞 channel 投递,队列满或 DB 不健康时直接丢弃(每 30 秒打印 metrics)。
- DB 健康状态机:写入失败立即标记 unhealthy,后台 ping 恢复后重新启用。
- ClickHouse `PrepareBatch` 失败会保留整批重试;逐项 `Append` 失败只保留明确失败项,已成功追加项继续发送。`Send` 返回错误时提交结果可能不明,系统不会自动重发该批,避免静默重复,并通过 `ambiguous_send` 计数暴露可能丢失;进程内 best-effort 队列不承诺分布式 exactly-once。
- 启动会校验现有 `proxy_logs` 的列、引擎、分区键和排序键;不兼容 schema 会保持数据库 unhealthy,不自动重建数据表。
- 所有 batch 路径都会执行清理;`Abort`/关闭失败会与原始错误合并记录,但不能使模糊提交变得可判定。
+65
View File
@@ -0,0 +1,65 @@
services:
tokenthief:
build:
context: .
dockerfile: Dockerfile
image: tokenthief:0.1.0
container_name: tokenthief
restart: unless-stopped
depends_on:
thief_clickhouse:
condition: service_healthy
ports:
- "${LISTEN_PORT:-8080}:8080"
environment:
LISTEN_ADDR: ":8080"
UPSTREAM_URL: "${UPSTREAM_URL:?set UPSTREAM_URL}"
UPSTREAM_TLS_INSECURE_SKIP_VERIFY: "${UPSTREAM_TLS_INSECURE_SKIP_VERIFY:-false}"
CLICKHOUSE_URL: "${CLICKHOUSE_URL:?set CLICKHOUSE_URL with URL-encoded credentials}"
MAX_BODY_BYTES: "${MAX_BODY_BYTES:-1048576}"
LOG_QUEUE_SIZE: "${LOG_QUEUE_SIZE:-256}"
LOG_QUEUE_BYTES: "${LOG_QUEUE_BYTES:-67108864}"
LOG_BATCH_SIZE: "${LOG_BATCH_SIZE:-50}"
LOG_BATCH_INTERVAL: "${LOG_BATCH_INTERVAL:-2s}"
LOG_WORKERS: "${LOG_WORKERS:-2}"
FILTER_FILE: "/etc/tokenthief/filter.yaml"
DB_RECONNECT_INTERVAL: "${DB_RECONNECT_INTERVAL:-10s}"
READ_TIMEOUT: "${READ_TIMEOUT:-30s}"
WRITE_TIMEOUT: "${WRITE_TIMEOUT:-10m}"
IDLE_TIMEOUT: "${IDLE_TIMEOUT:-5m}"
UPSTREAM_TIMEOUT: "${UPSTREAM_TIMEOUT:-30s}"
UPSTREAM_RESPONSE_TIMEOUT: "${UPSTREAM_RESPONSE_TIMEOUT:-30s}"
UPSTREAM_STREAM_IDLE_TIMEOUT: "${UPSTREAM_STREAM_IDLE_TIMEOUT:-2m}"
TRUSTED_PROXIES: "${TRUSTED_PROXIES:-}"
TZ: "${TZ:-Asia/Shanghai}"
volumes:
- ./filter.yaml:/etc/tokenthief/filter.yaml:ro
networks:
- tokenthief
thief_clickhouse:
image: clickhouse/clickhouse-server:25.3.3.42-alpine
container_name: tokenthief-clickhouse
restart: unless-stopped
environment:
CLICKHOUSE_USER: "${CLICKHOUSE_USER:-tokenthief}"
CLICKHOUSE_PASSWORD: "${CLICKHOUSE_PASSWORD:?set a strong CLICKHOUSE_PASSWORD}"
CLICKHOUSE_DB: "${CLICKHOUSE_DB:-tokenthief}"
TZ: "${TZ:-Asia/Shanghai}"
volumes:
- clickhouse_data:/var/lib/clickhouse
healthcheck:
test: ["CMD-SHELL", "clickhouse-client --user \"$${CLICKHOUSE_USER}\" --password \"$${CLICKHOUSE_PASSWORD}\" --database \"$${CLICKHOUSE_DB}\" --query 'SELECT 1'"]
interval: 10s
timeout: 5s
retries: 5
start_period: 10s
networks:
- tokenthief
volumes:
clickhouse_data:
networks:
tokenthief:
driver: bridge
+245
View File
@@ -0,0 +1,245 @@
package config
import (
"errors"
"fmt"
"net/netip"
"net/url"
"os"
"strconv"
"strings"
"time"
"git.misaka.ren/M1saka/token_thief/db"
)
// Config 保存从环境变量加载的全部运行配置。
type Config struct {
ListenAddr string
UpstreamURL *url.URL
ClickHouseURL string
MaxBodyBytes int64
LogQueueSize int
LogQueueBytes int64
LogBatchSize int
LogBatchInterval time.Duration
LogWorkers int
FilterFile string
DBReconnectInterval time.Duration
ReadTimeout time.Duration
WriteTimeout time.Duration
IdleTimeout time.Duration
UpstreamTimeout time.Duration
UpstreamResponseTimeout time.Duration
UpstreamStreamIdleTimeout time.Duration
UpstreamTLSInsecureSkipVerify bool
TrustedProxies []netip.Prefix
}
// Load 从进程环境读取配置;缺少必填字段时返回错误。
func Load() (*Config, error) {
maxBodyBytes, err := getEnvInt64("MAX_BODY_BYTES", 1<<20)
if err != nil {
return nil, err
}
logQueueSize, err := getEnvInt("LOG_QUEUE_SIZE", 256)
if err != nil {
return nil, err
}
logQueueBytes, err := getEnvInt64("LOG_QUEUE_BYTES", 64<<20)
if err != nil {
return nil, err
}
logBatchSize, err := getEnvInt("LOG_BATCH_SIZE", 50)
if err != nil {
return nil, err
}
logBatchInterval, err := getEnvDuration("LOG_BATCH_INTERVAL", 2*time.Second)
if err != nil {
return nil, err
}
logWorkers, err := getEnvInt("LOG_WORKERS", 2)
if err != nil {
return nil, err
}
dbReconnectInterval, err := getEnvDuration("DB_RECONNECT_INTERVAL", 10*time.Second)
if err != nil {
return nil, err
}
readTimeout, err := getEnvDuration("READ_TIMEOUT", 30*time.Second)
if err != nil {
return nil, err
}
writeTimeout, err := getEnvDuration("WRITE_TIMEOUT", 10*time.Minute)
if err != nil {
return nil, err
}
idleTimeout, err := getEnvDuration("IDLE_TIMEOUT", 5*time.Minute)
if err != nil {
return nil, err
}
upstreamTimeout, err := getEnvDuration("UPSTREAM_TIMEOUT", 30*time.Second)
if err != nil {
return nil, err
}
upstreamResponseTimeout, err := getEnvDuration("UPSTREAM_RESPONSE_TIMEOUT", 30*time.Second)
if err != nil {
return nil, err
}
upstreamStreamIdleTimeout, err := getEnvDuration("UPSTREAM_STREAM_IDLE_TIMEOUT", 2*time.Minute)
if err != nil {
return nil, err
}
upstreamTLSInsecureSkipVerify, err := getEnvBool("UPSTREAM_TLS_INSECURE_SKIP_VERIFY", false)
if err != nil {
return nil, err
}
trustedProxies, err := getEnvPrefixes("TRUSTED_PROXIES")
if err != nil {
return nil, err
}
cfg := &Config{
ListenAddr: getEnv("LISTEN_ADDR", ":8080"),
ClickHouseURL: os.Getenv("CLICKHOUSE_URL"),
MaxBodyBytes: maxBodyBytes,
LogQueueSize: logQueueSize,
LogQueueBytes: logQueueBytes,
LogBatchSize: logBatchSize,
LogBatchInterval: logBatchInterval,
LogWorkers: logWorkers,
FilterFile: getEnv("FILTER_FILE", "./filter.yaml"),
DBReconnectInterval: dbReconnectInterval,
ReadTimeout: readTimeout,
WriteTimeout: writeTimeout,
IdleTimeout: idleTimeout,
UpstreamTimeout: upstreamTimeout,
UpstreamResponseTimeout: upstreamResponseTimeout,
UpstreamStreamIdleTimeout: upstreamStreamIdleTimeout,
UpstreamTLSInsecureSkipVerify: upstreamTLSInsecureSkipVerify,
TrustedProxies: trustedProxies,
}
upstream := os.Getenv("UPSTREAM_URL")
if upstream == "" {
return nil, errors.New("UPSTREAM_URL is required")
}
u, err := url.Parse(upstream)
if err != nil {
return nil, fmt.Errorf("invalid UPSTREAM_URL: %w", err)
}
if u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid UPSTREAM_URL: %q", upstream)
}
cfg.UpstreamURL = u
if cfg.ClickHouseURL == "" {
return nil, errors.New("CLICKHOUSE_URL is required")
}
if _, err := db.ClickHouseOptions(cfg.ClickHouseURL); err != nil {
return nil, err
}
positiveValues := []struct {
key string
value int64
}{
{key: "MAX_BODY_BYTES", value: cfg.MaxBodyBytes},
{key: "LOG_QUEUE_SIZE", value: int64(cfg.LogQueueSize)},
{key: "LOG_QUEUE_BYTES", value: cfg.LogQueueBytes},
{key: "LOG_BATCH_SIZE", value: int64(cfg.LogBatchSize)},
{key: "LOG_BATCH_INTERVAL", value: int64(cfg.LogBatchInterval)},
{key: "LOG_WORKERS", value: int64(cfg.LogWorkers)},
{key: "DB_RECONNECT_INTERVAL", value: int64(cfg.DBReconnectInterval)},
{key: "READ_TIMEOUT", value: int64(cfg.ReadTimeout)},
{key: "WRITE_TIMEOUT", value: int64(cfg.WriteTimeout)},
{key: "IDLE_TIMEOUT", value: int64(cfg.IdleTimeout)},
{key: "UPSTREAM_TIMEOUT", value: int64(cfg.UpstreamTimeout)},
{key: "UPSTREAM_RESPONSE_TIMEOUT", value: int64(cfg.UpstreamResponseTimeout)},
{key: "UPSTREAM_STREAM_IDLE_TIMEOUT", value: int64(cfg.UpstreamStreamIdleTimeout)},
}
for _, item := range positiveValues {
if item.value <= 0 {
return nil, fmt.Errorf("%s must be greater than zero", item.key)
}
}
return cfg, nil
}
func getEnv(key, def string) string {
if v := os.Getenv(key); v != "" {
return v
}
return def
}
func getEnvInt(key string, def int) (int, error) {
if v := os.Getenv(key); v != "" {
n, err := strconv.Atoi(v)
if err != nil {
return 0, fmt.Errorf("invalid %s: %w", key, err)
}
return n, nil
}
return def, nil
}
func getEnvInt64(key string, def int64) (int64, error) {
if v := os.Getenv(key); v != "" {
n, err := strconv.ParseInt(v, 10, 64)
if err != nil {
return 0, fmt.Errorf("invalid %s: %w", key, err)
}
return n, nil
}
return def, nil
}
func getEnvDuration(key string, def time.Duration) (time.Duration, error) {
if v := os.Getenv(key); v != "" {
d, err := time.ParseDuration(v)
if err != nil {
return 0, fmt.Errorf("invalid %s: %w", key, err)
}
return d, nil
}
return def, nil
}
func getEnvBool(key string, def bool) (bool, error) {
if v := os.Getenv(key); v != "" {
b, err := strconv.ParseBool(v)
if err != nil {
return false, fmt.Errorf("invalid %s: %w", key, err)
}
return b, nil
}
return def, nil
}
func getEnvPrefixes(key string) ([]netip.Prefix, error) {
v := os.Getenv(key)
if v == "" {
return nil, nil
}
parts := strings.Split(v, ",")
prefixes := make([]netip.Prefix, 0, len(parts))
for _, part := range parts {
value := strings.TrimSpace(part)
if value == "" {
return nil, fmt.Errorf("invalid %s: empty IP or CIDR", key)
}
if addr, err := netip.ParseAddr(value); err == nil {
prefixes = append(prefixes, netip.PrefixFrom(addr, addr.BitLen()))
continue
}
prefix, err := netip.ParsePrefix(value)
if err != nil || !prefix.IsValid() || prefix != prefix.Masked() {
return nil, fmt.Errorf("invalid %s entry %q", key, value)
}
prefixes = append(prefixes, prefix)
}
return prefixes, nil
}
+134
View File
@@ -0,0 +1,134 @@
package config
import (
"fmt"
"os"
"regexp"
"strings"
"gopkg.in/yaml.v3"
)
// FilterMode 决定如何应用 patterns。
type FilterMode string
const (
FilterWhitelist FilterMode = "whitelist"
FilterBlacklist FilterMode = "blacklist"
FilterDisabled FilterMode = "disabled"
)
// Filter 表示路径匹配规则。
type Filter struct {
Mode FilterMode `yaml:"mode"`
Patterns []string `yaml:"patterns"`
compiled []*regexp.Regexp
}
// LoadFilter 从 yaml 文件加载过滤配置;文件不存在则默认全部记录。
//
// Pattern 语法(glob 风格):
// - `*` 匹配单个路径段内除 `/` 之外的任意字符(包括零个)。
// - `**` 匹配任意字符,含 `/`,可跨段。
// - `?` 匹配单个非 `/` 字符。
// - 其它字符按字面匹配。
//
// 示例:
// - `/v1/audio/*` 匹配 /v1/audio/speech、/v1/audio/transcriptions
// - `/v1/videos/**` 匹配 /v1/videos/任意子路径
// - `/v1beta/models/*:generateContent` 匹配 Gemini 风格端点
func LoadFilter(filePath string) (*Filter, error) {
data, err := os.ReadFile(filePath)
if err != nil {
if os.IsNotExist(err) {
return &Filter{Mode: FilterDisabled}, nil
}
return nil, err
}
var f Filter
if err := yaml.Unmarshal(data, &f); err != nil {
return nil, fmt.Errorf("parse filter yaml: %w", err)
}
switch f.Mode {
case FilterWhitelist, FilterBlacklist, FilterDisabled:
case "":
f.Mode = FilterDisabled
default:
return nil, fmt.Errorf("unknown filter mode: %q", f.Mode)
}
if err := f.compile(); err != nil {
return nil, err
}
return &f, nil
}
// NewFilter 程序化构造一个 Filter(主要供测试使用)。
func NewFilter(mode FilterMode, patterns []string) (*Filter, error) {
f := &Filter{Mode: mode, Patterns: append([]string(nil), patterns...)}
if err := f.compile(); err != nil {
return nil, err
}
return f, nil
}
func (f *Filter) compile() error {
f.compiled = make([]*regexp.Regexp, 0, len(f.Patterns))
for _, p := range f.Patterns {
re, err := CompileGlob(p)
if err != nil {
return fmt.Errorf("invalid pattern %q: %w", p, err)
}
f.compiled = append(f.compiled, re)
}
return nil
}
// ShouldLog 决定一个请求 path 是否需要被记录。
func (f *Filter) ShouldLog(reqPath string) bool {
if f == nil || f.Mode == FilterDisabled {
return true
}
matched := false
for _, re := range f.compiled {
if re.MatchString(reqPath) {
matched = true
break
}
}
switch f.Mode {
case FilterWhitelist:
return matched
case FilterBlacklist:
return !matched
default:
return true
}
}
// CompileGlob 将 glob 风格 pattern 转换为 anchored 正则表达式。
func CompileGlob(pattern string) (*regexp.Regexp, error) {
var sb strings.Builder
sb.WriteString("^")
for i := 0; i < len(pattern); i++ {
c := pattern[i]
switch c {
case '*':
if i+1 < len(pattern) && pattern[i+1] == '*' {
sb.WriteString(".*")
i++
} else {
sb.WriteString("[^/]*")
}
case '?':
sb.WriteString("[^/]")
case '.', '+', '(', ')', '|', '^', '$', '{', '}', '[', ']', '\\':
sb.WriteByte('\\')
sb.WriteByte(c)
default:
sb.WriteByte(c)
}
}
sb.WriteString("$")
return regexp.Compile(sb.String())
}
+221
View File
@@ -0,0 +1,221 @@
package db
import (
"context"
"fmt"
"log"
"net"
"net/url"
"strconv"
"strings"
"sync"
"time"
"github.com/ClickHouse/clickhouse-go/v2"
)
// Pool 封装 ClickHouse 连接并维护健康状态,DB 故障时不阻塞调用方。
type Pool struct {
dsn string
reconnectInterval time.Duration
mu sync.RWMutex
conn clickhouse.Conn
healthy bool
}
// NewPool 创建 Pool。即使首次连接失败也返回非 nil 实例,后台会持续重试。
func NewPool(ctx context.Context, dsn string, reconnectInterval time.Duration) *Pool {
p := &Pool{dsn: dsn, reconnectInterval: reconnectInterval}
if err := p.connect(ctx); err != nil {
log.Printf("[db] initial connect failed: %v (service continues without DB)", err)
}
go p.watch(ctx)
return p
}
func (p *Pool) connect(ctx context.Context) error {
cctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
opts, err := ClickHouseOptions(p.dsn)
if err != nil {
return err
}
conn, err := clickhouse.Open(opts)
if err != nil {
return err
}
if err := conn.Ping(cctx); err != nil {
_ = conn.Close()
return err
}
// 先 migrate,再 swap:避免新连接 migrate 失败时取代掉旧的可用连接。
mctx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if err := migrate(mctx, conn); err != nil {
_ = conn.Close()
log.Printf("[db] migrate failed: %v", err)
return err
}
p.mu.Lock()
if p.conn != nil {
_ = p.conn.Close()
}
p.conn = conn
p.healthy = true
p.mu.Unlock()
log.Printf("[db] connected and migrated")
return nil
}
// ClickHouseOptions converts the supported CLICKHOUSE_URL subset into driver options.
func ClickHouseOptions(raw string) (*clickhouse.Options, error) {
if !strings.Contains(raw, "://") {
if err := validateHostPort(raw); err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
return &clickhouse.Options{Addr: []string{raw}}, nil
}
u, err := url.Parse(raw)
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
if u.Scheme != "clickhouse" && u.Scheme != "clickhouses" {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: unsupported scheme %q", u.Scheme)
}
if u.Fragment != "" {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: fragment is not allowed")
}
if err := validateHostPort(u.Host); err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
query, err := url.ParseQuery(u.RawQuery)
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: query: %w", err)
}
for key, values := range query {
switch key {
case "secure", "skip_verify", "compress":
default:
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: unsupported parameter %q", key)
}
if len(values) != 1 {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: parameter %q must occur once", key)
}
}
wantSecure := u.Scheme == "clickhouses"
if value, ok := query["secure"]; ok {
secure, err := strconv.ParseBool(value[0])
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: secure: %w", err)
}
if secure != wantSecure {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: secure conflicts with %s", u.Scheme)
}
}
if _, ok := query["skip_verify"]; ok && !wantSecure {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: skip_verify requires clickhouses")
}
if value, ok := query["skip_verify"]; ok && value[0] != "" {
if _, err := strconv.ParseBool(value[0]); err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: skip_verify: %w", err)
}
}
if value, ok := query["compress"]; ok {
switch value[0] {
case "true", "false", "none", "zstd", "lz4", "lz4hc", "gzip", "deflate", "br":
default:
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: unsupported compression %q", value[0])
}
}
// Let the driver decode userinfo/path and interpret compression and TLS values.
if wantSecure {
query.Set("secure", "true")
u.RawQuery = query.Encode()
}
opts, err := clickhouse.ParseDSN(u.String())
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
return opts, nil
}
func validateHostPort(address string) error {
if address == "" {
return fmt.Errorf("missing host and port")
}
host, port, err := net.SplitHostPort(address)
if err != nil {
return fmt.Errorf("expected host:port: %w", err)
}
if host == "" {
return fmt.Errorf("missing host")
}
n, err := strconv.ParseUint(port, 10, 16)
if err != nil || n == 0 {
return fmt.Errorf("invalid port %q", port)
}
return nil
}
// watch 定期探活;不健康时尝试重连。
func (p *Pool) watch(ctx context.Context) {
t := time.NewTicker(p.reconnectInterval)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
if p.Healthy() {
if conn := p.Get(); conn != nil {
pctx, cancel := context.WithTimeout(ctx, 3*time.Second)
if err := conn.Ping(pctx); err != nil {
log.Printf("[db] ping failed, marking unhealthy: %v", err)
p.MarkUnhealthy()
}
cancel()
}
continue
}
if err := p.connect(ctx); err != nil {
log.Printf("[db] reconnect failed: %v", err)
}
}
}
}
// Healthy 报告连接是否可用。
func (p *Pool) Healthy() bool {
p.mu.RLock()
defer p.mu.RUnlock()
return p.healthy
}
// MarkUnhealthy 由调用方在写入失败后调用。
func (p *Pool) MarkUnhealthy() {
p.mu.Lock()
p.healthy = false
p.mu.Unlock()
}
// Get 返回当前连接,可能为 nil。
func (p *Pool) Get() clickhouse.Conn {
p.mu.RLock()
defer p.mu.RUnlock()
return p.conn
}
// Close 释放底层连接。
func (p *Pool) Close() {
p.mu.Lock()
defer p.mu.Unlock()
if p.conn != nil {
_ = p.conn.Close()
p.conn = nil
}
p.healthy = false
}
+127
View File
@@ -0,0 +1,127 @@
package db
import (
"reflect"
"testing"
)
func TestClickHouseOptionsHostPort(t *testing.T) {
opts, err := ClickHouseOptions("localhost:9000")
if err != nil {
t.Fatalf("ClickHouseOptions: %v", err)
}
if len(opts.Addr) != 1 || opts.Addr[0] != "localhost:9000" {
t.Fatalf("Addr=%v want [localhost:9000]", opts.Addr)
}
if opts.Auth.Username != "" || opts.Auth.Password != "" || opts.Auth.Database != "" {
t.Fatalf("Auth=%+v want empty", opts.Auth)
}
}
func TestClickHouseOptionsURLWithAuthAndDatabase(t *testing.T) {
opts, err := ClickHouseOptions("clickhouse://user%40name:p%2Fass@localhost:9000/token%20thief?compress=zstd")
if err != nil {
t.Fatalf("ClickHouseOptions: %v", err)
}
if len(opts.Addr) != 1 || opts.Addr[0] != "localhost:9000" {
t.Fatalf("Addr=%v want [localhost:9000]", opts.Addr)
}
if opts.Auth.Username != "user@name" {
t.Fatalf("Username=%q want user@name", opts.Auth.Username)
}
if opts.Auth.Password != "p/ass" {
t.Fatalf("Password=%q want p/ass", opts.Auth.Password)
}
if opts.Auth.Database != "token thief" {
t.Fatalf("Database=%q want token thief", opts.Auth.Database)
}
if opts.Compression == nil {
t.Fatal("Compression=nil want enabled")
}
}
func TestClickHouseOptionsURLWithDatabaseOnly(t *testing.T) {
opts, err := ClickHouseOptions("clickhouse://localhost:9000/tokenthief")
if err != nil {
t.Fatalf("ClickHouseOptions: %v", err)
}
if len(opts.Addr) != 1 || opts.Addr[0] != "localhost:9000" {
t.Fatalf("Addr=%v want [localhost:9000]", opts.Addr)
}
if opts.Auth.Database != "tokenthief" {
t.Fatalf("Database=%q want tokenthief", opts.Auth.Database)
}
}
func TestClickHouseOptionsTLS(t *testing.T) {
opts, err := ClickHouseOptions("clickhouses://localhost:9440/tokenthief?skip_verify=true")
if err != nil {
t.Fatalf("ClickHouseOptions: %v", err)
}
if opts.TLS == nil || !opts.TLS.InsecureSkipVerify {
t.Fatalf("TLS=%+v want InsecureSkipVerify", opts.TLS)
}
}
func TestClickHouseOptionsRejectsInvalidDSN(t *testing.T) {
tests := []string{
"localhost",
"localhost:http",
"clickhouse://localhost/tokenthief",
"clickhouse:///tokenthief",
"clickhouse://localhost:9000/db#fragment",
"clickhouse://localhost:9000/db?unknown=true",
"clickhouse://localhost:9000/db?compress=snappy",
"clickhouse://localhost:9000/db?compress=lz4&compress=zstd",
"clickhouse://localhost:9000/db?compress=%zz",
"clickhouse://localhost:9000/db?secure=true",
"clickhouses://localhost:9440/db?secure=false",
"clickhouses://localhost:9440/db?skip_verify=maybe",
"clickhouse://localhost:9000/db?skip_verify=true",
"https://localhost:9440/db",
}
for _, dsn := range tests {
t.Run(dsn, func(t *testing.T) {
if _, err := ClickHouseOptions(dsn); err == nil {
t.Fatalf("ClickHouseOptions(%q) succeeded", dsn)
}
})
}
}
func TestValidateProxyLogsSchema(t *testing.T) {
columns := append([]schemaColumn(nil), proxyLogsColumns...)
table := schemaTable{
engine: "MergeTree",
partitionKey: "toYYYYMM(started_at)",
sortingKey: "started_at, request_id",
}
if err := validateProxyLogsSchema(columns, table); err != nil {
t.Fatalf("validateProxyLogsSchema: %v", err)
}
badColumns := append([]schemaColumn(nil), columns...)
badColumns[8].typ = "UInt16"
tests := []struct {
name string
columns []schemaColumn
table schemaTable
}{
{"missing column", columns[:16], table},
{"wrong type", badColumns, table},
{"wrong engine", columns, schemaTable{"ReplacingMergeTree", table.partitionKey, table.sortingKey}},
{"wrong partition", columns, schemaTable{table.engine, "toDate(started_at)", table.sortingKey}},
{"wrong sorting", columns, schemaTable{table.engine, table.partitionKey, "request_id"}},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if err := validateProxyLogsSchema(test.columns, test.table); err == nil {
t.Fatal("validateProxyLogsSchema succeeded")
}
})
}
if !reflect.DeepEqual(columns, proxyLogsColumns) {
t.Fatal("schema validation mutated columns")
}
}
+134
View File
@@ -0,0 +1,134 @@
package db
import (
"context"
"fmt"
"strings"
"github.com/ClickHouse/clickhouse-go/v2"
)
const schemaSQL = `
CREATE TABLE IF NOT EXISTS proxy_logs (
request_id String,
method String,
path String,
query String,
client_ip String,
request_headers String,
request_body String,
request_truncated Bool DEFAULT false,
status_code Int32,
response_headers String,
response_body String,
response_truncated Bool DEFAULT false,
is_stream Bool DEFAULT false,
latency_ms Int64,
started_at DateTime64(3),
finished_at DateTime64(3),
error String
)
ENGINE = MergeTree
PARTITION BY toYYYYMM(started_at)
ORDER BY (started_at, request_id)
`
func migrate(ctx context.Context, conn clickhouse.Conn) error {
if err := conn.Exec(ctx, schemaSQL); err != nil {
return fmt.Errorf("create proxy_logs: %w", err)
}
rows, err := conn.Query(ctx, `
SELECT name, type
FROM system.columns
WHERE database = currentDatabase() AND table = 'proxy_logs'
ORDER BY position`)
if err != nil {
return fmt.Errorf("query proxy_logs columns: %w", err)
}
var columns []schemaColumn
for rows.Next() {
var column schemaColumn
if err := rows.Scan(&column.name, &column.typ); err != nil {
rows.Close()
return fmt.Errorf("scan proxy_logs columns: %w", err)
}
columns = append(columns, column)
}
if err := rows.Err(); err != nil {
rows.Close()
return fmt.Errorf("read proxy_logs columns: %w", err)
}
rows.Close()
var table schemaTable
err = conn.QueryRow(ctx, `
SELECT engine, partition_key, sorting_key
FROM system.tables
WHERE database = currentDatabase() AND name = 'proxy_logs'`).Scan(
&table.engine, &table.partitionKey, &table.sortingKey,
)
if err != nil {
return fmt.Errorf("query proxy_logs table: %w", err)
}
if err := validateProxyLogsSchema(columns, table); err != nil {
return fmt.Errorf("incompatible proxy_logs schema: %w", err)
}
return nil
}
type schemaColumn struct {
name string
typ string
}
type schemaTable struct {
engine string
partitionKey string
sortingKey string
}
var proxyLogsColumns = []schemaColumn{
{"request_id", "String"},
{"method", "String"},
{"path", "String"},
{"query", "String"},
{"client_ip", "String"},
{"request_headers", "String"},
{"request_body", "String"},
{"request_truncated", "Bool"},
{"status_code", "Int32"},
{"response_headers", "String"},
{"response_body", "String"},
{"response_truncated", "Bool"},
{"is_stream", "Bool"},
{"latency_ms", "Int64"},
{"started_at", "DateTime64(3)"},
{"finished_at", "DateTime64(3)"},
{"error", "String"},
}
func validateProxyLogsSchema(columns []schemaColumn, table schemaTable) error {
if len(columns) != len(proxyLogsColumns) {
return fmt.Errorf("got %d columns, want %d", len(columns), len(proxyLogsColumns))
}
for i, want := range proxyLogsColumns {
if columns[i] != want {
return fmt.Errorf("column %d is %s %s, want %s %s", i+1, columns[i].name, columns[i].typ, want.name, want.typ)
}
}
if table.engine != "MergeTree" {
return fmt.Errorf("engine is %q, want MergeTree", table.engine)
}
if compactExpression(table.partitionKey) != "toYYYYMM(started_at)" {
return fmt.Errorf("partition key is %q, want toYYYYMM(started_at)", table.partitionKey)
}
if compactExpression(table.sortingKey) != "started_at,request_id" {
return fmt.Errorf("sorting key is %q, want started_at, request_id", table.sortingKey)
}
return nil
}
func compactExpression(value string) string {
return strings.Join(strings.Fields(value), "")
}
@@ -0,0 +1,66 @@
# ClickHouse Migration 安全修复计划
> [!NOTE]
> This document may not reflect the current implementation.
> See the final report for up-to-date state:
> [Final Report](../reports/reliability-security-fixes.md)
**变更规模:** 大型跨模块修复。现有迁移主体复用,仅重新实施受本次审查影响的任务及其集成依赖。
## 全局约束
- 保留非阻塞 best-effort 日志语义,并通过 Git 提交保留可审查的变更证据。
- 先用失败测试固定根因,再实施最小修复。
- `Send` 模糊失败不重试;不宣称分布式 exactly-once。
- WebSocket 只管理连接并记录 101 元数据,不采集帧。
### Task 1: 严格配置与安全默认
**文件:** `config/config.go``tests/config/config_test.go`
- [ ] 严格解析整数、布尔值和 duration,拒绝非正 body/队列/timeout。
- [ ] 增加 `LOG_QUEUE_BYTES`、响应总/idle timeout、`TRUSTED_PROXIES`
- [ ] 在配置加载阶段校验 ClickHouse DSN,并补齐边界测试。
### Task 2: Queue 生命周期、字节预算与确定性写入
**文件:** `logger/queue.go``tests/logger/queue_test.go`
- [ ] 用同步状态机消除 Submit/Stop send-close 竞态,Stop 使用单一 context 总预算。
- [ ] 增加条数/字节双预算并在所有消费、丢弃和 shutdown 路径释放预留。
- [ ] 区分 Prepare、Append、Send 错误;Append 失败清理 batch,Send 模糊失败不重试。
- [ ] 增加并发关闭、预算、清理、部分失败和模糊提交测试。
### Task 3: ClickHouse DSN 与 schema 验证
**文件:** `db/clickhouse.go``db/migrate.go``db/clickhouse_test.go`
- [ ] 严格解析 scheme、TLS 和白名单 query 参数。
- [ ] 创建 schema 后校验列、引擎、分区和排序键;不兼容时保持 unhealthy。
- [ ] 增加 DSN 与 schema 元数据校验测试。
### Task 4: HTTP/SSE 代理完整性与可信来源
**文件:** `proxy/proxy.go``proxy/writer.go``proxy/capture.go``proxy/sse.go``tests/proxy/*`
- [ ] 请求体读取失败 fail closed;502 使用固定客户端文本。
- [ ] 记录响应写错误/短写,响应中断不提交完整日志。
- [ ] 仅对真正 SSE 按完整事件检测终止,并等待正常返回后提交。
- [ ] 为普通响应总 timeout 与 SSE idle timeout 包装 upstream Body。
- [ ] 默认忽略转发头,仅按可信代理链提取客户端 IP。
### Task 5: WebSocket 元数据与统一 shutdown
**文件:** `proxy/writer.go``proxy/proxy.go``tests/proxy/websocket_test.go``main.go`
- [ ] Hijack 成功后登记连接,记录 101 握手元数据并在结束时注销。
- [ ] 提供 Handler shutdown,关闭受管升级连接。
- [ ] server 启动错误和 signal 共用清理路径,所有关闭步骤受单一总预算约束。
### Task 6: Compose、文档与最终验证
**文件:** `compose.yml``.env.example``README.md``docs/compose/reports/*`
- [ ] 固定镜像、取消默认 ClickHouse 端口发布、强制显式密码/DSN并同步新配置。
- [ ] 记录 best-effort、模糊提交、schema 和 WebSocket 取舍。
- [ ] 运行 `gofmt``go test ./...``go vet ./...` 并生成报告。
@@ -0,0 +1,5 @@
# Reliability And Security Fixes Plan
The canonical dated plan is [2026-07-09-clickhouse-migration.md](2026-07-09-clickhouse-migration.md).
This stable path is the review entry point for the reliability and security fixes.
@@ -0,0 +1,32 @@
# TokenThief 可靠性与安全修复验证报告
## 结果
本次修订覆盖队列并发关闭和单一总 shutdown deadline、响应完整性、请求体 fail-closed、ClickHouse 部分/模糊提交、64 MiB 默认队列字节预算、严格正值配置、响应体总/idle timeout、SSE 协议边界、WebSocket 101 与 shutdown、DSN/TLS、schema 校验、Compose 安全默认、固定 502、可信代理、batch 清理和 server 启动统一清理。
Compose 使用固定 ClickHouse 镜像 `clickhouse/clickhouse-server:25.3.3.42-alpine`,默认不发布 ClickHouse 端口,并在配置展开阶段拒绝空的 `UPSTREAM_URL``CLICKHOUSE_URL``CLICKHOUSE_PASSWORD``.env.example` 不提供密码或 DSN 默认值。
## 关键取舍
- 日志仍是内存中、非阻塞、best-effort。队列满、字节预算不足或数据库不可用时允许丢弃并计数。
- ClickHouse `PrepareBatch` 的确定失败保留整批;逐项 `Append` 的确定失败只保留失败项,成功项照常发送。`Send` 错误无法从单机客户端确认服务端是否已提交,因此不自动重试该批并计入 `ambiguous_send`,优先避免静默重复。跨进程 exactly-once 需要持久化 outbox 和服务端幂等协议,未在本次引入。
- WebSocket 仅记录 101 握手元数据并关闭受管连接,不解析帧或承诺跨进程迁移。
- schema 会自动创建缺失表,并校验现有表的列类型、MergeTree 引擎、分区键和排序键;不兼容时拒绝标记健康,不执行有数据风险的自动重建。
- Compose 不默认发布 ClickHouse 端口,强制显式密码和独立 DSN,以同时保留原始服务端密码和 URL 编码凭据。
## 行为边界
- 仅当请求体完整可重放且 `ReverseProxy` 正常结束时提交 HTTP 日志;读取失败的请求不转发,响应中断不提交不完整日志。
- SSE 仅由 `Content-Type: text/event-stream` 判定;看到 `[DONE]` 不会提前提交,仍等待上游正常结束。普通响应使用总 timeout,SSE 使用可重置 idle timeout。
- WebSocket 仅记录成功 `101` 的握手元数据,不采集帧;进程关闭会关闭当前进程管理的升级连接,但不提供跨进程连接迁移或协调。
- `502` 对客户端使用固定错误文本,内部连接错误只进入服务端日志。转发客户端 IP 仅在 TCP 对端属于 `TRUSTED_PROXIES` 时生效。
## 验证
- `gofmt -w .`:通过。
- `go test ./...`:通过。
- `go vet ./...`:通过。
- 新增 `tests/deployment` 契约测试,覆盖必填部署参数、固定 ClickHouse 镜像、不发布数据库端口和示例凭据留空。
- `docker compose config`:未执行,当前环境没有可用 Docker daemon/CLICompose 安全约束由 Go 契约测试覆盖。
- `go test -race`:不属于验收命令,未执行。
- 未连接真实 ClickHouse 做断链模糊提交和旧 schema 集成测试;相关阶段语义通过 fake batch 与纯 schema 校验单测覆盖。
@@ -0,0 +1,55 @@
---
feature: reliability-security-fixes
status: delivered
specs:
- docs/compose/specs/reliability-security-fixes.md
- docs/compose/specs/2026-07-09-clickhouse-migration.md
plans:
- docs/compose/plans/reliability-security-fixes.md
- docs/compose/plans/2026-07-09-clickhouse-migration.md
branch: main
---
# 可靠性与安全修复 - 最终报告
## What Was Built
本轮完成了 token_thief 的可靠性与安全加固。异步日志队列现在安全处理并发 `Submit`/`Stop`、使用条目数和字节双预算、在统一 shutdown deadline 内排空,并区分 ClickHouse 的可安全重试、确定失败和提交结果不明三类写入结果。
反向代理现在对请求体读取、上游响应复制、普通响应超时、SSE idle timeout、WebSocket 101 元数据和升级连接关闭实施完整性保护。客户端错误文本不再泄露内部信息,转发来源头仅在直接对端属于 `TRUSTED_PROXIES` 时参与客户端 IP 判定。
配置、ClickHouse DSN/TLS、数据库 schema 和 Compose 部署均采用 fail-closed 校验与更安全默认;服务启动错误和信号关闭共用资源清理路径。
## Architecture
`config/config.go` 在启动前严格解析正值预算、duration、布尔值、可信代理和 DSN。`proxy/` 在转发前完整读取请求体,通过 capture writer 和 response body wrapper 判断响应是否完整,并管理 hijacked 连接。`logger/queue.go` 提供非阻塞有界队列,按 `PrepareBatch``Append``Send` 阶段决定重试或丢弃;`db/` 只在连接、迁移和 schema 校验均成功后标记健康。`main.go` 用一个 30 秒 shutdown context 依次关闭 HTTP、升级连接、队列和数据库。
### Design Decisions
- 选择仅重试 `PrepareBatch` 失败,因为此时可确定没有提交;`Send` 失败计为 ambiguous 且不重放,避免静默重复。
- 选择进程内有界 best-effort 队列,因为当前项目没有持久化 outbox;该策略明确不承诺跨进程 exactly-once。
- 选择 WebSocket 仅记录 101 握手元数据并管理连接,不采集帧,以保持代理边界和关闭行为可测试。
- 选择拒绝不兼容 ClickHouse schema,而不是自动破坏性迁移或重建表。
## Usage
必须设置 `UPSTREAM_URL``CLICKHOUSE_URL`,Compose 部署还必须显式设置强 `CLICKHOUSE_PASSWORD`。关键新增配置包括 `LOG_QUEUE_BYTES``UPSTREAM_RESPONSE_TIMEOUT``UPSTREAM_STREAM_IDLE_TIMEOUT``TRUSTED_PROXIES`;所有预算和 timeout 必须大于零。`CLICKHOUSE_URL` 支持严格的 `host:port``clickhouse://``clickhouses://`TLS 与压缩参数受白名单校验。
## Verification
迭代 1 验证全部通过:`gofmt -l .` 无输出,`go test -json ./...` 共 118 个测试通过、0 失败、5 个测试包通过,`go vet ./...` 无诊断,`go build ./` 成功。针对 config、queue、ClickHouse、HTTP/SSE/WebSocket 和 Compose 的失败路径均有测试覆盖。
## Journey Log
> Brief notes on what informed the final design. Not required reading.
- [lesson] 迭代 1:网络写入的 `Send` 错误无法证明提交与否;最小安全策略是不自动重试、显式统计 ambiguous,并将连接标记为不健康。
- [lesson] 迭代 1:进程内队列只能提供有界 best-effort;跨进程幂等需要持久化 outbox 和下游幂等协议,不能由本地重试可靠模拟。
- [pivot] 迭代 1:SSE 完成判定限定为真正的 `text/event-stream` 且等待代理正常返回,避免终止标记导致提前记录不完整响应。
## Source Materials
| File | Role | Notes |
|------|------|-------|
| `docs/compose/specs/2026-07-09-clickhouse-migration.md` | 安全修复规格 | 定义本轮行为边界与取舍 |
| `docs/compose/plans/2026-07-09-clickhouse-migration.md` | 实施计划 | 覆盖跨模块修复和验证 |
@@ -0,0 +1,54 @@
# ClickHouse Migration 安全修复规格
> [!NOTE]
> This document may not reflect the current implementation.
> See the final report for up-to-date state:
> [Final Report](../reports/reliability-security-fixes.md)
## 修订范围
本规格是现有 ClickHouse Migration 的增量修订。保留当前单表 `proxy_logs`、内存异步队列、`httputil.ReverseProxy` 和 best-effort 审计模型,不引入 PostgreSQL 兼容层、持久化 outbox 或 WebSocket 帧采集。
## 行为要求
### 队列与关闭
- `Submit` 始终非阻塞;与 `Stop` 并发、停止后提交和重复停止均不得 panic。
- 队列由条数和字节双重预算约束。默认 `LOG_QUEUE_SIZE=256``LOG_QUEUE_BYTES=67108864``MAX_BODY_BYTES=1048576`;任一预算不足即丢弃并计数。
- `Stop(ctx)` 使用调用方提供的单一总 deadline 排空;deadline 到期取消所有 worker I/O、丢弃剩余条目并返回错误。
- worker flush 后清空 batch 指针并释放条目的字节预留。
### ClickHouse 写入
- `PrepareBatch` 错误是确定未提交错误,可有限重试整个 batch。
- `Append` 错误是确定未提交错误,必须中止/关闭 batch,并允许逐条隔离坏记录;成功记录继续写入,失败记录明确计数。
- `Send` 错误视为提交结果不明,不自动重试或逐条重发,避免静默重复;该批计为 ambiguous drop 并标记连接不健康。
- 进程内 best-effort 方案不承诺跨进程 exactly-once。需要更强保证时必须另行引入持久化 outbox 与幂等协议。
### 代理完整性
- 请求体预读失败时返回固定 400,不调用 upstream,不转发已损坏请求。
- 502 客户端响应只包含固定 `bad gateway` 和 request ID;内部网络错误仅写服务端日志和 `LogEntry.Error`
- 只有 ReverseProxy 正常完成的普通/SSE 响应才提交完整日志。下游写失败、短写、客户端断开或响应复制中断不得提交成完整成功日志。
- SSE 终止检测只在响应 `Content-Type``text/event-stream` 时按完整 event 边界解析。终止事件仅用于状态判断,不触发提前提交;提交仍等待代理正常返回。
- 普通 HTTP 响应体使用可配置总 timeoutSSE 使用可配置 idle timeout,成功读取数据后重置;timeout 必须能关闭阻塞中的上游 Body。
- WebSocket 成功升级时记录一条 101 握手元数据日志,body 为空。Handler 跟踪已 hijack 连接,并提供受 context 限制的 shutdown 关闭能力;不采集帧。
- 默认不信任 `X-Forwarded-For`/`X-Real-IP`。仅当直接对端命中 `TRUSTED_PROXIES` CIDR 时,从 XFF 右向左剥离可信代理并选择最近的不可信地址。
### 配置与数据库边界
- 所有整数、布尔值和 duration 环境变量格式错误时启动失败;要求正值的 body、队列、batch、worker、重连和 timeout 配置必须严格大于零。
- `CLICKHOUSE_URL` 支持严格 `host:port``clickhouse://``clickhouses://`。仅允许 `secure``skip_verify``compress` 查询参数;未知参数、冲突 TLS 配置、fragment、空 host/port 均失败。
- 建表后校验 `system.columns``system.tables` 的必需列类型、引擎、分区键和排序键;不兼容 schema 拒绝标记健康,不自动重建或改类型。
### 启动与部署
- server 启动错误与 signal 进入同一清理路径;禁止 goroutine 内 `log.Fatalf` 绕过 defer。
- shutdown 总预算依次覆盖 HTTP、升级连接、日志队列、DB 和根 context。
- Compose 固定 ClickHouse 明确版本,不使用 `latest`;默认不发布 ClickHouse 端口;`CLICKHOUSE_PASSWORD` 必须显式设置,应用 DSN 使用独立 `CLICKHOUSE_URL`,避免弱默认和 URL 编码冲突。
## 验证
- 针对上述失败路径增加 config、logger、proxy、db 单元测试。
- 运行 `gofmt``go test ./...``go vet ./...`
- 使用 Git 提交保留完整变更证据;审查应以提交及其 diff 验证实现和验收项。
@@ -0,0 +1,5 @@
# Reliability And Security Fixes Specification
The canonical dated specification is [2026-07-09-clickhouse-migration.md](2026-07-09-clickhouse-migration.md).
This stable path is the review entry point for the reliability and security fixes.
+25
View File
@@ -0,0 +1,25 @@
# TokenThief 路径过滤配置(聊天与补全接口)
#
# 仅记录聊天(Chat)和补全(Completions)相关请求。
#
# mode:
# whitelist - 仅记录命中 patterns 的请求
# blacklist - 命中 patterns 的请求不记录
# disabled - 全部记录(忽略 patterns
#
# Pattern 语法(glob):
# * 匹配单个段内除 / 之外的任意字符
# ** 跨段匹配任意字符(包括 /)
# ? 匹配单个非 / 字符
mode: whitelist
patterns:
# ===== 聊天(Chat =====
- /v1/chat/completions # 原生 OpenAI ChatCompletions(也是 Gemini 图片生成的 OpenAI 入口)
- /v1/responses # 原生 OpenAI Responses 格式
- /v1/messages # 原生 Claude (Anthropic) Messages
- /v1beta/models/*:generateContent # 原生 Gemini generateContent(含媒体识别 / TTS / 图片生成)
- /v1beta/models/*:generateContent/ # 同上,Gemini 图片生成页面所示带尾斜杠形式
- /v1beta/models/*:streamGenerateContent # 原生 Gemini 流式 generateContent
# ===== 补全(Completions =====
- /v1/completions # 原生 OpenAI 文本补全
+29
View File
@@ -0,0 +1,29 @@
module git.misaka.ren/M1saka/token_thief
go 1.25.0
require (
github.com/ClickHouse/clickhouse-go/v2 v2.47.0
gopkg.in/yaml.v3 v3.0.1
)
require (
github.com/ClickHouse/ch-go v0.73.0 // indirect
github.com/andybalholm/brotli v1.2.1 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/go-faster/city v1.0.1 // indirect
github.com/go-faster/errors v0.7.1 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/klauspost/compress v1.18.6 // indirect
github.com/kr/pretty v0.3.0 // indirect
github.com/paulmach/orb v0.13.0 // indirect
github.com/pierrec/lz4/v4 v4.1.27 // indirect
github.com/rogpeppe/go-internal v1.14.1 // indirect
github.com/segmentio/asm v1.2.1 // indirect
github.com/shopspring/decimal v1.4.0 // indirect
go.opentelemetry.io/otel v1.44.0 // indirect
go.opentelemetry.io/otel/trace v1.44.0 // indirect
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/sys v0.46.0 // indirect
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect
)
+61
View File
@@ -0,0 +1,61 @@
github.com/ClickHouse/ch-go v0.73.0 h1:jsHiGRbQ3sz+gekvDFJF29LWDo5dzbJm5s1h8TWVP2M=
github.com/ClickHouse/ch-go v0.73.0/go.mod h1:wkFIxrqlXeRJ9cn3r5Fz5Qen9jl5aTMPuGZeuJpANNY=
github.com/ClickHouse/clickhouse-go/v2 v2.47.0 h1:ZDAzrnKSOPTIsm4tdUNfrii2yc8dk4SVRLC77BR7Z5Q=
github.com/ClickHouse/clickhouse-go/v2 v2.47.0/go.mod h1:sPj7C7UYQ2MWHcfX+4eGN6nwnCqwUKfgO6PcwKpd6K8=
github.com/andybalholm/brotli v1.2.1 h1:R+f5xP285VArJDRgowrfb9DqL18yVK0gKAW/F+eTWro=
github.com/andybalholm/brotli v1.2.1/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY=
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/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/go-faster/city v1.0.1 h1:4WAxSZ3V2Ws4QRDrscLEDcibJY8uf41H6AhXDrNDcGw=
github.com/go-faster/city v1.0.1/go.mod h1:jKcUJId49qdW3L1qKHH/3wPeUstCVpVSXTM6vO3VcTw=
github.com/go-faster/errors v0.7.1 h1:MkJTnDoEdi9pDabt1dpWf7AA8/BaSYZqibYyhZ20AYg=
github.com/go-faster/errors v0.7.1/go.mod h1:5ySTjWFiphBs07IKuiL69nxdfd5+fzh1u7FPGZP2quo=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao=
github.com/klauspost/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORNo=
github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI=
github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0=
github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk=
github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ=
github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/paulmach/orb v0.13.0 h1:r7n7mQGGF+cj/CbcivEj9J3HGK+XR+yXnvzRdq9saIw=
github.com/paulmach/orb v0.13.0/go.mod h1:6scRWINywA2Jf05dcjOfLfxrUIMECvTSG2MVbRLxu/k=
github.com/pierrec/lz4/v4 v4.1.27 h1:+PhzhWDrjRj89TH2sw43nE3+4+W8lSxIuQadEHZyjUk=
github.com/pierrec/lz4/v4 v4.1.27/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/rogpeppe/go-internal v1.6.1/go.mod h1:xXDCJY+GAPziupqXw64V24skbSoqbTEfhy4qGm1nDQc=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc=
github.com/segmentio/asm v1.2.1 h1:DTNbBqs57ioxAD4PrArqftgypG4/qNpXoJx8TVXxPR0=
github.com/segmentio/asm v1.2.1/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs=
github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k=
github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU=
github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E=
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
+24
View File
@@ -0,0 +1,24 @@
package logger
import "time"
// LogEntry 表示一次代理请求的完整记录。
type LogEntry struct {
RequestID string
Method string
Path string
Query string
ClientIP string
RequestHeaders []byte // JSON
RequestBody []byte
RequestTruncated bool
StatusCode int
ResponseHeaders []byte // JSON
ResponseBody []byte
ResponseTruncated bool
IsStream bool
LatencyMS int64
StartedAt time.Time
FinishedAt time.Time
Error string
}
+447
View File
@@ -0,0 +1,447 @@
package logger
import (
"context"
"errors"
"log"
"math/rand"
"sync"
"sync/atomic"
"time"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
"git.misaka.ren/M1saka/token_thief/db"
)
const (
maxAttempts = 3
BaseBackoff = 200 * time.Millisecond
)
type Batch interface {
Append(v ...any) error
Send() error
Abort() error
}
type BatchPreparer interface {
PrepareBatch(ctx context.Context, query string) (Batch, error)
}
type clickHouseBatchPreparer struct {
conn driver.Conn
}
// The explicit adapter keeps the local Batch contract aligned with the real driver.
func (p clickHouseBatchPreparer) PrepareBatch(ctx context.Context, query string) (Batch, error) {
return p.conn.PrepareBatch(ctx, query)
}
type Stats struct {
Enqueued uint64
Dropped uint64
Failed uint64
Ambiguous uint64
Bytes int64
}
// FlushResult separates retry-safe Prepare failures from final or ambiguous failures.
type FlushResult struct {
Retry []*LogEntry
Failed int
Ambiguous int
Err error
}
// Queue is an asynchronous, bounded logger. Submit never waits for database work.
type Queue struct {
ch chan *LogEntry
pool *db.Pool
batchSize int
batchInterval time.Duration
workers int
byteBudget int64
dropped atomic.Uint64
failed atomic.Uint64
ambiguous atomic.Uint64
enq atomic.Uint64
bytes atomic.Int64
mu sync.RWMutex
started bool
stopped bool
shutdownCtx context.Context
reporterCancel context.CancelFunc
workCancel context.CancelFunc
wg sync.WaitGroup
done chan struct{}
doneOnce sync.Once
}
// NewQueue accepts an optional byte budget. A non-positive or omitted budget disables byte limiting.
func NewQueue(pool *db.Pool, queueSize, batchSize, workers int, batchInterval time.Duration, byteBudget ...int64) *Queue {
var budget int64
if len(byteBudget) > 0 {
budget = byteBudget[0]
}
if queueSize < 0 {
queueSize = 0
}
if batchSize < 1 {
batchSize = 1
}
if workers < 1 {
workers = 1
}
if batchInterval <= 0 {
batchInterval = time.Second
}
return &Queue{
ch: make(chan *LogEntry, queueSize),
pool: pool,
batchSize: batchSize,
batchInterval: batchInterval,
workers: workers,
byteBudget: budget,
done: make(chan struct{}),
}
}
// EstimatedBytes covers all variable-size strings and byte slices retained by an entry.
func EstimatedBytes(e *LogEntry) int64 {
if e == nil {
return 0
}
return int64(len(e.RequestID) + len(e.Method) + len(e.Path) + len(e.Query) + len(e.ClientIP) + len(e.Error) +
len(e.RequestHeaders) + len(e.RequestBody) + len(e.ResponseHeaders) + len(e.ResponseBody))
}
func (q *Queue) Stats() Stats {
return Stats{
Enqueued: q.enq.Load(),
Dropped: q.dropped.Load(),
Failed: q.failed.Load(),
Ambiguous: q.ambiguous.Load(),
Bytes: q.bytes.Load(),
}
}
// Start starts workers once. The root context controls reporting only; Stop owns worker shutdown.
func (q *Queue) Start(root context.Context) {
q.mu.Lock()
if q.started || q.stopped {
q.mu.Unlock()
return
}
q.started = true
workCtx, workCancel := context.WithCancel(context.Background())
reportCtx, reporterCancel := context.WithCancel(root)
q.workCancel = workCancel
q.reporterCancel = reporterCancel
for i := 0; i < q.workers; i++ {
q.wg.Add(1)
go q.run(workCtx)
}
q.wg.Add(1)
go q.reportLoop(reportCtx)
q.mu.Unlock()
}
// Stop closes submissions once and waits under the caller's single total deadline.
// The variadic form permits legacy Stop() calls while new callers should pass a context.
func (q *Queue) Stop(contexts ...context.Context) error {
ctx := context.Background()
if len(contexts) > 0 && contexts[0] != nil {
ctx = contexts[0]
}
q.mu.Lock()
if !q.stopped {
q.stopped = true
q.shutdownCtx = ctx
close(q.ch)
if q.reporterCancel != nil {
q.reporterCancel()
}
if q.started {
go func() {
q.wg.Wait()
q.doneOnce.Do(func() { close(q.done) })
}()
} else {
q.discardQueued()
q.doneOnce.Do(func() { close(q.done) })
}
}
done := q.done
workCancel := q.workCancel
q.mu.Unlock()
select {
case <-done:
return nil
case <-ctx.Done():
if workCancel != nil {
workCancel()
}
return ctx.Err()
}
}
// Submit reserves memory and enqueues without waiting; full, over-budget, and stopped queues drop.
func (q *Queue) Submit(e *LogEntry) {
if e == nil {
q.dropped.Add(1)
return
}
size := EstimatedBytes(e)
if !q.mu.TryRLock() {
q.dropped.Add(1)
return
}
if q.stopped || !q.reserve(size) {
q.mu.RUnlock()
q.dropped.Add(1)
return
}
select {
case q.ch <- e:
q.enq.Add(1)
default:
q.release(size)
q.dropped.Add(1)
}
q.mu.RUnlock()
}
func (q *Queue) reserve(size int64) bool {
if q.byteBudget <= 0 {
q.bytes.Add(size)
return true
}
for {
used := q.bytes.Load()
if size > q.byteBudget-used {
return false
}
if q.bytes.CompareAndSwap(used, used+size) {
return true
}
}
}
func (q *Queue) release(size int64) {
q.bytes.Add(-size)
}
func (q *Queue) discardQueued() {
for e := range q.ch {
q.release(EstimatedBytes(e))
q.failed.Add(1)
}
}
func (q *Queue) run(workCtx context.Context) {
defer q.wg.Done()
batch := make([]*LogEntry, 0, q.batchSize)
ticker := time.NewTicker(q.batchInterval)
defer ticker.Stop()
flush := func() {
if len(batch) == 0 {
return
}
q.flush(q.flushContext(workCtx), batch)
for _, e := range batch {
q.release(EstimatedBytes(e))
}
clear(batch)
batch = batch[:0]
}
for {
select {
case e, ok := <-q.ch:
if !ok {
flush()
return
}
batch = append(batch, e)
if len(batch) >= q.batchSize {
flush()
}
case <-ticker.C:
flush()
case <-workCtx.Done():
flush()
q.discardQueued()
return
}
}
}
func (q *Queue) flushContext(workCtx context.Context) context.Context {
q.mu.RLock()
defer q.mu.RUnlock()
if q.shutdownCtx != nil {
return q.shutdownCtx
}
return workCtx
}
func (q *Queue) flush(ctx context.Context, entries []*LogEntry) {
if q.pool == nil || !q.pool.Healthy() {
q.failed.Add(uint64(len(entries)))
return
}
retry := entries
var lastErr error
for attempt := 1; attempt <= maxAttempts && len(retry) > 0; attempt++ {
if err := ctx.Err(); err != nil {
lastErr = err
break
}
conn := q.pool.Get()
if conn == nil {
lastErr = errors.New("pool nil")
break
}
result := Flush(ctx, clickHouseBatchPreparer{conn: conn}, retry)
q.failed.Add(uint64(result.Failed))
q.ambiguous.Add(uint64(result.Ambiguous))
if result.Ambiguous > 0 {
q.pool.MarkUnhealthy()
}
lastErr = result.Err
retry = result.Retry
if len(retry) == 0 {
return
}
if attempt < maxAttempts && !waitBackoff(ctx, attempt) {
lastErr = ctx.Err()
break
}
}
if len(retry) > 0 {
q.failed.Add(uint64(len(retry)))
q.pool.MarkUnhealthy()
log.Printf("[logger] giving up %d retry-safe rows after %d attempts: %v", len(retry), maxAttempts, lastErr)
}
}
func waitBackoff(ctx context.Context, attempt int) bool {
wait := BaseBackoff
for i := 1; i < attempt; i++ {
wait *= 3
}
jitter := time.Duration((rand.Float64()*0.4 - 0.2) * float64(wait))
timer := time.NewTimer(wait + jitter)
defer timer.Stop()
select {
case <-timer.C:
return true
case <-ctx.Done():
return false
}
}
const insertStatement = `INSERT INTO proxy_logs (
request_id, method, path, query, client_ip,
request_headers, request_body, request_truncated,
status_code, response_headers, response_body, response_truncated,
is_stream, latency_ms, started_at, finished_at, error
)`
// Flush executes one batch. Only Prepare failures are retryable. Append failures are isolated by row;
// Send failures are ambiguous and therefore never replayed.
func Flush(ctx context.Context, conn BatchPreparer, entries []*LogEntry) FlushResult {
if len(entries) == 0 {
return FlushResult{}
}
result := flushOnce(ctx, conn, entries)
if result.stage != flushAppend || len(entries) == 1 {
return result.public(entries)
}
combined := FlushResult{}
for _, entry := range entries {
single := flushOnce(ctx, conn, []*LogEntry{entry}).public([]*LogEntry{entry})
combined.Retry = append(combined.Retry, single.Retry...)
combined.Failed += single.Failed
combined.Ambiguous += single.Ambiguous
combined.Err = errors.Join(combined.Err, single.Err)
}
return combined
}
type flushStage uint8
const (
flushSuccess flushStage = iota
flushPrepare
flushAppend
flushSend
)
type flushAttempt struct {
stage flushStage
err error
}
func (r flushAttempt) public(entries []*LogEntry) FlushResult {
switch r.stage {
case flushSuccess:
return FlushResult{}
case flushPrepare:
return FlushResult{Retry: entries, Err: r.err}
case flushSend:
return FlushResult{Failed: len(entries), Ambiguous: len(entries), Err: r.err}
default:
return FlushResult{Failed: len(entries), Err: r.err}
}
}
func flushOnce(ctx context.Context, conn BatchPreparer, entries []*LogEntry) flushAttempt {
batch, err := conn.PrepareBatch(ctx, insertStatement)
if err != nil {
return flushAttempt{stage: flushPrepare, err: err}
}
abort := func(cause error) error {
return errors.Join(cause, batch.Abort())
}
for _, e := range entries {
if err := batch.Append(
e.RequestID, e.Method, e.Path, e.Query, e.ClientIP,
string(e.RequestHeaders), string(e.RequestBody), e.RequestTruncated,
int32(e.StatusCode), string(e.ResponseHeaders), string(e.ResponseBody), e.ResponseTruncated,
e.IsStream, e.LatencyMS, e.StartedAt, e.FinishedAt, e.Error,
); err != nil {
return flushAttempt{stage: flushAppend, err: abort(err)}
}
}
if err := batch.Send(); err != nil {
return flushAttempt{stage: flushSend, err: abort(err)}
}
return flushAttempt{stage: flushSuccess}
}
func (q *Queue) reportLoop(ctx context.Context) {
defer q.wg.Done()
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
stats := q.Stats()
healthy := q.pool != nil && q.pool.Healthy()
log.Printf("[logger] metrics: enq=%d dropped_queue_full=%d dropped_db_fail=%d ambiguous_send=%d queue_len=%d queue_bytes=%d db_healthy=%v",
stats.Enqueued, stats.Dropped, stats.Failed, stats.Ambiguous, len(q.ch), stats.Bytes, healthy)
}
}
}
+109
View File
@@ -0,0 +1,109 @@
package main
import (
"context"
"errors"
"fmt"
"log"
"net/http"
"os"
"os/signal"
"syscall"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/db"
"git.misaka.ren/M1saka/token_thief/logger"
"git.misaka.ren/M1saka/token_thief/proxy"
)
func main() {
log.SetFlags(log.LstdFlags | log.Lmicroseconds)
if err := run(); err != nil {
log.Printf("[main] fatal: %v", err)
os.Exit(1)
}
}
func run() error {
cfg, err := config.Load()
if err != nil {
return fmt.Errorf("config: %w", err)
}
if _, err := db.ClickHouseOptions(cfg.ClickHouseURL); err != nil {
return err
}
filter, err := config.LoadFilter(cfg.FilterFile)
if err != nil {
return fmt.Errorf("load filter: %w", err)
}
log.Printf("[main] filter mode=%s patterns=%d", filter.Mode, len(filter.Patterns))
rootCtx, cancel := context.WithCancel(context.Background())
defer cancel()
pool := db.NewPool(rootCtx, cfg.ClickHouseURL, cfg.DBReconnectInterval)
queue := logger.NewQueue(pool, cfg.LogQueueSize, cfg.LogBatchSize, cfg.LogWorkers, cfg.LogBatchInterval, cfg.LogQueueBytes)
queue.Start(rootCtx)
h := proxy.NewWithOptions(cfg.UpstreamURL, filter, queue, cfg.MaxBodyBytes, proxy.Options{
UpstreamTimeout: cfg.UpstreamTimeout,
ResponseTimeout: cfg.UpstreamResponseTimeout,
SSEIdleTimeout: cfg.UpstreamStreamIdleTimeout,
UpstreamTLSInsecureSkipVerify: cfg.UpstreamTLSInsecureSkipVerify,
TrustedProxies: cfg.TrustedProxies,
})
srv := &http.Server{
Addr: cfg.ListenAddr,
Handler: h,
ReadHeaderTimeout: 30 * time.Second,
ReadTimeout: cfg.ReadTimeout,
WriteTimeout: cfg.WriteTimeout,
IdleTimeout: cfg.IdleTimeout,
}
serverErr := make(chan error, 1)
go func() {
log.Printf("[main] listening on %s, upstream=%s", cfg.ListenAddr, cfg.UpstreamURL)
serverErr <- srv.ListenAndServe()
}()
stop := make(chan os.Signal, 1)
signal.Notify(stop, syscall.SIGINT, syscall.SIGTERM)
var runErr error
select {
case <-stop:
log.Printf("[main] shutdown signal received")
case err := <-serverErr:
if !errors.Is(err, http.ErrServerClosed) {
runErr = fmt.Errorf("server: %w", err)
}
}
signal.Stop(stop)
// HTTP、升级连接和日志排空共享同一个关闭总预算。
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second)
defer shutdownCancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
log.Printf("[main] http shutdown: %v", err)
}
if err := h.Shutdown(shutdownCtx); err != nil {
log.Printf("[main] upgraded connection shutdown: %v", err)
}
queueStopped := true
if err := queue.Stop(shutdownCtx); err != nil {
log.Printf("[main] queue shutdown: %v", err)
queueStopped = false
}
cancel()
if queueStopped {
pool.Close()
} else {
log.Printf("[main] skip db close while queue workers are still exiting")
}
log.Printf("[main] bye")
return runErr
}
+104
View File
@@ -0,0 +1,104 @@
package proxy
import (
"bytes"
"crypto/rand"
"encoding/hex"
"encoding/json"
"io"
"net/http"
"net/netip"
"strings"
)
// readRequestBody 在转发前完整读取请求体,确保读取失败时不会向上游发送损坏请求。
func readRequestBody(r *http.Request, max int64) (captured []byte, truncated bool, err error) {
if r.Body == nil || r.ContentLength == 0 {
return nil, false, nil
}
body, err := io.ReadAll(r.Body)
if err != nil {
return nil, false, err
}
if err := r.Body.Close(); err != nil {
return nil, false, err
}
r.Body = io.NopCloser(bytes.NewReader(body))
if int64(len(body)) > max {
return body[:max], true, nil
}
return body, false, nil
}
func headersJSON(h http.Header) []byte {
if len(h) == 0 {
return nil
}
b, err := json.Marshal(h)
if err != nil {
return nil
}
return b
}
func newRequestID() string {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "unknown"
}
return hex.EncodeToString(b[:])
}
// clientIP 从请求中提取客户端 IP。
func clientIP(r *http.Request, trusted []netip.Prefix) string {
peer, ok := parsePeerAddr(r.RemoteAddr)
if !ok {
return r.RemoteAddr
}
if !isTrusted(peer, trusted) {
return peer.String()
}
xff := strings.Split(r.Header.Get("X-Forwarded-For"), ",")
if len(xff) == 1 && strings.TrimSpace(xff[0]) == "" {
if realIP, err := netip.ParseAddr(strings.TrimSpace(r.Header.Get("X-Real-IP"))); err == nil {
return realIP.Unmap().String()
}
return peer.String()
}
chain := make([]netip.Addr, len(xff))
for i, raw := range xff {
addr, err := netip.ParseAddr(strings.TrimSpace(raw))
if err != nil {
return peer.String()
}
chain[i] = addr.Unmap()
}
client := peer
for i := len(chain) - 1; i >= 0 && isTrusted(client, trusted); i-- {
client = chain[i]
}
return client.String()
}
func parsePeerAddr(remote string) (netip.Addr, bool) {
if addrPort, err := netip.ParseAddrPort(remote); err == nil {
return addrPort.Addr().Unmap(), true
}
addr, err := netip.ParseAddr(remote)
return addr.Unmap(), err == nil
}
func isTrusted(addr netip.Addr, prefixes []netip.Prefix) bool {
for _, prefix := range prefixes {
if prefix.Contains(addr) {
return true
}
}
return false
}
// isStreamResponse 通过响应头判断是否为流式响应。
func isStreamResponse(h http.Header) bool {
ct := strings.ToLower(strings.TrimSpace(strings.SplitN(h.Get("Content-Type"), ";", 2)[0]))
return ct == "text/event-stream"
}
+339
View File
@@ -0,0 +1,339 @@
package proxy
import (
"context"
"crypto/tls"
"log"
"net"
"net/http"
"net/http/httputil"
"net/netip"
"net/url"
"strings"
"sync"
"sync/atomic"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/logger"
)
// LogSubmitter 是 proxy 唯一依赖的日志接收方接口。
type LogSubmitter interface {
Submit(*logger.LogEntry)
}
// Handler 构造反代 HTTP handler。
type Handler struct {
rp *httputil.ReverseProxy
filter *config.Filter
queue LogSubmitter
maxBodyBytes int64
trusted []netip.Prefix
connMu sync.Mutex
conns map[net.Conn]struct{}
connChanged chan struct{}
}
// Options 控制反代连接上游时的网络行为。
type Options struct {
UpstreamTimeout time.Duration
ResponseTimeout time.Duration
SSEIdleTimeout time.Duration
UpstreamTLSInsecureSkipVerify bool
TrustedProxies []netip.Prefix
}
// requestState 通过 context 在 ErrorHandler / ModifyResponse / 主 handler 之间共享状态。
type requestState struct {
requestID string
lastErr atomic.Pointer[string]
upgrade atomic.Pointer[upgradeResponse]
readFailed atomic.Bool
}
type upgradeResponse struct {
status int
header http.Header
}
type ctxKey struct{}
func newRequestState(id string) *requestState { return &requestState{requestID: id} }
func stateFromCtx(ctx context.Context) *requestState {
v, _ := ctx.Value(ctxKey{}).(*requestState)
return v
}
func New(upstream *url.URL, filter *config.Filter, queue LogSubmitter, maxBody int64, upstreamTimeout ...time.Duration) *Handler {
opts := Options{}
if len(upstreamTimeout) > 0 {
opts.UpstreamTimeout = upstreamTimeout[0]
}
return NewWithOptions(upstream, filter, queue, maxBody, opts)
}
func NewWithOptions(upstream *url.URL, filter *config.Filter, queue LogSubmitter, maxBody int64, opts Options) *Handler {
rp := httputil.NewSingleHostReverseProxy(upstream)
rp.FlushInterval = -1 // 让流式 chunk 立即转发
if opts.UpstreamTimeout > 0 || opts.UpstreamTLSInsecureSkipVerify {
transport := http.DefaultTransport.(*http.Transport).Clone()
if opts.UpstreamTimeout > 0 {
transport.DialContext = (&net.Dialer{Timeout: opts.UpstreamTimeout, KeepAlive: 30 * time.Second}).DialContext
transport.ResponseHeaderTimeout = opts.UpstreamTimeout
transport.TLSHandshakeTimeout = opts.UpstreamTimeout
}
if opts.UpstreamTLSInsecureSkipVerify {
transport.TLSClientConfig = &tls.Config{InsecureSkipVerify: true}
}
rp.Transport = transport
}
origDirector := rp.Director
rp.Director = func(r *http.Request) {
origDirector(r)
r.Host = upstream.Host
}
// ModifyResponse 在响应头写回客户端之前调用,确保 X-Request-Id 一定生效。
rp.ModifyResponse = func(resp *http.Response) error {
if st := stateFromCtx(resp.Request.Context()); st != nil {
resp.Header.Set("X-Request-Id", st.requestID)
if resp.StatusCode == http.StatusSwitchingProtocols &&
strings.EqualFold(resp.Request.Header.Get("Upgrade"), "websocket") &&
strings.EqualFold(resp.Header.Get("Upgrade"), "websocket") {
st.upgrade.Store(&upgradeResponse{status: resp.StatusCode, header: resp.Header.Clone()})
}
}
if timeout := responseBodyTimeout(resp, opts); timeout > 0 {
resp.Body = newTimeoutBody(resp.Body, timeout, isStreamResponse(resp.Header))
}
if st := stateFromCtx(resp.Request.Context()); st != nil && resp.StatusCode != http.StatusSwitchingProtocols {
resp.Body = &trackingBody{ReadCloser: resp.Body, failed: &st.readFailed}
}
return nil
}
rp.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) {
log.Printf("[proxy] upstream error for %s %s: %v", r.Method, r.URL.Path, err)
// 把错误暴露给主 handler,使其能写入日志。
if st := stateFromCtx(r.Context()); st != nil {
s := err.Error()
st.lastErr.Store(&s)
// ErrorHandler 路径下 ModifyResponse 不会被调用,这里手动写 X-Request-Id。
w.Header().Set("X-Request-Id", st.requestID)
}
http.Error(w, "bad gateway", http.StatusBadGateway)
}
return &Handler{
rp: rp,
filter: filter,
queue: queue,
maxBodyBytes: maxBody,
trusted: append([]netip.Prefix(nil), opts.TrustedProxies...),
conns: make(map[net.Conn]struct{}),
connChanged: make(chan struct{}),
}
}
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// 健康检查不参与反代与日志。
if r.URL.Path == "/healthz" {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("ok"))
return
}
shouldLog := h.filter.ShouldLog(r.URL.Path)
reqBody, reqTruncated, err := readRequestBody(r, h.maxBodyBytes)
if err != nil && !shouldLog {
log.Printf("[proxy] read request body failed: %v", err)
_ = r.Body.Close()
http.Error(w, "bad request", http.StatusBadRequest)
return
}
if !shouldLog {
log.Printf("[proxy] skip log by filter method=%s path=%s", r.Method, r.URL.Path)
cw := newCaptureWriter(w, 0)
cw.OnHijack(h.trackConn)
h.rp.ServeHTTP(cw, r)
return
}
started := time.Now()
reqID := newRequestID()
log.Printf("[proxy] capture start request_id=%s method=%s path=%s", reqID, r.Method, r.URL.Path)
st := newRequestState(reqID)
r = r.WithContext(context.WithValue(r.Context(), ctxKey{}, st))
cw := newCaptureWriter(w, h.maxBodyBytes)
if err != nil {
log.Printf("[proxy] read request body failed: %v", err)
_ = r.Body.Close()
s := "read request body: " + err.Error()
st.lastErr.Store(&s)
h.serveRequestBodyError(cw, r, st, started, reqID)
return
}
reqHeadersJSON := headersJSON(r.Header)
clientAddr := clientIP(r, h.trusted)
method := r.Method
path := r.URL.Path
query := r.URL.RawQuery
var submitOnce sync.Once
submit := func(finished time.Time) {
entry, ok := h.buildLogEntry(cw, st, logEntryInput{
requestID: reqID,
method: method,
path: path,
query: query,
clientAddr: clientAddr,
requestHeaders: reqHeadersJSON,
requestBody: reqBody,
requestTruncated: reqTruncated,
started: started,
finished: finished,
})
if !ok {
return
}
h.queue.Submit(entry)
log.Printf("[proxy] capture finish request_id=%s method=%s path=%s status=%d is_stream=%v latency_ms=%d",
reqID, method, path, entry.StatusCode, entry.IsStream, entry.LatencyMS)
}
cw.OnHijack(func(conn net.Conn) net.Conn {
tracked := h.trackConn(conn)
if upgrade := st.upgrade.Load(); upgrade != nil {
cw.SetHijackedResponse(upgrade.status, upgrade.header)
submitOnce.Do(func() { submit(time.Now()) })
}
return tracked
})
h.rp.ServeHTTP(cw, r)
finished := time.Now()
responseComplete := !isStreamResponse(cw.Header()) || cw.SSEComplete()
if cw.Complete() && responseComplete && !st.readFailed.Load() && r.Context().Err() == nil {
submitOnce.Do(func() { submit(finished) })
}
}
func (h *Handler) serveRequestBodyError(cw *captureWriter, r *http.Request, st *requestState, started time.Time, requestID string) {
cw.Header().Set("X-Request-Id", requestID)
http.Error(cw, "bad request", http.StatusBadRequest)
if !cw.Complete() {
return
}
entry, ok := h.buildLogEntry(cw, st, logEntryInput{
requestID: requestID,
method: r.Method,
path: r.URL.Path,
query: r.URL.RawQuery,
clientAddr: clientIP(r, h.trusted),
requestHeaders: headersJSON(r.Header),
started: started,
finished: time.Now(),
})
if ok {
h.queue.Submit(entry)
}
}
// Shutdown closes all active hijacked connections and waits for their release.
func (h *Handler) Shutdown(ctx context.Context) error {
for {
h.connMu.Lock()
if len(h.conns) == 0 {
h.connMu.Unlock()
return nil
}
conns := make([]net.Conn, 0, len(h.conns))
for conn := range h.conns {
conns = append(conns, conn)
}
changed := h.connChanged
h.connMu.Unlock()
for _, conn := range conns {
_ = conn.Close()
}
select {
case <-changed:
case <-ctx.Done():
return ctx.Err()
}
}
}
func (h *Handler) trackConn(conn net.Conn) net.Conn {
tracked := &trackedConn{Conn: conn}
tracked.onClose = func() {
h.connMu.Lock()
delete(h.conns, tracked)
close(h.connChanged)
h.connChanged = make(chan struct{})
h.connMu.Unlock()
}
h.connMu.Lock()
h.conns[tracked] = struct{}{}
close(h.connChanged)
h.connChanged = make(chan struct{})
h.connMu.Unlock()
return tracked
}
type logEntryInput struct {
requestID string
method string
path string
query string
clientAddr string
requestHeaders []byte
requestBody []byte
requestTruncated bool
started time.Time
finished time.Time
}
func (h *Handler) buildLogEntry(cw *captureWriter, st *requestState, in logEntryInput) (*logger.LogEntry, bool) {
if cw.Hijacked() && cw.Status() != http.StatusSwitchingProtocols {
return nil, false
}
var errMsg string
if p := st.lastErr.Load(); p != nil {
errMsg = *p
}
isStream := isStreamResponse(cw.Header())
responseBody := append([]byte(nil), cw.Body()...)
responseTruncated := cw.Truncated()
if isStream && !responseTruncated {
if assembled, ok := assembleSSEJSON(responseBody); ok {
responseBody = assembled
}
}
return &logger.LogEntry{
RequestID: in.requestID,
Method: in.method,
Path: in.path,
Query: in.query,
ClientIP: in.clientAddr,
RequestHeaders: in.requestHeaders,
RequestBody: in.requestBody,
RequestTruncated: in.requestTruncated,
StatusCode: cw.Status(),
ResponseHeaders: headersJSON(cw.Header()),
ResponseBody: responseBody,
ResponseTruncated: responseTruncated,
IsStream: isStream,
LatencyMS: in.finished.Sub(in.started).Milliseconds(),
StartedAt: in.started,
FinishedAt: in.finished,
Error: errMsg,
}, true
}
+661
View File
@@ -0,0 +1,661 @@
package proxy
import (
"bufio"
"bytes"
"encoding/json"
"strings"
)
func assembleSSEJSON(body []byte) ([]byte, bool) {
payloads := sseDataPayloads(body)
if len(payloads) == 0 {
return nil, false
}
if assembled, ok := assembleOpenAICompletionsSSE(payloads); ok {
return assembled, true
}
if assembled, ok := assembleOpenAIChatSSE(payloads); ok {
return assembled, true
}
if assembled, ok := assembleOpenAIResponsesSSE(payloads); ok {
return assembled, true
}
if assembled, ok := assembleAnthropicSSE(payloads); ok {
return assembled, true
}
if assembled, ok := assembleGeminiSSE(payloads); ok {
return assembled, true
}
return nil, false
}
func sseDataPayloads(body []byte) []string {
payloads := make([]string, 0)
scanner := bufio.NewScanner(bytes.NewReader(body))
scanner.Buffer(make([]byte, 0, 64*1024), len(body)+1)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if !strings.HasPrefix(line, "data:") {
continue
}
payload := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
if payload == "" || payload == "[DONE]" {
continue
}
payloads = append(payloads, payload)
}
return payloads
}
type openAICompletionChunk struct {
ID string `json:"id,omitempty"`
Object string `json:"object,omitempty"`
Created int64 `json:"created,omitempty"`
Model string `json:"model,omitempty"`
Usage json.RawMessage `json:"usage"`
Choices []struct {
Index int `json:"index"`
Text *string `json:"text"`
FinishReason *string `json:"finish_reason"`
} `json:"choices"`
}
type openAICompletionChoice struct {
index int
text strings.Builder
finishReason string
}
func assembleOpenAICompletionsSSE(payloads []string) ([]byte, bool) {
choices := map[int]*openAICompletionChoice{}
order := make([]int, 0, 1)
var id, object, model string
var created int64
var usage json.RawMessage
matched := false
for _, payload := range payloads {
var chunk openAICompletionChunk
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
return nil, false
}
if len(chunk.Choices) == 0 && len(chunk.Usage) == 0 {
return nil, false
}
for _, choice := range chunk.Choices {
if choice.Text == nil {
return nil, false
}
}
matched = true
if id == "" {
id = chunk.ID
}
if object == "" {
object = chunk.Object
}
if created == 0 {
created = chunk.Created
}
if model == "" {
model = chunk.Model
}
if len(chunk.Usage) > 0 && string(chunk.Usage) != "null" {
usage = chunk.Usage
}
for _, choice := range chunk.Choices {
assembled := choices[choice.Index]
if assembled == nil {
assembled = &openAICompletionChoice{index: choice.Index}
choices[choice.Index] = assembled
order = append(order, choice.Index)
}
assembled.text.WriteString(*choice.Text)
if choice.FinishReason != nil {
assembled.finishReason = *choice.FinishReason
}
}
}
if !matched {
return nil, false
}
assembled := struct {
ID string `json:"id,omitempty"`
Object string `json:"object,omitempty"`
Created int64 `json:"created,omitempty"`
Model string `json:"model,omitempty"`
Usage json.RawMessage `json:"usage,omitempty"`
Choices []struct {
Index int `json:"index"`
Text string `json:"text"`
FinishReason string `json:"finish_reason,omitempty"`
} `json:"choices"`
}{ID: id, Object: object, Created: created, Model: model, Usage: usage}
for _, index := range order {
choice := choices[index]
assembled.Choices = append(assembled.Choices, struct {
Index int `json:"index"`
Text string `json:"text"`
FinishReason string `json:"finish_reason,omitempty"`
}{Index: choice.index, Text: choice.text.String(), FinishReason: choice.finishReason})
}
data, err := json.Marshal(assembled)
return data, err == nil
}
type openAIChatChunk struct {
ID string `json:"id,omitempty"`
Object string `json:"object,omitempty"`
Created int64 `json:"created,omitempty"`
Model string `json:"model,omitempty"`
SystemFingerprint string `json:"system_fingerprint,omitempty"`
Usage json.RawMessage `json:"usage"`
Choices []struct {
Index int `json:"index"`
Delta struct {
Role string `json:"role,omitempty"`
Content string `json:"content,omitempty"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCalls []struct {
Index int `json:"index"`
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"`
} `json:"function,omitempty"`
} `json:"tool_calls,omitempty"`
} `json:"delta"`
FinishReason *string `json:"finish_reason"`
NativeFinishReason *string `json:"native_finish_reason"`
} `json:"choices"`
}
type openAIChatChoice struct {
index int
role string
content strings.Builder
reasoning strings.Builder
toolCalls map[int]*openAIToolCall
toolOrder []int
finishReason string
nativeFinish string
}
type openAIToolCall struct {
id string
callType string
name string
arguments strings.Builder
}
type openAIChatResponse struct {
ID string `json:"id,omitempty"`
Object string `json:"object,omitempty"`
Created int64 `json:"created,omitempty"`
Model string `json:"model,omitempty"`
SystemFingerprint string `json:"system_fingerprint,omitempty"`
Usage json.RawMessage `json:"usage,omitempty"`
Choices []openAIResponseChoice `json:"choices"`
}
type openAIResponseChoice struct {
Index int `json:"index"`
Message openAIResponseMessage `json:"message"`
FinishReason string `json:"finish_reason,omitempty"`
NativeFinishReason string `json:"native_finish_reason,omitempty"`
}
type openAIResponseMessage struct {
Role string `json:"role,omitempty"`
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCalls []openAIResponseTool `json:"tool_calls,omitempty"`
}
type openAIResponseTool struct {
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function openAIResponseToolFunction `json:"function"`
}
type openAIResponseToolFunction struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments"`
}
func assembleOpenAIChatSSE(payloads []string) ([]byte, bool) {
choices := map[int]*openAIChatChoice{}
order := make([]int, 0, 1)
var id, object string
var created int64
var model string
var systemFingerprint string
var usage json.RawMessage
matched := false
for _, payload := range payloads {
var chunk openAIChatChunk
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
return nil, false
}
if len(chunk.Choices) == 0 && len(chunk.Usage) == 0 {
return nil, false
}
matched = true
if id == "" {
id = chunk.ID
}
if object == "" {
object = strings.TrimSuffix(chunk.Object, ".chunk")
}
if created == 0 {
created = chunk.Created
}
if model == "" {
model = chunk.Model
}
if systemFingerprint == "" {
systemFingerprint = chunk.SystemFingerprint
}
if len(chunk.Usage) > 0 && string(chunk.Usage) != "null" {
usage = chunk.Usage
}
for _, choice := range chunk.Choices {
assembled := choices[choice.Index]
if assembled == nil {
assembled = &openAIChatChoice{index: choice.Index}
choices[choice.Index] = assembled
order = append(order, choice.Index)
}
if choice.Delta.Role != "" {
assembled.role = choice.Delta.Role
}
if choice.Delta.Content != "" {
assembled.content.WriteString(choice.Delta.Content)
}
if choice.Delta.ReasoningContent != "" {
assembled.reasoning.WriteString(choice.Delta.ReasoningContent)
}
for _, toolCall := range choice.Delta.ToolCalls {
if assembled.toolCalls == nil {
assembled.toolCalls = map[int]*openAIToolCall{}
}
assembledTool := assembled.toolCalls[toolCall.Index]
if assembledTool == nil {
assembledTool = &openAIToolCall{}
assembled.toolCalls[toolCall.Index] = assembledTool
assembled.toolOrder = append(assembled.toolOrder, toolCall.Index)
}
if toolCall.ID != "" {
assembledTool.id = toolCall.ID
}
if toolCall.Type != "" {
assembledTool.callType = toolCall.Type
}
if toolCall.Function.Name != "" {
assembledTool.name = toolCall.Function.Name
}
if toolCall.Function.Arguments != "" {
assembledTool.arguments.WriteString(toolCall.Function.Arguments)
}
}
if choice.FinishReason != nil {
assembled.finishReason = *choice.FinishReason
}
if choice.NativeFinishReason != nil {
assembled.nativeFinish = *choice.NativeFinishReason
}
}
}
if !matched {
return nil, false
}
assembled := openAIChatResponse{ID: id, Object: object, Created: created, Model: model, SystemFingerprint: systemFingerprint, Usage: usage}
for _, index := range order {
choice := choices[index]
out := openAIResponseChoice{Index: choice.index, FinishReason: choice.finishReason, NativeFinishReason: choice.nativeFinish}
out.Message.Role = choice.role
out.Message.Content = choice.content.String()
out.Message.ReasoningContent = choice.reasoning.String()
for _, toolIndex := range choice.toolOrder {
toolCall := choice.toolCalls[toolIndex]
outTool := openAIResponseTool{ID: toolCall.id, Type: toolCall.callType}
outTool.Function.Name = toolCall.name
outTool.Function.Arguments = toolCall.arguments.String()
out.Message.ToolCalls = append(out.Message.ToolCalls, outTool)
}
assembled.Choices = append(assembled.Choices, out)
}
data, err := json.Marshal(assembled)
return data, err == nil
}
type openAIResponsesEvent struct {
Type string `json:"type"`
Response json.RawMessage `json:"response"`
}
func assembleOpenAIResponsesSSE(payloads []string) ([]byte, bool) {
var completed json.RawMessage
matched := false
for _, payload := range payloads {
var event openAIResponsesEvent
if err := json.Unmarshal([]byte(payload), &event); err != nil {
return nil, false
}
if !strings.HasPrefix(event.Type, "response.") {
return nil, false
}
matched = true
if event.Type == "response.completed" && len(event.Response) > 0 {
completed = append(json.RawMessage(nil), event.Response...)
}
}
if !matched || len(completed) == 0 {
return nil, false
}
return completed, true
}
type anthropicEvent struct {
Type string `json:"type"`
Message *struct {
ID string `json:"id"`
Type string `json:"type"`
Role string `json:"role"`
Model string `json:"model,omitempty"`
StopReason string `json:"stop_reason"`
StopSequence string `json:"stop_sequence"`
Usage struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
} `json:"message"`
Index int `json:"index"`
ContentBlock *struct {
Type string `json:"type"`
Text string `json:"text"`
ID string `json:"id"`
Name string `json:"name"`
} `json:"content_block"`
Delta *struct {
StopReason string `json:"stop_reason"`
StopSequence string `json:"stop_sequence"`
Text string `json:"text"`
PartialJSON string `json:"partial_json"`
} `json:"delta"`
Usage *struct {
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
}
type anthropicContentBlock struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ID string `json:"id,omitempty"`
Name string `json:"name,omitempty"`
Input map[string]any `json:"input,omitempty"`
partialJSON strings.Builder
}
func assembleAnthropicSSE(payloads []string) ([]byte, bool) {
var assembled struct {
ID string `json:"id"`
Type string `json:"type"`
Role string `json:"role"`
Content []anthropicContentBlock `json:"content"`
Model string `json:"model,omitempty"`
StopReason string `json:"stop_reason,omitempty"`
StopSequence string `json:"stop_sequence,omitempty"`
Usage struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
}
contents := map[int]*anthropicContentBlock{}
order := make([]int, 0, 1)
matched := false
for _, payload := range payloads {
var event anthropicEvent
if err := json.Unmarshal([]byte(payload), &event); err != nil {
return nil, false
}
if !strings.HasPrefix(event.Type, "message_") && !strings.HasPrefix(event.Type, "content_block_") {
return nil, false
}
matched = true
if event.Message != nil {
assembled.ID = event.Message.ID
assembled.Type = event.Message.Type
assembled.Role = event.Message.Role
assembled.Model = event.Message.Model
assembled.StopReason = event.Message.StopReason
assembled.StopSequence = event.Message.StopSequence
assembled.Usage.InputTokens = event.Message.Usage.InputTokens
assembled.Usage.OutputTokens = event.Message.Usage.OutputTokens
}
if event.ContentBlock != nil {
block := contents[event.Index]
if block == nil {
block = &anthropicContentBlock{Type: event.ContentBlock.Type, ID: event.ContentBlock.ID, Name: event.ContentBlock.Name}
contents[event.Index] = block
order = append(order, event.Index)
}
if event.ContentBlock.ID != "" {
block.ID = event.ContentBlock.ID
}
if event.ContentBlock.Name != "" {
block.Name = event.ContentBlock.Name
}
block.Text += event.ContentBlock.Text
}
if event.Delta != nil {
if event.Delta.Text != "" {
block := contents[event.Index]
if block == nil {
block = &anthropicContentBlock{Type: "text"}
contents[event.Index] = block
order = append(order, event.Index)
}
block.Text += event.Delta.Text
}
if event.Delta.PartialJSON != "" {
block := contents[event.Index]
if block == nil {
block = &anthropicContentBlock{Type: "tool_use"}
contents[event.Index] = block
order = append(order, event.Index)
}
block.partialJSON.WriteString(event.Delta.PartialJSON)
}
if event.Delta.StopReason != "" {
assembled.StopReason = event.Delta.StopReason
}
if event.Delta.StopSequence != "" {
assembled.StopSequence = event.Delta.StopSequence
}
}
if event.Usage != nil {
assembled.Usage.OutputTokens = event.Usage.OutputTokens
}
}
if !matched || assembled.Type == "" {
return nil, false
}
for _, index := range order {
block := contents[index]
if block.partialJSON.Len() > 0 {
var input map[string]any
if err := json.Unmarshal([]byte(block.partialJSON.String()), &input); err != nil {
return nil, false
}
block.Input = input
}
assembled.Content = append(assembled.Content, *block)
}
data, err := json.Marshal(assembled)
return data, err == nil
}
type geminiChunk struct {
Raw map[string]json.RawMessage `json:"-"`
Candidates []struct {
Raw map[string]json.RawMessage `json:"-"`
Content struct {
Parts []struct {
Text string `json:"text"`
} `json:"parts"`
Role string `json:"role"`
} `json:"content"`
FinishReason string `json:"finishReason"`
Index int `json:"index"`
} `json:"candidates"`
UsageMetadata json.RawMessage `json:"usageMetadata"`
}
func (g *geminiChunk) UnmarshalJSON(data []byte) error {
type alias geminiChunk
var decoded alias
if err := json.Unmarshal(data, &decoded); err != nil {
return err
}
var raw map[string]json.RawMessage
if err := json.Unmarshal(data, &raw); err != nil {
return err
}
*g = geminiChunk(decoded)
g.Raw = raw
if candidatesRaw, ok := raw["candidates"]; ok {
var rawCandidates []map[string]json.RawMessage
if err := json.Unmarshal(candidatesRaw, &rawCandidates); err != nil {
return err
}
for i := range g.Candidates {
if i < len(rawCandidates) {
g.Candidates[i].Raw = rawCandidates[i]
}
}
}
return nil
}
type geminiCandidate struct {
raw map[string]json.RawMessage
index int
role string
text strings.Builder
partRaw map[string]json.RawMessage
}
func assembleGeminiSSE(payloads []string) ([]byte, bool) {
candidates := map[int]*geminiCandidate{}
order := make([]int, 0, 1)
var usage json.RawMessage
matched := false
for _, payload := range payloads {
var chunk geminiChunk
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
return nil, false
}
if len(chunk.Candidates) == 0 {
return nil, false
}
matched = true
for _, candidate := range chunk.Candidates {
assembled := candidates[candidate.Index]
if assembled == nil {
assembled = &geminiCandidate{index: candidate.Index}
candidates[candidate.Index] = assembled
order = append(order, candidate.Index)
}
if candidate.Raw != nil {
assembled.raw = cloneRawMap(candidate.Raw)
}
if candidate.Content.Role != "" {
assembled.role = candidate.Content.Role
}
if len(candidate.Content.Parts) > 0 {
for _, part := range candidate.Content.Parts {
assembled.text.WriteString(part.Text)
}
if len(assembled.partRaw) == 0 {
var contentRaw struct {
Parts []map[string]json.RawMessage `json:"parts"`
}
if rawContent, ok := candidate.Raw["content"]; ok && json.Unmarshal(rawContent, &contentRaw) == nil && len(contentRaw.Parts) > 0 {
assembled.partRaw = cloneRawMap(contentRaw.Parts[0])
}
}
}
}
if len(chunk.UsageMetadata) > 0 {
usage = chunk.UsageMetadata
}
}
if !matched {
return nil, false
}
assembled := map[string]any{"candidates": make([]any, 0, len(order))}
for _, index := range order {
candidate := candidates[index]
candidateRaw := cloneRawMap(candidate.raw)
candidateRaw["index"] = mustJSON(candidate.index)
contentRaw := map[string]any{"role": candidate.role, "parts": []any{map[string]any{"text": candidate.text.String()}}}
if len(candidate.partRaw) > 0 {
partRaw := cloneRawMap(candidate.partRaw)
partRaw["text"] = mustJSON(candidate.text.String())
contentRaw["parts"] = []any{rawMapToMap(partRaw)}
}
candidateRaw["content"] = mustJSON(contentRaw)
assembled["candidates"] = append(assembled["candidates"].([]any), rawMapToMap(candidateRaw))
}
if len(usage) > 0 {
assembled["usageMetadata"] = usage
}
data, err := json.Marshal(assembled)
return data, err == nil
}
func cloneRawMap(in map[string]json.RawMessage) map[string]json.RawMessage {
out := make(map[string]json.RawMessage, len(in))
for k, v := range in {
out[k] = append(json.RawMessage(nil), v...)
}
return out
}
func mustJSON(v any) json.RawMessage {
data, err := json.Marshal(v)
if err != nil {
return nil
}
return data
}
func rawMapToMap(in map[string]json.RawMessage) map[string]any {
out := make(map[string]any, len(in))
for k, v := range in {
var decoded any
if err := json.Unmarshal(v, &decoded); err != nil {
out[k] = string(v)
continue
}
out[k] = decoded
}
return out
}
+65
View File
@@ -0,0 +1,65 @@
package proxy
import (
"bytes"
"encoding/json"
"strings"
)
type sseEventTracker struct {
buf []byte
terminal bool
}
func (t *sseEventTracker) Write(p []byte) {
t.buf = append(t.buf, p...)
for {
end, separator := completeSSEEvent(t.buf)
if end < 0 {
return
}
event := t.buf[:end]
t.buf = t.buf[end+separator:]
if terminalSSEEvent(event) {
t.terminal = true
}
}
}
func (t *sseEventTracker) Complete() bool { return len(t.buf) == 0 }
func completeSSEEvent(buf []byte) (int, int) {
lf := bytes.Index(buf, []byte("\n\n"))
crlf := bytes.Index(buf, []byte("\r\n\r\n"))
if crlf >= 0 && (lf < 0 || crlf < lf) {
return crlf, 4
}
if lf >= 0 {
return lf, 2
}
return -1, 0
}
func terminalSSEEvent(event []byte) bool {
var data strings.Builder
for _, line := range strings.Split(strings.ReplaceAll(string(event), "\r\n", "\n"), "\n") {
if !strings.HasPrefix(line, "data:") {
continue
}
if data.Len() > 0 {
data.WriteByte('\n')
}
data.WriteString(strings.TrimSpace(strings.TrimPrefix(line, "data:")))
}
payload := data.String()
if payload == "[DONE]" {
return true
}
var envelope struct {
Type string `json:"type"`
}
if json.Unmarshal([]byte(payload), &envelope) != nil {
return false
}
return envelope.Type == "message_stop" || envelope.Type == "response.completed"
}
+101
View File
@@ -0,0 +1,101 @@
package proxy
import (
"io"
"net/http"
"sync"
"sync/atomic"
"time"
)
type trackingBody struct {
io.ReadCloser
failed *atomic.Bool
}
func (b *trackingBody) Read(p []byte) (int, error) {
n, err := b.ReadCloser.Read(p)
if err != nil && err != io.EOF {
b.failed.Store(true)
}
return n, err
}
func responseBodyTimeout(resp *http.Response, opts Options) time.Duration {
if resp.StatusCode == http.StatusSwitchingProtocols {
return 0
}
if isStreamResponse(resp.Header) {
return opts.SSEIdleTimeout
}
return opts.ResponseTimeout
}
type timeoutBody struct {
body io.ReadCloser
idle bool
timeout time.Duration
timer *time.Timer
mu sync.Mutex
done bool
sequence uint64
closeOnce sync.Once
closeErr error
}
func newTimeoutBody(body io.ReadCloser, timeout time.Duration, idle bool) *timeoutBody {
t := &timeoutBody{body: body, idle: idle, timeout: timeout}
t.resetLocked()
return t
}
func (b *timeoutBody) Read(p []byte) (int, error) {
n, err := b.body.Read(p)
b.mu.Lock()
if !b.done {
if err != nil {
b.done = true
b.timer.Stop()
} else if n > 0 && b.idle {
b.resetLocked()
}
}
b.mu.Unlock()
return n, err
}
func (b *timeoutBody) Close() error {
b.mu.Lock()
if !b.done {
b.done = true
b.sequence++
b.timer.Stop()
}
b.mu.Unlock()
return b.closeUnderlying()
}
func (b *timeoutBody) resetLocked() {
if b.timer != nil {
b.timer.Stop()
}
b.sequence++
sequence := b.sequence
b.timer = time.AfterFunc(b.timeout, func() { b.expire(sequence) })
}
func (b *timeoutBody) expire(sequence uint64) {
b.mu.Lock()
if b.done || sequence != b.sequence {
b.mu.Unlock()
return
}
b.done = true
b.mu.Unlock()
_ = b.closeUnderlying()
}
func (b *timeoutBody) closeUnderlying() error {
b.closeOnce.Do(func() { b.closeErr = b.body.Close() })
return b.closeErr
}
+128
View File
@@ -0,0 +1,128 @@
package proxy
import (
"bufio"
"bytes"
"errors"
"net"
"net/http"
"sync"
)
// captureWriter 包装 http.ResponseWriter,边转发边缓冲响应体。
// 实现 http.Flusher 与 http.Hijacker 以支持 SSE/chunked/WebSocket。
type captureWriter struct {
http.ResponseWriter
buf bytes.Buffer
max int64
written int64
truncated bool
status int
wroteHeader bool
hijacked bool
writeFailed bool
onHijack func(net.Conn) net.Conn
sse sseEventTracker
}
func newCaptureWriter(w http.ResponseWriter, max int64) *captureWriter {
return &captureWriter{ResponseWriter: w, max: max, status: http.StatusOK}
}
func (c *captureWriter) WriteHeader(code int) {
if c.wroteHeader {
return
}
c.status = code
c.wroteHeader = true
c.ResponseWriter.WriteHeader(code)
}
func (c *captureWriter) Write(p []byte) (int, error) {
if !c.wroteHeader {
c.wroteHeader = true
}
n, err := c.ResponseWriter.Write(p)
if err != nil || n != len(p) {
c.writeFailed = true
}
if n > len(p) {
n = len(p)
}
if n > 0 {
// 仅缓冲 max 字节以内的内容。
remaining := c.max - c.written
if remaining > 0 {
toBuf := n
if int64(toBuf) > remaining {
toBuf = int(remaining)
c.truncated = true
}
c.buf.Write(p[:toBuf])
} else if c.max > 0 {
c.truncated = true
}
c.written += int64(n)
if isStreamResponse(c.Header()) {
c.sse.Write(p[:n])
}
if f, ok := c.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
return n, err
}
func (c *captureWriter) Flush() {
if f, ok := c.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
func (c *captureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
h, ok := c.ResponseWriter.(http.Hijacker)
if !ok {
return nil, nil, errors.New("hijack not supported")
}
conn, rw, err := h.Hijack()
if err != nil {
c.writeFailed = true
return nil, nil, err
}
c.hijacked = true
if c.onHijack != nil {
conn = c.onHijack(conn)
}
return conn, rw, nil
}
func (c *captureWriter) Body() []byte { return c.buf.Bytes() }
func (c *captureWriter) Truncated() bool { return c.truncated }
func (c *captureWriter) Status() int { return c.status }
func (c *captureWriter) Hijacked() bool { return c.hijacked }
func (c *captureWriter) Complete() bool { return !c.writeFailed }
func (c *captureWriter) SSEComplete() bool { return c.sse.Complete() }
func (c *captureWriter) OnHijack(fn func(net.Conn) net.Conn) { c.onHijack = fn }
func (c *captureWriter) SetHijackedResponse(status int, header http.Header) {
c.status = status
for key := range c.Header() {
c.Header().Del(key)
}
for key, values := range header {
c.Header()[key] = append([]string(nil), values...)
}
}
type trackedConn struct {
net.Conn
closeOnce sync.Once
onClose func()
}
func (c *trackedConn) Close() error {
err := c.Conn.Close()
c.closeOnce.Do(c.onClose)
return err
}
+61
View File
@@ -0,0 +1,61 @@
# tests/
集中放置项目的测试代码。每个子目录对应被测包,使用 Go 的黑盒测试包模式(`package xxx_test`)通过被测包的导出 API 进行验证。
## 目录结构
```
tests/
├── config/ # 配置与 glob 过滤器测试
├── logger/ # 异步日志队列与重试逻辑测试
├── proxy/ # 反向代理测试(WebSocket、SSE、上游错误)
└── scripts/ # 手工端到端测试脚本与 DB 校验工具
```
## 运行
```bash
go test ./tests/...
```
或者运行某一个子包:
```bash
go test ./tests/proxy/...
```
## 端到端冒烟测试
先构建本地二进制:
```powershell
go build -o TokenThief.exe .
```
加载 `.env` 并设置 newapi Key
```powershell
. .\tests\scripts\load-env.ps1
$env:NEWAPI_KEY = "sk-..."
```
`.env` 需要包含可写的 `CLICKHOUSE_URL`,例如 `clickhouse://tokenthief:tokenthief@localhost:9000/tokenthief`
运行聊天端点冒烟测试:
```powershell
.\tests\scripts\smoke.ps1 -Model "gpt-5.4-mini"
```
按 request_id 打印数据库里的流式 `response_body`
```powershell
.\tests\scripts\dump-stream-body.ps1 -RequestID "<request_id>"
```
## 构建排除
- 测试目录中的 `_test.go` 文件不会参与 `go build`
- `tests/scripts/` 只放手工测试脚本,不被主程序 import。
- `Dockerfile` 仅构建 main 包(`./`),不会触及 `tests/`
- 仓库根目录的 `.dockerignore``tests/` 整体排除在 build context 之外,镜像中不会包含测试代码或测试脚本。
+283
View File
@@ -0,0 +1,283 @@
package config_test
import (
"net/netip"
"strings"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/config"
)
func TestLoadTimeoutDefaults(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.ReadTimeout != 30*time.Second {
t.Fatalf("ReadTimeout=%s want 30s", cfg.ReadTimeout)
}
if cfg.WriteTimeout != 10*time.Minute {
t.Fatalf("WriteTimeout=%s want 10m", cfg.WriteTimeout)
}
if cfg.IdleTimeout != 5*time.Minute {
t.Fatalf("IdleTimeout=%s want 5m", cfg.IdleTimeout)
}
if cfg.UpstreamTimeout != 30*time.Second {
t.Fatalf("UpstreamTimeout=%s want 30s", cfg.UpstreamTimeout)
}
if cfg.UpstreamResponseTimeout != 30*time.Second {
t.Fatalf("UpstreamResponseTimeout=%s want 30s", cfg.UpstreamResponseTimeout)
}
if cfg.UpstreamStreamIdleTimeout != 2*time.Minute {
t.Fatalf("UpstreamStreamIdleTimeout=%s want 2m", cfg.UpstreamStreamIdleTimeout)
}
}
func TestLoadQueueDefaults(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.MaxBodyBytes != 1<<20 {
t.Fatalf("MaxBodyBytes=%d want %d", cfg.MaxBodyBytes, 1<<20)
}
if cfg.LogQueueSize != 256 {
t.Fatalf("LogQueueSize=%d want 256", cfg.LogQueueSize)
}
if cfg.LogQueueBytes != 64<<20 {
t.Fatalf("LogQueueBytes=%d want %d", cfg.LogQueueBytes, 64<<20)
}
}
func TestLoadTimeoutOverrides(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
t.Setenv("READ_TIMEOUT", "10s")
t.Setenv("WRITE_TIMEOUT", "20s")
t.Setenv("IDLE_TIMEOUT", "30s")
t.Setenv("UPSTREAM_TIMEOUT", "40s")
t.Setenv("UPSTREAM_RESPONSE_TIMEOUT", "50s")
t.Setenv("UPSTREAM_STREAM_IDLE_TIMEOUT", "60s")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.ReadTimeout != 10*time.Second {
t.Fatalf("ReadTimeout=%s want 10s", cfg.ReadTimeout)
}
if cfg.WriteTimeout != 20*time.Second {
t.Fatalf("WriteTimeout=%s want 20s", cfg.WriteTimeout)
}
if cfg.IdleTimeout != 30*time.Second {
t.Fatalf("IdleTimeout=%s want 30s", cfg.IdleTimeout)
}
if cfg.UpstreamTimeout != 40*time.Second {
t.Fatalf("UpstreamTimeout=%s want 40s", cfg.UpstreamTimeout)
}
if cfg.UpstreamResponseTimeout != 50*time.Second {
t.Fatalf("UpstreamResponseTimeout=%s want 50s", cfg.UpstreamResponseTimeout)
}
if cfg.UpstreamStreamIdleTimeout != 60*time.Second {
t.Fatalf("UpstreamStreamIdleTimeout=%s want 60s", cfg.UpstreamStreamIdleTimeout)
}
}
func TestLoadRejectsInvalidTypedValues(t *testing.T) {
tests := []struct {
name string
key string
value string
}{
{name: "int", key: "LOG_QUEUE_SIZE", value: "not-an-int"},
{name: "int64", key: "MAX_BODY_BYTES", value: "1.5"},
{name: "duration", key: "READ_TIMEOUT", value: "30"},
{name: "bool", key: "UPSTREAM_TLS_INSECURE_SKIP_VERIFY", value: "yes"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
t.Setenv(tt.key, tt.value)
_, err := config.Load()
if err == nil {
t.Fatalf("Load succeeded with %s=%q", tt.key, tt.value)
}
if !strings.Contains(err.Error(), tt.key) {
t.Fatalf("err=%q does not identify %s", err, tt.key)
}
})
}
}
func TestLoadRejectsNonPositiveNumericAndDurationValues(t *testing.T) {
tests := []struct {
key string
value string
}{
{key: "MAX_BODY_BYTES", value: "0"},
{key: "LOG_QUEUE_SIZE", value: "-1"},
{key: "LOG_QUEUE_BYTES", value: "0"},
{key: "LOG_BATCH_SIZE", value: "0"},
{key: "LOG_BATCH_INTERVAL", value: "-1s"},
{key: "LOG_WORKERS", value: "0"},
{key: "DB_RECONNECT_INTERVAL", value: "0s"},
{key: "READ_TIMEOUT", value: "0s"},
{key: "WRITE_TIMEOUT", value: "0s"},
{key: "IDLE_TIMEOUT", value: "0s"},
{key: "UPSTREAM_TIMEOUT", value: "0s"},
{key: "UPSTREAM_RESPONSE_TIMEOUT", value: "0s"},
{key: "UPSTREAM_STREAM_IDLE_TIMEOUT", value: "0s"},
}
for _, tt := range tests {
t.Run(tt.key, func(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
t.Setenv(tt.key, tt.value)
_, err := config.Load()
if err == nil {
t.Fatalf("Load succeeded with %s=%q", tt.key, tt.value)
}
if !strings.Contains(err.Error(), tt.key) {
t.Fatalf("err=%q does not identify %s", err, tt.key)
}
})
}
}
func TestLoadTrustedProxies(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
t.Setenv("TRUSTED_PROXIES", "10.0.0.0/8, 192.0.2.1,2001:db8::/32, 2001:db8::1")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
want := []netip.Prefix{
netip.MustParsePrefix("10.0.0.0/8"),
netip.MustParsePrefix("192.0.2.1/32"),
netip.MustParsePrefix("2001:db8::/32"),
netip.MustParsePrefix("2001:db8::1/128"),
}
if len(cfg.TrustedProxies) != len(want) {
t.Fatalf("TrustedProxies=%v want %v", cfg.TrustedProxies, want)
}
for i := range want {
if cfg.TrustedProxies[i] != want[i] {
t.Fatalf("TrustedProxies[%d]=%v want %v", i, cfg.TrustedProxies[i], want[i])
}
}
}
func TestLoadRejectsInvalidTrustedProxies(t *testing.T) {
for _, value := range []string{"not-an-ip", "10.0.0.0/33", "10.0.0.1/8", "10.0.0.1,,192.0.2.1"} {
t.Run(value, func(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
t.Setenv("TRUSTED_PROXIES", value)
_, err := config.Load()
if err == nil {
t.Fatalf("Load succeeded with TRUSTED_PROXIES=%q", value)
}
})
}
}
func TestLoadUpstreamTLSInsecureSkipVerifyDefault(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.UpstreamTLSInsecureSkipVerify {
t.Fatal("UpstreamTLSInsecureSkipVerify=true want false")
}
}
func TestLoadUpstreamTLSInsecureSkipVerifyOverride(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouse://user:pass@localhost:9000/db")
t.Setenv("UPSTREAM_TLS_INSECURE_SKIP_VERIFY", "true")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if !cfg.UpstreamTLSInsecureSkipVerify {
t.Fatal("UpstreamTLSInsecureSkipVerify=false want true")
}
}
func TestLoadRequiresClickHouseURL(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("DATABASE_URL", "postgres://user:pass@localhost/db")
_, err := config.Load()
if err == nil {
t.Fatal("Load succeeded without CLICKHOUSE_URL")
}
if err.Error() != "CLICKHOUSE_URL is required" {
t.Fatalf("err=%q want CLICKHOUSE_URL is required", err.Error())
}
}
func TestLoadValidatesClickHouseURL(t *testing.T) {
tests := []struct {
name string
dsn string
}{
{name: "missing port", dsn: "clickhouse://user:pass@localhost/db"},
{name: "unsupported scheme", dsn: "https://localhost:9440/db"},
{name: "plaintext skip verify", dsn: "clickhouse://localhost:9000/db?skip_verify=true"},
{name: "conflicting secure parameter", dsn: "clickhouses://localhost:9440/db?secure=false"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", tt.dsn)
_, err := config.Load()
if err == nil {
t.Fatalf("Load succeeded with CLICKHOUSE_URL=%q", tt.dsn)
}
if !strings.Contains(err.Error(), "CLICKHOUSE_URL") {
t.Fatalf("err=%q does not identify CLICKHOUSE_URL", err)
}
})
}
}
func TestLoadAcceptsSecureClickHouseURL(t *testing.T) {
t.Setenv("UPSTREAM_URL", "https://example.com")
t.Setenv("CLICKHOUSE_URL", "clickhouses://user:pass@localhost:9440/db?skip_verify=false")
cfg, err := config.Load()
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.ClickHouseURL != "clickhouses://user:pass@localhost:9440/db?skip_verify=false" {
t.Fatalf("ClickHouseURL=%q", cfg.ClickHouseURL)
}
}
+164
View File
@@ -0,0 +1,164 @@
package config_test
import (
"os"
"path/filepath"
"testing"
"git.misaka.ren/M1saka/token_thief/config"
)
func TestCompileGlobMatch(t *testing.T) {
cases := []struct {
pattern string
target string
want bool
}{
// 单段 *
{"/v1/audio/*", "/v1/audio/speech", true},
{"/v1/audio/*", "/v1/audio/transcriptions", true},
{"/v1/audio/*", "/v1/audio/sub/x", false},
// 字面匹配
{"/v1/chat/completions", "/v1/chat/completions", true},
{"/v1/chat/completions", "/v1/chat/completions/extra", false},
// 跨段 **
{"/v1/videos/**", "/v1/videos/", true},
{"/v1/videos/**", "/v1/videos/abc", true},
{"/v1/videos/**", "/v1/videos/abc/content", true},
// Gemini 风格 :generateContent
{"/v1beta/models/*:generateContent", "/v1beta/models/gemini-pro:generateContent", true},
{"/v1beta/models/*:generateContent", "/v1beta/models/gemini-1.5-flash:generateContent", true},
{"/v1beta/models/*:generateContent", "/v1beta/models/x/y:generateContent", false},
// engines 嵌套
{"/v1/engines/*/embeddings", "/v1/engines/text-embedding-ada-002/embeddings", true},
{"/v1/engines/*/embeddings", "/v1/engines/a/b/embeddings", false},
// 含正则元字符
{"/v1/models/*", "/v1/models/gpt-4.1", true},
}
for _, c := range cases {
re, err := config.CompileGlob(c.pattern)
if err != nil {
t.Fatalf("compile %q: %v", c.pattern, err)
}
got := re.MatchString(c.target)
if got != c.want {
t.Errorf("%q vs %q: got %v want %v (regex=%s)", c.pattern, c.target, got, c.want, re.String())
}
}
}
func TestFilterShouldLog(t *testing.T) {
f, err := config.NewFilter(config.FilterWhitelist, []string{"/v1/chat/completions", "/v1/videos/*"})
if err != nil {
t.Fatalf("NewFilter: %v", err)
}
if !f.ShouldLog("/v1/chat/completions") {
t.Error("whitelist must allow /v1/chat/completions")
}
if !f.ShouldLog("/v1/videos/abc") {
t.Error("whitelist must allow /v1/videos/abc")
}
if f.ShouldLog("/v1/embeddings") {
t.Error("whitelist must reject /v1/embeddings")
}
fb, err := config.NewFilter(config.FilterBlacklist, []string{"/v1/chat/completions", "/v1/videos/*"})
if err != nil {
t.Fatalf("NewFilter blacklist: %v", err)
}
if fb.ShouldLog("/v1/chat/completions") {
t.Error("blacklist must reject /v1/chat/completions")
}
if !fb.ShouldLog("/v1/embeddings") {
t.Error("blacklist must allow /v1/embeddings")
}
fd, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatalf("NewFilter disabled: %v", err)
}
if !fd.ShouldLog("/anything") {
t.Error("disabled must allow everything")
}
}
// TestRepoFilterOnlyAllowsChatAndCompletions 加载仓库根目录的 filter.yaml
// 验证默认过滤器只记录聊天与补全端点。
func TestRepoFilterOnlyAllowsChatAndCompletions(t *testing.T) {
cwd, err := os.Getwd()
if err != nil {
t.Fatalf("getwd: %v", err)
}
// 测试目录位于 <repo>/tests/configfilter.yaml 在 <repo>/filter.yaml
path := filepath.Join(cwd, "..", "..", "filter.yaml")
if _, err := os.Stat(path); err != nil {
t.Fatalf("filter.yaml not found: %v", err)
}
f, err := config.LoadFilter(path)
if err != nil {
t.Fatalf("load filter: %v", err)
}
allowed := []string{
// Chat
"/v1/chat/completions",
"/v1/responses",
"/v1/messages",
"/v1beta/models/gemini-1.5-pro:generateContent",
"/v1beta/models/gemini-1.5-pro:generateContent/",
"/v1beta/models/gemini-1.5-pro:streamGenerateContent",
// Completions
"/v1/completions",
}
for _, p := range allowed {
if !f.ShouldLog(p) {
t.Errorf("filter.yaml should match chat/completion endpoint %q", p)
}
}
blocked := []string{
// Models
"/v1/models",
"/v1beta/models",
// Embeddings
"/v1/embeddings",
"/v1/engines/text-embedding-ada-002/embeddings",
// Moderations / Rerank / Realtime
"/v1/moderations",
"/v1/rerank",
"/v1/realtime",
// Audio
"/v1/audio/speech",
"/v1/audio/transcriptions",
"/v1/audio/translations",
// Images
"/v1/images/generations",
"/v1/images/generations/",
"/v1/images/edits",
"/v1/images/edits/",
// Videos - 通用
"/v1/video/generations",
"/v1/video/generations/task_abc",
// Videos - Sora
"/v1/videos",
"/v1/videos/task_abc",
"/v1/videos/task_abc/content",
// Videos - 即梦
"/jimeng/",
// Videos - Kling
"/kling/v1/videos/text2video",
"/kling/v1/videos/text2video/task_abc",
"/kling/v1/videos/image2video",
"/kling/v1/videos/image2video/task_abc",
}
for _, p := range blocked {
if f.ShouldLog(p) {
t.Errorf("filter.yaml should not match non-chat/completion endpoint %q", p)
}
}
}
+66
View File
@@ -0,0 +1,66 @@
package deployment_test
import (
"os"
"path/filepath"
"regexp"
"strings"
"testing"
)
func TestComposeUsesSafeDeploymentDefaults(t *testing.T) {
root := filepath.Join("..", "..")
compose := readFile(t, filepath.Join(root, "compose.yml"))
envExample := readFile(t, filepath.Join(root, ".env.example"))
for _, required := range []string{
`${UPSTREAM_URL:?set UPSTREAM_URL}`,
`${CLICKHOUSE_URL:?set CLICKHOUSE_URL with URL-encoded credentials}`,
`${CLICKHOUSE_PASSWORD:?set a strong CLICKHOUSE_PASSWORD}`,
} {
if !strings.Contains(compose, required) {
t.Errorf("compose.yml must reject an empty required setting with %q", required)
}
}
if !regexp.MustCompile(`(?m)^\s+image: clickhouse/clickhouse-server:\d+\.\d+\.\d+\.\d+-alpine\s*$`).MatchString(compose) {
t.Error("ClickHouse image must use an exact version tag")
}
serviceStart := strings.LastIndex(compose, " thief_clickhouse:")
if serviceStart < 0 {
t.Fatal("compose.yml is missing the thief_clickhouse service")
}
clickhouseService := compose[serviceStart:]
if strings.Contains(clickhouseService, "\n ports:") {
t.Error("ClickHouse must not publish host ports by default")
}
for _, emptySecret := range []string{"CLICKHOUSE_URL=\n", "CLICKHOUSE_PASSWORD=\n"} {
if !strings.Contains(envExample, emptySecret) {
t.Errorf(".env.example must leave %q empty", strings.TrimSpace(emptySecret))
}
}
}
func TestReviewDocumentsUseStablePaths(t *testing.T) {
root := filepath.Join("..", "..")
for _, path := range []string{
filepath.Join(root, "docs", "compose", "specs", "reliability-security-fixes.md"),
filepath.Join(root, "docs", "compose", "plans", "reliability-security-fixes.md"),
} {
content := readFile(t, path)
if !strings.Contains(content, "2026-07-09-clickhouse-migration.md") {
t.Errorf("%s must link to the dated source document", path)
}
}
}
func readFile(t *testing.T, path string) string {
t.Helper()
b, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
return string(b)
}
+323
View File
@@ -0,0 +1,323 @@
package logger_test
import (
"context"
"errors"
"sync"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/logger"
)
type batchScript struct {
prepareErr error
appendAt int
appendErr error
sendErr error
}
type fakeBatch struct {
script batchScript
rows [][]any
appendCall int
abortCalls int
}
func (f *fakeBatch) Append(v ...any) error {
f.appendCall++
if f.script.appendErr != nil && f.appendCall == f.script.appendAt {
return f.script.appendErr
}
f.rows = append(f.rows, append([]any(nil), v...))
return nil
}
func (f *fakeBatch) Send() error { return f.script.sendErr }
func (f *fakeBatch) Abort() error { f.abortCalls++; return nil }
type fakePool struct {
calls int
queries []string
batches []*fakeBatch
scripts []batchScript
}
func (p *fakePool) PrepareBatch(_ context.Context, query string) (logger.Batch, error) {
idx := p.calls
p.calls++
p.queries = append(p.queries, query)
var script batchScript
if idx < len(p.scripts) {
script = p.scripts[idx]
}
if script.prepareErr != nil {
return nil, script.prepareErr
}
b := &fakeBatch{script: script}
p.batches = append(p.batches, b)
return b, nil
}
func TestFlushSuccessUsesClickHouseTypes(t *testing.T) {
p := &fakePool{}
started := time.Date(2026, 7, 9, 1, 2, 3, 4_000_000, time.UTC)
entry := &logger.LogEntry{
RequestID: "rid", Method: "POST", Path: "/v1", Query: "a=b", ClientIP: "127.0.0.1",
RequestHeaders: []byte(`{"x":"y"}`), RequestBody: []byte("request"), RequestTruncated: true,
StatusCode: 201, ResponseHeaders: []byte(`{"h":"v"}`), ResponseBody: []byte("response"),
ResponseTruncated: true, IsStream: true, LatencyMS: 150, StartedAt: started,
FinishedAt: started.Add(150 * time.Millisecond),
}
result := logger.Flush(context.Background(), p, []*logger.LogEntry{entry})
if result.Err != nil || result.Failed != 0 || len(result.Retry) != 0 {
t.Fatalf("unexpected result: %+v", result)
}
if p.calls != 1 || len(p.batches) != 1 || len(p.batches[0].rows) != 1 {
t.Fatalf("calls=%d batches=%d rows=%d", p.calls, len(p.batches), len(p.batches[0].rows))
}
if contains(p.queries[0], "VALUES") {
t.Fatalf("query contains VALUES: %s", p.queries[0])
}
row := p.batches[0].rows[0]
if len(row) != 17 {
t.Fatalf("columns=%d want 17", len(row))
}
if row[5] != string(entry.RequestHeaders) || row[6] != string(entry.RequestBody) || row[9] != string(entry.ResponseHeaders) || row[10] != string(entry.ResponseBody) {
t.Fatalf("byte fields were not converted to strings: %#v", row)
}
if _, ok := row[8].(int32); !ok {
t.Fatalf("status code type=%T want int32", row[8])
}
if p.batches[0].abortCalls != 0 {
t.Fatalf("abort calls=%d want 0", p.batches[0].abortCalls)
}
}
func TestFlushEmptyBatchDoesNotPrepare(t *testing.T) {
p := &fakePool{}
result := logger.Flush(context.Background(), p, nil)
if result.Err != nil || result.Failed != 0 || result.Ambiguous != 0 || len(result.Retry) != 0 {
t.Fatalf("unexpected result: %+v", result)
}
if p.calls != 0 {
t.Fatalf("PrepareBatch calls=%d want 0", p.calls)
}
}
func TestFlushPrepareFailureIsRetryable(t *testing.T) {
p := &fakePool{scripts: []batchScript{{prepareErr: errors.New("prepare")}}}
entries := []*logger.LogEntry{{RequestID: "a"}, {RequestID: "b"}}
result := logger.Flush(context.Background(), p, entries)
if result.Err == nil || result.Failed != 0 || len(result.Retry) != 2 {
t.Fatalf("unexpected result: %+v", result)
}
if len(p.batches) != 0 {
t.Fatalf("prepare failure created %d batches", len(p.batches))
}
}
func TestFlushAppendFailureIsolatesOnlyBadRow(t *testing.T) {
appendErr := errors.New("bad row")
p := &fakePool{scripts: []batchScript{
{appendAt: 2, appendErr: appendErr},
{},
{appendAt: 1, appendErr: appendErr},
{},
}}
entries := []*logger.LogEntry{{RequestID: "a"}, {RequestID: "bad"}, {RequestID: "c"}}
result := logger.Flush(context.Background(), p, entries)
if result.Err == nil || result.Failed != 1 || len(result.Retry) != 0 {
t.Fatalf("unexpected result: %+v", result)
}
if p.calls != 4 {
t.Fatalf("PrepareBatch calls=%d want 4", p.calls)
}
if p.batches[0].abortCalls != 1 || p.batches[2].abortCalls != 1 {
t.Fatalf("abort calls initial=%d bad-row=%d want 1 each", p.batches[0].abortCalls, p.batches[2].abortCalls)
}
if p.batches[1].abortCalls != 0 || p.batches[3].abortCalls != 0 {
t.Fatalf("successful batches were aborted")
}
}
func TestFlushAppendFailureCanFullyRecover(t *testing.T) {
p := &fakePool{scripts: []batchScript{
{appendAt: 2, appendErr: errors.New("batch append")},
{},
{},
}}
entries := []*logger.LogEntry{{RequestID: "a"}, {RequestID: "b"}}
result := logger.Flush(context.Background(), p, entries)
if result.Err != nil || result.Failed != 0 || len(result.Retry) != 0 {
t.Fatalf("unexpected result: %+v", result)
}
if p.calls != 3 || p.batches[0].abortCalls != 1 {
t.Fatalf("calls=%d abort=%d", p.calls, p.batches[0].abortCalls)
}
}
func TestFlushSendFailureIsAmbiguousAndNotRetried(t *testing.T) {
p := &fakePool{scripts: []batchScript{{sendErr: errors.New("connection lost")}}}
entries := []*logger.LogEntry{{RequestID: "a"}, {RequestID: "b"}}
result := logger.Flush(context.Background(), p, entries)
if result.Err == nil || result.Failed != 2 || len(result.Retry) != 0 {
t.Fatalf("unexpected result: %+v", result)
}
if result.Ambiguous != 2 {
t.Fatalf("Ambiguous=%d want 2", result.Ambiguous)
}
if p.calls != 1 {
t.Fatalf("PrepareBatch calls=%d want 1", p.calls)
}
if p.batches[0].abortCalls != 1 {
t.Fatalf("abort calls=%d want 1", p.batches[0].abortCalls)
}
}
func TestFlushAppendRecoveryCountsAmbiguousSingleSend(t *testing.T) {
p := &fakePool{scripts: []batchScript{
{appendAt: 1, appendErr: errors.New("bad batch")},
{},
{sendErr: errors.New("ack lost")},
}}
entries := []*logger.LogEntry{{RequestID: "a"}, {RequestID: "b"}}
result := logger.Flush(context.Background(), p, entries)
if result.Failed != 1 || result.Ambiguous != 1 || len(result.Retry) != 0 {
t.Fatalf("unexpected result: %+v", result)
}
}
func TestEstimatedBytesCoversStringsAndByteSlices(t *testing.T) {
entry := &logger.LogEntry{
RequestID: "1", Method: "22", Path: "333", Query: "4444", ClientIP: "55555", Error: "666666",
RequestHeaders: []byte("7777777"), RequestBody: []byte("88888888"),
ResponseHeaders: []byte("999999999"), ResponseBody: []byte("0000000000"),
}
if got, want := logger.EstimatedBytes(entry), int64(55); got != want {
t.Fatalf("EstimatedBytes=%d want %d", got, want)
}
}
func TestQueueByteBudgetReleasedAfterFinalDiscard(t *testing.T) {
entry := &logger.LogEntry{
RequestID: "request", Method: "POST", Path: "/path", Query: "q=1", ClientIP: "ip", Error: "error",
RequestHeaders: []byte("rh"), RequestBody: []byte("rb"), ResponseHeaders: []byte("sh"), ResponseBody: []byte("sb"),
}
budget := logger.EstimatedBytes(entry)
q := logger.NewQueue(nil, 4, 2, 1, time.Hour, budget)
q.Submit(entry)
q.Submit(entry)
stats := q.Stats()
if stats.Enqueued != 1 || stats.Dropped != 1 || stats.Bytes != budget {
t.Fatalf("before drain: %+v budget=%d", stats, budget)
}
q.Start(context.Background())
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := q.Stop(ctx); err != nil {
t.Fatalf("Stop: %v", err)
}
stats = q.Stats()
if stats.Bytes != 0 || stats.Failed != 1 {
t.Fatalf("after drain: %+v", stats)
}
}
func TestQueueStopBeforeStartReleasesBudget(t *testing.T) {
entry := &logger.LogEntry{RequestID: "queued"}
q := logger.NewQueue(nil, 2, 2, 1, time.Hour, logger.EstimatedBytes(entry))
q.Submit(entry)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := q.Stop(ctx); err != nil {
t.Fatalf("Stop: %v", err)
}
if stats := q.Stats(); stats.Bytes != 0 || stats.Failed != 1 {
t.Fatalf("stats=%+v", stats)
}
}
func TestQueueCanceledStopEventuallyReleasesAllBudget(t *testing.T) {
entry := &logger.LogEntry{RequestID: "queued"}
q := logger.NewQueue(nil, 32, 32, 1, time.Hour, 32*logger.EstimatedBytes(entry))
q.Start(context.Background())
for i := 0; i < 32; i++ {
q.Submit(entry)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
_ = q.Stop(ctx)
deadline := time.Now().Add(time.Second)
for q.Stats().Bytes != 0 && time.Now().Before(deadline) {
time.Sleep(time.Millisecond)
}
if stats := q.Stats(); stats.Bytes != 0 || stats.Failed != 32 {
t.Fatalf("stats=%+v", stats)
}
}
func TestQueueStopAndSubmitAreConcurrentAndRepeatSafe(t *testing.T) {
q := logger.NewQueue(nil, 32, 8, 2, time.Millisecond, 1<<20)
q.Start(context.Background())
var submitters sync.WaitGroup
for i := 0; i < 8; i++ {
submitters.Add(1)
go func() {
defer submitters.Done()
for j := 0; j < 2_000; j++ {
q.Submit(&logger.LogEntry{RequestID: "x"})
}
}()
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
if err := q.Stop(ctx); err != nil {
t.Fatalf("first Stop: %v", err)
}
submitters.Wait()
if err := q.Stop(ctx); err != nil {
t.Fatalf("second Stop: %v", err)
}
q.Submit(&logger.LogEntry{RequestID: "after-stop"})
if stats := q.Stats(); stats.Bytes != 0 {
t.Fatalf("bytes after Stop=%d want 0", stats.Bytes)
}
}
func TestStopIsIndependentFromStartContext(t *testing.T) {
root, cancelRoot := context.WithCancel(context.Background())
q := logger.NewQueue(nil, 4, 4, 1, time.Hour, 1024)
q.Start(root)
q.Submit(&logger.LogEntry{RequestID: "x"})
cancelRoot()
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := q.Stop(ctx); err != nil {
t.Fatalf("Stop after root cancellation: %v", err)
}
if stats := q.Stats(); stats.Failed != 1 || stats.Bytes != 0 {
t.Fatalf("stats=%+v", stats)
}
}
func contains(s, substr string) bool {
for i := 0; i+len(substr) <= len(s); i++ {
if s[i:i+len(substr)] == substr {
return true
}
}
return false
}
+139
View File
@@ -0,0 +1,139 @@
package proxy_test
import (
"io"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/logger"
"git.misaka.ren/M1saka/token_thief/proxy"
)
type sliceSubmitter struct{ entries []*logger.LogEntry }
func (s *sliceSubmitter) Submit(e *logger.LogEntry) { s.entries = append(s.entries, e) }
// TestUpstreamErrorRecorded 验证上游不可达时 502 响应被记录、错误信息进入 LogEntry.Error。
func TestUpstreamErrorRecorded(t *testing.T) {
// 指向一个一定不可用的端口
badURL, _ := url.Parse("http://127.0.0.1:1") // port 1 几乎肯定 connection refused
sub := &sliceSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(badURL, filter, sub, 1024)
srv := httptest.NewServer(h)
defer srv.Close()
resp, err := http.Get(srv.URL + "/v1/chat/completions")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusBadGateway {
t.Errorf("status=%d want 502, body=%q", resp.StatusCode, body)
}
if resp.Header.Get("X-Request-Id") == "" {
t.Errorf("missing X-Request-Id header")
}
deadline := time.Now().Add(2 * time.Second)
for len(sub.entries) == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if len(sub.entries) == 0 {
t.Fatal("no log entry captured")
}
e := sub.entries[0]
if e.StatusCode != http.StatusBadGateway {
t.Errorf("LogEntry.StatusCode=%d want 502", e.StatusCode)
}
if e.Error == "" {
t.Errorf("LogEntry.Error should be set, got empty")
}
if string(e.ResponseBody) != "bad gateway\n" {
t.Errorf("response_body should be fixed bad gateway, got %q", e.ResponseBody)
}
if e.RequestID == "" {
t.Errorf("LogEntry.RequestID empty")
}
}
// TestModifyResponseSetsRequestID 验证正常上游响应也会带上 X-Request-Id。
func TestModifyResponseSetsRequestID(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(200)
_, _ = w.Write([]byte(`{"ok":true}`))
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &sliceSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 1024)
srv := httptest.NewServer(h)
defer srv.Close()
resp, err := http.Get(srv.URL + "/v1/anything")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
rid := resp.Header.Get("X-Request-Id")
if len(rid) != 32 {
t.Errorf("X-Request-Id length=%d want 32, value=%q", len(rid), rid)
}
deadline := time.Now().Add(2 * time.Second)
for len(sub.entries) == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if len(sub.entries) == 0 {
t.Fatal("no log entry")
}
if sub.entries[0].RequestID != rid {
t.Errorf("LogEntry.RequestID=%q response header=%q (should match)", sub.entries[0].RequestID, rid)
}
}
func TestUpstreamTLSInsecureSkipVerifyAllowsSelfSigned(t *testing.T) {
upstream := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"ok":true}`))
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &sliceSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.NewWithOptions(u, filter, sub, 1024, proxy.Options{UpstreamTLSInsecureSkipVerify: true})
srv := httptest.NewServer(h)
defer srv.Close()
resp, err := http.Get(srv.URL + "/v1/anything")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
t.Fatalf("status=%d want 200, body=%q", resp.StatusCode, body)
}
}
+336
View File
@@ -0,0 +1,336 @@
package proxy_test
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strings"
"sync/atomic"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/proxy"
)
type failingBody struct{ err error }
func (b failingBody) Read([]byte) (int, error) { return 0, b.err }
func (failingBody) Close() error { return nil }
type lateFailingBody struct {
remaining int
err error
}
func (b *lateFailingBody) Read(p []byte) (int, error) {
if b.remaining == 0 {
return 0, b.err
}
n := min(len(p), b.remaining)
for i := range p[:n] {
p[i] = 'x'
}
b.remaining -= n
return n, nil
}
func (*lateFailingBody) Close() error { return nil }
type failingResponseWriter struct {
header http.Header
short bool
}
func (w *failingResponseWriter) Header() http.Header { return w.header }
func (*failingResponseWriter) WriteHeader(int) {}
func (w *failingResponseWriter) Write(p []byte) (int, error) {
if w.short {
return len(p) - 1, nil
}
return 0, errors.New("write failed")
}
func newTestHandler(t *testing.T, upstream *url.URL, sub *captureSubmitter, opts proxy.Options) *proxy.Handler {
t.Helper()
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
return proxy.NewWithOptions(upstream, filter, sub, 1024*1024, opts)
}
func TestRequestBodyReadFailureReturnsFixed400WithoutUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/v1/chat", nil)
req.Body = failingBody{err: errors.New("secret read failure")}
req.ContentLength = 1
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestRequestBodyReadFailureAfterCaptureLimitDoesNotReachUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
_, _ = io.Copy(io.Discard, r.Body)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/v1/chat", nil)
req.Body = &lateFailingBody{remaining: 1024*1024 + 1, err: errors.New("late read failure")}
req.ContentLength = 1024*1024 + 2
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestFilteredRequestBodyReadFailureDoesNotReachUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
filter, err := config.NewFilter(config.FilterBlacklist, []string{"/ignored"})
if err != nil {
t.Fatal(err)
}
h := proxy.NewWithOptions(u, filter, &captureSubmitter{}, 1024, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/ignored", nil)
req.Body = failingBody{err: errors.New("read failure")}
req.ContentLength = 1
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestBadGatewayResponseDoesNotLeakUpstreamError(t *testing.T) {
badURL, _ := url.Parse("http://127.0.0.1:1")
h := newTestHandler(t, badURL, &captureSubmitter{}, proxy.Options{})
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if rec.Code != http.StatusBadGateway || rec.Body.String() != "bad gateway\n" {
t.Fatalf("response=(%d, %q), want fixed 502 bad gateway", rec.Code, rec.Body.String())
}
}
func TestResponseWriteFailurePreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, "response")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
for _, short := range []bool{false, true} {
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
w := &failingResponseWriter{header: make(http.Header), short: short}
h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatalf("short=%v: failed response write was committed", short)
}
}
}
func TestTrustedProxyClientIP(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, "ok")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
trusted := netip.MustParsePrefix("10.0.0.0/8")
tests := []struct {
name string
peer string
xff string
wantIP string
}{
{name: "untrusted peer ignores xff", peer: "203.0.113.9:1234", xff: "198.51.100.1", wantIP: "203.0.113.9"},
{name: "strip trusted from right", peer: "10.0.0.2:1234", xff: "198.51.100.7, 10.0.0.3", wantIP: "198.51.100.7"},
{name: "ipv6", peer: "[2001:db8::2]:1234", xff: "198.51.100.7", wantIP: "2001:db8::2"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{TrustedProxies: []netip.Prefix{trusted}})
req := httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil)
req.RemoteAddr = tc.peer
req.Header.Set("X-Forwarded-For", tc.xff)
h.ServeHTTP(httptest.NewRecorder(), req)
if sub.Len() != 1 || sub.Entry(0).ClientIP != tc.wantIP {
t.Fatalf("ClientIP=%q, want %q", sub.Entry(0).ClientIP, tc.wantIP)
}
})
}
}
func TestTrustedProxyUsesValidXRealIPWithoutXFF(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
trusted := netip.MustParsePrefix("192.0.2.0/24")
h := newTestHandler(t, u, sub, proxy.Options{TrustedProxies: []netip.Prefix{trusted}})
req := httptest.NewRequest(http.MethodGet, "http://proxy.test/x", nil)
req.RemoteAddr = "192.0.2.10:1234"
req.Header.Set("X-Real-IP", "198.51.100.20")
h.ServeHTTP(httptest.NewRecorder(), req)
if sub.Len() != 1 || sub.Entry(0).ClientIP != "198.51.100.20" {
t.Fatalf("entries=%d", sub.Len())
}
}
func TestResponseBodyTimeoutPreventsCommit(t *testing.T) {
release := make(chan struct{})
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
w.(http.Flusher).Flush()
<-release
}))
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{ResponseTimeout: 50 * time.Millisecond})
srv := httptest.NewServer(h)
resp, err := http.Get(srv.URL + "/v1/test")
if err == nil {
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
}
time.Sleep(50 * time.Millisecond)
if sub.Len() != 0 {
t.Fatal("timed out response must not be committed")
}
close(release)
srv.Close()
upstream.Close()
}
func TestSSEIdleTimeoutResetsAfterReads(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
f := w.(http.Flusher)
for _, event := range []string{"data: one\n\n", "data: two\n\n", "data: [DONE]\n\n"} {
_, _ = io.WriteString(w, event)
f.Flush()
time.Sleep(30 * time.Millisecond)
}
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{SSEIdleTimeout: 60 * time.Millisecond})
srv := httptest.NewServer(h)
defer srv.Close()
resp, err := http.Get(srv.URL + "/v1/test")
if err != nil {
t.Fatal(err)
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
if sub.Len() != 1 {
t.Fatalf("entries=%d, want completed SSE commit", sub.Len())
}
}
func TestIncompleteSSEEventPreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = io.WriteString(w, "data: partial")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("SSE ending mid-event must not be committed as complete")
}
}
func TestIncompleteSSEEventAfterTerminalEventPreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = io.WriteString(w, "data: [DONE]\n\n")
w.(http.Flusher).Flush()
_, _ = io.WriteString(w, "data: partial")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("SSE ending mid-event after a terminal event must not be committed")
}
}
func TestShutdownReturnsWithoutHijackedConnections(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := h.Shutdown(ctx); err != nil {
t.Fatal(err)
}
}
func TestTerminalTextInNonSSEBodyDoesNotAffectCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"text":"data: [DONE]"}`)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", strings.NewReader("")))
if sub.Len() != 1 {
t.Fatalf("entries=%d, want 1", sub.Len())
}
}
+824
View File
@@ -0,0 +1,824 @@
package proxy_test
import (
"bufio"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/logger"
"git.misaka.ren/M1saka/token_thief/proxy"
)
type captureSubmitter struct {
mu sync.Mutex
entries []*logger.LogEntry
}
func (c *captureSubmitter) Submit(e *logger.LogEntry) {
c.mu.Lock()
defer c.mu.Unlock()
c.entries = append(c.entries, e)
}
func (c *captureSubmitter) Len() int {
c.mu.Lock()
defer c.mu.Unlock()
return len(c.entries)
}
func (c *captureSubmitter) Entry(i int) *logger.LogEntry {
c.mu.Lock()
defer c.mu.Unlock()
return c.entries[i]
}
// fakeOpenAIStreamUpstream 模拟一个 OpenAI 兼容的 SSE 上游:
// 分 5 次往响应里 write 一行 SSE 数据,每次都 Flush。
func fakeOpenAIStreamUpstream() *httptest.Server {
chunks := []string{
`data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"你"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"好"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}` + "\n\n",
"data: [DONE]\n\n",
}
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.WriteHeader(http.StatusOK)
flusher := w.(http.Flusher)
for _, ch := range chunks {
_, _ = io.WriteString(w, ch)
flusher.Flush()
time.Sleep(5 * time.Millisecond)
}
}))
}
func fakeAnthropicStreamUpstream() *httptest.Server {
chunks := []string{
`event: message_start` + "\n" + `data: {"type":"message_start","message":{"id":"msg-1","type":"message","role":"assistant","content":[],"model":"claude-3-5-sonnet","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}` + "\n\n",
`event: content_block_start` + "\n" + `data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}` + "\n\n",
`event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"你"}}` + "\n\n",
`event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"好"}}` + "\n\n",
`event: message_delta` + "\n" + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":3}}` + "\n\n",
`event: message_stop` + "\n" + `data: {"type":"message_stop"}` + "\n\n",
}
return newSSEUpstream(chunks)
}
func fakeGeminiStreamUpstream() *httptest.Server {
chunks := []string{
`data: {"candidates":[{"content":{"parts":[{"text":"你"}],"role":"model"},"index":0}]}` + "\n\n",
`data: {"candidates":[{"content":{"parts":[{"text":"好"}],"role":"model"},"finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":2,"totalTokenCount":4}}` + "\n\n",
}
return newSSEUpstream(chunks)
}
func fakeUnknownStreamUpstream() *httptest.Server {
return newSSEUpstream([]string{
`event: custom` + "\n" + `data: not-json` + "\n\n",
})
}
func fakeOpenAIStreamUpstreamThatStaysOpen(release <-chan struct{}) *httptest.Server {
chunks := []string{
`data: {"id":"chatcmpl-hang","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-hang","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-hang","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}` + "\n\n",
"data: [DONE]\n\n",
}
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
w.WriteHeader(http.StatusOK)
flusher := w.(http.Flusher)
for _, ch := range chunks {
_, _ = io.WriteString(w, ch)
flusher.Flush()
}
<-release
}))
}
func newSSEUpstream(chunks []string) *httptest.Server {
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.WriteHeader(http.StatusOK)
flusher := w.(http.Flusher)
for _, ch := range chunks {
_, _ = io.WriteString(w, ch)
flusher.Flush()
time.Sleep(5 * time.Millisecond)
}
}))
}
func TestSSEChunkAssembly(t *testing.T) {
upstream := fakeOpenAIStreamUpstream()
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 1024*1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
// 客户端走原始 TCP,逐字节读 + 打印,验证流式实时到达
pu, _ := url.Parse(proxySrv.URL)
conn, err := net.DialTimeout("tcp", pu.Host, 3*time.Second)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
fmt.Fprintf(conn, "POST /v1/chat/completions HTTP/1.1\r\nHost: %s\r\nContent-Length: 0\r\n\r\n", pu.Host)
_ = conn.SetReadDeadline(time.Now().Add(3 * time.Second))
br := bufio.NewReader(conn)
resp, err := http.ReadResponse(br, nil)
if err != nil {
t.Fatal(err)
}
clientBody, _ := io.ReadAll(resp.Body)
// 等待 proxy.ServeHTTP 返回并把 entry 提交
deadline := time.Now().Add(2 * time.Second)
for sub.Len() == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if sub.Len() == 0 {
t.Fatal("no log entry captured")
}
e := sub.Entry(0)
t.Logf("\n========== 客户端收到的字节 ==========\n%s", clientBody)
t.Logf("\n========== 数据库 response_body 字段(按字节原样存储)==========\n%s", e.ResponseBody)
t.Logf("\n========== 元数据 ==========")
t.Logf("is_stream = %v", e.IsStream)
t.Logf("status_code = %d", e.StatusCode)
t.Logf("len(body) = %d bytes", len(e.ResponseBody))
t.Logf("response_truncated = %v", e.ResponseTruncated)
if !strings.Contains(string(clientBody), `"content":"你"`) ||
!strings.Contains(string(clientBody), `"content":"好"`) ||
!strings.Contains(string(clientBody), "[DONE]") {
t.Errorf("client response body 缺少预期 chunk 内容")
}
var captured struct {
Choices []struct {
Message struct {
Role string `json:"role"`
Content string `json:"content"`
} `json:"message"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("stream response_body should be assembled JSON: %v; body=%q", err, e.ResponseBody)
}
if len(captured.Choices) != 1 {
t.Fatalf("assembled JSON choices length=%d, want 1", len(captured.Choices))
}
if captured.Choices[0].Message.Role != "assistant" {
t.Errorf("assembled role=%q, want assistant", captured.Choices[0].Message.Role)
}
if captured.Choices[0].Message.Content != "你好" {
t.Errorf("assembled content=%q, want 你好", captured.Choices[0].Message.Content)
}
if captured.Choices[0].FinishReason != "stop" {
t.Errorf("assembled finish_reason=%q, want stop", captured.Choices[0].FinishReason)
}
if !e.IsStream {
t.Errorf("is_stream 应为 true")
}
}
func TestOpenAIStreamPreservesReasoningAndMetadata(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"think "},"finish_reason":null,"native_finish_reason":null}]}` + "\n\n",
`data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{"reasoning_content":"hard"},"finish_reason":null,"native_finish_reason":null}]}` + "\n\n",
`data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{"content":"final"},"finish_reason":null,"native_finish_reason":null}]}` + "\n\n",
`data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{},"finish_reason":"stop","native_finish_reason":"stop"}]}` + "\n\n",
"data: [DONE]\n\n",
})
e := requestStreamEntry(t, upstream, "/v1/chat/completions")
var captured struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []struct {
Index int `json:"index"`
Message struct {
Role string `json:"role"`
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content"`
} `json:"message"`
FinishReason string `json:"finish_reason"`
NativeFinishReason string `json:"native_finish_reason"`
} `json:"choices"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("openai stream response_body should be JSON: %v; body=%q", err, e.ResponseBody)
}
if captured.ID != "resp-1" || captured.Object != "chat.completion" || captured.Created != 1779335544 || captured.Model != "gpt-5.4-mini-2026-03-17" {
t.Fatalf("unexpected metadata: %+v", captured)
}
if len(captured.Choices) != 1 {
t.Fatalf("choices length=%d, want 1", len(captured.Choices))
}
choice := captured.Choices[0]
if choice.Message.Role != "assistant" || choice.Message.Content != "final" || choice.Message.ReasoningContent != "think hard" {
t.Fatalf("unexpected message: %+v", choice.Message)
}
if choice.FinishReason != "stop" || choice.NativeFinishReason != "stop" {
t.Fatalf("unexpected finish reasons: %+v", choice)
}
}
func TestOpenAIStreamDoneDoesNotSubmitBeforeUpstreamCloses(t *testing.T) {
release := make(chan struct{})
upstream := fakeOpenAIStreamUpstreamThatStaysOpen(release)
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 1024*1024)
proxySrv := httptest.NewServer(h)
released := false
defer func() {
if !released {
close(release)
}
proxySrv.Close()
upstream.Close()
}()
clientDone := make(chan error, 1)
go func() {
resp, err := http.Post(proxySrv.URL+"/v1/chat/completions", "application/json", strings.NewReader(`{"stream":true}`))
if err != nil {
clientDone <- err
return
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
clientDone <- nil
}()
time.Sleep(100 * time.Millisecond)
if sub.Len() != 0 {
t.Fatal("terminal SSE event must not submit before ReverseProxy returns")
}
close(release)
released = true
select {
case err := <-clientDone:
if err != nil {
t.Fatal(err)
}
if sub.Len() != 1 {
t.Fatalf("stream should be submitted once, got %d entries", sub.Len())
}
case <-time.After(2 * time.Second):
t.Fatal("client did not finish after upstream closed")
}
}
func TestOpenAIStreamPreservesUsageChunk(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"id":"resp-usage","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","system_fingerprint":"fp_123","choices":[{"index":0,"delta":{"role":"assistant","content":"ok"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"resp-usage","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","system_fingerprint":"fp_123","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":null}` + "\n\n",
`data: {"id":"resp-usage","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","system_fingerprint":"fp_123","choices":[],"usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7}}` + "\n\n",
"data: [DONE]\n\n",
})
e := requestStreamEntry(t, upstream, "/v1/chat/completions")
var captured struct {
ID string `json:"id"`
SystemFingerprint string `json:"system_fingerprint"`
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
Usage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("openai stream with usage should be assembled JSON: %v; body=%q", err, e.ResponseBody)
}
if captured.ID != "resp-usage" || captured.SystemFingerprint != "fp_123" {
t.Fatalf("metadata not preserved: %+v", captured)
}
if len(captured.Choices) != 1 || captured.Choices[0].Message.Content != "ok" {
t.Fatalf("choices not assembled: %+v", captured.Choices)
}
if captured.Usage.PromptTokens != 5 || captured.Usage.CompletionTokens != 2 || captured.Usage.TotalTokens != 7 {
t.Fatalf("usage not preserved: %+v", captured.Usage)
}
}
func TestOpenAIStreamAssemblesToolCalls(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"id":"resp-tools","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"lookup","arguments":"{\"q\":"}}]},"finish_reason":null}]}` + "\n\n",
`data: {"id":"resp-tools","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"weather\"}"}}]},"finish_reason":null}]}` + "\n\n",
`data: {"id":"resp-tools","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}` + "\n\n",
"data: [DONE]\n\n",
})
e := requestStreamEntry(t, upstream, "/v1/chat/completions")
var captured struct {
Choices []struct {
Message struct {
Role string `json:"role"`
Content string `json:"content"`
ToolCalls []struct {
ID string `json:"id"`
Type string `json:"type"`
Function struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
} `json:"function"`
} `json:"tool_calls"`
} `json:"message"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("openai tool stream should be assembled JSON: %v; body=%q", err, e.ResponseBody)
}
if len(captured.Choices) != 1 || captured.Choices[0].FinishReason != "tool_calls" {
t.Fatalf("unexpected choices: %+v", captured.Choices)
}
message := captured.Choices[0].Message
if message.Role != "assistant" || message.Content != "" {
t.Fatalf("unexpected message basics: %+v", message)
}
if len(message.ToolCalls) != 1 {
t.Fatalf("tool_calls length=%d, want 1; body=%s", len(message.ToolCalls), e.ResponseBody)
}
tool := message.ToolCalls[0]
if tool.ID != "call_1" || tool.Type != "function" || tool.Function.Name != "lookup" || tool.Function.Arguments != `{"q":"weather"}` {
t.Fatalf("unexpected tool call: %+v", tool)
}
}
func TestOpenAIResponsesStreamAssemblesCompletedResponse(t *testing.T) {
upstream := newSSEUpstream([]string{
`event: response.created` + "\n" + `data: {"type":"response.created","response":{"id":"resp-1","object":"response","status":"in_progress","model":"gpt-5.4-mini","output":[]}}` + "\n\n",
`event: response.output_text.delta` + "\n" + `data: {"type":"response.output_text.delta","item_id":"msg-1","output_index":0,"content_index":0,"delta":"hello"}` + "\n\n",
`event: response.completed` + "\n" + `data: {"type":"response.completed","response":{"id":"resp-1","object":"response","status":"completed","model":"gpt-5.4-mini","output":[{"id":"msg-1","type":"message","role":"assistant","content":[{"type":"output_text","text":"hello"}]}],"usage":{"input_tokens":3,"output_tokens":1,"total_tokens":4}}}` + "\n\n",
})
e := requestStreamEntry(t, upstream, "/v1/responses")
var captured struct {
ID string `json:"id"`
Object string `json:"object"`
Status string `json:"status"`
Model string `json:"model"`
Output []struct {
Type string `json:"type"`
Role string `json:"role"`
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"content"`
} `json:"output"`
Usage struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("responses stream should be completed response JSON: %v; body=%q", err, e.ResponseBody)
}
if captured.ID != "resp-1" || captured.Object != "response" || captured.Status != "completed" || captured.Model != "gpt-5.4-mini" {
t.Fatalf("unexpected response metadata: %+v", captured)
}
if len(captured.Output) != 1 || len(captured.Output[0].Content) != 1 || captured.Output[0].Content[0].Text != "hello" {
t.Fatalf("unexpected response output: %+v", captured.Output)
}
if captured.Usage.TotalTokens != 4 {
t.Fatalf("usage not preserved: %+v", captured.Usage)
}
}
func TestOpenAICompletionsStreamAssemblesText(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"id":"cmpl-1","object":"text_completion","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"text":"hello","finish_reason":null}]}` + "\n\n",
`data: {"id":"cmpl-1","object":"text_completion","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"text":" world","finish_reason":null}]}` + "\n\n",
`data: {"id":"cmpl-1","object":"text_completion","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"text":"","finish_reason":"stop"}],"usage":{"prompt_tokens":2,"completion_tokens":2,"total_tokens":4}}` + "\n\n",
"data: [DONE]\n\n",
})
e := requestStreamEntry(t, upstream, "/v1/completions")
var captured struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []struct {
Index int `json:"index"`
Text string `json:"text"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
Usage struct {
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("completions stream should be assembled JSON: %v; body=%q", err, e.ResponseBody)
}
if captured.ID != "cmpl-1" || captured.Object != "text_completion" || captured.Model != "gpt-5.4-mini" || captured.Created != 1779335544 {
t.Fatalf("unexpected metadata: %+v", captured)
}
if len(captured.Choices) != 1 || captured.Choices[0].Text != "hello world" || captured.Choices[0].FinishReason != "stop" {
t.Fatalf("unexpected choices: %+v", captured.Choices)
}
if captured.Usage.TotalTokens != 4 {
t.Fatalf("usage not preserved: %+v", captured.Usage)
}
}
func TestAnthropicSSEAssemblesNativeMessageJSON(t *testing.T) {
e := requestStreamEntry(t, fakeAnthropicStreamUpstream(), "/v1/messages")
var captured struct {
ID string `json:"id"`
Type string `json:"type"`
Role string `json:"role"`
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"content"`
StopReason string `json:"stop_reason"`
Usage struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("anthropic stream response_body should be native JSON: %v; body=%q", err, e.ResponseBody)
}
if captured.ID != "msg-1" || captured.Type != "message" || captured.Role != "assistant" {
t.Fatalf("unexpected anthropic message metadata: %+v", captured)
}
if len(captured.Content) != 1 || captured.Content[0].Type != "text" || captured.Content[0].Text != "你好" {
t.Fatalf("unexpected anthropic content: %+v", captured.Content)
}
if captured.StopReason != "end_turn" {
t.Errorf("stop_reason=%q, want end_turn", captured.StopReason)
}
if captured.Usage.InputTokens != 10 || captured.Usage.OutputTokens != 3 {
t.Errorf("usage=%+v, want input=10 output=3", captured.Usage)
}
}
func TestAnthropicSSEAssemblesToolUseContent(t *testing.T) {
upstream := newSSEUpstream([]string{
`event: message_start` + "\n" + `data: {"type":"message_start","message":{"id":"msg-tool","type":"message","role":"assistant","content":[],"model":"claude-3-5-sonnet","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}` + "\n\n",
`event: content_block_start` + "\n" + `data: {"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"toolu_1","name":"lookup","input":{}}}` + "\n\n",
`event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"q\":"}}` + "\n\n",
`event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"\"weather\"}"}}` + "\n\n",
`event: message_delta` + "\n" + `data: {"type":"message_delta","delta":{"stop_reason":"tool_use","stop_sequence":null},"usage":{"output_tokens":8}}` + "\n\n",
`event: message_stop` + "\n" + `data: {"type":"message_stop"}` + "\n\n",
})
e := requestStreamEntry(t, upstream, "/v1/messages")
var captured struct {
Content []struct {
Type string `json:"type"`
ID string `json:"id"`
Name string `json:"name"`
Input map[string]any `json:"input"`
} `json:"content"`
StopReason string `json:"stop_reason"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("anthropic tool stream should be JSON: %v; body=%q", err, e.ResponseBody)
}
if len(captured.Content) != 1 {
t.Fatalf("content length=%d, want 1; body=%s", len(captured.Content), e.ResponseBody)
}
tool := captured.Content[0]
if tool.Type != "tool_use" || tool.ID != "toolu_1" || tool.Name != "lookup" || tool.Input["q"] != "weather" {
t.Fatalf("unexpected tool content: %+v", tool)
}
if captured.StopReason != "tool_use" {
t.Fatalf("stop_reason=%q, want tool_use", captured.StopReason)
}
}
func TestGeminiSSEAssemblesNativeGenerateContentJSON(t *testing.T) {
e := requestStreamEntry(t, fakeGeminiStreamUpstream(), "/v1beta/models/gemini-1.5-pro:generateContent")
var captured struct {
Candidates []struct {
Content struct {
Role string `json:"role"`
Parts []struct {
Text string `json:"text"`
} `json:"parts"`
} `json:"content"`
FinishReason string `json:"finishReason"`
Index int `json:"index"`
} `json:"candidates"`
UsageMetadata struct {
PromptTokenCount int `json:"promptTokenCount"`
CandidatesTokenCount int `json:"candidatesTokenCount"`
TotalTokenCount int `json:"totalTokenCount"`
} `json:"usageMetadata"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("gemini stream response_body should be native JSON: %v; body=%q", err, e.ResponseBody)
}
if len(captured.Candidates) != 1 {
t.Fatalf("candidates length=%d, want 1", len(captured.Candidates))
}
candidate := captured.Candidates[0]
if candidate.Content.Role != "model" || len(candidate.Content.Parts) != 1 || candidate.Content.Parts[0].Text != "你好" {
t.Fatalf("unexpected gemini content: %+v", candidate.Content)
}
if candidate.FinishReason != "STOP" {
t.Errorf("finishReason=%q, want STOP", candidate.FinishReason)
}
if captured.UsageMetadata.TotalTokenCount != 4 {
t.Errorf("usageMetadata=%+v, want totalTokenCount=4", captured.UsageMetadata)
}
}
func TestGeminiStreamPreservesSafetyRatingsAndUsageMetadata(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"candidates":[{"content":{"parts":[{"text":"你"}],"role":"model"},"finishReason":null,"index":0,"safetyRatings":[{"category":"HARM_CATEGORY_HARASSMENT","probability":"NEGLIGIBLE"}]}],"usageMetadata":{"promptTokenCount":8,"toolUsePromptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":8,"thoughtsTokenCount":10}}` + "\n\n",
`data: {"candidates":[{"content":{"parts":[{"text":"好"}],"role":"model"},"finishReason":"STOP","index":0,"safetyRatings":[{"category":"HARM_CATEGORY_HARASSMENT","probability":"NEGLIGIBLE"}]}],"usageMetadata":{"promptTokenCount":8,"toolUsePromptTokenCount":0,"candidatesTokenCount":2,"totalTokenCount":20,"thoughtsTokenCount":10}}` + "\n\n",
})
e := requestStreamEntry(t, upstream, "/v1beta/models/gemini-1.5-pro:streamGenerateContent")
var captured struct {
Candidates []struct {
Content struct {
Parts []struct {
Text string `json:"text"`
} `json:"parts"`
} `json:"content"`
SafetyRatings []struct {
Category string `json:"category"`
Probability string `json:"probability"`
} `json:"safetyRatings"`
} `json:"candidates"`
UsageMetadata struct {
PromptTokenCount int `json:"promptTokenCount"`
ToolUsePromptTokenCount int `json:"toolUsePromptTokenCount"`
CandidatesTokenCount int `json:"candidatesTokenCount"`
TotalTokenCount int `json:"totalTokenCount"`
ThoughtsTokenCount int `json:"thoughtsTokenCount"`
} `json:"usageMetadata"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("gemini stream response_body should be JSON: %v; body=%q", err, e.ResponseBody)
}
if len(captured.Candidates) != 1 || len(captured.Candidates[0].SafetyRatings) != 1 {
t.Fatalf("expected safetyRatings to be preserved, got %+v", captured.Candidates)
}
if captured.Candidates[0].SafetyRatings[0].Category != "HARM_CATEGORY_HARASSMENT" {
t.Fatalf("unexpected safetyRatings: %+v", captured.Candidates[0].SafetyRatings)
}
if captured.UsageMetadata.ToolUsePromptTokenCount != 0 || captured.UsageMetadata.ThoughtsTokenCount != 10 || captured.UsageMetadata.TotalTokenCount != 20 {
t.Fatalf("usageMetadata fields not preserved: %+v", captured.UsageMetadata)
}
}
func TestGeminiStreamPreservesFunctionCallParts(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"candidates":[{"content":{"parts":[{"functionCall":{"name":"lookup","args":{"q":"weather"}}}],"role":"model"},"finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":2,"totalTokenCount":4}}` + "\n\n",
})
e := requestStreamEntry(t, upstream, "/v1beta/models/gemini-1.5-pro:streamGenerateContent")
var captured struct {
Candidates []struct {
Content struct {
Parts []struct {
FunctionCall struct {
Name string `json:"name"`
Args map[string]any `json:"args"`
} `json:"functionCall"`
} `json:"parts"`
} `json:"content"`
} `json:"candidates"`
}
if err := json.Unmarshal(e.ResponseBody, &captured); err != nil {
t.Fatalf("gemini functionCall stream should be JSON: %v; body=%q", err, e.ResponseBody)
}
call := captured.Candidates[0].Content.Parts[0].FunctionCall
if call.Name != "lookup" || call.Args["q"] != "weather" {
t.Fatalf("functionCall not preserved: %+v; body=%s", call, e.ResponseBody)
}
}
func TestUnknownSSEKeepsRawBody(t *testing.T) {
e := requestStreamEntry(t, fakeUnknownStreamUpstream(), "/v1/chat/completions")
if string(e.ResponseBody) != "event: custom\ndata: not-json\n\n" {
t.Fatalf("unknown stream should keep raw body, got %q", e.ResponseBody)
}
}
func TestTruncatedSSEKeepsCapturedRawBody(t *testing.T) {
upstream := fakeOpenAIStreamUpstream()
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 520)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
resp, err := http.Post(proxySrv.URL+"/v1/chat/completions", "application/json", nil)
if err != nil {
t.Fatal(err)
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
deadline := time.Now().Add(2 * time.Second)
for sub.Len() == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if sub.Len() == 0 {
t.Fatal("no log entry captured")
}
e := sub.Entry(0)
if !e.ResponseTruncated {
t.Fatal("response should be marked truncated")
}
if !strings.HasPrefix(string(e.ResponseBody), "data: ") {
t.Fatalf("truncated stream should keep captured raw body, got %q", e.ResponseBody)
}
if json.Valid(e.ResponseBody) {
t.Fatalf("truncated stream should not be assembled as JSON, got %q", e.ResponseBody)
}
}
func TestMultimodalStreamRequestBodyIsCaptured(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":null}]}` + "\n\n",
`data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}` + "\n\n",
"data: [DONE]\n\n",
})
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 1024*1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
reqBody := `{"model":"gpt-5.4-mini","messages":[{"role":"user","content":[{"type":"text","text":"describe"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA"}}]}],"stream":true}`
resp, err := http.Post(proxySrv.URL+"/v1/chat/completions", "application/json", strings.NewReader(reqBody))
if err != nil {
t.Fatal(err)
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
deadline := time.Now().Add(2 * time.Second)
for sub.Len() == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if sub.Len() == 0 {
t.Fatal("no log entry captured")
}
e := sub.Entry(0)
if e.RequestTruncated {
t.Fatal("multimodal request should not be truncated")
}
requestText := string(e.RequestBody)
if !strings.Contains(requestText, `"image_url"`) || !strings.Contains(requestText, `data:image/png;base64,AAAA`) {
t.Fatalf("request_body should contain image input, got %q", requestText)
}
if !e.IsStream {
t.Fatal("response should still be marked stream")
}
}
func TestChunkedMultimodalRequestBodyIsCaptured(t *testing.T) {
upstream := newSSEUpstream([]string{
`data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":"stop"}]}` + "\n\n",
"data: [DONE]\n\n",
})
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 1024*1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
reqBody := `{"model":"gpt-5.4-mini","messages":[{"role":"user","content":[{"type":"text","text":"describe"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA"}}]}],"stream":true}`
req, err := http.NewRequest(http.MethodPost, proxySrv.URL+"/v1/chat/completions", strings.NewReader(reqBody))
if err != nil {
t.Fatal(err)
}
req.ContentLength = -1
req.Header.Set("Content-Type", "application/json")
req.TransferEncoding = []string{"chunked"}
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
deadline := time.Now().Add(2 * time.Second)
for sub.Len() == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if sub.Len() == 0 {
t.Fatal("no log entry captured")
}
e := sub.Entry(0)
requestText := string(e.RequestBody)
if !strings.Contains(requestText, `"image_url"`) || !strings.Contains(requestText, `data:image/png;base64,AAAA`) {
t.Fatalf("chunked request_body should contain image input, got %q", requestText)
}
}
func requestStreamEntry(t *testing.T, upstream *httptest.Server, path string) *logger.LogEntry {
t.Helper()
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, sub, 1024*1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
resp, err := http.Post(proxySrv.URL+path, "application/json", nil)
if err != nil {
t.Fatal(err)
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
deadline := time.Now().Add(2 * time.Second)
for sub.Len() == 0 && time.Now().Before(deadline) {
time.Sleep(10 * time.Millisecond)
}
if sub.Len() == 0 {
t.Fatal("no log entry captured")
}
return sub.Entry(0)
}
+299
View File
@@ -0,0 +1,299 @@
package proxy_test
import (
"bufio"
"context"
"io"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/logger"
"git.misaka.ren/M1saka/token_thief/proxy"
)
type noopSubmitter struct{}
func (noopSubmitter) Submit(*logger.LogEntry) {}
type chanSubmitter chan *logger.LogEntry
func (c chanSubmitter) Submit(e *logger.LogEntry) { c <- e }
// fakeUpstream 模拟一个最简 WebSocket 升级:
// 收到 GET + Upgrade: websocket 后回 101,然后做字节回声直到对端关闭。
func fakeUpstream(t *testing.T) *httptest.Server {
t.Helper()
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.ToLower(r.Header.Get("Upgrade")) != "websocket" {
http.Error(w, "expected websocket upgrade", http.StatusBadRequest)
return
}
hj, ok := w.(http.Hijacker)
if !ok {
http.Error(w, "no hijack", http.StatusInternalServerError)
return
}
conn, brw, err := hj.Hijack()
if err != nil {
t.Errorf("upstream hijack: %v", err)
return
}
defer conn.Close()
// 直接回 101 握手响应(简化版,不做真正 Sec-WebSocket-Accept 计算)
_, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
"Upgrade: websocket\r\n" +
"Connection: Upgrade\r\n" +
"\r\n")
_ = brw.Flush()
// echo
buf := make([]byte, 1024)
for {
n, err := conn.Read(buf)
if err != nil {
return
}
if _, err := conn.Write(buf[:n]); err != nil {
return
}
}
}))
return srv
}
func TestWebSocketHandshakeCapturedAndShutdownClosesConnection(t *testing.T) {
upstream := fakeUpstream(t)
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
entries := make(chanSubmitter, 1)
h := proxy.New(u, filter, entries, 1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
pu, _ := url.Parse(proxySrv.URL)
conn, err := net.DialTimeout("tcp", pu.Host, time.Second)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
_, _ = io.WriteString(conn, "GET /v1/realtime HTTP/1.1\r\nHost: "+pu.Host+"\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n")
br := bufio.NewReader(conn)
resp, err := http.ReadResponse(br, nil)
if err != nil {
t.Fatal(err)
}
if resp.StatusCode != http.StatusSwitchingProtocols {
t.Fatalf("status=%d", resp.StatusCode)
}
select {
case entry := <-entries:
if entry.StatusCode != http.StatusSwitchingProtocols || len(entry.ResponseBody) != 0 {
t.Fatalf("handshake entry status=%d body=%q", entry.StatusCode, entry.ResponseBody)
}
case <-time.After(time.Second):
t.Fatal("websocket handshake was not captured")
}
payload := []byte("frame-data-must-not-be-logged")
if _, err := conn.Write(payload); err != nil {
t.Fatal(err)
}
echo := make([]byte, len(payload))
if _, err := io.ReadFull(br, echo); err != nil {
t.Fatalf("read websocket payload: %v", err)
}
if string(echo) != string(payload) {
t.Fatalf("echo mismatch: got %q want %q", echo, payload)
}
select {
case entry := <-entries:
t.Fatalf("websocket frame produced an extra log entry: %+v", entry)
case <-time.After(50 * time.Millisecond):
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := h.Shutdown(ctx); err != nil {
t.Fatal(err)
}
_ = conn.SetReadDeadline(time.Now().Add(time.Second))
if _, err := conn.Read(make([]byte, 1)); err == nil {
t.Fatal("connection remains open after Shutdown")
}
if err := h.Shutdown(ctx); err != nil {
t.Fatalf("second Shutdown: %v", err)
}
}
func TestWebSocketNaturalCloseUnregistersConnection(t *testing.T) {
upstream := fakeUpstream(t)
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, noopSubmitter{}, 1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
pu, _ := url.Parse(proxySrv.URL)
conn, err := net.DialTimeout("tcp", pu.Host, time.Second)
if err != nil {
t.Fatal(err)
}
_, _ = io.WriteString(conn, "GET /v1/realtime HTTP/1.1\r\nHost: "+pu.Host+"\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n")
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
if err != nil {
t.Fatal(err)
}
if resp.StatusCode != http.StatusSwitchingProtocols {
t.Fatalf("status=%d", resp.StatusCode)
}
if err := conn.Close(); err != nil {
t.Fatal(err)
}
deadline := time.Now().Add(time.Second)
for {
ctx, cancel := context.WithCancel(context.Background())
cancel()
err := h.Shutdown(ctx)
if err == nil {
break
}
if time.Now().After(deadline) {
t.Fatalf("connection was not unregistered after natural close: %v", err)
}
time.Sleep(10 * time.Millisecond)
}
}
func TestNonWebSocketUpgradeIsManagedButNotLogged(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hj := w.(http.Hijacker)
conn, brw, err := hj.Hijack()
if err != nil {
t.Errorf("upstream hijack: %v", err)
return
}
defer conn.Close()
_, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\nUpgrade: test-protocol\r\nConnection: Upgrade\r\n\r\n")
_ = brw.Flush()
_, _ = io.Copy(io.Discard, conn)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
entries := make(chanSubmitter, 1)
h := proxy.New(u, filter, entries, 1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
pu, _ := url.Parse(proxySrv.URL)
conn, err := net.DialTimeout("tcp", pu.Host, time.Second)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
_, _ = io.WriteString(conn, "GET /upgrade HTTP/1.1\r\nHost: "+pu.Host+"\r\nUpgrade: test-protocol\r\nConnection: Upgrade\r\n\r\n")
resp, err := http.ReadResponse(bufio.NewReader(conn), nil)
if err != nil {
t.Fatal(err)
}
if resp.StatusCode != http.StatusSwitchingProtocols {
t.Fatalf("status=%d", resp.StatusCode)
}
select {
case entry := <-entries:
t.Fatalf("non-WebSocket upgrade was logged: %+v", entry)
case <-time.After(50 * time.Millisecond):
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := h.Shutdown(ctx); err != nil {
t.Fatal(err)
}
}
// TestWebSocketProxyPassthrough 验证 Upgrade 请求能正确透传,
// 确认 captureWriter 的 Hijacker 实现没破坏 ReverseProxy 的 WS 行为。
func TestWebSocketProxyPassthrough(t *testing.T) {
upstream := fakeUpstream(t)
defer upstream.Close()
u, err := url.Parse(upstream.URL)
if err != nil {
t.Fatal(err)
}
// disabled 模式下 ShouldLog 返回 true(全量记录),会进入捕获分支。
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
h := proxy.New(u, filter, noopSubmitter{}, 1024)
proxySrv := httptest.NewServer(h)
defer proxySrv.Close()
// 建立到代理的 TCP 连接,手写 Upgrade 请求
pu, _ := url.Parse(proxySrv.URL)
d := net.Dialer{Timeout: 3 * time.Second}
conn, err := d.DialContext(context.Background(), "tcp", pu.Host)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
req := "GET /v1/realtime HTTP/1.1\r\n" +
"Host: " + pu.Host + "\r\n" +
"Upgrade: websocket\r\n" +
"Connection: Upgrade\r\n" +
"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n" +
"Sec-WebSocket-Version: 13\r\n" +
"\r\n"
if _, err := io.WriteString(conn, req); err != nil {
t.Fatal(err)
}
_ = conn.SetReadDeadline(time.Now().Add(3 * time.Second))
br := bufio.NewReader(conn)
resp, err := http.ReadResponse(br, nil)
if err != nil {
t.Fatalf("read upgrade response: %v", err)
}
if resp.StatusCode != http.StatusSwitchingProtocols {
t.Fatalf("expected 101, got %d", resp.StatusCode)
}
if !strings.EqualFold(resp.Header.Get("Upgrade"), "websocket") {
t.Fatalf("expected Upgrade: websocket, got %q", resp.Header.Get("Upgrade"))
}
// echo 测试
payload := "hello-websocket"
if _, err := io.WriteString(conn, payload); err != nil {
t.Fatal(err)
}
got := make([]byte, len(payload))
if _, err := io.ReadFull(br, got); err != nil {
t.Fatalf("read echo: %v", err)
}
if string(got) != payload {
t.Fatalf("echo mismatch: got %q want %q", got, payload)
}
}
+109
View File
@@ -0,0 +1,109 @@
// Standalone helper for smoke.ps1. It reads CHECK_RIDS (JSON array) and verifies rows in proxy_logs.
package main
import (
"context"
"encoding/json"
"fmt"
"log"
"os"
"time"
"github.com/ClickHouse/clickhouse-go/v2"
"git.misaka.ren/M1saka/token_thief/db"
)
type row struct {
RequestID string
Method string
Path string
StatusCode int
RequestTruncated bool
ResponseTruncated bool
IsStream bool
LatencyMS int64
ReqBodyLen int64
RespBodyLen int64
ReqHeaders string
ErrorMsg string
}
func main() {
log.SetFlags(0)
dsn := os.Getenv("CLICKHOUSE_URL")
if dsn == "" {
log.Fatalf("CLICKHOUSE_URL not set")
}
ridsRaw := os.Getenv("CHECK_RIDS")
if ridsRaw == "" {
log.Fatalf("CHECK_RIDS not set")
}
var rids []string
if err := json.Unmarshal([]byte(ridsRaw), &rids); err != nil {
log.Fatalf("parse CHECK_RIDS: %v", err)
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
opts, err := db.ClickHouseOptions(dsn)
if err != nil {
log.Fatalf("connect: %v", err)
}
conn, err := clickhouse.Open(opts)
if err != nil {
log.Fatalf("connect: %v", err)
}
defer conn.Close()
const q = `
SELECT request_id, method, path, status_code,
request_truncated, response_truncated, is_stream, latency_ms,
toInt64(length(request_body)),
toInt64(length(response_body)),
request_headers,
error
FROM proxy_logs WHERE request_id = ?
ORDER BY started_at DESC LIMIT 1`
ok := 0
for _, rid := range rids {
var r row
err := conn.QueryRow(ctx, q, rid).Scan(
&r.RequestID, &r.Method, &r.Path, &r.StatusCode,
&r.RequestTruncated, &r.ResponseTruncated, &r.IsStream, &r.LatencyMS,
&r.ReqBodyLen, &r.RespBodyLen, &r.ReqHeaders, &r.ErrorMsg,
)
if err != nil {
fmt.Printf(" FAIL rid=%s: not found in db (%v)\n", rid, err)
continue
}
ok++
fmt.Printf(" OK rid=%s\n", rid)
fmt.Printf(" method=%s path=%s status=%d latency_ms=%d\n", r.Method, r.Path, r.StatusCode, r.LatencyMS)
fmt.Printf(" req_body=%d B (truncated=%v) resp_body=%d B (truncated=%v) is_stream=%v\n",
r.ReqBodyLen, r.RequestTruncated, r.RespBodyLen, r.ResponseTruncated, r.IsStream)
hasAuth := false
var hm map[string][]string
if err := json.Unmarshal([]byte(r.ReqHeaders), &hm); err == nil {
_, hasAuth = hm["Authorization"]
}
fmt.Printf(" headers has Authorization=%v\n", hasAuth)
if r.ErrorMsg != "" {
fmt.Printf(" error=%s\n", trunc(r.ErrorMsg, 200))
}
}
fmt.Printf("\n found %d / %d\n", ok, len(rids))
if ok != len(rids) {
os.Exit(1)
}
}
func trunc(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n] + "..."
}
+54
View File
@@ -0,0 +1,54 @@
param([string]$RequestID)
if (-not $RequestID) { throw "RequestID is required" }
. .\tests\scripts\load-env.ps1 | Out-Null
$env:DUMP_RID = $RequestID
@'
package main
import (
"context"
"fmt"
"log"
"os"
"time"
"github.com/ClickHouse/clickhouse-go/v2"
"git.misaka.ren/M1saka/token_thief/db"
)
func main() {
log.SetFlags(0)
dsn := os.Getenv("CLICKHOUSE_URL")
if dsn == "" {
log.Fatal("CLICKHOUSE_URL not set")
}
rid := os.Getenv("DUMP_RID")
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
opts, err := db.ClickHouseOptions(dsn)
if err != nil {
log.Fatal(err)
}
conn, err := clickhouse.Open(opts)
if err != nil {
log.Fatal(err)
}
defer conn.Close()
var body string
if err := conn.QueryRow(ctx,
`SELECT response_body FROM proxy_logs WHERE request_id = ? ORDER BY started_at DESC LIMIT 1`,
rid,
).Scan(&body); err != nil {
log.Fatal(err)
}
fmt.Println(body)
}
'@ | Set-Content -Path tmp_dump.go -Encoding UTF8 -NoNewline
go run tmp_dump.go
Remove-Item tmp_dump.go -Force
Remove-Item Env:DUMP_RID -ErrorAction SilentlyContinue
+25
View File
@@ -0,0 +1,25 @@
# Load KEY=VALUE pairs from .env into the current PowerShell process.
# Usage: . .\tests\scripts\load-env.ps1
param(
[string]$Path = ".env"
)
if (-not (Test-Path $Path)) {
Write-Error "env file not found: $Path"
return
}
Get-Content $Path | ForEach-Object {
$line = $_.Trim()
if ($line -eq "" -or $line.StartsWith("#")) { return }
$idx = $line.IndexOf("=")
if ($idx -lt 1) { return }
$key = $line.Substring(0, $idx).Trim()
$val = $line.Substring($idx + 1).Trim()
if (($val.StartsWith('"') -and $val.EndsWith('"')) -or
($val.StartsWith("'") -and $val.EndsWith("'"))) {
$val = $val.Substring(1, $val.Length - 2)
}
[Environment]::SetEnvironmentVariable($key, $val, "Process")
Write-Host " loaded $key"
}
+135
View File
@@ -0,0 +1,135 @@
# End-to-end smoke test:
# start proxy -> run chat cases -> wait batch flush -> verify ClickHouse -> stop proxy.
param(
[string]$Model = "gpt-5.4-mini",
[string]$ApiKey = $env:NEWAPI_KEY,
[string]$ProxyBase = "http://127.0.0.1:8080"
)
if (-not $ApiKey) { throw "NEWAPI_KEY not set" }
$ErrorActionPreference = "Stop"
function Section($name) {
Write-Host ""
Write-Host ("=" * 70) -ForegroundColor Cyan
Write-Host $name -ForegroundColor Cyan
Write-Host ("=" * 70) -ForegroundColor Cyan
}
Section "Start proxy"
if (Test-Path proxy.log) { Remove-Item proxy.log -Force }
if (Test-Path proxy.err.log) { Remove-Item proxy.err.log -Force }
$proxy = Start-Process -FilePath .\TokenThief.exe -PassThru -RedirectStandardOutput proxy.log -RedirectStandardError proxy.err.log -WindowStyle Hidden
Write-Host "proxy pid=$($proxy.Id)"
Start-Sleep -Seconds 2
$health = & curl.exe -s -o NUL -w "%{http_code}" "$ProxyBase/healthz"
Write-Host "healthz: $health"
if ($health -ne "200") {
Get-Content proxy.log -Tail 30 | Write-Host
Get-Content proxy.err.log -Tail 30 | Write-Host
throw "proxy did not start"
}
$results = @{}
function CallChat {
param([string]$Url, [string]$JsonBody)
$bodyTmp = New-TemporaryFile
$headTmp = New-TemporaryFile
$outTmp = New-TemporaryFile
# PowerShell 5.1 Set-Content -Encoding UTF8 writes a BOM; newapi rejects BOM-prefixed JSON.
$utf8NoBom = New-Object System.Text.UTF8Encoding($false)
[System.IO.File]::WriteAllText($bodyTmp.FullName, $JsonBody, $utf8NoBom)
$code = & curl.exe -s -X POST $Url `
-H "Content-Type: application/json" `
-H "Authorization: Bearer $ApiKey" `
-D $headTmp.FullName `
-o $outTmp.FullName `
--data-binary "@$($bodyTmp.FullName)" `
-w "%{http_code}"
$rid = ""
foreach ($line in Get-Content $headTmp.FullName) {
if ($line -match '^X-Request-Id:\s*(.+)$') {
$rid = $matches[1].Trim()
break
}
}
$body = Get-Content $outTmp.FullName -Raw -ErrorAction SilentlyContinue
if (-not $body) { $body = "" }
Remove-Item $bodyTmp.FullName, $headTmp.FullName, $outTmp.FullName -Force -ErrorAction SilentlyContinue
return @{ StatusCode = $code; RequestID = $rid; Body = $body }
}
try {
Section "T1: non-stream chat completions"
$t1Body = '{"model":"' + $Model + '","messages":[{"role":"user","content":"Say hello in one short sentence."}],"stream":false}'
$r = CallChat "$ProxyBase/v1/chat/completions" $t1Body
Write-Host " status=$($r.StatusCode) rid=$($r.RequestID)"
$preview = if ($r.Body.Length -gt 200) { $r.Body.Substring(0, 200) } else { $r.Body }
Write-Host " body(first 200): $preview"
$results.T1 = $r
Section "T2: stream chat completions"
$t2Body = '{"model":"' + $Model + '","messages":[{"role":"user","content":"Say exactly: one two three four five"}],"stream":true}'
$r = CallChat "$ProxyBase/v1/chat/completions" $t2Body
Write-Host " status=$($r.StatusCode) rid=$($r.RequestID)"
$chunkCount = ([regex]::Matches($r.Body, "^data:", "Multiline")).Count
$hasDone = $r.Body.Contains("[DONE]")
Write-Host " SSE chunk lines=$chunkCount contains [DONE]=$hasDone body_len=$($r.Body.Length)"
$results.T2 = $r
Section "T3: large body (request_truncated should be true)"
$bigContent = "x" * (12 * 1024 * 1024)
$t3Body = '{"model":"' + $Model + '","messages":[{"role":"user","content":"' + $bigContent + '"}],"stream":false,"max_tokens":5}'
Write-Host " request body size: $($t3Body.Length) bytes"
$r = CallChat "$ProxyBase/v1/chat/completions" $t3Body
Write-Host " status=$($r.StatusCode) rid=$($r.RequestID)"
$preview = if ($r.Body.Length -gt 200) { $r.Body.Substring(0, 200) } else { $r.Body }
Write-Host " body(first 200): $preview"
$results.T3 = $r
Section "T4: nonexistent model (upstream error response should be logged)"
$t4Body = '{"model":"definitely-not-a-real-model-xyz","messages":[{"role":"user","content":"hi"}]}'
$r = CallChat "$ProxyBase/v1/chat/completions" $t4Body
Write-Host " status=$($r.StatusCode) rid=$($r.RequestID)"
$preview = if ($r.Body.Length -gt 200) { $r.Body.Substring(0, 200) } else { $r.Body }
Write-Host " body(first 200): $preview"
$results.T4 = $r
Section "T5: healthz (should not be logged)"
$code = & curl.exe -s -o NUL -w "%{http_code}" "$ProxyBase/healthz"
Write-Host " /healthz status=$code"
Section "Wait batch flush (4s)"
Start-Sleep -Seconds 4
Section "T6: ClickHouse verification"
$rids = @()
foreach ($k in @("T1","T2","T3","T4")) {
if ($results[$k].RequestID) { $rids += $results[$k].RequestID }
}
Write-Host " request_ids: $($rids -join ', ')"
$env:CHECK_RIDS = ($rids | ConvertTo-Json -Compress)
if ($rids.Count -eq 1) { $env:CHECK_RIDS = "[`"$($rids[0])`"]" }
& go run .\tests\scripts\dbcheck\main.go
Remove-Item Env:CHECK_RIDS -ErrorAction SilentlyContinue
} finally {
Section "Stop proxy"
Stop-Process -Id $proxy.Id -Force
Write-Host "tail proxy.log:"
Get-Content proxy.log -Tail 30 | Write-Host
if (Test-Path proxy.err.log) {
$err = Get-Content proxy.err.log -ErrorAction SilentlyContinue
if ($err) {
Write-Host "proxy.err.log:"
$err | Write-Host
}
}
}