honcho/src/utils/files.py

457 lines
15 KiB
Python

# ruff: noqa: I001
import asyncio
import datetime
import logging
import subprocess
import tempfile
from io import BytesIO
from pathlib import Path
from typing import Any, Protocol
from fastapi import UploadFile
from nanoid import generate as generate_nanoid
from sqlalchemy import Integer, select
from sqlalchemy.ext.asyncio import AsyncSession
from src import schemas
from src.config import settings
from src.exceptions import FileProcessingError, UnsupportedFileTypeError, ValidationException
from src.schemas import Message
from src.utils.clients import CLIENTS, transcribe_audio
logger = logging.getLogger(__name__)
SUPPORTED_AUDIO_CONTENT_TYPES = {
"audio/mpeg",
"audio/mp3",
"audio/wave",
"audio/wav",
"audio/x-wav",
}
SUPPORTED_AUDIO_EXTENSIONS = {".mp3", ".wav"}
AUDIO_EXTENSION_CONTENT_TYPES = {
".mp3": "audio/mpeg",
".wav": "audio/wav",
}
UPLOAD_VALIDATION_CHUNK_BYTES = 1024 * 1024
GENERIC_CONTENT_TYPES = {
"",
"application/octet-stream",
}
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)
class FileProcessor(Protocol):
async def extract_text(self, content: bytes) -> ExtractedFileText: ...
def supports_file_type(self, content_type: str) -> bool: ...
class PDFProcessor:
def supports_file_type(self, content_type: str) -> bool:
return content_type == "application/pdf"
async def extract_text(self, content: bytes) -> ExtractedFileText:
import pdfplumber
with pdfplumber.open(BytesIO(content)) as pdf_reader:
text_parts: list[str] = []
for page_num, page in enumerate(pdf_reader.pages):
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))
class TextProcessor:
def supports_file_type(self, content_type: str) -> bool:
return content_type.startswith("text/")
async def extract_text(self, content: bytes) -> ExtractedFileText:
# Try different encodings
for encoding in ["utf-8", "utf-16", "latin-1"]:
try:
return ExtractedFileText(text=content.decode(encoding))
except UnicodeDecodeError:
continue
raise ValueError("Could not decode text file")
class JSONProcessor:
def supports_file_type(self, content_type: str) -> bool:
return content_type == "application/json"
async def extract_text(self, content: bytes) -> ExtractedFileText:
import json
try:
decoded_content = content.decode("utf-8")
except UnicodeDecodeError as exc:
raise ValidationException("JSON uploads must be UTF-8 encoded") from exc
if not decoded_content.strip():
return ExtractedFileText(text="")
try:
data = json.loads(decoded_content)
except json.JSONDecodeError as exc:
raise ValidationException("Uploaded JSON is invalid") from exc
# Convert JSON to readable text format
return ExtractedFileText(text=json.dumps(data, ensure_ascii=False))
class AudioProcessor:
def supports_file_type(self, content_type: str) -> bool:
return content_type in SUPPORTED_AUDIO_CONTENT_TYPES
def supports_filename(self, filename: str | None) -> bool:
if not filename:
return False
return Path(filename).suffix.lower() in SUPPORTED_AUDIO_EXTENSIONS
def supports_content(self, *, filename: str | None, content_type: str) -> bool:
if self.supports_file_type(content_type):
return True
if content_type not in GENERIC_CONTENT_TYPES:
return False
return self.supports_filename(filename)
def supports_upload(self, file: UploadFile) -> bool:
return self.supports_content(
filename=file.filename,
content_type=file.content_type or "",
)
async def extract_text(
self,
content: bytes,
*,
filename: str | None,
content_type: str,
) -> ExtractedFileText:
if not filename:
raise ValidationException("Audio upload requires a filename")
if not content:
raise ValidationException("Audio upload is empty")
suffix = self.get_output_suffix(filename, content_type)
normalized_filename = self.ensure_audio_filename(filename, suffix)
normalized_content_type = self.normalize_content_type(filename, content_type)
await asyncio.to_thread(
self._probe_audio_duration_seconds,
content,
suffix,
)
text = await transcribe_audio(
content,
filename=normalized_filename,
content_type=normalized_content_type,
)
return ExtractedFileText(
text=text,
metadata={
"processing_type": "audio_transcription",
"audio_segment_count": 1,
"transcription_provider": settings.AUDIO.PROVIDER,
},
)
def normalize_content_type(self, filename: str, content_type: str) -> str:
if content_type in SUPPORTED_AUDIO_CONTENT_TYPES:
return content_type
extension = Path(filename).suffix.lower()
return AUDIO_EXTENSION_CONTENT_TYPES.get(extension, content_type)
def get_output_suffix(self, filename: str, content_type: str) -> str:
suffix = Path(filename).suffix.lower()
if suffix in SUPPORTED_AUDIO_EXTENSIONS:
return suffix
if content_type in {"audio/wave", "audio/wav", "audio/x-wav"}:
return ".wav"
return ".mp3"
def ensure_audio_filename(self, filename: str, suffix: str) -> str:
path = Path(filename)
if path.suffix.lower() == suffix:
return filename
return f"{filename}{suffix}"
def _probe_audio_duration_seconds(self, content: bytes, suffix: str) -> float:
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as temp_file:
temp_file.write(content)
temp_path = Path(temp_file.name)
try:
return self.probe_audio_duration_seconds_from_path(temp_path)
finally:
temp_path.unlink(missing_ok=True)
def probe_audio_duration_seconds_from_path(self, path: Path) -> float:
try:
command = [
"ffprobe",
"-v",
"error",
"-show_entries",
"format=duration",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(path),
]
result = subprocess.run(
command,
capture_output=True,
text=True,
check=True,
)
return max(float(result.stdout.strip()), 0.0)
except FileNotFoundError as exc:
raise ValidationException(
"Audio uploads require ffmpeg and ffprobe to be installed on the server"
) from exc
except (subprocess.CalledProcessError, ValueError) as exc:
raise ValidationException("Uploaded audio is invalid or unreadable") from exc
def is_audio_upload(file: UploadFile) -> bool:
return AudioProcessor().supports_upload(file)
def is_audio_transcription_enabled() -> bool:
return settings.AUDIO.PROVIDER == "openai" and "openai" in CLIENTS
async def is_validated_audio_upload(file: UploadFile) -> bool:
processor = AudioProcessor()
if not processor.supports_upload(file):
return False
filename = file.filename
if not filename:
return False
content_type = processor.normalize_content_type(filename, file.content_type or "")
suffix = processor.get_output_suffix(filename, content_type)
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as temp_file:
temp_path = Path(temp_file.name)
try:
while chunk := await file.read(UPLOAD_VALIDATION_CHUNK_BYTES):
temp_file.write(chunk)
finally:
await file.seek(0)
try:
await asyncio.to_thread(
processor.probe_audio_duration_seconds_from_path,
temp_path,
)
return True
except ValidationException as exc:
if str(exc) == "Uploaded audio is invalid or unreadable":
return False
raise
finally:
temp_path.unlink(missing_ok=True)
class FileProcessingService:
def __init__(self):
self.audio_processor: AudioProcessor = AudioProcessor()
self.processors: list[FileProcessor] = [
PDFProcessor(),
TextProcessor(),
JSONProcessor(),
# Add more processors as needed
]
async def extract_text_from_upload(self, file: UploadFile) -> ExtractedFileText:
"""Extract text from uploaded file without saving to disk."""
content = await file.read()
await file.seek(0)
normalized_content_type = file.content_type or ""
if self.audio_processor.supports_content(
filename=file.filename,
content_type=normalized_content_type,
):
if "openai" not in CLIENTS:
raise ValidationException(
"Audio uploads require OpenAI transcription credentials"
)
return await self.audio_processor.extract_text(
content,
filename=file.filename,
content_type=normalized_content_type,
)
processor = self._get_processor(normalized_content_type)
if not processor:
raise UnsupportedFileTypeError(
f"Unsupported file type: {file.content_type}. Supported types: {[p.__class__.__name__ for p in self.processors]}"
)
return await processor.extract_text(content)
def _get_processor(self, content_type: str) -> FileProcessor | None:
for processor in self.processors:
if processor.supports_file_type(content_type):
return processor
return None
def split_text_into_chunks(text: str, max_chars: int = 49500) -> list[str]:
"""Split text into chunks that fit within message limits."""
if len(text) <= max_chars:
return [text]
chunks: list[str] = []
current_pos = 0
while current_pos < len(text):
# Try to break at paragraph boundaries first
end_pos = current_pos + max_chars
if end_pos >= len(text):
chunks.append(text[current_pos:])
break
# Look for good break points (paragraph, sentence, word)
break_pos = end_pos
for delimiter in ["\n\n", "\n", ". ", " "]:
last_delimiter = text.rfind(delimiter, current_pos, end_pos)
if last_delimiter > current_pos:
break_pos = last_delimiter + len(delimiter)
break
chunks.append(text[current_pos:break_pos])
current_pos = break_pos
return chunks
async def get_file_messages(
db: AsyncSession,
workspace_name: str,
file_id: str,
session_name: str | None = None,
) -> list[Message]:
"""Get all messages for a specific document, ordered by chunk_index."""
from sqlalchemy import and_, func
from src.models import Message
query = select(Message).where(
and_(
Message.workspace_name == workspace_name,
func.jsonb_extract_path_text(Message.internal_metadata, "file_id")
== file_id,
)
)
if session_name:
query = query.where(Message.session_name == session_name)
# Order by chunk_index
query = query.order_by(
func.jsonb_extract_path_text(Message.internal_metadata, "chunk_index").cast(
Integer
)
)
result = await db.execute(query)
return list(result.scalars().all())
async def process_file_uploads_for_messages(
file: UploadFile,
peer_id: str,
max_chars: int = settings.MAX_MESSAGE_SIZE,
metadata: dict[str, Any] | None = None,
configuration: schemas.MessageConfiguration | None = None,
created_at: datetime.datetime | None = None,
) -> list[dict[str, Any]]:
"""
Process an uploaded file and prepare message creation data.
This function extracts text from a file, splits it into chunks, and prepares
the data needed to create messages.
Args:
file: Uploaded file to process
peer_id: ID of the peer creating the messages
max_chars: Maximum characters per message chunk
metadata: Optional metadata to associate with all messages created from this file
configuration: Optional configuration to associate with all messages created from this file
created_at: Optional created_at timestamp to use for all messages created from this file
Returns:
List of dictionaries containing message_create and file_metadata
Raises:
HTTPException: If file processing fails
"""
file_processor = FileProcessingService()
all_message_data: list[dict[str, Any]] = []
extracted_text = await file_processor.extract_text_from_upload(file)
# Split into chunks and create messages
chunks = split_text_into_chunks(extracted_text, max_chars=max_chars)
file_id = generate_nanoid()
for i, chunk in enumerate(chunks):
# Build message content properly handling empty files
message_content = chunk or ""
# Create message with optional metadata, configuration, and created_at
message_create = schemas.MessageCreate(
content=message_content,
peer_id=peer_id,
metadata=metadata,
configuration=configuration,
created_at=created_at,
)
# Store file metadata separately to add to internal_metadata later
file_metadata = {
"file_id": file_id,
"filename": file.filename,
"chunk_index": i,
"total_chunks": len(chunks),
"original_file_size": file.size,
"content_type": file.content_type,
"chunk_character_range": [
i * max_chars,
min((i + 1) * max_chars, len(extracted_text)),
],
}
file_metadata.update(extracted_text.metadata)
all_message_data.append(
{
"message_create": message_create,
"file_metadata": file_metadata,
}
)
if not all_message_data:
raise FileProcessingError()
return all_message_data