Files
TextNLPClassifierApp/tests/test_models.py
T

115 lines
3.6 KiB
Python

"""Unit tests for ECP models, schema validation, and structured error handling."""
import pytest
from src.models import (
ECPSnapshot,
RelatedEntity,
ClassificationResult,
ClassificationError,
DecisionCategory,
ErrorCode,
)
from src.parser import strip_markdown, extract_sentences, extract_evidence_snippets
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]