mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-07 16:19:09 +00:00
The head-tree fingerprint keeps the security win (approval binds to the exact workflow files and blob SHAs at the pushed head), but the guard was firing on every push because it no longer compared against a base. This change re-adds change detection using a base fetched from the authenticated remote at guard time: - Fetches the pushed branch from the remote; if it does not exist (new branch), fetches the remote's default branch via ls-remote and a fallback chain. - Compares the head workflow tree (ls-tree) against the freshly fetched base (FETCH_HEAD), not against any local refs/remotes/origin/* ref. - Returns None (no guard) when the workflow trees are identical, so code-only pushes to a branch that already contains workflow files are not blocked. - Updated test_non_workflow_push_runs_without_approval to use a non-empty-but unchanged workflow tree. Refs: 98
694 lines
24 KiB
Python
694 lines
24 KiB
Python
"""Gate workflow-file pushes on human approval."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import json
|
|
import logging
|
|
import os
|
|
import re
|
|
import shlex
|
|
import threading
|
|
from collections.abc import Awaitable, Callable, Mapping
|
|
from dataclasses import dataclass
|
|
from typing import Any
|
|
|
|
from langchain.agents.middleware.types import AgentMiddleware, AgentState
|
|
from langchain_core.messages import ToolMessage
|
|
from langgraph.config import get_config
|
|
from langgraph.prebuilt.tool_node import ToolCallRequest
|
|
from langgraph.types import Command
|
|
|
|
from ..dashboard.workflow_approval import (
|
|
ensure_workflow_push_pending,
|
|
find_workflow_push_approval,
|
|
mark_workflow_push_notified,
|
|
workflow_push_approved,
|
|
workflow_push_rejected,
|
|
)
|
|
from ..tools.slack_thread_reply import build_workflow_approval_blocks
|
|
from ..utils.github_app import (
|
|
BASE_RUNTIME_PROXY_TOKEN_PERMISSIONS,
|
|
RUNTIME_PROXY_TOKEN_PERMISSIONS,
|
|
WORKFLOW_RUNTIME_PROXY_TOKEN_PERMISSIONS,
|
|
)
|
|
from ..utils.github_proxy import refresh_proxy_token
|
|
from ..utils.sandbox_state import SANDBOX_BACKENDS
|
|
from ..utils.slack import post_slack_thread_reply_with_ts
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_WORKFLOW_PREFIX = ".github/workflows/"
|
|
_SHELL_OPERATORS = {";", "|", "||", "&"}
|
|
_REF_NAME = re.compile(r"^[A-Za-z0-9._/@+-]+$")
|
|
_GIT_OBJECT_ID = re.compile(r"^[0-9a-fA-F]{40,64}$")
|
|
_UNSAFE_RAW_COMMAND = re.compile(r"[;|`$<>\n\r]")
|
|
|
|
|
|
class _BlockedGitPush:
|
|
"""Sentinel returned by the parser when a git push command is unsafe or unrecognized."""
|
|
|
|
__slots__ = ("reason",)
|
|
|
|
def __init__(self, reason: str) -> None:
|
|
self.reason = reason
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ParsedGitPush:
|
|
repo_dir: str | None
|
|
remote: str
|
|
local_ref: str
|
|
remote_ref: str
|
|
set_upstream: bool = False
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class WorkflowPushChange:
|
|
fingerprint: str
|
|
repo: str
|
|
branch: str
|
|
files: list[str]
|
|
head_sha: str
|
|
remote: str
|
|
local_ref: str
|
|
remote_ref: str
|
|
fixed_command: str
|
|
base_sha: str = ""
|
|
blocked: bool = False
|
|
blocked_reason: str = ""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class GitInspectResult:
|
|
output: str
|
|
ok: bool
|
|
|
|
|
|
def _tool_name(request: ToolCallRequest) -> str | None:
|
|
tool_call = getattr(request, "tool_call", None)
|
|
if isinstance(tool_call, Mapping):
|
|
name = tool_call.get("name")
|
|
return name if isinstance(name, str) else None
|
|
return None
|
|
|
|
|
|
def _tool_args(request: ToolCallRequest) -> dict[str, Any]:
|
|
tool_call = getattr(request, "tool_call", None)
|
|
args = tool_call.get("args") if isinstance(tool_call, Mapping) else None
|
|
return dict(args) if isinstance(args, Mapping) else {}
|
|
|
|
|
|
def _tool_call_id(request: ToolCallRequest) -> str | None:
|
|
tool_call = getattr(request, "tool_call", None)
|
|
if isinstance(tool_call, Mapping):
|
|
value = tool_call.get("id")
|
|
return value if isinstance(value, str) else None
|
|
return None
|
|
|
|
|
|
def _config(request: ToolCallRequest) -> Mapping[str, Any]:
|
|
runtime_config = getattr(getattr(request, "runtime", None), "config", None)
|
|
if isinstance(runtime_config, Mapping):
|
|
return runtime_config
|
|
try:
|
|
config = get_config()
|
|
except Exception:
|
|
return {}
|
|
return config if isinstance(config, Mapping) else {}
|
|
|
|
|
|
def _configurable(request: ToolCallRequest) -> Mapping[str, Any]:
|
|
config = _config(request)
|
|
configurable = config.get("configurable")
|
|
return configurable if isinstance(configurable, Mapping) else {}
|
|
|
|
|
|
def _thread_id(request: ToolCallRequest) -> str | None:
|
|
thread_id = _configurable(request).get("thread_id")
|
|
return thread_id if isinstance(thread_id, str) and thread_id else None
|
|
|
|
|
|
def _backend(thread_id: str | None) -> Any | None:
|
|
return SANDBOX_BACKENDS.get(thread_id) if thread_id else None
|
|
|
|
|
|
def _response_output(response: Any) -> str:
|
|
output = getattr(response, "output", None)
|
|
if isinstance(output, str):
|
|
return output
|
|
if isinstance(response, Mapping):
|
|
value = response.get("output")
|
|
if isinstance(value, str):
|
|
return value
|
|
return str(response or "")
|
|
|
|
|
|
def _response_ok(response: Any) -> bool:
|
|
exit_code = getattr(response, "exit_code", None)
|
|
if isinstance(exit_code, int):
|
|
return exit_code == 0
|
|
if isinstance(response, Mapping):
|
|
value = response.get("exit_code")
|
|
if isinstance(value, int):
|
|
return value == 0
|
|
return True
|
|
|
|
|
|
def _parse_git_push(command: str) -> ParsedGitPush | _BlockedGitPush | None:
|
|
stripped = command.strip()
|
|
if _UNSAFE_RAW_COMMAND.search(stripped) or "&" in stripped.replace("&&", ""):
|
|
return _BlockedGitPush("unsafe shell characters in git push command")
|
|
try:
|
|
tokens = shlex.split(stripped)
|
|
except ValueError:
|
|
return _BlockedGitPush("unparseable git push command")
|
|
if not tokens:
|
|
return None
|
|
|
|
if len(tokens) >= 4 and tokens[0] == "cd" and tokens[2] == "&&":
|
|
if any(token in _SHELL_OPERATORS or token == "&&" for token in tokens[3:]):
|
|
return _BlockedGitPush("chained shell commands in git push")
|
|
return _parse_git_tokens(tokens[3:], repo_dir=tokens[1])
|
|
|
|
if any(token in _SHELL_OPERATORS or token == "&&" for token in tokens):
|
|
return _BlockedGitPush("chained shell commands")
|
|
return _parse_git_tokens(tokens, repo_dir=None)
|
|
|
|
|
|
def _parse_git_tokens(
|
|
tokens: list[str], *, repo_dir: str | None
|
|
) -> ParsedGitPush | _BlockedGitPush | None:
|
|
if not tokens or tokens[0] != "git":
|
|
return None
|
|
i = 1
|
|
while i < len(tokens) and tokens[i] != "push":
|
|
if tokens[i] == "-C" and i + 1 < len(tokens):
|
|
repo_dir = tokens[i + 1]
|
|
i += 2
|
|
continue
|
|
# Any non-push git command (e.g. git status) is not a push, so let it pass.
|
|
return None
|
|
if i >= len(tokens) or tokens[i] != "push":
|
|
return None
|
|
return _parse_push_args(tokens[i + 1 :], repo_dir=repo_dir)
|
|
|
|
|
|
def _parse_push_args(
|
|
tokens: list[str], *, repo_dir: str | None
|
|
) -> ParsedGitPush | _BlockedGitPush | None:
|
|
set_upstream = False
|
|
while tokens and tokens[0] in {"-u", "--set-upstream"}:
|
|
set_upstream = True
|
|
tokens = tokens[1:]
|
|
if len(tokens) != 2 or tokens[0] != "origin":
|
|
return _BlockedGitPush("unrecognized or unsafe git push arguments")
|
|
parsed = _parse_refspec(tokens[1])
|
|
if parsed is None:
|
|
return _BlockedGitPush("unrecognized or unsafe git push refspec")
|
|
local_ref, remote_ref = parsed
|
|
return ParsedGitPush(
|
|
repo_dir=repo_dir,
|
|
remote="origin",
|
|
local_ref=local_ref,
|
|
remote_ref=remote_ref,
|
|
set_upstream=set_upstream,
|
|
)
|
|
|
|
|
|
def _parse_refspec(refspec: str) -> tuple[str, str] | None:
|
|
if refspec.startswith("-") or ".." in refspec:
|
|
return None
|
|
if ":" in refspec:
|
|
parts = refspec.split(":")
|
|
if len(parts) != 2 or not parts[0] or not parts[1]:
|
|
return None
|
|
local_ref, remote_ref = parts
|
|
else:
|
|
local_ref = remote_ref = refspec
|
|
if not _safe_ref(local_ref, allow_head=True) or not _safe_ref(remote_ref, allow_head=False):
|
|
return None
|
|
return local_ref, remote_ref
|
|
|
|
|
|
def _safe_ref(ref: str, *, allow_head: bool) -> bool:
|
|
if allow_head and ref == "HEAD":
|
|
return True
|
|
if ref == "HEAD" or not _REF_NAME.fullmatch(ref):
|
|
return False
|
|
return not any(part in {"", ".", ".."} for part in ref.split("/"))
|
|
|
|
|
|
def _git_command(repo_dir: str | None, args: str) -> str:
|
|
if repo_dir:
|
|
return f"git -C {shlex.quote(repo_dir)} {args}"
|
|
return f"git {args}"
|
|
|
|
|
|
def _run_git(backend: Any, repo_dir: str | None, args: str) -> GitInspectResult:
|
|
try:
|
|
response = backend.execute(_git_command(repo_dir, args), timeout=30)
|
|
except Exception:
|
|
logger.debug("workflow push inspection failed for git %s", args, exc_info=True)
|
|
return GitInspectResult("", False)
|
|
return GitInspectResult(_response_output(response).strip(), _response_ok(response))
|
|
|
|
|
|
def _first_line(text: str) -> str:
|
|
for line in text.splitlines():
|
|
stripped = line.strip()
|
|
if stripped:
|
|
return stripped
|
|
return ""
|
|
|
|
|
|
def _normalize_remote(remote: str) -> str:
|
|
value = remote.strip()
|
|
if value.endswith(".git"):
|
|
value = value[:-4]
|
|
value = re.sub(r"^https://[^/@]+@github\.com/", "https://github.com/", value)
|
|
value = re.sub(r"^git@github\.com:", "https://github.com/", value)
|
|
return value
|
|
|
|
|
|
def _fingerprint(payload: Mapping[str, Any]) -> str:
|
|
encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
|
return hashlib.sha256(encoded).hexdigest()
|
|
|
|
|
|
def _run_coroutine_sync(coro: Awaitable[ToolMessage | Command]) -> ToolMessage | Command:
|
|
try:
|
|
asyncio.get_running_loop()
|
|
except RuntimeError:
|
|
return asyncio.run(coro)
|
|
|
|
result: dict[str, ToolMessage | Command | BaseException] = {}
|
|
|
|
def target() -> None:
|
|
try:
|
|
result["value"] = asyncio.run(coro)
|
|
except BaseException as exc: # noqa: BLE001
|
|
result["value"] = exc
|
|
|
|
thread = threading.Thread(target=target)
|
|
thread.start()
|
|
thread.join()
|
|
value = result["value"]
|
|
if isinstance(value, BaseException):
|
|
raise value
|
|
return value
|
|
|
|
|
|
def _parse_ls_tree(output: str) -> list[tuple[str, str]]:
|
|
"""Parse `git ls-tree -r` output into a sorted list of (sha, path) tuples."""
|
|
entries: list[tuple[str, str]] = []
|
|
for line in output.splitlines():
|
|
line = line.strip()
|
|
if not line:
|
|
continue
|
|
# Format: "<mode> <type> <sha>\t<path>"
|
|
meta, _, path = line.partition("\t")
|
|
if not path:
|
|
continue
|
|
parts = meta.split()
|
|
if len(parts) < 3:
|
|
continue
|
|
entries.append((parts[2], path))
|
|
entries.sort(key=lambda item: item[1])
|
|
return entries
|
|
|
|
|
|
def _workflow_tree_at_ref(
|
|
backend: Any, repo_dir: str | None, ref: str
|
|
) -> tuple[list[tuple[str, str]], str] | None:
|
|
"""Return the sorted workflow tree entries and a stable hash for the given ref.
|
|
|
|
Returns an empty list when the ref has no workflow files, so callers can detect
|
|
additions and deletions against a base ref.
|
|
"""
|
|
ls_tree = _run_git(backend, repo_dir, f"ls-tree -r {shlex.quote(ref)} -- .github/workflows")
|
|
if not ls_tree.ok:
|
|
return None
|
|
entries = _parse_ls_tree(ls_tree.output)
|
|
content_hash = _fingerprint({"tree": entries})
|
|
return entries, content_hash
|
|
|
|
|
|
def _fetch_remote_base(backend: Any, repo_dir: str | None, remote: str, branch: str) -> str | None:
|
|
"""Fetch the base ref from the authenticated remote and return a local alias for it.
|
|
|
|
Fetches the pushed branch first; if it does not exist on the remote, fetches the
|
|
remote's default branch. Uses `FETCH_HEAD` so the base is bound to the freshly-fetched
|
|
remote tip, not a local `refs/remotes/origin/*` ref that the sandbox could rewrite.
|
|
"""
|
|
fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(remote)} {shlex.quote(branch)}")
|
|
if fetch.ok:
|
|
return "FETCH_HEAD"
|
|
|
|
default = _run_git(backend, repo_dir, "ls-remote --symref origin HEAD")
|
|
default_branch = ""
|
|
if default.ok:
|
|
default_branch = _first_line(default.output).removeprefix("ref: refs/heads/").split("\t")[0]
|
|
|
|
for fallback in (default_branch, "dev", "main"):
|
|
if not fallback or fallback == branch:
|
|
continue
|
|
fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(remote)} {shlex.quote(fallback)}")
|
|
if fetch.ok:
|
|
return "FETCH_HEAD"
|
|
return None
|
|
|
|
|
|
def _workflow_change_for_push(backend: Any, parsed: ParsedGitPush) -> WorkflowPushChange | None:
|
|
root_result = _run_git(backend, parsed.repo_dir, "rev-parse --show-toplevel")
|
|
if not root_result.ok:
|
|
return None
|
|
root = _first_line(root_result.output)
|
|
if not root:
|
|
return None
|
|
|
|
branch = _run_git(backend, root, "rev-parse --abbrev-ref HEAD")
|
|
branch_name = _first_line(branch.output) if branch.ok else ""
|
|
if not branch_name or branch_name == "HEAD" or parsed.remote_ref != branch_name:
|
|
return None
|
|
if parsed.local_ref not in {"HEAD", branch_name}:
|
|
return None
|
|
|
|
target_sha = _run_git(backend, root, f"rev-parse {shlex.quote(parsed.local_ref)}")
|
|
head = _first_line(target_sha.output) if target_sha.ok else ""
|
|
if not head or not _GIT_OBJECT_ID.fullmatch(head):
|
|
return None
|
|
|
|
head_tree = _workflow_tree_at_ref(backend, root, head)
|
|
if head_tree is None:
|
|
return None
|
|
|
|
base_ref = _fetch_remote_base(backend, root, parsed.remote, branch_name)
|
|
if base_ref is not None:
|
|
base_tree = _workflow_tree_at_ref(backend, root, base_ref)
|
|
if base_tree is not None and head_tree[0] == base_tree[0]:
|
|
# No workflow change against the trusted remote base.
|
|
return None
|
|
|
|
head_entries, head_content_hash = head_tree
|
|
files = sorted({path for _sha, path in head_entries})
|
|
|
|
remote = _run_git(backend, root, "config --get remote.origin.url")
|
|
repo = _normalize_remote(_first_line(remote.output)) if remote.ok else ""
|
|
fixed_refspec = f"{head}:refs/heads/{parsed.remote_ref}"
|
|
fixed_args = ["push"]
|
|
if parsed.set_upstream:
|
|
fixed_args.append("--set-upstream")
|
|
fixed_args.extend([parsed.remote, fixed_refspec])
|
|
fixed_command = _git_command(root, " ".join(shlex.quote(arg) for arg in fixed_args))
|
|
content_payload = {
|
|
"repo": repo,
|
|
"branch": branch_name,
|
|
"files": files,
|
|
"content_hash": head_content_hash,
|
|
}
|
|
return WorkflowPushChange(
|
|
fingerprint=_fingerprint(content_payload),
|
|
repo=repo,
|
|
branch=branch_name,
|
|
files=files,
|
|
head_sha=head,
|
|
remote=parsed.remote,
|
|
local_ref=parsed.local_ref,
|
|
remote_ref=parsed.remote_ref,
|
|
fixed_command=fixed_command,
|
|
)
|
|
|
|
|
|
def _blocked_message(
|
|
change: WorkflowPushChange,
|
|
*,
|
|
already_rejected: bool = False,
|
|
stale: bool = False,
|
|
blocked: bool = False,
|
|
) -> ToolMessage:
|
|
status = "rejected" if already_rejected else "approval_required"
|
|
if blocked:
|
|
error = change.blocked_reason or "This git push command is not recognized as safe."
|
|
error_type = "WorkflowPushBlocked"
|
|
elif stale:
|
|
error = (
|
|
"This git push includes GitHub workflow file changes. A previous approval "
|
|
"exists for the same branch and workflow files, but the workflow content "
|
|
"at the pushed head has changed since that approval (for example, a rebase "
|
|
"that changed the workflow files or an amend that edited them). The thread "
|
|
"owner must re-approve the new fingerprint before Open SWE can push it."
|
|
)
|
|
error_type = "WorkflowPushApprovalRequired"
|
|
else:
|
|
error = (
|
|
"This git push includes GitHub workflow file changes and requires human "
|
|
"approval before Open SWE can push it. Retry the same standalone git push "
|
|
"after the thread owner approves the workflow files."
|
|
)
|
|
error_type = "WorkflowPushApprovalRequired"
|
|
content = {
|
|
"status": "error",
|
|
"error_type": error_type,
|
|
"error": error,
|
|
"workflow_approval_status": status if not blocked else "blocked",
|
|
"fingerprint": change.fingerprint,
|
|
"files": change.files,
|
|
"repo": change.repo,
|
|
"branch": change.branch,
|
|
}
|
|
return ToolMessage(content=json.dumps(content), tool_call_id="", status="error")
|
|
|
|
|
|
def _tool_message_for_request(message: ToolMessage, request: ToolCallRequest) -> ToolMessage:
|
|
message.tool_call_id = _tool_call_id(request)
|
|
return message
|
|
|
|
|
|
def _override_execute_command(request: ToolCallRequest, command: str) -> ToolCallRequest:
|
|
tool_call = getattr(request, "tool_call", None)
|
|
if not isinstance(tool_call, Mapping):
|
|
return request
|
|
args = dict(_tool_args(request))
|
|
args["command"] = command
|
|
return request.override(tool_call={**dict(tool_call), "args": args})
|
|
|
|
|
|
def _approval_slack_message(change: WorkflowPushChange) -> str:
|
|
files = "\n".join(f"• `{path}`" for path in change.files[:10])
|
|
if len(change.files) > 10:
|
|
files += f"\n• …and {len(change.files) - 10} more"
|
|
repo = change.repo or "the repository"
|
|
branch = change.branch or "the current branch"
|
|
return (
|
|
"*Workflow file approval required*\n"
|
|
f"Open SWE is trying to push the workflow files below to `{repo}` on `{branch}`.\n\n"
|
|
f"*Files at the pushed head:*\n{files}\n\n"
|
|
f"*Fingerprint:* `{change.fingerprint}`\n\n"
|
|
"Approval covers the exact workflow files and content listed above at the pushed head, "
|
|
"including future rebases or amends that replay the same workflow tree. If the set of "
|
|
"workflow files, the branch, or the workflow-file content at the pushed head changes, "
|
|
"a new fingerprint will be required."
|
|
)
|
|
|
|
|
|
async def _post_slack_approval_if_needed(
|
|
request: ToolCallRequest, change: WorkflowPushChange, record: Mapping[str, Any]
|
|
) -> None:
|
|
if record.get("notified") is True:
|
|
return
|
|
configurable = _configurable(request)
|
|
slack_thread = configurable.get("slack_thread")
|
|
if not isinstance(slack_thread, Mapping):
|
|
return
|
|
channel_id = slack_thread.get("channel_id")
|
|
thread_ts = slack_thread.get("thread_ts")
|
|
if not isinstance(channel_id, str) or not isinstance(thread_ts, str):
|
|
return
|
|
message = _approval_slack_message(change)
|
|
message_ts, error = await post_slack_thread_reply_with_ts(
|
|
channel_id,
|
|
thread_ts,
|
|
message,
|
|
blocks=build_workflow_approval_blocks(message, change.fingerprint),
|
|
)
|
|
if message_ts and not error:
|
|
thread_id = _thread_id(request)
|
|
if thread_id:
|
|
await mark_workflow_push_notified(thread_id, change.fingerprint)
|
|
|
|
|
|
async def _approval_state(request: ToolCallRequest, change: WorkflowPushChange) -> str:
|
|
thread_id = _thread_id(request)
|
|
if not thread_id:
|
|
return "missing_thread"
|
|
try:
|
|
if await workflow_push_approved(thread_id, change.fingerprint):
|
|
return "approved"
|
|
if await workflow_push_rejected(thread_id, change.fingerprint):
|
|
return "rejected"
|
|
|
|
# If the exact identity fingerprint is not approved, check whether a prior
|
|
# approval covers the same (repo, branch, files) identity. If so, the workflow
|
|
# tree at the pushed head changed underneath the prior approval (rebase/amend
|
|
# that edited workflow files), so we surface a loud re-approval message rather
|
|
# than a fresh silent pending record.
|
|
prior = await find_workflow_push_approval(
|
|
thread_id,
|
|
repo=change.repo,
|
|
branch=change.branch,
|
|
files=change.files,
|
|
)
|
|
if prior is not None:
|
|
return "stale_approval"
|
|
|
|
record, _created = await ensure_workflow_push_pending(
|
|
thread_id,
|
|
fingerprint=change.fingerprint,
|
|
repo=change.repo,
|
|
branch=change.branch,
|
|
base_sha=change.base_sha,
|
|
head_sha=change.head_sha,
|
|
files=change.files,
|
|
)
|
|
await _post_slack_approval_if_needed(request, change, record)
|
|
return str(record.get("status") or "pending")
|
|
except Exception:
|
|
logger.exception("Failed to read or write workflow push approval state")
|
|
return "approval_error"
|
|
|
|
|
|
async def _run_with_workflow_token(
|
|
thread_id: str,
|
|
request: ToolCallRequest,
|
|
run: Callable[[], Awaitable[ToolMessage | Command]],
|
|
) -> ToolMessage | Command:
|
|
sandbox_type = os.getenv("SANDBOX_TYPE", "langsmith")
|
|
if sandbox_type != "langsmith":
|
|
return await run()
|
|
|
|
elevated = await refresh_proxy_token(
|
|
thread_id, permissions=WORKFLOW_RUNTIME_PROXY_TOKEN_PERMISSIONS
|
|
)
|
|
if not elevated:
|
|
logger.error(
|
|
"Workflow push approved for thread %s, but proxy token elevation to workflows:write "
|
|
"failed; the sandbox cannot push workflow files without an elevated token.",
|
|
thread_id,
|
|
)
|
|
error_message = ToolMessage(
|
|
content=json.dumps(
|
|
{
|
|
"status": "error",
|
|
"error_type": "WorkflowPushElevationFailed",
|
|
"error": (
|
|
"Workflow push approved, but the sandbox could not obtain a "
|
|
"workflows-scoped token. Please retry the push or check the "
|
|
"GitHub proxy / token minting configuration."
|
|
),
|
|
}
|
|
),
|
|
tool_call_id="",
|
|
status="error",
|
|
)
|
|
return _tool_message_for_request(error_message, request)
|
|
try:
|
|
return await run()
|
|
finally:
|
|
restored = await refresh_proxy_token(thread_id, permissions=RUNTIME_PROXY_TOKEN_PERMISSIONS)
|
|
if not restored:
|
|
await refresh_proxy_token(thread_id, permissions=BASE_RUNTIME_PROXY_TOKEN_PERMISSIONS)
|
|
|
|
|
|
class WorkflowPushGuardMiddleware(AgentMiddleware):
|
|
"""Require approval before pushing `.github/workflows` changes."""
|
|
|
|
state_schema = AgentState
|
|
|
|
def _change_for_request(self, request: ToolCallRequest) -> WorkflowPushChange | None:
|
|
if _tool_name(request) != "execute":
|
|
return None
|
|
command = _tool_args(request).get("command")
|
|
if not isinstance(command, str):
|
|
return None
|
|
parsed = _parse_git_push(command)
|
|
if parsed is None:
|
|
return None
|
|
if isinstance(parsed, _BlockedGitPush):
|
|
return WorkflowPushChange(
|
|
fingerprint="",
|
|
repo="",
|
|
branch="",
|
|
files=[],
|
|
head_sha="",
|
|
remote="origin",
|
|
local_ref="",
|
|
remote_ref="",
|
|
fixed_command="",
|
|
blocked=True,
|
|
blocked_reason=parsed.reason,
|
|
)
|
|
backend = _backend(_thread_id(request))
|
|
if backend is None:
|
|
return None
|
|
return _workflow_change_for_push(backend, parsed)
|
|
|
|
async def _handle_change_async(
|
|
self,
|
|
request: ToolCallRequest,
|
|
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
|
|
change: WorkflowPushChange,
|
|
) -> ToolMessage | Command:
|
|
if change.blocked:
|
|
return _tool_message_for_request(_blocked_message(change, blocked=True), request)
|
|
thread_id = _thread_id(request)
|
|
state = await _approval_state(request, change)
|
|
if state == "approved" and thread_id:
|
|
safe_request = _override_execute_command(request, change.fixed_command)
|
|
return await _run_with_workflow_token(thread_id, request, lambda: handler(safe_request))
|
|
if state == "stale_approval":
|
|
record, _created = await ensure_workflow_push_pending(
|
|
thread_id,
|
|
fingerprint=change.fingerprint,
|
|
repo=change.repo,
|
|
branch=change.branch,
|
|
base_sha=change.base_sha,
|
|
head_sha=change.head_sha,
|
|
files=change.files,
|
|
)
|
|
await _post_slack_approval_if_needed(request, change, record)
|
|
return _tool_message_for_request(
|
|
_blocked_message(
|
|
change,
|
|
already_rejected=state == "rejected",
|
|
stale=state == "stale_approval",
|
|
),
|
|
request,
|
|
)
|
|
|
|
def wrap_tool_call(
|
|
self,
|
|
request: ToolCallRequest,
|
|
handler: Callable[[ToolCallRequest], ToolMessage | Command],
|
|
) -> ToolMessage | Command:
|
|
change = self._change_for_request(request)
|
|
if change is None:
|
|
return handler(request)
|
|
|
|
async def run_handler() -> ToolMessage | Command:
|
|
return handler(request)
|
|
|
|
return _run_coroutine_sync(
|
|
self._handle_change_async(request, lambda _request: run_handler(), change)
|
|
)
|
|
|
|
async def awrap_tool_call(
|
|
self,
|
|
request: ToolCallRequest,
|
|
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
|
|
) -> ToolMessage | Command:
|
|
change = self._change_for_request(request)
|
|
if change is None:
|
|
return await handler(request)
|
|
return await self._handle_change_async(request, handler, change)
|