Files
Genarrative/scripts/gitea-runner-fetch-gate.py
T
lhk229 d7ddf5fbc1
Project CI / AI game creator shell Rust lane 1/2 (pull_request) Has been cancelled
Project CI / AI game creator shell Rust lane 2/2 (pull_request) Has been cancelled
Project CI / AI game creator shell Rust crates (pull_request) Has been cancelled
Project CI / Backend tests (pull_request) Has been cancelled
Project CI / Native shell tests (pull_request) Has been cancelled
Project CI / Frontend tests (pull_request) Has been cancelled
Project CI / Repository checks (pull_request) Has been cancelled
Project CI / AI game creator shell web tests (pull_request) Has been cancelled
修复 CI 领取冲突导致网关永久锁定
识别 Gitea 已回滚的任务分配冲突,允许 runner 正常重试。
保留未知响应与未完成领取的保护,增加脱敏原因日志。
补充冲突重试及错误边界测试,同步部署文档与共享记忆。
2026-10-07 09:20:04 +00:00

481 lines
19 KiB
Python

#!/usr/bin/env python3
"""Gitea Runner 领取屏障;控制 socket 仅供宿主缓存维护任务使用。"""
import http.client
import http.server
import gzip
import io
import json
import os
from pathlib import Path
import re
import socketserver
import sys
import threading
import time
import urllib.parse
import zlib
MAX_BODY = 32 * 1024 * 1024
HOP_HEADERS = {
"connection", "keep-alive", "proxy-authenticate", "proxy-authorization",
"te", "trailer", "transfer-encoding", "upgrade", "host", "content-length",
}
RPC_PREFIX = "/api/actions/runner.v1.RunnerService/"
def fetch_conflict_rolled_back(status, body, headers):
"""仅识别 Gitea 1.26.4 在任务分配事务提交前返回的 run 更新冲突。"""
# PickTask/WithTx -> CreateTaskForRunner -> UpdateRunJob -> UpdateRun。
# 该错误会回滚任务分配;其它 5xx、截断响应和未知协议仍须阻断领取。
if (status != 500 or len(body) > 4096
or headers.get("Content-Encoding", "identity").lower() != "identity"
or headers.get("Content-Type", "").split(";", 1)[0] != "application/json"):
return False
try:
error = json.loads(body)
except (ValueError, UnicodeError):
return False
return (isinstance(error, dict) and error.get("code") == "unknown"
and isinstance(error.get("message"), str)
and re.fullmatch(
r"rpc error: code = Internal desc = pick task: CreateTaskForRunner: "
r"update run [1-9][0-9]*: run has changed", error["message"]) is not None)
def diagnostic(message):
# 仅传固定分类,不输出上游正文、URL 或 RPC 凭据。
print(f"[runner-fetch-gate] {message}", file=sys.stderr, flush=True)
def protobuf_fields(data):
"""只解码已核对的 actions-proto-go v0.4.1 字段,不引入 protobuf 运行时。"""
fields = {}
offset = 0
def varint():
nonlocal offset
value = 0
for shift in range(0, 70, 7):
if offset >= len(data):
raise ValueError("truncated protobuf")
byte = data[offset]
offset += 1
value |= (byte & 127) << shift
if byte < 128:
return value
raise ValueError("invalid protobuf varint")
while offset < len(data):
tag = varint()
number, wire = tag >> 3, tag & 7
if not number:
raise ValueError("invalid protobuf tag")
if wire == 0:
value = varint()
elif wire in (1, 2, 5):
length = varint() if wire == 2 else (8 if wire == 1 else 4)
if offset + length > len(data):
raise ValueError("truncated protobuf field")
value = data[offset:offset + length]
offset += length
else:
raise ValueError("unsupported protobuf wire type")
fields.setdefault(number, []).append(value)
return fields
def single(fields, number, default=None):
values = fields.get(number, [])
if len(values) > 1:
raise ValueError("ambiguous tracking field")
return values[0] if values else default
def rpc_fields(body, headers):
encoding = headers.get("Content-Encoding", "identity").lower()
if encoding == "gzip":
with gzip.GzipFile(fileobj=io.BytesIO(body)) as stream:
body = stream.read(MAX_BODY + 1)
elif encoding != "identity":
raise ValueError("unsupported RPC encoding")
if len(body) > MAX_BODY:
raise ValueError("decoded RPC too large")
if headers.get("Content-Type", "").split(";", 1)[0] != "application/proto":
raise ValueError("task tracking requires Connect protobuf")
return protobuf_fields(body)
def task_state(fields):
nested = single(fields, 1)
if not isinstance(nested, bytes):
raise ValueError("missing task state")
state = protobuf_fields(nested)
task_id = positive_id(single(state, 1))
result = single(state, 2, 0)
if type(result) is not int or result not in range(5):
raise ValueError("unknown task result")
return task_id, result
def positive_id(value):
if type(value) is not int or not 0 < value < 2 ** 63:
raise ValueError("invalid task ID")
return str(value)
class Gate:
def __init__(self, directory):
self.directory = Path(directory)
self.directory.mkdir(parents=True, exist_ok=True)
self.lock = threading.Lock()
self.inflight = 0
self.last_fetch_peer = None
self.last_fetch_at = None
self.tasks = {}
self.uncertain = (self.directory / "uncertain").exists()
try:
if (self.directory / "tasks.json").exists():
tasks = json.loads((self.directory / "tasks.json").read_text())
if (not isinstance(tasks, dict) or any(
positive_id(int(key)) != key or type(value) is not bool
for key, value in tasks.items())):
raise ValueError("invalid task ledger")
self.tasks = tasks
except (ValueError, TypeError, OSError):
self.uncertain = True
self.mark("uncertain")
if (self.directory / "inflight").exists():
self.uncertain = True
self.mark("uncertain")
self.paused = (self.directory / "paused").exists() or self.uncertain
def mark(self, name):
with (self.directory / name).open("w", encoding="ascii") as stream:
stream.write("1\n")
stream.flush()
os.fsync(stream.fileno())
self.sync_directory()
def sync_directory(self):
descriptor = os.open(self.directory, os.O_RDONLY | os.O_DIRECTORY)
try:
os.fsync(descriptor)
finally:
os.close(descriptor)
def snapshot(self):
return {"paused": self.paused, "inflight": self.inflight,
"active_tasks": len(self.tasks), "task_ids": sorted(self.tasks),
"uncertain": self.uncertain, "last_fetch_peer": self.last_fetch_peer,
"last_fetch_at": self.last_fetch_at}
def save_tasks(self):
temporary = self.directory / "tasks.json.tmp"
with temporary.open("w", encoding="ascii") as stream:
json.dump(self.tasks, stream, sort_keys=True)
stream.flush()
os.fsync(stream.fileno())
os.replace(temporary, self.directory / "tasks.json")
self.sync_directory()
def fail_closed(self, reason):
with self.lock:
if not self.uncertain:
diagnostic(f"uncertain: {reason}; operator recovery required")
self.uncertain = self.paused = True
self.mark("uncertain")
def assigned(self, fields):
task = single(fields, 1)
if task is None:
return
if not isinstance(task, bytes):
raise ValueError("invalid fetched task")
task_id = positive_id(single(protobuf_fields(task), 1))
with self.lock:
self.tasks[task_id] = False
# 必须先持久化,再把领取结果交给 runner;仅保存 ID,不保存 secrets。
self.save_tasks()
def reported(self, method, request, response):
if method == "UpdateLog":
task_id = positive_id(single(request, 1))
index = single(request, 2, 0)
ack = single(response, 1, 0)
no_more = single(request, 4, 0)
if (type(index) is not int or type(ack) is not int
or index < 0 or ack < 0 or no_more not in (0, 1)):
raise ValueError("invalid log acknowledgement")
finalized = no_more == 1 and ack == index + len(request.get(3, []))
else:
task_id, result = task_state(request)
response_id, response_result = task_state(response)
if response_id != task_id:
raise ValueError("mismatched task acknowledgement")
# 取消响应可能出现在任务执行途中,必须等 runner 自己报告终态。
finalized = result != 0 and response_result != 0
output_keys = {single(protobuf_fields(entry), 1, b"")
for entry in request.get(2, [])}
finalized = finalized and output_keys.issubset(set(response.get(2, [])))
with self.lock:
if task_id not in self.tasks:
# 客户端可能未读到已成功发出的终态响应而重试;v2.0.0 仅在
# executor 清理后的 Close 中发送终态,不为幂等重报增加墓碑账本。
if method == "UpdateTask" and finalized:
return
raise ValueError("report for untracked task; idle bootstrap required")
if method == "UpdateLog" and finalized:
self.tasks[task_id] = True
self.save_tasks()
elif method == "UpdateTask" and finalized and self.tasks[task_id]:
# act_runner Reporter.Close 在 executor 清理之后先封存日志再报终态。
del self.tasks[task_id]
self.save_tasks()
def record_fetch(self, peer):
with self.lock:
self.last_fetch_peer = peer
self.last_fetch_at = time.time()
def control(self, action):
with self.lock:
if action == "pause":
self.mark("paused")
self.paused = True
elif action == "resume":
if self.uncertain:
return {**self.snapshot(), "error": "upstream completion uncertain; operator recovery required"}
(self.directory / "paused").unlink(missing_ok=True)
self.sync_directory()
self.paused = False
elif action != "status":
return {**self.snapshot(), "error": "unknown action"}
return self.snapshot()
def enter(self):
with self.lock:
if self.paused or self.uncertain:
return False
# 必须先落盘再转发;崩溃后不能把遗留的领取请求误认为已完成。
self.mark("inflight")
self.inflight += 1
return True
def leave(self, completed):
with self.lock:
self.inflight -= 1
if not completed:
if not self.uncertain:
diagnostic("uncertain: incomplete FetchTask response; operator recovery required")
self.uncertain = self.paused = True
self.mark("uncertain")
if self.inflight == 0 and not self.uncertain:
(self.directory / "inflight").unlink(missing_ok=True)
self.sync_directory()
def read_body(stream, headers):
"""解码 HTTP 请求;拒绝含糊 framing,不依赖下游连接关闭。"""
encodings = headers.get_all("Transfer-Encoding", [])
lengths = headers.get_all("Content-Length", [])
if encodings and lengths:
raise ValueError("ambiguous request framing")
if len(lengths) > 1 or len(encodings) > 1:
raise ValueError("duplicate request framing")
def exact(length):
data = stream.read(length)
if len(data) != length:
raise ValueError("incomplete request body")
return data
if not encodings:
length = int(lengths[0]) if lengths else 0
if not 0 <= length <= MAX_BODY:
raise ValueError("request too large")
return exact(length)
if encodings[0].strip().lower() != "chunked":
raise ValueError("unsupported transfer encoding")
body = bytearray()
while True:
line = stream.readline(8193)
if len(line) > 8192 or not line.endswith(b"\r\n"):
raise ValueError("invalid chunk header")
length = int(line.split(b";", 1)[0].strip(), 16)
if length < 0 or len(body) + length > MAX_BODY:
raise ValueError("request too large")
if length == 0:
trailer_size = 0
while True:
line = stream.readline(8193)
trailer_size += len(line)
if trailer_size > 8192 or not line.endswith(b"\r\n"):
raise ValueError("invalid request trailers")
if line == b"\r\n":
return bytes(body)
body.extend(exact(length))
if exact(2) != b"\r\n":
raise ValueError("invalid chunk terminator")
class Proxy(http.server.BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def log_message(self, *_args):
pass # 不输出 RPC 认证头、请求内容或带认证信息的 URL。
def reply(self, status, body, headers=()):
delivered = False
try:
self.send_response(status)
excluded = HOP_HEADERS | {
item.strip().lower() for key, value in headers if key.lower() == "connection"
for item in value.split(",")
}
for key, value in headers:
if key.lower() not in excluded:
self.send_header(key, value)
self.send_header("Content-Length", str(len(body)))
self.send_header("Connection", "close")
self.end_headers()
self.wfile.write(body)
self.wfile.flush()
delivered = True
except (OSError, ValueError):
pass
self.close_connection = True
return delivered
def do_POST(self):
path = urllib.parse.urlsplit(self.path)
if (path.scheme or path.netloc or not path.path.startswith("/api/actions/")
or "%" in path.path or any(p in (".", "..") for p in path.path.split("/"))):
self.reply(404, b"runner RPC only\n")
return
try:
body = read_body(self.rfile, self.headers)
except (ValueError, OSError):
self.reply(400, b"invalid request body\n")
return
method = path.path.removeprefix(RPC_PREFIX) if path.path.startswith(RPC_PREFIX) else ""
is_fetch = method == "FetchTask"
if is_fetch:
# 包括暂停时被拒绝的请求;只记录网络来源与时间以证明 daemon 路由。
self.server.gate.record_fetch(self.client_address[0])
if not self.server.gate.enter():
self.reply(503, b"runner maintenance\n")
return
completed = False
connection = None
try:
upstream = self.server.upstream
cls = http.client.HTTPSConnection if upstream.scheme == "https" else http.client.HTTPConnection
connection = cls(upstream.hostname, upstream.port, timeout=None)
try:
connection.connect()
except OSError:
# TCP/TLS 连接阶段尚未发送 RPC,不存在服务端分配事务。
completed = True
raise
excluded = HOP_HEADERS | {
name.strip().lower() for name in self.headers.get("Connection", "").split(",")
}
headers = {key: value for key, value in self.headers.items() if key.lower() not in excluded}
headers["Content-Length"] = str(len(body))
connection.request("POST", self.path, body=body, headers=headers)
response = connection.getresponse()
result = bytearray()
oversized = False
# 客户端即使断开,也继续读取上游完整响应,再释放领取屏障。
while block := response.read(65536):
if len(result) + len(block) <= MAX_BODY and not oversized:
result.extend(block)
else:
oversized = True
result.clear()
if response.length not in (None, 0):
raise http.client.IncompleteRead(bytes(result), response.length)
completed = True
if is_fetch:
try:
if not oversized and fetch_conflict_rolled_back(response.status, result, response.headers):
diagnostic("FetchTask transaction rolled back: run update conflict; runner may retry")
elif oversized or response.status != 200:
raise ValueError("unknown FetchTask result")
else:
self.server.gate.assigned(rpc_fields(bytes(result), response.headers))
except (ValueError, TypeError, OSError, EOFError, zlib.error):
self.server.gate.fail_closed("unrecognized FetchTask response or assignment tracking failure")
self.server.gate.leave(True)
is_fetch = False
if oversized:
self.reply(502, b"upstream response too large\n")
else:
# 日志封存先记账再返回,避免 runner 立即发终态时抢先读到旧账本。
if method == "UpdateLog" and response.status == 200:
self.observe_report(method, body, bytes(result), response.headers)
delivered = self.reply(response.status, bytes(result), response.getheaders())
if method == "UpdateTask" and response.status == 200 and delivered:
self.observe_report(method, body, bytes(result), response.headers)
except (OSError, http.client.HTTPException, ValueError):
self.reply(502, b"runner upstream unavailable\n")
finally:
if connection is not None:
connection.close()
if is_fetch:
self.server.gate.leave(completed)
def observe_report(self, method, request, response, headers):
try:
self.server.gate.reported(method, rpc_fields(request, self.headers),
rpc_fields(response, headers))
except (ValueError, TypeError, OSError, EOFError, zlib.error):
self.server.gate.fail_closed("task report tracking failure")
class Control(socketserver.StreamRequestHandler):
def handle(self):
try:
line = self.rfile.readline(4097)
if len(line) > 4096 or not line.endswith(b"\n"):
raise ValueError("invalid control request")
request = json.loads(line)
result = self.server.gate.control(request["action"])
except (ValueError, KeyError, TypeError, OSError):
result = {"error": "invalid control request"}
self.wfile.write(json.dumps(result).encode("utf-8") + b"\n")
class ControlServer(socketserver.ThreadingUnixStreamServer):
daemon_threads = True
def create_proxy(address, upstream, gate):
parsed = urllib.parse.urlsplit(upstream)
if (parsed.scheme not in ("http", "https") or not parsed.hostname
or parsed.username or parsed.password or parsed.path not in ("", "/")
or parsed.query or parsed.fragment):
raise ValueError("upstream must be an HTTP(S) origin")
server = http.server.ThreadingHTTPServer(address, Proxy)
server.upstream = parsed
server.gate = gate
return server
def main():
gate = Gate(os.environ.get("GITEA_GATE_CONTROL_DIR", "/control"))
socket_path = gate.directory / "gate.sock"
socket_path.unlink(missing_ok=True)
control = ControlServer(str(socket_path), Control)
control.gate = gate
os.chmod(socket_path, 0o600)
proxy = create_proxy(("0.0.0.0", 8080), os.environ.get(
"GITEA_RUNNER_UPSTREAM", "http://gitea:3000"), gate)
threading.Thread(target=control.serve_forever, daemon=True).start()
proxy.serve_forever()
if __name__ == "__main__":
main()