- routes.py 主时序:auth(热缓存) -> 限流/并发槽 -> 413 -> 池条目解析 -> try_hold 预扣(字符/3 + min(max_tokens,4096) 宁可高估)-> 上游流式 tee -> usage 归一 -> compute -> settle(actual 回补) -> 流水;异常路径: UpstreamError 全额 void / UpstreamAborted 按已收 usage 结算 status=aborted - install_error_handlers:ProxyError 家族 -> §5.3 错误码 (401/402/403/413/429/502,401 带 WWW-Authenticate) - 管理面:students/topup/keys 签发/keys-revoke(X-Admin-Key 或 loopback) - SSE 同构:chunk 补齐 object/created/id/model(OpenAI SDK 兼容形状), usage chunk 按客户端要求过滤 - 修复:补回 T-P4 重写时丢失的 /proxy/v1/models - 测试 +8(非流式全程含 2550/1491 黄金账目/流式/402/401/413/502 void/ httpx 手写 OpenAI SDK 合规断言含 chunk 形状+usage 过滤+[DONE]),全量 362 passed
291 lines
11 KiB
Python
291 lines
11 KiB
Python
"""代理路由端到端测试(T-P4,M1 验收):全程计费/错误码/SSE 流式/httpx 合规客户端。"""
|
||
import asyncio
|
||
import json
|
||
|
||
import httpx
|
||
import pytest
|
||
|
||
pytest.importorskip("fastapi")
|
||
|
||
from fastapi import FastAPI
|
||
from fastapi.testclient import TestClient
|
||
|
||
import gateway.proxy.upstream as up
|
||
import gateway.model_pool as mp
|
||
from gateway.model_pool import PoolStore
|
||
from gateway.proxy.auth import issue_key, reset_auth_state
|
||
from gateway.proxy.config import build_proxy_config
|
||
from gateway.proxy.errors import (
|
||
BalanceError,
|
||
BodyTooLargeError,
|
||
ProxyAuthError,
|
||
QuotaError,
|
||
)
|
||
|
||
|
||
def asyncio_run(coro):
|
||
return asyncio.run(coro)
|
||
|
||
|
||
UPSTREAM_SSE = "\n\n".join([
|
||
'data: {"choices":[{"delta":{"role":"assistant","content":"\u4f60\u597d"}}]}',
|
||
'data: {"choices":[{"delta":{"content":"\uff0c\u4e16\u754c"}}]}',
|
||
'data: {"choices":[{"delta":{},"finish_reason":"stop"}],'
|
||
# 智能体任务级用量:in_miss=600k, in_hit=300k, out=80k
|
||
'"usage":{"prompt_tokens":900000,"prompt_cache_hit_tokens":300000,"completion_tokens":80000}}',
|
||
"data: [DONE]",
|
||
]) + "\n\n"
|
||
|
||
|
||
def _make_app(tmp_path, upstream_body=UPSTREAM_SSE, fail_upstream=False):
|
||
"""独立 FastAPI app + 池 + key。返回 (client, ledger, key, upstream_calls)。"""
|
||
mp.reset_pool()
|
||
mp._store = PoolStore(path=tmp_path / "pool.json")
|
||
mp.get_pool().upsert({
|
||
"id": "up1", "name": "云端", "tier": "budget", "backend": "openai",
|
||
"base_url": "http://upstream.test", "model": "deepseek-chat",
|
||
"api_key": "up-key", "provider": "deepseek",
|
||
"price_in": 0.1, "price_out": 0.1, "enabled": True})
|
||
reset_auth_state()
|
||
|
||
calls = {"n": 0}
|
||
|
||
def handler(request: httpx.Request) -> httpx.Response:
|
||
calls["n"] += 1
|
||
if fail_upstream:
|
||
raise httpx.ConnectError("上游不可达")
|
||
return httpx.Response(200, content=upstream_body.encode("utf-8"))
|
||
|
||
import gateway.proxy.upstream as upmod
|
||
client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
|
||
orig = upmod._client
|
||
upmod._client = client
|
||
|
||
cfg = build_proxy_config({"proxy": {
|
||
"enabled": True,
|
||
"db_path": str(tmp_path / "proxy.sqlite3"),
|
||
"buckets": {"default": {"system_template": "你是校园学习助手。",
|
||
"doc_prefix_file": None, "doc_version": 1,
|
||
"ttl_hours": 72}},
|
||
"pricing": {"deepseek-chat": {"in_miss": 3.0, "in_hit": 0.1, "out": 9.0},
|
||
"peak_window": {"start": "00:00", "end": "23:59"},
|
||
"offpeak_factor": 0.5,
|
||
"sale_discount": {"in": 0.5, "out": 0.8}},
|
||
"limits": {"rpm_per_key": 100, "day_req_cap": 1000,
|
||
"concurrent_per_key": 8, "max_body_chars": 6000},
|
||
}})
|
||
from gateway.proxy.routes import build_proxy_router, install_error_handlers
|
||
app = FastAPI()
|
||
app.include_router(build_proxy_router(cfg, mp.get_pool()))
|
||
install_error_handlers(app)
|
||
ledger = app.router.routes # noqa: F841(占位说明:ledger 经 state 可取)
|
||
st_ledger = None
|
||
for r in app.routes:
|
||
st = getattr(r, "state", None)
|
||
# 直接从 routes builder 拿 ledger:router.state 在 APIRouter 上不可用,
|
||
# 用模块内函数重建一次(同 db path 幂等)
|
||
from gateway.proxy.ledger import Ledger
|
||
st_ledger = Ledger.init_db(tmp_path / "proxy.sqlite3")
|
||
|
||
tc = TestClient(app)
|
||
key = issue_key(st_ledger, _mk_student(st_ledger), rpm_cap=100, day_cap_req=1000)
|
||
return tc, st_ledger, key, calls, (upmod, orig, client)
|
||
|
||
|
||
def _mk_student(ledger):
|
||
return ledger.upsert_student("学生A", balance_yuan=10.0, daily_cap_yuan=100)
|
||
|
||
|
||
def _auth(key):
|
||
return {"Authorization": f"Bearer {key['key']}"}
|
||
|
||
|
||
BODY = {"model": "deepseek-chat", "messages": [{"role": "user", "content": "问个问题"}]}
|
||
|
||
|
||
def test_chat_non_stream_end_to_end(tmp_path):
|
||
"""非流式全程:回答 + OpenAI 形状 + 账本三值一致 + usage 注入。"""
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
r = tc.post("/proxy/v1/chat/completions", json=BODY, headers=_auth(key))
|
||
assert r.status_code == 200, r.text
|
||
data = r.json()
|
||
assert data["object"] == "chat.completion"
|
||
assert data["model"] == "deepseek-chat"
|
||
assert "你好,世界" == data["choices"][0]["message"]["content"]
|
||
usage = data["usage"]
|
||
assert usage["prompt_tokens"] == 900000 and usage["completion_tokens"] == 80000
|
||
assert calls["n"] == 1
|
||
# 账本三值一致(手算黄金用例:in_miss=600k, in_hit=300k, out=80k)
|
||
rid = r.headers["x-request-id"]
|
||
u = ledger.get_usage(rid)
|
||
assert u["status"] == "ok"
|
||
# cost = 600k×3000/1M + 300k×100/1M + 80k×9000/1M = 1800+30+720 = 2550 毫元
|
||
# charged = 900 + 15 + 576 = 1491 毫元(in 5 折 / out 8 折;分项 round 后求和)
|
||
assert u["upstream_cost_milli"] == 2550
|
||
assert u["charged_milli"] == 1491
|
||
assert u["margin_milli"] == u["charged_milli"] - u["upstream_cost_milli"]
|
||
assert u["in_miss_tok"] == 600000 and u["in_hit_tok"] == 300000
|
||
assert u["out_tok"] == 80000
|
||
finally:
|
||
upmod._client = orig
|
||
try:
|
||
asyncio_run(hclient.aclose())
|
||
except Exception:
|
||
pass
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_chat_stream_end_to_end(tmp_path):
|
||
"""流式全程:SSE 逐块、usage chunk 被过滤(客户端未要求)、终态 DONE。"""
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
with tc.stream("POST", "/proxy/v1/chat/completions",
|
||
json={**BODY, "stream": True},
|
||
headers=_auth(key)) as resp:
|
||
assert resp.status_code == 200
|
||
assert resp.headers["content-type"].startswith("text/event-stream")
|
||
lines = [l for l in resp.iter_lines() if l.strip()]
|
||
assert lines[-1] == "data: [DONE]"
|
||
text = "\n".join(lines)
|
||
assert "usage" not in text # 客户端未要求 -> 过滤
|
||
assert "\u4f60\u597d" in text
|
||
rid = resp.headers["x-request-id"]
|
||
u = ledger.get_usage(rid)
|
||
assert u["status"] == "ok" and u["charged_milli"] > 0
|
||
finally:
|
||
upmod._client = orig
|
||
try:
|
||
asyncio_run(hclient.aclose())
|
||
except Exception:
|
||
pass
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_402_when_balance_insufficient(tmp_path):
|
||
"""余额不足 -> 402(预扣失败)。"""
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
sid = ledger.list_usage() and None
|
||
# 直接清空该 key 学生余额
|
||
row = ledger.find_key(key["key"] and __import__("hashlib").sha256(
|
||
key["key"].encode()).hexdigest())
|
||
ledger.topup(row["student_id"], -10.0) # 余额归零
|
||
r = tc.post("/proxy/v1/chat/completions", json=BODY, headers=_auth(key))
|
||
assert r.status_code == 402
|
||
assert calls["n"] == 0 # 未打上游
|
||
finally:
|
||
upmod._client = orig
|
||
try:
|
||
asyncio_run(hclient.aclose())
|
||
except Exception:
|
||
pass
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_401_invalid_key(tmp_path):
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
r = tc.post("/proxy/v1/chat/completions", json=BODY,
|
||
headers={"Authorization": "Bearer sk-campus-wrong"})
|
||
assert r.status_code == 401
|
||
finally:
|
||
upmod._client = orig
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_429_day_cap(tmp_path):
|
||
"""日请求上限 -> 429(配置 cap=2)。"""
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
ok = 0
|
||
for _ in range(3):
|
||
r = tc.post("/proxy/v1/chat/completions", json=BODY, headers=_auth(key))
|
||
if r.status_code == 200:
|
||
ok += 1
|
||
assert ok >= 1
|
||
finally:
|
||
upmod._client = orig
|
||
try:
|
||
asyncio_run(hclient.aclose())
|
||
except Exception:
|
||
pass
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_413_body_too_large(tmp_path):
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
big = {"model": "deepseek-chat",
|
||
"messages": [{"role": "user", "content": "x" * 7000}]}
|
||
r = tc.post("/proxy/v1/chat/completions", json=big, headers=_auth(key))
|
||
assert r.status_code == 413
|
||
finally:
|
||
upmod._client = orig
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_502_upstream_fail_voids_hold(tmp_path):
|
||
"""上游全挂 -> 502 且预扣全额退(余额不变)。"""
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path,
|
||
fail_upstream=True)
|
||
try:
|
||
row = ledger.find_key(__import__("hashlib").sha256(
|
||
key["key"].encode()).hexdigest())
|
||
before = ledger.get_student(row["student_id"])["balance_milli"]
|
||
r = tc.post("/proxy/v1/chat/completions", json=BODY, headers=_auth(key))
|
||
assert r.status_code == 502
|
||
after = ledger.get_student(row["student_id"])["balance_milli"]
|
||
assert before == after # void 全退
|
||
finally:
|
||
upmod._client = orig
|
||
try:
|
||
asyncio_run(hclient.aclose())
|
||
except Exception:
|
||
pass
|
||
mp.reset_pool()
|
||
reset_auth_state()
|
||
|
||
|
||
def test_openai_protocol_compliance_via_httpx(tmp_path):
|
||
"""httpx 手写 OpenAI SDK 合规断言:流式可被标准解析器消费、形状正确。"""
|
||
tc, ledger, key, calls, (upmod, orig, hclient) = _make_app(tmp_path)
|
||
try:
|
||
# 模拟 OpenAI SDK 的流式解析路径
|
||
with tc.stream("POST", "/proxy/v1/chat/completions",
|
||
json={**BODY, "stream": True,
|
||
"stream_options": {"include_usage": True}},
|
||
headers=_auth(key)) as resp:
|
||
collected = []
|
||
usage_seen = False
|
||
done = False
|
||
for line in resp.iter_lines():
|
||
if not line.startswith("data:"):
|
||
continue
|
||
payload = line[5:].strip()
|
||
if payload == "[DONE]":
|
||
done = True
|
||
continue
|
||
obj = json.loads(payload)
|
||
assert obj.get("object") == "chat.completion.chunk"
|
||
for choice in obj.get("choices", []):
|
||
collected.append(choice.get("delta", {}).get("content") or "")
|
||
if obj.get("usage"):
|
||
usage_seen = True
|
||
assert done and usage_seen
|
||
assert "".join(collected) == "你好,世界"
|
||
finally:
|
||
upmod._client = orig
|
||
try:
|
||
asyncio_run(hclient.aclose())
|
||
except Exception:
|
||
pass
|
||
mp.reset_pool()
|
||
reset_auth_state()
|