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
50 changes: 44 additions & 6 deletions packages/api/src/api/controllers/devices.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from core.use_cases.check_device_status import CheckDeviceStatusUseCase
from core.use_cases.execute_component_action import ExecuteComponentActionUseCase
from core.use_cases.get_component_actions import GetComponentActionsUseCase
from core.use_cases.get_local_firmware_releases import GetLocalFirmwareReleases
from core.use_cases.scan_devices import ScanDevicesUseCase
from core.use_cases.update_device_from_local import UpdateDeviceFromLocal
from litestar import Router, get, post
Expand Down Expand Up @@ -274,13 +275,8 @@ async def update_device(
update_device_from_local_interactor = _require(
"update_device_from_local_interactor", update_device_from_local_interactor
)
if channel is not UpdateChannel.STABLE:
raise HTTPException(
status_code=400,
detail="Local updates support the stable channel only",
)
result = await update_device_from_local_interactor.execute(
BaseDeviceRequest(device_ip=ip)
BaseDeviceRequest(device_ip=ip), channel=channel.value
)
else:
execute_component_action_interactor = _require(
Expand All @@ -307,6 +303,47 @@ async def update_device(
}


@get(
"/{ip:str}/firmware-releases",
tags=["Devices"],
summary="Preview Firmware Releases for Local Update",
)
async def get_firmware_releases(
ip: str,
local_firmware_releases_interactor: GetLocalFirmwareReleases | None = None,
) -> dict:
"""
Show what firmware a local (via manager) update would install, per channel.

The versions come from the manager's own index lookup, so they are
accurate for devices without internet access. Channels the index
publishes nothing for are null, as is every channel for a device local
updates cannot serve (Gen1 or no reported app name).

Args:
ip: Device IP address (e.g., "192.168.1.100")

Returns:
dict: Release metadata by channel name, e.g.
{"stable": {"version": ..., "build_id": ...}, "beta": null}
"""
local_firmware_releases_interactor = _require(
"local_firmware_releases_interactor", local_firmware_releases_interactor
)

releases = await local_firmware_releases_interactor.execute(
BaseDeviceRequest(device_ip=ip)
)
return {
channel: (
None
if release is None
else {"version": release.version, "build_id": release.build_id}
)
for channel, release in releases.items()
}


@post("/{ip:str}/reboot", status_code=200, tags=["Devices"], summary="Reboot Device")
async def reboot_device(
ip: str,
Expand Down Expand Up @@ -463,6 +500,7 @@ async def bulk_apply_config(
execute_component_action,
get_device_status,
update_device,
get_firmware_releases,
reboot_device,
execute_bulk_operations,
bulk_export_config,
Expand Down
4 changes: 4 additions & 0 deletions packages/api/src/api/dependencies/container.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,10 @@ def get_dependencies(container: APIContainer) -> dict:
lambda: container.get_update_device_from_local_interactor(),
sync_to_thread=False,
),
"local_firmware_releases_interactor": Provide(
lambda: container.get_local_firmware_releases_interactor(),
sync_to_thread=False,
),
"manage_firmware_use_case": Provide(
lambda: container.get_manage_firmware_interactor(),
sync_to_thread=False,
Expand Down
113 changes: 94 additions & 19 deletions packages/api/tests/unit/controllers/test_devices.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
execute_component_action,
get_component_actions,
get_device_status,
get_firmware_releases,
scan_devices,
update_device,
)
Expand All @@ -20,6 +21,7 @@
from core.use_cases.check_device_status import CheckDeviceStatusUseCase
from core.use_cases.execute_component_action import ExecuteComponentActionUseCase
from core.use_cases.get_component_actions import GetComponentActionsUseCase
from core.use_cases.get_local_firmware_releases import GetLocalFirmwareReleases
from core.use_cases.scan_devices import ScanDevicesUseCase
from core.use_cases.update_device_from_local import UpdateDeviceFromLocal
from litestar.di import Provide
Expand Down Expand Up @@ -188,8 +190,9 @@ class MockUpdateDeviceFromLocal(UpdateDeviceFromLocal):
def __init__(self):
pass

async def execute(self, request):
async def execute(self, request, channel="stable"):
captured["request"] = request
captured["channel"] = channel
return ActionResult(
device_ip=request.device_ip,
success=True,
Expand All @@ -213,27 +216,24 @@ async def execute(self, request):
assert data["source"] == "local"
assert data["channel"] == "stable"
assert captured["request"].device_ip == "192.168.1.100"
assert captured["channel"] == "stable"

def test_update_device_rejects_an_unknown_source(self):
with create_test_client(route_handlers=[update_device]) as client:
response = client.post("/192.168.1.100/update", json={"source": "usb"})

assert response.status_code == 400

def test_update_device_rejects_an_unknown_channel(self):
with create_test_client(
route_handlers=[update_device],
exception_handlers=EXCEPTION_HANDLERS,
) as client:
response = client.post("/192.168.1.100/update", json={"channel": "nightly"})

assert response.status_code == 400
def test_update_device_from_local_source_passes_the_channel(self):
captured = {}

def test_update_device_rejects_a_beta_local_update(self):
class MockUpdateDeviceFromLocal(UpdateDeviceFromLocal):
def __init__(self):
pass

async def execute(self, request, channel="stable"):
captured["channel"] = channel
return ActionResult(
device_ip=request.device_ip,
success=True,
message="Update executed successfully on shelly",
action_type="shelly.Update",
)

with create_test_client(
route_handlers=[update_device],
dependencies={
Expand All @@ -247,6 +247,81 @@ def __init__(self):
json={"source": "local", "channel": "beta"},
)

assert response.status_code == 200
assert response.json()["channel"] == "beta"
assert captured["channel"] == "beta"

def test_get_firmware_releases_reports_each_channel(self):
from core.domain.value_objects.firmware_release import FirmwareRelease

class MockGetLocalFirmwareReleases(GetLocalFirmwareReleases):
def __init__(self):
pass

async def execute(self, request):
return {
"stable": FirmwareRelease(
app_name="Plus2PM",
version="1.8.0",
build_id="20250611-100000/1.8.0-g1234567",
download_url="https://fwcdn.example.test/Plus2PM.zip",
),
"beta": None,
}

with create_test_client(
route_handlers=[get_firmware_releases],
dependencies={
"local_firmware_releases_interactor": Provide(
lambda: MockGetLocalFirmwareReleases(), sync_to_thread=False
)
},
) as client:
response = client.get("/192.168.1.100/firmware-releases")

assert response.status_code == 200
assert response.json() == {
"stable": {
"version": "1.8.0",
"build_id": "20250611-100000/1.8.0-g1234567",
},
"beta": None,
}

def test_get_firmware_releases_maps_an_index_failure_to_422(self):
class MockGetLocalFirmwareReleases(GetLocalFirmwareReleases):
def __init__(self):
pass

async def execute(self, request):
raise FirmwareError("Firmware index request failed for Plus2PM")

with create_test_client(
route_handlers=[get_firmware_releases],
dependencies={
"local_firmware_releases_interactor": Provide(
lambda: MockGetLocalFirmwareReleases(), sync_to_thread=False
)
},
exception_handlers=EXCEPTION_HANDLERS,
) as client:
response = client.get("/192.168.1.100/firmware-releases")

assert response.status_code == 422

def test_update_device_rejects_an_unknown_source(self):
with create_test_client(route_handlers=[update_device]) as client:
response = client.post("/192.168.1.100/update", json={"source": "usb"})

assert response.status_code == 400

def test_update_device_rejects_an_unknown_channel(self):
with create_test_client(
route_handlers=[update_device],
exception_handlers=EXCEPTION_HANDLERS,
) as client:
response = client.post("/192.168.1.100/update", json={"channel": "nightly"})

assert response.status_code == 400

def test_update_device_maps_firmware_misconfiguration_to_500(self):
Expand All @@ -256,7 +331,7 @@ class MockUpdateDeviceFromLocal(UpdateDeviceFromLocal):
def __init__(self):
pass

async def execute(self, request):
async def execute(self, request, channel="stable"):
raise FirmwareConfigurationError(
"Local updates need SHELLY_FIRMWARE_ADVERTISED_BASE_URL set"
)
Expand All @@ -280,8 +355,8 @@ class MockUpdateDeviceFromLocal(UpdateDeviceFromLocal):
def __init__(self):
pass

async def execute(self, request):
raise FirmwareError("No firmware published for app 'Plus2PM'")
async def execute(self, request, channel="stable"):
raise FirmwareError("No stable firmware published for app 'Plus2PM'")

with create_test_client(
route_handlers=[update_device],
Expand Down
11 changes: 11 additions & 0 deletions packages/core/src/core/dependencies/container_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
from core.use_cases.check_device_status import CheckDeviceStatusUseCase
from core.use_cases.execute_component_action import ExecuteComponentActionUseCase
from core.use_cases.get_component_actions import GetComponentActionsUseCase
from core.use_cases.get_local_firmware_releases import GetLocalFirmwareReleases
from core.use_cases.manage_backup_schedules import ManageBackupSchedulesUseCase
from core.use_cases.manage_firmware import ManageFirmware
from core.use_cases.manage_provisioning_profiles import (
Expand Down Expand Up @@ -82,6 +83,7 @@ def __init__(self) -> None:
self._firmware_gateway: ShellyCloudFirmwareGateway | None = None
self._acquire_firmware_interactor: AcquireFirmware | None = None
self._update_device_from_local_interactor: UpdateDeviceFromLocal | None = None
self._local_firmware_releases_interactor: GetLocalFirmwareReleases | None = None
self._manage_firmware_interactor: ManageFirmware | None = None
# Every slot above is a device-scoped cache cleared by
# _reset_device_caches(); slots below survive close().
Expand Down Expand Up @@ -180,6 +182,7 @@ def get_scan_interactor(self) -> ScanDevicesUseCase:
device_gateway=self.get_device_gateway(),
mdns_client=self.get_mdns_client(),
auth_state_cache=self.get_auth_state_cache(),
firmware_gateway=self.get_firmware_gateway(),
)
return self._scan_interactor

Expand Down Expand Up @@ -230,6 +233,14 @@ def get_update_device_from_local_interactor(self) -> UpdateDeviceFromLocal:
)
return self._update_device_from_local_interactor

def get_local_firmware_releases_interactor(self) -> GetLocalFirmwareReleases:
if self._local_firmware_releases_interactor is None:
self._local_firmware_releases_interactor = GetLocalFirmwareReleases(
device_gateway=self.get_device_gateway(),
firmware_gateway=self.get_firmware_gateway(),
)
return self._local_firmware_releases_interactor

def get_manage_firmware_interactor(self) -> ManageFirmware:
if self._manage_firmware_interactor is None:
self._manage_firmware_interactor = ManageFirmware(
Expand Down
3 changes: 3 additions & 0 deletions packages/core/src/core/domain/entities/discovered_device.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@ class DiscoveredDevice(BaseModel):
status: Status = Field(..., description="Current device status")
device_id: str | None = Field(None, description="Unique device identifier")
device_type: str | None = Field(None, description="Device model/type")
app_name: str | None = Field(
None, description="Device app name firmware is looked up by, e.g. 'Plus2PM'"
)
firmware_version: str | None = Field(None, description="Current firmware version")
device_name: str | None = Field(None, description="User-defined device name")
auth_required: bool = Field(
Expand Down
16 changes: 16 additions & 0 deletions packages/core/src/core/domain/value_objects/firmware_release.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,3 +13,19 @@ class FirmwareRelease(BaseModel):
build_id: str = Field(..., description="Full build identifier")
download_url: str = Field(..., description="Direct download URL for the bundle")
channel: str = Field(default="stable", description="Release channel (stable/beta)")

def is_installed_on(self, firmware_version: str | None) -> bool:
"""Whether a device reporting this fw_id already runs exactly this build.

A fw_id and the index's build id name one exact build, so they decide
this whenever the device reports one. Comparing versions instead would
call a device current when the same version has been republished under
a new build, which is how a reissued fix would never install. A device
that reports a bare version has no build to compare, so its version has
to settle it, otherwise it would be reflashed on every run.
"""
if not firmware_version:
return False
if "/" in firmware_version:
return firmware_version == self.build_id
return firmware_version.split("-g", 1)[0] == self.version
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,7 @@ async def discover_device(
status=Status.DETECTED,
device_id=device_data.get("id"),
device_type=device_data.get("model"),
app_name=device_data.get("app"),
device_name=device_data.get("name"),
firmware_version=device_data.get("fw_id"),
auth_required=auth_required,
Expand Down
6 changes: 4 additions & 2 deletions packages/core/src/core/gateways/firmware/firmware.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,10 @@

class FirmwareGateway(ABC):
@abstractmethod
async def get_latest(self, app_name: str) -> FirmwareRelease | None:
"""Latest stable release for an app, or ``None`` when the index has none."""
async def get_latest(
self, app_name: str, channel: str = "stable"
) -> FirmwareRelease | None:
"""Latest release for an app on a channel, or ``None`` when the index has none."""
pass

@abstractmethod
Expand Down
Loading
Loading