Investment Plans workspace
Open raw ↗
"""Pytest fixtures and configuration for HAD Digital MVP test suite.

Provides:
- Shared session-level and function-level server fixtures
- Isolated temporary SQLite databases with WAL mode
- Pre-authenticated role clients (patient, oncologist, nurse, admin)
- Direct SQLite query helper
"""

import json
import os
import shutil
import socket
import sqlite3
import subprocess
import sys
import tempfile
import time
from pathlib import Path
import pytest
import requests

PROJECT_ROOT = Path(__file__).resolve().parent.parent
MVP_DIR = PROJECT_ROOT / "MVP"


def allocate_free_port() -> int:
    """Find a free ephemeral port on localhost."""
    s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    s.bind(("127.0.0.1", 0))
    port = s.getsockname()[1]
    s.close()
    return port


class TestServerInstance:
    """Represents a running test server instance."""

    def __init__(self, port: int, temp_dir: str, db_path: str, proc: subprocess.Popen):
        self.port = port
        self.temp_dir = temp_dir
        self.db_path = db_path
        self.proc = proc
        self.base_url = f"http://127.0.0.1:{port}"

    def query_db(self, query: str, params: tuple = ()) -> list[dict]:
        """Execute a query directly against the test SQLite database."""
        conn = sqlite3.connect(self.db_path)
        conn.row_factory = sqlite3.Row
        cur = conn.cursor()
        cur.execute(query, params)
        rows = [dict(r) for r in cur.fetchall()]
        conn.close()
        return rows

    def execute_db(self, statement: str, params: tuple = ()):
        """Execute an insert/update statement directly against SQLite database."""
        conn = sqlite3.connect(self.db_path)
        cur = conn.cursor()
        cur.execute(statement, params)
        conn.commit()
        conn.close()


def launch_server_instance(timeout: int = 20) -> TestServerInstance:
    """Spawn an instance of MVP/app.py on an ephemeral port with an isolated DB."""
    port = allocate_free_port()
    temp_dir = tempfile.mkdtemp(prefix="had_pytest_")
    db_path = os.path.join(temp_dir, "test_had.db")

    env = os.environ.copy()
    env["HAD_DB_PATH"] = str(db_path)
    env["HAD_PORT"] = str(port)
    env["HAD_HOST"] = "127.0.0.1"
    env["HAD_DEBUG"] = "false"

    cmd = [sys.executable, str(MVP_DIR / "app.py"), "--port", str(port)]
    proc = subprocess.Popen(
        cmd, cwd=str(PROJECT_ROOT), env=env,
        stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True
    )

    # Health check wait
    start = time.time()
    online = False
    while time.time() - start < timeout:
        time.sleep(0.2)
        if proc.poll() is not None:
            stdout, stderr = proc.communicate()
            shutil.rmtree(temp_dir, ignore_errors=True)
            raise RuntimeError(f"Server crashed on startup: code {proc.returncode}\n{stderr}\n{stdout}")
        try:
            r = requests.get(f"http://127.0.0.1:{port}/api/whoami", timeout=1)
            if r.status_code in (200, 401):
                online = True
                break
        except Exception:
            pass

    if not online:
        proc.kill()
        proc.wait(2)
        shutil.rmtree(temp_dir, ignore_errors=True)
        raise TimeoutError(f"Server on port {port} failed to respond within {timeout}s")

    return TestServerInstance(port, temp_dir, db_path, proc)


def stop_server_instance(server: TestServerInstance):
    """Cleanly stop a test server instance and delete temp files."""
    if server.proc:
        try:
            server.proc.terminate()
            server.proc.wait(timeout=5)
        except Exception:
            try:
                server.proc.kill()
                server.proc.wait(timeout=2)
            except Exception:
                pass
    if os.path.exists(server.temp_dir):
        shutil.rmtree(server.temp_dir, ignore_errors=True)


@pytest.fixture(scope="session")
def session_server():
    """Session-level shared server for fast non-destructive test execution."""
    server = launch_server_instance(timeout=25)
    yield server
    stop_server_instance(server)


@pytest.fixture(scope="function")
def isolated_server():
    """Function-level isolated server for state-altering, lockout, or stress tests."""
    server = launch_server_instance(timeout=20)
    yield server
    stop_server_instance(server)


class HADClient:
    """Helper client for interacting with HAD Digital API endpoints."""

    def __init__(self, base_url: str):
        self.base_url = base_url
        self.session = requests.Session()
        self.current_user = None

    def login(self, username: str, password: str) -> requests.Response:
        """Authenticate user and retain session cookie."""
        resp = self.session.post(
            f"{self.base_url}/api/login",
            json={"username": username, "password": password}
        )
        if resp.status_code == 200:
            data = resp.json()
            self.current_user = data.get("user")
        return resp

    def logout(self) -> requests.Response:
        """Log out current user."""
        resp = self.session.post(f"{self.base_url}/api/logout")
        self.current_user = None
        return resp

    def get(self, path: str, **kwargs) -> requests.Response:
        url = f"{self.base_url}{path}" if path.startswith("/") else f"{self.base_url}/{path}"
        return self.session.get(url, **kwargs)

    def post(self, path: str, **kwargs) -> requests.Response:
        url = f"{self.base_url}{path}" if path.startswith("/") else f"{self.base_url}/{path}"
        return self.session.post(url, **kwargs)


@pytest.fixture
def api_client(session_server):
    """Unauthenticated client bound to session server."""
    return HADClient(session_server.base_url)


@pytest.fixture
def patient_client(session_server):
    """Client authenticated as patient.durand."""
    client = HADClient(session_server.base_url)
    resp = client.login("patient.durand", "demo123")
    assert resp.status_code == 200, f"Patient login failed: {resp.text}"
    return client


@pytest.fixture
def oncologist_client(session_server):
    """Client authenticated as dr.martin."""
    client = HADClient(session_server.base_url)
    resp = client.login("dr.martin", "demo123")
    assert resp.status_code == 200, f"Oncologist login failed: {resp.text}"
    return client


@pytest.fixture
def nurse_client(session_server):
    """Client authenticated as inf.moret."""
    client = HADClient(session_server.base_url)
    resp = client.login("inf.moret", "demo123")
    assert resp.status_code == 200, f"Nurse login failed: {resp.text}"
    return client


@pytest.fixture
def admin_client(session_server):
    """Client authenticated as admin."""
    client = HADClient(session_server.base_url)
    resp = client.login("admin", "admin123")
    assert resp.status_code == 200, f"Admin login failed: {resp.text}"
    return client