from datetime import datetime, timedelta from typing import Dict, List, Optional from app.models import InferenceJob, JobResult, JobStatus, JobType, WorkerState, PrivacyMode from app.security import generate_job_id from app.storage import save_job_event from app.session_store import list_workers, get_session _jobs_by_id: Dict[str, InferenceJob] = {} def create_job(session_id: str, job_type: JobType, payload: dict, privacy_mode: PrivacyMode, constraints: dict) -> InferenceJob: job = InferenceJob( job_id=generate_job_id(), session_id=session_id, job_type=job_type, privacy_mode=privacy_mode, payload=payload, constraints=constraints, ) _jobs_by_id[job.job_id] = job save_job_event({ "event": "created", "job_id": job.job_id, "session_id": session_id, "job_type": job_type.value, "timestamp": datetime.utcnow().isoformat(), }) return job def enqueue_job(job: InferenceJob) -> None: # v1: immediate assignment attempt; v2 can use real queue pass def select_worker_for_job(session_id: str, job: InferenceJob) -> Optional[str]: workers = list_workers(session_id) if not workers: return None best = None best_score = -1.0 for w in workers: score = _score_worker_for_job(w, job) if score > best_score: best_score = score best = w.worker_id return best def _score_worker_for_job(worker: WorkerState, job: InferenceJob) -> float: score = 0.0 # Capability match for cap in worker.capabilities: if cap.capability_name == job.job_type.value: score += 100.0 break # Battery if worker.battery_level is not None and worker.battery_level > 0.2: score += 20.0 # Thermal if worker.thermal_state in (None, "nominal", "fair"): score += 20.0 # Heartbeat freshness age = (datetime.utcnow() - worker.last_heartbeat).total_seconds() if age < 10: score += 30.0 elif age < 30: score += 10.0 return score def assign_job_to_worker(job_id: str, worker_id: str) -> bool: job = _jobs_by_id.get(job_id) if not job: return False job.worker_id = worker_id job.status = JobStatus.ASSIGNED job.assigned_at = datetime.utcnow() save_job_event({ "event": "assigned", "job_id": job_id, "worker_id": worker_id, "timestamp": datetime.utcnow().isoformat(), }) return True def mark_job_running(job_id: str) -> bool: job = _jobs_by_id.get(job_id) if job: job.status = JobStatus.RUNNING return True return False def complete_job(job_id: str, result: JobResult) -> bool: job = _jobs_by_id.get(job_id) if not job: return False job.status = JobStatus.COMPLETED job.result = result.output job.completed_at = datetime.utcnow() save_job_event({ "event": "completed", "job_id": job_id, "worker_id": result.worker_id, "latency_ms": result.latency_ms, "timestamp": datetime.utcnow().isoformat(), }) return True def fail_job(job_id: str, reason: str) -> bool: job = _jobs_by_id.get(job_id) if job: job.status = JobStatus.FAILED job.error_reason = reason job.completed_at = datetime.utcnow() save_job_event({ "event": "failed", "job_id": job_id, "reason": reason, "timestamp": datetime.utcnow().isoformat(), }) return True return False def reject_job(job_id: str, reason: str) -> bool: job = _jobs_by_id.get(job_id) if job: job.status = JobStatus.REJECTED job.error_reason = reason job.completed_at = datetime.utcnow() return True return False def expire_old_jobs() -> None: from app.config import Settings ttl = Settings.get_job_ttl_seconds() now = datetime.utcnow() expired = [] for jid, job in _jobs_by_id.items(): if job.created_at + timedelta(seconds=ttl) < now and job.status in (JobStatus.QUEUED, JobStatus.ASSIGNED): expired.append(jid) for jid in expired: fail_job(jid, "expired") def get_job(job_id: str) -> Optional[InferenceJob]: expire_old_jobs() return _jobs_by_id.get(job_id) def list_jobs_for_session(session_id: str) -> List[InferenceJob]: expire_old_jobs() return [j for j in _jobs_by_id.values() if j.session_id == session_id]