Files
Genarrative/scripts/gitea-runner-fetch-gate.py
T
lhk229 705bb1e6b3
Project CI / AI game creator shell Rust lane 1/2 (push) Has been cancelled
Project CI / AI game creator shell Rust lane 2/2 (push) Has been cancelled
Project CI / AI game creator shell Rust smoke (push) Has been cancelled
Project CI / AI game creator shell Rust crates (push) Has been cancelled
Project CI / Backend tests (push) Has been cancelled
Project CI / Native shell tests (push) Has been cancelled
Project CI / Frontend tests (push) Has been cancelled
Project CI / Repository checks (push) Has been cancelled
Project CI / AI game creator shell web tests (push) Has been cancelled
CI缓存自动维护与清理 (#462)
Reviewed-on: https://git.genarrative.world/git/GenarrativeAI/Genarrative/pulls/462
2026-09-22 17:19:11 +08:00

448 lines
18 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 socketserver
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 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):
with self.lock:
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:
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 oversized or response.status != 200:
raise ValueError("unknown FetchTask result")
self.server.gate.assigned(rpc_fields(bytes(result), response.headers))
except (ValueError, TypeError, OSError, EOFError, zlib.error):
self.server.gate.fail_closed()
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()
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()