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")