flyteplugins-dbt 0.0.0a0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- flyteplugins/dbt/__init__.py +12 -0
- flyteplugins/dbt/resolver.py +61 -0
- flyteplugins/dbt/runner.py +308 -0
- flyteplugins/dbt/task.py +251 -0
- flyteplugins_dbt-0.0.0a0.dist-info/METADATA +93 -0
- flyteplugins_dbt-0.0.0a0.dist-info/RECORD +8 -0
- flyteplugins_dbt-0.0.0a0.dist-info/WHEEL +5 -0
- flyteplugins_dbt-0.0.0a0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
from flyteplugins.dbt.resolver import DbtTaskResolver
|
|
2
|
+
from flyteplugins.dbt.runner import DbtEventCallback, DbtInvocationError, DbtNodeResult, invoke_dbt
|
|
3
|
+
from flyteplugins.dbt.task import DbtTask
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
"DbtEventCallback",
|
|
7
|
+
"DbtInvocationError",
|
|
8
|
+
"DbtNodeResult",
|
|
9
|
+
"DbtTask",
|
|
10
|
+
"DbtTaskResolver",
|
|
11
|
+
"invoke_dbt",
|
|
12
|
+
]
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import pathlib
|
|
5
|
+
|
|
6
|
+
from flyte._task import TaskTemplate
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class DbtTaskResolver:
|
|
10
|
+
"""Reconstructs a DbtTask in the remote task container."""
|
|
11
|
+
|
|
12
|
+
@property
|
|
13
|
+
def import_path(self) -> str:
|
|
14
|
+
return "flyteplugins.dbt.resolver.DbtTaskResolver"
|
|
15
|
+
|
|
16
|
+
def load_task(self, loader_args: list[str]) -> TaskTemplate:
|
|
17
|
+
from flyteplugins.dbt.task import DbtTask
|
|
18
|
+
|
|
19
|
+
it = iter(loader_args)
|
|
20
|
+
args_dict: dict[str, str] = {}
|
|
21
|
+
for key in it:
|
|
22
|
+
try:
|
|
23
|
+
args_dict[key] = next(it)
|
|
24
|
+
except StopIteration:
|
|
25
|
+
raise ValueError(f"Odd number of loader args: missing value for key '{key}'")
|
|
26
|
+
|
|
27
|
+
callbacks = json.loads(args_dict.get("callbacks_json", "[]"))
|
|
28
|
+
|
|
29
|
+
return DbtTask(
|
|
30
|
+
name=args_dict["name"],
|
|
31
|
+
task_environment=None,
|
|
32
|
+
project_dir=args_dict.get("project_dir") or None,
|
|
33
|
+
profiles_dir=args_dict.get("profiles_dir") or None,
|
|
34
|
+
profile=args_dict.get("profile") or None,
|
|
35
|
+
target_path=args_dict.get("target_path") or None,
|
|
36
|
+
callbacks=callbacks,
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
def loader_args(self, task: TaskTemplate, root_dir: pathlib.Path | None = None) -> list[str]:
|
|
40
|
+
from flyteplugins.dbt.runner import callback_import_paths
|
|
41
|
+
from flyteplugins.dbt.task import DbtTask
|
|
42
|
+
|
|
43
|
+
if not isinstance(task, DbtTask):
|
|
44
|
+
raise TypeError(f"DbtTaskResolver only handles DbtTask, got {type(task)}")
|
|
45
|
+
|
|
46
|
+
callback_paths = callback_import_paths(task.callbacks, source_dir=root_dir)
|
|
47
|
+
|
|
48
|
+
return [
|
|
49
|
+
"name",
|
|
50
|
+
task.name,
|
|
51
|
+
"project_dir",
|
|
52
|
+
task.project_dir or "",
|
|
53
|
+
"profiles_dir",
|
|
54
|
+
task.profiles_dir or "",
|
|
55
|
+
"profile",
|
|
56
|
+
task.profile or "",
|
|
57
|
+
"target_path",
|
|
58
|
+
task.target_path or "",
|
|
59
|
+
"callbacks_json",
|
|
60
|
+
json.dumps(callback_paths),
|
|
61
|
+
]
|
|
@@ -0,0 +1,308 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import html
|
|
4
|
+
import importlib
|
|
5
|
+
import pathlib
|
|
6
|
+
from collections.abc import Callable, Sequence
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from typing import Any, Optional
|
|
9
|
+
|
|
10
|
+
_MAX_DBT_INVOCATION_ERROR_RESULTS = 10
|
|
11
|
+
|
|
12
|
+
DbtEventCallback = Callable[[Any], None]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass
|
|
16
|
+
class DbtNodeResult:
|
|
17
|
+
"""Small serializable summary of one dbt node result."""
|
|
18
|
+
|
|
19
|
+
unique_id: str
|
|
20
|
+
name: str
|
|
21
|
+
resource_type: str
|
|
22
|
+
status: str
|
|
23
|
+
message: Optional[str] = None
|
|
24
|
+
failures: Optional[int] = None
|
|
25
|
+
execution_time: Optional[float] = None
|
|
26
|
+
relation_name: Optional[str] = None
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class DbtInvocationError(RuntimeError):
|
|
30
|
+
"""Raised when dbt finishes cleanly but reports failed node results."""
|
|
31
|
+
|
|
32
|
+
def __init__(self, cli_args: list[str], results: list[DbtNodeResult]):
|
|
33
|
+
self.cli_args = cli_args
|
|
34
|
+
self.results = results
|
|
35
|
+
super().__init__(_format_dbt_invocation_error(cli_args, results))
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _format_dbt_invocation_error(cli_args: list[str], results: list[DbtNodeResult]) -> str:
|
|
39
|
+
if not results:
|
|
40
|
+
return f"dbt invocation failed for cli_args={cli_args!r}"
|
|
41
|
+
|
|
42
|
+
failed_results = [result for result in results if result.status.lower() not in {"pass", "success"}]
|
|
43
|
+
if not failed_results:
|
|
44
|
+
failed_results = results
|
|
45
|
+
|
|
46
|
+
node_summaries = []
|
|
47
|
+
for result in failed_results[:_MAX_DBT_INVOCATION_ERROR_RESULTS]:
|
|
48
|
+
details = [f"status={result.status!r}"]
|
|
49
|
+
if result.failures is not None:
|
|
50
|
+
details.append(f"failures={result.failures}")
|
|
51
|
+
if result.message:
|
|
52
|
+
details.append(f"message={result.message!r}")
|
|
53
|
+
node_summaries.append(f"{result.unique_id or result.name} ({', '.join(details)})")
|
|
54
|
+
|
|
55
|
+
suffix = ""
|
|
56
|
+
if len(failed_results) > _MAX_DBT_INVOCATION_ERROR_RESULTS:
|
|
57
|
+
suffix = f"; and {len(failed_results) - _MAX_DBT_INVOCATION_ERROR_RESULTS} more"
|
|
58
|
+
|
|
59
|
+
return f"dbt invocation failed for cli_args={cli_args!r}: {', '.join(node_summaries)}{suffix}"
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _node_name(node: Any) -> str:
|
|
63
|
+
return str(getattr(node, "name", "") or getattr(node, "unique_id", ""))
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _node_resource_type(node: Any) -> str:
|
|
67
|
+
resource_type = getattr(node, "resource_type", "")
|
|
68
|
+
if hasattr(resource_type, "value"):
|
|
69
|
+
return str(resource_type.value)
|
|
70
|
+
return str(resource_type)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def _summarize_node_result(result: Any) -> DbtNodeResult:
|
|
74
|
+
node = getattr(result, "node", None)
|
|
75
|
+
return DbtNodeResult(
|
|
76
|
+
unique_id=str(getattr(result, "unique_id", "") or getattr(node, "unique_id", "")),
|
|
77
|
+
name=_node_name(node),
|
|
78
|
+
resource_type=_node_resource_type(node),
|
|
79
|
+
status=str(getattr(result, "status", "")),
|
|
80
|
+
message=getattr(result, "message", None),
|
|
81
|
+
failures=getattr(result, "failures", None),
|
|
82
|
+
execution_time=getattr(result, "execution_time", None),
|
|
83
|
+
relation_name=getattr(result, "relation_name", None),
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def _is_node_result(result: Any) -> bool:
|
|
88
|
+
return hasattr(result, "node") or hasattr(result, "unique_id") or hasattr(result, "status")
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def _raw_node_results(runner_result: Any) -> list[Any]:
|
|
92
|
+
raw_results = getattr(runner_result, "result", None)
|
|
93
|
+
if raw_results is None:
|
|
94
|
+
return []
|
|
95
|
+
|
|
96
|
+
if hasattr(raw_results, "results"):
|
|
97
|
+
raw_results = raw_results.results
|
|
98
|
+
|
|
99
|
+
if isinstance(raw_results, tuple):
|
|
100
|
+
raw_results = list(raw_results)
|
|
101
|
+
elif not isinstance(raw_results, list):
|
|
102
|
+
raw_results = []
|
|
103
|
+
|
|
104
|
+
return raw_results
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def summarize_dbt_runner_result(runner_result: Any) -> list[DbtNodeResult]:
|
|
108
|
+
return [_summarize_node_result(result) for result in _raw_node_results(runner_result) if _is_node_result(result)]
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
_REPORT_CSS = """
|
|
112
|
+
<style>
|
|
113
|
+
.dbt-report {
|
|
114
|
+
font-family: Inter, -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
|
|
115
|
+
margin: 0 0 1rem;
|
|
116
|
+
}
|
|
117
|
+
.dbt-report h2 {
|
|
118
|
+
font-size: 1rem;
|
|
119
|
+
font-weight: 650;
|
|
120
|
+
margin: 0 0 0.75rem;
|
|
121
|
+
}
|
|
122
|
+
.dbt-report table {
|
|
123
|
+
border-collapse: collapse;
|
|
124
|
+
font-size: 0.875rem;
|
|
125
|
+
width: 100%;
|
|
126
|
+
}
|
|
127
|
+
.dbt-report th,
|
|
128
|
+
.dbt-report td {
|
|
129
|
+
border-bottom: 1px solid #d9dee7;
|
|
130
|
+
padding: 0.5rem 0.625rem;
|
|
131
|
+
text-align: left;
|
|
132
|
+
vertical-align: top;
|
|
133
|
+
}
|
|
134
|
+
.dbt-report th {
|
|
135
|
+
background: #f7f8fa;
|
|
136
|
+
color: #394150;
|
|
137
|
+
font-weight: 650;
|
|
138
|
+
}
|
|
139
|
+
.dbt-report .status-pass,
|
|
140
|
+
.dbt-report .status-success {
|
|
141
|
+
color: #087443;
|
|
142
|
+
font-weight: 650;
|
|
143
|
+
}
|
|
144
|
+
.dbt-report .status-fail,
|
|
145
|
+
.dbt-report .status-error {
|
|
146
|
+
color: #b42318;
|
|
147
|
+
font-weight: 650;
|
|
148
|
+
}
|
|
149
|
+
.dbt-report .muted {
|
|
150
|
+
color: #697386;
|
|
151
|
+
}
|
|
152
|
+
</style>
|
|
153
|
+
""".strip()
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def _html_cell(value: Any) -> str:
|
|
157
|
+
if value is None or value == "":
|
|
158
|
+
return '<span class="muted">-</span>'
|
|
159
|
+
return html.escape(str(value))
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def dbt_results_to_html(results: list[DbtNodeResult]) -> str:
|
|
163
|
+
rows = []
|
|
164
|
+
for result in results:
|
|
165
|
+
status = result.status.lower()
|
|
166
|
+
status_class = "status-" + "".join(ch if ch.isalnum() else "-" for ch in status)
|
|
167
|
+
rows.append(
|
|
168
|
+
"<tr>"
|
|
169
|
+
f"<td>{_html_cell(result.name or result.unique_id)}</td>"
|
|
170
|
+
f"<td>{_html_cell(result.resource_type)}</td>"
|
|
171
|
+
f'<td class="{html.escape(status_class)}">{_html_cell(result.status)}</td>'
|
|
172
|
+
f"<td>{_html_cell(result.failures)}</td>"
|
|
173
|
+
f"<td>{_html_cell(result.execution_time)}</td>"
|
|
174
|
+
f"<td>{_html_cell(result.message)}</td>"
|
|
175
|
+
"</tr>"
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
body = (
|
|
179
|
+
"<p>No dbt node results were returned for this command.</p>"
|
|
180
|
+
if not rows
|
|
181
|
+
else (
|
|
182
|
+
"<table>"
|
|
183
|
+
"<thead><tr>"
|
|
184
|
+
"<th>Node</th><th>Type</th><th>Status</th><th>Failures</th><th>Execution Time</th><th>Message</th>"
|
|
185
|
+
"</tr></thead>"
|
|
186
|
+
f"<tbody>{''.join(rows)}</tbody>"
|
|
187
|
+
"</table>"
|
|
188
|
+
)
|
|
189
|
+
)
|
|
190
|
+
return f'{_REPORT_CSS}<section class="dbt-report"><h2>dbt node results</h2>{body}</section>'
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def write_dbt_report(results: list[DbtNodeResult]) -> None:
|
|
194
|
+
try:
|
|
195
|
+
import flyte.report
|
|
196
|
+
|
|
197
|
+
flyte.report.get_tab("dbt").replace(dbt_results_to_html(results))
|
|
198
|
+
flyte.report.flush()
|
|
199
|
+
except Exception:
|
|
200
|
+
from flyte._logging import logger
|
|
201
|
+
|
|
202
|
+
logger.debug("Failed to write dbt report.", exc_info=True)
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
def callback_import_path(callback: DbtEventCallback, source_dir: pathlib.Path | None = None) -> str:
|
|
206
|
+
name = getattr(callback, "__name__", None)
|
|
207
|
+
qualname = getattr(callback, "__qualname__", None)
|
|
208
|
+
if not qualname or "<locals>" in qualname or name == "<lambda>":
|
|
209
|
+
raise ValueError("dbt callbacks used in DbtTask must be importable functions when running remotely.")
|
|
210
|
+
|
|
211
|
+
if source_dir is None:
|
|
212
|
+
module = getattr(callback, "__module__", None)
|
|
213
|
+
if not module:
|
|
214
|
+
raise ValueError("dbt callbacks used in DbtTask must be importable functions when running remotely.")
|
|
215
|
+
else:
|
|
216
|
+
from flyte._module import extract_obj_module
|
|
217
|
+
|
|
218
|
+
module, _ = extract_obj_module(callback, source_dir=source_dir)
|
|
219
|
+
|
|
220
|
+
import_path = f"{module}:{qualname}"
|
|
221
|
+
if callback.__module__ == "__main__" and source_dir is not None:
|
|
222
|
+
# The workflow file is imported under its file name in the container.
|
|
223
|
+
# Importing it here would load a second copy, so identity cannot match.
|
|
224
|
+
return import_path
|
|
225
|
+
try:
|
|
226
|
+
resolved_callback = import_callback(import_path)
|
|
227
|
+
except (AttributeError, ModuleNotFoundError):
|
|
228
|
+
if source_dir is not None:
|
|
229
|
+
return import_path
|
|
230
|
+
raise
|
|
231
|
+
if resolved_callback is not callback:
|
|
232
|
+
raise ValueError(f"dbt callback {import_path!r} must resolve to the original callback object when imported.")
|
|
233
|
+
return import_path
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def _import_dotted_path(import_path: str) -> Any:
|
|
237
|
+
parts = import_path.split(".")
|
|
238
|
+
for module_end in range(len(parts), 0, -1):
|
|
239
|
+
module_name = ".".join(parts[:module_end])
|
|
240
|
+
try:
|
|
241
|
+
value: Any = importlib.import_module(module_name)
|
|
242
|
+
except ModuleNotFoundError:
|
|
243
|
+
continue
|
|
244
|
+
for attr in parts[module_end:]:
|
|
245
|
+
value = getattr(value, attr)
|
|
246
|
+
return value
|
|
247
|
+
raise ModuleNotFoundError(f"No module found in dbt callback import path: {import_path!r}")
|
|
248
|
+
|
|
249
|
+
|
|
250
|
+
def import_callback(import_path: str) -> DbtEventCallback:
|
|
251
|
+
if ":" in import_path:
|
|
252
|
+
module_name, attr_path = import_path.split(":", 1)
|
|
253
|
+
if not module_name or not attr_path:
|
|
254
|
+
raise ValueError(f"Invalid dbt callback import path: {import_path!r}")
|
|
255
|
+
|
|
256
|
+
value: Any = importlib.import_module(module_name)
|
|
257
|
+
for attr in attr_path.split("."):
|
|
258
|
+
value = getattr(value, attr)
|
|
259
|
+
else:
|
|
260
|
+
value = _import_dotted_path(import_path)
|
|
261
|
+
|
|
262
|
+
if not callable(value):
|
|
263
|
+
raise TypeError(f"dbt callback import path does not resolve to a callable: {import_path!r}")
|
|
264
|
+
return value
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def resolve_callbacks(callbacks: Sequence[DbtEventCallback | str] | None) -> list[DbtEventCallback]:
|
|
268
|
+
resolved = []
|
|
269
|
+
for callback in callbacks or []:
|
|
270
|
+
if isinstance(callback, str):
|
|
271
|
+
resolved.append(import_callback(callback))
|
|
272
|
+
else:
|
|
273
|
+
resolved.append(callback)
|
|
274
|
+
return resolved
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def callback_import_paths(
|
|
278
|
+
callbacks: Sequence[DbtEventCallback | str] | None,
|
|
279
|
+
source_dir: pathlib.Path | None = None,
|
|
280
|
+
) -> list[str]:
|
|
281
|
+
paths = []
|
|
282
|
+
for callback in callbacks or []:
|
|
283
|
+
if isinstance(callback, str):
|
|
284
|
+
import_callback(callback)
|
|
285
|
+
paths.append(callback)
|
|
286
|
+
else:
|
|
287
|
+
paths.append(callback_import_path(callback, source_dir=source_dir))
|
|
288
|
+
return paths
|
|
289
|
+
|
|
290
|
+
|
|
291
|
+
def invoke_dbt(
|
|
292
|
+
cli_args: list[str],
|
|
293
|
+
callbacks: Sequence[DbtEventCallback | str] | None = None,
|
|
294
|
+
) -> list[DbtNodeResult]:
|
|
295
|
+
"""Run exactly one dbtRunner invocation and return a serializable summary."""
|
|
296
|
+
from dbt.cli.main import dbtRunner
|
|
297
|
+
|
|
298
|
+
args = list(cli_args)
|
|
299
|
+
event_callbacks = resolve_callbacks(callbacks)
|
|
300
|
+
|
|
301
|
+
runner_result = dbtRunner(callbacks=event_callbacks).invoke(args)
|
|
302
|
+
results = summarize_dbt_runner_result(runner_result)
|
|
303
|
+
write_dbt_report(results)
|
|
304
|
+
if not runner_result.success:
|
|
305
|
+
if runner_result.exception is not None:
|
|
306
|
+
raise runner_result.exception
|
|
307
|
+
raise DbtInvocationError(args, results)
|
|
308
|
+
return results
|
flyteplugins/dbt/task.py
ADDED
|
@@ -0,0 +1,251 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import inspect
|
|
4
|
+
import shlex
|
|
5
|
+
import weakref
|
|
6
|
+
from dataclasses import dataclass, field
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import TYPE_CHECKING, Any, Optional
|
|
9
|
+
|
|
10
|
+
from flyte.extend import RuntimeTaskTemplate
|
|
11
|
+
from flyte.models import NativeInterface
|
|
12
|
+
|
|
13
|
+
from flyteplugins.dbt.runner import (
|
|
14
|
+
DbtEventCallback,
|
|
15
|
+
DbtNodeResult,
|
|
16
|
+
invoke_dbt,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
if TYPE_CHECKING:
|
|
20
|
+
from flyte import CacheRequest, TaskEnvironment
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
_MANAGED_DBT_FLAGS = {
|
|
24
|
+
"--project-dir",
|
|
25
|
+
"--profiles-dir",
|
|
26
|
+
"--profile",
|
|
27
|
+
"--target",
|
|
28
|
+
"--target-path",
|
|
29
|
+
"--select",
|
|
30
|
+
"--exclude",
|
|
31
|
+
}
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def _validate_cache(cache: CacheRequest) -> CacheRequest:
|
|
35
|
+
from flyte._cache.cache import cache_from_request
|
|
36
|
+
|
|
37
|
+
if cache_from_request(cache).behavior != "disable":
|
|
38
|
+
raise ValueError("DbtTask does not support caching yet. Set cache='disable' or remove cache configuration.")
|
|
39
|
+
return cache
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _validate_project_dir(project_dir: str | None) -> None:
|
|
43
|
+
if project_dir is None:
|
|
44
|
+
return
|
|
45
|
+
if not (Path(project_dir) / "dbt_project.yml").exists():
|
|
46
|
+
raise ValueError(f"dbt project_dir {project_dir!r} must contain a dbt_project.yml file.")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _validate_extra_args(extra_args: list[str] | None) -> list[str]:
|
|
50
|
+
args = list(extra_args or [])
|
|
51
|
+
managed_flags = sorted(set(args) & _MANAGED_DBT_FLAGS)
|
|
52
|
+
if managed_flags:
|
|
53
|
+
raise ValueError("dbt extra_args cannot include flags managed by DbtTask: " + ", ".join(managed_flags))
|
|
54
|
+
return args
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _parse_command(command: str | list[str]) -> list[str]:
|
|
58
|
+
if isinstance(command, str):
|
|
59
|
+
args = shlex.split(command)
|
|
60
|
+
else:
|
|
61
|
+
args = list(command)
|
|
62
|
+
if not args:
|
|
63
|
+
raise ValueError("dbt command must include at least one argument.")
|
|
64
|
+
return args
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _build_cli_args(
|
|
68
|
+
*,
|
|
69
|
+
command: str | list[str],
|
|
70
|
+
project_dir: str | None,
|
|
71
|
+
profiles_dir: str | None,
|
|
72
|
+
profile: str | None,
|
|
73
|
+
target_path: str | None,
|
|
74
|
+
select: list[str] | None,
|
|
75
|
+
exclude: list[str] | None,
|
|
76
|
+
target: str | None,
|
|
77
|
+
extra_args: list[str] | None,
|
|
78
|
+
) -> list[str]:
|
|
79
|
+
_validate_project_dir(project_dir)
|
|
80
|
+
args = _parse_command(command)
|
|
81
|
+
if project_dir:
|
|
82
|
+
args.extend(["--project-dir", project_dir])
|
|
83
|
+
if profiles_dir:
|
|
84
|
+
args.extend(["--profiles-dir", profiles_dir])
|
|
85
|
+
if profile:
|
|
86
|
+
args.extend(["--profile", profile])
|
|
87
|
+
if target:
|
|
88
|
+
args.extend(["--target", target])
|
|
89
|
+
if target_path:
|
|
90
|
+
args.extend(["--target-path", target_path])
|
|
91
|
+
if select:
|
|
92
|
+
args.extend(["--select", *select])
|
|
93
|
+
if exclude:
|
|
94
|
+
args.extend(["--exclude", *exclude])
|
|
95
|
+
args.extend(_validate_extra_args(extra_args))
|
|
96
|
+
return args
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
@dataclass(kw_only=True)
|
|
100
|
+
class DbtTask(RuntimeTaskTemplate):
|
|
101
|
+
"""A Flyte task that maps one dbtRunner.invoke(...) call to one task."""
|
|
102
|
+
|
|
103
|
+
project_dir: str | None = None
|
|
104
|
+
profiles_dir: str | None = None
|
|
105
|
+
profile: str | None = None
|
|
106
|
+
target_path: str | None = None
|
|
107
|
+
callbacks: list[DbtEventCallback | str] = field(default_factory=list)
|
|
108
|
+
|
|
109
|
+
def __init__(
|
|
110
|
+
self,
|
|
111
|
+
*,
|
|
112
|
+
name: str,
|
|
113
|
+
task_environment: Optional[TaskEnvironment] = None,
|
|
114
|
+
project_dir: str | None = None,
|
|
115
|
+
profiles_dir: str | None = None,
|
|
116
|
+
profile: str | None = None,
|
|
117
|
+
target_path: str | None = None,
|
|
118
|
+
callbacks: list[DbtEventCallback | str] | None = None,
|
|
119
|
+
**kwargs: Any,
|
|
120
|
+
):
|
|
121
|
+
project_dir = kwargs.pop("project_dir", project_dir)
|
|
122
|
+
profiles_dir = kwargs.pop("profiles_dir", profiles_dir)
|
|
123
|
+
profile = kwargs.pop("profile", profile)
|
|
124
|
+
target_path = kwargs.pop("target_path", target_path)
|
|
125
|
+
self.project_dir = project_dir
|
|
126
|
+
self.profiles_dir = profiles_dir
|
|
127
|
+
self.profile = profile
|
|
128
|
+
self.target_path = target_path
|
|
129
|
+
self.callbacks = list(callbacks or [])
|
|
130
|
+
from flyteplugins.dbt.resolver import DbtTaskResolver
|
|
131
|
+
|
|
132
|
+
task_name = f"{task_environment.name}.{name}" if task_environment else name
|
|
133
|
+
interface = kwargs.pop(
|
|
134
|
+
"interface",
|
|
135
|
+
NativeInterface(
|
|
136
|
+
inputs={
|
|
137
|
+
"command": (str | list[str], inspect.Parameter.empty),
|
|
138
|
+
"select": (Optional[list[str]], None),
|
|
139
|
+
"exclude": (Optional[list[str]], None),
|
|
140
|
+
"target": (Optional[str], None),
|
|
141
|
+
"extra_args": (Optional[list[str]], None),
|
|
142
|
+
},
|
|
143
|
+
outputs={"results": list[DbtNodeResult]},
|
|
144
|
+
),
|
|
145
|
+
)
|
|
146
|
+
parent_env = kwargs.pop("parent_env", weakref.ref(task_environment) if task_environment else None)
|
|
147
|
+
parent_env_name = kwargs.pop("parent_env_name", task_environment.name if task_environment else None)
|
|
148
|
+
cache = kwargs.pop("cache", task_environment.cache if task_environment else "disable")
|
|
149
|
+
cache = _validate_cache(cache)
|
|
150
|
+
|
|
151
|
+
super().__init__(
|
|
152
|
+
name=task_name,
|
|
153
|
+
interface=interface,
|
|
154
|
+
image=kwargs.pop("image", task_environment.image if task_environment else "auto"),
|
|
155
|
+
resources=kwargs.pop("resources", task_environment.resources if task_environment else None),
|
|
156
|
+
cache=cache,
|
|
157
|
+
reusable=kwargs.pop("reusable", task_environment.reusable if task_environment else None),
|
|
158
|
+
env_vars=kwargs.pop("env_vars", task_environment.env_vars if task_environment else None),
|
|
159
|
+
secrets=kwargs.pop("secrets", task_environment.secrets if task_environment else None),
|
|
160
|
+
service_account=kwargs.pop(
|
|
161
|
+
"service_account", task_environment.service_account if task_environment else None
|
|
162
|
+
),
|
|
163
|
+
pod_template=kwargs.pop("pod_template", task_environment.pod_template if task_environment else None),
|
|
164
|
+
report=kwargs.pop("report", False),
|
|
165
|
+
queue=kwargs.pop("queue", task_environment.queue if task_environment else None),
|
|
166
|
+
interruptible=kwargs.pop("interruptible", task_environment.interruptible if task_environment else False),
|
|
167
|
+
short_name=kwargs.pop("short_name", name if task_environment else ""),
|
|
168
|
+
task_type=kwargs.pop("task_type", "dbt"),
|
|
169
|
+
_call_as_synchronous=kwargs.pop("_call_as_synchronous", True),
|
|
170
|
+
parent_env=parent_env,
|
|
171
|
+
parent_env_name=parent_env_name,
|
|
172
|
+
task_resolver=kwargs.pop("task_resolver", DbtTaskResolver()),
|
|
173
|
+
**kwargs,
|
|
174
|
+
)
|
|
175
|
+
|
|
176
|
+
if task_environment is not None:
|
|
177
|
+
task_environment._tasks[task_name] = self
|
|
178
|
+
|
|
179
|
+
def forward(self, *args: Any, **kwargs: Any) -> list[DbtNodeResult]:
|
|
180
|
+
kwargs = self.interface.convert_to_kwargs(*args, **kwargs)
|
|
181
|
+
cli_args = _build_cli_args(
|
|
182
|
+
command=kwargs["command"],
|
|
183
|
+
project_dir=self.project_dir,
|
|
184
|
+
profiles_dir=self.profiles_dir,
|
|
185
|
+
profile=self.profile,
|
|
186
|
+
target_path=self.target_path,
|
|
187
|
+
select=kwargs.get("select"),
|
|
188
|
+
exclude=kwargs.get("exclude"),
|
|
189
|
+
target=kwargs.get("target"),
|
|
190
|
+
extra_args=kwargs.get("extra_args"),
|
|
191
|
+
)
|
|
192
|
+
return invoke_dbt(
|
|
193
|
+
cli_args,
|
|
194
|
+
callbacks=self.callbacks,
|
|
195
|
+
)
|
|
196
|
+
|
|
197
|
+
async def aio(
|
|
198
|
+
self,
|
|
199
|
+
command: str | list[str],
|
|
200
|
+
*,
|
|
201
|
+
select: list[str] | None = None,
|
|
202
|
+
exclude: list[str] | None = None,
|
|
203
|
+
target: str | None = None,
|
|
204
|
+
extra_args: list[str] | None = None,
|
|
205
|
+
) -> list[DbtNodeResult]:
|
|
206
|
+
return await super().aio(
|
|
207
|
+
command=command,
|
|
208
|
+
select=select,
|
|
209
|
+
exclude=exclude,
|
|
210
|
+
target=target,
|
|
211
|
+
extra_args=extra_args,
|
|
212
|
+
)
|
|
213
|
+
|
|
214
|
+
def __call__(
|
|
215
|
+
self,
|
|
216
|
+
command: str | list[str],
|
|
217
|
+
*,
|
|
218
|
+
select: list[str] | None = None,
|
|
219
|
+
exclude: list[str] | None = None,
|
|
220
|
+
target: str | None = None,
|
|
221
|
+
extra_args: list[str] | None = None,
|
|
222
|
+
) -> list[DbtNodeResult]:
|
|
223
|
+
return super().__call__(
|
|
224
|
+
command=command,
|
|
225
|
+
select=select,
|
|
226
|
+
exclude=exclude,
|
|
227
|
+
target=target,
|
|
228
|
+
extra_args=extra_args,
|
|
229
|
+
)
|
|
230
|
+
|
|
231
|
+
async def execute(self, *args: Any, **kwargs: Any) -> list[DbtNodeResult]:
|
|
232
|
+
kwargs = self.interface.convert_to_kwargs(*args, **kwargs)
|
|
233
|
+
cli_args = _build_cli_args(
|
|
234
|
+
command=kwargs["command"],
|
|
235
|
+
project_dir=self.project_dir,
|
|
236
|
+
profiles_dir=self.profiles_dir,
|
|
237
|
+
profile=self.profile,
|
|
238
|
+
target_path=self.target_path,
|
|
239
|
+
select=kwargs.get("select"),
|
|
240
|
+
exclude=kwargs.get("exclude"),
|
|
241
|
+
target=kwargs.get("target"),
|
|
242
|
+
extra_args=kwargs.get("extra_args"),
|
|
243
|
+
)
|
|
244
|
+
|
|
245
|
+
from flyte._utils.asyncify import run_sync_in_thread
|
|
246
|
+
|
|
247
|
+
return await run_sync_in_thread(
|
|
248
|
+
invoke_dbt,
|
|
249
|
+
cli_args,
|
|
250
|
+
self.callbacks,
|
|
251
|
+
)
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: flyteplugins-dbt
|
|
3
|
+
Version: 0.0.0a0
|
|
4
|
+
Summary: dbt task plugin for flyte
|
|
5
|
+
Author-email: Jan Fiedler <jan@union.ai>
|
|
6
|
+
Requires-Python: >=3.10
|
|
7
|
+
Description-Content-Type: text/markdown
|
|
8
|
+
Requires-Dist: dbt-core
|
|
9
|
+
Requires-Dist: flyte
|
|
10
|
+
|
|
11
|
+
# flyteplugins-dbt
|
|
12
|
+
|
|
13
|
+
Run dbt CLI invocations as Flyte v2 tasks.
|
|
14
|
+
|
|
15
|
+
`DbtTask` maps one `dbtRunner.invoke(...)` call to one Flyte task. Project-level dbt paths are configured on the task, and invocation-level options are task inputs.
|
|
16
|
+
|
|
17
|
+
Install the dbt adapter required by your project, such as `dbt-duckdb`, `dbt-bigquery`, or `dbt-snowflake`, in the task image alongside this plugin.
|
|
18
|
+
|
|
19
|
+
```python
|
|
20
|
+
from pathlib import Path
|
|
21
|
+
|
|
22
|
+
import flyte
|
|
23
|
+
from flyteplugins.dbt import DbtTask
|
|
24
|
+
|
|
25
|
+
DBT_PROJECT_DIR = "jaffle_shop"
|
|
26
|
+
DBT_PROFILES_DIR = "dbt-profiles"
|
|
27
|
+
|
|
28
|
+
env = flyte.TaskEnvironment(
|
|
29
|
+
name="dbt",
|
|
30
|
+
image=flyte.Image.from_debian_base()
|
|
31
|
+
.with_requirements("requirements.txt")
|
|
32
|
+
.with_source_folder(Path(DBT_PROJECT_DIR))
|
|
33
|
+
.with_source_folder(Path(DBT_PROFILES_DIR)),
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
dbt_test = DbtTask(
|
|
37
|
+
name="dbt-test",
|
|
38
|
+
task_environment=env,
|
|
39
|
+
project_dir=DBT_PROJECT_DIR,
|
|
40
|
+
profiles_dir=DBT_PROFILES_DIR,
|
|
41
|
+
profile=DBT_PROJECT_DIR,
|
|
42
|
+
report=True,
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@env.task
|
|
47
|
+
async def main():
|
|
48
|
+
return await dbt_test.aio(command="test", select=["stg_orders+"], target="prod")
|
|
49
|
+
```
|
|
50
|
+
|
|
51
|
+
Multi-token dbt commands can be passed as a shell-style string or as explicit tokens:
|
|
52
|
+
|
|
53
|
+
```python
|
|
54
|
+
await dbt_test.aio(command="docs generate")
|
|
55
|
+
await dbt_test.aio(command=["source", "freshness"])
|
|
56
|
+
```
|
|
57
|
+
|
|
58
|
+
The task returns a list of `DbtNodeResult` values summarized from dbt node results. If dbt reports failure, the task raises the dbt exception when one is available, otherwise it raises `DbtInvocationError` with the summarized node results.
|
|
59
|
+
|
|
60
|
+
When `report=True`, the task writes a dbt report tab with node names, resource types, statuses, failures, execution times, and messages. The report is written before dbt failures are raised, so failed dbt runs can still show succeeded and failed node results.
|
|
61
|
+
|
|
62
|
+
## Event callbacks
|
|
63
|
+
|
|
64
|
+
Custom dbt event callbacks can be passed to `DbtTask`.
|
|
65
|
+
|
|
66
|
+
```python
|
|
67
|
+
def log_dbt_node(event):
|
|
68
|
+
data = getattr(event, "data", None)
|
|
69
|
+
node_info = getattr(data, "node_info", None)
|
|
70
|
+
if node_info is None:
|
|
71
|
+
return
|
|
72
|
+
|
|
73
|
+
node_name = getattr(node_info, "node_name", None)
|
|
74
|
+
node_status = getattr(node_info, "node_status", None)
|
|
75
|
+
print(f"dbt node={node_name} status={node_status}")
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
dbt_test = DbtTask(
|
|
79
|
+
name="dbt-test",
|
|
80
|
+
task_environment=env,
|
|
81
|
+
callbacks=[log_dbt_node],
|
|
82
|
+
)
|
|
83
|
+
```
|
|
84
|
+
|
|
85
|
+
For remote execution, callbacks must be importable functions. Import-path strings are also supported:
|
|
86
|
+
|
|
87
|
+
```python
|
|
88
|
+
dbt_test = DbtTask(
|
|
89
|
+
name="dbt-test",
|
|
90
|
+
task_environment=env,
|
|
91
|
+
callbacks=["my_project.callbacks.log_dbt_node"],
|
|
92
|
+
)
|
|
93
|
+
```
|
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
flyteplugins/dbt/__init__.py,sha256=xNbFKiTelhEPnkilDGSHmCbacrUHQpRRgZHXPiUGqZw,338
|
|
2
|
+
flyteplugins/dbt/resolver.py,sha256=x7fQUQF0yVE_o9qxSESIO3dctzj04_s2ndTXpEjAPHA,1963
|
|
3
|
+
flyteplugins/dbt/runner.py,sha256=-e-iVP-7sMSCvB9NsQTbT__MC_VysfNrWSW2w98nzko,10125
|
|
4
|
+
flyteplugins/dbt/task.py,sha256=_huzdAx8NHKD2twmKfLBVIijgDVbHPlz2i21GZeQTmA,8822
|
|
5
|
+
flyteplugins_dbt-0.0.0a0.dist-info/METADATA,sha256=7T2LvdlxTSGIj1HkkUeuSibEDjcdczggxQ8Nnbf0Aeo,2768
|
|
6
|
+
flyteplugins_dbt-0.0.0a0.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
|
|
7
|
+
flyteplugins_dbt-0.0.0a0.dist-info/top_level.txt,sha256=cgd779rPu9EsvdtuYgUxNHHgElaQvPn74KhB5XSeMBE,13
|
|
8
|
+
flyteplugins_dbt-0.0.0a0.dist-info/RECORD,,
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
flyteplugins
|