# Generated from ActuatorFD specification SHA-256: c8b7736cbd0e5225e8ed78790c3b3191ae08096323ec087c59f0ffc017c6db46
"""Strict little-endian ActuatorFD codecs derived from the wire specification."""

import json
import math
import struct
from dataclasses import dataclass
from typing import cast

CAN_FD_LENGTHS = (0, 1, 2, 3, 4, 5, 6, 7, 8, 12, 16, 20, 24, 32, 48, 64)
SCALARS = {"u8": "B", "u16": "H", "u32": "I", "f32": "f"}
WireValue = int | float | bytes


@dataclass(frozen=True)
class WireField:
    name: str
    type: str
    unit: str = ""

    @property
    def format(self) -> str:
        if self.type in SCALARS:
            return SCALARS[self.type]
        if self.type.startswith("bytes") and self.type[5:].isdigit():
            count = int(self.type[5:])
            if 1 <= count <= 64:
                return f"{count}s"
        raise ValueError(f"Unsupported field type: {self.type}")

    @property
    def size(self) -> int:
        return struct.calcsize("<" + self.format)


@dataclass(frozen=True)
class MessageSchema:
    name: str
    service: int
    priority: int
    fields: tuple[WireField, ...]

    def __post_init__(self) -> None:
        if (
            not self.name.isascii()
            or not self.name.isidentifier()
            or type(self.service) is not int
            or not 0 <= self.service <= 255
        ):
            raise ValueError("Invalid service name or identifier")
        if type(self.priority) is not int or not 0 <= self.priority <= 7:
            raise ValueError("CAN priority must be in [0, 7]")
        if len({field.name for field in self.fields}) != len(self.fields):
            raise ValueError("Duplicate wire field")
        if not isinstance(self.fields, tuple) or any(
            not isinstance(field, WireField) for field in self.fields
        ):
            raise ValueError("Fields must be an immutable WireField tuple")
        if any(not field.name.isascii() or not field.name.isidentifier() for field in self.fields):
            raise ValueError("Invalid wire field name")
        if not 0 < self.payload_size <= 64:
            raise ValueError("CAN-FD message exceeds payload capacity")

    @property
    def payload_size(self) -> int:
        return sum(field.size for field in self.fields)

    @property
    def wire_size(self) -> int:
        return next(length for length in CAN_FD_LENGTHS if length >= self.payload_size)

    def encode(self, values: dict[str, WireValue]) -> bytes:
        if set(values) != {field.name for field in self.fields}:
            raise ValueError(f"{self.name}: missing or unknown wire field")
        packed = bytearray()
        for field in self.fields:
            value = values[field.name]
            if field.type == "f32":
                if isinstance(value, bool) or not isinstance(value, (int, float)):
                    raise ValueError(f"{field.name}: expected finite float")
                if not math.isfinite(value):
                    raise ValueError(f"{field.name}: nonfinite float")
            elif field.type.startswith("bytes"):
                if not isinstance(value, bytes) or len(value) != field.size:
                    raise ValueError(f"{field.name}: expected {field.size} bytes")
            elif isinstance(value, bool) or not isinstance(value, int):
                raise ValueError(f"{field.name}: expected integer")
            try:
                packed.extend(struct.pack("<" + field.format, value))
            except (struct.error, OverflowError) as error:
                raise ValueError(f"{field.name}: outside wire type range") from error
        return bytes(packed).ljust(self.wire_size, b"\0")

    def decode(self, payload: bytes) -> dict[str, WireValue]:
        if len(payload) != self.wire_size or any(payload[self.payload_size :]):
            raise ValueError(f"{self.name}: invalid DLC or nonzero padding")
        output: dict[str, WireValue] = {}
        offset = 0
        for field in self.fields:
            value: WireValue = struct.unpack_from("<" + field.format, payload, offset)[0]
            if isinstance(value, float) and not math.isfinite(value):
                raise ValueError(f"{field.name}: nonfinite float")
            output[field.name] = value
            offset += field.size
        return output


def encode_identifier(schema: MessageSchema, source: int, destination: int, major: int = 1) -> int:
    if any(type(value) is not int for value in (source, destination, major)) or not (
        0 <= source <= 127 and 0 <= destination <= 127 and 1 <= major <= 15
    ):
        raise ValueError("Invalid CAN node ID or protocol major")
    return (
        (schema.priority << 26)
        | (schema.service << 18)
        | (source << 11)
        | (destination << 4)
        | major
    )


def decode_identifier(identifier: int) -> tuple[int, int, int, int, int]:
    if type(identifier) is not int or not 0 <= identifier < 1 << 29:
        raise ValueError("Expected extended 29-bit CAN identifier")
    major = identifier & 15
    if major != 1:
        raise ValueError("Unsupported ActuatorFD major version")
    return (
        identifier >> 26,
        (identifier >> 18) & 255,
        (identifier >> 11) & 127,
        (identifier >> 4) & 127,
        major,
    )


@dataclass(frozen=True, slots=True)
class ProtocolDefinition:
    name: str
    major: int
    minor: int
    arbitration_bps: int
    data_bps: int
    messages: tuple[MessageSchema, ...]

    def __post_init__(self) -> None:
        if self.name != "ActuatorFD" or self.major != 1 or self.minor < 0:
            raise ValueError("Unsupported protocol version")
        if min(self.arbitration_bps, self.data_bps) <= 0:
            raise ValueError("Protocol bit rates must be positive")
        if len({message.name for message in self.messages}) != len(self.messages) or len(
            {message.service for message in self.messages}
        ) != len(self.messages):
            raise ValueError("Duplicate message names or service identifiers")

    def message(self, name: str) -> MessageSchema:
        return next(message for message in self.messages if message.name == name)


def _object(value: object) -> dict[str, object]:
    if not isinstance(value, dict) or any(not isinstance(key, str) for key in value):
        raise ValueError("Expected a JSON object")
    return cast(dict[str, object], value)


def _string(value: object) -> str:
    if not isinstance(value, str):
        raise ValueError("Expected a string")
    return value


def _int(value: object) -> int:
    if type(value) is not int:
        raise ValueError("Expected an integer")
    return value


def _array(value: object) -> list[object]:
    if not isinstance(value, list):
        raise ValueError("Expected a JSON array")
    return cast(list[object], value)


def _pairs(pairs: list[tuple[str, object]]) -> dict[str, object]:
    if len({key for key, _ in pairs}) != len(pairs):
        raise ValueError("Duplicate JSON keys")
    return dict(pairs)


def protocol_from_json(payload: str) -> ProtocolDefinition:
    data = _object(json.loads(payload, object_pairs_hook=_pairs))
    if data.get("endianness") != "little":
        raise ValueError("Unsupported byte order")
    messages = []
    for item in _array(data["messages"]):
        message = _object(item)
        fields = []
        for field_item in _array(message["fields"]):
            field = _object(field_item)
            fields.append(
                WireField(
                    _string(field["name"]), _string(field["type"]), _string(field.get("unit", ""))
                )
            )
        messages.append(
            MessageSchema(
                _string(message["name"]),
                _int(message["service"]),
                _int(message["priority"]),
                tuple(fields),
            )
        )
    return ProtocolDefinition(
        _string(data["protocol"]),
        _int(data["major"]),
        _int(data["minor"]),
        _int(data["arbitration_bps"]),
        _int(data["data_bps"]),
        tuple(messages),
    )

SPEC_SHA256 = 'c8b7736cbd0e5225e8ed78790c3b3191ae08096323ec087c59f0ffc017c6db46'
SCHEMAS = {
    'emergency_disable': MessageSchema('emergency_disable', 0, 0, (WireField('reason', 'u32', ''),)),
    'fault': MessageSchema('fault', 1, 1, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('fault_bits', 'u32', ''), WireField('latched_bits', 'u32', ''), WireField('time_us', 'u32', ''),)),
    'sync': MessageSchema('sync', 2, 2, (WireField('session', 'u32', ''), WireField('time_us', 'u32', ''), WireField('cycle', 'u32', ''),)),
    'command': MessageSchema('command', 3, 3, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('execute_us', 'u32', ''), WireField('valid_until_us', 'u32', ''), WireField('position_rad', 'f32', 'rad'), WireField('velocity_rad_s', 'f32', 'rad/s'), WireField('feedforward_nm', 'f32', 'Nm'), WireField('kp_nm_rad', 'f32', 'Nm/rad'), WireField('kd_nm_s_rad', 'f32', 'Nm s/rad'), WireField('mode', 'u8', ''), WireField('flags', 'u8', ''),)),
    'state': MessageSchema('state', 4, 4, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('time_us', 'u32', ''), WireField('output_position_rad', 'f32', 'rad'), WireField('motor_position_rad', 'f32', 'rad'), WireField('velocity_rad_s', 'f32', 'rad/s'), WireField('iq_a', 'f32', 'A'), WireField('estimated_torque_nm', 'f32', 'Nm'), WireField('temperature_c', 'f32', 'degC'), WireField('fault_bits', 'u32', ''),)),
    'health': MessageSchema('health', 5, 5, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('bus_voltage_v', 'f32', 'V'), WireField('winding_temperature_c', 'f32', 'degC'), WireField('inverter_temperature_c', 'f32', 'degC'), WireField('uptime_ms', 'u32', ''), WireField('rx_errors', 'u32', ''), WireField('tx_errors', 'u32', ''), WireField('mode', 'u8', ''),)),
    'discover': MessageSchema('discover', 16, 6, (WireField('nonce', 'u32', ''), WireField('slot_count', 'u16', ''), WireField('slot_duration_us', 'u32', ''),)),
    'identity': MessageSchema('identity', 17, 6, (WireField('nonce', 'u32', ''), WireField('unique_id', 'bytes16', ''), WireField('hardware_revision', 'u16', ''), WireField('firmware_version', 'u32', ''), WireField('node_id', 'u8', ''),)),
    'assign_address': MessageSchema('assign_address', 18, 6, (WireField('nonce', 'u32', ''), WireField('unique_id', 'bytes16', ''), WireField('node_id', 'u8', ''),)),
    'capabilities': MessageSchema('capabilities', 19, 6, (WireField('protocol_minor', 'u16', ''), WireField('capability_bits', 'u32', ''), WireField('continuous_torque_nm', 'f32', 'Nm'), WireField('peak_torque_nm', 'f32', 'Nm'), WireField('maximum_velocity_rad_s', 'f32', 'rad/s'), WireField('watchdog_us', 'u32', ''), WireField('encoder_status', 'u32', ''),)),
    'session_start': MessageSchema('session_start', 20, 6, (WireField('session', 'u32', ''), WireField('unique_id', 'bytes16', ''), WireField('watchdog_us', 'u32', ''),)),
    'ack': MessageSchema('ack', 21, 6, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('request_service', 'u8', ''), WireField('result', 'u16', ''), WireField('detail', 'u32', ''),)),
    'configure': MessageSchema('configure', 32, 7, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('key', 'u16', ''), WireField('operation', 'u8', ''), WireField('value', 'bytes32', ''),)),
    'calibrate': MessageSchema('calibrate', 33, 7, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('routine', 'u16', ''), WireField('operation', 'u8', ''),)),
    'clear_faults': MessageSchema('clear_faults', 34, 7, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('mask', 'u32', ''),)),
    'firmware_begin': MessageSchema('firmware_begin', 48, 7, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('image_bytes', 'u32', ''), WireField('image_sha256', 'bytes32', ''), WireField('target_slot', 'u8', ''),)),
    'firmware_chunk': MessageSchema('firmware_chunk', 49, 7, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('offset', 'u32', ''), WireField('valid_bytes', 'u8', ''), WireField('data', 'bytes32', ''),)),
    'firmware_commit': MessageSchema('firmware_commit', 50, 7, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('image_sha256', 'bytes32', ''),)),
    'enable': MessageSchema('enable', 35, 3, (WireField('session', 'u32', ''), WireField('sequence', 'u32', ''), WireField('enabled', 'u8', ''),)),
}

def encode(name: str, values: dict[str, WireValue]) -> bytes:
    return SCHEMAS[name].encode(values)

def decode(name: str, payload: bytes) -> dict[str, WireValue]:
    return SCHEMAS[name].decode(payload)
