# -*- coding: utf-8 -*- """移动端交接班摘要生成单元测试(issue #53 / PRD 场景C / 5.4「④ LLM 网关」)。 覆盖: 1. 配置校验:合法配置通过、默认值与字段别名、sections 去重与枚举校验、 maxEvents/maxTodos/collapseThreshold/fontSize/requireConfirm 各类非法情况; 2. ``load_handover_config`` 校验失败抛 ``HandoverConfigError`` 并携带全部错误; 3. ``normalize_shift_record``:合法 dict → ShiftRecord;缺必填/非法结构被拒; 4. ``build_llm_input``:占位符填充正确、含班次/事件/能耗/待办; 5. ``render_deterministic_brief``:章节与配置 sections 联动; 6. ``generate_handover_brief``:注入 LLM 用其输出;LLM 抛异常/返回空 → 自动降级; 7. ``render_handover_brief``:props 结构正确、截断/溢出计数/高利害强制确认; 8. ``within_generation_budget``:2 分钟验收线判定。 """ from __future__ import annotations import os import sys import unittest sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import _bootstrap # noqa: F401 把 core/shift-handover 挂载为 shift_handover 包 from shift_handover.handover import ( # type: ignore[import-not-found] DEFAULT_PROMPT_TEMPLATE, DEFAULT_PROMPT_VERSION, HANDOVER_CONFIG_SCHEMA_ID, HANDOVER_GENERATION_BUDGET_MS, HandoverBriefConfig, HandoverConfigError, HandoverConfigValidationResult, ShiftEvent, ShiftRecord, ShiftTodo, build_llm_input, generate_handover_brief, load_handover_config, normalize_shift_record, render_deterministic_brief, render_handover_brief, validate_handover_config, within_generation_budget, ) def _valid_config() -> dict: """返回一份合法的交接班摘要配置(Ti 模板风格)。""" return { "$schema": HANDOVER_CONFIG_SCHEMA_ID, "promptTemplate": "shift_handover", "promptVersion": "1.0.1", "sections": ["overview", "abnormal", "safety", "energy", "todos"], "maxEvents": 8, "maxTodos": 5, "collapseThreshold": 4, "fontSize": "md", "requireConfirm": False, } def _full_record_dict() -> dict: return { "shift": "夜班 2026-08-05 00:00~08:00", "operator": "张工", "overview": "TiCl4 产量 36.2t,质量达成率 98.6%。", "events": [ {"time": "01:20", "title": "炉层温度越上限", "severity": "P1", "detail": "CLF-01 第3层 942℃ → 已调风量"}, {"time": "03:05", "title": "夜巡正常", "severity": "info"}, ], "todos": [ {"title": "白班复测 3 层温度趋势", "priority": "high", "due": "接班后 1h"}, {"title": "补录 LIMS 03:00 批次", "priority": "medium"}, ], "energy": "总电耗 12.4 万 kWh", "safety": "夜班无安全事件;注意 3 层高温区巡检。", "next_shift": "李工", } # --------------------------------------------------------------------------- # 1) 配置校验 # --------------------------------------------------------------------------- class TestValidateConfig(unittest.TestCase): def test_valid_config_passes(self): res = validate_handover_config(_valid_config()) self.assertTrue(res.ok, msg=res.errors) self.assertEqual(res.errors, []) self.assertEqual(res.normalized["$schema"], HANDOVER_CONFIG_SCHEMA_ID) def test_defaults_and_aliases(self): # snake_case 与 camelCase 都接受;缺省给默认值 cfg = {"prompt_template": "shift_handover", "prompt_version": "1.0.0"} res = validate_handover_config(cfg) self.assertTrue(res.ok, msg=res.errors) n = res.normalized self.assertEqual(n["maxEvents"], 8) self.assertEqual(n["maxTodos"], 5) self.assertEqual(n["fontSize"], "md") self.assertEqual(n["collapseThreshold"], 4) self.assertFalse(n["requireConfirm"]) # sections 默认为全部合法章节 self.assertEqual( n["sections"], ["overview", "abnormal", "safety", "energy", "todos"], ) def test_wrong_schema_rejected(self): cfg = _valid_config() cfg["$schema"] = "something-else" res = validate_handover_config(cfg) self.assertFalse(res.ok) self.assertTrue(any("$schema" in e for e in res.errors)) def test_non_dict_rejected(self): res = validate_handover_config(["not", "a", "dict"]) self.assertFalse(res.ok) self.assertEqual(len(res.errors), 1) def test_sections_validation(self): # 非法章节、重复、空列表、非列表 cases = [ (["overview", "unknown"], "非法枚举"), (["overview", "overview"], "重复"), ([], "空列表"), ("not-a-list", "非列表"), ] for bad_sections, hint in cases: cfg = _valid_config() cfg["sections"] = bad_sections res = validate_handover_config(cfg) self.assertFalse(res.ok, msg=f"{hint} 应被拒绝: {res.errors}") def test_numeric_bounds(self): for field, bad in [ ("maxEvents", 0), ("maxEvents", 51), ("maxEvents", True), ("maxTodos", 0), ("maxTodos", 31), ("maxTodos", "5"), ("collapseThreshold", -1), ("collapseThreshold", True), ]: cfg = _valid_config() cfg[field] = bad res = validate_handover_config(cfg) self.assertFalse(res.ok, msg=f"{field}={bad!r} 应被拒绝: {res.errors}") def test_font_size_enum(self): cfg = _valid_config() cfg["fontSize"] = "xl" res = validate_handover_config(cfg) self.assertFalse(res.ok) def test_require_confirm_type(self): cfg = _valid_config() cfg["requireConfirm"] = "yes" res = validate_handover_config(cfg) self.assertFalse(res.ok) def test_aggregates_multiple_errors(self): bad = { "promptTemplate": "", "sections": ["oops"], "maxEvents": 0, "fontSize": "xxl", "requireConfirm": 1, } res = validate_handover_config(bad) self.assertFalse(res.ok) # 至少 5 条字段级错误被聚合 self.assertGreaterEqual(len(res.errors), 5) # --------------------------------------------------------------------------- # 2) load_handover_config # --------------------------------------------------------------------------- class TestLoadConfig(unittest.TestCase): def test_load_valid(self): cfg = load_handover_config(_valid_config()) self.assertIsInstance(cfg, HandoverBriefConfig) self.assertEqual(cfg.prompt_template, "shift_handover") self.assertEqual(cfg.prompt_version, "1.0.1") self.assertEqual(cfg.font_size, "md") def test_load_invalid_raises_with_errors(self): with self.assertRaises(HandoverConfigError) as ctx: load_handover_config({"promptTemplate": "", "maxEvents": 0}) # 异常携带错误清单 self.assertGreaterEqual(len(ctx.exception.errors), 2) # --------------------------------------------------------------------------- # 3) normalize_shift_record # --------------------------------------------------------------------------- class TestNormalizeRecord(unittest.TestCase): def test_full_record(self): rec = normalize_shift_record(_full_record_dict()) self.assertIsInstance(rec, ShiftRecord) self.assertEqual(rec.shift, "夜班 2026-08-05 00:00~08:00") self.assertEqual(rec.operator, "张工") self.assertEqual(len(rec.events), 2) self.assertEqual(rec.events[0].severity, "P1") self.assertEqual(rec.todos[0].priority, "high") self.assertEqual(rec.next_shift, "李工") def test_minimal_record(self): rec = normalize_shift_record({"shift": "白班", "operator": "王五"}) self.assertEqual(rec.events, []) self.assertEqual(rec.todos, []) self.assertEqual(rec.overview, "") def test_missing_required_rejected(self): with self.assertRaises(HandoverConfigError): normalize_shift_record({"shift": "白班"}) # 缺 operator def test_bad_event_rejected(self): raw = _full_record_dict() raw["events"][0] = {"time": "01:00"} # 缺 title with self.assertRaises(HandoverConfigError): normalize_shift_record(raw) def test_bad_todo_rejected(self): raw = _full_record_dict() raw["todos"] = [{"priority": "high"}] # 缺 title with self.assertRaises(HandoverConfigError): normalize_shift_record(raw) def test_non_dict_rejected(self): with self.assertRaises(HandoverConfigError): normalize_shift_record("not a dict") # --------------------------------------------------------------------------- # 4) build_llm_input # --------------------------------------------------------------------------- class TestBuildLLMInput(unittest.TestCase): def test_renders_prompt_with_placeholders(self): rec = normalize_shift_record(_full_record_dict()) name, text = build_llm_input(rec) self.assertEqual(name, DEFAULT_PROMPT_TEMPLATE) # 正文以 shift_handover 提示词开头 self.assertTrue(text.startswith("生成交接班摘要")) # 班次与关键字段都被填入 {query} self.assertIn("夜班 2026-08-05", text) self.assertIn("张工", text) self.assertIn("炉层温度越上限", text) self.assertIn("总电耗", text) self.assertIn("复测 3 层温度趋势", text) def test_respects_config_template_name(self): rec = normalize_shift_record({"shift": "夜班", "operator": "张工"}) cfg = HandoverBriefConfig(prompt_template="custom_tmpl") name, _ = build_llm_input(rec, cfg) self.assertEqual(name, "custom_tmpl") # --------------------------------------------------------------------------- # 5) render_deterministic_brief # --------------------------------------------------------------------------- class TestDeterministicBrief(unittest.TestCase): def test_contains_all_sections(self): rec = normalize_shift_record(_full_record_dict()) text = render_deterministic_brief(rec) self.assertIn("交接班摘要", text) self.assertIn("生产概况", text) self.assertIn("异常事项", text) self.assertIn("安全注意事项", text) self.assertIn("能耗", text) self.assertIn("待办", text) def test_sections_filter(self): rec = normalize_shift_record(_full_record_dict()) cfg = HandoverBriefConfig(sections=["overview"]) text = render_deterministic_brief(rec, cfg) self.assertIn("生产概况", text) # 只开 overview,异常/能耗等章节不应出现为标题 self.assertNotIn("## 异常事项", text) self.assertNotIn("## 能耗", text) # --------------------------------------------------------------------------- # 6) generate_handover_brief(LLM 注入 + 降级) # --------------------------------------------------------------------------- class TestGenerateBrief(unittest.TestCase): def test_uses_llm_output(self): rec = normalize_shift_record(_full_record_dict()) def llm(prompt: str) -> str: return "LLM 摘要:本班平稳,3 层温度曾越限已处置。" out = generate_handover_brief(rec, llm_generate=llm) self.assertEqual(out, "LLM 摘要:本班平稳,3 层温度曾越限已处置。") def test_falls_back_when_no_llm(self): rec = normalize_shift_record(_full_record_dict()) out = generate_handover_brief(rec) # llm_generate=None self.assertIn("交接班摘要", out) self.assertIn("炉层温度越上限", out) def test_falls_back_on_llm_exception(self): rec = normalize_shift_record(_full_record_dict()) def broken(prompt: str) -> str: raise RuntimeError("LLM 网关不可达") out = generate_handover_brief(rec, llm_generate=broken) # LLM 故障 → 降级为确定性摘要 self.assertIn("交接班摘要", out) def test_falls_back_on_empty_llm_output(self): rec = normalize_shift_record(_full_record_dict()) def empty(prompt: str) -> str: return " " out = generate_handover_brief(rec, llm_generate=empty) self.assertIn("交接班摘要", out) # --------------------------------------------------------------------------- # 7) render_handover_brief(移动端 props) # --------------------------------------------------------------------------- class TestRenderBriefProps(unittest.TestCase): def test_props_structure(self): rec = normalize_shift_record(_full_record_dict()) props = render_handover_brief(rec) self.assertEqual(props["schema"], HANDOVER_CONFIG_SCHEMA_ID) self.assertEqual(props["shift"], "夜班 2026-08-05 00:00~08:00") self.assertEqual(props["operator"], "张工") self.assertEqual(props["nextShift"], "李工") self.assertEqual(props["promptTemplate"], DEFAULT_PROMPT_TEMPLATE) self.assertEqual(props["promptVersion"], DEFAULT_PROMPT_VERSION) self.assertEqual(len(props["events"]), 2) self.assertEqual(len(props["todos"]), 2) self.assertEqual(props["eventOverflow"], 0) self.assertEqual(props["todoOverflow"], 0) self.assertEqual(props["fontSize"], "md") # generatedAt 是 ISO 时间戳 self.assertRegex(props["generatedAt"], r"^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}Z$") def test_truncation_and_overflow(self): raw = { "shift": "夜班", "operator": "张工", "events": [{"time": f"0{i}:00", "title": f"事件{i}"} for i in range(12)], "todos": [{"title": f"待办{i}"} for i in range(8)], } rec = normalize_shift_record(raw) cfg = HandoverBriefConfig(max_events=5, max_todos=3) props = render_handover_brief(rec, cfg) self.assertEqual(len(props["events"]), 5) self.assertEqual(len(props["todos"]), 3) self.assertEqual(props["eventOverflow"], 7) self.assertEqual(props["todoOverflow"], 5) def test_critical_event_forces_confirm(self): raw = { "shift": "夜班", "operator": "张工", "events": [{"time": "01:00", "title": "严重告警", "severity": "P0"}], } rec = normalize_shift_record(raw) cfg = HandoverBriefConfig(require_confirm=False) props = render_handover_brief(rec, cfg) # 含 P0 事件 → requireConfirm 被强制为 true(PRD 高利害人工确认) self.assertTrue(props["requireConfirm"]) def test_safety_text_forces_confirm(self): raw = { "shift": "夜班", "operator": "张工", "safety": "注意 3 层高温区巡检。", } rec = normalize_shift_record(raw) cfg = HandoverBriefConfig(require_confirm=False) props = render_handover_brief(rec, cfg) self.assertTrue(props["requireConfirm"]) def test_section_switches(self): rec = normalize_shift_record(_full_record_dict()) cfg = HandoverBriefConfig(sections=["overview", "abnormal"]) props = render_handover_brief(rec, cfg) self.assertTrue(props["showOverview"]) self.assertTrue(props["showAbnormal"]) self.assertFalse(props["showSafety"]) self.assertFalse(props["showEnergy"]) self.assertFalse(props["showTodos"]) def test_injected_summary_passes_through(self): rec = normalize_shift_record({"shift": "夜班", "operator": "张工"}) props = render_handover_brief(rec, summary="自定义摘要") self.assertEqual(props["summary"], "自定义摘要") # --------------------------------------------------------------------------- # 8) 验收线 # --------------------------------------------------------------------------- class TestGenerationBudget(unittest.TestCase): def test_within_budget(self): self.assertTrue(within_generation_budget(60_000)) # 1 分钟 self.assertTrue(within_generation_budget(HANDOVER_GENERATION_BUDGET_MS)) self.assertFalse(within_generation_budget(HANDOVER_GENERATION_BUDGET_MS + 1)) if __name__ == "__main__": unittest.main(verbosity=2)