diff --git a/src/routers/messages.py b/src/routers/messages.py index 3acdc1e5..ef2e3068 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -18,7 +18,7 @@ from sqlalchemy.orm.attributes import flag_modified from src import crud, schemas from src.config import settings -from src.dependencies import db +from src.dependencies import db, tracked_db from src.deriver import enqueue from src.exceptions import FileTooLargeError, ResourceNotFoundException from src.security import require_auth @@ -141,7 +141,6 @@ async def create_messages_with_file( session_id: str = Path(...), form_data: schemas.MessageUploadCreate = Depends(parse_upload_form), file: UploadFile = File(...), - db: AsyncSession = db, ): """Create messages from uploaded files. Files are converted to text and split into multiple messages.""" @@ -160,22 +159,23 @@ async def create_messages_with_file( created_at=form_data.created_at, ) - # Create messages - message_creates = [item["message_create"] for item in all_message_data] - created_messages = await crud.create_messages( - db, - messages=message_creates, - workspace_name=workspace_id, - session_name=session_id, - ) + async with tracked_db("messages.upload") as db: + # Create messages + message_creates = [item["message_create"] for item in all_message_data] + created_messages = await crud.create_messages( + db, + messages=message_creates, + workspace_name=workspace_id, + session_name=session_id, + ) - # Update internal_metadata for file-related messages - for i, message in enumerate(created_messages): - file_metadata = all_message_data[i]["file_metadata"] - message.internal_metadata.update(file_metadata) - flag_modified(message, "internal_metadata") + # Update internal_metadata for file-related messages + for i, message in enumerate(created_messages): + file_metadata = all_message_data[i]["file_metadata"] + message.internal_metadata.update(file_metadata) + flag_modified(message, "internal_metadata") - await db.commit() + await db.commit() # Enqueue for processing (same as regular messages) payloads = [ diff --git a/tests/conftest.py b/tests/conftest.py index 3c9b8e63..d03a989b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -787,6 +787,7 @@ def mock_tracked_db(request: pytest.FixtureRequest): patch("src.deriver.queue_manager.tracked_db", mock_tracked_db_context), patch("src.deriver.consumer.tracked_db", mock_tracked_db_context), patch("src.deriver.enqueue.tracked_db", mock_tracked_db_context), + patch("src.routers.messages.tracked_db", mock_tracked_db_context), patch("src.routers.peers.tracked_db", mock_tracked_db_context), patch("src.crud.representation.tracked_db", mock_tracked_db_context), patch("src.dreamer.orchestrator.tracked_db", mock_tracked_db_context), diff --git a/tests/routes/test_files.py b/tests/routes/test_files.py index 3af9f60b..097a744c 100644 --- a/tests/routes/test_files.py +++ b/tests/routes/test_files.py @@ -1,6 +1,7 @@ # File upload tests for session endpoints import io import json +from contextlib import asynccontextmanager from typing import Any import pytest @@ -9,9 +10,12 @@ from nanoid import generate as generate_nanoid from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession -from src import models +from src import models, schemas from src.config import settings +from src.dependencies import get_db +from src.main import app from src.models import Peer, Workspace +from src.routers import messages as messages_router async def _create_test_session( @@ -161,6 +165,65 @@ async def test_create_messages_with_empty_json_file( assert data[0]["session_id"] == session_name +@pytest.mark.asyncio +async def test_file_upload_extracts_before_opening_db_session( + client: TestClient, + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """File extraction can perform remote OCR, so DB work must start after it.""" + test_workspace, test_peer = sample_data + test_session = await _create_test_session(db_session, test_workspace) + events: list[str] = [] + + async def fake_process_file_uploads_for_messages(*_args: Any, **_kwargs: Any): + events.append("extract") + return [ + { + "message_create": schemas.MessageCreate( + content="extracted text", + peer_id=test_peer.name, + ), + "file_metadata": { + "file_id": "test-file", + "filename": "test.txt", + "chunk_index": 0, + "total_chunks": 1, + "original_file_size": 13, + "content_type": "text/plain", + "chunk_character_range": [0, 13], + }, + } + ] + + async def override_get_db(): + events.append("db_open") + yield db_session + + @asynccontextmanager + async def fake_tracked_db(_: str | None = None): + events.append("db_open") + yield db_session + + monkeypatch.setattr( + messages_router, + "process_file_uploads_for_messages", + fake_process_file_uploads_for_messages, + ) + monkeypatch.setattr(messages_router, "tracked_db", fake_tracked_db, raising=False) + monkeypatch.setitem(app.dependency_overrides, get_db, override_get_db) + + response = client.post( + _get_upload_url(test_workspace.name, test_session.name), + files={"file": ("test.txt", io.BytesIO(b"test content"), "text/plain")}, + data={"peer_id": test_peer.name}, + ) + + assert response.status_code == 201 + assert events == ["extract", "db_open"] + + @pytest.mark.asyncio async def test_create_messages_with_unsupported_file_type( client: TestClient,