diff --git a/README.md b/README.md index ae7b054..7de65b9 100644 --- a/README.md +++ b/README.md @@ -74,6 +74,38 @@ Homelable continuously monitors your nodes and displays their live status (onlin --- +## Zigbee2MQTT Import + +Homelable can connect directly to your MQTT broker and import your Zigbee network topology from **Zigbee2MQTT**, placing each device on the canvas as a typed node. + +### Prerequisites + +- A running **MQTT broker** (e.g. Mosquitto) accessible from the Homelable host +- **Zigbee2MQTT** connected to the broker with at least one device paired + +### Usage + +1. Click **Zigbee Import** in the left sidebar (below "Scan Network") +2. Enter your broker host, port (default `1883`), optional credentials, and base topic (default `zigbee2mqtt`) +3. Click **Test Connection** to verify reachability, then **Fetch Devices** +4. Select the devices you want from the grouped list (Coordinator / Router / End Device) +5. Click **Add N to Canvas** — devices are placed in a grid with IoT edges + +### Node Types + +| Type | Z2M Device | Icon | +|------|-----------|------| +| `zigbee_coordinator` | Coordinator | Network hub | +| `zigbee_router` | Router (mains-powered) | Radio | +| `zigbee_enddevice` | End Device (battery) | Antenna | + +Hierarchy is set automatically: coordinator → routers → end devices (`parent_id`). +LQI (Link Quality Indicator) is stored as a node property. + +> **Full documentation:** [docs/zigbee-import.md](./docs/zigbee-import.md) + +--- + ## Live View (read-only public canvas) Live View lets you share a read-only snapshot of your canvas with anyone on your network — no login required. It is disabled by default. diff --git a/backend/app/api/routes/scan.py b/backend/app/api/routes/scan.py index 5c77bf1..476ec23 100644 --- a/backend/app/api/routes/scan.py +++ b/backend/app/api/routes/scan.py @@ -11,7 +11,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.api.deps import get_current_user from app.core.config import settings from app.db.database import AsyncSessionLocal, get_db -from app.db.models import Node, PendingDevice, ScanRun +from app.db.models import Edge, Node, PendingDevice, PendingDeviceLink, ScanRun from app.schemas.nodes import NodeCreate from app.schemas.scan import PendingDeviceResponse, ScanRunResponse from app.services.scanner import request_cancel, run_scan @@ -126,23 +126,31 @@ async def bulk_approve_devices( for device in devices: device.status = "approved" node = Node( - label=device.hostname or device.ip, + label=device.hostname or device.friendly_name or device.ip or "device", type=device.suggested_type or "generic", ip=device.ip, hostname=device.hostname, status="unknown", services=device.services or [], + ieee_address=device.ieee_address, ) db.add(node) created_nodes.append(node) await db.flush() # populates node.id from Python-side default before reading node_ids = [n.id for n in created_nodes] approved_device_ids = [d.id for d in devices] + + all_edges: list[dict[str, str]] = [] + for device in devices: + all_edges.extend(await _resolve_pending_links_for_ieee(db, device.ieee_address)) + await db.commit() return { "approved": len(node_ids), "node_ids": node_ids, "device_ids": approved_device_ids, + "edges_created": len(all_edges), + "edges": all_edges, "skipped": len(payload.device_ids) - len(node_ids), } @@ -166,6 +174,41 @@ async def bulk_hide_devices( return {"hidden": len(devices), "skipped": len(payload.device_ids) - len(devices)} +@router.post("/pending/{device_id}/restore", response_model=dict) +async def restore_device( + device_id: str, + db: AsyncSession = Depends(get_db), + _: str = Depends(get_current_user), +) -> dict[str, Any]: + device = await db.get(PendingDevice, device_id) + if not device: + raise HTTPException(status_code=404, detail="Device not found") + if device.status != "hidden": + raise HTTPException(status_code=409, detail="Device is not hidden") + device.status = "pending" + await db.commit() + return {"restored": True, "device_id": device_id} + + +@router.post("/pending/bulk-restore", response_model=dict) +async def bulk_restore_devices( + payload: BulkActionRequest, + db: AsyncSession = Depends(get_db), + _: str = Depends(get_current_user), +) -> dict[str, Any]: + result = await db.execute( + select(PendingDevice).where( + PendingDevice.id.in_(payload.device_ids), + PendingDevice.status == "hidden", + ) + ) + devices = result.scalars().all() + for device in devices: + device.status = "pending" + await db.commit() + return {"restored": len(devices), "skipped": len(payload.device_ids) - len(devices)} + + @router.post("/pending/{device_id}/approve", response_model=dict) async def approve_device( device_id: str, @@ -186,12 +229,102 @@ async def approve_device( hostname=node_data.hostname, status=node_data.status, services=node_data.services or [], + ieee_address=device.ieee_address, ) db.add(node) await db.flush() node_id = node.id + + edges = await _resolve_pending_links_for_ieee(db, device.ieee_address) + await db.commit() - return {"approved": True, "node_id": node_id} + return { + "approved": True, + "node_id": node_id, + "edges_created": len(edges), + "edges": edges, + } + + +async def _resolve_pending_links_for_ieee( + db: AsyncSession, ieee: str | None +) -> list[dict[str, str]]: + """Materialize edges for any pending_device_links involving ``ieee``. + + For each link where the other endpoint already exists as a canvas Node + (matched by ``Node.ieee_address``), create the Edge and drop the link + row. Links where the other endpoint is still pending are kept so they + can resolve when that endpoint is approved later. + """ + if not ieee: + return [] + + links_q = await db.execute( + select(PendingDeviceLink).where( + (PendingDeviceLink.source_ieee == ieee) + | (PendingDeviceLink.target_ieee == ieee) + ) + ) + links = list(links_q.scalars().all()) + if not links: + return [] + + # Map every relevant ieee → Node (single query). + other_ieees = { + link.target_ieee if link.source_ieee == ieee else link.source_ieee + for link in links + } + other_ieees.add(ieee) + nodes_q = await db.execute( + select(Node).where(Node.ieee_address.in_(other_ieees)) + ) + by_ieee = {n.ieee_address: n for n in nodes_q.scalars().all() if n.ieee_address} + + self_node = by_ieee.get(ieee) + if self_node is None: + return [] + + # Pre-fetch existing edges between these node ids so we don't create dups + # if the user re-approves a device or had drawn the link manually. + candidate_node_ids = [n.id for n in by_ieee.values()] + existing_q = await db.execute( + select(Edge).where( + Edge.source.in_(candidate_node_ids), + Edge.target.in_(candidate_node_ids), + ) + ) + existing_pairs = {(e.source, e.target) for e in existing_q.scalars().all()} + + created: list[dict[str, str]] = [] + for link in links: + other_ieee = ( + link.target_ieee if link.source_ieee == ieee else link.source_ieee + ) + other_node = by_ieee.get(other_ieee) + if other_node is None: + continue + if link.source_ieee == ieee: + src_id, tgt_id = self_node.id, other_node.id + else: + src_id, tgt_id = other_node.id, self_node.id + # Skip if either direction already exists. + if (src_id, tgt_id) in existing_pairs or (tgt_id, src_id) in existing_pairs: + await db.delete(link) + continue + edge = Edge( + source=src_id, + target=tgt_id, + type="iot", + source_handle="bottom", + target_handle="top-t", + ) + db.add(edge) + await db.flush() + existing_pairs.add((src_id, tgt_id)) + created.append({"id": edge.id, "source": src_id, "target": tgt_id}) + await db.delete(link) + + return created @router.post("/pending/{device_id}/hide") diff --git a/backend/app/api/routes/zigbee.py b/backend/app/api/routes/zigbee.py new file mode 100644 index 0000000..af8c9c9 --- /dev/null +++ b/backend/app/api/routes/zigbee.py @@ -0,0 +1,260 @@ +"""FastAPI router for Zigbee2MQTT import.""" + +import logging +from datetime import datetime, timezone +from typing import Any + +from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException +from sqlalchemy import delete as sa_delete +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.api.deps import get_current_user +from app.db.database import AsyncSessionLocal, get_db +from app.db.models import Node, PendingDevice, PendingDeviceLink, ScanRun +from app.schemas.scan import ScanRunResponse +from app.schemas.zigbee import ( + ZigbeeCoordinatorOut, + ZigbeeEdgeOut, + ZigbeeImportPendingResponse, + ZigbeeImportRequest, + ZigbeeImportResponse, + ZigbeeNodeOut, + ZigbeeTestConnectionRequest, + ZigbeeTestConnectionResponse, +) +from app.services.zigbee_service import fetch_networkmap, test_mqtt_connection + +logger = logging.getLogger(__name__) +router = APIRouter() + + +@router.post("/import", response_model=ZigbeeImportResponse) +async def import_zigbee_network( + payload: ZigbeeImportRequest, + _: str = Depends(get_current_user), +) -> ZigbeeImportResponse: + """Fetch the Zigbee2MQTT network map and return nodes + edges ready for canvas drop. + + Connects to the specified MQTT broker, publishes a networkmap request to + ``/bridge/request/networkmap``, and waits up to 60 s for the + response (large meshes can take 30 s+). The devices are returned as typed homelable nodes with a + coordinator → router → end-device hierarchy. + """ + try: + nodes_raw, edges_raw = await fetch_networkmap( + mqtt_host=payload.mqtt_host, + mqtt_port=payload.mqtt_port, + base_topic=payload.base_topic, + username=payload.mqtt_username, + password=payload.mqtt_password, + tls=payload.mqtt_tls, + tls_insecure=payload.mqtt_tls_insecure, + ) + except ImportError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc + except ConnectionError as exc: + raise HTTPException(status_code=502, detail=str(exc)) from exc + except TimeoutError as exc: + raise HTTPException(status_code=504, detail=str(exc)) from exc + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc + except Exception as exc: + logger.exception("Unexpected error during Zigbee import") + raise HTTPException(status_code=500, detail="Unexpected error during Zigbee import") from exc + + nodes = [ZigbeeNodeOut(**n) for n in nodes_raw] + edges = [ZigbeeEdgeOut(**e) for e in edges_raw] + return ZigbeeImportResponse(nodes=nodes, edges=edges, device_count=len(nodes)) + + +@router.post("/import-pending", response_model=ScanRunResponse) +async def import_zigbee_to_pending( + payload: ZigbeeImportRequest, + background_tasks: BackgroundTasks, + db: AsyncSession = Depends(get_db), + _: str = Depends(get_current_user), +) -> ScanRun: + """Queue a Zigbee2MQTT pending import as a background scan run. + + Returns the ScanRun row immediately so the UI can close the import + modal and surface progress under Scan History (kind=zigbee). The + actual MQTT fetch + pending upsert happens in the background. + """ + run = ScanRun( + status="running", + kind="zigbee", + ranges=[f"{payload.mqtt_host}:{payload.mqtt_port}"], + ) + db.add(run) + await db.commit() + await db.refresh(run) + background_tasks.add_task(_background_zigbee_import, run.id, payload) + return run + + +async def _background_zigbee_import(run_id: str, payload: ZigbeeImportRequest) -> None: + async with AsyncSessionLocal() as db: + try: + nodes_raw, edges_raw = await fetch_networkmap( + mqtt_host=payload.mqtt_host, + mqtt_port=payload.mqtt_port, + base_topic=payload.base_topic, + username=payload.mqtt_username, + password=payload.mqtt_password, + tls=payload.mqtt_tls, + tls_insecure=payload.mqtt_tls_insecure, + ) + result = await _persist_pending_import(db, nodes_raw, edges_raw) + run = await db.get(ScanRun, run_id) + if run: + run.status = "done" + run.devices_found = result.device_count + run.finished_at = datetime.now(timezone.utc) + await db.commit() + except Exception as exc: + logger.exception("Zigbee import %s failed", run_id) + await db.rollback() + run = await db.get(ScanRun, run_id) + if run: + run.status = "error" + run.error = str(exc)[:500] + run.finished_at = datetime.now(timezone.utc) + await db.commit() + + +async def _persist_pending_import( + db: AsyncSession, + nodes_raw: list[dict[str, Any]], + edges_raw: list[dict[str, Any]], +) -> ZigbeeImportPendingResponse: + """Upsert nodes/edges into pending_devices + pending_device_links. + + Coordinator auto-approves to a canvas Node. Other devices upsert by IEEE. + All zigbee-source links are wiped and re-inserted from the new map. + """ + coordinator_out: ZigbeeCoordinatorOut | None = None + coordinator_existed = False + pending_created = 0 + pending_updated = 0 + + for n in nodes_raw: + ieee = n.get("ieee_address") + if not ieee: + continue + if n.get("device_type") == "Coordinator": + existing = await db.execute(select(Node).where(Node.ieee_address == ieee)) + existing_node = existing.scalar_one_or_none() + if existing_node: + coordinator_out = ZigbeeCoordinatorOut( + id=existing_node.id, + label=existing_node.label, + ieee_address=ieee, + ) + coordinator_existed = True + continue + label = n.get("friendly_name") or ieee + node = Node( + label=label, + type=n.get("type") or "zigbee_coordinator", + status="unknown", + ieee_address=ieee, + services=[], + ) + db.add(node) + await db.flush() + coordinator_out = ZigbeeCoordinatorOut( + id=node.id, label=label, ieee_address=ieee + ) + continue + + result = await db.execute( + select(PendingDevice).where(PendingDevice.ieee_address == ieee) + ) + pending = result.scalar_one_or_none() + if pending is None: + db.add( + PendingDevice( + ieee_address=ieee, + friendly_name=n.get("friendly_name"), + hostname=n.get("friendly_name"), + suggested_type=n.get("type"), + device_subtype=n.get("device_type"), + model=n.get("model"), + vendor=n.get("vendor"), + lqi=n.get("lqi"), + status="pending", + discovery_source="zigbee", + ) + ) + pending_created += 1 + else: + pending.friendly_name = n.get("friendly_name") or pending.friendly_name + pending.suggested_type = n.get("type") or pending.suggested_type + pending.device_subtype = n.get("device_type") or pending.device_subtype + pending.model = n.get("model") or pending.model + pending.vendor = n.get("vendor") or pending.vendor + if n.get("lqi") is not None: + pending.lqi = n.get("lqi") + if pending.status == "hidden": + # Re-imported a hidden device → leave it hidden, just refresh fields. + pass + pending_updated += 1 + + # Replace all zigbee-source links with the freshly discovered set. + await db.execute( + sa_delete(PendingDeviceLink).where(PendingDeviceLink.discovery_source == "zigbee") + ) + + links_recorded = 0 + seen: set[tuple[str, str]] = set() + for e in edges_raw: + src = e.get("source") + tgt = e.get("target") + if not src or not tgt or (src, tgt) in seen: + continue + seen.add((src, tgt)) + db.add( + PendingDeviceLink( + source_ieee=src, + target_ieee=tgt, + discovery_source="zigbee", + ) + ) + links_recorded += 1 + + await db.commit() + + return ZigbeeImportPendingResponse( + pending_created=pending_created, + pending_updated=pending_updated, + coordinator=coordinator_out, + coordinator_already_existed=coordinator_existed, + links_recorded=links_recorded, + device_count=len(nodes_raw), + ) + + +@router.post("/test-connection", response_model=ZigbeeTestConnectionResponse) +async def test_zigbee_connection( + payload: ZigbeeTestConnectionRequest, + _: str = Depends(get_current_user), +) -> ZigbeeTestConnectionResponse: + """Quick MQTT ping to validate broker connection before importing.""" + try: + await test_mqtt_connection( + mqtt_host=payload.mqtt_host, + mqtt_port=payload.mqtt_port, + username=payload.mqtt_username, + password=payload.mqtt_password, + tls=payload.mqtt_tls, + tls_insecure=payload.mqtt_tls_insecure, + ) + return ZigbeeTestConnectionResponse(connected=True, message="Connection successful") + except ImportError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc + except (ConnectionError, TimeoutError) as exc: + return ZigbeeTestConnectionResponse(connected=False, message=str(exc)) + except Exception: + logger.exception("Unexpected error during connection test") + return ZigbeeTestConnectionResponse(connected=False, message="Unexpected error") diff --git a/backend/app/db/database.py b/backend/app/db/database.py index 73c1c61..e0eca20 100644 --- a/backend/app/db/database.py +++ b/backend/app/db/database.py @@ -5,13 +5,30 @@ from contextlib import suppress from pathlib import Path from sqlalchemy.exc import OperationalError -from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.orm import DeclarativeBase from app.core.config import APP_VERSION, settings logger = logging.getLogger(__name__) + +async def _try_migrate(conn: AsyncConnection, sql: str, *, label: str) -> None: + """Run an idempotent migration statement, logging any error. + + Distinguishes 'already applied' errors (debug) from genuine failures + (warning) so silent corruption is avoided. Used for new in-commit + migrations; existing legacy ALTERs above remain wrapped in suppress. + """ + try: + await conn.exec_driver_sql(sql) + except OperationalError as exc: + msg = str(exc).lower() + if "duplicate column" in msg or "already exists" in msg: + logger.debug("Migration %s skipped (already applied): %s", label, exc) + else: + logger.warning("Migration %s failed: %s", label, exc) + # Ensure the data directory exists before SQLite tries to open the file Path(settings.sqlite_path).parent.mkdir(parents=True, exist_ok=True) @@ -80,6 +97,77 @@ async def init_db() -> None: await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN bottom_handles INTEGER NOT NULL DEFAULT 1") with suppress(OperationalError): await conn.exec_driver_sql("ALTER TABLE pending_devices ADD COLUMN discovery_source TEXT") + with suppress(OperationalError): + await conn.exec_driver_sql("ALTER TABLE scan_runs ADD COLUMN kind TEXT NOT NULL DEFAULT 'ip'") + # --- Zigbee schema migrations (logged variant per CLAUDE.md feedback) --- + zigbee_migrations: list[tuple[str, str]] = [ + ("nodes.ieee_address", "ALTER TABLE nodes ADD COLUMN ieee_address TEXT"), + ( + "nodes.ieee_address.index", + "CREATE INDEX IF NOT EXISTS ix_nodes_ieee_address ON nodes(ieee_address)", + ), + ("pending_devices.ieee_address", "ALTER TABLE pending_devices ADD COLUMN ieee_address TEXT"), + ( + "pending_devices.ieee_address.index", + "CREATE INDEX IF NOT EXISTS ix_pending_devices_ieee_address " + "ON pending_devices(ieee_address)", + ), + ("pending_devices.friendly_name", "ALTER TABLE pending_devices ADD COLUMN friendly_name TEXT"), + ("pending_devices.device_subtype", "ALTER TABLE pending_devices ADD COLUMN device_subtype TEXT"), + ("pending_devices.model", "ALTER TABLE pending_devices ADD COLUMN model TEXT"), + ("pending_devices.vendor", "ALTER TABLE pending_devices ADD COLUMN vendor TEXT"), + ("pending_devices.lqi", "ALTER TABLE pending_devices ADD COLUMN lqi INTEGER"), + ] + for label, sql in zigbee_migrations: + await _try_migrate(conn, sql, label=label) + # Drop NOT NULL on pending_devices.ip (Zigbee devices have no IP). + # SQLite can't ALTER column nullability — rebuild the table if needed. + try: + info = await conn.exec_driver_sql("PRAGMA table_info(pending_devices)") + cols = info.fetchall() + ip_col = next((c for c in cols if c[1] == "ip"), None) + # PRAGMA table_info row layout: (cid, name, type, notnull, dflt, pk) + if ip_col and ip_col[3] == 1: + logger.info("Migrating pending_devices: dropping NOT NULL on ip column") + await conn.exec_driver_sql("PRAGMA foreign_keys = OFF") + await conn.exec_driver_sql( + "CREATE TABLE pending_devices_new (" + "id VARCHAR PRIMARY KEY," + "ip VARCHAR," + "mac VARCHAR, hostname VARCHAR, os VARCHAR, services JSON," + "suggested_type VARCHAR," + "status VARCHAR," + "discovery_source VARCHAR," + "ieee_address VARCHAR," + "friendly_name VARCHAR," + "device_subtype VARCHAR," + "model VARCHAR," + "vendor VARCHAR," + "lqi INTEGER," + "discovered_at DATETIME" + ")" + ) + await conn.exec_driver_sql( + "INSERT INTO pending_devices_new " + "(id, ip, mac, hostname, os, services, suggested_type, status, " + "discovery_source, ieee_address, friendly_name, device_subtype, " + "model, vendor, lqi, discovered_at) " + "SELECT id, ip, mac, hostname, os, services, suggested_type, status, " + "discovery_source, ieee_address, friendly_name, device_subtype, " + "model, vendor, lqi, discovered_at FROM pending_devices" + ) + await conn.exec_driver_sql("DROP TABLE pending_devices") + await conn.exec_driver_sql( + "ALTER TABLE pending_devices_new RENAME TO pending_devices" + ) + await conn.exec_driver_sql( + "CREATE INDEX IF NOT EXISTS ix_pending_devices_ieee_address " + "ON pending_devices(ieee_address)" + ) + await conn.exec_driver_sql("PRAGMA foreign_keys = ON") + except OperationalError as exc: + logger.warning("pending_devices ip-nullable rebuild failed: %s", exc) + # --- end Zigbee schema migrations ------------------------------------- with suppress(OperationalError): await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN waypoints JSON") with suppress(OperationalError): diff --git a/backend/app/db/models.py b/backend/app/db/models.py index 9eb0c88..203c36f 100644 --- a/backend/app/db/models.py +++ b/backend/app/db/models.py @@ -46,6 +46,7 @@ class Node(Base): width: Mapped[float | None] = mapped_column(Float, nullable=True) height: Mapped[float | None] = mapped_column(Float, nullable=True) bottom_handles: Mapped[int] = mapped_column(Integer, default=1) + ieee_address: Mapped[str | None] = mapped_column(String, index=True, nullable=True) last_seen: Mapped[datetime | None] = mapped_column(DateTime(timezone=True)) response_time_ms: Mapped[int | None] = mapped_column(Integer) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now) @@ -86,7 +87,7 @@ class PendingDevice(Base): __tablename__ = "pending_devices" id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid) - ip: Mapped[str] = mapped_column(String, nullable=False) + ip: Mapped[str | None] = mapped_column(String, nullable=True) mac: Mapped[str | None] = mapped_column(String) hostname: Mapped[str | None] = mapped_column(String) os: Mapped[str | None] = mapped_column(String) @@ -94,6 +95,31 @@ class PendingDevice(Base): suggested_type: Mapped[str | None] = mapped_column(String) status: Mapped[str] = mapped_column(String, default="pending") discovery_source: Mapped[str | None] = mapped_column(String) + ieee_address: Mapped[str | None] = mapped_column(String, index=True, nullable=True, unique=True) + friendly_name: Mapped[str | None] = mapped_column(String, nullable=True) + device_subtype: Mapped[str | None] = mapped_column(String, nullable=True) + model: Mapped[str | None] = mapped_column(String, nullable=True) + vendor: Mapped[str | None] = mapped_column(String, nullable=True) + lqi: Mapped[int | None] = mapped_column(Integer, nullable=True) + discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now) + + +class PendingDeviceLink(Base): + """Link between two Zigbee endpoints discovered during import. + + Endpoints are addressed by IEEE (stable across re-imports). Either side may + already exist as a canvas Node (resolved via Node.ieee_address) or still be + a PendingDevice. On approval, the matching Edge is auto-created when both + endpoints exist as canvas Nodes. + """ + + __tablename__ = "pending_device_links" + + id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid) + source_ieee: Mapped[str] = mapped_column(String, nullable=False, index=True) + target_ieee: Mapped[str] = mapped_column(String, nullable=False, index=True) + lqi: Mapped[int | None] = mapped_column(Integer, nullable=True) + discovery_source: Mapped[str] = mapped_column(String, nullable=False, default="zigbee") discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now) @@ -102,6 +128,7 @@ class ScanRun(Base): id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid) status: Mapped[str] = mapped_column(String, default="running") + kind: Mapped[str] = mapped_column(String, default="ip", server_default="ip") ranges: Mapped[list[str]] = mapped_column(JSON, default=list) devices_found: Mapped[int] = mapped_column(Integer, default=0) started_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now) diff --git a/backend/app/main.py b/backend/app/main.py index 7446e89..b954546 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -7,7 +7,7 @@ from typing import Any from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware -from app.api.routes import auth, canvas, edges, liveview, nodes, scan, status +from app.api.routes import auth, canvas, edges, liveview, nodes, scan, status, zigbee from app.api.routes import settings as settings_routes from app.core.config import settings from app.core.scheduler import start_scheduler, stop_scheduler @@ -55,6 +55,7 @@ app.include_router(scan.router, prefix="/api/v1/scan", tags=["scan"]) app.include_router(status.router, prefix="/api/v1/status", tags=["status"]) app.include_router(settings_routes.router, prefix="/api/v1/settings", tags=["settings"]) app.include_router(liveview.router, prefix="/api/v1/liveview", tags=["liveview"]) +app.include_router(zigbee.router, prefix="/api/v1/zigbee", tags=["zigbee"]) @app.get("/api/v1/health") diff --git a/backend/app/schemas/scan.py b/backend/app/schemas/scan.py index 5b7314e..92151bf 100644 --- a/backend/app/schemas/scan.py +++ b/backend/app/schemas/scan.py @@ -6,7 +6,7 @@ from pydantic import BaseModel class PendingDeviceResponse(BaseModel): id: str - ip: str + ip: str | None mac: str | None hostname: str | None os: str | None @@ -14,6 +14,12 @@ class PendingDeviceResponse(BaseModel): suggested_type: str | None status: str discovery_source: str | None + ieee_address: str | None = None + friendly_name: str | None = None + device_subtype: str | None = None + model: str | None = None + vendor: str | None = None + lqi: int | None = None discovered_at: datetime model_config = {"from_attributes": True} @@ -22,6 +28,7 @@ class PendingDeviceResponse(BaseModel): class ScanRunResponse(BaseModel): id: str status: str + kind: str = "ip" ranges: list[str] devices_found: int started_at: datetime diff --git a/backend/app/schemas/zigbee.py b/backend/app/schemas/zigbee.py new file mode 100644 index 0000000..fe5a270 --- /dev/null +++ b/backend/app/schemas/zigbee.py @@ -0,0 +1,95 @@ +"""Pydantic v2 schemas for Zigbee2MQTT import.""" + +from pydantic import BaseModel, Field, model_validator + + +class ZigbeeImportRequest(BaseModel): + mqtt_host: str = Field(..., description="MQTT broker hostname or IP address") + mqtt_port: int = Field(1883, ge=1, le=65535, description="MQTT broker port") + mqtt_username: str | None = Field(None, description="MQTT username (optional)") + mqtt_password: str | None = Field(None, description="MQTT password (optional)") + base_topic: str = Field("zigbee2mqtt", description="Zigbee2MQTT base topic") + mqtt_tls: bool = Field(False, description="Enable TLS (typically port 8883)") + mqtt_tls_insecure: bool = Field( + False, description="Skip TLS certificate verification (self-signed only)" + ) + + @model_validator(mode="after") + def _insecure_requires_tls(self) -> "ZigbeeImportRequest": + if self.mqtt_tls_insecure and not self.mqtt_tls: + raise ValueError("mqtt_tls_insecure requires mqtt_tls=true") + return self + + +class ZigbeeTestConnectionRequest(BaseModel): + mqtt_host: str + mqtt_port: int = Field(1883, ge=1, le=65535) + mqtt_username: str | None = None + mqtt_password: str | None = None + mqtt_tls: bool = False + mqtt_tls_insecure: bool = False + + @model_validator(mode="after") + def _insecure_requires_tls(self) -> "ZigbeeTestConnectionRequest": + if self.mqtt_tls_insecure and not self.mqtt_tls: + raise ValueError("mqtt_tls_insecure requires mqtt_tls=true") + return self + + +class ZigbeeDeviceData(BaseModel): + ieee_address: str + friendly_name: str + device_type: str # Coordinator, Router, EndDevice + model: str | None = None + vendor: str | None = None + description: str | None = None + lqi: int | None = None + last_seen: str | None = None + + +class ZigbeeNodeOut(BaseModel): + """A homelable-ready node representation of a Zigbee device.""" + + id: str + label: str + type: str # zigbee_coordinator | zigbee_router | zigbee_enddevice + ieee_address: str + friendly_name: str + device_type: str + model: str | None = None + vendor: str | None = None + lqi: int | None = None + parent_id: str | None = None + + +class ZigbeeEdgeOut(BaseModel): + source: str + target: str + + +class ZigbeeImportResponse(BaseModel): + nodes: list[ZigbeeNodeOut] + edges: list[ZigbeeEdgeOut] + device_count: int + + +class ZigbeeTestConnectionResponse(BaseModel): + connected: bool + message: str + + +class ZigbeeCoordinatorOut(BaseModel): + id: str + label: str + ieee_address: str + + +class ZigbeeImportPendingResponse(BaseModel): + """Result of importing a Z2M network into the pending section.""" + + pending_created: int + pending_updated: int + coordinator: ZigbeeCoordinatorOut | None = None + coordinator_already_existed: bool = False + links_recorded: int + device_count: int diff --git a/backend/app/services/zigbee_service.py b/backend/app/services/zigbee_service.py new file mode 100644 index 0000000..4b6222a --- /dev/null +++ b/backend/app/services/zigbee_service.py @@ -0,0 +1,326 @@ +"""Zigbee2MQTT service: connects to MQTT broker and fetches the network map.""" + +from __future__ import annotations + +import asyncio +import json +import logging +import ssl +from typing import Any + +logger = logging.getLogger(__name__) + +try: + import aiomqtt +except ImportError: # pragma: no cover + aiomqtt = None # type: ignore[assignment] + +_NETWORKMAP_REQUEST_TOPIC = "{base_topic}/bridge/request/networkmap" +_NETWORKMAP_RESPONSE_TOPIC = "{base_topic}/bridge/response/networkmap" +_CONNECTION_TIMEOUT = 5.0 # seconds to verify broker reachability +_NETWORKMAP_TIMEOUT = 300.0 # seconds to wait for the networkmap response (large meshes can be slow) + + +def _sanitize_mqtt_error(exc: BaseException) -> str: + """Return a generic, credential-free message for an MQTT error. + + The raw aiomqtt/paho error string can include the broker URI with + embedded credentials (e.g. ``mqtt://user:pass@host``) or auth-related + detail that should not leak to API clients. Map known patterns to + coarse categories; default to a generic failure message. The original + exception is logged at WARNING level for operator debugging. + """ + logger.warning("MQTT error (sanitized for client): %r", exc) + raw = str(exc).lower() + if "not authoriz" in raw or "bad user" in raw or "bad username" in raw: + return "Authentication failed" + if "refused" in raw: + return "Connection refused by broker" + if "name or service not known" in raw or "getaddrinfo" in raw or "nodename nor servname" in raw: + return "Broker hostname could not be resolved" + if "ssl" in raw or "tls" in raw or "certificate" in raw: + return "TLS handshake failed" + if "timed out" in raw or "timeout" in raw: + return "Connection to broker timed out" + return "MQTT connection failed" + + +def _build_tls_context(insecure: bool) -> ssl.SSLContext: + """Build an SSL context for MQTT TLS. If insecure, skip verification.""" + ctx = ssl.create_default_context() + if insecure: + logger.warning( + "MQTT TLS certificate verification is DISABLED — " + "use only with self-signed brokers on trusted networks." + ) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + return ctx + + +def _z2m_type_to_homelable(device_type: str) -> str: + """Map a Z2M device type string to a homelable node type.""" + mapping = { + "Coordinator": "zigbee_coordinator", + "Router": "zigbee_router", + "EndDevice": "zigbee_enddevice", + } + return mapping.get(device_type, "zigbee_enddevice") + + +def _node_from_z2m(raw: dict[str, Any]) -> dict[str, Any] | None: + """Build a homelable node dict from a Z2M raw networkmap node entry.""" + ieee: str = raw.get("ieeeAddr") or raw.get("ieee_address") or "" + if not ieee: + return None + device_type: str = raw.get("type") or "EndDevice" + friendly_name: str = ( + raw.get("friendlyName") or raw.get("friendly_name") or ieee + ) + definition: dict[str, Any] = raw.get("definition") or {} + model: str | None = ( + raw.get("modelID") + or raw.get("model") + or definition.get("model") + or None + ) + vendor: str | None = raw.get("vendor") or definition.get("vendor") or None + return { + "id": ieee, + "label": friendly_name, + "type": _z2m_type_to_homelable(device_type), + "ieee_address": ieee, + "friendly_name": friendly_name, + "device_type": device_type, + "model": model, + "vendor": vendor, + "lqi": None, + "parent_id": None, + } + + +def parse_networkmap( + payload: dict[str, Any], +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + """Parse a Z2M ``bridge/response/networkmap`` payload into node + edge lists. + + Z2M raw response shape:: + + { + "data": { + "type": "raw", + "routes": false, + "value": { + "nodes": [{"ieeeAddr": ..., "type": "Coordinator|Router|EndDevice", + "friendlyName": ..., "definition": {"model": ..., "vendor": ...}}], + "links": [{"source": {"ieeeAddr": ...}, "target": {"ieeeAddr": ...}, + "lqi": 200, "depth": 1}] + } + }, + "status": "ok" + } + + Older or alternate shapes may put nodes/links directly under ``data``. + Both are accepted. + """ + data: dict[str, Any] = payload.get("data") or {} + value = data.get("value") + container: dict[str, Any] = value if isinstance(value, dict) else data + + raw_nodes: list[dict[str, Any]] = container.get("nodes") or [] + raw_links: list[dict[str, Any]] = container.get("links") or [] + + if not isinstance(raw_nodes, list): + raise ValueError("Malformed networkmap: 'nodes' is not a list") + if not isinstance(raw_links, list): + raise ValueError("Malformed networkmap: 'links' is not a list") + + nodes_list: list[dict[str, Any]] = [] + seen_ids: set[str] = set() + coordinator_id: str | None = None + + for entry in raw_nodes: + if not isinstance(entry, dict): + continue + node = _node_from_z2m(entry) + if node is None or node["id"] in seen_ids: + continue + seen_ids.add(node["id"]) + nodes_list.append(node) + if node["device_type"] == "Coordinator": + coordinator_id = node["id"] + + # Z2M `links` is bidirectional/mesh: every pair appears twice and routers + # carry sibling-mesh paths. Walk it only to extract LQI per device and to + # resolve which router an end device hangs off; do NOT emit edges directly + # from links. The final edge set is the strict parent→child tree built + # from parent_id below — that avoids duplicate edges and keeps the visual + # flow consistent (parent bottom → child top). + raw_edges: list[dict[str, Any]] = [] + lqi_by_id: dict[str, int] = {} + + for link in raw_links: + if not isinstance(link, dict): + continue + src_obj = link.get("source") or {} + tgt_obj = link.get("target") or {} + src = src_obj.get("ieeeAddr") if isinstance(src_obj, dict) else None + tgt = tgt_obj.get("ieeeAddr") if isinstance(tgt_obj, dict) else None + if not src or not tgt: + continue + if src not in seen_ids or tgt not in seen_ids: + continue + raw_edges.append({"source": src, "target": tgt}) + lqi = link.get("lqi") or link.get("linkquality") + if isinstance(lqi, int) and tgt not in lqi_by_id: + lqi_by_id[tgt] = lqi + + for node in nodes_list: + if node["id"] in lqi_by_id: + node["lqi"] = lqi_by_id[node["id"]] + + # Build parent_id hierarchy: coordinator → routers → end devices + if coordinator_id: + router_ids = {n["id"] for n in nodes_list if n["device_type"] == "Router"} + for node in nodes_list: + if node["device_type"] == "Router": + node["parent_id"] = coordinator_id + elif node["device_type"] == "EndDevice": + parent = _find_parent_router(node["id"], router_ids, raw_edges) + node["parent_id"] = parent or coordinator_id + + # Final edges = strict parent → child tree (one edge per non-coordinator) + edges_list: list[dict[str, Any]] = [ + {"source": node["parent_id"], "target": node["id"]} + for node in nodes_list + if node.get("parent_id") + ] + + return nodes_list, edges_list + + +def _find_parent_router( + device_id: str, + router_ids: set[str], + edges: list[dict[str, Any]], +) -> str | None: + """Return the first router that has a direct edge to device_id.""" + for edge in edges: + src: str = edge["source"] + tgt: str = edge["target"] + if tgt == device_id and src in router_ids: + return src + if src == device_id and tgt in router_ids: + return tgt + return None + + +async def fetch_networkmap( + mqtt_host: str, + mqtt_port: int, + base_topic: str, + username: str | None = None, + password: str | None = None, + tls: bool = False, + tls_insecure: bool = False, +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + """Connect to the MQTT broker, request the Z2M networkmap, and return (nodes, edges). + + Raises: + TimeoutError: if the broker does not respond in time. + ConnectionError: if the broker cannot be reached. + ValueError: if the response payload is malformed. + """ + if aiomqtt is None: # pragma: no cover + raise ImportError( + "aiomqtt is required for Zigbee import. " + "Install it with: pip install aiomqtt" + ) + + request_topic = _NETWORKMAP_REQUEST_TOPIC.format(base_topic=base_topic) + response_topic = _NETWORKMAP_RESPONSE_TOPIC.format(base_topic=base_topic) + + response_payload: dict[str, Any] = {} + + tls_context = _build_tls_context(tls_insecure) if tls else None + + try: + async with aiomqtt.Client( + hostname=mqtt_host, + port=mqtt_port, + username=username, + password=password, + timeout=_CONNECTION_TIMEOUT, + tls_context=tls_context, + ) as client: + await client.subscribe(response_topic) + # Give the broker a brief window to register the subscription + # before we publish the request. Without this, brokers that + # race SUBACK with our PUBLISH may deliver the response before + # the subscription is active and we'd hang until timeout. + await asyncio.sleep(0.1) + await client.publish( + request_topic, + json.dumps({"type": "raw", "routes": False}), + ) + + async def _wait_for_response() -> None: + async for message in client.messages: + if str(message.topic) != response_topic: + continue + raw = message.payload + try: + payload_str = ( + raw.decode() if isinstance(raw, bytes | bytearray) else str(raw) + ) + response_payload.update(json.loads(payload_str)) + except (json.JSONDecodeError, TypeError) as exc: + raise ValueError( + f"Malformed networkmap response: {exc}" + ) from exc + return + + await asyncio.wait_for(_wait_for_response(), timeout=_NETWORKMAP_TIMEOUT) + + except aiomqtt.MqttError as exc: + raise ConnectionError(_sanitize_mqtt_error(exc)) from exc + except asyncio.TimeoutError as exc: + raise TimeoutError("Timed out waiting for networkmap response") from exc + + if not response_payload: + raise ValueError("Empty networkmap response received") + + return parse_networkmap(response_payload) + + +async def test_mqtt_connection( + mqtt_host: str, + mqtt_port: int, + username: str | None = None, + password: str | None = None, + tls: bool = False, + tls_insecure: bool = False, +) -> bool: + """Attempt a quick MQTT connection to verify broker reachability. + + Returns True on success, raises ConnectionError on failure. + """ + if aiomqtt is None: # pragma: no cover + raise ImportError("aiomqtt is required") + + tls_context = _build_tls_context(tls_insecure) if tls else None + + try: + async with aiomqtt.Client( + hostname=mqtt_host, + port=mqtt_port, + username=username, + password=password, + timeout=_CONNECTION_TIMEOUT, + tls_context=tls_context, + ): + return True + except aiomqtt.MqttError as exc: + raise ConnectionError(_sanitize_mqtt_error(exc)) from exc + except asyncio.TimeoutError as exc: + raise TimeoutError("Connection to broker timed out") from exc diff --git a/backend/requirements.txt b/backend/requirements.txt index 3bd01ad..e504051 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -9,7 +9,7 @@ pydantic-settings==2.5.2 python-jose[cryptography]==3.5.0 passlib[bcrypt]==1.7.4 bcrypt==4.0.1 -python-multipart==0.0.26 +python-multipart==0.0.27 apscheduler==3.10.4 python-nmap==0.7.1 pyyaml==6.0.2 @@ -17,6 +17,7 @@ types-PyYAML==6.0.12.20240917 websockets==13.1 httpx==0.27.2 zeroconf==0.131.0 +aiomqtt==2.3.0 # Dev ruff==0.6.9 diff --git a/backend/tests/test_scan.py b/backend/tests/test_scan.py index 86ce554..3fdf851 100644 --- a/backend/tests/test_scan.py +++ b/backend/tests/test_scan.py @@ -140,6 +140,49 @@ async def test_hide_device(client: AsyncClient, headers, pending_device): assert len(hidden_res.json()) == 1 +# --- Restore hidden device --- + +@pytest.mark.asyncio +async def test_restore_device(client: AsyncClient, headers, pending_device): + # Hide first + await client.post(f"/api/v1/scan/pending/{pending_device.id}/hide", headers=headers) + + # Restore + res = await client.post(f"/api/v1/scan/pending/{pending_device.id}/restore", headers=headers) + assert res.status_code == 200 + assert res.json()["restored"] is True + + # Now back in pending, gone from hidden + pending_res = await client.get("/api/v1/scan/pending", headers=headers) + assert len(pending_res.json()) == 1 + hidden_res = await client.get("/api/v1/scan/hidden", headers=headers) + assert hidden_res.json() == [] + + +@pytest.mark.asyncio +async def test_restore_device_rejects_non_hidden(client: AsyncClient, headers, pending_device): + res = await client.post(f"/api/v1/scan/pending/{pending_device.id}/restore", headers=headers) + assert res.status_code == 409 + + +@pytest.mark.asyncio +async def test_bulk_restore_devices(client: AsyncClient, headers, pending_device): + # Hide + await client.post(f"/api/v1/scan/pending/{pending_device.id}/hide", headers=headers) + + res = await client.post( + "/api/v1/scan/pending/bulk-restore", + headers=headers, + json={"device_ids": [pending_device.id]}, + ) + assert res.status_code == 200 + assert res.json()["restored"] == 1 + assert res.json()["skipped"] == 0 + + pending_res = await client.get("/api/v1/scan/pending", headers=headers) + assert len(pending_res.json()) == 1 + + # --- Ignore device --- @pytest.mark.asyncio @@ -542,3 +585,201 @@ async def test_bulk_hide_requires_auth(client: AsyncClient, two_pending_devices) ids = [d.id for d in two_pending_devices] res = await client.post("/api/v1/scan/pending/bulk-hide", json={"device_ids": ids}) assert res.status_code == 401 + + +# --------------------------------------------------------------------------- +# Approve auto-creates Edges from pending_device_links (Zigbee flow) +# --------------------------------------------------------------------------- + + +async def _seed_zigbee_pending_pair(db_session): + """Create a coordinator Node + a pending device + a link between them.""" + from app.db.models import Node, PendingDevice, PendingDeviceLink + + coord = Node( + label="Coordinator", + type="zigbee_coordinator", + status="unknown", + ieee_address="0xCOORD", + ) + db_session.add(coord) + + pending = PendingDevice( + ieee_address="0xR1", + friendly_name="router_1", + suggested_type="zigbee_router", + device_subtype="Router", + status="pending", + discovery_source="zigbee", + ) + db_session.add(pending) + + db_session.add( + PendingDeviceLink( + source_ieee="0xCOORD", + target_ieee="0xR1", + discovery_source="zigbee", + ) + ) + await db_session.commit() + return coord, pending + + +@pytest.mark.asyncio +async def test_approve_zigbee_creates_edge_when_other_endpoint_is_node( + client: AsyncClient, headers, db_session +): + from sqlalchemy import select + + from app.db.models import Edge + + coord, pending = await _seed_zigbee_pending_pair(db_session) + + res = await client.post( + f"/api/v1/scan/pending/{pending.id}/approve", + json={ + "label": "router_1", + "type": "zigbee_router", + "ip": None, + "status": "unknown", + "services": [], + }, + headers=headers, + ) + assert res.status_code == 200 + data = res.json() + assert data["approved"] is True + assert data["edges_created"] == 1 + + edges = (await db_session.execute(select(Edge))).scalars().all() + assert len(edges) == 1 + assert edges[0].source == coord.id + assert edges[0].target == data["node_id"] + assert edges[0].source_handle == "bottom" + assert edges[0].target_handle == "top-t" + assert edges[0].type == "iot" + + +@pytest.mark.asyncio +async def test_approve_zigbee_skips_duplicate_edge( + client: AsyncClient, headers, db_session +): + """Re-running the resolution does not create a second edge for the same pair.""" + from sqlalchemy import select + + from app.db.models import Edge, PendingDevice, PendingDeviceLink + + coord, pending = await _seed_zigbee_pending_pair(db_session) + body = {"label": "router_1", "type": "zigbee_router", "ip": None, "status": "unknown", "services": []} + await client.post(f"/api/v1/scan/pending/{pending.id}/approve", json=body, headers=headers) + + # Simulate a second pending row + link between same coord and a new device, + # but keep an existing edge in place to verify dedupe also handles + # the swapped-direction case. + new_pending = PendingDevice( + ieee_address="0xR1B", + friendly_name="r1b", + suggested_type="zigbee_router", + status="pending", + discovery_source="zigbee", + ) + db_session.add(new_pending) + db_session.add( + PendingDeviceLink(source_ieee="0xCOORD", target_ieee="0xR1B", discovery_source="zigbee") + ) + await db_session.commit() + res = await client.post( + f"/api/v1/scan/pending/{new_pending.id}/approve", json=body, headers=headers + ) + assert res.json()["edges_created"] == 1 # only the new pair + edges = (await db_session.execute(select(Edge))).scalars().all() + assert len(edges) == 2 # original + new, no duplicate + + +@pytest.mark.asyncio +async def test_approve_zigbee_skips_when_other_endpoint_still_pending( + client: AsyncClient, headers, db_session +): + """Both endpoints pending → no edge yet, link row preserved for later.""" + from sqlalchemy import select + + from app.db.models import Edge, PendingDevice, PendingDeviceLink + + a = PendingDevice( + ieee_address="0xA", + friendly_name="a", + suggested_type="zigbee_router", + status="pending", + discovery_source="zigbee", + ) + b = PendingDevice( + ieee_address="0xB", + friendly_name="b", + suggested_type="zigbee_enddevice", + status="pending", + discovery_source="zigbee", + ) + db_session.add_all([a, b]) + db_session.add( + PendingDeviceLink(source_ieee="0xA", target_ieee="0xB", discovery_source="zigbee") + ) + await db_session.commit() + + res = await client.post( + f"/api/v1/scan/pending/{a.id}/approve", + json={ + "label": "a", + "type": "zigbee_router", + "ip": None, + "status": "unknown", + "services": [], + }, + headers=headers, + ) + assert res.status_code == 200 + assert res.json()["edges_created"] == 0 + + edges = (await db_session.execute(select(Edge))).scalars().all() + assert edges == [] + links = (await db_session.execute(select(PendingDeviceLink))).scalars().all() + assert len(links) == 1 # preserved for later resolution + + +@pytest.mark.asyncio +async def test_approve_zigbee_resolves_link_after_second_approval( + client: AsyncClient, headers, db_session +): + """First approval keeps link; second approval creates the edge.""" + from sqlalchemy import select + + from app.db.models import Edge, PendingDevice, PendingDeviceLink + + a = PendingDevice( + ieee_address="0xA", + friendly_name="a", + suggested_type="zigbee_router", + status="pending", + discovery_source="zigbee", + ) + b = PendingDevice( + ieee_address="0xB", + friendly_name="b", + suggested_type="zigbee_enddevice", + status="pending", + discovery_source="zigbee", + ) + db_session.add_all([a, b]) + db_session.add( + PendingDeviceLink(source_ieee="0xA", target_ieee="0xB", discovery_source="zigbee") + ) + await db_session.commit() + + body = {"label": "x", "type": "zigbee_router", "ip": None, "status": "unknown", "services": []} + await client.post(f"/api/v1/scan/pending/{a.id}/approve", json=body, headers=headers) + res = await client.post(f"/api/v1/scan/pending/{b.id}/approve", json=body, headers=headers) + assert res.json()["edges_created"] == 1 + + edges = (await db_session.execute(select(Edge))).scalars().all() + assert len(edges) == 1 + links = (await db_session.execute(select(PendingDeviceLink))).scalars().all() + assert links == [] # consumed diff --git a/backend/tests/test_zigbee_router.py b/backend/tests/test_zigbee_router.py new file mode 100644 index 0000000..2fc463d --- /dev/null +++ b/backend/tests/test_zigbee_router.py @@ -0,0 +1,417 @@ +"""API endpoint tests for /api/v1/zigbee/*.""" + +from __future__ import annotations + +from unittest.mock import patch + +import pytest +from httpx import AsyncClient + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +async def headers(client: AsyncClient): + res = await client.post("/api/v1/auth/login", json={"username": "admin", "password": "admin"}) + token = res.json()["access_token"] + return {"Authorization": f"Bearer {token}"} + + +# --------------------------------------------------------------------------- +# /api/v1/zigbee/test-connection +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_test_connection_success(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.test_mqtt_connection") as mock_conn: + mock_conn.return_value = True + res = await client.post( + "/api/v1/zigbee/test-connection", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 200 + data = res.json() + assert data["connected"] is True + assert "success" in data["message"].lower() + + +@pytest.mark.asyncio +async def test_test_connection_failure(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.test_mqtt_connection") as mock_conn: + mock_conn.side_effect = ConnectionError("Connection refused") + res = await client.post( + "/api/v1/zigbee/test-connection", + json={"mqtt_host": "bad-host", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 200 + data = res.json() + assert data["connected"] is False + assert "refused" in data["message"].lower() + + +@pytest.mark.asyncio +async def test_test_connection_requires_auth(client: AsyncClient) -> None: + res = await client.post( + "/api/v1/zigbee/test-connection", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + ) + assert res.status_code == 401 + + +@pytest.mark.asyncio +async def test_test_connection_invalid_port(client: AsyncClient, headers: dict) -> None: + res = await client.post( + "/api/v1/zigbee/test-connection", + json={"mqtt_host": "localhost", "mqtt_port": 99999}, + headers=headers, + ) + assert res.status_code == 422 # pydantic validation error + + +# --------------------------------------------------------------------------- +# /api/v1/zigbee/import +# --------------------------------------------------------------------------- + +_SAMPLE_NODES = [ + { + "id": "0x00000000", + "label": "Coordinator", + "type": "zigbee_coordinator", + "ieee_address": "0x00000000", + "friendly_name": "Coordinator", + "device_type": "Coordinator", + "model": None, + "vendor": None, + "lqi": None, + "parent_id": None, + }, + { + "id": "0x00000001", + "label": "router_1", + "type": "zigbee_router", + "ieee_address": "0x00000001", + "friendly_name": "router_1", + "device_type": "Router", + "model": "CC2530", + "vendor": "Texas Instruments", + "lqi": 230, + "parent_id": "0x00000000", + }, +] + +_SAMPLE_EDGES = [ + {"source": "0x00000000", "target": "0x00000001"}, +] + + +@pytest.mark.asyncio +async def test_import_success(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.return_value = (_SAMPLE_NODES, _SAMPLE_EDGES) + res = await client.post( + "/api/v1/zigbee/import", + json={ + "mqtt_host": "localhost", + "mqtt_port": 1883, + "base_topic": "zigbee2mqtt", + }, + headers=headers, + ) + + assert res.status_code == 200 + data = res.json() + assert data["device_count"] == 2 + assert len(data["nodes"]) == 2 + assert len(data["edges"]) == 1 + coordinator = next(n for n in data["nodes"] if n["type"] == "zigbee_coordinator") + assert coordinator["ieee_address"] == "0x00000000" + + +@pytest.mark.asyncio +async def test_import_with_credentials(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.return_value = ([], []) + res = await client.post( + "/api/v1/zigbee/import", + json={ + "mqtt_host": "localhost", + "mqtt_port": 1883, + "mqtt_username": "admin", + "mqtt_password": "secret", + "base_topic": "z2m", + }, + headers=headers, + ) + assert res.status_code == 200 + mock_fetch.assert_called_once_with( + mqtt_host="localhost", + mqtt_port=1883, + base_topic="z2m", + username="admin", + password="secret", + tls=False, + tls_insecure=False, + ) + + +@pytest.mark.asyncio +async def test_import_connection_error_returns_502(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.side_effect = ConnectionError("broker unreachable") + res = await client.post( + "/api/v1/zigbee/import", + json={"mqtt_host": "bad-host", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 502 + assert "broker unreachable" in res.json()["detail"] + + +@pytest.mark.asyncio +async def test_import_timeout_returns_504(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.side_effect = TimeoutError("timed out") + res = await client.post( + "/api/v1/zigbee/import", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 504 + + +@pytest.mark.asyncio +async def test_import_malformed_payload_returns_422(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.side_effect = ValueError("malformed response") + res = await client.post( + "/api/v1/zigbee/import", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 422 + + +@pytest.mark.asyncio +async def test_import_requires_auth(client: AsyncClient) -> None: + res = await client.post( + "/api/v1/zigbee/import", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + ) + assert res.status_code == 401 + + +@pytest.mark.asyncio +async def test_import_empty_network(client: AsyncClient, headers: dict) -> None: + """An empty Zigbee network (coordinator only) is a valid response.""" + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.return_value = ([], []) + res = await client.post( + "/api/v1/zigbee/import", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 200 + data = res.json() + assert data["device_count"] == 0 + assert data["nodes"] == [] + assert data["edges"] == [] + + +@pytest.mark.asyncio +async def test_import_missing_mqtt_host(client: AsyncClient, headers: dict) -> None: + res = await client.post( + "/api/v1/zigbee/import", + json={"mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 422 + + +@pytest.mark.asyncio +async def test_import_with_tls_passes_flags(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch: + mock_fetch.return_value = ([], []) + res = await client.post( + "/api/v1/zigbee/import", + json={ + "mqtt_host": "broker.example.com", + "mqtt_port": 8883, + "mqtt_tls": True, + }, + headers=headers, + ) + assert res.status_code == 200 + kwargs = mock_fetch.call_args.kwargs + assert kwargs["tls"] is True + assert kwargs["tls_insecure"] is False + + +@pytest.mark.asyncio +async def test_import_tls_insecure_requires_tls(client: AsyncClient, headers: dict) -> None: + res = await client.post( + "/api/v1/zigbee/import", + json={ + "mqtt_host": "broker.example.com", + "mqtt_port": 1883, + "mqtt_tls": False, + "mqtt_tls_insecure": True, + }, + headers=headers, + ) + assert res.status_code == 422 + + +# --------------------------------------------------------------------------- +# /api/v1/zigbee/import-pending +# --------------------------------------------------------------------------- + +_PENDING_NODES = [ + { + "id": "0xCOORD", + "label": "Coordinator", + "type": "zigbee_coordinator", + "ieee_address": "0xCOORD", + "friendly_name": "Coordinator", + "device_type": "Coordinator", + "model": None, + "vendor": None, + "lqi": None, + "parent_id": None, + }, + { + "id": "0xR1", + "label": "router_1", + "type": "zigbee_router", + "ieee_address": "0xR1", + "friendly_name": "router_1", + "device_type": "Router", + "model": "CC2530", + "vendor": "TI", + "lqi": 220, + "parent_id": "0xCOORD", + }, + { + "id": "0xE1", + "label": "bulb_kitchen", + "type": "zigbee_enddevice", + "ieee_address": "0xE1", + "friendly_name": "bulb_kitchen", + "device_type": "EndDevice", + "model": "TRADFRI", + "vendor": "IKEA", + "lqi": 180, + "parent_id": "0xR1", + }, +] + +_PENDING_EDGES = [ + {"source": "0xCOORD", "target": "0xR1"}, + {"source": "0xR1", "target": "0xE1"}, +] + + +@pytest.mark.asyncio +async def test_import_pending_endpoint_creates_zigbee_scan_run( + client: AsyncClient, headers: dict +) -> None: + """Endpoint returns a ScanRun (kind=zigbee, status=running) immediately; + the actual networkmap fetch + pending persist runs in the background.""" + from unittest.mock import AsyncMock + + with patch( + "app.api.routes.zigbee._background_zigbee_import", + new_callable=AsyncMock, + ): + res = await client.post( + "/api/v1/zigbee/import-pending", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + headers=headers, + ) + assert res.status_code == 200 + run = res.json() + assert run["kind"] == "zigbee" + assert run["status"] == "running" + assert run["ranges"] == ["localhost:1883"] + + +@pytest.mark.asyncio +async def test_persist_pending_import_creates_coordinator_and_pending( + db_session, +) -> None: + from app.api.routes.zigbee import _persist_pending_import + + result = await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES) + assert result.device_count == 3 + assert result.pending_created == 2 + assert result.pending_updated == 0 + assert result.coordinator is not None + assert result.coordinator.ieee_address == "0xCOORD" + assert result.coordinator_already_existed is False + assert result.links_recorded == 2 + + +@pytest.mark.asyncio +async def test_persist_pending_import_idempotent_updates_existing( + db_session, +) -> None: + from app.api.routes.zigbee import _persist_pending_import + + await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES) + + bumped = [dict(n) for n in _PENDING_NODES] + bumped[1]["lqi"] = 99 + result = await _persist_pending_import(db_session, bumped, _PENDING_EDGES) + + assert result.pending_created == 0 + assert result.pending_updated == 2 + assert result.coordinator_already_existed is True + assert result.links_recorded == 2 + + +@pytest.mark.asyncio +async def test_persist_pending_import_replaces_links(db_session) -> None: + from sqlalchemy import select + + from app.api.routes.zigbee import _persist_pending_import + from app.db.models import PendingDeviceLink + + await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES) + + new_edges = [{"source": "0xCOORD", "target": "0xR1"}] + await _persist_pending_import(db_session, _PENDING_NODES[:2], new_edges) + + rows = (await db_session.execute(select(PendingDeviceLink))).scalars().all() + assert len(rows) == 1 + assert (rows[0].source_ieee, rows[0].target_ieee) == ("0xCOORD", "0xR1") + + +@pytest.mark.asyncio +async def test_import_pending_requires_auth(client: AsyncClient) -> None: + res = await client.post( + "/api/v1/zigbee/import-pending", + json={"mqtt_host": "localhost", "mqtt_port": 1883}, + ) + assert res.status_code == 401 + + +@pytest.mark.asyncio +async def test_test_connection_with_tls(client: AsyncClient, headers: dict) -> None: + with patch("app.api.routes.zigbee.test_mqtt_connection") as mock_conn: + mock_conn.return_value = True + res = await client.post( + "/api/v1/zigbee/test-connection", + json={ + "mqtt_host": "broker.example.com", + "mqtt_port": 8883, + "mqtt_tls": True, + "mqtt_tls_insecure": True, + }, + headers=headers, + ) + assert res.status_code == 200 + kwargs = mock_conn.call_args.kwargs + assert kwargs["tls"] is True + assert kwargs["tls_insecure"] is True diff --git a/backend/tests/test_zigbee_service.py b/backend/tests/test_zigbee_service.py new file mode 100644 index 0000000..39ee260 --- /dev/null +++ b/backend/tests/test_zigbee_service.py @@ -0,0 +1,573 @@ +"""Unit tests for zigbee_service: parser and hierarchy builder.""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import patch + +import aiomqtt # noqa: F401 +import pytest + +from app.services.zigbee_service import ( + _find_parent_router, + _z2m_type_to_homelable, + fetch_networkmap, + parse_networkmap, +) +from app.services.zigbee_service import ( + test_mqtt_connection as _test_mqtt_connection, +) + +# --------------------------------------------------------------------------- +# Helper builders — real Z2M `bridge/response/networkmap` shape +# (data.value.nodes + data.value.links) +# --------------------------------------------------------------------------- + +def _make_node( + ieee: str, + device_type: str = "EndDevice", + friendly_name: str | None = None, + model: str | None = None, + vendor: str | None = None, +) -> dict[str, Any]: + entry: dict[str, Any] = { + "ieeeAddr": ieee, + "type": device_type, + "friendlyName": friendly_name or ieee, + } + if model or vendor: + entry["definition"] = {"model": model, "vendor": vendor} + return entry + + +def _make_link(source_ieee: str, target_ieee: str, lqi: int = 200) -> dict[str, Any]: + return { + "source": {"ieeeAddr": source_ieee}, + "target": {"ieeeAddr": target_ieee}, + "lqi": lqi, + } + + +def _wrap(nodes: list[dict[str, Any]], links: list[dict[str, Any]] | None = None) -> dict[str, Any]: + return { + "data": { + "type": "raw", + "routes": False, + "value": {"nodes": nodes, "links": links or []}, + }, + "status": "ok", + } + + +# --------------------------------------------------------------------------- +# _z2m_type_to_homelable +# --------------------------------------------------------------------------- + +class TestZ2mTypeToHomelable: + def test_coordinator(self) -> None: + assert _z2m_type_to_homelable("Coordinator") == "zigbee_coordinator" + + def test_router(self) -> None: + assert _z2m_type_to_homelable("Router") == "zigbee_router" + + def test_enddevice(self) -> None: + assert _z2m_type_to_homelable("EndDevice") == "zigbee_enddevice" + + def test_unknown_defaults_to_enddevice(self) -> None: + assert _z2m_type_to_homelable("Unknown") == "zigbee_enddevice" + + +# --------------------------------------------------------------------------- +# parse_networkmap +# --------------------------------------------------------------------------- + +class TestParseNetworkmap: + def test_empty_payload(self) -> None: + nodes, edges = parse_networkmap({}) + assert nodes == [] + assert edges == [] + + def test_empty_value(self) -> None: + nodes, edges = parse_networkmap(_wrap([], [])) + assert nodes == [] + assert edges == [] + + def test_coordinator_only(self) -> None: + payload = _wrap([_make_node("0x0000000000000000", "Coordinator", "Coordinator")]) + nodes, edges = parse_networkmap(payload) + assert len(nodes) == 1 + assert nodes[0]["type"] == "zigbee_coordinator" + assert nodes[0]["ieee_address"] == "0x0000000000000000" + assert edges == [] + + def test_coordinator_router_enddevice(self) -> None: + coord_ieee = "0x0000000000000000" + router_ieee = "0x0000000000000001" + end_ieee = "0x0000000000000002" + + payload = _wrap( + nodes=[ + _make_node(coord_ieee, "Coordinator", "Coordinator"), + _make_node(router_ieee, "Router", "my_router"), + _make_node(end_ieee, "EndDevice"), + ], + links=[ + _make_link(coord_ieee, router_ieee), + _make_link(router_ieee, end_ieee), + ], + ) + + nodes, edges = parse_networkmap(payload) + node_by_id = {n["id"]: n for n in nodes} + + assert coord_ieee in node_by_id + assert router_ieee in node_by_id + assert end_ieee in node_by_id + + assert node_by_id[coord_ieee]["type"] == "zigbee_coordinator" + assert node_by_id[router_ieee]["type"] == "zigbee_router" + assert node_by_id[end_ieee]["type"] == "zigbee_enddevice" + + # Parent hierarchy + assert node_by_id[router_ieee]["parent_id"] == coord_ieee + assert node_by_id[end_ieee]["parent_id"] == router_ieee + assert len(edges) == 2 + + def test_no_duplicate_nodes(self) -> None: + ieee = "0x0000000000000001" + payload = _wrap( + nodes=[_make_node(ieee, "Router"), _make_node(ieee, "Router")], + ) + nodes, _ = parse_networkmap(payload) + assert len(nodes) == 1 + + def test_edges_built_correctly(self) -> None: + coord = "0x0000" + router = "0x0001" + payload = _wrap( + nodes=[_make_node(coord, "Coordinator"), _make_node(router, "Router")], + links=[_make_link(coord, router)], + ) + _, edges = parse_networkmap(payload) + assert len(edges) == 1 + assert edges[0]["source"] == coord + assert edges[0]["target"] == router + + def test_friendly_name_used_as_label(self) -> None: + payload = _wrap([_make_node("0xABCD", "EndDevice", "Living Room Sensor")]) + nodes, _ = parse_networkmap(payload) + assert nodes[0]["label"] == "Living Room Sensor" + + def test_enddevice_falls_back_to_coordinator_when_no_router(self) -> None: + coord = "0x0000" + end = "0x0003" + payload = _wrap([_make_node(coord, "Coordinator"), _make_node(end, "EndDevice")]) + nodes, _ = parse_networkmap(payload) + end_node = next(n for n in nodes if n["id"] == end) + assert end_node["parent_id"] == coord + + def test_missing_ieee_skipped(self) -> None: + payload = _wrap([{"type": "EndDevice"}]) # no ieeeAddr + nodes, edges = parse_networkmap(payload) + assert nodes == [] + assert edges == [] + + def test_lqi_propagated_from_link_to_target_node(self) -> None: + coord = "0x0000" + end = "0x0001" + payload = _wrap( + nodes=[_make_node(coord, "Coordinator"), _make_node(end, "EndDevice")], + links=[_make_link(coord, end, lqi=180)], + ) + nodes, _ = parse_networkmap(payload) + end_node = next(n for n in nodes if n["id"] == end) + assert end_node["lqi"] == 180 + + def test_definition_model_and_vendor_extracted(self) -> None: + payload = _wrap([ + _make_node("0xAA", "EndDevice", "Sensor", model="WSDCGQ11LM", vendor="Aqara"), + ]) + nodes, _ = parse_networkmap(payload) + assert nodes[0]["model"] == "WSDCGQ11LM" + assert nodes[0]["vendor"] == "Aqara" + + def test_legacy_shape_without_value_wrapper(self) -> None: + """Some Z2M variants put nodes/links directly under data.""" + payload = {"data": {"nodes": [_make_node("0x01", "Coordinator")], "links": []}} + nodes, _ = parse_networkmap(payload) + assert len(nodes) == 1 + assert nodes[0]["type"] == "zigbee_coordinator" + + def test_routes_bool_is_ignored(self) -> None: + """`routes: false` echo from the request must not crash the parser.""" + payload = {"data": {"routes": False, "type": "raw", "value": {"nodes": [], "links": []}}} + nodes, edges = parse_networkmap(payload) + assert nodes == [] + assert edges == [] + + def test_malformed_nodes_not_list_raises(self) -> None: + with pytest.raises(ValueError, match="not a list"): + parse_networkmap({"data": {"value": {"nodes": "oops", "links": []}}}) + + def test_link_to_unknown_node_dropped(self) -> None: + payload = _wrap( + nodes=[_make_node("0x01", "Coordinator")], + links=[_make_link("0x01", "0xDEAD")], # 0xDEAD not in nodes + ) + _, edges = parse_networkmap(payload) + assert edges == [] + + def test_bidirectional_links_yield_single_edge(self) -> None: + """Z2M links are bidirectional — every pair appears twice. The output + must collapse to a single parent→child edge (no back-link, no dup).""" + coord = "0x0000" + router = "0x0001" + payload = _wrap( + nodes=[_make_node(coord, "Coordinator"), _make_node(router, "Router")], + links=[ + _make_link(coord, router), + _make_link(router, coord), # reverse direction + ], + ) + _, edges = parse_networkmap(payload) + assert edges == [{"source": coord, "target": router}] + + def test_router_mesh_siblings_dropped(self) -> None: + """Router↔router mesh paths in `links` must NOT produce sibling edges + in the final tree. Each router gets exactly one edge from coordinator.""" + coord = "0x0000" + r1 = "0x0001" + r2 = "0x0002" + payload = _wrap( + nodes=[ + _make_node(coord, "Coordinator"), + _make_node(r1, "Router"), + _make_node(r2, "Router"), + ], + links=[ + _make_link(coord, r1), + _make_link(coord, r2), + _make_link(r1, r2), # mesh sibling — must be dropped + _make_link(r2, r1), + ], + ) + _, edges = parse_networkmap(payload) + pairs = {(e["source"], e["target"]) for e in edges} + assert pairs == {(coord, r1), (coord, r2)} + + def test_coordinator_has_no_incoming_edge(self) -> None: + coord = "0x0000" + end = "0x0001" + payload = _wrap( + nodes=[_make_node(coord, "Coordinator"), _make_node(end, "EndDevice")], + links=[_make_link(end, coord)], # back-edge from end to coord + ) + _, edges = parse_networkmap(payload) + # No edge should target the coordinator + assert all(e["target"] != coord for e in edges) + assert edges == [{"source": coord, "target": end}] + + +# --------------------------------------------------------------------------- +# _find_parent_router +# --------------------------------------------------------------------------- + +class TestFindParentRouter: + def test_finds_router_as_source(self) -> None: + router_ids = {"r1"} + edges = [{"source": "r1", "target": "e1"}] + assert _find_parent_router("e1", router_ids, edges) == "r1" + + def test_finds_router_as_target(self) -> None: + router_ids = {"r1"} + edges = [{"source": "e1", "target": "r1"}] + assert _find_parent_router("e1", router_ids, edges) == "r1" + + def test_returns_none_when_no_router(self) -> None: + router_ids: set[str] = set() + edges = [{"source": "e1", "target": "e2"}] + assert _find_parent_router("e1", router_ids, edges) is None + + def test_returns_none_empty_edges(self) -> None: + assert _find_parent_router("e1", {"r1"}, []) is None + + +# --------------------------------------------------------------------------- +# fetch_networkmap (integration-style with mocked aiomqtt) +# --------------------------------------------------------------------------- + +SAMPLE_RESPONSE_PAYLOAD = { + "data": { + "type": "raw", + "routes": False, + "value": { + "nodes": [ + { + "ieeeAddr": "0x00000000", + "type": "Coordinator", + "friendlyName": "Coordinator", + }, + { + "ieeeAddr": "0x00000001", + "type": "Router", + "friendlyName": "router_1", + }, + ], + "links": [ + { + "source": {"ieeeAddr": "0x00000000"}, + "target": {"ieeeAddr": "0x00000001"}, + "lqi": 230, + } + ], + }, + }, + "status": "ok", +} + + +@pytest.mark.asyncio +async def test_fetch_networkmap_success() -> None: + """fetch_networkmap returns parsed nodes/edges when MQTT responds normally.""" + + class _FakeMessage: + topic = "zigbee2mqtt/bridge/response/networkmap" + payload = json.dumps(SAMPLE_RESPONSE_PAYLOAD).encode() + _yielded = False + + def __aiter__(self): + return self + + async def __anext__(self): + if self._yielded: + raise StopAsyncIteration + self._yielded = True + return self + + class _FakeClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_): + pass + + async def subscribe(self, _topic: str) -> None: + pass + + async def publish(self, _topic: str, _payload: str) -> None: + pass + + @property + def messages(self): + return _FakeMessage() + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + nodes, edges = await fetch_networkmap( + mqtt_host="localhost", + mqtt_port=1883, + base_topic="zigbee2mqtt", + ) + + assert any(n["type"] == "zigbee_coordinator" for n in nodes) + assert any(n["type"] == "zigbee_router" for n in nodes) + + +@pytest.mark.asyncio +async def test_fetch_networkmap_connection_error() -> None: + """fetch_networkmap raises ConnectionError when MQTT broker is unreachable.""" + + class _FakeClient: + async def __aenter__(self): + raise Exception("Connection refused") + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + with pytest.raises(ConnectionError): + await fetch_networkmap( + mqtt_host="bad-host", + mqtt_port=1883, + base_topic="zigbee2mqtt", + ) + + +@pytest.mark.asyncio +async def test_test_mqtt_connection_success() -> None: + class _FakeClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + result = await _test_mqtt_connection("localhost", 1883) + assert result is True + + +@pytest.mark.asyncio +async def test_test_mqtt_connection_failure() -> None: + class _FakeClient: + async def __aenter__(self): + raise Exception("refused") + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + with pytest.raises(ConnectionError): + await _test_mqtt_connection("bad-host", 1883) + + +# --------------------------------------------------------------------------- +# TLS context +# --------------------------------------------------------------------------- + +import ssl # noqa: E402 + +from app.services.zigbee_service import _build_tls_context # noqa: E402 + + +def test_build_tls_context_secure_verifies_cert() -> None: + ctx = _build_tls_context(insecure=False) + assert ctx.check_hostname is True + assert ctx.verify_mode == ssl.CERT_REQUIRED + + +def test_build_tls_context_insecure_disables_verification() -> None: + ctx = _build_tls_context(insecure=True) + assert ctx.check_hostname is False + assert ctx.verify_mode == ssl.CERT_NONE + + +@pytest.mark.asyncio +async def test_test_mqtt_connection_passes_tls_context() -> None: + class _FakeClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + await _test_mqtt_connection("host", 8883, tls=True) + kwargs = mock_aiomqtt.Client.call_args.kwargs + assert kwargs["tls_context"] is not None + assert kwargs["tls_context"].verify_mode == ssl.CERT_REQUIRED + + +@pytest.mark.asyncio +async def test_test_mqtt_connection_no_tls_context_when_disabled() -> None: + class _FakeClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + await _test_mqtt_connection("host", 1883, tls=False) + assert mock_aiomqtt.Client.call_args.kwargs["tls_context"] is None + + +@pytest.mark.asyncio +async def test_test_mqtt_connection_insecure_passes_no_verify_context() -> None: + class _FakeClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + await _test_mqtt_connection("host", 8883, tls=True, tls_insecure=True) + ctx = mock_aiomqtt.Client.call_args.kwargs["tls_context"] + assert ctx.verify_mode == ssl.CERT_NONE + assert ctx.check_hostname is False + + +# --------------------------------------------------------------------------- +# Sanitize MQTT errors +# --------------------------------------------------------------------------- + +from app.services.zigbee_service import _sanitize_mqtt_error # noqa: E402 + + +def test_sanitize_auth_error_does_not_leak_credentials() -> None: + msg = _sanitize_mqtt_error( + Exception("Not authorized: bad username or password for user=admin pwd=secret") + ) + assert msg == "Authentication failed" + assert "admin" not in msg + assert "secret" not in msg + + +def test_sanitize_refused() -> None: + assert _sanitize_mqtt_error(Exception("Connection refused")) == "Connection refused by broker" + + +def test_sanitize_dns_failure_strips_host() -> None: + msg = _sanitize_mqtt_error( + Exception("[Errno 8] nodename nor servname provided, or not known: broker.internal.lan") + ) + assert msg == "Broker hostname could not be resolved" + assert "broker.internal.lan" not in msg + + +def test_sanitize_tls_error() -> None: + assert _sanitize_mqtt_error( + Exception("[SSL: CERTIFICATE_VERIFY_FAILED] certificate verify failed") + ) == "TLS handshake failed" + + +def test_sanitize_unknown_falls_back_to_generic() -> None: + msg = _sanitize_mqtt_error(Exception("mqtt://admin:hunter2@broker:1883 weird state")) + assert msg == "MQTT connection failed" + assert "hunter2" not in msg + assert "admin" not in msg + + +@pytest.mark.asyncio +async def test_fetch_networkmap_does_not_leak_creds_in_connection_error() -> None: + class _FakeClient: + async def __aenter__(self): + raise Exception("Not authorized: rejected mqtt://admin:hunter2@host") + + async def __aexit__(self, *_): + pass + + with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt: + mock_aiomqtt.Client.return_value = _FakeClient() + mock_aiomqtt.MqttError = Exception + + with pytest.raises(ConnectionError) as ei: + await fetch_networkmap( + mqtt_host="host", mqtt_port=1883, base_topic="zigbee2mqtt" + ) + msg = str(ei.value) + assert "hunter2" not in msg + assert "admin" not in msg + assert msg == "Authentication failed" diff --git a/docs/zigbee-import.md b/docs/zigbee-import.md new file mode 100644 index 0000000..1de33ee --- /dev/null +++ b/docs/zigbee-import.md @@ -0,0 +1,130 @@ +# Zigbee2MQTT Network Map Importer + +This feature lets you connect Homelable to your MQTT broker, fetch the Zigbee2MQTT network topology, and drop all Zigbee devices onto the canvas as typed nodes with proper hierarchy. + +--- + +## Feature Overview + +- **Automatic device discovery** — Requests the Z2M networkmap via the MQTT bridge API and parses the full device list +- **Typed nodes** — Devices are mapped to three homelable node types: + - `zigbee_coordinator` — The Zigbee coordinator (hub) + - `zigbee_router` — Mains-powered router devices + - `zigbee_enddevice` — Battery-powered end devices (sensors, bulbs, etc.) +- **Hierarchy** — `parent_id` is set automatically: coordinator → routers → end devices +- **LQI display** — Link Quality Indicator is stored as a node property +- **IoT edges** — Links between devices are added as `IoT / Zigbee` edge type + +--- + +## Prerequisites + +1. A running **MQTT broker** (e.g. Mosquitto) accessible from your Homelable host +2. **Zigbee2MQTT** connected to the broker and running +3. Z2M must respond to networkmap requests on: + - **Request topic:** `/bridge/request/networkmap` + - **Response topic:** `/bridge/response/networkmap` + - The default base topic is `zigbee2mqtt` + +--- + +## Step-by-step Usage + +### 1. Open the Zigbee Import dialog + +Click **Zigbee Import** in the left sidebar (below "Scan Network"). + +### 2. Configure the MQTT connection + +| Field | Default | Description | +|---|---|---| +| Broker Host | — | IP or hostname of your MQTT broker | +| Port | 1883 | MQTT broker port | +| Base Topic | `zigbee2mqtt` | Zigbee2MQTT base topic | +| Username | _(optional)_ | MQTT username if authentication is enabled | +| Password | _(optional)_ | MQTT password | + +### 3. Test the connection (optional) + +Click **Test Connection** to verify broker reachability before fetching devices. +A green indicator confirms success; red shows the error message from the broker. + +### 4. Fetch devices + +Click **Fetch Devices**. Homelable will: +1. Connect to the broker +2. Subscribe to the response topic +3. Publish `{"type": "raw", "routes": false}` to the request topic +4. Wait up to 60 seconds for the network map response (large meshes can take 30 s+) +5. Parse and group devices by type + +### 5. Select and add to canvas + +Devices are grouped by type (Coordinator / Router / End Device). +Use the checkboxes to select which devices to add, then click **Add N to Canvas**. + +> **Tip:** All devices are selected by default. Uncheck any you don't want. + +### 6. Arrange on the canvas + +Devices are placed in a grid at the top-right of the canvas. +Use **Auto Layout** (toolbar) to re-arrange the full canvas, or drag nodes manually. + +--- + +## MQTT Configuration Tips + +### Mosquitto without authentication + +``` +listener 1883 +allow_anonymous true +``` + +### Mosquitto with password file + +``` +listener 1883 +password_file /etc/mosquitto/passwd +``` + +Create a user: +```bash +mosquitto_passwd -c /etc/mosquitto/passwd +``` + +### Zigbee2MQTT `configuration.yaml` + +```yaml +mqtt: + base_topic: zigbee2mqtt + server: mqtt://localhost:1883 + # user: mqtt_user + # password: mqtt_password +``` + +--- + +## Supported Z2M Versions + +The networkmap bridge API is available in **Zigbee2MQTT 1.x and 2.x**. +Tested against Z2M 1.35+ and 2.x. + +The importer uses the `raw` topology format (`routes: false`) which is the most widely supported mode. + +--- + +## Troubleshooting + +| Symptom | Cause | Fix | +|---|---|---| +| "Connection refused" | Broker unreachable | Check host/port, firewall rules | +| "Timed out waiting for networkmap" | Z2M not running or wrong base_topic | Verify Z2M is connected, check base_topic setting | +| 0 devices returned | Z2M has no devices paired | Pair at least one device first | +| "Malformed networkmap response" | Z2M returned unexpected format | Check Z2M version; open an issue | + +--- + +## Screenshots + +_(Screenshots will be added in a future release)_ diff --git a/frontend/package-lock.json b/frontend/package-lock.json index 00ebfd7..bafecf0 100644 --- a/frontend/package-lock.json +++ b/frontend/package-lock.json @@ -1,12 +1,12 @@ { "name": "frontend", - "version": "1.10.2", + "version": "1.13.0", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "frontend", - "version": "1.10.2", + "version": "1.13.0", "dependencies": { "@base-ui/react": "^1.2.0", "@dagrejs/dagre": "^2.0.4", @@ -16,7 +16,7 @@ "@radix-ui/react-tooltip": "^1.2.8", "@types/js-yaml": "^4.0.9", "@xyflow/react": "^12.10.1", - "axios": "^1.13.6", + "axios": "^1.15.2", "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "dagre": "^0.8.5", @@ -1646,9 +1646,9 @@ } }, "node_modules/@hono/node-server": { - "version": "1.19.12", - "resolved": "https://registry.npmjs.org/@hono/node-server/-/node-server-1.19.12.tgz", - "integrity": "sha512-txsUW4SQ1iilgE0l9/e9VQWmELXifEFvmdA1j6WFh/aFPj99hIntrSsq/if0UWyGVkmrRPKA1wCeP+UCr1B9Uw==", + "version": "1.19.14", + "resolved": "https://registry.npmjs.org/@hono/node-server/-/node-server-1.19.14.tgz", + "integrity": "sha512-GwtvgtXxnWsucXvbQXkRgqksiH2Qed37H9xHZocE5sA3N8O8O8/8FA3uclQXxXVzc9XBZuEOMK7+r02FmSpHtw==", "license": "MIT", "engines": { "node": ">=18.14.1" @@ -2520,9 +2520,6 @@ "arm" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2537,9 +2534,6 @@ "arm" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -2554,9 +2548,6 @@ "arm64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2571,9 +2562,6 @@ "arm64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -2588,9 +2576,6 @@ "loong64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2605,9 +2590,6 @@ "loong64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -2622,9 +2604,6 @@ "ppc64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2639,9 +2618,6 @@ "ppc64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -2656,9 +2632,6 @@ "riscv64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2673,9 +2646,6 @@ "riscv64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -2690,9 +2660,6 @@ "s390x" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2707,9 +2674,6 @@ "x64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2724,9 +2688,6 @@ "x64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -2975,9 +2936,6 @@ "arm64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -2995,9 +2953,6 @@ "arm64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -3015,9 +2970,6 @@ "x64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MIT", "optional": true, "os": [ @@ -3035,9 +2987,6 @@ "x64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MIT", "optional": true, "os": [ @@ -4186,12 +4135,12 @@ "license": "MIT" }, "node_modules/axios": { - "version": "1.14.0", - "resolved": "https://registry.npmjs.org/axios/-/axios-1.14.0.tgz", - "integrity": "sha512-3Y8yrqLSwjuzpXuZ0oIYZ/XGgLwUIBU3uLvbcpb0pidD9ctpShJd43KSlEEkVQg6DS0G9NKyzOvBfUtDKEyHvQ==", + "version": "1.16.0", + "resolved": "https://registry.npmjs.org/axios/-/axios-1.16.0.tgz", + "integrity": "sha512-6hp5CwvTPlN2A31g5dxnwAX0orzM7pmCRDLnZSX772mv8WDqICwFjowHuPs04Mc8deIld1+ejhtaMn5vp6b+1w==", "license": "MIT", "dependencies": { - "follow-redirects": "^1.15.11", + "follow-redirects": "^1.16.0", "form-data": "^4.0.5", "proxy-from-env": "^2.1.0" } @@ -5602,12 +5551,12 @@ } }, "node_modules/express-rate-limit": { - "version": "8.3.2", - "resolved": "https://registry.npmjs.org/express-rate-limit/-/express-rate-limit-8.3.2.tgz", - "integrity": "sha512-77VmFeJkO0/rvimEDuUC5H30oqUC4EyOhyGccfqoLebB0oiEYfM7nwPrsDsBL1gsTpwfzX8SFy2MT3TDyRq+bg==", + "version": "8.5.1", + "resolved": "https://registry.npmjs.org/express-rate-limit/-/express-rate-limit-8.5.1.tgz", + "integrity": "sha512-5O6KYmyJEpuPJV5hNTXKbAHWRqrzyu+OI3vUnSd2kXFubIVpG7ezpgxQy76Zo5GQZtrQBg86hF+CM/NX+cioiQ==", "license": "MIT", "dependencies": { - "ip-address": "10.1.0" + "ip-address": "^10.2.0" }, "engines": { "node": ">= 16" @@ -5693,9 +5642,9 @@ "license": "MIT" }, "node_modules/fast-uri": { - "version": "3.1.0", - "resolved": "https://registry.npmjs.org/fast-uri/-/fast-uri-3.1.0.tgz", - "integrity": "sha512-iPeeDKJSWf4IEOasVVrknXpaBV0IApz/gp7S2bb7Z4Lljbl2MGJRqInZiUrQwV16cpzw/D3S5j5Julj/gT52AA==", + "version": "3.1.2", + "resolved": "https://registry.npmjs.org/fast-uri/-/fast-uri-3.1.2.tgz", + "integrity": "sha512-rVjf7ArG3LTk+FS6Yw81V1DLuZl1bRbNrev6Tmd/9RaroeeRRJhAt7jg/6YFxbvAQXUCavSoZhPPj6oOx+5KjQ==", "funding": [ { "type": "github", @@ -5857,9 +5806,9 @@ "license": "ISC" }, "node_modules/follow-redirects": { - "version": "1.15.11", - "resolved": "https://registry.npmjs.org/follow-redirects/-/follow-redirects-1.15.11.tgz", - "integrity": "sha512-deG2P0JfjrTxl50XGCDyfI97ZGVCxIpfKYmfyrQ54n5FO/0gfIES8C/Psl6kWVDolizcaaxZJnTS0QSMxvnsBQ==", + "version": "1.16.0", + "resolved": "https://registry.npmjs.org/follow-redirects/-/follow-redirects-1.16.0.tgz", + "integrity": "sha512-y5rN/uOsadFT/JfYwhxRS5R7Qce+g3zG97+JrtFZlC9klX/W5hD7iiLzScI4nZqUS7DNUdhPgw4xI8W2LuXlUw==", "funding": [ { "type": "individual", @@ -6196,9 +6145,9 @@ } }, "node_modules/hono": { - "version": "4.12.11", - "resolved": "https://registry.npmjs.org/hono/-/hono-4.12.11.tgz", - "integrity": "sha512-r4xbIa3mGGGoH9nN4A14DOg2wx7y2oQyJEb5O57C/xzETG/qx4c7CVDQ5WMeKHZ7ORk2W0hZ/sQKXTav3cmYBA==", + "version": "4.12.18", + "resolved": "https://registry.npmjs.org/hono/-/hono-4.12.18.tgz", + "integrity": "sha512-RWzP96k/yv0PQfyXnWjs6zot20TqfpfsNXhOnev8d1InAxubW93L11/oNUc3tQqn2G0bSdAOBpX+2uDFHV7kdQ==", "license": "MIT", "engines": { "node": ">=16.9.0" @@ -6354,9 +6303,9 @@ "license": "ISC" }, "node_modules/ip-address": { - "version": "10.1.0", - "resolved": "https://registry.npmjs.org/ip-address/-/ip-address-10.1.0.tgz", - "integrity": "sha512-XXADHxXmvT9+CRxhXg56LJovE+bmWnEWB78LB83VZTprKTmaC5QfruXocxzTZ2Kl0DNwKuBdlIhjL8LeY8Sf8Q==", + "version": "10.2.0", + "resolved": "https://registry.npmjs.org/ip-address/-/ip-address-10.2.0.tgz", + "integrity": "sha512-/+S6j4E9AHvW9SWMSEY9Xfy66O5PWvVEJ08O0y5JGyEKQpojb0K0GKpz/v5HJ/G0vi3D2sjGK78119oXZeE0qA==", "license": "MIT", "engines": { "node": ">= 12" @@ -6935,9 +6884,6 @@ "arm64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MPL-2.0", "optional": true, "os": [ @@ -6959,9 +6905,6 @@ "arm64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MPL-2.0", "optional": true, "os": [ @@ -6983,9 +6926,6 @@ "x64" ], "dev": true, - "libc": [ - "glibc" - ], "license": "MPL-2.0", "optional": true, "os": [ @@ -7007,9 +6947,6 @@ "x64" ], "dev": true, - "libc": [ - "musl" - ], "license": "MPL-2.0", "optional": true, "os": [ @@ -7878,9 +7815,9 @@ } }, "node_modules/postcss": { - "version": "8.5.8", - "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.8.tgz", - "integrity": "sha512-OW/rX8O/jXnm82Ey1k44pObPtdblfiuWnrd8X7GJ7emImCOstunGbXUpp7HdBrFQX6rJzn3sPT397Wp5aCwCHg==", + "version": "8.5.14", + "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.14.tgz", + "integrity": "sha512-SoSL4+OSEtR99LHFZQiJLkT59C5B1amGO1NzTwj7TT1qCUgUO6hxOvzkOYxD+vMrXBM3XJIKzokoERdqQq/Zmg==", "funding": [ { "type": "opencollective", diff --git a/frontend/package.json b/frontend/package.json index fd6d973..b45d482 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -22,7 +22,7 @@ "@radix-ui/react-tooltip": "^1.2.8", "@types/js-yaml": "^4.0.9", "@xyflow/react": "^12.10.1", - "axios": "^1.13.6", + "axios": "^1.15.2", "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "dagre": "^0.8.5", diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index 477f2e4..c74c62b 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -19,9 +19,11 @@ import { LoginPage } from '@/components/LoginPage' import { NodeModal } from '@/components/modals/NodeModal' import { EdgeModal } from '@/components/modals/EdgeModal' import { ScanConfigModal } from '@/components/modals/ScanConfigModal' +import { ZigbeeImportModal } from '@/components/zigbee/ZigbeeImportModal' import { GroupRectModal, type GroupRectFormData } from '@/components/modals/GroupRectModal' import { ThemeModal } from '@/components/modals/ThemeModal' import { SearchModal } from '@/components/modals/SearchModal' +import { PendingDevicesModal } from '@/components/modals/PendingDevicesModal' import { ShortcutsModal } from '@/components/modals/ShortcutsModal' import { useCanvasStore } from '@/stores/canvasStore' import { useAuthStore } from '@/stores/authStore' @@ -30,6 +32,7 @@ import { canvasApi } from '@/api/client' import { demoNodes, demoEdges } from '@/utils/demoData' import { useStatusPolling } from '@/hooks/useStatusPolling' import type { NodeData, EdgeData, CustomStyleDef } from '@/types' +import type { ZigbeeNode, ZigbeeEdge } from '@/components/zigbee/types' const STANDALONE = import.meta.env.VITE_STANDALONE === 'true' const STANDALONE_STORAGE_KEY = 'homelable_canvas' @@ -45,8 +48,16 @@ export default function App() { const [themeModalOpen, setThemeModalOpen] = useState(false) const [searchOpen, setSearchOpen] = useState(false) - const [sidebarForceView, setSidebarForceView] = useState<'pending' | 'history' | undefined>(undefined) - const [highlightPendingId, setHighlightPendingId] = useState(undefined) + const [sidebarForceView, setSidebarForceView] = useState<'history' | undefined>(undefined) + const [pendingModalOpen, setPendingModalOpen] = useState(false) + const [pendingModalStatus, setPendingModalStatus] = useState<'pending' | 'hidden'>('pending') + const [pendingHighlightId, setPendingHighlightId] = useState(undefined) + const openPendingModal = useCallback((deviceId?: string, status: 'pending' | 'hidden' = 'pending') => { + setPendingHighlightId(undefined) + setPendingModalStatus(status) + setPendingModalOpen(true) + if (deviceId) setTimeout(() => setPendingHighlightId(deviceId), 0) + }, []) const [shortcutsOpen, setShortcutsOpen] = useState(false) const [addNodeOpen, setAddNodeOpen] = useState(false) const [addGroupRectOpen, setAddGroupRectOpen] = useState(false) @@ -55,6 +66,7 @@ export default function App() { const [editEdgeId, setEditEdgeId] = useState(null) const [scanConfigOpen, setScanConfigOpen] = useState(false) const [exportModalOpen, setExportModalOpen] = useState(false) + const [zigbeeImportOpen, setZigbeeImportOpen] = useState(false) // Declare handleSave before the Ctrl+S effect so it is in scope const handleSave = useCallback(async () => { @@ -315,6 +327,54 @@ export default function App() { setExportModalOpen(true) }, []) + const handleZigbeeAddToCanvas = useCallback((zigbeeNodes: ZigbeeNode[], zigbeeEdges: ZigbeeEdge[]) => { + snapshotHistory() + // Place nodes in a grid starting at x=500, y=100 + const COLS = 4 + const SPACING_X = 170 + const SPACING_Y = 100 + zigbeeNodes.forEach((zn, i) => { + const id = zn.id + const col = i % COLS + const row = Math.floor(i / COLS) + const position = { x: 500 + col * SPACING_X, y: 100 + row * SPACING_Y } + const newNode: import('@xyflow/react').Node = { + id, + type: zn.type, + position, + data: { + label: zn.friendly_name, + type: zn.type as NodeData['type'], + status: 'unknown' as const, + services: [], + ...(zn.lqi != null ? { properties: [{ key: 'LQI', value: String(zn.lqi), icon: 'signal', visible: true }] } : {}), + ...(zn.model ? { os: zn.model } : {}), + ...(zn.parent_id ? { parent_id: zn.parent_id } : {}), + }, + } + addNode(newNode) + }) + // Add IoT edges between Zigbee devices: parent bottom -> child top + zigbeeEdges.forEach((ze) => { + onConnect({ + source: ze.source, + sourceHandle: 'bottom', + target: ze.target, + targetHandle: 'top-t', + type: 'iot', + } as unknown as import('@xyflow/react').Connection) + }) + // Auto-select only the freshly imported nodes so the user can drag the + // whole subtree as a group. + const importedIds = new Set(zigbeeNodes.map((zn) => zn.id)) + useCanvasStore.setState((state) => ({ + nodes: state.nodes.map((n) => ({ ...n, selected: importedIds.has(n.id) })), + selectedNodeIds: Array.from(importedIds), + selectedNodeId: importedIds.size === 1 ? Array.from(importedIds)[0] : null, + })) + markUnsaved() + }, [addNode, onConnect, snapshotHistory, markUnsaved]) + const handleEdgeConnect = useCallback((connection: Connection) => { setPendingConnection(connection) }, []) @@ -384,10 +444,10 @@ export default function App() { onAddNode={() => setAddNodeOpen(true)} onAddGroupRect={() => setAddGroupRectOpen(true)} onScan={() => setScanConfigOpen(true)} + onZigbeeImport={() => setZigbeeImportOpen(true)} onSave={handleSave} - onNodeApproved={setEditNodeId} forceView={sidebarForceView} - highlightPendingId={highlightPendingId} + onOpenPending={openPendingModal} />
{ - setHighlightPendingId(undefined) - setSidebarForceView(undefined) - setTimeout(() => { - setHighlightPendingId(deviceId) - setSidebarForceView('pending') - }, 0) - }} + onOpenPending={(deviceId) => openPendingModal(deviceId)} />
{(selectedNodeId || selectedNodeIds.length > 1) && } @@ -483,6 +536,18 @@ export default function App() { /> )} + {!STANDALONE && ( + setZigbeeImportOpen(false)} + onAddToCanvas={handleZigbeeAddToCanvas} + onPendingImported={() => { + setSidebarForceView(undefined) + setTimeout(() => setSidebarForceView('history'), 0) + }} + /> + )} + setAddGroupRectOpen(false)} @@ -528,17 +593,17 @@ export default function App() { setSearchOpen(false)} - onOpenPending={(deviceId) => { - setHighlightPendingId(undefined) - setSidebarForceView(undefined) - setTimeout(() => { - setHighlightPendingId(deviceId) - setSidebarForceView('pending') - }, 0) - }} + onOpenPending={(deviceId) => openPendingModal(deviceId)} /> setShortcutsOpen(false)} /> + setPendingModalOpen(false)} + highlightId={pendingHighlightId} + initialStatus={pendingModalStatus} + /> + setExportModalOpen(false)} diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index b0ccdce..1c4485a 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -58,11 +58,27 @@ export const scanApi = { hidden: () => api.get('/scan/hidden'), runs: () => api.get('/scan/runs'), clearPending: () => api.delete('/scan/pending'), - approve: (id: string, nodeData: object) => api.post(`/scan/pending/${id}/approve`, nodeData), + approve: (id: string, nodeData: object) => + api.post<{ + approved: boolean + node_id: string + edges_created: number + edges: { id: string; source: string; target: string }[] + }>(`/scan/pending/${id}/approve`, nodeData), hide: (id: string) => api.post(`/scan/pending/${id}/hide`), ignore: (id: string) => api.post(`/scan/pending/${id}/ignore`), - bulkApprove: (ids: string[]) => api.post<{ approved: number; node_ids: string[]; device_ids: string[]; skipped: number }>('/scan/pending/bulk-approve', { device_ids: ids }), + bulkApprove: (ids: string[]) => + api.post<{ + approved: number + node_ids: string[] + device_ids: string[] + edges_created: number + edges: { id: string; source: string; target: string }[] + skipped: number + }>('/scan/pending/bulk-approve', { device_ids: ids }), bulkHide: (ids: string[]) => api.post<{ hidden: number; skipped: number }>('/scan/pending/bulk-hide', { device_ids: ids }), + restore: (id: string) => api.post<{ restored: boolean; device_id: string }>(`/scan/pending/${id}/restore`), + bulkRestore: (ids: string[]) => api.post<{ restored: number; skipped: number }>('/scan/pending/bulk-restore', { device_ids: ids }), stop: (runId: string) => api.post(`/scan/${runId}/stop`), getConfig: () => api.get<{ ranges: string[] }>('/scan/config'), saveConfig: (data: { ranges: string[] }) => api.post('/scan/config', data), @@ -72,3 +88,50 @@ export const settingsApi = { get: () => api.get<{ interval_seconds: number }>('/settings'), save: (data: { interval_seconds: number }) => api.post<{ interval_seconds: number }>('/settings', data), } + +export const zigbeeApi = { + testConnection: (data: { + mqtt_host: string + mqtt_port: number + mqtt_username?: string + mqtt_password?: string + mqtt_tls?: boolean + mqtt_tls_insecure?: boolean + }) => + api.post<{ connected: boolean; message: string }>('/zigbee/test-connection', data), + + importNetwork: (data: { + mqtt_host: string + mqtt_port: number + mqtt_username?: string + mqtt_password?: string + base_topic?: string + mqtt_tls?: boolean + mqtt_tls_insecure?: boolean + }) => + api.post<{ + nodes: import('@/components/zigbee/types').ZigbeeNode[] + edges: import('@/components/zigbee/types').ZigbeeEdge[] + device_count: number + }>('/zigbee/import', data), + + importToPending: (data: { + mqtt_host: string + mqtt_port: number + mqtt_username?: string + mqtt_password?: string + base_topic?: string + mqtt_tls?: boolean + mqtt_tls_insecure?: boolean + }) => + api.post<{ + id: string + status: string + kind: string + ranges: string[] + devices_found: number + started_at: string + finished_at: string | null + error: string | null + }>('/zigbee/import-pending', data), +} diff --git a/frontend/src/components/canvas/SearchBar.tsx b/frontend/src/components/canvas/SearchBar.tsx index beac37b..5dca55c 100644 --- a/frontend/src/components/canvas/SearchBar.tsx +++ b/frontend/src/components/canvas/SearchBar.tsx @@ -57,8 +57,10 @@ export function SearchBar({ onOpenPending }: SearchBarProps) { const pendingResults = q ? pendingDevices.filter((d) => - d.ip.toLowerCase().includes(q) || + d.ip?.toLowerCase().includes(q) || d.hostname?.toLowerCase().includes(q) || + d.friendly_name?.toLowerCase().includes(q) || + d.ieee_address?.toLowerCase().includes(q) || d.services.some((s) => s.service_name?.toLowerCase().includes(q) || s.category?.toLowerCase().includes(q) @@ -196,10 +198,10 @@ export function SearchBar({ onOpenPending }: SearchBarProps) { > pending - {d.hostname ?? d.ip} + {d.friendly_name ?? d.hostname ?? d.ip ?? d.ieee_address ?? 'device'} - {serviceName ?? d.ip} + {serviceName ?? d.ip ?? d.ieee_address ?? ''} ) diff --git a/frontend/src/components/canvas/nodes/index.tsx b/frontend/src/components/canvas/nodes/index.tsx index cfb6604..6bf6eff 100644 --- a/frontend/src/components/canvas/nodes/index.tsx +++ b/frontend/src/components/canvas/nodes/index.tsx @@ -1,7 +1,7 @@ import { type NodeProps, type Node } from '@xyflow/react' import { Globe, Router, Network, Server, Layers, Box, Container, - HardDrive, Cpu, Wifi, Circle, Cctv, Printer, Monitor, PlugZap, Anchor, Package, Flame, + HardDrive, Cpu, Wifi, Circle, Cctv, Printer, Monitor, PlugZap, Anchor, Package, Flame, Radio, Antenna, } from 'lucide-react' import { BaseNode } from './BaseNode' import type { NodeData } from '@/types' @@ -26,3 +26,7 @@ export const CplNode = (props: N) => export const DockerHostNode = (props: N) => export const DockerContainerNode = (props: N) => export const GenericNode = (props: N) => +// Zigbee node types +export const ZigbeeCoordinatorNode = (props: N) => +export const ZigbeeRouterNode = (props: N) => +export const ZigbeeEndDeviceNode = (props: N) => diff --git a/frontend/src/components/canvas/nodes/nodeTypes.ts b/frontend/src/components/canvas/nodes/nodeTypes.ts index 6adba16..c390c4a 100644 --- a/frontend/src/components/canvas/nodes/nodeTypes.ts +++ b/frontend/src/components/canvas/nodes/nodeTypes.ts @@ -1,4 +1,4 @@ -import { IspNode, RouterNode, FirewallNode, SwitchNode, ServerNode, VmNode, LxcNode, NasNode, IotNode, ApNode, CameraNode, PrinterNode, ComputerNode, CplNode, DockerHostNode, DockerContainerNode, GenericNode } from './index' +import { IspNode, RouterNode, FirewallNode, SwitchNode, ServerNode, VmNode, LxcNode, NasNode, IotNode, ApNode, CameraNode, PrinterNode, ComputerNode, CplNode, DockerHostNode, DockerContainerNode, GenericNode, ZigbeeCoordinatorNode, ZigbeeRouterNode, ZigbeeEndDeviceNode } from './index' import { ProxmoxGroupNode } from './ProxmoxGroupNode' import { GroupRectNode } from './GroupRectNode' import { GroupNode } from './GroupNode' @@ -24,4 +24,7 @@ export const nodeTypes = { generic: GenericNode, groupRect: GroupRectNode, group: GroupNode, + zigbee_coordinator: ZigbeeCoordinatorNode, + zigbee_router: ZigbeeRouterNode, + zigbee_enddevice: ZigbeeEndDeviceNode, } diff --git a/frontend/src/components/modals/PendingDeviceModal.tsx b/frontend/src/components/modals/PendingDeviceModal.tsx index cbad9c7..aa76069 100644 --- a/frontend/src/components/modals/PendingDeviceModal.tsx +++ b/frontend/src/components/modals/PendingDeviceModal.tsx @@ -12,7 +12,7 @@ interface Service { export interface PendingDevice { id: string - ip: string + ip: string | null mac: string | null hostname: string | null os: string | null @@ -20,6 +20,12 @@ export interface PendingDevice { suggested_type: string | null status: string discovery_source: string | null + ieee_address?: string | null + friendly_name?: string | null + device_subtype?: string | null + model?: string | null + vendor?: string | null + lqi?: number | null discovered_at: string } @@ -77,6 +83,8 @@ export function PendingDeviceModal({ device, onClose, onApprove, onHide, onIgnor if (!device) return null const TypeIcon = TYPE_ICONS[device.suggested_type ?? 'generic'] ?? Circle + const isZigbee = device.discovery_source === 'zigbee' + const titleLabel = device.friendly_name ?? device.hostname ?? device.ip ?? device.ieee_address ?? 'Pending device' const handleApprove = () => { onApprove(device) } const handleHide = () => { onHide(device); onClose() } @@ -88,17 +96,30 @@ export function PendingDeviceModal({ device, onClose, onApprove, onHide, onIgnor - {device.hostname ?? device.ip} + {titleLabel} + {isZigbee && ( + + Zigbee + + )}
{/* Device info */}
- + {device.ip && } {device.hostname && } {device.mac && } {device.os && } + {device.ieee_address && } + {device.friendly_name && device.friendly_name !== device.hostname && ( + + )} + {device.vendor && } + {device.model && } + {device.device_subtype && } + {device.lqi != null && } {device.suggested_type && ( )} @@ -108,8 +129,8 @@ export function PendingDeviceModal({ device, onClose, onApprove, onHide, onIgnor
- {/* Services */} -
+ {/* Services (skipped for Zigbee devices — they don't have IP services) */} + {!isZigbee &&

Services found ({device.services.length})

@@ -138,7 +159,7 @@ export function PendingDeviceModal({ device, onClose, onApprove, onHide, onIgnor ))}
)} -
+
} {/* Actions */}
diff --git a/frontend/src/components/modals/PendingDevicesModal.tsx b/frontend/src/components/modals/PendingDevicesModal.tsx new file mode 100644 index 0000000..41a86ee --- /dev/null +++ b/frontend/src/components/modals/PendingDevicesModal.tsx @@ -0,0 +1,677 @@ +import { useState, useEffect, useCallback, useRef, useMemo } from 'react' +import { + Globe, Router, Server, Layers, Box, Container, HardDrive, Cpu, Wifi, Circle, Network, + Search, RefreshCw, X, CheckCircle2, EyeOff, Trash2, Loader2, +} from 'lucide-react' +import { Dialog, DialogContent, DialogHeader, DialogTitle } from '@/components/ui/dialog' +import { scanApi } from '@/api/client' +import { useCanvasStore } from '@/stores/canvasStore' +import { toast } from 'sonner' +import { PendingDeviceModal, type PendingDevice } from '@/components/modals/PendingDeviceModal' +import type { NodeType, ServiceInfo } from '@/types' + +interface PendingDevicesModalProps { + open: boolean + onClose: () => void + highlightId?: string + initialStatus?: 'pending' | 'hidden' +} + +const PORT_COLORS: Record = { + 22: '#a855f7', // SSH purple + 80: '#00d4ff', // HTTP cyan + 443: '#39d353', // HTTPS green + 53: '#e3b341', // DNS amber + 3306: '#a855f7', // MySQL + 5432: '#a855f7', // Postgres + 6379: '#f85149', // Redis + 9090: '#e3b341', // Prometheus + 3000: '#00d4ff', // Grafana/dev + 8080: '#00d4ff', + 8443: '#39d353', +} + +const CATEGORY_COLORS: Record = { + hypervisor: '#ff6e00', + nas: '#39d353', + automation: '#a855f7', + containers: '#00d4ff', + network: '#39d353', + security: '#f85149', + monitoring: '#e3b341', + database: '#a855f7', + web: '#00d4ff', + media: '#ff6e00', + iot: '#e3b341', +} + +function serviceColor(port: number | null | undefined, category?: string | null): string { + if (port != null && PORT_COLORS[port]) return PORT_COLORS[port] + if (category && CATEGORY_COLORS[category.toLowerCase()]) return CATEGORY_COLORS[category.toLowerCase()] + return '#8b949e' +} + +const TYPE_ICONS: Record = { + isp: Globe, + router: Router, + server: Server, + proxmox: Layers, + vm: Box, + lxc: Container, + nas: HardDrive, + iot: Cpu, + ap: Wifi, + switch: Network, + generic: Circle, +} + +type SourceFilter = 'all' | 'ip' | 'zigbee' +type StatusFilter = 'pending' | 'hidden' + +function inferSource(d: PendingDevice): 'zigbee' | 'ip' { + if (d.discovery_source === 'zigbee' || d.ieee_address) return 'zigbee' + return 'ip' +} + +const COMMON_PORTS = new Set([22, 80, 443]) + +function specialServiceName(d: PendingDevice): string | undefined { + const candidates = (d.services ?? []).filter( + (s) => s.category != null && s.port != null && !COMMON_PORTS.has(s.port) && s.service_name, + ) + // Deprioritize generic web category so apps like home assistant / jellyfin win + const nonWeb = candidates.find((s) => s.category?.toLowerCase() !== 'web') + return (nonWeb ?? candidates[0])?.service_name ?? undefined +} + +function deviceLabel(d: PendingDevice): string { + return d.friendly_name ?? d.hostname ?? specialServiceName(d) ?? d.ip ?? d.ieee_address ?? 'device' +} + +function injectAutoEdges(edges: { id: string; source: string; target: string }[] | undefined) { + if (!edges || edges.length === 0) return + useCanvasStore.setState((state) => ({ + edges: [ + ...state.edges, + ...edges.map((e) => ({ + id: e.id, + source: e.source, + target: e.target, + sourceHandle: 'bottom', + targetHandle: 'top-t', + type: 'iot', + data: { type: 'iot' as const }, + })), + ], + hasUnsavedChanges: true, + })) +} + +export function PendingDevicesModal({ open, onClose, highlightId, initialStatus = 'pending' }: PendingDevicesModalProps) { + const [devices, setDevices] = useState([]) + const [loading, setLoading] = useState(false) + const [selected, setSelected] = useState(null) + const [selectMode, setSelectMode] = useState(false) + const [selectedIds, setSelectedIds] = useState>(new Set()) + const [search, setSearch] = useState('') + const [sourceFilter, setSourceFilter] = useState('all') + const [typeFilter, setTypeFilter] = useState('all') + const [statusFilter, setStatusFilter] = useState(initialStatus) + const { addNode, scanEventTs } = useCanvasStore() + const highlightRef = useRef(null) + + const load = useCallback(async () => { + setLoading(true) + try { + const res = statusFilter === 'pending' ? await scanApi.pending() : await scanApi.hidden() + setDevices(res.data) + } catch { + toast.error(`Failed to load ${statusFilter} devices`) + } finally { + setLoading(false) + } + }, [statusFilter]) + + useEffect(() => { if (open) load() }, [open, load]) + useEffect(() => { if (open && scanEventTs > 0) load() }, [scanEventTs, open, load]) + + // Reset transient state when reopening + useEffect(() => { + if (!open) { + setSelectMode(false) + setSelectedIds(new Set()) + setSearch('') + } else { + setStatusFilter(initialStatus) + } + }, [open, initialStatus]) + + const distinctTypes = useMemo(() => { + const set = new Set() + devices.forEach((d) => { if (d.suggested_type) set.add(d.suggested_type) }) + return [...set].sort() + }, [devices]) + + const filtered = useMemo(() => { + const q = search.trim().toLowerCase() + return devices.filter((d) => { + if (sourceFilter !== 'all' && inferSource(d) !== sourceFilter) return false + if (typeFilter !== 'all' && d.suggested_type !== typeFilter) return false + if (q) { + const hay = [ + d.friendly_name, d.hostname, d.ip, d.mac, d.ieee_address, d.vendor, d.model, + ...d.services.map((s) => s.service_name), + ].filter(Boolean).join(' ').toLowerCase() + if (!hay.includes(q)) return false + } + return true + }) + }, [devices, search, sourceFilter, typeFilter]) + + useEffect(() => { + if (!highlightId || loading || !open) return + highlightRef.current?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }) + }, [highlightId, loading, open, filtered]) + + const toggleSelect = (id: string) => { + setSelectedIds((prev) => { + const next = new Set(prev) + if (next.has(id)) next.delete(id); else next.add(id) + return next + }) + } + + const handleCardClick = (d: PendingDevice) => { + if (selectMode) { toggleSelect(d.id); return } + if (statusFilter === 'hidden') { handleRestore(d); return } + setSelected(d) + } + + const handleRestore = async (device: PendingDevice) => { + try { + await scanApi.restore(device.id) + setDevices((prev) => prev.filter((d) => d.id !== device.id)) + toast.success(`Restored ${deviceLabel(device)}`) + } catch { + toast.error('Failed to restore device') + } + } + + const handleBulkRestore = async () => { + const ids = [...selectedIds] + if (ids.length === 0) return + try { + const res = await scanApi.bulkRestore(ids) + setDevices((prev) => prev.filter((d) => !ids.includes(d.id))) + setSelectedIds(new Set()) + toast.success(`Restored ${res.data.restored} device${res.data.restored !== 1 ? 's' : ''}`) + } catch { + toast.error('Failed to bulk restore devices') + } + } + + const enterSelectMode = () => { + setSelectMode(true) + } + + const exitSelectMode = () => { + setSelectMode(false) + setSelectedIds(new Set()) + } + + const selectAllVisible = () => { + setSelectedIds(new Set(filtered.map((d) => d.id))) + } + + const handleClearAll = async () => { + const targets = filtered + if (targets.length === 0) return + const filtersActive = targets.length !== devices.length + try { + if (filtersActive) { + const results = await Promise.allSettled(targets.map((d) => scanApi.ignore(d.id))) + const failed = results.filter((r) => r.status === 'rejected').length + const removedIds = new Set( + targets.filter((_, i) => results[i].status === 'fulfilled').map((d) => d.id) + ) + setDevices((prev) => prev.filter((d) => !removedIds.has(d.id))) + setSelectedIds(new Set()) + if (failed > 0) toast.error(`Removed ${removedIds.size}, ${failed} failed`) + else toast.success(`Removed ${removedIds.size} device${removedIds.size !== 1 ? 's' : ''}`) + } else { + await scanApi.clearPending() + setDevices([]) + setSelectedIds(new Set()) + toast.success('Pending devices cleared') + } + } catch { + toast.error('Failed to clear pending devices') + } + } + + const handleApprove = async (device: PendingDevice) => { + try { + const fallbackLabel = deviceLabel(device) + const nodeData = { + label: fallbackLabel, + type: (device.suggested_type ?? 'generic') as NodeType, + ip: device.ip ?? undefined, + hostname: device.hostname ?? undefined, + status: 'unknown', + services: (device.services ?? []) as ServiceInfo[], + } + const res = await scanApi.approve(device.id, nodeData) + const nodeId = res.data.node_id + addNode({ + id: nodeId, + type: nodeData.type, + position: { x: 400, y: 300 }, + data: { ...nodeData, status: 'unknown' as const }, + }) + injectAutoEdges(res.data.edges) + const extra = res.data.edges_created > 0 ? ` (+${res.data.edges_created} link${res.data.edges_created !== 1 ? 's' : ''})` : '' + toast.success(`Approved ${nodeData.label}${extra}`) + setDevices((prev) => prev.filter((d) => d.id !== device.id)) + setSelected(null) + } catch { + toast.error('Failed to approve device') + } + } + + const handleHide = async (device: PendingDevice) => { + try { + await scanApi.hide(device.id) + setDevices((prev) => prev.filter((d) => d.id !== device.id)) + setSelected(null) + toast.success('Device hidden') + } catch { + toast.error('Failed to hide device') + } + } + + const handleIgnore = async (device: PendingDevice) => { + try { + await scanApi.ignore(device.id) + setDevices((prev) => prev.filter((d) => d.id !== device.id)) + setSelected(null) + } catch { + toast.error('Failed to remove device') + } + } + + const handleBulkApprove = async () => { + const ids = [...selectedIds] + if (ids.length === 0) return + try { + const res = await scanApi.bulkApprove(ids) + const deviceToNode: Record = {} + res.data.device_ids.forEach((did, i) => { deviceToNode[did] = res.data.node_ids[i] }) + const approvedDevices = devices.filter((d) => ids.includes(d.id)) + approvedDevices.forEach((d, i) => { + const nodeId = deviceToNode[d.id] + if (!nodeId) return + addNode({ + id: nodeId, + type: (d.suggested_type ?? 'generic') as NodeType, + position: { x: 400 + (i % 4) * 160, y: 300 + Math.floor(i / 4) * 100 }, + data: { + label: deviceLabel(d), + type: (d.suggested_type ?? 'generic') as NodeType, + ip: d.ip ?? undefined, + hostname: d.hostname ?? undefined, + status: 'unknown' as const, + services: (d.services ?? []) as ServiceInfo[], + }, + }) + }) + injectAutoEdges(res.data.edges) + setDevices((prev) => prev.filter((d) => !ids.includes(d.id))) + setSelectedIds(new Set()) + const linkExtra = res.data.edges_created > 0 ? ` (+${res.data.edges_created} link${res.data.edges_created !== 1 ? 's' : ''})` : '' + toast.success(`Approved ${res.data.approved} device${res.data.approved !== 1 ? 's' : ''}${linkExtra}`) + } catch { + toast.error('Failed to bulk approve devices') + } + } + + const handleBulkHide = async () => { + const ids = [...selectedIds] + if (ids.length === 0) return + try { + const res = await scanApi.bulkHide(ids) + setDevices((prev) => prev.filter((d) => !ids.includes(d.id))) + setSelectedIds(new Set()) + toast.success(`Hidden ${res.data.hidden} device${res.data.hidden !== 1 ? 's' : ''}`) + } catch { + toast.error('Failed to bulk hide devices') + } + } + + // Keyboard shortcuts: 's' select-mode, 'a' select-all-visible, Esc clears selection or closes, '/' focuses search + const searchRef = useRef(null) + useEffect(() => { + if (!open) return + const handler = (e: KeyboardEvent) => { + const target = e.target as HTMLElement | null + const inField = target && (target.tagName === 'INPUT' || target.tagName === 'TEXTAREA' || target.tagName === 'SELECT') + if (e.key === 'Escape') { + if (selectMode && selectedIds.size > 0) { e.preventDefault(); setSelectedIds(new Set()) } + return + } + if (inField) return + if (e.key === '/') { e.preventDefault(); searchRef.current?.focus() } + else if (e.key.toLowerCase() === 's') { e.preventDefault(); if (selectMode) exitSelectMode(); else enterSelectMode() } + else if (e.key.toLowerCase() === 'a' && selectMode) { e.preventDefault(); selectAllVisible() } + else if (e.key === 'Enter' && selectMode && selectedIds.size > 0) { e.preventDefault(); handleBulkApprove() } + } + window.addEventListener('keydown', handler) + return () => window.removeEventListener('keydown', handler) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [open, selectMode, selectedIds, filtered]) + + return ( + <> + { if (!v) onClose() }}> + + +
+ + {statusFilter === 'pending' ? 'Pending Devices' : 'Hidden Devices'} + + ({filtered.length}{filtered.length !== devices.length && ` of ${devices.length}`}) + + +
+ + {statusFilter === 'pending' && devices.length > 0 && ( + + )} + +
+
+
+ + {/* Toolbar */} +
+
+ + setSearch(e.target.value)} + placeholder="Search name, IP, MAC, IEEE, service…" + className="w-full text-xs bg-[#0d1117] border border-border rounded px-7 py-1.5 outline-none focus:border-[#00d4ff]/50" + /> +
+
+ + + +
+ +
+ + +
+ +
+ + {/* Body */} +
+ {loading && ( +
+ +
+ )} + {!loading && filtered.length === 0 && ( +

+ {devices.length === 0 ? `No ${statusFilter} devices` : 'No devices match filters'} +

+ )} + {!loading && filtered.length > 0 && ( +
+ {filtered.map((d) => ( + handleCardClick(d)} + cardRef={d.id === highlightId ? highlightRef : undefined} + /> + ))} +
+ )} +
+ + {/* Selection action bar */} + {selectMode && ( +
+ + {selectedIds.size} selected + + + +
+ {statusFilter === 'pending' && ( + <> + + + + )} + {statusFilter === 'hidden' && ( + + )} +
+ )} + +
+ + setSelected(null)} + onApprove={handleApprove} + onHide={handleHide} + onIgnore={handleIgnore} + /> + + ) +} + +interface DeviceCardProps { + device: PendingDevice + selected: boolean + selectMode: boolean + highlighted: boolean + onClick: () => void + cardRef?: React.Ref +} + +function DeviceCard({ device, selected, selectMode, highlighted, onClick, cardRef }: DeviceCardProps) { + const source = inferSource(device) + const Icon = TYPE_ICONS[device.suggested_type ?? 'generic'] ?? Circle + const label = deviceLabel(device) + const sourceColor = source === 'zigbee' ? '#00d4ff' : '#a855f7' + const sourceLabel = source === 'zigbee' ? 'ZIGBEE' : (device.discovery_source ?? 'IP').toUpperCase() + const services = device.services ?? [] + const visibleServices = services.slice(0, 4) + const moreServices = services.length - visibleServices.length + + const borderClass = highlighted + ? 'border-[#e3b341] bg-[#2d3748]' + : selected + ? 'border-[#00d4ff] bg-[#00d4ff]/5 shadow-[0_0_0_1px_rgba(0,212,255,0.4)] scale-[1.02]' + : 'border-border bg-[#161b22] hover:border-[#30363d] hover:bg-[#21262d]' + + return ( + + ) +} + +function InfoLine({ label, value }: { label: string; value: string }) { + return ( +
+ {label} + {value} +
+ ) +} diff --git a/frontend/src/components/modals/SearchModal.tsx b/frontend/src/components/modals/SearchModal.tsx index 10f4b5c..630a22d 100644 --- a/frontend/src/components/modals/SearchModal.tsx +++ b/frontend/src/components/modals/SearchModal.tsx @@ -33,8 +33,10 @@ export function SearchModal({ open, onClose, onOpenPending }: SearchModalProps) ).slice(0, 6) const pendingResults = q.length === 0 ? [] : pendingDevices.filter((d) => - d.ip.toLowerCase().includes(q) || + d.ip?.toLowerCase().includes(q) || d.hostname?.toLowerCase().includes(q) || + d.friendly_name?.toLowerCase().includes(q) || + d.ieee_address?.toLowerCase().includes(q) || d.services.some((s) => s.service_name?.toLowerCase().includes(q) || s.category?.toLowerCase().includes(q) diff --git a/frontend/src/components/modals/__tests__/PendingDevicesModal.test.tsx b/frontend/src/components/modals/__tests__/PendingDevicesModal.test.tsx new file mode 100644 index 0000000..6164950 --- /dev/null +++ b/frontend/src/components/modals/__tests__/PendingDevicesModal.test.tsx @@ -0,0 +1,216 @@ +import { describe, it, expect, beforeEach, vi } from 'vitest' +import { render, screen, fireEvent, waitFor } from '@testing-library/react' +import { PendingDevicesModal } from '../PendingDevicesModal' +import { useCanvasStore } from '@/stores/canvasStore' + +vi.mock('@/stores/canvasStore') + +const mockBulkApprove = vi.fn() +const mockBulkHide = vi.fn() +const mockRestore = vi.fn() +const mockBulkRestore = vi.fn() +const mockApprove = vi.fn() +const mockHide = vi.fn() +const mockPending = vi.fn() +const mockHidden = vi.fn() + +vi.mock('@/api/client', () => ({ + scanApi: { + pending: (...a: unknown[]) => mockPending(...a), + hidden: (...a: unknown[]) => mockHidden(...a), + clearPending: vi.fn().mockResolvedValue({}), + approve: (...a: unknown[]) => mockApprove(...a), + hide: (...a: unknown[]) => mockHide(...a), + ignore: vi.fn().mockResolvedValue({}), + bulkApprove: (...a: unknown[]) => mockBulkApprove(...a), + bulkHide: (...a: unknown[]) => mockBulkHide(...a), + restore: (...a: unknown[]) => mockRestore(...a), + bulkRestore: (...a: unknown[]) => mockBulkRestore(...a), + }, +})) + +vi.mock('sonner', () => ({ toast: { success: vi.fn(), error: vi.fn() } })) + +vi.mock('@/components/modals/PendingDeviceModal', () => ({ + PendingDeviceModal: ({ device }: { device: unknown }) => + device ?
: null, +})) + +const DEVICE_IP = { + id: 'dev-a', + ip: '192.168.1.10', + hostname: 'host-a', + mac: 'aa:bb:cc:dd:ee:01', + os: null, + services: [{ port: 80, protocol: 'tcp', service_name: 'http' }], + suggested_type: 'server', + status: 'pending', + discovery_source: 'arp', + discovered_at: '2026-01-01T00:00:00Z', +} + +const DEVICE_ZIGBEE = { + id: 'dev-b', + ip: null, + hostname: null, + mac: null, + os: null, + services: [], + suggested_type: 'iot', + status: 'pending', + discovery_source: 'zigbee', + ieee_address: '0x00124b001234abcd', + friendly_name: 'living-room-bulb', + vendor: 'Philips', + model: 'Hue White', + discovered_at: '2026-01-02T00:00:00Z', +} + +beforeEach(() => { + vi.clearAllMocks() + vi.mocked(useCanvasStore).mockReturnValue({ + addNode: vi.fn(), + scanEventTs: 0, + } as unknown as ReturnType) + // setState is used by injectAutoEdges + ;(useCanvasStore as unknown as { setState: (fn: unknown) => void }).setState = vi.fn() + mockPending.mockResolvedValue({ data: [DEVICE_IP, DEVICE_ZIGBEE] }) + mockHidden.mockResolvedValue({ data: [] }) + mockApprove.mockResolvedValue({ data: { node_id: 'n1', edges: [], edges_created: 0 } }) + mockHide.mockResolvedValue({ data: {} }) + mockBulkApprove.mockResolvedValue({ + data: { approved: 2, node_ids: ['n1', 'n2'], device_ids: ['dev-a', 'dev-b'], edges: [], edges_created: 0 }, + }) + mockBulkHide.mockResolvedValue({ data: { hidden: 2, skipped: 0 } }) + mockRestore.mockResolvedValue({ data: { restored: true, device_id: 'dev-a' } }) + mockBulkRestore.mockResolvedValue({ data: { restored: 1, skipped: 0 } }) +}) + +const baseProps = { + open: true, + onClose: vi.fn(), +} + +describe('PendingDevicesModal', () => { + it('loads and renders pending devices on open', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + expect(screen.getByText('living-room-bulb')).toBeInTheDocument() + }) + + it('shows source chip ZIGBEE for zigbee device', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + expect(screen.getByText('ZIGBEE')).toBeInTheDocument() + }) + + it('filters by search query', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.change(screen.getByPlaceholderText(/Search/), { target: { value: 'living' } }) + expect(screen.queryByTestId('pending-card-dev-a')).not.toBeInTheDocument() + expect(screen.getByTestId('pending-card-dev-b')).toBeInTheDocument() + }) + + it('filters by source (zigbee only)', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Zigbee' })) + expect(screen.queryByTestId('pending-card-dev-a')).not.toBeInTheDocument() + expect(screen.getByTestId('pending-card-dev-b')).toBeInTheDocument() + }) + + it('filters by suggested type', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.change(screen.getByLabelText('Type filter'), { target: { value: 'server' } }) + expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument() + expect(screen.queryByTestId('pending-card-dev-b')).not.toBeInTheDocument() + }) + + it('switches to hidden status loads hidden devices', async () => { + mockHidden.mockResolvedValue({ + data: [{ ...DEVICE_IP, id: 'h1', hostname: 'hidden-host', status: 'hidden' }], + }) + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Hidden' })) + await waitFor(() => expect(screen.getByTestId('pending-card-h1')).toBeInTheDocument()) + expect(mockHidden).toHaveBeenCalled() + }) + + it('opens approval modal when card is clicked outside select mode', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByTestId('pending-card-dev-a')) + expect(screen.getByTestId('approval-modal')).toBeInTheDocument() + }) + + it('toggles selection in select mode instead of opening approval', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Select mode' })) + fireEvent.click(screen.getByTestId('pending-card-dev-a')) + expect(screen.queryByTestId('approval-modal')).not.toBeInTheDocument() + expect(screen.getByText('1 selected')).toBeInTheDocument() + }) + + it('select all visible selects only filtered devices', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Select mode' })) + fireEvent.change(screen.getByPlaceholderText(/Search/), { target: { value: 'host-a' } }) + fireEvent.click(screen.getByRole('button', { name: /Select all visible/ })) + expect(screen.getByText('1 selected')).toBeInTheDocument() + }) + + it('bulk approve calls API with selected ids', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Select mode' })) + fireEvent.click(screen.getByTestId('pending-card-dev-a')) + fireEvent.click(screen.getByTestId('pending-card-dev-b')) + fireEvent.click(screen.getByRole('button', { name: /Approve \(2\)/ })) + await waitFor(() => expect(mockBulkApprove).toHaveBeenCalledWith(['dev-a', 'dev-b'])) + }) + + it('bulk hide calls API with selected ids', async () => { + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Select mode' })) + fireEvent.click(screen.getByTestId('pending-card-dev-a')) + fireEvent.click(screen.getByRole('button', { name: /Hide \(1\)/ })) + await waitFor(() => expect(mockBulkHide).toHaveBeenCalledWith(['dev-a'])) + }) + + it('does not load when closed', () => { + render() + expect(mockPending).not.toHaveBeenCalled() + }) + + it('respects initialStatus=hidden', async () => { + mockHidden.mockResolvedValue({ data: [{ ...DEVICE_IP, hostname: 'hidden-host', status: 'hidden' }] }) + render() + await waitFor(() => expect(mockHidden).toHaveBeenCalled()) + expect(mockPending).not.toHaveBeenCalled() + }) + + it('clicking a hidden card restores it instead of opening approval', async () => { + mockHidden.mockResolvedValue({ data: [{ ...DEVICE_IP, status: 'hidden' }] }) + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByTestId('pending-card-dev-a')) + await waitFor(() => expect(mockRestore).toHaveBeenCalledWith('dev-a')) + expect(screen.queryByTestId('approval-modal')).not.toBeInTheDocument() + }) + + it('bulk restore in hidden mode calls API with selected ids', async () => { + mockHidden.mockResolvedValue({ data: [{ ...DEVICE_IP, status: 'hidden' }] }) + render() + await waitFor(() => expect(screen.getByTestId('pending-card-dev-a')).toBeInTheDocument()) + fireEvent.click(screen.getByRole('button', { name: 'Select mode' })) + fireEvent.click(screen.getByTestId('pending-card-dev-a')) + fireEvent.click(screen.getByRole('button', { name: /Restore \(1\)/ })) + await waitFor(() => expect(mockBulkRestore).toHaveBeenCalledWith(['dev-a'])) + }) +}) diff --git a/frontend/src/components/panels/Sidebar.tsx b/frontend/src/components/panels/Sidebar.tsx index 4325938..9a0f889 100644 --- a/frontend/src/components/panels/Sidebar.tsx +++ b/frontend/src/components/panels/Sidebar.tsx @@ -1,5 +1,5 @@ import { useState, useCallback, useEffect, useRef } from 'react' -import { Plus, Save, ScanLine, ChevronLeft, ChevronRight, LayoutDashboard, Clock, EyeOff, Trash2, RefreshCw, Loader2, Square, Eye, Settings, StopCircle, X, LogOut } from 'lucide-react' +import { Plus, Save, ScanLine, ChevronLeft, ChevronRight, LayoutDashboard, Clock, EyeOff, RefreshCw, Loader2, Square, Eye, Settings, StopCircle, LogOut, Network } from 'lucide-react' import { Logo } from '@/components/ui/Logo' import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip' import { useCanvasStore } from '@/stores/canvasStore' @@ -8,23 +8,19 @@ import { scanApi, settingsApi } from '@/api/client' import { toast } from 'sonner' import { useLatestRelease } from '@/hooks/useLatestRelease' -import { PendingDeviceModal, type PendingDevice } from '@/components/modals/PendingDeviceModal' - const STANDALONE = import.meta.env.VITE_STANDALONE === 'true' -type SidebarView = 'canvas' | 'pending' | 'hidden' | 'history' | 'settings' +type SidebarView = 'canvas' | 'history' | 'settings' -const ALL_VIEWS = [ - { id: 'canvas' as SidebarView, icon: LayoutDashboard, label: 'Canvas' }, - { id: 'pending' as SidebarView, icon: ScanLine, label: 'Pending Devices' }, - { id: 'hidden' as SidebarView, icon: EyeOff, label: 'Hidden Devices' }, - { id: 'history' as SidebarView, icon: Clock, label: 'Scan History' }, +const PENDING_TRIGGERS: { kind: 'pending' | 'hidden'; icon: typeof ScanLine; label: string }[] = [ + { kind: 'pending', icon: ScanLine, label: 'Pending Devices' }, + { kind: 'hidden', icon: EyeOff, label: 'Hidden Devices' }, ] -const VIEWS = STANDALONE ? ALL_VIEWS.slice(0, 1) : ALL_VIEWS interface ScanRun { id: string status: string + kind?: string ranges: string[] devices_found: number started_at: string @@ -36,13 +32,13 @@ interface SidebarProps { onAddNode: () => void onAddGroupRect: () => void onScan: () => void + onZigbeeImport: () => void onSave: () => void - onNodeApproved: (nodeId: string) => void forceView?: SidebarView - highlightPendingId?: string + onOpenPending: (deviceId?: string, status?: 'pending' | 'hidden') => void } -export function Sidebar({ onAddNode, onAddGroupRect, onScan, onSave, onNodeApproved, forceView, highlightPendingId }: SidebarProps) { +export function Sidebar({ onAddNode, onAddGroupRect, onScan, onZigbeeImport, onSave, forceView, onOpenPending }: SidebarProps) { const [collapsed, setCollapsed] = useState(false) const [activeView, setActiveView] = useState(forceView ?? 'canvas') const [prevForceView, setPrevForceView] = useState(forceView) @@ -87,23 +83,36 @@ export function Sidebar({ onAddNode, onAddGroupRect, onScan, onSave, onNodeAppro {/* Views */} {/* View content (only when expanded) */} {!collapsed && activeView !== 'canvas' && (
- {activeView === 'pending' && } - {activeView === 'hidden' && } {activeView === 'history' && } {activeView === 'settings' && }
@@ -137,6 +146,7 @@ export function Sidebar({ onAddNode, onAddGroupRect, onScan, onSave, onNodeAppro {!STANDALONE && } + {!STANDALONE && } void; highlightId?: string }) { - const [devices, setDevices] = useState([]) - const [loading, setLoading] = useState(false) - const [selected, setSelected] = useState(null) - const [checkedIds, setCheckedIds] = useState>(new Set()) - const { addNode, scanEventTs } = useCanvasStore() - const highlightRef = useRef(null) - - const allChecked = devices.length > 0 && checkedIds.size === devices.length - const someChecked = checkedIds.size > 0 - - const toggleCheck = (id: string, e: React.MouseEvent) => { - e.stopPropagation() - setCheckedIds((prev) => { - const next = new Set(prev) - if (next.has(id)) next.delete(id); else next.add(id) - return next - }) - } - - const toggleAll = () => { - setCheckedIds(allChecked ? new Set() : new Set(devices.map((d) => d.id))) - } - - const load = useCallback(async () => { - setLoading(true) - try { - const res = await scanApi.pending() - setDevices(res.data) - } catch { - toast.error('Failed to load pending devices') - } finally { - setLoading(false) - } - }, []) - - const handleClearAll = async () => { - try { - await scanApi.clearPending() - setDevices([]) - setCheckedIds(new Set()) - toast.success('Pending devices cleared') - } catch { - toast.error('Failed to clear pending devices') - } - } - - const handleBulkApprove = async () => { - const ids = [...checkedIds] - try { - const res = await scanApi.bulkApprove(ids) - const deviceToNode: Record = {} - res.data.device_ids.forEach((did, i) => { deviceToNode[did] = res.data.node_ids[i] }) - const approvedDevices = devices.filter((d) => ids.includes(d.id)) - approvedDevices.forEach((d, i) => { - const nodeId = deviceToNode[d.id] - if (!nodeId) return - addNode({ - id: nodeId, - type: (d.suggested_type ?? 'generic') as import('@/types').NodeType, - position: { x: 400 + (i % 4) * 160, y: 300 + Math.floor(i / 4) * 100 }, - data: { - label: d.hostname ?? d.ip, - type: (d.suggested_type ?? 'generic') as import('@/types').NodeType, - ip: d.ip, - hostname: d.hostname ?? undefined, - status: 'unknown' as const, - services: (d.services ?? []) as import('@/types').ServiceInfo[], - }, - }) - onNodeApproved(nodeId) - }) - setDevices((prev) => prev.filter((d) => !ids.includes(d.id))) - setCheckedIds(new Set()) - toast.success(`Approved ${res.data.approved} device${res.data.approved !== 1 ? 's' : ''}`) - } catch { - toast.error('Failed to bulk approve devices') - } - } - - const handleBulkHide = async () => { - const ids = [...checkedIds] - try { - const res = await scanApi.bulkHide(ids) - setDevices((prev) => prev.filter((d) => !ids.includes(d.id))) - setCheckedIds(new Set()) - toast.success(`Hidden ${res.data.hidden} device${res.data.hidden !== 1 ? 's' : ''}`) - } catch { - toast.error('Failed to bulk hide devices') - } - } - - useEffect(() => { load() }, [load]) - - useEffect(() => { - if (scanEventTs > 0) load() - }, [scanEventTs, load]) - - useEffect(() => { - if (!highlightId || loading) return - highlightRef.current?.scrollIntoView({ behavior: 'smooth', block: 'nearest' }) - }, [highlightId, loading]) - - const handleApprove = async (device: PendingDevice) => { - try { - const nodeData = { - label: device.hostname ?? device.ip, - type: (device.suggested_type ?? 'generic') as import('@/types').NodeType, - ip: device.ip, - hostname: device.hostname ?? undefined, - status: 'unknown', - services: (device.services ?? []) as import('@/types').ServiceInfo[], - } - const res = await scanApi.approve(device.id, nodeData) - const nodeId = res.data.node_id - addNode({ - id: nodeId, - type: nodeData.type, - position: { x: 400, y: 300 }, - data: { ...nodeData, status: 'unknown' as const }, - }) - toast.success(`Approved ${nodeData.label}`) - setDevices((prev) => prev.filter((d) => d.id !== device.id)) - setSelected(null) - onNodeApproved(nodeId) - } catch { - toast.error('Failed to approve device') - } - } - - const handleHide = async (device: PendingDevice) => { - try { - await scanApi.hide(device.id) - setDevices((prev) => prev.filter((d) => d.id !== device.id)) - toast.success('Device hidden') - } catch { - toast.error('Failed to hide device') - } - } - - const handleIgnore = async (device: PendingDevice) => { - try { - await scanApi.ignore(device.id) - setDevices((prev) => prev.filter((d) => d.id !== device.id)) - } catch { - toast.error('Failed to ignore device') - } - } - - return ( - <> -
-
-
- {devices.length > 0 && ( - { if (el) el.indeterminate = someChecked && !allChecked }} - onChange={toggleAll} - className="w-3 h-3 accent-[#00d4ff] cursor-pointer" - title="Select all" - /> - )} - Pending -
-
- - {devices.length > 0 && ( - - )} -
-
- {someChecked && ( -
- - -
- )} - {loading && } - {!loading && devices.length === 0 && ( -

No pending devices

- )} - {devices.map((d) => { - const namedService = d.services.find((s) => s.category != null && s.port != null && !COMMON_PORTS.has(s.port)) - const titleService = namedService - ?? d.services.find((s) => s.port === 80) - ?? d.services.find((s) => s.port === 443) - ?? d.services.find((s) => s.port === 22) - const title = titleService?.service_name ?? d.hostname ?? d.ip - const showIpBelow = title !== d.ip - const hasSsh = d.services.some((s) => s.port === 22) - const hasHttp = d.services.some((s) => s.port === 80) - const hasHttps = d.services.some((s) => s.port === 443) - const otherCount = d.services.filter((s) => s.port !== 22 && s.port !== 80 && s.port !== 443).length - const virtualBadge = detectVirtualBadge(d.mac) - const sourceColor = d.discovery_source === 'mdns' ? '#a855f7' : '#8b949e' - const sourceLabel = d.discovery_source === 'mdns' ? 'mDNS' : d.discovery_source === 'arp' ? 'ARP' : null - const isHighlighted = d.id === highlightId - return ( - - ) - })} -
- - setSelected(null)} - onApprove={handleApprove} - onHide={handleHide} - onIgnore={handleIgnore} - /> - - ) -} - -function HiddenDevicesPanel() { - const [devices, setDevices] = useState([]) - const [loading, setLoading] = useState(false) - - const load = useCallback(async () => { - setLoading(true) - try { - const res = await scanApi.hidden() - setDevices(res.data) - } catch { - toast.error('Failed to load hidden devices') - } finally { - setLoading(false) - } - }, []) - - useEffect(() => { load() }, [load]) - - const handleIgnore = async (id: string) => { - try { - await scanApi.ignore(id) - setDevices((prev) => prev.filter((d) => d.id !== id)) - } catch { - toast.error('Failed to remove device') - } - } - - return ( -
-
- Hidden - -
- {loading && } - {!loading && devices.length === 0 && ( -

No hidden devices

- )} - {devices.map((d) => ( -
-
{d.ip}
- {d.hostname &&
{d.hostname}
} -
- handleIgnore(d.id)} /> -
-
- ))} -
- ) -} function ScanHistoryPanel() { const [runs, setRuns] = useState([]) @@ -507,12 +198,19 @@ function ScanHistoryPanel() { const res = await scanApi.runs() const next: ScanRun[] = res.data - // Toast when a run transitions from running → error + // Surface transitions and refresh dependent UI for (const run of next) { const prev = prevRunsRef.current.find((r) => r.id === run.id) if (prev?.status === 'running' && run.status === 'error') { toast.error(`Scan failed: ${run.error ?? 'unknown error'}`) } + if (prev?.status === 'running' && run.status === 'done') { + if (run.kind === 'zigbee') { + toast.success(`Zigbee import done — ${run.devices_found} device${run.devices_found !== 1 ? 's' : ''}`) + } + // Notify pending modal/canvas to refresh + useCanvasStore.getState().notifyScanDeviceFound() + } } prevRunsRef.current = next setRuns(next) @@ -573,6 +271,14 @@ function ScanHistoryPanel() { {r.status} {r.status === 'running' && } + + {r.kind === 'zigbee' ? 'ZIG' : 'IP'} + {r.devices_found} found {r.status === 'running' && ( @@ -693,55 +399,6 @@ function VersionBadge() { ) } -const MAC_OUI: Record = { - '52:54:00': { label: 'QEMU', title: 'QEMU/KVM Virtual Machine' }, - 'bc:24:11': { label: 'PVE', title: 'Proxmox Virtual Machine or LXC' }, - '00:50:56': { label: 'VMware', title: 'VMware Virtual Machine' }, - '00:0c:29': { label: 'VMware', title: 'VMware Virtual Machine' }, - '08:00:27': { label: 'VBox', title: 'VirtualBox Virtual Machine' }, - '00:15:5d': { label: 'Hyper-V', title: 'Hyper-V Virtual Machine' }, -} - -function detectVirtualBadge(mac: string | null) { - if (!mac) return null - return MAC_OUI[mac.toLowerCase().slice(0, 8)] ?? null -} - -function ServiceBadge({ label, color }: { label: string; color: string }) { - return ( - - {label} - - ) -} - -interface ActionButtonProps { - icon: React.ElementType - label: string - color?: 'green' | 'red' - onClick: () => void -} - -function ActionButton({ icon: Icon, label, color, onClick }: ActionButtonProps) { - const colorClass = - color === 'green' ? 'text-[#39d353] hover:bg-[#39d353]/10' : - color === 'red' ? 'text-[#f85149] hover:bg-[#f85149]/10' : - 'text-muted-foreground hover:text-foreground hover:bg-[#30363d]' - return ( - - - - - {label} - - ) -} - interface SidebarItemProps { icon: React.ElementType label: string diff --git a/frontend/src/components/panels/__tests__/Sidebar.test.tsx b/frontend/src/components/panels/__tests__/Sidebar.test.tsx index 5488f1a..91bacea 100644 --- a/frontend/src/components/panels/__tests__/Sidebar.test.tsx +++ b/frontend/src/components/panels/__tests__/Sidebar.test.tsx @@ -11,22 +11,11 @@ import type { NodeData } from '@/types' vi.mock('@/stores/canvasStore') vi.mock('@/stores/authStore') -const mockBulkApprove = vi.fn() -const mockBulkHide = vi.fn() - vi.mock('@/api/client', () => ({ scanApi: { trigger: vi.fn().mockResolvedValue({}), - pending: vi.fn().mockResolvedValue({ data: [] }), - hidden: vi.fn().mockResolvedValue({ data: [] }), runs: vi.fn().mockResolvedValue({ data: [] }), stop: vi.fn().mockResolvedValue({}), - clearPending: vi.fn().mockResolvedValue({}), - approve: vi.fn().mockResolvedValue({ data: { approved: true, node_id: 'new-node-1' } }), - hide: vi.fn().mockResolvedValue({ data: { hidden: true } }), - ignore: vi.fn().mockResolvedValue({ data: { ignored: true } }), - bulkApprove: (...args: unknown[]) => mockBulkApprove(...args), - bulkHide: (...args: unknown[]) => mockBulkHide(...args), }, settingsApi: { get: vi.fn().mockResolvedValue({ data: { interval_seconds: 60 } }), @@ -48,10 +37,6 @@ vi.mock('@/components/ui/tooltip', () => ({ TooltipContent: () => null, })) -vi.mock('@/components/modals/PendingDeviceModal', () => ({ - PendingDeviceModal: () => null, -})) - // ── Helpers ─────────────────────────────────────────────────────────────────── const makeNode = (id: string, status: NodeData['status'], type: NodeData['type'] = 'server'): Node => ({ @@ -86,8 +71,9 @@ const defaultProps = { onAddNode: vi.fn(), onAddGroupRect: vi.fn(), onScan: vi.fn(), + onZigbeeImport: vi.fn(), onSave: vi.fn(), - onNodeApproved: vi.fn(), + onOpenPending: vi.fn(), } // ── Tests ───────────────────────────────────────────────────────────────────── @@ -129,26 +115,22 @@ describe('Sidebar', () => { ], }) render() - // Total (excludes groupRect) expect(screen.getByText('4')).toBeInTheDocument() - // Online expect(screen.getByText('2')).toBeInTheDocument() - // Offline expect(screen.getByText('1')).toBeInTheDocument() }) it('excludes groupRect nodes from stats', () => { mockStore({ nodes: [ - makeNode('n1', 'unknown'), // 1 real node, not online/offline + makeNode('n1', 'unknown'), makeNode('zone', 'unknown', 'groupRect'), ], }) render() - // Total row shows 1 (groupRect excluded), online/offline both 0 const totalRow = screen.getByText('Total').closest('div')! expect(totalRow).toHaveTextContent('1') - expect(screen.getAllByText('0')).toHaveLength(2) // online=0, offline=0 + expect(screen.getAllByText('0')).toHaveLength(2) }) // ── Collapse ─────────────────────────────────────────────────────────────── @@ -225,7 +207,6 @@ describe('Sidebar', () => { it('shows unsaved badge dot on Save Canvas when hasUnsavedChanges', () => { mockStore({ hasUnsavedChanges: true }) render() - // The badge is a span sibling of the Save Canvas button icon const saveBtn = screen.getByText('Save Canvas').closest('button')! const badge = saveBtn.querySelector('span.rounded-full') expect(badge).toBeInTheDocument() @@ -241,24 +222,24 @@ describe('Sidebar', () => { // ── Scan action ──────────────────────────────────────────────────────────── - it('calls onScan prop when Scan Network is clicked (scan trigger moved to ScanConfigModal)', () => { + it('calls onScan prop when Scan Network is clicked', () => { render() fireEvent.click(screen.getByText('Scan Network')) expect(defaultProps.onScan).toHaveBeenCalledOnce() }) - // ── Navigation ───────────────────────────────────────────────────────────── + // ── Pending / Hidden open modal ──────────────────────────────────────────── - it('shows Pending panel when Pending Devices nav item is clicked', async () => { + it('calls onOpenPending with pending status when Pending Devices is clicked', () => { render() fireEvent.click(screen.getByText('Pending Devices')) - await waitFor(() => expect(screen.getByText('No pending devices')).toBeInTheDocument()) + expect(defaultProps.onOpenPending).toHaveBeenCalledWith(undefined, 'pending') }) - it('shows Hidden panel when Hidden Devices nav item is clicked', async () => { + it('calls onOpenPending with hidden status when Hidden Devices is clicked', () => { render() fireEvent.click(screen.getByText('Hidden Devices')) - await waitFor(() => expect(screen.getByText('No hidden devices')).toBeInTheDocument()) + expect(defaultProps.onOpenPending).toHaveBeenCalledWith(undefined, 'hidden') }) it('shows History panel when Scan History nav item is clicked', async () => { @@ -267,16 +248,13 @@ describe('Sidebar', () => { await waitFor(() => expect(screen.getByText('No scans yet')).toBeInTheDocument()) }) - // Regression: forceView used to override local state on every render, freezing - // the sidebar on whichever view the parent forced (e.g. 'history' after a scan). + // Regression: forceView must not freeze local state across rerenders. it('allows switching views after forceView is set by parent', async () => { const { rerender } = render() await waitFor(() => expect(screen.getByText('No scans yet')).toBeInTheDocument()) - // Parent keeps forceView as 'history'; user clicks another nav item. rerender() - fireEvent.click(screen.getByText('Pending Devices')) - await waitFor(() => expect(screen.getByText('No pending devices')).toBeInTheDocument()) - expect(screen.queryByText('No scans yet')).not.toBeInTheDocument() + fireEvent.click(screen.getByText('Canvas')) + await waitFor(() => expect(screen.queryByText('No scans yet')).not.toBeInTheDocument()) }) it('toggles Settings panel on Settings click', async () => { @@ -285,7 +263,6 @@ describe('Sidebar', () => { await waitFor(() => expect(screen.getByText('Status check interval (s)')).toBeInTheDocument(), ) - // Click the nav button again to close (use role to avoid matching the panel heading) fireEvent.click(screen.getByRole('button', { name: 'Settings' })) expect(screen.queryByText('Status check interval (s)')).not.toBeInTheDocument() }) @@ -303,101 +280,3 @@ describe('Sidebar', () => { expect(mockLogout).toHaveBeenCalledOnce() }) }) - -// ── PendingDevicesPanel — bulk select ───────────────────────────────────────── - -const DEVICE_A = { - id: 'dev-a', - ip: '192.168.1.10', - hostname: 'host-a', - mac: null, - os: null, - services: [], - suggested_type: 'generic', - status: 'pending', - discovery_source: 'arp', -} - -const DEVICE_B = { - id: 'dev-b', - ip: '192.168.1.11', - hostname: 'host-b', - mac: null, - os: null, - services: [], - suggested_type: 'generic', - status: 'pending', - discovery_source: 'arp', -} - -describe('PendingDevicesPanel — bulk select', () => { - beforeEach(() => { - mockStore() - mockAuth() - vi.clearAllMocks() - mockBulkApprove.mockResolvedValue({ - data: { approved: 2, node_ids: ['n1', 'n2'], device_ids: ['dev-a', 'dev-b'], skipped: 0 }, - }) - mockBulkHide.mockResolvedValue({ data: { hidden: 2, skipped: 0 } }) - }) - - async function renderWithDevices() { - const { scanApi } = await import('@/api/client') - vi.mocked(scanApi.pending).mockResolvedValue({ data: [DEVICE_A, DEVICE_B] } as never) - render() - await waitFor(() => expect(screen.getByText('host-a')).toBeInTheDocument()) - } - - it('renders checkboxes for each device', async () => { - await renderWithDevices() - const checkboxes = screen.getAllByRole('checkbox') - // select-all + 2 device checkboxes - expect(checkboxes.length).toBe(3) - }) - - it('shows bulk action bar when a device is checked', async () => { - await renderWithDevices() - const [, firstDeviceCheckbox] = screen.getAllByRole('checkbox') - fireEvent.click(firstDeviceCheckbox) - await waitFor(() => expect(screen.getByText(/Approve \(1\)/)).toBeInTheDocument()) - expect(screen.getByText(/Hide \(1\)/)).toBeInTheDocument() - }) - - it('hides bulk action bar when no device is checked', async () => { - await renderWithDevices() - expect(screen.queryByText(/Approve \(/)).not.toBeInTheDocument() - }) - - it('select-all checks all devices', async () => { - await renderWithDevices() - const [selectAll] = screen.getAllByRole('checkbox') - fireEvent.click(selectAll) - await waitFor(() => expect(screen.getByText(/Approve \(2\)/)).toBeInTheDocument()) - }) - - it('select-all unchecks all when all are selected', async () => { - await renderWithDevices() - const [selectAll] = screen.getAllByRole('checkbox') - fireEvent.click(selectAll) // select all - fireEvent.click(selectAll) // deselect all - await waitFor(() => expect(screen.queryByText(/Approve \(/)).not.toBeInTheDocument()) - }) - - it('calls bulkApprove with checked ids and removes devices from list', async () => { - await renderWithDevices() - const [selectAll] = screen.getAllByRole('checkbox') - fireEvent.click(selectAll) - fireEvent.click(screen.getByText(/Approve \(2\)/)) - await waitFor(() => expect(mockBulkApprove).toHaveBeenCalledWith(['dev-a', 'dev-b'])) - await waitFor(() => expect(screen.queryByText('host-a')).not.toBeInTheDocument()) - }) - - it('calls bulkHide with checked ids and removes devices from list', async () => { - await renderWithDevices() - const [selectAll] = screen.getAllByRole('checkbox') - fireEvent.click(selectAll) - fireEvent.click(screen.getByText(/Hide \(2\)/)) - await waitFor(() => expect(mockBulkHide).toHaveBeenCalledWith(['dev-a', 'dev-b'])) - await waitFor(() => expect(screen.queryByText('host-b')).not.toBeInTheDocument()) - }) -}) diff --git a/frontend/src/components/zigbee/ZigbeeImportModal.tsx b/frontend/src/components/zigbee/ZigbeeImportModal.tsx new file mode 100644 index 0000000..9c452ee --- /dev/null +++ b/frontend/src/components/zigbee/ZigbeeImportModal.tsx @@ -0,0 +1,450 @@ +import { useState } from 'react' +import { Network, Router, Cpu, CheckCircle2, XCircle, Loader2, Plus } from 'lucide-react' +import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogFooter } from '@/components/ui/dialog' +import { Button } from '@/components/ui/button' +import { Input } from '@/components/ui/input' +import { Label } from '@/components/ui/label' +import { zigbeeApi } from '@/api/client' +import { toast } from 'sonner' +import type { ZigbeeNode, ZigbeeEdge } from './types' + +interface ZigbeeImportModalProps { + open: boolean + onClose: () => void + onAddToCanvas: (nodes: ZigbeeNode[], edges: ZigbeeEdge[]) => void + onPendingImported?: ( + coordinator?: { id: string; label: string; ieee_address: string } | null, + ) => void +} + +type ImportMode = 'pending' | 'canvas' + +interface ConnectionForm { + mqtt_host: string + mqtt_port: string + mqtt_username: string + mqtt_password: string + base_topic: string + mqtt_tls: boolean + mqtt_tls_insecure: boolean + port_user_edited: boolean +} + +const DEFAULT_FORM: ConnectionForm = { + mqtt_host: '', + mqtt_port: '1883', + mqtt_username: '', + mqtt_password: '', + base_topic: 'zigbee2mqtt', + mqtt_tls: false, + mqtt_tls_insecure: false, + port_user_edited: false, +} + +const DEVICE_TYPE_ICON = { + zigbee_coordinator: Network, + zigbee_router: Router, + zigbee_enddevice: Cpu, +} as const + +const DEVICE_TYPE_LABEL = { + zigbee_coordinator: 'Coordinator', + zigbee_router: 'Router', + zigbee_enddevice: 'End Device', +} as const + +const DEVICE_TYPE_COLOR = { + zigbee_coordinator: '#00d4ff', + zigbee_router: '#39d353', + zigbee_enddevice: '#e3b341', +} as const + +export function ZigbeeImportModal({ open, onClose, onAddToCanvas, onPendingImported }: ZigbeeImportModalProps) { + const [form, setForm] = useState(DEFAULT_FORM) + const [connectionStatus, setConnectionStatus] = useState<'idle' | 'testing' | 'ok' | 'fail'>('idle') + const [connectionMsg, setConnectionMsg] = useState('') + const [loading, setLoading] = useState(false) + const [devices, setDevices] = useState([]) + const [edges, setEdges] = useState([]) + const [checked, setChecked] = useState>(new Set()) + const [importMode, setImportMode] = useState('pending') + + const updateField = (field: keyof ConnectionForm, value: string) => + setForm((f) => ({ + ...f, + [field]: value, + ...(field === 'mqtt_port' ? { port_user_edited: true } : {}), + })) + + const toggleTls = (next: boolean) => + setForm((f) => { + const port = f.port_user_edited + ? f.mqtt_port + : next + ? '8883' + : '1883' + return { + ...f, + mqtt_tls: next, + mqtt_tls_insecure: next ? f.mqtt_tls_insecure : false, + mqtt_port: port, + } + }) + + const buildPayload = () => ({ + mqtt_host: form.mqtt_host.trim(), + mqtt_port: Number(form.mqtt_port) || (form.mqtt_tls ? 8883 : 1883), + mqtt_username: form.mqtt_username.trim() || undefined, + mqtt_password: form.mqtt_password || undefined, + base_topic: form.base_topic.trim() || 'zigbee2mqtt', + mqtt_tls: form.mqtt_tls, + mqtt_tls_insecure: form.mqtt_tls_insecure, + }) + + const handleTestConnection = async () => { + if (!form.mqtt_host.trim()) { toast.error('Enter a broker hostname'); return } + setConnectionStatus('testing') + try { + const res = await zigbeeApi.testConnection({ + mqtt_host: form.mqtt_host.trim(), + mqtt_port: Number(form.mqtt_port) || (form.mqtt_tls ? 8883 : 1883), + mqtt_username: form.mqtt_username.trim() || undefined, + mqtt_password: form.mqtt_password || undefined, + mqtt_tls: form.mqtt_tls, + mqtt_tls_insecure: form.mqtt_tls_insecure, + }) + if (res.data.connected) { + setConnectionStatus('ok') + setConnectionMsg(res.data.message) + } else { + setConnectionStatus('fail') + setConnectionMsg(res.data.message) + } + } catch { + setConnectionStatus('fail') + setConnectionMsg('Request failed — check broker address') + } + } + + const extractError = (err: unknown): string | undefined => { + if (err && typeof err === 'object' && 'response' in err) { + return (err as { response?: { data?: { detail?: string } } }).response?.data?.detail + } + return undefined + } + + const handleFetchDevices = async () => { + if (!form.mqtt_host.trim()) { toast.error('Enter a broker hostname'); return } + setLoading(true) + try { + if (importMode === 'pending') { + await zigbeeApi.importToPending(buildPayload()) + toast.success('Zigbee import started — track progress in Scan History') + onPendingImported?.(null) + handleClose() + } else { + const res = await zigbeeApi.importNetwork(buildPayload()) + setDevices(res.data.nodes) + setEdges(res.data.edges) + setChecked(new Set(res.data.nodes.map((n) => n.id))) + if (res.data.device_count === 0) { + toast.info('No Zigbee devices found in the network map') + } else { + toast.success(`Found ${res.data.device_count} device${res.data.device_count !== 1 ? 's' : ''}`) + } + } + } catch (err: unknown) { + toast.error(extractError(err) ?? 'Failed to fetch Zigbee devices') + } finally { + setLoading(false) + } + } + + const toggleCheck = (id: string) => + setChecked((prev) => { + const next = new Set(prev) + if (next.has(id)) next.delete(id); else next.add(id) + return next + }) + + const toggleAll = () => { + setChecked(checked.size === devices.length ? new Set() : new Set(devices.map((d) => d.id))) + } + + const handleAddToCanvas = () => { + const selectedDevices = devices.filter((d) => checked.has(d.id)) + const selectedIds = new Set(selectedDevices.map((d) => d.id)) + const selectedEdges = edges.filter((e) => selectedIds.has(e.source) && selectedIds.has(e.target)) + onAddToCanvas(selectedDevices, selectedEdges) + toast.success(`Added ${selectedDevices.length} device${selectedDevices.length !== 1 ? 's' : ''} to canvas`) + onClose() + } + + const handleClose = () => { + setDevices([]) + setEdges([]) + setChecked(new Set()) + setConnectionStatus('idle') + setConnectionMsg('') + setImportMode('pending') + onClose() + } + + const groupedDevices = { + zigbee_coordinator: devices.filter((d) => d.type === 'zigbee_coordinator'), + zigbee_router: devices.filter((d) => d.type === 'zigbee_router'), + zigbee_enddevice: devices.filter((d) => d.type === 'zigbee_enddevice'), + } as const + + return ( + !v && handleClose()}> + + + + + Zigbee2MQTT Import + + + +
+ {/* Connection Form */} +
+
+
+ + updateField('mqtt_host', e.target.value)} + placeholder="192.168.1.x or mqtt.local" + className="font-mono text-sm bg-[#0d1117] border-border" + /> +
+
+ + updateField('mqtt_port', e.target.value)} + placeholder="1883" + type="number" + className="font-mono text-sm bg-[#0d1117] border-border" + /> +
+
+ + updateField('base_topic', e.target.value)} + placeholder="zigbee2mqtt" + className="font-mono text-sm bg-[#0d1117] border-border" + /> +
+
+ + updateField('mqtt_username', e.target.value)} + placeholder="mqtt_user" + className="text-sm bg-[#0d1117] border-border" + /> +
+
+ + updateField('mqtt_password', e.target.value)} + placeholder="••••••••" + type="password" + autoComplete="new-password" + className="text-sm bg-[#0d1117] border-border" + /> +
+
+ + +
+
+ + {/* Connection status indicator */} + {connectionStatus !== 'idle' && ( +
+ {connectionStatus === 'testing' && } + {connectionStatus === 'ok' && } + {connectionStatus === 'fail' && } + {connectionStatus === 'testing' ? 'Testing…' : connectionMsg} +
+ )} + +
+ Send devices to: + + +
+
+ + +
+

+ Fetching the network map can take several minutes on large meshes. +

+
+ + {/* Device List */} + {devices.length > 0 && ( +
+
+
+ { if (el) el.indeterminate = checked.size > 0 && checked.size < devices.length }} + onChange={toggleAll} + className="w-3 h-3 accent-[#00d4ff] cursor-pointer" + title="Select all" + /> + + Devices ({checked.size}/{devices.length} selected) + +
+
+ + {(Object.entries(groupedDevices) as [keyof typeof groupedDevices, ZigbeeNode[]][]) + .filter(([, group]) => group.length > 0) + .map(([type, group]) => { + const Icon = DEVICE_TYPE_ICON[type] + const color = DEVICE_TYPE_COLOR[type] + return ( +
+
+ + + {DEVICE_TYPE_LABEL[type]} ({group.length}) + +
+ {group.map((device) => ( +
toggleCheck(device.id)} + > + toggleCheck(device.id)} + onClick={(e) => e.stopPropagation()} + className="w-3 h-3 mt-0.5 accent-[#00d4ff] cursor-pointer shrink-0" + /> +
+
{device.friendly_name}
+
{device.ieee_address}
+ {(device.model || device.vendor) && ( +
+ {[device.vendor, device.model].filter(Boolean).join(' · ')} +
+ )} +
+ {device.lqi != null && ( + + LQI {device.lqi} + + )} +
+ ))} +
+ ) + })} +
+ )} +
+ + + + {devices.length > 0 && ( + + )} + +
+
+ ) +} diff --git a/frontend/src/components/zigbee/__tests__/ZigbeeImportModal.test.tsx b/frontend/src/components/zigbee/__tests__/ZigbeeImportModal.test.tsx new file mode 100644 index 0000000..c7f2088 --- /dev/null +++ b/frontend/src/components/zigbee/__tests__/ZigbeeImportModal.test.tsx @@ -0,0 +1,223 @@ +import { describe, it, expect, vi, beforeEach } from 'vitest' +import { render, screen, fireEvent, waitFor } from '@testing-library/react' +import { ZigbeeImportModal } from '../ZigbeeImportModal' + +vi.mock('@/api/client', () => ({ + zigbeeApi: { + testConnection: vi.fn(), + importNetwork: vi.fn(), + importToPending: vi.fn(), + }, +})) +vi.mock('sonner', () => ({ toast: { success: vi.fn(), error: vi.fn(), info: vi.fn() } })) + +import { zigbeeApi } from '@/api/client' +import { toast } from 'sonner' + +const defaultProps = { + open: true, + onClose: vi.fn(), + onAddToCanvas: vi.fn(), +} + +const sampleNodes = [ + { + id: '0x0000', + label: 'Coordinator', + type: 'zigbee_coordinator' as const, + ieee_address: '0x0000', + friendly_name: 'Coordinator', + device_type: 'Coordinator', + model: null, + vendor: null, + lqi: null, + parent_id: null, + }, + { + id: '0x0001', + label: 'router_1', + type: 'zigbee_router' as const, + ieee_address: '0x0001', + friendly_name: 'router_1', + device_type: 'Router', + model: 'CC2530', + vendor: 'TI', + lqi: 200, + parent_id: '0x0000', + }, +] + +describe('ZigbeeImportModal', () => { + beforeEach(() => { + vi.mocked(zigbeeApi.testConnection).mockReset() + vi.mocked(zigbeeApi.importNetwork).mockReset() + vi.mocked(zigbeeApi.importToPending).mockReset() + vi.mocked(toast.success).mockReset() + vi.mocked(toast.error).mockReset() + vi.mocked(toast.info).mockReset() + defaultProps.onClose.mockReset() + defaultProps.onAddToCanvas.mockReset() + }) + + it('renders nothing when closed', () => { + const { container } = render() + expect(container.querySelector('[role="dialog"]')).toBeNull() + }) + + it('renders the modal with form fields when open', () => { + render() + expect(screen.getByText('Zigbee2MQTT Import')).toBeDefined() + expect(screen.getByPlaceholderText('192.168.1.x or mqtt.local')).toBeDefined() + expect(screen.getByPlaceholderText('1883')).toBeDefined() + expect(screen.getByPlaceholderText('zigbee2mqtt')).toBeDefined() + }) + + it('shows error toast when testing connection without a host', async () => { + render() + fireEvent.click(screen.getByRole('button', { name: /test connection/i })) + await waitFor(() => { + expect(toast.error).toHaveBeenCalledWith('Enter a broker hostname') + }) + expect(zigbeeApi.testConnection).not.toHaveBeenCalled() + }) + + it('shows success status when connection test passes', async () => { + vi.mocked(zigbeeApi.testConnection).mockResolvedValue({ + data: { connected: true, message: 'Connection successful' }, + } as never) + + render() + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '192.168.1.100' } }) + fireEvent.click(screen.getByRole('button', { name: /test connection/i })) + + await waitFor(() => { + expect(screen.getByText('Connection successful')).toBeDefined() + }) + }) + + it('shows failure status when connection test fails', async () => { + vi.mocked(zigbeeApi.testConnection).mockResolvedValue({ + data: { connected: false, message: 'Connection refused' }, + } as never) + + render() + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '10.0.0.1' } }) + fireEvent.click(screen.getByRole('button', { name: /test connection/i })) + + await waitFor(() => { + expect(screen.getByText('Connection refused')).toBeDefined() + }) + }) + + const selectCanvasMode = () => { + fireEvent.click(screen.getByRole('radio', { name: /canvas directly/i })) + } + + it('fetches devices and renders them grouped by type', async () => { + vi.mocked(zigbeeApi.importNetwork).mockResolvedValue({ + data: { nodes: sampleNodes, edges: [], device_count: 2 }, + } as never) + + render() + selectCanvasMode() + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '192.168.1.100' } }) + fireEvent.click(screen.getByRole('button', { name: /fetch devices/i })) + + await waitFor(() => { + expect(screen.getByText('Coordinator')).toBeDefined() + expect(screen.getByText('router_1')).toBeDefined() + }) + expect(toast.success).toHaveBeenCalledWith('Found 2 devices') + }) + + it('shows info toast when no devices found', async () => { + vi.mocked(zigbeeApi.importNetwork).mockResolvedValue({ + data: { nodes: [], edges: [], device_count: 0 }, + } as never) + + render() + selectCanvasMode() + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '192.168.1.100' } }) + fireEvent.click(screen.getByRole('button', { name: /fetch devices/i })) + + await waitFor(() => { + expect(toast.info).toHaveBeenCalledWith('No Zigbee devices found in the network map') + }) + }) + + it('calls onAddToCanvas with selected devices and closes modal', async () => { + vi.mocked(zigbeeApi.importNetwork).mockResolvedValue({ + data: { nodes: sampleNodes, edges: [{ source: '0x0000', target: '0x0001' }], device_count: 2 }, + } as never) + + render() + selectCanvasMode() + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '192.168.1.100' } }) + fireEvent.click(screen.getByRole('button', { name: /fetch devices/i })) + + await waitFor(() => screen.getByText('Coordinator')) + + // Click "Add N to Canvas" button + const addBtn = screen.getByRole('button', { name: /add.*canvas/i }) + fireEvent.click(addBtn) + + await waitFor(() => { + expect(defaultProps.onAddToCanvas).toHaveBeenCalledOnce() + expect(defaultProps.onClose).toHaveBeenCalledOnce() + }) + }) + + it('calls onClose when Cancel is clicked', () => { + render() + fireEvent.click(screen.getByRole('button', { name: 'Cancel' })) + expect(defaultProps.onClose).toHaveBeenCalledOnce() + }) + + it('imports to pending by default and notifies parent', async () => { + vi.mocked(zigbeeApi.importToPending).mockResolvedValue({ + data: { + id: 'run-1', + status: 'running', + kind: 'zigbee', + ranges: ['192.168.1.100:1883'], + devices_found: 0, + started_at: '2026-01-01T00:00:00Z', + finished_at: null, + error: null, + }, + } as never) + const onPendingImported = vi.fn() + + render() + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '192.168.1.100' } }) + fireEvent.click(screen.getByRole('button', { name: /import to pending/i })) + + await waitFor(() => { + expect(zigbeeApi.importToPending).toHaveBeenCalled() + expect(onPendingImported).toHaveBeenCalled() + expect(defaultProps.onClose).toHaveBeenCalled() + }) + expect(zigbeeApi.importNetwork).not.toHaveBeenCalled() + }) + + it('switching to canvas mode calls importNetwork and not importToPending', async () => { + vi.mocked(zigbeeApi.importNetwork).mockResolvedValue({ + data: { nodes: sampleNodes, edges: [], device_count: 2 }, + } as never) + + render() + fireEvent.click(screen.getByRole('radio', { name: /canvas directly/i })) + const hostInput = screen.getByPlaceholderText('192.168.1.x or mqtt.local') + fireEvent.change(hostInput, { target: { value: '192.168.1.100' } }) + fireEvent.click(screen.getByRole('button', { name: /fetch devices/i })) + + await waitFor(() => expect(zigbeeApi.importNetwork).toHaveBeenCalled()) + expect(zigbeeApi.importToPending).not.toHaveBeenCalled() + }) +}) diff --git a/frontend/src/components/zigbee/types.ts b/frontend/src/components/zigbee/types.ts new file mode 100644 index 0000000..c730bc6 --- /dev/null +++ b/frontend/src/components/zigbee/types.ts @@ -0,0 +1,37 @@ +/** Shared Zigbee type definitions for the frontend. */ + +export interface ZigbeeNode { + id: string + label: string + type: 'zigbee_coordinator' | 'zigbee_router' | 'zigbee_enddevice' + ieee_address: string + friendly_name: string + device_type: string + model?: string | null + vendor?: string | null + lqi?: number | null + parent_id?: string | null +} + +export interface ZigbeeEdge { + source: string + target: string +} + +export interface ZigbeeImportResponse { + nodes: ZigbeeNode[] + edges: ZigbeeEdge[] + device_count: number +} + +export interface ZigbeeTestConnectionRequest { + mqtt_host: string + mqtt_port: number + mqtt_username?: string + mqtt_password?: string +} + +export interface ZigbeeTestConnectionResponse { + connected: boolean + message: string +} diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 7b326fd..76291a8 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -19,6 +19,9 @@ export type NodeType = | 'generic' | 'groupRect' | 'group' + | 'zigbee_coordinator' + | 'zigbee_router' + | 'zigbee_enddevice' export type TextPosition = | 'top-left' @@ -136,6 +139,9 @@ export const NODE_TYPE_LABELS: Record = { generic: 'Generic Device', groupRect: 'Group Rectangle', group: 'Node Group', + zigbee_coordinator: 'Zigbee Coordinator', + zigbee_router: 'Zigbee Router', + zigbee_enddevice: 'Zigbee End Device', } export const STATUS_COLORS: Record = { diff --git a/frontend/src/utils/nodeIcons.ts b/frontend/src/utils/nodeIcons.ts index 5c329c8..d57f11f 100644 --- a/frontend/src/utils/nodeIcons.ts +++ b/frontend/src/utils/nodeIcons.ts @@ -136,6 +136,9 @@ export const NODE_TYPE_DEFAULT_ICONS: Record = { cpl: PlugZap, docker_host: Anchor, docker_container: Package, + zigbee_coordinator: Radio, + zigbee_router: Zap, + zigbee_enddevice: Lightbulb, generic: Circle, group: Circle, groupRect: Circle, diff --git a/frontend/src/utils/themes.ts b/frontend/src/utils/themes.ts index 202c0dd..768a583 100644 --- a/frontend/src/utils/themes.ts +++ b/frontend/src/utils/themes.ts @@ -59,6 +59,9 @@ export const THEMES: Record = { cpl: { border: '#e3b341', icon: '#e3b341' }, docker_host: { border: '#2496ED', icon: '#2496ED' }, docker_container: { border: '#0ea5e9', icon: '#0ea5e9' }, + zigbee_coordinator:{ border: '#ff6e00', icon: '#ff6e00' }, + zigbee_router: { border: '#e3b341', icon: '#e3b341' }, + zigbee_enddevice: { border: '#a855f7', icon: '#a855f7' }, generic: { border: '#8b949e', icon: '#8b949e' }, groupRect: { border: '#00d4ff', icon: '#00d4ff' }, group: { border: '#00d4ff', icon: '#00d4ff' }, @@ -116,6 +119,9 @@ export const THEMES: Record = { cpl: { border: '#fbbf24', icon: '#fbbf24' }, docker_host: { border: '#2496ED', icon: '#2496ED' }, docker_container: { border: '#38bdf8', icon: '#38bdf8' }, + zigbee_coordinator:{ border: '#fb923c', icon: '#fb923c' }, + zigbee_router: { border: '#fbbf24', icon: '#fbbf24' }, + zigbee_enddevice: { border: '#c084fc', icon: '#c084fc' }, generic: { border: '#94a3b8', icon: '#94a3b8' }, groupRect: { border: '#22d3ee', icon: '#22d3ee' }, group: { border: '#22d3ee', icon: '#22d3ee' }, @@ -173,6 +179,9 @@ export const THEMES: Record = { cpl: { border: '#b45309', icon: '#b45309' }, docker_host: { border: '#2496ED', icon: '#2496ED' }, docker_container: { border: '#0369a1', icon: '#0369a1' }, + zigbee_coordinator:{ border: '#ea580c', icon: '#ea580c' }, + zigbee_router: { border: '#b45309', icon: '#b45309' }, + zigbee_enddevice: { border: '#7c3aed', icon: '#7c3aed' }, generic: { border: '#6b7280', icon: '#6b7280' }, groupRect: { border: '#0284c7', icon: '#0284c7' }, group: { border: '#0284c7', icon: '#0284c7' }, @@ -230,6 +239,9 @@ export const THEMES: Record = { cpl: { border: '#ffff00', icon: '#ffff00' }, docker_host: { border: '#00aaff', icon: '#00aaff' }, docker_container: { border: '#00ddff', icon: '#00ddff' }, + zigbee_coordinator:{ border: '#ff8800', icon: '#ff8800' }, + zigbee_router: { border: '#ffff00', icon: '#ffff00' }, + zigbee_enddevice: { border: '#ff00ff', icon: '#ff00ff' }, generic: { border: '#8888ff', icon: '#8888ff' }, groupRect: { border: '#00ffff', icon: '#00ffff' }, group: { border: '#00ffff', icon: '#00ffff' }, @@ -287,6 +299,9 @@ export const THEMES: Record = { cpl: { border: '#66ff33', icon: '#66ff33' }, docker_host: { border: '#00cc88', icon: '#00cc88' }, docker_container: { border: '#00aacc', icon: '#00aacc' }, + zigbee_coordinator:{ border: '#33ff66', icon: '#33ff66' }, + zigbee_router: { border: '#66ff33', icon: '#66ff33' }, + zigbee_enddevice: { border: '#008822', icon: '#008822' }, generic: { border: '#006600', icon: '#006600' }, groupRect: { border: '#00ff41', icon: '#00ff41' }, group: { border: '#00ff41', icon: '#00ff41' }, @@ -344,6 +359,9 @@ export const THEMES: Record = { cpl: { border: '#e3b341', icon: '#e3b341' }, docker_host: { border: '#2496ED', icon: '#2496ED' }, docker_container: { border: '#0ea5e9', icon: '#0ea5e9' }, + zigbee_coordinator:{ border: '#ff6e00', icon: '#ff6e00' }, + zigbee_router: { border: '#e3b341', icon: '#e3b341' }, + zigbee_enddevice: { border: '#a855f7', icon: '#a855f7' }, generic: { border: '#8b949e', icon: '#8b949e' }, groupRect: { border: '#00d4ff', icon: '#00d4ff' }, group: { border: '#00d4ff', icon: '#00d4ff' },