Spaces:
Paused
Paused
| 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] | |