115 lines
3.6 KiB
Python
115 lines
3.6 KiB
Python
"""Unit tests for ECP models, schema validation, and structured error handling."""
|
|
|
|
import pytest
|
|
|
|
from src.tools.models import (
|
|
ClassificationError,
|
|
ClassificationResult,
|
|
DecisionCategory,
|
|
ECPSnapshot,
|
|
ErrorCode,
|
|
)
|
|
from src.tools.parser import extract_evidence_snippets, strip_markdown
|
|
|
|
|
|
def test_ecp_snapshot_valid():
|
|
data = {
|
|
"target_entity_id": "ent_123",
|
|
"target_name": "Petrobras",
|
|
"aliases": ["Petróleo Brasileiro S.A.", "Petrobras"],
|
|
"domain": "Oil & Gas",
|
|
"anchors": ["pré-sal", "refinaria", "combustíveis"],
|
|
"negative_anchors": ["petrobras posto pirata"],
|
|
"graph_version": "1.0.0",
|
|
"related_entities": [
|
|
{
|
|
"entity_id": "ent_456",
|
|
"name": "Transpetro",
|
|
"relation_type": "SUBSIDIARY_OF",
|
|
"weight": 0.9,
|
|
"aliases": ["Transpetro Logística"],
|
|
"scope": "logistics",
|
|
"confidence": 0.95,
|
|
}
|
|
],
|
|
}
|
|
snapshot = ECPSnapshot.from_dict(data)
|
|
assert snapshot.target_entity_id == "ent_123"
|
|
assert snapshot.target_name == "Petrobras"
|
|
assert len(snapshot.aliases) == 2
|
|
assert len(snapshot.related_entities) == 1
|
|
assert snapshot.related_entities[0].name == "Transpetro"
|
|
assert snapshot.related_entities[0].weight == 0.9
|
|
|
|
|
|
def test_ecp_snapshot_defaults():
|
|
data = {
|
|
"target_entity_id": "ent_123",
|
|
"target_name": "Petrobras",
|
|
"aliases": ["Petrobras"],
|
|
"domain": "Oil & Gas",
|
|
"anchors": ["energia"],
|
|
}
|
|
snapshot = ECPSnapshot.from_dict(data)
|
|
assert snapshot.negative_anchors == []
|
|
assert snapshot.graph_version == "1.0.0"
|
|
assert snapshot.related_entities == []
|
|
|
|
|
|
def test_ecp_snapshot_missing_required():
|
|
data = {
|
|
"target_entity_id": "ent_123",
|
|
"aliases": ["Petrobras"],
|
|
"domain": "Oil & Gas",
|
|
"anchors": ["energia"],
|
|
}
|
|
with pytest.raises(ValueError, match="Missing required field"):
|
|
ECPSnapshot.from_dict(data)
|
|
|
|
|
|
def test_classification_result_serialization():
|
|
res = ClassificationResult(
|
|
decision=DecisionCategory.DIRECT_INHERENT,
|
|
is_inherent=True,
|
|
confidence=0.95,
|
|
detected_language="pt",
|
|
matched_anchors=["Petrobras"],
|
|
evidence=["Petrobras anunciou investimentos no pré-sal."],
|
|
rationale="Match forte da entidade alvo.",
|
|
)
|
|
d = res.to_dict()
|
|
assert d["decision"] == "DIRECT_INHERENT"
|
|
assert d["is_inherent"] is True
|
|
assert d["confidence"] == 0.95
|
|
assert d["detected_language"] == "pt"
|
|
assert "Petrobras" in d["matched_anchors"]
|
|
|
|
|
|
def test_classification_error_serialization():
|
|
err = ClassificationError(
|
|
error_code=ErrorCode.INVALID_ECP_JSON,
|
|
message="Malformed JSON syntax",
|
|
details={"path": "snapshot.json"},
|
|
)
|
|
d = err.to_dict()
|
|
assert d["error_code"] == "invalid_ecp_json"
|
|
assert d["message"] == "Malformed JSON syntax"
|
|
assert d["details"]["path"] == "snapshot.json"
|
|
|
|
|
|
def test_parser_strip_markdown():
|
|
md = "# Title\n\nThis is **bold** text and [link](https://example.com).\n- item 1\n- item 2"
|
|
plain = strip_markdown(md)
|
|
assert "Title" in plain
|
|
assert "bold text" in plain
|
|
assert "link" in plain
|
|
assert "[" not in plain
|
|
assert "*" not in plain
|
|
|
|
|
|
def test_extract_evidence_snippets():
|
|
md = "O pré-sal brasileiro é uma das maiores reservas de petróleo. A Petrobras lidera a exploração técnica."
|
|
snippets = extract_evidence_snippets(md, ["Petrobras"])
|
|
assert len(snippets) > 0
|
|
assert "Petrobras lidera" in snippets[0]
|