From ddead8d46e0d4ec88277b61f262e4cfd82c7ea59 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Andreas=20Fredh=C3=B8i?= Date: Mon, 14 Sep 2026 21:54:50 +0200 Subject: [PATCH] fix: include node name in execution error messages MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Signed-off-by: Andreas Fredhøi --- sqlmesh/core/console.py | 23 +++++++++++++++++------ sqlmesh/core/plan/evaluator.py | 2 +- sqlmesh/core/snapshot/evaluator.py | 2 +- sqlmesh/utils/concurrency.py | 6 ++++++ tests/core/test_console.py | 13 +++++++++++++ tests/core/test_plan_evaluator.py | 24 +++++++++++++++++++++++- tests/utils/test_concurrency.py | 7 +++++++ 7 files changed, 68 insertions(+), 9 deletions(-) diff --git a/sqlmesh/core/console.py b/sqlmesh/core/console.py index f9a758b405..d0a39288a2 100644 --- a/sqlmesh/core/console.py +++ b/sqlmesh/core/console.py @@ -3279,7 +3279,8 @@ def log_skipped_models(self, snapshot_names: t.Set[str]) -> None: super().log_skipped_models(snapshot_names) def log_failed_models(self, errors: t.List[NodeExecutionFailedError]) -> None: - self._errors.extend([str(ex) for ex in errors if str(ex) not in self._errors]) + failed_model_errors = [_format_failed_model_error(error) for error in errors] + self._errors.extend(error for error in failed_model_errors if error not in self._errors) super().log_failed_models(errors) 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: def log_failed_models(self, errors: t.List[NodeExecutionFailedError]) -> None: if errors: + failed_model_errors = [_format_failed_model_error(error) for error in errors] + self._errors.extend(error for error in failed_model_errors if error not in self._errors) self._print("**Failed models**") error_messages = _format_node_errors(errors) @@ -4195,11 +4198,7 @@ def _format_node_error(ex: NodeExecutionFailedError) -> str: num_fails = len(errors) for i, error in enumerate(errors): - node_name = "" - if isinstance(error.node, SnapshotId): - node_name = error.node.name - elif hasattr(error.node, "snapshot_name"): - node_name = error.node.snapshot_name + node_name = _node_name(error) msg = _format_node_error(error) msg = " " + msg.replace("\n", "\n ") @@ -4211,6 +4210,18 @@ def _format_node_error(ex: NodeExecutionFailedError) -> str: return error_messages +def _node_name(error: NodeExecutionFailedError) -> str: + if isinstance(error.node, SnapshotId): + return error.node.name + if hasattr(error.node, "snapshot_name"): + return error.node.snapshot_name + return str(error.node) + + +def _format_failed_model_error(error: NodeExecutionFailedError) -> str: + return f"{_node_name(error)}: {error.__cause__ or error}" + + def _format_audits_errors(error: NodeAuditsErrors) -> str: error_messages = [] for err in error.errors: diff --git a/sqlmesh/core/plan/evaluator.py b/sqlmesh/core/plan/evaluator.py index f2f432a97e..94011e862c 100644 --- a/sqlmesh/core/plan/evaluator.py +++ b/sqlmesh/core/plan/evaluator.py @@ -381,7 +381,7 @@ def visit_migrate_schemas_stage( deployability_index=stage.deployability_index, ) except NodeExecutionFailedError as ex: - raise PlanError(str(ex.__cause__) if ex.__cause__ else str(ex)) + raise PlanError(str(ex)) from ex def visit_unpause_stage(self, stage: stages.UnpauseStage, plan: EvaluatablePlan) -> None: self.state_sync.unpause_snapshots(stage.promoted_snapshots, plan.end) diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 11b3fd1f33..bdd99585f9 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -103,7 +103,7 @@ class SnapshotCreationFailedError(SQLMeshError): def __init__( self, errors: t.List[NodeExecutionFailedError[SnapshotId]], skipped: t.List[SnapshotId] ): - messages = "\n\n".join(f"{error}\n {error.__cause__}" for error in errors) + messages = "\n\n".join(str(error) for error in errors) super().__init__(f"Physical table creation failed:\n\n{messages}") self.errors = errors self.skipped = skipped diff --git a/sqlmesh/utils/concurrency.py b/sqlmesh/utils/concurrency.py index c5f78645f6..52eac0147c 100644 --- a/sqlmesh/utils/concurrency.py +++ b/sqlmesh/utils/concurrency.py @@ -17,6 +17,12 @@ def __init__(self, node: H): self.node = node super().__init__(f"Execution failed for node {node}") + def __str__(self) -> str: + message = super().__str__() + if self.__cause__: + return f"{message}: {self.__cause__}" + return message + class ConcurrentDAGExecutor(t.Generic[H]): """Concurrently traverses the given DAG in topological order while applying a function to each node. diff --git a/tests/core/test_console.py b/tests/core/test_console.py index f899713235..a3e36ff3df 100644 --- a/tests/core/test_console.py +++ b/tests/core/test_console.py @@ -1,4 +1,6 @@ from sqlmesh.core.console import MarkdownConsole +from sqlmesh.core.snapshot import SnapshotId +from sqlmesh.utils.concurrency import NodeExecutionFailedError def test_markdown_console_warning_block(): @@ -129,3 +131,14 @@ def test_markdown_console_error_block(): ) assert console.consume_captured_errors() == "" + + +def test_markdown_console_failed_model_includes_node_in_captured_error(): + error = NodeExecutionFailedError(SnapshotId(name="model", identifier="snapshot")) + error.__cause__ = RuntimeError("driver error") + console = MarkdownConsole() + + console.log_failed_models([error]) + + assert "model: driver error" in console.consume_captured_errors() + assert "* `model`" in console.consume_captured_output() diff --git a/tests/core/test_plan_evaluator.py b/tests/core/test_plan_evaluator.py index 575f5ae742..c3a0f21277 100644 --- a/tests/core/test_plan_evaluator.py +++ b/tests/core/test_plan_evaluator.py @@ -10,7 +10,9 @@ PlanBuilder, stages as plan_stages, ) -from sqlmesh.core.snapshot import SnapshotChangeCategory +from sqlmesh.core.snapshot import SnapshotChangeCategory, SnapshotId +from sqlmesh.utils.concurrency import NodeExecutionFailedError +from sqlmesh.utils.errors import PlanError @pytest.fixture @@ -82,3 +84,23 @@ def test_builtin_evaluator_push(sushi_context: Context, make_snapshot): ) assert sushi_context.engine_adapter.table_exists(new_model_snapshot.table_name()) assert sushi_context.engine_adapter.table_exists(new_view_model_snapshot.table_name()) + + +def test_migrate_schema_failure_includes_node_context(mocker: MockerFixture): + error = NodeExecutionFailedError(SnapshotId(name="model", identifier="snapshot")) + error.__cause__ = RuntimeError("driver error") + + snapshot_evaluator = mocker.Mock() + snapshot_evaluator.migrate.side_effect = error + evaluator = BuiltInPlanEvaluator( + state_sync=mocker.Mock(), + snapshot_evaluator=snapshot_evaluator, + create_scheduler=mocker.Mock(), + default_catalog=None, + console=mocker.Mock(), + ) + + with pytest.raises(PlanError, match="model.*driver error") as ex: + evaluator.visit_migrate_schemas_stage(mocker.Mock(), mocker.Mock()) + + assert ex.value.__cause__ is error diff --git a/tests/utils/test_concurrency.py b/tests/utils/test_concurrency.py index 5e1e4326f7..8e2779c3d9 100644 --- a/tests/utils/test_concurrency.py +++ b/tests/utils/test_concurrency.py @@ -66,6 +66,13 @@ def raise_(): ) +def test_node_execution_failed_error_includes_cause(): + error = NodeExecutionFailedError(SnapshotId(name="model", identifier="snapshot")) + error.__cause__ = RuntimeError("driver error") + + assert str(error) == ("Execution failed for node SnapshotId: driver error") + + @pytest.mark.parametrize("tasks_num", [1, 2]) def test_concurrent_apply_to_snapshots_return_failed_skipped(mocker: MockerFixture, tasks_num: int): snapshot_a = mocker.Mock()