Avoid DB session during upload extraction

This commit is contained in:
adavyas 2026-04-26 20:46:36 -07:00
parent 8399cc1ecf
commit 8342448a84
3 changed files with 81 additions and 17 deletions

View File

@ -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 = [

View File

@ -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),

View File

@ -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,