Files
group_xinghuo_jinrong/tests/test_main.py
T
GaoYiYuan_0626 01ec5fce32 feat(chat): 新增 SSE 流式对话端点 POST /api/chat/stream(方案 C)
契约(OpenAI 兼容 chunk,AI SDK / fetch-event-source 可直接接):
首帧 meta(session_id/trace_id/disclaimer)→ delta → finish_reason=stop → [DONE]。
前端侧:免责声明由首帧下发、前端常驻渲染;落库文本仍按原口径拼尾部。

设计拍板:
1. 新增独立端点,原 POST /api/chat 契约与既有测试零影响;
2. 鉴权/限流/输入防护全部在返回 StreamingResponse 之前完成(SSE 一开就改不了
   状态码),401/403/404/409/429/400 仍是普通 JSON;
3. 整轮一次性落库:中途异常/断连不落消息(Tool 留痕已落可审计),不产生
   半截内容污染历史窗口。

实现:
- agent_service:抽 needs_disclaimer/_degraded_reply/_base_state;新增 stream_chat
  生成器(Tool 节点同步跑完再推 LLM 文本,无 key 走降级整块);
- api/chat:抽 _guard_request(准入→空白→限流→注入拦截)与 _prepare_turn
  (归属+会话解析/创建),同步与流式共用,守卫零偏差;
- session_repository:新增 insert_turn——user+assistant 同事务落库 + 事务内
  取 seq,修掉评审 P0(落库失败会半截落且前端收不到 [DONE] 挂起)与 P1
  (两次写非原子);同步端点一并改用。

测试:新增 9 例(契约/免责/降级/异常不落库/落库失败/鉴权边界/超长/closed/Tool
留痕),全量 pytest 494→503 绿。遗留:无心跳帧(长生成空隙靠反代 timeout 配置),
断连留空会话待清理策略。
2026-09-08 14:23:33 +08:00

124 lines
4.3 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.
"""main 集成冒烟(B7):路由挂载 / trace 中间件 / lifespan 启动期校验 / 统一错误体。
main app 全路由经 TestClient 走真实 lifespan(dev 环境);仓储 monkeypatch
注入 sqlite(audit_log 供 401 留痕);Redis 网关注入 fake(不依赖本机 Redis)。
全链路 trace 一致性与集成测试归 B8 conftest,此处只验中间件行为本身。
"""
import re
import pytest
from fastapi.testclient import TestClient
from sqlalchemy import text
from _ddl import create_sqlite_engine
from app.config.settings import settings
from app.main import app
from app.repository.risk_repository import RiskRepository
from app.service.risk import redis_gateway
TRACE_HEADER_PATTERN = r"^trc-[0-9a-f]{16}$"
class FakeGateway:
def __init__(self):
self.messages = []
self.deletes = []
def publish(self, channel, payload):
self.messages.append((channel, payload))
def delete(self, *keys):
self.deletes.append(keys)
@pytest.fixture()
def client(monkeypatch):
engine = create_sqlite_engine() # DDL 单一事实源(B4 评审 P3-12)
repo = RiskRepository(engine=engine)
from app.api import deps as deps_mod
from app.api import risk as risk_api
from app.api import simulate as simulate_mod
monkeypatch.setattr(risk_api, "_repo", lambda: repo)
monkeypatch.setattr(simulate_mod, "_repo", lambda: repo)
monkeypatch.setattr(deps_mod, "RiskRepository", lambda: repo)
with TestClient(app) as c: # lifespan:dev 放行;Redis 惰性连接不触网
monkeypatch.setattr(redis_gateway, "_gateway", FakeGateway())
yield c
engine.dispose()
def test_health(client):
r = client.get("/health")
assert r.status_code == 200 and r.json()["status"] == "ok"
def test_all_routers_mounted(client):
paths = client.get("/openapi.json").json()["paths"]
assert set(paths) == {
"/health",
"/api/risk/alerts",
"/api/risk/alerts/{alert_id}/handle",
"/api/risk/suitability/check",
"/api/risk/aml/scan",
"/api/simulate/trade",
"/api/chat",
# 方案 B:前端拉侧三端点
"/api/chat/sessions",
"/api/chat/sessions/{session_id}/messages",
"/api/chat/sessions/{session_id}/close",
# 方案 C:SSE 流式对话
"/api/chat/stream",
}
def test_trace_header_generated(client):
r = client.get("/health")
assert r.headers["X-Trace-Id"] and re.fullmatch(TRACE_HEADER_PATTERN, r.headers["X-Trace-Id"])
def test_trace_header_passthrough(client):
tid = "trc-abc123def45678"
assert client.get("/health", headers={"X-Trace-Id": tid}).headers["X-Trace-Id"] == tid
def test_trace_header_invalid_regenerated(client):
bad = "bad id!"
tid = client.get("/health", headers={"X-Trace-Id": bad}).headers["X-Trace-Id"]
assert tid != bad and re.fullmatch(TRACE_HEADER_PATTERN, tid)
def test_unified_error_body_401_with_trace(client):
"""手册 §10 错误体四键 + trace_id/request_id 分别对齐响应头(T-02 独立双 ID)。"""
r = client.get("/api/risk/alerts") # 无 debug 头
assert r.status_code == 401
body = r.json()
assert body["error_code"] == "AUTH_401_MISSING_DEBUG_HEADERS"
assert body["message"]
assert set(body) == {"error_code", "message", "trace_id", "request_id"}
assert body["trace_id"] == r.headers["X-Trace-Id"]
assert body["request_id"] == r.headers["X-Request-Id"]
assert body["trace_id"] != body["request_id"]
def test_lifespan_rejects_debug_factory_in_non_dev(monkeypatch):
"""挂账⑤(T-01 后兜底分支):debug 工厂在场 + 非 dev → 拒绝启动。"""
monkeypatch.setattr(settings, "app_env", "production")
from app.api import deps as deps_mod
monkeypatch.setattr(deps_mod, "AUTH_FACTORY_IS_DEBUG", True)
with pytest.raises(RuntimeError, match="debug auth factory is wired"):
with TestClient(app):
pass
def test_lifespan_rejects_jwt_not_ready_in_non_dev(monkeypatch):
"""T-01:非 dev 且 RS256 公钥未配置(HS256 dev secret)→ 拒绝启动。"""
monkeypatch.setattr(settings, "app_env", "production")
monkeypatch.setattr(settings, "jwt_public_key_path", "")
with pytest.raises(RuntimeError, match="JWT auth not ready"):
with TestClient(app):
pass