Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions datajunction-server/datajunction_server/api/sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)

Expand Down
8 changes: 8 additions & 0 deletions datajunction-server/datajunction_server/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 []
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ def parse_filter(filter_str: str) -> ast.Expression:
# Returns ast.BinaryOp with comparison
"""
# Parse as "SELECT 1 WHERE <filter>" 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
Expand Down
2 changes: 2 additions & 0 deletions datajunction-server/datajunction_server/internal/sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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

Expand All @@ -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:
Expand Down
78 changes: 77 additions & 1 deletion datajunction-server/tests/sql/parsing/backends/antlr4_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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
)
Loading