Compare commits

..

1 Commits

Author SHA1 Message Date
Pouzor 701c9c5bb9 fix: force frontend builder stage to native platform, fixes QEMU arm64 npm crash 2026-03-28 14:23:21 +01:00
233 changed files with 3519 additions and 28423 deletions
-12
View File
@@ -1,7 +1,6 @@
# Backend - server-side only (NEVER commit .env)
SECRET_KEY=change_me_in_production
SQLITE_PATH=./data/homelab.db
# Set this to the URL(s) you use to access Homelable in your browser.
CORS_ORIGINS=["http://localhost:5173","http://localhost:3000"]
# Auth — default credentials: admin / admin
@@ -23,14 +22,3 @@ STATUS_CHECKER_INTERVAL=60
# Generate keys: python3 -c "import secrets; print(secrets.token_hex(32))"
MCP_API_KEY=mcp_sk_changeme
MCP_SERVICE_KEY=svc_changeme
# Live view — read-only public canvas at /view?key=<value>
# Off by default. Set to a random secret to enable.
# Generate: python3 -c "import secrets; print(secrets.token_urlsafe(32))"
# LIVEVIEW_KEY=
# Gethomepage widget — read-only stats at /api/v1/stats/summary
# Off by default. Set to a random secret to enable; clients must send
# the same value in the `X-API-Key` header.
# Generate: python3 -c "import secrets; print(secrets.token_urlsafe(32))"
# HOMEPAGE_API_KEY=
-15
View File
@@ -1,15 +0,0 @@
# These are supported funding model platforms
github: # Replace with up to 4 GitHub Sponsors-enabled usernames e.g., [user1, user2]
patreon: # Replace with a single Patreon username
open_collective: # Replace with a single Open Collective username
ko_fi: pouzor
tidelift: # Replace with a single Tidelift platform-name/package-name e.g., npm/babel
community_bridge: # Replace with a single Community Bridge project-name e.g., cloud-foundry
liberapay: # Replace with a single Liberapay username
issuehunt: # Replace with a single IssueHunt username
lfx_crowdfunding: # Replace with a single LFX Crowdfunding project-name e.g., cloud-foundry
polar: # Replace with a single Polar username
buy_me_a_coffee: # Replace with a single Buy Me a Coffee username
thanks_dev: # Replace with a single thanks.dev username
custom: # Replace with up to 4 custom sponsorship URLs e.g., ['link1', 'link2']
-3
View File
@@ -6,9 +6,6 @@ on:
pull_request:
branches: [main]
permissions:
contents: read
jobs:
smoke-and-integration:
runs-on: ubuntu-latest
+2 -9
View File
@@ -16,21 +16,14 @@ jobs:
matrix:
include:
- image: ghcr.io/pouzor/homelable-backend
context: .
dockerfile: Dockerfile.backend
build_args: ""
- image: ghcr.io/pouzor/homelable-frontend
context: .
dockerfile: Dockerfile.frontend
build_args: ""
- image: ghcr.io/pouzor/homelable-frontend-standalone
context: .
dockerfile: Dockerfile.frontend
build_args: "VITE_STANDALONE=true"
- image: ghcr.io/pouzor/homelable-mcp
context: ./mcp
dockerfile: Dockerfile.mcp
build_args: ""
steps:
- uses: actions/checkout@v4
@@ -62,8 +55,8 @@ jobs:
- name: Build and push
uses: docker/build-push-action@v6
with:
context: ${{ matrix.context }}
file: ${{ matrix.context }}/${{ matrix.dockerfile }}
context: .
file: ${{ matrix.dockerfile }}
platforms: linux/amd64,linux/arm64
push: true
tags: ${{ steps.meta.outputs.tags }}
-3
View File
@@ -6,9 +6,6 @@ on:
pull_request:
branches: [main]
permissions:
contents: read
jobs:
lint-scripts:
runs-on: ubuntu-latest
-3
View File
@@ -8,9 +8,6 @@ on:
schedule:
- cron: '0 9 * * 1' # Weekly on Monday
permissions:
contents: read
jobs:
secrets-scan:
runs-on: ubuntu-latest
-2
View File
@@ -45,8 +45,6 @@ htmlcov/
*.db
*.db-shm
*.db-wal
*.db.back
*.db.back-*
# Docker
.docker/
-276
View File
@@ -1,276 +0,0 @@
# Contributing to Homelable
Thanks for taking the time to contribute! This document covers everything you need to get started.
---
## Table of Contents
- [Ways to Contribute](#ways-to-contribute)
- [Reporting Bugs](#reporting-bugs)
- [Suggesting Features](#suggesting-features)
- [Development Setup](#development-setup)
- [Project Structure](#project-structure)
- [Coding Standards](#coding-standards)
- [Testing](#testing)
- [Submitting a Pull Request](#submitting-a-pull-request)
- [Commit Message Format](#commit-message-format)
---
## Ways to Contribute
- Report bugs or unexpected behavior
- Suggest new features or improvements
- Fix open issues (check the [issue tracker](https://github.com/Pouzor/homelable/issues))
- Improve documentation
- Add service signatures to `service_signatures.json`
---
## Reporting Bugs
Before opening an issue, search existing ones to avoid duplicates.
When filing a bug, include:
- **Homelable version** (visible in the sidebar bottom-left)
- **Deployment method** (Docker Compose, Proxmox LXC, source)
- **Steps to reproduce**
- **Expected vs actual behavior**
- **Relevant logs** (`docker compose logs backend` / `docker compose logs frontend`)
- **Browser console errors** if it's a UI issue
---
## Suggesting Features
Open an issue with the `enhancement` label. Describe:
- The problem you're trying to solve
- Your proposed solution
- Any alternatives you considered
For large changes, discuss first before writing code — it avoids wasted effort.
---
## Development Setup
### Prerequisites
- **Node.js 20+** and **npm**
- **Python 3.113.13** (3.14 not yet supported by all dependencies)
- **nmap** installed on your system (required for scanner)
- **Docker + Docker Compose** (optional, for full-stack testing)
### 1. Clone the repo
```bash
git clone https://github.com/Pouzor/homelable.git
cd homelable
```
### 2. Backend
```bash
cd backend
python3.13 -m venv .venv
source .venv/bin/activate # Windows: .venv\Scripts\activate
pip install -r requirements.txt
# Copy and configure environment
cp .env.example .env # edit AUTH_PASSWORD_HASH, SECRET_KEY, etc.
# Start the backend (auto-reloads on change)
uvicorn app.main:app --reload --port 8000
```
API docs available at `http://localhost:8000/docs`.
### 3. Frontend
```bash
cd frontend
npm install
npm run dev # http://localhost:5173
```
Vite proxies `/api` to `localhost:8000` — the backend must be running.
### 4. Verify tooling
```bash
./scripts/verify-tooling.sh
```
---
## Project Structure
```
homelable/
├── frontend/src/
│ ├── components/
│ │ ├── canvas/ # React Flow canvas, custom nodes & edges
│ │ ├── panels/ # Sidebar, detail panel, toolbar
│ │ ├── modals/ # Add/edit node, scan config, pending devices
│ │ └── ui/ # Shadcn/ui base components
│ ├── stores/ # Zustand state (canvas, auth, scan)
│ ├── hooks/ # Custom React hooks
│ ├── types/ # TypeScript interfaces & enums
│ ├── api/ # Axios client & typed endpoints
│ └── utils/ # Layout, export, color helpers
├── backend/app/
│ ├── api/routes/ # FastAPI route handlers
│ ├── services/ # Scanner, status checker, canvas service
│ ├── db/ # SQLAlchemy models, Alembic migrations
│ ├── schemas/ # Pydantic request/response schemas
│ └── core/ # Config, JWT, scheduler
├── docker/ # Nginx configs
├── scripts/ # LXC bootstrap, dev helpers
└── mcp/ # MCP server (AI integration)
```
---
## Coding Standards
### General
- No untested code merged — every feature or fix must include tests
- Keep changes focused — one concern per PR
### Frontend (TypeScript + React)
- Strict TypeScript — no `any`, no type assertions unless truly necessary
- React Flow node domain fields go in `node.data`, never on the node root
- State management via Zustand stores — no prop drilling beyond 2 levels
- Styling via TailwindCSS utility classes — follow the existing [design system](#design-system)
- Run before committing:
```bash
cd frontend
npm run lint
npm run typecheck
npm test
```
### Backend (Python + FastAPI)
- Python 3.11+ syntax
- Pydantic v2 schemas for all request/response types
- SQLAlchemy async sessions — never block the event loop
- Scanner logic runs in a background thread — never in an async route directly
- All schema changes via Alembic migrations — never modify tables directly
- Run before committing:
```bash
cd backend
source .venv/bin/activate
ruff check .
pytest
```
### Design System
| Token | Value |
|---|---|
| Background | `#0d1117` |
| Surface | `#161b22` |
| Card | `#21262d` |
| Accent cyan | `#00d4ff` |
| Online | `#39d353` |
| Offline | `#f85149` |
| Pending | `#e3b341` |
| Font (UI) | Inter |
| Font (IPs/ports) | JetBrains Mono |
---
## Testing
Tests run automatically via a pre-commit hook when frontend or backend files are staged.
### Frontend
```bash
cd frontend
npm test # run all tests
npm run test:coverage # with coverage report
```
Test files live in `__tests__/` next to their module, named `*.test.ts(x)`.
**What to test:** Zustand store actions, utility functions, non-trivial component logic.
### Backend
```bash
cd backend
source .venv/bin/activate
pytest # run all tests
pytest -v tests/test_nodes.py # single file
```
Test files live in `backend/tests/test_*.py`.
**What to test:** all API routes (happy path + error cases), auth flows, service logic.
Use the `client` and `headers` fixtures from `conftest.py` — they provide an in-memory SQLite database so tests are isolated and fast.
---
## Submitting a Pull Request
1. **Fork** the repo and create a branch from `main`:
```bash
git checkout -b feat/my-feature
```
2. **Make your changes** — include tests.
3. **Run the full test suite** (frontend + backend) and make sure everything passes.
4. **Open a PR** against `main`:
- Use a clear title (see commit format below)
- Describe what changed and why
- Reference any related issues (`Closes #123`)
- Include screenshots for UI changes
5. Keep the PR focused — one feature or fix per PR. Large refactors should be discussed in an issue first.
---
## Commit Message Format
Follow [Conventional Commits](https://www.conventionalcommits.org/):
```
<type>: <short description>
[optional body]
```
| Type | When to use |
|---|---|
| `feat` | New feature |
| `fix` | Bug fix |
| `docs` | Documentation only |
| `refactor` | Code change with no behavior change |
| `test` | Adding or fixing tests |
| `chore` | Build, deps, tooling |
**Examples:**
```
feat: add logout button to sidebar
fix: stop click propagation on pending device checkbox
docs: add CONTRIBUTING.md
```
---
## Questions?
Open a [GitHub Discussion](https://github.com/Pouzor/homelable/discussions) or drop a comment on a relevant issue.
+2 -3
View File
@@ -2,14 +2,13 @@ FROM python:3.13-slim
WORKDIR /app
# Install nmap for network scanning + iputils-ping for ping-based status checks + curl for the health check
RUN apt-get update && apt-get install -y --no-install-recommends nmap iputils-ping curl && rm -rf /var/lib/apt/lists/*
# Install nmap for network scanning + iputils-ping for ping-based status checks
RUN apt-get update && apt-get install -y --no-install-recommends nmap iputils-ping && rm -rf /var/lib/apt/lists/*
COPY backend/requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
COPY backend/ .
COPY VERSION /app/VERSION
# Create data directory (volume mount point)
RUN mkdir -p /app/data
+1 -3
View File
@@ -1,8 +1,7 @@
# Stage 1: build
# Use the native build platform so npm ci never runs under QEMU emulation.
# The build output (static HTML/JS/CSS) is platform-independent.
# node:20-slim (Debian/glibc) avoids lightningcss musl binary resolution issues on Alpine.
FROM --platform=$BUILDPLATFORM node:20-slim AS builder
FROM --platform=$BUILDPLATFORM node:20-alpine AS builder
ARG VITE_STANDALONE=false
ENV VITE_STANDALONE=$VITE_STANDALONE
@@ -12,7 +11,6 @@ COPY frontend/package*.json ./
RUN npm ci
COPY frontend/ .
COPY VERSION ../VERSION
RUN npm run build
# Stage 2: serve
+32 -17
View File
@@ -11,18 +11,9 @@ Open **http://localhost:3000** — login with `admin` / `admin`.
> Change the password before exposing to a network: edit `.env` and update `AUTH_USERNAME` / `AUTH_PASSWORD_HASH`.
>
Generate a new hash:
```bash
docker compose exec backend python -c "from passlib.context import CryptContext; print(CryptContext(schemes=['bcrypt']).hash('yourpassword'))"
```
⚠️ **bcrypt hashes contain `$` characters** — how to handle them depends on where you set the value:
- **`.env` file** (recommended): wrap the hash in single quotes → `AUTH_PASSWORD_HASH='$2b$12$...'`
- **`docker-compose.yml` `environment:` block**: escape every `$` as `$$` — use this command to generate a pre-escaped hash:
```bash
docker compose exec backend python -c "from passlib.context import CryptContext; print(CryptContext(schemes=['bcrypt']).hash('yourpassword').replace('\$', '\$\$'))"
```
> Generate a new hash: `docker compose exec backend python -c "from passlib.context import CryptContext; print(CryptContext(schemes=['bcrypt']).hash('yourpassword'))"`
>
> ⚠️ Keep the single quotes around the hash value in `.env` — bcrypt hashes contain `$` characters that Docker Compose would otherwise misinterpret.
## Quick Start — Frontend only
@@ -53,13 +44,37 @@ docker compose up -d
## Proxmox LXC Install
You can now install Homelable with community-scripts (proxmox-VE) :
`https://community-scripts.org/scripts/homelable`
Run this **on the Proxmox host** — it creates a Debian 12 LXC container and installs Homelable inside automatically:
```bash
bash -c "$(curl -fsSL https://raw.githubusercontent.com/community-scripts/ProxmoxVE/main/ct/homelable.sh)"
bash <(curl -fsSL https://raw.githubusercontent.com/Pouzor/homelable/main/scripts/install-proxmox.sh)
```
Default container settings: 2 cores, 1 GB RAM, 8 GB disk, DHCP on `vmbr0`. Override before running:
```bash
CTID=150 RAM=2048 STORAGE=local-zfs bash <(curl -fsSL .../install-proxmox.sh)
```
The backend runs as a systemd service, the frontend is served via nginx on port 80.
> To install manually inside an existing Debian/Ubuntu machine or LXC:
> ```bash
> bash <(curl -fsSL https://raw.githubusercontent.com/Pouzor/homelable/main/scripts/lxc-install.sh)
> ```
### Update (LXC)
Run the update script inside the container (pulls latest code, rebuilds frontend, restarts services — `.env` and database are never touched):
```bash
sudo bash /opt/homelable/scripts/update.sh
```
Or directly from GitHub:
```bash
sudo bash <(curl -fsSL https://raw.githubusercontent.com/Pouzor/homelable/main/scripts/update.sh)
```
---
-21
View File
@@ -1,21 +0,0 @@
MIT License
Copyright (c) 2026 Remy Jardinet
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+5 -125
View File
@@ -1,15 +1,13 @@
# Homelable
Homelable is a self-hosted infrastructure visualization solution. It provides a network/zigbee scanning feature to accelerate the identification of machines, devices and services deployed on your local infrastructure.
Homelable is a self-hosted infrastructure visualization solution. It provides a network scanning feature to accelerate the identification of machines and services deployed on your local infrastructure.
Homelable also offers a healthcheck system through multiple methods (ping/TCP, /health API, etc.) to get a global overview of online/offline services.
Homelable also offers a healthcheck system (WIP) through multiple methods (ping/TCP, /health API, etc.) to get a global overview of online/offline services.
You can also select some pre-built design styles, or personalize each device in your diagram.
If you just like the design, you can only run the frontend and export your design as PNG.
If you are running <img width="35" height="35" align="middle" alt="New_Home_Assistant_logo" src="https://github.com/user-attachments/assets/3bb17686-c706-40ce-a2d3-57e02378f37c" /> Homeassistant, check the [Homelable HA version](https://github.com/Pouzor/homelable-hacs) (via HACS)
---
@@ -18,9 +16,8 @@ If you are running <img width="35" height="35" align="middle" alt="New_Home_Ass
<p align="center">
<img src="docs/homelable1.png" alt="Homelable canvas overview" width="100%" />
<img src="docs/homelable2.png" alt="Homelable node detail" width="100%" />
<img src="docs/homelable4.png" alt="Homelable edit pannel" width="48%" />
<img width="48%" alt="Homelable Zigbee Network" src="https://github.com/user-attachments/assets/06caab68-6637-4dda-ab16-7e83f63d3972" />
<img src="docs/homelable3.png" alt="Homelable sidebar and scan" width="40%" />
<img src="docs/homelable4.png" alt="Homelable edit pannel" width="40%" />
</p>
---
@@ -77,118 +74,7 @@ Homelable continuously monitors your nodes and displays their live status (onlin
---
## Zigbee2MQTT Import
Homelable can connect directly to your MQTT broker and import your Zigbee network topology from **Zigbee2MQTT**, placing each device on the canvas as a typed node.
### Prerequisites
- A running **MQTT broker** (e.g. Mosquitto) accessible from the Homelable host
- **Zigbee2MQTT** connected to the broker with at least one device paired
### Usage
1. Click **Zigbee Import** in the left sidebar (below "Scan Network")
2. Enter your broker host, port (default `1883`), optional credentials, and base topic (default `zigbee2mqtt`)
3. Click **Test Connection** to verify reachability, then **Fetch Devices**
4. Select the devices you want from the grouped list (Coordinator / Router / End Device)
5. Click **Add N to Canvas** — devices are placed in a grid with IoT edges
### Node Types
| Type | Z2M Device | Icon |
|------|-----------|------|
| `zigbee_coordinator` | Coordinator | Network hub |
| `zigbee_router` | Router (mains-powered) | Radio |
| `zigbee_enddevice` | End Device (battery) | Antenna |
Hierarchy is set automatically: coordinator → routers → end devices (`parent_id`).
LQI (Link Quality Indicator) is stored as a node property.
> **Full documentation:** [docs/zigbee-import.md](./docs/zigbee-import.md)
---
## Live View (read-only public canvas)
Live View lets you share a read-only snapshot of your canvas with anyone on your network — no login required. It is disabled by default.
### Activation
Add LIVEVIEW_KEY to your .env:
`LIVEVIEW_KEY=your-secret-key`
Then restart the backend:
`docker compose restart backend`
### Usage
Use this URL to view your canvas:
http://<your-homelab-ip>/view?key=your-secret-key
The page shows your canvas in pan/zoom-only mode — no editing, no credentials needed. Clicking a node that has an IP opens it in a new tab.
---
## Gethomepage Widget (read-only stats)
Homelable can expose a small JSON stats endpoint that [gethomepage](https://gethomepage.dev) consumes through its built-in `customapi` widget. Disabled by default.
### Activation
Add `HOMEPAGE_API_KEY` to your `.env`:
`HOMEPAGE_API_KEY=your-secret-key`
Restart the backend (`docker compose restart backend`).
### Endpoint
`GET /api/v1/stats/summary` — requires header `X-API-Key: your-secret-key`. Returns:
```json
{
"nodes": 12,
"online": 9,
"offline": 2,
"unknown": 1,
"pending_devices": 3,
"zigbee_devices": 5,
"last_scan_at": "2026-05-14T10:00:00+00:00"
}
```
### gethomepage `services.yaml` snippet
```yaml
- Homelab:
- Homelable:
icon: mdi-lan
href: http://homelable.local:3000
widget:
type: customapi
url: http://homelable.local:8000/api/v1/stats/summary
method: GET
headers:
X-API-Key: your-secret-key
mappings:
- field: nodes ; label: Nodes
- field: online ; label: Online
- field: offline ; label: Offline
- field: pending_devices ; label: Pending
- field: zigbee_devices ; label: Zigbee
- field: last_scan_at ; label: Last scan
```
The backend port (`8000`) must be reachable from your gethomepage container.
---
## MCP Server (AI Integration) (optional)
## MCP Server (AI Integration) (optionnal)
Homelable can exposes a [Model Context Protocol](https://modelcontextprotocol.io) server so any MCP-compatible AI client (Claude Code, Claude Desktop, Open WebUI…) can read your homelab topology and act on it.
@@ -223,12 +109,6 @@ docker compose up -d mcp
# MCP server is now listening on http://<your-homelab-ip>:8001
```
> **Proxmox LXC / bare-metal (no Docker):** create the LXC via
> [community-scripts/ProxmoxVE](https://github.com/community-scripts/ProxmoxVE) (or any
> Debian/Ubuntu LXC), then inside it run `sudo bash scripts/lxc-mcp-install.sh`.
> Installs a `homelable-mcp` systemd service, prompts for `MCP_API_KEY` / `MCP_SERVICE_KEY`
> (auto-generated if you press Enter), and skips prompts if `mcp/.env` already exists.
**3. Configure your AI client:**
**Claude Code** — run this command in your terminal:
-1
View File
@@ -1 +0,0 @@
2.5.1
+18 -45
View File
@@ -1,14 +1,13 @@
import uuid
from datetime import datetime, timezone
from typing import Any
from fastapi import APIRouter, Depends, Query
from fastapi import APIRouter, Depends
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.deps import get_current_user
from app.db.database import get_db
from app.db.models import CanvasState, Design, Edge, Node
from app.db.models import CanvasState, Edge, Node
from app.schemas.canvas import CanvasSaveRequest, CanvasStateResponse
from app.schemas.edges import EdgeResponse
from app.schemas.nodes import NodeResponse
@@ -17,54 +16,33 @@ router = APIRouter()
@router.get("", response_model=CanvasStateResponse)
async def load_canvas(
design_id: str | None = Query(None, description="Design ID to load; uses first design if omitted"),
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> CanvasStateResponse:
if design_id is None:
first = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
design_id = first.id if first else None
if design_id is None:
return CanvasStateResponse(nodes=[], edges=[], viewport={"x": 0, "y": 0, "zoom": 1}, custom_style=None)
nodes = (await db.execute(select(Node).where(Node.design_id == design_id))).scalars().all()
edges = (await db.execute(select(Edge).where(Edge.design_id == design_id))).scalars().all()
state = await db.get(CanvasState, design_id)
async def load_canvas(db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)) -> CanvasStateResponse:
nodes = (await db.execute(select(Node))).scalars().all()
edges = (await db.execute(select(Edge))).scalars().all()
state = await db.get(CanvasState, 1)
viewport: dict[str, Any] = state.viewport if state else {"x": 0, "y": 0, "zoom": 1}
return CanvasStateResponse(
nodes=[NodeResponse.model_validate(n) for n in nodes],
edges=[EdgeResponse.model_validate(e) for e in edges],
viewport=viewport,
custom_style=state.custom_style if state else None,
)
@router.post("/save")
async def save_canvas(
body: CanvasSaveRequest, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)
) -> dict[str, bool | str]:
design_id = body.design_id
if design_id is None:
first = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
design_id = first.id if first else None
if design_id is None:
new_design = Design(id=str(uuid.uuid4()), name="Network Topology", design_type="network")
db.add(new_design)
await db.flush()
design_id = new_design.id
) -> dict[str, bool]:
incoming_node_ids = {n.id for n in body.nodes}
incoming_edge_ids = {e.id for e in body.edges}
# Delete nodes removed from canvas (only within this design)
existing_nodes = (await db.execute(select(Node).where(Node.design_id == design_id))).scalars().all()
# Delete nodes removed from canvas
existing_nodes = (await db.execute(select(Node))).scalars().all()
for node in existing_nodes:
if node.id not in incoming_node_ids:
await db.delete(node)
# Delete edges removed from canvas (only within this design)
existing_edges = (await db.execute(select(Edge).where(Edge.design_id == design_id))).scalars().all()
# Delete edges removed from canvas
existing_edges = (await db.execute(select(Edge))).scalars().all()
for edge in existing_edges:
if edge.id not in incoming_edge_ids:
await db.delete(edge)
@@ -74,33 +52,28 @@ async def save_canvas(
# Upsert nodes
for node_data in body.nodes:
db_node = await db.get(Node, node_data.id)
payload = node_data.model_dump()
payload["design_id"] = design_id
if db_node:
for field, value in payload.items():
for field, value in node_data.model_dump().items():
setattr(db_node, field, value)
else:
db.add(Node(**payload))
db.add(Node(**node_data.model_dump()))
# Upsert edges
for edge_data in body.edges:
db_edge = await db.get(Edge, edge_data.id)
payload = edge_data.model_dump()
payload["design_id"] = design_id
if db_edge:
for field, value in payload.items():
for field, value in edge_data.model_dump().items():
setattr(db_edge, field, value)
else:
db.add(Edge(**payload))
db.add(Edge(**edge_data.model_dump()))
# Upsert viewport + custom style
state = await db.get(CanvasState, design_id)
# Upsert viewport
state = await db.get(CanvasState, 1)
if state:
state.viewport = body.viewport
state.custom_style = body.custom_style
state.saved_at = datetime.now(timezone.utc)
else:
db.add(CanvasState(design_id=design_id, viewport=body.viewport, custom_style=body.custom_style))
db.add(CanvasState(id=1, viewport=body.viewport))
await db.commit()
return {"saved": True}
-81
View File
@@ -1,81 +0,0 @@
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.deps import get_current_user
from app.db.database import get_db
from app.db.models import CanvasState, Design, Edge, Node
from app.schemas.designs import DesignCreate, DesignResponse, DesignUpdate
router = APIRouter()
@router.get("", response_model=list[DesignResponse])
async def list_designs(
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> list[DesignResponse]:
designs = (await db.execute(select(Design).order_by(Design.created_at))).scalars().all()
return [DesignResponse.model_validate(d) for d in designs]
@router.post("", response_model=DesignResponse, status_code=201)
async def create_design(
body: DesignCreate,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> DesignResponse:
design = Design(name=body.name, design_type=body.design_type, icon=body.icon)
db.add(design)
await db.flush()
# Create empty canvas state for the new design
db.add(CanvasState(design_id=design.id))
await db.commit()
await db.refresh(design)
return DesignResponse.model_validate(design)
@router.put("/{design_id}", response_model=DesignResponse)
async def update_design(
design_id: str,
body: DesignUpdate,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> DesignResponse:
design = await db.get(Design, design_id)
if not design:
raise HTTPException(404, "Design not found")
if body.name is not None:
design.name = body.name
if body.icon is not None:
design.icon = body.icon
await db.commit()
await db.refresh(design)
return DesignResponse.model_validate(design)
@router.delete("/{design_id}", status_code=204)
async def delete_design(
design_id: str,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> None:
design = await db.get(Design, design_id)
if not design:
raise HTTPException(404, "Design not found")
# Count remaining designs — prevent deleting the last one
count = (await db.execute(select(Design))).scalars().all()
if len(count) <= 1:
raise HTTPException(400, "Cannot delete the only design")
# Delete associated canvas state, edges, nodes
cs = await db.get(CanvasState, design_id)
if cs:
await db.delete(cs)
edges = (await db.execute(select(Edge).where(Edge.design_id == design_id))).scalars().all()
for e in edges:
await db.delete(e)
nodes = (await db.execute(select(Node).where(Node.design_id == design_id))).scalars().all()
for n in nodes:
await db.delete(n)
await db.delete(design)
await db.commit()
-72
View File
@@ -1,72 +0,0 @@
import hmac
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.deps import get_current_user
from app.core.config import settings
from app.db.database import get_db
from app.db.models import CanvasState, Design, Edge, Node
from app.schemas.canvas import CanvasStateResponse
from app.schemas.edges import EdgeResponse
from app.schemas.nodes import NodeResponse
router = APIRouter()
class LiveViewConfigResponse(BaseModel):
"""Whether live view is enabled, plus the key (admin-only) to build share links."""
enabled: bool
key: str | None = None
@router.get("/config", response_model=LiveViewConfigResponse)
async def liveview_config(
_: str = Depends(get_current_user),
) -> LiveViewConfigResponse:
"""Authenticated: expose the configured live view key so the UI can build a
ready-to-use share link (e.g. /view?key=...&design=<id>).
Only reachable by a logged-in user — the key is never exposed publicly.
"""
key = settings.liveview_key or None
return LiveViewConfigResponse(enabled=bool(key), key=key)
@router.get("", response_model=CanvasStateResponse)
async def liveview_canvas(
key: str | None = Query(default=None),
design_id: str | None = Query(default=None, description="Design to show; uses first if omitted"),
db: AsyncSession = Depends(get_db),
) -> CanvasStateResponse:
"""Read-only public canvas endpoint.
Disabled by default — requires LIVEVIEW_KEY to be set in .env.
Always returns 403 when disabled, regardless of the key provided.
"""
if not settings.liveview_key:
raise HTTPException(status_code=403, detail="Live view is disabled")
if not key or not hmac.compare_digest(key, settings.liveview_key):
raise HTTPException(status_code=403, detail="Invalid live view key")
if design_id is None:
first = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
design_id = first.id if first else None
if design_id is None:
return CanvasStateResponse(nodes=[], edges=[], viewport={"x": 0, "y": 0, "zoom": 1}, custom_style=None)
nodes = (await db.execute(select(Node).where(Node.design_id == design_id))).scalars().all()
edges = (await db.execute(select(Edge).where(Edge.design_id == design_id))).scalars().all()
state = await db.get(CanvasState, design_id)
viewport: dict[str, Any] = state.viewport if state else {"x": 0, "y": 0, "zoom": 1}
custom_style: dict[str, Any] | None = state.custom_style if state else None
return CanvasStateResponse(
nodes=[NodeResponse.model_validate(n) for n in nodes],
edges=[EdgeResponse.model_validate(e) for e in edges],
viewport=viewport,
custom_style=custom_style,
)
+25 -348
View File
@@ -1,69 +1,23 @@
import ipaddress
import logging
import uuid
from typing import Any
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
from pydantic import BaseModel, field_validator
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.deps import get_current_user
from app.core.config import settings
from app.db.database import AsyncSessionLocal, get_db
from app.db.models import Design, Edge, Node, PendingDevice, PendingDeviceLink, ScanRun
from app.db.models import Node, PendingDevice, ScanRun
from app.schemas.nodes import NodeCreate
from app.schemas.scan import PendingDeviceResponse, ScanRunResponse
from app.services.scanner import request_cancel, run_scan
from app.services.zigbee_service import build_zigbee_properties
_ZIGBEE_TYPES = {"zigbee_coordinator", "zigbee_router", "zigbee_enddevice"}
def build_mac_property(mac: str | None) -> list[dict[str, Any]]:
"""Build a NodeProperty list carrying a device MAC address.
Shape matches the frontend ``NodeProperty`` type
(``{key, value, icon, visible}``). Hidden by default — the user opts in to
showing it on the canvas card from the right panel. Returns an empty list
when no MAC is known.
"""
if not mac:
return []
return [{"key": "MAC", "value": mac, "icon": None, "visible": False}]
def merge_mac_property(
props: list[dict[str, Any]] | None, mac: str | None
) -> list[dict[str, Any]]:
"""Append a MAC NodeProperty to ``props`` unless one is already present.
Preserves any user-supplied properties (and an existing MAC row's
visibility) untouched. Used on approve so the scanned MAC is not lost.
"""
out = [dict(p) for p in (props or [])]
if not mac or any(p.get("key") == "MAC" for p in out):
return out
out.append({"key": "MAC", "value": mac, "icon": None, "visible": False})
return out
class BulkActionRequest(BaseModel):
device_ids: list[str]
from app.services.scanner import run_scan
class ScanConfig(BaseModel):
ranges: list[str]
@field_validator("ranges")
@classmethod
def validate_cidr(cls, v: list[str]) -> list[str]:
for r in v:
try:
ipaddress.ip_network(r, strict=False)
except ValueError as exc:
raise ValueError(f"Invalid CIDR range: {r!r}") from exc
return v
interval_seconds: int
logger = logging.getLogger(__name__)
@@ -72,15 +26,7 @@ router = APIRouter()
async def _background_scan(run_id: str, ranges: list[str]) -> None:
async with AsyncSessionLocal() as db:
try:
await run_scan(ranges, db, run_id)
except Exception:
logger.exception("Scan run %s failed unexpectedly", run_id)
await db.rollback()
run = await db.get(ScanRun, run_id)
if run and run.status == "running":
run.status = "failed"
await db.commit()
await run_scan(ranges, db, run_id)
@router.post("/trigger", response_model=ScanRunResponse)
@@ -98,162 +44,18 @@ async def trigger_scan(
return run
@router.post("/{run_id}/stop", response_model=dict)
async def stop_scan(
run_id: str,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> dict[str, bool]:
try:
uuid.UUID(run_id)
except ValueError:
raise HTTPException(status_code=400, detail="Invalid run_id format") from None
run = await db.get(ScanRun, run_id)
if not run:
raise HTTPException(status_code=404, detail="Scan run not found")
if run.status != "running":
raise HTTPException(status_code=409, detail="Scan is not running")
request_cancel(run_id)
return {"stopping": True}
@router.get("/pending", response_model=list[PendingDeviceResponse])
async def list_pending(db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)) -> list[PendingDevice]:
result = await db.execute(select(PendingDevice).where(PendingDevice.status == "pending"))
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]:
from sqlalchemy import delete as sa_delete
result = await db.execute(sa_delete(PendingDevice).where(PendingDevice.status == "pending"))
await db.commit()
return {"deleted": result.rowcount}
@router.get("/hidden", response_model=list[PendingDeviceResponse])
async def list_hidden(db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)) -> list[PendingDevice]:
result = await db.execute(select(PendingDevice).where(PendingDevice.status == "hidden"))
return list(result.scalars().all())
@router.post("/pending/bulk-approve", response_model=dict)
async def bulk_approve_devices(
payload: BulkActionRequest,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> dict[str, Any]:
# Determine target design (use first design as fallback)
first_design = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
default_design_id = first_design.id if first_design else None
result = await db.execute(
select(PendingDevice).where(
PendingDevice.id.in_(payload.device_ids),
PendingDevice.status == "pending",
)
)
devices = result.scalars().all()
created_nodes: list[Node] = []
for device in devices:
device.status = "approved"
node_type = device.suggested_type or "generic"
is_zigbee = node_type in _ZIGBEE_TYPES
node = Node(
label=device.hostname or device.friendly_name or device.ip or "device",
type=node_type,
ip=device.ip,
mac=device.mac,
hostname=device.hostname,
status="online" if is_zigbee else "unknown",
services=device.services or [],
ieee_address=device.ieee_address,
properties=build_zigbee_properties(
device.ieee_address, device.vendor, device.model, device.lqi
) if is_zigbee else build_mac_property(device.mac),
# Default to ping so the status checker actually polls the new node.
# Without this the scheduler skips it (check_method NULL → no check).
check_method="none" if is_zigbee else ("ping" if device.ip else None),
design_id=default_design_id,
)
db.add(node)
created_nodes.append(node)
await db.flush() # populates node.id from Python-side default before reading
node_ids = [n.id for n in created_nodes]
approved_device_ids = [d.id for d in devices]
all_edges: list[dict[str, str]] = []
for device in devices:
all_edges.extend(await _resolve_pending_links_for_ieee(db, device.ieee_address))
await db.commit()
return {
"approved": len(node_ids),
"node_ids": node_ids,
"device_ids": approved_device_ids,
"edges_created": len(all_edges),
"edges": all_edges,
"skipped": len(payload.device_ids) - len(node_ids),
}
@router.post("/pending/bulk-hide", response_model=dict)
async def bulk_hide_devices(
payload: BulkActionRequest,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> dict[str, Any]:
result = await db.execute(
select(PendingDevice).where(
PendingDevice.id.in_(payload.device_ids),
PendingDevice.status == "pending",
)
)
devices = result.scalars().all()
for device in devices:
device.status = "hidden"
await db.commit()
return {"hidden": len(devices), "skipped": len(payload.device_ids) - len(devices)}
@router.post("/pending/{device_id}/restore", response_model=dict)
async def restore_device(
device_id: str,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> dict[str, Any]:
device = await db.get(PendingDevice, device_id)
if not device:
raise HTTPException(status_code=404, detail="Device not found")
if device.status != "hidden":
raise HTTPException(status_code=409, detail="Device is not hidden")
device.status = "pending"
await db.commit()
return {"restored": True, "device_id": device_id}
@router.post("/pending/bulk-restore", response_model=dict)
async def bulk_restore_devices(
payload: BulkActionRequest,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> dict[str, Any]:
result = await db.execute(
select(PendingDevice).where(
PendingDevice.id.in_(payload.device_ids),
PendingDevice.status == "hidden",
)
)
devices = result.scalars().all()
for device in devices:
device.status = "pending"
await db.commit()
return {"restored": len(devices), "skipped": len(payload.device_ids) - len(devices)}
@router.post("/pending/{device_id}/approve", response_model=dict)
async def approve_device(
device_id: str,
@@ -261,138 +63,14 @@ async def approve_device(
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> dict[str, Any]:
# Determine target design
node_design_id = node_data.design_id
if node_design_id is None:
first = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
node_design_id = first.id if first else None
device = await db.get(PendingDevice, device_id)
if not device:
raise HTTPException(status_code=404, detail="Device not found")
if device.status != "pending":
raise HTTPException(status_code=409, detail="Device already processed")
device.status = "approved"
_is_zigbee = node_data.type in _ZIGBEE_TYPES
# Prefer the MAC discovered during the scan (stored on the pending device);
# fall back to whatever the approve payload carried.
_mac = device.mac or node_data.mac
node = Node(
label=node_data.label,
type=node_data.type,
ip=node_data.ip,
mac=_mac,
hostname=node_data.hostname,
status="online" if _is_zigbee else node_data.status,
services=node_data.services or [],
ieee_address=device.ieee_address,
properties=build_zigbee_properties(
device.ieee_address, device.vendor, device.model, device.lqi
) if _is_zigbee else merge_mac_property(node_data.properties, _mac),
check_method="none" if _is_zigbee else (node_data.check_method or ("ping" if node_data.ip else None)),
check_target=None if _is_zigbee else node_data.check_target,
design_id=node_design_id,
)
db.add(node)
await db.flush()
node_id = node.id
edges = await _resolve_pending_links_for_ieee(db, device.ieee_address)
await db.commit()
return {
"approved": True,
"node_id": node_id,
"edges_created": len(edges),
"edges": edges,
}
async def _resolve_pending_links_for_ieee(
db: AsyncSession, ieee: str | None
) -> list[dict[str, str]]:
"""Materialize edges for any pending_device_links involving ``ieee``.
For each link where the other endpoint already exists as a canvas Node
(matched by ``Node.ieee_address``), create the Edge and drop the link
row. Links where the other endpoint is still pending are kept so they
can resolve when that endpoint is approved later.
"""
if not ieee:
return []
links_q = await db.execute(
select(PendingDeviceLink).where(
(PendingDeviceLink.source_ieee == ieee)
| (PendingDeviceLink.target_ieee == ieee)
)
)
links = list(links_q.scalars().all())
if not links:
return []
# Map every relevant ieee → Node (single query).
other_ieees = {
link.target_ieee if link.source_ieee == ieee else link.source_ieee
for link in links
}
other_ieees.add(ieee)
nodes_q = await db.execute(
select(Node).where(Node.ieee_address.in_(other_ieees))
)
by_ieee = {n.ieee_address: n for n in nodes_q.scalars().all() if n.ieee_address}
self_node = by_ieee.get(ieee)
if self_node is None:
return []
# Pre-fetch existing edges between these node ids so we don't create dups
# if the user re-approves a device or had drawn the link manually.
candidate_node_ids = [n.id for n in by_ieee.values()]
existing_q = await db.execute(
select(Edge).where(
Edge.source.in_(candidate_node_ids),
Edge.target.in_(candidate_node_ids),
)
)
existing_pairs = {(e.source, e.target) for e in existing_q.scalars().all()}
created: list[dict[str, str]] = []
for link in links:
other_ieee = (
link.target_ieee if link.source_ieee == ieee else link.source_ieee
)
other_node = by_ieee.get(other_ieee)
if other_node is None:
continue
if link.source_ieee == ieee:
src_id, tgt_id = self_node.id, other_node.id
else:
src_id, tgt_id = other_node.id, self_node.id
# Skip if either direction already exists.
if (src_id, tgt_id) in existing_pairs or (tgt_id, src_id) in existing_pairs:
await db.delete(link)
continue
# Use the source node's design_id for the edge
edge_design_id = self_node.design_id if self_node else None
if edge_design_id is None:
first = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
edge_design_id = first.id if first else None
edge = Edge(
source=src_id,
target=tgt_id,
type="iot",
source_handle="bottom",
target_handle="top-t",
design_id=edge_design_id,
)
db.add(edge)
await db.flush()
existing_pairs.add((src_id, tgt_id))
created.append({"id": edge.id, "source": src_id, "target": tgt_id})
await db.delete(link)
return created
if device:
device.status = "approved"
node = Node(**node_data.model_dump())
db.add(node)
await db.commit()
return {"approved": True, "node_id": node.id}
return {"approved": False}
@router.post("/pending/{device_id}/hide")
@@ -400,10 +78,9 @@ async def hide_device(
device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)
) -> dict[str, bool]:
device = await db.get(PendingDevice, device_id)
if not device:
raise HTTPException(status_code=404, detail="Device not found")
device.status = "hidden"
await db.commit()
if device:
device.status = "hidden"
await db.commit()
return {"hidden": True}
@@ -412,10 +89,9 @@ async def ignore_device(
device_id: str, db: AsyncSession = Depends(get_db), _: str = Depends(get_current_user)
) -> dict[str, bool]:
device = await db.get(PendingDevice, device_id)
if not device:
raise HTTPException(status_code=404, detail="Device not found")
await db.delete(device)
await db.commit()
if device:
await db.delete(device)
await db.commit()
return {"ignored": True}
@@ -427,17 +103,18 @@ async def list_runs(db: AsyncSession = Depends(get_db), _: str = Depends(get_cur
@router.get("/config", response_model=ScanConfig)
async def get_scan_config(_: str = Depends(get_current_user)) -> ScanConfig:
return ScanConfig(ranges=settings.scanner_ranges)
return ScanConfig(
ranges=settings.scanner_ranges,
interval_seconds=settings.status_checker_interval,
)
@router.post("/config", response_model=ScanConfig)
async def update_scan_config(payload: ScanConfig, _: str = Depends(get_current_user)) -> ScanConfig:
previous = settings.scanner_ranges
settings.scanner_ranges = payload.ranges
try:
settings.scanner_ranges = payload.ranges
settings.status_checker_interval = payload.interval_seconds
settings.save_overrides()
return payload
except Exception as exc:
settings.scanner_ranges = previous
logger.error("Failed to save scan config: %s", exc)
raise HTTPException(status_code=500, detail="Failed to save scan config") from exc
raise HTTPException(status_code=500, detail=str(exc)) from exc
-42
View File
@@ -1,42 +0,0 @@
"""App-level settings (status checker interval, etc.)."""
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
from app.api.deps import get_current_user
from app.core.config import settings
from app.core.scheduler import reschedule_service_checks, set_service_checks_enabled
router = APIRouter()
class AppSettings(BaseModel):
interval_seconds: int
service_check_enabled: bool = False
service_check_interval: int = Field(default=300, ge=30)
@router.get("", response_model=AppSettings)
async def get_settings(_: str = Depends(get_current_user)) -> AppSettings:
return AppSettings(
interval_seconds=settings.status_checker_interval,
service_check_enabled=settings.service_check_enabled,
service_check_interval=settings.service_check_interval,
)
@router.post("", response_model=AppSettings)
async def update_settings(
payload: AppSettings, _: str = Depends(get_current_user)
) -> AppSettings:
try:
settings.status_checker_interval = payload.interval_seconds
settings.service_check_enabled = payload.service_check_enabled
settings.service_check_interval = payload.service_check_interval
settings.save_overrides()
# Apply the service-check schedule live.
set_service_checks_enabled(payload.service_check_enabled)
if payload.service_check_enabled:
reschedule_service_checks(payload.service_check_interval)
return payload
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
-64
View File
@@ -1,64 +0,0 @@
import hmac
from fastapi import APIRouter, Depends, Header, HTTPException
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.db.database import get_db
from app.db.models import Node, PendingDevice, ScanRun
router = APIRouter()
def _check_key(x_api_key: str | None) -> None:
if not settings.homepage_api_key:
raise HTTPException(status_code=403, detail="Stats endpoint is disabled")
if not x_api_key or not hmac.compare_digest(x_api_key, settings.homepage_api_key):
raise HTTPException(status_code=403, detail="Invalid API key")
@router.get("/summary")
async def summary(
x_api_key: str | None = Header(default=None, alias="X-API-Key"),
db: AsyncSession = Depends(get_db),
) -> dict[str, object]:
"""Read-only stats payload for the gethomepage `customapi` widget.
Disabled unless HOMEPAGE_API_KEY is set. Caller must send the same
value in the `X-API-Key` header.
"""
_check_key(x_api_key)
status_rows = (
await db.execute(select(Node.status, func.count()).group_by(Node.status))
).all()
counts = {row[0]: row[1] for row in status_rows}
pending = (
await db.execute(
select(func.count())
.select_from(PendingDevice)
.where(PendingDevice.status == "pending")
)
).scalar_one()
zigbee = (
await db.execute(
select(func.count()).select_from(Node).where(Node.ieee_address.isnot(None))
)
).scalar_one()
last_scan_at = (
await db.execute(select(func.max(ScanRun.finished_at)))
).scalar_one()
return {
"nodes": sum(counts.values()),
"online": counts.get("online", 0),
"offline": counts.get("offline", 0),
"unknown": counts.get("unknown", 0),
"pending_devices": pending,
"zigbee_devices": zigbee,
"last_scan_at": last_scan_at.isoformat() if last_scan_at else None,
}
+2 -22
View File
@@ -1,4 +1,3 @@
import contextlib
import json
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
@@ -11,12 +10,6 @@ router = APIRouter()
_connections: list[WebSocket] = []
def _drop(websocket: WebSocket) -> None:
"""Remove a connection if still present — idempotent, never raises."""
with contextlib.suppress(ValueError):
_connections.remove(websocket)
@router.websocket("/ws/status")
async def ws_status(websocket: WebSocket) -> None:
# Accept first so we can send a close frame with a reason code
@@ -40,11 +33,7 @@ async def ws_status(websocket: WebSocket) -> None:
while True:
await websocket.receive_text()
except WebSocketDisconnect:
pass
finally:
# Any error (disconnect or otherwise) must release the slot, else the
# dead socket lingers in the broadcast pool.
_drop(websocket)
_connections.remove(websocket)
async def _broadcast(payload: str) -> None:
@@ -52,7 +41,7 @@ async def _broadcast(payload: str) -> None:
try:
await conn.send_text(payload)
except Exception:
_drop(conn)
_connections.remove(conn)
async def broadcast_status(node_id: str, status: str, checked_at: str, response_time_ms: int | None = None) -> None:
@@ -65,15 +54,6 @@ async def broadcast_status(node_id: str, status: str, checked_at: str, response_
}))
async def broadcast_service_status(node_id: str, services: list[dict[str, object]], checked_at: str) -> None:
await _broadcast(json.dumps({
"type": "service_status",
"node_id": node_id,
"services": services,
"checked_at": checked_at,
}))
async def broadcast_scan_update(run_id: str, devices_found: int) -> None:
await _broadcast(json.dumps({
"type": "scan_device_found",
-298
View File
@@ -1,298 +0,0 @@
"""FastAPI router for Zigbee2MQTT import."""
import logging
from datetime import datetime, timezone
from typing import Any
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
from sqlalchemy import delete as sa_delete
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.deps import get_current_user
from app.db.database import AsyncSessionLocal, get_db
from app.db.models import Design, Node, PendingDevice, PendingDeviceLink, ScanRun
from app.schemas.scan import ScanRunResponse
from app.schemas.zigbee import (
ZigbeeCoordinatorOut,
ZigbeeEdgeOut,
ZigbeeImportPendingResponse,
ZigbeeImportRequest,
ZigbeeImportResponse,
ZigbeeNodeOut,
ZigbeeTestConnectionRequest,
ZigbeeTestConnectionResponse,
)
from app.services.zigbee_service import (
build_zigbee_properties,
fetch_networkmap,
merge_zigbee_properties,
test_mqtt_connection,
)
logger = logging.getLogger(__name__)
router = APIRouter()
@router.post("/import", response_model=ZigbeeImportResponse)
async def import_zigbee_network(
payload: ZigbeeImportRequest,
_: str = Depends(get_current_user),
) -> ZigbeeImportResponse:
"""Fetch the Zigbee2MQTT network map and return nodes + edges ready for canvas drop.
Connects to the specified MQTT broker, publishes a networkmap request to
``<base_topic>/bridge/request/networkmap``, and waits up to 60 s for the
response (large meshes can take 30 s+). The devices are returned as typed homelable nodes with a
coordinator → router → end-device hierarchy.
"""
try:
nodes_raw, edges_raw = await fetch_networkmap(
mqtt_host=payload.mqtt_host,
mqtt_port=payload.mqtt_port,
base_topic=payload.base_topic,
username=payload.mqtt_username,
password=payload.mqtt_password,
tls=payload.mqtt_tls,
tls_insecure=payload.mqtt_tls_insecure,
)
except ImportError as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
except ConnectionError as exc:
raise HTTPException(status_code=502, detail=str(exc)) from exc
except TimeoutError as exc:
raise HTTPException(status_code=504, detail=str(exc)) from exc
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
except Exception as exc:
logger.exception("Unexpected error during Zigbee import")
raise HTTPException(status_code=500, detail="Unexpected error during Zigbee import") from exc
nodes = [ZigbeeNodeOut(**n) for n in nodes_raw]
edges = [ZigbeeEdgeOut(**e) for e in edges_raw]
return ZigbeeImportResponse(nodes=nodes, edges=edges, device_count=len(nodes))
@router.post("/import-pending", response_model=ScanRunResponse)
async def import_zigbee_to_pending(
payload: ZigbeeImportRequest,
background_tasks: BackgroundTasks,
db: AsyncSession = Depends(get_db),
_: str = Depends(get_current_user),
) -> ScanRun:
"""Queue a Zigbee2MQTT pending import as a background scan run.
Returns the ScanRun row immediately so the UI can close the import
modal and surface progress under Scan History (kind=zigbee). The
actual MQTT fetch + pending upsert happens in the background.
"""
run = ScanRun(
status="running",
kind="zigbee",
ranges=[f"{payload.mqtt_host}:{payload.mqtt_port}"],
)
db.add(run)
await db.commit()
await db.refresh(run)
background_tasks.add_task(_background_zigbee_import, run.id, payload)
return run
async def _background_zigbee_import(run_id: str, payload: ZigbeeImportRequest) -> None:
async with AsyncSessionLocal() as db:
try:
nodes_raw, edges_raw = await fetch_networkmap(
mqtt_host=payload.mqtt_host,
mqtt_port=payload.mqtt_port,
base_topic=payload.base_topic,
username=payload.mqtt_username,
password=payload.mqtt_password,
tls=payload.mqtt_tls,
tls_insecure=payload.mqtt_tls_insecure,
)
result = await _persist_pending_import(db, nodes_raw, edges_raw)
run = await db.get(ScanRun, run_id)
if run:
run.status = "done"
run.devices_found = result.device_count
run.finished_at = datetime.now(timezone.utc)
await db.commit()
except Exception as exc:
logger.exception("Zigbee import %s failed", run_id)
await db.rollback()
run = await db.get(ScanRun, run_id)
if run:
run.status = "error"
run.error = str(exc)[:500]
run.finished_at = datetime.now(timezone.utc)
await db.commit()
async def _persist_pending_import(
db: AsyncSession,
nodes_raw: list[dict[str, Any]],
edges_raw: list[dict[str, Any]],
) -> ZigbeeImportPendingResponse:
"""Upsert nodes/edges into pending_devices + pending_device_links.
Coordinator auto-approves to a canvas Node. Other devices upsert by IEEE.
All zigbee-source links are wiped and re-inserted from the new map.
"""
# Determine target design (use first design as fallback)
first_design = (await db.execute(select(Design).order_by(Design.created_at).limit(1))).scalar()
default_design_id = first_design.id if first_design else None
coordinator_out: ZigbeeCoordinatorOut | None = None
coordinator_existed = False
pending_created = 0
pending_updated = 0
for n in nodes_raw:
ieee = n.get("ieee_address")
if not ieee:
continue
props = build_zigbee_properties(
ieee, n.get("vendor"), n.get("model"), n.get("lqi")
)
if n.get("device_type") == "Coordinator":
existing = await db.execute(select(Node).where(Node.ieee_address == ieee))
existing_node = existing.scalar_one_or_none()
if existing_node:
existing_node.properties = merge_zigbee_properties(
existing_node.properties, props
)
coordinator_out = ZigbeeCoordinatorOut(
id=existing_node.id,
label=existing_node.label,
ieee_address=ieee,
)
coordinator_existed = True
continue
label = n.get("friendly_name") or ieee
node = Node(
label=label,
type=n.get("type") or "zigbee_coordinator",
status="online",
check_method="none",
ieee_address=ieee,
services=[],
properties=props,
design_id=default_design_id,
)
db.add(node)
await db.flush()
coordinator_out = ZigbeeCoordinatorOut(
id=node.id, label=label, ieee_address=ieee
)
continue
# If the device has already been approved as a canvas Node, refresh
# its properties and skip creating a pending row (keeps approved
# devices out of pending/hidden modals on re-import).
existing_node_q = await db.execute(
select(Node).where(Node.ieee_address == ieee)
)
existing_node = existing_node_q.scalar_one_or_none()
if existing_node:
existing_node.properties = merge_zigbee_properties(
existing_node.properties, props
)
continue
result = await db.execute(
select(PendingDevice).where(PendingDevice.ieee_address == ieee)
)
pending = result.scalar_one_or_none()
if pending is None:
db.add(
PendingDevice(
ieee_address=ieee,
friendly_name=n.get("friendly_name"),
hostname=n.get("friendly_name"),
suggested_type=n.get("type"),
device_subtype=n.get("device_type"),
model=n.get("model"),
vendor=n.get("vendor"),
lqi=n.get("lqi"),
status="pending",
discovery_source="zigbee",
)
)
pending_created += 1
else:
pending.friendly_name = n.get("friendly_name") or pending.friendly_name
pending.suggested_type = n.get("type") or pending.suggested_type
pending.device_subtype = n.get("device_type") or pending.device_subtype
pending.model = n.get("model") or pending.model
pending.vendor = n.get("vendor") or pending.vendor
if n.get("lqi") is not None:
pending.lqi = n.get("lqi")
if pending.status == "approved":
# The device was approved earlier but its canvas Node no longer
# exists (no Node matched the IEEE above) — it was deleted. Revive
# the row to "pending" so it reappears in the Pending list on
# re-import instead of being silently swallowed. (Issue #167)
pending.status = "pending"
elif pending.status == "hidden":
# Re-imported a hidden device → leave it hidden, just refresh fields.
pass
pending_updated += 1
# Replace all zigbee-source links with the freshly discovered set.
await db.execute(
sa_delete(PendingDeviceLink).where(PendingDeviceLink.discovery_source == "zigbee")
)
links_recorded = 0
seen: set[tuple[str, str]] = set()
for e in edges_raw:
src = e.get("source")
tgt = e.get("target")
if not src or not tgt or (src, tgt) in seen:
continue
seen.add((src, tgt))
db.add(
PendingDeviceLink(
source_ieee=src,
target_ieee=tgt,
discovery_source="zigbee",
)
)
links_recorded += 1
await db.commit()
return ZigbeeImportPendingResponse(
pending_created=pending_created,
pending_updated=pending_updated,
coordinator=coordinator_out,
coordinator_already_existed=coordinator_existed,
links_recorded=links_recorded,
device_count=len(nodes_raw),
)
@router.post("/test-connection", response_model=ZigbeeTestConnectionResponse)
async def test_zigbee_connection(
payload: ZigbeeTestConnectionRequest,
_: str = Depends(get_current_user),
) -> ZigbeeTestConnectionResponse:
"""Quick MQTT ping to validate broker connection before importing."""
try:
await test_mqtt_connection(
mqtt_host=payload.mqtt_host,
mqtt_port=payload.mqtt_port,
username=payload.mqtt_username,
password=payload.mqtt_password,
tls=payload.mqtt_tls,
tls_insecure=payload.mqtt_tls_insecure,
)
return ZigbeeTestConnectionResponse(connected=True, message="Connection successful")
except ImportError as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
except (ConnectionError, TimeoutError) as exc:
return ZigbeeTestConnectionResponse(connected=False, message=str(exc))
except Exception:
logger.exception("Unexpected error during connection test")
return ZigbeeTestConnectionResponse(connected=False, message="Unexpected error")
-46
View File
@@ -1,23 +1,8 @@
import json
import logging
from pathlib import Path
from pydantic import model_validator
from pydantic_settings import BaseSettings, SettingsConfigDict
logger = logging.getLogger(__name__)
def _read_version() -> str:
for candidate in [
Path(__file__).parent.parent.parent.parent / "VERSION", # repo root (dev)
Path("/app/VERSION"), # Docker image
]:
if candidate.exists():
return candidate.read_text().strip()
return "unknown"
APP_VERSION = _read_version()
class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8")
@@ -34,42 +19,17 @@ class Settings(BaseSettings):
auth_username: str = "admin"
auth_password_hash: str = ""
@model_validator(mode="after")
def check_password_hash(self) -> "Settings":
h = self.auth_password_hash
if h and not h.startswith("$2"):
logger.error(
"AUTH_PASSWORD_HASH looks invalid (does not start with '$2b$'). "
"bcrypt hashes contain '$' signs — wrap the value in single quotes "
"in your .env file: AUTH_PASSWORD_HASH='$2b$12$...'"
)
return self
# Scanner
scanner_ranges: list[str] = ["192.168.1.0/24"]
# Status checker
status_checker_interval: int = 60
# Per-service status checker (independent of node checks). Off by default.
service_check_enabled: bool = False
service_check_interval: int = 300
# MCP service key — set MCP_SERVICE_KEY in .env
# Used by the MCP server to authenticate against the backend without a user password.
# Leave empty to disable MCP service key auth.
mcp_service_key: str = ""
# Live view — optional read-only public canvas endpoint.
# Set to a random secret string to enable /api/v1/liveview?key=<value>.
# Leave unset (or empty) to keep the feature disabled (default).
liveview_key: str | None = None
# Homepage widget — optional read-only stats endpoint for gethomepage.
# Set to a random secret to enable /api/v1/stats/summary (X-API-Key header).
# Leave empty to keep the feature disabled (default).
homepage_api_key: str = ""
def _override_path(self) -> Path:
return Path(self.sqlite_path).parent / "scan_config.json"
@@ -81,10 +41,6 @@ class Settings(BaseSettings):
self.scanner_ranges = data["scanner_ranges"]
if "status_checker_interval" in data:
self.status_checker_interval = int(data["status_checker_interval"])
if "service_check_enabled" in data:
self.service_check_enabled = bool(data["service_check_enabled"])
if "service_check_interval" in data:
self.service_check_interval = int(data["service_check_interval"])
except Exception:
pass
@@ -94,8 +50,6 @@ class Settings(BaseSettings):
self._override_path().write_text(json.dumps({
"scanner_ranges": self.scanner_ranges,
"status_checker_interval": self.status_checker_interval,
"service_check_enabled": self.service_check_enabled,
"service_check_interval": self.service_check_interval,
}))
+23 -146
View File
@@ -1,5 +1,4 @@
"""APScheduler setup for background scan and status check jobs."""
import asyncio
import logging
from datetime import datetime, timezone
@@ -9,172 +8,50 @@ from sqlalchemy import select
from app.core.config import settings
from app.db.database import AsyncSessionLocal
from app.db.models import Node
from app.services.status_checker import check_node, check_services
from app.services.status_checker import check_node
logger = logging.getLogger(__name__)
scheduler: AsyncIOScheduler = AsyncIOScheduler()
async def _check_single_node(
node_id: str,
check_method: str,
check_target: str | None,
ip: str | None,
) -> tuple[str, dict[str, object] | None]:
"""Run a single node check; returns (node_id, result_or_None).
Accepts plain scalars — not an ORM object — so there is no risk of
DetachedInstanceError when the originating session has already closed.
"""
async def _run_status_checks() -> None:
"""Check all nodes and broadcast results via WebSocket."""
from app.api.routes.status import broadcast_status # avoid circular import
try:
check_result = await check_node(check_method, check_target, ip)
now = datetime.now(timezone.utc)
async with AsyncSessionLocal() as db:
n = await db.get(Node, node_id)
if n:
n.status = check_result["status"]
n.response_time_ms = check_result["response_time_ms"]
if check_result["status"] == "online":
n.last_seen = now
await db.commit()
await broadcast_status(
node_id=node_id,
status=check_result["status"],
checked_at=now.isoformat(),
response_time_ms=check_result["response_time_ms"],
)
return node_id, check_result
except Exception as exc:
logger.error("Status check failed for node %s: %s", node_id, exc)
return node_id, None
async def _run_status_checks() -> None:
"""Check all nodes concurrently and broadcast results via WebSocket."""
async with AsyncSessionLocal() as db:
result = await db.execute(select(Node))
nodes = result.scalars().all()
# Extract scalars while the session is open to avoid DetachedInstanceError
checkable = [
(n.id, n.check_method, n.check_target, n.ip)
for n in nodes
if n.check_method
]
if not checkable:
return
await asyncio.gather(*[
_check_single_node(node_id, method, target, ip)
for node_id, method, target, ip in checkable
])
def _node_host(ip: str | None, hostname: str | None) -> str | None:
"""Pick the address to probe services on: first IP, else hostname."""
if ip:
first = ip.split(",")[0].strip()
if first:
return first
return hostname or None
async def _run_service_checks() -> None:
"""Check every service of every node and broadcast per-service results."""
if not settings.service_check_enabled:
return
from app.api.routes.status import broadcast_service_status # avoid circular import
async with AsyncSessionLocal() as db:
result = await db.execute(select(Node))
nodes = result.scalars().all()
checkable = [
(n.id, _node_host(n.ip, n.hostname), list(n.services or []))
for n in nodes
if n.services
]
now = datetime.now(timezone.utc).isoformat()
for node_id, host, services in checkable:
for node in nodes:
if not node.check_method:
continue
try:
statuses = await check_services(host, services)
await broadcast_service_status(node_id=node_id, services=statuses, checked_at=now)
check_result = await check_node(node.check_method, node.check_target, node.ip)
async with AsyncSessionLocal() as db:
n = await db.get(Node, node.id)
if n:
n.status = check_result["status"]
n.response_time_ms = check_result["response_time_ms"]
n.last_seen = datetime.now(timezone.utc) if check_result["status"] == "online" else n.last_seen
await db.commit()
await broadcast_status(
node_id=node.id,
status=check_result["status"],
checked_at=datetime.now(timezone.utc).isoformat(),
response_time_ms=check_result["response_time_ms"],
)
except Exception as exc:
logger.error("Service checks failed for node %s: %s", node_id, exc)
def _add_service_check_job() -> None:
scheduler.add_job(
_run_service_checks,
"interval",
seconds=settings.service_check_interval,
id="service_checks",
max_instances=1,
coalesce=True,
)
logger.error("Status check failed for node %s: %s", node.id, exc)
def start_scheduler() -> None:
global scheduler
if scheduler.running:
try:
scheduler.shutdown(wait=False)
except Exception as exc:
logger.warning("Failed to shut down previous scheduler instance: %s", exc)
scheduler = AsyncIOScheduler()
scheduler.add_job(
_run_status_checks,
"interval",
seconds=settings.status_checker_interval,
id="status_checks",
max_instances=1,
coalesce=True,
)
if settings.service_check_enabled:
_add_service_check_job()
scheduler.add_job(_run_status_checks, "interval", seconds=settings.status_checker_interval, id="status_checks")
scheduler.start()
logger.info("Scheduler started — status checks every %ds", settings.status_checker_interval)
def reschedule_status_checks(interval_seconds: int) -> None:
"""Update the status check interval on the running scheduler."""
if interval_seconds < 10:
raise ValueError(f"interval_seconds must be >= 10, got {interval_seconds}")
if not scheduler.running:
logger.warning("Scheduler not running, skipping reschedule")
return
scheduler.reschedule_job("status_checks", trigger="interval", seconds=interval_seconds)
logger.info("Status checks rescheduled to every %ds", interval_seconds)
def reschedule_service_checks(interval_seconds: int) -> None:
"""Update the service-check interval on the running scheduler (if enabled)."""
if interval_seconds < 30:
raise ValueError(f"interval_seconds must be >= 30, got {interval_seconds}")
if not scheduler.running:
logger.warning("Scheduler not running, skipping reschedule")
return
if scheduler.get_job("service_checks"):
scheduler.reschedule_job("service_checks", trigger="interval", seconds=interval_seconds)
logger.info("Service checks rescheduled to every %ds", interval_seconds)
def set_service_checks_enabled(enabled: bool) -> None:
"""Add or remove the service-check job on the running scheduler."""
if not scheduler.running:
return
job = scheduler.get_job("service_checks")
if enabled and not job:
_add_service_check_job()
logger.info("Service checks enabled — every %ds", settings.service_check_interval)
elif not enabled and job:
scheduler.remove_job("service_checks")
logger.info("Service checks disabled")
def stop_scheduler() -> None:
if scheduler.running:
scheduler.shutdown(wait=False)
scheduler.shutdown(wait=False)
+5 -8
View File
@@ -1,22 +1,19 @@
from datetime import datetime, timedelta, timezone
import bcrypt
from jose import JWTError, jwt
from passlib.context import CryptContext
from app.core.config import settings
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
def verify_password(plain: str, hashed: str) -> bool:
if not plain or not hashed:
return False
try:
return bcrypt.checkpw(plain.encode("utf-8"), hashed.encode("utf-8"))
except (ValueError, TypeError):
return False
return bool(pwd_context.verify(plain, hashed))
def hash_password(password: str) -> str:
return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8")
return str(pwd_context.hash(password))
def create_access_token(subject: str) -> str:
+17 -260
View File
@@ -1,35 +1,11 @@
import json as _json
import logging
import shutil
import uuid as _uuid_mod
from collections.abc import AsyncGenerator
from contextlib import suppress
from pathlib import Path
from sqlalchemy.exc import OperationalError
from sqlalchemy.ext.asyncio import AsyncConnection, AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import DeclarativeBase
from app.core.config import APP_VERSION, settings
logger = logging.getLogger(__name__)
async def _try_migrate(conn: AsyncConnection, sql: str, *, label: str) -> None:
"""Run an idempotent migration statement, logging any error.
Distinguishes 'already applied' errors (debug) from genuine failures
(warning) so silent corruption is avoided. Used for new in-commit
migrations; existing legacy ALTERs above remain wrapped in suppress.
"""
try:
await conn.exec_driver_sql(sql)
except OperationalError as exc:
msg = str(exc).lower()
if "duplicate column" in msg or "already exists" in msg:
logger.debug("Migration %s skipped (already applied): %s", label, exc)
else:
logger.warning("Migration %s failed: %s", label, exc)
from app.core.config import settings
# Ensure the data directory exists before SQLite tries to open the file
Path(settings.sqlite_path).parent.mkdir(parents=True, exist_ok=True)
@@ -46,259 +22,40 @@ class Base(DeclarativeBase):
pass
def _backup_db() -> None:
db_path = Path(settings.sqlite_path)
if not db_path.exists():
return
backup_path = db_path.with_suffix(f".db.back-{APP_VERSION}")
if backup_path.exists():
return
try:
shutil.copy2(db_path, backup_path)
logger.info("DB backup created: %s", backup_path.name)
except OSError:
logger.warning("Could not create DB backup at %s", backup_path)
async def init_db() -> None:
_backup_db()
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
# Add columns introduced after initial schema (idempotent)
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN container_mode BOOLEAN NOT NULL DEFAULT 0")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN custom_colors JSON")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN custom_color TEXT")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN path_style TEXT")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN custom_icon TEXT")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN source_handle TEXT")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN target_handle TEXT")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN animated BOOLEAN NOT NULL DEFAULT 0")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN cpu_count INTEGER")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN cpu_model TEXT")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN ram_gb REAL")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN disk_gb REAL")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN show_hardware BOOLEAN NOT NULL DEFAULT 0")
with suppress(OperationalError):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN show_port_numbers BOOLEAN NOT NULL DEFAULT 0")
with suppress(OperationalError):
with suppress(Exception):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN width REAL")
with suppress(OperationalError):
with suppress(Exception):
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")
with suppress(OperationalError):
await conn.exec_driver_sql("ALTER TABLE scan_runs ADD COLUMN kind TEXT NOT NULL DEFAULT 'ip'")
# --- Zigbee schema migrations (logged variant per CLAUDE.md feedback) ---
zigbee_migrations: list[tuple[str, str]] = [
("nodes.ieee_address", "ALTER TABLE nodes ADD COLUMN ieee_address TEXT"),
(
"nodes.ieee_address.index",
"CREATE INDEX IF NOT EXISTS ix_nodes_ieee_address ON nodes(ieee_address)",
),
("pending_devices.ieee_address", "ALTER TABLE pending_devices ADD COLUMN ieee_address TEXT"),
(
"pending_devices.ieee_address.index",
"CREATE INDEX IF NOT EXISTS ix_pending_devices_ieee_address "
"ON pending_devices(ieee_address)",
),
("pending_devices.friendly_name", "ALTER TABLE pending_devices ADD COLUMN friendly_name TEXT"),
("pending_devices.device_subtype", "ALTER TABLE pending_devices ADD COLUMN device_subtype TEXT"),
("pending_devices.model", "ALTER TABLE pending_devices ADD COLUMN model TEXT"),
("pending_devices.vendor", "ALTER TABLE pending_devices ADD COLUMN vendor TEXT"),
("pending_devices.lqi", "ALTER TABLE pending_devices ADD COLUMN lqi INTEGER"),
]
for label, sql in zigbee_migrations:
await _try_migrate(conn, sql, label=label)
# Drop NOT NULL on pending_devices.ip (Zigbee devices have no IP).
# SQLite can't ALTER column nullability — rebuild the table if needed.
try:
info = await conn.exec_driver_sql("PRAGMA table_info(pending_devices)")
cols = info.fetchall()
ip_col = next((c for c in cols if c[1] == "ip"), None)
# PRAGMA table_info row layout: (cid, name, type, notnull, dflt, pk)
if ip_col and ip_col[3] == 1:
logger.info("Migrating pending_devices: dropping NOT NULL on ip column")
await conn.exec_driver_sql("PRAGMA foreign_keys = OFF")
await conn.exec_driver_sql(
"CREATE TABLE pending_devices_new ("
"id VARCHAR PRIMARY KEY,"
"ip VARCHAR,"
"mac VARCHAR, hostname VARCHAR, os VARCHAR, services JSON,"
"suggested_type VARCHAR,"
"status VARCHAR,"
"discovery_source VARCHAR,"
"ieee_address VARCHAR,"
"friendly_name VARCHAR,"
"device_subtype VARCHAR,"
"model VARCHAR,"
"vendor VARCHAR,"
"lqi INTEGER,"
"discovered_at DATETIME"
")"
)
await conn.exec_driver_sql(
"INSERT INTO pending_devices_new "
"(id, ip, mac, hostname, os, services, suggested_type, status, "
"discovery_source, ieee_address, friendly_name, device_subtype, "
"model, vendor, lqi, discovered_at) "
"SELECT id, ip, mac, hostname, os, services, suggested_type, status, "
"discovery_source, ieee_address, friendly_name, device_subtype, "
"model, vendor, lqi, discovered_at FROM pending_devices"
)
await conn.exec_driver_sql("DROP TABLE pending_devices")
await conn.exec_driver_sql(
"ALTER TABLE pending_devices_new RENAME TO pending_devices"
)
await conn.exec_driver_sql(
"CREATE INDEX IF NOT EXISTS ix_pending_devices_ieee_address "
"ON pending_devices(ieee_address)"
)
await conn.exec_driver_sql("PRAGMA foreign_keys = ON")
except OperationalError as exc:
logger.warning("pending_devices ip-nullable rebuild failed: %s", exc)
# --- end Zigbee schema migrations -------------------------------------
# --- Electrical designs schema migrations -----------------------------
# Create designs table (idempotent)
await _try_migrate(
conn,
"CREATE TABLE IF NOT EXISTS designs ("
"id VARCHAR PRIMARY KEY,"
"name VARCHAR NOT NULL,"
"design_type VARCHAR NOT NULL DEFAULT 'network',"
"created_at DATETIME,"
"updated_at DATETIME"
")",
label="designs.table",
)
# Add user-chosen icon to designs (idempotent), then backfill existing rows
# so legacy designs keep a sensible icon based on their original type.
await _try_migrate(
conn, "ALTER TABLE designs ADD COLUMN icon VARCHAR", label="designs.icon",
)
with suppress(OperationalError):
await conn.exec_driver_sql(
"UPDATE designs SET icon = 'zap' WHERE icon IS NULL AND design_type = 'electrical'"
)
with suppress(OperationalError):
await conn.exec_driver_sql(
"UPDATE designs SET icon = 'dashboard' WHERE icon IS NULL"
)
# Seed default Network Topology design if designs table is empty
_default_design_id = str(_uuid_mod.uuid4())
row = await conn.exec_driver_sql("SELECT COUNT(*) FROM designs")
count_row = row.fetchone()
count = count_row[0] if count_row else 0
if count == 0:
await conn.exec_driver_sql(
"INSERT INTO designs (id, name, design_type, icon, created_at, updated_at) "
"VALUES (?, 'Network Topology', 'network', 'dashboard', datetime('now'), datetime('now'))",
(_default_design_id,),
)
else:
row2 = await conn.exec_driver_sql("SELECT id FROM designs WHERE design_type = 'network' LIMIT 1")
default = row2.fetchone()
_default_design_id = default[0] if default else _default_design_id
# Add design_id to nodes
await _try_migrate(
conn, "ALTER TABLE nodes ADD COLUMN design_id VARCHAR REFERENCES designs(id)",
label="nodes.design_id",
)
# Assign existing nodes to default design
await conn.exec_driver_sql(
"UPDATE nodes SET design_id = ? WHERE design_id IS NULL", (_default_design_id,),
)
# Add design_id to edges
await _try_migrate(
conn, "ALTER TABLE edges ADD COLUMN design_id VARCHAR REFERENCES designs(id)",
label="edges.design_id",
)
# Assign existing edges to default design
await conn.exec_driver_sql(
"UPDATE edges SET design_id = ? WHERE design_id IS NULL", (_default_design_id,),
)
# Migrate canvas_state from id=1 to design_id PK (SQLite rebuild)
try:
info = await conn.exec_driver_sql("PRAGMA table_info(canvas_state)")
cols = info.fetchall()
has_design_id = any(c[1] == "design_id" for c in cols)
if not has_design_id:
logger.info("Migrating canvas_state: switching to design_id primary key")
await conn.exec_driver_sql("PRAGMA foreign_keys = OFF")
await conn.exec_driver_sql(
"CREATE TABLE canvas_state_new ("
"design_id VARCHAR PRIMARY KEY REFERENCES designs(id) ON DELETE CASCADE,"
"viewport JSON,"
"custom_style JSON,"
"saved_at DATETIME"
")"
)
# Copy existing row(s), mapping id=1 to default design_id
old_rows = await conn.exec_driver_sql("SELECT id, viewport, custom_style, saved_at FROM canvas_state")
for old in old_rows.fetchall():
cs_id, viewport, custom_style, saved_at = old
target_design = _default_design_id
await conn.exec_driver_sql(
"INSERT INTO canvas_state_new (design_id, viewport, custom_style, saved_at) "
"VALUES (?, ?, ?, ?)",
(target_design, viewport, custom_style, saved_at),
)
await conn.exec_driver_sql("DROP TABLE canvas_state")
await conn.exec_driver_sql("ALTER TABLE canvas_state_new RENAME TO canvas_state")
await conn.exec_driver_sql("PRAGMA foreign_keys = ON")
except OperationalError as exc:
logger.warning("canvas_state migration failed: %s", exc)
# --- end Electrical designs schema migrations --------------------------
with suppress(OperationalError):
await conn.exec_driver_sql("ALTER TABLE edges ADD COLUMN waypoints JSON")
with suppress(OperationalError):
await conn.exec_driver_sql("ALTER TABLE nodes ADD COLUMN properties JSON")
# Migrate hardware columns → properties JSON (idempotent: only runs on nodes where properties IS NULL)
with suppress(OperationalError):
rows = await conn.exec_driver_sql(
"SELECT id, cpu_model, cpu_count, ram_gb, disk_gb, show_hardware "
"FROM nodes WHERE properties IS NULL"
)
for r in rows.fetchall():
node_id, cpu_model, cpu_count, ram_gb, disk_gb, show_hardware = r
props = []
visible = bool(show_hardware)
if cpu_model:
props.append({"key": "CPU Model", "value": str(cpu_model), "icon": "Cpu", "visible": visible})
if cpu_count is not None:
props.append({"key": "CPU Cores", "value": str(cpu_count), "icon": "Cpu", "visible": visible})
if ram_gb is not None:
props.append({"key": "RAM", "value": f"{ram_gb} GB", "icon": "MemoryStick", "visible": visible})
if disk_gb is not None:
props.append({"key": "Disk", "value": f"{disk_gb} GB", "icon": "HardDrive", "visible": visible})
await conn.exec_driver_sql(
"UPDATE nodes SET properties = ? WHERE id = ?",
(_json.dumps(props), node_id),
)
# Migrate animated column from boolean (0/1) to string ('none'/'snake')
with suppress(OperationalError):
await conn.exec_driver_sql("UPDATE edges SET animated = 'snake' WHERE animated = '1' OR animated = 1")
with suppress(OperationalError):
sql = "UPDATE edges SET animated = 'none' WHERE animated = '0' OR animated = 0 OR animated IS NULL"
await conn.exec_driver_sql(sql)
async def get_db() -> AsyncGenerator[AsyncSession, None]:
+5 -50
View File
@@ -16,24 +16,12 @@ def _uuid() -> str:
return str(uuid.uuid4())
class Design(Base):
__tablename__ = "designs"
id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid)
name: Mapped[str] = mapped_column(String, nullable=False)
design_type: Mapped[str] = mapped_column(String, nullable=False, default="network")
icon: Mapped[str | None] = mapped_column(String, nullable=True, default="dashboard")
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now, onupdate=_now)
class Node(Base):
__tablename__ = "nodes"
id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid)
type: Mapped[str] = mapped_column(String, nullable=False)
label: Mapped[str] = mapped_column(String, nullable=False)
design_id: Mapped[str | None] = mapped_column(String, ForeignKey("designs.id", ondelete="SET NULL"), nullable=True)
hostname: Mapped[str | None] = mapped_column(String)
ip: Mapped[str | None] = mapped_column(String)
mac: Mapped[str | None] = mapped_column(String)
@@ -45,7 +33,7 @@ class Node(Base):
notes: Mapped[str | None] = mapped_column(Text)
pos_x: Mapped[float] = mapped_column(Float, default=0)
pos_y: Mapped[float] = mapped_column(Float, default=0)
parent_id: Mapped[str | None] = mapped_column(String, ForeignKey("nodes.id", ondelete="CASCADE"))
parent_id: Mapped[str | None] = mapped_column(String, ForeignKey("nodes.id"))
container_mode: Mapped[bool] = mapped_column(Boolean, default=False)
custom_colors: Mapped[dict[str, Any] | None] = mapped_column(JSON, nullable=True)
custom_icon: Mapped[str | None] = mapped_column(String, nullable=True)
@@ -54,16 +42,13 @@ class Node(Base):
ram_gb: Mapped[float | None] = mapped_column(Float, nullable=True)
disk_gb: Mapped[float | None] = mapped_column(Float, nullable=True)
show_hardware: Mapped[bool] = mapped_column(Boolean, default=False)
show_port_numbers: Mapped[bool] = mapped_column(Boolean, default=False)
properties: Mapped[list[Any]] = mapped_column(JSON, default=list)
width: Mapped[float | None] = mapped_column(Float, nullable=True)
height: Mapped[float | None] = mapped_column(Float, nullable=True)
bottom_handles: Mapped[int] = mapped_column(Integer, default=1)
ieee_address: Mapped[str | None] = mapped_column(String, index=True, nullable=True)
last_seen: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
response_time_ms: Mapped[int | None] = mapped_column(Integer)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now, onupdate=_now)
children: Mapped[list["Node"]] = relationship("Node", back_populates="parent")
parent: Mapped["Node | None"] = relationship("Node", back_populates="children", remote_side=[id])
@@ -74,26 +59,23 @@ class Edge(Base):
id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid)
source: Mapped[str] = mapped_column(String, ForeignKey("nodes.id", ondelete="CASCADE"))
target: Mapped[str] = mapped_column(String, ForeignKey("nodes.id", ondelete="CASCADE"))
design_id: Mapped[str | None] = mapped_column(String, ForeignKey("designs.id", ondelete="SET NULL"), nullable=True)
type: Mapped[str] = mapped_column(String, default="ethernet")
label: Mapped[str | None] = mapped_column(String)
vlan_id: Mapped[int | None] = mapped_column(Integer)
speed: Mapped[str | None] = mapped_column(String)
custom_color: Mapped[str | None] = mapped_column(String)
path_style: Mapped[str | None] = mapped_column(String)
animated: Mapped[str] = mapped_column(String, nullable=False, default='none')
animated: Mapped[bool] = mapped_column(Boolean, default=False)
source_handle: Mapped[str | None] = mapped_column(String)
target_handle: Mapped[str | None] = mapped_column(String)
waypoints: Mapped[list[dict[str, float]] | None] = mapped_column(JSON, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
class CanvasState(Base):
__tablename__ = "canvas_state"
design_id: Mapped[str] = mapped_column(String, ForeignKey("designs.id", ondelete="CASCADE"), primary_key=True)
id: Mapped[int] = mapped_column(Integer, primary_key=True, default=1)
viewport: Mapped[dict[str, Any]] = mapped_column(JSON, default=dict)
custom_style: Mapped[dict[str, Any] | None] = mapped_column(JSON, nullable=True)
saved_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
@@ -101,39 +83,13 @@ class PendingDevice(Base):
__tablename__ = "pending_devices"
id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid)
ip: Mapped[str | None] = mapped_column(String, nullable=True)
ip: Mapped[str] = mapped_column(String, nullable=False)
mac: Mapped[str | None] = mapped_column(String)
hostname: Mapped[str | None] = mapped_column(String)
os: Mapped[str | None] = mapped_column(String)
services: Mapped[list[Any]] = mapped_column(JSON, default=list)
suggested_type: Mapped[str | None] = mapped_column(String)
status: Mapped[str] = mapped_column(String, default="pending")
discovery_source: Mapped[str | None] = mapped_column(String)
ieee_address: Mapped[str | None] = mapped_column(String, index=True, nullable=True, unique=True)
friendly_name: Mapped[str | None] = mapped_column(String, nullable=True)
device_subtype: Mapped[str | None] = mapped_column(String, nullable=True)
model: Mapped[str | None] = mapped_column(String, nullable=True)
vendor: Mapped[str | None] = mapped_column(String, nullable=True)
lqi: Mapped[int | None] = mapped_column(Integer, nullable=True)
discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
class PendingDeviceLink(Base):
"""Link between two Zigbee endpoints discovered during import.
Endpoints are addressed by IEEE (stable across re-imports). Either side may
already exist as a canvas Node (resolved via Node.ieee_address) or still be
a PendingDevice. On approval, the matching Edge is auto-created when both
endpoints exist as canvas Nodes.
"""
__tablename__ = "pending_device_links"
id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid)
source_ieee: Mapped[str] = mapped_column(String, nullable=False, index=True)
target_ieee: Mapped[str] = mapped_column(String, nullable=False, index=True)
lqi: Mapped[int | None] = mapped_column(Integer, nullable=True)
discovery_source: Mapped[str] = mapped_column(String, nullable=False, default="zigbee")
discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
@@ -142,7 +98,6 @@ class ScanRun(Base):
id: Mapped[str] = mapped_column(String, primary_key=True, default=_uuid)
status: Mapped[str] = mapped_column(String, default="running")
kind: Mapped[str] = mapped_column(String, default="ip", server_default="ip")
ranges: Mapped[list[str]] = mapped_column(JSON, default=list)
devices_found: Mapped[int] = mapped_column(Integer, default=0)
started_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_now)
+2 -20
View File
@@ -1,5 +1,3 @@
import logging
import logging.config
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from typing import Any
@@ -7,8 +5,7 @@ from typing import Any
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from app.api.routes import auth, canvas, designs, edges, liveview, nodes, scan, stats, status, zigbee
from app.api.routes import settings as settings_routes
from app.api.routes import auth, canvas, edges, nodes, scan, status
from app.core.config import settings
from app.core.scheduler import start_scheduler, stop_scheduler
from app.db.database import init_db
@@ -16,16 +13,6 @@ from app.db.database import init_db
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
# Ensure app logs are visible: attach a handler to the root logger if none
# exists (uvicorn only installs handlers on its own loggers, not the root).
root_logger = logging.getLogger()
if not root_logger.handlers:
handler = logging.StreamHandler()
handler.setFormatter(logging.Formatter("%(levelname)s:%(name)s:%(message)s"))
root_logger.addHandler(handler)
root_logger.setLevel(logging.INFO)
logging.getLogger("app").setLevel(logging.INFO)
logging.getLogger("app.services.scanner").setLevel(logging.INFO)
await init_db()
settings.load_overrides()
start_scheduler()
@@ -35,7 +22,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
app = FastAPI(
title="Homelable API",
version="1.9.0",
version="1.3.3",
lifespan=lifespan,
)
@@ -51,13 +38,8 @@ app.include_router(auth.router, prefix="/api/v1/auth", tags=["auth"])
app.include_router(nodes.router, prefix="/api/v1/nodes", tags=["nodes"])
app.include_router(edges.router, prefix="/api/v1/edges", tags=["edges"])
app.include_router(canvas.router, prefix="/api/v1/canvas", tags=["canvas"])
app.include_router(designs.router, prefix="/api/v1/designs", tags=["designs"])
app.include_router(scan.router, prefix="/api/v1/scan", tags=["scan"])
app.include_router(status.router, prefix="/api/v1/status", tags=["status"])
app.include_router(settings_routes.router, prefix="/api/v1/settings", tags=["settings"])
app.include_router(liveview.router, prefix="/api/v1/liveview", tags=["liveview"])
app.include_router(zigbee.router, prefix="/api/v1/zigbee", tags=["zigbee"])
app.include_router(stats.router, prefix="/api/v1/stats", tags=["stats"])
@app.get("/api/v1/health")
+2 -15
View File
@@ -1,10 +1,9 @@
from typing import Any
from pydantic import BaseModel, field_validator
from pydantic import BaseModel
from app.schemas.edges import EdgeResponse
from app.schemas.nodes import NodeResponse
from app.schemas.utils import normalize_animated
class NodeSave(BaseModel):
@@ -29,11 +28,8 @@ class NodeSave(BaseModel):
ram_gb: float | None = None
disk_gb: float | None = None
show_hardware: bool = False
show_port_numbers: bool = False
properties: list[Any] = []
width: float | None = None
height: float | None = None
bottom_handles: int = 1
pos_x: float = 0
pos_y: float = 0
@@ -48,27 +44,18 @@ class EdgeSave(BaseModel):
speed: str | None = None
custom_color: str | None = None
path_style: str | None = None
animated: str = 'none'
animated: bool = False
source_handle: str | None = None
target_handle: str | None = None
waypoints: list[dict[str, float]] | None = None
@field_validator('animated', mode='before')
@classmethod
def validate_animated(cls, v: object) -> str:
return normalize_animated(v)
class CanvasSaveRequest(BaseModel):
nodes: list[NodeSave] = []
edges: list[EdgeSave] = []
viewport: dict[str, Any] = {}
custom_style: dict[str, Any] | None = None
design_id: str | None = None
class CanvasStateResponse(BaseModel):
nodes: list[NodeResponse]
edges: list[EdgeResponse]
viewport: dict[str, Any]
custom_style: dict[str, Any] | None = None
-27
View File
@@ -1,27 +0,0 @@
from datetime import datetime
from pydantic import BaseModel
class DesignCreate(BaseModel):
name: str
icon: str = "dashboard"
# Vestigial: kept for backward compatibility. The UI no longer branches on it;
# the chosen icon now drives presentation. Defaults to a generic canvas.
design_type: str = "network"
class DesignUpdate(BaseModel):
name: str | None = None
icon: str | None = None
class DesignResponse(BaseModel):
id: str
name: str
design_type: str
icon: str | None = None
created_at: datetime
updated_at: datetime
model_config = {"from_attributes": True}
+4 -20
View File
@@ -1,8 +1,6 @@
from datetime import datetime
from pydantic import BaseModel, field_validator
from app.schemas.utils import normalize_animated
from pydantic import BaseModel
class EdgeBase(BaseModel):
@@ -14,19 +12,13 @@ class EdgeBase(BaseModel):
speed: str | None = None
custom_color: str | None = None
path_style: str | None = None
animated: str = 'none'
animated: bool = False
source_handle: str | None = None
target_handle: str | None = None
waypoints: list[dict[str, float]] | None = None
@field_validator('animated', mode='before')
@classmethod
def validate_animated(cls, v: object) -> str:
return normalize_animated(v)
class EdgeCreate(EdgeBase):
design_id: str | None = None
pass
class EdgeUpdate(BaseModel):
@@ -36,17 +28,9 @@ class EdgeUpdate(BaseModel):
speed: str | None = None
custom_color: str | None = None
path_style: str | None = None
animated: str | None = None
animated: bool | None = None
source_handle: str | None = None
target_handle: str | None = None
waypoints: list[dict[str, float]] | None = None
@field_validator('animated', mode='before')
@classmethod
def validate_animated(cls, v: object) -> str | None:
if v is None:
return None
return normalize_animated(v)
class EdgeResponse(EdgeBase):
+1 -9
View File
@@ -27,15 +27,12 @@ class NodeBase(BaseModel):
ram_gb: float | None = None
disk_gb: float | None = None
show_hardware: bool = False
show_port_numbers: bool = False
properties: list[dict[str, Any]] = []
width: float | None = None
height: float | None = None
bottom_handles: int = 1
class NodeCreate(NodeBase):
design_id: str | None = None
pass
class NodeUpdate(BaseModel):
@@ -61,17 +58,12 @@ class NodeUpdate(BaseModel):
ram_gb: float | None = None
disk_gb: float | None = None
show_hardware: bool | None = None
show_port_numbers: bool | None = None
properties: list[dict[str, Any]] | None = None
width: float | None = None
height: float | None = None
bottom_handles: int | None = None
class NodeResponse(NodeBase):
id: str
design_id: str | None = None
ieee_address: str | None = None
last_seen: datetime | None = None
response_time_ms: int | None = None
created_at: datetime
+1 -9
View File
@@ -6,20 +6,13 @@ from pydantic import BaseModel
class PendingDeviceResponse(BaseModel):
id: str
ip: str | None
ip: str
mac: str | None
hostname: str | None
os: str | None
services: list[Any]
suggested_type: str | None
status: str
discovery_source: str | None
ieee_address: str | None = None
friendly_name: str | None = None
device_subtype: str | None = None
model: str | None = None
vendor: str | None = None
lqi: int | None = None
discovered_at: datetime
model_config = {"from_attributes": True}
@@ -28,7 +21,6 @@ class PendingDeviceResponse(BaseModel):
class ScanRunResponse(BaseModel):
id: str
status: str
kind: str = "ip"
ranges: list[str]
devices_found: int
started_at: datetime
-9
View File
@@ -1,9 +0,0 @@
def normalize_animated(v: object) -> str:
"""Normalize legacy bool/int animated values to string mode ('none'/'snake'/'flow')."""
if v is True or v == 1 or v == '1':
return 'snake'
if v is False or v == 0 or v == '0' or v is None or v == 'none':
return 'none'
if v in ('snake', 'flow', 'basic'):
return str(v)
return 'none'
-95
View File
@@ -1,95 +0,0 @@
"""Pydantic v2 schemas for Zigbee2MQTT import."""
from pydantic import BaseModel, Field, model_validator
class ZigbeeImportRequest(BaseModel):
mqtt_host: str = Field(..., description="MQTT broker hostname or IP address")
mqtt_port: int = Field(1883, ge=1, le=65535, description="MQTT broker port")
mqtt_username: str | None = Field(None, description="MQTT username (optional)")
mqtt_password: str | None = Field(None, description="MQTT password (optional)")
base_topic: str = Field("zigbee2mqtt", description="Zigbee2MQTT base topic")
mqtt_tls: bool = Field(False, description="Enable TLS (typically port 8883)")
mqtt_tls_insecure: bool = Field(
False, description="Skip TLS certificate verification (self-signed only)"
)
@model_validator(mode="after")
def _insecure_requires_tls(self) -> "ZigbeeImportRequest":
if self.mqtt_tls_insecure and not self.mqtt_tls:
raise ValueError("mqtt_tls_insecure requires mqtt_tls=true")
return self
class ZigbeeTestConnectionRequest(BaseModel):
mqtt_host: str
mqtt_port: int = Field(1883, ge=1, le=65535)
mqtt_username: str | None = None
mqtt_password: str | None = None
mqtt_tls: bool = False
mqtt_tls_insecure: bool = False
@model_validator(mode="after")
def _insecure_requires_tls(self) -> "ZigbeeTestConnectionRequest":
if self.mqtt_tls_insecure and not self.mqtt_tls:
raise ValueError("mqtt_tls_insecure requires mqtt_tls=true")
return self
class ZigbeeDeviceData(BaseModel):
ieee_address: str
friendly_name: str
device_type: str # Coordinator, Router, EndDevice
model: str | None = None
vendor: str | None = None
description: str | None = None
lqi: int | None = None
last_seen: str | None = None
class ZigbeeNodeOut(BaseModel):
"""A homelable-ready node representation of a Zigbee device."""
id: str
label: str
type: str # zigbee_coordinator | zigbee_router | zigbee_enddevice
ieee_address: str
friendly_name: str
device_type: str
model: str | None = None
vendor: str | None = None
lqi: int | None = None
parent_id: str | None = None
class ZigbeeEdgeOut(BaseModel):
source: str
target: str
class ZigbeeImportResponse(BaseModel):
nodes: list[ZigbeeNodeOut]
edges: list[ZigbeeEdgeOut]
device_count: int
class ZigbeeTestConnectionResponse(BaseModel):
connected: bool
message: str
class ZigbeeCoordinatorOut(BaseModel):
id: str
label: str
ieee_address: str
class ZigbeeImportPendingResponse(BaseModel):
"""Result of importing a Z2M network into the pending section."""
pending_created: int
pending_updated: int
coordinator: ZigbeeCoordinatorOut | None = None
coordinator_already_existed: bool = False
links_recorded: int
device_count: int
+9 -49
View File
@@ -65,46 +65,14 @@ def fingerprint_ports(open_ports: list[dict[str, Any]]) -> list[dict[str, Any]]:
return results
# Known OUI prefixes lowercase, colon-separated, first 3 octets
# Known OUI prefixes for virtual machines / hypervisors (lowercase, colon-separated)
_MAC_OUI_TYPES: dict[str, str] = {
# Hypervisors / VMs
"52:54:00": "vm", # QEMU/KVM (Proxmox VMs)
"bc:24:11": "vm", # Proxmox official OUI (VMs and LXC, 7.3+)
"52:54:00": "vm", # QEMU/KVM (used by Proxmox VMs)
"bc:24:11": "vm", # Proxmox official OUI (VMs and LXC, Proxmox 7.3+)
"00:50:56": "vm", # VMware
"00:0c:29": "vm", # VMware Workstation / Fusion
"08:00:27": "vm", # VirtualBox
"00:15:5d": "vm", # Hyper-V
# Shelly
"34:94:54": "iot",
"84:f3:eb": "iot",
"ec:fa:bc": "iot",
"30:c6:f7": "iot",
# Espressif (ESP8266 / ESP32 — used by Sonoff, many DIY IoT)
"a0:20:a6": "iot",
"24:62:ab": "iot",
"30:ae:a4": "iot",
"cc:50:e3": "iot",
"ac:67:b2": "iot",
"b4:e6:2d": "iot",
"3c:71:bf": "iot",
"8c:aa:b5": "iot",
# Sonoff / ITEAD
"dc:4f:22": "iot",
"e8:db:84": "iot",
# Tapo / TP-Link smart home
"b0:a7:b9": "iot",
"50:c7:bf": "iot",
"1c:3b:f3": "iot",
"10:27:f5": "iot",
# Philips Hue
"00:17:88": "iot",
"ec:b5:fa": "iot",
# IKEA Tradfri
"34:13:e8": "iot",
"00:21:2e": "iot",
# Tuya / Smart Life (widely used chip in many brands)
"d8:f1:5b": "iot",
"68:57:2d": "iot",
}
@@ -133,13 +101,10 @@ _PORT_TYPE_HINTS: dict[int, str] = {
37777: "camera", # Dahua
34567: "camera", # Amcrest
2020: "camera", # Tapo
# Smart-home / MQTT / CoAP → iot
# Smart-home / MQTT → iot
1883: "iot",
8883: "iot",
6052: "iot", # ESPHome dashboard
4915: "iot", # Shelly CoIoT
5683: "iot", # CoAP (Shelly Gen1, many IoT devices)
5684: "iot", # CoAP DTLS
6052: "iot", # ESPHome
# AP / wireless
8880: "ap", # UniFi HTTP
8443: "ap", # UniFi HTTPS
@@ -150,13 +115,8 @@ _PORT_TYPE_HINTS: dict[int, str] = {
def suggest_node_type(open_ports: list[dict[str, Any]], mac: str | None = None) -> str:
"""Suggest a node type based on matched signatures, port hints, and MAC OUI."""
# IoT vendor MACs are a strong, unambiguous signal — don't let generic HTTP ports override
mac_type = suggest_type_from_mac(mac)
if mac_type == "iot":
return "iot"
priority = ["proxmox", "nas", "router", "lxc", "vm", "ap", "camera", "iot", "server", "switch"]
"""Suggest a node type based on matched signatures and MAC OUI."""
priority = ["proxmox", "nas", "router", "lxc", "vm", "server", "ap", "camera", "iot", "switch"]
found: set[str] = set()
for p in open_ports:
port = p["port"]
@@ -166,10 +126,10 @@ def suggest_node_type(open_ports: list[dict[str, Any]], mac: str | None = None)
found.add(sig["suggested_node_type"])
if port in _PORT_TYPE_HINTS:
found.add(_PORT_TYPE_HINTS[port])
# MAC OUI is a lower-priority hint — only used if ports give no better answer
mac_type = suggest_type_from_mac(mac)
if mac_type:
found.add(mac_type)
for t in priority:
if t in found:
return t
+92 -432
View File
@@ -1,48 +1,17 @@
"""Network scanner: ARP sweep + nmap service detection + mDNS discovery."""
"""Network scanner: ARP sweep + nmap service detection."""
import asyncio
import ipaddress
import logging
import os
import re
import socket
import subprocess
import threading
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.db.models import Node, PendingDevice, ScanRun
from app.db.models import PendingDevice, ScanRun
from app.services.fingerprint import fingerprint_ports, suggest_node_type
logger = logging.getLogger(__name__)
# Run IDs that have been requested to cancel (thread-safe via lock)
_cancelled_runs: set[str] = set()
_cancelled_lock = threading.Lock()
# Port list for service detection (Phase 2)
_EXTRA_PORTS = (
"80,443,22,21,23,25,53,110,143,161,162,179,389,445,548,"
"554,636,873,1883,1880,1935,2020,2375,2376,3000,3001,3306,"
"3389,4711,4915,5000,5001,5432,5601,5683,5684,5900,5984,"
"6052,6379,6432,6443,6767,6789,6800,7878,8000,8006,8080,"
"8081,8086,8088,8090,8096,8112,8123,8200,8291,8428,8443,"
"8554,8686,8789,8843,8880,8883,8971,8989,9000,9001,9090,"
"9091,9092,9093,9100,9117,9200,9300,9411,9443,9696,10051,"
"16686,34567,37777,51413,64738"
)
_MDNS_SERVICE_TYPES = [
"_http._tcp.local.",
"_shelly._tcp.local.",
"_esphomelib._tcp.local.",
"_hap._tcp.local.", # HomeKit Accessory Protocol
"_mqtt._tcp.local.",
"_device-info._tcp.local.",
]
try:
import nmap
_NMAP_AVAILABLE = True
@@ -50,24 +19,50 @@ except ImportError:
_NMAP_AVAILABLE = False
logger.warning("python-nmap not available — scanner will run in mock mode")
try:
from zeroconf import ServiceStateChange
from zeroconf.asyncio import AsyncServiceBrowser, AsyncServiceInfo, AsyncZeroconf
_ZEROCONF_AVAILABLE = True
except ImportError:
_ZEROCONF_AVAILABLE = False
logger.warning("zeroconf not available — mDNS discovery disabled")
def _nmap_scan(target: str) -> list[dict[str, Any]]:
"""Run nmap -sV --open on target, return list of host dicts."""
if not _NMAP_AVAILABLE:
return _mock_scan(target)
def request_cancel(run_id: str) -> None:
"""Signal a running scan to stop early."""
with _cancelled_lock:
_cancelled_runs.add(run_id)
nm = nmap.PortScanner()
try:
# Home lab port range: standard top-1000 + common self-hosted service ports
extra_ports = (
"80,443,22,21,23,25,53,110,143,161,162,179,389,445,548,"
"554,636,873,1883,1880,1935,2020,2375,2376,3000,3001,3306,"
"3389,4711,5000,5001,5432,5601,5900,5984,6052,6379,6432,6443,"
"6767,6789,6800,7878,8000,8006,8080,8081,8086,8088,8090,8096,"
"8112,8123,8200,8291,8428,8443,8554,8686,8789,8843,8880,8883,"
"8971,8989,9000,9001,9090,9091,9092,9093,9100,9117,9200,9300,"
"9411,9443,9696,10051,16686,34567,37777,51413,64738"
)
nm.scan(hosts=target, arguments=f"-sV --open -T4 --host-timeout 120s -p {extra_ports}")
except Exception as exc:
logger.error("nmap scan failed: %s", exc)
raise RuntimeError(str(exc)) from exc
def _is_cancelled(run_id: str) -> bool:
with _cancelled_lock:
return run_id in _cancelled_runs
hosts = []
for host in nm.all_hosts():
if nm[host].state() != "up":
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", ""),
})
hosts.append({
"ip": host,
"hostname": _resolve_hostname(host),
"mac": nm[host].get("addresses", {}).get("mac"),
"os": _extract_os(nm, host),
"open_ports": open_ports,
})
return hosts
def _resolve_hostname(ip: str) -> str | None:
@@ -87,278 +82,6 @@ def _extract_os(nm: object, host: str) -> str | None:
return None
def _arp_table_hosts(network: str) -> dict[str, dict[str, Any]]:
"""
Read the OS ARP cache for recently-seen hosts in the target network.
Works without root on both Linux (/proc/net/arp) and macOS (arp -a).
Supplements nmap discovery — catches IoT and devices with all ports filtered.
"""
try:
net = ipaddress.ip_network(network, strict=False)
found: dict[str, dict[str, Any]] = {}
# Linux: parse /proc/net/arp — present on any Linux kernel (including Docker)
proc_arp = "/proc/net/arp"
try:
with open(proc_arp) as f:
for line in f.readlines()[1:]: # skip header row
parts = line.split()
if len(parts) >= 4:
ip, mac = parts[0], parts[3]
if mac == "00:00:00:00:00:00":
continue
try:
if ipaddress.ip_address(ip) in net:
found[ip] = {
"ip": ip, "mac": mac,
"hostname": _resolve_hostname(ip),
"os": None, "open_ports": [],
}
except ValueError:
pass
# /proc/net/arp opened successfully — return whatever we found (may be empty)
# Don't fall through to `arp -a` since we're on Linux
return found
except FileNotFoundError:
pass # Not Linux — fall through to macOS `arp -a`
# macOS: parse `arp -a` output
result = subprocess.run(["arp", "-a"], capture_output=True, text=True, timeout=5)
for line in result.stdout.splitlines():
m = re.search(r"\((\d+\.\d+\.\d+\.\d+)\)\s+at\s+([0-9a-f:]+)", line)
if not m:
continue
ip, mac = m.group(1), m.group(2)
if mac in ("(incomplete)", "ff:ff:ff:ff:ff:ff"):
continue
try:
if ipaddress.ip_address(ip) in net:
found[ip] = {"ip": ip, "mac": mac, "hostname": _resolve_hostname(ip), "os": None, "open_ports": []}
except ValueError:
pass
return found
except Exception as exc:
logger.warning("[Phase 1] ARP cache lookup failed: %s", exc)
return {}
async def _ping_sweep(target: str) -> dict[str, dict[str, Any]]:
"""
Phase 1: Concurrent ICMP ping sweep + ARP cache.
Pings all IPs in the CIDR in parallel (up to 50 at once, 1s timeout each).
Supplements with the OS ARP cache to catch devices that block ICMP.
Works in Docker with CAP_NET_RAW — no nmap, no false positives.
"""
net = ipaddress.ip_network(target, strict=False)
all_ips = [str(ip) for ip in net.hosts()]
logger.info("[Phase 1] Pinging %d hosts in %s ...", len(all_ips), target)
sem = asyncio.Semaphore(50)
async def _ping(ip: str) -> str | None:
async with sem:
try:
proc = await asyncio.create_subprocess_exec(
"ping", "-c", "1", "-W", "1", ip,
stdout=asyncio.subprocess.DEVNULL,
stderr=asyncio.subprocess.DEVNULL,
)
await proc.wait()
return ip if proc.returncode == 0 else None
except Exception:
return None
ping_results = await asyncio.gather(*[_ping(ip) for ip in all_ips])
alive_ips: set[str] = {ip for ip in ping_results if ip is not None}
logger.info("[Phase 1] %d/%d hosts responded to ping", len(alive_ips), len(all_ips))
# ARP cache: catch devices that block ICMP but were recently active,
# and enrich ping-alive hosts with their MAC addresses.
arp_cache = await asyncio.to_thread(_arp_table_hosts, target)
alive: dict[str, dict[str, Any]] = {}
for ip in alive_ips:
mac = arp_cache.get(ip, {}).get("mac")
hostname = await asyncio.to_thread(_resolve_hostname, ip)
logger.info("[Phase 1] %s mac=%s hostname=%s (ping)", ip, mac or "n/a", hostname or "n/a")
alive[ip] = {"ip": ip, "mac": mac, "hostname": hostname, "os": None, "open_ports": []}
for ip, host in arp_cache.items():
if ip not in alive:
logger.info(
"[Phase 1] %s mac=%s hostname=%s (ARP cache only)",
ip, host.get("mac") or "n/a", host.get("hostname") or "n/a",
)
alive[ip] = host
return alive
def _nmap_scan_single(host_dict: dict[str, Any]) -> dict[str, Any]:
"""
Phase 2 — single-IP port scan with service detection.
Runs in a thread (blocking). Returns the host dict enriched with open_ports.
"""
ip = host_dict["ip"]
logger.info("[Phase 2] Scanning %s ...", ip)
if not _NMAP_AVAILABLE:
logger.warning("[Phase 2] nmap not available, skipping %s", ip)
return host_dict
is_root = os.geteuid() == 0
if is_root:
# SYN scan + version detection (fastest, most accurate)
scan_args = f"-sS -sV --open -T4 -Pn --host-timeout 60s -p {_EXTRA_PORTS}"
else:
# TCP connect scan (-sT) — no raw sockets needed, works without root.
# nmap auto-selects -sT without root but being explicit avoids edge cases.
scan_args = f"-sT -sV --open -T4 -Pn --host-timeout 60s -p {_EXTRA_PORTS}"
logger.debug("[Phase 2] %s args: %s", ip, scan_args)
nm = nmap.PortScanner()
try:
nm.scan(hosts=ip, arguments=scan_args)
except Exception as exc:
logger.warning("[Phase 2] nmap FAILED for %s (%s: %s) — skipping port scan", ip, type(exc).__name__, exc)
return host_dict
all_scanned = nm.all_hosts()
logger.debug("[Phase 2] %s — nmap returned %d host(s) in results", ip, len(all_scanned))
if ip not in all_scanned:
logger.info("[Phase 2] %s — no open ports found (all closed/filtered or nmap had no results)", ip)
return host_dict
open_ports = []
for proto in nm[ip].all_protocols():
for port, info in nm[ip][proto].items():
if info["state"] == "open":
banner = (info.get("product", "") + " " + info.get("version", "")).strip()
open_ports.append({"port": port, "protocol": proto, "banner": banner})
if open_ports:
port_summary = ", ".join(
f"{p['port']}/{p['protocol']} ({p['banner'] or 'unknown'})" for p in open_ports
)
logger.info("[Phase 2] %s%d open port(s): %s", ip, len(open_ports), port_summary)
else:
logger.info("[Phase 2] %s — 0 open ports detected", ip)
host_dict["open_ports"] = open_ports
if not host_dict["mac"]:
host_dict["mac"] = nm[ip].get("addresses", {}).get("mac")
host_dict["os"] = _extract_os(nm, ip)
return host_dict
async def _nmap_port_scan(alive: dict[str, dict[str, Any]]) -> list[dict[str, Any]]:
"""
Phase 2: Per-IP service detection with bounded concurrency.
Each host is scanned independently in a thread — no inter-host timeout interference.
Up to 10 hosts scanned concurrently.
"""
if not alive:
return []
logger.info("[Phase 2] Starting per-IP port scan for %d host(s)", len(alive))
semaphore = asyncio.Semaphore(10)
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
async def _nmap_scan(target: str) -> list[dict[str, Any]]:
"""
Two-phase scan for a CIDR range.
Phase 1: Concurrent ping sweep to find alive hosts (fast, no false positives).
Phase 2: Per-IP nmap port scan with service detection (bounded concurrency, 10 at a time).
"""
logger.info("[Scan] Starting scan for %s — nmap available: %s", target, _NMAP_AVAILABLE)
if not _NMAP_AVAILABLE:
logger.warning("[Scan] nmap not available — returning mock data")
return _mock_scan(target)
try:
alive = await _ping_sweep(target)
logger.info("[Phase 1] Found %d alive host(s) in %s: %s",
len(alive), target, ", ".join(sorted(alive.keys())))
except Exception as exc:
logger.error("Phase 1 ping sweep failed: %s", exc)
raise RuntimeError(str(exc)) from exc
return await _nmap_port_scan(alive)
async def _mdns_discover(timeout: float = 4.0) -> list[dict[str, Any]]:
"""
Passive mDNS/Bonjour sweep.
Returns devices advertising on _shelly._tcp, _esphomelib._tcp, _hap._tcp, etc.
Runs for `timeout` seconds then returns what it found.
"""
if not _ZEROCONF_AVAILABLE:
return []
import ipaddress
found_services: list[tuple[str, str]] = []
def _on_change(
zeroconf: Any,
service_type: str,
name: str,
state_change: Any,
) -> None:
if state_change == ServiceStateChange.Added:
found_services.append((service_type, name))
discovered: dict[str, dict[str, Any]] = {}
try:
async with AsyncZeroconf() as azc:
browser = AsyncServiceBrowser(
azc.zeroconf, _MDNS_SERVICE_TYPES, handlers=[_on_change]
)
await asyncio.sleep(timeout)
await browser.async_cancel()
for service_type, name in found_services:
try:
info = AsyncServiceInfo(service_type, name)
await info.async_request(azc.zeroconf, 3000)
if not info.addresses:
continue
ip = str(ipaddress.IPv4Address(info.addresses[0]))
if ip in discovered:
continue
discovered[ip] = {
"ip": ip,
"hostname": info.server,
"mac": None,
"os": None,
"open_ports": (
[{"port": info.port, "protocol": "tcp", "banner": ""}]
if info.port else []
),
}
except Exception as exc:
logger.debug("mDNS resolution failed for %s: %s", name, exc)
except Exception as exc:
logger.warning("mDNS discovery error: %s", exc)
logger.info("mDNS discovery found %d device(s)", len(discovered))
return list(discovered.values())
def _mock_scan(target: str) -> list[dict[str, Any]]:
"""Return fake results for dev/test environments without nmap."""
return [
@@ -377,136 +100,73 @@ def _mock_scan(target: str) -> list[dict[str, Any]]:
async def run_scan(ranges: list[str], db: AsyncSession, run_id: str) -> None:
"""Execute scan for given CIDR ranges and populate pending_devices."""
# Avoid circular import
from sqlalchemy import select
from app.api.routes.status import broadcast_scan_update
devices_found = 0
mdns_task: asyncio.Task[list[dict[str, Any]]] | None = None
try:
# Validate all ranges are valid CIDRs before passing anything to nmap
for r in ranges:
try:
ipaddress.ip_network(r, strict=False)
except ValueError:
raise ValueError(f"Invalid CIDR range: {r!r}") from None
# Pre-fetch canvas IPs and hidden IPs once — avoids N+1 queries per host
canvas_ips_result = await db.execute(select(Node.ip).where(Node.ip.isnot(None)))
canvas_ips: set[str] = {row[0] for row in canvas_ips_result.fetchall()}
hidden_ips_result = await db.execute(
select(PendingDevice.ip).where(PendingDevice.status == "hidden")
)
hidden_ips: set[str] = {row[0] for row in hidden_ips_result.fetchall()}
# Clean up stale pending devices whose IPs are already in the canvas
if canvas_ips:
from sqlalchemy import delete as sa_delete
await db.execute(
sa_delete(PendingDevice).where(
PendingDevice.status == "pending",
PendingDevice.ip.in_(canvas_ips),
)
)
await db.commit()
# Start mDNS discovery in the background while nmap scans run
mdns_task = asyncio.create_task(_mdns_discover())
# Track IPs found by nmap so mDNS doesn't duplicate them
nmap_ips: set[str] = set()
async def _process_host(host: dict[str, Any], discovery_source: str = "arp") -> None:
nonlocal devices_found
ip = host["ip"]
# Skip canvas nodes and user-hidden devices (sets pre-fetched before loop)
if ip in canvas_ips:
logger.debug("Skipping %s — already in canvas", ip)
return
if ip in hidden_ips:
logger.debug("Skipping %s — hidden by user", ip)
return
services = fingerprint_ports(host["open_ports"])
suggested_type = suggest_node_type(host["open_ports"], host.get("mac"))
existing_result = await db.execute(
select(PendingDevice).where(
PendingDevice.ip == ip,
PendingDevice.status == "pending",
)
)
existing = existing_result.scalar_one_or_none()
if existing:
existing.mac = host.get("mac") or existing.mac
existing.hostname = host.get("hostname") or existing.hostname
existing.os = host.get("os") or existing.os
existing.services = services
existing.suggested_type = suggested_type
else:
db.add(PendingDevice(
ip=ip,
mac=host.get("mac"),
hostname=host.get("hostname"),
os=host.get("os"),
services=services,
suggested_type=suggested_type,
status="pending",
discovery_source=discovery_source,
))
devices_found += 1
await db.commit()
await broadcast_scan_update(run_id=run_id, devices_found=devices_found)
# nmap scan per CIDR — results stream in progressively
for cidr in ranges:
if _is_cancelled(run_id):
break
hosts = await _nmap_scan(cidr)
# Run nmap in a thread pool — does not block the event loop
hosts = await asyncio.to_thread(_nmap_scan, cidr)
for host in hosts:
if _is_cancelled(run_id):
break
nmap_ips.add(host["ip"])
await _process_host(host)
services = fingerprint_ports(host["open_ports"])
suggested_type = suggest_node_type(host["open_ports"], host.get("mac"))
# Update ScanRun count once after all CIDR ranges
# Update existing pending device or create a new one
existing_result = await db.execute(
select(PendingDevice).where(
PendingDevice.ip == host["ip"],
PendingDevice.status == "pending",
)
)
existing = existing_result.scalar_one_or_none()
if existing:
existing.mac = host.get("mac") or existing.mac
existing.hostname = host.get("hostname") or existing.hostname
existing.os = host.get("os") or existing.os
existing.services = services
existing.suggested_type = suggested_type
else:
device = PendingDevice(
ip=host["ip"],
mac=host.get("mac"),
hostname=host.get("hostname"),
os=host.get("os"),
services=services,
suggested_type=suggested_type,
status="pending",
)
db.add(device)
devices_found += 1
# Commit immediately so the device is visible right away
await db.commit()
# Update running count on the scan run record
run = await db.get(ScanRun, run_id)
if run:
run.devices_found = devices_found
await db.commit()
# Push WS event so the frontend refreshes pending panel
await broadcast_scan_update(run_id=run_id, devices_found=devices_found)
# Mark scan as done
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):
mdns_hosts = await mdns_task
for host in mdns_hosts:
if _is_cancelled(run_id):
break
if host["ip"] in nmap_ips:
continue # already processed with richer nmap data
await _process_host(host, discovery_source="mdns")
else:
mdns_task.cancel()
# Mark scan as done or cancelled
run = await db.get(ScanRun, run_id)
if run:
run.status = "cancelled" if _is_cancelled(run_id) else "done"
run.status = "done"
run.devices_found = devices_found
run.finished_at = datetime.now(timezone.utc)
await db.commit()
except Exception as exc:
logger.error("Scan failed: %s", exc)
if mdns_task is not None and not mdns_task.done():
mdns_task.cancel()
run = await db.get(ScanRun, run_id)
if run:
run.status = "error"
run.error = str(exc)
run.finished_at = datetime.now(timezone.utc)
await db.commit()
finally:
with _cancelled_lock:
_cancelled_runs.discard(run_id)
+2 -110
View File
@@ -2,7 +2,6 @@
import asyncio
import logging
import socket
import sys
import time
from typing import Any
@@ -19,16 +18,9 @@ async def check_node(check_method: str, target: str | None, ip: str | None) -> d
if check_method == "none":
return {"status": "online", "response_time_ms": None}
# Use only the first IP when the field contains comma-separated addresses
raw_ip = ip.split(",")[0].strip() if ip else None
host = target or raw_ip
host = target or ip
if not host:
return {"status": "unknown", "response_time_ms": None}
# Reject hostnames that look like CLI flags — defends ping/tcp invocations
# against arg-injection if a malicious admin sets target like "-O".
if host.startswith("-"):
logger.warning("Rejecting check target that starts with '-': %r", host)
return {"status": "unknown", "response_time_ms": None}
start = time.monotonic()
try:
@@ -64,37 +56,9 @@ async def check_node(check_method: str, target: str | None, ip: str | None) -> d
return {"status": "offline", "response_time_ms": None}
def _is_ipv6(host: str) -> bool:
"""True if host is a literal IPv6 address (bracketed or bare)."""
try:
socket.inet_pton(socket.AF_INET6, host.strip("[]"))
return True
except OSError:
return False
async def _ping(host: str) -> bool:
# Send 2 probes with a ~2s timeout so a single dropped packet or a slow
# device (ESPHome, IoT) doesn't flap a node offline. Success = any reply.
#
# -W flag units differ by OS:
# Linux: seconds (-W 2 = 2s)
# macOS: milliseconds (-W 2000 = 2s)
# Windows: -w in ms (-w 2000 = 2s)
#
# IPv6-only hosts (e.g. Alexa) never answer IPv4 ping, so target the right
# stack: macOS ships a separate ping6; Linux/Windows take a -6 flag.
ipv6 = _is_ipv6(host)
if sys.platform == "win32":
family = ["-6"] if ipv6 else ["-4"]
args = ["ping", *family, "-n", "2", "-w", "2000", host]
elif sys.platform == "darwin":
args = ["ping6", "-c", "2", host] if ipv6 else ["ping", "-c", "2", "-W", "2000", host]
else:
family = ["-6"] if ipv6 else []
args = ["ping", *family, "-c", "2", "-W", "2", host]
proc = await asyncio.create_subprocess_exec(
*args,
"ping", "-c", "1", "-W", "1", host,
stdout=asyncio.subprocess.DEVNULL,
stderr=asyncio.subprocess.DEVNULL,
)
@@ -118,75 +82,3 @@ async def _tcp_connect(host: str, port: int) -> bool:
return True
except (TimeoutError, OSError, socket.gaierror):
return False
# --- Per-service status checks ---
# Ports that are not HTTP/web. These get NO status check — a service here stays
# grey (unknown) rather than going red. An open TCP socket doesn't prove the
# service is healthy, and a closed one flaps red misleadingly (e.g. SSH on a
# box that simply firewalls 22). Only HTTP(S)-reachable services are checked.
_NON_HTTP_PORTS = frozenset({
22, 21, 23, 25, 465, 587, 53, 110, 143, 993, 995, 389, 636, 445, 514,
1433, 3306, 5432, 5672, 6379, 9092, 11211, 27017, 27018,
})
_HTTPS_PORTS = frozenset({443, 8443})
def _service_host(svc: dict[str, Any], host: str) -> str:
"""Bracket bare IPv6 literals for use in a URL."""
return f"[{host}]" if _is_ipv6(host) else host
async def check_service(svc: dict[str, Any], host: str | None) -> str:
"""Check a single service. Returns 'online' | 'offline' | 'unknown'.
Only HTTP(S)-reachable services get a real check (an HTTP GET). Everything
else — SSH, databases, mail, DNS, raw TCP, UDP, port-less — stays 'unknown'
so it keeps its category colour instead of flashing red. An open TCP socket
doesn't prove a non-web service is healthy, so we don't pretend it does.
"""
if not host or host.startswith("-"):
return "unknown"
if str(svc.get("protocol", "")).lower() == "udp":
return "unknown"
port = svc.get("port")
port = int(port) if isinstance(port, int) or (isinstance(port, str) and port.isdigit()) else None
# Non-HTTP ports (SSH 22, DB, mail, …) are never checked — keep them grey.
if port is not None and port in _NON_HTTP_PORTS:
return "unknown"
name = str(svc.get("service_name", "")).lower()
is_web = port is not None or "http" in name
if not is_web:
return "unknown"
try:
scheme = "https" if (
port in _HTTPS_PORTS or "https" in name or "ssl" in name or "tls" in name
) else "http"
url_host = _service_host(svc, host)
url = f"{scheme}://{url_host}" + (f":{port}" if port is not None else "")
return "online" if await _http_get(url, verify=False) else "offline"
except Exception as exc:
logger.debug("Service check failed for %s:%s (%s)", host, port, exc)
return "offline"
async def check_services(
host: str | None, services: list[dict[str, Any]], concurrency: int = 10
) -> list[dict[str, Any]]:
"""Check every service against host concurrently (bounded).
Returns a list of {port, protocol, status} dicts, one per input service.
"""
sem = asyncio.Semaphore(concurrency)
async def _one(svc: dict[str, Any]) -> dict[str, Any]:
async with sem:
status = await check_service(svc, host)
return {"port": svc.get("port"), "protocol": svc.get("protocol"), "status": status}
return await asyncio.gather(*[_one(s) for s in services]) if services else []
-374
View File
@@ -1,374 +0,0 @@
"""Zigbee2MQTT service: connects to MQTT broker and fetches the network map."""
from __future__ import annotations
import asyncio
import json
import logging
import ssl
from typing import Any
logger = logging.getLogger(__name__)
try:
import aiomqtt
except ImportError: # pragma: no cover
aiomqtt = None # type: ignore[assignment]
_NETWORKMAP_REQUEST_TOPIC = "{base_topic}/bridge/request/networkmap"
_NETWORKMAP_RESPONSE_TOPIC = "{base_topic}/bridge/response/networkmap"
_CONNECTION_TIMEOUT = 5.0 # seconds to verify broker reachability
_NETWORKMAP_TIMEOUT = 300.0 # seconds to wait for the networkmap response (large meshes can be slow)
def _sanitize_mqtt_error(exc: BaseException) -> str:
"""Return a generic, credential-free message for an MQTT error.
The raw aiomqtt/paho error string can include the broker URI with
embedded credentials (e.g. ``mqtt://user:pass@host``) or auth-related
detail that should not leak to API clients. Map known patterns to
coarse categories; default to a generic failure message. The original
exception is logged at WARNING level for operator debugging.
"""
logger.warning("MQTT error (sanitized for client): %r", exc)
raw = str(exc).lower()
if "not authoriz" in raw or "bad user" in raw or "bad username" in raw:
return "Authentication failed"
if "refused" in raw:
return "Connection refused by broker"
if "name or service not known" in raw or "getaddrinfo" in raw or "nodename nor servname" in raw:
return "Broker hostname could not be resolved"
if "ssl" in raw or "tls" in raw or "certificate" in raw:
return "TLS handshake failed"
if "timed out" in raw or "timeout" in raw:
return "Connection to broker timed out"
return "MQTT connection failed"
def _build_tls_context(insecure: bool) -> ssl.SSLContext:
"""Build an SSL context for MQTT TLS. If insecure, skip verification."""
ctx = ssl.create_default_context()
if insecure:
logger.warning(
"MQTT TLS certificate verification is DISABLED — "
"use only with self-signed brokers on trusted networks."
)
ctx.check_hostname = False
ctx.verify_mode = ssl.CERT_NONE
return ctx
def build_zigbee_properties(
ieee: str | None,
vendor: str | None,
model: str | None,
lqi: int | None,
) -> list[dict[str, Any]]:
"""Build a NodeProperty list for a Zigbee device (IEEE, Vendor, Model, LQI).
Only includes a row when the value is non-empty. Shape matches the
frontend ``NodeProperty`` type: ``{key, value, icon, visible}``.
New props default to ``visible=False`` — users opt in to showing them on
the canvas card from the right panel.
"""
props: list[dict[str, Any]] = []
if ieee:
props.append({"key": "IEEE", "value": ieee, "icon": None, "visible": False})
if vendor:
props.append({"key": "Vendor", "value": vendor, "icon": None, "visible": False})
if model:
props.append({"key": "Model", "value": model, "icon": None, "visible": False})
if lqi is not None:
props.append({"key": "LQI", "value": str(lqi), "icon": None, "visible": False})
return props
def merge_zigbee_properties(
existing: list[dict[str, Any]] | None,
new_props: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Merge fresh zigbee props into an existing property list.
For keys already present: update ``value`` but preserve the user's
``visible`` choice. New keys are appended with whatever visibility the
caller gave them (hidden by default per ``build_zigbee_properties``).
Non-zigbee custom properties are preserved untouched.
"""
out = [dict(p) for p in (existing or [])]
by_key = {p.get("key"): p for p in out}
for np in new_props:
key = np.get("key")
if key in by_key:
by_key[key]["value"] = np.get("value")
else:
out.append(dict(np))
return out
def _z2m_type_to_homelable(device_type: str) -> str:
"""Map a Z2M device type string to a homelable node type."""
mapping = {
"Coordinator": "zigbee_coordinator",
"Router": "zigbee_router",
"EndDevice": "zigbee_enddevice",
}
return mapping.get(device_type, "zigbee_enddevice")
def _node_from_z2m(raw: dict[str, Any]) -> dict[str, Any] | None:
"""Build a homelable node dict from a Z2M raw networkmap node entry."""
ieee: str = raw.get("ieeeAddr") or raw.get("ieee_address") or ""
if not ieee:
return None
device_type: str = raw.get("type") or "EndDevice"
friendly_name: str = (
raw.get("friendlyName") or raw.get("friendly_name") or ieee
)
definition: dict[str, Any] = raw.get("definition") or {}
model: str | None = (
raw.get("modelID")
or raw.get("model")
or definition.get("model")
or None
)
vendor: str | None = raw.get("vendor") or definition.get("vendor") or None
return {
"id": ieee,
"label": friendly_name,
"type": _z2m_type_to_homelable(device_type),
"ieee_address": ieee,
"friendly_name": friendly_name,
"device_type": device_type,
"model": model,
"vendor": vendor,
"lqi": None,
"parent_id": None,
}
def parse_networkmap(
payload: dict[str, Any],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
"""Parse a Z2M ``bridge/response/networkmap`` payload into node + edge lists.
Z2M raw response shape::
{
"data": {
"type": "raw",
"routes": false,
"value": {
"nodes": [{"ieeeAddr": ..., "type": "Coordinator|Router|EndDevice",
"friendlyName": ..., "definition": {"model": ..., "vendor": ...}}],
"links": [{"source": {"ieeeAddr": ...}, "target": {"ieeeAddr": ...},
"lqi": 200, "depth": 1}]
}
},
"status": "ok"
}
Older or alternate shapes may put nodes/links directly under ``data``.
Both are accepted.
"""
data: dict[str, Any] = payload.get("data") or {}
value = data.get("value")
container: dict[str, Any] = value if isinstance(value, dict) else data
raw_nodes: list[dict[str, Any]] = container.get("nodes") or []
raw_links: list[dict[str, Any]] = container.get("links") or []
if not isinstance(raw_nodes, list):
raise ValueError("Malformed networkmap: 'nodes' is not a list")
if not isinstance(raw_links, list):
raise ValueError("Malformed networkmap: 'links' is not a list")
nodes_list: list[dict[str, Any]] = []
seen_ids: set[str] = set()
coordinator_id: str | None = None
for entry in raw_nodes:
if not isinstance(entry, dict):
continue
node = _node_from_z2m(entry)
if node is None or node["id"] in seen_ids:
continue
seen_ids.add(node["id"])
nodes_list.append(node)
if node["device_type"] == "Coordinator":
coordinator_id = node["id"]
# Z2M `links` is bidirectional/mesh: every pair appears twice and routers
# carry sibling-mesh paths. Walk it only to extract LQI per device and to
# resolve which router an end device hangs off; do NOT emit edges directly
# from links. The final edge set is the strict parent→child tree built
# from parent_id below — that avoids duplicate edges and keeps the visual
# flow consistent (parent bottom → child top).
raw_edges: list[dict[str, Any]] = []
lqi_by_id: dict[str, int] = {}
for link in raw_links:
if not isinstance(link, dict):
continue
src_obj = link.get("source") or {}
tgt_obj = link.get("target") or {}
src = src_obj.get("ieeeAddr") if isinstance(src_obj, dict) else None
tgt = tgt_obj.get("ieeeAddr") if isinstance(tgt_obj, dict) else None
if not src or not tgt:
continue
if src not in seen_ids or tgt not in seen_ids:
continue
raw_edges.append({"source": src, "target": tgt})
lqi = link.get("lqi") or link.get("linkquality")
if isinstance(lqi, int) and tgt not in lqi_by_id:
lqi_by_id[tgt] = lqi
for node in nodes_list:
if node["id"] in lqi_by_id:
node["lqi"] = lqi_by_id[node["id"]]
# Build parent_id hierarchy: coordinator → routers → end devices
if coordinator_id:
router_ids = {n["id"] for n in nodes_list if n["device_type"] == "Router"}
for node in nodes_list:
if node["device_type"] == "Router":
node["parent_id"] = coordinator_id
elif node["device_type"] == "EndDevice":
parent = _find_parent_router(node["id"], router_ids, raw_edges)
node["parent_id"] = parent or coordinator_id
# Final edges = strict parent → child tree (one edge per non-coordinator)
edges_list: list[dict[str, Any]] = [
{"source": node["parent_id"], "target": node["id"]}
for node in nodes_list
if node.get("parent_id")
]
return nodes_list, edges_list
def _find_parent_router(
device_id: str,
router_ids: set[str],
edges: list[dict[str, Any]],
) -> str | None:
"""Return the first router that has a direct edge to device_id."""
for edge in edges:
src: str = edge["source"]
tgt: str = edge["target"]
if tgt == device_id and src in router_ids:
return src
if src == device_id and tgt in router_ids:
return tgt
return None
async def fetch_networkmap(
mqtt_host: str,
mqtt_port: int,
base_topic: str,
username: str | None = None,
password: str | None = None,
tls: bool = False,
tls_insecure: bool = False,
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
"""Connect to the MQTT broker, request the Z2M networkmap, and return (nodes, edges).
Raises:
TimeoutError: if the broker does not respond in time.
ConnectionError: if the broker cannot be reached.
ValueError: if the response payload is malformed.
"""
if aiomqtt is None: # pragma: no cover
raise ImportError(
"aiomqtt is required for Zigbee import. "
"Install it with: pip install aiomqtt"
)
request_topic = _NETWORKMAP_REQUEST_TOPIC.format(base_topic=base_topic)
response_topic = _NETWORKMAP_RESPONSE_TOPIC.format(base_topic=base_topic)
response_payload: dict[str, Any] = {}
tls_context = _build_tls_context(tls_insecure) if tls else None
try:
async with aiomqtt.Client(
hostname=mqtt_host,
port=mqtt_port,
username=username,
password=password,
timeout=_CONNECTION_TIMEOUT,
tls_context=tls_context,
) as client:
await client.subscribe(response_topic)
# Give the broker a brief window to register the subscription
# before we publish the request. Without this, brokers that
# race SUBACK with our PUBLISH may deliver the response before
# the subscription is active and we'd hang until timeout.
await asyncio.sleep(0.1)
await client.publish(
request_topic,
json.dumps({"type": "raw", "routes": False}),
)
async def _wait_for_response() -> None:
async for message in client.messages:
if str(message.topic) != response_topic:
continue
raw = message.payload
try:
payload_str = (
raw.decode() if isinstance(raw, bytes | bytearray) else str(raw)
)
response_payload.update(json.loads(payload_str))
except (json.JSONDecodeError, TypeError) as exc:
raise ValueError(
f"Malformed networkmap response: {exc}"
) from exc
return
await asyncio.wait_for(_wait_for_response(), timeout=_NETWORKMAP_TIMEOUT)
except aiomqtt.MqttError as exc:
raise ConnectionError(_sanitize_mqtt_error(exc)) from exc
except asyncio.TimeoutError as exc:
raise TimeoutError("Timed out waiting for networkmap response") from exc
if not response_payload:
raise ValueError("Empty networkmap response received")
return parse_networkmap(response_payload)
async def test_mqtt_connection(
mqtt_host: str,
mqtt_port: int,
username: str | None = None,
password: str | None = None,
tls: bool = False,
tls_insecure: bool = False,
) -> bool:
"""Attempt a quick MQTT connection to verify broker reachability.
Returns True on success, raises ConnectionError on failure.
"""
if aiomqtt is None: # pragma: no cover
raise ImportError("aiomqtt is required")
tls_context = _build_tls_context(tls_insecure) if tls else None
try:
async with aiomqtt.Client(
hostname=mqtt_host,
port=mqtt_port,
username=username,
password=password,
timeout=_CONNECTION_TIMEOUT,
tls_context=tls_context,
):
return True
except aiomqtt.MqttError as exc:
raise ConnectionError(_sanitize_mqtt_error(exc)) from exc
except asyncio.TimeoutError as exc:
raise TimeoutError("Connection to broker timed out") from exc
-1
View File
@@ -2,4 +2,3 @@
*.db-shm
*.db-wal
scan_config.json
homelab.db.*
+146
View File
@@ -0,0 +1,146 @@
[
{"port": 8006, "protocol": "tcp", "banner_regex": null, "service_name": "Proxmox VE", "icon": "layers", "category": "hypervisor", "suggested_node_type": "proxmox"},
{"port": 5000, "protocol": "tcp", "banner_regex": "synology|DSM", "service_name": "Synology DSM", "icon": "hard-drive", "category": "nas", "suggested_node_type": "nas"},
{"port": 5001, "protocol": "tcp", "banner_regex": null, "service_name": "Synology DSM HTTPS", "icon": "hard-drive", "category": "nas", "suggested_node_type": "nas"},
{"port": 5006, "protocol": "tcp", "banner_regex": null, "service_name": "Synology DSM Mobile", "icon": "hard-drive", "category": "nas", "suggested_node_type": "nas"},
{"port": 8080, "protocol": "tcp", "banner_regex": "QNAP|qnap|QTS", "service_name": "QNAP NAS", "icon": "hard-drive", "category": "nas", "suggested_node_type": "nas"},
{"port": 5005, "protocol": "tcp", "banner_regex": null, "service_name": "TrueNAS", "icon": "hard-drive", "category": "nas", "suggested_node_type": "nas"},
{"port": 445, "protocol": "tcp", "banner_regex": null, "service_name": "SMB / CIFS", "icon": "share-2", "category": "storage", "suggested_node_type": "nas"},
{"port": 2049, "protocol": "tcp", "banner_regex": null, "service_name": "NFS", "icon": "share-2", "category": "storage", "suggested_node_type": "nas"},
{"port": 548, "protocol": "tcp", "banner_regex": null, "service_name": "AFP (Apple Filing)", "icon": "share-2", "category": "storage", "suggested_node_type": "nas"},
{"port": 873, "protocol": "tcp", "banner_regex": null, "service_name": "rsync", "icon": "refresh-cw", "category": "storage", "suggested_node_type": "nas"},
{"port": 32400, "protocol": "tcp", "banner_regex": null, "service_name": "Plex Media Server", "icon": "play-circle", "category": "media", "suggested_node_type": "server"},
{"port": 32469, "protocol": "tcp", "banner_regex": null, "service_name": "Plex DLNA", "icon": "play-circle", "category": "media", "suggested_node_type": "server"},
{"port": 8096, "protocol": "tcp", "banner_regex": "Jellyfin", "service_name": "Jellyfin", "icon": "play-circle", "category": "media", "suggested_node_type": "server"},
{"port": 8096, "protocol": "tcp", "banner_regex": "Emby", "service_name": "Emby", "icon": "play-circle", "category": "media", "suggested_node_type": "server"},
{"port": 8096, "protocol": "tcp", "banner_regex": null, "service_name": "Jellyfin / Emby", "icon": "play-circle", "category": "media", "suggested_node_type": "server"},
{"port": 8920, "protocol": "tcp", "banner_regex": null, "service_name": "Jellyfin HTTPS", "icon": "play-circle", "category": "media", "suggested_node_type": "server"},
{"port": 8181, "protocol": "tcp", "banner_regex": null, "service_name": "Tautulli", "icon": "bar-chart", "category": "media", "suggested_node_type": "server"},
{"port": 8013, "protocol": "tcp", "banner_regex": null, "service_name": "Komga", "icon": "book-open", "category": "media", "suggested_node_type": "server"},
{"port": 1935, "protocol": "tcp", "banner_regex": null, "service_name": "RTMP (Stream)", "icon": "video", "category": "media", "suggested_node_type": "server"},
{"port": 8989, "protocol": "tcp", "banner_regex": null, "service_name": "Sonarr", "icon": "tv", "category": "media", "suggested_node_type": "server"},
{"port": 7878, "protocol": "tcp", "banner_regex": null, "service_name": "Radarr", "icon": "film", "category": "media", "suggested_node_type": "server"},
{"port": 8686, "protocol": "tcp", "banner_regex": null, "service_name": "Lidarr", "icon": "music", "category": "media", "suggested_node_type": "server"},
{"port": 9696, "protocol": "tcp", "banner_regex": null, "service_name": "Prowlarr", "icon": "search", "category": "media", "suggested_node_type": "server"},
{"port": 8787, "protocol": "tcp", "banner_regex": null, "service_name": "Readarr", "icon": "book", "category": "media", "suggested_node_type": "server"},
{"port": 6767, "protocol": "tcp", "banner_regex": null, "service_name": "Bazarr", "icon": "subtitles", "category": "media", "suggested_node_type": "server"},
{"port": 5055, "protocol": "tcp", "banner_regex": null, "service_name": "Overseerr / Jellyseerr", "icon": "search", "category": "media", "suggested_node_type": "server"},
{"port": 9117, "protocol": "tcp", "banner_regex": null, "service_name": "Jackett", "icon": "search", "category": "media", "suggested_node_type": "server"},
{"port": 6969, "protocol": "tcp", "banner_regex": null, "service_name": "Whisparr", "icon": "film", "category": "media", "suggested_node_type": "server"},
{"port": 5454, "protocol": "tcp", "banner_regex": null, "service_name": "Notifiarr", "icon": "bell", "category": "media", "suggested_node_type": "server"},
{"port": 8191, "protocol": "tcp", "banner_regex": null, "service_name": "FlareSolverr", "icon": "shield", "category": "network", "suggested_node_type": "server"},
{"port": 9091, "protocol": "tcp", "banner_regex": "Transmission", "service_name": "Transmission", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 8112, "protocol": "tcp", "banner_regex": null, "service_name": "Deluge", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 6789, "protocol": "tcp", "banner_regex": null, "service_name": "NZBGet", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 6800, "protocol": "tcp", "banner_regex": null, "service_name": "Aria2 RPC", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 51413, "protocol": "tcp", "banner_regex": null, "service_name": "Transmission BitTorrent", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 6881, "protocol": "tcp", "banner_regex": null, "service_name": "BitTorrent Peer", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 8123, "protocol": "tcp", "banner_regex": null, "service_name": "Home Assistant", "icon": "home", "category": "automation", "suggested_node_type": "iot"},
{"port": 1883, "protocol": "tcp", "banner_regex": null, "service_name": "MQTT Broker", "icon": "radio", "category": "iot", "suggested_node_type": "iot"},
{"port": 8883, "protocol": "tcp", "banner_regex": null, "service_name": "MQTT Broker TLS", "icon": "radio", "category": "iot", "suggested_node_type": "iot"},
{"port": 6052, "protocol": "tcp", "banner_regex": null, "service_name": "ESPHome", "icon": "cpu", "category": "iot", "suggested_node_type": "iot"},
{"port": 1880, "protocol": "tcp", "banner_regex": null, "service_name": "Node-RED", "icon": "git-branch", "category": "automation", "suggested_node_type": "iot"},
{"port": 8971, "protocol": "tcp", "banner_regex": null, "service_name": "Frigate NVR", "icon": "camera", "category": "nvr", "suggested_node_type": "camera"},
{"port": 10443, "protocol": "tcp", "banner_regex": null, "service_name": "Scrypted", "icon": "camera", "category": "nvr", "suggested_node_type": "camera"},
{"port": 5000, "protocol": "tcp", "banner_regex": "frigate", "service_name": "Frigate NVR", "icon": "camera", "category": "nvr", "suggested_node_type": "camera"},
{"port": 8081, "protocol": "tcp", "banner_regex": "iobroker|ioBroker", "service_name": "ioBroker", "icon": "cpu", "category": "automation", "suggested_node_type": "iot"},
{"port": 8080, "protocol": "tcp", "banner_regex": "Domoticz|domoticz", "service_name": "Domoticz", "icon": "home", "category": "automation", "suggested_node_type": "iot"},
{"port": 5683, "protocol": "udp", "banner_regex": null, "service_name": "CoAP (IoT)", "icon": "radio", "category": "iot", "suggested_node_type": "iot"},
{"port": 554, "protocol": "tcp", "banner_regex": null, "service_name": "RTSP (Camera)", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 8554, "protocol": "tcp", "banner_regex": null, "service_name": "RTSP Alt (Camera)", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 37777, "protocol": "tcp", "banner_regex": null, "service_name": "Dahua Camera SDK", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 34567, "protocol": "tcp", "banner_regex": null, "service_name": "Amcrest / Dahua Camera", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 8000, "protocol": "tcp", "banner_regex": "[Hh]ikvision|[Dd]ahua", "service_name": "IP Camera SDK", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 2020, "protocol": "tcp", "banner_regex": null, "service_name": "TP-Link Tapo Camera", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 9000, "protocol": "tcp", "banner_regex": "[Rr]eolink", "service_name": "Reolink Camera", "icon": "camera", "category": "camera", "suggested_node_type": "camera"},
{"port": 8291, "protocol": "tcp", "banner_regex": null, "service_name": "MikroTik Winbox", "icon": "router", "category": "network", "suggested_node_type": "router"},
{"port": 8880, "protocol": "tcp", "banner_regex": null, "service_name": "UniFi HTTP Portal", "icon": "wifi", "category": "network", "suggested_node_type": "ap"},
{"port": 8443, "protocol": "tcp", "banner_regex": "[Uu]ni[Ff]i", "service_name": "UniFi Controller", "icon": "wifi", "category": "network", "suggested_node_type": "ap"},
{"port": 4711, "protocol": "tcp", "banner_regex": null, "service_name": "Pi-hole API", "icon": "shield", "category": "network", "suggested_node_type": "router"},
{"port": 3000, "protocol": "tcp", "banner_regex": "[Aa]d[Gg]uard", "service_name": "AdGuard Home", "icon": "shield", "category": "network", "suggested_node_type": "router"},
{"port": 81, "protocol": "tcp", "banner_regex": null, "service_name": "Nginx Proxy Manager", "icon": "arrow-right", "category": "network", "suggested_node_type": "router"},
{"port": 23, "protocol": "tcp", "banner_regex": null, "service_name": "Telnet", "icon": "terminal", "category": "network", "suggested_node_type": "switch"},
{"port": 161, "protocol": "udp", "banner_regex": null, "service_name": "SNMP", "icon": "activity", "category": "network", "suggested_node_type": "switch"},
{"port": 8200, "protocol": "tcp", "banner_regex": null, "service_name": "HashiCorp Vault", "icon": "lock", "category": "security", "suggested_node_type": "server"},
{"port": 389, "protocol": "tcp", "banner_regex": null, "service_name": "LDAP", "icon": "users", "category": "auth", "suggested_node_type": "server"},
{"port": 636, "protocol": "tcp", "banner_regex": null, "service_name": "LDAPS", "icon": "users", "category": "auth", "suggested_node_type": "server"},
{"port": 9091, "protocol": "tcp", "banner_regex": "[Aa]uthelia", "service_name": "Authelia", "icon": "shield", "category": "security", "suggested_node_type": "server"},
{"port": 9000, "protocol": "tcp", "banner_regex": "[Aa]uthentik", "service_name": "Authentik", "icon": "shield", "category": "security", "suggested_node_type": "server"},
{"port": 8080, "protocol": "tcp", "banner_regex": "[Kk]eycloak", "service_name": "Keycloak", "icon": "shield", "category": "auth", "suggested_node_type": "server"},
{"port": 3000, "protocol": "tcp", "banner_regex": "[Gg]rafana", "service_name": "Grafana", "icon": "bar-chart-2", "category": "monitoring", "suggested_node_type": "server"},
{"port": 9090, "protocol": "tcp", "banner_regex": null, "service_name": "Prometheus", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 9093, "protocol": "tcp", "banner_regex": null, "service_name": "Alertmanager", "icon": "bell", "category": "monitoring", "suggested_node_type": "server"},
{"port": 9100, "protocol": "tcp", "banner_regex": null, "service_name": "Node Exporter", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 8086, "protocol": "tcp", "banner_regex": null, "service_name": "InfluxDB", "icon": "database", "category": "monitoring", "suggested_node_type": "server"},
{"port": 3100, "protocol": "tcp", "banner_regex": null, "service_name": "Grafana Loki", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 8428, "protocol": "tcp", "banner_regex": null, "service_name": "VictoriaMetrics", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 19999, "protocol": "tcp", "banner_regex": null, "service_name": "Netdata", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 3001, "protocol": "tcp", "banner_regex": null, "service_name": "Uptime Kuma", "icon": "heart", "category": "monitoring", "suggested_node_type": "server"},
{"port": 8581, "protocol": "tcp", "banner_regex": null, "service_name": "Uptime Kuma", "icon": "heart", "category": "monitoring", "suggested_node_type": "server"},
{"port": 10051, "protocol": "tcp", "banner_regex": null, "service_name": "Zabbix Server", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 9411, "protocol": "tcp", "banner_regex": null, "service_name": "Zipkin", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 16686, "protocol": "tcp", "banner_regex": null, "service_name": "Jaeger UI", "icon": "activity", "category": "monitoring", "suggested_node_type": "server"},
{"port": 5601, "protocol": "tcp", "banner_regex": null, "service_name": "Kibana", "icon": "bar-chart-2", "category": "monitoring", "suggested_node_type": "server"},
{"port": 9443, "protocol": "tcp", "banner_regex": "[Pp]ortainer", "service_name": "Portainer HTTPS", "icon": "box", "category": "containers", "suggested_node_type": "lxc"},
{"port": 9000, "protocol": "tcp", "banner_regex": "[Pp]ortainer", "service_name": "Portainer", "icon": "box", "category": "containers", "suggested_node_type": "lxc"},
{"port": 2375, "protocol": "tcp", "banner_regex": null, "service_name": "Docker API", "icon": "box", "category": "containers", "suggested_node_type": "server"},
{"port": 2376, "protocol": "tcp", "banner_regex": null, "service_name": "Docker API TLS", "icon": "box", "category": "containers", "suggested_node_type": "server"},
{"port": 6443, "protocol": "tcp", "banner_regex": null, "service_name": "Kubernetes API", "icon": "layers", "category": "containers", "suggested_node_type": "server"},
{"port": 3306, "protocol": "tcp", "banner_regex": null, "service_name": "MySQL / MariaDB", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 5432, "protocol": "tcp", "banner_regex": null, "service_name": "PostgreSQL", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 6379, "protocol": "tcp", "banner_regex": null, "service_name": "Redis", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 27017, "protocol": "tcp", "banner_regex": null, "service_name": "MongoDB", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 9200, "protocol": "tcp", "banner_regex": null, "service_name": "Elasticsearch", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 9300, "protocol": "tcp", "banner_regex": null, "service_name": "Elasticsearch Transport", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 5984, "protocol": "tcp", "banner_regex": null, "service_name": "CouchDB", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 1521, "protocol": "tcp", "banner_regex": null, "service_name": "Oracle DB", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 6432, "protocol": "tcp", "banner_regex": null, "service_name": "PgBouncer", "icon": "database", "category": "database", "suggested_node_type": "server"},
{"port": 22, "protocol": "tcp", "banner_regex": null, "service_name": "SSH", "icon": "terminal", "category": "remote", "suggested_node_type": "server"},
{"port": 21, "protocol": "tcp", "banner_regex": null, "service_name": "FTP", "icon": "upload", "category": "storage", "suggested_node_type": "server"},
{"port": 25, "protocol": "tcp", "banner_regex": null, "service_name": "SMTP", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 110, "protocol": "tcp", "banner_regex": null, "service_name": "POP3", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 143, "protocol": "tcp", "banner_regex": null, "service_name": "IMAP", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 465, "protocol": "tcp", "banner_regex": null, "service_name": "SMTPS", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 587, "protocol": "tcp", "banner_regex": null, "service_name": "SMTP Submission", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 993, "protocol": "tcp", "banner_regex": null, "service_name": "IMAPS", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 995, "protocol": "tcp", "banner_regex": null, "service_name": "POP3S", "icon": "mail", "category": "mail", "suggested_node_type": "server"},
{"port": 3389, "protocol": "tcp", "banner_regex": null, "service_name": "RDP", "icon": "monitor", "category": "remote", "suggested_node_type": "server"},
{"port": 5900, "protocol": "tcp", "banner_regex": null, "service_name": "VNC", "icon": "monitor", "category": "remote", "suggested_node_type": "server"},
{"port": 5800, "protocol": "tcp", "banner_regex": null, "service_name": "VNC (HTTP)", "icon": "monitor", "category": "remote", "suggested_node_type": "server"},
{"port": 8888, "protocol": "tcp", "banner_regex": null, "service_name": "Jupyter Notebook", "icon": "code", "category": "dev", "suggested_node_type": "server"},
{"port": 3000, "protocol": "tcp", "banner_regex": "[Gg]itea", "service_name": "Gitea", "icon": "git-branch", "category": "dev", "suggested_node_type": "server"},
{"port": 80, "protocol": "tcp", "banner_regex": null, "service_name": "HTTP", "icon": "globe", "category": "web", "suggested_node_type": "server"},
{"port": 443, "protocol": "tcp", "banner_regex": null, "service_name": "HTTPS", "icon": "lock", "category": "web", "suggested_node_type": "server"},
{"port": 8080, "protocol": "tcp", "banner_regex": null, "service_name": "HTTP Alt", "icon": "globe", "category": "web", "suggested_node_type": "server"},
{"port": 8443, "protocol": "tcp", "banner_regex": null, "service_name": "HTTPS Alt", "icon": "lock", "category": "web", "suggested_node_type": "server"},
{"port": 8008, "protocol": "tcp", "banner_regex": null, "service_name": "HTTP Alt", "icon": "globe", "category": "web", "suggested_node_type": "server"},
{"port": 3000, "protocol": "tcp", "banner_regex": null, "service_name": "Web service", "icon": "globe", "category": "web", "suggested_node_type": "server"},
{"port": 9091, "protocol": "tcp", "banner_regex": null, "service_name": "Transmission", "icon": "download", "category": "download", "suggested_node_type": "server"},
{"port": 9000, "protocol": "tcp", "banner_regex": null, "service_name": "Web service", "icon": "globe", "category": "web", "suggested_node_type": "server"},
{"port": 9443, "protocol": "tcp", "banner_regex": null, "service_name": "HTTPS Alt", "icon": "lock", "category": "web", "suggested_node_type": "server"},
{"port": 5000, "protocol": "tcp", "banner_regex": null, "service_name": "Web service", "icon": "globe", "category": "web", "suggested_node_type": "server"},
{"port": 8448, "protocol": "tcp", "banner_regex": null, "service_name": "Matrix (Synapse)", "icon": "message-square", "category": "communication", "suggested_node_type": "server"},
{"port": 64738, "protocol": "tcp", "banner_regex": null, "service_name": "Mumble", "icon": "mic", "category": "communication", "suggested_node_type": "server"},
{"port": 25565, "protocol": "tcp", "banner_regex": null, "service_name": "Minecraft Server", "icon": "cpu", "category": "gaming", "suggested_node_type": "server"},
{"port": 51820, "protocol": "udp", "banner_regex": null, "service_name": "WireGuard", "icon": "shield", "category": "vpn", "suggested_node_type": "router"},
{"port": 1194, "protocol": "udp", "banner_regex": null, "service_name": "OpenVPN", "icon": "shield", "category": "vpn", "suggested_node_type": "router"},
{"port": 500, "protocol": "udp", "banner_regex": null, "service_name": "IPsec IKE", "icon": "shield", "category": "vpn", "suggested_node_type": "router"},
{"port": 53, "protocol": "udp", "banner_regex": null, "service_name": "DNS", "icon": "search", "category": "network", "suggested_node_type": "router"},
{"port": 67, "protocol": "udp", "banner_regex": null, "service_name": "DHCP", "icon": "wifi", "category": "network", "suggested_node_type": "router"}
]
-2
View File
@@ -25,8 +25,6 @@ addopts = "--tb=short -q"
[tool.coverage.run]
source = ["app"]
omit = ["*/migrations/*", "*/tests/*"]
concurrency = ["thread"]
core = "sysmon"
[tool.coverage.report]
skip_empty = true
+5 -6
View File
@@ -7,20 +7,19 @@ alembic==1.13.3
pydantic==2.9.2
pydantic-settings==2.5.2
python-jose[cryptography]==3.5.0
bcrypt==4.2.1
python-multipart==0.0.31
passlib[bcrypt]==1.7.4
bcrypt==4.0.1
python-multipart==0.0.22
apscheduler==3.10.4
python-nmap==0.7.1
pyyaml==6.0.2
types-PyYAML==6.0.12.20240917
websockets==13.1
httpx==0.27.2
zeroconf==0.149.12
aiomqtt==2.3.0
# Dev
ruff==0.6.9
mypy==1.11.2
pytest==9.0.3
pytest-asyncio==1.3.0
pytest==8.3.3
pytest-asyncio==0.24.0
pytest-cov==5.0.0
+5 -3
View File
@@ -1,11 +1,13 @@
"""Generate a bcrypt password hash for the AUTH_PASSWORD_HASH env var."""
"""Generate a bcrypt password hash for config.yml."""
import sys
import bcrypt
from passlib.context import CryptContext
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
if len(sys.argv) < 2:
print("Usage: python scripts/hash_password.py <password>")
sys.exit(1)
password = sys.argv[1]
print(bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8"))
print(pwd_context.hash(password))
+4 -2
View File
@@ -5,21 +5,23 @@ os.environ.setdefault("SECRET_KEY", "test-only-secret-key-not-for-production")
import pytest
from httpx import ASGITransport, AsyncClient
from passlib.context import CryptContext
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from app.core.security import hash_password
from app.db.database import Base, get_db
from app.main import app
TEST_DB_URL = "sqlite+aiosqlite:///:memory:"
_pwd_ctx = CryptContext(schemes=["bcrypt"], deprecated="auto")
@pytest.fixture(autouse=True, scope="session")
def test_credentials():
"""Configure test auth credentials directly on settings."""
from app.core.config import settings
settings.auth_username = "admin"
settings.auth_password_hash = hash_password("admin")
settings.auth_password_hash = _pwd_ctx.hash("admin")
@pytest.fixture
-80
View File
@@ -56,83 +56,3 @@ async def test_service_key_disabled_when_not_configured(client: AsyncClient):
settings.mcp_service_key = ""
res = await client.get("/api/v1/nodes", headers={"X-MCP-Service-Key": "any-key"})
assert res.status_code == 401
async def test_login_with_malformed_hash_returns_401_not_500(client: AsyncClient):
"""Malformed hash (e.g. $ stripped by shell) must not crash with 500."""
from app.core.config import settings
original = settings.auth_password_hash
settings.auth_password_hash = "2b12RtMbyw17l4N5UGzeXMNAWu" # $ signs stripped
try:
res = await client.post("/api/v1/auth/login", json={"username": "admin", "password": "admin"})
assert res.status_code == 401
finally:
settings.auth_password_hash = original
# --- JWT-level cases ---
async def test_expired_token_rejected(client: AsyncClient):
"""A JWT whose `exp` is in the past must be refused."""
from datetime import datetime, timedelta, timezone
from jose import jwt
from app.core.config import settings
payload = {
"sub": "admin",
"exp": datetime.now(timezone.utc) - timedelta(minutes=1),
}
token = jwt.encode(payload, settings.secret_key, algorithm=settings.algorithm)
res = await client.get("/api/v1/nodes", headers={"Authorization": f"Bearer {token}"})
assert res.status_code == 401
async def test_malformed_token_rejected(client: AsyncClient):
res = await client.get("/api/v1/nodes", headers={"Authorization": "Bearer not-a-jwt"})
assert res.status_code == 401
async def test_token_signed_with_wrong_secret_rejected(client: AsyncClient):
"""A token signed with a different key must not be accepted."""
from datetime import datetime, timedelta, timezone
from jose import jwt
from app.core.config import settings
payload = {
"sub": "admin",
"exp": datetime.now(timezone.utc) + timedelta(minutes=5),
}
forged = jwt.encode(payload, "different-secret", algorithm=settings.algorithm)
res = await client.get("/api/v1/nodes", headers={"Authorization": f"Bearer {forged}"})
assert res.status_code == 401
async def test_missing_authorization_header_rejected(client: AsyncClient):
res = await client.get("/api/v1/nodes")
assert res.status_code == 401
async def test_empty_password_does_not_pass_when_hash_empty(client: AsyncClient):
"""No credentials configured server-side must not authenticate an empty password."""
from app.core.config import settings
original_hash = settings.auth_password_hash
settings.auth_password_hash = ""
try:
res = await client.post("/api/v1/auth/login", json={"username": "admin", "password": ""})
assert res.status_code == 401
finally:
settings.auth_password_hash = original_hash
# --- Password helper ---
def test_verify_password_handles_empty_inputs():
"""verify_password must be safe against empty plain / empty hash without raising."""
from app.core.security import hash_password, verify_password
h = hash_password("hunter2")
assert verify_password("hunter2", h) is True
assert verify_password("", h) is False
assert verify_password("hunter2", "") is False
assert verify_password("", "") is False
-377
View File
@@ -104,25 +104,6 @@ async def test_save_canvas_persists_custom_colors(client: AsyncClient, headers:
assert canvas["nodes"][0]["custom_colors"] == {"border": "#ff0000", "icon": "#00ff00"}
async def test_save_canvas_persists_zone_label_position_and_text_size(client: AsyncClient, headers: dict):
"""label_position and text_size are stored in custom_colors and returned unchanged."""
n1 = node_payload(custom_colors={
"border": "#00d4ff",
"border_style": "solid",
"border_width": 3,
"label_position": "outside",
"text_size": 16,
"text_color": "#e6edf3",
})
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
cc = canvas["nodes"][0]["custom_colors"]
assert cc["label_position"] == "outside"
assert cc["text_size"] == 16
assert cc["border_width"] == 3
async def test_save_canvas_persists_edge_custom_color_and_path_style(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
@@ -199,24 +180,6 @@ async def test_save_canvas_show_hardware_defaults_false(client: AsyncClient, hea
assert canvas["nodes"][0]["show_hardware"] is False
# Regression (#184): show_port_numbers was dropped by the save schema, so the
# toggle reset on every reload.
async def test_save_canvas_persists_show_port_numbers(client: AsyncClient, headers: dict):
n1 = node_payload(show_port_numbers=True)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["show_port_numbers"] is True
async def test_save_canvas_show_port_numbers_defaults_false(client: AsyncClient, headers: dict):
n1 = node_payload()
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["show_port_numbers"] is False
async def test_save_canvas_hardware_fields_cleared_on_update(client: AsyncClient, headers: dict):
n1 = node_payload(cpu_count=8, ram_gb=32.0)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
@@ -274,343 +237,3 @@ async def test_save_canvas_dimensions_cleared_when_null(client: AsyncClient, hea
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["width"] is None
assert canvas["nodes"][0]["height"] is None
# ── properties ────────────────────────────────────────────────────────────────
async def test_save_canvas_properties_default_empty(client: AsyncClient, headers: dict):
n1 = node_payload()
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["properties"] == []
async def test_save_canvas_persists_properties(client: AsyncClient, headers: dict):
props = [
{"key": "RAM", "value": "32 GB", "icon": "MemoryStick", "visible": True},
{"key": "CPU", "value": "Intel i9", "icon": "Cpu", "visible": False},
]
n1 = node_payload(properties=props)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
returned = canvas["nodes"][0]["properties"]
assert len(returned) == 2
assert returned[0] == {"key": "RAM", "value": "32 GB", "icon": "MemoryStick", "visible": True}
assert returned[1] == {"key": "CPU", "value": "Intel i9", "icon": "Cpu", "visible": False}
async def test_save_canvas_properties_updated_on_second_save(client: AsyncClient, headers: dict):
n1 = node_payload(properties=[{"key": "RAM", "value": "16 GB", "icon": None, "visible": True}])
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
n1_updated = {**n1, "properties": [
{"key": "RAM", "value": "64 GB", "icon": "MemoryStick", "visible": True},
{"key": "Disk", "value": "2 TB", "icon": "HardDrive", "visible": True},
]}
await client.post("/api/v1/canvas/save", json={"nodes": [n1_updated], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
props = canvas["nodes"][0]["properties"]
assert len(props) == 2
assert props[0]["value"] == "64 GB"
assert props[1]["key"] == "Disk"
async def test_save_canvas_properties_with_null_icon(client: AsyncClient, headers: dict):
props = [{"key": "Note", "value": "custom rack", "icon": None, "visible": True}]
n1 = node_payload(properties=props)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["properties"][0]["icon"] is None
async def test_save_canvas_properties_cleared_to_empty(client: AsyncClient, headers: dict):
n1 = node_payload(properties=[{"key": "RAM", "value": "32 GB", "icon": None, "visible": True}])
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
n1_cleared = {**n1, "properties": []}
await client.post("/api/v1/canvas/save", json={"nodes": [n1_cleared], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["properties"] == []
# ── edge waypoints & handles ──────────────────────────────────────────────────
async def test_save_canvas_edge_waypoints_default_null(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"])
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["edges"][0]["waypoints"] is None
async def test_save_canvas_persists_waypoints_on_edge(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
waypoints = [{"x": 100.0, "y": 200.0}, {"x": 300.0, "y": 150.0}]
e1 = edge_payload(n1["id"], n2["id"], waypoints=waypoints)
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
returned = canvas["edges"][0]["waypoints"]
assert returned == [{"x": 100.0, "y": 200.0}, {"x": 300.0, "y": 150.0}]
async def test_save_canvas_waypoints_updated_on_second_save(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"], waypoints=[{"x": 10.0, "y": 20.0}])
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
e1_updated = {**e1, "waypoints": [{"x": 50.0, "y": 60.0}, {"x": 70.0, "y": 80.0}]}
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1_updated], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["edges"][0]["waypoints"] == [{"x": 50.0, "y": 60.0}, {"x": 70.0, "y": 80.0}]
async def test_save_canvas_persists_edge_handles(client: AsyncClient, headers: dict):
n1 = node_payload(bottom_handles=3)
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"], source_handle="bottom-1", target_handle="top")
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
edge = canvas["edges"][0]
assert edge["source_handle"] == "bottom-1"
assert edge["target_handle"] == "top"
async def test_save_canvas_persists_animated_edge(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"], animated="snake")
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["edges"][0]["animated"] == "snake"
async def test_save_canvas_persists_animated_basic(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"], animated="basic")
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["edges"][0]["animated"] == "basic"
# ── node fields ───────────────────────────────────────────────────────────────
async def test_save_canvas_persists_all_node_fields(client: AsyncClient, headers: dict):
n1 = node_payload(
type="server",
label="Main Server",
hostname="server.local",
ip="192.168.1.10",
mac="aa:bb:cc:dd:ee:ff",
os="Ubuntu 22.04",
status="online",
check_method="http",
check_target="http://192.168.1.10",
services=[{"name": "nginx", "port": 80}],
notes="Primary web server",
pos_x=150.0,
pos_y=250.0,
bottom_handles=2,
)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
node = canvas["nodes"][0]
assert node["hostname"] == "server.local"
assert node["ip"] == "192.168.1.10"
assert node["mac"] == "aa:bb:cc:dd:ee:ff"
assert node["os"] == "Ubuntu 22.04"
assert node["status"] == "online"
assert node["check_method"] == "http"
assert node["check_target"] == "http://192.168.1.10"
assert node["services"] == [{"name": "nginx", "port": 80}]
assert node["notes"] == "Primary web server"
assert node["pos_x"] == 150.0
assert node["pos_y"] == 250.0
assert node["bottom_handles"] == 2
async def test_save_canvas_persists_bottom_handles(client: AsyncClient, headers: dict):
n1 = node_payload(bottom_handles=4)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["bottom_handles"] == 4
async def test_save_canvas_bottom_handles_defaults_one(client: AsyncClient, headers: dict):
n1 = node_payload()
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["bottom_handles"] == 1
async def test_save_canvas_persists_services_and_notes(client: AsyncClient, headers: dict):
services = [{"name": "ssh", "port": 22}, {"name": "http", "port": 80}]
n1 = node_payload(services=services, notes="My NAS device")
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
node = canvas["nodes"][0]
assert node["services"] == services
assert node["notes"] == "My NAS device"
async def test_save_canvas_persists_service_paths(client: AsyncClient, headers: dict):
services = [{"service_name": "Grafana", "protocol": "tcp", "port": 3000, "path": "/login"}]
n1 = node_payload(ip="192.168.1.50:8080", services=services)
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"][0]["services"] == services
async def test_save_canvas_persists_check_fields(client: AsyncClient, headers: dict):
n1 = node_payload(check_method="ping", check_target="192.168.1.1")
await client.post("/api/v1/canvas/save", json={"nodes": [n1], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
node = canvas["nodes"][0]
assert node["check_method"] == "ping"
assert node["check_target"] == "192.168.1.1"
# ── parent/child nodes ────────────────────────────────────────────────────────
async def test_save_canvas_persists_parent_child_nodes(client: AsyncClient, headers: dict):
parent = node_payload(type="proxmox", label="PVE Host")
child = node_payload(type="vm", label="VM-100", parent_id=parent["id"])
await client.post("/api/v1/canvas/save", json={"nodes": [parent, child], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
node_map = {n["id"]: n for n in canvas["nodes"]}
assert node_map[child["id"]]["parent_id"] == parent["id"]
assert node_map[parent["id"]]["parent_id"] is None
async def test_save_canvas_child_removed_with_parent(client: AsyncClient, headers: dict):
parent = node_payload(type="proxmox", label="PVE Host")
child = node_payload(type="lxc", label="LXC-101", parent_id=parent["id"])
await client.post("/api/v1/canvas/save", json={"nodes": [parent, child], "edges": [], "viewport": {}}, headers=headers)
# Remove both parent and child
await client.post("/api/v1/canvas/save", json={"nodes": [], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["nodes"] == []
# ── groupRect / group node ────────────────────────────────────────────────────
async def test_save_canvas_persists_group_node(client: AsyncClient, headers: dict):
group = node_payload(type="group", label="Network Zone", width=400.0, height=300.0)
member = node_payload(type="server", label="Member", parent_id=group["id"])
await client.post("/api/v1/canvas/save", json={"nodes": [group, member], "edges": [], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
node_map = {n["id"]: n for n in canvas["nodes"]}
assert node_map[group["id"]]["type"] == "group"
assert node_map[group["id"]]["width"] == 400.0
assert node_map[group["id"]]["height"] == 300.0
assert node_map[member["id"]]["parent_id"] == group["id"]
# ── viewport ──────────────────────────────────────────────────────────────────
async def test_load_canvas_returns_default_viewport_when_no_state(client: AsyncClient, headers: dict):
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["viewport"] == {"x": 0, "y": 0, "zoom": 1}
async def test_save_canvas_updates_existing_canvas_state(client: AsyncClient, headers: dict):
"""Second save updates the existing CanvasState row (exercises the state.viewport branch)."""
await client.post("/api/v1/canvas/save", json={"nodes": [], "edges": [], "viewport": {"x": 1, "y": 2, "zoom": 1}}, headers=headers)
await client.post("/api/v1/canvas/save", json={"nodes": [], "edges": [], "viewport": {"x": 99, "y": 88, "zoom": 0.75}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["viewport"] == {"x": 99, "y": 88, "zoom": 0.75}
# ── edge types ────────────────────────────────────────────────────────────────
async def test_save_canvas_persists_edge_type_vlan(client: AsyncClient, headers: dict):
n1 = node_payload()
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"], type="vlan", vlan_id=10, label="VLAN 10")
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
edge = canvas["edges"][0]
assert edge["type"] == "vlan"
assert edge["vlan_id"] == 10
assert edge["label"] == "VLAN 10"
async def test_save_canvas_edge_update_existing(client: AsyncClient, headers: dict):
"""Second save updates an existing edge (exercises the db_edge branch)."""
n1 = node_payload()
n2 = node_payload()
e1 = edge_payload(n1["id"], n2["id"], label="original")
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1], "viewport": {}}, headers=headers)
e1_updated = {**e1, "label": "updated", "custom_color": "#ff0000"}
await client.post("/api/v1/canvas/save", json={"nodes": [n1, n2], "edges": [e1_updated], "viewport": {}}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
edge = canvas["edges"][0]
assert edge["label"] == "updated"
assert edge["custom_color"] == "#ff0000"
# ── custom_style ──────────────────────────────────────────────────────────────
async def test_save_and_load_custom_style(client: AsyncClient, headers: dict):
custom_style = {
"nodes": {
"server": {"borderColor": "#ff0000", "borderOpacity": 0.8, "bgColor": "#000000", "bgOpacity": 1, "iconColor": "#ff0000", "iconOpacity": 1, "width": 200, "height": 80},
},
"edges": {
"ethernet": {"color": "#00ff00", "opacity": 1, "pathStyle": "bezier", "animated": "none"},
},
}
payload = {"nodes": [], "edges": [], "viewport": {"theme_id": "custom"}, "custom_style": custom_style}
res = await client.post("/api/v1/canvas/save", json=payload, headers=headers)
assert res.status_code == 200
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert canvas["custom_style"] is not None
assert canvas["custom_style"]["nodes"]["server"]["borderColor"] == "#ff0000"
assert canvas["custom_style"]["edges"]["ethernet"]["color"] == "#00ff00"
async def test_load_canvas_custom_style_null_by_default(client: AsyncClient, headers: dict):
res = await client.get("/api/v1/canvas", headers=headers)
assert res.status_code == 200
assert res.json()["custom_style"] is None
async def test_save_canvas_custom_style_overwrite(client: AsyncClient, headers: dict):
style_v1 = {"nodes": {"server": {"borderColor": "#aabbcc", "borderOpacity": 1, "bgColor": "#000000", "bgOpacity": 1, "iconColor": "#aabbcc", "iconOpacity": 1, "width": 0, "height": 0}}, "edges": {}}
style_v2 = {"nodes": {"proxmox": {"borderColor": "#ff6e00", "borderOpacity": 1, "bgColor": "#111111", "bgOpacity": 1, "iconColor": "#ff6e00", "iconOpacity": 1, "width": 0, "height": 0}}, "edges": {}}
await client.post("/api/v1/canvas/save", json={"nodes": [], "edges": [], "viewport": {}, "custom_style": style_v1}, headers=headers)
await client.post("/api/v1/canvas/save", json={"nodes": [], "edges": [], "viewport": {}, "custom_style": style_v2}, headers=headers)
canvas = (await client.get("/api/v1/canvas", headers=headers)).json()
assert "proxmox" in canvas["custom_style"]["nodes"]
assert "server" not in canvas["custom_style"]["nodes"]
-57
View File
@@ -1,57 +0,0 @@
"""
Tests for automatic DB backup before migrations.
"""
import os
os.environ.setdefault("SECRET_KEY", "test-only-secret-key-not-for-production")
from pathlib import Path
from unittest.mock import patch
import pytest
from app.db.database import _backup_db
@pytest.fixture()
def tmp_db(tmp_path: Path):
db = tmp_path / "homelab.db"
db.write_bytes(b"SQLite placeholder")
return db
def test_backup_created_when_db_exists(tmp_db: Path):
with patch("app.db.database.settings") as mock_settings, \
patch("app.db.database.APP_VERSION", "1.9"):
mock_settings.sqlite_path = str(tmp_db)
_backup_db()
backup = tmp_db.parent / "homelab.db.back-1.9"
assert backup.exists()
assert backup.read_bytes() == b"SQLite placeholder"
def test_backup_skipped_when_db_missing(tmp_path: Path):
with patch("app.db.database.settings") as mock_settings, \
patch("app.db.database.APP_VERSION", "1.9"):
mock_settings.sqlite_path = str(tmp_path / "nonexistent.db")
_backup_db()
assert not any(tmp_path.glob("*.back-*"))
def test_backup_idempotent_second_call_no_overwrite(tmp_db: Path):
with patch("app.db.database.settings") as mock_settings, \
patch("app.db.database.APP_VERSION", "1.9"):
mock_settings.sqlite_path = str(tmp_db)
_backup_db()
backup = tmp_db.parent / "homelab.db.back-1.9"
backup.write_bytes(b"original backup")
_backup_db()
assert backup.read_bytes() == b"original backup"
def test_backup_version_in_filename(tmp_db: Path):
with patch("app.db.database.settings") as mock_settings, \
patch("app.db.database.APP_VERSION", "2.0"):
mock_settings.sqlite_path = str(tmp_db)
_backup_db()
assert (tmp_db.parent / "homelab.db.back-2.0").exists()
-165
View File
@@ -1,165 +0,0 @@
import uuid
import pytest
from httpx import AsyncClient
@pytest.fixture
async def headers(client: AsyncClient):
res = await client.post("/api/v1/auth/login", json={"username": "admin", "password": "admin"})
return {"Authorization": f"Bearer {res.json()['access_token']}"}
def node_payload(**kwargs):
return {"id": str(uuid.uuid4()), "type": "server", "label": "N", "status": "unknown", "pos_x": 0, "pos_y": 0, **kwargs}
def edge_payload(src, tgt, **kwargs):
return {"id": str(uuid.uuid4()), "source": src, "target": tgt, "type": "ethernet", **kwargs}
async def _create(client: AsyncClient, headers: dict, **body) -> dict:
res = await client.post("/api/v1/designs", json={"name": "D", **body}, headers=headers)
assert res.status_code == 201, res.text
return res.json()
# ── auth ──────────────────────────────────────────────────────────────────────
async def test_list_designs_requires_auth(client: AsyncClient):
res = await client.get("/api/v1/designs")
assert res.status_code == 401
async def test_create_design_requires_auth(client: AsyncClient):
res = await client.post("/api/v1/designs", json={"name": "X"})
assert res.status_code == 401
# ── list / create ─────────────────────────────────────────────────────────────
async def test_list_designs_empty(client: AsyncClient, headers: dict):
res = await client.get("/api/v1/designs", headers=headers)
assert res.status_code == 200
assert res.json() == []
async def test_create_design_defaults(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Workshop")
assert design["name"] == "Workshop"
assert design["design_type"] == "network"
assert design["icon"] == "dashboard"
assert "id" in design and design["id"]
async def test_create_design_explicit_type(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Net", design_type="network")
assert design["design_type"] == "network"
async def test_create_design_with_custom_icon(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Power", icon="zap")
assert design["icon"] == "zap"
async def test_update_design_changes_icon(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="D", icon="dashboard")
res = await client.put(f"/api/v1/designs/{design['id']}", json={"icon": "server"}, headers=headers)
assert res.status_code == 200
assert res.json()["icon"] == "server"
# Name left untouched when only icon is sent.
assert res.json()["name"] == "D"
async def test_update_design_name_and_icon_together(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Old", icon="dashboard")
res = await client.put(
f"/api/v1/designs/{design['id']}", json={"name": "New", "icon": "network"}, headers=headers,
)
assert res.status_code == 200
body = res.json()
assert body["name"] == "New"
assert body["icon"] == "network"
async def test_create_design_creates_empty_canvas_state(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Has Canvas")
# Loading the new design returns an (empty) canvas without falling back to another design.
res = await client.get("/api/v1/canvas", params={"design_id": design["id"]}, headers=headers)
assert res.status_code == 200
body = res.json()
assert body["nodes"] == []
assert body["edges"] == []
async def test_list_returns_created_designs_ordered(client: AsyncClient, headers: dict):
a = await _create(client, headers, name="First")
b = await _create(client, headers, name="Second")
listed = (await client.get("/api/v1/designs", headers=headers)).json()
ids = [d["id"] for d in listed]
assert ids == [a["id"], b["id"]]
# ── update ────────────────────────────────────────────────────────────────────
async def test_update_design_renames(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Old Name")
res = await client.put(f"/api/v1/designs/{design['id']}", json={"name": "New Name"}, headers=headers)
assert res.status_code == 200
assert res.json()["name"] == "New Name"
async def test_update_design_missing_returns_404(client: AsyncClient, headers: dict):
res = await client.put(f"/api/v1/designs/{uuid.uuid4()}", json={"name": "X"}, headers=headers)
assert res.status_code == 404
# ── delete ────────────────────────────────────────────────────────────────────
async def test_delete_last_design_blocked(client: AsyncClient, headers: dict):
design = await _create(client, headers, name="Only One")
res = await client.delete(f"/api/v1/designs/{design['id']}", headers=headers)
assert res.status_code == 400
async def test_delete_design_missing_returns_404(client: AsyncClient, headers: dict):
# Need >1 design so we get past nothing; 404 path is checked before the count guard.
await _create(client, headers, name="Keep")
res = await client.delete(f"/api/v1/designs/{uuid.uuid4()}", headers=headers)
assert res.status_code == 404
async def test_delete_design_removes_its_nodes_edges_and_canvas(client: AsyncClient, headers: dict):
keep = await _create(client, headers, name="Keep")
victim = await _create(client, headers, name="Victim")
# Populate the victim design with nodes + an edge via canvas save.
n1 = node_payload(label="A")
n2 = node_payload(label="B")
e1 = edge_payload(n1["id"], n2["id"])
save = await client.post(
"/api/v1/canvas/save",
json={"nodes": [n1, n2], "edges": [e1], "viewport": {}, "design_id": victim["id"]},
headers=headers,
)
assert save.status_code == 200
# Populate the kept design too, to prove scoping.
k1 = node_payload(label="K")
await client.post(
"/api/v1/canvas/save",
json={"nodes": [k1], "edges": [], "viewport": {}, "design_id": keep["id"]},
headers=headers,
)
res = await client.delete(f"/api/v1/designs/{victim['id']}", headers=headers)
assert res.status_code == 204
# Victim gone from list.
listed = (await client.get("/api/v1/designs", headers=headers)).json()
assert [d["id"] for d in listed] == [keep["id"]]
# Kept design's node survives untouched.
kept_canvas = (await client.get("/api/v1/canvas", params={"design_id": keep["id"]}, headers=headers)).json()
assert len(kept_canvas["nodes"]) == 1
assert kept_canvas["nodes"][0]["label"] == "K"
-42
View File
@@ -131,45 +131,3 @@ def test_suggest_node_type_camera_from_signature():
]):
result = suggest_node_type([{"port": 554, "protocol": "tcp"}])
assert result == "camera"
# ── IoT detection ─────────────────────────────────────────────────────────────
def test_suggest_node_type_iot_from_mqtt_port():
result = suggest_node_type([{"port": 1883, "protocol": "tcp"}])
assert result == "iot"
def test_suggest_node_type_iot_from_coap_port():
result = suggest_node_type([{"port": 5683, "protocol": "tcp"}])
assert result == "iot"
def test_suggest_node_type_iot_from_esphome_port():
result = suggest_node_type([{"port": 6052, "protocol": "tcp"}])
assert result == "iot"
def test_suggest_node_type_shelly_mac_overrides_http_port():
# Shelly exposes port 80 (would suggest "server") but MAC identifies it as IoT
result = suggest_node_type([{"port": 80, "protocol": "tcp"}], mac="34:94:54:aa:bb:cc")
assert result == "iot"
def test_suggest_node_type_espressif_mac_returns_iot():
result = suggest_node_type([], mac="a0:20:a6:11:22:33")
assert result == "iot"
def test_suggest_node_type_tuya_mac_returns_iot():
result = suggest_node_type([{"port": 80, "protocol": "tcp"}], mac="d8:f1:5b:aa:bb:cc")
assert result == "iot"
def test_suggest_node_type_iot_wins_over_server_when_mqtt_present():
# MQTT port + HTTP port → iot wins (iot is higher priority than server now)
result = suggest_node_type([
{"port": 80, "protocol": "tcp"},
{"port": 1883, "protocol": "tcp"},
])
assert result == "iot"
-229
View File
@@ -1,229 +0,0 @@
"""
Tests for the /api/v1/liveview read-only canvas endpoint.
The endpoint is:
- Disabled by default (LIVEVIEW_KEY not set) → 403
- Returns 403 for missing or wrong key even when enabled
- Returns canvas data for a valid key (no JWT required)
"""
import pytest
from httpx import AsyncClient
from app.core.config import settings
@pytest.fixture(autouse=True)
def reset_liveview_key():
"""Restore liveview_key after each test so tests are isolated."""
original = settings.liveview_key
yield
settings.liveview_key = original
# ── Disabled (no key configured) ─────────────────────────────────────────────
@pytest.mark.asyncio
async def test_liveview_disabled_by_default(client: AsyncClient):
settings.liveview_key = None
res = await client.get("/api/v1/liveview?key=anything")
assert res.status_code == 403
assert res.json()["detail"] == "Live view is disabled"
@pytest.mark.asyncio
async def test_liveview_disabled_when_key_empty(client: AsyncClient):
settings.liveview_key = ""
res = await client.get("/api/v1/liveview?key=anything")
assert res.status_code == 403
assert res.json()["detail"] == "Live view is disabled"
# ── Enabled but wrong / missing key ──────────────────────────────────────────
@pytest.mark.asyncio
async def test_liveview_wrong_key(client: AsyncClient):
settings.liveview_key = "correct-secret"
res = await client.get("/api/v1/liveview?key=wrong-key")
assert res.status_code == 403
assert res.json()["detail"] == "Invalid live view key"
@pytest.mark.asyncio
async def test_liveview_missing_key_param(client: AsyncClient):
settings.liveview_key = "correct-secret"
res = await client.get("/api/v1/liveview")
assert res.status_code == 403
assert res.json()["detail"] == "Invalid live view key"
# ── Valid key — no JWT needed ────────────────────────────────────────────────
@pytest.mark.asyncio
async def test_liveview_valid_key_returns_canvas(client: AsyncClient):
settings.liveview_key = "my-secret-key"
res = await client.get("/api/v1/liveview?key=my-secret-key")
assert res.status_code == 200
data = res.json()
assert "nodes" in data
assert "edges" in data
assert "viewport" in data
assert isinstance(data["nodes"], list)
assert isinstance(data["edges"], list)
@pytest.mark.asyncio
async def test_liveview_does_not_require_jwt(client: AsyncClient):
"""Accessing without Authorization header must work when key is correct."""
settings.liveview_key = "open-sesame"
# client has no auth headers set here
res = await client.get("/api/v1/liveview?key=open-sesame")
assert res.status_code == 200
@pytest.mark.asyncio
async def test_liveview_returns_saved_canvas(client: AsyncClient, auth_headers):
"""Canvas saved via POST /canvas/save appears in liveview response."""
settings.liveview_key = "test-key"
headers = await auth_headers()
# Save a canvas with one node
payload = {
"nodes": [{
"id": "lv-node-1",
"type": "server",
"label": "Live Node",
"status": "online",
"services": [],
"pos_x": 10,
"pos_y": 20,
}],
"edges": [],
"viewport": {"x": 0, "y": 0, "zoom": 1},
}
await client.post("/api/v1/canvas/save", json=payload, headers=headers)
# Liveview should return the same node
res = await client.get("/api/v1/liveview?key=test-key")
assert res.status_code == 200
nodes = res.json()["nodes"]
assert len(nodes) == 1
assert nodes[0]["id"] == "lv-node-1"
assert nodes[0]["label"] == "Live Node"
# ── custom_style + theme propagation ─────────────────────────────────────────
@pytest.mark.asyncio
async def test_liveview_returns_custom_style_and_theme(client: AsyncClient, auth_headers):
"""custom_style and viewport.theme_id from a saved canvas surface in liveview."""
settings.liveview_key = "test-key"
headers = await auth_headers()
payload = {
"nodes": [],
"edges": [],
"viewport": {"x": 0, "y": 0, "zoom": 1, "theme_id": "matrix"},
"custom_style": {"fontFamily": "Inter", "nodeRadius": 12},
}
await client.post("/api/v1/canvas/save", json=payload, headers=headers)
res = await client.get("/api/v1/liveview?key=test-key")
assert res.status_code == 200
body = res.json()
assert body["viewport"].get("theme_id") == "matrix"
assert body["custom_style"] == {"fontFamily": "Inter", "nodeRadius": 12}
# ── Re-disable after enabling ─────────────────────────────────────────────────
@pytest.mark.asyncio
async def test_liveview_disabled_after_key_cleared(client: AsyncClient):
settings.liveview_key = "was-enabled"
res = await client.get("/api/v1/liveview?key=was-enabled")
assert res.status_code == 200
settings.liveview_key = None
res = await client.get("/api/v1/liveview?key=was-enabled")
assert res.status_code == 403
assert res.json()["detail"] == "Live view is disabled"
# ── /config (authenticated) — key used to build share links ──────────────────
@pytest.mark.asyncio
async def test_liveview_config_requires_auth(client: AsyncClient):
"""The config endpoint exposes the key, so it must reject unauthenticated calls."""
settings.liveview_key = "secret"
res = await client.get("/api/v1/liveview/config")
assert res.status_code == 401
@pytest.mark.asyncio
async def test_liveview_config_returns_key_when_enabled(client: AsyncClient, auth_headers):
settings.liveview_key = "share-me"
headers = await auth_headers()
res = await client.get("/api/v1/liveview/config", headers=headers)
assert res.status_code == 200
body = res.json()
assert body == {"enabled": True, "key": "share-me"}
@pytest.mark.asyncio
async def test_liveview_config_disabled_hides_key(client: AsyncClient, auth_headers):
settings.liveview_key = None
headers = await auth_headers()
res = await client.get("/api/v1/liveview/config", headers=headers)
assert res.status_code == 200
assert res.json() == {"enabled": False, "key": None}
@pytest.mark.asyncio
async def test_liveview_config_empty_key_disabled(client: AsyncClient, auth_headers):
settings.liveview_key = ""
headers = await auth_headers()
res = await client.get("/api/v1/liveview/config", headers=headers)
assert res.status_code == 200
assert res.json() == {"enabled": False, "key": None}
# ── design_id selects which canvas is rendered ───────────────────────────────
@pytest.mark.asyncio
async def test_liveview_design_id_selects_canvas(client: AsyncClient, auth_headers):
"""?design_id=<id> renders that design's canvas, not the first one."""
settings.liveview_key = "test-key"
headers = await auth_headers()
# Create two designs
d1 = (await client.post("/api/v1/designs", json={"name": "Network"}, headers=headers)).json()
d2 = (await client.post("/api/v1/designs", json={"name": "Electrical"}, headers=headers)).json()
# Save a distinct node into each design
for design, node_id, label in ((d1, "n-net", "Net Node"), (d2, "n-elec", "Elec Node")):
payload = {
"nodes": [{
"id": node_id,
"type": "server",
"label": label,
"status": "online",
"services": [],
"pos_x": 0,
"pos_y": 0,
}],
"edges": [],
"viewport": {"x": 0, "y": 0, "zoom": 1},
"design_id": design["id"],
}
await client.post("/api/v1/canvas/save", json=payload, headers=headers)
# Requesting d2 returns only the electrical node
res = await client.get(f"/api/v1/liveview?key=test-key&design_id={d2['id']}")
assert res.status_code == 200
nodes = res.json()["nodes"]
assert [n["id"] for n in nodes] == ["n-elec"]
# Requesting d1 returns only the network node
res = await client.get(f"/api/v1/liveview?key=test-key&design_id={d1['id']}")
assert res.status_code == 200
nodes = res.json()["nodes"]
assert [n["id"] for n in nodes] == ["n-net"]
-134
View File
@@ -1,134 +0,0 @@
"""Backward-compatibility tests for the legacy → multi-design migration.
Simulates a database created by a pre-"designs" version of the app and asserts
that running init_db() adopts all existing nodes/edges/canvas into a single
default "Network Topology" design with no data loss. The rest of the test suite
builds the *current* schema via create_all and never exercises this upgrade
path, so this file guards real users upgrading in place.
"""
import os
os.environ.setdefault("SECRET_KEY", "test-only-secret-key-not-for-production")
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
import app.db.database as database
@pytest.fixture
def legacy_engine(tmp_path, monkeypatch):
"""Point the module-global engine + sqlite_path at a throwaway legacy DB."""
db_path = tmp_path / "legacy.db"
monkeypatch.setattr(database.settings, "sqlite_path", str(db_path))
engine = create_async_engine(f"sqlite+aiosqlite:///{db_path}")
monkeypatch.setattr(database, "engine", engine)
return db_path, engine
async def _build_legacy_schema(engine) -> None:
"""Create the pre-designs schema (no design_id, integer canvas_state PK)."""
async with engine.begin() as conn:
await conn.exec_driver_sql(
"CREATE TABLE nodes (id VARCHAR PRIMARY KEY, type VARCHAR, label VARCHAR, "
"status VARCHAR, services JSON, pos_x FLOAT, pos_y FLOAT)"
)
await conn.exec_driver_sql(
"CREATE TABLE edges (id VARCHAR PRIMARY KEY, source VARCHAR, target VARCHAR, type VARCHAR)"
)
await conn.exec_driver_sql(
"CREATE TABLE canvas_state (id INTEGER PRIMARY KEY, viewport JSON, "
"custom_style JSON, saved_at DATETIME)"
)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, type, label, status, services, pos_x, pos_y) "
"VALUES ('n1','server','Old Server','online','[]',10,20)"
)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, type, label, status, services, pos_x, pos_y) "
"VALUES ('n2','router','Old Router','offline','[]',30,40)"
)
await conn.exec_driver_sql(
"INSERT INTO edges (id, source, target, type) VALUES ('e1','n1','n2','ethernet')"
)
await conn.exec_driver_sql(
"INSERT INTO canvas_state (id, viewport, custom_style, saved_at) "
"VALUES (1, '{\"x\":5,\"y\":6,\"zoom\":2}', NULL, '2024-01-01 00:00:00')"
)
async def test_legacy_canvas_migrates_into_default_design(legacy_engine):
db_path, engine = legacy_engine
await _build_legacy_schema(engine)
await database.init_db()
check = create_async_engine(f"sqlite+aiosqlite:///{db_path}")
try:
async with check.begin() as conn:
# Exactly one seeded default design.
designs = (await conn.exec_driver_sql(
"SELECT id, name, design_type, icon FROM designs"
)).fetchall()
assert len(designs) == 1
did, name, dtype, icon = designs[0]
assert name == "Network Topology"
assert dtype == "network"
assert icon == "dashboard"
# Every legacy node adopted into the default design, data preserved.
nodes = (await conn.exec_driver_sql(
"SELECT id, label, status, design_id FROM nodes ORDER BY id"
)).fetchall()
assert [(n[0], n[1], n[2]) for n in nodes] == [
("n1", "Old Server", "online"),
("n2", "Old Router", "offline"),
]
assert all(n[3] == did for n in nodes)
# Legacy edge adopted too.
edge = (await conn.exec_driver_sql(
"SELECT design_id FROM edges WHERE id='e1'"
)).fetchone()
assert edge[0] == did
# canvas_state rebuilt with design_id PK; the old id=1 row maps to the
# default design and the viewport survives.
cs = (await conn.exec_driver_sql(
"SELECT design_id, viewport FROM canvas_state"
)).fetchall()
assert len(cs) == 1
assert cs[0][0] == did
assert "zoom" in (cs[0][1] or "")
finally:
await check.dispose()
await engine.dispose()
async def test_migration_is_idempotent(legacy_engine):
"""Running init_db twice must not duplicate the design or drop any data."""
db_path, engine = legacy_engine
await _build_legacy_schema(engine)
await database.init_db()
await database.init_db() # second boot — should be a no-op
check = create_async_engine(f"sqlite+aiosqlite:///{db_path}")
try:
async with check.begin() as conn:
designs = (await conn.exec_driver_sql("SELECT id FROM designs")).fetchall()
assert len(designs) == 1
did = designs[0][0]
nodes = (await conn.exec_driver_sql(
"SELECT design_id FROM nodes"
)).fetchall()
assert len(nodes) == 2
assert all(n[0] == did for n in nodes)
cs = (await conn.exec_driver_sql("SELECT design_id FROM canvas_state")).fetchall()
assert len(cs) == 1
assert cs[0][0] == did
finally:
await check.dispose()
await engine.dispose()
-94
View File
@@ -115,97 +115,3 @@ async def test_update_node_parent_id(client: AsyncClient, headers: dict):
async def test_create_node_requires_auth(client: AsyncClient):
res = await client.post("/api/v1/nodes", json={"type": "server", "label": "N", "status": "unknown"})
assert res.status_code == 401
# --- Properties tests ---
async def test_create_node_default_properties_empty(client: AsyncClient, headers: dict):
"""New node has an empty properties list by default."""
res = await client.post("/api/v1/nodes", json={"type": "server", "label": "Srv", "status": "unknown"}, headers=headers)
assert res.status_code == 201
assert res.json()["properties"] == []
async def test_create_node_with_properties(client: AsyncClient, headers: dict):
"""Node created with properties round-trips correctly."""
props = [
{"key": "CPU Model", "value": "i7-12700K", "icon": "Cpu", "visible": True},
{"key": "RAM", "value": "32 GB", "icon": "MemoryStick", "visible": False},
]
res = await client.post(
"/api/v1/nodes",
json={"type": "server", "label": "Srv", "status": "unknown", "properties": props},
headers=headers,
)
assert res.status_code == 201
assert res.json()["properties"] == props
async def test_patch_node_properties(client: AsyncClient, headers: dict):
"""PATCH with properties replaces the full properties array."""
create = await client.post("/api/v1/nodes", json={"type": "server", "label": "Srv", "status": "unknown"}, headers=headers)
node_id = create.json()["id"]
props = [{"key": "Disk", "value": "2 TB", "icon": "HardDrive", "visible": True}]
res = await client.patch(f"/api/v1/nodes/{node_id}", json={"properties": props}, headers=headers)
assert res.status_code == 200
assert res.json()["properties"] == props
async def test_patch_node_without_properties_does_not_wipe(client: AsyncClient, headers: dict):
"""PATCH that omits properties leaves existing properties untouched."""
props = [{"key": "GPU", "value": "RTX 4090", "icon": "Monitor", "visible": True}]
create = await client.post(
"/api/v1/nodes",
json={"type": "server", "label": "Srv", "status": "unknown", "properties": props},
headers=headers,
)
node_id = create.json()["id"]
# PATCH only the label — properties must survive
res = await client.patch(f"/api/v1/nodes/{node_id}", json={"label": "Updated"}, headers=headers)
assert res.status_code == 200
assert res.json()["properties"] == props
assert res.json()["label"] == "Updated"
async def test_patch_node_clears_properties_with_empty_array(client: AsyncClient, headers: dict):
"""PATCH with properties=[] explicitly clears all properties."""
props = [{"key": "CPU Model", "value": "i5", "icon": "Cpu", "visible": True}]
create = await client.post(
"/api/v1/nodes",
json={"type": "server", "label": "Srv", "status": "unknown", "properties": props},
headers=headers,
)
node_id = create.json()["id"]
res = await client.patch(f"/api/v1/nodes/{node_id}", json={"properties": []}, headers=headers)
assert res.status_code == 200
assert res.json()["properties"] == []
async def test_get_node_returns_properties(client: AsyncClient, headers: dict):
"""GET /nodes/:id returns the properties field."""
props = [{"key": "OS", "value": "Debian 12", "icon": "Server", "visible": True}]
create = await client.post(
"/api/v1/nodes",
json={"type": "server", "label": "Srv", "status": "unknown", "properties": props},
headers=headers,
)
node_id = create.json()["id"]
res = await client.get(f"/api/v1/nodes/{node_id}", headers=headers)
assert res.status_code == 200
assert res.json()["properties"] == props
async def test_properties_icon_can_be_null(client: AsyncClient, headers: dict):
"""A property with icon=null is valid and round-trips correctly."""
props = [{"key": "Notes", "value": "custom value", "icon": None, "visible": False}]
create = await client.post(
"/api/v1/nodes",
json={"type": "generic", "label": "G", "status": "unknown", "properties": props},
headers=headers,
)
assert create.status_code == 201
assert create.json()["properties"] == props
-176
View File
@@ -1,176 +0,0 @@
"""
Tests for the hardware → properties migration logic.
We test the migration function directly against an in-memory SQLite database
so we can set up legacy rows (with hardware columns, NULL properties) and
verify the migration produces the expected properties JSON.
"""
import json
import os
os.environ.setdefault("SECRET_KEY", "test-only-secret-key-not-for-production")
import pytest
from sqlalchemy.ext.asyncio import create_async_engine
TEST_DB_URL = "sqlite+aiosqlite:///:memory:"
async def _setup_legacy_table(conn):
"""Create a minimal nodes table that mimics the pre-migration schema."""
await conn.exec_driver_sql("""
CREATE TABLE IF NOT EXISTS nodes (
id TEXT PRIMARY KEY,
type TEXT NOT NULL DEFAULT 'generic',
label TEXT NOT NULL DEFAULT '',
cpu_model TEXT,
cpu_count INTEGER,
ram_gb REAL,
disk_gb REAL,
show_hardware BOOLEAN NOT NULL DEFAULT 0,
properties JSON
)
""")
async def _run_migration(conn):
"""Run only the properties migration portion (extracted from init_db)."""
rows = await conn.exec_driver_sql(
"SELECT id, cpu_model, cpu_count, ram_gb, disk_gb, show_hardware "
"FROM nodes WHERE properties IS NULL"
)
for row in rows.fetchall():
node_id, cpu_model, cpu_count, ram_gb, disk_gb, show_hardware = row
props = []
visible = bool(show_hardware)
if cpu_model:
props.append({"key": "CPU Model", "value": str(cpu_model), "icon": "Cpu", "visible": visible})
if cpu_count is not None:
props.append({"key": "CPU Cores", "value": str(cpu_count), "icon": "Cpu", "visible": visible})
if ram_gb is not None:
props.append({"key": "RAM", "value": f"{ram_gb} GB", "icon": "MemoryStick", "visible": visible})
if disk_gb is not None:
props.append({"key": "Disk", "value": f"{disk_gb} GB", "icon": "HardDrive", "visible": visible})
await conn.exec_driver_sql(
"UPDATE nodes SET properties = ? WHERE id = ?",
(json.dumps(props), node_id),
)
async def _get_properties(conn, node_id: str) -> list:
rows = await conn.exec_driver_sql("SELECT properties FROM nodes WHERE id = ?", (node_id,))
raw = rows.fetchone()[0]
return json.loads(raw) if raw else []
@pytest.mark.asyncio
async def test_migration_full_hardware():
"""Node with all 4 hardware fields → 4 property entries with correct icons."""
engine = create_async_engine(TEST_DB_URL)
async with engine.begin() as conn:
await _setup_legacy_table(conn)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, cpu_model, cpu_count, ram_gb, disk_gb, show_hardware) "
"VALUES (?, ?, ?, ?, ?, ?)",
("node-1", "i7-12700K", 12, 32.0, 2000.0, 1),
)
await _run_migration(conn)
props = await _get_properties(conn, "node-1")
assert len(props) == 4
assert props[0] == {"key": "CPU Model", "value": "i7-12700K", "icon": "Cpu", "visible": True}
assert props[1] == {"key": "CPU Cores", "value": "12", "icon": "Cpu", "visible": True}
assert props[2] == {"key": "RAM", "value": "32.0 GB", "icon": "MemoryStick", "visible": True}
assert props[3] == {"key": "Disk", "value": "2000.0 GB", "icon": "HardDrive", "visible": True}
await engine.dispose()
@pytest.mark.asyncio
async def test_migration_partial_hardware():
"""Node with only cpu_model and ram_gb → 2 property entries."""
engine = create_async_engine(TEST_DB_URL)
async with engine.begin() as conn:
await _setup_legacy_table(conn)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, cpu_model, ram_gb, show_hardware) VALUES (?, ?, ?, ?)",
("node-2", "Ryzen 5 5600", 16.0, 0),
)
await _run_migration(conn)
props = await _get_properties(conn, "node-2")
assert len(props) == 2
assert props[0]["key"] == "CPU Model"
assert props[0]["visible"] is False
assert props[1]["key"] == "RAM"
assert props[1]["icon"] == "MemoryStick"
await engine.dispose()
@pytest.mark.asyncio
async def test_migration_no_hardware():
"""Node with no hardware fields → empty properties array."""
engine = create_async_engine(TEST_DB_URL)
async with engine.begin() as conn:
await _setup_legacy_table(conn)
await conn.exec_driver_sql(
"INSERT INTO nodes (id) VALUES (?)",
("node-3",),
)
await _run_migration(conn)
props = await _get_properties(conn, "node-3")
assert props == []
await engine.dispose()
@pytest.mark.asyncio
async def test_migration_idempotent():
"""Running migration twice does not duplicate properties."""
engine = create_async_engine(TEST_DB_URL)
async with engine.begin() as conn:
await _setup_legacy_table(conn)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, cpu_model, show_hardware) VALUES (?, ?, ?)",
("node-4", "Core i5", 1),
)
await _run_migration(conn)
await _run_migration(conn) # second pass — node already has properties, should be skipped
props = await _get_properties(conn, "node-4")
assert len(props) == 1
await engine.dispose()
@pytest.mark.asyncio
async def test_migration_show_hardware_false_sets_visible_false():
"""show_hardware=0 means all migrated properties have visible=False."""
engine = create_async_engine(TEST_DB_URL)
async with engine.begin() as conn:
await _setup_legacy_table(conn)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, cpu_model, ram_gb, show_hardware) VALUES (?, ?, ?, ?)",
("node-5", "ARM Cortex-A72", 4.0, 0),
)
await _run_migration(conn)
props = await _get_properties(conn, "node-5")
assert all(p["visible"] is False for p in props)
await engine.dispose()
@pytest.mark.asyncio
async def test_migration_already_migrated_node_not_touched():
"""Node that already has properties is skipped — existing properties preserved."""
existing = [{"key": "GPU", "value": "RTX 4090", "icon": "Monitor", "visible": True}]
engine = create_async_engine(TEST_DB_URL)
async with engine.begin() as conn:
await _setup_legacy_table(conn)
await conn.exec_driver_sql(
"INSERT INTO nodes (id, cpu_model, ram_gb, show_hardware, properties) VALUES (?, ?, ?, ?, ?)",
("node-6", "i9-13900K", 64.0, 1, json.dumps(existing)),
)
await _run_migration(conn)
props = await _get_properties(conn, "node-6")
assert props == existing
await engine.dispose()
+5 -890
View File
@@ -1,4 +1,4 @@
"""Tests for scan routes: trigger, pending devices, approve/hide/ignore, stop."""
"""Tests for scan routes: trigger, pending devices, approve/hide/ignore."""
import uuid
from unittest.mock import AsyncMock, patch
@@ -7,8 +7,8 @@ from httpx import AsyncClient
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.db.models import Node, PendingDevice, ScanRun
from app.services.scanner import _cancelled_runs, request_cancel, run_scan
from app.db.models import PendingDevice, ScanRun
from app.services.scanner import run_scan
@pytest.fixture
@@ -37,95 +37,6 @@ async def pending_device(db_session):
return device
# --- _background_scan error handling ---
@pytest.fixture
async def mem_db():
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from app.db.database import Base
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
factory = async_sessionmaker(engine, expire_on_commit=False)
yield factory
await engine.dispose()
@pytest.mark.asyncio
async def test_background_scan_marks_run_failed_on_exception(mem_db):
"""If run_scan() raises, the ScanRun must transition running → failed and the
session rollback path must execute without a follow-on exception."""
from app.api.routes.scan import _background_scan
async with mem_db() as session:
run = ScanRun(status="running", ranges=["10.0.0.0/24"])
session.add(run)
await session.commit()
run_id = run.id
with (
patch("app.api.routes.scan.AsyncSessionLocal", mem_db),
patch(
"app.api.routes.scan.run_scan",
new_callable=AsyncMock,
side_effect=RuntimeError("boom"),
),
):
await _background_scan(run_id, ["10.0.0.0/24"])
async with mem_db() as session:
refreshed = await session.get(ScanRun, run_id)
assert refreshed is not None
assert refreshed.status == "failed"
@pytest.mark.asyncio
async def test_background_scan_leaves_non_running_status_alone(mem_db):
"""If the run was already stopped/cancelled before run_scan failed, _background_scan
must NOT overwrite that terminal status with 'failed'."""
from app.api.routes.scan import _background_scan
async with mem_db() as session:
run = ScanRun(status="cancelled", ranges=["10.0.0.0/24"])
session.add(run)
await session.commit()
run_id = run.id
with (
patch("app.api.routes.scan.AsyncSessionLocal", mem_db),
patch(
"app.api.routes.scan.run_scan",
new_callable=AsyncMock,
side_effect=RuntimeError("boom"),
),
):
await _background_scan(run_id, ["10.0.0.0/24"])
async with mem_db() as session:
refreshed = await session.get(ScanRun, run_id)
assert refreshed is not None
assert refreshed.status == "cancelled"
@pytest.mark.asyncio
async def test_background_scan_success_path_invokes_run_scan(mem_db):
from app.api.routes.scan import _background_scan
async with mem_db() as session:
run = ScanRun(status="running", ranges=["10.0.0.0/24"])
session.add(run)
await session.commit()
run_id = run.id
with (
patch("app.api.routes.scan.AsyncSessionLocal", mem_db),
patch("app.api.routes.scan.run_scan", new_callable=AsyncMock) as mock_run_scan,
):
await _background_scan(run_id, ["10.0.0.0/24"])
mock_run_scan.assert_awaited_once()
# --- Trigger scan ---
@pytest.mark.asyncio
@@ -209,7 +120,8 @@ async def test_approve_nonexistent_device(client: AsyncClient, headers):
json=node_payload,
headers=headers,
)
assert res.status_code == 404
assert res.status_code == 200
assert res.json()["approved"] is False
# --- Hide device ---
@@ -229,49 +141,6 @@ async def test_hide_device(client: AsyncClient, headers, pending_device):
assert len(hidden_res.json()) == 1
# --- Restore hidden device ---
@pytest.mark.asyncio
async def test_restore_device(client: AsyncClient, headers, pending_device):
# Hide first
await client.post(f"/api/v1/scan/pending/{pending_device.id}/hide", headers=headers)
# Restore
res = await client.post(f"/api/v1/scan/pending/{pending_device.id}/restore", headers=headers)
assert res.status_code == 200
assert res.json()["restored"] is True
# Now back in pending, gone from hidden
pending_res = await client.get("/api/v1/scan/pending", headers=headers)
assert len(pending_res.json()) == 1
hidden_res = await client.get("/api/v1/scan/hidden", headers=headers)
assert hidden_res.json() == []
@pytest.mark.asyncio
async def test_restore_device_rejects_non_hidden(client: AsyncClient, headers, pending_device):
res = await client.post(f"/api/v1/scan/pending/{pending_device.id}/restore", headers=headers)
assert res.status_code == 409
@pytest.mark.asyncio
async def test_bulk_restore_devices(client: AsyncClient, headers, pending_device):
# Hide
await client.post(f"/api/v1/scan/pending/{pending_device.id}/hide", headers=headers)
res = await client.post(
"/api/v1/scan/pending/bulk-restore",
headers=headers,
json={"device_ids": [pending_device.id]},
)
assert res.status_code == 200
assert res.json()["restored"] == 1
assert res.json()["skipped"] == 0
pending_res = await client.get("/api/v1/scan/pending", headers=headers)
assert len(pending_res.json()) == 1
# --- Ignore device ---
@pytest.mark.asyncio
@@ -330,213 +199,6 @@ async def test_run_scan_creates_new_pending_device(db_session: AsyncSession):
assert device.suggested_type == "server"
@pytest.mark.asyncio
async def test_run_scan_purges_stale_pending_for_canvas_nodes(db_session: AsyncSession):
"""Pending devices that were already in canvas before scan starts must be removed."""
node = Node(
id=str(uuid.uuid4()),
label="Existing Server",
type="server",
ip="192.168.1.50",
status="online",
services=[],
pos_x=0.0,
pos_y=0.0,
)
stale = PendingDevice(
id=str(uuid.uuid4()),
ip="192.168.1.50",
mac=None,
hostname=None,
os=None,
services=[],
suggested_type="generic",
status="pending",
)
db_session.add(node)
db_session.add(stale)
await db_session.commit()
run_id = str(uuid.uuid4())
run = ScanRun(id=run_id, status="running", ranges=["192.168.1.0/24"])
db_session.add(run)
await db_session.commit()
with (
patch("app.services.scanner._nmap_scan", return_value=[]),
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock),
):
await run_scan(["192.168.1.0/24"], db_session, run_id)
result = await db_session.execute(
select(PendingDevice).where(PendingDevice.ip == "192.168.1.50")
)
assert result.scalar_one_or_none() is None
@pytest.mark.asyncio
async def test_run_scan_skips_ip_already_in_canvas(db_session: AsyncSession):
"""Devices whose IP already exists as a canvas Node must not appear in pending."""
node = Node(
id=str(uuid.uuid4()),
label="Existing Server",
type="server",
ip="192.168.1.50",
status="online",
services=[],
pos_x=0.0,
pos_y=0.0,
)
db_session.add(node)
await db_session.commit()
run_id = str(uuid.uuid4())
run = ScanRun(id=run_id, status="running", ranges=["192.168.1.0/24"])
db_session.add(run)
await db_session.commit()
with (
patch("app.services.scanner._nmap_scan", return_value=[MOCK_HOST]),
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock),
):
await run_scan(["192.168.1.0/24"], db_session, run_id)
result = await db_session.execute(
select(PendingDevice).where(PendingDevice.ip == "192.168.1.50")
)
assert result.scalar_one_or_none() is None
@pytest.mark.asyncio
async def test_run_scan_skips_hidden_device(db_session: AsyncSession):
"""Devices previously hidden by the user must not re-appear in pending on re-scan."""
hidden = PendingDevice(
id=str(uuid.uuid4()),
ip="192.168.1.50",
mac=None,
hostname=None,
os=None,
services=[],
suggested_type="generic",
status="hidden",
)
db_session.add(hidden)
await db_session.commit()
run_id = str(uuid.uuid4())
run = ScanRun(id=run_id, status="running", ranges=["192.168.1.0/24"])
db_session.add(run)
await db_session.commit()
with (
patch("app.services.scanner._nmap_scan", return_value=[MOCK_HOST]),
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock),
):
await run_scan(["192.168.1.0/24"], db_session, run_id)
result = await db_session.execute(
select(PendingDevice).where(
PendingDevice.ip == "192.168.1.50",
PendingDevice.status == "pending",
)
)
assert result.scalar_one_or_none() is None
# --- Stop scan ---
@pytest.mark.asyncio
async def test_stop_scan_requires_auth(client: AsyncClient):
res = await client.post("/api/v1/scan/fake-id/stop")
assert res.status_code == 401
@pytest.mark.asyncio
async def test_stop_scan_not_found(client: AsyncClient, headers):
import uuid as _uuid
res = await client.post(f"/api/v1/scan/{_uuid.uuid4()}/stop", headers=headers)
assert res.status_code == 404
@pytest.mark.asyncio
async def test_stop_scan_not_running(client: AsyncClient, headers, db_session: AsyncSession):
run = ScanRun(id=str(uuid.uuid4()), status="done", ranges=["192.168.1.0/24"])
db_session.add(run)
await db_session.commit()
res = await client.post(f"/api/v1/scan/{run.id}/stop", headers=headers)
assert res.status_code == 409
@pytest.mark.asyncio
async def test_stop_scan_success(client: AsyncClient, headers, db_session: AsyncSession):
run = ScanRun(id=str(uuid.uuid4()), status="running", ranges=["192.168.1.0/24"])
db_session.add(run)
await db_session.commit()
res = await client.post(f"/api/v1/scan/{run.id}/stop", headers=headers)
assert res.status_code == 200
assert res.json() == {"stopping": True}
# run_id added to cancel set
assert run.id in _cancelled_runs
# cleanup for other tests
_cancelled_runs.discard(run.id)
# --- run_scan cancellation ---
@pytest.mark.asyncio
async def test_run_scan_cancelled_marks_status(db_session: AsyncSession):
"""When cancel is requested before the scan starts, status becomes 'cancelled'."""
run_id = str(uuid.uuid4())
run = ScanRun(id=run_id, status="running", ranges=["192.168.1.0/24"])
db_session.add(run)
await db_session.commit()
request_cancel(run_id)
with (
patch("app.services.scanner._nmap_scan", return_value=[MOCK_HOST]) as mock_nmap,
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock),
):
await run_scan(["192.168.1.0/24"], db_session, run_id)
# nmap should not have been called — cancelled before first range
mock_nmap.assert_not_called()
await db_session.refresh(run)
assert run.status == "cancelled"
assert run.finished_at is not None
@pytest.mark.asyncio
async def test_run_scan_cancelled_mid_scan_skips_remaining_cidrs(db_session: AsyncSession):
"""Cancel flag set after first CIDR is started prevents processing of the second CIDR."""
run_id = str(uuid.uuid4())
run = ScanRun(id=run_id, status="running", ranges=["10.0.0.0/24", "10.0.1.0/24"])
db_session.add(run)
await db_session.commit()
call_count = 0
def nmap_side_effect(target: str):
nonlocal call_count
call_count += 1
# Signal cancellation after the first CIDR scan completes
if call_count == 1:
request_cancel(run_id)
return []
with (
patch("app.services.scanner._nmap_scan", side_effect=nmap_side_effect),
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock),
):
await run_scan(["10.0.0.0/24", "10.0.1.0/24"], db_session, run_id)
assert call_count == 1 # second CIDR was skipped
await db_session.refresh(run)
assert run.status == "cancelled"
@pytest.mark.asyncio
async def test_run_scan_updates_existing_pending_device(db_session: AsyncSession):
"""Re-scanning the same IP updates services instead of creating a duplicate."""
@@ -575,550 +237,3 @@ async def test_run_scan_updates_existing_pending_device(db_session: AsyncSession
# Services and hostname should be updated
assert device.hostname == "myhost.lan"
assert any(s["port"] == 8096 for s in device.services)
# --- Bulk approve ---
@pytest.fixture
async def two_pending_devices(db_session):
devices = []
for i in range(2):
d = PendingDevice(
id=str(uuid.uuid4()),
ip=f"192.168.1.{10 + i}",
mac=None,
hostname=f"host-{i}",
os=None,
services=[],
suggested_type="generic",
status="pending",
)
db_session.add(d)
devices.append(d)
await db_session.commit()
for d in devices:
await db_session.refresh(d)
return devices
@pytest.mark.asyncio
async def test_bulk_approve_approves_devices(client: AsyncClient, headers, two_pending_devices):
ids = [d.id for d in two_pending_devices]
res = await client.post("/api/v1/scan/pending/bulk-approve", json={"device_ids": ids}, headers=headers)
assert res.status_code == 200
data = res.json()
assert data["approved"] == 2
assert len(data["node_ids"]) == 2
assert all(nid is not None for nid in data["node_ids"]), "node_ids must be non-null UUIDs"
assert len(data["device_ids"]) == 2
assert data["skipped"] == 0
# Pending list should now be empty
pending_res = await client.get("/api/v1/scan/pending", headers=headers)
assert pending_res.json() == []
@pytest.fixture
async def zigbee_pending_device(db_session):
device = PendingDevice(
id=str(uuid.uuid4()),
ip=None,
mac=None,
hostname=None,
friendly_name="bulb_1",
services=[],
suggested_type="zigbee_enddevice",
device_subtype="EndDevice",
ieee_address="0xABCDEF",
vendor="IKEA",
model="TRADFRI",
lqi=180,
status="pending",
discovery_source="zigbee",
)
db_session.add(device)
await db_session.commit()
await db_session.refresh(device)
return device
@pytest.mark.asyncio
async def test_approve_zigbee_device_populates_properties(
client: AsyncClient, headers, zigbee_pending_device, db_session
):
"""Approving a zigbee device must populate IEEE/Vendor/Model/LQI in properties."""
from sqlalchemy import select
from app.db.models import Node as NodeModel
payload = {
"label": "bulb_1",
"type": "zigbee_enddevice",
"status": "online",
"services": [],
"check_method": "none",
}
res = await client.post(
f"/api/v1/scan/pending/{zigbee_pending_device.id}/approve",
json=payload,
headers=headers,
)
assert res.status_code == 200
node = (
await db_session.execute(select(NodeModel).where(NodeModel.ieee_address == "0xABCDEF"))
).scalar_one()
keys = {p["key"]: p["value"] for p in node.properties}
assert keys == {
"IEEE": "0xABCDEF",
"Vendor": "IKEA",
"Model": "TRADFRI",
"LQI": "180",
}
@pytest.mark.asyncio
async def test_bulk_approve_zigbee_populates_properties(
client: AsyncClient, headers, zigbee_pending_device, db_session
):
from sqlalchemy import select
from app.db.models import Node as NodeModel
res = await client.post(
"/api/v1/scan/pending/bulk-approve",
json={"device_ids": [zigbee_pending_device.id]},
headers=headers,
)
assert res.status_code == 200
node = (
await db_session.execute(select(NodeModel).where(NodeModel.ieee_address == "0xABCDEF"))
).scalar_one()
keys = {p["key"]: p["value"] for p in node.properties}
assert keys["IEEE"] == "0xABCDEF"
assert keys["Vendor"] == "IKEA"
assert keys["Model"] == "TRADFRI"
assert keys["LQI"] == "180"
assert node.check_method == "none"
# --- MAC address propagation on approve (issue #168) ---
def test_build_mac_property_returns_hidden_row():
from app.api.routes.scan import build_mac_property
assert build_mac_property("aa:bb:cc:dd:ee:ff") == [
{"key": "MAC", "value": "aa:bb:cc:dd:ee:ff", "icon": None, "visible": False}
]
def test_build_mac_property_empty_when_no_mac():
from app.api.routes.scan import build_mac_property
assert build_mac_property(None) == []
assert build_mac_property("") == []
def test_merge_mac_property_appends_when_absent():
from app.api.routes.scan import merge_mac_property
existing = [{"key": "Custom", "value": "x", "icon": None, "visible": True}]
merged = merge_mac_property(existing, "aa:bb:cc:dd:ee:ff")
assert {"key": "MAC", "value": "aa:bb:cc:dd:ee:ff", "icon": None, "visible": False} in merged
# Existing prop preserved untouched.
assert existing[0] in merged
def test_merge_mac_property_idempotent_and_preserves_visibility():
from app.api.routes.scan import merge_mac_property
existing = [{"key": "MAC", "value": "aa:bb:cc:dd:ee:ff", "icon": None, "visible": True}]
merged = merge_mac_property(existing, "aa:bb:cc:dd:ee:ff")
# No duplicate MAC row; user's visible=True choice kept.
macs = [p for p in merged if p["key"] == "MAC"]
assert len(macs) == 1
assert macs[0]["visible"] is True
def test_merge_mac_property_noop_without_mac():
from app.api.routes.scan import merge_mac_property
existing = [{"key": "Custom", "value": "x", "icon": None, "visible": True}]
assert merge_mac_property(existing, None) == existing
@pytest.mark.asyncio
async def test_approve_device_copies_mac_to_node_and_properties(
client: AsyncClient, headers, pending_device, db_session
):
"""Approving a scanned device must carry its MAC onto the node + properties."""
from sqlalchemy import select
from app.db.models import Node as NodeModel
# Payload intentionally omits mac — it must come from the pending device.
res = await client.post(
f"/api/v1/scan/pending/{pending_device.id}/approve",
json={"label": "My Server", "type": "server", "ip": "192.168.1.100", "status": "unknown", "services": []},
headers=headers,
)
assert res.status_code == 200
node = (
await db_session.execute(select(NodeModel).where(NodeModel.ip == "192.168.1.100"))
).scalar_one()
assert node.mac == "aa:bb:cc:dd:ee:ff"
mac_props = [p for p in node.properties if p["key"] == "MAC"]
assert mac_props == [
{"key": "MAC", "value": "aa:bb:cc:dd:ee:ff", "icon": None, "visible": False}
]
@pytest.mark.asyncio
async def test_approve_device_does_not_duplicate_mac_property(
client: AsyncClient, headers, pending_device, db_session
):
"""If the approve payload already carries a MAC prop, don't add a second one."""
from sqlalchemy import select
from app.db.models import Node as NodeModel
res = await client.post(
f"/api/v1/scan/pending/{pending_device.id}/approve",
json={
"label": "My Server",
"type": "server",
"ip": "192.168.1.100",
"status": "unknown",
"services": [],
"properties": [
{"key": "MAC", "value": "aa:bb:cc:dd:ee:ff", "icon": None, "visible": True}
],
},
headers=headers,
)
assert res.status_code == 200
node = (
await db_session.execute(select(NodeModel).where(NodeModel.ip == "192.168.1.100"))
).scalar_one()
mac_props = [p for p in node.properties if p["key"] == "MAC"]
assert len(mac_props) == 1
# User's visibility choice is preserved.
assert mac_props[0]["visible"] is True
@pytest.mark.asyncio
async def test_bulk_approve_copies_mac_to_node_and_properties(
client: AsyncClient, headers, db_session
):
"""Bulk approve must also propagate the scanned MAC to node + properties."""
from sqlalchemy import select
from app.db.models import Node as NodeModel
device = PendingDevice(
id=str(uuid.uuid4()),
ip="192.168.1.55",
mac="11:22:33:44:55:66",
hostname="host-mac",
services=[],
suggested_type="generic",
status="pending",
)
db_session.add(device)
await db_session.commit()
res = await client.post(
"/api/v1/scan/pending/bulk-approve",
json={"device_ids": [device.id]},
headers=headers,
)
assert res.status_code == 200
node = (
await db_session.execute(select(NodeModel).where(NodeModel.ip == "192.168.1.55"))
).scalar_one()
assert node.mac == "11:22:33:44:55:66"
mac_props = [p for p in node.properties if p["key"] == "MAC"]
assert mac_props == [
{"key": "MAC", "value": "11:22:33:44:55:66", "icon": None, "visible": False}
]
@pytest.mark.asyncio
async def test_bulk_approve_sets_default_check_method(client: AsyncClient, headers, two_pending_devices, db_session):
"""Approved devices with an IP must default to ping; otherwise scheduler skips them."""
from sqlalchemy import select
from app.db.models import Node as NodeModel
ids = [d.id for d in two_pending_devices]
res = await client.post("/api/v1/scan/pending/bulk-approve", json={"device_ids": ids}, headers=headers)
assert res.status_code == 200
nodes = (await db_session.execute(select(NodeModel))).scalars().all()
for n in nodes:
if n.ip:
assert n.check_method == "ping", f"node {n.id} created without check_method"
@pytest.mark.asyncio
async def test_approve_device_sets_default_check_method(client: AsyncClient, headers, pending_device, db_session):
from sqlalchemy import select
from app.db.models import Node as NodeModel
res = await client.post(
f"/api/v1/scan/pending/{pending_device.id}/approve",
json={"label": "h", "type": "generic", "ip": "192.168.1.10", "status": "unknown", "services": []},
headers=headers,
)
assert res.status_code == 200
node = (await db_session.execute(select(NodeModel))).scalars().first()
assert node is not None
assert node.check_method == "ping"
@pytest.mark.asyncio
async def test_bulk_approve_skips_already_approved(client: AsyncClient, headers, two_pending_devices):
ids = [d.id for d in two_pending_devices]
# Approve first device individually first
await client.post(
f"/api/v1/scan/pending/{ids[0]}/approve",
json={"label": "h", "type": "generic", "ip": "192.168.1.10", "status": "unknown", "services": []},
headers=headers,
)
# Bulk approve both — first one is already approved (not pending), should be skipped
res = await client.post("/api/v1/scan/pending/bulk-approve", json={"device_ids": ids}, headers=headers)
assert res.status_code == 200
data = res.json()
assert data["approved"] == 1
assert data["skipped"] == 1
@pytest.mark.asyncio
async def test_bulk_approve_requires_auth(client: AsyncClient, two_pending_devices):
ids = [d.id for d in two_pending_devices]
res = await client.post("/api/v1/scan/pending/bulk-approve", json={"device_ids": ids})
assert res.status_code == 401
# --- Bulk hide ---
@pytest.mark.asyncio
async def test_bulk_hide_hides_devices(client: AsyncClient, headers, two_pending_devices):
ids = [d.id for d in two_pending_devices]
res = await client.post("/api/v1/scan/pending/bulk-hide", json={"device_ids": ids}, headers=headers)
assert res.status_code == 200
data = res.json()
assert data["hidden"] == 2
assert data["skipped"] == 0
# Should appear in hidden list
hidden_res = await client.get("/api/v1/scan/hidden", headers=headers)
assert len(hidden_res.json()) == 2
@pytest.mark.asyncio
async def test_bulk_hide_skips_non_pending(client: AsyncClient, headers, two_pending_devices):
ids = [d.id for d in two_pending_devices]
# Hide first device individually first
await client.post(f"/api/v1/scan/pending/{ids[0]}/hide", headers=headers)
# Bulk hide both — first is already hidden (not pending anymore)
res = await client.post("/api/v1/scan/pending/bulk-hide", json={"device_ids": ids}, headers=headers)
assert res.status_code == 200
data = res.json()
assert data["hidden"] == 1
assert data["skipped"] == 1
@pytest.mark.asyncio
async def test_bulk_hide_requires_auth(client: AsyncClient, two_pending_devices):
ids = [d.id for d in two_pending_devices]
res = await client.post("/api/v1/scan/pending/bulk-hide", json={"device_ids": ids})
assert res.status_code == 401
# ---------------------------------------------------------------------------
# Approve auto-creates Edges from pending_device_links (Zigbee flow)
# ---------------------------------------------------------------------------
async def _seed_zigbee_pending_pair(db_session):
"""Create a coordinator Node + a pending device + a link between them."""
from app.db.models import Node, PendingDevice, PendingDeviceLink
coord = Node(
label="Coordinator",
type="zigbee_coordinator",
status="unknown",
ieee_address="0xCOORD",
)
db_session.add(coord)
pending = PendingDevice(
ieee_address="0xR1",
friendly_name="router_1",
suggested_type="zigbee_router",
device_subtype="Router",
status="pending",
discovery_source="zigbee",
)
db_session.add(pending)
db_session.add(
PendingDeviceLink(
source_ieee="0xCOORD",
target_ieee="0xR1",
discovery_source="zigbee",
)
)
await db_session.commit()
return coord, pending
@pytest.mark.asyncio
async def test_approve_zigbee_creates_edge_when_other_endpoint_is_node(
client: AsyncClient, headers, db_session
):
from sqlalchemy import select
from app.db.models import Edge
coord, pending = await _seed_zigbee_pending_pair(db_session)
res = await client.post(
f"/api/v1/scan/pending/{pending.id}/approve",
json={
"label": "router_1",
"type": "zigbee_router",
"ip": None,
"status": "unknown",
"services": [],
},
headers=headers,
)
assert res.status_code == 200
data = res.json()
assert data["approved"] is True
assert data["edges_created"] == 1
edges = (await db_session.execute(select(Edge))).scalars().all()
assert len(edges) == 1
assert edges[0].source == coord.id
assert edges[0].target == data["node_id"]
assert edges[0].source_handle == "bottom"
assert edges[0].target_handle == "top-t"
assert edges[0].type == "iot"
@pytest.mark.asyncio
async def test_approve_zigbee_skips_duplicate_edge(
client: AsyncClient, headers, db_session
):
"""Re-running the resolution does not create a second edge for the same pair."""
from sqlalchemy import select
from app.db.models import Edge, PendingDevice, PendingDeviceLink
coord, pending = await _seed_zigbee_pending_pair(db_session)
body = {"label": "router_1", "type": "zigbee_router", "ip": None, "status": "unknown", "services": []}
await client.post(f"/api/v1/scan/pending/{pending.id}/approve", json=body, headers=headers)
# Simulate a second pending row + link between same coord and a new device,
# but keep an existing edge in place to verify dedupe also handles
# the swapped-direction case.
new_pending = PendingDevice(
ieee_address="0xR1B",
friendly_name="r1b",
suggested_type="zigbee_router",
status="pending",
discovery_source="zigbee",
)
db_session.add(new_pending)
db_session.add(
PendingDeviceLink(source_ieee="0xCOORD", target_ieee="0xR1B", discovery_source="zigbee")
)
await db_session.commit()
res = await client.post(
f"/api/v1/scan/pending/{new_pending.id}/approve", json=body, headers=headers
)
assert res.json()["edges_created"] == 1 # only the new pair
edges = (await db_session.execute(select(Edge))).scalars().all()
assert len(edges) == 2 # original + new, no duplicate
@pytest.mark.asyncio
async def test_approve_zigbee_skips_when_other_endpoint_still_pending(
client: AsyncClient, headers, db_session
):
"""Both endpoints pending → no edge yet, link row preserved for later."""
from sqlalchemy import select
from app.db.models import Edge, PendingDevice, PendingDeviceLink
a = PendingDevice(
ieee_address="0xA",
friendly_name="a",
suggested_type="zigbee_router",
status="pending",
discovery_source="zigbee",
)
b = PendingDevice(
ieee_address="0xB",
friendly_name="b",
suggested_type="zigbee_enddevice",
status="pending",
discovery_source="zigbee",
)
db_session.add_all([a, b])
db_session.add(
PendingDeviceLink(source_ieee="0xA", target_ieee="0xB", discovery_source="zigbee")
)
await db_session.commit()
res = await client.post(
f"/api/v1/scan/pending/{a.id}/approve",
json={
"label": "a",
"type": "zigbee_router",
"ip": None,
"status": "unknown",
"services": [],
},
headers=headers,
)
assert res.status_code == 200
assert res.json()["edges_created"] == 0
edges = (await db_session.execute(select(Edge))).scalars().all()
assert edges == []
links = (await db_session.execute(select(PendingDeviceLink))).scalars().all()
assert len(links) == 1 # preserved for later resolution
@pytest.mark.asyncio
async def test_approve_zigbee_resolves_link_after_second_approval(
client: AsyncClient, headers, db_session
):
"""First approval keeps link; second approval creates the edge."""
from sqlalchemy import select
from app.db.models import Edge, PendingDevice, PendingDeviceLink
a = PendingDevice(
ieee_address="0xA",
friendly_name="a",
suggested_type="zigbee_router",
status="pending",
discovery_source="zigbee",
)
b = PendingDevice(
ieee_address="0xB",
friendly_name="b",
suggested_type="zigbee_enddevice",
status="pending",
discovery_source="zigbee",
)
db_session.add_all([a, b])
db_session.add(
PendingDeviceLink(source_ieee="0xA", target_ieee="0xB", discovery_source="zigbee")
)
await db_session.commit()
body = {"label": "x", "type": "zigbee_router", "ip": None, "status": "unknown", "services": []}
await client.post(f"/api/v1/scan/pending/{a.id}/approve", json=body, headers=headers)
res = await client.post(f"/api/v1/scan/pending/{b.id}/approve", json=body, headers=headers)
assert res.json()["edges_created"] == 1
edges = (await db_session.execute(select(Edge))).scalars().all()
assert len(edges) == 1
links = (await db_session.execute(select(PendingDeviceLink))).scalars().all()
assert links == [] # consumed
-535
View File
@@ -1,535 +0,0 @@
"""Tests for scanner: two-phase nmap, mDNS discovery, run_scan integration."""
import uuid
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy import select as sa_select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.pool import StaticPool
from app.db.database import Base
from app.db.models import Node, PendingDevice, ScanRun
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_run_id() -> str:
return str(uuid.uuid4())
@pytest.fixture
async def mem_db():
engine = create_async_engine(
"sqlite+aiosqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
yield factory
await engine.dispose()
def _make_scan_run(run_id: str) -> ScanRun:
return ScanRun(id=run_id, status="running", ranges=["192.168.1.0/24"])
# ---------------------------------------------------------------------------
# _ping_sweep
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_ping_sweep_returns_alive_hosts():
from app.services.scanner import _ping_sweep
async def fake_ping(ip: str) -> str | None:
return ip if ip in {"192.168.1.1", "192.168.1.2"} else None
with patch("app.services.scanner._ping_sweep", wraps=None):
pass # just ensure import is fine
# Patch asyncio.create_subprocess_exec to simulate ping responses
responding = {"192.168.1.1", "192.168.1.2"}
async def mock_subprocess(*args, **kwargs):
ip = args[-1]
proc = MagicMock()
proc.returncode = 0 if ip in responding else 1
proc.wait = AsyncMock(return_value=proc.returncode)
return proc
with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \
patch("app.services.scanner._arp_table_hosts", return_value={}), \
patch("app.services.scanner._resolve_hostname", return_value=None):
result = await _ping_sweep("192.168.1.0/30") # .1 .2 only in /30
assert "192.168.1.1" in result
assert "192.168.1.2" in result
for host in result.values():
assert host["open_ports"] == []
@pytest.mark.asyncio
async def test_ping_sweep_excludes_non_responding():
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
with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \
patch("app.services.scanner._arp_table_hosts", return_value={}), \
patch("app.services.scanner._resolve_hostname", return_value=None):
result = await _ping_sweep("192.168.1.0/30")
assert "192.168.1.1" in result
assert "192.168.1.2" not in result
@pytest.mark.asyncio
async def test_ping_sweep_supplements_with_arp_cache():
"""Devices that block ICMP but appear in ARP cache should still be discovered."""
from app.services.scanner import _ping_sweep
async def mock_subprocess(*args, **kwargs):
proc = MagicMock()
proc.returncode = 1 # all pings fail
proc.wait = AsyncMock(return_value=1)
return proc
arp_extra = {
"192.168.1.10": {"ip": "192.168.1.10", "mac": "aa:bb:cc:dd:ee:10", "hostname": None, "os": None, "open_ports": []},
}
with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \
patch("app.services.scanner._arp_table_hosts", return_value=arp_extra), \
patch("app.services.scanner._resolve_hostname", return_value=None):
result = await _ping_sweep("192.168.1.0/24")
assert "192.168.1.10" in result
assert result["192.168.1.10"]["mac"] == "aa:bb:cc:dd:ee:10"
@pytest.mark.asyncio
async def test_ping_sweep_enriches_mac_from_arp_cache():
"""Ping-alive hosts with no ARP entry get their MAC from the ARP cache."""
from app.services.scanner import _ping_sweep
async def mock_subprocess(*args, **kwargs):
ip = args[-1]
proc = MagicMock()
proc.returncode = 0 if ip == "192.168.1.1" else 1
proc.wait = AsyncMock(return_value=proc.returncode)
return proc
arp_extra = {
"192.168.1.1": {"ip": "192.168.1.1", "mac": "de:ad:be:ef:00:01", "hostname": None, "os": None, "open_ports": []},
}
with patch("asyncio.create_subprocess_exec", side_effect=mock_subprocess), \
patch("app.services.scanner._arp_table_hosts", return_value=arp_extra), \
patch("app.services.scanner._resolve_hostname", return_value=None):
result = await _ping_sweep("192.168.1.0/30")
assert result["192.168.1.1"]["mac"] == "de:ad:be:ef:00:01"
# ---------------------------------------------------------------------------
# _arp_table_hosts
# ---------------------------------------------------------------------------
def test_arp_table_hosts_parses_proc_net_arp():
import io # noqa: PLC0415
from app.services.scanner import _arp_table_hosts
arp_content = (
"IP address HW type Flags HW address Mask Device\n"
"192.168.1.1 0x1 0x2 aa:bb:cc:dd:ee:01 * eth0\n"
"192.168.1.50 0x1 0x2 aa:bb:cc:dd:ee:02 * eth0\n"
"10.0.0.1 0x1 0x2 aa:bb:cc:dd:ee:03 * eth0\n" # outside subnet
"192.168.1.99 0x1 0x2 00:00:00:00:00:00 * eth0\n" # incomplete
)
mock_file = MagicMock()
mock_file.__enter__ = MagicMock(return_value=io.StringIO(arp_content))
mock_file.__exit__ = MagicMock(return_value=False)
with patch("builtins.open", return_value=mock_file), \
patch("app.services.scanner._resolve_hostname", return_value=None):
result = _arp_table_hosts("192.168.1.0/24")
assert "192.168.1.1" in result
assert "192.168.1.50" in result
assert "10.0.0.1" not in result # outside target subnet
assert "192.168.1.99" not in result # zero MAC skipped
def test_arp_table_hosts_parses_macos_arp_output():
from app.services.scanner import _arp_table_hosts
arp_output = (
"router.lan (192.168.1.1) at aa:bb:cc:dd:ee:01 on en0 ifscope [ethernet]\n"
"device.lan (192.168.1.20) at aa:bb:cc:dd:ee:02 on en0 ifscope [ethernet]\n"
"? (192.168.1.99) at (incomplete) on en0 ifscope [ethernet]\n"
"? (10.0.0.1) at aa:bb:cc:dd:ee:04 on en0 ifscope [ethernet]\n" # outside subnet
)
mock_result = MagicMock()
mock_result.stdout = arp_output
with patch("builtins.open", side_effect=FileNotFoundError), \
patch("subprocess.run", return_value=mock_result), \
patch("app.services.scanner._resolve_hostname", return_value=None):
result = _arp_table_hosts("192.168.1.0/24")
assert "192.168.1.1" in result
assert "192.168.1.20" in result
assert "192.168.1.99" not in result # incomplete MAC
assert "10.0.0.1" not in result # outside subnet
# ---------------------------------------------------------------------------
# _nmap_scan_single (Phase 2 per-IP worker)
# ---------------------------------------------------------------------------
def test_nmap_scan_single_detects_open_ports():
from app.services.scanner import _nmap_scan_single
host = {"ip": "192.168.1.10", "hostname": None, "mac": None, "os": None, "open_ports": []}
# Build a realistic host entry: protocols → ports → port info
port_info = {80: {"state": "open", "product": "nginx", "version": "1.24"}}
mock_host = MagicMock()
mock_host.all_protocols.return_value = ["tcp"]
mock_host.__getitem__ = MagicMock(return_value=port_info)
mock_host.get.return_value = {}
mock_nm = MagicMock()
mock_nm.all_hosts.return_value = ["192.168.1.10"]
mock_nm.__getitem__ = MagicMock(return_value=mock_host)
with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm), \
patch("app.services.scanner._extract_os", return_value=None):
result = _nmap_scan_single(host)
assert len(result["open_ports"]) == 1
assert result["open_ports"][0]["port"] == 80
assert result["open_ports"][0]["banner"] == "nginx 1.24"
def test_nmap_scan_single_returns_host_unchanged_on_error():
from app.services.scanner import _nmap_scan_single
host = {"ip": "192.168.1.20", "hostname": None, "mac": None, "os": None, "open_ports": []}
mock_nm = MagicMock()
mock_nm.scan.side_effect = Exception("nmap error")
with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm):
result = _nmap_scan_single(host)
assert result["ip"] == "192.168.1.20"
assert result["open_ports"] == []
def test_nmap_scan_single_returns_host_unchanged_when_no_results():
"""Host confirmed alive in Phase 1 but all ports filtered — keep it with empty ports."""
from app.services.scanner import _nmap_scan_single
host = {"ip": "192.168.1.30", "hostname": "shelly1.lan", "mac": "34:94:54:aa:bb:cc", "os": None, "open_ports": []}
mock_nm = MagicMock()
mock_nm.all_hosts.return_value = [] # no results
with patch("app.services.scanner.nmap.PortScanner", return_value=mock_nm):
result = _nmap_scan_single(host)
assert result["ip"] == "192.168.1.30"
assert result["open_ports"] == []
assert result["mac"] == "34:94:54:aa:bb:cc" # preserved from Phase 1
# ---------------------------------------------------------------------------
# _nmap_scan
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_nmap_scan_uses_mock_when_nmap_unavailable():
from app.services.scanner import _nmap_scan
with patch("app.services.scanner._NMAP_AVAILABLE", False):
result = await _nmap_scan("192.168.1.0/24")
assert len(result) == 1
assert result[0]["ip"] == "192.168.1.99"
@pytest.mark.asyncio
async def test_nmap_scan_raises_on_sweep_error():
from app.services.scanner import _nmap_scan
with patch("app.services.scanner._ping_sweep", side_effect=Exception("ping sweep failed")), \
pytest.raises(RuntimeError, match="ping sweep failed"):
await _nmap_scan("192.168.1.0/24")
# ---------------------------------------------------------------------------
# _mdns_discover
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_mdns_discover_returns_empty_when_zeroconf_unavailable():
from app.services.scanner import _mdns_discover
with patch("app.services.scanner._ZEROCONF_AVAILABLE", False):
result = await _mdns_discover()
assert result == []
@pytest.mark.asyncio
async def test_mdns_discover_returns_devices():
from app.services.scanner import _mdns_discover
mock_info = MagicMock()
mock_info.addresses = [b"\xc0\xa8\x01\x50"] # 192.168.1.80
mock_info.server = "shelly1.local."
mock_info.port = 80
mock_info.async_request = AsyncMock(return_value=True)
mock_browser = AsyncMock()
mock_browser.async_cancel = AsyncMock()
# Simulate a service being found during the sleep
captured_handler: list = []
def fake_browser(zc, types, handlers):
captured_handler.extend(handlers)
return mock_browser
from zeroconf import ServiceStateChange
async def fake_sleep(t):
# Fire the handler as if a device was discovered
for h in captured_handler:
h(None, "_shelly._tcp.local.", "Shelly1._shelly._tcp.local.", ServiceStateChange.Added)
mock_azc = AsyncMock()
mock_azc.__aenter__ = AsyncMock(return_value=mock_azc)
mock_azc.__aexit__ = AsyncMock(return_value=None)
mock_azc.zeroconf = MagicMock()
with patch("app.services.scanner._ZEROCONF_AVAILABLE", True), \
patch("app.services.scanner.AsyncZeroconf", return_value=mock_azc), \
patch("app.services.scanner.AsyncServiceBrowser", side_effect=fake_browser), \
patch("app.services.scanner.AsyncServiceInfo", return_value=mock_info), \
patch("asyncio.sleep", side_effect=fake_sleep):
result = await _mdns_discover(timeout=0.01)
assert len(result) == 1
assert result[0]["ip"] == "192.168.1.80"
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
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_run_scan_adds_nmap_devices_as_pending(mem_db):
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))
await session.commit()
nmap_hosts = [{"ip": "192.168.1.5", "hostname": "device.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))
devices = result.scalars().all()
assert any(d.ip == "192.168.1.5" for d in devices)
@pytest.mark.asyncio
async def test_run_scan_mdns_only_device_added(mem_db):
"""Devices found only by mDNS (not nmap) should appear in pending_devices."""
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))
await session.commit()
mdns_hosts = [{"ip": "192.168.1.80", "hostname": "shelly1.local.", "mac": None, "os": None, "open_ports": [{"port": 80, "protocol": "tcp", "banner": ""}]}]
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=mdns_hosts), \
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.80"))
device = result.scalar_one_or_none()
assert device is not None
assert device.status == "pending"
assert device.discovery_source == "mdns"
@pytest.mark.asyncio
async def test_run_scan_mdns_skipped_if_already_in_nmap(mem_db):
"""If nmap and mDNS both find the same IP, it should not be double-counted."""
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))
await session.commit()
shared_host = {"ip": "192.168.1.10", "hostname": "device.lan", "mac": None, "os": None, "open_ports": []}
async with mem_db() as session:
with patch("app.services.scanner._nmap_scan", return_value=[shared_host]), \
patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[shared_host]), \
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.10"))
devices = result.scalars().all()
assert len(devices) == 1 # not duplicated
@pytest.mark.asyncio
async def test_run_scan_skips_canvas_nodes(mem_db):
"""Hosts already approved onto the canvas must be skipped."""
from app.services.scanner import run_scan
run_id = _make_run_id()
async with mem_db() as session:
session.add(_make_scan_run(run_id))
canvas_node = Node(
id=str(uuid.uuid4()), label="PVE", type="proxmox",
ip="192.168.1.100", status="online",
)
session.add(canvas_node)
await session.commit()
nmap_hosts = [{"ip": "192.168.1.100", "hostname": "pve.lan", "mac": None, "os": None, "open_ports": []}]
async with mem_db() as session:
with patch("app.services.scanner._nmap_scan", return_value=nmap_hosts), \
patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[]), \
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock):
await run_scan(["192.168.1.0/24"], session, run_id)
async with mem_db() as session:
result = await session.execute(sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.100"))
assert result.scalar_one_or_none() is None
@pytest.mark.asyncio
async def test_run_scan_skips_hidden_devices(mem_db):
"""Hosts hidden by the user must not re-appear in pending."""
from app.services.scanner import run_scan
run_id = _make_run_id()
async with mem_db() as session:
session.add(_make_scan_run(run_id))
hidden = PendingDevice(ip="192.168.1.55", status="hidden")
session.add(hidden)
await session.commit()
nmap_hosts = [{"ip": "192.168.1.55", "hostname": None, "mac": None, "os": None, "open_ports": []}]
async with mem_db() as session:
with patch("app.services.scanner._nmap_scan", return_value=nmap_hosts), \
patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[]), \
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock):
await run_scan(["192.168.1.0/24"], session, run_id)
async with mem_db() as session:
result = await session.execute(
sa_select(PendingDevice).where(PendingDevice.ip == "192.168.1.55", PendingDevice.status == "pending")
)
assert result.scalar_one_or_none() is None
@pytest.mark.asyncio
async def test_run_scan_cancelled_marks_status_cancelled(mem_db):
"""Cancelling a running scan sets the ScanRun status to 'cancelled'."""
from app.services.scanner import request_cancel, run_scan
run_id = _make_run_id()
async with mem_db() as session:
session.add(_make_scan_run(run_id))
await session.commit()
request_cancel(run_id)
async with mem_db() as session:
with patch("app.services.scanner._nmap_scan", return_value=[]), \
patch("app.services.scanner._mdns_discover", new_callable=AsyncMock, return_value=[]), \
patch("app.api.routes.status.broadcast_scan_update", new_callable=AsyncMock):
await run_scan(["192.168.1.0/24"], session, run_id)
async with mem_db() as session:
run = await session.get(ScanRun, run_id)
assert run is not None
assert run.status == "cancelled"
+1 -95
View File
@@ -5,13 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from app.core.scheduler import (
_run_service_checks,
_run_status_checks,
set_service_checks_enabled,
start_scheduler,
stop_scheduler,
)
from app.core.scheduler import _run_status_checks, start_scheduler, stop_scheduler
from app.db.database import Base
from app.db.models import Node
@@ -147,7 +141,6 @@ def test_scheduler_uses_settings_interval():
with patch("app.core.scheduler.settings") as mock_settings, \
patch("app.core.scheduler.AsyncIOScheduler", return_value=mock_sched):
mock_settings.status_checker_interval = 45
mock_settings.service_check_enabled = False
start_scheduler()
_, kwargs = mock_sched.add_job.call_args
assert kwargs["seconds"] == 45
@@ -162,90 +155,3 @@ def test_start_and_stop_scheduler():
mock_sched.add_job.assert_called_once()
mock_sched.start.assert_called_once()
mock_sched.shutdown.assert_called_once()
# ---------------------------------------------------------------------------
# Service checks
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_run_service_checks_disabled_does_nothing(mem_db):
async with mem_db() as session:
session.add(_make_node(services=[{"port": 80, "protocol": "tcp", "service_name": "http"}]))
await session.commit()
with patch("app.core.scheduler.settings") as mock_settings, \
patch("app.core.scheduler.AsyncSessionLocal", mem_db), \
patch("app.services.status_checker.check_services", new_callable=AsyncMock) as mock_cs:
mock_settings.service_check_enabled = False
await _run_service_checks()
mock_cs.assert_not_called()
@pytest.mark.asyncio
async def test_run_service_checks_broadcasts_per_node(mem_db):
async with mem_db() as session:
node = _make_node(
ip="10.0.0.5",
services=[{"port": 80, "protocol": "tcp", "service_name": "http"}],
)
session.add(node)
await session.commit()
node_id = node.id
statuses = [{"port": 80, "protocol": "tcp", "status": "offline"}]
with patch("app.core.scheduler.settings") as mock_settings, \
patch("app.core.scheduler.AsyncSessionLocal", mem_db), \
patch("app.core.scheduler.check_services", new_callable=AsyncMock, return_value=statuses), \
patch("app.api.routes.status.broadcast_service_status", new_callable=AsyncMock) as mock_bcast:
mock_settings.service_check_enabled = True
await _run_service_checks()
mock_bcast.assert_awaited_once()
_, kwargs = mock_bcast.call_args
assert kwargs["node_id"] == node_id
assert kwargs["services"] == statuses
@pytest.mark.asyncio
async def test_run_service_checks_skips_nodes_without_services(mem_db):
async with mem_db() as session:
session.add(_make_node(ip="10.0.0.6", services=[]))
await session.commit()
with patch("app.core.scheduler.settings") as mock_settings, \
patch("app.core.scheduler.AsyncSessionLocal", mem_db), \
patch("app.core.scheduler.check_services", new_callable=AsyncMock) as mock_cs:
mock_settings.service_check_enabled = True
await _run_service_checks()
mock_cs.assert_not_called()
def test_set_service_checks_enabled_adds_and_removes_job():
mock_sched = MagicMock()
mock_sched.running = True
with patch("app.core.scheduler.scheduler", mock_sched), \
patch("app.core.scheduler.settings") as mock_settings:
mock_settings.service_check_interval = 300
# Enable: no existing job -> add
mock_sched.get_job.return_value = None
set_service_checks_enabled(True)
mock_sched.add_job.assert_called_once()
# Disable: existing job -> remove
mock_sched.get_job.return_value = MagicMock()
set_service_checks_enabled(False)
mock_sched.remove_job.assert_called_once_with("service_checks")
def test_start_scheduler_adds_service_job_when_enabled():
mock_sched = MagicMock()
with patch("app.core.scheduler.settings") as mock_settings, \
patch("app.core.scheduler.AsyncIOScheduler", return_value=mock_sched):
mock_settings.status_checker_interval = 60
mock_settings.service_check_enabled = True
mock_settings.service_check_interval = 300
start_scheduler()
job_ids = [kw.get("id") for _, kw in mock_sched.add_job.call_args_list]
assert "status_checks" in job_ids
assert "service_checks" in job_ids
-86
View File
@@ -1,86 +0,0 @@
"""Tests for GET/POST /api/v1/settings."""
from unittest.mock import patch
import pytest
from httpx import AsyncClient
@pytest.fixture
async def headers(client: AsyncClient):
res = await client.post("/api/v1/auth/login", json={"username": "admin", "password": "admin"})
token = res.json()["access_token"]
return {"Authorization": f"Bearer {token}"}
@pytest.mark.asyncio
async def test_get_settings_requires_auth(client: AsyncClient):
res = await client.get("/api/v1/settings")
assert res.status_code == 401
@pytest.mark.asyncio
async def test_get_settings_returns_interval(client: AsyncClient, headers):
res = await client.get("/api/v1/settings", headers=headers)
assert res.status_code == 200
data = res.json()
assert "interval_seconds" in data
assert isinstance(data["interval_seconds"], int)
@pytest.mark.asyncio
async def test_update_settings_saves_interval(client: AsyncClient, headers):
with patch("app.api.routes.settings.settings") as mock_settings:
mock_settings.status_checker_interval = 60
mock_settings.save_overrides = lambda: None
res = await client.post(
"/api/v1/settings",
json={"interval_seconds": 120},
headers=headers,
)
assert res.status_code == 200
assert res.json()["interval_seconds"] == 120
@pytest.mark.asyncio
async def test_update_settings_requires_auth(client: AsyncClient):
res = await client.post("/api/v1/settings", json={"interval_seconds": 30})
assert res.status_code == 401
@pytest.mark.asyncio
async def test_get_settings_returns_service_check_fields(client: AsyncClient, headers):
res = await client.get("/api/v1/settings", headers=headers)
data = res.json()
assert "service_check_enabled" in data
assert "service_check_interval" in data
assert isinstance(data["service_check_enabled"], bool)
assert isinstance(data["service_check_interval"], int)
@pytest.mark.asyncio
async def test_update_settings_saves_service_check_fields(client: AsyncClient, headers):
with patch("app.api.routes.settings.settings") as mock_settings:
mock_settings.save_overrides = lambda: None
res = await client.post(
"/api/v1/settings",
json={
"interval_seconds": 60,
"service_check_enabled": True,
"service_check_interval": 600,
},
headers=headers,
)
assert res.status_code == 200
body = res.json()
assert body["service_check_enabled"] is True
assert body["service_check_interval"] == 600
@pytest.mark.asyncio
async def test_update_settings_rejects_too_short_service_interval(client: AsyncClient, headers):
res = await client.post(
"/api/v1/settings",
json={"interval_seconds": 60, "service_check_enabled": True, "service_check_interval": 5},
headers=headers,
)
assert res.status_code == 422
-100
View File
@@ -1,100 +0,0 @@
"""API tests for /api/v1/stats/* (gethomepage widget)."""
from __future__ import annotations
from datetime import datetime, timezone
import pytest
from httpx import AsyncClient
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.db.models import Node, PendingDevice, ScanRun
@pytest.fixture(autouse=True)
def _reset_homepage_key():
original = settings.homepage_api_key
settings.homepage_api_key = ""
yield
settings.homepage_api_key = original
@pytest.mark.asyncio
async def test_summary_disabled_when_key_unset(client: AsyncClient) -> None:
res = await client.get("/api/v1/stats/summary")
assert res.status_code == 403
assert "disabled" in res.json()["detail"].lower()
@pytest.mark.asyncio
async def test_summary_rejects_missing_header(client: AsyncClient) -> None:
settings.homepage_api_key = "topsecret"
res = await client.get("/api/v1/stats/summary")
assert res.status_code == 403
@pytest.mark.asyncio
async def test_summary_rejects_wrong_key(client: AsyncClient) -> None:
settings.homepage_api_key = "topsecret"
res = await client.get(
"/api/v1/stats/summary", headers={"X-API-Key": "wrong"}
)
assert res.status_code == 403
@pytest.mark.asyncio
async def test_summary_empty_db(client: AsyncClient) -> None:
settings.homepage_api_key = "topsecret"
res = await client.get(
"/api/v1/stats/summary", headers={"X-API-Key": "topsecret"}
)
assert res.status_code == 200
body = res.json()
assert body == {
"nodes": 0,
"online": 0,
"offline": 0,
"unknown": 0,
"pending_devices": 0,
"zigbee_devices": 0,
"last_scan_at": None,
}
@pytest.mark.asyncio
async def test_summary_aggregates_counts(
client: AsyncClient, db_session: AsyncSession
) -> None:
settings.homepage_api_key = "topsecret"
finished = datetime(2026, 5, 14, 10, 0, tzinfo=timezone.utc)
db_session.add_all([
Node(type="server", label="A", status="online"),
Node(type="server", label="B", status="online"),
Node(type="server", label="C", status="offline"),
Node(type="server", label="D", status="unknown"),
Node(type="iot", label="Z1", status="online", ieee_address="0x1"),
Node(type="iot", label="Z2", status="online", ieee_address="0x2"),
PendingDevice(ip="10.0.0.1", status="pending"),
PendingDevice(ip="10.0.0.2", status="pending"),
PendingDevice(ip="10.0.0.3", status="hidden"), # excluded
ScanRun(status="success", finished_at=finished),
ScanRun(status="success",
finished_at=datetime(2026, 5, 13, 10, 0, tzinfo=timezone.utc)),
])
await db_session.commit()
res = await client.get(
"/api/v1/stats/summary", headers={"X-API-Key": "topsecret"}
)
assert res.status_code == 200
body = res.json()
assert body["nodes"] == 6
assert body["online"] == 4
assert body["offline"] == 1
assert body["unknown"] == 1
assert body["pending_devices"] == 2
assert body["zigbee_devices"] == 2
# SQLite returns naive datetimes; compare prefix only.
assert body["last_scan_at"] is not None
assert body["last_scan_at"].startswith("2026-05-14T10:00:00")
+1 -67
View File
@@ -5,13 +5,7 @@ import pytest
from fastapi.testclient import TestClient
from starlette.websockets import WebSocketDisconnect
from app.api.routes.status import (
_connections,
_drop,
broadcast_scan_update,
broadcast_service_status,
broadcast_status,
)
from app.api.routes.status import _connections, broadcast_scan_update, broadcast_status
from app.main import app
# ---------------------------------------------------------------------------
@@ -161,63 +155,3 @@ async def test_broadcast_no_connections():
assert len(_connections) == 0
await broadcast_status(node_id="n", status="online", checked_at="t")
await broadcast_scan_update(run_id="r", devices_found=0)
# ---------------------------------------------------------------------------
# broadcast_service_status
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_broadcast_service_status_payload():
received: list[str] = []
class FakeWS:
async def send_text(self, text: str) -> None:
received.append(text)
fake = FakeWS()
_connections.append(fake)
try:
await broadcast_service_status(
node_id="node-7",
services=[{"port": 80, "protocol": "tcp", "status": "offline"}],
checked_at="2024-01-01T00:00:00",
)
finally:
_drop(fake)
msg = json.loads(received[0])
assert msg["type"] == "service_status"
assert msg["node_id"] == "node-7"
assert msg["services"] == [{"port": 80, "protocol": "tcp", "status": "offline"}]
# ---------------------------------------------------------------------------
# _drop — idempotent connection removal (regression for double-remove crash)
# ---------------------------------------------------------------------------
def test_drop_is_idempotent():
"""Dropping a connection twice must not raise (was a ValueError crash)."""
class FakeWS:
pass
fake = FakeWS()
_connections.append(fake)
_drop(fake)
_drop(fake) # second drop must be a no-op
assert fake not in _connections
@pytest.mark.asyncio
async def test_broadcast_dead_connection_dropped_once_safely():
"""A send failure removes the dead socket without a double-remove crash."""
class DeadWS:
async def send_text(self, _: str) -> None:
raise RuntimeError("disconnected")
dead = DeadWS()
_connections.append(dead)
await broadcast_status(node_id="n", status="online", checked_at="t")
# A second broadcast must not raise even though dead is already gone.
await broadcast_status(node_id="n", status="online", checked_at="t")
assert dead not in _connections
+1 -287
View File
@@ -3,13 +3,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from app.services.status_checker import (
_ping,
_tcp_connect,
check_node,
check_service,
check_services,
)
from app.services.status_checker import _tcp_connect, check_node
# --- check_node dispatcher ---
@@ -155,172 +149,6 @@ async def test_check_node_exception_returns_offline():
assert result["response_time_ms"] is None
# --- _ping platform args ---
@pytest.mark.asyncio
async def test_ping_uses_unix_args_on_non_windows():
captured = {}
async def fake_exec(*args, **kwargs):
captured["args"] = args
proc = MagicMock()
proc.returncode = 0
proc.wait = AsyncMock()
return proc
with patch("app.services.status_checker.sys.platform", "linux"), \
patch("asyncio.create_subprocess_exec", side_effect=fake_exec):
await _ping("192.168.1.1")
assert "-c" in captured["args"]
assert "-W" in captured["args"]
assert "-n" not in captured["args"]
# 2 probes so a single dropped packet doesn't flap the node offline
c_idx = captured["args"].index("-c")
assert captured["args"][c_idx + 1] == "2"
# Linux: -W is in seconds; 2s is the intended timeout
w_idx = captured["args"].index("-W")
assert captured["args"][w_idx + 1] == "2"
# IPv4 target → no -6 flag
assert "-6" not in captured["args"]
@pytest.mark.asyncio
async def test_ping_uses_macos_millisecond_timeout():
"""macOS ping(8) -W is milliseconds, not seconds. 1ms would fail any RTT >1ms."""
captured = {}
async def fake_exec(*args, **kwargs):
captured["args"] = args
proc = MagicMock()
proc.returncode = 0
proc.wait = AsyncMock()
return proc
with patch("app.services.status_checker.sys.platform", "darwin"), \
patch("asyncio.create_subprocess_exec", side_effect=fake_exec):
await _ping("192.168.1.1")
assert "-c" in captured["args"]
assert "-W" in captured["args"]
w_idx = captured["args"].index("-W")
assert captured["args"][w_idx + 1] == "2000"
@pytest.mark.asyncio
async def test_ping_uses_windows_args_on_win32():
captured = {}
async def fake_exec(*args, **kwargs):
captured["args"] = args
proc = MagicMock()
proc.returncode = 0
proc.wait = AsyncMock()
return proc
with patch("app.services.status_checker.sys.platform", "win32"), \
patch("asyncio.create_subprocess_exec", side_effect=fake_exec):
await _ping("192.168.1.1")
assert "-n" in captured["args"]
assert "-w" in captured["args"]
assert "-c" not in captured["args"]
# --- _ping IPv6 support ---
@pytest.mark.asyncio
async def test_ping_ipv6_linux_uses_dash6():
"""IPv6-only devices (e.g. Alexa) need ping -6 on Linux."""
captured = {}
async def fake_exec(*args, **kwargs):
captured["args"] = args
proc = MagicMock()
proc.returncode = 0
proc.wait = AsyncMock()
return proc
with patch("app.services.status_checker.sys.platform", "linux"), \
patch("asyncio.create_subprocess_exec", side_effect=fake_exec):
await _ping("fe80::1")
assert "-6" in captured["args"]
assert captured["args"][-1] == "fe80::1"
@pytest.mark.asyncio
async def test_ping_ipv6_macos_uses_ping6():
"""macOS ships a separate ping6 binary for IPv6 targets."""
captured = {}
async def fake_exec(*args, **kwargs):
captured["args"] = args
proc = MagicMock()
proc.returncode = 0
proc.wait = AsyncMock()
return proc
with patch("app.services.status_checker.sys.platform", "darwin"), \
patch("asyncio.create_subprocess_exec", side_effect=fake_exec):
await _ping("2001:db8::1")
assert captured["args"][0] == "ping6"
@pytest.mark.asyncio
async def test_ping_ipv6_windows_uses_dash6():
captured = {}
async def fake_exec(*args, **kwargs):
captured["args"] = args
proc = MagicMock()
proc.returncode = 0
proc.wait = AsyncMock()
return proc
with patch("app.services.status_checker.sys.platform", "win32"), \
patch("asyncio.create_subprocess_exec", side_effect=fake_exec):
await _ping("2001:db8::1")
assert "-6" in captured["args"]
def test_is_ipv6_detection():
from app.services.status_checker import _is_ipv6
assert _is_ipv6("fe80::1") is True
assert _is_ipv6("2001:db8::1") is True
assert _is_ipv6("[2001:db8::1]") is True
assert _is_ipv6("192.168.1.1") is False
assert _is_ipv6("example.local") is False
# --- check_node target validation ---
@pytest.mark.asyncio
async def test_check_node_rejects_flag_like_target():
"""A target starting with '-' must never reach subprocess invocation."""
from app.services.status_checker import check_node
with patch("asyncio.create_subprocess_exec") as mock_exec:
result = await check_node("ping", "-O", None)
mock_exec.assert_not_called()
assert result["status"] == "unknown"
@pytest.mark.asyncio
async def test_check_node_rejects_flag_like_ip():
from app.services.status_checker import check_node
with patch("asyncio.create_subprocess_exec") as mock_exec:
result = await check_node("ping", None, "-O")
mock_exec.assert_not_called()
assert result["status"] == "unknown"
# --- _tcp_connect ---
@pytest.mark.asyncio
@@ -348,117 +176,3 @@ async def test_tcp_connect_os_error():
with patch("asyncio.open_connection", new_callable=AsyncMock, side_effect=OSError("refused")):
result = await _tcp_connect("192.168.1.1", 9999)
assert result is False
# --- check_service ---
@pytest.mark.asyncio
async def test_check_service_no_host_is_unknown():
assert await check_service({"port": 80, "protocol": "tcp", "service_name": "http"}, None) == "unknown"
@pytest.mark.asyncio
async def test_check_service_flag_host_is_unknown():
assert await check_service({"port": 80, "protocol": "tcp", "service_name": "http"}, "-O") == "unknown"
@pytest.mark.asyncio
async def test_check_service_udp_is_unknown():
assert await check_service({"port": 53, "protocol": "udp", "service_name": "dns"}, "10.0.0.1") == "unknown"
@pytest.mark.asyncio
async def test_check_service_portless_non_web_is_unknown():
svc = {"protocol": "tcp", "service_name": "thing"}
assert await check_service(svc, "10.0.0.1") == "unknown"
@pytest.mark.asyncio
async def test_check_service_web_uses_http_get():
captured = {}
async def fake_http_get(url, verify=False):
captured["url"] = url
return True
svc = {"port": 8080, "protocol": "tcp", "service_name": "http"}
with patch("app.services.status_checker._http_get", side_effect=fake_http_get):
result = await check_service(svc, "10.0.0.1")
assert result == "online"
assert captured["url"] == "http://10.0.0.1:8080"
@pytest.mark.asyncio
async def test_check_service_https_port_uses_https_scheme():
captured = {}
async def fake_http_get(url, verify=False):
captured["url"] = url
return True
svc = {"port": 443, "protocol": "tcp", "service_name": "web"}
with patch("app.services.status_checker._http_get", side_effect=fake_http_get):
await check_service(svc, "10.0.0.1")
assert captured["url"].startswith("https://")
@pytest.mark.asyncio
async def test_check_service_web_offline_when_http_fails():
svc = {"port": 80, "protocol": "tcp", "service_name": "http"}
with patch("app.services.status_checker._http_get", new_callable=AsyncMock, return_value=False):
assert await check_service(svc, "10.0.0.1") == "offline"
@pytest.mark.asyncio
async def test_check_service_non_http_port_is_unknown():
"""Non-HTTP ports (DB, mail, …) stay grey — no TCP check, no red flap."""
svc = {"port": 5432, "protocol": "tcp", "service_name": "postgres"}
with patch("app.services.status_checker._tcp_connect", new_callable=AsyncMock) as mock_tcp, \
patch("app.services.status_checker._http_get", new_callable=AsyncMock) as mock_http:
result = await check_service(svc, "10.0.0.1")
assert result == "unknown"
mock_tcp.assert_not_called()
mock_http.assert_not_called()
@pytest.mark.asyncio
async def test_check_service_ssh_port_22_is_unknown():
"""SSH (port 22) is never checked — keep it grey, not red/green."""
svc = {"port": 22, "protocol": "tcp", "service_name": "ssh"}
with patch("app.services.status_checker._tcp_connect", new_callable=AsyncMock) as mock_tcp:
result = await check_service(svc, "10.0.0.1")
assert result == "unknown"
mock_tcp.assert_not_called()
@pytest.mark.asyncio
async def test_check_service_ipv6_brackets_url_host():
captured = {}
async def fake_http_get(url, verify=False):
captured["url"] = url
return True
svc = {"port": 80, "protocol": "tcp", "service_name": "http"}
with patch("app.services.status_checker._http_get", side_effect=fake_http_get):
await check_service(svc, "2001:db8::1")
assert captured["url"] == "http://[2001:db8::1]:80"
@pytest.mark.asyncio
async def test_check_services_returns_status_per_service():
services = [
{"port": 80, "protocol": "tcp", "service_name": "http"},
{"port": 5432, "protocol": "tcp", "service_name": "postgres"},
]
with patch("app.services.status_checker._http_get", new_callable=AsyncMock, return_value=True):
results = await check_services("10.0.0.1", services)
assert results == [
{"port": 80, "protocol": "tcp", "status": "online"},
{"port": 5432, "protocol": "tcp", "status": "unknown"},
]
@pytest.mark.asyncio
async def test_check_services_empty_list():
assert await check_services("10.0.0.1", []) == []
-647
View File
@@ -1,647 +0,0 @@
"""API endpoint tests for /api/v1/zigbee/*."""
from __future__ import annotations
from unittest.mock import patch
import pytest
from httpx import AsyncClient
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
async def headers(client: AsyncClient):
res = await client.post("/api/v1/auth/login", json={"username": "admin", "password": "admin"})
token = res.json()["access_token"]
return {"Authorization": f"Bearer {token}"}
# ---------------------------------------------------------------------------
# /api/v1/zigbee/test-connection
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_test_connection_success(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.test_mqtt_connection") as mock_conn:
mock_conn.return_value = True
res = await client.post(
"/api/v1/zigbee/test-connection",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 200
data = res.json()
assert data["connected"] is True
assert "success" in data["message"].lower()
@pytest.mark.asyncio
async def test_test_connection_failure(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.test_mqtt_connection") as mock_conn:
mock_conn.side_effect = ConnectionError("Connection refused")
res = await client.post(
"/api/v1/zigbee/test-connection",
json={"mqtt_host": "bad-host", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 200
data = res.json()
assert data["connected"] is False
assert "refused" in data["message"].lower()
@pytest.mark.asyncio
async def test_test_connection_requires_auth(client: AsyncClient) -> None:
res = await client.post(
"/api/v1/zigbee/test-connection",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
)
assert res.status_code == 401
@pytest.mark.asyncio
async def test_test_connection_invalid_port(client: AsyncClient, headers: dict) -> None:
res = await client.post(
"/api/v1/zigbee/test-connection",
json={"mqtt_host": "localhost", "mqtt_port": 99999},
headers=headers,
)
assert res.status_code == 422 # pydantic validation error
# ---------------------------------------------------------------------------
# /api/v1/zigbee/import
# ---------------------------------------------------------------------------
_SAMPLE_NODES = [
{
"id": "0x00000000",
"label": "Coordinator",
"type": "zigbee_coordinator",
"ieee_address": "0x00000000",
"friendly_name": "Coordinator",
"device_type": "Coordinator",
"model": None,
"vendor": None,
"lqi": None,
"parent_id": None,
},
{
"id": "0x00000001",
"label": "router_1",
"type": "zigbee_router",
"ieee_address": "0x00000001",
"friendly_name": "router_1",
"device_type": "Router",
"model": "CC2530",
"vendor": "Texas Instruments",
"lqi": 230,
"parent_id": "0x00000000",
},
]
_SAMPLE_EDGES = [
{"source": "0x00000000", "target": "0x00000001"},
]
@pytest.mark.asyncio
async def test_import_success(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.return_value = (_SAMPLE_NODES, _SAMPLE_EDGES)
res = await client.post(
"/api/v1/zigbee/import",
json={
"mqtt_host": "localhost",
"mqtt_port": 1883,
"base_topic": "zigbee2mqtt",
},
headers=headers,
)
assert res.status_code == 200
data = res.json()
assert data["device_count"] == 2
assert len(data["nodes"]) == 2
assert len(data["edges"]) == 1
coordinator = next(n for n in data["nodes"] if n["type"] == "zigbee_coordinator")
assert coordinator["ieee_address"] == "0x00000000"
@pytest.mark.asyncio
async def test_import_with_credentials(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.return_value = ([], [])
res = await client.post(
"/api/v1/zigbee/import",
json={
"mqtt_host": "localhost",
"mqtt_port": 1883,
"mqtt_username": "admin",
"mqtt_password": "secret",
"base_topic": "z2m",
},
headers=headers,
)
assert res.status_code == 200
mock_fetch.assert_called_once_with(
mqtt_host="localhost",
mqtt_port=1883,
base_topic="z2m",
username="admin",
password="secret",
tls=False,
tls_insecure=False,
)
@pytest.mark.asyncio
async def test_import_connection_error_returns_502(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.side_effect = ConnectionError("broker unreachable")
res = await client.post(
"/api/v1/zigbee/import",
json={"mqtt_host": "bad-host", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 502
assert "broker unreachable" in res.json()["detail"]
@pytest.mark.asyncio
async def test_import_timeout_returns_504(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.side_effect = TimeoutError("timed out")
res = await client.post(
"/api/v1/zigbee/import",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 504
@pytest.mark.asyncio
async def test_import_malformed_payload_returns_422(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.side_effect = ValueError("malformed response")
res = await client.post(
"/api/v1/zigbee/import",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 422
@pytest.mark.asyncio
async def test_import_requires_auth(client: AsyncClient) -> None:
res = await client.post(
"/api/v1/zigbee/import",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
)
assert res.status_code == 401
@pytest.mark.asyncio
async def test_import_empty_network(client: AsyncClient, headers: dict) -> None:
"""An empty Zigbee network (coordinator only) is a valid response."""
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.return_value = ([], [])
res = await client.post(
"/api/v1/zigbee/import",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 200
data = res.json()
assert data["device_count"] == 0
assert data["nodes"] == []
assert data["edges"] == []
@pytest.mark.asyncio
async def test_import_missing_mqtt_host(client: AsyncClient, headers: dict) -> None:
res = await client.post(
"/api/v1/zigbee/import",
json={"mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 422
@pytest.mark.asyncio
async def test_import_with_tls_passes_flags(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.fetch_networkmap") as mock_fetch:
mock_fetch.return_value = ([], [])
res = await client.post(
"/api/v1/zigbee/import",
json={
"mqtt_host": "broker.example.com",
"mqtt_port": 8883,
"mqtt_tls": True,
},
headers=headers,
)
assert res.status_code == 200
kwargs = mock_fetch.call_args.kwargs
assert kwargs["tls"] is True
assert kwargs["tls_insecure"] is False
@pytest.mark.asyncio
async def test_import_tls_insecure_requires_tls(client: AsyncClient, headers: dict) -> None:
res = await client.post(
"/api/v1/zigbee/import",
json={
"mqtt_host": "broker.example.com",
"mqtt_port": 1883,
"mqtt_tls": False,
"mqtt_tls_insecure": True,
},
headers=headers,
)
assert res.status_code == 422
# ---------------------------------------------------------------------------
# /api/v1/zigbee/import-pending
# ---------------------------------------------------------------------------
_PENDING_NODES = [
{
"id": "0xCOORD",
"label": "Coordinator",
"type": "zigbee_coordinator",
"ieee_address": "0xCOORD",
"friendly_name": "Coordinator",
"device_type": "Coordinator",
"model": None,
"vendor": None,
"lqi": None,
"parent_id": None,
},
{
"id": "0xR1",
"label": "router_1",
"type": "zigbee_router",
"ieee_address": "0xR1",
"friendly_name": "router_1",
"device_type": "Router",
"model": "CC2530",
"vendor": "TI",
"lqi": 220,
"parent_id": "0xCOORD",
},
{
"id": "0xE1",
"label": "bulb_kitchen",
"type": "zigbee_enddevice",
"ieee_address": "0xE1",
"friendly_name": "bulb_kitchen",
"device_type": "EndDevice",
"model": "TRADFRI",
"vendor": "IKEA",
"lqi": 180,
"parent_id": "0xR1",
},
]
_PENDING_EDGES = [
{"source": "0xCOORD", "target": "0xR1"},
{"source": "0xR1", "target": "0xE1"},
]
@pytest.mark.asyncio
async def test_import_pending_endpoint_creates_zigbee_scan_run(
client: AsyncClient, headers: dict
) -> None:
"""Endpoint returns a ScanRun (kind=zigbee, status=running) immediately;
the actual networkmap fetch + pending persist runs in the background."""
from unittest.mock import AsyncMock
with patch(
"app.api.routes.zigbee._background_zigbee_import",
new_callable=AsyncMock,
):
res = await client.post(
"/api/v1/zigbee/import-pending",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
headers=headers,
)
assert res.status_code == 200
run = res.json()
assert run["kind"] == "zigbee"
assert run["status"] == "running"
assert run["ranges"] == ["localhost:1883"]
@pytest.mark.asyncio
async def test_persist_pending_import_creates_coordinator_and_pending(
db_session,
) -> None:
from app.api.routes.zigbee import _persist_pending_import
result = await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES)
assert result.device_count == 3
assert result.pending_created == 2
assert result.pending_updated == 0
assert result.coordinator is not None
assert result.coordinator.ieee_address == "0xCOORD"
assert result.coordinator_already_existed is False
assert result.links_recorded == 2
@pytest.mark.asyncio
async def test_persist_pending_import_idempotent_updates_existing(
db_session,
) -> None:
from app.api.routes.zigbee import _persist_pending_import
await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES)
bumped = [dict(n) for n in _PENDING_NODES]
bumped[1]["lqi"] = 99
result = await _persist_pending_import(db_session, bumped, _PENDING_EDGES)
assert result.pending_created == 0
assert result.pending_updated == 2
assert result.coordinator_already_existed is True
assert result.links_recorded == 2
@pytest.mark.asyncio
async def test_persist_pending_import_replaces_links(db_session) -> None:
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import PendingDeviceLink
await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES)
new_edges = [{"source": "0xCOORD", "target": "0xR1"}]
await _persist_pending_import(db_session, _PENDING_NODES[:2], new_edges)
rows = (await db_session.execute(select(PendingDeviceLink))).scalars().all()
assert len(rows) == 1
assert (rows[0].source_ieee, rows[0].target_ieee) == ("0xCOORD", "0xR1")
@pytest.mark.asyncio
async def test_persist_pending_import_sets_coordinator_properties(db_session) -> None:
"""Coordinator Node is created with IEEE/Vendor/Model/LQI in properties."""
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import Node
nodes_with_meta = [dict(n) for n in _PENDING_NODES]
nodes_with_meta[0]["vendor"] = "TI"
nodes_with_meta[0]["model"] = "CC2652"
await _persist_pending_import(db_session, nodes_with_meta, _PENDING_EDGES)
coord = (
await db_session.execute(select(Node).where(Node.ieee_address == "0xCOORD"))
).scalar_one()
keys = {p["key"]: p["value"] for p in coord.properties}
assert keys == {"IEEE": "0xCOORD", "Vendor": "TI", "Model": "CC2652"}
# New zigbee props default to hidden — user opts in from the right panel.
assert all(p["visible"] is False for p in coord.properties)
@pytest.mark.asyncio
async def test_persist_pending_import_skips_pending_for_approved_node(
db_session,
) -> None:
"""A device already approved as a canvas Node must not reappear in pending.
Its properties must still be refreshed with the latest Vendor/Model/LQI.
"""
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import Node, PendingDevice
# Simulate: router was approved earlier → exists as a canvas Node.
approved = Node(
label="router_1",
type="zigbee_router",
status="online",
check_method="none",
ieee_address="0xR1",
services=[],
properties=[],
)
db_session.add(approved)
await db_session.commit()
bumped = [dict(n) for n in _PENDING_NODES]
bumped[1]["lqi"] = 250 # new LQI from re-import
await _persist_pending_import(db_session, bumped, _PENDING_EDGES)
# No PendingDevice row was created for the approved router.
pendings = (
await db_session.execute(
select(PendingDevice).where(PendingDevice.ieee_address == "0xR1")
)
).scalars().all()
assert pendings == []
# Node properties got refreshed.
refreshed = (
await db_session.execute(select(Node).where(Node.ieee_address == "0xR1"))
).scalar_one()
keys = {p["key"]: p["value"] for p in refreshed.properties}
assert keys == {"IEEE": "0xR1", "Vendor": "TI", "Model": "CC2530", "LQI": "250"}
# Brand-new props on an existing Node start hidden.
assert all(p["visible"] is False for p in refreshed.properties)
@pytest.mark.asyncio
async def test_persist_pending_import_revives_orphaned_approved_device(
db_session,
) -> None:
"""Regression for #167: approve → delete node → re-import must re-list device.
When a device was approved (PendingDevice.status="approved") and its canvas
Node was later deleted, the orphaned "approved" row must be reset to
"pending" on re-import so it shows up in the Pending list again — instead of
being silently swallowed (re-import reports "found" but Pending stays empty).
"""
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import PendingDevice
# Simulate prior approve: a PendingDevice marked approved, but NO matching
# Node exists (the user deleted the canvas node afterwards).
orphan = PendingDevice(
ieee_address="0xR1",
friendly_name="router_1",
hostname="router_1",
suggested_type="zigbee_router",
device_subtype="Router",
model="CC2530",
vendor="TI",
lqi=220,
status="approved",
discovery_source="zigbee",
)
db_session.add(orphan)
await db_session.commit()
result = await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES)
# No new row created for 0xR1 — the existing one was updated/revived.
revived = (
await db_session.execute(
select(PendingDevice).where(PendingDevice.ieee_address == "0xR1")
)
).scalar_one()
assert revived.status == "pending"
# End device 0xE1 is brand new → created as pending; router was updated.
assert result.pending_created == 1
assert result.pending_updated == 1
# It is now visible to the Pending list (status filter == "pending").
listed = (
await db_session.execute(
select(PendingDevice).where(PendingDevice.status == "pending")
)
).scalars().all()
assert {p.ieee_address for p in listed} == {"0xR1", "0xE1"}
@pytest.mark.asyncio
async def test_persist_pending_import_keeps_hidden_hidden_on_reimport(
db_session,
) -> None:
"""A user-hidden device must stay hidden on re-import (not revived like #167)."""
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import PendingDevice
hidden = PendingDevice(
ieee_address="0xR1",
friendly_name="router_1",
suggested_type="zigbee_router",
device_subtype="Router",
status="hidden",
discovery_source="zigbee",
)
db_session.add(hidden)
await db_session.commit()
await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES)
still_hidden = (
await db_session.execute(
select(PendingDevice).where(PendingDevice.ieee_address == "0xR1")
)
).scalar_one()
assert still_hidden.status == "hidden"
@pytest.mark.asyncio
async def test_persist_pending_import_preserves_user_visibility(db_session) -> None:
"""If user has already made props visible, re-import must not flip them back."""
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import Node
approved = Node(
label="router_1",
type="zigbee_router",
status="online",
check_method="none",
ieee_address="0xR1",
services=[],
properties=[
{"key": "IEEE", "value": "0xR1", "icon": None, "visible": True},
{"key": "Vendor", "value": "TI", "icon": None, "visible": True},
{"key": "Custom", "value": "kept", "icon": None, "visible": True},
],
)
db_session.add(approved)
await db_session.commit()
bumped = [dict(n) for n in _PENDING_NODES]
bumped[1]["lqi"] = 99
bumped[1]["model"] = "CC2530"
await _persist_pending_import(db_session, bumped, _PENDING_EDGES)
refreshed = (
await db_session.execute(select(Node).where(Node.ieee_address == "0xR1"))
).scalar_one()
by_key = {p["key"]: p for p in refreshed.properties}
# Existing keys keep their visibility (True).
assert by_key["IEEE"]["visible"] is True
assert by_key["Vendor"]["visible"] is True
# New key arrives hidden.
assert by_key["Model"]["visible"] is False
assert by_key["LQI"]["visible"] is False
assert by_key["LQI"]["value"] == "99"
# Non-zigbee user-added prop is preserved untouched.
assert by_key["Custom"]["value"] == "kept"
assert by_key["Custom"]["visible"] is True
@pytest.mark.asyncio
async def test_persist_pending_import_refreshes_existing_coordinator_properties(
db_session,
) -> None:
from sqlalchemy import select
from app.api.routes.zigbee import _persist_pending_import
from app.db.models import Node
await _persist_pending_import(db_session, _PENDING_NODES, _PENDING_EDGES)
bumped = [dict(n) for n in _PENDING_NODES]
bumped[0]["vendor"] = "TI"
bumped[0]["model"] = "CC2652"
await _persist_pending_import(db_session, bumped, _PENDING_EDGES)
coord = (
await db_session.execute(select(Node).where(Node.ieee_address == "0xCOORD"))
).scalar_one()
keys = {p["key"]: p["value"] for p in coord.properties}
assert keys["Vendor"] == "TI"
assert keys["Model"] == "CC2652"
# Newly added keys on re-import default to hidden.
by_key = {p["key"]: p for p in coord.properties}
assert by_key["Vendor"]["visible"] is False
assert by_key["Model"]["visible"] is False
@pytest.mark.asyncio
async def test_import_pending_requires_auth(client: AsyncClient) -> None:
res = await client.post(
"/api/v1/zigbee/import-pending",
json={"mqtt_host": "localhost", "mqtt_port": 1883},
)
assert res.status_code == 401
@pytest.mark.asyncio
async def test_test_connection_with_tls(client: AsyncClient, headers: dict) -> None:
with patch("app.api.routes.zigbee.test_mqtt_connection") as mock_conn:
mock_conn.return_value = True
res = await client.post(
"/api/v1/zigbee/test-connection",
json={
"mqtt_host": "broker.example.com",
"mqtt_port": 8883,
"mqtt_tls": True,
"mqtt_tls_insecure": True,
},
headers=headers,
)
assert res.status_code == 200
kwargs = mock_conn.call_args.kwargs
assert kwargs["tls"] is True
assert kwargs["tls_insecure"] is True
-573
View File
@@ -1,573 +0,0 @@
"""Unit tests for zigbee_service: parser and hierarchy builder."""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import patch
import aiomqtt # noqa: F401
import pytest
from app.services.zigbee_service import (
_find_parent_router,
_z2m_type_to_homelable,
fetch_networkmap,
parse_networkmap,
)
from app.services.zigbee_service import (
test_mqtt_connection as _test_mqtt_connection,
)
# ---------------------------------------------------------------------------
# Helper builders — real Z2M `bridge/response/networkmap` shape
# (data.value.nodes + data.value.links)
# ---------------------------------------------------------------------------
def _make_node(
ieee: str,
device_type: str = "EndDevice",
friendly_name: str | None = None,
model: str | None = None,
vendor: str | None = None,
) -> dict[str, Any]:
entry: dict[str, Any] = {
"ieeeAddr": ieee,
"type": device_type,
"friendlyName": friendly_name or ieee,
}
if model or vendor:
entry["definition"] = {"model": model, "vendor": vendor}
return entry
def _make_link(source_ieee: str, target_ieee: str, lqi: int = 200) -> dict[str, Any]:
return {
"source": {"ieeeAddr": source_ieee},
"target": {"ieeeAddr": target_ieee},
"lqi": lqi,
}
def _wrap(nodes: list[dict[str, Any]], links: list[dict[str, Any]] | None = None) -> dict[str, Any]:
return {
"data": {
"type": "raw",
"routes": False,
"value": {"nodes": nodes, "links": links or []},
},
"status": "ok",
}
# ---------------------------------------------------------------------------
# _z2m_type_to_homelable
# ---------------------------------------------------------------------------
class TestZ2mTypeToHomelable:
def test_coordinator(self) -> None:
assert _z2m_type_to_homelable("Coordinator") == "zigbee_coordinator"
def test_router(self) -> None:
assert _z2m_type_to_homelable("Router") == "zigbee_router"
def test_enddevice(self) -> None:
assert _z2m_type_to_homelable("EndDevice") == "zigbee_enddevice"
def test_unknown_defaults_to_enddevice(self) -> None:
assert _z2m_type_to_homelable("Unknown") == "zigbee_enddevice"
# ---------------------------------------------------------------------------
# parse_networkmap
# ---------------------------------------------------------------------------
class TestParseNetworkmap:
def test_empty_payload(self) -> None:
nodes, edges = parse_networkmap({})
assert nodes == []
assert edges == []
def test_empty_value(self) -> None:
nodes, edges = parse_networkmap(_wrap([], []))
assert nodes == []
assert edges == []
def test_coordinator_only(self) -> None:
payload = _wrap([_make_node("0x0000000000000000", "Coordinator", "Coordinator")])
nodes, edges = parse_networkmap(payload)
assert len(nodes) == 1
assert nodes[0]["type"] == "zigbee_coordinator"
assert nodes[0]["ieee_address"] == "0x0000000000000000"
assert edges == []
def test_coordinator_router_enddevice(self) -> None:
coord_ieee = "0x0000000000000000"
router_ieee = "0x0000000000000001"
end_ieee = "0x0000000000000002"
payload = _wrap(
nodes=[
_make_node(coord_ieee, "Coordinator", "Coordinator"),
_make_node(router_ieee, "Router", "my_router"),
_make_node(end_ieee, "EndDevice"),
],
links=[
_make_link(coord_ieee, router_ieee),
_make_link(router_ieee, end_ieee),
],
)
nodes, edges = parse_networkmap(payload)
node_by_id = {n["id"]: n for n in nodes}
assert coord_ieee in node_by_id
assert router_ieee in node_by_id
assert end_ieee in node_by_id
assert node_by_id[coord_ieee]["type"] == "zigbee_coordinator"
assert node_by_id[router_ieee]["type"] == "zigbee_router"
assert node_by_id[end_ieee]["type"] == "zigbee_enddevice"
# Parent hierarchy
assert node_by_id[router_ieee]["parent_id"] == coord_ieee
assert node_by_id[end_ieee]["parent_id"] == router_ieee
assert len(edges) == 2
def test_no_duplicate_nodes(self) -> None:
ieee = "0x0000000000000001"
payload = _wrap(
nodes=[_make_node(ieee, "Router"), _make_node(ieee, "Router")],
)
nodes, _ = parse_networkmap(payload)
assert len(nodes) == 1
def test_edges_built_correctly(self) -> None:
coord = "0x0000"
router = "0x0001"
payload = _wrap(
nodes=[_make_node(coord, "Coordinator"), _make_node(router, "Router")],
links=[_make_link(coord, router)],
)
_, edges = parse_networkmap(payload)
assert len(edges) == 1
assert edges[0]["source"] == coord
assert edges[0]["target"] == router
def test_friendly_name_used_as_label(self) -> None:
payload = _wrap([_make_node("0xABCD", "EndDevice", "Living Room Sensor")])
nodes, _ = parse_networkmap(payload)
assert nodes[0]["label"] == "Living Room Sensor"
def test_enddevice_falls_back_to_coordinator_when_no_router(self) -> None:
coord = "0x0000"
end = "0x0003"
payload = _wrap([_make_node(coord, "Coordinator"), _make_node(end, "EndDevice")])
nodes, _ = parse_networkmap(payload)
end_node = next(n for n in nodes if n["id"] == end)
assert end_node["parent_id"] == coord
def test_missing_ieee_skipped(self) -> None:
payload = _wrap([{"type": "EndDevice"}]) # no ieeeAddr
nodes, edges = parse_networkmap(payload)
assert nodes == []
assert edges == []
def test_lqi_propagated_from_link_to_target_node(self) -> None:
coord = "0x0000"
end = "0x0001"
payload = _wrap(
nodes=[_make_node(coord, "Coordinator"), _make_node(end, "EndDevice")],
links=[_make_link(coord, end, lqi=180)],
)
nodes, _ = parse_networkmap(payload)
end_node = next(n for n in nodes if n["id"] == end)
assert end_node["lqi"] == 180
def test_definition_model_and_vendor_extracted(self) -> None:
payload = _wrap([
_make_node("0xAA", "EndDevice", "Sensor", model="WSDCGQ11LM", vendor="Aqara"),
])
nodes, _ = parse_networkmap(payload)
assert nodes[0]["model"] == "WSDCGQ11LM"
assert nodes[0]["vendor"] == "Aqara"
def test_legacy_shape_without_value_wrapper(self) -> None:
"""Some Z2M variants put nodes/links directly under data."""
payload = {"data": {"nodes": [_make_node("0x01", "Coordinator")], "links": []}}
nodes, _ = parse_networkmap(payload)
assert len(nodes) == 1
assert nodes[0]["type"] == "zigbee_coordinator"
def test_routes_bool_is_ignored(self) -> None:
"""`routes: false` echo from the request must not crash the parser."""
payload = {"data": {"routes": False, "type": "raw", "value": {"nodes": [], "links": []}}}
nodes, edges = parse_networkmap(payload)
assert nodes == []
assert edges == []
def test_malformed_nodes_not_list_raises(self) -> None:
with pytest.raises(ValueError, match="not a list"):
parse_networkmap({"data": {"value": {"nodes": "oops", "links": []}}})
def test_link_to_unknown_node_dropped(self) -> None:
payload = _wrap(
nodes=[_make_node("0x01", "Coordinator")],
links=[_make_link("0x01", "0xDEAD")], # 0xDEAD not in nodes
)
_, edges = parse_networkmap(payload)
assert edges == []
def test_bidirectional_links_yield_single_edge(self) -> None:
"""Z2M links are bidirectional — every pair appears twice. The output
must collapse to a single parent→child edge (no back-link, no dup)."""
coord = "0x0000"
router = "0x0001"
payload = _wrap(
nodes=[_make_node(coord, "Coordinator"), _make_node(router, "Router")],
links=[
_make_link(coord, router),
_make_link(router, coord), # reverse direction
],
)
_, edges = parse_networkmap(payload)
assert edges == [{"source": coord, "target": router}]
def test_router_mesh_siblings_dropped(self) -> None:
"""Router↔router mesh paths in `links` must NOT produce sibling edges
in the final tree. Each router gets exactly one edge from coordinator."""
coord = "0x0000"
r1 = "0x0001"
r2 = "0x0002"
payload = _wrap(
nodes=[
_make_node(coord, "Coordinator"),
_make_node(r1, "Router"),
_make_node(r2, "Router"),
],
links=[
_make_link(coord, r1),
_make_link(coord, r2),
_make_link(r1, r2), # mesh sibling — must be dropped
_make_link(r2, r1),
],
)
_, edges = parse_networkmap(payload)
pairs = {(e["source"], e["target"]) for e in edges}
assert pairs == {(coord, r1), (coord, r2)}
def test_coordinator_has_no_incoming_edge(self) -> None:
coord = "0x0000"
end = "0x0001"
payload = _wrap(
nodes=[_make_node(coord, "Coordinator"), _make_node(end, "EndDevice")],
links=[_make_link(end, coord)], # back-edge from end to coord
)
_, edges = parse_networkmap(payload)
# No edge should target the coordinator
assert all(e["target"] != coord for e in edges)
assert edges == [{"source": coord, "target": end}]
# ---------------------------------------------------------------------------
# _find_parent_router
# ---------------------------------------------------------------------------
class TestFindParentRouter:
def test_finds_router_as_source(self) -> None:
router_ids = {"r1"}
edges = [{"source": "r1", "target": "e1"}]
assert _find_parent_router("e1", router_ids, edges) == "r1"
def test_finds_router_as_target(self) -> None:
router_ids = {"r1"}
edges = [{"source": "e1", "target": "r1"}]
assert _find_parent_router("e1", router_ids, edges) == "r1"
def test_returns_none_when_no_router(self) -> None:
router_ids: set[str] = set()
edges = [{"source": "e1", "target": "e2"}]
assert _find_parent_router("e1", router_ids, edges) is None
def test_returns_none_empty_edges(self) -> None:
assert _find_parent_router("e1", {"r1"}, []) is None
# ---------------------------------------------------------------------------
# fetch_networkmap (integration-style with mocked aiomqtt)
# ---------------------------------------------------------------------------
SAMPLE_RESPONSE_PAYLOAD = {
"data": {
"type": "raw",
"routes": False,
"value": {
"nodes": [
{
"ieeeAddr": "0x00000000",
"type": "Coordinator",
"friendlyName": "Coordinator",
},
{
"ieeeAddr": "0x00000001",
"type": "Router",
"friendlyName": "router_1",
},
],
"links": [
{
"source": {"ieeeAddr": "0x00000000"},
"target": {"ieeeAddr": "0x00000001"},
"lqi": 230,
}
],
},
},
"status": "ok",
}
@pytest.mark.asyncio
async def test_fetch_networkmap_success() -> None:
"""fetch_networkmap returns parsed nodes/edges when MQTT responds normally."""
class _FakeMessage:
topic = "zigbee2mqtt/bridge/response/networkmap"
payload = json.dumps(SAMPLE_RESPONSE_PAYLOAD).encode()
_yielded = False
def __aiter__(self):
return self
async def __anext__(self):
if self._yielded:
raise StopAsyncIteration
self._yielded = True
return self
class _FakeClient:
async def __aenter__(self):
return self
async def __aexit__(self, *_):
pass
async def subscribe(self, _topic: str) -> None:
pass
async def publish(self, _topic: str, _payload: str) -> None:
pass
@property
def messages(self):
return _FakeMessage()
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
nodes, edges = await fetch_networkmap(
mqtt_host="localhost",
mqtt_port=1883,
base_topic="zigbee2mqtt",
)
assert any(n["type"] == "zigbee_coordinator" for n in nodes)
assert any(n["type"] == "zigbee_router" for n in nodes)
@pytest.mark.asyncio
async def test_fetch_networkmap_connection_error() -> None:
"""fetch_networkmap raises ConnectionError when MQTT broker is unreachable."""
class _FakeClient:
async def __aenter__(self):
raise Exception("Connection refused")
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
with pytest.raises(ConnectionError):
await fetch_networkmap(
mqtt_host="bad-host",
mqtt_port=1883,
base_topic="zigbee2mqtt",
)
@pytest.mark.asyncio
async def test_test_mqtt_connection_success() -> None:
class _FakeClient:
async def __aenter__(self):
return self
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
result = await _test_mqtt_connection("localhost", 1883)
assert result is True
@pytest.mark.asyncio
async def test_test_mqtt_connection_failure() -> None:
class _FakeClient:
async def __aenter__(self):
raise Exception("refused")
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
with pytest.raises(ConnectionError):
await _test_mqtt_connection("bad-host", 1883)
# ---------------------------------------------------------------------------
# TLS context
# ---------------------------------------------------------------------------
import ssl # noqa: E402
from app.services.zigbee_service import _build_tls_context # noqa: E402
def test_build_tls_context_secure_verifies_cert() -> None:
ctx = _build_tls_context(insecure=False)
assert ctx.check_hostname is True
assert ctx.verify_mode == ssl.CERT_REQUIRED
def test_build_tls_context_insecure_disables_verification() -> None:
ctx = _build_tls_context(insecure=True)
assert ctx.check_hostname is False
assert ctx.verify_mode == ssl.CERT_NONE
@pytest.mark.asyncio
async def test_test_mqtt_connection_passes_tls_context() -> None:
class _FakeClient:
async def __aenter__(self):
return self
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
await _test_mqtt_connection("host", 8883, tls=True)
kwargs = mock_aiomqtt.Client.call_args.kwargs
assert kwargs["tls_context"] is not None
assert kwargs["tls_context"].verify_mode == ssl.CERT_REQUIRED
@pytest.mark.asyncio
async def test_test_mqtt_connection_no_tls_context_when_disabled() -> None:
class _FakeClient:
async def __aenter__(self):
return self
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
await _test_mqtt_connection("host", 1883, tls=False)
assert mock_aiomqtt.Client.call_args.kwargs["tls_context"] is None
@pytest.mark.asyncio
async def test_test_mqtt_connection_insecure_passes_no_verify_context() -> None:
class _FakeClient:
async def __aenter__(self):
return self
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
await _test_mqtt_connection("host", 8883, tls=True, tls_insecure=True)
ctx = mock_aiomqtt.Client.call_args.kwargs["tls_context"]
assert ctx.verify_mode == ssl.CERT_NONE
assert ctx.check_hostname is False
# ---------------------------------------------------------------------------
# Sanitize MQTT errors
# ---------------------------------------------------------------------------
from app.services.zigbee_service import _sanitize_mqtt_error # noqa: E402
def test_sanitize_auth_error_does_not_leak_credentials() -> None:
msg = _sanitize_mqtt_error(
Exception("Not authorized: bad username or password for user=admin pwd=secret")
)
assert msg == "Authentication failed"
assert "admin" not in msg
assert "secret" not in msg
def test_sanitize_refused() -> None:
assert _sanitize_mqtt_error(Exception("Connection refused")) == "Connection refused by broker"
def test_sanitize_dns_failure_strips_host() -> None:
msg = _sanitize_mqtt_error(
Exception("[Errno 8] nodename nor servname provided, or not known: broker.internal.lan")
)
assert msg == "Broker hostname could not be resolved"
assert "broker.internal.lan" not in msg
def test_sanitize_tls_error() -> None:
assert _sanitize_mqtt_error(
Exception("[SSL: CERTIFICATE_VERIFY_FAILED] certificate verify failed")
) == "TLS handshake failed"
def test_sanitize_unknown_falls_back_to_generic() -> None:
msg = _sanitize_mqtt_error(Exception("mqtt://admin:hunter2@broker:1883 weird state"))
assert msg == "MQTT connection failed"
assert "hunter2" not in msg
assert "admin" not in msg
@pytest.mark.asyncio
async def test_fetch_networkmap_does_not_leak_creds_in_connection_error() -> None:
class _FakeClient:
async def __aenter__(self):
raise Exception("Not authorized: rejected mqtt://admin:hunter2@host")
async def __aexit__(self, *_):
pass
with patch("app.services.zigbee_service.aiomqtt") as mock_aiomqtt:
mock_aiomqtt.Client.return_value = _FakeClient()
mock_aiomqtt.MqttError = Exception
with pytest.raises(ConnectionError) as ei:
await fetch_networkmap(
mqtt_host="host", mqtt_port=1883, base_topic="zigbee2mqtt"
)
msg = str(ei.value)
assert "hunter2" not in msg
assert "admin" not in msg
assert msg == "Authentication failed"
-14
View File
@@ -24,20 +24,6 @@ services:
networks:
- homelable
mcp:
image: ghcr.io/pouzor/homelable-mcp:latest
restart: unless-stopped
ports:
- "8001:8001"
env_file:
- .env
environment:
BACKEND_URL: "http://backend:8000"
depends_on:
- backend
networks:
- homelable
volumes:
backend_data:
+2 -1
View File
@@ -7,8 +7,9 @@ services:
env_file:
- .env
environment:
# Override env_file: SQLite path must point inside the container volume
# Override env_file values that differ in Docker
SQLITE_PATH: /app/data/homelab.db
CORS_ORIGINS: '["http://localhost:3000"]'
volumes:
- backend_data:/app/data
networks:
-27
View File
@@ -1,27 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="128" height="128" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.7 KiB

-27
View File
@@ -1,27 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="16" height="16" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.7 KiB

-27
View File
@@ -1,27 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="256" height="256" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.7 KiB

-27
View File
@@ -1,27 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="32" height="32" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.7 KiB

-27
View File
@@ -1,27 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="512" height="512" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.7 KiB

-27
View File
@@ -1,27 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" width="64" height="64" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.7 KiB

-46
View File
@@ -1,46 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<!-- Background -->
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<!-- House body -->
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<!-- Floor line (subtle) -->
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<!-- Door -->
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<!-- Network lines (drawn under nodes) -->
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<!-- Center hub glow -->
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<!-- Center hub -->
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<!-- Left node -->
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<!-- Right node -->
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<!-- Top node -->
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</svg>

Before

Width:  |  Height:  |  Size: 1.9 KiB

-33
View File
@@ -1,33 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 280 72" width="640" height="164" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<g transform="translate(4, 4)">
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</g>
<text x="80" y="42" font-family="Inter, system-ui, -apple-system, sans-serif" font-weight="600" font-size="30" letter-spacing="-0.5">
<tspan fill="#e6edf3">Home</tspan><tspan fill="#00d4ff">lable</tspan>
</text>
<text x="81" y="58" font-family="Inter, system-ui, -apple-system, sans-serif" font-weight="400" font-size="11" fill="#8b949e" letter-spacing="0.5">HomeLab Visualizer</text>
</svg>

Before

Width:  |  Height:  |  Size: 2.1 KiB

-33
View File
@@ -1,33 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 280 72" width="360" height="92" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<g transform="translate(4, 4)">
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</g>
<text x="80" y="42" font-family="Inter, system-ui, -apple-system, sans-serif" font-weight="600" font-size="30" letter-spacing="-0.5">
<tspan fill="#e6edf3">Home</tspan><tspan fill="#00d4ff">lable</tspan>
</text>
<text x="81" y="58" font-family="Inter, system-ui, -apple-system, sans-serif" font-weight="400" font-size="11" fill="#8b949e" letter-spacing="0.5">HomeLab Visualizer</text>
</svg>

Before

Width:  |  Height:  |  Size: 2.1 KiB

-33
View File
@@ -1,33 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 280 72" width="200" height="51" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<g transform="translate(4, 4)">
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</g>
<text x="80" y="42" font-family="Inter, system-ui, -apple-system, sans-serif" font-weight="600" font-size="30" letter-spacing="-0.5">
<tspan fill="#e6edf3">Home</tspan><tspan fill="#00d4ff">lable</tspan>
</text>
<text x="81" y="58" font-family="Inter, system-ui, -apple-system, sans-serif" font-weight="400" font-size="11" fill="#8b949e" letter-spacing="0.5">HomeLab Visualizer</text>
</svg>

Before

Width:  |  Height:  |  Size: 2.1 KiB

-48
View File
@@ -1,48 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 280 72" fill="none">
<defs>
<radialGradient id="bg-glow" cx="50%" cy="35%" r="55%">
<stop offset="0%" stop-color="#00d4ff" stop-opacity="0.08"/>
<stop offset="100%" stop-color="#0d1117" stop-opacity="0"/>
</radialGradient>
<filter id="node-glow">
<feGaussianBlur stdDeviation="1.5" result="blur"/>
<feMerge><feMergeNode in="blur"/><feMergeNode in="SourceGraphic"/></feMerge>
</filter>
</defs>
<!-- Icon (72×72, scaled from 64 viewBox) -->
<g transform="translate(4, 4) scale(1)">
<circle cx="32" cy="32" r="32" fill="#0d1117"/>
<circle cx="32" cy="32" r="32" fill="url(#bg-glow)"/>
<path d="M32 11 L53 30 L48 30 L48 53 L16 53 L16 30 L11 30 Z"
fill="#161b22" stroke="#00d4ff" stroke-width="1.5" stroke-linejoin="round"/>
<line x1="16" y1="30" x2="48" y2="30" stroke="#00d4ff" stroke-width="0.5" opacity="0.25"/>
<rect x="27" y="40" width="10" height="13" rx="1.5"
fill="#0d1117" stroke="#00d4ff" stroke-width="1" opacity="0.9"/>
<line x1="32" y1="23" x2="32" y2="30" stroke="#a855f7" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="21" y1="38" x2="29" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<line x1="43" y1="38" x2="35" y2="33" stroke="#39d353" stroke-width="1.2" opacity="0.7" stroke-linecap="round"/>
<circle cx="32" cy="33" r="5" fill="#00d4ff" opacity="0.12"/>
<circle cx="32" cy="33" r="3" fill="#00d4ff" filter="url(#node-glow)"/>
<circle cx="21" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="43" cy="38" r="2.5" fill="#39d353" filter="url(#node-glow)"/>
<circle cx="32" cy="23" r="2" fill="#a855f7" filter="url(#node-glow)"/>
</g>
<!-- Text -->
<text x="80" y="42"
font-family="Inter, system-ui, -apple-system, sans-serif"
font-weight="600"
font-size="30"
letter-spacing="-0.5">
<tspan fill="#e6edf3">Home</tspan><tspan fill="#00d4ff">lable</tspan>
</text>
<!-- Subtitle -->
<text x="81" y="58"
font-family="Inter, system-ui, -apple-system, sans-serif"
font-weight="400"
font-size="11"
fill="#8b949e"
letter-spacing="0.5">HomeLab Visualizer</text>
</svg>

Before

Width:  |  Height:  |  Size: 2.3 KiB

-130
View File
@@ -1,130 +0,0 @@
# Zigbee2MQTT Network Map Importer
This feature lets you connect Homelable to your MQTT broker, fetch the Zigbee2MQTT network topology, and drop all Zigbee devices onto the canvas as typed nodes with proper hierarchy.
---
## Feature Overview
- **Automatic device discovery** — Requests the Z2M networkmap via the MQTT bridge API and parses the full device list
- **Typed nodes** — Devices are mapped to three homelable node types:
- `zigbee_coordinator` — The Zigbee coordinator (hub)
- `zigbee_router` — Mains-powered router devices
- `zigbee_enddevice` — Battery-powered end devices (sensors, bulbs, etc.)
- **Hierarchy** — `parent_id` is set automatically: coordinator → routers → end devices
- **LQI display** — Link Quality Indicator is stored as a node property
- **IoT edges** — Links between devices are added as `IoT / Zigbee` edge type
---
## Prerequisites
1. A running **MQTT broker** (e.g. Mosquitto) accessible from your Homelable host
2. **Zigbee2MQTT** connected to the broker and running
3. Z2M must respond to networkmap requests on:
- **Request topic:** `<base_topic>/bridge/request/networkmap`
- **Response topic:** `<base_topic>/bridge/response/networkmap`
- The default base topic is `zigbee2mqtt`
---
## Step-by-step Usage
### 1. Open the Zigbee Import dialog
Click **Zigbee Import** in the left sidebar (below "Scan Network").
### 2. Configure the MQTT connection
| Field | Default | Description |
|---|---|---|
| Broker Host | — | IP or hostname of your MQTT broker |
| Port | 1883 | MQTT broker port |
| Base Topic | `zigbee2mqtt` | Zigbee2MQTT base topic |
| Username | _(optional)_ | MQTT username if authentication is enabled |
| Password | _(optional)_ | MQTT password |
### 3. Test the connection (optional)
Click **Test Connection** to verify broker reachability before fetching devices.
A green indicator confirms success; red shows the error message from the broker.
### 4. Fetch devices
Click **Fetch Devices**. Homelable will:
1. Connect to the broker
2. Subscribe to the response topic
3. Publish `{"type": "raw", "routes": false}` to the request topic
4. Wait up to 60 seconds for the network map response (large meshes can take 30 s+)
5. Parse and group devices by type
### 5. Select and add to canvas
Devices are grouped by type (Coordinator / Router / End Device).
Use the checkboxes to select which devices to add, then click **Add N to Canvas**.
> **Tip:** All devices are selected by default. Uncheck any you don't want.
### 6. Arrange on the canvas
Devices are placed in a grid at the top-right of the canvas.
Use **Auto Layout** (toolbar) to re-arrange the full canvas, or drag nodes manually.
---
## MQTT Configuration Tips
### Mosquitto without authentication
```
listener 1883
allow_anonymous true
```
### Mosquitto with password file
```
listener 1883
password_file /etc/mosquitto/passwd
```
Create a user:
```bash
mosquitto_passwd -c /etc/mosquitto/passwd <username>
```
### Zigbee2MQTT `configuration.yaml`
```yaml
mqtt:
base_topic: zigbee2mqtt
server: mqtt://localhost:1883
# user: mqtt_user
# password: mqtt_password
```
---
## Supported Z2M Versions
The networkmap bridge API is available in **Zigbee2MQTT 1.x and 2.x**.
Tested against Z2M 1.35+ and 2.x.
The importer uses the `raw` topology format (`routes: false`) which is the most widely supported mode.
---
## Troubleshooting
| Symptom | Cause | Fix |
|---|---|---|
| "Connection refused" | Broker unreachable | Check host/port, firewall rules |
| "Timed out waiting for networkmap" | Z2M not running or wrong base_topic | Verify Z2M is connected, check base_topic setting |
| 0 devices returned | Z2M has no devices paired | Pair at least one device first |
| "Malformed networkmap response" | Z2M returned unexpected format | Check Z2M version; open an issue |
---
## Screenshots
_(Screenshots will be added in a future release)_
+906 -1231
View File
File diff suppressed because it is too large Load Diff
+4 -10
View File
@@ -1,7 +1,7 @@
{
"name": "frontend",
"private": true,
"version": "2.5.1",
"version": "1.3.3",
"type": "module",
"scripts": {
"dev": "vite",
@@ -19,10 +19,9 @@
"@fontsource-variable/geist": "^5.2.8",
"@fontsource-variable/inter": "^5.2.8",
"@fontsource/jetbrains-mono": "^5.2.8",
"@radix-ui/react-tooltip": "^1.2.8",
"@types/js-yaml": "^4.0.9",
"@xyflow/react": "^12.10.1",
"axios": "^1.15.2",
"axios": "^1.13.6",
"class-variance-authority": "^0.7.1",
"clsx": "^2.1.1",
"dagre": "^0.8.5",
@@ -37,11 +36,6 @@
"tw-animate-css": "^1.4.0",
"zustand": "^5.0.11"
},
"overrides": {
"hono": "^4.12.21",
"esbuild": "^0.28.1",
"form-data": "^4.0.6"
},
"devDependencies": {
"@eslint/js": "^9.39.1",
"@tailwindcss/vite": "^4.2.1",
@@ -59,11 +53,11 @@
"eslint-plugin-react-refresh": "^0.4.24",
"globals": "^16.5.0",
"jsdom": "^28.1.0",
"lucide-react": "^1.7.0",
"lucide-react": "^0.577.0",
"tailwindcss": "^4.2.1",
"typescript": "~5.9.3",
"typescript-eslint": "^8.48.0",
"vite": "^7.3.5",
"vite": "^7.3.1",
"vitest": "^4.0.18"
}
}
@@ -1,27 +0,0 @@
#!/usr/bin/env node
// Regenerate frontend/src/data/dashboardIcons.json from the upstream
// homarr-labs/dashboard-icons repo. Run manually to refresh the manifest.
//
// node scripts/fetch-dashboard-icons.mjs
import { writeFileSync, mkdirSync } from 'node:fs'
import { dirname, resolve } from 'node:path'
import { fileURLToPath } from 'node:url'
const TREE_URL = 'https://raw.githubusercontent.com/homarr-labs/dashboard-icons/main/tree.json'
const OUT = resolve(dirname(fileURLToPath(import.meta.url)), '../src/data/dashboardIcons.json')
const res = await fetch(TREE_URL)
if (!res.ok) {
console.error(`fetch failed: ${res.status} ${res.statusText}`)
process.exit(1)
}
const tree = await res.json()
const slugs = (tree.svg ?? [])
.filter((f) => f.endsWith('.svg'))
.map((f) => f.slice(0, -4))
.sort()
mkdirSync(dirname(OUT), { recursive: true })
writeFileSync(OUT, JSON.stringify(slugs))
console.log(`wrote ${slugs.length} slugs → ${OUT}`)
+75 -461
View File
@@ -4,9 +4,8 @@ import { type Node } from '@xyflow/react'
import { applyDagreLayout } from '@/utils/layout'
import { serializeNode, serializeEdge, deserializeApiNode, deserializeApiEdge, type ApiNode, type ApiEdge } from '@/utils/canvasSerializer'
import { generateUUID } from '@/utils/uuid'
import { resolveVirtualEdgeParent } from '@/utils/virtualEdgeParent'
import { generateMarkdownTable } from '@/utils/exportMarkdown'
import { ExportModal } from '@/components/modals/ExportModal'
import { exportToPng } from '@/utils/export'
import { exportCanvasToYaml, downloadYaml } from '@/utils/exportYaml'
import { parseYamlToCanvas } from '@/utils/importYaml'
import { TooltipProvider } from '@/components/ui/tooltip'
@@ -20,142 +19,70 @@ import { LoginPage } from '@/components/LoginPage'
import { NodeModal } from '@/components/modals/NodeModal'
import { EdgeModal } from '@/components/modals/EdgeModal'
import { ScanConfigModal } from '@/components/modals/ScanConfigModal'
import { SettingsModal } from '@/components/modals/SettingsModal'
import { ZigbeeImportModal } from '@/components/zigbee/ZigbeeImportModal'
import { GroupRectModal, type GroupRectFormData } from '@/components/modals/GroupRectModal'
import { TextModal, type TextFormData } from '@/components/modals/TextModal'
import { ThemeModal } from '@/components/modals/ThemeModal'
import { SearchModal } from '@/components/modals/SearchModal'
import { PendingDevicesModal } from '@/components/modals/PendingDevicesModal'
import { ScanHistoryModal } from '@/components/modals/ScanHistoryModal'
import { ShortcutsModal } from '@/components/modals/ShortcutsModal'
import { ConfirmAddToGroupModal } from '@/components/modals/ConfirmAddToGroupModal'
import { useCanvasStore } from '@/stores/canvasStore'
import { useDesignStore } from '@/stores/designStore'
import { useAuthStore } from '@/stores/authStore'
import { useThemeStore } from '@/stores/themeStore'
import { canvasApi, designsApi, liveviewApi } from '@/api/client'
import { canvasApi } from '@/api/client'
import { demoNodes, demoEdges } from '@/utils/demoData'
import { useStatusPolling } from '@/hooks/useStatusPolling'
import type { NodeData, EdgeData, CustomStyleDef } from '@/types'
import type { ZigbeeNode, ZigbeeEdge } from '@/components/zigbee/types'
import type { NodeData, EdgeData } from '@/types'
const STANDALONE = import.meta.env.VITE_STANDALONE === 'true'
const STANDALONE_STORAGE_KEY = 'homelable_canvas'
export default function App() {
const { loadCanvas, markSaved, markUnsaved, selectedNodeId, selectedNodeIds, addNode, updateNode, deleteNode, onConnect, updateEdge, deleteEdge, setProxmoxContainerMode, setNodeZIndex, editingGroupRectId, setEditingGroupRectId, editingTextId, setEditingTextId, nodes, edges, snapshotHistory, undo, redo, addToGroup, addToContainer } = useCanvasStore()
const { loadCanvas, markSaved, markUnsaved, selectedNodeId, addNode, updateNode, deleteNode, onConnect, updateEdge, deleteEdge, setProxmoxContainerMode, setNodeZIndex, editingGroupRectId, setEditingGroupRectId, nodes, edges, snapshotHistory, undo, redo, copySelectedNodes, pasteNodes } = useCanvasStore()
const canvasRef = useRef<HTMLDivElement>(null)
const { isAuthenticated } = useAuthStore()
const { activeTheme, setTheme, customStyle, setCustomStyle } = useThemeStore()
const { activeDesignId, setDesigns, setActiveDesign } = useDesignStore()
const { activeTheme, setTheme } = useThemeStore()
useStatusPolling()
const [themeModalOpen, setThemeModalOpen] = useState(false)
const [searchOpen, setSearchOpen] = useState(false)
const [scanHistoryOpen, setScanHistoryOpen] = useState(false)
const [pendingModalOpen, setPendingModalOpen] = useState(false)
const [pendingModalStatus, setPendingModalStatus] = useState<'pending' | 'hidden'>('pending')
const [pendingHighlightId, setPendingHighlightId] = useState<string | undefined>(undefined)
const openPendingModal = useCallback((deviceId?: string, status: 'pending' | 'hidden' = 'pending') => {
setPendingHighlightId(undefined)
setPendingModalStatus(status)
setPendingModalOpen(true)
if (deviceId) setTimeout(() => setPendingHighlightId(deviceId), 0)
}, [])
const [shortcutsOpen, setShortcutsOpen] = useState(false)
const [addNodeOpen, setAddNodeOpen] = useState(false)
const [addGroupRectOpen, setAddGroupRectOpen] = useState(false)
const [addTextOpen, setAddTextOpen] = useState(false)
const [editNodeId, setEditNodeId] = useState<string | null>(null)
const [pendingConnection, setPendingConnection] = useState<Connection | null>(null)
const [pendingGroupAdd, setPendingGroupAdd] = useState<{ nodeId: string; groupId: string } | null>(null)
const [pendingContainerAdd, setPendingContainerAdd] = useState<{ nodeId: string; containerId: string } | null>(null)
const [editEdgeId, setEditEdgeId] = useState<string | null>(null)
const [scanConfigOpen, setScanConfigOpen] = useState(false)
const [settingsOpen, setSettingsOpen] = useState(false)
const [exportModalOpen, setExportModalOpen] = useState(false)
const [zigbeeImportOpen, setZigbeeImportOpen] = useState(false)
// Declare handleSave before the Ctrl+S effect so it is in scope.
// Returns true on success, false on failure — the design-switch effect relies
// on this to avoid loading (and clobbering) the canvas when a save fails.
const handleSave = useCallback(async (designIdOverride?: string): Promise<boolean> => {
// Declare handleSave before the Ctrl+S effect so it is in scope
const handleSave = useCallback(async () => {
try {
const saveDesignId = designIdOverride ?? activeDesignId
if (STANDALONE) {
localStorage.setItem(STANDALONE_STORAGE_KEY, JSON.stringify({ nodes, edges, theme_id: activeTheme, custom_style: customStyle }))
localStorage.setItem(STANDALONE_STORAGE_KEY, JSON.stringify({ nodes, edges, theme_id: activeTheme }))
markSaved()
toast.success('Canvas saved')
return true
return
}
const nodesToSave = nodes.map(serializeNode)
const edgesToSave = edges.map(serializeEdge)
await canvasApi.save({ nodes: nodesToSave, edges: edgesToSave, viewport: { theme_id: activeTheme }, custom_style: customStyle, design_id: saveDesignId })
await canvasApi.save({ nodes: nodesToSave, edges: edgesToSave, viewport: { theme_id: activeTheme } })
markSaved()
toast.success('Canvas saved')
return true
} catch {
toast.error('Save failed')
return false
}
}, [nodes, edges, markSaved, activeTheme, customStyle, activeDesignId])
}, [nodes, edges, markSaved, activeTheme])
// Keep a ref so the keydown handler always calls the latest version
const handleSaveRef = useRef(handleSave)
useEffect(() => { handleSaveRef.current = handleSave }, [handleSave])
const loadCanvasFromApi = useCallback(async (designId?: string) => {
try {
const res = await canvasApi.load(designId)
const { nodes: apiNodes, edges: apiEdges } = res.data
if (apiNodes.length > 0) {
const proxmoxContainerMap = new Map<string, boolean>(
(apiNodes as ApiNode[])
.filter((n) => n.type === 'group' || n.container_mode === true)
.map((n) => [n.id, true])
)
const rfNodes = (apiNodes as ApiNode[]).map((n) => deserializeApiNode(n, proxmoxContainerMap))
const rfEdges = (apiEdges as ApiEdge[]).map(deserializeApiEdge)
const savedTheme = res.data.viewport?.theme_id
if (savedTheme) setTheme(savedTheme)
if (res.data.custom_style) setCustomStyle(res.data.custom_style as CustomStyleDef)
loadCanvas(rfNodes, rfEdges)
} else {
loadCanvas(demoNodes, demoEdges)
}
} catch {
loadCanvas(demoNodes, demoEdges)
}
}, [loadCanvas, setTheme, setCustomStyle])
const loadDesignsAndCanvas = useCallback(async () => {
if (STANDALONE) return
try {
const res = await designsApi.list()
const loadedDesigns = res.data
setDesigns(loadedDesigns)
const targetId = activeDesignId ?? loadedDesigns[0]?.id
if (targetId) {
setActiveDesign(targetId)
await loadCanvasFromApi(targetId)
}
} catch {
// If API fails (e.g. fresh DB with no designs), fall back to demo data
loadCanvas(demoNodes, demoEdges)
}
}, [setDesigns, setActiveDesign, loadCanvasFromApi, activeDesignId, loadCanvas])
// Load canvas on auth (or immediately in standalone mode)
useEffect(() => {
if (STANDALONE) {
try {
const saved = localStorage.getItem(STANDALONE_STORAGE_KEY)
if (saved) {
const { nodes: savedNodes, edges: savedEdges, theme_id, custom_style } = JSON.parse(saved)
const { nodes: savedNodes, edges: savedEdges, theme_id } = JSON.parse(saved)
if (theme_id) setTheme(theme_id)
if (custom_style) setCustomStyle(custom_style)
loadCanvas(savedNodes, savedEdges)
} else {
loadCanvas(demoNodes, demoEdges)
@@ -166,59 +93,37 @@ export default function App() {
return
}
if (!isAuthenticated) return
loadDesignsAndCanvas()
}, [isAuthenticated, loadCanvas, setTheme, setCustomStyle]) // only on auth change, not design change
// Reload canvas when active design changes (after initial load)
const initialLoadDone = useRef(false)
const prevDesignRef = useRef<string | null>(null)
// Set while we programmatically revert activeDesignId after a failed save, so
// the re-entrant effect run skips save/load and just re-syncs the refs.
const revertingRef = useRef(false)
useEffect(() => {
if (revertingRef.current) {
revertingRef.current = false
prevDesignRef.current = activeDesignId
return
}
if (!STANDALONE && isAuthenticated && activeDesignId && initialLoadDone.current) {
const oldId = prevDesignRef.current
// If the previous design was deleted (no longer in the list), don't try to
// save into it — just load the newly-selected design.
const oldStillExists = oldId ? useDesignStore.getState().designs.some((d) => d.id === oldId) : false
if (oldId && oldId !== activeDesignId && oldStillExists) {
// Save current (old) canvas data under the old design ID before switching.
// We call handleSave directly (not via ref) so it runs in this effect's
// closure where activeDesignId is already the NEW value — the override
// ensures data is stored under the correct design_id.
const targetId = activeDesignId
handleSave(oldId).then((ok) => {
if (ok) {
loadCanvasFromApi(targetId)
} else {
// Save failed: don't load the new design — that would overwrite the
// unsaved in-memory canvas. Revert the selection back to the old
// design so the UI matches the data still on screen.
toast.error('Switch cancelled — unsaved changes kept')
revertingRef.current = true
setActiveDesign(oldId)
}
})
} else {
loadCanvasFromApi(activeDesignId)
}
}
if (activeDesignId) {
prevDesignRef.current = activeDesignId
initialLoadDone.current = true
}
}, [activeDesignId])
canvasApi.load()
.then((res) => {
const { nodes: apiNodes, edges: apiEdges } = res.data
if (apiNodes.length > 0) {
// Build a map of proxmox container mode to know if children should be nested
const proxmoxContainerMap = new Map<string, boolean>(
(apiNodes as ApiNode[])
.filter((n) => n.type === 'proxmox')
.map((n) => [n.id, n.container_mode !== false])
)
const rfNodes = (apiNodes as ApiNode[]).map((n) => deserializeApiNode(n, proxmoxContainerMap))
const rfEdges = (apiEdges as ApiEdge[]).map(deserializeApiEdge)
const savedTheme = res.data.viewport?.theme_id
if (savedTheme) setTheme(savedTheme)
loadCanvas(rfNodes, rfEdges)
} else {
loadCanvas(demoNodes, demoEdges)
}
})
.catch(() => loadCanvas(demoNodes, demoEdges))
}, [isAuthenticated, loadCanvas, setTheme])
// Keep refs for store actions so keydown handler is always up-to-date without re-registering
const undoRef = useRef(undo)
const redoRef = useRef(redo)
const copyRef = useRef(copySelectedNodes)
const pasteRef = useRef(pasteNodes)
useEffect(() => { undoRef.current = undo }, [undo])
useEffect(() => { redoRef.current = redo }, [redo])
useEffect(() => { copyRef.current = copySelectedNodes }, [copySelectedNodes])
useEffect(() => { pasteRef.current = pasteNodes }, [pasteNodes])
// Global keyboard shortcuts
useEffect(() => {
@@ -232,8 +137,8 @@ export default function App() {
if (ctrl && e.key === 'z') { e.preventDefault(); undoRef.current(); return }
if (ctrl && (e.key === 'y' || (e.shiftKey && e.key === 'z'))) { e.preventDefault(); redoRef.current(); return }
if (ctrl && e.key === 'k') { e.preventDefault(); setSearchOpen(true); return }
// Copy/paste (Ctrl/Cmd+C/V) handled in CanvasContainer so paste can place
// nodes under the cursor / viewport center.
if (ctrl && e.key === 'c' && !isInput) { copyRef.current(); return }
if (ctrl && e.key === 'v' && !isInput) { pasteRef.current(); return }
if (e.key === '?' && !isInput) { setShortcutsOpen(true); return }
}
window.addEventListener('keydown', handler)
@@ -243,18 +148,11 @@ export default function App() {
const handleAddNode = useCallback((data: Partial<NodeData>) => {
snapshotHistory()
const id = generateUUID()
const isContainerNode = data.container_mode === true
const isProxmox = data.type === 'proxmox'
const parentNode = data.parent_id ? nodes.find((n) => n.id === data.parent_id) : null
// Only nest when the parent is an actual container. For a non-container
// parent the LXC/VM stays a free node (linked by a virtual edge) — setting
// extent:'parent' on a non-container would trap it inside the parent's tiny
// bounding box with no way to drag it out (issue #205 follow-up).
const nestInParent = !!parentNode?.data.container_mode
// Seed an ABSOLUTE position near the container's top-left; addNode converts
// it to container-relative. addNode is the single authority for parentId /
// extent, so we don't set them here.
const position = nestInParent && parentNode
? { x: parentNode.position.x + 20, y: parentNode.position.y + 50 }
// Children position is relative to parent; place near top-left with padding
const position = parentNode
? { x: 20, y: 50 }
: { x: 300, y: 300 }
const newNode: Node<NodeData> = {
@@ -262,7 +160,8 @@ export default function App() {
type: data.type ?? 'generic',
position,
data: { status: 'unknown', services: [], ...data } as NodeData,
...(isContainerNode ? { width: 300, height: 200 } : {}),
...(data.parent_id ? { parentId: data.parent_id, extent: 'parent' as const } : {}),
...(isProxmox ? { width: 300, height: 200 } : {}),
}
addNode(newNode)
toast.success(`Added "${data.label}"`)
@@ -283,12 +182,9 @@ export default function App() {
custom_colors: {
border: data.border_color,
border_style: data.border_style,
border_width: data.border_width,
background: data.background_color,
text_color: data.text_color,
text_position: data.text_position,
text_size: data.text_size,
label_position: data.label_position,
font: data.font,
z_order: data.z_order,
},
@@ -302,7 +198,6 @@ export default function App() {
const handleUpdateGroupRect = useCallback((data: GroupRectFormData) => {
if (!editingGroupRectId) return
snapshotHistory()
const existing = nodes.find((n) => n.id === editingGroupRectId)
updateNode(editingGroupRectId, {
label: data.label,
@@ -310,80 +205,16 @@ export default function App() {
...existing?.data.custom_colors,
border: data.border_color,
border_style: data.border_style,
border_width: data.border_width,
background: data.background_color,
text_color: data.text_color,
text_position: data.text_position,
text_size: data.text_size,
label_position: data.label_position,
font: data.font,
z_order: data.z_order,
},
})
setNodeZIndex(editingGroupRectId, data.z_order - 10)
setEditingGroupRectId(null)
}, [editingGroupRectId, nodes, updateNode, setNodeZIndex, setEditingGroupRectId, snapshotHistory])
const handleAddText = useCallback((data: TextFormData) => {
snapshotHistory()
const id = generateUUID()
const newNode: Node<NodeData> = {
id,
// Text lives in `label` because the API serializer only persists top-level
// node fields; text_content is not in the schema and was lost on reload.
// TextNode and the edit modal both already fall back to label.
type: 'text',
position: { x: 250, y: 250 },
data: {
label: data.text,
type: 'text',
status: 'unknown',
services: [],
custom_colors: {
border: data.border_color,
border_style: data.border_style,
border_width: data.border_width,
background: data.background_color,
text_color: data.text_color,
text_size: data.text_size,
font: data.font,
},
},
width: 200,
height: 60,
}
addNode(newNode)
}, [addNode, snapshotHistory])
const handleUpdateText = useCallback((data: TextFormData) => {
if (!editingTextId) return
snapshotHistory()
const existing = nodes.find((n) => n.id === editingTextId)
updateNode(editingTextId, {
label: data.text,
// Clear stale text_content if present from older builds — label is the
// source of truth now.
text_content: undefined,
custom_colors: {
...existing?.data.custom_colors,
border: data.border_color,
border_style: data.border_style,
border_width: data.border_width,
background: data.background_color,
text_color: data.text_color,
text_size: data.text_size,
font: data.font,
},
})
setEditingTextId(null)
}, [editingTextId, nodes, updateNode, setEditingTextId, snapshotHistory])
const handleDeleteText = useCallback(() => {
if (!editingTextId) return
snapshotHistory()
deleteNode(editingTextId)
setEditingTextId(null)
}, [editingTextId, deleteNode, setEditingTextId, snapshotHistory])
}, [editingGroupRectId, nodes, updateNode, setNodeZIndex, setEditingGroupRectId])
const handleDeleteGroupRect = useCallback(() => {
if (!editingGroupRectId) return
@@ -401,13 +232,13 @@ export default function App() {
snapshotHistory()
const existingNode = nodes.find((n) => n.id === editNodeId)
updateNode(editNodeId, data)
// If container_mode changed, apply structural changes (children parentId, node dimensions)
if (typeof data.container_mode === 'boolean') {
// If proxmox container_mode changed, apply structural changes (children parentId, node dimensions)
if (data.type === 'proxmox' && typeof data.container_mode === 'boolean') {
setProxmoxContainerMode(editNodeId, data.container_mode)
}
// Sync virtual edge when parent_id changes on an LXC/VM node
const nodeType = data.type ?? existingNode?.data.type
if ((nodeType === 'lxc' || nodeType === 'vm' || nodeType === 'docker_container') && 'parent_id' in data) {
if ((nodeType === 'lxc' || nodeType === 'vm') && 'parent_id' in data) {
const oldParentId = existingNode?.data.parent_id ?? null
const newParentId = data.parent_id ?? null
if (oldParentId !== newParentId) {
@@ -420,13 +251,10 @@ export default function App() {
)
if (oldEdge) deleteEdge(oldEdge.id)
}
// Create virtual edge only when parent is NOT in container mode
// (container mode shows containment visually — no edge needed)
// Create new virtual edge: LXC top → Proxmox bottom
if (newParentId) {
const parentNode = nodes.find((n) => n.id === newParentId)
if (!parentNode?.data.container_mode) {
onConnect({ source: editNodeId, sourceHandle: 'top', target: newParentId, targetHandle: 'bottom', type: 'virtual' } as unknown as Connection)
}
// Pass type as extra field — canvasStore.onConnect casts to Connection & Partial<EdgeData>
onConnect({ source: editNodeId, sourceHandle: 'top', target: newParentId, targetHandle: 'bottom', type: 'virtual' } as unknown as Connection)
}
}
}
@@ -465,82 +293,17 @@ export default function App() {
}
}, [nodes, edges, snapshotHistory, loadCanvas, markUnsaved])
// Open the read-only live view of the currently active design in a new tab.
// Standalone has no backend/key — it reads localStorage, so just open /view.
// Otherwise fetch the configured live view key and build /view?key=...&design=<id>.
const handleViewOnly = useCallback(async () => {
if (STANDALONE) {
window.open('/view', '_blank', 'noopener,noreferrer')
return
}
try {
const res = await liveviewApi.getConfig()
if (!res.data.enabled || !res.data.key) {
toast.error('Live view is disabled — set LIVEVIEW_KEY in the backend .env')
return
}
const params = new URLSearchParams({ key: res.data.key })
if (activeDesignId) params.set('design', activeDesignId)
window.open(`/view?${params.toString()}`, '_blank', 'noopener,noreferrer')
} catch {
toast.error('Failed to open live view')
}
}, [activeDesignId])
const handleExport = useCallback(() => {
const handleExport = useCallback(async () => {
const el = canvasRef.current?.querySelector<HTMLElement>('.react-flow')
if (!el) { toast.error('Canvas not ready'); return }
setExportModalOpen(true)
try {
await exportToPng(el)
toast.success('Exported as PNG')
} catch {
toast.error('Export failed')
}
}, [])
const handleZigbeeAddToCanvas = useCallback((zigbeeNodes: ZigbeeNode[], zigbeeEdges: ZigbeeEdge[]) => {
snapshotHistory()
// Place nodes in a grid starting at x=500, y=100
const COLS = 4
const SPACING_X = 170
const SPACING_Y = 100
zigbeeNodes.forEach((zn, i) => {
const id = zn.id
const col = i % COLS
const row = Math.floor(i / COLS)
const position = { x: 500 + col * SPACING_X, y: 100 + row * SPACING_Y }
const newNode: import('@xyflow/react').Node<NodeData> = {
id,
type: zn.type,
position,
data: {
label: zn.friendly_name,
type: zn.type as NodeData['type'],
status: 'unknown' as const,
services: [],
...(zn.lqi != null ? { properties: [{ key: 'LQI', value: String(zn.lqi), icon: 'signal', visible: true }] } : {}),
...(zn.model ? { os: zn.model } : {}),
...(zn.parent_id ? { parent_id: zn.parent_id } : {}),
},
}
addNode(newNode)
})
// Add IoT edges between Zigbee devices: parent bottom -> child top
zigbeeEdges.forEach((ze) => {
onConnect({
source: ze.source,
sourceHandle: 'bottom',
target: ze.target,
targetHandle: 'top-t',
type: 'iot',
} as unknown as import('@xyflow/react').Connection)
})
// Auto-select only the freshly imported nodes so the user can drag the
// whole subtree as a group.
const importedIds = new Set(zigbeeNodes.map((zn) => zn.id))
useCanvasStore.setState((state) => ({
nodes: state.nodes.map((n) => ({ ...n, selected: importedIds.has(n.id) })),
selectedNodeIds: Array.from(importedIds),
selectedNodeId: importedIds.size === 1 ? Array.from(importedIds)[0] : null,
}))
markUnsaved()
}, [addNode, onConnect, snapshotHistory, markUnsaved])
const handleEdgeConnect = useCallback((connection: Connection) => {
setPendingConnection(connection)
}, [])
@@ -549,18 +312,16 @@ export default function App() {
if (!pendingConnection) return
snapshotHistory()
onConnect({ ...pendingConnection, ...edgeData } as unknown as Connection)
// When a virtual edge is drawn between a child node and a container node, sync parent_id
// When a virtual edge is drawn between LXC/VM (top) and Proxmox (bottom), sync parent_id
if (edgeData.type === 'virtual') {
const src = nodes.find((n) => n.id === pendingConnection.source)
const tgt = nodes.find((n) => n.id === pendingConnection.target)
if (src && tgt) {
const assignment = resolveVirtualEdgeParent(
{ id: src.id, type: src.data.type as NodeData['type'] },
{ id: tgt.id, type: tgt.data.type as NodeData['type'] },
)
if (assignment) {
updateNode(assignment.childId, { parent_id: assignment.parentId })
}
const srcType = src?.data.type
const tgtType = tgt?.data.type
if ((srcType === 'lxc' || srcType === 'vm') && tgtType === 'proxmox') {
updateNode(pendingConnection.source, { parent_id: pendingConnection.target })
} else if (srcType === 'proxmox' && (tgtType === 'lxc' || tgtType === 'vm')) {
updateNode(pendingConnection.target, { parent_id: pendingConnection.source })
}
}
setPendingConnection(null)
@@ -570,15 +331,6 @@ export default function App() {
setEditEdgeId(edge.id)
}, [])
const handleNodeDoubleClick = useCallback((node: Node<NodeData>) => {
// 'group' uses inline rename (pencil button in header). Opening the
// generic NodeModal would clobber the group's height (via the
// properties-clears-height rule in updateNode) and lose its children.
// 'groupRect' has its own onDoubleClick that already routes to GroupRectModal.
if (node.data.type === 'group' || node.data.type === 'groupRect') return
handleEditNode(node.id)
}, [handleEditNode])
const handleEdgeUpdate = useCallback((data: EdgeData) => {
if (!editEdgeId) return
snapshotHistory()
@@ -593,13 +345,6 @@ export default function App() {
setEditEdgeId(null)
}, [editEdgeId, deleteEdge, snapshotHistory])
const handleClearWaypoints = useCallback(() => {
if (!editEdgeId) return
snapshotHistory()
updateEdge(editEdgeId, { waypoints: [] })
setEditEdgeId(null)
}, [editEdgeId, updateEdge, snapshotHistory])
const editNode = editNodeId ? nodes.find((n) => n.id === editNodeId) : null
const editEdge = editEdgeId ? edges.find((e) => e.id === editEdgeId) : null
@@ -612,13 +357,9 @@ export default function App() {
<Sidebar
onAddNode={() => setAddNodeOpen(true)}
onAddGroupRect={() => setAddGroupRectOpen(true)}
onAddText={() => setAddTextOpen(true)}
onScan={() => setScanConfigOpen(true)}
onZigbeeImport={() => setZigbeeImportOpen(true)}
onSave={handleSave}
onOpenSettings={() => setSettingsOpen(true)}
onOpenHistory={() => setScanHistoryOpen(true)}
onOpenPending={openPendingModal}
onNodeApproved={setEditNodeId}
/>
<div className="flex flex-col flex-1 min-w-0">
<Toolbar
@@ -632,32 +373,22 @@ export default function App() {
onExportMd={handleExportMd}
onExportYaml={handleExportYaml}
onImportYaml={handleImportYaml}
onViewOnly={handleViewOnly}
/>
<div className="flex flex-1 min-h-0">
<div ref={canvasRef} className="flex-1 min-w-0 h-full">
<CanvasContainer
onConnect={handleEdgeConnect}
onEdgeDoubleClick={handleEdgeDoubleClick}
onNodeDoubleClick={handleNodeDoubleClick}
onNodeDragStart={snapshotHistory}
onRequestAddToGroup={setPendingGroupAdd}
onRequestAddToContainer={setPendingContainerAdd}
onOpenPending={(deviceId) => openPendingModal(deviceId)}
/>
<CanvasContainer onConnect={handleEdgeConnect} onEdgeDoubleClick={handleEdgeDoubleClick} onNodeDragStart={snapshotHistory} />
</div>
{(selectedNodeId || selectedNodeIds.length > 1) && <DetailPanel onEdit={handleEditNode} />}
{selectedNodeId && <DetailPanel onEdit={handleEditNode} />}
</div>
</div>
</div>
<NodeModal
key={addNodeOpen ? 'add-open' : 'add-closed'}
open={addNodeOpen}
onClose={() => setAddNodeOpen(false)}
onSubmit={handleAddNode}
title="Add Node"
parentCandidates={nodes.map((n) => ({ id: n.id, label: n.data.label ?? n.id, type: n.data.type, container_mode: n.data.container_mode }))}
proxmoxNodes={nodes.filter((n) => n.type === 'proxmox').map((n) => ({ id: n.id, label: n.data.label }))}
/>
{/* key forces re-mount when editing a different node, resetting form state */}
@@ -668,25 +399,7 @@ export default function App() {
onSubmit={handleUpdateNode}
initial={editNode?.data}
title="Edit Node"
parentCandidates={(() => {
const descendants = new Set<string>()
if (editNodeId) {
const queue = [editNodeId]
while (queue.length) {
const id = queue.shift()!
for (const n of nodes) {
if (n.data.parent_id === id && !descendants.has(n.id)) {
descendants.add(n.id)
queue.push(n.id)
}
}
}
}
return nodes
.filter((n) => !descendants.has(n.id))
.map((n) => ({ id: n.id, label: n.data.label ?? n.id, type: n.data.type, container_mode: n.data.container_mode }))
})()}
currentNodeId={editNodeId ?? undefined}
proxmoxNodes={nodes.filter((n) => n.type === 'proxmox').map((n) => ({ id: n.id, label: n.data.label }))}
/>
<EdgeModal
@@ -707,7 +420,6 @@ export default function App() {
onClose={() => setEditEdgeId(null)}
onSubmit={handleEdgeUpdate}
onDelete={handleEdgeDelete}
onClearWaypoints={handleClearWaypoints}
initial={editEdge?.data}
title="Edit Link"
/>
@@ -716,27 +428,7 @@ export default function App() {
<ScanConfigModal
open={scanConfigOpen}
onClose={() => setScanConfigOpen(false)}
onScanNow={() => {
toast.success('Network scan started — check Scan History for results')
}}
/>
)}
{!STANDALONE && (
<ZigbeeImportModal
open={zigbeeImportOpen}
onClose={() => setZigbeeImportOpen(false)}
onAddToCanvas={handleZigbeeAddToCanvas}
onPendingImported={() => {
toast.success('Zigbee import started — check Scan History for results')
}}
/>
)}
{!STANDALONE && (
<ScanHistoryModal
open={scanHistoryOpen}
onClose={() => setScanHistoryOpen(false)}
onScanNow={() => toast.success('Scan triggered')}
/>
)}
@@ -744,7 +436,7 @@ export default function App() {
open={addGroupRectOpen}
onClose={() => setAddGroupRectOpen(false)}
onSubmit={handleAddGroupRect}
title="Add Zone"
title="Add Rectangle"
/>
{/* key forces re-mount when editing a different rect */}
@@ -765,45 +457,11 @@ export default function App() {
text_position: rc.text_position ?? 'top-left',
border_color: rc.border ?? '#00d4ff',
border_style: rc.border_style ?? 'solid',
border_width: rc.border_width ?? 2,
background_color: rc.background ?? '#00d4ff0d',
text_size: rc.text_size ?? 12,
label_position: rc.label_position ?? 'inside',
z_order: rc.z_order ?? 1,
}
})()}
title="Edit Zone"
/>
<TextModal
open={addTextOpen}
onClose={() => setAddTextOpen(false)}
onSubmit={handleAddText}
title="Add Text"
/>
<TextModal
key={editingTextId ?? 'text-edit'}
open={!!editingTextId}
onClose={() => setEditingTextId(null)}
onSubmit={handleUpdateText}
onDelete={handleDeleteText}
initial={(() => {
const n = editingTextId ? nodes.find((nd) => nd.id === editingTextId) : null
if (!n) return undefined
const rc = n.data.custom_colors ?? {}
return {
text: n.data.text_content ?? n.data.label ?? '',
font: rc.font ?? 'inter',
text_color: rc.text_color ?? '#e6edf3',
text_size: rc.text_size ?? 14,
border_color: rc.border ?? '#30363d',
border_style: (rc.border_style ?? 'none') as TextFormData['border_style'],
border_width: rc.border_width ?? 1,
background_color: rc.background ?? '#00000000',
}
})()}
title="Edit Text"
title="Edit Rectangle"
/>
{/* key forces re-mount on open so useState captures current theme as original */}
@@ -813,53 +471,9 @@ export default function App() {
onClose={() => setThemeModalOpen(false)}
/>
<SearchModal
open={searchOpen}
onClose={() => setSearchOpen(false)}
onOpenPending={(deviceId) => openPendingModal(deviceId)}
/>
<SearchModal open={searchOpen} onClose={() => setSearchOpen(false)} />
<ShortcutsModal open={shortcutsOpen} onClose={() => setShortcutsOpen(false)} />
<ConfirmAddToGroupModal
open={!!pendingGroupAdd}
nodeLabel={pendingGroupAdd ? (nodes.find((n) => n.id === pendingGroupAdd.nodeId)?.data.label ?? '') : ''}
targetLabel={pendingGroupAdd ? (nodes.find((n) => n.id === pendingGroupAdd.groupId)?.data.label ?? '') : ''}
onConfirm={() => {
if (pendingGroupAdd) addToGroup(pendingGroupAdd.groupId, pendingGroupAdd.nodeId)
setPendingGroupAdd(null)
}}
onCancel={() => setPendingGroupAdd(null)}
/>
<ConfirmAddToGroupModal
open={!!pendingContainerAdd}
variant="container"
nodeLabel={pendingContainerAdd ? (nodes.find((n) => n.id === pendingContainerAdd.nodeId)?.data.label ?? '') : ''}
targetLabel={pendingContainerAdd ? (nodes.find((n) => n.id === pendingContainerAdd.containerId)?.data.label ?? '') : ''}
onConfirm={() => {
if (pendingContainerAdd) addToContainer(pendingContainerAdd.containerId, pendingContainerAdd.nodeId)
setPendingContainerAdd(null)
}}
onCancel={() => setPendingContainerAdd(null)}
/>
{!STANDALONE && (
<SettingsModal open={settingsOpen} onClose={() => setSettingsOpen(false)} />
)}
<PendingDevicesModal
open={pendingModalOpen}
onClose={() => setPendingModalOpen(false)}
highlightId={pendingHighlightId}
initialStatus={pendingModalStatus}
/>
<ExportModal
open={exportModalOpen}
onClose={() => setExportModalOpen(false)}
getElement={() => canvasRef.current?.querySelector<HTMLElement>('.react-flow') ?? null}
/>
<Toaster theme="dark" position="bottom-right" />
</ReactFlowProvider>
</TooltipProvider>
-220
View File
@@ -1,220 +0,0 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
type Interceptor<T> = {
fulfilled?: (v: T) => T | Promise<T>
rejected?: (e: unknown) => unknown
}
interface MockInstance {
defaults: { baseURL?: string }
interceptors: {
request: { use: (f: Interceptor<unknown>['fulfilled'], r?: Interceptor<unknown>['rejected']) => void }
response: { use: (f: Interceptor<unknown>['fulfilled'], r?: Interceptor<unknown>['rejected']) => void }
}
get: ReturnType<typeof vi.fn>
post: ReturnType<typeof vi.fn>
patch: ReturnType<typeof vi.fn>
delete: ReturnType<typeof vi.fn>
__req: Interceptor<{ headers: Record<string, string> }>
__res: Interceptor<unknown>
}
const hoisted = vi.hoisted(() => ({ instances: [] as unknown[] }))
const instances = hoisted.instances as MockInstance[]
vi.mock('axios', () => {
return {
default: {
create: (cfg: { baseURL?: string }) => {
const inst: MockInstance = {
defaults: { baseURL: cfg?.baseURL },
interceptors: {
request: { use: (f: unknown, r?: unknown) => { inst.__req = { fulfilled: f as never, rejected: r as never } } },
response: { use: (f: unknown, r?: unknown) => { inst.__res = { fulfilled: f as never, rejected: r as never } } },
},
get: vi.fn(() => Promise.resolve({ data: {} })),
post: vi.fn(() => Promise.resolve({ data: {} })),
patch: vi.fn(() => Promise.resolve({ data: {} })),
delete: vi.fn(() => Promise.resolve({ data: {} })),
__req: {},
__res: {},
}
hoisted.instances.push(inst)
return inst
},
},
}
})
import { useAuthStore } from '@/stores/authStore'
import * as clientModule from '../client'
describe('api/client', () => {
const mod = clientModule
const [api, publicApi] = instances
beforeEach(() => {
useAuthStore.setState({ token: null, isAuthenticated: false })
api.get.mockClear()
api.post.mockClear()
api.patch.mockClear()
api.delete.mockClear()
publicApi.get.mockClear()
publicApi.post.mockClear()
})
it('creates two axios instances with /api/v1 baseURL', () => {
expect(instances).toHaveLength(2)
expect(api.defaults.baseURL).toBe('/api/v1')
expect(publicApi.defaults.baseURL).toBe('/api/v1')
})
it('exports `api` matching the first created instance', () => {
expect(mod.api).toBe(api)
})
it('request interceptor adds Authorization header when token present', () => {
useAuthStore.setState({ token: 'tok-123', isAuthenticated: true })
const cfg = { headers: {} as Record<string, string> }
const out = api.__req.fulfilled!(cfg)
expect((out as typeof cfg).headers.Authorization).toBe('Bearer tok-123')
})
it('request interceptor leaves headers untouched when no token', () => {
const cfg = { headers: {} as Record<string, string> }
const out = api.__req.fulfilled!(cfg)
expect((out as typeof cfg).headers.Authorization).toBeUndefined()
})
it('response interceptor passes through 2xx responses', () => {
const r = { status: 200, data: { ok: true } }
expect(api.__res.fulfilled!(r)).toBe(r)
})
it('response interceptor calls logout on 401', async () => {
const logout = vi.spyOn(useAuthStore.getState(), 'logout')
useAuthStore.setState({ token: 't', isAuthenticated: true, logout })
const err = { response: { status: 401 } }
await expect(api.__res.rejected!(err)).rejects.toBe(err)
expect(logout).toHaveBeenCalled()
})
it('response interceptor does not call logout on non-401', async () => {
const logout = vi.fn()
useAuthStore.setState({ token: 't', isAuthenticated: true, logout })
const err = { response: { status: 500 } }
await expect(api.__res.rejected!(err)).rejects.toBe(err)
expect(logout).not.toHaveBeenCalled()
})
it('response interceptor handles error with no response object', async () => {
const logout = vi.fn()
useAuthStore.setState({ logout })
const err = { message: 'network down' }
await expect(api.__res.rejected!(err)).rejects.toBe(err)
expect(logout).not.toHaveBeenCalled()
})
it('publicApi has no request/response interceptors registered', () => {
expect(publicApi.__req.fulfilled).toBeUndefined()
expect(publicApi.__res.fulfilled).toBeUndefined()
})
it('authApi.login posts to /auth/login', () => {
mod.authApi.login('u', 'p')
expect(api.post).toHaveBeenCalledWith('/auth/login', { username: 'u', password: 'p' })
})
it('canvasApi.load GETs /canvas', () => {
mod.canvasApi.load()
expect(api.get).toHaveBeenCalledWith('/canvas', expect.objectContaining({}))
})
it('canvasApi.save POSTs to /canvas/save with payload', () => {
const payload = { nodes: [], edges: [], viewport: {} }
mod.canvasApi.save(payload)
expect(api.post).toHaveBeenCalledWith('/canvas/save', payload)
})
it('nodesApi CRUD calls correct endpoints', () => {
mod.nodesApi.create({ a: 1 })
expect(api.post).toHaveBeenCalledWith('/nodes', { a: 1 })
mod.nodesApi.update('n1', { b: 2 })
expect(api.patch).toHaveBeenCalledWith('/nodes/n1', { b: 2 })
mod.nodesApi.delete('n1')
expect(api.delete).toHaveBeenCalledWith('/nodes/n1')
})
it('edgesApi CRUD calls correct endpoints', () => {
mod.edgesApi.create({ s: 'a', t: 'b' })
expect(api.post).toHaveBeenCalledWith('/edges', { s: 'a', t: 'b' })
mod.edgesApi.delete('e1')
expect(api.delete).toHaveBeenCalledWith('/edges/e1')
})
it('liveviewApi.load uses publicApi with key param', () => {
mod.liveviewApi.load('k-1')
expect(publicApi.get).toHaveBeenCalledWith('/liveview', { params: { key: 'k-1' } })
expect(api.get).not.toHaveBeenCalled()
})
it('liveviewApi.load forwards design as design_id when provided', () => {
mod.liveviewApi.load('k-1', 'design-9')
expect(publicApi.get).toHaveBeenCalledWith('/liveview', { params: { key: 'k-1', design_id: 'design-9' } })
})
it('liveviewApi.getConfig hits the authenticated config endpoint', () => {
mod.liveviewApi.getConfig()
expect(api.get).toHaveBeenCalledWith('/liveview/config')
})
it('scanApi endpoints route correctly', () => {
mod.scanApi.trigger()
expect(api.post).toHaveBeenCalledWith('/scan/trigger')
mod.scanApi.pending()
expect(api.get).toHaveBeenCalledWith('/scan/pending')
mod.scanApi.hidden()
expect(api.get).toHaveBeenCalledWith('/scan/hidden')
mod.scanApi.runs()
expect(api.get).toHaveBeenCalledWith('/scan/runs')
mod.scanApi.clearPending()
expect(api.delete).toHaveBeenCalledWith('/scan/pending')
mod.scanApi.approve('d1', { foo: 'bar' })
expect(api.post).toHaveBeenCalledWith('/scan/pending/d1/approve', { foo: 'bar' })
mod.scanApi.hide('d1')
expect(api.post).toHaveBeenCalledWith('/scan/pending/d1/hide')
mod.scanApi.ignore('d1')
expect(api.post).toHaveBeenCalledWith('/scan/pending/d1/ignore')
mod.scanApi.bulkApprove(['a', 'b'])
expect(api.post).toHaveBeenCalledWith('/scan/pending/bulk-approve', { device_ids: ['a', 'b'] })
mod.scanApi.bulkHide(['a'])
expect(api.post).toHaveBeenCalledWith('/scan/pending/bulk-hide', { device_ids: ['a'] })
mod.scanApi.restore('d1')
expect(api.post).toHaveBeenCalledWith('/scan/pending/d1/restore')
mod.scanApi.bulkRestore(['a'])
expect(api.post).toHaveBeenCalledWith('/scan/pending/bulk-restore', { device_ids: ['a'] })
mod.scanApi.stop('run-1')
expect(api.post).toHaveBeenCalledWith('/scan/run-1/stop')
mod.scanApi.getConfig()
expect(api.get).toHaveBeenCalledWith('/scan/config')
mod.scanApi.saveConfig({ ranges: ['1.0/24'] })
expect(api.post).toHaveBeenCalledWith('/scan/config', { ranges: ['1.0/24'] })
})
it('settingsApi get/save', () => {
mod.settingsApi.get()
expect(api.get).toHaveBeenCalledWith('/settings')
mod.settingsApi.save({ interval_seconds: 30, service_check_enabled: true, service_check_interval: 600 })
expect(api.post).toHaveBeenCalledWith('/settings', { interval_seconds: 30, service_check_enabled: true, service_check_interval: 600 })
})
it('zigbeeApi.testConnection/importNetwork/importToPending', () => {
const cfg = { mqtt_host: 'h', mqtt_port: 1883 }
mod.zigbeeApi.testConnection(cfg)
expect(api.post).toHaveBeenCalledWith('/zigbee/test-connection', cfg)
mod.zigbeeApi.importNetwork(cfg)
expect(api.post).toHaveBeenCalledWith('/zigbee/import', cfg)
mod.zigbeeApi.importToPending(cfg)
expect(api.post).toHaveBeenCalledWith('/zigbee/import-pending', cfg)
})
})
+4 -105
View File
@@ -5,9 +5,6 @@ export const api = axios.create({
baseURL: '/api/v1',
})
// Unauthenticated axios instance — no JWT, no 401 redirect (used for public endpoints)
const publicApi = axios.create({ baseURL: '/api/v1' })
api.interceptors.request.use((config) => {
const token = useAuthStore.getState().token
if (token) config.headers.Authorization = `Bearer ${token}`
@@ -28,16 +25,11 @@ export const authApi = {
}
export const canvasApi = {
load: (design_id?: string) => {
const params = design_id ? { design_id } : {}
return api.get('/canvas', { params })
},
load: () => api.get('/canvas'),
save: (payload: {
nodes: object[]
edges: object[]
viewport: object
custom_style?: object | null
design_id?: string | null
}) => api.post('/canvas/save', payload),
}
@@ -52,107 +44,14 @@ export const edgesApi = {
delete: (id: string) => api.delete(`/edges/${id}`),
}
export const liveviewApi = {
load: (key: string, design?: string) =>
publicApi.get('/liveview', { params: { key, ...(design ? { design_id: design } : {}) } }),
getConfig: () => api.get<{ enabled: boolean; key: string | null }>('/liveview/config'),
}
export const scanApi = {
trigger: () => api.post('/scan/trigger'),
pending: () => api.get('/scan/pending'),
hidden: () => api.get('/scan/hidden'),
runs: () => api.get('/scan/runs'),
clearPending: () => api.delete('/scan/pending'),
approve: (id: string, nodeData: object) =>
api.post<{
approved: boolean
node_id: string
edges_created: number
edges: { id: string; source: string; target: string }[]
}>(`/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`),
ignore: (id: string) => api.post(`/scan/pending/${id}/ignore`),
bulkApprove: (ids: string[]) =>
api.post<{
approved: number
node_ids: string[]
device_ids: string[]
edges_created: number
edges: { id: string; source: string; target: string }[]
skipped: number
}>('/scan/pending/bulk-approve', { device_ids: ids }),
bulkHide: (ids: string[]) => api.post<{ hidden: number; skipped: number }>('/scan/pending/bulk-hide', { device_ids: ids }),
restore: (id: string) => api.post<{ restored: boolean; device_id: string }>(`/scan/pending/${id}/restore`),
bulkRestore: (ids: string[]) => api.post<{ restored: number; skipped: number }>('/scan/pending/bulk-restore', { device_ids: ids }),
stop: (runId: string) => api.post(`/scan/${runId}/stop`),
getConfig: () => api.get<{ ranges: string[] }>('/scan/config'),
saveConfig: (data: { ranges: string[] }) => api.post('/scan/config', data),
}
export interface AppSettings {
interval_seconds: number
service_check_enabled: boolean
service_check_interval: number
}
export const settingsApi = {
get: () => api.get<AppSettings>('/settings'),
save: (data: AppSettings) => api.post<AppSettings>('/settings', data),
}
export const designsApi = {
list: () => api.get<import('@/types').Design[]>('/designs'),
create: (data: { name: string; icon?: string; design_type?: string }) =>
api.post<import('@/types').Design>('/designs', data),
update: (id: string, data: { name?: string; icon?: string }) =>
api.put<import('@/types').Design>(`/designs/${id}`, data),
delete: (id: string) => api.delete(`/designs/${id}`),
}
export const zigbeeApi = {
testConnection: (data: {
mqtt_host: string
mqtt_port: number
mqtt_username?: string
mqtt_password?: string
mqtt_tls?: boolean
mqtt_tls_insecure?: boolean
}) =>
api.post<{ connected: boolean; message: string }>('/zigbee/test-connection', data),
importNetwork: (data: {
mqtt_host: string
mqtt_port: number
mqtt_username?: string
mqtt_password?: string
base_topic?: string
mqtt_tls?: boolean
mqtt_tls_insecure?: boolean
}) =>
api.post<{
nodes: import('@/components/zigbee/types').ZigbeeNode[]
edges: import('@/components/zigbee/types').ZigbeeEdge[]
device_count: number
}>('/zigbee/import', data),
importToPending: (data: {
mqtt_host: string
mqtt_port: number
mqtt_username?: string
mqtt_password?: string
base_topic?: string
mqtt_tls?: boolean
mqtt_tls_insecure?: boolean
}) =>
api.post<{
id: string
status: string
kind: string
ranges: string[]
devices_found: number
started_at: string
finished_at: string | null
error: string | null
}>('/zigbee/import-pending', data),
getConfig: () => api.get<{ ranges: string[]; interval_seconds: number }>('/scan/config'),
saveConfig: (data: { ranges: string[]; interval_seconds: number }) => api.post('/scan/config', data),
}
-189
View File
@@ -1,189 +0,0 @@
/**
* LiveView — read-only canvas accessible at /view?key=<LIVEVIEW_KEY>.
*
* - Non-standalone: fetches canvas from /api/v1/liveview?key=... (no JWT needed).
* Returns 403 when the feature is disabled or the key is wrong.
* - Standalone: loads canvas from localStorage directly (no key required,
* since there is no backend to validate against).
*
* Pan and zoom work. Editing is fully disabled.
* Clicking a node with an IP opens http://<ip> in a new tab.
*/
import { useCallback, useEffect, useMemo, useState } from 'react'
import {
ReactFlowProvider,
ReactFlow,
Background,
BackgroundVariant,
Controls,
ConnectionMode,
useReactFlow,
type Node,
} from '@xyflow/react'
import '@xyflow/react/dist/style.css'
import { useCanvasStore } from '@/stores/canvasStore'
import { useThemeStore } from '@/stores/themeStore'
import { THEMES } from '@/utils/themes'
import { nodeTypes } from '@/components/canvas/nodes/nodeTypes'
import { edgeTypes } from '@/components/canvas/edges/edgeTypes'
import { deserializeApiNode, deserializeApiEdge, type ApiNode, type ApiEdge } from '@/utils/canvasSerializer'
import { computeCollapseInfo, rewireEdgesForCollapse } from '@/utils/collapseFilter'
import { liveviewApi } from '@/api/client'
import type { NodeData, CustomStyleDef } from '@/types'
const STANDALONE = import.meta.env.VITE_STANDALONE === 'true'
const STORAGE_KEY = 'homelable_canvas'
type ViewState = 'loading' | 'disabled' | 'invalid-key' | 'no-key' | 'network-error' | 'ready'
function LiveViewCanvas() {
const { nodes, edges, loadCanvas, fitViewPending, clearFitViewPending } = useCanvasStore()
const { fitView } = useReactFlow()
const activeTheme = useThemeStore((s) => s.activeTheme)
const setTheme = useThemeStore((s) => s.setTheme)
const setCustomStyle = useThemeStore((s) => s.setCustomStyle)
const theme = THEMES[activeTheme]
// Derive initial view state synchronously (avoids calling setState inside an effect):
// - standalone → always ready (localStorage, no key required)
// - non-standalone, no ?key= → no-key error immediately
// - non-standalone, key present → loading (API call below)
const [viewState, setViewState] = useState<ViewState>(() => {
if (STANDALONE) return 'ready'
return new URLSearchParams(window.location.search).get('key') ? 'loading' : 'no-key'
})
useEffect(() => {
if (STANDALONE) {
try {
const saved = localStorage.getItem(STORAGE_KEY)
if (saved) {
const { nodes: savedNodes, edges: savedEdges } = JSON.parse(saved)
loadCanvas(savedNodes, savedEdges)
}
} catch {
// empty canvas on parse error — show empty canvas
}
return
}
// Already handled synchronously in useState initializer
const search = new URLSearchParams(window.location.search)
const key = search.get('key')
if (!key) return
// Optional ?design=<id> selects which canvas to render; backend falls back
// to the first design when omitted.
const design = search.get('design') ?? undefined
liveviewApi.load(key, design)
.then((res) => {
const { nodes: apiNodes, edges: apiEdges } = res.data
const proxmoxMap = new Map<string, boolean>(
(apiNodes as ApiNode[])
.filter((n: ApiNode) => n.type === 'group' || n.container_mode === true)
.map((n: ApiNode) => [n.id, true])
)
const savedTheme = res.data.viewport?.theme_id
if (savedTheme) setTheme(savedTheme)
if (res.data.custom_style) setCustomStyle(res.data.custom_style as CustomStyleDef)
loadCanvas(
(apiNodes as ApiNode[]).map((n) => deserializeApiNode(n, proxmoxMap)),
(apiEdges as ApiEdge[]).map(deserializeApiEdge),
)
setViewState('ready')
})
.catch((err) => {
if (!err.response) { setViewState('network-error'); return }
const detail: string = err.response.data?.detail ?? ''
setViewState(detail === 'Live view is disabled' ? 'disabled' : 'invalid-key')
})
}, [loadCanvas, setTheme, setCustomStyle])
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 ip = node.data.ip
if (ip) window.open(`http://${ip}`, '_blank', 'noopener,noreferrer')
}, [])
// Apply collapse-state filtering — same pipeline the editor canvas uses,
// so a collapsed group/zone hides its contents in live view too.
const collapseInfo = useMemo(() => computeCollapseInfo(nodes), [nodes])
const visibleNodes = useMemo(
() => nodes.filter((n) => collapseInfo.visibleIds.has(n.id)),
[nodes, collapseInfo],
)
const visibleEdges = useMemo(
() => rewireEdgesForCollapse(edges, nodes, collapseInfo.visibleIds, collapseInfo.hiddenBy),
[edges, nodes, collapseInfo],
)
if (viewState === 'loading') {
return (
<div className="flex h-screen w-screen items-center justify-center bg-[#0d1117] text-[#8b949e]">
Loading
</div>
)
}
if (viewState !== 'ready') {
const messages: Record<Exclude<ViewState, 'loading' | 'ready'>, string> = {
disabled: 'Live view is disabled on this instance.',
'invalid-key': 'Invalid or expired live view key.',
'no-key': 'Missing key — use ?key=your-secret in the URL.',
'network-error': 'Could not reach the server. Check your connection.',
}
return (
<div className="flex h-screen w-screen items-center justify-center bg-[#0d1117]">
<div className="text-center space-y-2">
<p className="text-[#f85149] text-lg font-medium">Access Denied</p>
<p className="text-[#8b949e] text-sm">{messages[viewState]}</p>
</div>
</div>
)
}
return (
<div className="w-full h-screen" style={{ background: theme.colors.canvasBackground }}>
<ReactFlow
nodes={visibleNodes}
edges={visibleEdges}
nodeTypes={nodeTypes}
edgeTypes={edgeTypes}
nodesDraggable={false}
nodesConnectable={false}
elementsSelectable={false}
panOnDrag
zoomOnScroll
minZoom={0.25}
maxZoom={2.5}
colorMode={theme.colors.reactFlowColorMode}
connectionMode={ConnectionMode.Loose}
onNodeClick={onNodeClick}
>
<Background
variant={BackgroundVariant.Dots}
gap={24}
size={1}
color={theme.colors.canvasDotColor}
/>
<Controls showInteractive={false} />
</ReactFlow>
</div>
)
}
export default function LiveView() {
return (
<ReactFlowProvider>
<LiveViewCanvas />
</ReactFlowProvider>
)
}
+3 -4
View File
@@ -20,9 +20,8 @@ export function LoginPage() {
try {
const res = await authApi.login(username, password)
login(res.data.access_token)
} catch (err: unknown) {
const hasResponse = err && typeof err === 'object' && 'response' in err
setError(hasResponse ? 'Invalid username or password' : 'Could not reach the server — check your CORS_ORIGINS setting')
} catch {
setError('Invalid username or password')
} finally {
setLoading(false)
}
@@ -96,7 +95,7 @@ export function LoginPage() {
</form>
<p className="text-center text-[10px] text-muted-foreground/40 mt-4">
Credentials configured in <span className="font-mono">.env</span>
Credentials configured in <span className="font-mono">config.yml</span>
</p>
</div>
</div>
@@ -1,104 +0,0 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { render, waitFor } from '@testing-library/react'
import type { Node, Edge } from '@xyflow/react'
import type { NodeData, EdgeData } from '@/types'
// ── Capture the props ReactFlow is rendered with ──────────────────────────
const rfPropsSpy = vi.fn()
vi.mock('@xyflow/react', () => ({
ReactFlowProvider: ({ children }: { children: React.ReactNode }) => <>{children}</>,
ReactFlow: (props: unknown) => {
rfPropsSpy(props)
return <div data-testid="react-flow" />
},
Background: () => null,
Controls: () => null,
BackgroundVariant: { Dots: 'dots' },
ConnectionMode: { Loose: 'loose' },
Position: { Top: 'top', Right: 'right', Bottom: 'bottom', Left: 'left' },
useReactFlow: () => ({ fitView: vi.fn() }),
}))
vi.mock('@xyflow/react/dist/style.css', () => ({}))
vi.mock('@/api/client', () => ({ liveviewApi: { load: vi.fn() } }))
import { liveviewApi } from '@/api/client'
import LiveView from '../LiveView'
function setSearch(params: string) {
Object.defineProperty(window, 'location', {
writable: true,
value: { ...window.location, search: params, pathname: '/view' },
})
}
/** Build a /liveview API response with the given nodes/edges. */
const apiResponse = (nodes: unknown[], edges: unknown[] = []) => ({
data: { nodes, edges, viewport: { x: 0, y: 0, zoom: 1 } },
})
const apiNode = (
id: string,
parent_id?: string,
collapsed?: boolean,
type = 'server',
) => ({
id,
type,
label: id,
status: 'online',
services: [],
pos_x: 0,
pos_y: 0,
parent_id: parent_id ?? null,
container_mode: type === 'group',
custom_colors: collapsed !== undefined ? { collapsed } : null,
created_at: '2024-01-01T00:00:00Z',
updated_at: '2024-01-01T00:00:00Z',
})
describe('LiveView — applies collapse filter to the rendered canvas', () => {
beforeEach(() => {
rfPropsSpy.mockClear()
setSearch('?key=valid')
vi.mocked(liveviewApi.load).mockReset()
})
it('hides children of a collapsed group container in view-only mode', async () => {
vi.mocked(liveviewApi.load).mockResolvedValue(
apiResponse([apiNode('g1', undefined, true, 'group'), apiNode('c1', 'g1')]),
)
render(<LiveView />)
await waitFor(() => {
const last = rfPropsSpy.mock.calls[rfPropsSpy.mock.calls.length - 1]?.[0] as
| { nodes: Node<NodeData>[] }
| undefined
expect(last?.nodes.length).toBeGreaterThan(0)
})
const last = rfPropsSpy.mock.calls[rfPropsSpy.mock.calls.length - 1][0] as {
nodes: Node<NodeData>[]
edges: Edge<EdgeData>[]
}
const ids = last.nodes.map((n) => n.id)
expect(ids).toContain('g1')
expect(ids).not.toContain('c1')
})
it('shows children when the group is expanded', async () => {
vi.mocked(liveviewApi.load).mockResolvedValue(
apiResponse([apiNode('g1', undefined, false, 'group'), apiNode('c1', 'g1')]),
)
render(<LiveView />)
await waitFor(() => {
const last = rfPropsSpy.mock.calls[rfPropsSpy.mock.calls.length - 1]?.[0] as
| { nodes: Node<NodeData>[] }
| undefined
expect(last?.nodes.length).toBeGreaterThan(1)
})
const last = rfPropsSpy.mock.calls[rfPropsSpy.mock.calls.length - 1][0] as {
nodes: Node<NodeData>[]
}
const ids = last.nodes.map((n) => n.id)
expect(ids).toContain('g1')
expect(ids).toContain('c1')
})
})
@@ -1,283 +0,0 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { render, screen, waitFor } from '@testing-library/react'
import { useCanvasStore } from '@/stores/canvasStore'
import { useThemeStore } from '@/stores/themeStore'
// ── Mock heavy dependencies ────────────────────────────────────────────────
// Capture props passed to ReactFlow so we can assert zoom bounds etc.
let rfProps: Record<string, unknown> = {}
vi.mock('@xyflow/react', () => ({
ReactFlowProvider: ({ children }: { children: React.ReactNode }) => <>{children}</>,
ReactFlow: (props: Record<string, unknown>) => {
rfProps = props
return <div data-testid="react-flow" />
},
Background: () => null,
Controls: () => null,
BackgroundVariant: { Dots: 'dots' },
ConnectionMode: { Loose: 'loose' },
Position: { Top: 'top', Right: 'right', Bottom: 'bottom', Left: 'left' },
useReactFlow: () => ({ fitView: vi.fn() }),
}))
vi.mock('@xyflow/react/dist/style.css', () => ({}))
vi.mock('@/api/client', () => ({
liveviewApi: { load: vi.fn() },
}))
import { liveviewApi } from '@/api/client'
import LiveView from '../LiveView'
// ── Helpers ────────────────────────────────────────────────────────────────
function setSearch(params: string) {
Object.defineProperty(window, 'location', {
writable: true,
value: { ...window.location, search: params, pathname: '/view' },
})
}
const canvasPayload = {
data: {
nodes: [{
id: 'n1', type: 'server', label: 'CI Node', status: 'online',
services: [], pos_x: 0, pos_y: 0,
created_at: '2024-01-01T00:00:00Z', updated_at: '2024-01-01T00:00:00Z',
}],
edges: [],
viewport: { x: 0, y: 0, zoom: 1 },
},
}
// ── Tests ──────────────────────────────────────────────────────────────────
describe('LiveView (non-standalone)', () => {
beforeEach(() => {
rfProps = {}
vi.mocked(liveviewApi.load).mockReset()
useCanvasStore.setState({ nodes: [], edges: [] })
})
afterEach(() => { setSearch('') })
// ── No key ────────────────────────────────────────────────────────────────
it('shows no-key error when ?key= is missing', async () => {
setSearch('')
render(<LiveView />)
await waitFor(() => {
expect(screen.getByText('Access Denied')).toBeDefined()
expect(screen.getByText(/Missing key/)).toBeDefined()
})
expect(liveviewApi.load).not.toHaveBeenCalled()
})
// ── Disabled ──────────────────────────────────────────────────────────────
it('shows disabled error when backend returns "Live view is disabled"', async () => {
setSearch('?key=anything')
vi.mocked(liveviewApi.load).mockRejectedValue({
response: { data: { detail: 'Live view is disabled' } },
})
render(<LiveView />)
await waitFor(() => {
expect(screen.getByText(/disabled on this instance/)).toBeDefined()
})
})
// ── Invalid key ───────────────────────────────────────────────────────────
it('shows invalid-key error when backend returns "Invalid live view key"', async () => {
setSearch('?key=wrong')
vi.mocked(liveviewApi.load).mockRejectedValue({
response: { data: { detail: 'Invalid live view key' } },
})
render(<LiveView />)
await waitFor(() => {
expect(screen.getByText(/Invalid or expired/)).toBeDefined()
})
})
it('shows network-error for non-response errors (offline, CORS, 500)', async () => {
setSearch('?key=anything')
vi.mocked(liveviewApi.load).mockRejectedValue(new Error('network'))
render(<LiveView />)
await waitFor(() => {
expect(screen.getByText(/Could not reach the server/)).toBeDefined()
})
})
// ── Valid key → canvas rendered ───────────────────────────────────────────
it('renders the canvas on valid key', async () => {
setSearch('?key=correct-key')
vi.mocked(liveviewApi.load).mockResolvedValue(canvasPayload as never)
render(<LiveView />)
await waitFor(() => {
expect(screen.getByTestId('react-flow')).toBeDefined()
})
expect(liveviewApi.load).toHaveBeenCalledWith('correct-key', undefined)
})
it('forwards ?design=<id> to the API so a specific canvas is loaded', async () => {
setSearch('?key=correct-key&design=elec-123')
vi.mocked(liveviewApi.load).mockResolvedValue(canvasPayload as never)
render(<LiveView />)
await waitFor(() => expect(screen.getByTestId('react-flow')).toBeDefined())
expect(liveviewApi.load).toHaveBeenCalledWith('correct-key', 'elec-123')
})
it('allows zooming out to 0.25 so large infra fits (matches the editor)', async () => {
setSearch('?key=correct-key')
vi.mocked(liveviewApi.load).mockResolvedValue(canvasPayload as never)
render(<LiveView />)
await waitFor(() => expect(screen.getByTestId('react-flow')).toBeDefined())
// Without an explicit minZoom, React Flow defaults to 0.5 and big canvases
// can't zoom out far enough to fit.
expect(rfProps.minZoom).toBe(0.25)
expect(rfProps.maxZoom).toBe(2.5)
})
it('loads nodes into the canvas store on success', async () => {
setSearch('?key=secret')
vi.mocked(liveviewApi.load).mockResolvedValue(canvasPayload as never)
render(<LiveView />)
await waitFor(() => {
expect(screen.getByTestId('react-flow')).toBeDefined()
})
const { nodes } = useCanvasStore.getState()
expect(nodes.find((n) => n.id === 'n1')).toBeDefined()
})
// ── Nested children (docker_container inside docker_host) ────────────────
it('nests docker_container under docker_host parent (container_mode=true)', async () => {
setSearch('?key=valid')
const nestedPayload = {
data: {
nodes: [
{
id: 'host', type: 'docker', label: 'Docker Host', status: 'online',
services: [], pos_x: 0, pos_y: 0, container_mode: true,
created_at: '2024-01-01T00:00:00Z', updated_at: '2024-01-01T00:00:00Z',
},
{
id: 'ctr', type: 'docker_container', label: 'nginx', status: 'online',
services: [], pos_x: 20, pos_y: 30, parent_id: 'host',
created_at: '2024-01-01T00:00:00Z', updated_at: '2024-01-01T00:00:00Z',
},
],
edges: [],
viewport: { x: 0, y: 0, zoom: 1 },
},
}
vi.mocked(liveviewApi.load).mockResolvedValue(nestedPayload as never)
render(<LiveView />)
await waitFor(() => expect(screen.getByTestId('react-flow')).toBeDefined())
const ctr = useCanvasStore.getState().nodes.find((n) => n.id === 'ctr')
expect(ctr?.parentId).toBe('host')
expect(ctr?.extent).toBe('parent')
})
// ── Theme + custom_style applied from payload ────────────────────────────
it('applies viewport.theme_id and custom_style from the payload', async () => {
setSearch('?key=valid')
const styledPayload = {
data: {
nodes: [],
edges: [],
viewport: { x: 0, y: 0, zoom: 1, theme_id: 'matrix' },
custom_style: { fontFamily: 'Inter', nodeRadius: 12 },
},
}
vi.mocked(liveviewApi.load).mockResolvedValue(styledPayload as never)
render(<LiveView />)
await waitFor(() => expect(screen.getByTestId('react-flow')).toBeDefined())
expect(useThemeStore.getState().activeTheme).toBe('matrix')
expect(useThemeStore.getState().customStyle).toEqual({ fontFamily: 'Inter', nodeRadius: 12 })
})
// ── No editing props passed ───────────────────────────────────────────────
it('does not show any Access Denied when key is valid', async () => {
setSearch('?key=valid')
vi.mocked(liveviewApi.load).mockResolvedValue(canvasPayload as never)
render(<LiveView />)
await waitFor(() => expect(screen.getByTestId('react-flow')).toBeDefined())
expect(screen.queryByText('Access Denied')).toBeNull()
})
})
// ── 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' },
Position: { Top: 'top', Right: 'right', Bottom: 'bottom', Left: 'left' },
useReactFlow: () => ({ fitView: vi.fn() }),
}
describe('LiveView (standalone — localStorage)', () => {
beforeEach(() => {
localStorage.clear()
useCanvasStore.setState({ nodes: [], edges: [] })
})
afterEach(() => {
setSearch('')
vi.unstubAllEnvs()
})
it('loads canvas from localStorage without calling the API', async () => {
const stored = {
nodes: [{
id: 'ls-node', type: 'router',
position: { x: 10, y: 20 },
data: { label: 'Router', type: 'router', status: 'unknown', services: [] },
}],
edges: [],
}
localStorage.setItem('homelable_canvas', JSON.stringify(stored))
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()
})
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()
})
})
@@ -51,7 +51,7 @@ describe('LoginPage', () => {
})
it('shows a generic error message — no credential enumeration', async () => {
vi.mocked(authApi.login).mockRejectedValue({ response: { status: 401 } })
vi.mocked(authApi.login).mockRejectedValue(new Error('401'))
render(<LoginPage />)
fireEvent.change(screen.getByLabelText('Username'), { target: { value: 'admin' } })
fireEvent.change(screen.getByLabelText('Password'), { target: { value: 'wrongpass' } })
@@ -65,21 +65,10 @@ describe('LoginPage', () => {
expect(errors[0].textContent).toBe('Invalid username or password')
})
it('shows a network error message when no response (e.g. CORS misconfiguration)', async () => {
vi.mocked(authApi.login).mockRejectedValue(new Error('Network Error'))
render(<LoginPage />)
fireEvent.change(screen.getByLabelText('Username'), { target: { value: 'admin' } })
fireEvent.change(screen.getByLabelText('Password'), { target: { value: 'admin' } })
fireEvent.submit(screen.getByRole('button', { name: /sign in/i }).closest('form')!)
await waitFor(() => {
expect(screen.getByText(/Could not reach the server/)).toBeDefined()
})
})
it('clears previous error before each new attempt', async () => {
vi.mocked(authApi.login)
.mockRejectedValueOnce({ response: { status: 401 } })
.mockRejectedValueOnce({ response: { status: 401 } })
.mockRejectedValueOnce(new Error('401'))
.mockRejectedValueOnce(new Error('401'))
render(<LoginPage />)
const form = screen.getByRole('button', { name: /sign in/i }).closest('form')!
fireEvent.change(screen.getByLabelText('Username'), { target: { value: 'admin' } })
@@ -1,70 +0,0 @@
import { useViewport } from '@xyflow/react'
import type { Guide } from '@/utils/alignment'
interface AlignmentGuidesProps {
guides: Guide[]
color?: string
}
/**
* SVG overlay that draws alignment guide lines on top of the React Flow canvas.
* Coordinates are in canvas (flow) space; we read the viewport transform to
* project them into screen space so lines stay locked to nodes when the user
* pans or zooms.
*/
export function AlignmentGuides({ guides, color = '#00d4ff' }: AlignmentGuidesProps) {
const { x: vx, y: vy, zoom } = useViewport()
if (guides.length === 0) return null
return (
<svg
style={{
position: 'absolute',
inset: 0,
width: '100%',
height: '100%',
pointerEvents: 'none',
zIndex: 5,
overflow: 'visible',
}}
>
{guides.map((g, i) => {
if (g.axis === 'x') {
const x = g.position * zoom + vx
const y1 = g.start * zoom + vy
const y2 = g.end * zoom + vy
return (
<line
key={`x-${i}-${g.position}`}
x1={x}
y1={y1}
x2={x}
y2={y2}
stroke={color}
strokeWidth={1}
strokeDasharray="4 3"
shapeRendering="crispEdges"
/>
)
}
const y = g.position * zoom + vy
const x1 = g.start * zoom + vx
const x2 = g.end * zoom + vx
return (
<line
key={`y-${i}-${g.position}`}
x1={x1}
y1={y}
x2={x2}
y2={y}
stroke={color}
strokeWidth={1}
strokeDasharray="4 3"
shapeRendering="crispEdges"
/>
)
})}
</svg>
)
}
@@ -1,106 +1,40 @@
import { useCallback, useEffect, useMemo, useRef, useState } from 'react'
import { useCallback } from 'react'
import {
ReactFlow,
Background,
Controls,
ControlButton,
BackgroundVariant,
ConnectionMode,
SelectionMode,
useReactFlow,
type Node,
type Edge,
type Connection,
} from '@xyflow/react'
import { MousePointer2, Hand } from 'lucide-react'
import '@xyflow/react/dist/style.css'
import { useCanvasStore } from '@/stores/canvasStore'
import { useThemeStore } from '@/stores/themeStore'
import { THEMES } from '@/utils/themes'
import { computeCollapseInfo, rewireEdgesForCollapse } from '@/utils/collapseFilter'
import { nodeTypes } from './nodes/nodeTypes'
import { edgeTypes } from './edges/edgeTypes'
import { SearchBar } from './SearchBar'
import { AlignmentGuides } from './AlignmentGuides'
import { useAlignmentGuides } from '@/hooks/useAlignmentGuides'
import type { NodeData, EdgeData } from '@/types'
interface CanvasContainerProps {
onConnect?: (connection: Connection) => void
onEdgeDoubleClick?: (edge: Edge<EdgeData>) => void
onNodeDoubleClick?: (node: Node<NodeData>) => void
onNodeDragStart?: () => void
onRequestAddToGroup?: (payload: { nodeId: string; groupId: string }) => void
onRequestAddToContainer?: (payload: { nodeId: string; containerId: string }) => void
onOpenPending?: (deviceId: string) => void
}
export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, onNodeDoubleClick, onNodeDragStart, onRequestAddToGroup, onRequestAddToContainer, onOpenPending }: CanvasContainerProps) {
const [lassoMode, setLassoMode] = useState(true)
export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, onNodeDragStart }: CanvasContainerProps) {
const {
nodes, edges,
onNodesChange, onEdgesChange,
setSelectedNode, snapshotHistory,
fitViewPending, clearFitViewPending,
copySelectedNodes, pasteNodes,
setSelectedNode,
} = useCanvasStore()
const { fitView, screenToFlowPosition, getIntersectingNodes } = useReactFlow<Node<NodeData>>()
// Track the last cursor position over the canvas so paste lands under it.
const cursorRef = useRef<{ x: number; y: number } | null>(null)
const onMouseMove = useCallback((e: React.MouseEvent) => {
cursorRef.current = { x: e.clientX, y: e.clientY }
}, [])
// Copy / paste shortcuts. Registered here (inside ReactFlowProvider) so paste
// can project the cursor / viewport center into flow coordinates.
useEffect(() => {
const handler = (e: KeyboardEvent) => {
if (!(e.ctrlKey || e.metaKey)) return
const el = e.target as HTMLElement
const isInput = el.tagName === 'INPUT' || el.tagName === 'TEXTAREA' || el.isContentEditable
if (isInput) return
if (e.key === 'c') {
copySelectedNodes()
} else if (e.key === 'v') {
const screen = cursorRef.current ?? { x: window.innerWidth / 2, y: window.innerHeight / 2 }
pasteNodes(screenToFlowPosition(screen))
}
}
window.addEventListener('keydown', handler)
return () => window.removeEventListener('keydown', handler)
}, [copySelectedNodes, pasteNodes, screenToFlowPosition])
// Fit view after canvas loads (fitViewPending is set by loadCanvas)
useEffect(() => {
if (!fitViewPending || nodes.length === 0) return
const id = setTimeout(() => {
fitView({ padding: 0.12, duration: 350 })
clearFitViewPending()
}, 50)
return () => clearTimeout(id)
}, [fitViewPending, nodes.length, fitView, clearFitViewPending])
const activeTheme = useThemeStore((s) => s.activeTheme)
const theme = THEMES[activeTheme]
// Filter nodes and edges based on collapsed state (memoized — O(n)).
const collapseInfo = useMemo(() => computeCollapseInfo(nodes), [nodes])
const visibleNodes = useMemo(
() => nodes.filter((n) => collapseInfo.visibleIds.has(n.id)),
[nodes, collapseInfo],
)
const visibleEdges = useMemo(
() => rewireEdgesForCollapse(edges, nodes, collapseInfo.visibleIds, collapseInfo.hiddenBy),
[edges, nodes, collapseInfo],
)
const onNodeClick = useCallback((e: React.MouseEvent, node: Node<NodeData>) => {
if (e.ctrlKey || e.metaKey) {
setSelectedNode(null)
} else {
setSelectedNode(node.id)
}
const onNodeClick = useCallback((_: React.MouseEvent, node: Node<NodeData>) => {
setSelectedNode(node.id)
}, [setSelectedNode])
const onPaneClick = useCallback(() => {
@@ -111,89 +45,35 @@ export function CanvasContainer({ onConnect: onConnectProp, onEdgeDoubleClick, o
onEdgeDoubleClick?.(edge)
}, [onEdgeDoubleClick])
const handleNodeDoubleClick = useCallback((_: React.MouseEvent, node: Node<NodeData>) => {
onNodeDoubleClick?.(node)
}, [onNodeDoubleClick])
const handleBeforeDelete = useCallback(async () => {
snapshotHistory()
return true
}, [snapshotHistory])
const isValidConnection = useCallback(
(connection: { source: string | null; target: string | null }) => connection.source !== connection.target,
[]
)
const { guides, onNodeDrag, onNodeDragStop } = useAlignmentGuides()
// Drop a top-level node onto a group → ask App to confirm adding it. Runs
// before the alignment snap so detection uses the dropped position.
const handleNodeDragStop = useCallback<NonNullable<typeof onNodeDragStop>>((event, dragNode, dragNodes) => {
if (dragNode && !dragNode.parentId &&
dragNode.data.type !== 'group' && dragNode.data.type !== 'groupRect') {
const intersecting = getIntersectingNodes(dragNode)
const group = intersecting.find((n) => n.data.type === 'group')
if (group) {
onRequestAddToGroup?.({ nodeId: dragNode.id, groupId: group.id })
} else {
// Any node in container_mode (proxmox, docker_host, …) accepts children.
const container = intersecting.find((n) => n.id !== dragNode.id && n.data.container_mode === true)
if (container) onRequestAddToContainer?.({ nodeId: dragNode.id, containerId: container.id })
}
}
onNodeDragStop(event, dragNode, dragNodes)
}, [onRequestAddToGroup, onRequestAddToContainer, getIntersectingNodes, onNodeDragStop])
return (
<div className="w-full h-full" style={{ background: theme.colors.canvasBackground }} onMouseMove={onMouseMove}>
<div className="w-full h-full" style={{ background: theme.colors.canvasBackground }}>
<ReactFlow
nodes={visibleNodes}
edges={visibleEdges}
nodes={nodes}
edges={edges}
onNodesChange={onNodesChange}
onEdgesChange={onEdgesChange}
onConnect={onConnectProp}
onNodeClick={onNodeClick}
onPaneClick={onPaneClick}
onEdgeDoubleClick={handleEdgeDoubleClick}
onNodeDoubleClick={handleNodeDoubleClick}
onNodeDragStart={onNodeDragStart}
onNodeDrag={onNodeDrag}
onNodeDragStop={handleNodeDragStop}
nodeTypes={nodeTypes}
edgeTypes={edgeTypes}
deleteKeyCode={['Backspace', 'Delete']}
onBeforeDelete={handleBeforeDelete}
selectionOnDrag={lassoMode}
panOnDrag={lassoMode ? [1, 2] : true}
panActivationKeyCode="Space"
selectionMode={SelectionMode.Partial}
multiSelectionKeyCode={['Meta', 'Control']}
minZoom={0.25}
maxZoom={2.5}
snapToGrid
snapGrid={[8, 8]}
snapGrid={[16, 16]}
fitView
colorMode={theme.colors.reactFlowColorMode}
elevateNodesOnSelect={false}
connectionMode={ConnectionMode.Loose}
isValidConnection={isValidConnection}
isValidConnection={(connection) => connection.source !== connection.target}
>
<Background
variant={BackgroundVariant.Dots}
gap={16}
gap={24}
size={1}
color={theme.colors.canvasDotColor}
/>
<SearchBar onOpenPending={onOpenPending} />
<AlignmentGuides guides={guides} />
<Controls>
<ControlButton
onClick={() => setLassoMode((m) => !m)}
title={lassoMode ? 'Switch to pan mode (Space to pan)' : 'Switch to lasso mode'}
>
{lassoMode ? <MousePointer2 size={12} /> : <Hand size={12} />}
</ControlButton>
</Controls>
<Controls />
</ReactFlow>
</div>
)
@@ -1,220 +0,0 @@
import { useState, useEffect, useRef } from 'react'
import { useReactFlow } from '@xyflow/react'
import { Search, X } from 'lucide-react'
import { useCanvasStore } from '@/stores/canvasStore'
import { scanApi } from '@/api/client'
import { NODE_TYPE_LABELS } from '@/types'
import type { PendingDevice } from '@/components/modals/PendingDeviceModal'
interface SearchBarProps {
onOpenPending?: (deviceId: string) => void
}
export function SearchBar({ onOpenPending }: SearchBarProps) {
const [open, setOpen] = useState(false)
const [query, setQuery] = useState('')
const [pendingDevices, setPendingDevices] = useState<PendingDevice[]>([])
const inputRef = useRef<HTMLInputElement>(null)
const { nodes, setSelectedNode } = useCanvasStore()
const { setCenter } = useReactFlow()
useEffect(() => {
if (!open) return
scanApi.pending().then((res) => setPendingDevices(res.data)).catch(() => {})
}, [open])
useEffect(() => {
const handler = (e: KeyboardEvent) => {
if ((e.ctrlKey || e.metaKey) && e.key === 'f') {
e.preventDefault()
setOpen(true)
}
if (e.key === 'Escape') {
setOpen(false)
setQuery('')
}
}
window.addEventListener('keydown', handler)
return () => window.removeEventListener('keydown', handler)
}, [])
useEffect(() => {
if (open) inputRef.current?.focus()
}, [open])
const q = query.toLowerCase().trim()
const nodeResults = q
? nodes.filter((n) => {
if (n.data.type === 'groupRect') return false
return (
n.data.label?.toLowerCase().includes(q) ||
n.data.ip?.toLowerCase().includes(q) ||
n.data.hostname?.toLowerCase().includes(q) ||
(n.data.services ?? []).some((s) => s.service_name?.toLowerCase().includes(q))
)
})
: []
const pendingResults = q
? pendingDevices.filter((d) =>
d.ip?.toLowerCase().includes(q) ||
d.hostname?.toLowerCase().includes(q) ||
d.friendly_name?.toLowerCase().includes(q) ||
d.ieee_address?.toLowerCase().includes(q) ||
d.services.some((s) =>
s.service_name?.toLowerCase().includes(q) ||
s.category?.toLowerCase().includes(q)
)
).slice(0, 4)
: []
const totalResults = nodeResults.length + pendingResults.length
const goToNode = (id: string) => {
const node = nodes.find((n) => n.id === id)
if (!node) return
setSelectedNode(id)
// For grouped nodes, add parent's absolute position
let absX = node.position.x
let absY = node.position.y
if (node.parentId) {
const parent = nodes.find((n) => n.id === node.parentId)
if (parent) { absX += parent.position.x; absY += parent.position.y }
}
const w = node.measured?.width ?? node.width ?? 200
const h = node.measured?.height ?? node.height ?? 80
setCenter(absX + w / 2, absY + h / 2, { zoom: 1.5, duration: 500 })
setOpen(false)
setQuery('')
}
if (!open) return null
return (
<div
className="nodrag nowheel"
style={{
position: 'absolute',
top: 16,
left: '50%',
transform: 'translateX(-50%)',
zIndex: 1000,
width: 360,
pointerEvents: 'all',
}}
>
<div style={{
background: '#161b22',
border: '1px solid #30363d',
borderRadius: 8,
boxShadow: '0 8px 24px rgba(0,0,0,0.6)',
overflow: 'hidden',
}}>
<div style={{ display: 'flex', alignItems: 'center', gap: 8, padding: '8px 12px' }}>
<Search size={14} style={{ color: '#8b949e', flexShrink: 0 }} />
<input
ref={inputRef}
value={query}
onChange={(e) => setQuery(e.target.value)}
placeholder="Search by name, IP, hostname or service…"
style={{
flex: 1,
background: 'transparent',
border: 'none',
outline: 'none',
color: '#e6edf3',
fontSize: 13,
}}
/>
{query && (
<span style={{ fontSize: 11, color: '#6e7681', flexShrink: 0 }}>
{totalResults} result{totalResults !== 1 ? 's' : ''}
</span>
)}
<button
onClick={() => { setOpen(false); setQuery('') }}
aria-label="Close search"
style={{ color: '#8b949e', background: 'none', border: 'none', cursor: 'pointer', padding: 2 }}
>
<X size={14} />
</button>
</div>
{totalResults > 0 && (
<div style={{ borderTop: '1px solid #30363d', maxHeight: 260, overflowY: 'auto' }}>
{nodeResults.map((n) => (
<button
key={n.id}
onClick={() => goToNode(n.id)}
style={{
width: '100%',
display: 'flex',
alignItems: 'center',
gap: 10,
padding: '7px 12px',
background: 'none',
border: 'none',
cursor: 'pointer',
textAlign: 'left',
}}
onMouseEnter={(e) => (e.currentTarget.style.background = '#21262d')}
onMouseLeave={(e) => (e.currentTarget.style.background = 'none')}
>
<span style={{ fontSize: 12, fontWeight: 600, color: '#e6edf3', flex: 1, overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap' }}>
{n.data.label}
</span>
{n.data.ip && (
<span style={{ fontSize: 11, color: '#8b949e', fontFamily: 'JetBrains Mono, monospace', flexShrink: 0 }}>
{n.data.ip}
</span>
)}
<span style={{ fontSize: 10, color: '#6e7681', flexShrink: 0 }}>
{NODE_TYPE_LABELS[n.data.type] ?? n.data.type}
</span>
</button>
))}
{pendingResults.length > 0 && nodeResults.length > 0 && (
<div style={{ height: 1, background: '#30363d', margin: '2px 0' }} />
)}
{pendingResults.map((d) => {
const serviceName = d.services.find((s) => s.service_name)?.service_name
return (
<button
key={d.id}
onClick={() => { onOpenPending?.(d.id); setOpen(false); setQuery('') }}
style={{
width: '100%',
display: 'flex',
alignItems: 'center',
gap: 10,
padding: '7px 12px',
background: 'none',
border: 'none',
cursor: 'pointer',
textAlign: 'left',
}}
onMouseEnter={(e) => (e.currentTarget.style.background = '#21262d')}
onMouseLeave={(e) => (e.currentTarget.style.background = 'none')}
>
<span style={{ fontSize: 10, color: '#e3b341', fontFamily: 'JetBrains Mono, monospace', flexShrink: 0 }}>pending</span>
<span style={{ fontSize: 12, fontWeight: 600, color: '#e6edf3', flex: 1, overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap' }}>
{d.friendly_name ?? d.hostname ?? d.ip ?? d.ieee_address ?? 'device'}
</span>
<span style={{ fontSize: 11, color: '#8b949e', fontFamily: 'JetBrains Mono, monospace', flexShrink: 0 }}>
{serviceName ?? d.ip ?? d.ieee_address ?? ''}
</span>
</button>
)
})}
</div>
)}
{q && totalResults === 0 && (
<div style={{ borderTop: '1px solid #30363d', padding: '10px 12px', fontSize: 12, color: '#6e7681', textAlign: 'center' }}>
No results for &ldquo;{query}&rdquo;
</div>
)}
</div>
</div>
)
}
@@ -1,47 +0,0 @@
import { describe, it, expect, vi } from 'vitest'
import { render } from '@testing-library/react'
import { AlignmentGuides } from '../AlignmentGuides'
import type { Guide } from '@/utils/alignment'
vi.mock('@xyflow/react', () => ({
useViewport: () => ({ x: 50, y: 100, zoom: 2 }),
}))
describe('AlignmentGuides', () => {
it('renders nothing when no guides', () => {
const { container } = render(<AlignmentGuides guides={[]} />)
expect(container.querySelector('svg')).toBeNull()
})
it('projects an x-axis guide through the viewport transform', () => {
const guides: Guide[] = [{ axis: 'x', position: 100, start: 0, end: 200 }]
const { container } = render(<AlignmentGuides guides={guides} />)
const line = container.querySelector('line')!
// x = position * zoom + vx → 100*2 + 50 = 250
expect(line.getAttribute('x1')).toBe('250')
expect(line.getAttribute('x2')).toBe('250')
// y1 = start * zoom + vy → 0*2 + 100 = 100; y2 = 200*2 + 100 = 500
expect(line.getAttribute('y1')).toBe('100')
expect(line.getAttribute('y2')).toBe('500')
})
it('projects a y-axis guide horizontally', () => {
const guides: Guide[] = [{ axis: 'y', position: 50, start: 10, end: 60 }]
const { container } = render(<AlignmentGuides guides={guides} />)
const line = container.querySelector('line')!
// y = 50*2 + 100 = 200; x1 = 10*2 + 50 = 70; x2 = 60*2 + 50 = 170
expect(line.getAttribute('y1')).toBe('200')
expect(line.getAttribute('y2')).toBe('200')
expect(line.getAttribute('x1')).toBe('70')
expect(line.getAttribute('x2')).toBe('170')
})
it('renders one line per guide', () => {
const guides: Guide[] = [
{ axis: 'x', position: 100, start: 0, end: 200 },
{ axis: 'y', position: 50, start: 10, end: 60 },
]
const { container } = render(<AlignmentGuides guides={guides} />)
expect(container.querySelectorAll('line')).toHaveLength(2)
})
})
@@ -1,274 +0,0 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { render, screen } from '@testing-library/react'
import { Server } from 'lucide-react'
import { BaseNode } from '../nodes/BaseNode'
import type { NodeData } from '@/types'
import type { Node } from '@xyflow/react'
let mockZoom = 1
vi.mock('@xyflow/react', () => ({
Handle: () => null,
Position: { Top: 'top', Bottom: 'bottom' },
NodeResizer: () => null,
useUpdateNodeInternals: () => vi.fn(),
useViewport: () => ({ zoom: mockZoom }),
}))
vi.mock('@/stores/themeStore', () => ({
useThemeStore: (sel: (s: { activeTheme: string }) => unknown) => sel({ activeTheme: 'dark' }),
}))
vi.mock('@/stores/canvasStore', () => ({
useCanvasStore: (sel: (s: { hideIp: boolean; serviceStatuses: Record<string, string> }) => unknown) =>
sel({ hideIp: false, serviceStatuses: {} }),
serviceStatusKey: (nodeId: string, port?: number, protocol?: string) => `${nodeId}:${port ?? ''}/${protocol ?? ''}`,
}))
vi.mock('@/utils/themes', () => ({
THEMES: {
dark: {
colors: {
statusColors: { online: '#39d353', offline: '#f85149', pending: '#e3b341', unknown: '#8b949e' },
nodeSubtextColor: '#8b949e',
nodeLabelColor: '#e6edf3',
nodeIconBackground: '#21262d',
handleBackground: '#30363d',
handleBorder: '#30363d',
},
},
},
}))
vi.mock('@/utils/nodeColors', () => ({
resolveNodeColors: () => ({ background: '#161b22', border: '#30363d', icon: '#00d4ff' }),
}))
vi.mock('@/utils/nodeIcons', () => ({
resolveNodeIcon: (_typeIcon: unknown) => _typeIcon,
isBrandIconKey: (k: string | undefined) => !!k && k.startsWith('brand:'),
}))
vi.mock('@/utils/maskIp', () => ({
maskIp: (ip: string) => ip,
splitIps: (ip: string) => ip ? ip.split(',').map((s: string) => s.trim()).filter(Boolean) : [],
primaryIp: (ip: string) => ip ? ip.split(',')[0].trim() : '',
}))
vi.mock('@/utils/propertyIcons', () => ({
resolvePropertyIcon: (icon: string | null) => icon ? Server : null,
}))
vi.mock('@/utils/handleUtils', () => ({
bottomHandleId: (idx: number) => idx === 0 ? 'bottom' : `bottom-${idx + 1}`,
bottomHandlePositions: (count: number) => {
const c = typeof count === 'number' && count > 0 ? Math.floor(count) : 1
return Array.from({ length: c }, (_, i) => ((i + 1) * 100) / (c + 1))
},
clampBottomHandles: (n: unknown) => typeof n === 'number' ? n : 1,
}))
beforeEach(() => { mockZoom = 1 })
function makeNode(data: Partial<NodeData>): Node<NodeData> {
return {
id: 'n1',
type: data.type ?? 'server',
position: { x: 0, y: 0 },
data: {
label: 'Test Node',
type: 'server',
status: 'online',
services: [],
...data,
},
}
}
function renderBaseNode(data: Partial<NodeData>) {
const node = makeNode(data)
return render(
<BaseNode
id={node.id}
data={node.data}
selected={false}
icon={Server}
type="server"
dragging={false}
zIndex={0}
isConnectable={true}
positionAbsoluteX={0}
positionAbsoluteY={0}
/>
)
}
describe('BaseNode — borderWidth zoom scaling', () => {
beforeEach(() => { mockZoom = 1 })
it('borderWidth is 1px at zoom=1', () => {
mockZoom = 1
const { container } = renderBaseNode({})
expect((container.firstChild as HTMLElement).style.borderWidth).toBe('1px')
})
it('borderWidth scales to 2px at zoom=0.5', () => {
mockZoom = 0.5
const { container } = renderBaseNode({})
expect((container.firstChild as HTMLElement).style.borderWidth).toBe('2px')
})
it('borderWidth is clamped to 1px at zoom=2', () => {
mockZoom = 2
const { container } = renderBaseNode({})
expect((container.firstChild as HTMLElement).style.borderWidth).toBe('1px')
})
it('boxShadow glow ring uses borderWidth when selected + online at zoom=0.5', () => {
mockZoom = 0.5
const node = makeNode({ status: 'online' })
const { container } = render(
<BaseNode id={node.id} data={node.data} selected={true} icon={Server}
type="server" dragging={false} zIndex={0} isConnectable={true}
positionAbsoluteX={0} positionAbsoluteY={0} />
)
expect((container.firstChild as HTMLElement).style.boxShadow).toContain('0 0 0 2px')
})
})
describe('BaseNode — properties rendering', () => {
it('renders visible properties on the node', () => {
renderBaseNode({
properties: [
{ key: 'CPU Model', value: 'i7-12700K', icon: 'Cpu', visible: true },
{ key: 'RAM', value: '32 GB', icon: 'MemoryStick', visible: true },
],
})
expect(screen.getByText('CPU Model')).toBeDefined()
// Value is rendered with a middle-dot prefix: "· 32 GB"
expect(screen.getByText(/32 GB/)).toBeDefined()
})
it('does not render properties with visible=false', () => {
renderBaseNode({
properties: [
{ key: 'Secret', value: 'hidden', icon: null, visible: false },
],
})
expect(screen.queryByText('Secret')).toBeNull()
})
it('renders nothing when properties array is empty', () => {
const { container } = renderBaseNode({ properties: [] })
// No properties section — only the main node card
expect(container.querySelectorAll('.flex.flex-col.gap-1').length).toBe(0)
})
it('renders label and ip regardless of properties', () => {
renderBaseNode({
label: 'My Server',
ip: '192.168.1.10',
properties: [{ key: 'OS', value: 'Debian 12', icon: 'Server', visible: true }],
})
expect(screen.getByText('My Server')).toBeDefined()
expect(screen.getByText('192.168.1.10')).toBeDefined()
expect(screen.getByText('OS')).toBeDefined()
})
})
describe('BaseNode — port numbers (issue #20)', () => {
it('renders a number above each bottom handle when show_port_numbers is on', () => {
renderBaseNode({ bottom_handles: 4, show_port_numbers: true })
expect(screen.getByText('1')).toBeDefined()
expect(screen.getByText('2')).toBeDefined()
expect(screen.getByText('3')).toBeDefined()
expect(screen.getByText('4')).toBeDefined()
})
it('does not render port numbers when show_port_numbers is off', () => {
renderBaseNode({ bottom_handles: 4 })
expect(screen.queryByText('1')).toBeNull()
expect(screen.queryByText('4')).toBeNull()
})
it('numbers match the handle count', () => {
renderBaseNode({ bottom_handles: 2, show_port_numbers: true })
expect(screen.getByText('1')).toBeDefined()
expect(screen.getByText('2')).toBeDefined()
expect(screen.queryByText('3')).toBeNull()
})
})
describe('BaseNode — services visibility toggle', () => {
it('does not render service toggle button on the node', () => {
renderBaseNode({ services: [{ service_name: 'nginx', port: 80, protocol: 'tcp' }] })
expect(screen.queryByTitle('Show services')).toBeNull()
})
it('renders service rows when services are toggled on', () => {
renderBaseNode({
ip: '192.168.1.10',
custom_colors: { show_services: true },
services: [
{ service_name: 'nginx', port: 80, protocol: 'tcp' },
{ service_name: 'ssh', port: 22, protocol: 'tcp' },
],
})
expect(screen.getByText('nginx')).toBeDefined()
expect(screen.getByText('80')).toBeDefined()
expect(screen.getByText('ssh')).toBeDefined()
})
it('renders clickable service links for web services', () => {
renderBaseNode({
ip: '192.168.1.10',
custom_colors: { show_services: true },
services: [{ service_name: 'nginx', port: 80, protocol: 'tcp' }],
})
const link = screen.getByRole('link', { name: /nginx/i }) as HTMLAnchorElement
expect(link.getAttribute('href')).toBe('http://192.168.1.10:80')
})
it('keeps non-web services as non-clickable rows', () => {
renderBaseNode({
ip: '192.168.1.10',
custom_colors: { show_services: true },
services: [{ service_name: 'ssh', port: 22, protocol: 'tcp' }],
})
expect(screen.queryByRole('link', { name: /ssh/i })).toBeNull()
})
})
describe('BaseNode — legacy hardware fallback', () => {
it('renders legacy hardware when properties is undefined and show_hardware is true', () => {
renderBaseNode({
properties: undefined,
show_hardware: true,
cpu_model: 'Intel Xeon E5-2680',
ram_gb: 32,
})
expect(screen.getByText('Intel Xeon E5-2680')).toBeDefined()
})
it('does not render legacy hardware when properties array is present (even if empty)', () => {
renderBaseNode({
properties: [],
show_hardware: true,
cpu_model: 'Intel Xeon E5-2680',
})
// properties array exists → new system, legacy section skipped
expect(screen.queryByText('Intel Xeon E5-2680')).toBeNull()
})
it('does not render legacy hardware when show_hardware is false', () => {
renderBaseNode({
properties: undefined,
show_hardware: false,
cpu_model: 'Intel Xeon E5-2680',
})
expect(screen.queryByText('Intel Xeon E5-2680')).toBeNull()
})
})
@@ -9,9 +9,6 @@ import type { NodeData, EdgeData } from '@/types'
// Capture props passed to ReactFlow so we can test the callbacks
let rfProps: Record<string, unknown> = {}
// Hoisted holder so the mock factory can read the configurable intersection set.
const rf = vi.hoisted(() => ({ intersecting: [] as unknown[] }))
vi.mock('@xyflow/react', () => ({
ReactFlow: (props: Record<string, unknown>) => {
rfProps = props
@@ -19,18 +16,8 @@ vi.mock('@xyflow/react', () => ({
},
Background: () => null,
Controls: () => null,
ControlButton: () => null,
BackgroundVariant: { Dots: 'dots' },
ConnectionMode: { Loose: 'loose' },
SelectionMode: { Partial: 'partial' },
Position: { Top: 'top', Right: 'right', Bottom: 'bottom', Left: 'left' },
useReactFlow: () => ({
fitView: vi.fn(),
screenToFlowPosition: vi.fn(),
getIntersectingNodes: () => rf.intersecting,
setNodes: vi.fn(),
getNodes: () => [],
}),
}))
vi.mock('@xyflow/react/dist/style.css', () => ({}))
@@ -51,7 +38,6 @@ function makeEdge(id: string): Edge<EdgeData> {
describe('CanvasContainer', () => {
beforeEach(() => {
rfProps = {}
rf.intersecting = []
useCanvasStore.setState({ nodes: [], edges: [], selectedNodeId: null })
useThemeStore.setState({ activeTheme: 'default' })
})
@@ -115,24 +101,6 @@ describe('CanvasContainer', () => {
}).not.toThrow()
})
// ── Node double-click ─────────────────────────────────────────────────────
it('calls onNodeDoubleClick prop when a node is double-clicked', () => {
const onNodeDoubleClick = vi.fn()
const node = makeNode('n1')
render(<CanvasContainer onNodeDoubleClick={onNodeDoubleClick} />)
;(rfProps.onNodeDoubleClick as (...args: unknown[]) => unknown)({} as MouseEvent, node)
expect(onNodeDoubleClick).toHaveBeenCalledWith(node)
})
it('does not throw when onNodeDoubleClick is not provided', () => {
const node = makeNode('n1')
render(<CanvasContainer />)
expect(() => {
;(rfProps.onNodeDoubleClick as (...args: unknown[]) => unknown)({} as MouseEvent, node)
}).not.toThrow()
})
// ── Connection validation ─────────────────────────────────────────────────
it('isValidConnection returns false for self-connections', () => {
@@ -164,93 +132,6 @@ describe('CanvasContainer', () => {
expect(rfProps.onNodeDragStart).toBe(onNodeDragStart)
})
// ── Drag onto group → onRequestAddToGroup ─────────────────────────────────
function groupNode(id: string): Node<NodeData> {
return { id, type: 'group', position: { x: 0, y: 0 }, data: { label: id, type: 'group', status: 'unknown', services: [] } }
}
it('fires onRequestAddToGroup when a node is dropped over a group', () => {
const onRequestAddToGroup = vi.fn()
const node = makeNode('n1')
const group = groupNode('g1')
rf.intersecting = [group]
render(<CanvasContainer onRequestAddToGroup={onRequestAddToGroup} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToGroup).toHaveBeenCalledWith({ nodeId: 'n1', groupId: 'g1' })
})
it('does not fire onRequestAddToGroup when no group is under the node', () => {
const onRequestAddToGroup = vi.fn()
const node = makeNode('n1')
rf.intersecting = [makeNode('n2')]
render(<CanvasContainer onRequestAddToGroup={onRequestAddToGroup} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToGroup).not.toHaveBeenCalled()
})
it('does not fire onRequestAddToGroup for an already-parented node', () => {
const onRequestAddToGroup = vi.fn()
const node = { ...makeNode('n1'), parentId: 'gOther' }
rf.intersecting = [groupNode('g1')]
render(<CanvasContainer onRequestAddToGroup={onRequestAddToGroup} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToGroup).not.toHaveBeenCalled()
})
it('does not fire onRequestAddToGroup when the dragged node is itself a group', () => {
const onRequestAddToGroup = vi.fn()
const node = groupNode('g2')
rf.intersecting = [groupNode('g1')]
render(<CanvasContainer onRequestAddToGroup={onRequestAddToGroup} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToGroup).not.toHaveBeenCalled()
})
// ── Drag onto container node → onRequestAddToContainer ────────────────────
function containerNode(id: string, type: NodeData['type'] = 'proxmox'): Node<NodeData> {
return { id, type, position: { x: 0, y: 0 }, data: { label: id, type, status: 'unknown', services: [], container_mode: true } }
}
it('fires onRequestAddToContainer when a node is dropped over a container_mode node', () => {
const onRequestAddToContainer = vi.fn()
const node = makeNode('n1')
rf.intersecting = [containerNode('px1')]
render(<CanvasContainer onRequestAddToContainer={onRequestAddToContainer} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToContainer).toHaveBeenCalledWith({ nodeId: 'n1', containerId: 'px1' })
})
it('prefers a group over a container when both intersect', () => {
const onRequestAddToGroup = vi.fn()
const onRequestAddToContainer = vi.fn()
const node = makeNode('n1')
rf.intersecting = [containerNode('px1'), groupNode('g1')]
render(<CanvasContainer onRequestAddToGroup={onRequestAddToGroup} onRequestAddToContainer={onRequestAddToContainer} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToGroup).toHaveBeenCalledWith({ nodeId: 'n1', groupId: 'g1' })
expect(onRequestAddToContainer).not.toHaveBeenCalled()
})
it('does not fire onRequestAddToContainer for an already-parented node', () => {
const onRequestAddToContainer = vi.fn()
const node = { ...makeNode('n1'), parentId: 'pxOther' }
rf.intersecting = [containerNode('px1')]
render(<CanvasContainer onRequestAddToContainer={onRequestAddToContainer} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToContainer).not.toHaveBeenCalled()
})
it('does not fire onRequestAddToContainer when the target node is not in container_mode', () => {
const onRequestAddToContainer = vi.fn()
const node = makeNode('n1')
rf.intersecting = [makeNode('n2')]
render(<CanvasContainer onRequestAddToContainer={onRequestAddToContainer} />)
;(rfProps.onNodeDragStop as (...args: unknown[]) => unknown)({} as MouseEvent, node, [node])
expect(onRequestAddToContainer).not.toHaveBeenCalled()
})
// ── Canvas settings ───────────────────────────────────────────────────────
it('enables snapToGrid', () => {
@@ -258,75 +139,8 @@ describe('CanvasContainer', () => {
expect(rfProps.snapToGrid).toBe(true)
})
it('sets snapGrid to [8, 8]', () => {
it('sets snapGrid to [16, 16]', () => {
render(<CanvasContainer />)
expect(rfProps.snapGrid).toEqual([8, 8])
})
// ── Delete key ────────────────────────────────────────────────────────────
it('sets deleteKeyCode to include both Backspace and Delete', () => {
render(<CanvasContainer />)
expect(rfProps.deleteKeyCode).toEqual(['Backspace', 'Delete'])
})
// ── Lasso / multi-select ──────────────────────────────────────────────────
it('enables selectionOnDrag for lasso selection', () => {
render(<CanvasContainer />)
expect(rfProps.selectionOnDrag).toBe(true)
})
it('sets panActivationKeyCode to Space', () => {
render(<CanvasContainer />)
expect(rfProps.panActivationKeyCode).toBe('Space')
})
it('sets panOnDrag to [1, 2]', () => {
render(<CanvasContainer />)
expect(rfProps.panOnDrag).toEqual([1, 2])
})
it('sets selectionMode to Partial', () => {
render(<CanvasContainer />)
expect(rfProps.selectionMode).toBe('partial')
})
it('sets multiSelectionKeyCode to Meta and Control', () => {
render(<CanvasContainer />)
expect(rfProps.multiSelectionKeyCode).toEqual(['Meta', 'Control'])
})
it('clears selectedNode (sets null) on Ctrl+click instead of selecting', () => {
const node = makeNode('n1')
useCanvasStore.setState({ nodes: [node], selectedNodeId: 'n1' })
render(<CanvasContainer />)
;(rfProps.onNodeClick as (...args: unknown[]) => unknown)(
{ ctrlKey: true, metaKey: false } as unknown as MouseEvent,
node,
)
expect(useCanvasStore.getState().selectedNodeId).toBeNull()
})
it('clears selectedNode (sets null) on Cmd+click', () => {
const node = makeNode('n1')
useCanvasStore.setState({ nodes: [node], selectedNodeId: 'n1' })
render(<CanvasContainer />)
;(rfProps.onNodeClick as (...args: unknown[]) => unknown)(
{ ctrlKey: false, metaKey: true } as unknown as MouseEvent,
node,
)
expect(useCanvasStore.getState().selectedNodeId).toBeNull()
})
// ── onBeforeDelete snapshot ───────────────────────────────────────────────
it('onBeforeDelete calls snapshotHistory and returns true', async () => {
const snapshotHistory = vi.fn()
useCanvasStore.setState({ snapshotHistory } as unknown as Parameters<typeof useCanvasStore.setState>[0])
render(<CanvasContainer />)
const result = await (rfProps.onBeforeDelete as () => Promise<boolean>)()
expect(snapshotHistory).toHaveBeenCalledOnce()
expect(result).toBe(true)
expect(rfProps.snapGrid).toEqual([16, 16])
})
})
@@ -1,185 +0,0 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { fireEvent, render, screen } from '@testing-library/react'
import { GroupNode } from '../nodes/GroupNode'
import * as canvasStore from '@/stores/canvasStore'
import type { Node } from '@xyflow/react'
import type { NodeData } from '@/types'
vi.mock('@/stores/canvasStore')
vi.mock('@xyflow/react', () => ({
NodeResizer: ({ isVisible }: { isVisible: boolean }) => (
<div data-testid="node-resizer" data-visible={isVisible} />
),
Handle: () => null,
Position: { Top: 'top', Right: 'right', Bottom: 'bottom', Left: 'left' },
useReactFlow: () => ({}),
}))
vi.mock('@xyflow/react/dist/style.css', () => ({}))
function makeGroupNode(overrides: Partial<NodeData> = {}): Node<NodeData> {
return {
id: 'g1',
type: 'group',
position: { x: 0, y: 0 },
width: 400,
height: 250,
data: {
label: 'My Group',
type: 'group',
status: 'unknown',
services: [],
custom_colors: { show_border: true },
...overrides,
},
}
}
function renderGroupNode(props: Partial<Parameters<typeof GroupNode>[0]> = {}, storeNodes: unknown[] = []) {
const node = makeGroupNode(props.data)
vi.mocked(canvasStore.useCanvasStore).mockReturnValue({
nodes: storeNodes,
updateNode: vi.fn(),
snapshotHistory: vi.fn(),
toggleNodeCollapsed: vi.fn(),
} as unknown as ReturnType<typeof canvasStore.useCanvasStore>)
return render(
<GroupNode
id="g1"
data={node.data}
selected={false}
dragging={false}
zIndex={1}
isConnectable={true}
positionAbsoluteX={0}
positionAbsoluteY={0}
{...props}
/>,
)
}
describe('GroupNode', () => {
beforeEach(() => {
vi.clearAllMocks()
})
it('renders the group label when show_border is true', () => {
renderGroupNode()
expect(screen.getByText('My Group')).toBeDefined()
})
it('hides the header when show_border is false and not selected', () => {
renderGroupNode({ data: makeGroupNode({ custom_colors: { show_border: false } }).data, selected: false })
expect(screen.queryByText('My Group')).toBeNull()
})
it('shows header when show_border is false but node is selected', () => {
renderGroupNode({ data: makeGroupNode({ custom_colors: { show_border: false } }).data, selected: true })
expect(screen.getByText('My Group')).toBeDefined()
})
it('shows NodeResizer only when selected', () => {
const { rerender } = renderGroupNode({ selected: false })
expect(screen.getByTestId('node-resizer').getAttribute('data-visible')).toBe('false')
vi.mocked(canvasStore.useCanvasStore).mockReturnValue({
nodes: [],
updateNode: vi.fn(),
snapshotHistory: vi.fn(),
} as unknown as ReturnType<typeof canvasStore.useCanvasStore>)
rerender(
<GroupNode
id="g1"
data={makeGroupNode().data}
selected={true}
dragging={false}
zIndex={1}
isConnectable={true}
positionAbsoluteX={0}
positionAbsoluteY={0}
/>,
)
expect(screen.getByTestId('node-resizer').getAttribute('data-visible')).toBe('true')
})
it('allows dragging from the header while keeping rename controls nodrag', () => {
renderGroupNode({ selected: true })
expect(screen.getByText('My Group').closest('div')).not.toHaveClass('nodrag')
const renameButton = screen.getByTitle('Rename group')
expect(renameButton).toHaveClass('nodrag')
fireEvent.click(renameButton)
expect(screen.getByDisplayValue('My Group')).toHaveClass('nodrag')
})
it('shows online/offline status summary from children', () => {
const storeNodes = [
{ id: 'c1', parentId: 'g1', data: { status: 'online' } },
{ id: 'c2', parentId: 'g1', data: { status: 'offline' } },
{ id: 'c3', parentId: 'other', data: { status: 'online' } }, // different group — excluded
]
renderGroupNode({}, storeNodes)
// Two status indicators: one online, one offline (c3 excluded — wrong parent)
const statusSpans = screen.getAllByText(/● \d+/)
expect(statusSpans).toHaveLength(2)
})
it('does not show status summary when group has no children', () => {
renderGroupNode()
expect(screen.queryByText(/●/)).toBeNull()
})
it('renders a collapse toggle when the group has parentId children', () => {
const storeNodes = [
{ id: 'c1', parentId: 'g1', data: { status: 'online' } },
{ id: 'c2', parentId: 'g1', data: { status: 'online' } },
]
renderGroupNode({}, storeNodes)
expect(screen.getByTitle('Hide 2 items')).toBeDefined()
})
it('flips the toggle title when collapsed', () => {
const storeNodes = [
{ id: 'c1', parentId: 'g1', data: { status: 'online' } },
]
renderGroupNode({ data: makeGroupNode({ collapsed: true }).data }, storeNodes)
expect(screen.getByTitle('Show 1 hidden items')).toBeDefined()
})
it('calls toggleNodeCollapsed when the toggle is clicked', () => {
const toggleNodeCollapsed = vi.fn()
const storeNodes = [{ id: 'c1', parentId: 'g1', data: { status: 'online' } }]
vi.mocked(canvasStore.useCanvasStore).mockReturnValue({
nodes: storeNodes,
updateNode: vi.fn(),
snapshotHistory: vi.fn(),
toggleNodeCollapsed,
} as unknown as ReturnType<typeof canvasStore.useCanvasStore>)
render(
<GroupNode
id="g1"
data={makeGroupNode().data}
selected={false}
dragging={false}
zIndex={1}
isConnectable={true}
positionAbsoluteX={0}
positionAbsoluteY={0}
/>,
)
fireEvent.click(screen.getByTitle('Hide 1 items'))
expect(toggleNodeCollapsed).toHaveBeenCalledWith('g1')
})
it('does not render the toggle when the group has no children', () => {
renderGroupNode()
expect(screen.queryByTitle(/Hide.*items|Show.*hidden/)).toBeNull()
})
})
@@ -1,77 +0,0 @@
import { describe, it, expect, vi } from 'vitest'
import { render, screen } from '@testing-library/react'
import { GroupRectNode } from '../nodes/GroupRectNode'
import type { NodeData } from '@/types'
import type { Node } from '@xyflow/react'
vi.mock('@xyflow/react', () => ({
Handle: ({ id, type }: { id: string; type: string }) => <div data-testid={`handle-${id}`} data-type={type} />,
Position: { Top: 'top', Right: 'right', Bottom: 'bottom', Left: 'left' },
NodeResizer: () => null,
}))
vi.mock('@/stores/canvasStore', () => ({
useCanvasStore: (sel: (s: { setEditingGroupRectId: () => void }) => unknown) =>
sel({ setEditingGroupRectId: vi.fn() }),
}))
function makeNode(overrides: Partial<NodeData> = {}): Node<NodeData> {
return {
id: 'zone1',
type: 'groupRect',
position: { x: 0, y: 0 },
data: { label: 'My Zone', type: 'groupRect', status: 'unknown', services: [], ...overrides },
}
}
function renderZone(overrides: Partial<NodeData> = {}) {
const node = makeNode(overrides)
return render(
<GroupRectNode
id={node.id}
data={node.data}
selected={false}
type="groupRect"
dragging={false}
zIndex={0}
isConnectable={true}
positionAbsoluteX={0}
positionAbsoluteY={0}
/>
)
}
describe('GroupRectNode — handles', () => {
it('renders source handles on all four sides', () => {
renderZone()
expect(screen.getByTestId('handle-zone-top')).toBeDefined()
expect(screen.getByTestId('handle-zone-right')).toBeDefined()
expect(screen.getByTestId('handle-zone-bottom')).toBeDefined()
expect(screen.getByTestId('handle-zone-left')).toBeDefined()
})
it('renders target handles on all four sides', () => {
renderZone()
expect(screen.getByTestId('handle-zone-top-t')).toBeDefined()
expect(screen.getByTestId('handle-zone-right-t')).toBeDefined()
expect(screen.getByTestId('handle-zone-bottom-t')).toBeDefined()
expect(screen.getByTestId('handle-zone-left-t')).toBeDefined()
})
it('renders 8 handles total (4 source + 4 target)', () => {
renderZone()
expect(screen.getAllByTestId(/^handle-zone-/).length).toBe(8)
})
})
describe('GroupRectNode — label', () => {
it('renders inside label by default', () => {
renderZone({ label: 'DMZ' })
expect(screen.getByText('DMZ')).toBeDefined()
})
it('renders no label when label is empty', () => {
renderZone({ label: '' })
expect(screen.queryByText('DMZ')).toBeNull()
})
})
@@ -1,146 +0,0 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { render, screen, fireEvent } from '@testing-library/react'
import { SearchBar } from '../SearchBar'
import * as canvasStore from '@/stores/canvasStore'
vi.mock('@/stores/canvasStore')
vi.mock('@xyflow/react', () => ({
useReactFlow: () => ({ setCenter: vi.fn() }),
}))
function makeNode(id: string, overrides = {}) {
return {
id,
type: 'server',
position: { x: 0, y: 0 },
data: { label: id, type: 'server', status: 'online', services: [], ip: null, hostname: null },
...overrides,
}
}
function setupStore(nodes: unknown[] = []) {
vi.mocked(canvasStore.useCanvasStore).mockReturnValue({
nodes,
setSelectedNode: vi.fn(),
} as unknown as ReturnType<typeof canvasStore.useCanvasStore>)
}
function openSearch() {
fireEvent.keyDown(window, { key: 'f', ctrlKey: true })
}
describe('SearchBar', () => {
beforeEach(() => {
setupStore([])
vi.clearAllMocks()
})
it('is hidden by default', () => {
render(<SearchBar />)
expect(screen.queryByPlaceholderText(/search/i)).toBeNull()
})
it('opens on Ctrl+F', () => {
render(<SearchBar />)
openSearch()
expect(screen.getByPlaceholderText(/search/i)).toBeDefined()
})
it('opens on Cmd+F', () => {
render(<SearchBar />)
fireEvent.keyDown(window, { key: 'f', metaKey: true })
expect(screen.getByPlaceholderText(/search/i)).toBeDefined()
})
it('closes on Escape', () => {
render(<SearchBar />)
openSearch()
fireEvent.keyDown(window, { key: 'Escape' })
expect(screen.queryByPlaceholderText(/search/i)).toBeNull()
})
it('closes when X button is clicked', () => {
render(<SearchBar />)
openSearch()
fireEvent.click(screen.getByLabelText('Close search'))
expect(screen.queryByPlaceholderText(/search/i)).toBeNull()
})
it('filters by label', () => {
setupStore([
makeNode('n1', { data: { label: 'My Router', type: 'router', status: 'online', services: [], ip: null, hostname: null } }),
makeNode('n2', { data: { label: 'My NAS', type: 'nas', status: 'online', services: [], ip: null, hostname: null } }),
])
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: 'router' } })
expect(screen.getByText('My Router')).toBeDefined()
expect(screen.queryByText('My NAS')).toBeNull()
})
it('filters by IP', () => {
setupStore([
makeNode('n1', { data: { label: 'Server A', type: 'server', status: 'online', services: [], ip: '192.168.1.10', hostname: null } }),
makeNode('n2', { data: { label: 'Server B', type: 'server', status: 'online', services: [], ip: '10.0.0.1', hostname: null } }),
])
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: '192.168' } })
expect(screen.getByText('Server A')).toBeDefined()
expect(screen.queryByText('Server B')).toBeNull()
})
it('filters by service name', () => {
setupStore([
makeNode('n1', { data: { label: 'Web Server', type: 'server', status: 'online', services: [{ service_name: 'nginx', port: 80, protocol: 'tcp' }], ip: null, hostname: null } }),
makeNode('n2', { data: { label: 'DB Server', type: 'server', status: 'online', services: [{ service_name: 'mysql', port: 3306, protocol: 'tcp' }], ip: null, hostname: null } }),
])
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: 'nginx' } })
expect(screen.getByText('Web Server')).toBeDefined()
expect(screen.queryByText('DB Server')).toBeNull()
})
it('excludes groupRect nodes from results', () => {
setupStore([
makeNode('gr1', { data: { label: 'DMZ Zone', type: 'groupRect', status: 'unknown', services: [], ip: null, hostname: null } }),
])
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: 'dmz' } })
expect(screen.queryByText('DMZ Zone')).toBeNull()
})
it('shows no-results message when query has no matches', () => {
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: 'zzznomatch' } })
expect(screen.getByText(/no results/i)).toBeDefined()
})
it('calls setSelectedNode when a result is clicked', () => {
const setSelectedNode = vi.fn()
vi.mocked(canvasStore.useCanvasStore).mockReturnValue({
nodes: [makeNode('n1', { data: { label: 'My Server', type: 'server', status: 'online', services: [], ip: null, hostname: null } })],
setSelectedNode,
} as unknown as ReturnType<typeof canvasStore.useCanvasStore>)
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: 'my server' } })
fireEvent.click(screen.getByText('My Server'))
expect(setSelectedNode).toHaveBeenCalledWith('n1')
})
it('shows result count', () => {
setupStore([
makeNode('n1', { data: { label: 'Alpha', type: 'server', status: 'online', services: [], ip: null, hostname: null } }),
makeNode('n2', { data: { label: 'Beta', type: 'server', status: 'online', services: [], ip: null, hostname: null } }),
])
render(<SearchBar />)
openSearch()
fireEvent.change(screen.getByPlaceholderText(/search/i), { target: { value: 'a' } })
expect(screen.getByText(/2 results/i)).toBeDefined()
})
})
@@ -1,77 +0,0 @@
import { describe, it, expect } from 'vitest'
import { render } from '@testing-library/react'
import { ReactFlowProvider } from '@xyflow/react'
import type { EdgeProps, Edge } from '@xyflow/react'
import { HomelableEdge } from '../index'
import type { EdgeData } from '@/types'
/**
* Regression: edge flow animations must use CSS, never SVG SMIL <animate>.
*
* SMIL <animate> keeps running while the tab is hidden and leaks memory in
* Chrome over time (RAM climbed only when the canvas tab was backgrounded).
* CSS animations pause when the tab is hidden and don't leak — so the rendered
* output must contain a CSS `animation` on the path and zero <animate> nodes.
*/
function renderEdge(data: Partial<EdgeData> = {}) {
const props = {
id: 'e1',
source: 'a',
target: 'b',
sourceX: 0,
sourceY: 0,
targetX: 100,
targetY: 100,
sourcePosition: 'bottom',
targetPosition: 'top',
data: { type: 'ethernet', ...data } as EdgeData,
selected: false,
} as unknown as EdgeProps<Edge<EdgeData>>
return render(
<ReactFlowProvider>
<svg>
<HomelableEdge {...props} />
</svg>
</ReactFlowProvider>,
)
}
describe('HomelableEdge animation', () => {
it('renders snake animation as CSS, not SMIL <animate>', () => {
const { container } = renderEdge({ animated: 'snake' })
expect(container.querySelector('animate')).toBeNull()
const animated = Array.from(container.querySelectorAll('path')).find((p) =>
(p.getAttribute('style') ?? '').includes('homelable-snake'),
)
expect(animated).toBeTruthy()
})
it('renders flow animation as CSS, not SMIL <animate>', () => {
const { container } = renderEdge({ animated: 'flow' })
expect(container.querySelector('animate')).toBeNull()
const animated = Array.from(container.querySelectorAll('path')).find((p) =>
(p.getAttribute('style') ?? '').includes('homelable-flow'),
)
expect(animated).toBeTruthy()
})
it('legacy animated:true maps to snake CSS animation', () => {
const { container } = renderEdge({ animated: true })
expect(container.querySelector('animate')).toBeNull()
const animated = Array.from(container.querySelectorAll('path')).find((p) =>
(p.getAttribute('style') ?? '').includes('homelable-snake'),
)
expect(animated).toBeTruthy()
})
it('non-animated edge has no flow animation and no <animate>', () => {
const { container } = renderEdge({ animated: false })
expect(container.querySelector('animate')).toBeNull()
const animated = Array.from(container.querySelectorAll('path')).find((p) => {
const s = p.getAttribute('style') ?? ''
return s.includes('homelable-snake') || s.includes('homelable-flow')
})
expect(animated).toBeUndefined()
})
})
@@ -1,70 +0,0 @@
import { describe, it, expect, vi } from 'vitest'
import { render } from '@testing-library/react'
import { ReactFlowProvider } from '@xyflow/react'
import type { EdgeProps, Edge } from '@xyflow/react'
import type { EdgeData } from '@/types'
/**
* Issue #183 — connection labels must support multiple lines.
*
* The label is a free-text string; newlines entered in the EdgeModal textarea
* are stored verbatim. The rendered label div must preserve those newlines
* (`whitespace-pre-line`) instead of collapsing them into a single line.
*
* <EdgeLabelRenderer> normally portals into a node that only exists inside a
* full <ReactFlow> host, so we stub it to a passthrough to render the label
* markup directly.
*/
vi.mock('@xyflow/react', async (importOriginal) => {
const actual = await importOriginal<typeof import('@xyflow/react')>()
return {
...actual,
EdgeLabelRenderer: ({ children }: { children: React.ReactNode }) => <>{children}</>,
}
})
const { HomelableEdge } = await import('../index')
function renderEdge(data: Partial<EdgeData> = {}) {
const props = {
id: 'e1',
source: 'a',
target: 'b',
sourceX: 0,
sourceY: 0,
targetX: 100,
targetY: 100,
sourcePosition: 'bottom',
targetPosition: 'top',
data: { type: 'ethernet', ...data } as EdgeData,
selected: false,
} as unknown as EdgeProps<Edge<EdgeData>>
return render(
<ReactFlowProvider>
<svg>
<HomelableEdge {...props} />
</svg>
</ReactFlowProvider>,
)
}
describe('HomelableEdge label', () => {
it('renders the label text', () => {
const { getByText } = renderEdge({ label: 'uplink' })
expect(getByText('uplink')).toBeTruthy()
})
it('preserves newlines in the rendered label (issue #183)', () => {
const { container } = renderEdge({ label: 'line one\nline two' })
const label = Array.from(container.querySelectorAll('div.whitespace-pre-line')).find((d) =>
d.textContent === 'line one\nline two',
)
expect(label).toBeTruthy()
})
it('renders no label div when label is empty', () => {
const { container } = renderEdge({ label: undefined })
expect(container.querySelector('div.whitespace-pre-line')).toBeNull()
})
})
@@ -1,17 +0,0 @@
import { describe, it, expect } from 'vitest'
import { edgeTypes } from '../edgeTypes'
import { EDGE_TYPE_LABELS, type EdgeType } from '@/types'
describe('edgeTypes registry', () => {
// Regression (issue #21): an EdgeType missing here makes React Flow fall back
// to its built-in default edge — grey, unstyled, ignoring custom_color.
it('registers a component for every EdgeType', () => {
for (const type of Object.keys(EDGE_TYPE_LABELS) as EdgeType[]) {
expect(edgeTypes[type as keyof typeof edgeTypes]).toBeDefined()
}
})
it('registers fibre', () => {
expect(edgeTypes.fibre).toBeDefined()
})
})

Some files were not shown because too many files have changed in this diff Show More