同步移动端工程、设备控制改进与发布资料

This commit is contained in:
2026-09-26 19:21:12 +08:00
parent 912ba432cf
commit 98eadc5c8b
137 changed files with 9274 additions and 85 deletions
@@ -0,0 +1 @@
"""Versioned mobile control transport, independent of web presentation."""
@@ -0,0 +1,252 @@
"""BlueZ GATT adapter. Failure is isolated from the display service."""
import asyncio
import logging
import threading
from dbus_next import Variant, DBusError, BusType, PropertyAccess
from dbus_next.aio import MessageBus
from dbus_next.service import ServiceInterface, method, dbus_property
from .protocol import fragments
log = logging.getLogger(__name__)
ROOT = '/org/qimiaoscreen/mobile'
SERVICE = ROOT + '/service0'
AD = ROOT + '/advertisement0'
SERVICE_UUID = '9f57a001-6c31-4c58-bc22-1f728b641001'
RX_UUID = '9f57a002-6c31-4c58-bc22-1f728b641001'
TX_UUID = '9f57a003-6c31-4c58-bc22-1f728b641001'
class ObjectManager(ServiceInterface):
def __init__(self, objects):
super().__init__('org.freedesktop.DBus.ObjectManager')
self.objects = objects
@method()
def GetManagedObjects(self) -> 'a{oa{sa{sv}}}':
return {path: {obj.name: obj.properties()} for path, obj in self.objects.items()}
class GattService(ServiceInterface):
def __init__(self):
super().__init__('org.bluez.GattService1')
@dbus_property(access=PropertyAccess.READ)
def UUID(self) -> 's':
return SERVICE_UUID
@dbus_property(access=PropertyAccess.READ)
def Primary(self) -> 'b':
return True
def properties(self):
return dict(UUID=Variant('s', self.UUID), Primary=Variant('b', True))
class Characteristic(ServiceInterface):
def __init__(self, uuid, runtime):
super().__init__('org.bluez.GattCharacteristic1')
self.uuid, self.runtime = uuid, runtime
self.notifying = False
@dbus_property(access=PropertyAccess.READ)
def UUID(self) -> 's':
return self.uuid
@dbus_property(access=PropertyAccess.READ)
def Service(self) -> 'o':
return SERVICE
@dbus_property(access=PropertyAccess.READ)
def Flags(self) -> 'as':
return ['write'] if self.uuid == RX_UUID else ['notify']
@dbus_property(access=PropertyAccess.READ)
def Value(self) -> 'ay':
return b''
@dbus_property(access=PropertyAccess.READ)
def Notifying(self) -> 'b':
return self.notifying
@method()
async def WriteValue(self, value: 'ay', options: 'a{sv}'):
if self.uuid != RX_UUID:
raise DBusError('org.bluez.Error.NotPermitted', 'Not permitted')
await self.runtime.write(value, options)
@method()
def StartNotify(self):
if self.uuid != TX_UUID:
raise DBusError('org.bluez.Error.NotSupported', 'Not supported')
self.notifying = True
@method()
def StopNotify(self):
self.notifying = False
def properties(self):
return dict(UUID=Variant('s', self.UUID), Service=Variant('o', SERVICE), Flags=Variant('as', self.Flags))
class Advertisement(ServiceInterface):
def __init__(self, short_id):
super().__init__('org.bluez.LEAdvertisement1')
self.local_name = 'QMS-' + short_id
@dbus_property(access=PropertyAccess.READ)
def Type(self) -> 's':
return 'peripheral'
@dbus_property(access=PropertyAccess.READ)
def ServiceUUIDs(self) -> 'as':
return [SERVICE_UUID]
@dbus_property(access=PropertyAccess.READ)
def LocalName(self) -> 's':
return self.local_name
@method()
def Release(self):
pass
class BluezRuntime:
def __init__(self, sessions, enabled=False):
self.sessions = sessions
self.enabled = enabled
self.stop_event = threading.Event()
self.thread = None
self.state = dict(available=False, reason='disabled')
self.bus = self.manager = self.ad_manager = None
self.advertising = False
self.tx = None
self.message_id = 0
self.write_lock = None
self.counters = dict(rx_fragments=0, tx_fragments=0, rx_bytes=0, tx_bytes=0, last_response_kind=0, mtu=23)
def status(self):
return {**self.state, **self.sessions.status(), 'transport': dict(self.counters)}
def start(self):
if not self.enabled or self.thread:
return
self.thread = threading.Thread(target=lambda: asyncio.run(self._run()), name='mobile-ble', daemon=True)
self.thread.start()
def close(self):
self.stop_event.set()
if self.thread:
self.thread.join(timeout=8)
async def _proxy(self, path):
introspection = await asyncio.wait_for(self.bus.introspect('org.bluez', path), 5)
return self.bus.get_proxy_object('org.bluez', path, introspection)
async def _disconnect(self, owner):
was_owner = self.sessions.owner == owner
try:
proxy = await self._proxy(owner)
await asyncio.wait_for(proxy.get_interface('org.bluez.Device1').call_disconnect(), 3)
except Exception:
pass
self.sessions.disconnect(owner)
if was_owner:
self.message_id = 0
async def write(self, value, options):
self.counters['rx_fragments'] += 1
self.counters['rx_bytes'] += len(value)
owner = options.get('device')
if not owner or options.get('offset', Variant('q', 0)).value != 0:
raise DBusError('org.bluez.Error.NotAuthorized', 'Invalid request')
owner = owner.value
if not self.sessions.connect(owner):
await self._disconnect(owner)
raise DBusError('org.bluez.Error.NotAuthorized', 'Busy')
async with self.write_lock:
try:
if not self.tx.notifying:
raise ValueError()
response = await asyncio.to_thread(self.sessions.accept, owner, bytes(value))
if response is None:
return
mtu = min(self.sessions.peer_mtu, max(23, options.get('mtu', Variant('q', 23)).value))
self.counters['mtu'] = mtu
self.counters['last_response_kind'] = response[0]
for fragment in fragments(response[0], self.message_id, response[1], mtu):
self.tx.emit_properties_changed({'Value': fragment})
self.counters['tx_fragments'] += 1
self.counters['tx_bytes'] += len(fragment)
await asyncio.sleep(0.002)
self.message_id += 1
except Exception:
await self._disconnect(owner)
raise DBusError('org.bluez.Error.NotAuthorized', 'Session rejected') from None
async def _run(self):
while not self.stop_event.is_set():
try:
await self._serve()
except Exception:
self.state = dict(available=False, reason='bluetooth_unavailable')
log.warning('Mobile BLE unavailable; display service remains active')
finally:
if self.sessions.owner is not None:
self.sessions.disconnect(self.sessions.owner)
self.message_id = 0
self.advertising = False
if self.bus:
self.bus.disconnect()
self.bus = None
for _ in range(5):
if self.stop_event.is_set():
return
await asyncio.sleep(1)
async def _serve(self):
self.bus = await asyncio.wait_for(MessageBus(bus_type=BusType.SYSTEM).connect(), 5)
root_proxy = await self._proxy('/')
self.manager = root_proxy.get_interface('org.freedesktop.DBus.ObjectManager')
objects = await asyncio.wait_for(self.manager.call_get_managed_objects(), 5)
adapters = [path for path, interfaces in objects.items() if 'org.bluez.GattManager1' in interfaces and 'org.bluez.LEAdvertisingManager1' in interfaces]
if not adapters:
raise RuntimeError('No GATT peripheral adapter')
adapter = await self._proxy(sorted(adapters)[0])
properties = adapter.get_interface('org.freedesktop.DBus.Properties')
await properties.call_set('org.bluez.Adapter1', 'Powered', Variant('b', True))
await properties.call_set('org.bluez.Adapter1', 'Pairable', Variant('b', False))
self.ad_manager = adapter.get_interface('org.bluez.LEAdvertisingManager1')
self.tx = Characteristic(TX_UUID, self)
exports = {SERVICE: GattService(), SERVICE + '/rx': Characteristic(RX_UUID, self), SERVICE + '/tx': self.tx}
self.bus.export(ROOT, ObjectManager(exports))
for path, interface in exports.items():
self.bus.export(path, interface)
self.bus.export(AD, Advertisement(self.sessions.identity()['short_id']))
await asyncio.wait_for(adapter.get_interface('org.bluez.GattManager1').call_register_application(ROOT, {}), 5)
self.write_lock = asyncio.Lock()
self.state = dict(available=True, reason=None)
while not self.stop_event.is_set():
objects = await asyncio.wait_for(self.manager.call_get_managed_objects(), 5)
connected = [path for path, interfaces in objects.items()
if interfaces.get('org.bluez.Device1', {}).get('Connected', Variant('b', False)).value]
owner = self.sessions.owner
if owner and owner not in connected:
self.sessions.disconnect(owner)
self.message_id = 0
for path in connected:
if not self.sessions.connect(path):
await self._disconnect(path)
if self.sessions.expired():
await self._disconnect(self.sessions.owner)
should_advertise = self.sessions.owner is None
if should_advertise != self.advertising:
if should_advertise:
await asyncio.wait_for(self.ad_manager.call_register_advertisement(AD, {}), 5)
else:
await asyncio.wait_for(self.ad_manager.call_unregister_advertisement(AD), 5)
self.advertising = should_advertise
await asyncio.sleep(0.5)
if self.sessions.owner:
await self._disconnect(self.sessions.owner)
@@ -0,0 +1,222 @@
"""Explicit mobile capability surface, sharing existing device business operations."""
from __future__ import annotations
import base64
import json
import time
import threading
from io import BytesIO
from pathlib import Path
from uuid import UUID, uuid4
from contextlib import nullcontext
from app.persistence import atomic_write_bytes
from app.demo_library import demo_template, demo_animation
from app.display.service import ANIMATION_PLAYBACK_SPEEDS, AnimationPlaybackConflictError
from app.templates.store import TemplateConflictError, TemplateNotFoundError
from .session import RpcError, checked_name
SETTINGS = ("brightness", "orientation", "matrix_refresh_rate_limit_hz", "performance_mode_enabled")
class MobileControl:
def __init__(self, control, network, templates, animations, resolve, apply, get_default, set_default,
status, storage, library_order=None):
self.control, self.network = control, network
self.templates, self.animations = templates, animations
self.resolve, self.apply = resolve, apply
self.get_default, self.set_default = get_default, set_default
self.status, self.storage = status, storage
self.tasks = {}
self.library_order = library_order
self.path = Path(control.store.data_dir) / "mobile/identity.json"
if self.path.exists():
identity = json.loads(self.path.read_text(encoding="utf-8"))
if set(identity) != {"schema", "device_id", "name"} or identity["schema"] != 1:
raise ValueError("Unsupported mobile identity")
UUID(identity["device_id"])
checked_name(identity["name"])
self.identity_record = identity
else:
self.identity_record = dict(schema=1, device_id=str(uuid4()), name="奇妙小屏幕")
self._persist_identity(self.identity_record)
def _persist_identity(self, identity):
atomic_write_bytes(self.path, (json.dumps(identity, ensure_ascii=False) + "\n").encode("utf-8"))
def identity(self):
record = self.identity_record
return dict(device_id=record["device_id"], device_name=record["name"],
short_id=record["device_id"].replace("-", "")[:8], protocol_major=1, protocol_minor=0,
capabilities=["status", "settings", "library", "playback", "frame", "device_name", "wifi", "wifi_scan"],
limits=dict(max_message_bytes=65536, library_page_size=50, preview_interval_ms=2000))
def settings(self):
config = self.control.store.config
values = {key: config[key] for key in SETTINGS}
values["prompt_delay_seconds"] = self.network.get_cached_status()["prompt_delay_seconds"]
return dict(revision=self.control.revision(values), values=values, writable_fields=list(values),
allowed=dict(orientation=[0, 90, 180, 270], matrix_refresh_rate_limit_hz=[15, 20, 30, 45, 60, 80, 100],
playback_speeds=list(ANIMATION_PLAYBACK_SPEEDS)))
def _content(self, params):
reference = {key: params.get(key) for key in ("type", "id")}
resolved = self.resolve(reference)
if params.get("revision") != resolved["revision"]:
raise RpcError("CONFLICT", "内容已改变,请刷新")
return resolved
@staticmethod
def _png(image):
output = BytesIO()
image.save(output, format="PNG")
return dict(mime="image/png", data_base64=base64.b64encode(output.getvalue()).decode("ascii"))
def dispatch(self, method, params):
try:
with self.control.lock, (self.network.configuration_lock if method in ('wifi.get', 'wifi.set') else nullcontext()):
return self._dispatch(method, params)
except (TemplateConflictError, AnimationPlaybackConflictError):
raise RpcError("CONFLICT", "内容或播放会话已改变,请刷新") from None
except TemplateNotFoundError:
raise RpcError("NOT_FOUND", "内容不存在") from None
except (ValueError, TypeError, KeyError):
raise RpcError("BAD_REQUEST", "参数无效,请检查后重试") from None
def _dispatch(self, method, params):
if method == 'wifi.scan':
return dict(networks=self.network.backend.scan(), scanned_at_ms=int(time.time() * 1000))
if method == 'wifi.get':
return self.wifi()
if method == 'wifi.set':
before = self.wifi()
if params.get('expected_revision') != before['revision']:
raise RpcError('CONFLICT', '网络设置已改变,请刷新')
action = params.get('password_action')
security = params.get('security')
if security not in ('open', 'wpa-psk') or action not in ('keep', 'replace', 'none'):
raise ValueError()
saved = before.get('saved') or {}
if security == 'open':
if action != 'none' or params.get('password'):
raise ValueError()
elif action == 'keep':
if params.get('ssid') != saved.get('ssid') or not saved.get('password_configured') or params.get('password'):
raise ValueError()
elif action != 'replace' or not params.get('password'):
raise ValueError()
values = {key: params[key] for key in ('ssid', 'security', 'ipv4_mode', 'address', 'prefix', 'gateway', 'dns_servers', 'activation') if key in params}
values['password'] = params.get('password') if action == 'replace' else None
result = self.network.save_settings(values)
task_id = result['operation_id']
if len(self.tasks) >= 64:
del self.tasks[next(iter(self.tasks))]
self.tasks[task_id] = dict(state='pending' if result['activation'] == 'immediate' else 'succeeded', stage='saved')
if result['activation'] == 'immediate':
threading.Thread(target=self._activate, args=(task_id,), name='mobile-wifi', daemon=True).start()
return dict(revision=self.wifi()['revision'], task_id=task_id)
if method == 'task.get':
task_id = params.get('task_id')
if not isinstance(task_id, str) or task_id not in self.tasks:
raise RpcError('NOT_FOUND', '任务不存在或已过期')
return dict(self.tasks[task_id])
if method == "device.rename":
identity = {**self.identity_record, "name": checked_name(params.get("name"))}
self._persist_identity(identity)
self.identity_record = identity
return self.identity()
if method == "status.get":
return self.status()
if method == "storage.get":
return self.storage()
if method == "settings.get":
return self.settings()
if method == "settings.patch":
if params.get("expected_revision") != self.settings()["revision"]:
raise RpcError("CONFLICT", "设置已改变,请刷新")
changes = params.get("changes")
if not isinstance(changes, dict) or not changes or set(changes) - {*SETTINGS, "prompt_delay_seconds"}:
raise RpcError("BAD_REQUEST", "不支持的设置字段")
# Separate NetworkManager persistence from display configuration.
if "prompt_delay_seconds" in changes and len(changes) != 1:
raise RpcError("BAD_REQUEST", "网络提示延迟须单独提交")
if "prompt_delay_seconds" in changes:
self.network.save_prompt_delay(changes["prompt_delay_seconds"])
else:
self.control.update_config(changes)
return self.settings()
if method == "library.list":
items = []
for kind, records in (("template", [demo_template()] + self.templates.list()["templates"]),
("animation", [demo_animation()] + self.animations.list()["animations"])):
for item in records:
items.append(dict(type=kind, id=item["id"], name=item["name"], revision=item["revision"],
is_demo=bool(item.get("is_demo", False)), playable=kind == "template" or item.get("frame_count", 0) > 0,
thumbnail_revision=item["revision"]))
if self.library_order is not None:
user_items = [item for item in items if not item['is_demo']]
order = self.library_order.get([{'type': item['type'], 'id': item['id']} for item in user_items])['items']
by_key = {(item['type'], item['id']): item for item in user_items}
items = [item for item in items if item['is_demo']] + [by_key[(ref['type'], ref['id'])] for ref in order]
revision = self.control.revision(items)
limit = params.get("limit", 20)
if type(limit) is not int or not 1 <= limit <= 50:
raise ValueError()
offset = 0
if params.get("cursor"):
cursor = params["cursor"].split(":")
if len(cursor) != 2 or cursor[0] != revision:
raise RpcError("CONFLICT", "内容列表已改变,请重新加载")
offset = int(cursor[1])
if not 0 <= offset <= len(items):
raise ValueError()
next_offset = offset + limit
return dict(library_revision=revision, items=items[offset:next_offset],
next_cursor=f"{revision}:{next_offset}" if next_offset < len(items) else None)
if method == "library.thumbnail":
content = self._content(params)
return self._png(content["image"] if content["type"] == "template" else content["frames"][0]["image"])
if method in ("content.play", "content.default.set"):
content = self._content(params)
if method == "content.default.set":
return self.set_default({"type": content["type"], "id": content["id"]})
self.apply(content, notify_activity=True)
return {key: content[key] for key in ("type", "id", "name", "revision", "is_demo")}
if method == "content.default.get":
content = self.get_default()
return {key: content[key] for key in ("type", "id", "name", "revision", "is_demo")}
if method == "playback.patch":
if set(params) - {"session_id", "paused", "position_ms", "speed"}:
raise ValueError()
playback = self.control.display.control_animation_playback(params.get("session_id"),
paused=params.get("paused"), position_ms=params.get("position_ms"), speed=params.get("speed"))
return dict(animation_playback=playback, state_revision=self.control.display.get_state_revision())
if method == "frame.get":
image = self.control.display.get_current_frame()
revision = self.control.revision(base64.b64encode(image.tobytes()).decode("ascii"))
if params.get("known_revision") == revision:
return dict(unchanged=True, frame_revision=revision)
return dict(unchanged=False, frame_revision=revision, sampled_at_ms=int(time.time() * 1000), **self._png(image))
raise RpcError("UNSUPPORTED_CAPABILITY", "设备不支持此操作")
def wifi(self):
status = self.network.get_status(include_secret=False)
saved = status.get('saved')
if saved:
saved = {key: value for key, value in saved.items() if key != 'password'}
saved.setdefault('security', 'wpa-psk' if saved.get('password_configured') else 'open')
revision = self.control.revision(dict(saved=saved, generation=self.network.settings_revision))
return {**status, 'saved': saved, 'revision': revision}
def _activate(self, task_id):
self.tasks[task_id] = dict(state='running', stage='connecting')
try:
self.network.activate_saved(task_id)
operation = self.network.get_cached_status().get('operation') or {}
state = operation.get('state') if operation.get('id') == task_id else 'failed'
self.tasks[task_id] = dict(state='succeeded' if state == 'succeeded' else 'failed', stage='finished')
if state != 'succeeded':
code = operation.get('error_code', 'UNKNOWN')
self.tasks[task_id]['error_code'] = code if code in {'AUTH_FAILED', 'NETWORK_UNAVAILABLE', 'IP_CONFIG_FAILED', 'UNKNOWN'} else 'UNKNOWN'
except Exception:
self.tasks[task_id] = dict(state='failed', stage='finished', error_code='UNKNOWN')
@@ -0,0 +1,165 @@
"""QMS BLE v1 framing and ephemeral authenticated encryption.
This encrypts traffic, but does not authenticate the identity of an App.
Each instance belongs to exactly one connection; discard it after any error.
"""
from __future__ import annotations
import base64
import hashlib
import json
import os
import struct
import time
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import ec
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
MAX_MESSAGE = 65536
HEADER = struct.Struct("!HBBIHHI")
LABEL = b"QMS-BLE-1"
class ProtocolError(ValueError):
"""Fixed public error: never include untrusted payload or crypto details."""
def fragments(kind: int, message_id: int, payload: bytes, mtu: int):
if kind not in (1, 2, 3, 4) or not 0 <= message_id <= 0xFFFFFFFF:
raise ProtocolError("BAD_REQUEST")
if not 1 <= len(payload) <= MAX_MESSAGE or not 23 <= mtu <= 517:
raise ProtocolError("BAD_REQUEST")
capacity = mtu - 19
count = (len(payload) + capacity - 1) // capacity
for index in range(count):
yield HEADER.pack(0x514D, 1, kind, message_id, index, count, len(payload)) + payload[index * capacity:(index + 1) * capacity]
class Reassembler:
def __init__(self, clock=time.monotonic):
self.clock = clock
self.next_message = 0
self.pending = None
self.buffer = bytearray()
self.index = 0
self.started = 0.0
def accept(self, part: bytes):
if not 16 < len(part) <= 514:
raise ProtocolError("BAD_REQUEST")
magic, wire, kind, mid, index, count, total = HEADER.unpack(part[:16])
if magic != 0x514D or wire != 1 or kind not in (1, 2, 3, 4) or not 1 <= count <= total <= MAX_MESSAGE:
raise ProtocolError("BAD_REQUEST")
signature = (kind, mid, count, total)
if self.pending is None:
if index != 0 or mid != self.next_message:
raise ProtocolError("BAD_REQUEST")
self.pending = signature
self.started = self.clock()
if self.clock() - self.started >= 15:
raise ProtocolError("TIMEOUT")
if self.pending != signature or index != self.index or index >= count:
raise ProtocolError("BAD_REQUEST")
self.buffer.extend(part[16:])
self.index += 1
if len(self.buffer) > total or (self.index < count and len(self.buffer) >= total):
raise ProtocolError("BAD_REQUEST")
if self.index != count:
return None
if len(self.buffer) != total:
raise ProtocolError("BAD_REQUEST")
result = kind, bytes(self.buffer)
self.pending = None
self.buffer.clear()
self.index = 0
self.next_message += 1
return result
def json_object(raw: bytes) -> dict:
def unique(pairs):
result = {}
for key, value in pairs:
if key in result:
raise ValueError("duplicate")
result[key] = value
return result
try:
value = json.loads(raw.decode("utf-8"), object_pairs_hook=unique,
parse_constant=lambda _: (_ for _ in ()).throw(ValueError()))
if not isinstance(value, dict):
raise ValueError()
return value
except (ValueError, UnicodeError, RecursionError):
raise ProtocolError("BAD_REQUEST") from None
class Handshake:
def __init__(self, *, private_key=None, random_bytes=None):
# Explicit values are exclusively for public interoperability fixtures.
self.key = private_key or ec.generate_private_key(ec.SECP256R1())
self.random = random_bytes if random_bytes is not None else os.urandom(32)
if len(self.random) != 32:
raise ValueError("random length")
public = self.key.public_key().public_bytes(serialization.Encoding.X962, serialization.PublicFormat.UncompressedPoint)
self.hello = json.dumps(dict(protocol_major=1, protocol_minor=0,
public_key=base64.b64encode(public).decode("ascii"),
random=base64.b64encode(self.random).decode("ascii")), separators=(",", ":")).encode("utf-8")
def finish(self, peer_hello: bytes, *, server: bool):
if len(peer_hello) > 1024:
raise ProtocolError("BAD_REQUEST")
peer = json_object(peer_hello)
if type(peer.get("protocol_major")) is not int or peer["protocol_major"] != 1:
raise ProtocolError("UNSUPPORTED_VERSION")
if type(peer.get("protocol_minor")) is not int or peer["protocol_minor"] < 0:
raise ProtocolError("BAD_REQUEST")
try:
public = base64.b64decode(peer["public_key"], validate=True)
random = base64.b64decode(peer["random"], validate=True)
if len(public) != 65 or public[0] != 4 or len(random) != 32:
raise ValueError()
peer_key = ec.EllipticCurvePublicKey.from_encoded_point(ec.SECP256R1(), public)
secret = self.key.exchange(ec.ECDH(), peer_key)
except (KeyError, TypeError, ValueError):
raise ProtocolError("BAD_REQUEST") from None
c, s = (peer_hello, self.hello) if server else (self.hello, peer_hello)
cr, sr = (random, self.random) if server else (self.random, random)
transcript = hashlib.sha256(struct.pack("!I", len(c)) + c + struct.pack("!I", len(s)) + s).digest()
material = HKDF(algorithm=hashes.SHA256(), length=72,
salt=hashlib.sha256(cr + sr).digest(), info=LABEL + transcript).derive(secret)
return Cipher(material, transcript, server=server)
class Cipher:
def __init__(self, material: bytes, transcript: bytes, *, server: bool):
self.tx_direction = 1 if server else 0
self.rx_direction = 1 - self.tx_direction
self.keys = (AESGCM(material[:32]), AESGCM(material[32:64]))
self.prefixes = (material[64:68], material[68:72])
self.transcript = transcript
self.tx_sequence = self.rx_sequence = 0
self.failed = False
def encrypt(self, payload: bytes) -> bytes:
if self.failed or not 1 <= len(payload) <= MAX_MESSAGE - 24 or self.tx_sequence >= 2**64:
raise ProtocolError("BAD_REQUEST")
seq = struct.pack("!Q", self.tx_sequence)
d = self.tx_direction
result = seq + self.keys[d].encrypt(self.prefixes[d] + seq, payload, LABEL + self.transcript + bytes([d]) + seq)
self.tx_sequence += 1
return result
def decrypt(self, record: bytes) -> bytes:
try:
if self.failed or not 25 <= len(record) <= MAX_MESSAGE or int.from_bytes(record[:8], "big") != self.rx_sequence:
raise ValueError()
seq, d = record[:8], self.rx_direction
result = self.keys[d].decrypt(self.prefixes[d] + seq, record[8:], LABEL + self.transcript + bytes([d]) + seq)
self.rx_sequence += 1
return result
except Exception:
self.failed = True
raise ProtocolError("BAD_REQUEST") from None
@@ -0,0 +1,133 @@
"""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)