File size: 4,291 Bytes
54eb2ce | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 | from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError, DBAPIError
from fastapi import HTTPException
from typing import List
from ..database.db import session_pool
from ..errors import DatabaseConnectionError
from ..model.source import Source
from ..schema.source import SourceCreate, SourceResponse
from ..core.logger import SingletonLogger
logger = SingletonLogger().get_logger()
async def create_sources(sources: List[SourceCreate]) -> List[SourceResponse]:
"""Create multiple sources for a message in bulk."""
if not sources:
return []
try:
async with session_pool() as session:
source_objects = [
Source(
message_id=source.message_id,
source_text=source.source_text,
source_type=source.source_type,
source_url=source.source_url,
source_metadata=source.metadata,
)
for source in sources
]
session.add_all(source_objects)
await session.flush() # Flush to get IDs
# Get IDs of created sources to re-query them
source_ids = [source_obj.id for source_obj in source_objects]
await session.commit()
# Re-query sources from database to ensure all fields are properly loaded
result = await session.execute(
select(Source).where(Source.id.in_(source_ids))
)
persisted_sources = result.scalars().all()
return [
SourceResponse(
id=source.id,
message_id=source.message_id,
source_text=source.source_text,
source_type=source.source_type,
source_url=source.source_url,
metadata=source.source_metadata,
created_at=source.created_at,
updated_at=source.updated_at,
)
for source in persisted_sources
]
except DBAPIError as e:
logger.exception(f"Database connection error creating sources: {str(e)}")
raise DatabaseConnectionError(str(e))
except SQLAlchemyError as e:
logger.error(f"Database error creating sources: {str(e)}")
raise HTTPException(status_code=500, detail="Failed to create sources")
except Exception as e:
logger.error(f"Unexpected error creating sources: {str(e)}")
raise HTTPException(status_code=500, detail="Internal server error")
async def get_sources_by_message(message_id: int) -> List[SourceResponse]:
"""Retrieve all sources for a specific message."""
try:
async with session_pool() as session:
result = await session.execute(
select(Source)
.where(Source.message_id == message_id)
.order_by(Source.created_at)
)
sources = result.scalars().all()
return [
SourceResponse(
id=source.id,
message_id=source.message_id,
source_text=source.source_text,
source_type=source.source_type,
source_url=source.source_url,
metadata=source.source_metadata,
created_at=source.created_at,
updated_at=source.updated_at,
)
for source in sources
]
except DBAPIError as e:
logger.exception(
f"Database connection error retrieving sources for message_id={message_id}: {str(e)}"
)
raise DatabaseConnectionError(str(e))
except SQLAlchemyError as e:
logger.error(
f"Database error retrieving sources for message_id={message_id}: {str(e)}"
)
raise HTTPException(status_code=500, detail="Failed to retrieve sources")
except Exception as e:
logger.error(
f"Unexpected error retrieving sources for message_id={message_id}: {str(e)}"
)
raise HTTPException(status_code=500, detail="Internal server error")
|