311 lines
11 KiB
Python
311 lines
11 KiB
Python
#!/usr/bin/env python3
|
|
"""Persist Gitea workflow_job events and optionally schedule VM sandboxes."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import hmac
|
|
import json
|
|
import logging
|
|
import os
|
|
import ssl
|
|
from pathlib import Path
|
|
|
|
import nats
|
|
from aiohttp import web
|
|
from nats.errors import TimeoutError as NatsTimeoutError
|
|
from nats.js.api import (
|
|
AckPolicy,
|
|
ConsumerConfig,
|
|
DiscardPolicy,
|
|
RetentionPolicy,
|
|
StorageType,
|
|
StreamConfig,
|
|
)
|
|
from nats.js.errors import NotFoundError
|
|
|
|
from .models import IdentityBinding, RunnerRequest
|
|
from .opensandbox import OpenSandboxClient
|
|
from .opensandbox_worker import OpenSandboxScheduler, RegistrationTokens
|
|
|
|
|
|
LOG = logging.getLogger(__name__)
|
|
WEBHOOK_SECRET_FILE = Path(
|
|
os.environ.get("WEBHOOK_SECRET_FILE", "/run/secrets/gitea/webhook-secret")
|
|
)
|
|
REGISTRATION_TOKEN_FILE = Path(
|
|
os.environ.get("REGISTRATION_TOKEN_FILE", "/run/secrets/gitea/registration-token")
|
|
)
|
|
OPENSANDBOX_API = os.environ.get("OPENSANDBOX_API", "http://10.60.0.13:8080")
|
|
OPENSANDBOX_API_KEY_FILE = Path(
|
|
os.environ.get(
|
|
"OPENSANDBOX_API_KEY_FILE",
|
|
"/run/secrets/opensandbox/api-key",
|
|
)
|
|
)
|
|
NATS_URL = os.environ.get("NATS_URL", "tls://nats.ad.ddupan.top:4222")
|
|
NATS_CA_FILE = os.environ.get("NATS_CA_FILE", "/etc/ssl/certs/ca-certificates.crt")
|
|
NATS_PRODUCER_USER = os.environ.get("NATS_PRODUCER_USER", "ci-producer")
|
|
NATS_PRODUCER_PASSWORD_FILE = Path(
|
|
os.environ.get("NATS_PRODUCER_PASSWORD_FILE", "/run/secrets/nats/producer-password")
|
|
)
|
|
NATS_WORKER_USER = os.environ.get("NATS_WORKER_USER", "ci-worker")
|
|
NATS_WORKER_PASSWORD_FILE = Path(
|
|
os.environ.get("NATS_WORKER_PASSWORD_FILE", "/run/secrets/nats/worker-password")
|
|
)
|
|
NATS_STREAM = os.environ.get("NATS_STREAM", "CI_RUNNER")
|
|
NATS_SUBJECT_PREFIX = os.environ.get("NATS_SUBJECT_PREFIX", "ci.runner")
|
|
VM_CONSUMER_ENABLED = os.environ.get("VM_CONSUMER_ENABLED", "false") == "true"
|
|
VM_CAPACITY = int(os.environ.get("VM_CAPACITY", "1"))
|
|
|
|
|
|
def accepts(payload: object) -> tuple[bool, str | None]:
|
|
request = RunnerRequest.from_webhook(payload)
|
|
return (request is not None, str(request.job_id) if request else None)
|
|
|
|
|
|
def valid_signature(body: bytes, signature: str) -> bool:
|
|
expected = hmac.new(
|
|
WEBHOOK_SECRET_FILE.read_bytes().strip(), body, hashlib.sha256
|
|
).hexdigest()
|
|
return hmac.compare_digest(signature.removeprefix("sha256="), expected)
|
|
|
|
|
|
async def webhook(request: web.Request) -> web.Response:
|
|
body = await request.read()
|
|
if not valid_signature(body, request.headers.get("X-Gitea-Signature", "")):
|
|
raise web.HTTPUnauthorized()
|
|
try:
|
|
payload = json.loads(body)
|
|
except json.JSONDecodeError as error:
|
|
raise web.HTTPBadRequest(text="invalid JSON\n") from error
|
|
context = _webhook_context(payload)
|
|
LOG.info("workflow_job webhook received %s", context)
|
|
|
|
if isinstance(payload, dict) and payload.get("action") == "completed":
|
|
job = payload.get("workflow_job")
|
|
runner_name = job.get("runner_name") if isinstance(job, dict) else None
|
|
scheduler = request.app.get("scheduler")
|
|
if (
|
|
isinstance(runner_name, str)
|
|
and runner_name.startswith("gitea-vm-")
|
|
and scheduler is not None
|
|
):
|
|
cleaned = await scheduler.complete(runner_name)
|
|
LOG.info(
|
|
"completed runner cleanup runner=%r cleaned=%s %s",
|
|
runner_name,
|
|
cleaned,
|
|
context,
|
|
)
|
|
return web.Response(status=204)
|
|
|
|
runner_request = RunnerRequest.from_webhook(payload)
|
|
if runner_request is None:
|
|
binding = IdentityBinding.from_webhook(payload)
|
|
if binding is None:
|
|
LOG.info("workflow_job webhook ignored %s", context)
|
|
return web.Response(status=204)
|
|
subject = f"{NATS_SUBJECT_PREFIX}.{binding.backend}.binding"
|
|
await request.app["js"].publish(
|
|
subject,
|
|
binding.to_json(),
|
|
headers={
|
|
"Nats-Msg-Id": f"gitea-workflow-job-{binding.job_id}-in-progress"
|
|
},
|
|
)
|
|
LOG.info("runner identity binding persisted subject=%s %s", subject, context)
|
|
return web.Response(status=202, text="binding queued\n")
|
|
|
|
subject = f"{NATS_SUBJECT_PREFIX}.{runner_request.backend}"
|
|
await request.app["js"].publish(
|
|
subject,
|
|
runner_request.to_json(),
|
|
headers={"Nats-Msg-Id": f"gitea-workflow-job-{runner_request.job_id}-queued"},
|
|
)
|
|
LOG.info("runner request persisted subject=%s %s", subject, context)
|
|
return web.Response(status=202, text="queued\n")
|
|
|
|
|
|
async def registration_token(request: web.Request) -> web.Response:
|
|
scheduler = request.app.get("scheduler")
|
|
if scheduler is None:
|
|
raise web.HTTPNotFound()
|
|
value = await scheduler.tokens.consume(request.match_info["nonce"])
|
|
if value is None:
|
|
raise web.HTTPNotFound()
|
|
return web.Response(body=value, headers={"Cache-Control": "no-store"})
|
|
|
|
|
|
def _webhook_context(payload: object) -> str:
|
|
if not isinstance(payload, dict):
|
|
return f"payload_type={type(payload).__name__}"
|
|
job = payload.get("workflow_job")
|
|
repository = payload.get("repository")
|
|
job = job if isinstance(job, dict) else {}
|
|
repository = repository if isinstance(repository, dict) else {}
|
|
return (
|
|
f"action={payload.get('action')!r} job_id={job.get('id')!r} "
|
|
f"run_id={job.get('run_id')!r} runner_name={job.get('runner_name')!r} "
|
|
f"repository={repository.get('full_name')!r} job_name={job.get('name')!r} "
|
|
f"labels={job.get('labels')!r}"
|
|
)
|
|
|
|
|
|
async def health(request: web.Request) -> web.Response:
|
|
connected = request.app["nc"].is_connected
|
|
return web.Response(
|
|
text="ok\n" if connected else "disconnected\n",
|
|
status=200 if connected else 503,
|
|
)
|
|
|
|
|
|
async def ensure_stream(js: object) -> None:
|
|
config = StreamConfig(
|
|
name=NATS_STREAM,
|
|
subjects=[f"{NATS_SUBJECT_PREFIX}.>"],
|
|
retention=RetentionPolicy.WORK_QUEUE,
|
|
storage=StorageType.FILE,
|
|
discard=DiscardPolicy.OLD,
|
|
max_age=24 * 60 * 60,
|
|
max_msgs=10_000,
|
|
max_bytes=256 * 1024 * 1024,
|
|
duplicate_window=24 * 60 * 60,
|
|
)
|
|
try:
|
|
await js.stream_info(NATS_STREAM)
|
|
except NotFoundError:
|
|
await js.add_stream(config=config)
|
|
else:
|
|
await js.update_stream(config=config)
|
|
|
|
|
|
async def run_message(message: object, scheduler: OpenSandboxScheduler) -> None:
|
|
try:
|
|
runner_request = RunnerRequest.from_json(message.data)
|
|
await scheduler.create(runner_request)
|
|
task = scheduler.active[runner_request.job_id]
|
|
try:
|
|
while not task.done():
|
|
try:
|
|
await asyncio.wait_for(asyncio.shield(task), timeout=30)
|
|
except asyncio.TimeoutError:
|
|
await message.in_progress()
|
|
await task
|
|
except asyncio.CancelledError:
|
|
# A completed webhook cancels the lifecycle monitor after its
|
|
# Lifecycle DELETE succeeds. That is successful message handling.
|
|
# During controller shutdown the scheduler still owns the job, so
|
|
# preserve the unacked message for redelivery.
|
|
if runner_request.job_id in scheduler.active:
|
|
raise
|
|
except ValueError as error:
|
|
LOG.info("runner request already active: %s", error)
|
|
await message.nak(delay=5)
|
|
except Exception:
|
|
LOG.exception("persistent runner request failed")
|
|
await message.nak(delay=15)
|
|
else:
|
|
await message.ack()
|
|
|
|
|
|
async def consume_vm_requests(js: object, scheduler: OpenSandboxScheduler) -> None:
|
|
subscription = await js.pull_subscribe(
|
|
f"{NATS_SUBJECT_PREFIX}.vm",
|
|
durable="vm",
|
|
stream=NATS_STREAM,
|
|
config=ConsumerConfig(
|
|
durable_name="vm",
|
|
filter_subject=f"{NATS_SUBJECT_PREFIX}.vm",
|
|
ack_policy=AckPolicy.EXPLICIT,
|
|
ack_wait=5 * 60,
|
|
max_ack_pending=VM_CAPACITY,
|
|
max_deliver=20,
|
|
),
|
|
)
|
|
active: set[asyncio.Task[None]] = set()
|
|
while True:
|
|
active = {task for task in active if not task.done()}
|
|
try:
|
|
messages = await subscription.fetch(
|
|
batch=max(1, VM_CAPACITY - len(active)), timeout=5
|
|
)
|
|
except NatsTimeoutError:
|
|
continue
|
|
for message in messages:
|
|
task = asyncio.create_task(run_message(message, scheduler))
|
|
active.add(task)
|
|
|
|
|
|
async def runtime_context(app: web.Application):
|
|
tls = ssl.create_default_context(cafile=NATS_CA_FILE)
|
|
producer = await nats.connect(
|
|
NATS_URL,
|
|
user=NATS_PRODUCER_USER,
|
|
password=NATS_PRODUCER_PASSWORD_FILE.read_text().strip(),
|
|
tls=tls,
|
|
name="opensandbox-runner-controller-producer",
|
|
)
|
|
app["nc"] = producer
|
|
app["js"] = producer.jetstream()
|
|
await ensure_stream(app["js"])
|
|
if not VM_CONSUMER_ENABLED:
|
|
try:
|
|
yield
|
|
finally:
|
|
await producer.drain()
|
|
return
|
|
|
|
if VM_CAPACITY < 1:
|
|
raise ValueError("VM_CAPACITY must be at least 1")
|
|
tokens = RegistrationTokens(REGISTRATION_TOKEN_FILE.read_bytes())
|
|
worker = await nats.connect(
|
|
NATS_URL,
|
|
user=NATS_WORKER_USER,
|
|
password=NATS_WORKER_PASSWORD_FILE.read_text().strip(),
|
|
tls=ssl.create_default_context(cafile=NATS_CA_FILE),
|
|
name="opensandbox-vm-runner-worker",
|
|
)
|
|
async with OpenSandboxClient(
|
|
api_url=OPENSANDBOX_API,
|
|
api_key_file=OPENSANDBOX_API_KEY_FILE,
|
|
) as client:
|
|
scheduler = OpenSandboxScheduler(client, tokens)
|
|
app["opensandbox_client"] = client
|
|
app["scheduler"] = scheduler
|
|
consumer = asyncio.create_task(
|
|
consume_vm_requests(worker.jetstream(), scheduler),
|
|
name="opensandbox-vm-consumer",
|
|
)
|
|
try:
|
|
yield
|
|
finally:
|
|
consumer.cancel()
|
|
await asyncio.gather(consumer, return_exceptions=True)
|
|
await scheduler.close()
|
|
await worker.drain()
|
|
await producer.drain()
|
|
|
|
|
|
def create_app() -> web.Application:
|
|
app = web.Application(client_max_size=1024 * 1024)
|
|
app.cleanup_ctx.append(runtime_context)
|
|
app.router.add_post("/webhook", webhook)
|
|
app.router.add_get("/token/{nonce}", registration_token)
|
|
app.router.add_get("/healthz", health)
|
|
return app
|
|
|
|
|
|
def main() -> None:
|
|
logging.basicConfig(level=os.environ.get("LOG_LEVEL", "INFO"))
|
|
web.run_app(
|
|
create_app(),
|
|
host=os.environ.get("LISTEN", "0.0.0.0"),
|
|
port=int(os.environ.get("PORT", "8787")),
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|