"""层 1 / T2:teacher 批量生成与缓存单测。
用假 OpenAI 客户端注入(TeacherClient 的测试口),验证思考段剥离、缓存契约
(与 data.attach_teacher_completions 的端到端闭环)、断点续传、失败汇总。
"""
from types import SimpleNamespace
import pytest
from datasets import Dataset
from ars_opd.configs import TeacherGenConfig
from ars_opd.data import attach_teacher_completions, prompt_key
from ars_opd.teacher import TeacherClient, _load_teacher_env, generate_completions
class FakeClient:
"""最小 OpenAI 客户端替身:chat.completions.create 按 responder 出内容。"""
def __init__(self, responder):
self.calls = []
self._responder = responder
self.chat = SimpleNamespace(completions=SimpleNamespace(create=self._create))
def _create(self, model, messages, **kwargs):
self.calls.append(messages)
content = self._responder(messages)
return SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content=content))]
)
def make_teacher(responder, **cfg_overrides):
cfg = TeacherGenConfig(**cfg_overrides)
return TeacherClient(cfg, client=FakeClient(responder), model="fake-m3")
def user(q):
return [{"role": "user", "content": q}]
# ---------------------------------------------------------------------------
# TeacherClient.generate
# ---------------------------------------------------------------------------
def test_剥离开头思考段():
teacher = make_teacher(lambda m: "心算一下\n答案是 42")
assert teacher.generate(user("q")) == "答案是 42"
def test_正文中的think字样不误删():
teacher = make_teacher(lambda m: "x正文提到 标签本身")
assert teacher.generate(user("q")) == "正文提到 标签本身"
def test_只剩思考段等于空解答_报错():
teacher = make_teacher(lambda m: "思考被截断在半途")
# 未闭合的 think 段剥不掉,但闭合后为空的要报错
teacher_empty = make_teacher(lambda m: "只有思考 ")
with pytest.raises(ValueError, match="空解答"):
teacher_empty.generate(user("q"))
# 未闭合时保留原文(宁可保留可疑内容也不静默删成空)
assert "" in teacher.generate(user("q"))
def test_关闭strip_think则原样保留():
teacher = make_teacher(lambda m: "ab", strip_think=False)
assert teacher.generate(user("q")) == "ab"
def test_system_prompt前置():
teacher = make_teacher(lambda m: "ok", system_prompt="你是数学助教")
teacher.generate(user("q"))
sent = teacher.client.calls[0]
assert sent[0] == {"role": "system", "content": "你是数学助教"}
assert sent[1]["role"] == "user"
def test_注入client但不给model报错():
with pytest.raises(ValueError, match="model"):
TeacherClient(TeacherGenConfig(), client=FakeClient(lambda m: "x"), model=None)
# ---------------------------------------------------------------------------
# generate_completions:缓存契约与断点续传
# ---------------------------------------------------------------------------
def test_端到端契约_生成的缓存能被attach消费(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("1+1=?"), user("2+2=?")]
teacher = make_teacher(lambda m: f"对「{m[-1]['content']}」的解答")
generate_completions(prompts, cache, teacher)
ds = Dataset.from_list([{"messages": p} for p in prompts])
out = attach_teacher_completions(ds, cache)
assert out[0]["messages"][-1]["content"] == "对「1+1=?」的解答"
assert out[1]["messages"][-1]["content"] == "对「2+2=?」的解答"
def test_断点续传_已缓存的不重新生成(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("q1"), user("q2")]
teacher = make_teacher(lambda m: "a")
generate_completions([prompts[0]], cache, teacher)
assert len(teacher.client.calls) == 1
generate_completions(prompts, cache, teacher) # q1 命中缓存
assert len(teacher.client.calls) == 2 # 只多了 q2 一次调用
def test_单条失败_其余落盘_结束时汇总报错(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("好题"), user("坏题")]
def responder(m):
if m[-1]["content"] == "坏题":
raise RuntimeError("网关 500")
return "解答"
teacher = make_teacher(responder)
with pytest.raises(RuntimeError, match="1/2"):
generate_completions(prompts, cache, teacher)
# 成功的那条已经在缓存里,重跑只会补坏题
from ars_opd.teacher import _cached_keys
from pathlib import Path
assert _cached_keys(Path(cache)) == {prompt_key(prompts[0])}
# ---------------------------------------------------------------------------
# .env 读取
# ---------------------------------------------------------------------------
def test_env缺失显式报错(monkeypatch):
for name in ("TEACHER_API_BASE", "TEACHER_API_KEY", "TEACHER_MODEL"):
monkeypatch.delenv(name, raising=False)
with pytest.raises(ValueError, match="TEACHER_API_BASE"):
_load_teacher_env(env_file="/不存在的路径/.env")