Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
237 changes: 117 additions & 120 deletions src/infuse_iot/database.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,10 @@

import base64
import binascii
import pathlib

from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import x25519

from infuse_iot.api_client import Client
from infuse_iot.api_client.api.key import get_shared_secret
Expand Down Expand Up @@ -42,67 +46,103 @@ class DeviceState:

def __init__(
self,
address: int,
infuse_id: int,
network_id: int | None = None,
device_id: int | None = None,
device_key_id: int | None = None,
):
self.address = address
self.infuse_id = infuse_id
self.network_id = network_id
self.device_id = device_id
self.device_key_id = device_key_id
self.bt_addr: InterfaceAddress.BluetoothLeAddr | None = None
self.public_key: bytes | None = None
self.device_public_key: bytes | None = None
self.shared_key: bytes | None = None
self.secondary_device_key_id: int | None = None
self.local_shared_key: bytes | None = None
self._tx_gatt_seq = 0

def gatt_sequence_num(self):
"""Persistent auto-incrementing sequence number for GATT"""
self._tx_gatt_seq += 1
return self._tx_gatt_seq

def __init__(self) -> None:
def __init__(self, local_root: pathlib.Path | None) -> None:
self.gateway: int | None = None
self.devices: dict[int, DeviceDatabase.DeviceState] = {}
self.bt_addr: dict[InterfaceAddress.BluetoothLeAddr, int] = {}
self._local_root: x25519.X25519PrivateKey | None = None
self._local_root_public: bytes | None = None
if local_root:
with local_root.open() as f:
private_key = serialization.load_pem_private_key(f.read().encode("utf-8"), password=None)
assert isinstance(private_key, x25519.X25519PrivateKey)
self._local_root = private_key
self._local_root_public = self._local_root.public_key().public_bytes_raw()

@property
def has_local_root(self) -> bool:
return self._local_root is not None

def is_local_root(self, public_key: bytes) -> bool:
return public_key == self._local_root_public

def observe_device(
self,
address: int,
infuse_id: int,
network_id: int | None = None,
device_id: int | None = None,
device_key_id: int | None = None,
bt_addr: InterfaceAddress.BluetoothLeAddr | None = None,
) -> None:
"""Update device state based on observed packet"""
if self.gateway is None:
self.gateway = address
if address not in self.devices:
self.devices[address] = self.DeviceState(address)
self.gateway = infuse_id
if infuse_id not in self.devices:
self.devices[infuse_id] = self.DeviceState(infuse_id)
dev = self.devices[infuse_id]
if network_id is not None:
self.devices[address].network_id = network_id
if device_id is not None:
if self.devices[address].device_id is not None and self.devices[address].device_id != device_id:
raise DeviceKeyChangedError(f"Device key for {address:016x} has changed")
self.devices[address].device_id = device_id
dev.network_id = network_id
if device_key_id is not None:
if device_key_id == dev.secondary_device_key_id:
pass
elif dev.device_key_id is not None and dev.device_key_id != device_key_id:
raise DeviceKeyChangedError(f"Device key for {infuse_id:016x} has changed")
else:
dev.device_key_id = device_key_id
if bt_addr is not None:
self.bt_addr[bt_addr] = address
self.devices[address].bt_addr = bt_addr

def observe_security_state(self, address: int, cloud_key: bytes, device_key: bytes, network_id: int) -> None:
self.bt_addr[bt_addr] = infuse_id
dev.bt_addr = bt_addr

def observe_secondary_remote_public_key(self, infuse_id: int, secondary_pub_key: bytes):
if not self.is_local_root(secondary_pub_key):
return
if infuse_id not in self.devices:
return
dev = self.devices[infuse_id]
assert self._local_root is not None
assert self._local_root_public is not None
assert dev.device_public_key is not None
device_public_key = x25519.X25519PublicKey.from_public_bytes(dev.device_public_key)
dev.secondary_device_key_id = binascii.crc32(self._local_root_public + dev.device_public_key) & 0xFFFFFF
dev.local_shared_key = self._local_root.exchange(device_public_key)

def observe_security_state(
self, infuse_id: int, cloud_pub_key: bytes, device_pub_key: bytes, network_id: int
) -> None:
"""Update device state based on security_state response"""
if address not in self.devices:
self.devices[address] = self.DeviceState(address)
device_id = binascii.crc32(cloud_key + device_key) & 0x00FFFFFF
self.devices[address].device_id = device_id
self.devices[address].network_id = network_id
self.devices[address].public_key = device_key
if infuse_id not in self.devices:
self.devices[infuse_id] = self.DeviceState(infuse_id)
device_key_id = binascii.crc32(cloud_pub_key + device_pub_key) & 0x00FFFFFF
self.devices[infuse_id].device_key_id = device_key_id
self.devices[infuse_id].network_id = network_id
self.devices[infuse_id].device_public_key = device_pub_key

client = Client(base_url="https://api.infuse-iot.com").with_headers({"x-api-key": f"Bearer {get_api_key()}"})

with client as client:
body = Key(base64.b64encode(device_key).decode("utf-8"))
body = Key(base64.b64encode(device_pub_key).decode("utf-8"))
response = get_shared_secret.sync(client=client, body=body)
if response is not None:
key = base64.b64decode(response.key)
self.devices[address].shared_key = key
self.devices[infuse_id].shared_key = key

def _network_key(self, network_id: int, interface: bytes, gps_time: int) -> bytes:
if network_id not in self._network_keys:
Expand All @@ -120,126 +160,83 @@ def _network_key(self, network_id: int, interface: bytes, gps_time: int) -> byte

return self._derived_keys[key_id]

def _serial_key(self, base: bytes, time_idx: int) -> bytes:
return hkdf_derive(base, time_idx.to_bytes(4, "little"), b"serial")

def _bt_adv_key(self, base: bytes, time_idx: int) -> bytes:
return hkdf_derive(base, time_idx.to_bytes(4, "little"), b"bt_adv")

def _bt_gatt_key(self, base: bytes, time_idx: int) -> bytes:
return hkdf_derive(base, time_idx.to_bytes(4, "little"), b"bt_gatt")

def _udp_key(self, base: bytes, time_idx: int) -> bytes:
return hkdf_derive(base, time_idx.to_bytes(4, "little"), b"udp")

def has_public_key(self, address: int) -> bool:
def has_public_key(self, infuse_id: int) -> bool:
"""Does the database have the public key for this device?"""
if address not in self.devices:
if infuse_id not in self.devices:
return False
return self.devices[address].public_key is not None
return self.devices[infuse_id].device_public_key is not None

def has_network_id(self, address: int) -> bool:
def has_network_id(self, infuse_id: int) -> bool:
"""Does the database know the network ID for this device?"""
if address not in self.devices:
if infuse_id not in self.devices:
return False
return self.devices[address].network_id is not None
return self.devices[infuse_id].network_id is not None

def infuse_id_from_bluetooth(self, bt_addr: InterfaceAddress.BluetoothLeAddr) -> int | None:
"""Get Bluetooth address associated with device"""
"""Get Bluetooth infuse_id associated with device"""
return self.bt_addr.get(bt_addr, None)

def serial_network_key(self, address: int, gps_time: int) -> bytes:
"""Network key for serial interface"""
if address not in self.devices:
def _get_network_key(self, infuse_id: int, name: bytes, gps_time: int) -> bytes:
if infuse_id not in self.devices:
raise DeviceUnknownNetworkKey
network_id = self.devices[address].network_id
network_id = self.devices[infuse_id].network_id
if network_id is None:
raise DeviceUnknownNetworkKey
return self._network_key(network_id, name, gps_time)

return self._network_key(network_id, b"serial", gps_time)

def serial_device_key(self, address: int, gps_time: int) -> bytes:
"""Device key for serial interface"""
if address not in self.devices:
def _get_device_key(
self, infuse_id: int, name: bytes, gps_time: int, key_id: int | None = None
) -> tuple[int, bytes]:
if infuse_id not in self.devices:
raise DeviceUnknownDeviceKey
d = self.devices[address]
if d.device_id is None:
d = self.devices[infuse_id]
if key_id is None:
if d.secondary_device_key_id:
key_id = d.secondary_device_key_id
base = d.local_shared_key
else:
key_id = d.device_key_id
base = d.shared_key
elif key_id == d.device_key_id:
base = d.shared_key
elif key_id == d.secondary_device_key_id:
base = d.local_shared_key
else:
raise DeviceUnknownDeviceKey
base = self.devices[address].shared_key
if base is None:
raise DeviceUnknownDeviceKey
assert key_id is not None
time_idx = gps_time // (60 * 60 * 24)
return key_id, hkdf_derive(base, time_idx.to_bytes(4, "little"), name)

def serial_network_key(self, infuse_id: int, gps_time: int) -> bytes:
"""Network key for serial interface"""
return self._get_network_key(infuse_id, b"serial", gps_time)

return self._serial_key(base, time_idx)
def serial_device_key(self, infuse_id: int, gps_time: int, key_id: int | None = None) -> tuple[int, bytes]:
"""Device key for serial interface"""
return self._get_device_key(infuse_id, b"serial", gps_time, key_id)

def bt_adv_network_key(self, address: int, gps_time: int) -> bytes:
def bt_adv_network_key(self, infuse_id: int, gps_time: int) -> bytes:
"""Network key for Bluetooth advertising interface"""
if address not in self.devices:
raise DeviceUnknownNetworkKey
network_id = self.devices[address].network_id
if network_id is None:
raise DeviceUnknownNetworkKey

return self._network_key(network_id, b"bt_adv", gps_time)
return self._get_network_key(infuse_id, b"bt_adv", gps_time)

def bt_adv_device_key(self, address: int, gps_time: int) -> bytes:
def bt_adv_device_key(self, infuse_id: int, gps_time: int, key_id: int | None = None) -> tuple[int, bytes]:
"""Device key for Bluetooth advertising interface"""
if address not in self.devices:
raise DeviceUnknownDeviceKey
d = self.devices[address]
if d.device_id is None:
raise DeviceUnknownDeviceKey
base = self.devices[address].shared_key
if base is None:
raise DeviceUnknownDeviceKey
time_idx = gps_time // (60 * 60 * 24)
return self._get_device_key(infuse_id, b"bt_adv", gps_time, key_id)

return self._bt_adv_key(base, time_idx)

def bt_gatt_network_key(self, address: int, gps_time: int) -> bytes:
def bt_gatt_network_key(self, infuse_id: int, gps_time: int) -> bytes:
"""Network key for Bluetooth advertising interface"""
if address not in self.devices:
raise DeviceUnknownNetworkKey
network_id = self.devices[address].network_id
if network_id is None:
raise DeviceUnknownNetworkKey

return self._network_key(network_id, b"bt_gatt", gps_time)
return self._get_network_key(infuse_id, b"bt_gatt", gps_time)

def bt_gatt_device_key(self, address: int, gps_time: int) -> bytes:
def bt_gatt_device_key(self, infuse_id: int, gps_time: int, key_id: int | None = None) -> tuple[int, bytes]:
"""Device key for Bluetooth advertising interface"""
if address not in self.devices:
raise DeviceUnknownDeviceKey
d = self.devices[address]
if d.device_id is None:
raise DeviceUnknownDeviceKey
base = self.devices[address].shared_key
if base is None:
raise DeviceUnknownDeviceKey
time_idx = gps_time // (60 * 60 * 24)

return self._bt_gatt_key(base, time_idx)
return self._get_device_key(infuse_id, b"bt_gatt", gps_time, key_id)

def udp_network_key(self, address: int, gps_time: int) -> bytes:
def udp_network_key(self, infuse_id: int, gps_time: int) -> bytes:
"""Network key for UDP interface"""
if address not in self.devices:
raise DeviceUnknownNetworkKey
network_id = self.devices[address].network_id
if network_id is None:
raise DeviceUnknownNetworkKey
return self._get_network_key(infuse_id, b"udp", gps_time)

return self._network_key(network_id, b"udp", gps_time)

def udp_device_key(self, address: int, gps_time: int) -> bytes:
def udp_device_key(self, infuse_id: int, gps_time: int, key_id: int | None = None) -> tuple[int, bytes]:
"""Device key for UDP interface"""
if address not in self.devices:
raise DeviceUnknownDeviceKey
d = self.devices[address]
if d.device_id is None:
raise DeviceUnknownDeviceKey
base = self.devices[address].shared_key
if base is None:
raise DeviceUnknownDeviceKey
time_idx = gps_time // (60 * 60 * 24)

return self._udp_key(base, time_idx)
return self._get_device_key(infuse_id, b"udp", gps_time, key_id)
19 changes: 9 additions & 10 deletions src/infuse_iot/epacket/packet.py
Original file line number Diff line number Diff line change
Expand Up @@ -258,8 +258,7 @@ def to_serial(self, database: DeviceDatabase) -> bytes:
key = database.serial_network_key(serial.infuse_id, gps_time)
else:
flags = Flags.ENCR_DEVICE
key_metadata = database.devices[serial.infuse_id].device_id
key = database.serial_device_key(serial.infuse_id, gps_time)
key_metadata, key = database.serial_device_key(serial.infuse_id, gps_time)

# Validation
assert key_metadata is not None
Expand Down Expand Up @@ -420,8 +419,8 @@ def hop_received(self) -> HopReceived:
def decrypt(cls, database: DeviceDatabase, frame: bytes):
header = cls.from_buffer_copy(frame)
if header.flags & Flags.ENCR_DEVICE:
database.observe_device(header.device_id, device_id=header.key_metadata)
key = database.serial_device_key(header.device_id, header.gps_time)
database.observe_device(header.device_id, device_key_id=header.key_metadata)
_, key = database.serial_device_key(header.device_id, header.gps_time, header.key_metadata)
else:
database.observe_device(header.device_id, network_id=header.key_metadata)
key = database.serial_network_key(header.device_id, header.gps_time)
Expand Down Expand Up @@ -460,11 +459,11 @@ def encrypt(
) -> bytes:
dev_state = database.devices[infuse_id]
gps_time = InfuseTime.gps_seconds_from_unix(int(time.time()))
key_meta: int | None
flags = 0

if auth == Auth.DEVICE:
key_meta = dev_state.device_id
key = database.bt_gatt_device_key(infuse_id, gps_time)
key_meta, key = database.bt_gatt_device_key(infuse_id, gps_time)
flags |= Flags.ENCR_DEVICE
else:
key_meta = dev_state.network_id
Expand Down Expand Up @@ -493,8 +492,8 @@ def encrypt(
def decrypt(cls, database: DeviceDatabase, bt_addr: Address.BluetoothLeAddr | None, frame: bytes):
header = cls.from_buffer_copy(frame)
if header.flags & Flags.ENCR_DEVICE:
database.observe_device(header.device_id, device_id=header.key_metadata, bt_addr=bt_addr)
key = database.bt_gatt_device_key(header.device_id, header.gps_time)
database.observe_device(header.device_id, device_key_id=header.key_metadata, bt_addr=bt_addr)
_, key = database.bt_gatt_device_key(header.device_id, header.gps_time, header.key_metadata)
else:
database.observe_device(header.device_id, network_id=header.key_metadata, bt_addr=bt_addr)
key = database.bt_gatt_network_key(header.device_id, header.gps_time)
Expand All @@ -508,8 +507,8 @@ class CtypeUdpFrame(CtypeV0UnversionedFrame):
def decrypt(cls, database: DeviceDatabase, frame: bytes):
header = cls.from_buffer_copy(frame)
if header.flags & Flags.ENCR_DEVICE:
database.observe_device(header.device_id, device_id=header.key_metadata)
key = database.udp_device_key(header.device_id, header.gps_time)
database.observe_device(header.device_id, device_key_id=header.key_metadata)
_, key = database.udp_device_key(header.device_id, header.gps_time, header.key_metadata)
else:
database.observe_device(header.device_id, network_id=header.key_metadata)
key = database.udp_network_key(header.device_id, header.gps_time)
Expand Down
Loading