"""Tool proxy between agent and MCP: whitelist, budget, prefix filter, trajectory log.""" import hashlib import json import re import time from . import localcheck 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"}, }, } def activation_messages(text, is_error=False): """Error messages (severity E) of a write result; the whole text if the call itself failed.""" try: obj = json.loads(text) except (TypeError, ValueError): return [str(text)[:300]] if is_error else [] if not isinstance(obj, dict): return [] msgs = (obj.get("activation") or {}).get("messages") or obj.get("messages") or [] out = [str(m.get("message", ""))[:300] for m in msgs if isinstance(m, dict) and m.get("severity") == "E"] if not out and (is_error or obj.get("success") is False): out = [str(obj.get("error") or obj.get("message") or text)[:300]] # {"success":false,"error":...} of a failed write return out 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, release=None): self.mcp = mcp self.release = release # task release target, for the local syntax check self.syntax_hints = 0 # failed writes that got local syntax messages 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.activation_failures = 0 # write calls that failed self.activation_errors = [] # unique error messages of failed writes self.same_push_streak = 0 # identical sap_push_source (object + source hash) in a row self._last_push = None self.same_call_streak = 0 # identical call (tool, arguments and result) in a row, any tool self._last_call = None self.log = open(log_path, "a") self.adt = None self.fallbacks = [] # writes that needed the ADT activation fallback self.last_source = {} # last main source pushed per PROG/FUNC, for the sap_activate 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 if tool == "sap_push_source": key = (name, str(args.get("includeType", "")), hashlib.md5(str(args.get("source", "")).encode()).hexdigest()) self.same_push_streak = self.same_push_streak + 1 if key == self._last_push else 1 self._last_push = key err, text = self.mcp.call(tool, args) if tool == "sap_push_source": if str(args.get("objectType", "")).upper() in ("PROG", "FUNC") and not args.get("includeType"): self.last_source[str(args.get("objectName", "")).upper()] = args.get("source", "") err, text = self._activation_fallback(args, err, text) elif tool == "sap_activate": # a write with activate=false, then sap_activate (G0014, G0124) src = self.last_source.get(str(args.get("objectName", "")).upper()) if src is not None: err, text = self._activation_fallback(dict(args, source=src), err, text) raw_text = text if tool in ("sap_push_source", "sap_push_element") and localcheck.is_bare_save_failure(text): text = self._add_syntax_messages(args, text) result = (err, self._filter(tool, text)) if text != raw_text: entry["raw_result"] = raw_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) if failed: self.activation_failures += 1 for m in activation_messages(text, err): if m not in self.activation_errors and len(self.activation_errors) < 30: self.activation_errors.append(m) if tool in MODEL_TOOLS and not result[1].startswith("Tool "): ck = (tool, hashlib.md5(json.dumps(args, sort_keys=True).encode()).hexdigest(), hashlib.md5(result[1].encode()).hexdigest()) self.same_call_streak = self.same_call_streak + 1 if ck == self._last_call else 1 self._last_call = ck entry["is_error"], entry["result"] = result[0], result[1][:20000] self.log.write(json.dumps(entry) + "\n") self.log.flush() return result def _add_syntax_messages(self, args, text): """SAP says only "save operation failed". The EPOD syntax check reads the stored version, not the rejected source, so the proxy checks the rejected source with local abaplint and adds the messages.""" msgs = localcheck.parser_messages(args.get("objectType"), args.get("objectName"), args.get("source"), args.get("includeType"), args.get("functionGroup"), self.release) if msgs: self.syntax_hints += 1 return localcheck.augment(text, msgs) 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 body = e.read().decode(errors="replace")[:600] if hasattr(e, "read") else "" self.fallbacks.append({"object": name, "error": str(e)[:200], "body": body}) 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()