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()