Files
projectAIpopular/tests/test_pipeline_tier_fn.py
T
tzt 204789aed6 feat(sense): T-G6 live 分流(pipeline tier_fn 三档钩子,D-G5 唯一受控改动)
- pipeline.py:+tier_fn 可选构造参数(None = 现行为逐字节不变);
  T1=fast path 入口 / T2=单次云端直答(跳过 brief/循环/review,失败落回
  完整管线 D-G1 阶梯)/ T3=强制完整管线(跳过快路径);同时兼容
  直接可调用与带 .decide 的 grader 对象
- 测试 +5:None 零回归 / T1 fast path / T2 直答跳过 brief/loop / T3 完整管线 /
  T2 失败落回,全量 410 passed(pipeline 既有测试零回归)
2026-09-05 14:25:46 +08:00

162 lines
5.9 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.
"""tier_fn 钩子测试(T-G6):None 零回归 / T1 / T2 单次直答 / T3 完整管线。"""
import sys
from pathlib import Path
import pytest
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from router_system.pipeline import CollaborativePipeline # noqa: E402
from router_system.architect import ArchitectClient # noqa: E402
from router_system.worker import WorkerLoop # noqa: E402
def asyncio_run(coro):
import asyncio
return asyncio.run(coro)
class FakeArchitect(ArchitectClient):
"""可脚本化的假架构师:记录 brief/直答调用。"""
def __init__(self, calls):
super().__init__(model="fake", api_key="k")
self.calls = calls
async def brief(self, query, ws):
self.calls.append("brief")
return {"goal": "g", "constraints": [], "tags": ["code"],
"acceptance": [], "plan": [{"id": "s1", "task": "实现加法",
"deps": [], "done_criteria": "c"}]}
async def decide(self, ws):
self.calls.append("decide")
return {"reply": "hint", "patch_plan": []}
async def final_review(self, ws):
self.calls.append("review")
return {"verdict": "done", "notes": "", "fix_issues": []}
async def _chat_once(self, ws, messages):
self.calls.append("chat_once")
return "直接答案"
class ScriptedWorker(WorkerLoop):
"""脚本化小模型:先给坏答案(触发自修/升级),再给好答案。"""
def __init__(self, calls):
import asyncio
calls_worker = calls
async def gen(prompt):
calls_worker.append("gen")
if len([c for c in calls_worker if c == "gen"]) == 1:
return "不够长的占位"
return ("```python\ndef add(a, b):\n"
" if not isinstance(a, int) or not isinstance(b, int):\n"
" raise ValueError\n"
" return a + b\n\n"
"def test_add():\n"
" assert add(2, 3) == 5\n```")
super().__init__(generate=gen, max_fix_attempts=2, model_used="fake-worker")
def _pipeline(calls, tier_fn=None, fast_path=True):
return CollaborativePipeline(architect=FakeArchitect(calls),
worker=ScriptedWorker(calls),
fast_path=fast_path, tier_fn=tier_fn)
def test_tier_fn_none_unchanged(tmp_path):
"""D-G5 零回归:tier_fn=None 现行为(快路径 miss -> 完整管线)。"""
calls = []
pipe = _pipeline(calls)
r = asyncio_run(pipe.run("讲讲排序", request_id="tf-none"))
assert "brief" in r.route and "reviewing" in r.route # 完整管线
assert r.status == "done"
def test_tier_fn_t1_fast_path(tmp_path):
"""T1 -> 入口为 fast path(本地直答命中)。"""
calls = []
class GoodWorker(WorkerLoop):
def __init__(self):
async def gen(prompt):
return "排序算法是把一组数据按照特定顺序重新排列的过程,广泛用于数据处理。"
super().__init__(generate=gen, max_fix_attempts=2, model_used="fake-worker")
async def tier_fn(query):
return "T1"
pipe = CollaborativePipeline(architect=FakeArchitect(calls),
worker=GoodWorker(), fast_path=True, tier_fn=tier_fn)
r = asyncio_run(pipe.run("讲讲排序", request_id="tf-t1"))
assert r.status == "fast_path"
assert any("fast_path@domain" in x for x in r.route)
def test_tier_fn_t2_single_cloud_direct(tmp_path):
"""T2 新档:单次云端直答(chat_once),跳过 brief/循环/review。"""
calls = []
class DirectArch(FakeArchitect):
async def _chat_once(self, ws, messages):
self.calls.append("chat_once")
return "T2 直答结果"
async def tier_fn(query):
return "T2"
pipe = CollaborativePipeline(architect=DirectArch(calls),
worker=ScriptedWorker(calls),
fast_path=True, tier_fn=tier_fn)
r = asyncio_run(pipe.run("实现一个函数", request_id="tf-t2"))
assert r.status == "fast_path" # T2 复用 fast_path 状态位
assert "t2:direct" in r.route
assert "brief" not in r.route # 跳过 brief
assert "loop" not in r.route # 跳过协作循环
assert r.model_used == "fake" # 云端模型
def test_tier_fn_t3_full_pipeline(tmp_path):
"""T3 -> 强制完整管线(跳过 fast path,走 brief/循环/review)。"""
calls = []
class DirectArch(FakeArchitect):
async def _chat_once(self, ws, messages):
self.calls.append("chat_once")
return "直接答案"
async def tier_fn(query):
return "T3"
pipe = CollaborativePipeline(architect=DirectArch(calls),
worker=ScriptedWorker(calls),
fast_path=True, tier_fn=tier_fn)
r = asyncio_run(pipe.run("实现加法函数", request_id="tf-t3"))
assert r.status in ("done", "escalated")
assert "brief" in r.route and "reviewing" in r.route
assert "fast_path@" not in " ".join(r.route) # 不走快路径
def test_tier_fn_t2_direct_miss_falls_back(tmp_path):
"""T2 直答失败(上游异常)-> 落回完整管线(D-G1 阶梯)。"""
calls = []
class FailArch(FakeArchitect):
async def _chat_once(self, ws, messages):
raise RuntimeError("云端不可用")
async def tier_fn(query):
return "T2"
pipe = CollaborativePipeline(architect=FailArch(calls),
worker=ScriptedWorker(calls),
fast_path=False, tier_fn=tier_fn)
r = asyncio_run(pipe.run("实现加法函数", request_id="tf-fb"))
# 直答失败 -> 落回完整管线(brief/loop/review
assert "brief" in r.route or r.status == "failed"