"""使用真实 HTTP/socket 验证领取屏障,不连接线上 Gitea。""" import http.client import http.server import gzip import importlib.util import json import os from pathlib import Path import socket import tempfile import threading import time import unittest if os.name != "posix": raise unittest.SkipTest("领取屏障使用 Linux Unix socket 和目录 fsync;在 Linux/WSL 运行") SPEC = importlib.util.spec_from_file_location( "gitea_cache_gate", Path(__file__).with_name("gitea-runner-fetch-gate.py")) MODULE = importlib.util.module_from_spec(SPEC) SPEC.loader.exec_module(MODULE) FETCH = "/api/actions/runner.v1.RunnerService/FetchTask" UPDATE = "/api/actions/runner.v1.RunnerService/UpdateTask" LOG = "/api/actions/runner.v1.RunnerService/UpdateLog" PING = "/api/actions/ping.v1.PingService/Ping" def varint(value): data = bytearray() while value > 127: data.append((value & 127) | 128) value >>= 7 data.append(value) return bytes(data) def field(number, value): # actions-proto-go v0.4.1: task/state ID=1; result=2; log task_id=1, # index=2, rows=3, no_more=4; log response ack_index=1. if isinstance(value, bytes): return varint(number * 8 + 2) + varint(len(value)) + value return varint(number * 8) + varint(value) def state(task_id, result=0): return field(1, field(1, task_id) + field(2, result)) def core_state(state): return {key: state[key] for key in ("paused", "inflight", "uncertain")} class Upstream(http.server.BaseHTTPRequestHandler): def log_message(self, *_args): pass def do_POST(self): data = self.rfile.read(int(self.headers.get("Content-Length", 0))) self.server.requests.append((self.path, data)) if self.path.endswith("/Redirect"): self.send_response(302) self.send_header("Location", "/api/actions/runner.v1.RunnerService/FetchTask") self.send_header("Content-Length", "0") self.end_headers() return if self.path == FETCH: self.server.entered.set() self.server.release.wait(10) if self.server.truncated and self.path == FETCH: self.send_response(200) self.send_header("Content-Length", "100") self.end_headers() self.wfile.write(b"incomplete") self.close_connection = True return body = self.server.responses.get(self.path, b"") if self.path == UPDATE and self.path not in self.server.responses: body = state(*map(int, MODULE.task_state(MODULE.protobuf_fields(data)))) if self.path == LOG and self.path not in self.server.responses: fields = MODULE.protobuf_fields(data) body = field(1, MODULE.single(fields, 2, 0) + len(fields.get(3, []))) self.send_response(self.server.statuses.get(self.path, 200)) self.send_header("Content-Type", self.server.content_type) if self.server.compressed: body = gzip.compress(body) self.send_header("Content-Encoding", "gzip") self.send_header("Content-Length", str(len(body))) self.end_headers() self.wfile.write(body) class GateTests(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.gate = MODULE.Gate(self.temp.name) self.upstream = http.server.ThreadingHTTPServer(("127.0.0.1", 0), Upstream) self.upstream.requests = [] self.upstream.entered = threading.Event() self.upstream.release = threading.Event() self.upstream.truncated = False self.upstream.responses = {} self.upstream.statuses = {} self.upstream.content_type = "application/proto" self.upstream.compressed = False self.proxy = MODULE.create_proxy(("127.0.0.1", 0), f"http://127.0.0.1:{self.upstream.server_port}", self.gate) self.servers = [self.upstream, self.proxy] for server in self.servers: threading.Thread(target=server.serve_forever, daemon=True).start() def tearDown(self): self.upstream.release.set() for server in reversed(self.servers): server.shutdown() server.server_close() self.temp.cleanup() def request(self, path, body=b"", chunked=False): connection = http.client.HTTPConnection("127.0.0.1", self.proxy.server_port, timeout=5) try: if chunked: connection.request("POST", path, body=[body], encode_chunked=True, headers={"Content-Type": "application/proto"}) else: connection.request("POST", path, body=body, headers={"Content-Type": "application/proto"}) response = connection.getresponse() result = response.status, response.read() return result finally: connection.close() def wait_for(self, predicate): deadline = time.monotonic() + 3 while time.monotonic() < deadline: if predicate(): return time.sleep(0.01) self.fail("condition did not become true") def start_fetch(self): self.result = [] thread = threading.Thread(target=lambda: self.result.append(self.request(FETCH))) thread.start() self.assertTrue(self.upstream.entered.wait(3)) return thread def test_pause_holds_inflight_and_allows_reporting(self): thread = self.start_fetch() self.assertEqual(core_state(self.gate.control("pause")), {"paused": True, "inflight": 1, "uncertain": False}) self.assertEqual(self.request(FETCH)[0], 503) self.assertEqual(self.request(PING)[0], 200) self.assertEqual(sum(path == FETCH for path, _ in self.upstream.requests), 1) self.upstream.release.set() thread.join(3) self.assertFalse(thread.is_alive()) self.assertEqual(self.result[0][0], 200) self.assertEqual(self.gate.control("status")["inflight"], 0) self.assertEqual(self.gate.control("resume")["paused"], False) self.assertEqual(self.request(FETCH)[0], 200) def test_disconnected_client_does_not_release_upstream_request(self): client = socket.create_connection(self.proxy.server_address, timeout=3) client.sendall(f"POST {FETCH} HTTP/1.1\r\nHost: localhost\r\nContent-Length: 2\r\n\r\n{{}}".encode()) self.assertTrue(self.upstream.entered.wait(3)) client.close() state = self.gate.control("pause") self.assertEqual(state["inflight"], 1) self.upstream.release.set() self.wait_for(lambda: self.gate.control("status")["inflight"] == 0) self.assertFalse(self.gate.control("status")["uncertain"]) def test_truncated_upstream_latches_uncertainty_across_restart(self): self.upstream.truncated = True self.upstream.release.set() self.assertEqual(self.request(FETCH)[0], 502) self.wait_for(lambda: self.gate.control("status")["uncertain"]) self.assertTrue(self.gate.control("resume")["paused"]) restored = MODULE.Gate(self.temp.name) self.assertTrue(restored.control("status")["uncertain"]) self.assertFalse(restored.enter()) def test_crashed_inflight_is_not_treated_as_successful_drain(self): self.assertTrue(self.gate.enter()) restored = MODULE.Gate(self.temp.name) self.assertTrue(restored.control("status")["uncertain"]) self.assertIn("error", restored.control("resume")) self.gate.leave(True) def test_pause_marker_survives_restart(self): self.gate.control("pause") restored = MODULE.Gate(self.temp.name) self.assertEqual(core_state(restored.control("status")), {"paused": True, "inflight": 0, "uncertain": False}) restored.control("resume") self.assertFalse(MODULE.Gate(self.temp.name).control("status")["paused"]) def test_chunked_body_is_decoded_and_non_rpc_path_rejected(self): self.assertEqual(self.request(PING, b"ping", chunked=True)[0], 200) self.assertEqual(self.upstream.requests[-1], (PING, b"ping")) self.assertEqual(self.request("/api/v1/repos")[0], 404) self.assertEqual(self.request("/api/actions/../v1/repos")[0], 404) def test_upstream_redirect_is_not_followed(self): self.assertEqual(self.request("/api/actions/runner.v1.RunnerService/Redirect")[0], 302) self.assertFalse(self.upstream.entered.is_set()) self.assertEqual(len(self.upstream.requests), 1) def test_connection_refused_before_rpc_does_not_latch_uncertainty(self): with socket.socket() as unused: unused.bind(("127.0.0.1", 0)) port = unused.getsockname()[1] self.proxy.upstream = MODULE.urllib.parse.urlsplit(f"http://127.0.0.1:{port}") self.assertEqual(self.request(FETCH)[0], 502) self.assertEqual(core_state(self.gate.control("status")), {"paused": False, "inflight": 0, "uncertain": False}) @unittest.skipUnless(hasattr(socket, "AF_UNIX"), "Unix control socket required") def test_control_socket_uses_newline_json(self): path = str(Path(self.temp.name) / "gate.sock") control = MODULE.ControlServer(path, MODULE.Control) control.gate = self.gate self.servers.append(control) threading.Thread(target=control.serve_forever, daemon=True).start() with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as client: client.connect(path) client.sendall(b'{"action":"pause"}\n') payload = client.makefile("rb").readline() self.assertEqual(core_state(json.loads(payload)), {"paused": True, "inflight": 0, "uncertain": False}) def test_paused_fetch_records_live_route_and_restart_discards_proof(self): self.gate.control("pause") before = time.time() self.assertEqual(self.request(FETCH)[0], 503) state = self.gate.control("status") self.assertEqual(state["last_fetch_peer"], "127.0.0.1") self.assertGreaterEqual(state["last_fetch_at"], before) self.assertEqual(self.upstream.requests, []) self.assertEqual(self.request(PING)[0], 200) self.assertEqual(self.gate.control("status")["last_fetch_at"], state["last_fetch_at"]) restarted = MODULE.Gate(self.temp.name).control("status") self.assertIsNone(restarted["last_fetch_peer"]) self.assertIsNone(restarted["last_fetch_at"]) def fetch_task(self, task_id=42): self.upstream.responses[FETCH] = field(1, field(1, task_id)) + field(2, 10) self.upstream.release.set() self.assertEqual(self.request(FETCH)[0], 200) self.assertEqual(self.gate.control("status")["task_ids"], [str(task_id)]) def finalize_log(self, task_id=42): self.assertEqual(self.request(LOG, field(1, task_id) + field(2, 3) + field(3, b"log row") + field(4, 1))[0], 200) self.wait_for(lambda: self.gate.tasks.get(str(task_id)) is True) def test_assigned_task_survives_pause_until_runner_finished_reporting(self): self.fetch_task() self.assertEqual(self.gate.control("pause")["active_tasks"], 1) self.assertEqual(self.request(FETCH)[0], 503) self.assertEqual(self.request(UPDATE, state(42, 1))[0], 200) self.assertEqual(self.gate.control("status")["active_tasks"], 1) self.finalize_log() self.assertEqual(self.gate.control("status")["active_tasks"], 1) self.assertEqual(self.request(UPDATE, state(42, 1))[0], 200) self.wait_for(lambda: self.gate.control("status")["active_tasks"] == 0) self.assertFalse(self.gate.control("status")["uncertain"]) # 终态响应在客户端超时后重试,不应再次阻断后续任务领取。 self.request(UPDATE, state(42, 1)) self.assertFalse(self.gate.control("status")["uncertain"]) def test_compressed_fetch_and_task_ledger_survive_restart(self): self.upstream.compressed = True self.fetch_task() restored = MODULE.Gate(self.temp.name) self.assertEqual(restored.control("status")["task_ids"], ["42"]) self.assertFalse(restored.control("status")["uncertain"]) self.proxy.gate = self.gate = restored self.finalize_log() self.request(UPDATE, state(42, 2)) self.wait_for(lambda: self.gate.control("status")["active_tasks"] == 0) self.assertEqual(MODULE.Gate(self.temp.name).control("status")["active_tasks"], 0) def test_server_cancellation_does_not_retire_executing_task(self): self.fetch_task() self.upstream.responses[UPDATE] = state(42, 3) self.request(UPDATE, state(42)) self.assertEqual(self.gate.control("status")["active_tasks"], 1) self.finalize_log() self.request(UPDATE, state(42)) self.assertEqual(self.gate.control("status")["active_tasks"], 1) self.request(UPDATE, state(42, 3)) self.wait_for(lambda: self.gate.control("status")["active_tasks"] == 0) def test_partial_final_log_acknowledgement_keeps_task_active(self): self.fetch_task() self.upstream.responses[LOG] = field(1, 3) self.request(LOG, field(1, 42) + field(2, 3) + field(3, b"row") + field(4, 1)) self.request(UPDATE, state(42, 1)) self.assertEqual(self.gate.control("status")["active_tasks"], 1) del self.upstream.responses[LOG] self.finalize_log() self.request(UPDATE, state(42, 1)) self.wait_for(lambda: self.gate.control("status")["active_tasks"] == 0) def test_final_state_waits_for_success_and_output_acknowledgement(self): self.fetch_task() self.finalize_log() self.upstream.statuses[UPDATE] = 500 self.request(UPDATE, state(42, 1)) self.assertEqual(self.gate.control("status")["active_tasks"], 1) self.upstream.statuses[UPDATE] = 200 output = field(2, field(1, b"key") + field(2, b"value")) self.request(UPDATE, state(42, 1) + output) self.assertEqual(self.gate.control("status")["active_tasks"], 1) self.upstream.responses[UPDATE] = state(42, 1) + field(2, b"key") self.request(UPDATE, state(42, 1) + output) self.wait_for(lambda: self.gate.control("status")["active_tasks"] == 0) def test_unknown_or_invalid_fetch_result_prevents_idle(self): self.upstream.release.set() self.upstream.responses[FETCH] = b"not protobuf" self.assertEqual(self.request(FETCH)[0], 200) self.assertTrue(self.gate.control("status")["uncertain"]) self.assertTrue(MODULE.Gate(self.temp.name).control("status")["uncertain"]) def test_task_report_without_observed_assignment_requires_idle_bootstrap(self): self.assertEqual(self.request(UPDATE, state(42))[0], 200) self.wait_for(lambda: self.gate.control("status")["uncertain"]) self.assertIn("error", self.gate.control("resume")) def test_disconnected_fetch_client_still_leaves_assigned_task_busy(self): self.upstream.responses[FETCH] = field(1, field(1, 42)) client = socket.create_connection(self.proxy.server_address, timeout=3) client.sendall(f"POST {FETCH} HTTP/1.1\r\nHost: localhost\r\nContent-Length: 0\r\n\r\n".encode()) self.assertTrue(self.upstream.entered.wait(3)) client.close() self.gate.control("pause") self.upstream.release.set() self.wait_for(lambda: self.gate.control("status")["inflight"] == 0) self.assertEqual(self.gate.control("status")["task_ids"], ["42"]) if __name__ == "__main__": unittest.main()