"""Data models and validation schemas for Multilingual NLP Entity Inherence Classifier.""" from __future__ import annotations import json from dataclasses import dataclass, field, asdict from enum import Enum from typing import Any, Dict, List, Optional class DecisionCategory(str, Enum): DIRECT_INHERENT = "DIRECT_INHERENT" CONTEXTUAL_INHERENT = "CONTEXTUAL_INHERENT" TANGENTIAL = "TANGENTIAL" NOT_RELATED = "NOT_RELATED" class ErrorCode(str, Enum): INVALID_ECP_JSON = "invalid_ecp_json" INVALID_MARKDOWN = "invalid_markdown" UNSUPPORTED_LANGUAGE = "unsupported_language" EMPTY_CONTENT = "empty_content" MISSING_REQUIRED_FIELD = "missing_required_field" @dataclass class RelatedEntity: entity_id: str name: str relation_type: str weight: float aliases: List[str] = field(default_factory=list) scope: str = "general" confidence: float = 1.0 @classmethod def from_dict(cls, data: Dict[str, Any]) -> RelatedEntity: if not isinstance(data, dict): raise ValueError("Related entity must be a JSON object") required = ["entity_id", "name", "relation_type", "weight"] for req in required: if req not in data or data[req] is None: raise ValueError(f"Missing required field in related entity: '{req}'") return cls( entity_id=str(data["entity_id"]), name=str(data["name"]), relation_type=str(data["relation_type"]), weight=float(data["weight"]), aliases=list(data.get("aliases", [])), scope=str(data.get("scope", "general")), confidence=float(data.get("confidence", 1.0)), ) @dataclass class ECPSnapshot: target_entity_id: str target_name: str aliases: List[str] domain: str anchors: List[str] negative_anchors: List[str] = field(default_factory=list) graph_version: str = "1.0.0" related_entities: List[RelatedEntity] = field(default_factory=list) @classmethod def from_dict(cls, data: Dict[str, Any]) -> ECPSnapshot: if not isinstance(data, dict): raise ValueError("ECP Snapshot payload must be a JSON object") required_fields = ["target_entity_id", "target_name", "aliases", "domain", "anchors"] for field_name in required_fields: if field_name not in data or data[field_name] is None: raise ValueError(f"Missing required field in ECP Snapshot: '{field_name}'") if not isinstance(data["aliases"], list): raise ValueError("Field 'aliases' must be a list of strings") if not isinstance(data["anchors"], list): raise ValueError("Field 'anchors' must be a list of strings") neg_anchors = data.get("negative_anchors", []) if neg_anchors is not None and not isinstance(neg_anchors, list): raise ValueError("Field 'negative_anchors' must be a list of strings if provided") related_data = data.get("related_entities", []) if related_data is not None and not isinstance(related_data, list): raise ValueError("Field 'related_entities' must be a list if provided") related_objs = [RelatedEntity.from_dict(item) for item in (related_data or [])] return cls( target_entity_id=str(data["target_entity_id"]), target_name=str(data["target_name"]), aliases=[str(a) for a in data["aliases"]], domain=str(data["domain"]), anchors=[str(a) for a in data["anchors"]], negative_anchors=[str(na) for na in (neg_anchors or [])], graph_version=str(data.get("graph_version", "1.0.0")), related_entities=related_objs, ) @classmethod def from_json_str(cls, json_str: str) -> ECPSnapshot: try: data = json.loads(json_str) except Exception as e: raise ValueError(f"Invalid JSON format: {e}") from e return cls.from_dict(data) @dataclass class MatchedGraphEntity: entity_id: str name: str relation_type: str weight: float @dataclass class ClassificationResult: decision: DecisionCategory is_inherent: bool confidence: float detected_language: str matched_anchors: List[str] = field(default_factory=list) negative_matches: List[str] = field(default_factory=list) graph_matches: List[Dict[str, Any]] = field(default_factory=list) evidence: List[str] = field(default_factory=list) rationale: str = "" warnings: List[str] = field(default_factory=list) def to_dict(self) -> Dict[str, Any]: return { "decision": self.decision.value if isinstance(self.decision, DecisionCategory) else str(self.decision), "is_inherent": bool(self.is_inherent), "confidence": round(float(self.confidence), 4), "detected_language": str(self.detected_language), "matched_anchors": list(self.matched_anchors), "negative_matches": list(self.negative_matches), "graph_matches": list(self.graph_matches), "evidence": list(self.evidence), "rationale": str(self.rationale), "warnings": list(self.warnings), } def to_json_str(self, indent: int = 2) -> str: return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False) @dataclass class ClassificationError: error_code: ErrorCode message: str details: Dict[str, Any] = field(default_factory=dict) def to_dict(self) -> Dict[str, Any]: return { "error_code": self.error_code.value if isinstance(self.error_code, ErrorCode) else str(self.error_code), "message": str(self.message), "details": dict(self.details), } def to_json_str(self, indent: int = 2) -> str: return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)