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
Reviewed-on: https://git.genarrative.world/git/GenarrativeAI/Genarrative/pulls/462
448 lines
18 KiB
Python
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()
|