161 lines
6.1 KiB
Python
161 lines
6.1 KiB
Python
"""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
|