#!/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()