from __future__ import annotations import argparse import os import time from .contracts import ( ExpansionRequest, FarmPolicy, RegulationIntent, ShardResult, WorkerPulse, utc_now, ) from .evaluator import evaluate_job from .regulation import create_pulse, tick_pulse from .store import FarmStore def maybe_request_expansion( store: FarmStore, job_run_id: str, generation: int, policy: FarmPolicy, reason: str, formation: str, ) -> None: pending = len(store.pending_jobs(formation)) active = store.active_worker_count(policy.pulse_ttl_seconds, formation) allowed = policy.effective_max_workers() if active >= allowed: return if pending <= max(1, active * 2): return requested = min(allowed - active, max(1, pending // 4)) request = ExpansionRequest( run_id=job_run_id, generation=generation, reason=reason, requested_workers=requested, active_workers=active, pending_jobs=pending, policy=policy.to_dict(), formation=formation, ) store.write_expansion_request(request) def primary_formation(policy: FarmPolicy) -> str: formations = tuple(policy.allowed_formations or ("cpu-worker",)) return formations[0] if formations else "cpu-worker" def worker_tick( store: FarmStore, worker_id: str, policy: FarmPolicy, pulse_state: dict[str, float], ) -> bool: formation = primary_formation(policy) claimed = store.claim_next(worker_id, tuple(policy.allowed_formations)) if claimed is None: store.write_worker_pulse( WorkerPulse(worker_id, None, None, "idle", formation=formation, pulse=pulse_state) ) return False job, claimed_path = claimed store.write_worker_pulse( WorkerPulse( worker_id, job.run_id, job.generation, "claimed", job.shard_id, formation=formation, pulse=pulse_state, ) ) children = job.split(policy) if len(children) > 1: store.put_jobs(children) result = ShardResult( run_id=job.run_id, generation=job.generation, shard_id=job.shard_id, worker_id=worker_id, status="refracted", node_results=[], started_at=utc_now(), message=f"split into {len(children)} child shards", ) store.complete_job(job, result, claimed_path) maybe_request_expansion( store, job.run_id, job.generation, policy, "job_refraction", job.formation ) return True result = evaluate_job(job, worker_id) store.complete_job(job, result, claimed_path) maybe_request_expansion(store, job.run_id, job.generation, policy, "pending_backlog", job.formation) return True def _drain_applies_to_worker( store: FarmStore, worker_id: str, control: RegulationIntent, policy: FarmPolicy, worker_formation: str, ) -> bool: if control.target_workers is None: return False if control.target_workers <= 0: return True active = store.active_worker_ids(policy.pulse_ttl_seconds, worker_formation) if len(active) <= control.target_workers: return False keepers = set(active[: control.target_workers]) return worker_id not in keepers def worker_exit_reason(store: FarmStore, worker_id: str, policy: FarmPolicy) -> str | None: control = store.read_regulation() worker_formation = primary_formation(policy) if control.formation is not None and control.formation != worker_formation: return None if control.mode == "shutdown": return f"shutdown:{control.reason}" if control.mode == "drain" and _drain_applies_to_worker( store, worker_id, control, policy, worker_formation ): return f"drain:{control.reason}" if policy.exit_when_pulse_lost and not store.controller_pulse_alive( policy.controller_pulse_ttl_seconds ): return "controller_pulse_lost" return None def worker_loop( store: FarmStore, worker_id: str, policy: FarmPolicy, *, poll_seconds: float = 2.0, once: bool = False, ) -> None: pulse = create_pulse(policy.pulse_rate) formation = primary_formation(policy) while True: pulse_state = tick_pulse(pulse) reason = worker_exit_reason(store, worker_id, policy) if reason is not None: store.write_worker_pulse( WorkerPulse( worker_id, None, None, f"exiting:{reason}", formation=formation, pulse=pulse_state, ) ) store.append_event("farm_worker_exiting", {"worker_id": worker_id, "reason": reason}) return store.requeue_expired_claims(policy) did_work = worker_tick(store, worker_id, policy, pulse_state) if once: return if not did_work: time.sleep(poll_seconds) def main() -> None: parser = argparse.ArgumentParser(description="KEY CPU farm worker") parser.add_argument("--root", default=os.environ.get("KEY_FARM_ROOT", "/data/farm")) parser.add_argument("--worker-id", default=os.environ.get("KEY_FARM_WORKER_ID")) parser.add_argument("--once", action="store_true") parser.add_argument("--poll-seconds", type=float, default=2.0) parser.add_argument("--target-batch-size", type=int, default=None) parser.add_argument("--max-workers", type=int, default=None) parser.add_argument("--budget-hourly-usd", type=float, default=None) parser.add_argument("--formation", default=os.environ.get("KEY_FARM_FORMATION", "cpu-worker")) args = parser.parse_args() worker_id = args.worker_id or f"worker_{os.getpid()}" policy_kwargs = {} if args.target_batch_size is not None: policy_kwargs["target_batch_size"] = args.target_batch_size if args.max_workers is not None: policy_kwargs["max_workers"] = args.max_workers if args.budget_hourly_usd is not None: policy_kwargs["budget_hourly_usd"] = args.budget_hourly_usd policy_kwargs["allowed_formations"] = (args.formation,) worker_loop( FarmStore(args.root), worker_id, FarmPolicy(**policy_kwargs), poll_seconds=args.poll_seconds, once=args.once, ) if __name__ == "__main__": main()