Files
truf-server/app/worker_api.py
T
2026-09-30 20:30:56 +03:00

1138 lines
47 KiB
Python

import argparse
import asyncio
import hashlib
import ipaddress
import json
import logging
import os
import re
import sys
from contextlib import asynccontextmanager
sys.dont_write_bytecode = True
if not sys.dont_write_bytecode:
raise RuntimeError('worker API could not disable bytecode writes')
from starlette.applications import Starlette
from starlette.concurrency import run_in_threadpool
from starlette.requests import Request
from starlette.responses import JSONResponse, Response
from starlette.routing import Route
from admin_api import AdminService, admin_routes
from capacity_model import MAX_RESULT_BUNDLE_BYTES, validate_remote_assignment_capacity
from host_agent_client import HostAgentClient
from host_agent_reconcile import (
fixed_result_directory_is_safe,
reconcile_pending_host_results,
)
from managed_files import ManagedFileTraversal, managed_file_root_registry_from_config
from result_bundle import (
BundleReservation,
ResultBundleError,
ResultBundleReader,
bundle_partial_path,
bundle_ready_path,
ensure_bundle_reservation_paths,
)
from runtime_security import (
durable_publish,
harden_private_file,
private_file_ready,
reject_reparse_components,
require_private_directory,
)
from scan_execution import (
WorkerBuildCompatibility,
validate_protocol1_remote_assignment, validate_protocol2_remote_assignment,
)
from scanner_db import (
PipelineCapacityUnavailable,
ScanEventConflictError,
ScannerDB,
WorkerProgressInactiveError,
utc_now_iso,
)
from worker_contracts import (
AssignmentOutcome,
DIAGNOSTIC_PROJECTION_VERSION,
MAX_DIAGNOSTICS_PER_ASSIGNMENT,
ScanOutcome,
decode_diagnostic_envelope,
encode_diagnostic_envelope,
ordered_diagnostic_uid_set_sha256,
)
REQUEST_ID_RE = re.compile(r'^[a-f0-9]{32,64}$')
DIGEST_RE = re.compile(r'^[a-f0-9]{64}$')
DEFAULT_MAX_JSON_BYTES = 16 * 1024
DEFAULT_MAX_BUNDLE_BYTES = MAX_RESULT_BUNDLE_BYTES
DEFAULT_BODY_IDLE_TIMEOUT_SECONDS = 30
DEFAULT_JSON_BODY_TIMEOUT_SECONDS = 60
DEFAULT_BUNDLE_BODY_TIMEOUT_SECONDS = 30 * 60
DEFAULT_CLAIM_RETRY_AFTER_SECONDS = 5
NO_WORK_REASONS = frozenset((
'empty_queue', 'assignment_cap', 'dispatch_paused', 'capacity',
'compatibility',
))
TERMINAL_FAILURE_CODES = frozenset((
'client_process_failed', 'client_storage_failed', 'client_cancelled',
))
logger = logging.getLogger(__name__)
class WorkerAPIError(RuntimeError):
def __init__(self, status_code, code, message):
super().__init__(message)
self.status_code = int(status_code)
self.code = str(code)
def _error(status_code, code, message):
return JSONResponse(
{'error': {'code': str(code), 'message': str(message)}},
status_code=int(status_code),
headers=(
{'WWW-Authenticate': 'Bearer'} if int(status_code) == 401 else None
),
)
def _db_instance(db_factory, db_url):
db = db_factory(db_url=db_url, initialize=False)
if not db.enabled:
db.close()
raise RuntimeError('worker API PostgreSQL connection is unavailable')
return db
def _hash_file(path):
digest = hashlib.sha256()
with open(path, 'rb', buffering=0) as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b''):
digest.update(chunk)
return digest.hexdigest()
def _log_remote_result_event(event, outcome, identity, reservation_id, transport):
target = str(
transport.get('normalized_target') or transport.get('target') or ''
)
correlation = {
'event': str(event),
'outcome': str(outcome),
'reservation_id': int(reservation_id),
'device_id': int(identity.get('device_id') or 0),
'queue_id': int(transport.get('queue_id') or 0),
'source': str(transport.get('source') or '')[:32],
'bundle_id': str(transport.get('bundle_id') or '')[:64],
'scan_event_id': str(transport.get('scan_event_id') or '')[:64],
'target_sha256': hashlib.sha256(target.encode('utf-8')).hexdigest(),
}
logger.info(
'remote result event %s',
json.dumps(correlation, ensure_ascii=True, sort_keys=True, separators=(',', ':')),
)
class WorkerService:
def __init__(
self, db_url, bundle_root, assignment_builder, *, db_factory=ScannerDB,
max_bundle_bytes=DEFAULT_MAX_BUNDLE_BYTES, reaper_batch_size=1000,
bundle_capacity_bytes=3 * 1024 * 1024 * 1024,
body_idle_timeout_seconds=DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
json_body_timeout_seconds=DEFAULT_JSON_BODY_TIMEOUT_SECONDS,
bundle_body_timeout_seconds=DEFAULT_BUNDLE_BODY_TIMEOUT_SECONDS,
claim_retry_after_seconds=DEFAULT_CLAIM_RETRY_AFTER_SECONDS,
):
self.db_url = str(db_url or '')
self.bundle_root = require_private_directory(
os.path.abspath(bundle_root), create=False,
)
if not callable(assignment_builder):
raise TypeError('worker assignment builder must be callable')
self.assignment_builder = assignment_builder
self.db_factory = db_factory
self.max_bundle_bytes = max(1, min(
DEFAULT_MAX_BUNDLE_BYTES, int(max_bundle_bytes),
))
self.bundle_capacity_bytes = max(1, int(bundle_capacity_bytes))
self.reaper_batch_size = max(1, min(1000, int(reaper_batch_size)))
self.body_idle_timeout_seconds = float(body_idle_timeout_seconds)
self.json_body_timeout_seconds = float(json_body_timeout_seconds)
self.bundle_body_timeout_seconds = float(bundle_body_timeout_seconds)
self.claim_retry_after_seconds = int(claim_retry_after_seconds)
if (
not 1 <= self.body_idle_timeout_seconds <= 120
or not 1 <= self.json_body_timeout_seconds <= 300
or not 30 <= self.bundle_body_timeout_seconds <= 24 * 60 * 60
or not 1 <= self.claim_retry_after_seconds <= 300
):
raise ValueError('worker API request body time bounds are invalid')
self._upload_locks = {}
self._upload_locks_guard = asyncio.Lock()
def authenticate(self, authorization):
prefix, separator, token = str(authorization or '').partition(' ')
if prefix.lower() != 'bearer' or not separator or not 16 <= len(token) <= 512:
raise WorkerAPIError(401, 'unauthorized', 'worker credentials are invalid')
token_sha256 = hashlib.sha256(token.encode('utf-8')).hexdigest()
db = _db_instance(self.db_factory, self.db_url)
try:
identity = db.authenticate_remote_worker(token_sha256)
finally:
db.close()
if not identity:
raise WorkerAPIError(401, 'unauthorized', 'worker credentials are invalid')
identity = dict(identity)
identity['token_sha256'] = token_sha256
return identity
def claim(self, identity, payload):
payload = dict(payload or {})
if set(payload) != {'request_id', 'build'}:
raise WorkerAPIError(400, 'invalid_request', 'claim request shape is invalid')
request_id = str(payload.get('request_id') or '').lower()
if not REQUEST_ID_RE.fullmatch(request_id):
raise WorkerAPIError(400, 'invalid_request', 'request_id must be a 128-bit or stronger lowercase hex value')
try:
compatibility = WorkerBuildCompatibility.from_mapping(payload.get('build'))
except (TypeError, ValueError) as exc:
raise WorkerAPIError(400, 'invalid_compatibility', str(exc)) from exc
assignment = self.assignment_builder(
dict(identity), request_id, compatibility.as_dict(),
)
if assignment is None:
if compatibility.protocol_version == 1:
raise WorkerAPIError(
409, 'incompatible_protocol',
'protocol-1 packages cannot receive new assignments',
)
return None
assignment = dict(assignment)
if set(assignment) == {'no_assignment'}:
if compatibility.protocol_version == 1:
raise WorkerAPIError(
409, 'incompatible_protocol',
'protocol-1 packages cannot receive new assignments',
)
no_assignment = assignment['no_assignment']
if (
not isinstance(no_assignment, dict)
or set(no_assignment) != {'reason'}
or no_assignment.get('reason') not in NO_WORK_REASONS
):
raise RuntimeError('assignment builder returned an invalid no-work reason')
return {'no_assignment': {'reason': no_assignment['reason']}}
if set(assignment) == {'resolution'}:
resolution = dict(assignment['resolution'] or {})
if (
int(resolution.get('reservation_id') or 0) <= 0
or resolution.get('resolution') not in {
'bundle_accepted', 'prebundle_report', 'expired',
}
or not DIGEST_RE.fullmatch(str(resolution.get('receipt_id') or ''))
or not REQUEST_ID_RE.fullmatch(str(resolution.get('bundle_id') or ''))
or not REQUEST_ID_RE.fullmatch(str(resolution.get('scan_event_id') or ''))
):
raise RuntimeError('assignment builder returned an invalid resolution receipt')
return {'resolution': resolution}
allowed = {
'reservation', 'deadlines', 'compatibility', 'scan_kwargs', 'event_scan_options',
'queue_policy', 'limits', 'scan_policy', 'execution_snapshot',
'execution_snapshot_sha256', 'execution_plan',
}
if set(assignment) != allowed:
raise RuntimeError('assignment builder returned an invalid payload shape')
reservation = dict(assignment['reservation'])
validator = (
validate_protocol1_remote_assignment
if compatibility.protocol_version == 1
else validate_protocol2_remote_assignment
)
validated = validator(assignment, compatibility)
required = validated['compatibility']
if (
int(reservation.get('remote_device_id') or 0) != int(identity['device_id'])
or str(reservation.get('assignment_kind') or '') != 'remote'
or not REQUEST_ID_RE.fullmatch(
str(reservation.get('reservation_token') or ''),
)
or str(reservation.get('remote_effective_config_sha256') or '')
!= required.effective_config_sha256
):
raise RuntimeError('assignment builder returned a conflicting remote identity')
return assignment
def status(self, identity, reservation_id):
db = _db_instance(self.db_factory, self.db_url)
try:
result = db.remote_assignment_status(
int(reservation_id), int(identity['device_id']),
str(identity['token_sha256']),
)
return result
finally:
db.close()
def progress(self, identity, reservation_id, payload):
db = _db_instance(self.db_factory, self.db_url)
try:
try:
stored = db.record_remote_worker_progress_event(
int(reservation_id), int(identity['device_id']),
str(identity['token_sha256']), payload,
)
except ValueError as exc:
raise WorkerAPIError(
400, 'invalid_progress', 'worker progress event is invalid'
) from exc
except WorkerProgressInactiveError as exc:
raise WorkerAPIError(
410, 'progress_stale',
'owned assignment is inactive for new progress',
) from exc
return {
'accepted': True,
'reservation_id': int(reservation_id),
'sequence': int(stored['sequence']),
'received_at': str(stored['received_at']),
'replayed': bool(stored.get('replayed')),
}
finally:
db.close()
def transport(self, identity, reservation_id):
db = _db_instance(self.db_factory, self.db_url)
try:
return db.remote_assignment_transport(
int(reservation_id), int(identity['device_id']),
str(identity['token_sha256']),
)
finally:
db.close()
def accept_ready(
self, identity, transport, effective_diagnostics, metadata, payload_sha256,
):
values = metadata.as_dict()
values['effective_diagnostic_count'] = len(effective_diagnostics)
values['effective_diagnostic_projection_version'] = (
DIAGNOSTIC_PROJECTION_VERSION
)
values['effective_diagnostic_uids_sha256'] = (
ordered_diagnostic_uid_set_sha256(effective_diagnostics)
)
values['relative_path'] = str(transport['ready_relative_path']).replace('\\', '/')
db = _db_instance(self.db_factory, self.db_url)
try:
try:
return db.mark_result_bundle_ready(
int(transport['reservation_id']), values,
remote_acceptance={
'device_id': int(identity['device_id']),
'payload_sha256': payload_sha256,
'token_sha256': str(identity['token_sha256']),
},
bundle_capacity_bytes=self.bundle_capacity_bytes,
)
except PipelineCapacityUnavailable as exc:
raise WorkerAPIError(
503, 'capacity_backpressure',
'bundle capacity is temporarily unavailable; retry the identical upload',
) from exc
finally:
db.close()
def report_terminal(self, identity, reservation_id, payload):
payload = dict(payload or {})
failure_code = str(payload.get('failure_code') or '')
detail = str(payload.get('detail') or '')
if (
set(payload) not in (
{'failure_code', 'detail'},
{'failure_code', 'detail', 'diagnostics'},
)
or failure_code not in TERMINAL_FAILURE_CODES or len(detail) > 1000
or not isinstance(payload.get('detail'), str)
):
raise WorkerAPIError(400, 'invalid_report', 'terminal report fields are invalid')
normalized = {'failure_code': failure_code, 'detail': detail}
if 'diagnostics' in payload:
diagnostics = payload['diagnostics']
if (
not isinstance(diagnostics, list)
or len(diagnostics) > MAX_DIAGNOSTICS_PER_ASSIGNMENT
):
raise WorkerAPIError(400, 'invalid_report', 'terminal diagnostics are invalid')
try:
normalized_diagnostics = []
diagnostic_uids = set()
for diagnostic in diagnostics:
envelope = decode_diagnostic_envelope(json.dumps(
diagnostic, ensure_ascii=True, sort_keys=True,
separators=(',', ':'),
).encode('ascii'))
if (
envelope.assignment_outcome
is not AssignmentOutcome.PREBUNDLE_FAILED
or envelope.scan_outcome is not ScanOutcome.UNAVAILABLE
):
raise ValueError('terminal diagnostic outcome is invalid')
if envelope.diagnostic_uid in diagnostic_uids:
raise ValueError('terminal diagnostic UID is duplicated')
diagnostic_uids.add(envelope.diagnostic_uid)
normalized_diagnostics.append(json.loads(
encode_diagnostic_envelope(envelope).decode('ascii')
))
normalized['diagnostics'] = normalized_diagnostics
except (TypeError, ValueError, UnicodeError) as exc:
raise WorkerAPIError(
400, 'invalid_report', 'terminal diagnostics are invalid'
) from exc
db = _db_instance(self.db_factory, self.db_url)
try:
return db.report_remote_prebundle_failure(
int(reservation_id), int(identity['device_id']),
str(identity['token_sha256']), normalized,
)
finally:
db.close()
def reap(self):
db = _db_instance(self.db_factory, self.db_url)
try:
receipts = db.reap_expired_remote_assignments(self.reaper_batch_size)
db.reconcile_runtime_drain()
try:
reconcile_pending_host_results(db)
except Exception as exc:
logger.warning(
'host operation result reconciliation unavailable (%s)',
type(exc).__name__,
)
return receipts
finally:
db.close()
@asynccontextmanager
async def upload_scope(self, reservation_id):
reservation_id = int(reservation_id)
async with self._upload_locks_guard:
entry = self._upload_locks.get(reservation_id)
if entry is None:
entry = [asyncio.Lock(), 0]
self._upload_locks[reservation_id] = entry
entry[1] += 1
await entry[0].acquire()
try:
yield
finally:
entry[0].release()
async with self._upload_locks_guard:
entry[1] -= 1
if entry[1] == 0:
self._upload_locks.pop(reservation_id, None)
async def _bounded_body_chunks(request, *, absolute_timeout, idle_timeout):
iterator = request.stream().__aiter__()
loop = asyncio.get_running_loop()
deadline = loop.time() + float(absolute_timeout)
while True:
remaining = deadline - loop.time()
if remaining <= 0:
raise WorkerAPIError(408, 'request_timeout', 'request body deadline elapsed')
try:
chunk = await asyncio.wait_for(
anext(iterator), timeout=min(float(idle_timeout), remaining),
)
except StopAsyncIteration:
return
except asyncio.TimeoutError as exc:
raise WorkerAPIError(408, 'request_timeout', 'request body deadline elapsed') from exc
yield chunk
async def _bounded_json(
request, max_bytes=DEFAULT_MAX_JSON_BYTES, *,
absolute_timeout=DEFAULT_JSON_BODY_TIMEOUT_SECONDS,
idle_timeout=DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
):
content_type = request.headers.get('content-type', '').split(';', 1)[0].strip().lower()
if content_type != 'application/json':
raise WorkerAPIError(415, 'unsupported_media_type', 'application/json is required')
body = bytearray()
async for chunk in _bounded_body_chunks(
request, absolute_timeout=absolute_timeout, idle_timeout=idle_timeout,
):
if len(body) + len(chunk) > max_bytes:
raise WorkerAPIError(413, 'request_too_large', 'JSON request exceeds its byte bound')
body.extend(chunk)
try:
def reject_duplicate_fields(pairs):
value = {}
for key, item in pairs:
if key in value:
raise ValueError('duplicate JSON field')
value[key] = item
return value
value = json.loads(
bytes(body).decode('utf-8', errors='strict'),
object_pairs_hook=reject_duplicate_fields,
)
except (UnicodeDecodeError, json.JSONDecodeError, ValueError) as exc:
raise WorkerAPIError(400, 'invalid_json', 'request body is not valid UTF-8 JSON') from exc
if not isinstance(value, dict):
raise WorkerAPIError(400, 'invalid_json', 'request body must be a JSON object')
return value
def _reservation_id(request):
try:
value = int(request.path_params['reservation_id'])
except (TypeError, ValueError, OverflowError):
raise WorkerAPIError(404, 'not_found', 'assignment was not found') from None
if value <= 0:
raise WorkerAPIError(404, 'not_found', 'assignment was not found')
return value
async def _identity(request):
return await run_in_threadpool(
request.app.state.worker_service.authenticate,
request.headers.get('authorization'),
)
async def claim(request):
identity = await _identity(request)
service = request.app.state.worker_service
payload = await _bounded_json(
request, absolute_timeout=getattr(
service, 'json_body_timeout_seconds', DEFAULT_JSON_BODY_TIMEOUT_SECONDS,
),
idle_timeout=getattr(
service, 'body_idle_timeout_seconds', DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
),
)
assignment = await run_in_threadpool(
request.app.state.worker_service.claim, identity, payload,
)
if assignment is None:
return Response(
status_code=204,
headers={'Retry-After': str(service.claim_retry_after_seconds)},
)
if set(assignment) == {'no_assignment'}:
return Response(
status_code=204,
headers={
'Retry-After': str(service.claim_retry_after_seconds),
'X-Truf-No-Work-Reason': assignment['no_assignment']['reason'],
},
)
if set(assignment) == {'resolution'}:
return JSONResponse(assignment, status_code=200)
return JSONResponse({'assignment': assignment}, status_code=201)
async def assignment_status(request):
identity = await _identity(request)
result = await run_in_threadpool(
request.app.state.worker_service.status, identity, _reservation_id(request),
)
if result is None:
raise WorkerAPIError(404, 'not_found', 'assignment was not found')
return JSONResponse(result)
async def assignment_progress(request):
identity = await _identity(request)
service = request.app.state.worker_service
payload = await _bounded_json(
request, absolute_timeout=getattr(
service, 'json_body_timeout_seconds', DEFAULT_JSON_BODY_TIMEOUT_SECONDS,
),
idle_timeout=getattr(
service, 'body_idle_timeout_seconds', DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
),
)
result = await run_in_threadpool(
service.progress, identity, _reservation_id(request), payload,
)
return JSONResponse(result)
async def terminal_report(request):
identity = await _identity(request)
service = request.app.state.worker_service
payload = await _bounded_json(
request, absolute_timeout=getattr(
service, 'json_body_timeout_seconds', DEFAULT_JSON_BODY_TIMEOUT_SECONDS,
),
idle_timeout=getattr(
service, 'body_idle_timeout_seconds', DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
),
)
result = await run_in_threadpool(
request.app.state.worker_service.report_terminal,
identity, _reservation_id(request), payload,
)
if result is None:
raise WorkerAPIError(409, 'assignment_not_active', 'assignment is no longer active')
return JSONResponse(result)
def _remove_private_regular(path):
if not os.path.lexists(path):
return
reject_reparse_components(path)
if not os.path.isfile(path) or os.path.islink(path):
raise ResultBundleError('bundle transport path is not a regular file')
os.remove(path)
async def upload_bundle(request):
identity = await _identity(request)
service = request.app.state.worker_service
reservation_id = _reservation_id(request)
supplied_digest = str(request.headers.get('x-truf-payload-sha256') or '').lower()
if not DIGEST_RE.fullmatch(supplied_digest):
raise WorkerAPIError(400, 'invalid_digest', 'X-Truf-Payload-SHA256 is required')
if request.headers.get('content-type', '').split(';', 1)[0].strip().lower() != 'application/octet-stream':
raise WorkerAPIError(415, 'unsupported_media_type', 'application/octet-stream is required')
try:
content_length = int(request.headers.get('content-length') or '')
except (TypeError, ValueError, OverflowError):
raise WorkerAPIError(411, 'length_required', 'a valid Content-Length is required') from None
async with service.upload_scope(reservation_id):
transport = await run_in_threadpool(service.transport, identity, reservation_id)
if transport is None:
raise WorkerAPIError(404, 'not_found', 'assignment was not found')
byte_bound = min(
int(transport['declared_bundle_bytes']), service.max_bundle_bytes,
)
if content_length <= 0 or content_length > byte_bound:
raise WorkerAPIError(413, 'bundle_too_large', 'bundle exceeds its assigned byte bound')
async def receive(handle=None):
digest = hashlib.sha256()
received = 0
async for chunk in _bounded_body_chunks(
request,
absolute_timeout=getattr(
service, 'bundle_body_timeout_seconds',
DEFAULT_BUNDLE_BODY_TIMEOUT_SECONDS,
),
idle_timeout=getattr(
service, 'body_idle_timeout_seconds',
DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
),
):
received += len(chunk)
if received > content_length or received > byte_bound:
raise WorkerAPIError(413, 'bundle_too_large', 'bundle exceeds its assigned byte bound')
if handle is not None:
handle.write(chunk)
digest.update(chunk)
if received != content_length:
raise WorkerAPIError(400, 'length_mismatch', 'bundle length does not match Content-Length')
if digest.hexdigest() != supplied_digest:
raise WorkerAPIError(400, 'digest_mismatch', 'bundle digest does not match its declaration')
receipt = transport.get('receipt')
if receipt is not None:
if (
str(transport.get('remote_resolution_kind')) != 'bundle_accepted'
or str(transport.get('remote_payload_sha256') or '') != supplied_digest
):
_log_remote_result_event(
'remote_result_rejected', 'resolution_conflict',
identity, reservation_id, transport,
)
raise WorkerAPIError(409, 'resolution_conflict', 'assignment already has a conflicting resolution')
await receive()
_log_remote_result_event(
'remote_result_replayed', 'original_acceptance_returned',
identity, reservation_id, transport,
)
return JSONResponse(receipt)
if str(transport.get('state')) != 'scanning' or str(
transport.get('remote_expires_at') or ''
) <= utc_now_iso():
_log_remote_result_event(
'remote_result_rejected', 'stale_assignment',
identity, reservation_id, transport,
)
raise WorkerAPIError(410, 'assignment_expired', 'assignment deadline has passed')
reservation = BundleReservation.from_mapping(transport)
await run_in_threadpool(ensure_bundle_reservation_paths, service.bundle_root, reservation)
partial_path = bundle_partial_path(
service.bundle_root, reservation.bundle_id, reservation.reservation_token,
)
ready_path = bundle_ready_path(service.bundle_root, reservation.bundle_id)
if os.path.lexists(ready_path):
await receive()
reader = ResultBundleReader(ready_path, max_event_bytes=byte_bound)
metadata = await run_in_threadpool(reader.validate)
effective_diagnostics = await run_in_threadpool(
reader.effective_diagnostics
)
existing_digest = await run_in_threadpool(_hash_file, ready_path)
if existing_digest != supplied_digest:
_log_remote_result_event(
'remote_result_rejected', 'bundle_conflict',
identity, reservation_id, transport,
)
raise WorkerAPIError(409, 'bundle_conflict', 'a different bundle already occupies the assigned path')
result = await run_in_threadpool(
service.accept_ready, identity, transport, effective_diagnostics, metadata,
supplied_digest,
)
if not result:
await run_in_threadpool(_remove_private_regular, ready_path)
_log_remote_result_event(
'remote_result_rejected', 'stale_assignment',
identity, reservation_id, transport,
)
raise WorkerAPIError(410, 'assignment_expired', 'assignment deadline has passed')
return JSONResponse(result)
await run_in_threadpool(_remove_private_regular, partial_path)
descriptor = os.open(
partial_path,
os.O_WRONLY | os.O_CREAT | os.O_EXCL | getattr(os, 'O_BINARY', 0),
0o600,
)
published = False
try:
os.close(descriptor)
descriptor = None
harden_private_file(partial_path)
with open(partial_path, 'wb', buffering=0) as handle:
await receive(handle)
handle.flush()
os.fsync(handle.fileno())
if not private_file_ready(partial_path):
raise ResultBundleError('bundle transport file is not private')
reader = ResultBundleReader(partial_path, max_event_bytes=byte_bound)
metadata = await run_in_threadpool(reader.validate)
effective_diagnostics = await run_in_threadpool(
reader.effective_diagnostics
)
await run_in_threadpool(durable_publish, partial_path, ready_path)
published = True
result = await run_in_threadpool(
service.accept_ready, identity, transport, effective_diagnostics, metadata,
supplied_digest,
)
if not result:
await run_in_threadpool(_remove_private_regular, ready_path)
published = False
_log_remote_result_event(
'remote_result_rejected', 'stale_assignment',
identity, reservation_id, transport,
)
raise WorkerAPIError(410, 'assignment_expired', 'assignment deadline has passed')
return JSONResponse(result, status_code=201)
except (WorkerAPIError, ResultBundleError):
if not published:
await run_in_threadpool(_remove_private_regular, partial_path)
raise
except ScanEventConflictError:
_log_remote_result_event(
'remote_result_rejected', 'ownership_conflict',
identity, reservation_id, transport,
)
await run_in_threadpool(
_remove_private_regular, ready_path if published else partial_path,
)
raise
except Exception:
if not published:
await run_in_threadpool(_remove_private_regular, partial_path)
raise
finally:
if descriptor is not None:
os.close(descriptor)
async def worker_api_error_handler(request, exc):
return _error(exc.status_code, exc.code, str(exc))
async def conflict_error_handler(request, exc):
return _error(409, 'reservation_conflict', str(exc))
async def bundle_error_handler(request, exc):
return _error(400, 'invalid_bundle', str(exc))
def create_worker_app(service, *, reaper_interval_seconds=60, admin_service=None):
interval = max(1, int(reaper_interval_seconds))
@asynccontextmanager
async def lifespan(app):
stopping = asyncio.Event()
traversal = None
app.state.managed_file_traversal = None
if admin_service is not None:
try:
traversal = await run_in_threadpool(
ManagedFileTraversal, admin_service.managed_file_roots,
)
app.state.managed_file_traversal = traversal
except Exception as exc:
logger.error(
'Managed files are unavailable: %s',
getattr(exc, 'category', type(exc).__name__),
)
async def reap_loop():
while not stopping.is_set():
try:
await run_in_threadpool(service.reap)
except Exception as exc:
logger.error('Remote assignment reaper failed: %s', type(exc).__name__)
try:
await asyncio.wait_for(stopping.wait(), timeout=interval)
except asyncio.TimeoutError:
continue
task = asyncio.create_task(reap_loop())
try:
yield
finally:
stopping.set()
try:
await task
finally:
app.state.managed_file_traversal = None
if traversal is not None:
try:
await run_in_threadpool(traversal.close)
except Exception as exc:
logger.error(
'Managed file traversal shutdown failed: %s',
type(exc).__name__,
)
routes = [
Route('/api/v1/worker/claim', claim, methods=['POST']),
Route('/api/v1/worker/assignments/{reservation_id:int}', assignment_status, methods=['GET']),
Route('/api/v1/worker/assignments/{reservation_id:int}/progress', assignment_progress, methods=['POST']),
Route('/api/v1/worker/assignments/{reservation_id:int}/bundle', upload_bundle, methods=['PUT']),
Route('/api/v1/worker/assignments/{reservation_id:int}/terminal', terminal_report, methods=['POST']),
]
if admin_service is not None:
routes.extend(admin_routes())
app = Starlette(
routes=routes,
exception_handlers={
WorkerAPIError: worker_api_error_handler,
ScanEventConflictError: conflict_error_handler,
ResultBundleError: bundle_error_handler,
},
lifespan=lifespan,
)
app.state.worker_service = service
if admin_service is not None:
app.state.admin_service = admin_service
return app
def _configured_auth_entry(source, source_config, secrets, configured_names):
from console_runner import auth_pool_entries
entries, _ = auth_pool_entries(source_config, secrets)
selected_name = str(configured_names.get(source) or '')
if selected_name:
matches = [entry for entry in entries if str(entry.get('name') or '') == selected_name]
if len(matches) != 1:
raise ValueError(f'worker API auth entry for {source} is unavailable or ambiguous')
return matches[0]
if len(entries) > 1:
raise ValueError(f'worker API requires an explicit auth entry for {source}')
return entries[0] if entries else None
def build_configured_worker_service(config_path, metadata, *, db_factory=ScannerDB):
from console_runner import (
apply_global_config,
build_args_from_source_config,
load_config,
load_secrets,
)
from db_backend import database_url_from_env
from worker_assignment import (
CORE_ASSIGNMENT_SOURCE_ADAPTERS,
PROTOCOL2_NEW_CLAIM_SOURCES,
RemoteAssignmentBuilder,
assignment_source_adapter,
)
config = load_config(config_path, managed_postgres=True, final_cutover=True)
global_config = config.get('global') or {}
supervisor_config = config.get('supervisor') or {}
worker_config = supervisor_config.get('worker_api') or {}
allowed_keys = {
'enabled', 'address', 'port', 'sources', 'auth_entries',
'compatibility_profiles', 'assignment_ttl_seconds',
'assignment_ttl_seconds_by_source',
'max_bundle_bytes', 'reaper_interval_seconds', 'reaper_batch_size',
'limit_concurrency', 'body_idle_timeout_seconds',
'json_body_timeout_seconds', 'bundle_body_timeout_seconds', 'admin',
}
unknown = sorted(set(worker_config) - allowed_keys)
if unknown:
raise ValueError('worker API configuration has unsupported keys: ' + ', '.join(unknown))
if worker_config.get('enabled') is not True:
raise ValueError('worker API runtime is not explicitly enabled')
admin_config = worker_config.get('admin', {})
if admin_config is None:
admin_config = {}
if not isinstance(admin_config, dict):
raise ValueError('worker API admin configuration must be a mapping')
admin_allowed_keys = {
'enabled', 'origin', 'edge_marker', 'max_body_bytes',
'snapshot_limit', 'requeue_limit', 'managed_file_roots',
}
admin_unknown = sorted(set(admin_config) - admin_allowed_keys)
if admin_unknown:
raise ValueError(
'worker API admin configuration has unsupported keys: '
+ ', '.join(admin_unknown)
)
if not isinstance(admin_config.get('enabled', False), bool):
raise ValueError('worker API admin enabled flag must be boolean')
managed_file_roots = managed_file_root_registry_from_config(config)
raw_sources = worker_config.get('sources') or []
if not isinstance(raw_sources, list):
raise ValueError('worker API sources must be a list')
sources = tuple(dict.fromkeys(
str(value or '').strip().lower() for value in raw_sources
)) if raw_sources else tuple(CORE_ASSIGNMENT_SOURCE_ADAPTERS)
if not sources or not set(sources) <= PROTOCOL2_NEW_CLAIM_SOURCES:
raise ValueError('worker API sources must use exact protocol-2 core sources')
assignment_ttl = worker_config.get('assignment_ttl_seconds', 86400)
if (
type(assignment_ttl) is not int
or not 60 <= assignment_ttl <= 7 * 24 * 60 * 60
):
raise ValueError('worker assignment lifetime must be between one minute and seven days')
assignment_ttl_by_source = worker_config.get(
'assignment_ttl_seconds_by_source', {},
)
if (
type(assignment_ttl_by_source) is not dict
or set(assignment_ttl_by_source) - set(PROTOCOL2_NEW_CLAIM_SOURCES)
):
raise ValueError('worker assignment lifetime overrides contain unsupported sources')
assignment_ttl_by_source = dict(assignment_ttl_by_source)
if any(
type(value) is not int or not 60 <= value <= 7 * 24 * 60 * 60
for value in assignment_ttl_by_source.values()
):
raise ValueError('worker assignment lifetime override is outside its bounds')
configured_names = worker_config.get('auth_entries') or {}
if (
not isinstance(configured_names, dict)
or set(configured_names) - (set(sources) | {'github'})
):
raise ValueError(
'worker API auth_entries must map only configured sources or legacy GitHub'
)
if set(configured_names) - {'github', 'gitlab'}:
raise ValueError(
'worker API auth_entries may select only GitHub or GitLab credentials'
)
profiles = worker_config.get('compatibility_profiles') or {}
if not isinstance(profiles, dict) or not profiles or len(profiles) > 16:
raise ValueError('worker API requires between one and sixteen compatibility profiles')
profiles = {name: dict(value or {}) for name, value in profiles.items()}
config_dir = os.path.dirname(os.path.abspath(config_path))
for profile in profiles.values():
package_manifest = profile.get('package_manifest')
if isinstance(package_manifest, str) and not os.path.isabs(package_manifest):
profile['package_manifest'] = os.path.join(config_dir, package_manifest)
address_text = str(worker_config.get('address') or '127.0.0.1').strip()
try:
address = ipaddress.ip_address(address_text)
except ValueError as exc:
raise ValueError('worker API address must be an IP literal') from exc
if (
address.is_unspecified or address.is_multicast
or not (address.is_loopback or address.is_private)
):
raise ValueError('worker API address must be loopback or private')
port = int(worker_config.get('port', 8766))
if not 1024 <= port <= 65535:
raise ValueError('worker API port must be between 1024 and 65535')
db_url = database_url_from_env()
if not db_url or str(global_config.get('database_url') or '') != db_url:
raise ValueError('worker API requires the canonical managed PostgreSQL DSN')
bundle_root = str(global_config.get('result_bundle_dir') or '')
if not bundle_root:
raise ValueError('worker API requires the canonical result bundle root')
instance_id = str((metadata or {}).get('instance_id') or '')
if not instance_id:
raise ValueError('worker API requires the authenticated supervisor identity')
apply_global_config(global_config)
secrets = load_secrets(config, config_path)
configured_sources = config.get('sources') or {}
source_args = {}
credential_refs = {}
runtime_sources = sources + (
('github',) if 'github' in configured_names and 'github' not in sources else ()
)
for source in runtime_sources:
adapter = assignment_source_adapter(source)
source_config = configured_sources.get(source)
if not isinstance(source_config, dict) or source_config.get('enabled') is not True:
raise ValueError(f'worker API source {source} is not explicitly enabled')
auth_entry = None
if adapter.planning_kind == 'exact_git_v1':
auth_entry = _configured_auth_entry(
source, source_config, secrets, configured_names,
)
args = build_args_from_source_config(
source, source_config, global_config, '', auth_entry=auth_entry,
)
if adapter.planning_kind != 'exact_git_v1':
args.token = ''
args.docker_username = ''
args.docker_token = ''
args.auth_name = None
adapter.validate_source_args(args)
source_args[source] = args
credential_refs[source] = str((auth_entry or {}).get('name') or '')
global_bundle_limit = int(global_config.get(
'result_bundle_max_event_bytes', DEFAULT_MAX_BUNDLE_BYTES,
))
validate_remote_assignment_capacity(global_config)
max_bundle_bytes = int(worker_config.get('max_bundle_bytes', global_bundle_limit))
if not 1024 * 1024 <= max_bundle_bytes <= min(DEFAULT_MAX_BUNDLE_BYTES, global_bundle_limit):
raise ValueError('worker API bundle limit exceeds the canonical event bound')
reaper_interval = int(worker_config.get('reaper_interval_seconds', 60))
if not 5 <= reaper_interval <= 3600:
raise ValueError('worker API reaper interval must be between 5 and 3600 seconds')
reaper_batch_size = int(worker_config.get('reaper_batch_size', 1000))
if not 1 <= reaper_batch_size <= 1000:
raise ValueError('worker API reaper batch size must be between 1 and 1000')
limit_concurrency = int(worker_config.get('limit_concurrency', 64))
if not 1 <= limit_concurrency <= 1024:
raise ValueError('worker API concurrency limit must be between 1 and 1024')
body_idle_timeout = int(worker_config.get(
'body_idle_timeout_seconds', DEFAULT_BODY_IDLE_TIMEOUT_SECONDS,
))
json_body_timeout = int(worker_config.get(
'json_body_timeout_seconds', DEFAULT_JSON_BODY_TIMEOUT_SECONDS,
))
bundle_body_timeout = worker_config.get(
'bundle_body_timeout_seconds', DEFAULT_BUNDLE_BODY_TIMEOUT_SECONDS,
)
if type(bundle_body_timeout) is not int or not 30 <= bundle_body_timeout <= 86400:
raise ValueError('worker result upload body timeout is outside its bounds')
for source in sources:
scan_timeout = getattr(source_args[source], 'timeout', 0)
if isinstance(scan_timeout, bool):
raise ValueError('worker source scan timeout is invalid')
try:
scan_timeout = int(scan_timeout)
except (TypeError, ValueError, OverflowError) as exc:
raise ValueError('worker source scan timeout is invalid') from exc
effective_ttl = assignment_ttl_by_source.get(source, assignment_ttl)
if effective_ttl < scan_timeout + bundle_body_timeout + 60:
raise ValueError(
f'worker assignment lifetime for {source} cannot cover scan and upload deadlines'
)
assignment_builder = RemoteAssignmentBuilder(
db_url, bundle_root, source_args, profiles, instance_id,
assignment_ttl_seconds=assignment_ttl,
assignment_ttl_seconds_by_source=assignment_ttl_by_source,
result_upload_body_timeout_seconds=bundle_body_timeout,
credential_refs=credential_refs,
db_factory=db_factory,
)
service = WorkerService(
db_url, bundle_root, assignment_builder, db_factory=db_factory,
max_bundle_bytes=max_bundle_bytes, reaper_batch_size=reaper_batch_size,
bundle_capacity_bytes=int(global_config['result_bundle_max_total_bytes']),
body_idle_timeout_seconds=body_idle_timeout,
json_body_timeout_seconds=json_body_timeout,
bundle_body_timeout_seconds=bundle_body_timeout,
)
service.admin_service = None
if admin_config.get('enabled') is True:
runtime_apply_provider = None
if HostAgentClient.is_available() and fixed_result_directory_is_safe():
runtime_apply_provider = HostAgentClient().dispatch
service.admin_service = AdminService(
db_url, admin_config.get('origin'), admin_config.get('edge_marker'),
db_factory=db_factory,
max_body_bytes=admin_config.get('max_body_bytes', 8 * 1024),
snapshot_limit=admin_config.get('snapshot_limit', 200),
requeue_limit=admin_config.get('requeue_limit', 100),
supervisor_metadata=metadata,
package_compatibility_provider=assignment_builder.compatibility_snapshot,
runtime_config_path=config_path,
managed_file_roots=managed_file_roots,
runtime_apply_provider=runtime_apply_provider,
)
return service, {
'address': str(address), 'port': port,
'reaper_interval_seconds': reaper_interval,
'limit_concurrency': limit_concurrency,
}
def parse_args():
parser = argparse.ArgumentParser(description='Authenticated remote scan Worker API')
parser.add_argument('--config', required=True)
return parser.parse_args()
def main():
from lifecycle_authority import require_active_supervisor_child
args = parse_args()
metadata = require_active_supervisor_child(
args.config, child_kind='worker-api', require_dsn=True,
)
service, runtime = build_configured_worker_service(args.config, metadata)
app = create_worker_app(
service, reaper_interval_seconds=runtime['reaper_interval_seconds'],
admin_service=getattr(service, 'admin_service', None),
)
import uvicorn
uvicorn.run(
app,
host=runtime['address'],
port=runtime['port'],
access_log=False,
proxy_headers=False,
server_header=False,
limit_concurrency=runtime['limit_concurrency'],
timeout_keep_alive=5,
workers=1,
)
if __name__ == '__main__':
main()