Files
group_fqcd_jr/tools/smoke_check.py
T

151 lines
5.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""底座接口冒烟测试:认证、越权、幂等、SSE 传输。
用法:先启动 `python -m uvicorn app.main:app --port 8099`,再运行本脚本。
"""
from __future__ import annotations
import datetime as dt
import uuid
from pathlib import Path
import httpx
import jwt
BASE_URL = "http://127.0.0.1:8099"
PRIVATE_KEY = Path("config/jwt/jwt-private.pem").read_text(encoding="utf-8")
_results: list[tuple[str, str, str]] = []
def record(name: str, expected: str, actual: str) -> None:
verdict = "PASS" if expected == actual else "FAIL"
_results.append((verdict, name, f"期望 {expected} / 实际 {actual}"))
def token(sub: str) -> str:
now = dt.datetime.now(dt.UTC)
return jwt.encode(
{
"sub": sub,
"iss": "jr-local",
"aud": "jr-agent-platform",
"exp": now + dt.timedelta(minutes=30),
"nbf": now - dt.timedelta(seconds=5),
"jti": str(uuid.uuid4()),
},
PRIVATE_KEY,
algorithm="RS256",
)
def headers(sub: str) -> dict[str, str]:
return {"Authorization": f"Bearer {token(sub)}"}
def payload(
agent_type: str,
key: str | None = None,
message: str = "冒烟测试消息",
session_id: str = "smoke-session",
) -> dict:
return {
"agent_type": agent_type,
"message": message,
"session_id": session_id,
"idempotency_key": key or uuid.uuid4().hex,
}
def main() -> None:
# trust_env=False:绕开 Windows 注册表代理设置,直连本地服务
with httpx.Client(base_url=BASE_URL, timeout=10, trust_env=False) as client:
r = client.post("/api/v1/agent-runs", json=payload("customer_service"))
record("无 token 调用", "401", str(r.status_code))
r = client.post(
"/api/v1/agent-runs",
json=payload("customer_service"),
headers={"Authorization": "Bearer not-a-jwt"},
)
record("无效 token", "401", str(r.status_code))
r = client.post(
"/api/v1/agent-runs",
json=payload("customer_service"),
headers=headers("9001"),
)
record("合法用户创建 customer_service run", "202", str(r.status_code))
run_id = r.json().get("run_id") if r.status_code == 202 else None
r = client.post("/api/v1/agent-runs", json=payload("risk"), headers=headers("9001"))
record("普通客户调用 risk Agent(越权)", "403", str(r.status_code))
if r.status_code == 202:
risk_run_id = r.json().get("run_id")
detail = client.get(f"/api/v1/agent-runs/{risk_run_id}", headers=headers("9001"))
record(
"越权 run 是否真的落库",
"未落库",
f"已落库 agent_type={detail.json().get('agent_type')}"
if detail.status_code == 200
else f"GET {detail.status_code}",
)
r = client.post(
"/api/v1/agent-runs", json=payload("no_such_agent"), headers=headers("9001")
)
record("未注册 agent_type", "400", str(r.status_code))
r = client.post(
"/api/v1/agent-runs", json=payload("customer_service"), headers=headers("not-a-number")
)
record("sub 非数字(应 4xx 而非 500)", "400", str(r.status_code))
if run_id:
r = client.get(f"/api/v1/agent-runs/{run_id}", headers=headers("9002"))
record("用户 2 读取用户 1 的 run", "404", str(r.status_code))
sse_events: list[str] = []
try:
with client.stream(
"GET",
f"/api/v1/agent-runs/{run_id}/events",
headers=headers("9001"),
timeout=8,
) as stream:
content_type = stream.headers.get("content-type", "").split(";")[0]
record("SSE Content-Type", "text/event-stream", content_type)
for line in stream.iter_lines():
if line.startswith("event:"):
sse_events.append(line.split(":", 1)[1].strip())
if len(sse_events) >= 2:
break
except httpx.TimeoutException:
pass
record("SSE 首个事件", "start", sse_events[0] if sse_events else "8 秒内无事件")
key = uuid.uuid4().hex
idem_session = f"smoke-idem-{key[:12]}"
client.post(
"/api/v1/agent-runs",
json=payload("customer_service", key, session_id=idem_session),
headers=headers("9001"),
)
r = client.post(
"/api/v1/agent-runs",
json=payload("customer_service", key, message="不同内容", session_id=idem_session),
headers=headers("9001"),
)
record("同幂等键不同请求体", "409", str(r.status_code))
width = max(len(name) for _, name, _ in _results)
print()
for verdict, name, detail in _results:
print(f"[{verdict}] {name.ljust(width)} {detail}")
failed = sum(1 for verdict, _, _ in _results if verdict == "FAIL")
print(f"\n合计 {len(_results)} 项,失败 {failed} 项")
if __name__ == "__main__":
main()