"""Tool proxy between agent and MCP: whitelist, budget, prefix filter, trajectory log.""" import json import re import time MODEL_TOOLS = { "sap_search_object", "sap_pull_source", "sap_object_structure", "sap_object_members", "sap_usage_references", "sap_element_info", "sap_inactive_objects", "sap_short_dumps", "sap_sql_query", "sap_create_object", "sap_push_source", "sap_push_element", "sap_push_message", "sap_activate", "sap_syntax_check", "sap_check_object", "sap_run_unit_test", "sap_atc_run", "sap_pretty_print", "sap_run_class", } WRITE_TOOLS = {"sap_create_object", "sap_push_source", "sap_push_element", "sap_push_message", "sap_activate"} # Tool schema variants (category K; small variation in training). Draft until abap-mcp-arayuz.md exists. # The model sees the variant names and argument names; the proxy translates them to EPOD. VARIANTS = { "generic_v0": { "tools": {"search_object": "sap_search_object", "read_source": "sap_pull_source", "object_structure": "sap_object_structure", "where_used": "sap_usage_references", "create_object": "sap_create_object", "write_source": "sap_push_source", "activate": "sap_activate", "syntax_check": "sap_syntax_check", "run_unit_tests": "sap_run_unit_test", "run_atc": "sap_atc_run", "object_members": "sap_object_members", "element_info": "sap_element_info", "inactive_objects": "sap_inactive_objects", "short_dumps": "sap_short_dumps", "sql_query": "sap_sql_query", "write_element": "sap_push_element", "write_message": "sap_push_message", "check_object": "sap_check_object", "pretty_print": "sap_pretty_print", "run_class": "sap_run_class"}, "args": {"objectName": "name", "objectType": "type", "functionGroup": "function_group", "packageName": "package", "includeType": "include", "className": "class_name"}, }, } RUN_PREFIX = re.compile(r"^Z\d[0-9A-Z]{5,6}_", re.I) class BudgetExceeded(Exception): pass class ToolProxy: def __init__(self, mcp, prefix, budget, log_path, tool_schema=None): self.mcp = mcp self.variant = VARIANTS.get(tool_schema) if tool_schema else None self.started = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()) # dump times are UTC self.prefix = prefix.upper() self.max_calls = budget.get("max_tool_calls", 60) self.max_activations = budget.get("max_activations", 15) self.calls = 0 self.activations = 0 self.fail_streak = 0 self.max_fail_streak = 0 self.log = open(log_path, "a") self.adt = None self.fallbacks = [] # writes that needed the ADT activation fallback def schemas(self): tools = [t for t in self.mcp.list_tools() if t["name"] in MODEL_TOOLS] if not self.variant: return tools back = {v: k for k, v in self.variant["tools"].items()} amap = self.variant["args"] out = [] for t in tools: if t["name"] not in back: continue sch = json.loads(json.dumps(t.get("inputSchema", {}))) sch["properties"] = {amap.get(k, k): v for k, v in sch.get("properties", {}).items()} if "required" in sch: sch["required"] = [amap.get(k, k) for k in sch["required"]] desc = t.get("description", "") for epod, alias in list(back.items()) + list(amap.items()): desc = re.sub(rf"\b{re.escape(epod)}\b", alias, desc) out.append(dict(t, name=back[t["name"]], description=desc, inputSchema=sch)) return out def _translate(self, tool, args): """Variant name and arguments to EPOD. EPOD names stay valid (the oracle uses them).""" if not self.variant or tool not in self.variant["tools"]: return tool, args back = {v: k for k, v in self.variant["args"].items()} return self.variant["tools"][tool], {back.get(k, k): v for k, v in (args or {}).items()} def _foreign(self, name): n = (name or "").upper() return bool(RUN_PREFIX.match(n)) and not n.startswith(self.prefix) def _filter(self, tool, text): if tool == "sap_short_dumps": # only dumps of this run: other runs (and mutants) also write dumps try: data = json.loads(text) except ValueError: return text if isinstance(data, dict) and isinstance(data.get("dumps"), list): data["dumps"] = [d for d in data["dumps"] if str(d.get("updated", "")) >= self.started] data["count"] = len(data["dumps"]) return json.dumps(data) return text if tool not in ("sap_search_object", "sap_usage_references"): return text try: data = json.loads(text) except ValueError: return text if isinstance(data, list): data = [d for d in data if not self._foreign(d.get("name") if isinstance(d, dict) else "")] return json.dumps(data) return text def call(self, tool, args): entry = {"t": time.time(), "tool": tool, "args": args} shown = tool tool, args = self._translate(tool, args) if shown != tool: entry.update(tool=tool, args=args, variant_tool=shown) if tool not in MODEL_TOOLS: result = (True, f"Tool {tool} is not available.") elif self.calls >= self.max_calls: raise BudgetExceeded(f"max_tool_calls={self.max_calls}") else: self.calls += 1 name = str(args.get("objectName", "")).upper() if tool in WRITE_TOOLS and self._foreign(name): result = (True, f"{name} is not available.") else: if tool in WRITE_TOOLS and (tool == "sap_activate" or args.get("activate", True)): self.activations += 1 err, text = self.mcp.call(tool, args) if tool == "sap_push_source": err, text = self._activation_fallback(args, err, text) result = (err, self._filter(tool, text)) if tool in WRITE_TOOLS: failed = err or '"success":false' in text.replace(" ", "") self.fail_streak = self.fail_streak + 1 if failed else 0 self.max_fail_streak = max(self.max_fail_streak, self.fail_streak) entry["is_error"], entry["result"] = result[0], result[1][:20000] self.log.write(json.dumps(entry) + "\n") self.log.flush() return result def _activation_fallback(self, args, err, text): """EPOD does not activate a second write of a PROG or FUNC (outcome notExecuted, and sap_activate cannot help). Write and activate the same source through ADT REST; the agent sees one normal result.""" otype = str(args.get("objectType", "")).upper() if otype not in ("PROG", "FUNC") or args.get("includeType") or args.get("activate") is False \ or '"outcome":"notExecuted"' not in text.replace(" ", ""): return err, text name = str(args.get("objectName", "")).lower() uri = (f"/sap/bc/adt/programs/programs/{name}" if otype == "PROG" else f"/sap/bc/adt/functions/groups/{str(args.get('functionGroup', '')).lower()}/fmodules/{name}") try: if self.adt is None: from .adt_client import AdtClient self.adt = AdtClient() done, msgs = self.adt.write_activate(uri, name.upper(), args.get("source", "")) except Exception as e: # noqa: BLE001 keep the EPOD result when the fallback fails self.fallbacks.append({"object": name, "error": str(e)[:200]}) return err, text self.fallbacks.append({"object": name, "activated": done}) out = {"success": done, "message": "Source written and activation requested." if done else "Source written. Activation failed. See the activation messages.", "activation": {"messages": [{"severity": sv, "message": tx} for sv, tx in msgs], "activationExecuted": done, "success": done}} return False, json.dumps(out) def note(self, kind, content): self.log.write(json.dumps({"t": time.time(), "kind": kind, "content": content}) + "\n") self.log.flush()