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)