diff --git a/datajunction-server/datajunction_server/api/sql.py b/datajunction-server/datajunction_server/api/sql.py index d39607b69..c2216df2a 100644 --- a/datajunction-server/datajunction_server/api/sql.py +++ b/datajunction-server/datajunction_server/api/sql.py @@ -23,6 +23,7 @@ from datajunction_server.database.user import User from datajunction_server.errors import DJInvalidInputException from datajunction_server.instrumentation.provider import get_metrics_provider +from datajunction_server.sql.parsing.backends.antlr4 import report_parse_cache_stats from datajunction_server.internal.access.authentication.http import SecureAPIRouter from datajunction_server.internal.caching.cachelib_cache import get_cache from datajunction_server.internal.caching.interface import Cache @@ -291,6 +292,7 @@ async def get_measures_sql_v3( _tags, ) get_metrics_provider().counter("dj.sql.requests", tags=_tags) + report_parse_cache_stats() if result.warnings: get_metrics_provider().counter("dj.sql.build_warnings", tags=_tags) @@ -529,6 +531,7 @@ async def get_combined_measures_sql_v3( _tags, ) get_metrics_provider().counter("dj.sql.requests", tags=_tags) + report_parse_cache_stats() if combined_result.warnings: get_metrics_provider().counter("dj.sql.build_warnings", tags=_tags) diff --git a/datajunction-server/datajunction_server/config.py b/datajunction-server/datajunction_server/config.py index 808aa552e..39941bca4 100644 --- a/datajunction-server/datajunction_server/config.py +++ b/datajunction-server/datajunction_server/config.py @@ -285,6 +285,14 @@ def validate_restrictive_scopes(cls, values: list[str]) -> list[str]: # Caps how many SQL rebuilds run simultaneously to avoid DB connection spikes. query_cache_max_concurrent_refreshes: int = 3 + # How many parsed node definitions to keep. Each entry retains its parser + # and token stream, per worker process. + definition_parse_cache_size: int = 512 + + # How many parsed filter and orderby clauses to keep. Sized for reuse + # within one request, not across them. + request_parse_cache_size: int = 64 + # Maximum amount of nodes to return for requests to list all nodes node_list_max: int = 10000 diff --git a/datajunction-server/datajunction_server/construction/build_v3/builder.py b/datajunction-server/datajunction_server/construction/build_v3/builder.py index 403641d3c..b1b1d7080 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/builder.py +++ b/datajunction-server/datajunction_server/construction/build_v3/builder.py @@ -182,7 +182,7 @@ def apply_orderby_limit( # Parse the orderby expressions orderby_str = ",".join(orderby) - parsed = parse(f"SELECT 1 ORDER BY {orderby_str}") + parsed = parse(f"SELECT 1 ORDER BY {orderby_str}", from_request=True) sort_items = ( parsed.select.organization.order if parsed.select.organization else [] ) diff --git a/datajunction-server/datajunction_server/construction/build_v3/filters.py b/datajunction-server/datajunction_server/construction/build_v3/filters.py index 2f092d425..74b882d96 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/filters.py +++ b/datajunction-server/datajunction_server/construction/build_v3/filters.py @@ -44,7 +44,7 @@ def parse_filter(filter_str: str) -> ast.Expression: # Returns ast.BinaryOp with comparison """ # Parse as "SELECT 1 WHERE " and extract the WHERE clause - query = parse(f"SELECT 1 WHERE {filter_str}") + query = parse(f"SELECT 1 WHERE {filter_str}", from_request=True) if query.select.where is None: # pragma: no cover raise DJInvalidInputException(f"Failed to parse filter: {filter_str}") return query.select.where diff --git a/datajunction-server/datajunction_server/internal/sql.py b/datajunction-server/datajunction_server/internal/sql.py index fd20ec0ca..d85240a63 100644 --- a/datajunction-server/datajunction_server/internal/sql.py +++ b/datajunction-server/datajunction_server/internal/sql.py @@ -37,6 +37,7 @@ from datajunction_server.database.node import Node, NodeRevision from datajunction_server.errors import DJException, DJInvalidInputException from datajunction_server.instrumentation.provider import get_metrics_provider +from datajunction_server.sql.parsing.backends.antlr4 import report_parse_cache_stats from datajunction_server.internal.access.authorization import ( AccessChecker, AccessDenialMode, @@ -281,6 +282,7 @@ async def generate_metrics_sql( provider = get_metrics_provider() provider.timer("dj.sql.build_latency_ms", elapsed_ms, _tags) provider.counter("dj.sql.requests", tags=_tags) + report_parse_cache_stats() if result.warnings: provider.counter("dj.sql.build_warnings", tags=_tags) logger.info( diff --git a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py index 5038d7e20..56212502b 100644 --- a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py +++ b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py @@ -200,15 +200,99 @@ def tree_to_strings(tree, indent=0): return result -def parse_rule(sql: str, rule: str) -> Union[ast.Node, "ColumnType"]: +@lru_cache(maxsize=1) +def _definition_tree_parser(): + """ + Build the node definition cache, sized from settings on first use. + """ + from datajunction_server.utils import get_settings # noqa: PLC0415 + + @lru_cache(maxsize=get_settings().definition_parse_cache_size) + def _parse(sql: str, rule: str): + return parse_sql_with_sll_fallback(sql, rule) + + return _parse + + +@lru_cache(maxsize=1) +def _request_tree_parser(): + """ + Build the request SQL cache, sized from settings on first use. + """ + from datajunction_server.utils import get_settings # noqa: PLC0415 + + @lru_cache(maxsize=get_settings().request_parse_cache_size) + def _parse(sql: str, rule: str): + return parse_sql_with_sll_fallback(sql, rule) + + return _parse + + +def cached_definition_tree(sql: str, rule: str): + """ + Parse a node definition into an ANTLR tree, caching the result. + + Node SQL repeats across requests, so this keeps the parse. The tree is + safe to share: ``visit`` only reads it, and the mutable DJ AST is rebuilt + on every call. + """ + return _definition_tree_parser()(sql, rule) + + +def cached_request_tree(sql: str, rule: str): + """ + Parse request-supplied SQL into an ANTLR tree, caching the result. + + One request parses the same filter many times over, so this removes work + the build would otherwise repeat. The tree is safe to share: ``visit`` + only reads it, and the mutable DJ AST is rebuilt on every call. + """ + return _request_tree_parser()(sql, rule) + + +def report_parse_cache_stats() -> None: + """ + Report parse cache sizes and hit counts, per cache. + + Misses rising while ``size`` sits at ``max_size`` means the cache is too + small and entries are being evicted. + """ + from datajunction_server.instrumentation.provider import ( # noqa: PLC0415 + get_metrics_provider, + ) + + provider = get_metrics_provider() + caches = { + "definitions": _definition_tree_parser(), + "requests": _request_tree_parser(), + } + for name, cache in caches.items(): + info = cache.cache_info() + tags = {"cache": name} + provider.gauge("dj.sql.parse_cache.hits", info.hits, tags) + provider.gauge("dj.sql.parse_cache.misses", info.misses, tags) + provider.gauge("dj.sql.parse_cache.size", info.currsize, tags) + provider.gauge("dj.sql.parse_cache.max_size", info.maxsize, tags) + + +def parse_rule( + sql: str, + rule: str, + from_request: bool = False, +) -> Union[ast.Node, "ColumnType"]: """ Parse a string into a DJ ast using the ANTLR4 backend. - Uses SLL mode first (faster), falls back to LL mode if needed. + Uses SLL mode first (faster), falls back to LL mode if needed. Set + ``from_request`` for filters and orderby clauses, which one request + parses repeatedly. """ - antlr_tree = parse_sql_with_sll_fallback(sql, rule) - ast_tree = visit(antlr_tree) - return ast_tree + antlr_tree = ( + cached_request_tree(sql, rule) + if from_request + else cached_definition_tree(sql, rule) + ) + return visit(antlr_tree) @lru_cache(maxsize=128) @@ -219,9 +303,11 @@ def _cached_parse(sql: str | None) -> ast.Query: return parse(sql) -def parse(sql: str | None) -> ast.Query: +def parse(sql: str | None, from_request: bool = False) -> ast.Query: """ Parse a string sql query into a DJ ast Query + + Set ``from_request`` for SQL that came from a request (filters, orderby). """ import time as _time # noqa: PLC0415 @@ -230,7 +316,7 @@ def parse(sql: str | None) -> ast.Query: if not sql: raise DJParseException("Empty query provided!") try: - return cast(ast.Query, parse_rule(sql, "singleStatement")) + return cast(ast.Query, parse_rule(sql, "singleStatement", from_request)) except SqlParsingError as exc: raise DJParseException(message=f"Error parsing SQL `{sql}`: {exc}") from exc except DJParseException: diff --git a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py index 6a6360914..b15051910 100644 --- a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py +++ b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py @@ -5,7 +5,14 @@ import pytest -from datajunction_server.sql.parsing.backends.antlr4 import ast, parse +from datajunction_server.sql.parsing.backends.antlr4 import ( + _definition_tree_parser, + _request_tree_parser, + ast, + cached_request_tree, + parse, + report_parse_cache_stats, +) from datajunction_server.sql.parsing.backends.exceptions import DJParseException @@ -290,3 +297,72 @@ def test_unsupported_grammar_branch_surfaces_djparse(): bad_sql = "SELECT * FROM foo CROSS JOIN ()" with pytest.raises(DJParseException): parse(bad_sql) + + +def test_request_sql_tree_is_cached_between_parses(): + """The same request filter reuses one ANTLR tree instead of re-parsing.""" + _request_tree_parser().cache_clear() + sql = "SELECT 1 WHERE colx = 'cached'" + + first = cached_request_tree(sql, "singleStatement") + second = cached_request_tree(sql, "singleStatement") + + assert first is second + assert _request_tree_parser().cache_info().hits == 1 + + +def test_node_definitions_use_their_own_cache(): + """Node SQL and request SQL are cached separately.""" + _definition_tree_parser().cache_clear() + _request_tree_parser().cache_clear() + + parse("SELECT defn_col FROM defn_tbl") + parse("SELECT 1 WHERE colx = 'a'", from_request=True) + + assert _definition_tree_parser().cache_info().currsize == 1 + assert _request_tree_parser().cache_info().currsize == 1 + + +def test_cached_request_tree_yields_independent_asts(): + """Each parse gets a fresh AST, so mutating one cannot corrupt the next.""" + sql = "SELECT 1 WHERE amount = 5" + + first = parse(sql, from_request=True) + first.select.where.right.value = 99 + + second = parse(sql, from_request=True) + assert str(second) == str(parse(sql, from_request=True)) + assert "99" not in str(second) + + +def test_report_parse_cache_stats_emits_gauges(mocker): + """The request cache reports hits, misses and size.""" + provider = mocker.MagicMock() + mocker.patch( + "datajunction_server.instrumentation.provider.get_metrics_provider", + return_value=provider, + ) + + report_parse_cache_stats() + + reported = { + (call.args[0], call.args[2]["cache"]) for call in provider.gauge.call_args_list + } + assert reported == { + (name, cache) + for name in ( + "dj.sql.parse_cache.hits", + "dj.sql.parse_cache.misses", + "dj.sql.parse_cache.size", + "dj.sql.parse_cache.max_size", + ) + for cache in ("definitions", "requests") + } + + +def test_request_parse_cache_size_comes_from_settings(settings): + """The cache is sized from settings, not a hardcoded constant.""" + _request_tree_parser.cache_clear() + assert ( + _request_tree_parser().cache_info().maxsize == settings.request_parse_cache_size + )