"""Configuration-driven query and answer evidence contracts.""" from __future__ import annotations import ast from decimal import Decimal, InvalidOperation from functools import lru_cache import json import os from pathlib import Path import re from typing import Any, Mapping ROOT = Path(__file__).resolve().parents[1] DEFAULT_CONTRACT_FILE = ROOT / "config" / "hmm_hr_query_contracts.json" def _contract_file() -> Path: configured = str( os.getenv("AI_WEB_AGENT_CONSOLE_QUERY_CONTRACTS_PATH") or os.getenv("POC4_QUERY_CONTRACTS_PATH") or "" ).strip() if not configured: return DEFAULT_CONTRACT_FILE path = Path(configured).expanduser() return path if path.is_absolute() else ROOT / path @lru_cache(maxsize=4) def _load_contract_file(path_text: str) -> tuple[Mapping[str, Any], ...]: path = Path(path_text) if not path.exists(): return () value = json.loads(path.read_text(encoding="utf-8")) contracts = value.get("contracts") if isinstance(value, Mapping) else None if not isinstance(contracts, list): raise ValueError("query contracts must contain a contracts array") return tuple(item for item in contracts if isinstance(item, Mapping)) def load_query_contracts() -> tuple[Mapping[str, Any], ...]: return _load_contract_file(str(_contract_file())) def matching_query_contracts( question: str, tool_name: str, ) -> tuple[Mapping[str, Any], ...]: normalized_question = " ".join(str(question or "").casefold().split()) normalized_tool = str(tool_name or "").strip() selected: list[Mapping[str, Any]] = [] for contract in load_query_contracts(): tools = contract.get("applies_to_tools") if isinstance(tools, list) and normalized_tool not in { str(item) for item in tools }: continue intent = contract.get("intent_match") if not isinstance(intent, Mapping): continue subject_terms = [ str(item).casefold() for item in intent.get("subject_terms_any", []) if str(item).strip() ] action_terms = [ str(item).casefold() for item in intent.get("action_terms_any", []) if str(item).strip() ] if subject_terms and not any( term in normalized_question for term in subject_terms ): continue if action_terms and not any( term in normalized_question for term in action_terms ): continue selected.append(contract) return tuple(selected) def query_contract_guidance(question: str, tool_name: str) -> str: contracts = matching_query_contracts(question, tool_name) if not contracts: return "" return json.dumps( {"query_contracts": contracts}, ensure_ascii=False, separators=(",", ":"), ) def append_query_contract_guidance( tool_query: str, *, original_question: str, tool_name: str, ) -> str: guidance = query_contract_guidance(original_question, tool_name) normalized = str(tool_query or "").strip() if not guidance: return normalized return ( f"{normalized}\n" "다음 질의 계약의 필드·계산·시간 기준을 반드시 지켜 결과를 반환하세요. " f"계약: {guidance}" ) def _json_data_result(value: str) -> Any: marker = "DATA_RESULT" position = value.find(marker) if position < 0: return None remainder = value[position + len(marker) :] object_position = remainder.find("{") array_position = remainder.find("[") positions = [ candidate for candidate in (object_position, array_position) if candidate >= 0 ] if not positions: return None candidate = remainder[min(positions) :] try: parsed, _ = json.JSONDecoder().raw_decode(candidate) except (TypeError, ValueError, json.JSONDecodeError): return None return parsed def _evidence_rows(value: Any) -> list[Mapping[str, Any]]: rows: list[Mapping[str, Any]] = [] if isinstance(value, Mapping): rows.append(value) for item in value.values(): rows.extend(_evidence_rows(item)) elif isinstance(value, list): for item in value: rows.extend(_evidence_rows(item)) elif isinstance(value, str): parsed = _json_data_result(value) if parsed is not None: rows.extend(_evidence_rows(parsed)) return rows def _decimal(value: Any) -> Decimal: if isinstance(value, bool) or value is None: raise InvalidOperation return Decimal(str(value)) def _evaluate_contract_expression( expression: str, row: Mapping[str, Any], ) -> Decimal: tree = ast.parse(expression, mode="eval") def evaluate(node: ast.AST) -> Decimal: if isinstance(node, ast.Expression): return evaluate(node.body) if isinstance(node, ast.Name): key = node.id.casefold() if key not in row: raise InvalidOperation return _decimal(row[key]) if isinstance(node, ast.Constant): return _decimal(node.value) if isinstance(node, ast.UnaryOp) and isinstance( node.op, (ast.UAdd, ast.USub) ): value = evaluate(node.operand) return value if isinstance(node.op, ast.UAdd) else -value if isinstance(node, ast.BinOp) and isinstance( node.op, (ast.Add, ast.Sub, ast.Mult, ast.Div) ): left = evaluate(node.left) right = evaluate(node.right) if isinstance(node.op, ast.Add): return left + right if isinstance(node.op, ast.Sub): return left - right if isinstance(node.op, ast.Mult): return left * right return left / right raise ValueError("unsupported contract expression") return evaluate(tree) def evidence_contract_report( contracts: tuple[Mapping[str, Any], ...], evidence: Any, ) -> list[dict[str, Any]]: rows = [ {str(key).strip().casefold(): value for key, value in row.items()} for row in _evidence_rows(evidence) ] reports: list[dict[str, Any]] = [] for contract in contracts: data_contract = contract.get("data_contract") if not isinstance(data_contract, Mapping): continue required = [ str(item).strip().casefold() for item in data_contract.get("required_fields", []) if str(item).strip() ] matching_row = next( (row for row in rows if required and all(key in row for key in required)), None, ) observed = sorted( { key for row in rows for key in row if not required or key in required } ) missing = [ key.upper() for key in required if matching_row is None or key not in matching_row ] computed_checks: list[dict[str, Any]] = [] computed_fields = data_contract.get("computed_fields") if matching_row is not None and isinstance(computed_fields, Mapping): for field, definition in computed_fields.items(): normalized_field = str(field).strip().casefold() expression = ( str(definition.get("expression") or "").strip() if isinstance(definition, Mapping) else "" ) if not normalized_field or not expression: continue try: actual = _decimal(matching_row.get(normalized_field)) expected = _evaluate_contract_expression( expression, matching_row, ) matches = actual == expected computed_checks.append( { "field": normalized_field.upper(), "expression": expression, "actual": str(actual), "expected": str(expected), "satisfied": matches, } ) except (InvalidOperation, ValueError, ZeroDivisionError): computed_checks.append( { "field": normalized_field.upper(), "expression": expression, "satisfied": False, } ) calculations_satisfied = all( bool(check.get("satisfied")) for check in computed_checks ) temporal_checks: list[dict[str, Any]] = [] temporal_contract = contract.get("temporal_contract") if matching_row is not None and isinstance(temporal_contract, Mapping): status_fields = [ str(item).strip().casefold() for item in temporal_contract.get("period_status_fields", []) if str(item).strip() ] past_status = str( temporal_contract.get("past_period_status") or "" ).strip() observed_status = next( ( str(matching_row.get(field) or "").strip() for field in status_fields if str(matching_row.get(field) or "").strip() ), "", ) if observed_status == past_status and past_status: decision_field = str( temporal_contract.get("past_period_decision_field") or "" ).strip().casefold() decision_value = str( temporal_contract.get("past_period_decision_value") or "" ).strip() actual_decision = str( matching_row.get(decision_field) or "" ).strip() temporal_checks.append( { "check": "past_period_decision", "field": decision_field.upper(), "expected": decision_value, "actual": actual_decision, "satisfied": bool( decision_field and decision_value and actual_decision == decision_value ), } ) forbidden_patterns = [ str(item).strip() for item in temporal_contract.get( "forbidden_past_period_decision_field_patterns", [] ) if str(item).strip() ] forbidden_fields = sorted( key.upper() for key in matching_row if any( re.search(pattern, key, flags=re.IGNORECASE) for pattern in forbidden_patterns ) ) temporal_checks.append( { "check": "no_past_period_yes_no_decision", "forbidden_fields": forbidden_fields, "satisfied": not forbidden_fields, } ) required_counts = [ str(item).strip().casefold() for item in temporal_contract.get( "validate_stated_days_against", [] ) if str(item).strip() ] temporal_checks.append( { "check": "date_range_counts", "required_fields": [ field.upper() for field in required_counts ], "missing_fields": [ field.upper() for field in required_counts if field not in matching_row ], "satisfied": all( field in matching_row for field in required_counts ), } ) temporal_satisfied = all( bool(check.get("satisfied")) for check in temporal_checks ) reports.append( { "contract_id": str(contract.get("id") or ""), "satisfied": bool( matching_row is not None and not missing and calculations_satisfied and temporal_satisfied ), "required_fields": [key.upper() for key in required], "observed_required_fields": [key.upper() for key in observed], "missing_fields": missing, "computed_field_checks": computed_checks, "temporal_contract_checks": temporal_checks, "validated_record": ( { key.upper(): matching_row.get(key) for key in required } if matching_row is not None else {} ), } ) return reports def missing_evidence_message( contracts: tuple[Mapping[str, Any], ...], ) -> str: for contract in contracts: answer_contract = contract.get("answer_contract") if not isinstance(answer_contract, Mapping): continue message = str(answer_contract.get("missing_evidence_message") or "").strip() if message: return message return "조회 결과가 답변 계약의 필수 근거를 충족하지 않아 결론을 제공할 수 없습니다." __all__ = [ "append_query_contract_guidance", "evidence_contract_report", "load_query_contracts", "matching_query_contracts", "missing_evidence_message", "query_contract_guidance", ]