DealDocumentScreening/tests/unit/test_prescreen_extractor_hybrid.py
2026-08-17 20:49:29 +03:00

161 lines
6.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""HybridMetaExtractor orchestrator matrix (plan Phase 4).
- high confidence → LLM skipped
- low confidence → LLM merged, version hybrid-llm-v1
- LLM failure → heuristic result kept + outcome=failed + error recorded
- fallback disabled → outcome=disabled, no provider call
- boolean merge semantics: OR across stages
"""
from __future__ import annotations
from typing import Any
from contract_check.core.metrics import prescreen_fallback_runs
from contract_check.worker_prescreen.extractor_heuristic import HeuristicExtractor
from contract_check.worker_prescreen.extractor_hybrid import HybridMetaExtractor
HIGH_CONF_TEXT = """\
ДОГОВОР ПОСТАВКИ № 42
г. Москва
Общество с ограниченной ответственностью «Продавец» и
Общество с ограниченной ответственностью «Покупатель» заключили договор.
1. Предмет договора
2.1. Стоимость товара составляет 1 250 000 рублей.
3.1. Договор вступает в силу с 01.09.2025 и действует по 31.08.2026.
4.1. Неустойка 0,1%.
5.1. Споры — в Арбитражном суде г. Москвы.
"""
LOW_CONF_TEXT = "Образец документа без опознаваемых полей."
class StubProvider:
def __init__(self, payload: dict[str, Any] | Exception) -> None:
self._payload = payload
self.calls = 0
async def analyze(self, text: str, *, checklist: str, extra_context: str = "") -> Any:
raise NotImplementedError
async def extract_prescreen(self, text: str) -> dict[str, Any]:
self.calls += 1
if isinstance(self._payload, Exception):
raise self._payload
return self._payload
async def aclose(self) -> None:
pass
LLM_DICT: dict[str, Any] = {
"contract_type": "services",
"party_a": "ИП Петров",
"party_b": "ООО «Клиент»",
"total_amount": 50000.0,
"currency": "RUB",
"has_penalty_clause": True,
}
def _counter(outcome: str) -> int:
return prescreen_fallback_runs.labels(outcome=outcome)._value.get() # type: ignore[no-any-return]
def _hybrid(
provider: StubProvider, *, enabled: bool, threshold: float = 0.75
) -> HybridMetaExtractor:
return HybridMetaExtractor(
provider=provider,
fallback_enabled=enabled,
fallback_threshold=threshold,
heuristic=HeuristicExtractor(),
)
async def test_disabled_never_calls_llm() -> None:
provider = StubProvider(LLM_DICT)
before = _counter("disabled")
result = await _hybrid(provider, enabled=False).extract(LOW_CONF_TEXT)
assert provider.calls == 0
assert result.extractor_version == "heuristic-v2"
assert result.llm_fallback_error is None
assert _counter("disabled") == before + 1
async def test_high_confidence_skips_llm() -> None:
provider = StubProvider(LLM_DICT)
before = _counter("skipped")
result = await _hybrid(provider, enabled=True).extract(HIGH_CONF_TEXT)
assert provider.calls == 0
assert result.extractor_version == "heuristic-v2"
assert result.meta.confidence_score >= 0.75
assert _counter("skipped") == before + 1
async def test_low_confidence_merges_llm() -> None:
provider = StubProvider(LLM_DICT)
before = _counter("used")
result = await _hybrid(provider, enabled=True).extract(LOW_CONF_TEXT)
assert provider.calls == 1
assert result.extractor_version == "hybrid-llm-v1"
assert result.llm_fallback_error is None
meta = result.meta
assert meta.contract_type == "services"
assert meta.party_a == "ИП Петров"
assert meta.party_b == "ООО «Клиент»"
assert meta.total_amount == 50000.0
assert meta.confidence_score >= 0.75
assert _counter("used") == before + 1
async def test_llm_failure_keeps_heuristic_result() -> None:
provider = StubProvider(RuntimeError("boom"))
before = _counter("failed")
result = await _hybrid(provider, enabled=True).extract(LOW_CONF_TEXT)
assert provider.calls == 1
assert result.extractor_version == "heuristic-v2"
assert result.llm_fallback_error is not None
assert "RuntimeError" in result.llm_fallback_error
assert result.meta.contract_type is None # heuristic result untouched
assert _counter("failed") == before + 1
async def test_llm_overrides_only_heuristic_nones() -> None:
provider = StubProvider({**LLM_DICT, "contract_type": None})
result = await _hybrid(provider, enabled=True).extract(
"Соглашение. ООО «Ромашка» действует.\nСумма 100 рублей."
)
# Heuristic found nothing for type; LLM had None too → stays None.
assert result.meta.contract_type is None
assert result.extractor_version == "hybrid-llm-v1"
async def test_booleans_or_merged() -> None:
# Heuristic: no clause words → False. LLM: penalty True → OR → True.
provider = StubProvider({**LLM_DICT, "has_penalty_clause": True})
result = await _hybrid(provider, enabled=True).extract(LOW_CONF_TEXT)
assert result.meta.has_penalty_clause is True
async def test_heuristic_boolean_survives_llm_false() -> None:
text = LOW_CONF_TEXT + " Стороны вправе расторгнуть договор."
provider = StubProvider({"has_termination_clause": False})
result = await _hybrid(provider, enabled=True).extract(text)
assert result.meta.has_termination_clause is True
async def test_threshold_boundary_equal_skips_llm() -> None:
provider = StubProvider(LLM_DICT)
# LOW_CONF_TEXT scores 0.0; threshold 0.0 → 0.0 >= 0.0 → skipped.
result = await _hybrid(provider, enabled=True, threshold=0.0).extract(LOW_CONF_TEXT)
assert provider.calls == 0
assert result.extractor_version == "heuristic-v2"
async def test_heuristic_only_failure_never_raised() -> None:
# The whole pipeline never raises from extraction stages on LLM errors.
provider = StubProvider(RuntimeError("quota exceeded"))
result = await _hybrid(provider, enabled=True).extract(HIGH_CONF_TEXT)
assert result.meta.confidence_score >= 0.75