File size: 4,016 Bytes
18ade12
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
109
110
111
112
113
114
115
116
117
118
119
"""
Database error handling utilities.

This module provides error handling for common database exceptions:
- OperationalError: Connection failures
- IntegrityError: Constraint violations
- TimeoutError: Query timeouts

Per contracts/database-interface.yaml error handling specification.
"""

import asyncio
import logging
from functools import wraps
from typing import Any, Callable, TypeVar

from sqlalchemy.exc import IntegrityError, OperationalError, SQLAlchemyError

logger = logging.getLogger(__name__)

T = TypeVar("T")


class DatabaseConnectionError(Exception):
    """Raised when database connection fails."""

    def __init__(self, message: str = "Unable to connect to database"):
        self.message = message
        super().__init__(self.message)


class DatabaseIntegrityError(Exception):
    """Raised when a database constraint is violated."""

    def __init__(self, message: str = "Database constraint violation"):
        self.message = message
        super().__init__(self.message)


class DatabaseTimeoutError(Exception):
    """Raised when a database operation times out."""

    def __init__(self, message: str = "Database operation timed out"):
        self.message = message
        super().__init__(self.message)


def handle_db_errors(func: Callable[..., T]) -> Callable[..., T]:
    """
    Decorator to handle database errors with appropriate logging.

    Catches SQLAlchemy exceptions and converts them to application-specific
    exceptions with sanitized error messages (no credential exposure).

    Usage:
        @handle_db_errors
        async def get_user(session: AsyncSession, user_id: str):
            ...
    """

    @wraps(func)
    async def wrapper(*args: Any, **kwargs: Any) -> T:
        try:
            return await func(*args, **kwargs)
        except OperationalError as e:
            logger.error(
                "Database connection error: Unable to connect to database. "
                "Please verify your DATABASE_URL configuration."
            )
            raise DatabaseConnectionError() from e
        except IntegrityError as e:
            # Extract constraint name if available, but not the full error
            constraint_info = "constraint violation"
            if hasattr(e, "orig") and e.orig:
                # Try to get constraint name without exposing table/column details
                orig_str = str(e.orig)
                if "unique" in orig_str.lower():
                    constraint_info = "unique constraint violation"
                elif "foreign key" in orig_str.lower():
                    constraint_info = "foreign key constraint violation"
                elif "not null" in orig_str.lower():
                    constraint_info = "not null constraint violation"
                elif "check" in orig_str.lower():
                    constraint_info = "check constraint violation"

            logger.error(f"Database integrity error: {constraint_info}")
            raise DatabaseIntegrityError(constraint_info) from e
        except asyncio.TimeoutError as e:
            logger.error("Database operation timed out")
            raise DatabaseTimeoutError() from e
        except SQLAlchemyError as e:
            logger.error(f"Database error: {type(e).__name__}")
            raise

    return wrapper


async def execute_with_timeout(coro: Any, timeout_seconds: float = 5.0) -> Any:
    """
    Execute an async database operation with a timeout.

    Args:
        coro: The coroutine to execute
        timeout_seconds: Maximum time to wait (default 5 seconds)

    Returns:
        The result of the coroutine

    Raises:
        DatabaseTimeoutError: If the operation times out
    """
    try:
        return await asyncio.wait_for(coro, timeout=timeout_seconds)
    except asyncio.TimeoutError:
        logger.error(f"Database operation timed out after {timeout_seconds} seconds")
        raise DatabaseTimeoutError(
            f"Database operation timed out after {timeout_seconds} seconds"
        )