Skip to content

Commit ddead8d

Browse files
committed
fix: include node name in execution error messages
Signed-off-by: Andreas Fredhøi <andreas.fredhoi@fresio.no>
1 parent cc30650 commit ddead8d

7 files changed

Lines changed: 68 additions & 9 deletions

File tree

‎sqlmesh/core/console.py‎

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3279,7 +3279,8 @@ def log_skipped_models(self, snapshot_names: t.Set[str]) -> None:
32793279
super().log_skipped_models(snapshot_names)
32803280

32813281
def log_failed_models(self, errors: t.List[NodeExecutionFailedError]) -> None:
3282-
self._errors.extend([str(ex) for ex in errors if str(ex) not in self._errors])
3282+
failed_model_errors = [_format_failed_model_error(error) for error in errors]
3283+
self._errors.extend(error for error in failed_model_errors if error not in self._errors)
32833284
super().log_failed_models(errors)
32843285

32853286
def _print(self, value: t.Any, **kwargs: t.Any) -> None:
@@ -3624,6 +3625,8 @@ def log_skipped_models(self, snapshot_names: t.Set[str]) -> None:
36243625

36253626
def log_failed_models(self, errors: t.List[NodeExecutionFailedError]) -> None:
36263627
if errors:
3628+
failed_model_errors = [_format_failed_model_error(error) for error in errors]
3629+
self._errors.extend(error for error in failed_model_errors if error not in self._errors)
36273630
self._print("**Failed models**")
36283631

36293632
error_messages = _format_node_errors(errors)
@@ -4195,11 +4198,7 @@ def _format_node_error(ex: NodeExecutionFailedError) -> str:
41954198

41964199
num_fails = len(errors)
41974200
for i, error in enumerate(errors):
4198-
node_name = ""
4199-
if isinstance(error.node, SnapshotId):
4200-
node_name = error.node.name
4201-
elif hasattr(error.node, "snapshot_name"):
4202-
node_name = error.node.snapshot_name
4201+
node_name = _node_name(error)
42034202

42044203
msg = _format_node_error(error)
42054204
msg = " " + msg.replace("\n", "\n ")
@@ -4211,6 +4210,18 @@ def _format_node_error(ex: NodeExecutionFailedError) -> str:
42114210
return error_messages
42124211

42134212

4213+
def _node_name(error: NodeExecutionFailedError) -> str:
4214+
if isinstance(error.node, SnapshotId):
4215+
return error.node.name
4216+
if hasattr(error.node, "snapshot_name"):
4217+
return error.node.snapshot_name
4218+
return str(error.node)
4219+
4220+
4221+
def _format_failed_model_error(error: NodeExecutionFailedError) -> str:
4222+
return f"{_node_name(error)}: {error.__cause__ or error}"
4223+
4224+
42144225
def _format_audits_errors(error: NodeAuditsErrors) -> str:
42154226
error_messages = []
42164227
for err in error.errors:

‎sqlmesh/core/plan/evaluator.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -381,7 +381,7 @@ def visit_migrate_schemas_stage(
381381
deployability_index=stage.deployability_index,
382382
)
383383
except NodeExecutionFailedError as ex:
384-
raise PlanError(str(ex.__cause__) if ex.__cause__ else str(ex))
384+
raise PlanError(str(ex)) from ex
385385

386386
def visit_unpause_stage(self, stage: stages.UnpauseStage, plan: EvaluatablePlan) -> None:
387387
self.state_sync.unpause_snapshots(stage.promoted_snapshots, plan.end)

‎sqlmesh/core/snapshot/evaluator.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -103,7 +103,7 @@ class SnapshotCreationFailedError(SQLMeshError):
103103
def __init__(
104104
self, errors: t.List[NodeExecutionFailedError[SnapshotId]], skipped: t.List[SnapshotId]
105105
):
106-
messages = "\n\n".join(f"{error}\n {error.__cause__}" for error in errors)
106+
messages = "\n\n".join(str(error) for error in errors)
107107
super().__init__(f"Physical table creation failed:\n\n{messages}")
108108
self.errors = errors
109109
self.skipped = skipped

‎sqlmesh/utils/concurrency.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,12 @@ def __init__(self, node: H):
1717
self.node = node
1818
super().__init__(f"Execution failed for node {node}")
1919

20+
def __str__(self) -> str:
21+
message = super().__str__()
22+
if self.__cause__:
23+
return f"{message}: {self.__cause__}"
24+
return message
25+
2026

2127
class ConcurrentDAGExecutor(t.Generic[H]):
2228
"""Concurrently traverses the given DAG in topological order while applying a function to each node.

‎tests/core/test_console.py‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,6 @@
11
from sqlmesh.core.console import MarkdownConsole
2+
from sqlmesh.core.snapshot import SnapshotId
3+
from sqlmesh.utils.concurrency import NodeExecutionFailedError
24

35

46
def test_markdown_console_warning_block():
@@ -129,3 +131,14 @@ def test_markdown_console_error_block():
129131
)
130132

131133
assert console.consume_captured_errors() == ""
134+
135+
136+
def test_markdown_console_failed_model_includes_node_in_captured_error():
137+
error = NodeExecutionFailedError(SnapshotId(name="model", identifier="snapshot"))
138+
error.__cause__ = RuntimeError("driver error")
139+
console = MarkdownConsole()
140+
141+
console.log_failed_models([error])
142+
143+
assert "model: driver error" in console.consume_captured_errors()
144+
assert "* `model`" in console.consume_captured_output()

‎tests/core/test_plan_evaluator.py‎

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,9 @@
1010
PlanBuilder,
1111
stages as plan_stages,
1212
)
13-
from sqlmesh.core.snapshot import SnapshotChangeCategory
13+
from sqlmesh.core.snapshot import SnapshotChangeCategory, SnapshotId
14+
from sqlmesh.utils.concurrency import NodeExecutionFailedError
15+
from sqlmesh.utils.errors import PlanError
1416

1517

1618
@pytest.fixture
@@ -82,3 +84,23 @@ def test_builtin_evaluator_push(sushi_context: Context, make_snapshot):
8284
)
8385
assert sushi_context.engine_adapter.table_exists(new_model_snapshot.table_name())
8486
assert sushi_context.engine_adapter.table_exists(new_view_model_snapshot.table_name())
87+
88+
89+
def test_migrate_schema_failure_includes_node_context(mocker: MockerFixture):
90+
error = NodeExecutionFailedError(SnapshotId(name="model", identifier="snapshot"))
91+
error.__cause__ = RuntimeError("driver error")
92+
93+
snapshot_evaluator = mocker.Mock()
94+
snapshot_evaluator.migrate.side_effect = error
95+
evaluator = BuiltInPlanEvaluator(
96+
state_sync=mocker.Mock(),
97+
snapshot_evaluator=snapshot_evaluator,
98+
create_scheduler=mocker.Mock(),
99+
default_catalog=None,
100+
console=mocker.Mock(),
101+
)
102+
103+
with pytest.raises(PlanError, match="model.*driver error") as ex:
104+
evaluator.visit_migrate_schemas_stage(mocker.Mock(), mocker.Mock())
105+
106+
assert ex.value.__cause__ is error

‎tests/utils/test_concurrency.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,13 @@ def raise_():
6666
)
6767

6868

69+
def test_node_execution_failed_error_includes_cause():
70+
error = NodeExecutionFailedError(SnapshotId(name="model", identifier="snapshot"))
71+
error.__cause__ = RuntimeError("driver error")
72+
73+
assert str(error) == ("Execution failed for node SnapshotId<model: snapshot>: driver error")
74+
75+
6976
@pytest.mark.parametrize("tasks_num", [1, 2])
7077
def test_concurrent_apply_to_snapshots_return_failed_skipped(mocker: MockerFixture, tasks_num: int):
7178
snapshot_a = mocker.Mock()

0 commit comments

Comments
 (0)