Skip to content

Commit 03c3593

Browse files
committed
WIP
1 parent 419f440 commit 03c3593

2 files changed

Lines changed: 24 additions & 34 deletions

File tree

src/dbt_core_interface/dbt_templater/__init__.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,23 +1,26 @@
11
"""Defines the hook endpoints for the dbt templater plugin."""
2+
23
import logging
34

45
from sqlfluff.core.plugin import hookimpl
56

67
from dbt_core_interface.dbt_templater.templater import DCIDbtTemplater
78

8-
99
LOGGER = logging.getLogger(__name__)
1010

1111

1212
@hookimpl
1313
def get_templaters():
1414
"""Get templaters."""
15+
1516
def create_templater(**kwargs):
1617
import dbt_core_interface.state
17-
assert dbt_core_interface.state.dbt_project_container is not None, "dbt_project_container is None"
18+
19+
assert dbt_core_interface.state.dbt_project_container is not None, (
20+
"dbt_project_container is None"
21+
)
1822
return DCIDbtTemplater(
19-
dbt_project_container=dbt_core_interface.state.dbt_project_container,
20-
**kwargs
23+
dbt_project_container=dbt_core_interface.state.dbt_project_container, **kwargs
2124
)
2225

2326
create_templater.name = DCIDbtTemplater.name

src/dbt_core_interface/project.py

Lines changed: 17 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,6 @@
99

1010
dbt.adapters.factory.get_adapter = lambda config: config.adapter # pyright: ignore[reportUnknownLambdaType]
1111

12-
import argparse
1312
import contextlib
1413
import json
1514
import logging
@@ -75,7 +74,8 @@ class DbtConfiguration:
7574
profiles_dir: str = DEFAULT_PROFILES_DIR
7675
target: str | None = None
7776
threads: int = 1
78-
vars: dict[str, t.Any] = {}
77+
vars: dict[str, t.Any] = field(default_factory=dict)
78+
profile: str | None = None
7979

8080
single_threaded: bool = True
8181
quiet: bool = True
@@ -84,7 +84,7 @@ class DbtConfiguration:
8484
partial_parse: bool = False
8585

8686
dependencies: list[str] = field(default_factory=list)
87-
which: str = "cupertino"
87+
which: str = "zezima was here"
8888
REQUIRE_RESOURCE_NAMES_WITHOUT_SPACES: bool = field(default_factory=bool)
8989

9090
def __post_init__(self) -> None:
@@ -237,33 +237,14 @@ def _initialize_adapter(self) -> None:
237237
self._adapter.debug_query()
238238

239239
logger.debug(f"Initialized adapter for {self.project_name}")
240-
self.runtime_config.adapter = self.adapter # pyright: ignore[reportAttributeAccessIssue]
240+
self.runtime_config.adapter = self._adapter # pyright: ignore[reportAttributeAccessIssue]
241241

242242
def _parse_project(self, write_manifest: bool = True) -> None:
243243
"""Parse the dbt project and load manifest."""
244-
set_from_args(
245-
argparse.Namespace(
246-
target=self._base_params.target,
247-
profiles_dir=self._base_params.profiles_dir,
248-
project_dir=self._base_params.project_dir,
249-
threads=self._base_params.threads,
250-
vars=self._base_params.vars,
251-
quiet=True,
252-
single_threaded=True,
253-
),
254-
None,
255-
)
244+
set_from_args(self._base_params, None) # pyright: ignore[reportArgumentType]
256245

257246
with self._manifest_lock:
258-
self._runtime_config = RuntimeConfig.from_args(
259-
argparse.Namespace(
260-
target=self._base_params.target,
261-
profiles_dir=self._base_params.profiles_dir,
262-
project_dir=self._base_params.project_dir,
263-
threads=self._base_params.threads,
264-
vars=self._base_params.vars,
265-
)
266-
)
247+
self._runtime_config = RuntimeConfig.from_args(self._base_params)
267248

268249
manifest_loader = ManifestLoader(
269250
self._runtime_config,
@@ -313,14 +294,17 @@ def generate_runtime_model_context(self, node: ManifestNode) -> dict[str, t.Any]
313294
self.manifest,
314295
)
315296

316-
def get_ref_node(
297+
def ref(
317298
self,
318299
model_name: str,
319300
target_package: str | None = None,
320301
model_version: int | None = None,
321302
source_node: ManifestNode | None = None,
322303
) -> ManifestNode | None:
323-
"""Get node by ref (model) name."""
304+
"""Look up a model node by name, package, and version.
305+
306+
Akin to using {{ ref() }} in SQL.
307+
"""
324308
candidates: list[str | None] = [self.project_name, None]
325309
if target_package:
326310
candidates.insert(0, target_package)
@@ -332,8 +316,11 @@ def get_ref_node(
332316
if node:
333317
return node
334318

335-
def get_source_node(self, source_name: str, table_name: str) -> SourceDefinition | None:
336-
"""Get source node by source and table name."""
319+
def source(self, source_name: str, table_name: str) -> SourceDefinition | None:
320+
"""Look up a source by name and table name.
321+
322+
Akin to using {{ source() }} in SQL.
323+
"""
337324
return self.manifest.source_lookup.find(f"{source_name}.{table_name}", None, self.manifest)
338325

339326
def execute_sql(self, sql: str) -> ExecutionResult:
@@ -409,7 +396,7 @@ def _create_temp_node(
409396

410397
def _cleanup() -> None:
411398
with contextlib.suppress(KeyError):
412-
del self.manifest.nodes[node_id]
399+
del self.manifest.nodes[sql_node.unique_id]
413400

414401
return sql_node, _cleanup
415402

0 commit comments

Comments
 (0)