diff --git a/src/routers/messages.py b/src/routers/messages.py index 02838c48..82f8a1d9 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -151,15 +151,15 @@ async def create_messages_with_file( # Validate file size max_file_size = settings.MAX_FILE_SIZE - if ( - file.size - and file.size > settings.MAX_FILE_SIZE - and is_audio_transcription_enabled() - and is_audio_upload(file) - and file.size <= settings.AUDIO.MAX_FILE_SIZE_BYTES - and await is_validated_audio_upload(file) - ): - max_file_size = settings.AUDIO.MAX_FILE_SIZE_BYTES + if file.size and file.size > settings.MAX_FILE_SIZE: + is_valid_audio_upload = ( + is_audio_transcription_enabled() + and is_audio_upload(file) + and file.size <= settings.AUDIO.MAX_FILE_SIZE_BYTES + and await is_validated_audio_upload(file) + ) + if is_valid_audio_upload: + max_file_size = settings.AUDIO.MAX_FILE_SIZE_BYTES if file.size and file.size > max_file_size: raise FileTooLargeError( diff --git a/src/utils/files.py b/src/utils/files.py index b57a691c..2f66761f 100644 --- a/src/utils/files.py +++ b/src/utils/files.py @@ -1,6 +1,7 @@ # ruff: noqa: I001 import asyncio import datetime +from dataclasses import dataclass, field import logging import subprocess import tempfile @@ -40,23 +41,14 @@ GENERIC_CONTENT_TYPES = { } -class ExtractedFileText(str): - metadata: dict[str, Any] - - def __new__( - cls, text: str, metadata: dict[str, Any] | None = None - ) -> "ExtractedFileText": - obj = str.__new__(cls, text) - obj.metadata = metadata or {} - return obj - - @property - def text(self) -> str: - return str(self) +@dataclass +class FileExtractionResult: + text: str + metadata: dict[str, Any] = field(default_factory=dict) class FileProcessor(Protocol): - async def extract_text(self, content: bytes) -> ExtractedFileText: ... + async def extract_text(self, content: bytes) -> str: ... def supports_file_type(self, content_type: str) -> bool: ... @@ -64,7 +56,7 @@ class PDFProcessor: def supports_file_type(self, content_type: str) -> bool: return content_type == "application/pdf" - async def extract_text(self, content: bytes) -> ExtractedFileText: + async def extract_text(self, content: bytes) -> str: import pdfplumber with pdfplumber.open(BytesIO(content)) as pdf_reader: @@ -73,18 +65,18 @@ class PDFProcessor: text = page.extract_text() if text and text.strip(): text_parts.append(f"[Page {page_num + 1}]\n{text}") - return ExtractedFileText(text="\n\n".join(text_parts)) + return "\n\n".join(text_parts) class TextProcessor: def supports_file_type(self, content_type: str) -> bool: return content_type.startswith("text/") - async def extract_text(self, content: bytes) -> ExtractedFileText: + async def extract_text(self, content: bytes) -> str: # Try different encodings for encoding in ["utf-8", "utf-16", "latin-1"]: try: - return ExtractedFileText(text=content.decode(encoding)) + return content.decode(encoding) except UnicodeDecodeError: continue raise ValueError("Could not decode text file") @@ -94,7 +86,7 @@ class JSONProcessor: def supports_file_type(self, content_type: str) -> bool: return content_type == "application/json" - async def extract_text(self, content: bytes) -> ExtractedFileText: + async def extract_text(self, content: bytes) -> str: import json try: @@ -103,7 +95,7 @@ class JSONProcessor: raise ValidationException("JSON uploads must be UTF-8 encoded") from exc if not decoded_content.strip(): - return ExtractedFileText(text="") + return "" try: data = json.loads(decoded_content) @@ -111,7 +103,7 @@ class JSONProcessor: raise ValidationException("Uploaded JSON is invalid") from exc # Convert JSON to readable text format - return ExtractedFileText(text=json.dumps(data, ensure_ascii=False)) + return json.dumps(data, ensure_ascii=False) class AudioProcessor: @@ -142,7 +134,7 @@ class AudioProcessor: *, filename: str | None, content_type: str, - ) -> ExtractedFileText: + ) -> FileExtractionResult: if not filename: raise ValidationException("Audio upload requires a filename") if not content: @@ -161,7 +153,7 @@ class AudioProcessor: filename=normalized_filename, content_type=normalized_content_type, ) - return ExtractedFileText( + return FileExtractionResult( text=text, metadata={ "processing_type": "audio_transcription", @@ -279,7 +271,7 @@ class FileProcessingService: # Add more processors as needed ] - async def extract_text_from_upload(self, file: UploadFile) -> ExtractedFileText: + async def extract_text_from_upload(self, file: UploadFile) -> FileExtractionResult: """Extract text from uploaded file without saving to disk.""" content = await file.read() await file.seek(0) @@ -305,7 +297,7 @@ class FileProcessingService: f"Unsupported file type: {file.content_type}. Supported types: {[p.__class__.__name__ for p in self.processors]}" ) - return await processor.extract_text(content) + return FileExtractionResult(text=await processor.extract_text(content)) def _get_processor(self, content_type: str) -> FileProcessor | None: for processor in self.processors: @@ -409,7 +401,8 @@ async def process_file_uploads_for_messages( file_processor = FileProcessingService() all_message_data: list[dict[str, Any]] = [] - extracted_text = await file_processor.extract_text_from_upload(file) + extracted = await file_processor.extract_text_from_upload(file) + extracted_text = extracted.text # Split into chunks and create messages chunks = split_text_into_chunks(extracted_text, max_chars=max_chars) @@ -441,7 +434,7 @@ async def process_file_uploads_for_messages( min((i + 1) * max_chars, len(extracted_text)), ], } - file_metadata.update(extracted_text.metadata) + file_metadata.update(extracted.metadata) all_message_data.append( { diff --git a/tests/routes/test_files.py b/tests/routes/test_files.py index 01d7f933..fb44e222 100644 --- a/tests/routes/test_files.py +++ b/tests/routes/test_files.py @@ -18,7 +18,7 @@ from src import models, schemas from src.config import settings from src.models import Peer, Workspace from src.routers.messages import create_messages_with_file -from src.utils.files import ExtractedFileText +from src.utils.files import FileExtractionResult async def _create_test_session( @@ -637,7 +637,7 @@ async def test_create_messages_with_mp3_file( files = {"file": ("call.mp3", file_data, "audio/mpeg")} form_data = {"peer_id": test_peer.name} - extracted = ExtractedFileText( + extracted = FileExtractionResult( text="First sentence.\nSecond sentence.", metadata={ "processing_type": "audio_transcription", @@ -686,7 +686,7 @@ async def test_audio_upload_over_generic_limit_uses_audio_size_limit_when_valida original_audio_max = settings.AUDIO.MAX_FILE_SIZE_BYTES settings.MAX_FILE_SIZE = 5 settings.AUDIO.MAX_FILE_SIZE_BYTES = 10 - extracted = ExtractedFileText( + extracted = FileExtractionResult( text="Transcribed audio", metadata={ "processing_type": "audio_transcription", @@ -736,7 +736,7 @@ async def test_extension_only_audio_upload_over_generic_limit_uses_audio_size_li original_audio_max = settings.AUDIO.MAX_FILE_SIZE_BYTES settings.MAX_FILE_SIZE = 5 settings.AUDIO.MAX_FILE_SIZE_BYTES = 10 - extracted = ExtractedFileText( + extracted = FileExtractionResult( text="Transcribed extension-only audio", metadata={ "processing_type": "audio_transcription", @@ -906,7 +906,7 @@ async def test_create_messages_with_wav_file_accepts_audio_wave_mime( files = {"file": ("call.wav", file_data, "audio/wave")} form_data = {"peer_id": test_peer.name} - extracted = ExtractedFileText( + extracted = FileExtractionResult( text="WAV transcript", metadata={ "processing_type": "audio_transcription", @@ -991,7 +991,7 @@ async def test_large_audio_upload_applies_audio_metadata_to_all_message_chunks( form_data = {"peer_id": test_peer.name} long_text = "Segment line. " * 4000 - extracted = ExtractedFileText( + extracted = FileExtractionResult( text=long_text, metadata={ "processing_type": "audio_transcription",