242 lines
13 KiB
Python
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()
|