Files
abap-llm/harness/proxy.py

242 lines
13 KiB
Python

"""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)
# The model sometimes names its own helper or test class with the prefix inside the name (ZCL_Z4CGT1HJ_JOB_COST_TEST). Such objects of
# other runs stay behind in A4H and show up in lists and searches (57 of 90 accepted trajectories saw them, 2026-10-06).
MID_PREFIX = re.compile(r"^[A-Z]{1,5}_(Z\d[0-9A-Z]{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()
if RUN_PREFIX.match(n):
return not n.startswith(self.prefix)
m = MID_PREFIX.match(n)
return bool(m) and m.group(1).upper() != 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", "sap_inactive_objects"):
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 or tool in ("sap_pull_source", "sap_object_structure", "sap_object_members", "sap_element_info",
"sap_run_unit_test", "sap_check_object", "sap_syntax_check", "sap_atc_run")) \
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()