ParthKulshreshtha's picture
Upload folder using huggingface_hub
324b1af verified
Raw
History Blame Contribute Delete
14.7 kB
"""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"