"""Deterministic post-processing of the LLM validation response.""" import json from datetime import datetime, timezone from typing import Any, Literal import config from src.logic.rule_confidence import ( CONFIDENCE_WEIGHTS, ConfidenceLevel, get_confidence, get_expected_rule_ids, get_lookup, ) ValidationStatus = Literal["FAIL", "PASS_WITH_WARNINGS", "PASS"] def _parse_model_response(raw: dict) -> dict: """Extract and JSON-decode the text payload from InvokeModel API response.""" return json.loads(raw["content"][0]["text"]) def _enrich_issues(issues: list[Any]) -> list[dict]: """Add a confidence field (from the rule lookup) to each issue.""" enriched = [] for issue in issues: if not isinstance(issue, dict): continue rule_id = issue.get("ruleId") confidence = get_confidence(rule_id) if rule_id else None enriched.append({**issue, "confidence": confidence}) return enriched def _unique_issue_rule_ids(issues: list[dict]) -> set[str]: return {i["ruleId"] for i in issues if i.get("ruleId")} def bucket_passed_rule_ids( passed_rule_ids: list[str], ) -> dict[ConfidenceLevel, list[str]]: """Bucket passed rule IDs by their confidence level.""" buckets: dict[ConfidenceLevel, list[str]] = { "HIGH": [], "MEDIUM": [], "LOW": [], } for rule_id in passed_rule_ids: level = get_confidence(rule_id) if level is None: continue buckets[level].append(rule_id) return buckets def confidence_score( passed_rule_ids: list[str], issue_rule_ids: list[str] ) -> int: """Return passing weight as a fraction of total evaluated weight. Does not rely on the LLM's own FAIL/WARNING classification — only on which rule IDs appear in passedRuleIds vs issues, cross-referenced with the deterministic confidence table. """ def _weight_sum(rule_ids: list[str]) -> int: return sum( CONFIDENCE_WEIGHTS[level] for rid in rule_ids if (level := get_confidence(rid)) is not None ) pass_weight = _weight_sum(passed_rule_ids) issue_weight = _weight_sum(issue_rule_ids) total = pass_weight + issue_weight if not total: return 0 return round(pass_weight / total * 100) def evaluation_coverage( returned_rule_ids: set[str], not_applicable_rule_ids: set[str] ) -> int: """Weighted coverage: evaluated weight as a fraction of total expected weight. Higher-confidence rules contribute proportionally more, so skipping a HIGH rule hurts coverage more than skipping a LOW rule. """ lookup = get_lookup() evaluated = returned_rule_ids | not_applicable_rule_ids expected_weight = sum(CONFIDENCE_WEIGHTS[level] for level in lookup.values()) if not expected_weight: return 0 evaluated_weight = sum( CONFIDENCE_WEIGHTS[lookup[rid]] for rid in evaluated if rid in lookup ) return round(evaluated_weight / expected_weight * 100) def derive_validation_status(issues: list[Any]) -> ValidationStatus: """Apply precedence FAIL > PASS_WITH_WARNINGS > PASS.""" statuses = {i.get("status") for i in issues if isinstance(i, dict)} if "FAIL" in statuses: return "FAIL" if "WARNING" in statuses: return "PASS_WITH_WARNINGS" return "PASS" def unevaluated_rule_ids( returned_rule_ids: set[str], not_applicable_rule_ids: set[str] ) -> list[str]: """Return expected rule IDs the LLM neither addressed nor marked N/A.""" return sorted( get_expected_rule_ids() - returned_rule_ids - not_applicable_rule_ids ) def _utc_timestamp() -> str: return datetime.now(timezone.utc).isoformat(timespec="seconds") def evaluate(llm_raw: dict, release_id: int | None = None) -> dict: """Produce the reliability-enriched validation result from model response. `llm_raw` is the full InvokeModel API message object. The inner JSON at content[0].text must contain: passedRuleIds, notApplicableRuleIds, issues (each with ruleId/scope/status/fieldName/message/affectedInstances), score, and rationale. """ llm = _parse_model_response(llm_raw) passed_rule_ids: list[str] = list(llm.get("passedRuleIds") or []) not_applicable_rule_ids: list[str] = list( llm.get("notApplicableRuleIds") or [] ) issues = _enrich_issues(list(llm.get("issues") or [])) not_applicable_set = set(not_applicable_rule_ids) issue_rule_ids = _unique_issue_rule_ids(issues) returned = set(passed_rule_ids) | issue_rule_ids failures = [i for i in issues if i.get("status") == "FAIL"] warnings = [i for i in issues if i.get("status") == "WARNING"] unevaluated = unevaluated_rule_ids(returned, not_applicable_set) return { "validationTimestamp": _utc_timestamp(), "validatorVersion": config.VALIDATOR_VERSION, "validationStatus": derive_validation_status(issues), "releaseId": release_id, "llmScore": llm.get("score"), "llmRationale": llm.get("rationale"), "confidenceScore": confidence_score( passed_rule_ids, list(issue_rule_ids) ), "evaluationCoverage": evaluation_coverage(returned, not_applicable_set), "issues": issues, "passedRuleIds": bucket_passed_rule_ids(passed_rule_ids), "notApplicableRuleIds": not_applicable_rule_ids, "unevaluatedRuleIds": unevaluated, "summary": { "totalRulesChecked": len(returned) + len(not_applicable_rule_ids), "passed": len(passed_rule_ids), "warnings": len(warnings), "failures": len(failures), "notApplicable": len(not_applicable_rule_ids), "unevaluated": len(unevaluated), }, }