Skip to content

Commit 92e9514

Browse files
authored
perf: avoid full snapshot mapping in _resolve_table/_resolve_tables (#6068)
1 parent 849d6ed commit 92e9514

5 files changed

Lines changed: 566 additions & 19 deletions

File tree

‎sqlmesh/core/renderer.py‎

Lines changed: 88 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,65 @@
4545
logger = logging.getLogger(__name__)
4646

4747

48+
class TableMapping(t.Dict[str, str]):
49+
"""A table name mapping that caches the dialect-normalized form of its keys.
50+
51+
`exp.replace_tables` normalizes every key of the mapping it's given, so resolving a single
52+
table against a mapping of every model in an environment costs O(N). Resolving it against
53+
this mapping costs a dictionary lookup, since each key is normalized once per dialect.
54+
"""
55+
56+
def __init__(self, *args: t.Any, **kwargs: t.Any):
57+
super().__init__(*args, **kwargs)
58+
self._normalized_keys: t.Dict[DialectType, t.Dict[str, str]] = {}
59+
60+
def normalized_keys(self, dialect: DialectType) -> t.Dict[str, str]:
61+
"""Returns a mapping from each normalized key to the last key that normalizes to it."""
62+
normalized_keys = self._normalized_keys.get(dialect)
63+
if normalized_keys is None:
64+
normalized_keys = {exp.normalize_table_name(key, dialect=dialect): key for key in self}
65+
self._normalized_keys[dialect] = normalized_keys
66+
return normalized_keys
67+
68+
def __setitem__(self, key: str, value: str) -> None:
69+
self._normalized_keys.clear()
70+
super().__setitem__(key, value)
71+
72+
def __delitem__(self, key: str) -> None:
73+
self._normalized_keys.clear()
74+
super().__delitem__(key)
75+
76+
def __ior__(self, other: t.Any) -> TableMapping: # type: ignore[override,misc]
77+
self._normalized_keys.clear()
78+
return super().__ior__(other)
79+
80+
def update(self, *args: t.Any, **kwargs: t.Any) -> None:
81+
self._normalized_keys.clear()
82+
super().update(*args, **kwargs)
83+
84+
def setdefault(self, key: str, default: str) -> str: # type: ignore[override]
85+
self._normalized_keys.clear()
86+
return super().setdefault(key, default)
87+
88+
def pop(self, key: str, *args: t.Any) -> t.Any:
89+
self._normalized_keys.clear()
90+
return super().pop(key, *args)
91+
92+
def popitem(self) -> t.Tuple[str, str]:
93+
self._normalized_keys.clear()
94+
return super().popitem()
95+
96+
def clear(self) -> None:
97+
self._normalized_keys.clear()
98+
super().clear()
99+
100+
101+
def _normalize_keys(mapping: t.Dict[str, str], dialect: DialectType) -> t.Dict[str, str]:
102+
if isinstance(mapping, TableMapping):
103+
return mapping.normalized_keys(dialect)
104+
return {exp.normalize_table_name(key, dialect=dialect): key for key in mapping}
105+
106+
48107
class BaseExpressionRenderer:
49108
def __init__(
50109
self,
@@ -325,20 +384,35 @@ def update_cache(self, expression: t.Optional[exp.Expr]) -> None:
325384

326385
def _resolve_table(
327386
self,
328-
table_name: str | exp.Expr,
387+
table_name: str,
329388
snapshots: t.Optional[t.Dict[str, Snapshot]] = None,
330389
table_mapping: t.Optional[t.Dict[str, str]] = None,
331390
deployability_index: t.Optional[DeployabilityIndex] = None,
332391
) -> exp.Table:
333-
table = exp.replace_tables(
334-
t.cast(exp.Table, exp.maybe_parse(table_name, into=exp.Table, dialect=self._dialect)),
335-
{
336-
**self._to_table_mapping((snapshots or {}).values(), deployability_index),
337-
**(table_mapping or {}),
338-
},
339-
dialect=self._dialect,
340-
copy=False,
392+
table = t.cast(
393+
exp.Table, exp.maybe_parse(table_name, into=exp.Table, dialect=self._dialect)
341394
)
395+
396+
mapping: t.Dict[str, str] = {}
397+
if table_mapping:
398+
# An explicit mapping takes precedence over snapshots, so when one of its keys matches
399+
# the table, that key alone decides the result. Among equivalent keys, the last wins.
400+
key = _normalize_keys(table_mapping, self._dialect).get(
401+
exp.normalize_table_name(table, dialect=self._dialect)
402+
)
403+
if key is not None:
404+
mapping = {key: table_mapping[key]}
405+
406+
if not mapping and snapshots:
407+
# An exact FQN match avoids scanning unrelated snapshots.
408+
snapshot = snapshots.get(table_name)
409+
# Keys normalized under different dialects may differ in casing or quoting.
410+
# Fall back to the full mapping so exp.replace_tables can reconcile them.
411+
mapping = self._to_table_mapping(
412+
[snapshot] if snapshot else snapshots.values(), deployability_index
413+
)
414+
415+
table = exp.replace_tables(table, mapping, dialect=self._dialect, copy=False)
342416
# We quote the table here to mimic the behavior of _resolve_tables, otherwise we may end
343417
# up normalizing twice, because _to_table_mapping returns the mapped names unquoted.
344418
return (
@@ -363,6 +437,11 @@ def _resolve_tables(
363437

364438
expression = expression.copy()
365439
with self._normalize_and_quote(expression) as expression:
440+
# An expression with no table (e.g. most session or virtual properties) has nothing
441+
# to expand or replace, so skip building the O(N) expand set and mapping.
442+
if not expression.find(exp.Table):
443+
return expression
444+
366445
snapshots = snapshots or {}
367446
table_mapping = table_mapping or {}
368447
mapping = {

‎sqlmesh/core/snapshot/definition.py‎

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
from sqlmesh.core.model import Model, ModelKindMixin, ModelKindName, ViewKind, CustomKind
2525
from sqlmesh.core.model.definition import _Model
2626
from sqlmesh.core.node import IntervalUnit, NodeType
27+
from sqlmesh.core.renderer import TableMapping
2728
from sqlmesh.utils import sanitize_name, unique
2829
from sqlmesh.utils.dag import DAG
2930
from sqlmesh.utils.date import (
@@ -2007,14 +2008,17 @@ def to_view_mapping(
20072008
environment_naming_info: EnvironmentNamingInfo,
20082009
default_catalog: t.Optional[str] = None,
20092010
dialect: t.Optional[str] = None,
2010-
) -> t.Dict[str, str]:
2011-
return {
2012-
snapshot.name: snapshot.display_name(
2013-
environment_naming_info, default_catalog=default_catalog, dialect=dialect
2011+
) -> TableMapping:
2012+
return TableMapping(
2013+
(
2014+
snapshot.name,
2015+
snapshot.display_name(
2016+
environment_naming_info, default_catalog=default_catalog, dialect=dialect
2017+
),
20142018
)
20152019
for snapshot in snapshots
20162020
if snapshot.is_model
2017-
}
2021+
)
20182022

20192023

20202024
def has_paused_forward_only(

‎sqlmesh/core/snapshot/evaluator.py‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -312,6 +312,8 @@ def promote(
312312
self._get_virtual_data_objects(target_snapshots, environment_naming_info)
313313

314314
deployability_index = deployability_index or DeployabilityIndex.all_deployable()
315+
# Renderers look snapshots up by model name, not by SnapshotId.
316+
snapshots_by_name = {s.name: s for s in (snapshots or {}).values()}
315317
with self.concurrent_context():
316318
concurrent_apply_to_snapshots(
317319
target_snapshots,
@@ -320,7 +322,7 @@ def promote(
320322
start=start,
321323
end=end,
322324
execution_time=execution_time,
323-
snapshots=snapshots,
325+
snapshots=snapshots_by_name,
324326
table_mapping=table_mapping,
325327
environment_naming_info=environment_naming_info,
326328
deployability_index=deployability_index, # type: ignore
@@ -1260,7 +1262,7 @@ def _promote_snapshot(
12601262
start: t.Optional[TimeLike] = None,
12611263
end: t.Optional[TimeLike] = None,
12621264
execution_time: t.Optional[TimeLike] = None,
1263-
snapshots: t.Optional[t.Dict[SnapshotId, Snapshot]] = None,
1265+
snapshots: t.Optional[t.Dict[str, Snapshot]] = None,
12641266
table_mapping: t.Optional[t.Dict[str, str]] = None,
12651267
) -> None:
12661268
if not snapshot.is_model:
@@ -1299,9 +1301,9 @@ def _promote_snapshot(
12991301
**render_kwargs,
13001302
)
13011303

1302-
snapshot_by_name = {s.name: s for s in (snapshots or {}).values()}
1303-
render_kwargs["snapshots"] = snapshot_by_name
1304-
adapter.execute(snapshot.model.render_on_virtual_update(**render_kwargs))
1304+
adapter.execute(
1305+
snapshot.model.render_on_virtual_update(snapshots=snapshots, **render_kwargs)
1306+
)
13051307

13061308
if on_complete is not None:
13071309
on_complete(snapshot)

0 commit comments

Comments
 (0)