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)