mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 15:03:16 +00:00
906 lines
34 KiB
Python
906 lines
34 KiB
Python
"""Gate workflow-file pushes on human approval."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import codecs
|
|
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]")
|
|
# Unquoted shell expansion/substitution (`$VAR`, `${...}`, `$(...)`, backticks) can produce
|
|
# a `git push` whose text we cannot see; its presence alongside a push forces fail-closed.
|
|
_EXPANSION = re.compile(r"[$`]")
|
|
# Unquoted command separators and grouping. Splitting the unquoted command skeleton on these
|
|
# yields the individual simple-commands the shell would run, each checked for a `git push`.
|
|
_SEGMENT_SPLIT = re.compile(r"[;|&(){}\n]+")
|
|
# 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"}
|
|
# Git global options that consume the following token as a separate value; skipped when
|
|
# locating the git subcommand so `git -c k=v push` is still recognized as a push.
|
|
_GIT_VALUE_OPTIONS = {"-C", "-c", "--namespace", "--git-dir", "--work-tree", "--exec-path"}
|
|
|
|
|
|
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 _tokens_invoke_git_push(tokens: list[str]) -> bool:
|
|
"""True if the token stream runs `git push` as a git subcommand (not merely mentions it).
|
|
|
|
Scans each `git` invocation (bareword or path-qualified, after any wrapper prefix),
|
|
skips `-C <dir>` and option flags, and checks whether the first positional argument is
|
|
`push`. This ignores the word "push" appearing inside a commit message or another
|
|
command, so `git commit -m "push it" && ls` is not treated as a push, while a genuine
|
|
`git ... push` anywhere in a chain is.
|
|
"""
|
|
i = 0
|
|
while i < len(tokens):
|
|
if tokens[i].rsplit("/", 1)[-1] != "git":
|
|
i += 1
|
|
continue
|
|
j = i + 1
|
|
while j < len(tokens):
|
|
tok = tokens[j]
|
|
if tok in _SHELL_OPERATORS or tok == "&&":
|
|
break
|
|
if (
|
|
tok in _GIT_VALUE_OPTIONS
|
|
and j + 1 < len(tokens)
|
|
and tokens[j + 1] not in _SHELL_OPERATORS
|
|
and tokens[j + 1] != "&&"
|
|
):
|
|
j += 2 # option consumes the next token as its value
|
|
continue
|
|
if tok.startswith("-"):
|
|
j += 1
|
|
continue
|
|
if tok == "push":
|
|
return True
|
|
break # first positional is a non-push subcommand
|
|
i = j + 1
|
|
return False
|
|
|
|
|
|
def _decode_ansi_c(body: str) -> str:
|
|
"""Best-effort ANSI-C (`$'...'`) escape decoding, so `$'\\x70ush'` / `$'\\160ush'` reveal
|
|
the literal `push` the shell would run. Falls back to the raw body if decoding fails."""
|
|
if "\\" not in body:
|
|
return body
|
|
try:
|
|
return codecs.decode(body, "unicode_escape")
|
|
except Exception:
|
|
return body
|
|
|
|
|
|
def _strip_quoted(command: str) -> str:
|
|
"""Return the command with quoted spans and escaped characters blanked to spaces.
|
|
|
|
Leaves only the unquoted shell structure, so metacharacters that are literal (inside
|
|
quotes or backslash-escaped — e.g. a `;` or `&` inside a commit message) do not look
|
|
like command separators. Plain `'...'`/`"..."` spans wrap arguments and are blanked;
|
|
`$'...'`/`$"..."` (ANSI-C / locale quoting) form a literal *word* (e.g. `git $'push'` runs
|
|
`git push`), so their content is kept — otherwise the push word would vanish. Unbalanced
|
|
quotes blank the remainder (fail-closed shape).
|
|
"""
|
|
out: list[str] = []
|
|
i, n = 0, len(command)
|
|
while i < n:
|
|
c = command[i]
|
|
if c in "'\"":
|
|
# `$'...'` / `$"..."` — the preceding `$` marks a literal word, so keep content
|
|
# and drop the `$` so adjacent pieces concatenate (`$'pus'$'h'` -> `push`).
|
|
word_quote = bool(out) and out[-1] == "$"
|
|
if word_quote:
|
|
out.pop()
|
|
if c == "'":
|
|
j = command.find("'", i + 1)
|
|
if j == -1:
|
|
return "".join(out) + " "
|
|
out.append(_decode_ansi_c(command[i + 1 : j]) if word_quote else " ")
|
|
i = j + 1
|
|
else:
|
|
k = i + 1
|
|
buf: list[str] = []
|
|
while k < n and command[k] != '"':
|
|
if command[k] == "\\" and k + 1 < n:
|
|
buf.append(command[k + 1])
|
|
k += 2
|
|
else:
|
|
buf.append(command[k])
|
|
k += 1
|
|
if k >= n:
|
|
return "".join(out) + " "
|
|
out.append("".join(buf) if word_quote else " ")
|
|
i = k + 1
|
|
elif c == "\\" and i + 1 < n:
|
|
out.append(" ")
|
|
i += 2
|
|
else:
|
|
out.append(c)
|
|
i += 1
|
|
return "".join(out)
|
|
|
|
|
|
def _maybe_obfuscated_git_push(command: str) -> bool:
|
|
"""True if the shell would run a `git push` that `_tokens_invoke_git_push` could not see.
|
|
|
|
Only reached when the precise token scan found no clean push. `shlex` splits on
|
|
whitespace and performs no expansion, so a push can hide behind a fused separator
|
|
(`true;git push`), subshell/grouping (`(git push ...)`), or expansion (`git${IFS}push`,
|
|
`git $(printf push)`). We work on the unquoted command skeleton so literal metacharacters
|
|
inside a commit message are ignored (no false positive on `git commit -m "push it" && ls`
|
|
or `git log | grep push`), then fail closed if any unquoted simple-command is a `git push`
|
|
or if unquoted expansion could produce one.
|
|
"""
|
|
skeleton = _strip_quoted(command)
|
|
if "push" not in skeleton:
|
|
# The only `push` text is inside quotes (a literal argument), so it cannot be a
|
|
# command word. (A push spelled purely by expansion with no literal `push` anywhere
|
|
# is left to the GitHub proxy token scope, which lacks workflows:write until approval.)
|
|
return False
|
|
if _EXPANSION.search(skeleton):
|
|
return True
|
|
return any(
|
|
_tokens_invoke_git_push(segment.split()) for segment in _SEGMENT_SPLIT.split(skeleton)
|
|
)
|
|
|
|
|
|
def _parse_git_push(command: str) -> ParsedGitPush | _BlockedGitPush | None:
|
|
stripped = command.strip()
|
|
try:
|
|
probe = shlex.split(stripped)
|
|
except ValueError:
|
|
probe = None
|
|
if probe is None or not _tokens_invoke_git_push(probe):
|
|
# No push the token scan can normalize. Fail closed if one is hidden behind shell
|
|
# metacharacters/expansion; otherwise leave the command untouched.
|
|
if _maybe_obfuscated_git_push(stripped):
|
|
return _BlockedGitPush(
|
|
"possible git push obscured by shell metacharacters; issue it as a plain "
|
|
"`git push origin <branch>`"
|
|
)
|
|
return None
|
|
if _UNSAFE_RAW_COMMAND.search(stripped) or "&" in stripped.replace("&&", ""):
|
|
return _BlockedGitPush("unsafe shell characters in git push command")
|
|
tokens = probe
|
|
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:
|
|
"""Parse tokens into a supported git push, a block sentinel, or None.
|
|
|
|
Recognizes bareword `git`, path invocations (`/usr/bin/git`), and wrapper-prefixed
|
|
forms (`command`/`env`/`nice`/... git). Anything git-push-shaped that cannot be reduced
|
|
to `git [-C <dir>] push [-u] origin <refspec>` fails closed to a block rather than a
|
|
silent pass-through. Non-push git commands and non-git commands return None.
|
|
"""
|
|
git_idx = next((idx for idx, tok in enumerate(tokens) if tok.rsplit("/", 1)[-1] == "git"), None)
|
|
if git_idx is None:
|
|
return None
|
|
git_args = tokens[git_idx + 1 :]
|
|
if "push" not in git_args:
|
|
# A non-push git command (e.g. `git status`, `git pull`).
|
|
return None
|
|
# Everything before the `git` executable must be a benign wrapper prefix (a known
|
|
# wrapper, an option flag, or a NAME=VALUE assignment). An unrecognized leading token
|
|
# (a value-taking wrapper option like the `10` in `nice -n 10`, an unknown wrapper, or
|
|
# `git` used as an argument to another program) is ambiguous, so fail closed.
|
|
if any(
|
|
not (tok in _GIT_WRAPPERS or tok.startswith("-") or "=" in tok) for tok in tokens[:git_idx]
|
|
):
|
|
return _BlockedGitPush(
|
|
"unrecognized wrapper before `git push`; use `git push origin <branch>`"
|
|
)
|
|
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.
|
|
|
|
Returns None (which makes the caller require approval) if the push destination is
|
|
ambiguous or rewritten: multiple `pushurl` entries (git pushes to ALL of them, so a
|
|
single base cannot represent the destination) or any `insteadOf`/`pushInsteadOf` URL
|
|
rewrite the sandbox could use to make the inspected URL differ from the push target.
|
|
"""
|
|
rewrites = _run_git(
|
|
backend, repo_dir, "config --get-regexp " + shlex.quote(r"url\..*\.(push)?insteadof")
|
|
)
|
|
if rewrites.ok and rewrites.output.strip():
|
|
# A URL rewrite means the inspected push URL may not be the real destination.
|
|
return None
|
|
|
|
push_url_res = _run_git(backend, repo_dir, f"remote get-url --push --all {shlex.quote(remote)}")
|
|
push_urls = (
|
|
[line.strip() for line in push_url_res.output.splitlines() if line.strip()]
|
|
if push_url_res.ok
|
|
else []
|
|
)
|
|
# `git push` sends to every configured push URL; if there is not exactly one we cannot
|
|
# represent the destination with a single base, so fail closed.
|
|
if len(push_urls) != 1:
|
|
return None
|
|
push_url = push_urls[0]
|
|
|
|
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)
|