diff --git a/app/api/audit_logs.py b/app/api/audit_logs.py index 3f6fa9a7..41a01bcd 100644 --- a/app/api/audit_logs.py +++ b/app/api/audit_logs.py @@ -20,16 +20,14 @@ logger = logging.getLogger(__name__) router = APIRouter() -# Module-level dependency singleton to satisfy Ruff B008 while maintaining default values for manual calls (e.g. in decorators). -_db_dep = Depends(get_db) -DbSession = Annotated[Session, _db_dep] +DbSession = Annotated[Session, Depends(get_db)] @router.get("/audit-logs") @require_login async def list_audit_logs( request: Request, - db: DbSession = _db_dep, + db: DbSession, action: Annotated[str | None, Query(description="Filter by action (exact match)")] = None, user: Annotated[str | None, Query(description="Filter by username")] = None, resource_type: Annotated[str | None, Query(description="Filter by resource type")] = None, @@ -75,7 +73,7 @@ async def list_audit_logs( @require_login async def list_distinct_actions( request: Request, - db: DbSession = _db_dep, + db: DbSession, ) -> list[str]: """Return the distinct action values present in the audit log.""" from app.models import AuditLog @@ -88,7 +86,7 @@ async def list_distinct_actions( @require_login async def list_distinct_users( request: Request, - db: DbSession = _db_dep, + db: DbSession, ) -> list[str]: """Return the distinct user values present in the audit log.""" from app.models import AuditLog diff --git a/app/api/billing.py b/app/api/billing.py index c28b6aa8..9c5582b0 100644 --- a/app/api/billing.py +++ b/app/api/billing.py @@ -168,9 +168,8 @@ async def create_checkout_session( checkout_session = client.checkout.sessions.create(params=session_params) logger.info( - "Created Stripe checkout session %s for user %s plan %s", + "Created Stripe checkout session %s for plan %s", checkout_session.id, - owner_id, body.plan_id, ) return {"checkout_url": checkout_session.url, "session_id": checkout_session.id} @@ -213,7 +212,7 @@ async def create_portal_session( } ) - logger.info("Created Stripe portal session for user %s", owner_id) + logger.info("Created Stripe portal session for user") return {"portal_url": portal.url} diff --git a/app/api/files.py b/app/api/files.py index 21e3c18a..64aa62b0 100644 --- a/app/api/files.py +++ b/app/api/files.py @@ -1497,7 +1497,7 @@ def claim_file(request: Request, file_id: int, db: DbSession): logger.exception(f"Error claiming file {file_id}: {e}") raise HTTPException(status_code=500, detail="Failed to claim document") - logger.info(f"File {file_id} claimed by user '{owner_id}'") + logger.info("File %d claimed by user", file_id) return {"status": "success", "message": "Document claimed successfully", "file_id": file_id, "owner_id": owner_id} @@ -1537,7 +1537,7 @@ def bulk_claim_files(request: Request, file_ids: list[int], db: DbSession): logger.exception(f"Error during bulk claim: {e}") raise HTTPException(status_code=500, detail="Failed to claim documents") - logger.info(f"Bulk claim by '{owner_id}': claimed={claimed}, skipped={[s['file_id'] for s in skipped]}") + logger.info("Bulk claim: claimed=%s, skipped=%s", claimed, [s["file_id"] for s in skipped]) return { "status": "success", "claimed_count": len(claimed), @@ -1590,8 +1590,7 @@ def assign_owner(request: Request, db: DbSession, owner_id: str = Query(...), fi logger.exception(f"Error assigning owner: {e}") raise HTTPException(status_code=500, detail="Failed to assign owner") - admin_name = get_current_owner_id(request) or "admin" - logger.info(f"Admin '{admin_name}' assigned owner_id='{owner_id}' to {updated} file(s)") + logger.info("Admin assigned owner to %d file(s)", updated) return { "status": "success", "message": f"Assigned owner to {updated} document(s)", diff --git a/app/utils/user_scope.py b/app/utils/user_scope.py index 05c5912b..a2d169c7 100644 --- a/app/utils/user_scope.py +++ b/app/utils/user_scope.py @@ -20,13 +20,31 @@ from app.models import FileRecord logger = logging.getLogger(__name__) +def _owner_id_from_user(user: dict) -> str | None: + """Extract the owner identifier from a user dict. + + Priority: ``sub`` (OAuth subject) → ``preferred_username`` → ``email`` → ``id``. + """ + return user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id") + + def get_current_owner_id(request: Request) -> str | None: """Extract the owner identifier for the current authenticated user. - The owner ID is derived from the user's session data. It uses the - ``sub`` claim (OAuth subject) when available, falling back to - ``preferred_username`` or ``email``. Returns ``None`` when no user - is authenticated. + The owner ID is derived from the user's session data or, when no session + is present, from a valid Bearer API token in the ``Authorization`` header. + This ensures that both browser-based (session cookie) and mobile/API + (Bearer token) requests are correctly identified. + + Priority for user resolution: + + 1. Session ``user`` dict (set by OAuth or local login). + 2. ``request.state.api_token_user`` (set by ``require_login`` or an + earlier call to this function during the same request). + 3. Direct Bearer token look-up against the database. + + Within the resolved user dict the owner ID is chosen as: + ``sub`` → ``preferred_username`` → ``email`` → ``id``. Args: request: The current FastAPI request with session data. @@ -34,11 +52,41 @@ def get_current_owner_id(request: Request) -> str | None: Returns: A stable string identifier for the user, or ``None``. """ + # 1. Session-based auth (most common for web UI) user = request.session.get("user") - if not user or not isinstance(user, dict): - return None - # Prefer 'sub' (OAuth subject), then 'preferred_username', then 'email', then 'id' - return user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id") + if user and isinstance(user, dict): + return _owner_id_from_user(user) + + # 2. Already-resolved API token user (cached by require_login or a + # prior dependency call during this request) + api_user = getattr(request.state, "api_token_user", None) + if isinstance(api_user, dict): + return _owner_id_from_user(api_user) + + # 3. Direct Bearer token resolution – necessary when this function is + # invoked as a FastAPI dependency (via Depends) which runs *before* + # the @require_login decorator wrapper has had a chance to resolve + # the token and populate request.state.api_token_user. + auth_header = request.headers.get("authorization", "") + if isinstance(auth_header, str) and auth_header.startswith("Bearer "): + try: + from app.auth import _resolve_bearer_user + from app.database import SessionLocal + + db = SessionLocal() + try: + resolved = _resolve_bearer_user(request, db) + finally: + db.close() + + if resolved: + # Cache so subsequent calls (and require_login) skip the DB + request.state.api_token_user = resolved + return _owner_id_from_user(resolved) + except Exception: + logger.debug("Bearer token resolution failed in get_current_owner_id", exc_info=True) + + return None def apply_owner_filter(query: Query, request: Request) -> Query: diff --git a/tests/test_coverage_polish.py b/tests/test_coverage_polish.py index 69857004..c5b7c4de 100644 --- a/tests/test_coverage_polish.py +++ b/tests/test_coverage_polish.py @@ -5,7 +5,7 @@ in files listed in the 90%+ coverage push issue. Each test class maps to a single source module. """ -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest from dropbox.exceptions import ApiError @@ -547,16 +547,23 @@ class TestURLUploadAdditionalCoverage: assert validate_file_type("", "noextfile") is False - @patch("app.api.url_upload.requests.get") + @patch("app.api.url_upload.httpx.AsyncClient.stream") @patch("app.api.url_upload.process_document") - def test_process_url_empty_url_path(self, mock_process, mock_get, client): + def test_process_url_empty_url_path(self, mock_process, mock_stream, client): """URL with empty path defaults to 'download' filename (line 197-202).""" - mock_response = Mock() + mock_response = AsyncMock() mock_response.status_code = 200 mock_response.headers = {"Content-Type": "application/pdf", "Content-Length": "50"} - mock_response.iter_content = Mock(return_value=[b"PDF"]) + + async def mock_aiter_bytes(chunk_size=None): + yield b"PDF" + + mock_response.aiter_bytes = mock_aiter_bytes mock_response.raise_for_status = Mock() - mock_get.return_value = mock_response + + mock_context = AsyncMock() + mock_context.__aenter__.return_value = mock_response + mock_stream.return_value = mock_context mock_task = Mock() mock_task.id = "task-empty-path" diff --git a/tests/test_endpoint_registration.py b/tests/test_endpoint_registration.py index 5d6719eb..d7cf6251 100644 --- a/tests/test_endpoint_registration.py +++ b/tests/test_endpoint_registration.py @@ -5,7 +5,7 @@ This test module serves as a regression prevention mechanism to ensure that endpoints remain accessible after code refactoring or reorganization. """ -from unittest.mock import Mock, patch +from unittest.mock import AsyncMock, Mock, patch import pytest @@ -17,17 +17,24 @@ TEST_URL = "https://example.com/test.pdf" class TestEndpointRegistration: """Verify that critical API endpoints are registered in the FastAPI app""" - @patch("app.api.url_upload.requests.get") + @patch("app.api.url_upload.httpx.AsyncClient.stream") @patch("app.api.url_upload.process_document") - def test_process_url_endpoint_exists(self, mock_process_document, mock_requests_get, client): + def test_process_url_endpoint_exists(self, mock_process_document, mock_stream, client): """Verify that /api/process-url endpoint is registered and accessible""" # Mock successful download to prevent actual HTTP requests - mock_response = Mock() + mock_response = AsyncMock() mock_response.status_code = 200 mock_response.headers = {"Content-Type": "application/pdf", "Content-Length": "1024"} - mock_response.iter_content = Mock(return_value=[b"PDF content"]) + + async def mock_aiter_bytes(chunk_size=None): + yield b"PDF content" + + mock_response.aiter_bytes = mock_aiter_bytes mock_response.raise_for_status = Mock() - mock_requests_get.return_value = mock_response + + mock_context = AsyncMock() + mock_context.__aenter__.return_value = mock_response + mock_stream.return_value = mock_context # Mock Celery task mock_task = Mock() @@ -46,17 +53,24 @@ class TestEndpointRegistration: "Verify that url_upload_router is included in app/api/__init__.py" ) - @patch("app.api.url_upload.requests.get") + @patch("app.api.url_upload.httpx.AsyncClient.stream") @patch("app.api.url_upload.process_document") - def test_process_url_endpoint_accepts_post(self, mock_process_document, mock_requests_get, client): + def test_process_url_endpoint_accepts_post(self, mock_process_document, mock_stream, client): """Verify that /api/process-url accepts POST requests""" # Mock successful download to prevent actual HTTP requests - mock_response = Mock() + mock_response = AsyncMock() mock_response.status_code = 200 mock_response.headers = {"Content-Type": "application/pdf", "Content-Length": "1024"} - mock_response.iter_content = Mock(return_value=[b"PDF content"]) + + async def mock_aiter_bytes(chunk_size=None): + yield b"PDF content" + + mock_response.aiter_bytes = mock_aiter_bytes mock_response.raise_for_status = Mock() - mock_requests_get.return_value = mock_response + + mock_context = AsyncMock() + mock_context.__aenter__.return_value = mock_response + mock_stream.return_value = mock_context # Mock Celery task mock_task = Mock() @@ -72,17 +86,24 @@ class TestEndpointRegistration: "Verify the endpoint is decorated with @router.post()" ) - @patch("app.api.url_upload.requests.get") + @patch("app.api.url_upload.httpx.AsyncClient.stream") @patch("app.api.url_upload.process_document") - def test_api_router_included_in_app(self, mock_process_document, mock_requests_get, client): + def test_api_router_included_in_app(self, mock_process_document, mock_stream, client): """Verify that the main API router is included in the FastAPI app""" # Mock successful download for /api/process-url test - mock_response = Mock() + mock_response = AsyncMock() mock_response.status_code = 200 mock_response.headers = {"Content-Type": "application/pdf", "Content-Length": "1024"} - mock_response.iter_content = Mock(return_value=[b"PDF content"]) + + async def mock_aiter_bytes(chunk_size=None): + yield b"PDF content" + + mock_response.aiter_bytes = mock_aiter_bytes mock_response.raise_for_status = Mock() - mock_requests_get.return_value = mock_response + + mock_context = AsyncMock() + mock_context.__aenter__.return_value = mock_response + mock_stream.return_value = mock_context # Mock Celery task mock_task = Mock() diff --git a/tests/test_multi_user.py b/tests/test_multi_user.py index 9f8fdab5..49456647 100644 --- a/tests/test_multi_user.py +++ b/tests/test_multi_user.py @@ -18,7 +18,7 @@ from sqlalchemy.pool import StaticPool from app.config import settings from app.database import Base -from app.models import FileRecord +from app.models import ApiToken, FileRecord # --------------------------------------------------------------------------- # Fixtures @@ -162,6 +162,122 @@ class TestGetCurrentOwnerId: request.session = {} assert get_current_owner_id(request) is None + @pytest.mark.unit + def test_resolves_from_api_token_user_state(self): + """get_current_owner_id should resolve from request.state.api_token_user.""" + from app.utils.user_scope import get_current_owner_id + + request = MagicMock() + request.session = {} + request.state.api_token_user = { + "id": "tok-owner", + "preferred_username": "tok-owner", + "email": "tok-owner", + } + assert get_current_owner_id(request) == "tok-owner" + + @pytest.mark.unit + def test_session_takes_precedence_over_api_token_user(self): + """Session auth should take precedence over api_token_user in state.""" + from app.utils.user_scope import get_current_owner_id + + request = MagicMock() + request.session = {"user": {"sub": "session-sub", "email": "session@example.com"}} + request.state.api_token_user = {"id": "tok-owner"} + assert get_current_owner_id(request) == "session-sub" + + @pytest.mark.unit + def test_resolves_bearer_token_directly(self, mu_engine, mu_session): + """get_current_owner_id should resolve a Bearer token when no session exists.""" + from types import SimpleNamespace + + from app.api.api_tokens import generate_api_token, hash_token + from app.utils.user_scope import get_current_owner_id + + # Create a token in the DB + plaintext = generate_api_token() + token_hash = hash_token(plaintext) + db_token = ApiToken( + owner_id="bearer-owner", + name="Test Bearer", + token_hash=token_hash, + token_prefix=plaintext[:12], + is_active=True, + ) + mu_session.add(db_token) + mu_session.commit() + + # Build a mock request with Bearer header but no session. + # SimpleNamespace starts with no attributes so getattr(..., None) works. + request = MagicMock() + request.session = {} + request.state = SimpleNamespace() + request.headers = {"authorization": f"Bearer {plaintext}"} + request.client.host = "127.0.0.1" + + # Provide the test session and make close() a no-op so the shared + # session is not torn down prematurely. + noop_close = MagicMock() + with patch("app.database.SessionLocal", return_value=mu_session), patch.object(mu_session, "close", noop_close): + result = get_current_owner_id(request) + + assert result == "bearer-owner" + # Verify the resolved user was cached in request.state + assert request.state.api_token_user["id"] == "bearer-owner" + + @pytest.mark.unit + def test_returns_none_for_invalid_bearer_token(self, mu_engine, mu_session): + """get_current_owner_id should return None for an invalid Bearer token.""" + from types import SimpleNamespace + + from app.utils.user_scope import get_current_owner_id + + request = MagicMock() + request.session = {} + request.state = SimpleNamespace() + request.headers = {"authorization": "Bearer de_invalid_token_value"} + request.client.host = "127.0.0.1" + + noop_close = MagicMock() + with patch("app.database.SessionLocal", return_value=mu_session), patch.object(mu_session, "close", noop_close): + result = get_current_owner_id(request) + + assert result is None + + +class TestOwnerIdFromUser: + """Tests for the _owner_id_from_user helper.""" + + @pytest.mark.unit + def test_prefers_sub(self): + from app.utils.user_scope import _owner_id_from_user + + assert _owner_id_from_user({"sub": "s", "preferred_username": "u", "email": "e"}) == "s" + + @pytest.mark.unit + def test_falls_back_to_preferred_username(self): + from app.utils.user_scope import _owner_id_from_user + + assert _owner_id_from_user({"preferred_username": "u", "email": "e"}) == "u" + + @pytest.mark.unit + def test_falls_back_to_email(self): + from app.utils.user_scope import _owner_id_from_user + + assert _owner_id_from_user({"email": "e"}) == "e" + + @pytest.mark.unit + def test_falls_back_to_id(self): + from app.utils.user_scope import _owner_id_from_user + + assert _owner_id_from_user({"id": "i"}) == "i" + + @pytest.mark.unit + def test_returns_none_for_empty_dict(self): + from app.utils.user_scope import _owner_id_from_user + + assert _owner_id_from_user({}) is None + class TestApplyOwnerFilter: """Tests for apply_owner_filter()."""