240 lines
9.1 KiB
Python
240 lines
9.1 KiB
Python
"""SETProtocol v2 firmware client over segmented classic CAN."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import struct
|
||
from dataclasses import dataclass
|
||
from time import perf_counter
|
||
|
||
from setprotocol.can import CanAddress, CanFrame, CanReassembler, segment
|
||
from setprotocol.core import (
|
||
Capabilities,
|
||
DeviceInfo,
|
||
Feature,
|
||
FirmwareBegin,
|
||
FirmwareFlag,
|
||
Frame as SetFrame,
|
||
FrameFlag,
|
||
MessageType,
|
||
SetProtocolError,
|
||
Status,
|
||
build_frame,
|
||
decode_datagram,
|
||
decode_response,
|
||
encode_firmware_data,
|
||
)
|
||
|
||
from . import transport as tr
|
||
from typing import Any
|
||
|
||
FirmwareImage = Any
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class CanFirmwareTarget:
|
||
"""Address and compatibility settings for a SETProtocol v2 target."""
|
||
|
||
node_id: int
|
||
device_class: int
|
||
hardware_min: int = 0
|
||
hardware_max: int = 0xFF
|
||
|
||
def __post_init__(self) -> None:
|
||
if not 0 <= self.node_id <= 0xFF:
|
||
raise ValueError("SETP node ID вне диапазона 0..255")
|
||
if not 0 <= self.device_class <= 0xFFFF:
|
||
raise ValueError("Device class вне диапазона u16")
|
||
if not 0 <= self.hardware_min <= self.hardware_max <= 0xFFFFFFFF:
|
||
raise ValueError("Диапазон hardware version задан неверно")
|
||
|
||
|
||
class CanFirmwareTransfer:
|
||
"""Stop-and-wait SETP v2 firmware transaction for a classic CAN link."""
|
||
|
||
HOST_NODE = 0
|
||
CHANNEL = 1
|
||
DEFAULT_BLOCK_SIZE = 64
|
||
|
||
def __init__(self, image: FirmwareImage, target: CanFirmwareTarget) -> None:
|
||
if not image.data:
|
||
raise ValueError("Образ прошивки пуст")
|
||
if len(image.data) > 512 * 1024:
|
||
raise ValueError("BALZAM поддерживает образ не более 512 КиБ")
|
||
self.image = image
|
||
self.target = target
|
||
self.offset = 0
|
||
self.block_size = self.DEFAULT_BLOCK_SIZE
|
||
self.stage = "ping"
|
||
self.finished = False
|
||
self._sequence = 0
|
||
self._expected_sequence = 0
|
||
self._expected_type = 0
|
||
self._response: SetFrame | None = None
|
||
self._reassembler = CanReassembler()
|
||
|
||
def _request(self, message_type: int, payload: bytes = b"") -> list[tr.Frame]:
|
||
self._sequence = (self._sequence + 1) & 0xFFFF
|
||
if self._sequence == 0:
|
||
self._sequence = 1
|
||
self._expected_sequence = self._sequence
|
||
self._expected_type = int(message_type)
|
||
packet = build_frame(
|
||
SetFrame(
|
||
message_type=message_type,
|
||
sequence=self._sequence,
|
||
payload=payload,
|
||
flags=FrameFlag.ACK_REQUIRED | FrameFlag.PRIORITY,
|
||
source=self.HOST_NODE,
|
||
destination=self.target.node_id,
|
||
)
|
||
)
|
||
address = CanAddress(
|
||
destination=self.target.node_id,
|
||
source=self.HOST_NODE,
|
||
priority=1,
|
||
channel=self.CHANNEL,
|
||
)
|
||
return [tr.build_frame(item.can_id, item.data, to_can=True)
|
||
for item in segment(packet, address)]
|
||
|
||
def start(self) -> list[tr.Frame]:
|
||
return self._request(MessageType.PING)
|
||
|
||
def abort(self) -> list[tr.Frame]:
|
||
return self._request(MessageType.FW_ABORT)
|
||
|
||
def accepts(self, frame: tr.Frame) -> bool:
|
||
if frame.to_can or not frame.ide:
|
||
return False
|
||
try:
|
||
address = CanAddress.unpack(frame.can_id)
|
||
except SetProtocolError:
|
||
return False
|
||
if address.source != self.target.node_id or address.destination != self.HOST_NODE:
|
||
return False
|
||
try:
|
||
packet = self._reassembler.feed(
|
||
CanFrame(frame.can_id, frame.data), int(perf_counter() * 1000)
|
||
)
|
||
if packet is None:
|
||
return False
|
||
response = decode_datagram(packet)
|
||
except SetProtocolError:
|
||
return False
|
||
if (
|
||
not response.flags & FrameFlag.RESPONSE
|
||
or response.source != self.target.node_id
|
||
or response.destination != self.HOST_NODE
|
||
or response.sequence != self._expected_sequence
|
||
or response.message_type != self._expected_type
|
||
):
|
||
return False
|
||
self._response = response
|
||
return True
|
||
|
||
def _begin(self) -> list[tr.Frame]:
|
||
begin = FirmwareBegin(
|
||
image_size=len(self.image.data),
|
||
image_crc32=self.image.crc32,
|
||
image_version=self.image.version,
|
||
base_address=self.image.base_address,
|
||
slot=0,
|
||
block_size=self.block_size,
|
||
sha256=bytes.fromhex(self.image.sha256),
|
||
flags=FirmwareFlag.RESUME | FirmwareFlag.ERASE_SLOT,
|
||
)
|
||
self.stage = "begin"
|
||
return self._request(MessageType.FW_BEGIN, begin.encode())
|
||
|
||
def _next_data_or_end(self) -> tuple[list[tr.Frame], str]:
|
||
if self.offset >= len(self.image.data):
|
||
self.stage = "end"
|
||
payload = struct.pack(
|
||
"<II32s", len(self.image.data), self.image.crc32,
|
||
bytes.fromhex(self.image.sha256),
|
||
)
|
||
return self._request(MessageType.FW_END, payload), "Проверка CRC32 и SHA-256"
|
||
self.stage = "data"
|
||
data = self.image.data[self.offset : self.offset + self.block_size]
|
||
return (
|
||
self._request(MessageType.FW_DATA, encode_firmware_data(self.offset, data)),
|
||
"Передача блоков SETProtocol v2",
|
||
)
|
||
|
||
def handle_status(self, _frame: tr.Frame) -> tuple[list[tr.Frame], str]:
|
||
response, self._response = self._response, None
|
||
if response is None:
|
||
return [], ""
|
||
status, body = decode_response(response)
|
||
if status != Status.OK:
|
||
try:
|
||
name = Status(status).name
|
||
except ValueError:
|
||
name = "0x%04X" % status
|
||
raise RuntimeError(name)
|
||
|
||
if self.stage == "ping":
|
||
if len(body) != 4:
|
||
raise RuntimeError("PING: неверная длина ответа")
|
||
self.stage = "device_info"
|
||
return self._request(MessageType.DEVICE_INFO), "Чтение информации об устройстве"
|
||
|
||
if self.stage == "device_info":
|
||
info = DeviceInfo.decode(body)
|
||
if self.target.device_class and info.device_class != self.target.device_class:
|
||
raise RuntimeError(
|
||
"device class 0x%04X вместо 0x%04X"
|
||
% (info.device_class, self.target.device_class)
|
||
)
|
||
if not self.target.hardware_min <= info.hardware_version <= self.target.hardware_max:
|
||
raise RuntimeError("hardware version устройства вне разрешённого диапазона")
|
||
self.stage = "capabilities"
|
||
return self._request(MessageType.CAPABILITIES), "Проверка возможностей устройства"
|
||
|
||
if self.stage == "capabilities":
|
||
capabilities = Capabilities.decode(body)
|
||
if not capabilities.features & Feature.FIRMWARE:
|
||
raise RuntimeError("устройство не объявило поддержку firmware update")
|
||
self.block_size = min(
|
||
self.DEFAULT_BLOCK_SIZE,
|
||
capabilities.max_payload - 12,
|
||
)
|
||
if self.block_size <= 0:
|
||
raise RuntimeError("устройство объявило слишком маленький CAN firmware MTU")
|
||
return self._begin(), "Начало SETProtocol v2 firmware session"
|
||
|
||
if self.stage == "begin":
|
||
if len(body) != 4:
|
||
raise RuntimeError("FW_BEGIN: неверный next_offset")
|
||
self.offset = int.from_bytes(body, "little")
|
||
if self.offset > len(self.image.data):
|
||
raise RuntimeError("FW_BEGIN: next_offset за пределами образа")
|
||
return self._next_data_or_end()
|
||
|
||
if self.stage == "data":
|
||
if len(body) != 4:
|
||
raise RuntimeError("FW_DATA: неверный next_offset")
|
||
next_offset = int.from_bytes(body, "little")
|
||
expected = min(self.offset + self.block_size, len(self.image.data))
|
||
if next_offset != expected:
|
||
raise RuntimeError(
|
||
"FW_DATA: подтверждён offset %d вместо %d" % (next_offset, expected)
|
||
)
|
||
self.offset = next_offset
|
||
return self._next_data_or_end()
|
||
|
||
if self.stage == "end":
|
||
if len(body) != 4 or int.from_bytes(body, "little") != len(self.image.data):
|
||
raise RuntimeError("FW_END: устройство не подтвердило полный образ")
|
||
self.stage = "activate"
|
||
return self._request(MessageType.FW_ACTIVATE), "Активация образа"
|
||
|
||
if self.stage == "activate":
|
||
self.finished = True
|
||
return [], "Прошивка по CAN (SETProtocol v2) завершена"
|
||
return [], ""
|
||
|
||
@property
|
||
def percent(self) -> int:
|
||
return min(100, int(self.offset * 100 / len(self.image.data)))
|