Merge pull request #40 from Pouzor/feat/bottom-handles

feat: configurable bottom handles, scanner rewrite, UI polish
This commit is contained in:
Remy
2026-04-04 23:23:20 +02:00
committed by GitHub
30 changed files with 1476 additions and 471 deletions
+40 -8
View File
@@ -1,8 +1,10 @@
import ipaddress
import logging import logging
import uuid
from typing import Any from typing import Any
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
from pydantic import BaseModel from pydantic import BaseModel, field_validator
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@@ -18,6 +20,16 @@ from app.services.scanner import request_cancel, run_scan
class ScanConfig(BaseModel): class ScanConfig(BaseModel):
ranges: list[str] 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__) logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
@@ -49,6 +61,10 @@ async def stop_scan(
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user), _: str = Depends(get_current_user),
) -> dict[str, bool]: ) -> 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) run = await db.get(ScanRun, run_id)
if not run: if not run:
raise HTTPException(status_code=404, detail="Scan run not found") 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()) 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]) @router.get("/hidden", response_model=list[PendingDeviceResponse])
async def list_hidden(db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)) -> list[PendingDevice]: 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")) 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) device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)
) -> dict[str, bool]: ) -> dict[str, bool]:
device = await db.get(PendingDevice, device_id) device = await db.get(PendingDevice, device_id)
if device: if not device:
device.status = "hidden" raise HTTPException(status_code=404, detail="Device not found")
await db.commit() device.status = "hidden"
await db.commit()
return {"hidden": True} return {"hidden": True}
@@ -103,9 +133,10 @@ async def ignore_device(
device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user) device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)
) -> dict[str, bool]: ) -> dict[str, bool]:
device = await db.get(PendingDevice, device_id) device = await db.get(PendingDevice, device_id)
if device: if not device:
await db.delete(device) raise HTTPException(status_code=404, detail="Device not found")
await db.commit() await db.delete(device)
await db.commit()
return {"ignored": True} return {"ignored": True}
@@ -127,4 +158,5 @@ async def update_scan_config(payload: ScanConfig, _: str = Depends(get_current_u
settings.save_overrides() settings.save_overrides()
return payload return payload
except Exception as exc: 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
+4
View File
@@ -57,6 +57,10 @@ async def init_db() -> None:
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN width REAL") await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN width REAL")
with suppress(OperationalError): with suppress(OperationalError):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN height REAL") 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') # Migrate animated column from boolean (0/1) to string ('none'/'snake')
with suppress(OperationalError): with suppress(OperationalError):
await conn.exec_driver_sql("UPDATE edges SET animated = 'snake' WHERE animated = '1' OR animated = 1") await conn.exec_driver_sql("UPDATE edges SET animated = 'snake' WHERE animated = '1' OR animated = 1")
+2
View File
@@ -44,6 +44,7 @@ class Node(Base):
show_hardware: Mapped[bool] = mapped_column(Boolean, default=False) show_hardware: Mapped[bool] = mapped_column(Boolean, default=False)
width: Mapped[float | None] = mapped_column(Float, nullable=True) width: Mapped[float | None] = mapped_column(Float, nullable=True)
height: 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)) last_seen: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
response_time_ms: Mapped[int | None] = mapped_column(Integer) response_time_ms: Mapped[int | None] = mapped_column(Integer)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now) 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) services: Mapped[list[Any]] = mapped_column(JSON, default=list)
suggested_type: Mapped[str | None] = mapped_column(String) suggested_type: Mapped[str | None] = mapped_column(String)
status: Mapped[str] = mapped_column(String, default="pending") 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) discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
+12
View File
@@ -1,3 +1,5 @@
import logging
import logging.config
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Any from typing import Any
@@ -14,6 +16,16 @@ from app.db.database import init_db
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: 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() await init_db()
settings.load_overrides() settings.load_overrides()
start_scheduler() start_scheduler()
+1
View File
@@ -31,6 +31,7 @@ class NodeSave(BaseModel):
show_hardware: bool = False show_hardware: bool = False
width: float | None = None width: float | None = None
height: float | None = None height: float | None = None
bottom_handles: int = 1
pos_x: float = 0 pos_x: float = 0
pos_y: float = 0 pos_y: float = 0
+2
View File
@@ -29,6 +29,7 @@ class NodeBase(BaseModel):
show_hardware: bool = False show_hardware: bool = False
width: float | None = None width: float | None = None
height: float | None = None height: float | None = None
bottom_handles: int = 1
class NodeCreate(NodeBase): class NodeCreate(NodeBase):
@@ -60,6 +61,7 @@ class NodeUpdate(BaseModel):
show_hardware: bool | None = None show_hardware: bool | None = None
width: float | None = None width: float | None = None
height: float | None = None height: float | None = None
bottom_handles: int | None = None
class NodeResponse(NodeBase): class NodeResponse(NodeBase):
+1
View File
@@ -13,6 +13,7 @@ class PendingDeviceResponse(BaseModel):
services: list[Any] services: list[Any]
suggested_type: str | None suggested_type: str | None
status: str status: str
discovery_source: str | None
discovered_at: datetime discovered_at: datetime
model_config = {"from_attributes": True} model_config = {"from_attributes": True}
+239 -90
View File
@@ -1,7 +1,12 @@
"""Network scanner: ARP sweep + nmap service detection + mDNS discovery.""" """Network scanner: ARP sweep + nmap service detection + mDNS discovery."""
import asyncio import asyncio
import ipaddress
import logging import logging
import os
import re
import socket import socket
import subprocess
import threading
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any from typing import Any
@@ -13,8 +18,9 @@ from app.services.fingerprint import fingerprint_ports, suggest_node_type
logger = logging.getLogger(__name__) 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_runs: set[str] = set()
_cancelled_lock = threading.Lock()
# Port list for service detection (Phase 2) # Port list for service detection (Phase 2)
_EXTRA_PORTS = ( _EXTRA_PORTS = (
@@ -55,11 +61,13 @@ except ImportError:
def request_cancel(run_id: str) -> None: def request_cancel(run_id: str) -> None:
"""Signal a running scan to stop early.""" """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: 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: def _resolve_hostname(ip: str) -> str | None:
@@ -79,80 +87,216 @@ def _extract_os(nm: object, host: str) -> str | None:
return 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. Read the OS ARP cache for recently-seen hosts in the target network.
Returns {ip: host_dict} for every host that responds. 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() try:
nm.scan(hosts=target, arguments="-sn -PR -PA80,443 --host-timeout 10s") 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]] = {} alive: dict[str, dict[str, Any]] = {}
for host in nm.all_hosts():
if nm[host].state() == "up": for ip in alive_ips:
alive[host] = { mac = arp_cache.get(ip, {}).get("mac")
"ip": host, hostname = await asyncio.to_thread(_resolve_hostname, ip)
"hostname": _resolve_hostname(host), logger.info("[Phase 1] %s mac=%s hostname=%s (ping)", ip, mac or "n/a", hostname or "n/a")
"mac": nm[host].get("addresses", {}).get("mac"), alive[ip] = {"ip": ip, "mac": mac, "hostname": hostname, "os": None, "open_ports": []}
"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 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. Phase 2 — single-IP port scan with service detection.
Mutates alive in-place with open_ports/os; returns all hosts including Runs in a thread (blocking). Returns the host dict enriched with open_ports.
those with zero open ports (IoT devices often have none). """
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: if not alive:
return [] 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(): logger.info("[Phase 2] Starting per-IP port scan for %d host(s)", len(alive))
if host not in alive: semaphore = asyncio.Semaphore(10)
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)
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. Two-phase scan for a CIDR range.
Phase 1: ARP sweep to find alive hosts (catches IoT with no open ports). Phase 1: Concurrent ping sweep to find alive hosts (fast, no false positives).
Phase 2: Service detection on alive hosts only. 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: if not _NMAP_AVAILABLE:
logger.warning("[Scan] nmap not available — returning mock data")
return _mock_scan(target) return _mock_scan(target)
try: 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: 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 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]]: 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 from app.api.routes.status import broadcast_scan_update
devices_found = 0 devices_found = 0
mdns_task: asyncio.Task[list[dict[str, Any]]] | None = None
try: 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_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()} 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: if canvas_ips:
stale_result = await db.execute( from sqlalchemy import delete as sa_delete
select(PendingDevice).where( await db.execute(
sa_delete(PendingDevice).where(
PendingDevice.status == "pending", PendingDevice.status == "pending",
PendingDevice.ip.in_(canvas_ips), PendingDevice.ip.in_(canvas_ips),
) )
) )
for stale in stale_result.scalars().all():
await db.delete(stale)
await db.commit() await db.commit()
# Start mDNS discovery in the background while nmap scans run # Start mDNS discovery in the background while nmap scans run
mdns_task: asyncio.Task[list[dict[str, Any]]] = asyncio.create_task( mdns_task = asyncio.create_task(_mdns_discover())
_mdns_discover()
)
# Track IPs found by nmap so mDNS doesn't duplicate them # Track IPs found by nmap so mDNS doesn't duplicate them
nmap_ips: set[str] = set() 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 nonlocal devices_found
ip = host["ip"] ip = host["ip"]
# Skip canvas nodes and user-hidden devices # Skip canvas nodes and user-hidden devices (sets pre-fetched before loop)
canvas_result = await db.execute(select(Node).where(Node.ip == ip)) if ip in canvas_ips:
if canvas_result.scalar_one_or_none() is not None:
logger.debug("Skipping %s — already in canvas", ip) logger.debug("Skipping %s — already in canvas", ip)
return return
hidden_result = await db.execute( if ip in hidden_ips:
select(PendingDevice).where(
PendingDevice.ip == ip,
PendingDevice.status == "hidden",
)
)
if hidden_result.scalar_one_or_none() is not None:
logger.debug("Skipping %s — hidden by user", ip) logger.debug("Skipping %s — hidden by user", ip)
return return
@@ -303,43 +452,40 @@ async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None:
services=services, services=services,
suggested_type=suggested_type, suggested_type=suggested_type,
status="pending", status="pending",
discovery_source=discovery_source,
)) ))
devices_found += 1 devices_found += 1
await db.commit() 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) await broadcast_scan_update(run_id=run_id, devices_found=devices_found)
# nmap scan per CIDR — results stream in progressively # nmap scan per CIDR — results stream in progressively
for cidr in ranges: for cidr in ranges:
if _is_cancelled(run_id): if _is_cancelled(run_id):
break break
hosts = await asyncio.to_thread(_nmap_scan, cidr) hosts = await _nmap_scan(cidr)
for host in hosts: for host in hosts:
if _is_cancelled(run_id): if _is_cancelled(run_id):
break break
nmap_ips.add(host["ip"]) nmap_ips.add(host["ip"])
await _process_host(host) 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): if not _is_cancelled(run_id):
try: mdns_hosts = await mdns_task
mdns_hosts = await asyncio.wait_for(mdns_task, timeout=1.0)
except asyncio.TimeoutError:
mdns_task.cancel()
mdns_hosts = []
for host in mdns_hosts: for host in mdns_hosts:
if _is_cancelled(run_id): if _is_cancelled(run_id):
break break
if host["ip"] in nmap_ips: if host["ip"] in nmap_ips:
continue # already processed with richer nmap data continue # already processed with richer nmap data
await _process_host(host) await _process_host(host, discovery_source="mdns")
else: else:
mdns_task.cancel() mdns_task.cancel()
@@ -353,6 +499,8 @@ async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None:
except Exception as exc: except Exception as exc:
logger.error("Scan failed: %s", 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) run = await db.get(ScanRun, run_id)
if run: if run:
run.status = "error" 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) run.finished_at = datetime.now(timezone.utc)
await db.commit() await db.commit()
finally: finally:
_cancelled_runs.discard(run_id) with _cancelled_lock:
_cancelled_runs.discard(run_id)
+2 -1
View File
@@ -322,7 +322,8 @@ async def test_stop_scan_requires_auth(client: AsyncClient):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_stop_scan_not_found(client: AsyncClient, headers): 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 assert res.status_code == 404
+316 -87
View File
@@ -3,10 +3,12 @@ import uuid
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from sqlalchemy import select as sa_select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine 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.database import Base
from app.db.models import PendingDevice, ScanRun from app.db.models import Node, PendingDevice, ScanRun
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Helpers # Helpers
@@ -18,7 +20,11 @@ def _make_run_id() -> str:
@pytest.fixture @pytest.fixture
async def mem_db(): 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: async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all) await conn.run_sync(Base.metadata.create_all)
factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) 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(): @pytest.mark.asyncio
from app.services.scanner import _nmap_arp_sweep async def test_ping_sweep_returns_alive_hosts():
from app.services.scanner import _ping_sweep
mock_nm = MagicMock() async def fake_ping(ip: str) -> str | None:
mock_nm.all_hosts.return_value = ["192.168.1.1", "192.168.1.2"] return ip if ip in {"192.168.1.1", "192.168.1.2"} else None
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,
)
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): 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(): 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(): @pytest.mark.asyncio
from app.services.scanner import _nmap_arp_sweep 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() with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \
mock_nm.all_hosts.return_value = list(states.keys()) patch("app.services.scanner._arp_table_hosts", return_value={}), \
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), \
patch("app.services.scanner._resolve_hostname", return_value=None): 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.1" in result
assert "192.168.1.2" not in result assert "192.168.1.2" not in result
# --------------------------------------------------------------------------- @pytest.mark.asyncio
# _nmap_port_scan 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(): async def mock_subprocess(*args, **kwargs):
from app.services.scanner import _nmap_port_scan proc = MagicMock()
proc.returncode = 1 # all pings fail
proc.wait = AsyncMock(return_value=1)
return proc
alive = { arp_extra = {
"192.168.1.10": {"ip": "192.168.1.10", "hostname": None, "mac": None, "os": None, "open_ports": []}, "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 = MagicMock()
mock_nm.all_hosts.return_value = ["192.168.1.10"] mock_nm.all_hosts.return_value = ["192.168.1.10"]
mock_nm.__getitem__ = lambda self, host: MagicMock( mock_nm.__getitem__ = MagicMock(return_value=mock_host)
all_protocols=lambda: ["tcp"],
**{"__getitem__": lambda self2, proto: {
80: {"state": "open", "product": "nginx", "version": "1.24"},
}},
get=lambda key, default=None: default,
)
with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \ with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \
patch("app.services.scanner._extract_os", return_value=None): patch("app.services.scanner._extract_os", return_value=None):
result = _nmap_port_scan(alive) result = _nmap_scan_single(host)
assert len(result) == 1 assert len(result["open_ports"]) == 1
assert result[0]["open_ports"][0]["port"] == 80 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(): def test_nmap_scan_single_returns_host_unchanged_on_error():
from app.services.scanner import _nmap_port_scan from app.services.scanner import _nmap_scan_single
alive = {
"192.168.1.20": {"ip": "192.168.1.20", "hostname": None, "mac": None, "os": None, "open_ports": []},
}
host = {"ip": "192.168.1.20", "hostname": None, "mac": None, "os": None, "open_ports": []}
mock_nm = MagicMock() mock_nm = MagicMock()
mock_nm.scan.side_effect = Exception("nmap error") mock_nm.scan.side_effect = Exception("nmap error")
with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm): 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 result["ip"] == "192.168.1.20"
assert len(result) == 1 assert result["open_ports"] == []
assert result[0]["ip"] == "192.168.1.20"
assert result[0]["open_ports"] == []
def test_nmap_port_scan_includes_hosts_with_no_open_ports(): def test_nmap_scan_single_returns_host_unchanged_when_no_results():
"""IoT devices found by ARP but with no open TCP ports must still be returned.""" """Host confirmed alive in Phase 1 but all ports filtered — keep it with empty ports."""
from app.services.scanner import _nmap_port_scan from app.services.scanner import _nmap_scan_single
alive = { host = {"ip": "192.168.1.30", "hostname": "shelly1.lan", "mac": "34:94:54:aa:bb:cc", "os": None, "open_ports": []}
"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)
mock_nm = MagicMock() mock_nm = MagicMock()
mock_nm.all_hosts.return_value = ["192.168.1.31"] mock_nm.all_hosts.return_value = [] # no results
mock_nm.__getitem__ = lambda self, host: MagicMock(
all_protocols=lambda: [],
get=lambda key, default=None: default,
)
with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \ with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm):
patch("app.services.scanner._extract_os", return_value=None): result = _nmap_scan_single(host)
result = _nmap_port_scan(alive)
ips = {h["ip"] for h in result} assert result["ip"] == "192.168.1.30"
assert "192.168.1.30" in ips, "ARP-found device with no open ports must still be returned" assert result["open_ports"] == []
assert "192.168.1.31" in ips 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 from app.services.scanner import _nmap_scan
with patch("app.services.scanner._NMAP_AVAILABLE", False): 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 len(result) == 1
assert result[0]["ip"] == "192.168.1.99" 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 # _mdns_discover
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -223,6 +335,47 @@ async def test_mdns_discover_returns_devices():
assert result[0]["hostname"] == "shelly1.local." 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 # 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) await run_scan(["192.168.1.0/24"], session, run_id)
async with mem_db() as session: async with mem_db() as session:
result = await session.execute( result = await session.execute(sa_select(PendingDevice))
__import__("sqlalchemy", fromlist=["select"]).select(PendingDevice)
)
devices = result.scalars().all() devices = result.scalars().all()
assert any(d.ip == "192.168.1.5" for d in devices) 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) await run_scan(["192.168.1.0/24"], session, run_id)
async with mem_db() as session: 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")) result = await session.execute(sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.80"))
device = result.scalar_one_or_none() device = result.scalar_one_or_none()
assert device is not None assert device is not None
assert device.status == "pending" assert device.status == "pending"
assert device.discovery_source == "mdns"
@pytest.mark.asyncio @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) await run_scan(["192.168.1.0/24"], session, run_id)
async with mem_db() as session: 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")) result = await session.execute(sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.10"))
devices = result.scalars().all() devices = result.scalars().all()
assert len(devices) == 1 # not duplicated 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"
+1
View File
@@ -56,6 +56,7 @@ export const scanApi = {
pending: () => api.get('/scan/pending'), pending: () => api.get('/scan/pending'),
hidden: () => api.get('/scan/hidden'), hidden: () => api.get('/scan/hidden'),
runs: () => api.get('/scan/runs'), 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(`/scan/pending/${id}/approve`, nodeData),
hide: (id: string) => api.post(`/scan/pending/${id}/hide`), hide: (id: string) => api.post(`/scan/pending/${id}/hide`),
ignore: (id: string) => api.post(`/scan/pending/${id}/ignore`), ignore: (id: string) => api.post(`/scan/pending/${id}/ignore`),
+12 -2
View File
@@ -18,6 +18,7 @@ import {
BackgroundVariant, BackgroundVariant,
Controls, Controls,
ConnectionMode, ConnectionMode,
useReactFlow,
type Node, type Node,
} from '@xyflow/react' } from '@xyflow/react'
import '@xyflow/react/dist/style.css' 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' type ViewState = 'loading' | 'disabled' | 'invalid-key' | 'no-key' | 'network-error' | 'ready'
function LiveViewCanvas() { function LiveViewCanvas() {
const { nodes, edges, loadCanvas } = useCanvasStore() const { nodes, edges, loadCanvas, fitViewPending, clearFitViewPending } = useCanvasStore()
const { fitView } = useReactFlow()
const activeTheme = useThemeStore((s) => s.activeTheme) const activeTheme = useThemeStore((s) => s.activeTheme)
const theme = THEMES[activeTheme] const theme = THEMES[activeTheme]
// Derive initial view state synchronously (avoids calling setState inside an effect): // Derive initial view state synchronously (avoids calling setState inside an effect):
@@ -87,6 +89,15 @@ function LiveViewCanvas() {
}) })
}, [loadCanvas]) }, [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<NodeData>) => { const onNodeClick = useCallback((_: React.MouseEvent, node: Node<NodeData>) => {
const ip = node.data.ip const ip = node.data.ip
if (ip) window.open(`http://${ip}`, '_blank', 'noopener,noreferrer') if (ip) window.open(`http://${ip}`, '_blank', 'noopener,noreferrer')
@@ -129,7 +140,6 @@ function LiveViewCanvas() {
elementsSelectable={false} elementsSelectable={false}
panOnDrag panOnDrag
zoomOnScroll zoomOnScroll
fitView
colorMode={theme.colors.reactFlowColorMode} colorMode={theme.colors.reactFlowColorMode}
connectionMode={ConnectionMode.Loose} connectionMode={ConnectionMode.Loose}
onNodeClick={onNodeClick} onNodeClick={onNodeClick}
@@ -11,6 +11,7 @@ vi.mock('@xyflow/react', () => ({
Controls: () => null, Controls: () => null,
BackgroundVariant: { Dots: 'dots' }, BackgroundVariant: { Dots: 'dots' },
ConnectionMode: { Loose: 'loose' }, ConnectionMode: { Loose: 'loose' },
useReactFlow: () => ({ fitView: vi.fn() }),
})) }))
vi.mock('@xyflow/react/dist/style.css', () => ({})) vi.mock('@xyflow/react/dist/style.css', () => ({}))
@@ -50,6 +51,8 @@ describe('LiveView (non-standalone)', () => {
useCanvasStore.setState({ nodes: [], edges: [] }) useCanvasStore.setState({ nodes: [], edges: [] })
}) })
afterEach(() => { setSearch('') })
// ── No key ──────────────────────────────────────────────────────────────── // ── No key ────────────────────────────────────────────────────────────────
it('shows no-key error when ?key= is missing', async () => { it('shows no-key error when ?key= is missing', async () => {
@@ -133,11 +136,25 @@ describe('LiveView (non-standalone)', () => {
// ── Standalone mode ──────────────────────────────────────────────────────── // ── Standalone mode ────────────────────────────────────────────────────────
const XYFLOW_MOCK = {
ReactFlowProvider: ({ children }: { children: React.ReactNode }) => <>{children}</>,
ReactFlow: () => <div data-testid="react-flow" />,
Background: () => null,
Controls: () => null,
BackgroundVariant: { Dots: 'dots' },
ConnectionMode: { Loose: 'loose' },
useReactFlow: () => ({ fitView: vi.fn() }),
}
describe('LiveView (standalone — localStorage)', () => { describe('LiveView (standalone — localStorage)', () => {
beforeEach(() => { beforeEach(() => {
localStorage.clear() localStorage.clear()
useCanvasStore.setState({ nodes: [], edges: [] }) useCanvasStore.setState({ nodes: [], edges: [] })
vi.mocked(liveviewApi.load).mockReset() })
afterEach(() => {
setSearch('')
vi.unstubAllEnvs()
}) })
it('loads canvas from localStorage without calling the API', async () => { 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)) 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(<LiveViewStandalone />)
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.stubEnv('VITE_STANDALONE', 'true')
vi.resetModules() 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') const { default: LiveViewStandalone } = await import('../LiveView')
setSearch('') setSearch('')
@@ -178,8 +182,24 @@ describe('LiveView (standalone — localStorage)', () => {
await waitFor(() => { await waitFor(() => {
expect(screen.getByTestId('react-flow')).toBeDefined() 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(<LiveViewStandalone />)
await waitFor(() => {
expect(screen.getByTestId('react-flow')).toBeDefined()
})
expect(mockLoad).not.toHaveBeenCalled()
}) })
}) })
@@ -1,4 +1,4 @@
import { useCallback, useState } from 'react' import { useCallback, useEffect, useState } from 'react'
import { import {
ReactFlow, ReactFlow,
Background, Background,
@@ -7,6 +7,7 @@ import {
BackgroundVariant, BackgroundVariant,
ConnectionMode, ConnectionMode,
SelectionMode, SelectionMode,
useReactFlow,
type Node, type Node,
type Edge, type Edge,
type Connection, type Connection,
@@ -33,7 +34,19 @@ export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, o
nodes, edges, nodes, edges,
onNodesChange, onEdgesChange, onNodesChange, onEdgesChange,
setSelectedNode, snapshotHistory, setSelectedNode, snapshotHistory,
fitViewPending, clearFitViewPending,
} = useCanvasStore() } = 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 activeTheme = useThemeStore((s) => s.activeTheme)
const theme = THEMES[activeTheme] const theme = THEMES[activeTheme]
@@ -77,7 +90,6 @@ export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, o
multiSelectionKeyCode={['Meta', 'Control']} multiSelectionKeyCode={['Meta', 'Control']}
snapToGrid snapToGrid
snapGrid={[16, 16]} snapGrid={[16, 16]}
fitView
colorMode={theme.colors.reactFlowColorMode} colorMode={theme.colors.reactFlowColorMode}
elevateNodesOnSelect={false} elevateNodesOnSelect={false}
connectionMode={ConnectionMode.Loose} connectionMode={ConnectionMode.Loose}
@@ -20,6 +20,7 @@ vi.mock('@xyflow/react', () => ({
BackgroundVariant: { Dots: 'dots' }, BackgroundVariant: { Dots: 'dots' },
ConnectionMode: { Loose: 'loose' }, ConnectionMode: { Loose: 'loose' },
SelectionMode: { Partial: 'partial' }, SelectionMode: { Partial: 'partial' },
useReactFlow: () => ({ fitView: vi.fn() }),
})) }))
vi.mock('@xyflow/react/dist/style.css', () => ({})) vi.mock('@xyflow/react/dist/style.css', () => ({}))
@@ -26,7 +26,7 @@ export function HomelableEdge({ id, source, target, sourceX, sourceY, targetX, t
const isBidirectional = sourceType === 'proxmox' && targetType === 'proxmox' const isBidirectional = sourceType === 'proxmox' && targetType === 'proxmox'
const pathArgs = { sourceX, sourceY, sourcePosition, targetX, targetY, targetPosition } 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 }) ? getSmoothStepPath({ ...pathArgs, borderRadius: 8 })
: getBezierPath(pathArgs) : getBezierPath(pathArgs)
@@ -95,9 +95,9 @@ export function HomelableEdge({ id, source, target, sourceX, sourceY, targetX, t
{data?.label && ( {data?.label && (
<EdgeLabelRenderer> <EdgeLabelRenderer>
<div <div
className="absolute pointer-events-none font-mono text-[10px] px-1 rounded" className="absolute pointer-events-none font-mono text-[10px] px-1.5 py-0.5 rounded"
style={{ style={{
transform: `translate(-50%, -50%) translate(${labelX}px, ${labelY}px)`, transform: `translate(-50%, -50%) translate(${labelX}px, ${(sourceY + targetY) / 2}px)`,
background: theme.colors.edgeLabelBackground, background: theme.colors.edgeLabelBackground,
color: theme.colors.edgeLabelColor, color: theme.colors.edgeLabelColor,
border: `1px solid ${theme.colors.edgeLabelBorder}`, border: `1px solid ${theme.colors.edgeLabelBorder}`,
@@ -1,5 +1,5 @@
import { createElement } from 'react' import { createElement, useEffect } from 'react'
import { Handle, Position, NodeResizer, type NodeProps, type Node } from '@xyflow/react' import { Handle, Position, NodeResizer, useUpdateNodeInternals, type NodeProps, type Node } from '@xyflow/react'
import { Cpu, MemoryStick, HardDrive, type LucideIcon } from 'lucide-react' import { Cpu, MemoryStick, HardDrive, type LucideIcon } from 'lucide-react'
import type { NodeData } from '@/types' import type { NodeData } from '@/types'
import { resolveNodeColors } from '@/utils/nodeColors' import { resolveNodeColors } from '@/utils/nodeColors'
@@ -8,6 +8,7 @@ import { useThemeStore } from '@/stores/themeStore'
import { THEMES } from '@/utils/themes' import { THEMES } from '@/utils/themes'
import { useCanvasStore } from '@/stores/canvasStore' import { useCanvasStore } from '@/stores/canvasStore'
import { maskIp } from '@/utils/maskIp' import { maskIp } from '@/utils/maskIp'
import { BOTTOM_HANDLE_IDS, BOTTOM_HANDLE_POSITIONS } from '@/utils/handleUtils'
interface BaseNodeProps extends NodeProps<Node<NodeData>> { interface BaseNodeProps extends NodeProps<Node<NodeData>> {
icon: LucideIcon icon: LucideIcon
@@ -18,7 +19,10 @@ function formatStorage(gb: number): string {
return `${gb} GB` 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 activeTheme = useThemeStore((s) => s.activeTheme)
const hideIp = useCanvasStore((s) => s.hideIp) const hideIp = useCanvasStore((s) => s.hideIp)
const theme = THEMES[activeTheme] const theme = THEMES[activeTheme]
@@ -141,13 +145,26 @@ export function BaseNode({ data, selected, icon: typeIcon, width, height }: Base
title={data.status} title={data.status}
/> />
<Handle {(BOTTOM_HANDLE_POSITIONS[data.bottom_handles ?? 1] ?? BOTTOM_HANDLE_POSITIONS[1]).map((leftPct, idx) => {
type="source" const sourceId = BOTTOM_HANDLE_IDS[idx]
position={Position.Bottom} const targetId = idx === 0 ? 'bottom-t' : `bottom-${idx + 1}-t`
id="bottom" return (
style={{ background: theme.colors.handleBackground, borderColor: theme.colors.handleBorder }} <span key={sourceId}>
/> <Handle
<Handle type="target" position={Position.Bottom} id="bottom-t" style={{ opacity: 0, width: 12, height: 12 }} /> type="source"
position={Position.Bottom}
id={sourceId}
style={{ left: `${leftPct}%`, background: theme.colors.handleBackground, borderColor: theme.colors.handleBorder }}
/>
<Handle
type="target"
position={Position.Bottom}
id={targetId}
style={{ left: `${leftPct}%`, opacity: 0, width: 12, height: 12 }}
/>
</span>
)
})}
</div> </div>
) )
} }
+81 -58
View File
@@ -7,7 +7,7 @@ import { Label } from '@/components/ui/label'
import { Select, SelectContent, SelectGroup, SelectItem, SelectLabel, SelectSeparator, SelectTrigger, SelectValue } from '@/components/ui/select' 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 { NODE_TYPE_LABELS, type NodeData, type NodeType, type CheckMethod } from '@/types'
import { resolveNodeColors } from '@/utils/nodeColors' 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[] }[] = [ const NODE_TYPE_GROUPS: { label: string; types: NodeType[] }[] = [
{ label: 'Hardware', types: ['isp', 'router', 'switch', 'server', 'nas', 'ap', 'printer'] }, { label: 'Hardware', types: ['isp', 'router', 'switch', 'server', 'nas', 'ap', 'printer'] },
@@ -75,11 +75,11 @@ export function NodeModal({ open, onClose, onSubmit, initial, title = 'Add Node'
<form onSubmit={handleSubmit} className="flex flex-col gap-4 mt-2"> <form onSubmit={handleSubmit} className="flex flex-col gap-4 mt-2">
<div className="grid grid-cols-2 gap-3"> <div className="grid grid-cols-2 gap-3">
{/* Type */} {/* Type + Icon on the same row */}
<div className="flex flex-col gap-1.5 col-span-2"> <div className="flex flex-col gap-1.5">
<Label className="text-xs text-muted-foreground">Type</Label> <Label className="text-xs text-muted-foreground">Type</Label>
<Select value={form.type} onValueChange={(v) => set('type', v as NodeType)}> <Select value={form.type} onValueChange={(v) => set('type', v as NodeType)}>
<SelectTrigger className="bg-[#21262d] border-[#30363d] text-sm h-8"> <SelectTrigger className="bg-[#21262d] border-[#30363d] text-sm h-8 w-full">
<SelectValue /> <SelectValue />
</SelectTrigger> </SelectTrigger>
<SelectContent className="bg-[#21262d] border-[#30363d]"> <SelectContent className="bg-[#21262d] border-[#30363d]">
@@ -103,7 +103,7 @@ export function NodeModal({ open, onClose, onSubmit, initial, title = 'Add Node'
</div> </div>
{/* Icon */} {/* Icon */}
<div className="flex flex-col gap-1.5 col-span-2"> <div className="flex flex-col gap-1.5">
<div className="flex items-center justify-between"> <div className="flex items-center justify-between">
<Label className="text-xs text-muted-foreground">Icon</Label> <Label className="text-xs text-muted-foreground">Icon</Label>
{form.custom_icon && ( {form.custom_icon && (
@@ -120,70 +120,72 @@ export function NodeModal({ open, onClose, onSubmit, initial, title = 'Add Node'
<button <button
type="button" type="button"
onClick={() => setIconPickerOpen((o) => !o)} onClick={() => setIconPickerOpen((o) => !o)}
className="flex items-center justify-between gap-2 h-8 px-3 rounded-md bg-[#21262d] border border-[#30363d] text-sm hover:border-[#8b949e] transition-colors" className="flex items-center justify-between gap-2 h-8 px-3 rounded-md bg-[#21262d] border border-[#30363d] text-sm hover:border-[#8b949e] transition-colors w-full"
> >
<span className="flex items-center gap-2"> <span className="flex items-center gap-2 min-w-0">
{(() => { {(() => {
const entry = ICON_REGISTRY.find((e) => e.key === form.custom_icon) const entry = ICON_REGISTRY.find((e) => e.key === form.custom_icon)
if (entry) { if (entry) {
return <>{createElement(entry.icon, { size: 13, className: 'text-[#00d4ff]' })}<span className="text-foreground">{entry.label}</span></> return <>{createElement(entry.icon, { size: 13, className: 'text-[#00d4ff] shrink-0' })}<span className="text-foreground truncate">{entry.label}</span></>
} }
return <span className="text-muted-foreground">Default (from type)</span> const defaultIcon = NODE_TYPE_DEFAULT_ICONS[form.type as NodeType] ?? NODE_TYPE_DEFAULT_ICONS.generic
return <>{createElement(defaultIcon, { size: 13, className: 'text-muted-foreground shrink-0' })}<span className="text-muted-foreground truncate">Default</span></>
})()} })()}
</span> </span>
<ChevronDown size={12} className="text-muted-foreground shrink-0" style={{ transform: iconPickerOpen ? 'rotate(180deg)' : undefined, transition: 'transform 0.15s' }} /> <ChevronDown size={12} className="text-muted-foreground shrink-0" style={{ transform: iconPickerOpen ? 'rotate(180deg)' : undefined, transition: 'transform 0.15s' }} />
</button> </button>
{/* Inline picker panel */}
{iconPickerOpen && (
<div className="flex flex-col gap-2 p-2.5 rounded-md bg-[#0d1117] border border-[#30363d]">
<Input
value={iconSearch}
onChange={(e) => setIconSearch(e.target.value)}
placeholder="Search icons…"
className="bg-[#21262d] border-[#30363d] text-xs h-7"
autoFocus
/>
<div className="flex flex-col gap-2 max-h-52 overflow-y-auto">
{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 (
<div key={cat}>
<p className="text-[9px] font-semibold text-muted-foreground/50 uppercase tracking-wider mb-1">{cat}</p>
<div className="grid grid-cols-7 gap-1">
{entries.map((entry) => {
const isSelected = form.custom_icon === entry.key
return (
<button
key={entry.key}
type="button"
title={entry.label}
onClick={() => { set('custom_icon', isSelected ? undefined : entry.key); setIconPickerOpen(false) }}
className="flex items-center justify-center w-7 h-7 rounded transition-colors"
style={{
background: isSelected ? '#00d4ff22' : 'transparent',
border: isSelected ? '1px solid #00d4ff88' : '1px solid transparent',
color: isSelected ? '#00d4ff' : '#8b949e',
}}
onMouseEnter={(e) => { if (!isSelected) (e.currentTarget as HTMLButtonElement).style.background = '#21262d' }}
onMouseLeave={(e) => { if (!isSelected) (e.currentTarget as HTMLButtonElement).style.background = 'transparent' }}
>
{createElement(entry.icon, { size: 13 })}
</button>
)
})}
</div>
</div>
)
})}
</div>
</div>
)}
</div> </div>
{/* Inline icon picker — full width, shown below the type+icon row */}
{iconPickerOpen && (
<div className="flex flex-col gap-2 p-2.5 rounded-md bg-[#0d1117] border border-[#30363d] col-span-2">
<Input
value={iconSearch}
onChange={(e) => setIconSearch(e.target.value)}
placeholder="Search icons…"
className="bg-[#21262d] border-[#30363d] text-xs h-7"
autoFocus
/>
<div className="flex flex-col gap-2 max-h-52 overflow-y-auto">
{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 (
<div key={cat}>
<p className="text-[9px] font-semibold text-muted-foreground/50 uppercase tracking-wider mb-1">{cat}</p>
<div className="grid grid-cols-7 gap-1">
{entries.map((entry) => {
const isSelected = form.custom_icon === entry.key
return (
<button
key={entry.key}
type="button"
title={entry.label}
onClick={() => { set('custom_icon', isSelected ? undefined : entry.key); setIconPickerOpen(false) }}
className="flex items-center justify-center w-7 h-7 rounded transition-colors"
style={{
background: isSelected ? '#00d4ff22' : 'transparent',
border: isSelected ? '1px solid #00d4ff88' : '1px solid transparent',
color: isSelected ? '#00d4ff' : '#8b949e',
}}
onMouseEnter={(e) => { if (!isSelected) (e.currentTarget as HTMLButtonElement).style.background = '#21262d' }}
onMouseLeave={(e) => { if (!isSelected) (e.currentTarget as HTMLButtonElement).style.background = 'transparent' }}
>
{createElement(entry.icon, { size: 13 })}
</button>
)
})}
</div>
</div>
)
})}
</div>
</div>
)}
{/* Label */} {/* Label */}
<div className="flex flex-col gap-1.5 col-span-2"> <div className="flex flex-col gap-1.5 col-span-2">
<Label className="text-xs text-muted-foreground">Label *</Label> <Label className="text-xs text-muted-foreground">Label *</Label>
@@ -414,6 +416,27 @@ export function NodeModal({ open, onClose, onSubmit, initial, title = 'Add Node'
</div> </div>
)} )}
{/* Bottom connection points (not for group containers) */}
{form.type !== 'groupRect' && form.type !== 'group' && (
<div className="flex flex-col gap-1.5 col-span-2">
<Label className="text-xs text-muted-foreground">Bottom Connection Points</Label>
<Select
value={String(form.bottom_handles ?? 1)}
onValueChange={(v) => set('bottom_handles', parseInt(v ?? '1', 10))}
>
<SelectTrigger className="bg-[#21262d] border-[#30363d] text-sm h-8">
<SelectValue />
</SelectTrigger>
<SelectContent className="bg-[#21262d] border-[#30363d]">
<SelectItem value="1" className="text-sm">1 center</SelectItem>
<SelectItem value="2" className="text-sm">2 left / right</SelectItem>
<SelectItem value="3" className="text-sm">3 left / center / right</SelectItem>
<SelectItem value="4" className="text-sm">4 evenly spaced</SelectItem>
</SelectContent>
</Select>
</div>
)}
{/* Notes */} {/* Notes */}
<div className="flex flex-col gap-1.5 col-span-2"> <div className="flex flex-col gap-1.5 col-span-2">
<Label className="text-xs text-muted-foreground">Notes</Label> <Label className="text-xs text-muted-foreground">Notes</Label>
@@ -19,6 +19,7 @@ export interface PendingDevice {
services: Service[] services: Service[]
suggested_type: string | null suggested_type: string | null
status: string status: string
discovery_source: string | null
discovered_at: string discovered_at: string
} }
@@ -101,6 +102,9 @@ export function PendingDeviceModal({ device, onClose, onApprove, onHide, onIgnor
{device.suggested_type && ( {device.suggested_type && (
<InfoRow label="Type" value={device.suggested_type} /> <InfoRow label="Type" value={device.suggested_type} />
)} )}
{device.discovery_source && (
<InfoRow label="Source" value={device.discovery_source.toUpperCase()} />
)}
<InfoRow label="Discovered" value={new Date(device.discovered_at).toLocaleString()} /> <InfoRow label="Discovered" value={new Date(device.discovered_at).toLocaleString()} />
</div> </div>
@@ -99,7 +99,6 @@ export function ScanConfigModal({ open, onClose, onScanNow }: ScanConfigModalPro
<DialogFooter className="gap-2"> <DialogFooter className="gap-2">
<Button variant="ghost" onClick={onClose}>Cancel</Button> <Button variant="ghost" onClick={onClose}>Cancel</Button>
<Button variant="outline" onClick={handleSave} disabled={saving}>Save</Button>
<Button <Button
onClick={handleScanNow} onClick={handleScanNow}
disabled={saving} disabled={saving}
@@ -1,170 +1,414 @@
import { describe, it, expect, vi } from 'vitest' import { describe, it, expect, vi } from 'vitest'
import { render, screen, fireEvent } from '@testing-library/react' import { render, screen, fireEvent } from '@testing-library/react'
import { NodeModal } from '../NodeModal' import { NodeModal } from '../NodeModal'
import type { NodeData } from '@/types'
// ── Mock Shadcn Select with native <select> for testability ───────────────
vi.mock('@/components/ui/select', () => ({
Select: ({ value, onValueChange, children }: {
value?: string; onValueChange?: (v: string) => void; children: React.ReactNode
}) => (
<select value={value} onChange={(e) => onValueChange?.(e.target.value)}>
{children}
</select>
),
SelectTrigger: ({ children }: { children: React.ReactNode }) => <>{children}</>,
SelectValue: () => null,
SelectContent: ({ children }: { children: React.ReactNode }) => <>{children}</>,
SelectGroup: ({ children }: { children: React.ReactNode }) => <>{children}</>,
SelectLabel: () => null,
SelectItem: ({ value, children }: { value: string; children: React.ReactNode }) => (
<option value={value}>{children}</option>
),
SelectSeparator: () => null,
}))
// ── Helpers ───────────────────────────────────────────────────────────────
function renderModal(props: Partial<Parameters<typeof NodeModal>[0]> = {}) {
const onClose = vi.fn()
const onSubmit = vi.fn()
render(<NodeModal open onClose={onClose} onSubmit={onSubmit} {...props} />)
return { onClose, onSubmit }
}
/** Get <select> elements in document order: [0]=Type, [1]=CheckMethod, [2]=BottomHandles */
function selects() { return screen.getAllByRole('combobox') as HTMLSelectElement[] }
const BASE: Partial<NodeData> = {
type: 'server', label: 'My Server', hostname: 'server.lan',
ip: '192.168.1.10', check_method: 'ping', services: [],
}
// ── Tests ─────────────────────────────────────────────────────────────────
describe('NodeModal', () => { describe('NodeModal', () => {
// ── Visibility ────────────────────────────────────────────────────────
it('renders nothing when closed', () => { it('renders nothing when closed', () => {
const { container } = render( const { container } = render(<NodeModal open={false} onClose={vi.fn()} onSubmit={vi.fn()} />)
<NodeModal open={false} onClose={vi.fn()} onSubmit={vi.fn()} />
)
expect(container.querySelector('[role="dialog"]')).toBeNull() expect(container.querySelector('[role="dialog"]')).toBeNull()
}) })
it('renders form fields when open', () => { it('renders form fields when open', () => {
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />) renderModal()
expect(screen.getByPlaceholderText('My Server')).toBeDefined() expect(screen.getByPlaceholderText('My Server')).toBeDefined()
expect(screen.getByText('Add Node')).toBeDefined() expect(screen.getByText('Add Node')).toBeDefined()
}) })
it('does not call onSubmit when label is empty and shows error', () => { it('shows "Add" button for default title', () => {
const onSubmit = vi.fn() renderModal()
render(<NodeModal open onClose={vi.fn()} onSubmit={onSubmit} />) expect(screen.getByRole('button', { name: 'Add' })).toBeDefined()
fireEvent.click(screen.getByText('Add')) })
it('shows "Save" button when title is Edit Node', () => {
renderModal({ title: 'Edit Node' })
expect(screen.getByRole('button', { name: 'Save' })).toBeDefined()
})
it('pre-fills form from initial prop', () => {
renderModal({ initial: BASE })
expect((screen.getByPlaceholderText('My Server') as HTMLInputElement).value).toBe('My Server')
expect((screen.getByPlaceholderText('server.lan') as HTMLInputElement).value).toBe('server.lan')
expect((screen.getByPlaceholderText('192.168.1.x') as HTMLInputElement).value).toBe('192.168.1.10')
})
// ── Cancel ────────────────────────────────────────────────────────────
it('calls onClose when Cancel is clicked', () => {
const { onClose } = renderModal()
fireEvent.click(screen.getByRole('button', { name: 'Cancel' }))
expect(onClose).toHaveBeenCalledOnce()
})
// ── Label validation ──────────────────────────────────────────────────
it('blocks submit and shows error when label is empty', () => {
const { onSubmit } = renderModal()
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect(onSubmit).not.toHaveBeenCalled() expect(onSubmit).not.toHaveBeenCalled()
expect(screen.getByText('Label is required')).toBeDefined() expect(screen.getByText('Label is required')).toBeDefined()
}) })
it('calls onSubmit with form data when label is filled', () => { it('blocks submit when label is whitespace only', () => {
const onSubmit = vi.fn() const { onSubmit } = renderModal()
const onClose = vi.fn() fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: ' ' } })
render(<NodeModal open onClose={onClose} onSubmit={onSubmit} />) fireEvent.click(screen.getByRole('button', { name: 'Add' }))
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'My NAS' } }) expect(onSubmit).not.toHaveBeenCalled()
fireEvent.click(screen.getByText('Add'))
expect(onSubmit).toHaveBeenCalledOnce()
expect(onSubmit.mock.calls[0][0].label).toBe('My NAS')
expect(onClose).toHaveBeenCalledOnce()
}) })
it('clears label error when user starts typing', () => { it('clears label error when user starts typing', () => {
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />) renderModal()
fireEvent.click(screen.getByText('Add')) fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect(screen.getByText('Label is required')).toBeDefined()
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'x' } }) fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'x' } })
expect(screen.queryByText('Label is required')).toBeNull() expect(screen.queryByText('Label is required')).toBeNull()
}) })
it('pre-fills form from initial prop', () => { // ── Form submission ───────────────────────────────────────────────────
render(
<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} initial={{ label: 'Pre-filled', ip: '10.0.0.1' }} />
)
const input = screen.getByPlaceholderText('My Server') as HTMLInputElement
expect(input.value).toBe('Pre-filled')
})
it('shows Save button text when title is Edit Node', () => { it('calls onSubmit and onClose with form data on valid submit', () => {
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} title="Edit Node" />) const { onSubmit, onClose } = renderModal({ initial: BASE })
expect(screen.getByText('Save')).toBeDefined() fireEvent.click(screen.getByRole('button', { name: 'Add' }))
}) expect(onSubmit).toHaveBeenCalledOnce()
it('calls onClose when Cancel is clicked', () => {
const onClose = vi.fn()
render(<NodeModal open onClose={onClose} onSubmit={vi.fn()} />)
fireEvent.click(screen.getByText('Cancel'))
expect(onClose).toHaveBeenCalledOnce() expect(onClose).toHaveBeenCalledOnce()
const data = onSubmit.mock.calls[0][0] as Partial<NodeData>
expect(data.label).toBe('My Server')
expect(data.type).toBe('server')
}) })
describe('Hardware section', () => { it('submits updated hostname, IP and notes', () => {
it('renders Hardware toggle button', () => { const { onSubmit } = renderModal({ initial: BASE })
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />) fireEvent.change(screen.getByPlaceholderText('server.lan'), { target: { value: 'nas.local' } })
expect(screen.getByText('Hardware')).toBeDefined() fireEvent.change(screen.getByPlaceholderText('192.168.1.x'), { target: { value: '10.0.0.1' } })
}) fireEvent.change(screen.getByPlaceholderText('Optional notes'), { target: { value: 'rack A' } })
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
const data = onSubmit.mock.calls[0][0] as Partial<NodeData>
expect(data.hostname).toBe('nas.local')
expect(data.ip).toBe('10.0.0.1')
expect(data.notes).toBe('rack A')
})
it('hardware fields are hidden by default', () => { it('submits check_target', () => {
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />) const { onSubmit } = renderModal({ initial: BASE })
expect(screen.queryByPlaceholderText('e.g. Intel Xeon E5-2680')).toBeNull() fireEvent.change(screen.getByPlaceholderText('http://...'), { target: { value: 'http://192.168.1.10:8080' } })
}) fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect((onSubmit.mock.calls[0][0] as Partial<NodeData>).check_target).toBe('http://192.168.1.10:8080')
})
it('expands hardware fields on toggle click', () => { // ── Type selector ─────────────────────────────────────────────────────
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />)
fireEvent.click(screen.getByText('Hardware'))
expect(screen.getByPlaceholderText('e.g. Intel Xeon E5-2680')).toBeDefined()
expect(screen.getByPlaceholderText('e.g. 8')).toBeDefined()
expect(screen.getByPlaceholderText('e.g. 32')).toBeDefined()
expect(screen.getByPlaceholderText('e.g. 500')).toBeDefined()
})
it('submits hardware fields when filled', () => { it('pre-fills type from initial', () => {
const onSubmit = vi.fn() renderModal({ initial: { ...BASE, type: 'router' } })
render(<NodeModal open onClose={vi.fn()} onSubmit={onSubmit} />) expect(selects()[0].value).toBe('router')
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'Homelab' } }) })
fireEvent.click(screen.getByText('Hardware'))
fireEvent.change(screen.getByPlaceholderText('e.g. Intel Xeon E5-2680'), { target: { value: 'Intel i7-12700K' } })
fireEvent.change(screen.getByPlaceholderText('e.g. 8'), { target: { value: '12' } })
fireEvent.change(screen.getByPlaceholderText('e.g. 32'), { target: { value: '64' } })
fireEvent.change(screen.getByPlaceholderText('e.g. 500'), { target: { value: '2000' } })
fireEvent.click(screen.getByText('Add'))
const submitted = onSubmit.mock.calls[0][0]
expect(submitted.cpu_model).toBe('Intel i7-12700K')
expect(submitted.cpu_count).toBe(12)
expect(submitted.ram_gb).toBe(64)
expect(submitted.disk_gb).toBe(2000)
})
it('auto-expands when initial has hardware data', () => { it('changes type and submits it', () => {
render( const { onSubmit } = renderModal({ initial: BASE })
<NodeModal fireEvent.change(selects()[0], { target: { value: 'nas' } })
open fireEvent.click(screen.getByRole('button', { name: 'Add' }))
onClose={vi.fn()} expect((onSubmit.mock.calls[0][0] as Partial<NodeData>).type).toBe('nas')
onSubmit={vi.fn()} })
initial={{ label: 'Server', cpu_count: 8, ram_gb: 32 }}
/>
)
expect(screen.getByPlaceholderText('e.g. Intel Xeon E5-2680')).toBeDefined()
})
it('hides hardware section for groupRect type', () => { // ── Check method ──────────────────────────────────────────────────────
render(
<NodeModal
open
onClose={vi.fn()}
onSubmit={vi.fn()}
initial={{ type: 'groupRect' }}
/>
)
expect(screen.queryByText('Hardware')).toBeNull()
})
it('show on node toggle is hidden when section is collapsed', () => { it('pre-fills check_method from initial', () => {
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />) renderModal({ initial: { ...BASE, check_method: 'http' } })
expect(screen.queryByText('Show on node')).toBeNull() expect(selects()[1].value).toBe('http')
}) })
it('show on node toggle appears when section is expanded', () => { it('changes check_method and submits it', () => {
render(<NodeModal open onClose={vi.fn()} onSubmit={vi.fn()} />) const { onSubmit } = renderModal({ initial: BASE })
fireEvent.click(screen.getByText('Hardware')) fireEvent.change(selects()[1], { target: { value: 'ssh' } })
expect(screen.getByText('Show on node')).toBeDefined() fireEvent.click(screen.getByRole('button', { name: 'Add' }))
}) expect((onSubmit.mock.calls[0][0] as Partial<NodeData>).check_method).toBe('ssh')
})
it('show_hardware defaults to false', () => { // ── Icon picker ───────────────────────────────────────────────────────
const onSubmit = vi.fn()
render(<NodeModal open onClose={vi.fn()} onSubmit={onSubmit} />)
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'Node' } })
fireEvent.click(screen.getByText('Add'))
expect(onSubmit.mock.calls[0][0].show_hardware).toBeFalsy()
})
it('toggling show on node sets show_hardware to true', () => { it('shows "Default" label when no custom icon', () => {
const onSubmit = vi.fn() renderModal({ initial: BASE })
render(<NodeModal open onClose={vi.fn()} onSubmit={onSubmit} />) expect(screen.getByText('Default')).toBeDefined()
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'Node' } }) })
fireEvent.click(screen.getByText('Hardware'))
fireEvent.click(screen.getByRole('switch'))
fireEvent.click(screen.getByText('Add'))
expect(onSubmit.mock.calls[0][0].show_hardware).toBe(true)
})
it('pre-fills show_hardware from initial prop', () => { it('opens icon picker on trigger button click', () => {
const onSubmit = vi.fn() renderModal({ initial: BASE })
render( expect(screen.queryByPlaceholderText('Search icons…')).toBeNull()
<NodeModal fireEvent.click(screen.getByText('Default'))
open expect(screen.getByPlaceholderText('Search icons…')).toBeDefined()
onClose={vi.fn()} })
onSubmit={onSubmit}
initial={{ label: 'Node', show_hardware: true, cpu_count: 8 }} it('closes picker and shows icon label after selecting an icon', () => {
/> renderModal({ initial: BASE })
) fireEvent.click(screen.getByText('Default'))
fireEvent.click(screen.getByText('Add')) fireEvent.click(screen.getByTitle('Database (SQL/NoSQL)'))
expect(onSubmit.mock.calls[0][0].show_hardware).toBe(true) expect(screen.queryByPlaceholderText('Search icons…')).toBeNull()
expect(screen.getByText('Database (SQL/NoSQL)')).toBeDefined()
})
it('submits custom_icon key after picking', () => {
const { onSubmit } = renderModal({ initial: BASE })
fireEvent.click(screen.getByText('Default'))
fireEvent.click(screen.getByTitle('Database (SQL/NoSQL)'))
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect((onSubmit.mock.calls[0][0] as Partial<NodeData>).custom_icon).toBe('database')
})
it('shows Reset button when custom_icon is set', () => {
renderModal({ initial: { ...BASE, custom_icon: 'database' } })
expect(screen.getByRole('button', { name: /Reset/i })).toBeDefined()
})
it('hides Reset button when no custom_icon', () => {
renderModal({ initial: BASE })
expect(screen.queryByRole('button', { name: /Reset/i })).toBeNull()
})
it('resets custom_icon and shows Default on Reset click', () => {
renderModal({ initial: { ...BASE, custom_icon: 'database' } })
fireEvent.click(screen.getByRole('button', { name: /Reset/i }))
expect(screen.getByText('Default')).toBeDefined()
})
it('filters icons by search query', () => {
renderModal({ initial: BASE })
fireEvent.click(screen.getByText('Default'))
fireEvent.change(screen.getByPlaceholderText('Search icons…'), { target: { value: 'grafana' } })
expect(screen.getByTitle('Grafana / Kibana')).toBeDefined()
expect(screen.queryByTitle('Router')).toBeNull()
})
// ── Container mode (proxmox only) ─────────────────────────────────────
it('shows Container Mode toggle for proxmox type', () => {
renderModal({ initial: { ...BASE, type: 'proxmox' } })
expect(screen.getByText('Container Mode')).toBeDefined()
})
it('hides Container Mode for non-proxmox types', () => {
renderModal({ initial: BASE })
expect(screen.queryByText('Container Mode')).toBeNull()
})
it('toggles container_mode on click', () => {
const { onSubmit } = renderModal({ initial: { ...BASE, type: 'proxmox', container_mode: true } })
fireEvent.click(screen.getByRole('switch'))
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect((onSubmit.mock.calls[0][0] as Partial<NodeData>).container_mode).toBe(false)
})
// ── Parent Proxmox (vm / lxc only) ───────────────────────────────────
it('shows Parent Proxmox for vm with proxmoxNodes', () => {
renderModal({
initial: { ...BASE, type: 'vm' },
proxmoxNodes: [{ id: 'px1', label: 'PVE-01' }],
}) })
expect(screen.getByText('Parent Proxmox')).toBeDefined()
expect(screen.getByText('PVE-01')).toBeDefined()
})
it('shows Parent Proxmox for lxc with proxmoxNodes', () => {
renderModal({
initial: { ...BASE, type: 'lxc' },
proxmoxNodes: [{ id: 'px1', label: 'PVE-01' }],
})
expect(screen.getByText('Parent Proxmox')).toBeDefined()
})
it('hides Parent Proxmox for server type', () => {
renderModal({ initial: BASE, proxmoxNodes: [{ id: 'px1', label: 'PVE-01' }] })
expect(screen.queryByText('Parent Proxmox')).toBeNull()
})
it('hides Parent Proxmox for vm when no proxmoxNodes', () => {
renderModal({ initial: { ...BASE, type: 'vm' } })
expect(screen.queryByText('Parent Proxmox')).toBeNull()
})
// ── Appearance ────────────────────────────────────────────────────────
it('renders 3 color swatch labels (border, background, icon)', () => {
renderModal({ initial: BASE })
expect(screen.getByText('border')).toBeDefined()
expect(screen.getByText('background')).toBeDefined()
expect(screen.getByText('icon')).toBeDefined()
})
it('shows default colors hint when no custom_colors', () => {
renderModal({ initial: BASE })
expect(screen.getByText(/Using default colors for/)).toBeDefined()
})
it('shows Reset to defaults when custom_colors are set', () => {
renderModal({ initial: { ...BASE, custom_colors: { border: '#ff0000' } } })
expect(screen.getByText('Reset to defaults')).toBeDefined()
})
it('resets custom_colors on Reset to defaults click', () => {
renderModal({ initial: { ...BASE, custom_colors: { border: '#ff0000' } } })
fireEvent.click(screen.getByText('Reset to defaults'))
expect(screen.queryByText('Reset to defaults')).toBeNull()
expect(screen.getByText(/Using default colors for/)).toBeDefined()
})
// ── Hardware section ──────────────────────────────────────────────────
it('renders Hardware toggle button', () => {
renderModal()
expect(screen.getByText('Hardware')).toBeDefined()
})
it('hardware fields are hidden by default', () => {
renderModal()
expect(screen.queryByPlaceholderText('e.g. Intel Xeon E5-2680')).toBeNull()
})
it('expands hardware fields on toggle click', () => {
renderModal()
fireEvent.click(screen.getByText('Hardware'))
expect(screen.getByPlaceholderText('e.g. Intel Xeon E5-2680')).toBeDefined()
expect(screen.getByPlaceholderText('e.g. 8')).toBeDefined()
expect(screen.getByPlaceholderText('e.g. 32')).toBeDefined()
expect(screen.getByPlaceholderText('e.g. 500')).toBeDefined()
})
it('auto-expands when initial has hardware data', () => {
renderModal({ initial: { ...BASE, cpu_count: 8, ram_gb: 32 } })
expect(screen.getByPlaceholderText('e.g. Intel Xeon E5-2680')).toBeDefined()
})
it('pre-fills hardware fields from initial', () => {
renderModal({ initial: { ...BASE, cpu_model: 'Intel i5', cpu_count: 4, ram_gb: 16, disk_gb: 500 } })
expect((screen.getByPlaceholderText('e.g. Intel Xeon E5-2680') as HTMLInputElement).value).toBe('Intel i5')
})
it('submits hardware fields when filled', () => {
const { onSubmit } = renderModal()
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'Homelab' } })
fireEvent.click(screen.getByText('Hardware'))
fireEvent.change(screen.getByPlaceholderText('e.g. Intel Xeon E5-2680'), { target: { value: 'Intel i7-12700K' } })
fireEvent.change(screen.getByPlaceholderText('e.g. 8'), { target: { value: '12' } })
fireEvent.change(screen.getByPlaceholderText('e.g. 32'), { target: { value: '64' } })
fireEvent.change(screen.getByPlaceholderText('e.g. 500'), { target: { value: '2000' } })
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
const data = onSubmit.mock.calls[0][0] as Partial<NodeData>
expect(data.cpu_model).toBe('Intel i7-12700K')
expect(data.cpu_count).toBe(12)
expect(data.ram_gb).toBe(64)
expect(data.disk_gb).toBe(2000)
})
it('hides Hardware section for groupRect type', () => {
renderModal({ initial: { type: 'groupRect' } })
expect(screen.queryByText('Hardware')).toBeNull()
})
it('show_hardware toggle hidden when section is collapsed', () => {
renderModal()
expect(screen.queryByText('Show on node')).toBeNull()
})
it('show_hardware toggle appears when section is expanded', () => {
renderModal()
fireEvent.click(screen.getByText('Hardware'))
expect(screen.getByText('Show on node')).toBeDefined()
})
it('show_hardware defaults to falsy', () => {
const { onSubmit } = renderModal()
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'Node' } })
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect(onSubmit.mock.calls[0][0].show_hardware).toBeFalsy()
})
it('toggling show_hardware sets it to true', () => {
const { onSubmit } = renderModal()
fireEvent.change(screen.getByPlaceholderText('My Server'), { target: { value: 'Node' } })
fireEvent.click(screen.getByText('Hardware'))
fireEvent.click(screen.getByRole('switch'))
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect(onSubmit.mock.calls[0][0].show_hardware).toBe(true)
})
it('pre-fills show_hardware from initial', () => {
const { onSubmit } = renderModal({ initial: { label: 'Node', show_hardware: true, cpu_count: 8 } })
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect(onSubmit.mock.calls[0][0].show_hardware).toBe(true)
})
// ── Bottom connection points ───────────────────────────────────────────
it('shows Bottom Connection Points for server type', () => {
renderModal({ initial: BASE })
expect(screen.getByText('Bottom Connection Points')).toBeDefined()
})
it('hides Bottom Connection Points for groupRect', () => {
renderModal({ initial: { ...BASE, type: 'groupRect' } })
expect(screen.queryByText('Bottom Connection Points')).toBeNull()
})
it('hides Bottom Connection Points for group', () => {
renderModal({ initial: { ...BASE, type: 'group' } })
expect(screen.queryByText('Bottom Connection Points')).toBeNull()
})
it('defaults bottom_handles to 1', () => {
renderModal({ initial: BASE })
expect(selects()[2].value).toBe('1')
})
it('pre-fills bottom_handles from initial', () => {
renderModal({ initial: { ...BASE, bottom_handles: 3 } })
expect(selects()[2].value).toBe('3')
})
it('submits updated bottom_handles', () => {
const { onSubmit } = renderModal({ initial: BASE })
fireEvent.change(selects()[2], { target: { value: '4' } })
fireEvent.click(screen.getByRole('button', { name: 'Add' }))
expect((onSubmit.mock.calls[0][0] as Partial<NodeData>).bottom_handles).toBe(4)
}) })
}) })
@@ -38,16 +38,6 @@ describe('ScanConfigModal', () => {
expect(input).toBeDefined() expect(input).toBeDefined()
}) })
it('saves only ranges (interval managed by settings endpoint)', async () => {
vi.mocked(scanApi.getConfig).mockResolvedValue({ data: { ranges: ['10.0.0.0/8'] } } as never)
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
await screen.findByDisplayValue('10.0.0.0/8')
fireEvent.click(screen.getByRole('button', { name: 'Save' }))
await waitFor(() => {
expect(scanApi.saveConfig).toHaveBeenCalledWith({ ranges: ['10.0.0.0/8'] })
})
})
it('adds a new empty range on "Add range" click', async () => { it('adds a new empty range on "Add range" click', async () => {
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />) render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
await screen.findByDisplayValue('192.168.1.0/24') await screen.findByDisplayValue('192.168.1.0/24')
@@ -59,53 +49,29 @@ describe('ScanConfigModal', () => {
it('delete button disabled when only one range', async () => { it('delete button disabled when only one range', async () => {
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />) render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
await screen.findByDisplayValue('192.168.1.0/24') await screen.findByDisplayValue('192.168.1.0/24')
// Only 1 range → delete button disabled
const trashButtons = document.querySelectorAll('button[disabled]') const trashButtons = document.querySelectorAll('button[disabled]')
expect(trashButtons.length).toBeGreaterThan(0) expect(trashButtons.length).toBeGreaterThan(0)
}) })
it('can remove a range when more than one exist', async () => { it('can remove a range when more than one exist', async () => {
vi.mocked(scanApi.getConfig).mockResolvedValue({ data: { ranges: ['192.168.1.0/24', '10.0.0.0/8'], } } as never) vi.mocked(scanApi.getConfig).mockResolvedValue({ data: { ranges: ['192.168.1.0/24', '10.0.0.0/8'] } } as never)
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />) render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
await screen.findByDisplayValue('192.168.1.0/24') await screen.findByDisplayValue('192.168.1.0/24')
// Both trash buttons should be enabled
const trashButtons = screen.getAllByRole('button').filter((b) => !b.hasAttribute('disabled') && b.querySelector('svg')) const trashButtons = screen.getAllByRole('button').filter((b) => !b.hasAttribute('disabled') && b.querySelector('svg'))
expect(trashButtons.length).toBeGreaterThanOrEqual(2) expect(trashButtons.length).toBeGreaterThanOrEqual(2)
}) })
it('shows error toast and does not save when all ranges are empty', async () => { it('shows error toast and does not save when all ranges are empty', async () => {
vi.mocked(scanApi.getConfig).mockResolvedValue({ data: { ranges: [''], } } as never) vi.mocked(scanApi.getConfig).mockResolvedValue({ data: { ranges: [''] } } as never)
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />) render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
await waitFor(() => expect(scanApi.getConfig).toHaveBeenCalled()) await waitFor(() => expect(scanApi.getConfig).toHaveBeenCalled())
fireEvent.click(screen.getByRole('button', { name: 'Save' })) fireEvent.click(screen.getByRole('button', { name: 'Scan Now' }))
await waitFor(() => { await waitFor(() => {
expect(toast.error).toHaveBeenCalledWith('Add at least one IP range') expect(toast.error).toHaveBeenCalledWith('Add at least one IP range')
}) })
expect(scanApi.saveConfig).not.toHaveBeenCalled() expect(scanApi.saveConfig).not.toHaveBeenCalled()
}) })
it('saves config and closes on Save click', async () => {
const onClose = vi.fn()
render(<ScanConfigModal open onClose={onClose} onScanNow={vi.fn()} />)
await screen.findByDisplayValue('192.168.1.0/24')
fireEvent.click(screen.getByRole('button', { name: 'Save' }))
await waitFor(() => {
expect(scanApi.saveConfig).toHaveBeenCalledWith({ ranges: ['192.168.1.0/24'] })
expect(toast.success).toHaveBeenCalledWith('Scan config saved')
expect(onClose).toHaveBeenCalledOnce()
})
})
it('shows error toast when save fails', async () => {
vi.mocked(scanApi.saveConfig).mockRejectedValue(new Error('network'))
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
await screen.findByDisplayValue('192.168.1.0/24')
fireEvent.click(screen.getByRole('button', { name: 'Save' }))
await waitFor(() => {
expect(toast.error).toHaveBeenCalledWith('Failed to save config')
})
})
it('calls onScanNow after saving on "Scan Now" click', async () => { it('calls onScanNow after saving on "Scan Now" click', async () => {
const onScanNow = vi.fn() const onScanNow = vi.fn()
const onClose = vi.fn() const onClose = vi.fn()
@@ -113,7 +79,7 @@ describe('ScanConfigModal', () => {
await screen.findByDisplayValue('192.168.1.0/24') await screen.findByDisplayValue('192.168.1.0/24')
fireEvent.click(screen.getByRole('button', { name: 'Scan Now' })) fireEvent.click(screen.getByRole('button', { name: 'Scan Now' }))
await waitFor(() => { await waitFor(() => {
expect(scanApi.saveConfig).toHaveBeenCalled() expect(scanApi.saveConfig).toHaveBeenCalledWith({ ranges: ['192.168.1.0/24'] })
expect(onScanNow).toHaveBeenCalledOnce() expect(onScanNow).toHaveBeenCalledOnce()
}) })
}) })
@@ -126,12 +92,11 @@ describe('ScanConfigModal', () => {
expect(onClose).toHaveBeenCalledOnce() expect(onClose).toHaveBeenCalledOnce()
}) })
it('strips whitespace from ranges before saving', async () => { it('strips whitespace from ranges before scanning', async () => {
render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />) render(<ScanConfigModal open onClose={vi.fn()} onScanNow={vi.fn()} />)
const input = await screen.findByDisplayValue('192.168.1.0/24') const input = await screen.findByDisplayValue('192.168.1.0/24')
// Type a range with surrounding whitespace
fireEvent.change(input, { target: { value: ' 10.0.0.0/8 ' } }) fireEvent.change(input, { target: { value: ' 10.0.0.0/8 ' } })
fireEvent.click(screen.getByRole('button', { name: 'Save' })) fireEvent.click(screen.getByRole('button', { name: 'Scan Now' }))
await waitFor(() => { await waitFor(() => {
expect(scanApi.saveConfig).toHaveBeenCalledWith( expect(scanApi.saveConfig).toHaveBeenCalledWith(
expect.objectContaining({ ranges: ['10.0.0.0/8'] }) expect.objectContaining({ ranges: ['10.0.0.0/8'] })
+25 -5
View File
@@ -1,5 +1,5 @@
import { useState, useCallback, useEffect, useRef } from 'react' import { useState, useCallback, useEffect, useRef } from 'react'
import { Plus, Save, ScanLine, ChevronLeft, ChevronRight, LayoutDashboard, Clock, EyeOff, Trash2, RefreshCw, Loader2, Square, Eye, Settings, StopCircle } from 'lucide-react' import { Plus, Save, ScanLine, ChevronLeft, ChevronRight, LayoutDashboard, Clock, EyeOff, Trash2, RefreshCw, Loader2, Square, Eye, Settings, StopCircle, X } from 'lucide-react'
import { Logo } from '@/components/ui/Logo' import { Logo } from '@/components/ui/Logo'
import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip' import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip'
import { useCanvasStore } from '@/stores/canvasStore' import { useCanvasStore } from '@/stores/canvasStore'
@@ -174,6 +174,16 @@ function PendingDevicesPanel({ onNodeApproved }: { onNodeApproved: (nodeId: stri
} }
}, []) }, [])
const handleClearAll = async () => {
try {
await scanApi.clearPending()
setDevices([])
toast.success('Pending devices cleared')
} catch {
toast.error('Failed to clear pending devices')
}
}
useEffect(() => { load() }, [load]) useEffect(() => { load() }, [load])
useEffect(() => { useEffect(() => {
@@ -230,9 +240,16 @@ function PendingDevicesPanel({ onNodeApproved }: { onNodeApproved: (nodeId: stri
<div className="p-2"> <div className="p-2">
<div className="flex items-center justify-between mb-2"> <div className="flex items-center justify-between mb-2">
<span className="text-xs font-medium text-muted-foreground uppercase tracking-wider">Pending</span> <span className="text-xs font-medium text-muted-foreground uppercase tracking-wider">Pending</span>
<button onClick={load} className="text-muted-foreground hover:text-foreground p-0.5"> <div className="flex items-center gap-1">
<RefreshCw size={12} /> <button onClick={load} className="text-muted-foreground hover:text-foreground p-0.5" title="Refresh">
</button> <RefreshCw size={12} />
</button>
{devices.length > 0 && (
<button onClick={handleClearAll} className="text-muted-foreground hover:text-[#f85149] p-0.5" title="Clear all pending">
<X size={12} />
</button>
)}
</div>
</div> </div>
{loading && <Loader2 size={14} className="animate-spin text-muted-foreground mx-auto my-4" />} {loading && <Loader2 size={14} className="animate-spin text-muted-foreground mx-auto my-4" />}
{!loading && devices.length === 0 && ( {!loading && devices.length === 0 && (
@@ -252,6 +269,8 @@ function PendingDevicesPanel({ onNodeApproved }: { onNodeApproved: (nodeId: stri
const hasHttps = d.services.some((s) => s.port === 443) 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 otherCount = d.services.filter((s) => s.port !== 22 && s.port !== 80 && s.port !== 443).length
const virtualBadge = detectVirtualBadge(d.mac) 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
return ( return (
<button <button
key={d.id} key={d.id}
@@ -265,8 +284,9 @@ function PendingDevicesPanel({ onNodeApproved }: { onNodeApproved: (nodeId: stri
{showIpBelow && ( {showIpBelow && (
<div className="font-mono text-muted-foreground truncate pl-3 text-[10px] mt-0.5">{d.ip}</div> <div className="font-mono text-muted-foreground truncate pl-3 text-[10px] mt-0.5">{d.ip}</div>
)} )}
{(hasSsh || hasHttp || hasHttps || otherCount > 0 || virtualBadge) && ( {(hasSsh || hasHttp || hasHttps || otherCount > 0 || virtualBadge || sourceLabel) && (
<div className="flex items-center gap-1 pl-3 mt-1.5 flex-wrap"> <div className="flex items-center gap-1 pl-3 mt-1.5 flex-wrap">
{sourceLabel && <ServiceBadge label={sourceLabel} color={sourceColor} />}
{virtualBadge && ( {virtualBadge && (
<Tooltip> <Tooltip>
<TooltipTrigger> <TooltipTrigger>
@@ -253,11 +253,11 @@ describe('canvasStore', () => {
expect(useCanvasStore.getState().selectedNodeIds).toEqual([]) expect(useCanvasStore.getState().selectedNodeIds).toEqual([])
}) })
it('setSelectedNode(id) preserves existing selectedNodeIds', () => { it('setSelectedNode(id) sets selectedNodeIds to [id], clearing multi-selection', () => {
useCanvasStore.setState({ selectedNodeIds: ['n1', 'n2'] }) useCanvasStore.setState({ selectedNodeIds: ['n1', 'n2'] })
useCanvasStore.getState().setSelectedNode('n1') useCanvasStore.getState().setSelectedNode('n1')
// does NOT wipe selectedNodeIds when setting a specific id // Single node click resets multi-selection to just the clicked node
expect(useCanvasStore.getState().selectedNodeIds).toEqual(['n1', 'n2']) expect(useCanvasStore.getState().selectedNodeIds).toEqual(['n1'])
}) })
// ── createGroup ─────────────────────────────────────────────────────────── // ── createGroup ───────────────────────────────────────────────────────────
@@ -618,4 +618,61 @@ describe('canvasStore', () => {
expect(stored?.width).toBeUndefined() expect(stored?.width).toBeUndefined()
expect(stored?.height).toBeUndefined() expect(stored?.height).toBeUndefined()
}) })
// ── bottom_handles edge remapping ──────────────────────────────────────────
it('remaps source edges to "bottom" when bottom_handles is reduced', () => {
const node = makeNode('n1', { bottom_handles: 4 })
const edge = { ...makeEdge('e1', 'n1', 'n2'), sourceHandle: 'bottom-3' }
useCanvasStore.setState({ nodes: [node, makeNode('n2')], edges: [edge] })
useCanvasStore.getState().updateNode('n1', { bottom_handles: 2 })
const updated = useCanvasStore.getState().edges.find((e) => e.id === 'e1')
expect(updated?.sourceHandle).toBe('bottom')
})
it('remaps target edges to "bottom" when bottom_handles is reduced', () => {
const node = makeNode('n2', { bottom_handles: 3 })
const edge = { ...makeEdge('e1', 'n1', 'n2'), targetHandle: 'bottom-3' }
useCanvasStore.setState({ nodes: [makeNode('n1'), node], edges: [edge] })
useCanvasStore.getState().updateNode('n2', { bottom_handles: 1 })
const updated = useCanvasStore.getState().edges.find((e) => e.id === 'e1')
expect(updated?.targetHandle).toBe('bottom')
})
it('does not remap edges that are on handles still present after reduction', () => {
const node = makeNode('n1', { bottom_handles: 4 })
const edge = { ...makeEdge('e1', 'n1', 'n2'), sourceHandle: 'bottom-2' }
useCanvasStore.setState({ nodes: [node, makeNode('n2')], edges: [edge] })
useCanvasStore.getState().updateNode('n1', { bottom_handles: 3 })
const updated = useCanvasStore.getState().edges.find((e) => e.id === 'e1')
expect(updated?.sourceHandle).toBe('bottom-2')
})
it('does not remap edges when bottom_handles increases', () => {
const node = makeNode('n1', { bottom_handles: 2 })
const edge = { ...makeEdge('e1', 'n1', 'n2'), sourceHandle: 'bottom' }
useCanvasStore.setState({ nodes: [node, makeNode('n2')], edges: [edge] })
useCanvasStore.getState().updateNode('n1', { bottom_handles: 4 })
const updated = useCanvasStore.getState().edges.find((e) => e.id === 'e1')
expect(updated?.sourceHandle).toBe('bottom')
})
it('never remaps the "bottom" handle itself', () => {
const node = makeNode('n1', { bottom_handles: 4 })
const edge = { ...makeEdge('e1', 'n1', 'n2'), sourceHandle: 'bottom' }
useCanvasStore.setState({ nodes: [node, makeNode('n2')], edges: [edge] })
useCanvasStore.getState().updateNode('n1', { bottom_handles: 1 })
const updated = useCanvasStore.getState().edges.find((e) => e.id === 'e1')
expect(updated?.sourceHandle).toBe('bottom')
})
}) })
+29 -9
View File
@@ -11,6 +11,7 @@ import {
} from '@xyflow/react' } from '@xyflow/react'
import type { NodeData, EdgeData } from '@/types' import type { NodeData, EdgeData } from '@/types'
import { generateUUID } from '@/utils/uuid' import { generateUUID } from '@/utils/uuid'
import { normalizeHandle, removedBottomHandleIds } from '@/utils/handleUtils'
type HistoryEntry = { nodes: Node<NodeData>[]; edges: Edge<EdgeData>[] } type HistoryEntry = { nodes: Node<NodeData>[]; edges: Edge<EdgeData>[] }
@@ -52,6 +53,8 @@ interface CanvasState {
markSaved: () => void markSaved: () => void
markUnsaved: () => void markUnsaved: () => void
loadCanvas: (nodes: Node<NodeData>[], edges: Edge<EdgeData>[]) => void loadCanvas: (nodes: Node<NodeData>[], edges: Edge<EdgeData>[]) => void
fitViewPending: boolean
clearFitViewPending: () => void
notifyScanDeviceFound: () => void notifyScanDeviceFound: () => void
hideIp: boolean hideIp: boolean
toggleHideIp: () => void toggleHideIp: () => void
@@ -66,6 +69,7 @@ export const useCanvasStore = create<CanvasState>((set) => ({
editingGroupRectId: null, editingGroupRectId: null,
hideIp: false, hideIp: false,
scanEventTs: 0, scanEventTs: 0,
fitViewPending: false,
past: [], past: [],
future: [], future: [],
@@ -149,10 +153,6 @@ export const useCanvasStore = create<CanvasState>((set) => ({
set((state) => { set((state) => {
const extra = connection as Connection & Partial<EdgeData> const extra = connection as Connection & Partial<EdgeData>
const edgeType = extra.type ?? 'ethernet' const edgeType = extra.type ?? 'ethernet'
// Normalize invisible stub handle IDs so React Flow can locate the handle
// and render the edge immediately (top-t / bottom-t are opacity:0 helpers).
const normalizeHandle = (h: string | null | undefined) =>
h === 'top-t' ? 'top' : h === 'bottom-t' ? 'bottom' : (h ?? null)
return { return {
edges: addEdge({ edges: addEdge({
...connection, ...connection,
@@ -165,10 +165,10 @@ export const useCanvasStore = create<CanvasState>((set) => ({
} }
}), }),
setSelectedNode: (id) => set((state) => ({ setSelectedNode: (id) => set({
selectedNodeId: id, selectedNodeId: id,
selectedNodeIds: id ? state.selectedNodeIds : [], selectedNodeIds: id ? [id] : [],
})), }),
addNode: (node) => addNode: (node) =>
set((state) => { set((state) => {
@@ -226,7 +226,25 @@ export const useCanvasStore = create<CanvasState>((set) => ({
const children = nodes.filter((n) => !!n.parentId) const children = nodes.filter((n) => !!n.parentId)
nodes = [...parents, ...children] nodes = [...parents, ...children]
} }
return { nodes, hasUnsavedChanges: true } // Remap edges when bottom_handles is reduced so no edge disappears
let edges = state.edges
if ('bottom_handles' in data && data.bottom_handles != null) {
const currentNode = state.nodes.find((n) => n.id === id)
const oldCount = currentNode?.data.bottom_handles ?? 1
const newCount = data.bottom_handles
if (newCount < oldCount) {
const removed = removedBottomHandleIds(oldCount, newCount)
edges = state.edges.map((e) => {
if (e.source === id && e.sourceHandle && removed.has(e.sourceHandle))
return { ...e, sourceHandle: 'bottom' }
if (e.target === id && e.targetHandle && removed.has(e.targetHandle))
return { ...e, targetHandle: 'bottom' }
return e
})
}
}
return { nodes, edges, hasUnsavedChanges: true }
}), }),
deleteNode: (id) => deleteNode: (id) =>
@@ -409,6 +427,8 @@ export const useCanvasStore = create<CanvasState>((set) => ({
// React Flow requires parents before children in the array // React Flow requires parents before children in the array
const parents = nodes.filter((n) => !n.parentId) const parents = nodes.filter((n) => !n.parentId)
const children = nodes.filter((n) => !!n.parentId) const children = nodes.filter((n) => !!n.parentId)
set({ nodes: [...parents, ...children], edges, hasUnsavedChanges: false, selectedNodeId: null, past: [], future: [], clipboard: [] }) set({ nodes: [...parents, ...children], edges, hasUnsavedChanges: false, selectedNodeId: null, past: [], future: [], clipboard: [], fitViewPending: true })
}, },
clearFitViewPending: () => set({ fitViewPending: false }),
})) }))
+1
View File
@@ -82,6 +82,7 @@ export interface NodeData extends Record<string, unknown> {
height?: number height?: number
} }
custom_icon?: string custom_icon?: string
bottom_handles?: number
} }
export type EdgePathStyle = 'bezier' | 'smooth' export type EdgePathStyle = 'bezier' | 'smooth'
@@ -0,0 +1,111 @@
import { describe, it, expect } from 'vitest'
import {
BOTTOM_HANDLE_IDS,
BOTTOM_HANDLE_POSITIONS,
normalizeHandle,
removedBottomHandleIds,
} from '../handleUtils'
describe('BOTTOM_HANDLE_IDS', () => {
it('first id is always "bottom" for backward compatibility', () => {
expect(BOTTOM_HANDLE_IDS[0]).toBe('bottom')
})
it('has ids for 14 handles', () => {
expect(BOTTOM_HANDLE_IDS).toHaveLength(4)
expect(BOTTOM_HANDLE_IDS).toEqual(['bottom', 'bottom-2', 'bottom-3', 'bottom-4'])
})
})
describe('BOTTOM_HANDLE_POSITIONS', () => {
it('1 handle is centered at 50%', () => {
expect(BOTTOM_HANDLE_POSITIONS[1]).toEqual([50])
})
it('2 handles are symmetric', () => {
const [a, b] = BOTTOM_HANDLE_POSITIONS[2]
expect(a).toBeLessThan(50)
expect(b).toBeGreaterThan(50)
expect(a + b).toBe(100)
})
it('3 handles include a center at 50%', () => {
expect(BOTTOM_HANDLE_POSITIONS[3]).toContain(50)
expect(BOTTOM_HANDLE_POSITIONS[3]).toHaveLength(3)
})
it('4 handles are evenly spaced', () => {
const pos = BOTTOM_HANDLE_POSITIONS[4]
expect(pos).toHaveLength(4)
// All values should be between 0 and 100 exclusive
pos.forEach((p) => {
expect(p).toBeGreaterThan(0)
expect(p).toBeLessThan(100)
})
// Positions should be strictly increasing
for (let i = 1; i < pos.length; i++) {
expect(pos[i]).toBeGreaterThan(pos[i - 1])
}
})
})
describe('normalizeHandle', () => {
it('returns null for null/undefined', () => {
expect(normalizeHandle(null)).toBeNull()
expect(normalizeHandle(undefined)).toBeNull()
})
it('maps top-t → top', () => {
expect(normalizeHandle('top-t')).toBe('top')
})
it('maps bottom-t → bottom', () => {
expect(normalizeHandle('bottom-t')).toBe('bottom')
})
it('maps bottom-2-t → bottom-2', () => {
expect(normalizeHandle('bottom-2-t')).toBe('bottom-2')
})
it('maps bottom-3-t → bottom-3', () => {
expect(normalizeHandle('bottom-3-t')).toBe('bottom-3')
})
it('maps bottom-4-t → bottom-4', () => {
expect(normalizeHandle('bottom-4-t')).toBe('bottom-4')
})
it('passes through non-stub handles unchanged', () => {
expect(normalizeHandle('top')).toBe('top')
expect(normalizeHandle('bottom')).toBe('bottom')
expect(normalizeHandle('bottom-2')).toBe('bottom-2')
expect(normalizeHandle('custom-handle')).toBe('custom-handle')
})
})
describe('removedBottomHandleIds', () => {
it('returns empty set when count does not decrease', () => {
expect(removedBottomHandleIds(2, 2).size).toBe(0)
expect(removedBottomHandleIds(1, 4).size).toBe(0)
})
it('4 → 1 removes bottom-2, bottom-3, bottom-4', () => {
const removed = removedBottomHandleIds(4, 1)
expect(removed).toEqual(new Set(['bottom-2', 'bottom-3', 'bottom-4']))
})
it('4 → 2 removes bottom-3, bottom-4', () => {
const removed = removedBottomHandleIds(4, 2)
expect(removed).toEqual(new Set(['bottom-3', 'bottom-4']))
})
it('3 → 2 removes only bottom-3', () => {
const removed = removedBottomHandleIds(3, 2)
expect(removed).toEqual(new Set(['bottom-3']))
})
it('never removes "bottom" (index 0)', () => {
const removed = removedBottomHandleIds(4, 1)
expect(removed.has('bottom')).toBe(false)
})
})
+3 -3
View File
@@ -1,5 +1,6 @@
import type { Node, Edge } from '@xyflow/react' import type { Node, Edge } from '@xyflow/react'
import type { NodeData, EdgeData } from '@/types' import type { NodeData, EdgeData } from '@/types'
import { normalizeHandle } from '@/utils/handleUtils'
// ── Types ──────────────────────────────────────────────────────────────────── // ── Types ────────────────────────────────────────────────────────────────────
@@ -29,6 +30,7 @@ export interface ApiNode extends Record<string, unknown> {
show_hardware?: boolean show_hardware?: boolean
width?: number | null width?: number | null
height?: number | null height?: number | null
bottom_handles?: number
} }
export interface ApiEdge { export interface ApiEdge {
@@ -99,14 +101,12 @@ export function serializeNode(n: Node<NodeData>): Record<string, unknown> {
show_hardware: n.data.show_hardware ?? false, show_hardware: n.data.show_hardware ?? false,
width: n.width ?? null, width: n.width ?? null,
height: n.height ?? null, height: n.height ?? null,
bottom_handles: n.data.bottom_handles ?? 1,
pos_x: n.position.x, pos_x: n.position.x,
pos_y: n.position.y, pos_y: n.position.y,
} }
} }
const normalizeHandle = (h: string | null | undefined): string | null =>
h === 'top-t' ? 'top' : h === 'bottom-t' ? 'bottom' : (h ?? null)
export function serializeEdge(e: Edge<EdgeData>): Record<string, unknown> { export function serializeEdge(e: Edge<EdgeData>): Record<string, unknown> {
return { return {
id: e.id, id: e.id,
+45
View File
@@ -0,0 +1,45 @@
/**
* Bottom handle configuration for multi-handle nodes.
*
* Handle IDs: index 0 = 'bottom' (always the default, backward-compatible)
* index 1 = 'bottom-2', index 2 = 'bottom-3', index 3 = 'bottom-4'
*
* Invisible target handles follow the same pattern with a '-t' suffix:
* 'bottom-t', 'bottom-2-t', 'bottom-3-t', 'bottom-4-t'
*/
export const BOTTOM_HANDLE_IDS = ['bottom', 'bottom-2', 'bottom-3', 'bottom-4'] as const
/** Left % position for each handle slot, per count. */
export const BOTTOM_HANDLE_POSITIONS: Record<number, number[]> = {
1: [50],
2: [25, 75],
3: [20, 50, 80],
4: [15, 38, 62, 85],
}
/**
* Normalize a raw handle ID coming from a React Flow connection event.
* Invisible target handles (e.g. 'bottom-2-t') are mapped to their source
* counterpart ('bottom-2') so the stored edge ID is stable and consistent.
*/
export function normalizeHandle(h: string | null | undefined): string | null {
if (!h) return null
if (h === 'top-t') return 'top'
// 'bottom-t' → 'bottom', 'bottom-2-t' → 'bottom-2', etc.
const m = h.match(/^(bottom(?:-\d+)?)-t$/)
if (m) return m[1]
return h
}
/**
* Returns the set of handle IDs that are removed when bottom_handles
* is reduced from `oldCount` to `newCount`.
*/
export function removedBottomHandleIds(oldCount: number, newCount: number): Set<string> {
const removed = new Set<string>()
for (let i = newCount; i < oldCount; i++) {
removed.add(BOTTOM_HANDLE_IDS[i])
}
return removed
}
+22
View File
@@ -1,3 +1,4 @@
import type { NodeType } from '@/types'
import { import {
// Infrastructure (node types) // Infrastructure (node types)
Globe, Router, Network, Server, Layers, Box, Container, HardDrive, Cpu, Wifi, Circle, Globe, Router, Network, Server, Layers, Box, Container, HardDrive, Cpu, Wifi, Circle,
@@ -116,6 +117,27 @@ export const ICON_MAP: Record<string, LucideIcon> = Object.fromEntries(
ICON_REGISTRY.map((e) => [e.key, e.icon]), ICON_REGISTRY.map((e) => [e.key, e.icon]),
) )
export const NODE_TYPE_DEFAULT_ICONS: Record<NodeType, LucideIcon> = {
isp: Globe,
router: Router,
switch: Network,
server: Server,
proxmox: Layers,
vm: Box,
lxc: Container,
nas: HardDrive,
iot: Cpu,
ap: Wifi,
camera: Cctv,
printer: Printer,
computer: Monitor,
cpl: PlugZap,
docker: Anchor,
generic: Circle,
group: Circle,
groupRect: Circle,
}
/** Resolve the display icon for a node — custom_icon takes priority over type default. */ /** Resolve the display icon for a node — custom_icon takes priority over type default. */
export function resolveNodeIcon( export function resolveNodeIcon(
typeIcon: LucideIcon, typeIcon: LucideIcon,