Add authenticated WiFi transport with BLE fallback
This commit is contained in:
@@ -51,8 +51,8 @@ class MobileControl:
|
||||
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=1,
|
||||
capabilities=["status", "settings", "library", "library_progressive", "playback", "frame", "device_name", "wifi", "wifi_scan"],
|
||||
short_id=record["device_id"].replace("-", "")[:8], protocol_major=1, protocol_minor=2,
|
||||
capabilities=["status", "settings", "library", "library_progressive", "playback", "frame", "device_name", "wifi", "wifi_scan", "wifi_transport"],
|
||||
limits=dict(max_message_bytes=65536, library_page_size=50, preview_interval_ms=2000))
|
||||
|
||||
def settings(self):
|
||||
|
||||
@@ -134,7 +134,8 @@ class Handshake:
|
||||
|
||||
|
||||
class Cipher:
|
||||
def __init__(self, material: bytes, transcript: bytes, *, server: bool):
|
||||
def __init__(self, material: bytes, transcript: bytes, *, server: bool, label=LABEL):
|
||||
self.label = label
|
||||
self.tx_direction = 1 if server else 0
|
||||
self.rx_direction = 1 - self.tx_direction
|
||||
self.keys = (AESGCM(material[:32]), AESGCM(material[32:64]))
|
||||
@@ -148,7 +149,7 @@ class Cipher:
|
||||
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)
|
||||
result = seq + self.keys[d].encrypt(self.prefixes[d] + seq, payload, self.label + self.transcript + bytes([d]) + seq)
|
||||
self.tx_sequence += 1
|
||||
return result
|
||||
|
||||
@@ -157,7 +158,7 @@ class Cipher:
|
||||
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)
|
||||
result = self.keys[d].decrypt(self.prefixes[d] + seq, record[8:], self.label + self.transcript + bytes([d]) + seq)
|
||||
self.rx_sequence += 1
|
||||
return result
|
||||
except Exception:
|
||||
|
||||
@@ -23,7 +23,7 @@ def checked_name(value):
|
||||
|
||||
|
||||
class SessionManager:
|
||||
def __init__(self, dispatch, identity, *, clock=time.monotonic):
|
||||
def __init__(self, dispatch, identity, *, clock=time.monotonic, wifi_address=lambda: None):
|
||||
self.dispatch = dispatch
|
||||
self.identity = identity
|
||||
self.clock = clock
|
||||
@@ -37,6 +37,8 @@ class SessionManager:
|
||||
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:
|
||||
@@ -55,6 +57,7 @@ class SessionManager:
|
||||
self.opened = False
|
||||
self.last_request = -1
|
||||
self.peer_mtu = 23
|
||||
self.wifi.close()
|
||||
self.generation += 1
|
||||
|
||||
def expired(self):
|
||||
@@ -68,7 +71,8 @@ class SessionManager:
|
||||
def status(self):
|
||||
with self.lock:
|
||||
return dict(connected=self.opened, client_name=self.client_name if self.opened else None,
|
||||
generation=self.generation)
|
||||
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:
|
||||
@@ -119,6 +123,11 @@ class SessionManager:
|
||||
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)
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
"""BLE-owned, single-use encrypted LAN channel. No independent WiFi session."""
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import ipaddress
|
||||
import json
|
||||
import secrets
|
||||
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
|
||||
from .protocol import Cipher, ProtocolError, json_object, MAX_MESSAGE
|
||||
|
||||
LABEL = b'QMS-WIFI-1'
|
||||
|
||||
|
||||
def wifi_cipher(secret, channel_id, *, server):
|
||||
transcript = hashlib.sha256(channel_id.encode('ascii')).digest()
|
||||
material = HKDF(algorithm=hashes.SHA256(), length=72, salt=transcript,
|
||||
info=LABEL + transcript).derive(secret)
|
||||
return Cipher(material, transcript, server=server, label=LABEL)
|
||||
|
||||
|
||||
def encode(value):
|
||||
return json.dumps(value, ensure_ascii=False, separators=(',', ':'), allow_nan=False).encode('utf-8')
|
||||
|
||||
|
||||
class WifiSession:
|
||||
# All methods run under the owning SessionManager's reentrant lock.
|
||||
def __init__(self, sessions, address):
|
||||
self.sessions, self.address = sessions, address
|
||||
self.counters = dict(requests=0, rx_bytes=0, tx_bytes=0)
|
||||
self.close()
|
||||
|
||||
def close(self):
|
||||
self.pending = None
|
||||
self.channel = None
|
||||
self.cipher = None
|
||||
self.last_request = -1
|
||||
self.last_activity = 0
|
||||
|
||||
def valid(self, channel=None):
|
||||
s = self.sessions
|
||||
return bool(self.channel and s.opened and not s.expired() and
|
||||
(channel is None or channel == self.channel) and
|
||||
s.clock() - self.last_activity < 10)
|
||||
|
||||
def offer(self):
|
||||
from .session import RpcError
|
||||
if self.valid():
|
||||
raise RpcError('BUSY', 'WiFi 通道已经启用')
|
||||
self.close()
|
||||
try:
|
||||
host = str(ipaddress.IPv4Address(self.address()))
|
||||
if ipaddress.ip_address(host).is_unspecified or ipaddress.ip_address(host).is_loopback:
|
||||
raise ValueError()
|
||||
except (ValueError, TypeError):
|
||||
raise RpcError('NOT_READY', '设备 WiFi 尚未取得地址') from None
|
||||
channel = secrets.token_hex(16)
|
||||
secret = secrets.token_bytes(32)
|
||||
self.pending = (channel, secret, self.sessions.clock() + 10)
|
||||
return dict(channel_id=channel, secret=base64.b64encode(secret).decode('ascii'),
|
||||
host=host, port=8080, path='/ws/mobile', expires_in_ms=10000,
|
||||
device_id=self.sessions.identity()['device_id'])
|
||||
|
||||
def attach(self, packet):
|
||||
if not isinstance(packet, bytes) or not 57 <= len(packet) <= MAX_MESSAGE + 32:
|
||||
raise ProtocolError('BAD_REQUEST')
|
||||
if not self.sessions.opened or self.sessions.expired() or not self.pending or self.valid():
|
||||
raise ProtocolError('NOT_READY')
|
||||
channel, secret, expires = self.pending
|
||||
if packet[:32] != channel.encode('ascii') or self.sessions.clock() >= expires:
|
||||
raise ProtocolError('NOT_READY')
|
||||
cipher = wifi_cipher(secret, channel, server=True)
|
||||
hello = json_object(cipher.decrypt(packet[32:]))
|
||||
device_id = self.sessions.identity()['device_id']
|
||||
if hello != dict(device_id=device_id):
|
||||
raise ProtocolError('BAD_REQUEST')
|
||||
self.pending = None
|
||||
self.channel, self.cipher = channel, cipher
|
||||
self.last_activity = self.sessions.clock()
|
||||
return channel, cipher.encrypt(encode(dict(device_id=device_id, ready=True)))
|
||||
|
||||
def accept(self, channel, packet):
|
||||
from .session import RpcError
|
||||
if not self.valid(channel):
|
||||
raise ProtocolError('NOT_READY')
|
||||
request = json_object(self.cipher.decrypt(packet))
|
||||
rid = request.get('id')
|
||||
if not isinstance(rid, str) or not 1 <= len(rid) <= 64 or not rid.isascii() or not rid.isdecimal():
|
||||
raise ProtocolError('BAD_REQUEST')
|
||||
response = dict(id=rid)
|
||||
try:
|
||||
if int(rid) <= self.last_request:
|
||||
raise RpcError('CONFLICT', '请求编号已使用')
|
||||
self.last_request = int(rid)
|
||||
method, params = request.get('method'), request.get('params')
|
||||
if not isinstance(method, str) or not isinstance(params, dict):
|
||||
raise RpcError('BAD_REQUEST', '请求格式无效')
|
||||
if method == 'transport.ping':
|
||||
result = dict(alive=True)
|
||||
elif method.startswith(('session.', 'transport.', 'wifi.')) or method == 'task.get':
|
||||
raise RpcError('BAD_REQUEST', '此操作须使用蓝牙')
|
||||
else:
|
||||
result = self.sessions.dispatch(method, params)
|
||||
response.update(ok=True, result=result)
|
||||
except RpcError as error:
|
||||
response.update(ok=False, error=dict(code=error.code, message=error.message, retryable=False))
|
||||
except Exception:
|
||||
response.update(ok=False, error=dict(code='DEVICE_ERROR', message='设备操作失败,请刷新状态', retryable=False))
|
||||
self.last_activity = self.sessions.clock()
|
||||
record = self.cipher.encrypt(encode(response))
|
||||
self.counters['requests'] += 1
|
||||
self.counters['rx_bytes'] += len(packet)
|
||||
self.counters['tx_bytes'] += len(record)
|
||||
return record
|
||||
|
||||
|
||||
def install_wifi_route(app):
|
||||
from fastapi import WebSocket
|
||||
|
||||
@app.websocket('/ws/mobile')
|
||||
async def mobile_wifi(socket: WebSocket):
|
||||
sessions = getattr(app.state, 'mobile_sessions', None)
|
||||
if sessions is None:
|
||||
await socket.close(code=1008)
|
||||
return
|
||||
channel = None
|
||||
receive = None
|
||||
await socket.accept()
|
||||
def locked(function, *args):
|
||||
with sessions.lock:
|
||||
return function(*args)
|
||||
try:
|
||||
packet = await asyncio.wait_for(socket.receive_bytes(), 5)
|
||||
channel, reply = await asyncio.to_thread(locked, sessions.wifi.attach, packet)
|
||||
await socket.send_bytes(reply)
|
||||
while await asyncio.to_thread(locked, sessions.wifi.valid, channel):
|
||||
if receive is None:
|
||||
receive = asyncio.create_task(socket.receive_bytes())
|
||||
done, _ = await asyncio.wait([receive], timeout=0.25)
|
||||
if not done:
|
||||
continue
|
||||
packet = receive.result()
|
||||
receive = None
|
||||
reply = await asyncio.to_thread(locked, sessions.wifi.accept, channel, packet)
|
||||
# A BLE disconnect while executing a command must not send stale data.
|
||||
if not await asyncio.to_thread(locked, sessions.wifi.valid, channel):
|
||||
break
|
||||
await socket.send_bytes(reply)
|
||||
except Exception:
|
||||
# Never log records, secrets, peer addresses or normal disconnections.
|
||||
pass
|
||||
finally:
|
||||
if receive is not None:
|
||||
receive.cancel()
|
||||
await asyncio.gather(receive, return_exceptions=True)
|
||||
def release():
|
||||
if channel is not None and sessions.wifi.channel == channel:
|
||||
sessions.wifi.close()
|
||||
await asyncio.to_thread(locked, release)
|
||||
try:
|
||||
await socket.close(code=1000)
|
||||
except Exception:
|
||||
pass
|
||||
Reference in New Issue
Block a user