File size: 5,071 Bytes
ab6b2ca | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | from __future__ import annotations
from dataclasses import dataclass
from io import BytesIO
from typing import TYPE_CHECKING, Any, Callable, Type
if TYPE_CHECKING:
from .adapters.base import ModelServerAdapter
class TorchSerializer:
@staticmethod
def to_bytes(data: Any) -> bytes:
import torch
buffer = BytesIO()
torch.save(data, buffer)
return buffer.getvalue()
@staticmethod
def from_bytes(data: bytes) -> Any:
import torch
buffer = BytesIO(data)
return torch.load(buffer, weights_only=False)
@dataclass
class EndpointHandler:
handler: Callable
requires_input: bool = True
class BaseInferenceServer:
"""Minimal REP server used by model adapters."""
def __init__(self, host: str = "*", port: int = 5555, api_token: str | None = None):
import zmq
self.running = True
self.context = zmq.Context()
self._zmq = zmq
self.socket = self.context.socket(zmq.REP)
self.socket.bind(f"tcp://{host}:{port}")
self._endpoints: dict[str, EndpointHandler] = {}
self.api_token = api_token
self.register_endpoint("ping", self._handle_ping, requires_input=False)
self.register_endpoint("kill", self._kill_server, requires_input=False)
def _kill_server(self) -> dict[str, str]:
self.running = False
return {"status": "ok", "message": "server will stop"}
def _handle_ping(self) -> dict[str, str]:
return {"status": "ok", "message": "Server is running"}
def register_endpoint(self, name: str, handler: Callable, requires_input: bool = True) -> None:
self._endpoints[name] = EndpointHandler(handler, requires_input)
def _validate_token(self, request: dict[str, Any]) -> bool:
if self.api_token is None:
return True
return request.get("api_token") == self.api_token
def run(self) -> None:
addr = self.socket.getsockopt_string(self._zmq.LAST_ENDPOINT)
print(f"Server is ready and listening on {addr}")
while self.running:
try:
message = self.socket.recv()
request = TorchSerializer.from_bytes(message)
if not self._validate_token(request):
self.socket.send(TorchSerializer.to_bytes({"error": "Unauthorized: Invalid API token"}))
continue
endpoint = request.get("endpoint", "select_action")
if endpoint not in self._endpoints:
raise ValueError(f"Unknown endpoint: {endpoint}")
handler = self._endpoints[endpoint]
result = (
handler.handler(request.get("data", {}))
if handler.requires_input
else handler.handler()
)
self.socket.send(TorchSerializer.to_bytes(result))
except Exception as exc:
print(f"Error in server: {exc}")
self.socket.send(TorchSerializer.to_bytes({"error": str(exc)}))
class ModelInferenceServer(BaseInferenceServer):
"""Standard server wrapper that exposes adapter endpoints."""
def __init__(
self,
adapter: ModelServerAdapter,
*,
host: str = "*",
port: int = 5555,
api_token: str | None = None,
):
super().__init__(host=host, port=port, api_token=api_token)
self.adapter = adapter
self.register_endpoint("metadata", adapter.metadata, requires_input=False)
self.register_endpoint("reset", adapter.reset, requires_input=False)
self.register_endpoint("select_action", adapter.select_action, requires_input=True)
self.register_endpoint("select_action_chunk", adapter.select_action_chunk, requires_input=True)
if hasattr(adapter, "select_action_chunk_rtc"):
self.register_endpoint("select_action_chunk_rtc", self._select_action_chunk_rtc, requires_input=True)
def _select_action_chunk_rtc(self, data: dict[str, Any]) -> dict[str, Any]:
return self.adapter.select_action_chunk_rtc(
data["observation"],
prev_chunk_leftover=data.get("prev_chunk_leftover"),
inference_delay=int(data.get("inference_delay", 0)),
execution_horizon=int(data["execution_horizon"]),
rtc_options=dict(data.get("rtc_options", {})),
)
_ADAPTER_REGISTRY: dict[str, Type[Any]] = {}
def register_adapter(adapter_cls: Type[Any]) -> Type[Any]:
if not adapter_cls.name:
raise ValueError("Adapter class must define a non-empty `name`.")
_ADAPTER_REGISTRY[adapter_cls.name] = adapter_cls
return adapter_cls
def get_adapter_class(name: str) -> Type[Any]:
try:
return _ADAPTER_REGISTRY[name]
except KeyError as exc:
known = ", ".join(sorted(_ADAPTER_REGISTRY)) or "<none>"
raise KeyError(f"Unknown adapter '{name}'. Available adapters: {known}") from exc
def list_adapters() -> list[str]:
return sorted(_ADAPTER_REGISTRY)
|