同步移动端工程、设备控制改进与发布资料
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user