Avoid DB session during upload extraction
This commit is contained in:
parent
8399cc1ecf
commit
8342448a84
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Reference in New Issue