97 lines
3.5 KiB
Python
97 lines
3.5 KiB
Python
"""BLE 传输层(bleak 封装)。硬件相关,单测以 FakeTransport 替代。"""
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
|
|
from bleak import BleakClient, BleakScanner
|
|
from bleak.exc import BleakError
|
|
|
|
from . import protocol as P
|
|
|
|
|
|
class TransportError(Exception):
|
|
pass
|
|
|
|
|
|
async def scan(timeout: float = 5.0, name_prefix: str = P.DEVICE_NAME_PREFIX):
|
|
"""扫描设备,返回 [{name, address, rssi}]。"""
|
|
found = await BleakScanner.discover(timeout=timeout, return_adv=True)
|
|
out = []
|
|
for addr, (dev, adv) in found.items():
|
|
name = dev.name or adv.local_name or ""
|
|
if name.startswith(name_prefix):
|
|
out.append({"name": name, "address": dev.address, "rssi": adv.rssi})
|
|
return out
|
|
|
|
|
|
class BLETransport:
|
|
"""与 FakeTransport 同接口:write_epd/write_rxtx/request_device_id/delay。"""
|
|
|
|
def __init__(self, address: str | None = None, timeout: float = 10.0,
|
|
name_prefix: str = P.DEVICE_NAME_PREFIX):
|
|
self.address = address
|
|
self.timeout = timeout
|
|
self.name_prefix = name_prefix
|
|
self._client: BleakClient | None = None
|
|
self._id_buf = bytearray()
|
|
self._id_event = asyncio.Event()
|
|
|
|
async def __aenter__(self):
|
|
await self.connect()
|
|
return self
|
|
|
|
async def __aexit__(self, *exc):
|
|
await self.close()
|
|
|
|
async def connect(self):
|
|
if not self.address:
|
|
devices = await scan(timeout=self.timeout, name_prefix=self.name_prefix)
|
|
if not devices:
|
|
raise TransportError(f"未找到名称前缀 {self.name_prefix!r} 的设备")
|
|
self.address = devices[0]["address"]
|
|
self._client = BleakClient(self.address, timeout=self.timeout)
|
|
try:
|
|
await self._client.connect()
|
|
except BleakError as e:
|
|
raise TransportError(f"连接失败: {e}") from e
|
|
await self._client.start_notify(P.RXTX_CHAR_UUID, self._on_notify)
|
|
|
|
async def close(self):
|
|
if self._client and self._client.is_connected:
|
|
try:
|
|
await self._client.stop_notify(P.RXTX_CHAR_UUID)
|
|
finally:
|
|
await self._client.disconnect()
|
|
|
|
def _on_notify(self, _sender, data: bytearray):
|
|
self._id_buf.extend(data)
|
|
if len(self._id_buf) >= 14:
|
|
self._id_event.set()
|
|
|
|
async def write_epd(self, data: bytes, response: bool = True):
|
|
await self._client.write_gatt_char(P.EPD_CHAR_UUID, data, response=response)
|
|
|
|
async def write_rxtx(self, data: bytes, response: bool = True):
|
|
await self._client.write_gatt_char(P.RXTX_CHAR_UUID, data, response=response)
|
|
|
|
async def read_rxtx(self) -> bytes:
|
|
return bytes(await self._client.read_gatt_char(P.RXTX_CHAR_UUID))
|
|
|
|
async def request_device_id(self, timeout: float = 8.0) -> str:
|
|
self._id_buf.clear()
|
|
self._id_event.clear()
|
|
await self.write_rxtx(P.request_id_frame())
|
|
try:
|
|
await asyncio.wait_for(self._id_event.wait(), timeout)
|
|
except asyncio.TimeoutError as e:
|
|
raise TransportError("等待设备 ID 应答超时") from e
|
|
return P.parse_device_id(bytes(self._id_buf[:14]))
|
|
|
|
async def delay(self, seconds: float):
|
|
await asyncio.sleep(seconds)
|
|
|
|
async def run_ota(self, image: bytes, on_progress=None):
|
|
"""SUOTA 固件升级(证据 analysis/apk/apk-analysis.md §e)。"""
|
|
from . import ota
|
|
return await ota.run_ota(self._client, image, on_progress)
|