"""层 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")