155 lines
6.6 KiB
Python
155 lines
6.6 KiB
Python
"""Single-owner encrypted RPC session; transport cannot bypass this gate."""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import threading
|
|
import time
|
|
|
|
from .protocol import Handshake, ProtocolError, Reassembler
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
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, wifi_address=lambda: None,
|
|
on_session_open=None):
|
|
self.dispatch = dispatch
|
|
self.identity = identity
|
|
self.clock = clock
|
|
self.on_session_open = on_session_open
|
|
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
|
|
from .wifi import WifiSession
|
|
self.wifi = WifiSession(self, wifi_address)
|
|
|
|
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.wifi.close()
|
|
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, active_transport=("wifi" if self.wifi.valid() else "ble") if self.opened else None,
|
|
wifi_counters=dict(self.wifi.counters))
|
|
|
|
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
|
|
if self.on_session_open is not None:
|
|
try:
|
|
self.on_session_open()
|
|
except Exception:
|
|
# A failed display notification must not invalidate
|
|
# an already established encrypted session.
|
|
logger.exception("Mobile session-open notification failed")
|
|
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", "会话已经建立")
|
|
elif method == "transport.offer":
|
|
result = self.wifi.offer()
|
|
elif method == "transport.close":
|
|
self.wifi.close()
|
|
result = dict(closed=True)
|
|
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)
|