- app/core/retrieval_fusion.py(新增): 两路候选加权融合,三层分数分离 - app/service/retrieval_rerank_service.py(新增): LLM 精排,只出 id / 越界剔除 / 失败静默回落 - customer_service.py: 接线,开关 CS_DUAL_ROUTE / CS_RERANK,未设置时一行都不执行 - 测试 61 条(融合 22 / 精排 22 / 接线 17) - agent_intent_config 补 5 条 active(1 -> 6) 铁律: vector_score 只用于出口判定,绝不改写;融合与精排只改「看哪些块、什么顺序」 门禁: 2608 passed / 3 skipped / 0 failed;金标 55 条逐项零差异,M-1 恒 100%
203 lines
7.7 KiB
Python
203 lines
7.7 KiB
Python
"""`app/service/retrieval_rerank_service.py` 的单测(`W33` 一期)。
|
||
|
||
重点不在"能不能排序",而在**它失败时不能造成损失**:
|
||
它是链路里的**增益**环节,任何一次异常、超时、格式违约都必须**静默回落原顺序**,
|
||
且**绝不能丢块、绝不能引入新块**(后者是 `INV-1` 档位不可越的一道防线)。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
from typing import Any
|
||
|
||
import pytest
|
||
|
||
from app.service.retrieval_rerank_service import RetrievalRerankService
|
||
|
||
|
||
class _Exec:
|
||
def __init__(self, text: str) -> None:
|
||
self.text = text
|
||
|
||
|
||
def _item(doc_id: str, title: str = "") -> dict[str, Any]:
|
||
return {"doc_id": doc_id, "title": title or f"标题-{doc_id}",
|
||
"content": f"正文-{doc_id}", "score": 0.6, "family_id": "F"}
|
||
|
||
|
||
def _service(reply: str | Exception, calls: list[str] | None = None, **kwargs: Any):
|
||
async def resolve() -> list[object]:
|
||
return [object()]
|
||
|
||
async def generate(endpoints: list[object], prompt: str): # noqa: ANN202
|
||
if calls is not None:
|
||
calls.append(prompt)
|
||
if isinstance(reply, Exception):
|
||
raise reply
|
||
return _Exec(reply)
|
||
|
||
return RetrievalRerankService(resolve, generate, **kwargs)
|
||
|
||
|
||
async def test_single_item_does_not_call_model() -> None:
|
||
"""只有 0/1 块 ⇒ 没有可排的余地,**不许花一次模型调用**。"""
|
||
calls: list[str] = []
|
||
svc = _service('{"ordered_ids": ["A"]}', calls)
|
||
out = await svc.rerank("问题", [_item("A")])
|
||
assert out.applied is False
|
||
assert out.reason == "nothing_to_rank"
|
||
assert calls == [], "单块也调了模型"
|
||
|
||
|
||
async def test_empty_evidence_is_noop() -> None:
|
||
calls: list[str] = []
|
||
svc = _service('{"ordered_ids": []}', calls)
|
||
out = await svc.rerank("问题", [])
|
||
assert out.evidence == []
|
||
assert calls == []
|
||
|
||
|
||
async def test_reorders() -> None:
|
||
"""正常路径:按模型给的 id 顺序重排。"""
|
||
svc = _service('{"ordered_ids": ["C", "A", "B"]}')
|
||
items = [_item("A"), _item("B"), _item("C")]
|
||
out = await svc.rerank("问题", items)
|
||
assert out.applied is True
|
||
assert [i["doc_id"] for i in out.evidence] == ["C", "A", "B"]
|
||
assert out.ordered_ids == ("C", "A", "B")
|
||
|
||
|
||
async def test_order_mode_never_drops_blocks() -> None:
|
||
"""🔴 **不变量守卫**:`order` 模式下输出必须是输入的一个**排列**(不丢块、不加块)。
|
||
|
||
少给一块证据就可能把 `M-4 事实正确率` 打下来,所以排序不允许顺手做取舍。
|
||
"""
|
||
svc = _service('{"ordered_ids": ["B", "A"]}') # 模型漏掉了 C
|
||
items = [_item("A"), _item("B"), _item("C")]
|
||
out = await svc.rerank("问题", items)
|
||
got = [i["doc_id"] for i in out.evidence]
|
||
assert sorted(got) == ["A", "B", "C"], f"块被丢了或多了:{got}"
|
||
assert got[:2] == ["B", "A"], "模型明确给出的顺序应被采纳"
|
||
assert got[2] == "C", "模型漏掉的块应追加到尾部(保持原相对顺序)"
|
||
|
||
|
||
async def test_out_of_scope_id_is_rejected() -> None:
|
||
"""🔴 **越界 id 必须剔除** —— 模型不得借精排引入候选之外的块(`INV-1` 防线)。"""
|
||
svc = _service('{"ordered_ids": ["A", "INTRUDER", "B"]}')
|
||
items = [_item("A"), _item("B")]
|
||
out = await svc.rerank("问题", items)
|
||
assert [i["doc_id"] for i in out.evidence] == ["A", "B"]
|
||
assert "INTRUDER" not in out.ordered_ids
|
||
assert out.applied is True
|
||
|
||
|
||
async def test_duplicate_ids_deduped() -> None:
|
||
"""重复 id 去重,不产生重复块。"""
|
||
svc = _service('{"ordered_ids": ["B", "B", "A"]}')
|
||
out = await svc.rerank("问题", [_item("A"), _item("B")])
|
||
assert [i["doc_id"] for i in out.evidence] == ["B", "A"]
|
||
assert len(out.evidence) == 2
|
||
|
||
|
||
@pytest.mark.parametrize("reply", [
|
||
"完全不是 JSON",
|
||
'{"ordered_ids": "A,B"}',
|
||
'{"ordered_ids": []}',
|
||
'{"other": ["A"]}',
|
||
"{}",
|
||
"",
|
||
])
|
||
async def test_unparsable_output_falls_back(reply: str) -> None:
|
||
"""非 JSON / 字段类型不对 / 空数组 ⇒ 静默回落原顺序。"""
|
||
items = [_item("A"), _item("B")]
|
||
svc = _service(reply)
|
||
out = await svc.rerank("问题", items)
|
||
assert out.applied is False
|
||
assert [i["doc_id"] for i in out.evidence] == ["A", "B"]
|
||
assert out.reason == "unparsable_output"
|
||
|
||
|
||
async def test_json_with_surrounding_prose() -> None:
|
||
"""模型多说了两句也不算违约(与 `E4` 的解析容错同取向)。"""
|
||
svc = _service('好的,排序如下:\n{"ordered_ids": ["B", "A"]}\n以上。')
|
||
out = await svc.rerank("问题", [_item("A"), _item("B")])
|
||
assert out.applied is True
|
||
assert [i["doc_id"] for i in out.evidence] == ["B", "A"]
|
||
|
||
|
||
async def test_code_fence_is_stripped() -> None:
|
||
"""带 ``` 围栏也认。"""
|
||
svc = _service('```json\n{"ordered_ids": ["B", "A"]}\n```')
|
||
out = await svc.rerank("问题", [_item("A"), _item("B")])
|
||
assert out.applied is True
|
||
assert [i["doc_id"] for i in out.evidence] == ["B", "A"]
|
||
|
||
|
||
async def test_timeout_falls_back() -> None:
|
||
"""超时 ⇒ 回落(并如实记录原因,供留痕统计)。"""
|
||
async def resolve() -> list[object]:
|
||
return [object()]
|
||
|
||
async def generate(endpoints: list[object], prompt: str): # noqa: ANN202
|
||
await asyncio.sleep(5)
|
||
return _Exec('{"ordered_ids": ["B", "A"]}')
|
||
|
||
svc = RetrievalRerankService(resolve, generate, timeout_seconds=0.05)
|
||
out = await svc.rerank("问题", [_item("A"), _item("B")])
|
||
assert out.applied is False
|
||
assert out.reason == "timeout"
|
||
assert [i["doc_id"] for i in out.evidence] == ["A", "B"]
|
||
|
||
|
||
async def test_exception_falls_back() -> None:
|
||
"""任何异常都只回落,绝不冒泡(增益调用不得把能答的问题变成答不出来)。"""
|
||
svc = _service(RuntimeError("模型不可用"))
|
||
out = await svc.rerank("问题", [_item("A"), _item("B")])
|
||
assert out.applied is False
|
||
assert out.reason == "call_failed"
|
||
assert [i["doc_id"] for i in out.evidence] == ["A", "B"]
|
||
|
||
|
||
async def test_missing_doc_id_falls_back() -> None:
|
||
"""有块缺 `doc_id` ⇒ 不排(排序需要可引用的 id,不能猜)。"""
|
||
svc = _service('{"ordered_ids": ["A"]}')
|
||
out = await svc.rerank("问题", [_item("A"), {"title": "没有 doc_id"}])
|
||
assert out.applied is False
|
||
assert out.reason == "missing_doc_id"
|
||
|
||
|
||
async def test_select_mode_respects_limit() -> None:
|
||
"""`select` 模式按 limit 截断(**默认不启用**,仅作为可度量的备选)。"""
|
||
svc = _service('{"ordered_ids": ["C", "B", "A"]}', mode="select")
|
||
out = await svc.rerank("问题", [_item("A"), _item("B"), _item("C")], limit=2)
|
||
assert out.applied is True
|
||
assert [i["doc_id"] for i in out.evidence] == ["C", "B"]
|
||
|
||
|
||
async def test_order_mode_ignores_limit() -> None:
|
||
"""`order` 模式下 limit 不生效 —— 排序模式不参与取舍。"""
|
||
svc = _service('{"ordered_ids": ["C", "B", "A"]}', mode="order")
|
||
out = await svc.rerank("问题", [_item("A"), _item("B"), _item("C")], limit=1)
|
||
assert len(out.evidence) == 3
|
||
|
||
|
||
def test_invalid_mode_rejected() -> None:
|
||
with pytest.raises(ValueError):
|
||
_service("{}", mode="bogus")
|
||
|
||
|
||
def test_invalid_timeout_rejected() -> None:
|
||
with pytest.raises(ValueError):
|
||
_service("{}", timeout_seconds=0)
|
||
|
||
|
||
async def test_prompt_carries_ids_and_no_evidence_injection() -> None:
|
||
"""提示词里必须出现候选 id(否则模型无从引用),且**不含任何要求它输出理由**的措辞。"""
|
||
calls: list[str] = []
|
||
svc = _service('{"ordered_ids": ["A"]}', calls)
|
||
await svc.rerank("用户问题", [_item("A"), _item("B")])
|
||
prompt = calls[0]
|
||
assert "id=A" in prompt and "id=B" in prompt
|
||
assert "用户问题" in prompt
|
||
assert "不要输出理由" in prompt
|