"""The 12-tool registry (docs/spec/tools.json) over an agent-visible Instance. Ground truth stays out by construction: this module takes Instance only and must never name the held-out carrier (tested via source inspection). """ from __future__ import annotations import ast import math import operator from collections.abc import Sequence from je_validation.envir.instance import Instance from je_validation.envir.state import VERDICTS, Disposition, EpisodeState from je_validation.envir.toolspecs import normalize_args _OPS = {ast.Add: operator.add, ast.Sub: operator.sub, ast.Mult: operator.mul, ast.Div: operator.truediv, ast.USub: operator.neg, ast.UAdd: operator.pos, ast.Pow: operator.pow, ast.Mod: operator.mod} class ToolError(Exception): pass MAX_EXPR_LEN = 200 MAX_EXPONENT = 64 MAX_ABS_VALUE = 1e30 # far above any ledger math; far below blowup _LISTABLE_ON_MISS = frozenset({"close_calendar"}) # Leak test (panel 2026-08-08): a type may list its ids on a miss ONLY if the # id set is identical for every defect placement. close_calendar passes # (periods are population facts). valid_combinations and chart_of_accounts # stay silent: their id sets interact with defect planting and # account_not_found probing. def _bounded(v): if isinstance(v, float) and not math.isfinite(v): raise ToolError("recompute: result is not finite") if abs(v) > MAX_ABS_VALUE: raise ToolError("recompute: intermediate result too large") return v def _safe_eval(expression: str) -> float: if len(expression) > MAX_EXPR_LEN: raise ToolError("recompute: expression too long") def ev(node): if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)): return node.value if isinstance(node, ast.BinOp) and type(node.op) in _OPS: left, right = ev(node.left), ev(node.right) if isinstance(node.op, ast.Pow) and abs(right) > MAX_EXPONENT: raise ToolError("recompute: exponent too large") return _bounded(_OPS[type(node.op)](left, right)) if isinstance(node, ast.UnaryOp) and type(node.op) in _OPS: return _bounded(_OPS[type(node.op)](ev(node.operand))) raise ToolError(f"recompute: unsupported syntax {ast.dump(node)[:40]}") try: return float(ev(ast.parse(expression, mode="eval").body)) except (SyntaxError, ZeroDivisionError, OverflowError) as e: raise ToolError(f"recompute failed: {e}") from e def _loose_eq(row_val, filter_val) -> bool: """Equality that survives the int-vs-string split: generators store e.g. period as int 7 while agents pass "7" (and vice versa). Exact match first; else compare canonical strings — but never across bools, and never when either side is None (missing field ≠ any value).""" if row_val == filter_val: return True if row_val is None or filter_val is None: return False if isinstance(row_val, bool) or isinstance(filter_val, bool): return False return str(row_val) == str(filter_val) class ToolRegistry: def __init__(self, instance: Instance, state: EpisodeState, enabled: Sequence[str], row_cap: int = 200, issue_types: Sequence[str] | None = None): self._inst = instance self._state = state self.enabled = tuple(enabled) self._entries = {e.entry_id: e for e in instance.entries} self._docs = {d.doc_id: d for d in instance.documents} fields = {"entry_id"} for e in instance.entries: fields.update(e.header) self._fields = frozenset(fields) self._row_cap = row_cap self._issue_types = tuple(issue_types) if issue_types else None self._impl = { "query_ledger": self._query_ledger, "get_entry": self._get_entry, "list_documents": self._list_documents, "open_document": self._open_document, "get_policy": self._get_policy, "get_master_data": self._get_master_data, "recompute": self._recompute, "aggregate": self._aggregate, "compare_period": self._compare_period, "request_info": self._request_info, "disposition": self._disposition, "submit": self._submit, } def call(self, name: str, args: dict) -> object: if name not in self._impl: raise ToolError(f"unknown tool: {name}") if name not in self.enabled: raise ToolError(f"tool withheld this run: {name}") if self._state.submitted: raise ToolError("episode already submitted") if not isinstance(args, dict): raise ToolError(f"{name}: args must be an object, " f"got {type(args).__name__}") try: args = normalize_args(name, args) except ValueError as e: raise ToolError(str(e)) from None try: return self._impl[name](args) except ToolError: raise except Exception as e: # structure was already validated above, so anything that still # blows up in an impl is an environment bug — never report it # as the agent's "bad arguments" raise ToolError( f"environment error in {name}: {type(e).__name__}: {e}") from e def _row(self, e): return {"entry_id": e.entry_id, **e.header} def _unknown_field(self, ctx: str, name: str) -> ToolError: return ToolError(f"{ctx}: unknown field '{name}'; available: " f"{', '.join(sorted(self._fields))}") def _select(self, filters: dict) -> list[dict]: for key in filters: base = key[:-3] if key.endswith("_gt") else key if base not in self._fields: raise self._unknown_field("filter", base) rows = [self._row(e) for e in self._inst.entries] for key, val in filters.items(): if key.endswith("_gt"): if isinstance(val, bool) or not isinstance(val, (int, float)): raise ToolError(f"{key}: greater-than filter needs a " f"number, got {val!r}") base = key[:-3] rows = [r for r in rows if isinstance(r.get(base), (int, float)) and not isinstance(r.get(base), bool) and r[base] > val] else: rows = [r for r in rows if _loose_eq(r.get(key), val)] return rows def _query_ledger(self, args): filters = {k: v for k, v in args.items() if k not in ("order_by", "offset", "limit")} rows = self._select(filters) # normalize_args dumps optional keys as explicit None — treat as absent order = str(args.get("order_by") or "entry_id") field, _, direction = order.partition(" ") if field not in self._fields: raise self._unknown_field("order_by", field) # type-aware key: a str() key sorts "1000" below "9" and corrupts any # numeric order_by; numbers sort numerically, the rest stringly after def sort_key(r): v = r.get(field) if isinstance(v, (int, float)) and not isinstance(v, bool): return (0, float(v), "") return (1, 0.0, str(v)) rows.sort(key=sort_key, reverse=direction == "desc") rows = rows[int(args.get("offset") or 0):] if args.get("limit") is not None: rows = rows[: int(args["limit"])] return rows[: self._row_cap] def _get_entry(self, args): e = self._entries.get(args.get("entry_id")) if e is None: raise ToolError(f"unknown entry: {args.get('entry_id')}") return {"entry_id": e.entry_id, **e.header, "lines": list(e.lines)} def _list_documents(self, args): if args.get("entry_id") not in self._entries: raise ToolError(f"unknown entry: {args.get('entry_id')}") return [{"doc_id": d.doc_id, "doc_type": d.doc_type, "filename": d.filename, "pages": len(d.pages)} for d in self._inst.documents if d.entry_id == args["entry_id"]] def _open_document(self, args): d = self._docs.get(args.get("doc_id")) if d is None: raise ToolError(f"unknown doc: {args.get('doc_id')}") page = args.get("page") if page is None: text = "\n".join(d.pages) else: try: idx = int(page) except (TypeError, ValueError): raise ToolError(f"page must be an integer 1..{len(d.pages)}, " f"got {page!r}") from None if not 1 <= idx <= len(d.pages): raise ToolError(f"page {idx} out of range 1..{len(d.pages)}") text = d.pages[idx - 1] self._state.opened_docs.add(d.doc_id) return text def _get_policy(self, args): topic = args.get("topic") if topic not in self._inst.policies: available = ", ".join(sorted(self._inst.policies)) or "none" raise ToolError(f"no policy on: {topic!r}; available topics: " f"{available}") return self._inst.policies[topic] def _get_master_data(self, args): tables = self._inst.master_data t = args.get("type") if t not in tables: available = ", ".join(sorted(tables)) or "none" raise ToolError(f"unknown master-data type: {t!r}; available: " f"{available}") records = tables[t] if args.get("id") not in records: if t in _LISTABLE_ON_MISS: raise ToolError(f"no {t} record: {args.get('id')}; " f"available ids: {', '.join(sorted(records))}") raise ToolError(f"no {t} record: {args.get('id')}") return records[args["id"]] def _recompute(self, args): return _safe_eval(str(args.get("expression", ""))) def _aggregate(self, args): group_by = args.get("group_by") fields = [group_by] if isinstance(group_by, str) else list(group_by) for f in fields: if f not in self._fields: raise self._unknown_field("group_by", str(f)) metric = args.get("metric", "count") filters = args.get("filters") or {} for k in ("order_by", "offset", "limit"): if k in filters: raise ToolError(f"aggregate: '{k}' is not a filter; aggregate " "always runs over the full population") if metric in ("sum", "first_digit") and "amount" not in self._fields: raise ToolError("aggregate: this ledger has no 'amount' field") rows = self._select(filters) # uncapped groups: dict[tuple, list] = {} for r in rows: key = tuple(str(r.get(f, "all")) for f in fields) groups.setdefault(key, []).append(r) def label(key): return key[0] if len(fields) == 1 else list(key) if metric == "first_digit": out = [] for key, rs in sorted(groups.items()): counts = {str(d): 0 for d in range(1, 10)} for r in rs: cents = round(abs(float(r.get("amount") or 0)) * 100) if cents: counts[str(cents)[0]] += 1 out.extend({"group": label(key), "digit": d, "count": n} for d, n in counts.items()) return out return [{"group": label(key), "value": len(rs) if metric == "count" else sum(r.get("amount", 0) for r in rs)} for key, rs in sorted(groups.items())] def _compare_period(self, args): for needed in ("period", "amount"): if needed not in self._fields: raise ToolError(f"compare_period: this ledger has no " f"'{needed}' field") account = args.get("account") if account is not None and "account" in self._fields: known = {str(e.header.get("account")) for e in self._inst.entries} if str(account) not in known: raise ToolError(f"unknown account: {account}") base = self._select({"account": account} if account is not None else {}) periods = sorted({str(r.get("period")) for r in (self._row(e) for e in self._inst.entries)}) def total(period): if str(period) not in periods: raise ToolError(f"unknown period {period!r}; available: " f"{', '.join(periods)}") return sum(r.get("amount", 0) for r in base if str(r.get("period")) == str(period)) a, b = args.get("period_a"), args.get("period_b") return {"period_a": a, "period_b": b, "delta": total(a) - total(b)} def _request_info(self, args): eid = args.get("entry_id") if eid not in self._entries: raise ToolError(f"unknown entry: {eid}") return self._inst.info_responses.get(eid, "UNAVAILABLE") def _disposition(self, args): entry_id, verdict = args.get("entry_id"), args.get("verdict") if entry_id not in self._entries: raise ToolError(f"unknown entry: {entry_id}") if verdict not in VERDICTS: raise ToolError(f"invalid verdict: {verdict}") issue = str(args.get("issue_type", "")) if self._issue_types is not None: if issue not in (*self._issue_types, "none", ""): raise ToolError( f"invalid issue_type: {issue!r}; must be one of: " f"{', '.join(self._issue_types)} (or 'none' when " "approving)") if verdict in ("flag", "reject") and issue in ("", "none"): raise ToolError( "flag/reject requires an issue_type; one of: " + ", ".join(self._issue_types)) if entry_id in self._state.dispositions: raise ToolError(f"already dispositioned: {entry_id}") self._state.dispositions[entry_id] = Disposition( entry_id, verdict, str(args.get("issue_type", "")), str(args.get("rationale", "")), tuple(args.get("evidence_ids", [])), str(args.get("counterpart_entry_id", ""))) return "ack" def _submit(self, args): self._state.submitted = True return "episode ends"