Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 17 additions & 6 deletions sqlmesh/core/console.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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 ")
Expand All @@ -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:
Expand Down
2 changes: 1 addition & 1 deletion sqlmesh/core/plan/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion sqlmesh/core/snapshot/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 6 additions & 0 deletions sqlmesh/utils/concurrency.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
13 changes: 13 additions & 0 deletions tests/core/test_console.py
Original file line number Diff line number Diff line change
@@ -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():
Expand Down Expand Up @@ -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()
24 changes: 23 additions & 1 deletion tests/core/test_plan_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
7 changes: 7 additions & 0 deletions tests/utils/test_concurrency.py
Original file line number Diff line number Diff line change
Expand Up @@ -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<model: snapshot>: 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()
Expand Down