165 lines
7.0 KiB
Python
165 lines
7.0 KiB
Python
"""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
|