From 38c5bcb60626e98a6e229e8c7f439fc567d2b714 Mon Sep 17 00:00:00 2001 From: Pouzor Date: Sat, 4 Apr 2026 23:05:14 +0200 Subject: [PATCH] feat: configurable bottom handles, scanner rewrite, UI polish MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## New features - Configurable bottom connection points per node (1–4 handles) - Fit view on load - LiveView improvements - Node modal: inline Type/Icon picker, default icon in trigger - Remove redundant Save button from ScanConfigModal ## Scanner fixes - Phase 1: replace nmap ARP sweep with concurrent asyncio ping sweep (50 parallel pings, 1s timeout). Zero false positives, works in any Docker network mode. Supplements with /proc/net/arp for ICMP-blocked devices. - Phase 2: explicit -sS (root) / -sT (non-root) scan type; bump host-timeout to 60s; gather(return_exceptions=True) so one failing host doesn't abort the batch - Fix 404 on missing device in hide/ignore - Validate CIDR ranges to prevent nmap injection - Thread-safe cancel set, pre-fetch canvas/hidden IPs (no N+1 queries) - Logging: attach StreamHandler to root logger so app.* logs are visible ## Tests - 21 backend scanner tests (ping sweep, ARP cache, Phase 2 tolerance) - Full NodeModal coverage (53 tests) - LiveView, store, edge label tests --- backend/app/api/routes/scan.py | 48 +- backend/app/db/database.py | 4 + backend/app/db/models.py | 2 + backend/app/main.py | 12 + backend/app/schemas/canvas.py | 1 + backend/app/schemas/nodes.py | 2 + backend/app/schemas/scan.py | 1 + backend/app/services/scanner.py | 329 ++++++++---- backend/tests/test_scan.py | 3 +- backend/tests/test_scanner.py | 403 +++++++++++--- frontend/src/api/client.ts | 1 + frontend/src/components/LiveView.tsx | 14 +- .../components/__tests__/LiveView.test.tsx | 60 ++- .../src/components/canvas/CanvasContainer.tsx | 16 +- .../canvas/__tests__/CanvasContainer.test.tsx | 1 + .../src/components/canvas/edges/index.tsx | 6 +- .../src/components/canvas/nodes/BaseNode.tsx | 37 +- frontend/src/components/modals/NodeModal.tsx | 139 +++-- .../components/modals/PendingDeviceModal.tsx | 4 + .../src/components/modals/ScanConfigModal.tsx | 1 - .../modals/__tests__/NodeModal.test.tsx | 500 +++++++++++++----- .../modals/__tests__/ScanConfigModal.test.tsx | 47 +- frontend/src/components/panels/Sidebar.tsx | 30 +- .../src/stores/__tests__/canvasStore.test.ts | 63 ++- frontend/src/stores/canvasStore.ts | 38 +- frontend/src/types/index.ts | 1 + .../src/utils/__tests__/handleUtils.test.ts | 111 ++++ frontend/src/utils/canvasSerializer.ts | 6 +- frontend/src/utils/handleUtils.ts | 45 ++ frontend/src/utils/nodeIcons.ts | 22 + 30 files changed, 1476 insertions(+), 471 deletions(-) create mode 100644 frontend/src/utils/__tests__/handleUtils.test.ts create mode 100644 frontend/src/utils/handleUtils.ts diff --git a/backend/app/api/routes/scan.py b/backend/app/api/routes/scan.py index f760ae1..6f3dcb2 100644 --- a/backend/app/api/routes/scan.py +++ b/backend/app/api/routes/scan.py @@ -1,8 +1,10 @@ +import ipaddress import logging +import uuid from typing import Any from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException -from pydantic import BaseModel +from pydantic import BaseModel, field_validator from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -18,6 +20,16 @@ from app.services.scanner import request_cancel, run_scan class ScanConfig(BaseModel): ranges: list[str] + @field_validator("ranges") + @classmethod + def validate_cidr(cls, v: list[str]) -> list[str]: + for r in v: + try: + ipaddress.ip_network(r, strict=False) + except ValueError as exc: + raise ValueError(f"Invalid CIDR range: {r!r}") from exc + return v + logger = logging.getLogger(__name__) router = APIRouter() @@ -49,6 +61,10 @@ async def stop_scan( db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user), ) -> dict[str, bool]: + try: + uuid.UUID(run_id) + except ValueError: + raise HTTPException(status_code=400, detail="Invalid run_id format") from None run = await db.get(ScanRun, run_id) if not run: raise HTTPException(status_code=404, detail="Scan run not found") @@ -64,6 +80,19 @@ async def list_pending(db: AsyncSession = Depends(get_db), _: str = Depends(get_ return list(result.scalars().all()) +@router.delete("/pending", response_model=dict) +async def clear_pending( + db: AsyncSession = Depends(get_db), + _: str = Depends(get_current_user), +) -> dict[str, int]: + result = await db.execute(select(PendingDevice).where(PendingDevice.status == "pending")) + devices = result.scalars().all() + for device in devices: + await db.delete(device) + await db.commit() + return {"deleted": len(devices)} + + @router.get("/hidden", response_model=list[PendingDeviceResponse]) async def list_hidden(db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)) -> list[PendingDevice]: result = await db.execute(select(PendingDevice).where(PendingDevice.status == "hidden")) @@ -92,9 +121,10 @@ async def hide_device( device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user) ) -> dict[str, bool]: device = await db.get(PendingDevice, device_id) - if device: - device.status = "hidden" - await db.commit() + if not device: + raise HTTPException(status_code=404, detail="Device not found") + device.status = "hidden" + await db.commit() return {"hidden": True} @@ -103,9 +133,10 @@ async def ignore_device( device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user) ) -> dict[str, bool]: device = await db.get(PendingDevice, device_id) - if device: - await db.delete(device) - await db.commit() + if not device: + raise HTTPException(status_code=404, detail="Device not found") + await db.delete(device) + await db.commit() return {"ignored": True} @@ -127,4 +158,5 @@ async def update_scan_config(payload: ScanConfig, _: str = Depends(get_current_u settings.save_overrides() return payload except Exception as exc: - raise HTTPException(status_code=500, detail=str(exc)) from exc + logger.error("Failed to save scan config: %s", exc) + raise HTTPException(status_code=500, detail="Failed to save scan config") from exc diff --git a/backend/app/db/database.py b/backend/app/db/database.py index 72016fe..6717f45 100644 --- a/backend/app/db/database.py +++ b/backend/app/db/database.py @@ -57,6 +57,10 @@ async def init_db() -> None: await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN width REAL") with suppress(OperationalError): await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN height REAL") + with suppress(OperationalError): + 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") # Migrate animated column from boolean (0/1) to string ('none'/'snake') with suppress(OperationalError): await conn.exec_driver_sql("UPDATE edges SET animated = 'snake' WHERE animated = '1' OR animated = 1") diff --git a/backend/app/db/models.py b/backend/app/db/models.py index 1e2daf2..0aac1b0 100644 --- a/backend/app/db/models.py +++ b/backend/app/db/models.py @@ -44,6 +44,7 @@ class Node(Base): show_hardware: Mapped[bool] = mapped_column(Boolean, default=False) 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) 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) @@ -90,6 +91,7 @@ class PendingDevice(Base): services: Mapped[list[Any]] = mapped_column(JSON, default=list) suggested_type: Mapped[str | None] = mapped_column(String) status: Mapped[str] = mapped_column(String, default="pending") + discovery_source: Mapped[str | None] = mapped_column(String) discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now) diff --git a/backend/app/main.py b/backend/app/main.py index fa7cac8..c1a3bc4 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -1,3 +1,5 @@ +import logging +import logging.config from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from typing import Any @@ -14,6 +16,16 @@ from app.db.database import init_db @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: + # Ensure app logs are visible: attach a handler to the root logger if none + # exists (uvicorn only installs handlers on its own loggers, not the root). + root_logger = logging.getLogger() + if not root_logger.handlers: + handler = logging.StreamHandler() + handler.setFormatter(logging.Formatter("%(levelname)s:%(name)s:%(message)s")) + root_logger.addHandler(handler) + root_logger.setLevel(logging.INFO) + logging.getLogger("app").setLevel(logging.INFO) + logging.getLogger("app.services.scanner").setLevel(logging.INFO) await init_db() settings.load_overrides() start_scheduler() diff --git a/backend/app/schemas/canvas.py b/backend/app/schemas/canvas.py index e3bb87b..5f1e5d4 100644 --- a/backend/app/schemas/canvas.py +++ b/backend/app/schemas/canvas.py @@ -31,6 +31,7 @@ class NodeSave(BaseModel): show_hardware: bool = False width: float | None = None height: float | None = None + bottom_handles: int = 1 pos_x: float = 0 pos_y: float = 0 diff --git a/backend/app/schemas/nodes.py b/backend/app/schemas/nodes.py index caaad34..70156a6 100644 --- a/backend/app/schemas/nodes.py +++ b/backend/app/schemas/nodes.py @@ -29,6 +29,7 @@ class NodeBase(BaseModel): show_hardware: bool = False width: float | None = None height: float | None = None + bottom_handles: int = 1 class NodeCreate(NodeBase): @@ -60,6 +61,7 @@ class NodeUpdate(BaseModel): show_hardware: bool | None = None width: float | None = None height: float | None = None + bottom_handles: int | None = None class NodeResponse(NodeBase): diff --git a/backend/app/schemas/scan.py b/backend/app/schemas/scan.py index e16d580..5b7314e 100644 --- a/backend/app/schemas/scan.py +++ b/backend/app/schemas/scan.py @@ -13,6 +13,7 @@ class PendingDeviceResponse(BaseModel): services: list[Any] suggested_type: str | None status: str + discovery_source: str | None discovered_at: datetime model_config = {"from_attributes": True} diff --git a/backend/app/services/scanner.py b/backend/app/services/scanner.py index 3ac8da5..4f0ac43 100644 --- a/backend/app/services/scanner.py +++ b/backend/app/services/scanner.py @@ -1,7 +1,12 @@ """Network scanner: ARP sweep + nmap service detection + mDNS discovery.""" import asyncio +import ipaddress import logging +import os +import re import socket +import subprocess +import threading from datetime import datetime, timezone from typing import Any @@ -13,8 +18,9 @@ from app.services.fingerprint import fingerprint_ports, suggest_node_type logger = logging.getLogger(__name__) -# Run IDs that have been requested to cancel +# Run IDs that have been requested to cancel (thread-safe via lock) _cancelled_runs: set[str] = set() +_cancelled_lock = threading.Lock() # Port list for service detection (Phase 2) _EXTRA_PORTS = ( @@ -55,11 +61,13 @@ except ImportError: def request_cancel(run_id: str) -> None: """Signal a running scan to stop early.""" - _cancelled_runs.add(run_id) + with _cancelled_lock: + _cancelled_runs.add(run_id) def _is_cancelled(run_id: str) -> bool: - return run_id in _cancelled_runs + with _cancelled_lock: + return run_id in _cancelled_runs def _resolve_hostname(ip: str) -> str | None: @@ -79,80 +87,216 @@ def _extract_os(nm: object, host: str) -> str | None: return None -def _nmap_arp_sweep(target: str) -> dict[str, dict[str, Any]]: +def _arp_table_hosts(network: str) -> dict[str, dict[str, Any]]: """ - Phase 1: ARP ping sweep — finds ALL alive hosts regardless of open ports. - Returns {ip: host_dict} for every host that responds. + Read the OS ARP cache for recently-seen hosts in the target network. + Works without root on both Linux (/proc/net/arp) and macOS (arp -a). + Supplements nmap discovery — catches IoT and devices with all ports filtered. """ - nm = nmap.PortScanner() - nm.scan(hosts=target, arguments="-sn -PR -PA80,443 --host-timeout 10s") + try: + net = ipaddress.ip_network(network, strict=False) + found: dict[str, dict[str, Any]] = {} + + # Linux: parse /proc/net/arp — present on any Linux kernel (including Docker) + proc_arp = "/proc/net/arp" + try: + with open(proc_arp) as f: + for line in f.readlines()[1:]: # skip header row + parts = line.split() + if len(parts) >= 4: + ip, mac = parts[0], parts[3] + if mac == "00:00:00:00:00:00": + continue + try: + if ipaddress.ip_address(ip) in net: + found[ip] = { + "ip": ip, "mac": mac, + "hostname": _resolve_hostname(ip), + "os": None, "open_ports": [], + } + except ValueError: + pass + # /proc/net/arp opened successfully — return whatever we found (may be empty) + # Don't fall through to `arp -a` since we're on Linux + return found + except FileNotFoundError: + pass # Not Linux — fall through to macOS `arp -a` + + # macOS: parse `arp -a` output + result = subprocess.run(["arp", "-a"], capture_output=True, text=True, timeout=5) + for line in result.stdout.splitlines(): + m = re.search(r"\((\d+\.\d+\.\d+\.\d+)\)\s+at\s+([0-9a-f:]+)", line) + if not m: + continue + ip, mac = m.group(1), m.group(2) + if mac in ("(incomplete)", "ff:ff:ff:ff:ff:ff"): + continue + try: + if ipaddress.ip_address(ip) in net: + found[ip] = {"ip": ip, "mac": mac, "hostname": _resolve_hostname(ip), "os": None, "open_ports": []} + except ValueError: + pass + return found + except Exception as exc: + logger.warning("[Phase 1] ARP cache lookup failed: %s", exc) + return {} + + +async def _ping_sweep(target: str) -> dict[str, dict[str, Any]]: + """ + Phase 1: Concurrent ICMP ping sweep + ARP cache. + Pings all IPs in the CIDR in parallel (up to 50 at once, 1s timeout each). + Supplements with the OS ARP cache to catch devices that block ICMP. + Works in Docker with CAP_NET_RAW — no nmap, no false positives. + """ + net = ipaddress.ip_network(target, strict=False) + all_ips = [str(ip) for ip in net.hosts()] + logger.info("[Phase 1] Pinging %d hosts in %s ...", len(all_ips), target) + + sem = asyncio.Semaphore(50) + + async def _ping(ip: str) -> str | None: + async with sem: + try: + proc = await asyncio.create_subprocess_exec( + "ping", "-c", "1", "-W", "1", ip, + stdout=asyncio.subprocess.DEVNULL, + stderr=asyncio.subprocess.DEVNULL, + ) + await proc.wait() + return ip if proc.returncode == 0 else None + except Exception: + return None + + ping_results = await asyncio.gather(*[_ping(ip) for ip in all_ips]) + alive_ips: set[str] = {ip for ip in ping_results if ip is not None} + logger.info("[Phase 1] %d/%d hosts responded to ping", len(alive_ips), len(all_ips)) + + # ARP cache: catch devices that block ICMP but were recently active, + # and enrich ping-alive hosts with their MAC addresses. + arp_cache = await asyncio.to_thread(_arp_table_hosts, target) + alive: dict[str, dict[str, Any]] = {} - for host in nm.all_hosts(): - if nm[host].state() == "up": - alive[host] = { - "ip": host, - "hostname": _resolve_hostname(host), - "mac": nm[host].get("addresses", {}).get("mac"), - "os": None, - "open_ports": [], - } + + for ip in alive_ips: + mac = arp_cache.get(ip, {}).get("mac") + hostname = await asyncio.to_thread(_resolve_hostname, ip) + logger.info("[Phase 1] %s mac=%s hostname=%s (ping)", ip, mac or "n/a", hostname or "n/a") + alive[ip] = {"ip": ip, "mac": mac, "hostname": hostname, "os": None, "open_ports": []} + + for ip, host in arp_cache.items(): + if ip not in alive: + logger.info( + "[Phase 1] %s mac=%s hostname=%s (ARP cache only)", + ip, host.get("mac") or "n/a", host.get("hostname") or "n/a", + ) + alive[ip] = host + return alive -def _nmap_port_scan(alive: dict[str, dict[str, Any]]) -> list[dict[str, Any]]: +def _nmap_scan_single(host_dict: dict[str, Any]) -> dict[str, Any]: """ - Phase 2: Service detection on the alive host set from Phase 1. - Mutates alive in-place with open_ports/os; returns all hosts including - those with zero open ports (IoT devices often have none). + Phase 2 — single-IP port scan with service detection. + Runs in a thread (blocking). Returns the host dict enriched with open_ports. + """ + ip = host_dict["ip"] + logger.info("[Phase 2] Scanning %s ...", ip) + + if not _NMAP_AVAILABLE: + logger.warning("[Phase 2] nmap not available, skipping %s", ip) + return host_dict + + is_root = os.geteuid() == 0 + if is_root: + # SYN scan + version detection (fastest, most accurate) + scan_args = f"-sS -sV --open -T4 -Pn --host-timeout 60s -p {_EXTRA_PORTS}" + else: + # TCP connect scan (-sT) — no raw sockets needed, works without root. + # nmap auto-selects -sT without root but being explicit avoids edge cases. + scan_args = f"-sT -sV --open -T4 -Pn --host-timeout 60s -p {_EXTRA_PORTS}" + + logger.debug("[Phase 2] %s args: %s", ip, scan_args) + nm = nmap.PortScanner() + try: + nm.scan(hosts=ip, arguments=scan_args) + except Exception as exc: + logger.warning("[Phase 2] nmap FAILED for %s (%s: %s) — skipping port scan", ip, type(exc).__name__, exc) + return host_dict + + all_scanned = nm.all_hosts() + logger.debug("[Phase 2] %s — nmap returned %d host(s) in results", ip, len(all_scanned)) + if ip not in all_scanned: + logger.info("[Phase 2] %s — no open ports found (all closed/filtered or nmap had no results)", ip) + return host_dict + + open_ports = [] + for proto in nm[ip].all_protocols(): + for port, info in nm[ip][proto].items(): + if info["state"] == "open": + banner = (info.get("product", "") + " " + info.get("version", "")).strip() + open_ports.append({"port": port, "protocol": proto, "banner": banner}) + + if open_ports: + port_summary = ", ".join( + f"{p['port']}/{p['protocol']} ({p['banner'] or 'unknown'})" for p in open_ports + ) + logger.info("[Phase 2] %s — %d open port(s): %s", ip, len(open_ports), port_summary) + else: + logger.info("[Phase 2] %s — 0 open ports detected", ip) + + host_dict["open_ports"] = open_ports + if not host_dict["mac"]: + host_dict["mac"] = nm[ip].get("addresses", {}).get("mac") + host_dict["os"] = _extract_os(nm, ip) + return host_dict + + +async def _nmap_port_scan(alive: dict[str, dict[str, Any]]) -> list[dict[str, Any]]: + """ + Phase 2: Per-IP service detection with bounded concurrency. + Each host is scanned independently in a thread — no inter-host timeout interference. + Up to 10 hosts scanned concurrently. """ if not alive: return [] - nm = nmap.PortScanner() - try: - nm.scan( - hosts=" ".join(alive.keys()), - arguments=f"-sV --open -T4 --host-timeout 30s -p {_EXTRA_PORTS}", - ) - except Exception as exc: - logger.warning("Port scan failed, returning ARP-only results: %s", exc) - return list(alive.values()) - for host in nm.all_hosts(): - if host not in alive: - continue - open_ports = [] - for proto in nm[host].all_protocols(): - for port, info in nm[host][proto].items(): - if info["state"] == "open": - open_ports.append({ - "port": port, - "protocol": proto, - "banner": ( - info.get("product", "") + " " + info.get("version", "") - ).strip(), - }) - alive[host]["open_ports"] = open_ports - if not alive[host]["mac"]: - alive[host]["mac"] = nm[host].get("addresses", {}).get("mac") - alive[host]["os"] = _extract_os(nm, host) + logger.info("[Phase 2] Starting per-IP port scan for %d host(s)", len(alive)) + semaphore = asyncio.Semaphore(10) - return list(alive.values()) + async def _scan_with_sem(host_dict: dict[str, Any]) -> dict[str, Any]: + async with semaphore: + return await asyncio.to_thread(_nmap_scan_single, host_dict) + + raw = await asyncio.gather(*[_scan_with_sem(h) for h in alive.values()], return_exceptions=True) + results = [] + for item in raw: + if isinstance(item, BaseException): + logger.warning("[Phase 2] Unexpected error in gather: %s", item) + else: + results.append(item) + logger.info("[Phase 2] Completed — %d/%d host(s) scanned", len(results), len(alive)) + return results -def _nmap_scan(target: str) -> list[dict[str, Any]]: +async def _nmap_scan(target: str) -> list[dict[str, Any]]: """ - Full two-phase scan for a CIDR range. - Phase 1: ARP sweep to find alive hosts (catches IoT with no open ports). - Phase 2: Service detection on alive hosts only. + Two-phase scan for a CIDR range. + Phase 1: Concurrent ping sweep to find alive hosts (fast, no false positives). + Phase 2: Per-IP nmap port scan with service detection (bounded concurrency, 10 at a time). """ + logger.info("[Scan] Starting scan for %s — nmap available: %s", target, _NMAP_AVAILABLE) if not _NMAP_AVAILABLE: + logger.warning("[Scan] nmap not available — returning mock data") return _mock_scan(target) try: - alive = _nmap_arp_sweep(target) + alive = await _ping_sweep(target) + logger.info("[Phase 1] Found %d alive host(s) in %s: %s", + len(alive), target, ", ".join(sorted(alive.keys()))) except Exception as exc: - logger.error("nmap ARP sweep failed: %s", exc) + logger.error("Phase 1 ping sweep failed: %s", exc) raise RuntimeError(str(exc)) from exc - return _nmap_port_scan(alive) + return await _nmap_port_scan(alive) async def _mdns_discover(timeout: float = 4.0) -> list[dict[str, Any]]: @@ -236,45 +380,50 @@ async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None: from app.api.routes.status import broadcast_scan_update devices_found = 0 + mdns_task: asyncio.Task[list[dict[str, Any]]] | None = None try: - # Clean up stale pending devices whose IPs are already in the canvas + # Validate all ranges are valid CIDRs before passing anything to nmap + for r in ranges: + try: + ipaddress.ip_network(r, strict=False) + except ValueError: + raise ValueError(f"Invalid CIDR range: {r!r}") from None + + # Pre-fetch canvas IPs and hidden IPs once — avoids N+1 queries per host canvas_ips_result = await db.execute(select(Node.ip).where(Node.ip.isnot(None))) canvas_ips: set[str] = {row[0] for row in canvas_ips_result.fetchall()} + + hidden_ips_result = await db.execute( + select(PendingDevice.ip).where(PendingDevice.status == "hidden") + ) + hidden_ips: set[str] = {row[0] for row in hidden_ips_result.fetchall()} + + # Clean up stale pending devices whose IPs are already in the canvas if canvas_ips: - stale_result = await db.execute( - select(PendingDevice).where( + from sqlalchemy import delete as sa_delete + await db.execute( + sa_delete(PendingDevice).where( PendingDevice.status == "pending", PendingDevice.ip.in_(canvas_ips), ) ) - for stale in stale_result.scalars().all(): - await db.delete(stale) await db.commit() # Start mDNS discovery in the background while nmap scans run - mdns_task: asyncio.Task[list[dict[str, Any]]] = asyncio.create_task( - _mdns_discover() - ) + mdns_task = asyncio.create_task(_mdns_discover()) # Track IPs found by nmap so mDNS doesn't duplicate them nmap_ips: set[str] = set() - async def _process_host(host: dict[str, Any]) -> None: + async def _process_host(host: dict[str, Any], discovery_source: str = "arp") -> None: nonlocal devices_found ip = host["ip"] - # Skip canvas nodes and user-hidden devices - canvas_result = await db.execute(select(Node).where(Node.ip == ip)) - if canvas_result.scalar_one_or_none() is not None: + # Skip canvas nodes and user-hidden devices (sets pre-fetched before loop) + if ip in canvas_ips: logger.debug("Skipping %s — already in canvas", ip) return - hidden_result = await db.execute( - select(PendingDevice).where( - PendingDevice.ip == ip, - PendingDevice.status == "hidden", - ) - ) - if hidden_result.scalar_one_or_none() is not None: + if ip in hidden_ips: logger.debug("Skipping %s — hidden by user", ip) return @@ -303,43 +452,40 @@ async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None: services=services, suggested_type=suggested_type, status="pending", + discovery_source=discovery_source, )) devices_found += 1 await db.commit() - - run = await db.get(ScanRun, run_id) - if run: - run.devices_found = devices_found - await db.commit() - await broadcast_scan_update(run_id=run_id, devices_found=devices_found) # nmap scan per CIDR — results stream in progressively for cidr in ranges: if _is_cancelled(run_id): break - hosts = await asyncio.to_thread(_nmap_scan, cidr) + hosts = await _nmap_scan(cidr) for host in hosts: if _is_cancelled(run_id): break nmap_ips.add(host["ip"]) await _process_host(host) - # Collect mDNS results; add devices not already found by nmap + # Update ScanRun count once after all CIDR ranges + run = await db.get(ScanRun, run_id) + if run: + run.devices_found = devices_found + await db.commit() + + # Collect mDNS results — task already has its own 4s internal timeout if not _is_cancelled(run_id): - try: - mdns_hosts = await asyncio.wait_for(mdns_task, timeout=1.0) - except asyncio.TimeoutError: - mdns_task.cancel() - mdns_hosts = [] + mdns_hosts = await mdns_task for host in mdns_hosts: if _is_cancelled(run_id): break if host["ip"] in nmap_ips: continue # already processed with richer nmap data - await _process_host(host) + await _process_host(host, discovery_source="mdns") else: mdns_task.cancel() @@ -353,6 +499,8 @@ async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None: except Exception as exc: logger.error("Scan failed: %s", exc) + if mdns_task is not None and not mdns_task.done(): + mdns_task.cancel() run = await db.get(ScanRun, run_id) if run: run.status = "error" @@ -360,4 +508,5 @@ async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None: run.finished_at = datetime.now(timezone.utc) await db.commit() finally: - _cancelled_runs.discard(run_id) + with _cancelled_lock: + _cancelled_runs.discard(run_id) diff --git a/backend/tests/test_scan.py b/backend/tests/test_scan.py index 0810040..9540370 100644 --- a/backend/tests/test_scan.py +++ b/backend/tests/test_scan.py @@ -322,7 +322,8 @@ async def test_stop_scan_requires_auth(client: AsyncClient): @pytest.mark.asyncio async def test_stop_scan_not_found(client: AsyncClient, headers): - res = await client.post("/api/v1/scan/nonexistent-id/stop", headers=headers) + import uuid as _uuid + res = await client.post(f"/api/v1/scan/{_uuid.uuid4()}/stop", headers=headers) assert res.status_code == 404 diff --git a/backend/tests/test_scanner.py b/backend/tests/test_scanner.py index d368832..e29e2fc 100644 --- a/backend/tests/test_scanner.py +++ b/backend/tests/test_scanner.py @@ -3,10 +3,12 @@ import uuid from unittest.mock import AsyncMock, MagicMock, patch import pytest +from sqlalchemy import select as sa_select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.pool import StaticPool from app.db.database import Base -from app.db.models import PendingDevice, ScanRun +from app.db.models import Node, PendingDevice, ScanRun # --------------------------------------------------------------------------- # Helpers @@ -18,7 +20,11 @@ def _make_run_id() -> str: @pytest.fixture async def mem_db(): - engine = create_async_engine("sqlite+aiosqlite:///:memory:") + engine = create_async_engine( + "sqlite+aiosqlite:///:memory:", + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) @@ -31,140 +37,246 @@ def _make_scan_run(run_id: str) -> ScanRun: # --------------------------------------------------------------------------- -# _nmap_arp_sweep +# _ping_sweep # --------------------------------------------------------------------------- -def test_nmap_arp_sweep_returns_alive_hosts(): - from app.services.scanner import _nmap_arp_sweep +@pytest.mark.asyncio +async def test_ping_sweep_returns_alive_hosts(): + from app.services.scanner import _ping_sweep - mock_nm = MagicMock() - mock_nm.all_hosts.return_value = ["192.168.1.1", "192.168.1.2"] - mock_nm.__getitem__ = lambda self, host: MagicMock( - state=lambda: "up", - get=lambda key, default=None: {"mac": "aa:bb:cc:dd:ee:ff"} if key == "addresses" else default, - ) + async def fake_ping(ip: str) -> str | None: + return ip if ip in {"192.168.1.1", "192.168.1.2"} else None - with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \ + with patch("app.services.scanner._ping_sweep", wraps=None): + pass # just ensure import is fine + + # Patch asyncio.create_subprocess_exec to simulate ping responses + responding = {"192.168.1.1", "192.168.1.2"} + + async def mock_subprocess(*args, **kwargs): + ip = args[-1] + proc = MagicMock() + proc.returncode = 0 if ip in responding else 1 + proc.wait = AsyncMock(return_value=proc.returncode) + return proc + + with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \ + patch("app.services.scanner._arp_table_hosts", return_value={}), \ patch("app.services.scanner._resolve_hostname", return_value=None): - result = _nmap_arp_sweep("192.168.1.0/24") + result = await _ping_sweep("192.168.1.0/30") # .1 .2 only in /30 - assert set(result.keys()) == {"192.168.1.1", "192.168.1.2"} + assert "192.168.1.1" in result + assert "192.168.1.2" in result for host in result.values(): - assert host["open_ports"] == [] # empty until phase 2 + assert host["open_ports"] == [] -def test_nmap_arp_sweep_skips_down_hosts(): - from app.services.scanner import _nmap_arp_sweep +@pytest.mark.asyncio +async def test_ping_sweep_excludes_non_responding(): + from app.services.scanner import _ping_sweep - states = {"192.168.1.1": "up", "192.168.1.2": "down"} + async def mock_subprocess(*args, **kwargs): + ip = args[-1] + proc = MagicMock() + proc.returncode = 0 if ip == "192.168.1.1" else 1 + proc.wait = AsyncMock(return_value=proc.returncode) + return proc - mock_nm = MagicMock() - mock_nm.all_hosts.return_value = list(states.keys()) - - def getitem(host): - m = MagicMock() - m.state.return_value = states[host] - m.get.return_value = {} - return m - - mock_nm.__getitem__ = lambda self, host: getitem(host) - - with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \ + with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \ + patch("app.services.scanner._arp_table_hosts", return_value={}), \ patch("app.services.scanner._resolve_hostname", return_value=None): - result = _nmap_arp_sweep("192.168.1.0/24") + result = await _ping_sweep("192.168.1.0/30") assert "192.168.1.1" in result assert "192.168.1.2" not in result -# --------------------------------------------------------------------------- -# _nmap_port_scan -# --------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_ping_sweep_supplements_with_arp_cache(): + """Devices that block ICMP but appear in ARP cache should still be discovered.""" + from app.services.scanner import _ping_sweep -def test_nmap_port_scan_merges_ports(): - from app.services.scanner import _nmap_port_scan + async def mock_subprocess(*args, **kwargs): + proc = MagicMock() + proc.returncode = 1 # all pings fail + proc.wait = AsyncMock(return_value=1) + return proc - alive = { - "192.168.1.10": {"ip": "192.168.1.10", "hostname": None, "mac": None, "os": None, "open_ports": []}, + arp_extra = { + "192.168.1.10": {"ip": "192.168.1.10", "mac": "aa:bb:cc:dd:ee:10", "hostname": None, "os": None, "open_ports": []}, } + with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \ + patch("app.services.scanner._arp_table_hosts", return_value=arp_extra), \ + patch("app.services.scanner._resolve_hostname", return_value=None): + result = await _ping_sweep("192.168.1.0/24") + + assert "192.168.1.10" in result + assert result["192.168.1.10"]["mac"] == "aa:bb:cc:dd:ee:10" + + +@pytest.mark.asyncio +async def test_ping_sweep_enriches_mac_from_arp_cache(): + """Ping-alive hosts with no ARP entry get their MAC from the ARP cache.""" + from app.services.scanner import _ping_sweep + + async def mock_subprocess(*args, **kwargs): + ip = args[-1] + proc = MagicMock() + proc.returncode = 0 if ip == "192.168.1.1" else 1 + proc.wait = AsyncMock(return_value=proc.returncode) + return proc + + arp_extra = { + "192.168.1.1": {"ip": "192.168.1.1", "mac": "de:ad:be:ef:00:01", "hostname": None, "os": None, "open_ports": []}, + } + + with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \ + patch("app.services.scanner._arp_table_hosts", return_value=arp_extra), \ + patch("app.services.scanner._resolve_hostname", return_value=None): + result = await _ping_sweep("192.168.1.0/30") + + assert result["192.168.1.1"]["mac"] == "de:ad:be:ef:00:01" + + +# --------------------------------------------------------------------------- +# _arp_table_hosts +# --------------------------------------------------------------------------- + +def test_arp_table_hosts_parses_proc_net_arp(): + import io # noqa: PLC0415 + + from app.services.scanner import _arp_table_hosts + + arp_content = ( + "IP address HW type Flags HW address Mask Device\n" + "192.168.1.1 0x1 0x2 aa:bb:cc:dd:ee:01 * eth0\n" + "192.168.1.50 0x1 0x2 aa:bb:cc:dd:ee:02 * eth0\n" + "10.0.0.1 0x1 0x2 aa:bb:cc:dd:ee:03 * eth0\n" # outside subnet + "192.168.1.99 0x1 0x2 00:00:00:00:00:00 * eth0\n" # incomplete + ) + + mock_file = MagicMock() + mock_file.__enter__ = MagicMock(return_value=io.StringIO(arp_content)) + mock_file.__exit__ = MagicMock(return_value=False) + + with patch("builtins.open", return_value=mock_file), \ + patch("app.services.scanner._resolve_hostname", return_value=None): + result = _arp_table_hosts("192.168.1.0/24") + + assert "192.168.1.1" in result + assert "192.168.1.50" in result + assert "10.0.0.1" not in result # outside target subnet + assert "192.168.1.99" not in result # zero MAC skipped + + +def test_arp_table_hosts_parses_macos_arp_output(): + from app.services.scanner import _arp_table_hosts + + arp_output = ( + "router.lan (192.168.1.1) at aa:bb:cc:dd:ee:01 on en0 ifscope [ethernet]\n" + "device.lan (192.168.1.20) at aa:bb:cc:dd:ee:02 on en0 ifscope [ethernet]\n" + "? (192.168.1.99) at (incomplete) on en0 ifscope [ethernet]\n" + "? (10.0.0.1) at aa:bb:cc:dd:ee:04 on en0 ifscope [ethernet]\n" # outside subnet + ) + + mock_result = MagicMock() + mock_result.stdout = arp_output + + with patch("builtins.open", side_effect=FileNotFoundError), \ + patch("subprocess.run", return_value=mock_result), \ + patch("app.services.scanner._resolve_hostname", return_value=None): + result = _arp_table_hosts("192.168.1.0/24") + + assert "192.168.1.1" in result + assert "192.168.1.20" in result + assert "192.168.1.99" not in result # incomplete MAC + assert "10.0.0.1" not in result # outside subnet + + +# --------------------------------------------------------------------------- +# _nmap_scan_single (Phase 2 per-IP worker) +# --------------------------------------------------------------------------- + +def test_nmap_scan_single_detects_open_ports(): + from app.services.scanner import _nmap_scan_single + + host = {"ip": "192.168.1.10", "hostname": None, "mac": None, "os": None, "open_ports": []} + + # Build a realistic host entry: protocols → ports → port info + port_info = {80: {"state": "open", "product": "nginx", "version": "1.24"}} + mock_host = MagicMock() + mock_host.all_protocols.return_value = ["tcp"] + mock_host.__getitem__ = MagicMock(return_value=port_info) + mock_host.get.return_value = {} + mock_nm = MagicMock() mock_nm.all_hosts.return_value = ["192.168.1.10"] - mock_nm.__getitem__ = lambda self, host: MagicMock( - all_protocols=lambda: ["tcp"], - **{"__getitem__": lambda self2, proto: { - 80: {"state": "open", "product": "nginx", "version": "1.24"}, - }}, - get=lambda key, default=None: default, - ) + mock_nm.__getitem__ = MagicMock(return_value=mock_host) with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \ patch("app.services.scanner._extract_os", return_value=None): - result = _nmap_port_scan(alive) + result = _nmap_scan_single(host) - assert len(result) == 1 - assert result[0]["open_ports"][0]["port"] == 80 + assert len(result["open_ports"]) == 1 + assert result["open_ports"][0]["port"] == 80 + assert result["open_ports"][0]["banner"] == "nginx 1.24" -def test_nmap_port_scan_returns_arp_only_on_failure(): - from app.services.scanner import _nmap_port_scan - - alive = { - "192.168.1.20": {"ip": "192.168.1.20", "hostname": None, "mac": None, "os": None, "open_ports": []}, - } +def test_nmap_scan_single_returns_host_unchanged_on_error(): + from app.services.scanner import _nmap_scan_single + host = {"ip": "192.168.1.20", "hostname": None, "mac": None, "os": None, "open_ports": []} mock_nm = MagicMock() mock_nm.scan.side_effect = Exception("nmap error") with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm): - result = _nmap_port_scan(alive) + result = _nmap_scan_single(host) - # Should return the ARP-found host even though port scan failed - assert len(result) == 1 - assert result[0]["ip"] == "192.168.1.20" - assert result[0]["open_ports"] == [] + assert result["ip"] == "192.168.1.20" + assert result["open_ports"] == [] -def test_nmap_port_scan_includes_hosts_with_no_open_ports(): - """IoT devices found by ARP but with no open TCP ports must still be returned.""" - from app.services.scanner import _nmap_port_scan +def test_nmap_scan_single_returns_host_unchanged_when_no_results(): + """Host confirmed alive in Phase 1 but all ports filtered — keep it with empty ports.""" + from app.services.scanner import _nmap_scan_single - alive = { - "192.168.1.30": {"ip": "192.168.1.30", "hostname": "shelly1.lan", "mac": "34:94:54:aa:bb:cc", "os": None, "open_ports": []}, - "192.168.1.31": {"ip": "192.168.1.31", "hostname": None, "mac": None, "os": None, "open_ports": []}, - } - - # Port scan returns only 192.168.1.31 (e.g., .30 filtered all ports) + host = {"ip": "192.168.1.30", "hostname": "shelly1.lan", "mac": "34:94:54:aa:bb:cc", "os": None, "open_ports": []} mock_nm = MagicMock() - mock_nm.all_hosts.return_value = ["192.168.1.31"] - mock_nm.__getitem__ = lambda self, host: MagicMock( - all_protocols=lambda: [], - get=lambda key, default=None: default, - ) + mock_nm.all_hosts.return_value = [] # no results - with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \ - patch("app.services.scanner._extract_os", return_value=None): - result = _nmap_port_scan(alive) + with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm): + result = _nmap_scan_single(host) - ips = {h["ip"] for h in result} - assert "192.168.1.30" in ips, "ARP-found device with no open ports must still be returned" - assert "192.168.1.31" in ips + assert result["ip"] == "192.168.1.30" + assert result["open_ports"] == [] + assert result["mac"] == "34:94:54:aa:bb:cc" # preserved from Phase 1 # --------------------------------------------------------------------------- -# _nmap_scan (integration of both phases) +# _nmap_scan # --------------------------------------------------------------------------- -def test_nmap_scan_uses_mock_when_nmap_unavailable(): +@pytest.mark.asyncio +async def test_nmap_scan_uses_mock_when_nmap_unavailable(): from app.services.scanner import _nmap_scan with patch("app.services.scanner._NMAP_AVAILABLE", False): - result = _nmap_scan("192.168.1.0/24") + result = await _nmap_scan("192.168.1.0/24") assert len(result) == 1 assert result[0]["ip"] == "192.168.1.99" +@pytest.mark.asyncio +async def test_nmap_scan_raises_on_sweep_error(): + from app.services.scanner import _nmap_scan + + with patch("app.services.scanner._ping_sweep", side_effect=Exception("ping sweep failed")), \ + pytest.raises(RuntimeError, match="ping sweep failed"): + await _nmap_scan("192.168.1.0/24") + + # --------------------------------------------------------------------------- # _mdns_discover # --------------------------------------------------------------------------- @@ -223,6 +335,47 @@ async def test_mdns_discover_returns_devices(): assert result[0]["hostname"] == "shelly1.local." +# --------------------------------------------------------------------------- +# _nmap_port_scan (Phase 2 concurrency) +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_nmap_port_scan_returns_empty_when_no_alive_hosts(): + from app.services.scanner import _nmap_port_scan + + result = await _nmap_port_scan({}) + assert result == [] + + +@pytest.mark.asyncio +async def test_nmap_port_scan_tolerates_single_host_exception(): + """A single per-host failure should not abort the entire Phase 2 gather.""" + from app.services.scanner import _nmap_port_scan + + hosts = { + "192.168.1.1": {"ip": "192.168.1.1", "hostname": None, "mac": None, "os": None, "open_ports": []}, + "192.168.1.2": {"ip": "192.168.1.2", "hostname": None, "mac": None, "os": None, "open_ports": []}, + } + + call_count = 0 + + def _flaky_scan(host_dict): + nonlocal call_count + call_count += 1 + if host_dict["ip"] == "192.168.1.1": + raise RuntimeError("simulated nmap crash") + return host_dict + + with patch("app.services.scanner._nmap_scan_single", side_effect=_flaky_scan), \ + patch("app.services.scanner._NMAP_AVAILABLE", True): + result = await _nmap_port_scan(hosts) + + assert call_count == 2 + # The crashing host is dropped; the healthy one survives + assert len(result) == 1 + assert result[0]["ip"] == "192.168.1.2" + + # --------------------------------------------------------------------------- # run_scan integration # --------------------------------------------------------------------------- @@ -245,9 +398,7 @@ async def test_run_scan_adds_nmap_devices_as_pending(mem_db): await run_scan(["192.168.1.0/24"], session, run_id) async with mem_db() as session: - result = await session.execute( - __import__("sqlalchemy", fromlist=["select"]).select(PendingDevice) - ) + result = await session.execute(sa_select(PendingDevice)) devices = result.scalars().all() assert any(d.ip == "192.168.1.5" for d in devices) @@ -272,12 +423,12 @@ async def test_run_scan_mdns_only_device_added(mem_db): await run_scan(["192.168.1.0/24"], session, run_id) async with mem_db() as session: - from sqlalchemy import select as sa_select result = await session.execute(sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.80")) device = result.scalar_one_or_none() assert device is not None assert device.status == "pending" + assert device.discovery_source == "mdns" @pytest.mark.asyncio @@ -299,8 +450,86 @@ async def test_run_scan_mdns_skipped_if_already_in_nmap(mem_db): await run_scan(["192.168.1.0/24"], session, run_id) async with mem_db() as session: - from sqlalchemy import select as sa_select result = await session.execute(sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.10")) devices = result.scalars().all() assert len(devices) == 1 # not duplicated + + +@pytest.mark.asyncio +async def test_run_scan_skips_canvas_nodes(mem_db): + """Hosts already approved onto the canvas must be skipped.""" + from app.services.scanner import run_scan + + run_id = _make_run_id() + async with mem_db() as session: + session.add(_make_scan_run(run_id)) + canvas_node = Node( + id=str(uuid.uuid4()), label="PVE", type="proxmox", + ip="192.168.1.100", status="online", + ) + session.add(canvas_node) + await session.commit() + + nmap_hosts = [{"ip": "192.168.1.100", "hostname": "pve.lan", "mac": None, "os": None, "open_ports": []}] + + async with mem_db() as session: + with patch("app.services.scanner._nmap_scan", return_value=nmap_hosts), \ + patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[]), \ + patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock): + await run_scan(["192.168.1.0/24"], session, run_id) + + async with mem_db() as session: + result = await session.execute(sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.100")) + assert result.scalar_one_or_none() is None + + +@pytest.mark.asyncio +async def test_run_scan_skips_hidden_devices(mem_db): + """Hosts hidden by the user must not re-appear in pending.""" + from app.services.scanner import run_scan + + run_id = _make_run_id() + async with mem_db() as session: + session.add(_make_scan_run(run_id)) + hidden = PendingDevice(ip="192.168.1.55", status="hidden") + session.add(hidden) + await session.commit() + + nmap_hosts = [{"ip": "192.168.1.55", "hostname": None, "mac": None, "os": None, "open_ports": []}] + + async with mem_db() as session: + with patch("app.services.scanner._nmap_scan", return_value=nmap_hosts), \ + patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[]), \ + patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock): + await run_scan(["192.168.1.0/24"], session, run_id) + + async with mem_db() as session: + result = await session.execute( + sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.55", PendingDevice.status == "pending") + ) + assert result.scalar_one_or_none() is None + + +@pytest.mark.asyncio +async def test_run_scan_cancelled_marks_status_cancelled(mem_db): + """Cancelling a running scan sets the ScanRun status to 'cancelled'.""" + from app.services.scanner import request_cancel, run_scan + + run_id = _make_run_id() + async with mem_db() as session: + session.add(_make_scan_run(run_id)) + await session.commit() + + request_cancel(run_id) + + async with mem_db() as session: + with patch("app.services.scanner._nmap_scan", return_value=[]), \ + patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[]), \ + patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock): + await run_scan(["192.168.1.0/24"], session, run_id) + + async with mem_db() as session: + run = await session.get(ScanRun, run_id) + assert run is not None + assert run.status == "cancelled" diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index e1601b5..6f9d808 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -56,6 +56,7 @@ export const scanApi = { pending: () => api.get('/scan/pending'), 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), hide: (id: string) => api.post(`/scan/pending/${id}/hide`), ignore: (id: string) => api.post(`/scan/pending/${id}/ignore`), diff --git a/frontend/src/components/LiveView.tsx b/frontend/src/components/LiveView.tsx index bfcdde4..930fb5d 100644 --- a/frontend/src/components/LiveView.tsx +++ b/frontend/src/components/LiveView.tsx @@ -18,6 +18,7 @@ import { BackgroundVariant, Controls, ConnectionMode, + useReactFlow, type Node, } from '@xyflow/react' import '@xyflow/react/dist/style.css' @@ -36,7 +37,8 @@ const STORAGE_KEY = 'homelable_canvas' type ViewState = 'loading' | 'disabled' | 'invalid-key' | 'no-key' | 'network-error' | 'ready' function LiveViewCanvas() { - const { nodes, edges, loadCanvas } = useCanvasStore() + const { nodes, edges, loadCanvas, fitViewPending, clearFitViewPending } = useCanvasStore() + const { fitView } = useReactFlow() const activeTheme = useThemeStore((s) => s.activeTheme) const theme = THEMES[activeTheme] // Derive initial view state synchronously (avoids calling setState inside an effect): @@ -87,6 +89,15 @@ function LiveViewCanvas() { }) }, [loadCanvas]) + useEffect(() => { + if (!fitViewPending || nodes.length === 0) return + const id = setTimeout(() => { + fitView({ padding: 0.12, duration: 350 }) + clearFitViewPending() + }, 50) + return () => clearTimeout(id) + }, [fitViewPending, nodes.length, fitView, clearFitViewPending]) + const onNodeClick = useCallback((_: React.MouseEvent, node: Node) => { const ip = node.data.ip if (ip) window.open(`http://${ip}`, '_blank', 'noopener,noreferrer') @@ -129,7 +140,6 @@ function LiveViewCanvas() { elementsSelectable={false} panOnDrag zoomOnScroll - fitView colorMode={theme.colors.reactFlowColorMode} connectionMode={ConnectionMode.Loose} onNodeClick={onNodeClick} diff --git a/frontend/src/components/__tests__/LiveView.test.tsx b/frontend/src/components/__tests__/LiveView.test.tsx index 8983796..f3c97de 100644 --- a/frontend/src/components/__tests__/LiveView.test.tsx +++ b/frontend/src/components/__tests__/LiveView.test.tsx @@ -11,6 +11,7 @@ vi.mock('@xyflow/react', () => ({ Controls: () => null, BackgroundVariant: { Dots: 'dots' }, ConnectionMode: { Loose: 'loose' }, + useReactFlow: () => ({ fitView: vi.fn() }), })) vi.mock('@xyflow/react/dist/style.css', () => ({})) @@ -50,6 +51,8 @@ describe('LiveView (non-standalone)', () => { useCanvasStore.setState({ nodes: [], edges: [] }) }) + afterEach(() => { setSearch('') }) + // ── No key ──────────────────────────────────────────────────────────────── it('shows no-key error when ?key= is missing', async () => { @@ -133,11 +136,25 @@ describe('LiveView (non-standalone)', () => { // ── Standalone mode ──────────────────────────────────────────────────────── +const XYFLOW_MOCK = { + ReactFlowProvider: ({ children }: { children: React.ReactNode }) => <>{children}, + ReactFlow: () =>
, + Background: () => null, + Controls: () => null, + BackgroundVariant: { Dots: 'dots' }, + ConnectionMode: { Loose: 'loose' }, + useReactFlow: () => ({ fitView: vi.fn() }), +} + describe('LiveView (standalone — localStorage)', () => { beforeEach(() => { localStorage.clear() useCanvasStore.setState({ nodes: [], edges: [] }) - vi.mocked(liveviewApi.load).mockReset() + }) + + afterEach(() => { + setSearch('') + vi.unstubAllEnvs() }) it('loads canvas from localStorage without calling the API', async () => { @@ -151,25 +168,12 @@ describe('LiveView (standalone — localStorage)', () => { } localStorage.setItem('homelable_canvas', JSON.stringify(stored)) - // Stub VITE_STANDALONE before re-importing - vi.stubEnv('VITE_STANDALONE', 'true') - vi.resetModules() - const { default: LiveViewStandalone } = await import('../LiveView') - - setSearch('') // no key needed in standalone - render() - - await waitFor(() => { - expect(screen.getByTestId('react-flow')).toBeDefined() - }) - expect(liveviewApi.load).not.toHaveBeenCalled() - - vi.unstubAllEnvs() - }) - - it('shows canvas (empty) when localStorage has no saved data', async () => { vi.stubEnv('VITE_STANDALONE', 'true') vi.resetModules() + const mockLoad = vi.fn() + vi.doMock('@xyflow/react', () => XYFLOW_MOCK) + vi.doMock('@xyflow/react/dist/style.css', () => ({})) + vi.doMock('@/api/client', () => ({ liveviewApi: { load: mockLoad } })) const { default: LiveViewStandalone } = await import('../LiveView') setSearch('') @@ -178,8 +182,24 @@ describe('LiveView (standalone — localStorage)', () => { await waitFor(() => { expect(screen.getByTestId('react-flow')).toBeDefined() }) - expect(liveviewApi.load).not.toHaveBeenCalled() + expect(mockLoad).not.toHaveBeenCalled() + }) - vi.unstubAllEnvs() + it('shows canvas (empty) when localStorage has no saved data', async () => { + vi.stubEnv('VITE_STANDALONE', 'true') + vi.resetModules() + const mockLoad = vi.fn() + vi.doMock('@xyflow/react', () => XYFLOW_MOCK) + vi.doMock('@xyflow/react/dist/style.css', () => ({})) + vi.doMock('@/api/client', () => ({ liveviewApi: { load: mockLoad } })) + const { default: LiveViewStandalone } = await import('../LiveView') + + setSearch('') + render() + + await waitFor(() => { + expect(screen.getByTestId('react-flow')).toBeDefined() + }) + expect(mockLoad).not.toHaveBeenCalled() }) }) diff --git a/frontend/src/components/canvas/CanvasContainer.tsx b/frontend/src/components/canvas/CanvasContainer.tsx index 330e3da..486afe1 100644 --- a/frontend/src/components/canvas/CanvasContainer.tsx +++ b/frontend/src/components/canvas/CanvasContainer.tsx @@ -1,4 +1,4 @@ -import { useCallback, useState } from 'react' +import { useCallback, useEffect, useState } from 'react' import { ReactFlow, Background, @@ -7,6 +7,7 @@ import { BackgroundVariant, ConnectionMode, SelectionMode, + useReactFlow, type Node, type Edge, type Connection, @@ -33,7 +34,19 @@ export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, o nodes, edges, onNodesChange, onEdgesChange, setSelectedNode, snapshotHistory, + fitViewPending, clearFitViewPending, } = useCanvasStore() + const { fitView } = useReactFlow() + + // Fit view after canvas loads (fitViewPending is set by loadCanvas) + useEffect(() => { + if (!fitViewPending || nodes.length === 0) return + const id = setTimeout(() => { + fitView({ padding: 0.12, duration: 350 }) + clearFitViewPending() + }, 50) + return () => clearTimeout(id) + }, [fitViewPending, nodes.length, fitView, clearFitViewPending]) const activeTheme = useThemeStore((s) => s.activeTheme) const theme = THEMES[activeTheme] @@ -77,7 +90,6 @@ export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, o multiSelectionKeyCode={['Meta', 'Control']} snapToGrid snapGrid={[16, 16]} - fitView colorMode={theme.colors.reactFlowColorMode} elevateNodesOnSelect={false} connectionMode={ConnectionMode.Loose} diff --git a/frontend/src/components/canvas/__tests__/CanvasContainer.test.tsx b/frontend/src/components/canvas/__tests__/CanvasContainer.test.tsx index afebdad..721e588 100644 --- a/frontend/src/components/canvas/__tests__/CanvasContainer.test.tsx +++ b/frontend/src/components/canvas/__tests__/CanvasContainer.test.tsx @@ -20,6 +20,7 @@ vi.mock('@xyflow/react', () => ({ BackgroundVariant: { Dots: 'dots' }, ConnectionMode: { Loose: 'loose' }, SelectionMode: { Partial: 'partial' }, + useReactFlow: () => ({ fitView: vi.fn() }), })) vi.mock('@xyflow/react/dist/style.css', () => ({})) diff --git a/frontend/src/components/canvas/edges/index.tsx b/frontend/src/components/canvas/edges/index.tsx index 6072b91..9917d25 100644 --- a/frontend/src/components/canvas/edges/index.tsx +++ b/frontend/src/components/canvas/edges/index.tsx @@ -26,7 +26,7 @@ export function HomelableEdge({ id, source, target, sourceX, sourceY, targetX, t const isBidirectional = sourceType === 'proxmox' && targetType === 'proxmox' const pathArgs = { sourceX, sourceY, sourcePosition, targetX, targetY, targetPosition } - const [edgePath, labelX, labelY] = data?.path_style === 'smooth' + const [edgePath, labelX] = data?.path_style === 'smooth' ? getSmoothStepPath({ ...pathArgs, borderRadius: 8 }) : getBezierPath(pathArgs) @@ -95,9 +95,9 @@ export function HomelableEdge({ id, source, target, sourceX, sourceY, targetX, t {data?.label && (
> { icon: LucideIcon @@ -18,7 +19,10 @@ function formatStorage(gb: number): string { return `${gb} GB` } -export function BaseNode({ data, selected, icon: typeIcon, width, height }: BaseNodeProps) { +export function BaseNode({ id, data, selected, icon: typeIcon, width, height }: BaseNodeProps) { + const updateNodeInternals = useUpdateNodeInternals() + useEffect(() => { updateNodeInternals(id) }, [data.bottom_handles, id, updateNodeInternals]) + const activeTheme = useThemeStore((s) => s.activeTheme) const hideIp = useCanvasStore((s) => s.hideIp) const theme = THEMES[activeTheme] @@ -141,13 +145,26 @@ export function BaseNode({ data, selected, icon: typeIcon, width, height }: Base title={data.status} /> - - + {(BOTTOM_HANDLE_POSITIONS[data.bottom_handles ?? 1] ?? BOTTOM_HANDLE_POSITIONS[1]).map((leftPct, idx) => { + const sourceId = BOTTOM_HANDLE_IDS[idx] + const targetId = idx === 0 ? 'bottom-t' : `bottom-${idx + 1}-t` + return ( + + + + + ) + })}
) } diff --git a/frontend/src/components/modals/NodeModal.tsx b/frontend/src/components/modals/NodeModal.tsx index 33ed7dc..f8417de 100644 --- a/frontend/src/components/modals/NodeModal.tsx +++ b/frontend/src/components/modals/NodeModal.tsx @@ -7,7 +7,7 @@ import { Label } from '@/components/ui/label' import { Select, SelectContent, SelectGroup, SelectItem, SelectLabel, SelectSeparator, SelectTrigger, SelectValue } from '@/components/ui/select' import { NODE_TYPE_LABELS, type NodeData, type NodeType, type CheckMethod } from '@/types' import { resolveNodeColors } from '@/utils/nodeColors' -import { ICON_REGISTRY, ICON_CATEGORIES } from '@/utils/nodeIcons' +import { ICON_REGISTRY, ICON_CATEGORIES, NODE_TYPE_DEFAULT_ICONS } from '@/utils/nodeIcons' const NODE_TYPE_GROUPS: { label: string; types: NodeType[] }[] = [ { label: 'Hardware', types: ['isp', 'router', 'switch', 'server', 'nas', 'ap', 'printer'] }, @@ -75,11 +75,11 @@ export function NodeModal({ open, onClose, onSubmit, initial, title = 'Add Node'
- {/* Type */} -
+ {/* Type + Icon on the same row */} +
setIconSearch(e.target.value)} - placeholder="Search icons…" - className="bg-[#21262d] border-[#30363d] text-xs h-7" - autoFocus - /> -
- {ICON_CATEGORIES.map((cat) => { - const entries = ICON_REGISTRY.filter( - (e) => e.category === cat && - (iconSearch === '' || e.label.toLowerCase().includes(iconSearch.toLowerCase()) || e.key.includes(iconSearch.toLowerCase())) - ) - if (entries.length === 0) return null - return ( -
-

{cat}

-
- {entries.map((entry) => { - const isSelected = form.custom_icon === entry.key - return ( - - ) - })} -
-
- ) - })} -
-
- )}
+ {/* Inline icon picker — full width, shown below the type+icon row */} + {iconPickerOpen && ( +
+ setIconSearch(e.target.value)} + placeholder="Search icons…" + className="bg-[#21262d] border-[#30363d] text-xs h-7" + autoFocus + /> +
+ {ICON_CATEGORIES.map((cat) => { + const entries = ICON_REGISTRY.filter( + (e) => e.category === cat && + (iconSearch === '' || e.label.toLowerCase().includes(iconSearch.toLowerCase()) || e.key.includes(iconSearch.toLowerCase())) + ) + if (entries.length === 0) return null + return ( +
+

{cat}

+
+ {entries.map((entry) => { + const isSelected = form.custom_icon === entry.key + return ( + + ) + })} +
+
+ ) + })} +
+
+ )} + {/* Label */}
@@ -414,6 +416,27 @@ export function NodeModal({ open, onClose, onSubmit, initial, title = 'Add Node'
)} + {/* Bottom connection points (not for group containers) */} + {form.type !== 'groupRect' && form.type !== 'group' && ( +
+ + +
+ )} + {/* Notes */}
diff --git a/frontend/src/components/modals/PendingDeviceModal.tsx b/frontend/src/components/modals/PendingDeviceModal.tsx index 34324cb..3044388 100644 --- a/frontend/src/components/modals/PendingDeviceModal.tsx +++ b/frontend/src/components/modals/PendingDeviceModal.tsx @@ -19,6 +19,7 @@ export interface PendingDevice { services: Service[] suggested_type: string | null status: string + discovery_source: string | null discovered_at: string } @@ -101,6 +102,9 @@ export function PendingDeviceModal({ device, onClose, onApprove, onHide, onIgnor {device.suggested_type && ( )} + {device.discovery_source && ( + + )}
diff --git a/frontend/src/components/modals/ScanConfigModal.tsx b/frontend/src/components/modals/ScanConfigModal.tsx index 0b33b66..bc184fb 100644 --- a/frontend/src/components/modals/ScanConfigModal.tsx +++ b/frontend/src/components/modals/ScanConfigModal.tsx @@ -99,7 +99,6 @@ export function ScanConfigModal({ open, onClose, onScanNow }: ScanConfigModalPro -