"""Single-owner encrypted RPC session; transport cannot bypass this gate.""" from __future__ import annotations import json import threading import time from .protocol import Handshake, ProtocolError, Reassembler class RpcError(Exception): def __init__(self, code, message): self.code, self.message = code, message super().__init__(code) def checked_name(value): if not isinstance(value, str) or not value.strip() or len(value) > 40: raise RpcError("BAD_REQUEST", "名称须为 1 至 40 个字符") if any(ord(char) < 32 or ord(char) == 127 for char in value): raise RpcError("BAD_REQUEST", "名称包含不支持的字符") return value.strip() class SessionManager: def __init__(self, dispatch, identity, *, clock=time.monotonic): self.dispatch = dispatch self.identity = identity self.clock = clock self.lock = threading.RLock() self.owner = None self.receiver = None self.cipher = None self.opened = False self.client_name = None self.last_request = -1 self.connected_at = self.last_activity = 0 self.generation = 0 self.peer_mtu = 23 def connect(self, owner): with self.lock: if self.owner is not None: return self.owner == owner self.owner = owner self.receiver = Reassembler(self.clock) self.connected_at = self.last_activity = self.clock() return True def disconnect(self, owner): with self.lock: if self.owner != owner: return self.owner = self.receiver = self.cipher = self.client_name = None self.opened = False self.last_request = -1 self.peer_mtu = 23 self.generation += 1 def expired(self): with self.lock: if self.owner is None: return False deadline = self.last_activity + 20 if self.opened else self.connected_at + 10 return self.clock() >= deadline or ( self.receiver.pending is not None and self.clock() - self.receiver.started >= 15) def status(self): with self.lock: return dict(connected=self.opened, client_name=self.client_name if self.opened else None, generation=self.generation) def accept(self, owner, fragment): with self.lock: if owner != self.owner or self.expired(): raise ProtocolError("NOT_READY") message = self.receiver.accept(fragment) if message is None: return None kind, payload = message if self.cipher is None: if kind != 1: raise ProtocolError("NOT_READY") handshake = Handshake() self.cipher = handshake.finish(payload, server=True) return 2, handshake.hello if kind != 3: raise ProtocolError("BAD_REQUEST") from .protocol import json_object request = json_object(self.cipher.decrypt(payload)) request_id = request.get("id") if not isinstance(request_id, str) or not 1 <= len(request_id) <= 64 or not request_id.isascii() or not request_id.isdecimal(): raise ProtocolError("BAD_REQUEST") number = int(request_id) response = dict(id=request_id) try: if number <= self.last_request: raise RpcError("CONFLICT", "请求编号已使用") self.last_request = number method, params = request.get("method"), request.get("params") if not isinstance(method, str) or not isinstance(params, dict): raise RpcError("BAD_REQUEST", "请求格式无效") if not self.opened: if method != "session.open": raise ProtocolError("NOT_READY") name = checked_name(params.get("client_name")) peer_mtu = params.get('receive_mtu', 23) if type(peer_mtu) is not int or not 23 <= peer_mtu <= 517: raise RpcError('BAD_REQUEST', '接收分片上限无效') result = self.identity() self.peer_mtu = peer_mtu self.client_name = name self.opened = True self.generation += 1 elif method == "session.ping": result = dict(alive=True) elif method == "session.rename": self.client_name = checked_name(params.get("client_name")) result = dict(client_name=self.client_name) elif method == "session.open": raise RpcError("CONFLICT", "会话已经建立") else: result = self.dispatch(method, params) response.update(ok=True, result=result) self.last_activity = self.clock() except RpcError as error: response.update(ok=False, error=dict(code=error.code, message=error.message, retryable=False)) except ProtocolError: raise except Exception: response.update(ok=False, error=dict(code="DEVICE_ERROR", message="设备操作失败,请刷新状态", retryable=False)) raw = json.dumps(response, ensure_ascii=False, separators=(",", ":"), allow_nan=False).encode("utf-8") return 3, self.cipher.encrypt(raw)