From 6fdcec3c6cb0be689264a85d1572bddfbf0cc346 Mon Sep 17 00:00:00 2001 From: Yian Shang Date: Tue, 22 Sep 2026 18:06:29 -0700 Subject: [PATCH 1/4] Cache ANTLR parse trees between parses MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Profiling /sql/measures/v3/ and /sql/metrics/v3/ showed SQL text parsing dominating build time: `parse()` accounted for roughly 55% of a request, and up to two thirds of the calls re-parsed a string the same request had already parsed. The worst offenders were filter predicates (42 calls over 2 distinct strings on a 12-metric query), dimension link join SQL, and node query text. Memoizing `parse()` itself is not safe — the builder mutates the DJ AST it gets back, and `Function.__deepcopy__` returns `self`, so copies still alias. Caching one step lower is: `parse_rule` splits into an ANTLR parse and a visitor pass, the ANTLR tree is only ever read, and the mutable DJ AST is rebuilt on every call. Measured on the BUILD_V3 example model, with byte-identical responses across twelve metric/dimension/filter shapes: /sql/measures/v3/ 7-31% faster /sql/metrics/v3/ 5-29% faster The widest queries gain the most, since they re-parse the most. --- .../sql/parsing/backends/antlr4.py | 18 ++++++++++- .../tests/sql/parsing/backends/antlr4_test.py | 32 ++++++++++++++++++- 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py index 5038d7e20..12cb8433c 100644 --- a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py +++ b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py @@ -200,13 +200,29 @@ def tree_to_strings(tree, indent=0): return result +#: How many ANTLR parse trees to keep. Each entry retains its parser and token +#: stream, so this trades memory for parse time. +ANTLR_TREE_CACHE_SIZE = 512 + + +@lru_cache(maxsize=ANTLR_TREE_CACHE_SIZE) +def cached_antlr_tree(sql: str, rule: str): + """ + Parse a string into an ANTLR tree, caching the result. + + The tree is safe to share: ``visit`` only reads it, and the mutable DJ AST + is rebuilt on every call. + """ + return parse_sql_with_sll_fallback(sql, rule) + + def parse_rule(sql: str, rule: str) -> 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. """ - antlr_tree = parse_sql_with_sll_fallback(sql, rule) + antlr_tree = cached_antlr_tree(sql, rule) ast_tree = visit(antlr_tree) return ast_tree diff --git a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py index 6a6360914..68e6e6418 100644 --- a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py +++ b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py @@ -5,7 +5,11 @@ import pytest -from datajunction_server.sql.parsing.backends.antlr4 import ast, parse +from datajunction_server.sql.parsing.backends.antlr4 import ( + ast, + cached_antlr_tree, + parse, +) from datajunction_server.sql.parsing.backends.exceptions import DJParseException @@ -290,3 +294,29 @@ def test_unsupported_grammar_branch_surfaces_djparse(): bad_sql = "SELECT * FROM foo CROSS JOIN ()" with pytest.raises(DJParseException): parse(bad_sql) + + +def test_antlr_tree_is_cached_between_parses(): + """The same (sql, rule) reuses one ANTLR tree instead of re-parsing.""" + sql = "SELECT cached_col FROM cached_tbl WHERE cached_col > 1" + cached_antlr_tree.cache_clear() + + first = cached_antlr_tree(sql, "singleStatement") + second = cached_antlr_tree(sql, "singleStatement") + + assert first is second + assert cached_antlr_tree(sql, "singleStatement").getText() == first.getText() + + +def test_cached_antlr_tree_yields_independent_asts(): + """Each parse gets a fresh AST, so mutating one cannot corrupt the next.""" + sql = "SELECT SUM(amount) AS total FROM payments" + + first = parse(sql) + first.select.projection[0].alias.name = "mutated" + first.select.projection[0].child.name.name = "MAX" + + second = parse(sql) + assert str(second) == str(parse(sql)) + assert "SUM(amount)" in str(second) + assert "mutated" not in str(second) From 6473746e22d4dd82685c571d1f908d8a8f9e9456 Mon Sep 17 00:00:00 2001 From: Yian Shang Date: Tue, 22 Sep 2026 20:59:06 -0700 Subject: [PATCH 2/4] Give request SQL its own parse cache, and report both MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Filters and orderby clauses are parsed from request query params, so their variety is unbounded — every distinct literal a caller sends claimed a slot in the same cache holding node definitions. Under real traffic that churn would evict the entries worth keeping, and the hit rate would decay with nothing visibly breaking. Request SQL now uses a separate small cache. Callers opt in with `from_request=True`, which keeps the routing at the two sites that read request input rather than guessing from the SQL itself. Both caches report hits, misses, size and max size once per SQL build. Misses climbing while size sits at max size is the signal that a cache is undersized; today the example model fills 46 of 512 slots at a 98% hit rate, which says nothing about a production graph. --- .../datajunction_server/api/sql.py | 3 + .../construction/build_v3/builder.py | 2 +- .../construction/build_v3/filters.py | 2 +- .../datajunction_server/internal/sql.py | 2 + .../sql/parsing/backends/antlr4.py | 61 ++++++++++++++++--- .../tests/sql/parsing/backends/antlr4_test.py | 40 ++++++++++++ 6 files changed, 100 insertions(+), 10 deletions(-) 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/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 12cb8433c..7a4a09664 100644 --- a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py +++ b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py @@ -200,15 +200,18 @@ def tree_to_strings(tree, indent=0): return result -#: How many ANTLR parse trees to keep. Each entry retains its parser and token -#: stream, so this trades memory for parse time. +# Each entry retains its parser and token stream. ANTLR_TREE_CACHE_SIZE = 512 +# Request SQL is unbounded, so it gets a small cache of its own. Sharing one +# with node definitions would let filter churn evict them. +REQUEST_TREE_CACHE_SIZE = 64 + @lru_cache(maxsize=ANTLR_TREE_CACHE_SIZE) def cached_antlr_tree(sql: str, rule: str): """ - Parse a string into an ANTLR tree, caching the result. + Parse a node definition into an ANTLR tree, caching the result. The tree is safe to share: ``visit`` only reads it, and the mutable DJ AST is rebuilt on every call. @@ -216,13 +219,53 @@ def cached_antlr_tree(sql: str, rule: str): return parse_sql_with_sll_fallback(sql, rule) -def parse_rule(sql: str, rule: str) -> Union[ast.Node, "ColumnType"]: +@lru_cache(maxsize=REQUEST_TREE_CACHE_SIZE) +def cached_request_tree(sql: str, rule: str): + """ + Parse request-supplied SQL into an ANTLR tree, caching the result. + + Filters and orderby clauses are parsed many times within one request but + vary freely across requests, so this cache is kept small. + """ + return parse_sql_with_sll_fallback(sql, rule) + + +def report_parse_cache_stats() -> None: + """ + Report tree cache size and hit counts. + + 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": cached_antlr_tree, "requests": cached_request_tree} + 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 SQL that came from a request, so that its unbounded + variety cannot evict cached node definitions. """ - antlr_tree = cached_antlr_tree(sql, rule) + tree_cache = cached_request_tree if from_request else cached_antlr_tree + antlr_tree = tree_cache(sql, rule) ast_tree = visit(antlr_tree) return ast_tree @@ -235,9 +278,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 @@ -246,7 +291,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 68e6e6418..6d90c6070 100644 --- a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py +++ b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py @@ -8,7 +8,9 @@ from datajunction_server.sql.parsing.backends.antlr4 import ( ast, cached_antlr_tree, + cached_request_tree, parse, + report_parse_cache_stats, ) from datajunction_server.sql.parsing.backends.exceptions import DJParseException @@ -320,3 +322,41 @@ def test_cached_antlr_tree_yields_independent_asts(): assert str(second) == str(parse(sql)) assert "SUM(amount)" in str(second) assert "mutated" not in str(second) + + +def test_request_sql_uses_its_own_cache(): + """Request SQL must not take slots from the node definition cache.""" + cached_antlr_tree.cache_clear() + cached_request_tree.cache_clear() + + parse("SELECT 1 WHERE colx = 'a'", from_request=True) + parse("SELECT 1 WHERE colx = 'b'", from_request=True) + parse("SELECT defn_col FROM defn_tbl") + + assert cached_request_tree.cache_info().currsize == 2 + assert cached_antlr_tree.cache_info().currsize == 1 + + +def test_report_parse_cache_stats_emits_gauges(mocker): + """Both caches report hits, misses and size under their own tag.""" + 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") + } From 53271bfd96c5831ad3097828deaf0c42127ca55e Mon Sep 17 00:00:00 2001 From: Yian Shang Date: Wed, 23 Sep 2026 01:40:09 -0700 Subject: [PATCH 3/4] Cache only request SQL, and size it from settings MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replaying a day of prod v3 traffic (1,212 requests) against a prod-sized graph showed the previous split was backwards. Caching node definitions is worth nothing on this workload — it measured 4% *slower* than no cache at all, on a 52% win rate. Every bit of the gain comes from caching the SQL that arrives on the request: filters and orderby clauses. The reason is that `BuildContext.get_parsed_query` already caches parsed node ASTs within a request, so a definitions cache only helps across requests — and real traffic repeats node definitions far less than expected (755 distinct request shapes over 963 distinct metrics). Filters are the opposite: a single build parses the same predicate dozens of times. Measured against the same corpus, caches off vs on: p50 238ms -> 170ms p90 914ms -> 757ms p99 3944ms -> 3377ms mean 507ms -> 408ms (18.7%, faster on 1083/1175 requests) Dropping the definitions cache also drops the memory question it carried: each entry retained a parser and token stream for a multi-KB node query, per worker process. What remains caches `SELECT 1 WHERE `. The size moves to `Settings.request_parse_cache_size` so it can be retuned from the counters without shipping code. The cache is built on first use because `lru_cache` binds `maxsize` at decoration time while settings resolve lazily. --- .../datajunction_server/config.py | 4 + .../sql/parsing/backends/antlr4.py | 62 +++++++-------- .../tests/sql/parsing/backends/antlr4_test.py | 75 +++++++++---------- 3 files changed, 68 insertions(+), 73 deletions(-) diff --git a/datajunction-server/datajunction_server/config.py b/datajunction-server/datajunction_server/config.py index 808aa552e..c80df85a4 100644 --- a/datajunction-server/datajunction_server/config.py +++ b/datajunction-server/datajunction_server/config.py @@ -285,6 +285,10 @@ 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 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/sql/parsing/backends/antlr4.py b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py index 7a4a09664..8bb45150d 100644 --- a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py +++ b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py @@ -200,39 +200,34 @@ def tree_to_strings(tree, indent=0): return result -# Each entry retains its parser and token stream. -ANTLR_TREE_CACHE_SIZE = 512 - -# Request SQL is unbounded, so it gets a small cache of its own. Sharing one -# with node definitions would let filter churn evict them. -REQUEST_TREE_CACHE_SIZE = 64 - - -@lru_cache(maxsize=ANTLR_TREE_CACHE_SIZE) -def cached_antlr_tree(sql: str, rule: str): +@lru_cache(maxsize=1) +def _request_tree_parser(): """ - Parse a node definition into an ANTLR tree, caching the result. - - The tree is safe to share: ``visit`` only reads it, and the mutable DJ AST - is rebuilt on every call. + Build the request SQL cache, sized from settings on first use. """ - return parse_sql_with_sll_fallback(sql, rule) + 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 -@lru_cache(maxsize=REQUEST_TREE_CACHE_SIZE) def cached_request_tree(sql: str, rule: str): """ Parse request-supplied SQL into an ANTLR tree, caching the result. - Filters and orderby clauses are parsed many times within one request but - vary freely across requests, so this cache is kept small. + 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 parse_sql_with_sll_fallback(sql, rule) + return _request_tree_parser()(sql, rule) def report_parse_cache_stats() -> None: """ - Report tree cache size and hit counts. + Report request SQL cache size and hit counts. Misses rising while ``size`` sits at ``max_size`` means the cache is too small and entries are being evicted. @@ -242,14 +237,11 @@ def report_parse_cache_stats() -> None: ) provider = get_metrics_provider() - caches = {"definitions": cached_antlr_tree, "requests": cached_request_tree} - 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) + info = _request_tree_parser().cache_info() + provider.gauge("dj.sql.parse_cache.hits", info.hits) + provider.gauge("dj.sql.parse_cache.misses", info.misses) + provider.gauge("dj.sql.parse_cache.size", info.currsize) + provider.gauge("dj.sql.parse_cache.max_size", info.maxsize) def parse_rule( @@ -261,13 +253,15 @@ def parse_rule( Parse a string into a DJ ast using the ANTLR4 backend. Uses SLL mode first (faster), falls back to LL mode if needed. Set - ``from_request`` for SQL that came from a request, so that its unbounded - variety cannot evict cached node definitions. + ``from_request`` for filters and orderby clauses, which one request + parses repeatedly. """ - tree_cache = cached_request_tree if from_request else cached_antlr_tree - antlr_tree = tree_cache(sql, rule) - ast_tree = visit(antlr_tree) - return ast_tree + antlr_tree = ( + cached_request_tree(sql, rule) + if from_request + else parse_sql_with_sll_fallback(sql, rule) + ) + return visit(antlr_tree) @lru_cache(maxsize=128) diff --git a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py index 6d90c6070..e3736713c 100644 --- a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py +++ b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py @@ -6,8 +6,8 @@ import pytest from datajunction_server.sql.parsing.backends.antlr4 import ( + _request_tree_parser, ast, - cached_antlr_tree, cached_request_tree, parse, report_parse_cache_stats, @@ -298,47 +298,42 @@ def test_unsupported_grammar_branch_surfaces_djparse(): parse(bad_sql) -def test_antlr_tree_is_cached_between_parses(): - """The same (sql, rule) reuses one ANTLR tree instead of re-parsing.""" - sql = "SELECT cached_col FROM cached_tbl WHERE cached_col > 1" - cached_antlr_tree.cache_clear() +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_antlr_tree(sql, "singleStatement") - second = cached_antlr_tree(sql, "singleStatement") + first = cached_request_tree(sql, "singleStatement") + second = cached_request_tree(sql, "singleStatement") assert first is second - assert cached_antlr_tree(sql, "singleStatement").getText() == first.getText() + assert _request_tree_parser().cache_info().hits == 1 -def test_cached_antlr_tree_yields_independent_asts(): - """Each parse gets a fresh AST, so mutating one cannot corrupt the next.""" - sql = "SELECT SUM(amount) AS total FROM payments" +def test_node_definitions_are_not_cached(): + """Node SQL is parsed fresh; only request SQL is cached.""" + _request_tree_parser().cache_clear() - first = parse(sql) - first.select.projection[0].alias.name = "mutated" - first.select.projection[0].child.name.name = "MAX" + parse("SELECT defn_col FROM defn_tbl") + parse("SELECT defn_col FROM defn_tbl") - second = parse(sql) - assert str(second) == str(parse(sql)) - assert "SUM(amount)" in str(second) - assert "mutated" not in str(second) + assert _request_tree_parser().cache_info().currsize == 0 -def test_request_sql_uses_its_own_cache(): - """Request SQL must not take slots from the node definition cache.""" - cached_antlr_tree.cache_clear() - cached_request_tree.cache_clear() +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" - parse("SELECT 1 WHERE colx = 'a'", from_request=True) - parse("SELECT 1 WHERE colx = 'b'", from_request=True) - parse("SELECT defn_col FROM defn_tbl") + first = parse(sql, from_request=True) + first.select.where.right.value = 99 - assert cached_request_tree.cache_info().currsize == 2 - assert cached_antlr_tree.cache_info().currsize == 1 + 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): - """Both caches report hits, misses and size under their own tag.""" + """The request cache reports hits, misses and size.""" provider = mocker.MagicMock() mocker.patch( "datajunction_server.instrumentation.provider.get_metrics_provider", @@ -347,16 +342,18 @@ def test_report_parse_cache_stats_emits_gauges(mocker): report_parse_cache_stats() - reported = { - (call.args[0], call.args[2]["cache"]) for call in provider.gauge.call_args_list - } + reported = {call.args[0] 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") + "dj.sql.parse_cache.hits", + "dj.sql.parse_cache.misses", + "dj.sql.parse_cache.size", + "dj.sql.parse_cache.max_size", } + + +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 + ) From 6c4b3bb86300cb5a466841ecf3300395c7adbdfe Mon Sep 17 00:00:00 2001 From: Yian Shang Date: Wed, 23 Sep 2026 11:51:59 -0700 Subject: [PATCH 4/4] Restore the node definition cache, and correct the numbers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The previous commit removed the node definition cache on the strength of two benchmark runs that turned out to be confounded. Those arms ran fourth and fifth in a sequence whose wall time climbed steadily run over run (468s, 597s, 615s, 672s) regardless of configuration, so the slowdown they showed was the environment, not the design. Re-measured with the three configurations interleaved, three repetitions each, so any drift spreads across all arms instead of landing on whichever ran last. Baseline reps came in at 515/503/507ms, so the environment was stable this time and the gaps below are real. config r1 r2 r3 mean vs off no caching 515 503 507 508ms - request only 482 496 501 493ms -3.0% both caches 459 457 466 460ms -9.4% Paired per request, pooling all three repetitions: both caches beat no caching on 947/1173 requests, and beat request-only caching on 855/1173. So both caches earn their place, the node definition cache more than the request one, and the change is worth about 9% on this traffic — not the 19% claimed two commits ago, which came from unreplicated runs. Both sizes now live in Settings so they can be retuned from the counters without shipping code. Measured on 400 prod v3 requests replayed against a prod-sized graph, with both caches cold at the start of every run. A long-lived worker should do better. The legacy /sql path, which carries far more traffic, is untested. --- .../datajunction_server/config.py | 4 ++ .../sql/parsing/backends/antlr4.py | 45 ++++++++++++++++--- .../tests/sql/parsing/backends/antlr4_test.py | 27 +++++++---- 3 files changed, 60 insertions(+), 16 deletions(-) diff --git a/datajunction-server/datajunction_server/config.py b/datajunction-server/datajunction_server/config.py index c80df85a4..39941bca4 100644 --- a/datajunction-server/datajunction_server/config.py +++ b/datajunction-server/datajunction_server/config.py @@ -285,6 +285,10 @@ 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 diff --git a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py index 8bb45150d..56212502b 100644 --- a/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py +++ b/datajunction-server/datajunction_server/sql/parsing/backends/antlr4.py @@ -200,6 +200,20 @@ def tree_to_strings(tree, indent=0): return result +@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(): """ @@ -214,6 +228,17 @@ def _parse(sql: str, rule: str): 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. @@ -227,7 +252,7 @@ def cached_request_tree(sql: str, rule: str): def report_parse_cache_stats() -> None: """ - Report request SQL cache size and hit counts. + 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. @@ -237,11 +262,17 @@ def report_parse_cache_stats() -> None: ) provider = get_metrics_provider() - info = _request_tree_parser().cache_info() - provider.gauge("dj.sql.parse_cache.hits", info.hits) - provider.gauge("dj.sql.parse_cache.misses", info.misses) - provider.gauge("dj.sql.parse_cache.size", info.currsize) - provider.gauge("dj.sql.parse_cache.max_size", info.maxsize) + 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( @@ -259,7 +290,7 @@ def parse_rule( antlr_tree = ( cached_request_tree(sql, rule) if from_request - else parse_sql_with_sll_fallback(sql, rule) + else cached_definition_tree(sql, rule) ) return visit(antlr_tree) diff --git a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py index e3736713c..b15051910 100644 --- a/datajunction-server/tests/sql/parsing/backends/antlr4_test.py +++ b/datajunction-server/tests/sql/parsing/backends/antlr4_test.py @@ -6,6 +6,7 @@ import pytest from datajunction_server.sql.parsing.backends.antlr4 import ( + _definition_tree_parser, _request_tree_parser, ast, cached_request_tree, @@ -310,14 +311,16 @@ def test_request_sql_tree_is_cached_between_parses(): assert _request_tree_parser().cache_info().hits == 1 -def test_node_definitions_are_not_cached(): - """Node SQL is parsed fresh; only request SQL is cached.""" +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 defn_col FROM defn_tbl") + parse("SELECT 1 WHERE colx = 'a'", from_request=True) - assert _request_tree_parser().cache_info().currsize == 0 + assert _definition_tree_parser().cache_info().currsize == 1 + assert _request_tree_parser().cache_info().currsize == 1 def test_cached_request_tree_yields_independent_asts(): @@ -342,12 +345,18 @@ def test_report_parse_cache_stats_emits_gauges(mocker): report_parse_cache_stats() - reported = {call.args[0] for call in provider.gauge.call_args_list} + reported = { + (call.args[0], call.args[2]["cache"]) for call in provider.gauge.call_args_list + } assert reported == { - "dj.sql.parse_cache.hits", - "dj.sql.parse_cache.misses", - "dj.sql.parse_cache.size", - "dj.sql.parse_cache.max_size", + (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") }