feat(session): per-owner printer/config storage via cookie-scoped bearer key

Printers and groups (Client) are now scoped to an Owner identified by an opaque
bearer key (secrets.token_urlsafe(32)) stored in an httponly cookie, defaulting
to temporary. First-visit modal offers backup-key download (marks permanent) or
temporary-only choice. /session/restore re-attaches a fresh browser to a saved
key. Every printer-facing route enforces ownership (404 on mismatch, not just
filtering) since printer IDs are sequential ints. Drivers stay global/shared.

On upgrade, pre-existing printer/client rows backfill to a synthetic legacy Owner;
its key is written to {DATA_DIR}/legacy_owner_key.txt for manual restore.

SECURITY: Added Origin/Referer same-origin check on POST /session/restore to
block login-CSRF/session-fixation attacks (cross-site form POST can't re-point
victim's cookie at attacker's Owner without hitting that check first).

Tests: 140 pass (2 deselected: pre-existing locale-flaky, unrelated to this change).
Verified live: modal on first visit, isolation between browsers, backup-key
download and restore flow work end-to-end.

Co-Authored-By: Claude Haiku 4.5 <noreply@anthropic.com>
This commit is contained in:
2026-08-04 11:29:43 +02:00
co-authored by Claude Haiku 4.5
parent 9e46fee312
commit ed41f7f520
27 changed files with 583 additions and 102 deletions
+9 -4
View File
@@ -14,6 +14,7 @@ pytest tests/ -k "test_name" # single test by name
**Run dev server:** **Run dev server:**
```bash ```bash
export DATA_DIR=/tmp/imptune_data export DATA_DIR=/tmp/imptune_data
export COOKIE_SECURE=false # plain HTTP — omit if serving behind TLS
uvicorn imptune.main:app --reload --port 8000 uvicorn imptune.main:app --reload --port 8000
``` ```
@@ -33,18 +34,21 @@ pip install -r requirements.txt -r requirements-dev.txt
ImpTune make printer deploy packages (`.intunewin` for Intune, `.zip` for NinjaRMM) from Windows driver ZIPs + web UI. No external services — single FastAPI + SQLite + Docker volume. ImpTune make printer deploy packages (`.intunewin` for Intune, `.zip` for NinjaRMM) from Windows driver ZIPs + web UI. No external services — single FastAPI + SQLite + Docker volume.
**Request flow:** **Request flow:**
1. Driver upload → `api/drivers.py``services/inf_parser.py` parse INF → `storage/driver_store.py` store by SHA256 → Peewee `Driver` record 1. Driver upload → `api/drivers.py``services/inf_parser.py` parse INF → `storage/driver_store.py` store by SHA256 → Peewee `Driver` record (shared/global — visible to every Owner)
2. Printer config → `api/printers.py``db/models.py` `Printer` record (links Driver FK) 2. Printer config → `api/printers.py``db/models.py` `Printer` record (links Driver FK, scoped to `request.state.owner`)
3. Icon upload → `api/icons.py` → Pillow validate PNG 256×256 → SHA256 storage → `Icon` record 3. Icon upload → `api/icons.py` → Pillow validate PNG 256×256 → SHA256 storage → `Icon` record
4. Package export → `api/packages.py``generators/script_generator.py` render Jinja2 PS1 templates → `generators/intunewin_builder.py` encrypt ZIP (AES-256-CBC + HMAC-SHA256) 4. Package export → `api/packages.py``generators/script_generator.py` render Jinja2 PS1 templates → `generators/intunewin_builder.py` encrypt ZIP (AES-256-CBC + HMAC-SHA256)
**Key modules:** **Key modules:**
- `imptune/config.py``DATA_DIR`, `DB_PATH`, `DRIVERS_DIR`, `ICONS_DIR` from env - `imptune/config.py``DATA_DIR`, `DB_PATH`, `DRIVERS_DIR`, `ICONS_DIR`, `COOKIE_SECURE` from env
- `imptune/db/database.py` — SQLite WAL mode + `foreign_keys=1`; all models inherit `BaseModel` - `imptune/db/database.py` — SQLite WAL mode + `foreign_keys=1`; all models inherit `BaseModel`; `init_db()` also backfills `owner_id` on pre-per-owner-scoping DBs into a synthetic legacy `Owner` (key written to `{DATA_DIR}/legacy_owner_key.txt`)
- `imptune/services/session.py``OwnerSessionMiddleware` resolves `request.state.owner` from the `imptune_owner_key` cookie, creating one on first visit (skips `/health`)
- `imptune/services/inf_parser.py` — auto-detect encoding (UTF-16/UTF-8/cp1252), resolve `%TOKEN%` from `[Strings]`, handle multi-model INFs - `imptune/services/inf_parser.py` — auto-detect encoding (UTF-16/UTF-8/cp1252), resolve `%TOKEN%` from `[Strings]`, handle multi-model INFs
- `imptune/generators/intunewin_builder.py` — Python-native `.intunewin` (ZIP-in-ZIP); IV 16 bytes (not 32); match reference tool `1.8.6.0` output - `imptune/generators/intunewin_builder.py` — Python-native `.intunewin` (ZIP-in-ZIP); IV 16 bytes (not 32); match reference tool `1.8.6.0` output
- `imptune/templates/scripts/` — Jinja2 templates for `install.ps1`, `uninstall.ps1`, `detect.ps1` - `imptune/templates/scripts/` — Jinja2 templates for `install.ps1`, `uninstall.ps1`, `detect.ps1`
**Per-owner storage:** `Printer`/`Client` (groups) are scoped to an `Owner` identified by an opaque bearer key in a cookie — no accounts. `Driver` stays global/shared. Every route taking a `printer_id`/`client_id` must filter/check `.owner == request.state.owner` (404, not 403, on mismatch) — printer IDs are small sequential ints, so a list-only filter isn't enough. Onboarding modal (`templates/base.html`, gated on `request.state.is_new_owner`) offers "download backup key" (`GET /session/key/download`, marks `Owner.is_permanent`) vs. temporary; `/session/restore` re-attaches a browser to a previously downloaded key. In tests, use the `owner` fixture (`tests/conftest.py`) when creating `Printer`/`Client` rows directly via the ORM so the `client` fixture's cookie-scoped requests can see them.
**UI stack:** Pico CSS + HTMX 2 + Alpine.js 3 + Jinja2 server-side templates. **UI stack:** Pico CSS + HTMX 2 + Alpine.js 3 + Jinja2 server-side templates.
**HTMX pattern:** Forms `hx-post`, swap `#driver-list` / `#printer-list` / `#client-list` targets. Errors return inline HTML fragments (HTTP 400/409) via `_error_response()`. Success return partials from `templates/partials/`. **HTMX pattern:** Forms `hx-post`, swap `#driver-list` / `#printer-list` / `#client-list` targets. Errors return inline HTML fragments (HTTP 400/409) via `_error_response()`. Success return partials from `templates/partials/`.
@@ -65,3 +69,4 @@ ImpTune make printer deploy packages (`.intunewin` for Intune, `.zip` for NinjaR
|-----|---------|---------| |-----|---------|---------|
| `DATA_DIR` | `/data` | Storage root (DB + drivers + icons) | | `DATA_DIR` | `/data` | Storage root (DB + drivers + icons) |
| `PORT` | `8000` | Server port | | `PORT` | `8000` | Server port |
| `COOKIE_SECURE` | `true` | Owner-session cookie `Secure` flag. Set `false` for local plain-HTTP dev (`uvicorn --reload`) or the browser drops the cookie and a new Owner is created on every request. |
+2 -2
View File
@@ -27,7 +27,7 @@ def _error_response(message: str, status_code: int = 400) -> HTMLResponse:
def _render_client_list(request: Request) -> HTMLResponse: def _render_client_list(request: Request) -> HTMLResponse:
"""Render the client list partial for HTMX swap.""" """Render the client list partial for HTMX swap."""
clients = list(Client.select().order_by(Client.name)) clients = list(Client.select().where(Client.owner == request.state.owner).order_by(Client.name))
return templates.TemplateResponse( return templates.TemplateResponse(
request=request, request=request,
name="partials/client_list.html", name="partials/client_list.html",
@@ -47,7 +47,7 @@ def create_client(request: Request, name: str = Form(...)) -> HTMLResponse:
return _error_response("Client name is required.") return _error_response("Client name is required.")
try: try:
Client.create(name=name) Client.create(name=name, owner=request.state.owner)
except IntegrityError: except IntegrityError:
return _error_response(f"Client '{name}' already exists.", status_code=409) return _error_response(f"Client '{name}' already exists.", status_code=409)
+6 -4
View File
@@ -5,7 +5,7 @@ import hashlib
import io import io
from pathlib import Path from pathlib import Path
from fastapi import APIRouter, UploadFile from fastapi import APIRouter, Request, UploadFile
from fastapi.responses import HTMLResponse from fastapi.responses import HTMLResponse
from PIL import Image from PIL import Image
@@ -18,7 +18,7 @@ MAX_ICON_BYTES = 750 * 1024 # 750 KB
@router.post("/{printer_id}/icon", response_class=HTMLResponse) @router.post("/{printer_id}/icon", response_class=HTMLResponse)
def upload_icon(printer_id: int, file: UploadFile) -> HTMLResponse: def upload_icon(request: Request, printer_id: int, file: UploadFile) -> HTMLResponse:
"""Accept a printer icon PNG, validate it, store it, and update the Icon record. """Accept a printer icon PNG, validate it, store it, and update the Icon record.
Validation rules: Validation rules:
@@ -29,8 +29,10 @@ def upload_icon(printer_id: int, file: UploadFile) -> HTMLResponse:
Replaces any previously uploaded icon for this printer. Replaces any previously uploaded icon for this printer.
Returns an HTMX-friendly HTML fragment. Returns an HTMX-friendly HTML fragment.
""" """
# Check printer exists # Check printer exists and belongs to this owner
printer = Printer.get_or_none(Printer.id == printer_id) printer = Printer.get_or_none(
(Printer.id == printer_id) & (Printer.owner == request.state.owner)
)
if printer is None: if printer is None:
return HTMLResponse( return HTMLResponse(
content="<p>Printer not found.</p>", content="<p>Printer not found.</p>",
+9 -9
View File
@@ -6,11 +6,11 @@ import shutil
import tempfile import tempfile
import zipfile import zipfile
from fastapi import APIRouter from fastapi import APIRouter, Request
from fastapi.responses import PlainTextResponse, Response from fastapi.responses import PlainTextResponse, Response
import imptune.config as cfg import imptune.config as cfg
from imptune.db.models import Icon, Printer from imptune.db.models import Icon, Owner, Printer
from imptune.generators.intunewin_builder import build_intunewin from imptune.generators.intunewin_builder import build_intunewin
from imptune.generators.script_generator import render_detect, render_install, render_uninstall from imptune.generators.script_generator import render_detect, render_install, render_uninstall
from imptune.storage.driver_store import DriverStore from imptune.storage.driver_store import DriverStore
@@ -18,9 +18,9 @@ from imptune.storage.driver_store import DriverStore
router = APIRouter(prefix="/printers") router = APIRouter(prefix="/printers")
def _get_printer_and_driver(printer_id: int): def _get_printer_and_driver(printer_id: int, owner: Owner):
"""Fetch printer and validate driver — returns (printer, driver, driver_name) or PlainTextResponse error.""" """Fetch printer (scoped to owner) and validate driver — returns (printer, driver, driver_name) or PlainTextResponse error."""
printer = Printer.get_or_none(Printer.id == printer_id) printer = Printer.get_or_none((Printer.id == printer_id) & (Printer.owner == owner))
if printer is None: if printer is None:
return None, PlainTextResponse("Printer not found", status_code=404) return None, PlainTextResponse("Printer not found", status_code=404)
@@ -48,9 +48,9 @@ def _get_driver_zip_path(driver) -> str:
@router.get("/{printer_id}/packages/ninja") @router.get("/{printer_id}/packages/ninja")
def get_ninja_package(printer_id: int): def get_ninja_package(request: Request, printer_id: int):
"""Download a NinjaRMM-ready ZIP containing install.ps1 and the driver files.""" """Download a NinjaRMM-ready ZIP containing install.ps1 and the driver files."""
result, error = _get_printer_and_driver(printer_id) result, error = _get_printer_and_driver(printer_id, request.state.owner)
if error is not None: if error is not None:
return error return error
@@ -96,9 +96,9 @@ def get_ninja_package(printer_id: int):
@router.get("/{printer_id}/packages/intunewin") @router.get("/{printer_id}/packages/intunewin")
def get_intunewin_package(printer_id: int): def get_intunewin_package(request: Request, printer_id: int):
"""Download a Microsoft Intune .intunewin deployment package.""" """Download a Microsoft Intune .intunewin deployment package."""
result, error = _get_printer_and_driver(printer_id) result, error = _get_printer_and_driver(printer_id, request.state.owner)
if error is not None: if error is not None:
return error return error
+17 -10
View File
@@ -17,12 +17,16 @@ templates = Jinja2Templates(directory=str(Path(__file__).parent.parent / "templa
def dashboard(request: Request): def dashboard(request: Request):
from imptune.db.models import Printer from imptune.db.models import Printer
owner = request.state.owner
recent_printers = list( recent_printers = list(
Printer.select().order_by(Printer.created_at.desc()).limit(5) Printer.select()
.where(Printer.owner == owner)
.order_by(Printer.created_at.desc())
.limit(5)
) )
recent_packages = list( recent_packages = list(
Printer.select() Printer.select()
.where(Printer.driver.is_null(False)) .where((Printer.owner == owner) & Printer.driver.is_null(False))
.order_by(Printer.created_at.desc()) .order_by(Printer.created_at.desc())
.limit(5) .limit(5)
) )
@@ -56,9 +60,11 @@ def drivers_page(request: Request):
def printers_page(request: Request): def printers_page(request: Request):
from imptune.db.models import Client, Driver, Printer from imptune.db.models import Client, Driver, Printer
owner = request.state.owner
query = ( query = (
Printer.select(Printer, Client) Printer.select(Printer, Client)
.join(Client, JOIN.LEFT_OUTER) .join(Client, JOIN.LEFT_OUTER)
.where(Printer.owner == owner)
.order_by(Client.name, Printer.name) .order_by(Client.name, Printer.name)
) )
grouped: dict[str, list] = defaultdict(list) grouped: dict[str, list] = defaultdict(list)
@@ -66,7 +72,7 @@ def printers_page(request: Request):
client_name = p.client.name if p.client_id else "Unassigned" client_name = p.client.name if p.client_id else "Unassigned"
grouped[client_name].append(p) grouped[client_name].append(p)
clients = list(Client.select().order_by(Client.name)) clients = list(Client.select().where(Client.owner == owner).order_by(Client.name))
all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc())) all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc()))
driver_data = [] driver_data = []
@@ -89,7 +95,7 @@ def printers_page(request: Request):
def printers_new_page(request: Request): def printers_new_page(request: Request):
from imptune.db.models import Client, Driver from imptune.db.models import Client, Driver
clients = list(Client.select().order_by(Client.name)) clients = list(Client.select().where(Client.owner == request.state.owner).order_by(Client.name))
all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc())) all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc()))
driver_data = [ driver_data = [
{"driver": d, "names": json.loads(d.driver_desc) if d.driver_desc else []} {"driver": d, "names": json.loads(d.driver_desc) if d.driver_desc else []}
@@ -111,7 +117,7 @@ def printer_detail(request: Request, printer_id: int):
.join(Client, JOIN.LEFT_OUTER) .join(Client, JOIN.LEFT_OUTER)
.switch(Printer) .switch(Printer)
.join(Driver, JOIN.LEFT_OUTER) .join(Driver, JOIN.LEFT_OUTER)
.where(Printer.id == printer_id) .where((Printer.id == printer_id) & (Printer.owner == request.state.owner))
.first() .first()
) )
if printer is None: if printer is None:
@@ -148,7 +154,7 @@ def printer_detail(request: Request, printer_id: int):
def clients_page(request: Request): def clients_page(request: Request):
from imptune.db.models import Client from imptune.db.models import Client
clients = list(Client.select().order_by(Client.name)) clients = list(Client.select().where(Client.owner == request.state.owner).order_by(Client.name))
return templates.TemplateResponse( return templates.TemplateResponse(
request=request, request=request,
name="clients.html", name="clients.html",
@@ -160,7 +166,8 @@ def clients_page(request: Request):
def client_detail(request: Request, client_id: int): def client_detail(request: Request, client_id: int):
from imptune.db.models import Client, Driver, Printer from imptune.db.models import Client, Driver, Printer
client = Client.get_or_none(Client.id == client_id) owner = request.state.owner
client = Client.get_or_none((Client.id == client_id) & (Client.owner == owner))
if client is None: if client is None:
return HTMLResponse( return HTMLResponse(
content="<h1>404 Not Found</h1><p>Client not found.</p>", content="<h1>404 Not Found</h1><p>Client not found.</p>",
@@ -170,12 +177,12 @@ def client_detail(request: Request, client_id: int):
query = ( query = (
Printer.select(Printer, Client) Printer.select(Printer, Client)
.join(Client, JOIN.LEFT_OUTER) .join(Client, JOIN.LEFT_OUTER)
.where(Printer.client == client_id) .where((Printer.client == client_id) & (Printer.owner == owner))
.order_by(Printer.name) .order_by(Printer.name)
) )
grouped = {client.name: list(query)} grouped = {client.name: list(query)}
clients = list(Client.select().order_by(Client.name)) clients = list(Client.select().where(Client.owner == owner).order_by(Client.name))
all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc())) all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc()))
driver_data = [ driver_data = [
{"driver": d, "names": json.loads(d.driver_desc) if d.driver_desc else []} {"driver": d, "names": json.loads(d.driver_desc) if d.driver_desc else []}
@@ -203,7 +210,7 @@ def packages_page(request: Request):
.join(Client, JOIN.LEFT_OUTER) .join(Client, JOIN.LEFT_OUTER)
.switch(Printer) .switch(Printer)
.join(Driver, JOIN.LEFT_OUTER) .join(Driver, JOIN.LEFT_OUTER)
.where(Printer.driver.is_null(False)) .where((Printer.owner == request.state.owner) & Printer.driver.is_null(False))
.order_by(Printer.name) .order_by(Printer.name)
) )
return templates.TemplateResponse( return templates.TemplateResponse(
+23 -5
View File
@@ -33,9 +33,11 @@ def _render_printer_list(request: Request) -> HTMLResponse:
"""Query printers with LEFT JOIN on client and render grouped partial.""" """Query printers with LEFT JOIN on client and render grouped partial."""
import json import json
owner = request.state.owner
query = ( query = (
Printer.select(Printer, Client) Printer.select(Printer, Client)
.join(Client, JOIN.LEFT_OUTER) .join(Client, JOIN.LEFT_OUTER)
.where(Printer.owner == owner)
.order_by(Client.name, Printer.name) .order_by(Client.name, Printer.name)
) )
grouped: dict[str, list[Printer]] = defaultdict(list) grouped: dict[str, list[Printer]] = defaultdict(list)
@@ -43,7 +45,7 @@ def _render_printer_list(request: Request) -> HTMLResponse:
client_name = p.client.name if p.client_id else "Unassigned" client_name = p.client.name if p.client_id else "Unassigned"
grouped[client_name].append(p) grouped[client_name].append(p)
clients = list(Client.select().order_by(Client.name)) clients = list(Client.select().where(Client.owner == owner).order_by(Client.name))
all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc())) all_drivers = list(Driver.select().order_by(Driver.uploaded_at.desc()))
driver_data = [ driver_data = [
{"driver": d, "names": json.loads(d.driver_desc) if d.driver_desc else []} {"driver": d, "names": json.loads(d.driver_desc) if d.driver_desc else []}
@@ -94,8 +96,13 @@ def create_printer(
color_mode_bool = color_mode == "on" color_mode_bool = color_mode == "on"
collate_bool = collate == "on" collate_bool = collate == "on"
# Resolve optional FK IDs owner = request.state.owner
# Resolve optional FK IDs — client must belong to this owner
client_fk = int(client_id) if client_id.strip() else None client_fk = int(client_id) if client_id.strip() else None
if client_fk is not None:
if Client.get_or_none((Client.id == client_fk) & (Client.owner == owner)) is None:
return _error_response(f"Client {client_fk} not found.", status_code=404)
driver_fk = int(driver_id) if driver_id.strip() else None driver_fk = int(driver_id) if driver_id.strip() else None
Printer.create( Printer.create(
@@ -106,6 +113,7 @@ def create_printer(
color_mode=color_mode_bool, color_mode=color_mode_bool,
paper_size=paper_size, paper_size=paper_size,
collate=collate_bool, collate=collate_bool,
owner=owner,
client=client_fk, client=client_fk,
driver=driver_fk, driver=driver_fk,
) )
@@ -116,7 +124,11 @@ def create_printer(
@router.delete("/{printer_id}", response_class=HTMLResponse) @router.delete("/{printer_id}", response_class=HTMLResponse)
def delete_printer(request: Request, printer_id: int) -> HTMLResponse: def delete_printer(request: Request, printer_id: int) -> HTMLResponse:
"""Delete a printer by ID. Returns updated printer list partial.""" """Delete a printer by ID. Returns updated printer list partial."""
deleted = Printer.delete().where(Printer.id == printer_id).execute() deleted = (
Printer.delete()
.where((Printer.id == printer_id) & (Printer.owner == request.state.owner))
.execute()
)
if not deleted: if not deleted:
return _error_response(f"Printer {printer_id} not found.", status_code=404) return _error_response(f"Printer {printer_id} not found.", status_code=404)
@@ -140,7 +152,8 @@ def update_printer(
"""Update an existing printer configuration in-place.""" """Update an existing printer configuration in-place."""
from datetime import UTC, datetime from datetime import UTC, datetime
printer = Printer.get_or_none(Printer.id == printer_id) owner = request.state.owner
printer = Printer.get_or_none((Printer.id == printer_id) & (Printer.owner == owner))
if printer is None: if printer is None:
return _error_response(f"Printer {printer_id} not found.", status_code=404) return _error_response(f"Printer {printer_id} not found.", status_code=404)
@@ -159,6 +172,11 @@ def update_printer(
if paper_size not in _VALID_PAPER: if paper_size not in _VALID_PAPER:
return _error_response(f"Invalid paper size: {paper_size}.") return _error_response(f"Invalid paper size: {paper_size}.")
client_fk = int(client_id) if client_id.strip() else None
if client_fk is not None:
if Client.get_or_none((Client.id == client_fk) & (Client.owner == owner)) is None:
return _error_response(f"Client {client_fk} not found.", status_code=404)
printer.name = name printer.name = name
printer.ip_address = ip_address printer.ip_address = ip_address
printer.port_name = port_name printer.port_name = port_name
@@ -166,7 +184,7 @@ def update_printer(
printer.color_mode = color_mode == "on" printer.color_mode = color_mode == "on"
printer.paper_size = paper_size printer.paper_size = paper_size
printer.collate = collate == "on" printer.collate = collate == "on"
printer.client = int(client_id) if client_id.strip() else None printer.client = client_fk
printer.driver = int(driver_id) if driver_id.strip() else None printer.driver = int(driver_id) if driver_id.strip() else None
printer.updated_at = datetime.now(UTC).replace(tzinfo=None) printer.updated_at = datetime.now(UTC).replace(tzinfo=None)
printer.save() printer.save()
+23 -23
View File
@@ -1,18 +1,18 @@
"""Script download endpoints — generates and serves PowerShell scripts for a printer.""" """Script download endpoints — generates and serves PowerShell scripts for a printer."""
import json import json
from fastapi import APIRouter from fastapi import APIRouter, Request
from fastapi.responses import PlainTextResponse from fastapi.responses import PlainTextResponse
from imptune.db.models import Printer from imptune.db.models import Owner, Printer
from imptune.generators.script_generator import render_detect, render_install, render_uninstall from imptune.generators.script_generator import render_detect, render_install, render_uninstall
router = APIRouter(prefix="/printers") router = APIRouter(prefix="/printers")
def _get_printer_and_driver(printer_id: int): def _get_printer_and_driver(printer_id: int, owner: Owner):
"""Fetch printer and validate driver — returns (printer, driver_name) or PlainTextResponse error.""" """Fetch printer (scoped to owner) and validate driver — returns (printer, driver_name) or PlainTextResponse error."""
printer = Printer.get_or_none(Printer.id == printer_id) printer = Printer.get_or_none((Printer.id == printer_id) & (Printer.owner == owner))
if printer is None: if printer is None:
return None, PlainTextResponse("Printer not found", status_code=404) return None, PlainTextResponse("Printer not found", status_code=404)
@@ -35,8 +35,8 @@ def _get_printer_and_driver(printer_id: int):
return (printer, driver, driver_name), None return (printer, driver, driver_name), None
def _install_response(printer_id: int): def _install_response(printer_id: int, owner: Owner):
result, error = _get_printer_and_driver(printer_id) result, error = _get_printer_and_driver(printer_id, owner)
if error is not None: if error is not None:
return error return error
printer, driver, driver_name = result printer, driver, driver_name = result
@@ -57,8 +57,8 @@ def _install_response(printer_id: int):
) )
def _uninstall_response(printer_id: int): def _uninstall_response(printer_id: int, owner: Owner):
result, error = _get_printer_and_driver(printer_id) result, error = _get_printer_and_driver(printer_id, owner)
if error is not None: if error is not None:
return error return error
printer, driver, driver_name = result printer, driver, driver_name = result
@@ -73,8 +73,8 @@ def _uninstall_response(printer_id: int):
) )
def _detect_response(printer_id: int): def _detect_response(printer_id: int, owner: Owner):
result, error = _get_printer_and_driver(printer_id) result, error = _get_printer_and_driver(printer_id, owner)
if error is not None: if error is not None:
return error return error
printer, driver, driver_name = result printer, driver, driver_name = result
@@ -86,36 +86,36 @@ def _detect_response(printer_id: int):
@router.get("/{printer_id}/scripts/install") @router.get("/{printer_id}/scripts/install")
def get_install_script(printer_id: int): def get_install_script(request: Request, printer_id: int):
"""Download the PowerShell install script for a printer.""" """Download the PowerShell install script for a printer."""
return _install_response(printer_id) return _install_response(printer_id, request.state.owner)
@router.get("/{printer_id}/scripts/install.ps1") @router.get("/{printer_id}/scripts/install.ps1")
def get_install_script_ps1(printer_id: int): def get_install_script_ps1(request: Request, printer_id: int):
"""Download the PowerShell install script for a printer (.ps1 alias).""" """Download the PowerShell install script for a printer (.ps1 alias)."""
return _install_response(printer_id) return _install_response(printer_id, request.state.owner)
@router.get("/{printer_id}/scripts/uninstall") @router.get("/{printer_id}/scripts/uninstall")
def get_uninstall_script(printer_id: int): def get_uninstall_script(request: Request, printer_id: int):
"""Download the PowerShell uninstall script for a printer.""" """Download the PowerShell uninstall script for a printer."""
return _uninstall_response(printer_id) return _uninstall_response(printer_id, request.state.owner)
@router.get("/{printer_id}/scripts/uninstall.ps1") @router.get("/{printer_id}/scripts/uninstall.ps1")
def get_uninstall_script_ps1(printer_id: int): def get_uninstall_script_ps1(request: Request, printer_id: int):
"""Download the PowerShell uninstall script for a printer (.ps1 alias).""" """Download the PowerShell uninstall script for a printer (.ps1 alias)."""
return _uninstall_response(printer_id) return _uninstall_response(printer_id, request.state.owner)
@router.get("/{printer_id}/scripts/detect") @router.get("/{printer_id}/scripts/detect")
def get_detect_script(printer_id: int): def get_detect_script(request: Request, printer_id: int):
"""Download the PowerShell detection script for a printer.""" """Download the PowerShell detection script for a printer."""
return _detect_response(printer_id) return _detect_response(printer_id, request.state.owner)
@router.get("/{printer_id}/scripts/detect.ps1") @router.get("/{printer_id}/scripts/detect.ps1")
def get_detect_script_ps1(printer_id: int): def get_detect_script_ps1(request: Request, printer_id: int):
"""Download the PowerShell detection script for a printer (.ps1 alias).""" """Download the PowerShell detection script for a printer (.ps1 alias)."""
return _detect_response(printer_id) return _detect_response(printer_id, request.state.owner)
+73
View File
@@ -0,0 +1,73 @@
"""Owner session routes — backup-key download and restore-on-new-browser."""
from __future__ import annotations
from pathlib import Path
from fastapi import APIRouter, Form, Request
from fastapi.responses import HTMLResponse, PlainTextResponse, RedirectResponse
from fastapi.templating import Jinja2Templates
import imptune.config as cfg
from imptune.db.models import Owner
from imptune.services.session import COOKIE_MAX_AGE, COOKIE_NAME, is_same_origin
router = APIRouter(prefix="/session")
templates = Jinja2Templates(
directory=str(Path(__file__).parent.parent / "templates")
)
@router.get("/key/download")
def download_key(request: Request) -> PlainTextResponse:
"""Mark the current owner permanent and hand back its key as a backup file."""
owner: Owner = request.state.owner
if not owner.is_permanent:
owner.is_permanent = True
owner.save()
return PlainTextResponse(
content=owner.key,
headers={"Content-Disposition": 'attachment; filename="imptune-backup-key.txt"'},
)
@router.get("/restore", response_class=HTMLResponse)
def restore_page(request: Request, error: str = "") -> HTMLResponse:
return templates.TemplateResponse(
request=request,
name="session_restore.html",
context={"error": error},
)
@router.post("/restore")
def restore_session(request: Request, key: str = Form(...)):
"""Re-associate this browser with a previously downloaded backup key."""
if not is_same_origin(request):
return templates.TemplateResponse(
request=request,
name="session_restore.html",
context={"error": "Request rejected — please submit this form directly from this site."},
status_code=403,
)
owner = Owner.get_or_none(Owner.key == key.strip())
if owner is None:
return templates.TemplateResponse(
request=request,
name="session_restore.html",
context={"error": "Key not found."},
status_code=404,
)
response = RedirectResponse(url="/", status_code=303)
response.set_cookie(
COOKIE_NAME,
owner.key,
max_age=COOKIE_MAX_AGE,
httponly=True,
samesite="lax",
secure=cfg.COOKIE_SECURE,
)
return response
+2
View File
@@ -7,6 +7,8 @@ load_dotenv()
DATA_DIR = os.environ.get("DATA_DIR", "/data") DATA_DIR = os.environ.get("DATA_DIR", "/data")
PORT = int(os.environ.get("PORT", "8000")) PORT = int(os.environ.get("PORT", "8000"))
# Set to "false" for local plain-HTTP dev — browsers drop Secure cookies over HTTP.
COOKIE_SECURE = os.environ.get("COOKIE_SECURE", "true").lower() != "false"
DB_PATH = str(Path(DATA_DIR) / "imptune.db") DB_PATH = str(Path(DATA_DIR) / "imptune.db")
DRIVERS_DIR = str(Path(DATA_DIR) / "drivers") DRIVERS_DIR = str(Path(DATA_DIR) / "drivers")
+37 -2
View File
@@ -16,7 +16,7 @@ def init_db() -> None:
Closes any existing connection before re-initializing so that test Closes any existing connection before re-initializing so that test
fixtures can monkeypatch DB_PATH between test runs. fixtures can monkeypatch DB_PATH between test runs.
""" """
from imptune.db.models import Client, Driver, Printer, Icon from imptune.db.models import Client, Driver, Owner, Printer, Icon
# Re-read DB_PATH each time so tests can patch imptune.config.DB_PATH # Re-read DB_PATH each time so tests can patch imptune.config.DB_PATH
import imptune.config as cfg import imptune.config as cfg
@@ -33,4 +33,39 @@ def init_db() -> None:
}, },
) )
db.connect(reuse_if_open=True) db.connect(reuse_if_open=True)
db.create_tables([Client, Driver, Printer, Icon], safe=True) db.create_tables([Owner, Client, Driver, Printer, Icon], safe=True)
_migrate_owner_column(cfg.DATA_DIR)
def _migrate_owner_column(data_dir: str) -> None:
"""Add owner_id to printer/client if missing, backfilling pre-existing rows.
Runs against a DB created before per-owner scoping existed. Idempotent:
a no-op once the column exists and no rows are left with a NULL owner_id.
"""
from pathlib import Path
from imptune.db.models import Client, Owner, Printer
from imptune.services.session import generate_key
for table in ("printer", "client"):
columns = {row[1] for row in db.execute_sql(f"PRAGMA table_info({table})")}
if "owner_id" not in columns:
db.execute_sql(
f"ALTER TABLE {table} ADD COLUMN owner_id INTEGER REFERENCES owner (id)"
)
orphaned = Printer.select().where(Printer.owner.is_null()).count() or Client.select().where(
Client.owner.is_null()
).count()
if not orphaned:
return
legacy_owner = Owner.create(key=generate_key(), is_permanent=True)
Printer.update(owner=legacy_owner.id).where(Printer.owner.is_null()).execute()
Client.update(owner=legacy_owner.id).where(Client.owner.is_null()).execute()
key_path = Path(data_dir) / "legacy_owner_key.txt"
key_path.write_text(legacy_owner.key, encoding="utf-8")
print(f"[imptune] Pre-existing printers/clients migrated to a legacy owner. "
f"Restore key written to {key_path} — paste it into /session/restore to reclaim them.")
+17 -3
View File
@@ -24,14 +24,27 @@ class BaseModel(Model):
database = db database = db
class Client(BaseModel): class Owner(BaseModel):
"""Represents a deployment target (AD client / OU).""" """Cookie-scoped identity — no accounts, just an opaque bearer key."""
name = CharField(unique=True) key = CharField(unique=True, index=True)
is_permanent = BooleanField(default=False)
created_at = DateTimeField(default=_utcnow)
class Meta:
table_name = "owner"
class Client(BaseModel):
"""Represents a deployment target (AD client / OU), scoped to an Owner."""
name = CharField()
owner = ForeignKeyField(Owner, backref="clients")
created_at = DateTimeField(default=_utcnow) created_at = DateTimeField(default=_utcnow)
class Meta: class Meta:
table_name = "client" table_name = "client"
indexes = ((("owner", "name"), True),)
class Driver(BaseModel): class Driver(BaseModel):
@@ -56,6 +69,7 @@ class Printer(BaseModel):
name = CharField() name = CharField()
ip_address = CharField() ip_address = CharField()
port_name = CharField() port_name = CharField()
owner = ForeignKeyField(Owner, backref="printers")
client = ForeignKeyField(Client, null=True, backref="printers") client = ForeignKeyField(Client, null=True, backref="printers")
driver = ForeignKeyField(Driver, null=True, backref="printers") driver = ForeignKeyField(Driver, null=True, backref="printers")
duplex_mode = CharField(default="OneSided") duplex_mode = CharField(default="OneSided")
+4 -1
View File
@@ -5,9 +5,10 @@ from pathlib import Path
from fastapi import FastAPI from fastapi import FastAPI
from fastapi.staticfiles import StaticFiles from fastapi.staticfiles import StaticFiles
from imptune.api import clients, drivers, health, icons, pages, packages, printers, scripts from imptune.api import clients, drivers, health, icons, pages, packages, printers, scripts, session
from imptune.config import DATA_DIR, DRIVERS_DIR, ICONS_DIR from imptune.config import DATA_DIR, DRIVERS_DIR, ICONS_DIR
from imptune.db.database import db, init_db from imptune.db.database import db, init_db
from imptune.services.session import OwnerSessionMiddleware
@asynccontextmanager @asynccontextmanager
@@ -23,6 +24,7 @@ async def lifespan(app: FastAPI):
app = FastAPI(title="ImpTune", lifespan=lifespan) app = FastAPI(title="ImpTune", lifespan=lifespan)
app.add_middleware(OwnerSessionMiddleware)
# Serve baked-in static assets (pico.min.css, htmx.min.js, alpine.min.js) # Serve baked-in static assets (pico.min.css, htmx.min.js, alpine.min.js)
_static_dir = Path(__file__).parent / "static" _static_dir = Path(__file__).parent / "static"
@@ -37,3 +39,4 @@ app.include_router(clients.router)
app.include_router(scripts.router) app.include_router(scripts.router)
app.include_router(packages.router) app.include_router(packages.router)
app.include_router(icons.router) app.include_router(icons.router)
app.include_router(session.router)
+84
View File
@@ -0,0 +1,84 @@
"""Cookie-scoped Owner session — opaque bearer key, no accounts."""
from __future__ import annotations
import secrets
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
import imptune.config as cfg
from imptune.db.models import Owner
COOKIE_NAME = "imptune_owner_key"
COOKIE_MAX_AGE = 10 * 365 * 24 * 60 * 60 # 10 years
def generate_key() -> str:
return secrets.token_urlsafe(32)
def is_same_origin(request: Request) -> bool:
"""Origin/Referer check — the CSRF guard for /session/restore.
Restoring a key re-points the cookie at a *different* Owner, so unlike
the rest of the app's unprotected POSTs (which only ever mutate the
caller's own data), a forged cross-site POST here is a login-CSRF /
session-fixation vector: an attacker who knows their own key can force
a victim's browser onto the attacker's Owner. Browsers always send
Origin (and usually Referer) on form POSTs, same-site or not, so
requiring a match — and rejecting when both are absent — blocks a plain
auto-submitting HTML form without needing a token.
"""
expected = f"{request.url.scheme}://{request.url.netloc}"
origin = request.headers.get("origin")
if origin is not None:
return origin == expected
referer = request.headers.get("referer")
if referer:
from urllib.parse import urlparse
parsed = urlparse(referer)
return f"{parsed.scheme}://{parsed.netloc}" == expected
return False
class OwnerSessionMiddleware(BaseHTTPMiddleware):
"""Resolves request.state.owner from a cookie, creating one on first visit."""
async def dispatch(self, request: Request, call_next):
if request.url.path == "/health":
return await call_next(request)
key = request.cookies.get(COOKIE_NAME)
owner = Owner.get_or_none(Owner.key == key) if key else None
is_new = owner is None
if owner is None:
owner = Owner.create(key=generate_key(), is_permanent=False)
request.state.owner = owner
request.state.is_new_owner = is_new
response = await call_next(request)
# Routes like /session/restore intentionally set this cookie themselves
# (to a *different* owner than the one this middleware just minted) —
# don't clobber that with the auto-provisioned one.
route_already_set_cookie = any(
header.lower() == b"set-cookie" and value.startswith(f"{COOKIE_NAME}=".encode())
for header, value in response.raw_headers
)
if is_new and not route_already_set_cookie:
response.set_cookie(
COOKIE_NAME,
owner.key,
max_age=COOKIE_MAX_AGE,
httponly=True,
samesite="lax",
secure=cfg.COOKIE_SECURE,
)
return response
+18
View File
@@ -315,6 +315,10 @@
<li><a href="/packages" {% if request.url.path == "/packages" %}class="active"{% endif %} <li><a href="/packages" {% if request.url.path == "/packages" %}class="active"{% endif %}
x-data x-text="$store.i18n.t('packages')">Packages</a></li> x-data x-text="$store.i18n.t('packages')">Packages</a></li>
</ul> </ul>
<ul class="sidebar-nav sidebar-nav-footer">
<li><a href="/session/key/download">Download backup key</a></li>
<li><a href="/session/restore" {% if request.url.path == "/session/restore" %}class="active"{% endif %}>Restore from backup key</a></li>
</ul>
</nav> </nav>
<div class="main-wrapper"> <div class="main-wrapper">
<header class="topbar"> <header class="topbar">
@@ -333,6 +337,20 @@
</div> </div>
</header> </header>
<main class="main-content"> <main class="main-content">
{% if request.state.is_new_owner %}
<dialog open id="session-choice-modal">
<article>
<h3>Keep your printers?</h3>
<p>Your printers, configs, and groups are private to this browser
(drivers stay shared with everyone). If you clear your cookies
without a backup key, they're gone for good.</p>
<footer>
<a role="button" href="/session/key/download">Store permanently (download backup key)</a>
<button class="secondary" onclick="this.closest('dialog').close()">Keep temporary, this browser only</button>
</footer>
</article>
</dialog>
{% endif %}
{% block content %}{% endblock %} {% block content %}{% endblock %}
</main> </main>
</div> </div>
+20
View File
@@ -0,0 +1,20 @@
{% extends "base.html" %}
{% block content %}
<h1>Restore backup key</h1>
<p>Paste the key from your <code>imptune-backup-key.txt</code> backup file to
recover your printers, configs, and groups on this browser.</p>
{% if error %}
<div class="error"><p>{{ error }}</p></div>
{% endif %}
<form method="post" action="/session/restore">
<label>
Backup key
<input type="text" name="key" placeholder="paste your key here" required autofocus>
</label>
<button type="submit">Restore</button>
</form>
{% endblock %}
+15
View File
@@ -12,6 +12,21 @@ def client(tmp_data_dir):
yield c yield c
@pytest.fixture
def owner(client):
"""The Owner the `client` fixture's cookie jar is scoped to.
Triggers the session middleware (any non-/health request creates the
cookie), then resolves the Owner record so tests can create Printer/Client
rows directly via the ORM that the same `client` can then see over HTTP.
"""
from imptune.db.models import Owner
from imptune.services.session import COOKIE_NAME
client.get("/")
return Owner.get(Owner.key == client.cookies[COOKIE_NAME])
@pytest.fixture @pytest.fixture
def tmp_data_dir(tmp_path, monkeypatch): def tmp_data_dir(tmp_path, monkeypatch):
"""Set DATA_DIR to a temp directory so tests don't write to /data.""" """Set DATA_DIR to a temp directory so tests don't write to /data."""
+32
View File
@@ -34,6 +34,10 @@ def live_server(tmp_path_factory):
os.makedirs(cfg.DRIVERS_DIR, exist_ok=True) os.makedirs(cfg.DRIVERS_DIR, exist_ok=True)
os.makedirs(cfg.ICONS_DIR, exist_ok=True) os.makedirs(cfg.ICONS_DIR, exist_ok=True)
# Plain HTTP loopback server — Secure-flagged cookies would be dropped
# inconsistently across browser engines, so disable that flag for E2E.
cfg.COOKIE_SECURE = False
# Also set env var so lifespan handler picks up correct dirs # Also set env var so lifespan handler picks up correct dirs
os.environ["DATA_DIR"] = str(data_dir) os.environ["DATA_DIR"] = str(data_dir)
@@ -68,3 +72,31 @@ def live_server(tmp_path_factory):
server.should_exit = True server.should_exit = True
thread.join(timeout=2.0) thread.join(timeout=2.0)
@pytest.fixture(scope="session")
def _e2e_owner_key(live_server):
"""One Owner shared by every E2E browser context.
E2E specs cover unrelated UI (theme, i18n, port autofill); without a
pre-existing cookie each fresh Playwright context would be treated as a
first-time visitor and blocked by the onboarding modal (a real <dialog>
that intercepts all pointer events until dismissed).
"""
from imptune.services.session import COOKIE_NAME
resp = httpx.get(f"{live_server}/")
return resp.cookies[COOKIE_NAME]
@pytest.fixture(autouse=True)
def _seed_owner_cookie(page, live_server, _e2e_owner_key):
"""Pre-seed the owner cookie so the first-visit modal never appears in specs unrelated to onboarding."""
from urllib.parse import urlparse
from imptune.services.session import COOKIE_NAME
hostname = urlparse(live_server).hostname
page.context.add_cookies(
[{"name": COOKIE_NAME, "value": _e2e_owner_key, "domain": hostname, "path": "/"}]
)
+13 -5
View File
@@ -4,12 +4,16 @@ from __future__ import annotations
import pytest import pytest
def test_printer_edit_modal_open_and_prefill(page, live_server: str) -> None: def test_printer_edit_modal_open_and_prefill(page, live_server: str, _e2e_owner_key: str) -> None:
"""Edit button opens modal with printer's current name pre-filled.""" """Edit button opens modal with printer's current name pre-filled."""
import httpx import httpx
# Create a printer via API from imptune.services.session import COOKIE_NAME
with httpx.Client(base_url=live_server, follow_redirects=True) as api:
# Create a printer via API, under the same owner the page's cookie is seeded with
with httpx.Client(
base_url=live_server, follow_redirects=True, cookies={COOKIE_NAME: _e2e_owner_key}
) as api:
api.post( api.post(
"/printers", "/printers",
data={ data={
@@ -30,11 +34,15 @@ def test_printer_edit_modal_open_and_prefill(page, live_server: str) -> None:
assert name_val == "EditTest Printer" assert name_val == "EditTest Printer"
def test_printer_edit_submit_updates_list(page, live_server: str) -> None: def test_printer_edit_submit_updates_list(page, live_server: str, _e2e_owner_key: str) -> None:
"""Submitting the edit form updates the printer name in the list (no page reload).""" """Submitting the edit form updates the printer name in the list (no page reload)."""
import httpx import httpx
with httpx.Client(base_url=live_server, follow_redirects=True) as api: from imptune.services.session import COOKIE_NAME
with httpx.Client(
base_url=live_server, follow_redirects=True, cookies={COOKIE_NAME: _e2e_owner_key}
) as api:
resp = api.post( resp = api.post(
"/printers", "/printers",
data={ data={
+1
View File
@@ -43,6 +43,7 @@ def test_create_tables(db_env):
init_db() init_db()
tables = db.get_tables() tables = db.get_tables()
assert "owner" in tables
assert "client" in tables assert "client" in tables
assert "driver" in tables assert "driver" in tables
assert "printer" in tables assert "printer" in tables
+12 -11
View File
@@ -28,7 +28,7 @@ def _make_jpeg(width: int = 256, height: int = 256) -> bytes:
return buf.getvalue() return buf.getvalue()
def _create_printer(client): def _create_printer(owner):
"""Create a test Printer record and return it.""" """Create a test Printer record and return it."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -36,15 +36,16 @@ def _create_printer(client):
name="Test Printer", name="Test Printer",
ip_address="10.0.0.1", ip_address="10.0.0.1",
port_name="IP_10.0.0.1", port_name="IP_10.0.0.1",
owner=owner,
) )
class TestIconUpload: class TestIconUpload:
def test_upload_valid_png(self, client, tmp_data_dir): def test_upload_valid_png(self, client, owner, tmp_data_dir):
"""POST a valid 256x256 PNG returns 200 and Icon record created.""" """POST a valid 256x256 PNG returns 200 and Icon record created."""
from imptune.db.models import Icon from imptune.db.models import Icon
printer = _create_printer(client) printer = _create_printer(owner)
png_data = _make_png(256, 256) png_data = _make_png(256, 256)
response = client.post( response = client.post(
f"/printers/{printer.id}/icon", f"/printers/{printer.id}/icon",
@@ -65,9 +66,9 @@ class TestIconUpload:
icon_file = Path(tmp_data_dir) / "icons" / sha256 icon_file = Path(tmp_data_dir) / "icons" / sha256
assert icon_file.exists() assert icon_file.exists()
def test_reject_non_png(self, client, tmp_data_dir): def test_reject_non_png(self, client, owner, tmp_data_dir):
"""POST with a JPEG file returns 422 with PNG format error.""" """POST with a JPEG file returns 422 with PNG format error."""
printer = _create_printer(client) printer = _create_printer(owner)
jpeg_data = _make_jpeg(256, 256) jpeg_data = _make_jpeg(256, 256)
response = client.post( response = client.post(
f"/printers/{printer.id}/icon", f"/printers/{printer.id}/icon",
@@ -76,9 +77,9 @@ class TestIconUpload:
assert response.status_code == 422 assert response.status_code == 422
assert "PNG" in response.text assert "PNG" in response.text
def test_reject_oversized(self, client, tmp_data_dir): def test_reject_oversized(self, client, owner, tmp_data_dir):
"""POST with PNG > 750KB returns 422 with 750 KB error.""" """POST with PNG > 750KB returns 422 with 750 KB error."""
printer = _create_printer(client) printer = _create_printer(owner)
# Craft oversized data: valid PNG bytes followed by padding # Craft oversized data: valid PNG bytes followed by padding
png_bytes = _make_png(256, 256) png_bytes = _make_png(256, 256)
oversized = png_bytes + b"\x00" * (750 * 1024 + 1 - len(png_bytes)) oversized = png_bytes + b"\x00" * (750 * 1024 + 1 - len(png_bytes))
@@ -89,9 +90,9 @@ class TestIconUpload:
assert response.status_code == 422 assert response.status_code == 422
assert "750" in response.text assert "750" in response.text
def test_reject_wrong_dimensions(self, client, tmp_data_dir): def test_reject_wrong_dimensions(self, client, owner, tmp_data_dir):
"""POST with 128x128 PNG returns 422 with 256x256 error.""" """POST with 128x128 PNG returns 422 with 256x256 error."""
printer = _create_printer(client) printer = _create_printer(owner)
png_data = _make_png(128, 128) png_data = _make_png(128, 128)
response = client.post( response = client.post(
f"/printers/{printer.id}/icon", f"/printers/{printer.id}/icon",
@@ -100,11 +101,11 @@ class TestIconUpload:
assert response.status_code == 422 assert response.status_code == 422
assert "256x256" in response.text assert "256x256" in response.text
def test_replace_existing_icon(self, client, tmp_data_dir): def test_replace_existing_icon(self, client, owner, tmp_data_dir):
"""Second upload for same printer replaces the Icon record.""" """Second upload for same printer replaces the Icon record."""
from imptune.db.models import Icon from imptune.db.models import Icon
printer = _create_printer(client) printer = _create_printer(owner)
# First upload # First upload
png1 = _make_png(256, 256) png1 = _make_png(256, 256)
+4 -2
View File
@@ -22,7 +22,7 @@ def driver_zip_bytes():
@pytest.fixture @pytest.fixture
def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes): def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes, owner):
"""Create Driver record (with ZIP on disk) and Printer record linked to it.""" """Create Driver record (with ZIP on disk) and Printer record linked to it."""
import hashlib import hashlib
import os import os
@@ -50,6 +50,7 @@ def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes):
ip_address="192.168.1.100", ip_address="192.168.1.100",
port_name="IP_192.168.1.100", port_name="IP_192.168.1.100",
driver=driver, driver=driver,
owner=owner,
duplex_mode="OneSided", duplex_mode="OneSided",
color_mode=True, color_mode=True,
paper_size="A4", paper_size="A4",
@@ -59,7 +60,7 @@ def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes):
@pytest.fixture @pytest.fixture
def printer_no_driver(tmp_data_dir): def printer_no_driver(tmp_data_dir, owner):
"""Create Printer record with no driver assigned.""" """Create Printer record with no driver assigned."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -68,6 +69,7 @@ def printer_no_driver(tmp_data_dir):
ip_address="10.0.0.1", ip_address="10.0.0.1",
port_name="IP_10.0.0.1", port_name="IP_10.0.0.1",
driver=None, driver=None,
owner=owner,
) )
+11 -6
View File
@@ -189,7 +189,7 @@ def test_create_printer_invalid_ip(client: TestClient) -> None:
assert resp.status_code == 400 assert resp.status_code == 400
def test_printer_detail_shows_driver(client: TestClient) -> None: def test_printer_detail_shows_driver(client: TestClient, owner) -> None:
"""GET /printers/{id} returns 200 with all printer fields and driver name.""" """GET /printers/{id} returns 200 with all printer fields and driver name."""
from imptune.db.models import Driver, Printer from imptune.db.models import Driver, Printer
@@ -204,6 +204,7 @@ def test_printer_detail_shows_driver(client: TestClient) -> None:
ip_address="10.0.1.1", ip_address="10.0.1.1",
port_name="IP_10_0_1_1", port_name="IP_10_0_1_1",
driver=driver_obj, driver=driver_obj,
owner=owner,
) )
resp = client.get(f"/printers/{printer.id}") resp = client.get(f"/printers/{printer.id}")
@@ -220,7 +221,7 @@ def test_printer_detail_not_found(client: TestClient) -> None:
assert resp.status_code == 404 assert resp.status_code == 404
def test_printer_detail_no_driver(client: TestClient) -> None: def test_printer_detail_no_driver(client: TestClient, owner) -> None:
"""GET /printers/{id} for printer with no driver returns 200 with 'No driver assigned'.""" """GET /printers/{id} for printer with no driver returns 200 with 'No driver assigned'."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -229,6 +230,7 @@ def test_printer_detail_no_driver(client: TestClient) -> None:
ip_address="10.0.1.2", ip_address="10.0.1.2",
port_name="IP_10_0_1_2", port_name="IP_10_0_1_2",
driver=None, driver=None,
owner=owner,
) )
resp = client.get(f"/printers/{printer.id}") resp = client.get(f"/printers/{printer.id}")
@@ -311,7 +313,7 @@ def test_printers_library_no_form(client: TestClient) -> None:
assert 'action="/printers" method="post"' not in html assert 'action="/printers" method="post"' not in html
def test_patch_printer(client: TestClient) -> None: def test_patch_printer(client: TestClient, owner) -> None:
"""PATCH /printers/{id} with updated name returns 200, updated name in response, DB updated.""" """PATCH /printers/{id} with updated name returns 200, updated name in response, DB updated."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -319,6 +321,7 @@ def test_patch_printer(client: TestClient) -> None:
name="Original Name", name="Original Name",
ip_address="10.0.2.1", ip_address="10.0.2.1",
port_name="IP_10_0_2_1", port_name="IP_10_0_2_1",
owner=owner,
) )
resp = client.patch( resp = client.patch(
@@ -339,18 +342,19 @@ def test_patch_printer_not_found(client: TestClient) -> None:
assert resp.status_code == 404 assert resp.status_code == 404
def test_client_detail_returns_200(client: TestClient) -> None: def test_client_detail_returns_200(client: TestClient, owner) -> None:
"""GET /clients/{id} returns 200 with client name and assigned printer name.""" """GET /clients/{id} returns 200 with client name and assigned printer name."""
from imptune.db.models import Client, Printer from imptune.db.models import Client, Printer
# Create client # Create client
cl = Client.create(name="Detail Client") cl = Client.create(name="Detail Client", owner=owner)
# Create printer assigned to that client # Create printer assigned to that client
Printer.create( Printer.create(
name="Client Printer", name="Client Printer",
ip_address="10.0.3.1", ip_address="10.0.3.1",
port_name="IP_10_0_3_1", port_name="IP_10_0_3_1",
client=cl, client=cl,
owner=owner,
) )
resp = client.get(f"/clients/{cl.id}") resp = client.get(f"/clients/{cl.id}")
@@ -366,7 +370,7 @@ def test_client_detail_not_found(client: TestClient) -> None:
assert resp.status_code == 404 assert resp.status_code == 404
def test_client_links_in_printer_list(client: TestClient) -> None: def test_client_links_in_printer_list(client: TestClient, owner) -> None:
"""GET /printers with a printer assigned to a client contains href to client detail.""" """GET /printers with a printer assigned to a client contains href to client detail."""
from imptune.db.models import Client, Printer from imptune.db.models import Client, Printer
@@ -382,6 +386,7 @@ def test_client_links_in_printer_list(client: TestClient) -> None:
ip_address="10.0.4.1", ip_address="10.0.4.1",
port_name="IP_10_0_4_1", port_name="IP_10_0_4_1",
client=cl, client=cl,
owner=owner,
) )
resp = client.get("/printers") resp = client.get("/printers")
+4 -2
View File
@@ -22,7 +22,7 @@ def driver_zip_bytes():
@pytest.fixture @pytest.fixture
def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes): def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes, owner):
"""Create Driver record (with ZIP on disk) and Printer record linked to it.""" """Create Driver record (with ZIP on disk) and Printer record linked to it."""
import hashlib import hashlib
import os import os
@@ -49,6 +49,7 @@ def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes):
ip_address="192.168.1.100", ip_address="192.168.1.100",
port_name="IP_192.168.1.100", port_name="IP_192.168.1.100",
driver=driver, driver=driver,
owner=owner,
duplex_mode="OneSided", duplex_mode="OneSided",
color_mode=True, color_mode=True,
paper_size="A4", paper_size="A4",
@@ -58,7 +59,7 @@ def setup_printer_with_driver(tmp_data_dir, driver_zip_bytes):
@pytest.fixture @pytest.fixture
def printer_no_driver(tmp_data_dir): def printer_no_driver(tmp_data_dir, owner):
"""Create Printer record with no driver assigned.""" """Create Printer record with no driver assigned."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -67,6 +68,7 @@ def printer_no_driver(tmp_data_dir):
ip_address="10.0.0.1", ip_address="10.0.0.1",
port_name="IP_10.0.0.1", port_name="IP_10.0.0.1",
driver=None, driver=None,
owner=owner,
) )
+10 -8
View File
@@ -3,7 +3,7 @@ import pytest
from imptune.generators.script_generator import render_detect, render_install, render_uninstall from imptune.generators.script_generator import render_detect, render_install, render_uninstall
def _create_test_driver_and_printer(): def _create_test_driver_and_printer(owner):
"""Helper: create a Driver + Printer for integration tests.""" """Helper: create a Driver + Printer for integration tests."""
from imptune.db.models import Driver, Printer from imptune.db.models import Driver, Printer
@@ -19,6 +19,7 @@ def _create_test_driver_and_printer():
ip_address="10.0.0.1", ip_address="10.0.0.1",
port_name="IP_10.0.0.1", port_name="IP_10.0.0.1",
driver=driver, driver=driver,
owner=owner,
duplex_mode="LongEdge", duplex_mode="LongEdge",
color_mode=True, color_mode=True,
paper_size="A4", paper_size="A4",
@@ -149,25 +150,25 @@ def test_render_detect():
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def test_install_endpoint(client): def test_install_endpoint(client, owner):
"""GET /printers/{id}/scripts/install returns 200 with pnputil in content.""" """GET /printers/{id}/scripts/install returns 200 with pnputil in content."""
_driver, printer = _create_test_driver_and_printer() _driver, printer = _create_test_driver_and_printer(owner)
response = client.get(f"/printers/{printer.id}/scripts/install") response = client.get(f"/printers/{printer.id}/scripts/install")
assert response.status_code == 200 assert response.status_code == 200
assert "pnputil" in response.text assert "pnputil" in response.text
def test_uninstall_endpoint(client): def test_uninstall_endpoint(client, owner):
"""GET /printers/{id}/scripts/uninstall returns 200 with Remove-Printer in content.""" """GET /printers/{id}/scripts/uninstall returns 200 with Remove-Printer in content."""
_driver, printer = _create_test_driver_and_printer() _driver, printer = _create_test_driver_and_printer(owner)
response = client.get(f"/printers/{printer.id}/scripts/uninstall") response = client.get(f"/printers/{printer.id}/scripts/uninstall")
assert response.status_code == 200 assert response.status_code == 200
assert "Remove-Printer" in response.text assert "Remove-Printer" in response.text
def test_detect_endpoint(client): def test_detect_endpoint(client, owner):
"""GET /printers/{id}/scripts/detect returns 200 with Write-Output in content.""" """GET /printers/{id}/scripts/detect returns 200 with Write-Output in content."""
_driver, printer = _create_test_driver_and_printer() _driver, printer = _create_test_driver_and_printer(owner)
response = client.get(f"/printers/{printer.id}/scripts/detect") response = client.get(f"/printers/{printer.id}/scripts/detect")
assert response.status_code == 200 assert response.status_code == 200
assert "Write-Output" in response.text assert "Write-Output" in response.text
@@ -179,7 +180,7 @@ def test_script_endpoint_missing_printer(client):
assert response.status_code == 404 assert response.status_code == 404
def test_script_endpoint_no_driver(client): def test_script_endpoint_no_driver(client, owner):
"""GET /printers/{id}/scripts/install returns 422 when no driver assigned.""" """GET /printers/{id}/scripts/install returns 422 when no driver assigned."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -188,6 +189,7 @@ def test_script_endpoint_no_driver(client):
ip_address="10.0.0.2", ip_address="10.0.0.2",
port_name="IP_10.0.0.2", port_name="IP_10.0.0.2",
driver=None, driver=None,
owner=owner,
duplex_mode="OneSided", duplex_mode="OneSided",
color_mode=True, color_mode=True,
paper_size="A4", paper_size="A4",
+127
View File
@@ -0,0 +1,127 @@
"""Tests for per-owner cookie-scoped session — first visit, isolation, restore."""
from fastapi.testclient import TestClient
from imptune.services.session import COOKIE_NAME
def test_first_visit_sets_cookie_and_new_owner_flag(client):
resp = client.get("/")
assert resp.status_code == 200
assert COOKIE_NAME in client.cookies
assert "session-choice-modal" in resp.text
# Second visit — cookie already present, modal must not reappear
resp2 = client.get("/")
assert "session-choice-modal" not in resp2.text
def test_printers_are_isolated_between_owners(tmp_data_dir):
from imptune.main import app
with TestClient(app) as client_a, TestClient(app) as client_b:
client_a.post(
"/printers",
data={"name": "Owner A Printer", "ip_address": "10.0.0.1", "port_name": "IP_A"},
follow_redirects=False,
)
page_a = client_a.get("/printers")
page_b = client_b.get("/printers")
assert "Owner A Printer" in page_a.text
assert "Owner A Printer" not in page_b.text
def test_owner_b_cannot_reach_owner_a_printer_by_id(tmp_data_dir):
from imptune.main import app
with TestClient(app) as client_a, TestClient(app) as client_b:
client_a.post(
"/printers",
data={"name": "Private Printer", "ip_address": "10.0.0.2", "port_name": "IP_B"},
follow_redirects=False,
)
from imptune.db.models import Printer
printer_id = Printer.get(Printer.name == "Private Printer").id
assert client_b.get(f"/printers/{printer_id}").status_code == 404
assert client_b.patch(f"/printers/{printer_id}", data={"name": "Hijacked"}).status_code == 404
assert client_b.delete(f"/printers/{printer_id}").status_code == 404
assert client_b.get(f"/printers/{printer_id}/scripts/install").status_code == 404
assert client_b.get(f"/printers/{printer_id}/packages/ninja").status_code == 404
def test_download_key_marks_permanent_and_returns_key(client):
from imptune.db.models import Owner
client.get("/")
resp = client.get("/session/key/download")
assert resp.status_code == 200
assert "attachment" in resp.headers["content-disposition"]
assert "imptune-backup-key.txt" in resp.headers["content-disposition"]
key = resp.text
assert len(key) > 20
owner = Owner.get(Owner.key == key)
assert owner.is_permanent is True
def test_restore_with_valid_key_reattaches_owner_data(tmp_data_dir):
from imptune.main import app
with TestClient(app) as client_a:
client_a.get("/")
client_a.post(
"/printers",
data={"name": "Backed Up Printer", "ip_address": "10.0.0.3", "port_name": "IP_C"},
follow_redirects=False,
)
key = client_a.get("/session/key/download").text
with TestClient(app) as client_new:
restore_resp = client_new.post(
"/session/restore",
data={"key": key},
headers={"origin": "http://testserver"},
follow_redirects=False,
)
assert restore_resp.status_code == 303
page = client_new.get("/printers")
assert "Backed Up Printer" in page.text
def test_restore_with_invalid_key_shows_error(client):
resp = client.post(
"/session/restore",
data={"key": "not-a-real-key"},
headers={"origin": "http://testserver"},
)
assert resp.status_code == 404
assert "Key not found" in resp.text
def test_restore_rejects_cross_origin_post(client):
"""CSRF guard: a forged cross-site form POST must not be able to re-point
the victim's cookie at an attacker-known key (login-CSRF / session fixation)."""
resp = client.post(
"/session/restore",
data={"key": "irrelevant"},
headers={"origin": "https://attacker.example"},
)
assert resp.status_code == 403
resp_no_header = client.post("/session/restore", data={"key": "irrelevant"})
assert resp_no_header.status_code == 403
def test_health_endpoint_does_not_create_owner_rows(client):
from imptune.db.models import Owner
before = Owner.select().count()
for _ in range(5):
client.get("/health")
after = Owner.select().count()
assert after == before
+6 -2
View File
@@ -53,7 +53,7 @@ def test_theme_toggle_present(client):
assert "cycle" in response.text or "store.theme" in response.text assert "cycle" in response.text or "store.theme" in response.text
def test_dashboard_shows_recent_printers(client): def test_dashboard_shows_recent_printers(client, owner):
"""Dashboard renders names of recently-created printers from DB.""" """Dashboard renders names of recently-created printers from DB."""
from imptune.db.models import Printer from imptune.db.models import Printer
@@ -61,11 +61,13 @@ def test_dashboard_shows_recent_printers(client):
name="TestPrinter-Alpha", name="TestPrinter-Alpha",
ip_address="10.0.0.1", ip_address="10.0.0.1",
port_name="IP_10.0.0.1", port_name="IP_10.0.0.1",
owner=owner,
) )
Printer.create( Printer.create(
name="TestPrinter-Beta", name="TestPrinter-Beta",
ip_address="10.0.0.2", ip_address="10.0.0.2",
port_name="IP_10.0.0.2", port_name="IP_10.0.0.2",
owner=owner,
) )
response = client.get("/") response = client.get("/")
@@ -85,7 +87,7 @@ def test_theme_toggle_present(client):
assert "cycle" in response.text or "store.theme" in response.text assert "cycle" in response.text or "store.theme" in response.text
def test_dashboard_shows_recent_packages(client): def test_dashboard_shows_recent_packages(client, owner):
"""Dashboard recent-packages section shows only printers with a driver assigned.""" """Dashboard recent-packages section shows only printers with a driver assigned."""
from imptune.db.models import Driver, Printer from imptune.db.models import Driver, Printer
@@ -100,11 +102,13 @@ def test_dashboard_shows_recent_packages(client):
ip_address="10.0.0.2", ip_address="10.0.0.2",
port_name="IP_10.0.0.2", port_name="IP_10.0.0.2",
driver=driver, driver=driver,
owner=owner,
) )
Printer.create( Printer.create(
name="PkgPrinter-NoDriver", name="PkgPrinter-NoDriver",
ip_address="10.0.0.3", ip_address="10.0.0.3",
port_name="IP_10.0.0.3", port_name="IP_10.0.0.3",
owner=owner,
) )
response = client.get("/") response = client.get("/")
+3 -2
View File
@@ -18,7 +18,7 @@ def driver_zip_bytes():
def test_upload_then_ninja_export_finds_driver_on_disk( def test_upload_then_ninja_export_finds_driver_on_disk(
client, tmp_data_dir, driver_zip_bytes client, tmp_data_dir, driver_zip_bytes, owner
): ):
from imptune.db.models import Client, Driver, Printer from imptune.db.models import Client, Driver, Printer
@@ -33,13 +33,14 @@ def test_upload_then_ninja_export_finds_driver_on_disk(
driver.driver_desc = json.dumps(["HP LaserJet Pro"]) driver.driver_desc = json.dumps(["HP LaserJet Pro"])
driver.save() driver.save()
tenant = Client.create(name="Acme Corp") tenant = Client.create(name="Acme Corp", owner=owner)
printer = Printer.create( printer = Printer.create(
name="Round Trip Printer", name="Round Trip Printer",
ip_address="192.168.1.50", ip_address="192.168.1.50",
port_name="IP_192.168.1.50", port_name="IP_192.168.1.50",
client=tenant, client=tenant,
driver=driver, driver=driver,
owner=owner,
) )
resp = client.get(f"/printers/{printer.id}/packages/ninja") resp = client.get(f"/printers/{printer.id}/packages/ninja")