feat: aggregate inference by question unit with pair AND

Reuse build_units/unit_correctness (pair contract single entry) in the
inference aggregation step: single questions count as one unit, AR pairs
collapse original+mirror into one unit scored by bidirectional AND. total/
correct/per_task_type are unit-grained; orphan pairs (missing one side) are
warned and dropped, not counted. Per-question predictions still land row by
row (traceability unchanged).
This commit is contained in:
2026-07-15 06:57:48 -04:00
parent 2429dad393
commit 730caa7e9a
3 changed files with 539 additions and 96 deletions
+94 -25
View File
@@ -20,6 +20,7 @@ from typing import TYPE_CHECKING, Any
from loguru import logger from loguru import logger
from app.harness.question_units import build_units, unit_correctness
from core.agent.loop import AgentLoop from core.agent.loop import AgentLoop
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -28,7 +29,7 @@ if TYPE_CHECKING:
from app.harness.log import HarnessLog from app.harness.log import HarnessLog
from core.agent.types import LoopResult from core.agent.types import LoopResult
from core.protocols import LLMProvider from core.protocols import LLMProvider
from core.types import GeneratedQuestion from core.types import GeneratedQuestion, QuestionUnit
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -182,23 +183,25 @@ def _zero_result(run_id: str) -> InferenceResult:
) )
def _group_by_task_type(records: list[dict[str, Any]]) -> dict[str, dict[str, Any]]: def _group_by_task_type(graded: list[tuple[QuestionUnit, bool]]) -> dict[str, dict[str, Any]]:
"""按 task_type 分组聚合正确率指标。 """按 task_type 分组聚合 unit 级正确率指标。
pair 单元整体计 1 个 unit,归入其 task_typesingle 单元计 1 个 unit。
参数: 参数:
records: 预测记录列表。 graded: (单元, 该单元是否整体正确) 元组列表。
返回: 返回:
{task_type: {accuracy, total, correct}} 映射。 {task_type: {accuracy, total, correct}} 映射unit 粒度)
""" """
task_groups: dict[str, list[dict[str, Any]]] = defaultdict(list) task_groups: dict[str, list[bool]] = defaultdict(list)
for r in records: for unit, is_correct in graded:
task_groups[r["task_type"]].append(r) task_groups[unit.task_type].append(is_correct)
per_task_type: dict[str, dict[str, Any]] = {} per_task_type: dict[str, dict[str, Any]] = {}
for task_type, group in task_groups.items(): for task_type, verdicts in task_groups.items():
t_total = len(group) t_total = len(verdicts)
t_correct = sum(1 for r in group if r["prediction"] == r["answer"]) t_correct = sum(verdicts)
per_task_type[task_type] = { per_task_type[task_type] = {
"accuracy": t_correct / t_total, "accuracy": t_correct / t_total,
"total": t_total, "total": t_total,
@@ -207,35 +210,101 @@ def _group_by_task_type(records: list[dict[str, Any]]) -> dict[str, dict[str, An
return per_task_type return per_task_type
def _aggregate_results(records: list[dict[str, Any]], run_id: str) -> InferenceResult: def _drop_orphan_pairs(questions: list[GeneratedQuestion]) -> list[GeneratedQuestion]:
"""从内存 records 聚合推理指标 """剔除收不齐 2 条 / 角色非法的孤儿 pair,告警不静默
TRM4 从 DB 回读 predictions 表聚合;TRM5 改为从内存直接聚合, 每条题目均会各答一次并逐题落库;能否合成 pair 单元仅取决于 questions
避免 DB 回读的同步开销和额外依赖。 是否同时含该 pair_id 的 original + mirror。收不齐者告警并整对剔除,使
后续 build_units 只面对合法孪生对(不触发 fail-fast),孤儿 unit 不计入
total(对齐设计 §8 聚合入口的"告警 + 剔除")。
参数: 参数:
records: _run_single_question 返回的 record 列表 questions: 待聚合的题目列表(可混含 single 与孪生对成员)
返回:
可安全交给 build_units 的题目列表(single 全保留,pair 仅保留成对者)。
"""
by_pair: dict[str, list[GeneratedQuestion]] = defaultdict(list)
singles: list[GeneratedQuestion] = []
for q in questions:
if q.pair_id:
by_pair[q.pair_id].append(q)
else:
singles.append(q)
kept_pairs: list[GeneratedQuestion] = []
for pair_id, group in by_pair.items():
originals = [q for q in group if q.question_role == "pair_original"]
mirrors = [q for q in group if q.question_role == "pair_mirror"]
if len(originals) == 1 and len(mirrors) == 1:
kept_pairs.extend(group)
else:
logger.warning(
"孤儿 pair {}:收不齐 2 条(original={} mirror={}),剔除该 unit 不计入 total",
pair_id,
len(originals),
len(mirrors),
)
return singles + kept_pairs
def _per_question_correctness(records: list[dict[str, Any]]) -> dict[str, bool]:
"""由逐题 record 构造 question_id → 该题作答是否正确 的映射。
prediction 为 None(作答异常)时与 answer 不相等 → False,天然计错。
参数:
records: _run_single_question 返回的逐题 record 列表。
返回:
{question_id: prediction == answer} 映射,供 unit_correctness 取值。
"""
return {r["question_id"]: r["prediction"] == r["answer"] for r in records}
def _aggregate_results(
records: list[dict[str, Any]],
questions: list[GeneratedQuestion],
run_id: str,
) -> InferenceResult:
"""从内存 records + 题目列表按 unit 粒度聚合推理指标。
逐题 record 保留逐题溯源(token/steps/stop_reason 诊断仍按 record 汇总);
正确率则按 unit 粒度计:single 计 1AR pair 经 build_units 收齐 original +
mirror 后走 unit_correctness 的双向 AND 判定,整对计 1 个 unit。孤儿 pair
在 _drop_orphan_pairs 中告警 + 剔除,不计入 total。
参数:
records: _run_single_question 返回的逐题 record 列表。
questions: 与 records 对应的题目列表(提供 pair_id/question_role 元数据)。
run_id: 当前运行标识。 run_id: 当前运行标识。
返回: 返回:
InferenceResult 冻结实例。 InferenceResult 冻结实例total/correct/per_task_type 为 unit 粒度)
""" """
total = len(records) if not records:
if total == 0:
return _zero_result(run_id) return _zero_result(run_id)
correct = sum(1 for r in records if r["prediction"] == r["answer"]) per_q = _per_question_correctness(records)
units = build_units(_drop_orphan_pairs(questions))
# unit 内任一题缺 record → unit_correctness 抛 KeyError(防静默兜底/读回校验)。
graded = [(unit, unit_correctness(unit, per_q)) for unit in units]
total = len(graded)
correct = sum(1 for _, is_correct in graded if is_correct)
stop_counts: dict[str, int] = defaultdict(int) stop_counts: dict[str, int] = defaultdict(int)
for r in records: for r in records:
stop_counts[r["stop_reason"]] += 1 stop_counts[r["stop_reason"]] += 1
n_records = len(records)
return InferenceResult( return InferenceResult(
run_id=run_id, run_id=run_id,
accuracy=correct / total, accuracy=correct / total if total else 0.0,
total=total, total=total,
correct=correct, correct=correct,
per_task_type=_group_by_task_type(records), per_task_type=_group_by_task_type(graded),
steps_mean=sum(r["steps_used"] for r in records) / total, steps_mean=sum(r["steps_used"] for r in records) / n_records,
token_usage={ token_usage={
"prompt_tokens": sum(r["prompt_tokens"] for r in records), "prompt_tokens": sum(r["prompt_tokens"] for r in records),
"completion_tokens": sum(r["completion_tokens"] for r in records), "completion_tokens": sum(r["completion_tokens"] for r in records),
@@ -399,7 +468,7 @@ async def run_inference(
if not questions: if not questions:
logger.info("题目列表为空,返回零值 InferenceResult") logger.info("题目列表为空,返回零值 InferenceResult")
return _aggregate_results([], run_id) return _aggregate_results([], [], run_id)
sem = asyncio.Semaphore(concurrency) sem = asyncio.Semaphore(concurrency)
total_count = len(questions) total_count = len(questions)
@@ -431,7 +500,7 @@ async def run_inference(
results = await asyncio.gather(*[_bounded(i, qa) for i, qa in enumerate(questions)]) results = await asyncio.gather(*[_bounded(i, qa) for i, qa in enumerate(questions)])
inference_result = _aggregate_results(list(results), run_id) inference_result = _aggregate_results(list(results), questions, run_id)
logger.info( logger.info(
"推理完成: accuracy={:.2%} ({}/{})", "推理完成: accuracy={:.2%} ({}/{})",
inference_result.accuracy, inference_result.accuracy,
+113 -71
View File
@@ -174,13 +174,37 @@ class TestToTextField:
assert "\\u" not in result assert "\\u" not in result
def _single_record(
question_id: str,
*,
prediction: str | None,
answer: str,
task_type: str,
steps_used: int,
prompt_tokens: int,
completion_tokens: int,
stop_reason: str,
) -> dict[str, Any]:
"""构造一条 single 题的 prediction record(含 question_id 供 unit 聚合)。"""
return {
"question_id": question_id,
"prediction": prediction,
"answer": answer,
"task_type": task_type,
"steps_used": steps_used,
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"stop_reason": stop_reason,
}
class TestAggregateResults: class TestAggregateResults:
"""_aggregate_results 内存聚合测试""" """_aggregate_results unit 粒度聚合测试(single 题:unit 数 = 题数)"""
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_empty_records(self) -> None: async def test_empty_records(self) -> None:
"""空列表返回零值 InferenceResult。""" """空列表返回零值 InferenceResult。"""
result = _aggregate_results([], "run-empty") result = _aggregate_results([], [], "run-empty")
assert result.run_id == "run-empty" assert result.run_id == "run-empty"
assert result.accuracy == 0.0 assert result.accuracy == 0.0
assert result.total == 0 assert result.total == 0
@@ -193,18 +217,20 @@ class TestAggregateResults:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_single_correct(self) -> None: async def test_single_correct(self) -> None:
"""单条正确记录 → accuracy=1.0。""" """单条正确记录 → accuracy=1.0。"""
questions = [_make_question(question_id="q1", task_type="AR", answer="B")]
records = [ records = [
{ _single_record(
"prediction": "B", "q1",
"answer": "B", prediction="B",
"task_type": "AR", answer="B",
"steps_used": 3, task_type="AR",
"prompt_tokens": 100, steps_used=3,
"completion_tokens": 50, prompt_tokens=100,
"stop_reason": "finished", completion_tokens=50,
} stop_reason="finished",
)
] ]
result = _aggregate_results(records, "run-1") result = _aggregate_results(records, questions, "run-1")
assert result.accuracy == 1.0 assert result.accuracy == 1.0
assert result.total == 1 assert result.total == 1
assert result.correct == 1 assert result.correct == 1
@@ -213,36 +239,44 @@ class TestAggregateResults:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mixed_correct_wrong(self) -> None: async def test_mixed_correct_wrong(self) -> None:
"""混合正确/错误 → 准确率与步数均正确聚合。""" """混合正确/错误 → 准确率与步数均正确聚合。"""
records = [ questions = [
{ _make_question(question_id="q1", task_type="AR", answer="B"),
"prediction": "B", _make_question(question_id="q2", task_type="AR", answer="A"),
"answer": "B", _make_question(question_id="q3", task_type="SP", answer="D"),
"task_type": "AR",
"steps_used": 2,
"prompt_tokens": 100,
"completion_tokens": 50,
"stop_reason": "finished",
},
{
"prediction": "C",
"answer": "A",
"task_type": "AR",
"steps_used": 4,
"prompt_tokens": 200,
"completion_tokens": 100,
"stop_reason": "budget_exceeded",
},
{
"prediction": "D",
"answer": "D",
"task_type": "SP",
"steps_used": 1,
"prompt_tokens": 50,
"completion_tokens": 25,
"stop_reason": "finished",
},
] ]
result = _aggregate_results(records, "run-mix") records = [
_single_record(
"q1",
prediction="B",
answer="B",
task_type="AR",
steps_used=2,
prompt_tokens=100,
completion_tokens=50,
stop_reason="finished",
),
_single_record(
"q2",
prediction="C",
answer="A",
task_type="AR",
steps_used=4,
prompt_tokens=200,
completion_tokens=100,
stop_reason="budget_exceeded",
),
_single_record(
"q3",
prediction="D",
answer="D",
task_type="SP",
steps_used=1,
prompt_tokens=50,
completion_tokens=25,
stop_reason="finished",
),
]
result = _aggregate_results(records, questions, "run-mix")
assert result.total == 3 assert result.total == 3
assert result.correct == 2 assert result.correct == 2
assert abs(result.accuracy - 2 / 3) < 1e-9 assert abs(result.accuracy - 2 / 3) < 1e-9
@@ -252,37 +286,45 @@ class TestAggregateResults:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_per_task_type_grouping(self) -> None: async def test_per_task_type_grouping(self) -> None:
"""按 task_type 分组聚合。""" """按 task_type 分组聚合unit 粒度)"""
records = [ questions = [
{ _make_question(question_id="q1", task_type="AR", answer="B"),
"prediction": "B", _make_question(question_id="q2", task_type="AR", answer="C"),
"answer": "B", _make_question(question_id="q3", task_type="SP", answer="D"),
"task_type": "AR",
"steps_used": 1,
"prompt_tokens": 10,
"completion_tokens": 5,
"stop_reason": "finished",
},
{
"prediction": "A",
"answer": "C",
"task_type": "AR",
"steps_used": 2,
"prompt_tokens": 20,
"completion_tokens": 10,
"stop_reason": "finished",
},
{
"prediction": "D",
"answer": "D",
"task_type": "SP",
"steps_used": 3,
"prompt_tokens": 30,
"completion_tokens": 15,
"stop_reason": "finished",
},
] ]
result = _aggregate_results(records, "run-task") records = [
_single_record(
"q1",
prediction="B",
answer="B",
task_type="AR",
steps_used=1,
prompt_tokens=10,
completion_tokens=5,
stop_reason="finished",
),
_single_record(
"q2",
prediction="A",
answer="C",
task_type="AR",
steps_used=2,
prompt_tokens=20,
completion_tokens=10,
stop_reason="finished",
),
_single_record(
"q3",
prediction="D",
answer="D",
task_type="SP",
steps_used=3,
prompt_tokens=30,
completion_tokens=15,
stop_reason="finished",
),
]
result = _aggregate_results(records, questions, "run-task")
assert "AR" in result.per_task_type assert "AR" in result.per_task_type
assert "SP" in result.per_task_type assert "SP" in result.per_task_type
assert result.per_task_type["AR"]["total"] == 2 assert result.per_task_type["AR"]["total"] == 2
+332
View File
@@ -0,0 +1,332 @@
"""inference pair-level 双向 AND 聚合单元测试(Task 6)。
覆盖 question-gen v3 Phase 1 Task 6 的核心契约:
- 逐题推理不变:每条 GeneratedQuestion 照常各答一次、per-question prediction
仍逐题落 predictions 表(保留逐题溯源)。
- pair 按 pair_id 收齐 original + mirror 后合成 1 条 unit-level 记录,
pair 正确 = (P.pred==P.answer) AND (Q.pred==Q.answer)(双向 AND)。
- InferenceResult.total / correct / per_task_type 全部按 unit 粒度
single 计 1pair 计 1)。
- 孤儿 pair(收不齐 2 条)→ 告警 + 剔除该 unit、不计入 total(不静默)。
"""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import AsyncMock
import pytest
from loguru import logger
from app.harness.inference import _aggregate_results, run_inference
from app.harness.log import HarnessLog
from core.types import GeneratedQuestion, LLMResponse
# ── 测试基础设施 ──────────────────────────────────────────────────
def _make_question(
question_id: str,
*,
task_type: str = "Action Reasoning",
answer: str = "B",
pair_id: str | None = None,
question_role: str = "single",
flip_axis: str | None = None,
video_id: str = "v1",
) -> GeneratedQuestion:
"""构造测试题目;pair_id 非空时视为孪生对成员。"""
return GeneratedQuestion(
question_id=question_id,
video_id=video_id,
task_type=task_type,
question="测试问题",
options=("A. 选项A", "B. 选项B", "C. 选项C", "D. 选项D"),
answer=answer,
source_nodes=("L1_001",),
difficulty="medium",
pair_id=pair_id,
question_role=question_role,
flip_axis=flip_axis if pair_id else None,
)
def _make_record(
question_id: str,
*,
prediction: str | None,
answer: str = "B",
task_type: str = "Action Reasoning",
steps_used: int = 2,
prompt_tokens: int = 100,
completion_tokens: int = 50,
stop_reason: str = "finished",
) -> dict[str, Any]:
"""构造与题目匹配的 prediction record(键与 _run_single_question 一致)。"""
return {
"question_id": question_id,
"prediction": prediction,
"answer": answer,
"task_type": task_type,
"steps_used": steps_used,
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"stop_reason": stop_reason,
}
def _make_llm_response(answer: str = "B") -> LLMResponse:
"""构造 submit_answer 场景的 LLMResponse。"""
content = json.dumps(
{
"reflect": {"observation": "找到答案"},
"plan": {"next_step": "提交"},
"action": {
"tool": "submit_answer",
"args": {"answer": answer, "evidence": "证据", "reasoning": "推理"},
},
}
)
return LLMResponse(
content=content,
thinking="思考",
model="test-model",
provider="test",
prompt_tokens=100,
completion_tokens=50,
latency_ms=200,
ttft_ms=30.0,
max_inter_token_ms=5.0,
cache_hit=False,
call_id="test-call-001",
)
async def _stub_tool_dispatch(
tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
) -> str:
"""测试用工具调度函数。"""
if tool_name == "submit_answer":
return "答案已提交"
raise ValueError(f"未知工具: {tool_name}")
def _stub_prompt_builder(qa: GeneratedQuestion) -> tuple[str, str]:
"""测试用 prompt 构建函数。"""
return "系统提示词", f"用户问题: {qa.question}"
@pytest.fixture
def harness_log(tmp_path: Any, request: Any) -> HarnessLog:
"""创建临时 HarnessLog 实例。"""
db_path = str(tmp_path / f"harness_{id(request)}.db")
log = HarnessLog(db_path, "test-run")
yield log
log.close()
# ── unit 粒度聚合(_aggregate_results 直测) ─────────────────────────
class TestUnitLevelAggregation:
"""_aggregate_results 按 unit 粒度聚合测试。"""
def test_single_units_counted_per_question(self) -> None:
"""全 singletotal = single 数,correct 逐题判定。"""
questions = [
_make_question("s1", answer="B"),
_make_question("s2", answer="A"),
]
records = [
_make_record("s1", prediction="B", answer="B"),
_make_record("s2", prediction="C", answer="A"),
]
result = _aggregate_results(records, questions, "run-single")
assert result.total == 2
assert result.correct == 1
assert abs(result.accuracy - 0.5) < 1e-9
def test_pair_both_correct_is_one_correct_unit(self) -> None:
"""pair 两题皆对 → 1 个 unit、correct=1。"""
questions = [
_make_question("po", answer="B", pair_id="p", question_role="pair_original"),
_make_question("pm", answer="A", pair_id="p", question_role="pair_mirror"),
]
records = [
_make_record("po", prediction="B", answer="B"),
_make_record("pm", prediction="A", answer="A"),
]
result = _aggregate_results(records, questions, "run-pair-ok")
assert result.total == 1
assert result.correct == 1
assert result.accuracy == 1.0
def test_pair_one_wrong_fails_by_and(self) -> None:
"""pair 一题错 → 双向 AND 判 unit 错,correct=0。"""
questions = [
_make_question("po", answer="B", pair_id="p", question_role="pair_original"),
_make_question("pm", answer="A", pair_id="p", question_role="pair_mirror"),
]
records = [
_make_record("po", prediction="B", answer="B"), # 对
_make_record("pm", prediction="D", answer="A"), # 错
]
result = _aggregate_results(records, questions, "run-pair-half")
assert result.total == 1
assert result.correct == 0
assert result.accuracy == 0.0
def test_mixed_single_and_pair_unit_total(self) -> None:
"""single + pair 混合:total = single 数 + pair 数(pair 计 1)。"""
questions = [
_make_question("s1", answer="B"),
_make_question("po", answer="B", pair_id="p", question_role="pair_original"),
_make_question("pm", answer="A", pair_id="p", question_role="pair_mirror"),
]
records = [
_make_record("s1", prediction="B", answer="B"),
_make_record("po", prediction="B", answer="B"),
_make_record("pm", prediction="A", answer="A"),
]
result = _aggregate_results(records, questions, "run-mixed")
assert result.total == 2 # 1 single + 1 pair
assert result.correct == 2
def test_per_task_type_by_unit(self) -> None:
"""per_task_type 按 unit 计数:pair 归入其 task_type 计 1 个 unit。"""
questions = [
_make_question("s1", task_type="SP", answer="B"),
_make_question(
"po", task_type="AR", answer="B", pair_id="p", question_role="pair_original"
),
_make_question(
"pm", task_type="AR", answer="A", pair_id="p", question_role="pair_mirror"
),
]
records = [
_make_record("s1", prediction="B", answer="B", task_type="SP"),
_make_record("po", prediction="B", answer="B", task_type="AR"),
_make_record("pm", prediction="C", answer="A", task_type="AR"), # pair 错
]
result = _aggregate_results(records, questions, "run-tt")
assert result.per_task_type["AR"]["total"] == 1 # pair 计 1 个 unit
assert result.per_task_type["AR"]["correct"] == 0
assert result.per_task_type["SP"]["total"] == 1
assert result.per_task_type["SP"]["correct"] == 1
def test_orphan_pair_dropped_and_warned(self) -> None:
"""孤儿 pair(收不齐 2 条)→ 告警 + 剔除、不计入 total。"""
questions = [
_make_question("s1", answer="B"),
_make_question(
"po", answer="B", pair_id="orphan", question_role="pair_original"
), # 缺 mirror
]
records = [
_make_record("s1", prediction="B", answer="B"),
_make_record("po", prediction="B", answer="B"),
]
captured: list[str] = []
sink_id = logger.add(captured.append, level="WARNING", format="{message}")
try:
result = _aggregate_results(records, questions, "run-orphan")
finally:
logger.remove(sink_id)
assert result.total == 1 # 仅 single,孤儿 pair 被剔除
assert result.correct == 1
assert any("orphan" in msg for msg in captured), "孤儿 pair 未告警(静默)"
def test_empty_records_returns_zero(self) -> None:
"""空 records/questions → 零值结果。"""
result = _aggregate_results([], [], "run-empty")
assert result.total == 0
assert result.correct == 0
assert result.accuracy == 0.0
assert result.per_task_type == {}
def test_token_and_steps_span_all_records(self) -> None:
"""token/steps 诊断字段覆盖全部 record(含 pair 两条)。"""
questions = [
_make_question("po", answer="B", pair_id="p", question_role="pair_original"),
_make_question("pm", answer="A", pair_id="p", question_role="pair_mirror"),
]
records = [
_make_record("po", prediction="B", answer="B", steps_used=3, prompt_tokens=100),
_make_record("pm", prediction="A", answer="A", steps_used=1, prompt_tokens=200),
]
result = _aggregate_results(records, questions, "run-diag")
assert result.token_usage["prompt_tokens"] == 300
assert abs(result.steps_mean - 2.0) < 1e-9 # (3+1)/2 record 粒度
# ── run_inference 端到端:逐题溯源 + pair 聚合 ──────────────────────
class TestRunInferencePairEndToEnd:
"""run_inference pair 端到端:逐题落库不变 + unit 级聚合。"""
@pytest.mark.asyncio
async def test_pair_predictions_persisted_per_question(
self, harness_log: HarnessLog
) -> None:
"""pair 两题各自逐题落 predictions(保留逐题溯源),聚合按 unit。"""
llm = AsyncMock()
llm.chat.return_value = _make_llm_response(answer="B")
questions = [
_make_question("po", answer="B", pair_id="p", question_role="pair_original"),
_make_question("pm", answer="B", pair_id="p", question_role="pair_mirror"),
]
result = await run_inference(
questions,
llm=llm,
tool_dispatch_fn=_stub_tool_dispatch,
prompt_builder=_stub_prompt_builder,
log=harness_log,
run_id="run-pair-e2e",
concurrency=2,
max_steps=10,
skill_mode="auto",
)
# unit 级:1 个 pair unit,两题皆对 → correct=1
assert result.total == 1
assert result.correct == 1
# 逐题溯源:predictions 表两条 record 都在
rows = harness_log.query("SELECT * FROM predictions WHERE run_id = ?", ("test-run",))
qids = {r["question_id"] for r in rows}
assert qids == {"po", "pm"}
@pytest.mark.asyncio
async def test_orphan_pair_excluded_single_survives(
self, harness_log: HarnessLog
) -> None:
"""run_inference 中孤儿 pair 被剔除、single 仍计入。"""
llm = AsyncMock()
llm.chat.return_value = _make_llm_response(answer="B")
questions = [
_make_question("s1", answer="B"),
_make_question(
"po", answer="B", pair_id="orphan", question_role="pair_original"
),
]
result = await run_inference(
questions,
llm=llm,
tool_dispatch_fn=_stub_tool_dispatch,
prompt_builder=_stub_prompt_builder,
log=harness_log,
run_id="run-orphan-e2e",
concurrency=2,
max_steps=10,
skill_mode="auto",
)
assert result.total == 1 # single 存活,孤儿剔除
# 逐题溯源:孤儿题仍逐题落库(推理不变)
rows = harness_log.query("SELECT * FROM predictions WHERE run_id = ?", ("test-run",))
assert {r["question_id"] for r in rows} == {"s1", "po"}