phonegpu-space / app /relay.py
josephrw's picture
Upload folder using huggingface_hub
d958e80 verified
Raw
History Blame Contribute Delete
10.6 kB
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)