mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-07 16:19:09 +00:00
Harden the workflow-push approval guard against three bypasses found by the security review: - H6: the push parser passed through (returned None, unguarded) git invoked via a path (`/usr/bin/git`), a wrapper (`command`/`env` ...), or with leading global options (`git -c`, `--git-dir`, `--no-pager`). Recognize wrapped and path-qualified git as pushes, and block pushes carrying unsupported global options instead of running them unguarded. - H1: the base was fetched via the `origin` remote name, which the sandbox can split from the push destination via `remote set-url --push`. Fetch the base from the effective push URL (`git remote get-url --push`) so the base and the push target are the same authenticated repo. - H3: an unreadable head workflow tree (`ls-tree` failure) skipped the guard; fail closed (block) instead, mirroring the base-read path. Adds tests for each. All confirmed guard bypasses are caught by the langsmith unelevated-token backstop today; these close the guard's own logic for non-langsmith providers too. Claude-Session: https://claude.ai/code/session_01GxSndB7VoGQyeS196eUr5E
751 lines
27 KiB
Python
751 lines
27 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]")
|
|
# Command wrappers that can prefix a `git` invocation (e.g. `command git push`,
|
|
# `env FOO=bar git push`, `/usr/bin/git push`). Recognized so a wrapped or path-qualified
|
|
# git push cannot slip past the guard as "not a git command".
|
|
_GIT_WRAPPERS = {"command", "env", "nice", "sudo", "stdbuf", "nohup", "time", "ionice", "setsid"}
|
|
|
|
|
|
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 _git_invocation_args(tokens: list[str]) -> list[str] | None:
|
|
"""Return the args following the `git` executable if these tokens invoke git, else None.
|
|
|
|
Recognizes bareword `git`, path invocations like `/usr/bin/git`, and simple command
|
|
wrappers (`command`/`env`/`nice`/`sudo`/...), including `env NAME=VALUE ...`. This keeps
|
|
a wrapped or path-qualified `git push` from being mistaken for a non-git command and
|
|
passed through the guard unchecked.
|
|
"""
|
|
idx = 0
|
|
while idx < len(tokens):
|
|
tok = tokens[idx]
|
|
if tok.rsplit("/", 1)[-1] == "git":
|
|
return tokens[idx + 1 :]
|
|
if tok in _GIT_WRAPPERS:
|
|
idx += 1
|
|
# Skip wrapper options and `NAME=VALUE` assignments (e.g. `env FOO=bar git ...`).
|
|
while idx < len(tokens) and (tokens[idx].startswith("-") or "=" in tokens[idx]):
|
|
idx += 1
|
|
continue
|
|
return None
|
|
return None
|
|
|
|
|
|
def _parse_git_tokens(
|
|
tokens: list[str], *, repo_dir: str | None
|
|
) -> ParsedGitPush | _BlockedGitPush | None:
|
|
git_args = _git_invocation_args(tokens)
|
|
if git_args is None:
|
|
return None
|
|
if "push" not in git_args:
|
|
# A non-push git command (e.g. `git status`) is not our concern.
|
|
return None
|
|
i = 0
|
|
while i < len(git_args) and git_args[i] != "push":
|
|
if git_args[i] == "-C" and i + 1 < len(git_args):
|
|
repo_dir = git_args[i + 1]
|
|
i += 2
|
|
continue
|
|
# A git push carrying unsupported global options (e.g. `git -c k=v push`,
|
|
# `git --git-dir=.git push`). We cannot safely normalize it, so fail closed.
|
|
return _BlockedGitPush(
|
|
"unsupported git options before `push`; use `git push origin <branch>`"
|
|
)
|
|
return _parse_push_args(git_args[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 exact URL the push will target and return a local alias.
|
|
|
|
The base is fetched from the remote's effective *push* URL (`git remote get-url --push`),
|
|
not the remote name, because git honors a separate `pushurl`: an untrusted sandbox can
|
|
point the fetch URL at attacker/local content while the push still lands on the real
|
|
repo. Fetching the base from the push URL keeps the base and the push destination the
|
|
same authenticated repo, so a split cannot hide a workflow change. Uses `FETCH_HEAD`
|
|
(freshly fetched), never a local `refs/remotes/origin/*` ref the sandbox could rewrite.
|
|
"""
|
|
push_url_res = _run_git(backend, repo_dir, f"remote get-url --push {shlex.quote(remote)}")
|
|
push_url = _first_line(push_url_res.output) if push_url_res.ok else ""
|
|
if not push_url:
|
|
return None
|
|
|
|
fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(push_url)} {shlex.quote(branch)}")
|
|
if fetch.ok:
|
|
return "FETCH_HEAD"
|
|
|
|
default = _run_git(backend, repo_dir, f"ls-remote --symref {shlex.quote(push_url)} 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(push_url)} {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:
|
|
# The workflow tree at head could not be read (an empty tree returns ([], hash), so
|
|
# None means the `ls-tree` read genuinely failed). Fail closed rather than skipping
|
|
# the guard, mirroring the base-read path.
|
|
return WorkflowPushChange(
|
|
fingerprint="",
|
|
repo="",
|
|
branch=branch_name,
|
|
files=[],
|
|
head_sha=head,
|
|
remote=parsed.remote,
|
|
local_ref=parsed.local_ref,
|
|
remote_ref=parsed.remote_ref,
|
|
fixed_command="",
|
|
blocked=True,
|
|
blocked_reason="could not read the workflow tree at the pushed head; blocking to be safe",
|
|
)
|
|
|
|
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)
|