From a80738f3491a9fa50c4ea1bc6acd12dd4173ad57 Mon Sep 17 00:00:00 2001 From: wind Date: Fri, 27 Feb 2026 21:23:57 +0100 Subject: [PATCH 1/2] =?UTF-8?q?feat:=20complete=20SQLite=20device=20invent?= =?UTF-8?q?ory=20=E2=80=94=20vendor=20field,=20unique=20constraints,=20ups?= =?UTF-8?q?ert=20helpers?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Device: add vendor column (arp-scan OUI lookup), UniqueConstraint on ip_address - Port: add UniqueConstraint on (device_id, port_number, protocol) - db.py: add upsert_device() and upsert_port() helpers — insert-or-update keyed on ip_address and (device_id, port, proto) respectively; None fields never overwrite existing data - Tests: 20 tests covering schema, CRUD, constraint enforcement, upsert create/update/idempotency/none-does-not-overwrite; total suite 57 tests at 99% Closes #5 --- backend/app/db.py | 97 +++++++++++- backend/app/models/device.py | 17 ++- backend/tests/test_db.py | 280 ++++++++++++++++++++++++++++++----- 3 files changed, 354 insertions(+), 40 deletions(-) diff --git a/backend/app/db.py b/backend/app/db.py index 7df8297..ce4446f 100644 --- a/backend/app/db.py +++ b/backend/app/db.py @@ -1,9 +1,15 @@ -"""Database connection and initialisation.""" +"""Database connection, initialisation, and session helpers.""" + +from __future__ import annotations import os +from typing import TYPE_CHECKING from sqlalchemy import create_engine -from sqlalchemy.orm import DeclarativeBase, sessionmaker +from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker + +if TYPE_CHECKING: + from app.models.device import Device DATABASE_URL = os.getenv("DATABASE_URL", "sqlite:///./networkcrawler.db") @@ -33,3 +39,90 @@ def get_db(): yield db finally: db.close() + + +def upsert_device( + session: Session, + *, + ip_address: str, + mac_address: str | None = None, + vendor: str | None = None, + hostname: str | None = None, + os_guess: str | None = None, +) -> "Device": + """Insert or update a Device row keyed on ip_address. + + If a Device with the given ip_address already exists, only non-None + fields are written so that richer data from a previous scan is never + overwritten with None. The caller is responsible for committing. + + Returns the Device instance (either existing or newly created). + """ + from app.models.device import Device + from sqlalchemy import select + + stmt = select(Device).where(Device.ip_address == ip_address) + device: Device | None = session.execute(stmt).scalar_one_or_none() + + if device is None: + device = Device( + ip_address=ip_address, + mac_address=mac_address, + vendor=vendor, + hostname=hostname, + os_guess=os_guess, + ) + session.add(device) + else: + if mac_address is not None: + device.mac_address = mac_address + if vendor is not None: + device.vendor = vendor + if hostname is not None: + device.hostname = hostname + if os_guess is not None: + device.os_guess = os_guess + + return device + + +def upsert_port( + session: Session, + *, + device_id: int, + port_number: int, + protocol: str = "tcp", + service_name: str | None = None, + version_banner: str | None = None, +) -> "Device": + """Insert or update a Port row keyed on (device_id, port_number, protocol). + + Non-None fields overwrite existing values. Caller must commit. + Returns the Port instance. + """ + from app.models.device import Port + from sqlalchemy import select + + stmt = select(Port).where( + Port.device_id == device_id, + Port.port_number == port_number, + Port.protocol == protocol, + ) + port = session.execute(stmt).scalar_one_or_none() + + if port is None: + port = Port( + device_id=device_id, + port_number=port_number, + protocol=protocol, + service_name=service_name, + version_banner=version_banner, + ) + session.add(port) + else: + if service_name is not None: + port.service_name = service_name + if version_banner is not None: + port.version_banner = version_banner + + return port diff --git a/backend/app/models/device.py b/backend/app/models/device.py index aa4b722..91273ea 100644 --- a/backend/app/models/device.py +++ b/backend/app/models/device.py @@ -1,17 +1,24 @@ -"""Device and Port SQLAlchemy models (stub — full implementation in Phase 2).""" +"""Device and Port SQLAlchemy ORM models.""" -from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, func +from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, UniqueConstraint, func from sqlalchemy.orm import relationship from app.db import Base class Device(Base): + """A network device discovered by arp-scan and/or nmap.""" + __tablename__ = "devices" + __table_args__ = ( + # One row per IP address — upserts update in place rather than inserting duplicates. + UniqueConstraint("ip_address", name="uq_devices_ip_address"), + ) id = Column(Integer, primary_key=True, index=True) ip_address = Column(String, nullable=False, index=True) mac_address = Column(String, nullable=True) + vendor = Column(String, nullable=True) # hardware vendor from arp-scan OUI lookup hostname = Column(String, nullable=True) os_guess = Column(String, nullable=True) first_seen = Column(DateTime, default=func.now()) @@ -21,7 +28,13 @@ class Device(Base): class Port(Base): + """An open TCP/UDP port observed on a Device during an nmap scan.""" + __tablename__ = "ports" + __table_args__ = ( + # One row per (device, port, protocol) tuple. + UniqueConstraint("device_id", "port_number", "protocol", name="uq_ports_device_port_proto"), + ) id = Column(Integer, primary_key=True, index=True) device_id = Column(Integer, ForeignKey("devices.id"), nullable=False) diff --git a/backend/tests/test_db.py b/backend/tests/test_db.py index b147007..77201a5 100644 --- a/backend/tests/test_db.py +++ b/backend/tests/test_db.py @@ -1,11 +1,13 @@ """ -Unit and integration tests for database initialisation and session management. +Unit and integration tests for database initialisation, session management, +ORM models, and upsert helpers. Markers: unit, integration """ import pytest from sqlalchemy import create_engine, inspect +from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import sessionmaker @@ -24,6 +26,17 @@ def in_memory_engine(): Base.metadata.drop_all(bind=engine) +@pytest.fixture +def session(in_memory_engine): + """Yield a transactional session that is rolled back after each test.""" + factory = sessionmaker(bind=in_memory_engine) + with factory() as s: + yield s + + +# ── Schema tests ────────────────────────────────────────────────────────────── + + @pytest.mark.unit def test_init_db_creates_devices_table(in_memory_engine): inspector = inspect(in_memory_engine) @@ -44,6 +57,7 @@ def test_devices_table_columns(in_memory_engine): "id", "ip_address", "mac_address", + "vendor", "hostname", "os_guess", "first_seen", @@ -58,6 +72,23 @@ def test_ports_table_columns(in_memory_engine): assert {"id", "device_id", "port_number", "protocol", "service_name", "version_banner"} <= cols +@pytest.mark.unit +def test_devices_unique_constraint_exists(in_memory_engine): + inspector = inspect(in_memory_engine) + unique_names = {uc["name"] for uc in inspector.get_unique_constraints("devices")} + assert "uq_devices_ip_address" in unique_names + + +@pytest.mark.unit +def test_ports_unique_constraint_exists(in_memory_engine): + inspector = inspect(in_memory_engine) + unique_names = {uc["name"] for uc in inspector.get_unique_constraints("ports")} + assert "uq_ports_device_port_proto" in unique_names + + +# ── Session helper tests ────────────────────────────────────────────────────── + + @pytest.mark.integration def test_get_db_yields_and_closes(in_memory_engine, monkeypatch): """get_db dependency yields a session and closes it after iteration.""" @@ -69,66 +100,243 @@ def test_get_db_yields_and_closes(in_memory_engine, monkeypatch): gen = db_module.get_db() session = next(gen) assert session is not None - # Exhaust the generator to trigger the finally block (close) try: next(gen) except StopIteration: pass +# ── Basic CRUD tests ────────────────────────────────────────────────────────── + + @pytest.mark.integration -def test_device_crud(in_memory_engine): - """Basic Device create/read round-trip against in-memory DB.""" +def test_device_crud(session): + """Basic Device create/read round-trip.""" from app.models.device import Device - session_factory = sessionmaker(bind=in_memory_engine) - with session_factory() as session: - device = Device(ip_address="192.168.1.1", hostname="router") - session.add(device) - session.commit() - session.refresh(device) + device = Device(ip_address="192.168.1.1", hostname="router") + session.add(device) + session.commit() + session.refresh(device) - fetched = session.get(Device, device.id) - assert fetched.ip_address == "192.168.1.1" - assert fetched.hostname == "router" + fetched = session.get(Device, device.id) + assert fetched.ip_address == "192.168.1.1" # type: ignore[union-attr] + assert fetched.hostname == "router" # type: ignore[union-attr] @pytest.mark.integration -def test_port_crud_with_device(in_memory_engine): +def test_device_vendor_field(session): + """vendor column is persisted and retrieved correctly.""" + from app.models.device import Device + + device = Device( + ip_address="192.168.1.10", mac_address="aa:bb:cc:dd:ee:ff", vendor="Raspberry Pi" + ) + session.add(device) + session.commit() + + fetched = session.get(Device, device.id) + assert fetched.vendor == "Raspberry Pi" # type: ignore[union-attr] + + +@pytest.mark.integration +def test_port_crud_with_device(session): """Port linked to Device creates FK relationship correctly.""" from app.models.device import Device, Port - session_factory = sessionmaker(bind=in_memory_engine) - with session_factory() as session: - device = Device(ip_address="192.168.1.2") - session.add(device) - session.flush() + device = Device(ip_address="192.168.1.2") + session.add(device) + session.flush() - port = Port(device_id=device.id, port_number=22, protocol="tcp", service_name="ssh") - session.add(port) - session.commit() + port = Port(device_id=device.id, port_number=22, protocol="tcp", service_name="ssh") + session.add(port) + session.commit() - fetched_port = session.get(Port, port.id) - assert fetched_port.port_number == 22 - assert fetched_port.device_id == device.id + fetched_port = session.get(Port, port.id) + assert fetched_port.port_number == 22 # type: ignore[union-attr] + assert fetched_port.device_id == device.id # type: ignore[union-attr] @pytest.mark.integration -def test_device_cascade_deletes_ports(in_memory_engine): +def test_device_cascade_deletes_ports(session): """Deleting a Device cascades to its Ports.""" from app.models.device import Device, Port - session_factory = sessionmaker(bind=in_memory_engine) - with session_factory() as session: - device = Device(ip_address="192.168.1.3") - session.add(device) - session.flush() - port = Port(device_id=device.id, port_number=80) - session.add(port) + device = Device(ip_address="192.168.1.3") + session.add(device) + session.flush() + port = Port(device_id=device.id, port_number=80) + session.add(port) + session.commit() + + port_id = port.id + session.delete(device) + session.commit() + + assert session.get(Port, port_id) is None + + +# ── UniqueConstraint enforcement tests ─────────────────────────────────────── + + +@pytest.mark.integration +def test_device_ip_unique_constraint_enforced(session): + """Inserting two Devices with the same IP raises IntegrityError.""" + from app.models.device import Device + + session.add(Device(ip_address="10.0.0.1")) + session.commit() + session.add(Device(ip_address="10.0.0.1")) + with pytest.raises(IntegrityError): session.commit() - port_id = port.id - session.delete(device) + +@pytest.mark.integration +def test_port_unique_constraint_enforced(session): + """Inserting two Ports with same (device_id, port_number, protocol) raises IntegrityError.""" + from app.models.device import Device, Port + + device = Device(ip_address="10.0.0.2") + session.add(device) + session.flush() + session.add(Port(device_id=device.id, port_number=443, protocol="tcp")) + session.commit() + session.add(Port(device_id=device.id, port_number=443, protocol="tcp")) + with pytest.raises(IntegrityError): session.commit() - assert session.get(Port, port_id) is None + +# ── upsert_device tests ─────────────────────────────────────────────────────── + + +@pytest.mark.integration +def test_upsert_device_creates_new(session): + """upsert_device inserts a new row when ip_address is not present.""" + from app.db import upsert_device + + device = upsert_device( + session, + ip_address="172.16.0.1", + mac_address="de:ad:be:ef:00:01", + vendor="Cisco", + hostname="gw", + ) + session.commit() + + assert device.id is not None + assert device.ip_address == "172.16.0.1" + assert device.vendor == "Cisco" + + +@pytest.mark.integration +def test_upsert_device_updates_existing(session): + """upsert_device updates non-None fields on an existing row.""" + from app.db import upsert_device + from app.models.device import Device + + session.add(Device(ip_address="172.16.0.2", hostname="old-name")) + session.commit() + + upsert_device(session, ip_address="172.16.0.2", hostname="new-name", vendor="Ubiquiti") + session.commit() + + from sqlalchemy import select + + device = session.execute(select(Device).where(Device.ip_address == "172.16.0.2")).scalar_one() + assert device.hostname == "new-name" + assert device.vendor == "Ubiquiti" + + +@pytest.mark.integration +def test_upsert_device_none_does_not_overwrite(session): + """upsert_device leaves existing values intact when update fields are None.""" + from app.db import upsert_device + from app.models.device import Device + + session.add(Device(ip_address="172.16.0.3", hostname="keep-me", vendor="Intel")) + session.commit() + + # Call with hostname=None — should NOT overwrite the stored "keep-me" + upsert_device(session, ip_address="172.16.0.3", hostname=None, vendor=None) + session.commit() + + from sqlalchemy import select + + device = session.execute(select(Device).where(Device.ip_address == "172.16.0.3")).scalar_one() + assert device.hostname == "keep-me" + assert device.vendor == "Intel" + + +@pytest.mark.integration +def test_upsert_device_idempotent(session): + """Calling upsert_device twice with the same IP returns the same row.""" + from app.db import upsert_device + + d1 = upsert_device(session, ip_address="172.16.0.4") + session.commit() + d2 = upsert_device(session, ip_address="172.16.0.4") + session.commit() + + assert d1.id == d2.id + + +# ── upsert_port tests ───────────────────────────────────────────────────────── + + +@pytest.mark.integration +def test_upsert_port_creates_new(session): + """upsert_port inserts a new Port row.""" + from app.db import upsert_device, upsert_port + + device = upsert_device(session, ip_address="10.10.0.1") + session.flush() + + port = upsert_port(session, device_id=device.id, port_number=22, service_name="ssh") + session.commit() + + assert port.id is not None + assert port.service_name == "ssh" + + +@pytest.mark.integration +def test_upsert_port_updates_existing(session): + """upsert_port updates service_name / version_banner on an existing Port.""" + from app.db import upsert_device, upsert_port + + device = upsert_device(session, ip_address="10.10.0.2") + session.flush() + upsert_port(session, device_id=device.id, port_number=80, service_name="http") + session.commit() + + upsert_port( + session, + device_id=device.id, + port_number=80, + service_name="http", + version_banner="nginx/1.25", + ) + session.commit() + + from app.models.device import Port + from sqlalchemy import select + + port = session.execute( + select(Port).where(Port.device_id == device.id, Port.port_number == 80) + ).scalar_one() + assert port.version_banner == "nginx/1.25" + + +@pytest.mark.integration +def test_upsert_port_idempotent(session): + """Calling upsert_port twice with same key returns the same row.""" + from app.db import upsert_device, upsert_port + + device = upsert_device(session, ip_address="10.10.0.3") + session.flush() + + p1 = upsert_port(session, device_id=device.id, port_number=443, protocol="tcp") + session.commit() + p2 = upsert_port(session, device_id=device.id, port_number=443, protocol="tcp") + session.commit() + + assert p1.id == p2.id From 77b27e8f5d0cb52d59358a24bcf4b99fbab13840 Mon Sep 17 00:00:00 2001 From: wind Date: Fri, 27 Feb 2026 21:25:00 +0100 Subject: [PATCH 2/2] =?UTF-8?q?fix:=20ruff=20UP037=20+=20I001=20in=20db.py?= =?UTF-8?q?=20=E2=80=94=20unquote=20return=20annotations,=20sort=20inline?= =?UTF-8?q?=20imports?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/db.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/backend/app/db.py b/backend/app/db.py index ce4446f..d450e2d 100644 --- a/backend/app/db.py +++ b/backend/app/db.py @@ -49,7 +49,7 @@ def upsert_device( vendor: str | None = None, hostname: str | None = None, os_guess: str | None = None, -) -> "Device": +) -> Device: """Insert or update a Device row keyed on ip_address. If a Device with the given ip_address already exists, only non-None @@ -58,9 +58,10 @@ def upsert_device( Returns the Device instance (either existing or newly created). """ - from app.models.device import Device from sqlalchemy import select + from app.models.device import Device + stmt = select(Device).where(Device.ip_address == ip_address) device: Device | None = session.execute(stmt).scalar_one_or_none() @@ -94,15 +95,16 @@ def upsert_port( protocol: str = "tcp", service_name: str | None = None, version_banner: str | None = None, -) -> "Device": +) -> Device: """Insert or update a Port row keyed on (device_id, port_number, protocol). Non-None fields overwrite existing values. Caller must commit. Returns the Port instance. """ - from app.models.device import Port from sqlalchemy import select + from app.models.device import Port + stmt = select(Port).where( Port.device_id == device_id, Port.port_number == port_number,