diff --git a/app/api/common.py b/app/api/common.py index 07e3283e..8cbb0535 100644 --- a/app/api/common.py +++ b/app/api/common.py @@ -9,21 +9,11 @@ from pathlib import Path from fastapi import HTTPException, status from app.config import settings -from app.database import SessionLocal # Set up logging logger = logging.getLogger(__name__) -def get_db(): - """Database dependency injection for routes""" - db = SessionLocal() - try: - yield db - finally: - db.close() - - def resolve_file_path(file_path: str, subfolder: str = None) -> str: """ Resolves a file path to an absolute path with path traversal protection. diff --git a/app/api/files.py b/app/api/files.py index 47d2b714..58e0d834 100644 --- a/app/api/files.py +++ b/app/api/files.py @@ -12,9 +12,9 @@ from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, Upl from sqlalchemy import asc, desc, or_ from sqlalchemy.orm import Session -from app.api.common import get_db from app.auth import require_login from app.config import settings +from app.database import get_db from app.models import FileRecord, ProcessingLog from app.tasks.convert_to_pdf import convert_to_pdf from app.tasks.process_document import process_document diff --git a/app/api/logs.py b/app/api/logs.py index 6b2c4cdb..7ef0c434 100644 --- a/app/api/logs.py +++ b/app/api/logs.py @@ -9,8 +9,8 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request from sqlalchemy import desc from sqlalchemy.orm import Session -from app.api.common import get_db from app.auth import require_login +from app.database import get_db from app.models import FileRecord, ProcessingLog # Set up logging diff --git a/app/views/base.py b/app/views/base.py index 087e88b4..45ffa7b9 100644 --- a/app/views/base.py +++ b/app/views/base.py @@ -11,7 +11,7 @@ from sqlalchemy.orm import Session # noqa: F401 from app.auth import require_login # noqa: F401 from app.config import settings -from app.database import SessionLocal +from app.database import get_db # noqa: F401 # Set up Jinja2 templates templates_dir = Path(__file__).parent.parent.parent / "frontend" / "templates" @@ -39,14 +39,3 @@ templates.TemplateResponse = template_response_with_version # Set up logging logger = logging.getLogger(__name__) - - -def get_db(): - """ - Dependency to get a database session. - """ - db = SessionLocal() - try: - yield db - finally: - db.close() diff --git a/tests/conftest.py b/tests/conftest.py index 9060c9de..dfe5fc50 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,10 +1,12 @@ """ Pytest configuration and shared fixtures for DocuElevate tests. """ + import os import tempfile -import pytest from typing import Generator + +import pytest from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker @@ -22,10 +24,11 @@ os.environ["WORKDIR"] = "/tmp" os.environ["AUTH_ENABLED"] = "False" os.environ["SESSION_SECRET"] = "test_secret_key_for_testing_must_be_at_least_32_characters_long" -from app.database import Base, get_db -from app.main import app as fastapi_app +from app.database import Base # noqa: E402 +from app.main import app as fastapi_app # noqa: E402 + # Import models to register them with SQLAlchemy Base -from app.models import DocumentMetadata, FileRecord, ProcessingLog +from app.models import DocumentMetadata, FileRecord, ProcessingLog # noqa: F401, E402 @pytest.fixture(scope="session") @@ -44,14 +47,14 @@ def db_session(): connect_args={"check_same_thread": False}, poolclass=StaticPool, ) - + # Create all tables Base.metadata.create_all(bind=engine) - + # Create a session TestingSessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) session = TestingSessionLocal() - + try: yield session finally: @@ -62,27 +65,24 @@ def db_session(): @pytest.fixture(scope="function") def client(db_session) -> TestClient: """Create a test client with a fresh database.""" - - # Import all the different get_db functions used across the app - from app.api.common import get_db as api_get_db - from app.views.base import get_db as views_get_db - + + # Import the canonical get_db function + from app.database import get_db + # Override the get_db dependency to use our test database def override_get_db(): try: yield db_session finally: pass - - # Override all variants of get_db + + # Override the single canonical get_db dependency fastapi_app.dependency_overrides[get_db] = override_get_db - fastapi_app.dependency_overrides[api_get_db] = override_get_db - fastapi_app.dependency_overrides[views_get_db] = override_get_db - + # Use base_url to satisfy TrustedHostMiddleware with TestClient(fastapi_app, base_url="http://localhost") as test_client: yield test_client - + # Clean up fastapi_app.dependency_overrides.clear() @@ -91,7 +91,7 @@ def client(db_session) -> TestClient: def sample_pdf_path(test_workdir) -> str: """Create a sample PDF file for testing.""" pdf_path = os.path.join(test_workdir, "test.pdf") - + # Create a minimal valid PDF pdf_content = b"""%PDF-1.4 1 0 obj @@ -129,10 +129,10 @@ startxref 197 %%EOF """ - - with open(pdf_path, 'wb') as f: + + with open(pdf_path, "wb") as f: f.write(pdf_content) - + return pdf_path @@ -140,10 +140,10 @@ startxref def sample_text_file(test_workdir) -> str: """Create a sample text file for testing.""" text_path = os.path.join(test_workdir, "test.txt") - - with open(text_path, 'w') as f: + + with open(text_path, "w") as f: f.write("This is a test document.\nWith multiple lines.\n") - + return text_path @@ -151,46 +151,29 @@ def sample_text_file(test_workdir) -> str: def mock_openai_response(): """Mock OpenAI API response for testing.""" return { - "choices": [{ - "message": { - "content": '{"document_type": "invoice", "summary": "Test invoice", "tags": ["test", "invoice"]}' + "choices": [ + { + "message": { + "content": '{"document_type": "invoice", "summary": "Test invoice", "tags": ["test", "invoice"]}' + } } - }] + ] } @pytest.fixture def mock_azure_response(): """Mock Azure Document Intelligence API response for testing.""" - return { - "analyzeResult": { - "content": "Test document content extracted by OCR", - "pages": [{"pageNumber": 1}] - } - } + return {"analyzeResult": {"content": "Test document content extracted by OCR", "pages": [{"pageNumber": 1}]}} # Markers for categorizing tests def pytest_configure(config): """Configure custom pytest markers.""" - config.addinivalue_line( - "markers", "unit: Unit tests for individual functions/methods" - ) - config.addinivalue_line( - "markers", "integration: Integration tests for API endpoints and workflows" - ) - config.addinivalue_line( - "markers", "slow: Tests that take significant time to run" - ) - config.addinivalue_line( - "markers", "security: Security-related tests" - ) - config.addinivalue_line( - "markers", "requires_external: Tests requiring external services" - ) - config.addinivalue_line( - "markers", "requires_db: Tests requiring database" - ) - config.addinivalue_line( - "markers", "requires_redis: Tests requiring Redis" - ) + config.addinivalue_line("markers", "unit: Unit tests for individual functions/methods") + config.addinivalue_line("markers", "integration: Integration tests for API endpoints and workflows") + config.addinivalue_line("markers", "slow: Tests that take significant time to run") + config.addinivalue_line("markers", "security: Security-related tests") + config.addinivalue_line("markers", "requires_external: Tests requiring external services") + config.addinivalue_line("markers", "requires_db: Tests requiring database") + config.addinivalue_line("markers", "requires_redis: Tests requiring Redis")