diff --git a/src/sentry/utils/db.py b/src/sentry/utils/db.py index d879bb6f27b0..d5616b57c966 100644 --- a/src/sentry/utils/db.py +++ b/src/sentry/utils/db.py @@ -1,16 +1,39 @@ import logging -from collections.abc import Generator, Sequence -from contextlib import ExitStack, contextmanager +import math +from collections.abc import Callable, Generator, Sequence +from contextlib import AbstractContextManager, ExitStack, contextmanager from datetime import timedelta from functools import wraps +from time import monotonic from django.db import DEFAULT_DB_ALIAS, connections, router, transaction from django.db.utils import OperationalError, ProgrammingError from sentry_sdk.integrations import Integration +from sentry.utils.env import in_test_environment from sentry.utils.tracing import set_span_data, start_span +class StatementTimeoutBudgetExceeded(OperationalError): + pass + + +def _has_open_transaction(alias: str) -> bool: + if in_test_environment(): + from sentry.testutils.hybrid_cloud import ( # NOQA:S007 + simulated_transaction_watermarks, + ) + + return ( + simulated_transaction_watermarks.connection_transaction_depth_above_watermark( + using=alias + ) + > 0 + ) + + return transaction.get_connection(alias).in_atomic_block + + @contextmanager def statement_timeout(alias: str, timeout: timedelta) -> Generator[None]: """ @@ -30,6 +53,38 @@ def statement_timeout(alias: str, timeout: timedelta) -> Generator[None]: yield +def make_statement_timeout_budget( + timeout: timedelta, *, using: str = DEFAULT_DB_ALIAS +) -> Callable[[], AbstractContextManager[None]]: + """Create a reusable context manager source bounded by one shared deadline. + + Each use opens a short transaction and applies the time remaining as a + server-side statement timeout. Keep one database statement in each context; + multiple statements would each receive the same remaining timeout. + + This cannot be used within an existing transaction because ``SET LOCAL`` + would remain active until that outer transaction ended. Catch timeout errors + outside the returned context so its transaction can roll back first. + """ + deadline = monotonic() + timeout.total_seconds() + + @contextmanager + def budgeted_statement() -> Generator[None]: + if _has_open_transaction(using): + raise RuntimeError("statement timeout budget cannot be used inside an open transaction") + + remaining_ms = math.floor((deadline - monotonic()) * 1000) + if remaining_ms < 1: + raise StatementTimeoutBudgetExceeded("statement timeout budget exhausted") + + with transaction.atomic(using=using): + with connections[using].cursor() as cursor: + cursor.execute("SET LOCAL statement_timeout = %s", [remaining_ms]) + yield + + return budgeted_statement + + def handle_db_failure(func, model, wrap_in_transaction: bool = True): @wraps(func) def wrapped(*args, **kwargs): diff --git a/tests/sentry/utils/test_db.py b/tests/sentry/utils/test_db.py new file mode 100644 index 000000000000..adeec672147d --- /dev/null +++ b/tests/sentry/utils/test_db.py @@ -0,0 +1,97 @@ +from datetime import timedelta +from unittest.mock import patch + +import pytest +from django.db import DEFAULT_DB_ALIAS, connections, transaction +from django.db.utils import OperationalError +from django.test.utils import CaptureQueriesContext + +from sentry.testutils.cases import TestCase +from sentry.utils.db import ( + StatementTimeoutBudgetExceeded, + make_statement_timeout_budget, +) + + +class MakeStatementTimeoutBudgetTest(TestCase): + def test_uses_remaining_budget_for_each_statement(self) -> None: + connection = connections[DEFAULT_DB_ALIAS] + + with patch("sentry.utils.db.monotonic", side_effect=[100, 101, 104]): + budgeted_statement = make_statement_timeout_budget(timedelta(seconds=10)) + with CaptureQueriesContext(connection) as queries: + with budgeted_statement(), connection.cursor() as cursor: + cursor.execute("SELECT 1") + with budgeted_statement(), connection.cursor() as cursor: + cursor.execute("SELECT 1") + + timeout_queries = [query["sql"] for query in queries if "statement_timeout" in query["sql"]] + assert timeout_queries == [ + "SET LOCAL statement_timeout = 9000", + "SET LOCAL statement_timeout = 6000", + ] + + def test_deadline_starts_when_budget_is_created(self) -> None: + connection = connections[DEFAULT_DB_ALIAS] + + with patch("sentry.utils.db.monotonic", side_effect=[100, 103]): + budgeted_statement = make_statement_timeout_budget(timedelta(seconds=10)) + with CaptureQueriesContext(connection) as queries: + with budgeted_statement(), connection.cursor() as cursor: + cursor.execute("SELECT 1") + + assert any("statement_timeout = 7000" in query["sql"] for query in queries) + + def test_raises_before_query_when_budget_is_exhausted(self) -> None: + connection = connections[DEFAULT_DB_ALIAS] + + with patch("sentry.utils.db.monotonic", side_effect=[100, 110]): + budgeted_statement = make_statement_timeout_budget(timedelta(seconds=10)) + with CaptureQueriesContext(connection) as queries: + with pytest.raises(StatementTimeoutBudgetExceeded): + with budgeted_statement(): + raise AssertionError("unreachable") + + assert list(queries) == [] + + def test_raises_instead_of_disabling_sub_millisecond_timeout(self) -> None: + with patch("sentry.utils.db.monotonic", side_effect=[100, 100.0005]): + budgeted_statement = make_statement_timeout_budget(timedelta(milliseconds=1)) + + with pytest.raises(StatementTimeoutBudgetExceeded): + with budgeted_statement(): + raise AssertionError("unreachable") + + def test_budget_exhaustion_is_an_operational_error(self) -> None: + assert issubclass(StatementTimeoutBudgetExceeded, OperationalError) + + def test_refuses_to_run_inside_an_open_transaction(self) -> None: + with patch("sentry.utils.db.monotonic", return_value=100): + budgeted_statement = make_statement_timeout_budget(timedelta(seconds=10)) + + with ( + transaction.atomic(using=DEFAULT_DB_ALIAS), + pytest.raises(RuntimeError, match="cannot be used inside an open transaction"), + ): + with budgeted_statement(): + raise AssertionError("unreachable") + + def test_server_side_timeout_leaves_connection_usable(self) -> None: + connection = connections[DEFAULT_DB_ALIAS] + budgeted_statement = make_statement_timeout_budget(timedelta(milliseconds=20)) + + with pytest.raises(OperationalError): + with budgeted_statement(), connection.cursor() as cursor: + cursor.execute("SELECT pg_sleep(0.1)") + + with connection.cursor() as cursor: + cursor.execute("SELECT 1") + assert cursor.fetchone() == (1,) + + def test_propagates_unrelated_exceptions(self) -> None: + with patch("sentry.utils.db.monotonic", return_value=100): + budgeted_statement = make_statement_timeout_budget(timedelta(seconds=10)) + + with pytest.raises(ValueError, match="query setup failed"): + with budgeted_statement(): + raise ValueError("query setup failed")