Files
matrix-screen-controller/核桃派软件源代码/app/mobile/session.py
T

134 lines
5.5 KiB
Python

"""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)