|
4 | 4 | from modelgauge.prompt import TextPrompt |
5 | 5 | from modelgauge.sut import SUTResponse |
6 | 6 |
|
7 | | -from modelplane.evaluator.dag import Composer, FailedDAGOutput |
| 7 | +from modelplane.evaluator.dag import Composer, FailedDAGOutput, NodeExecutionError |
8 | 8 | from modelplane.evaluator.safety import AnnotatorArbiter, Safety, SafetyDAGAnnotator |
9 | 9 | from modelplane.evaluator.verdict import Verdict |
10 | 10 |
|
@@ -50,13 +50,14 @@ def test_safety_dag_with_bad_node(sample_ctx, threshold_arbiter): |
50 | 50 | ) |
51 | 51 | dag_output = dag.run(sample_ctx) |
52 | 52 | assert isinstance(dag_output, FailedDAGOutput) |
53 | | - assert str(dag_output.error) == "I'm afraid I can't do that, Dave." |
| 53 | + assert str(dag_output.error.original_error) == "I'm afraid I can't do that, Dave." |
54 | 54 |
|
55 | 55 | dag_annotator = SafetyDAGAnnotator("safety_annotator", dag) |
56 | 56 | with pytest.raises( |
57 | | - type(dag_output.error), match="I'm afraid I can't do that, Dave." |
58 | | - ): |
| 57 | + NodeExecutionError, match="Error while executing node 'failing_node': I'm afraid I can't do that, Dave." |
| 58 | + ) as e: |
59 | 59 | dag_annotator.process( |
60 | 60 | prompt=TextPrompt(text=sample_ctx.prompt), |
61 | 61 | response=SUTResponse(text=sample_ctx.response), |
62 | 62 | ) |
| 63 | + assert type(e.value.original_error) == type(dag_output.error) |
0 commit comments