Skip to content
Closed
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
11 changes: 4 additions & 7 deletions rdagent/components/coder/CoSTEER/knowledge_management.py
Original file line number Diff line number Diff line change
Expand Up @@ -519,13 +519,10 @@ def analyze_error(
else:
error_list = []
for error_content in error_contents:
for error_node in all_error_nodes:
if error_content == error_node.content:
error_list.append(error_node)
else:
error_list.append(error_content)
if error_list[-1] in error_list[:-1]:
error_list.pop()
matched_node = self.knowledgebase.graph.find_node(content=error_content, label="error")
error_item = error_content if matched_node is None else matched_node
if error_item not in error_list:
error_list.append(error_item)

return error_list

Expand Down
60 changes: 60 additions & 0 deletions test/utils/coder/test_costeer_analyze_error.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
from types import SimpleNamespace

import pytest

from rdagent.components.coder.CoSTEER.knowledge_management import (
CoSTEERRAGStrategyV2,
)
from rdagent.components.knowledge_management.graph import (
UndirectedGraph,
UndirectedNode,
)

ROWS_ERROR = "The source dataframe and the ground truth dataframe have different rows count."
TOLERANCE_ERROR = "Some values differ by more than the tolerance of 1e-6."
EXECUTION_FEEDBACK = 'File "factor.py", line 3, in <module>\n x = 1 / 0\nZeroDivisionError: division by zero'
EXECUTION_ERROR = "ErrorType: ZeroDivisionError\nError line: x = 1 / 0"


def _error(content: str) -> UndirectedNode:
return UndirectedNode(content=content, label="error")


def _strategy(*nodes: UndirectedNode) -> CoSTEERRAGStrategyV2:
graph = UndirectedGraph()
graph.nodes = {node.id: node for node in nodes}
strategy = CoSTEERRAGStrategyV2.__new__(CoSTEERRAGStrategyV2)
strategy.knowledgebase = SimpleNamespace(graph=graph)
return strategy


@pytest.mark.offline
@pytest.mark.parametrize(
("feedback", "feedback_type", "content"),
[
(ROWS_ERROR, "value", ROWS_ERROR),
(EXECUTION_FEEDBACK, "execution", EXECUTION_ERROR),
("Execution timed out after 600 seconds.", "execution", "Undefined Error"),
],
ids=["value", "execution", "undefined"],
)
def test_analyze_error_returns_matched_node_once(feedback: str, feedback_type: str, content: str) -> None:
matched = _error(content)
strategy = _strategy(_error("A different previous error."), matched)
assert strategy.analyze_error(feedback, feedback_type=feedback_type) == [matched]


@pytest.mark.offline
@pytest.mark.parametrize("graph_order", ["parsed", "reversed"])
def test_analyze_error_orders_matched_nodes_by_feedback(graph_order: str) -> None:
rows, tolerance = _error(ROWS_ERROR), _error(TOLERANCE_ERROR)
strategy = _strategy(rows, tolerance) if graph_order == "parsed" else _strategy(tolerance, rows)
feedback = f"{ROWS_ERROR}\n{TOLERANCE_ERROR}\n{ROWS_ERROR}"
assert strategy.analyze_error(feedback, feedback_type="value") == [rows, tolerance]


@pytest.mark.offline
def test_analyze_error_keeps_unmatched_error_in_order() -> None:
rows = _error(ROWS_ERROR)
strategy = _strategy(_error("A different previous error."), rows)
assert strategy.analyze_error(f"{TOLERANCE_ERROR}\n{ROWS_ERROR}", feedback_type="value") == [TOLERANCE_ERROR, rows]
Loading