| """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 |
|
|
| _LISTABLE_ON_MISS = frozenset({"close_calendar"}) |
| |
| |
| |
| |
| |
|
|
|
|
| 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: |
| |
| |
| |
| 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) |
| |
| 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) |
| |
| |
| 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) |
| 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" |
|
|