Spaces:
Paused
Paused
| import json | |
| from datetime import datetime | |
| from typing import Dict, Optional | |
| from app.models import JobStatus | |
| from app.session_store import attach_worker, detach_worker, update_worker_heartbeat, get_session | |
| from app.job_queue import ( | |
| assign_job_to_worker, mark_job_running, complete_job, fail_job, get_job, | |
| ) | |
| from app.receipts import create_receipt | |
| from app.capabilities import register_capabilities | |
| _worker_sockets: Dict[str, any] = {} # key: session_id:worker_id -> WebSocket | |
| _client_sockets: Dict[str, any] = {} # key: session_id -> WebSocket (v1: one client per session) | |
| _mac_node_sockets: Dict[str, any] = {} # key: node_id -> WebSocket for Mac DMG nodes | |
| async def connect_worker_websocket(session_id: str, worker_id: str, websocket) -> bool: | |
| key = f"{session_id}:{worker_id}" | |
| _worker_sockets[key] = websocket | |
| return True | |
| async def connect_client_websocket(session_id: str, websocket) -> bool: | |
| _client_sockets[session_id] = websocket | |
| return True | |
| async def disconnect_worker(session_id: str, worker_id: str) -> None: | |
| key = f"{session_id}:{worker_id}" | |
| _worker_sockets.pop(key, None) | |
| detach_worker(session_id, worker_id) | |
| async def disconnect_client(session_id: str) -> None: | |
| _client_sockets.pop(session_id, None) | |
| async def send_to_worker(session_id: str, worker_id: str, message: dict) -> bool: | |
| key = f"{session_id}:{worker_id}" | |
| ws = _worker_sockets.get(key) | |
| if ws: | |
| try: | |
| await ws.send_json(message) | |
| return True | |
| except Exception: | |
| return False | |
| return False | |
| async def send_to_client(session_id: str, message: dict) -> bool: | |
| ws = _client_sockets.get(session_id) | |
| if ws: | |
| try: | |
| await ws.send_json(message) | |
| return True | |
| except Exception: | |
| return False | |
| return False | |
| async def broadcast_session_update(session_id: str) -> None: | |
| session = get_session(session_id) | |
| if not session: | |
| return | |
| msg = { | |
| "op": "session_update", | |
| "session_id": session_id, | |
| "worker_count": len(session.workers), | |
| "job_count": len([j for j in session.workers.values()]), | |
| } | |
| await send_to_client(session_id, msg) | |
| async def route_job_offer_to_worker(job) -> bool: | |
| if not job.worker_id: | |
| return False | |
| msg = { | |
| "op": "job_offer", | |
| "job_id": job.job_id, | |
| "job_type": job.job_type.value, | |
| "payload": job.payload, | |
| "privacy_mode": job.privacy_mode.value, | |
| } | |
| return await send_to_worker(job.session_id, job.worker_id, msg) | |
| async def route_job_result_to_client(result) -> bool: | |
| job = get_job(result.job_id) | |
| if not job: | |
| return False | |
| msg = { | |
| "op": "job_result", | |
| "job_id": result.job_id, | |
| "status": "completed", | |
| "output": result.output, | |
| "latency_ms": result.latency_ms, | |
| } | |
| return await send_to_client(job.session_id, msg) | |
| async def handle_worker_message(session_id: str, worker_id: str, message: dict) -> None: | |
| op = message.get("op") | |
| if op == "worker_hello": | |
| from app.models import WorkerState, WorkerRuntimeType | |
| runtime = WorkerRuntimeType(message.get("runtime_type", "safari_wasm")) | |
| worker = WorkerState( | |
| worker_id=worker_id, | |
| session_id=session_id, | |
| runtime_type=runtime, | |
| device_public_key=message.get("device_public_key"), | |
| ) | |
| attach_worker(session_id, worker) | |
| await send_to_worker(session_id, worker_id, {"op": "worker_welcome", "worker_id": worker_id}) | |
| elif op == "capabilities": | |
| register_capabilities(session_id, worker_id, message.get("capabilities", [])) | |
| await broadcast_session_update(session_id) | |
| elif op == "heartbeat": | |
| update_worker_heartbeat(session_id, worker_id) | |
| elif op == "job_accept": | |
| mark_job_running(message.get("job_id")) | |
| elif op == "job_progress": | |
| job = get_job(message.get("job_id")) | |
| if job: | |
| await send_to_client(job.session_id, { | |
| "op": "job_progress", | |
| "job_id": message.get("job_id"), | |
| "progress": message.get("progress", 0.0), | |
| }) | |
| elif op == "job_result": | |
| from app.models import JobResult | |
| from app.session_store import list_workers | |
| result = JobResult( | |
| job_id=message["job_id"], | |
| worker_id=worker_id, | |
| output=message.get("output", {}), | |
| latency_ms=message.get("latency_ms", 0), | |
| input_hash=message.get("input_hash"), | |
| output_hash=message.get("output_hash"), | |
| device_signature=message.get("device_signature"), | |
| ) | |
| job = get_job(result.job_id) | |
| if job: | |
| complete_job(result.job_id, result) | |
| workers = list_workers(session_id) | |
| worker = next((w for w in workers if w.worker_id == worker_id), None) | |
| if worker: | |
| create_receipt(job, worker, result) | |
| await route_job_result_to_client(result) | |
| elif op == "disconnect": | |
| await disconnect_worker(session_id, worker_id) | |
| async def handle_client_message(session_id: str, message: dict) -> None: | |
| op = message.get("op") | |
| if op == "client_hello": | |
| await send_to_client(session_id, {"op": "client_welcome", "session_id": session_id}) | |
| elif op == "session_update": | |
| await broadcast_session_update(session_id) | |
| elif op == "chat_submit": | |
| # iPhone client submits a prompt; route to MacBook node if available | |
| from app.models import JobType, PrivacyMode | |
| from app.job_queue import create_job, select_worker_for_job, assign_job_to_worker | |
| job = create_job( | |
| session_id=session_id, | |
| job_type=JobType.CHAT_COMPLETION, | |
| payload={"prompt": message.get("prompt", ""), "max_tokens": message.get("max_tokens", 512), "temperature": message.get("temperature", 0.7)}, | |
| privacy_mode=PrivacyMode.LOCAL_ONLY, | |
| constraints={}, | |
| ) | |
| # Prefer MacBook node; fallback to iPhone worker | |
| mac_nodes = list(_mac_node_sockets.keys()) | |
| if mac_nodes: | |
| node_id = mac_nodes[0] # Pick first available MacBook node | |
| assign_job_to_worker(job.job_id, node_id) | |
| await send_to_client(session_id, {"op": "job_assigned", "job_id": job.job_id, "worker": "MacBook"}) | |
| await route_job_to_mac_node(job, node_id) | |
| else: | |
| # Fallback to iPhone worker scoring | |
| worker_id = select_worker_for_job(session_id, job) | |
| if worker_id: | |
| assign_job_to_worker(job.job_id, worker_id) | |
| await send_to_client(session_id, {"op": "job_assigned", "job_id": job.job_id, "worker": "iPhone"}) | |
| await route_job_offer_to_worker(job) | |
| else: | |
| await send_to_client(session_id, {"op": "error", "error": "No compute workers available. Connect your MacBook DMG or iPhone worker."}) | |
| # Mac Node relay functions | |
| async def connect_mac_node(node_id: str, websocket) -> bool: | |
| _mac_node_sockets[node_id] = websocket | |
| return True | |
| async def disconnect_mac_node(node_id: str) -> None: | |
| _mac_node_sockets.pop(node_id, None) | |
| async def send_to_mac_node(node_id: str, message: dict) -> bool: | |
| ws = _mac_node_sockets.get(node_id) | |
| if ws: | |
| try: | |
| await ws.send_json(message) | |
| return True | |
| except Exception: | |
| return False | |
| return False | |
| async def broadcast_to_mac_nodes(message: dict) -> bool: | |
| if not _mac_node_sockets: | |
| return False | |
| sent = False | |
| for node_id, ws in _mac_node_sockets.items(): | |
| try: | |
| await ws.send_json(message) | |
| sent = True | |
| except Exception: | |
| await disconnect_mac_node(node_id) | |
| return sent | |
| async def route_job_to_mac_node(job, node_id: str) -> bool: | |
| msg = { | |
| "op": "inference.request", | |
| "job_id": job.job_id, | |
| "job_type": job.job_type.value, | |
| "payload": job.payload, | |
| "privacy_mode": job.privacy_mode.value, | |
| } | |
| return await send_to_mac_node(node_id, msg) | |
| async def handle_mac_node_message(node_id: str, message: dict) -> None: | |
| op = message.get("op") | |
| if op == "node.register": | |
| await send_to_mac_node(node_id, {"op": "node.registered", "node_id": node_id}) | |
| elif op == "inference.chunk": | |
| # Stream chunk to client | |
| job = get_job(message.get("job_id")) | |
| if job: | |
| await send_to_client(job.session_id, { | |
| "op": "job_stream", | |
| "job_id": message.get("job_id"), | |
| "chunk": message.get("chunk"), | |
| }) | |
| elif op == "inference.complete": | |
| from app.models import JobResult | |
| from app.session_store import list_workers | |
| result = JobResult( | |
| job_id=message["job_id"], | |
| worker_id=node_id, | |
| output=message.get("output", {}), | |
| latency_ms=message.get("latency_ms", 0), | |
| input_hash=message.get("input_hash"), | |
| output_hash=message.get("output_hash"), | |
| device_signature=message.get("device_signature"), | |
| ) | |
| job = get_job(result.job_id) | |
| if job: | |
| complete_job(result.job_id, result) | |
| workers = list_workers(job.session_id) | |
| worker = next((w for w in workers if w.worker_id == node_id), None) | |
| if worker: | |
| create_receipt(job, worker, result) | |
| await send_to_client(job.session_id, { | |
| "op": "job_complete", | |
| "job_id": result.job_id, | |
| "output": result.output, | |
| "latency_ms": result.latency_ms, | |
| "worker": "MacBook", | |
| "tokens": message.get("tokens_generated", 0), | |
| "tps": message.get("tokens_per_second", 0), | |
| }) | |
| await broadcast_session_update(job.session_id) | |
| elif op == "node.heartbeat": | |
| from app.mac_relay import update_mac_heartbeat | |
| update_mac_heartbeat(node_id) | |
| # Store health metrics from heartbeat payload | |
| if node_id in _mac_node_meta: | |
| _mac_node_meta[node_id]["health"] = { | |
| "cpu_usage": message.get("cpu_usage", 0.0), | |
| "ram_used_mb": message.get("ram_used_mb", 0.0), | |
| "thermal_state": message.get("thermal_state", "nominal"), | |
| "timestamp": datetime.utcnow().isoformat(), | |
| } | |
| elif op == "disconnect": | |
| await disconnect_mac_node(node_id) | |