Skip to content
Merged
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
41 changes: 41 additions & 0 deletions graphcore/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -348,7 +348,48 @@ def to_return(state: StateT) -> PureFunctionGenerator:
return {}
return to_return

class NoToolCallsError(RuntimeError):
"""The model kept ending its turn without a tool call, past the point of scolding it.

The scolding node exists because a turn with no tool call cannot advance the graph:
it tells the model to call one and routes back for another try. A model that answers
every scolding the same way turns that into a cycle whose only exit is langgraph's
recursion limit, hundreds of paid calls later. One scolding is enough for a model that
is going to recover, so past ``MAX_CONSECUTIVE_NO_TOOL_TURNS`` we stop and say so.
"""


#: Consecutive no-tool-call AI turns tolerated before :class:`NoToolCallsError`.
MAX_CONSECUTIVE_NO_TOOL_TURNS = 3


def _consecutive_no_tool_turns(messages: Iterable[AnyMessage]) -> int:
"""Count the AI turns at the tail that ended without a tool call.

Only the scoldings this node inserts may sit between them; anything else means the
run advanced in between, so the count starts over.
"""
count = 0
for m in reversed(list(messages)):
if isinstance(m, AIMessage):
if m.tool_calls:
break
count += 1
elif isinstance(m, HumanMessage) and getattr(m, "display_tag", None) == "scolding":
continue
else:
break
return count


def _scolding_node(state: MessagesState) -> dict[str, list[BaseMessage]]:
turns = _consecutive_no_tool_turns(state["messages"])
if turns >= MAX_CONSECUTIVE_NO_TOOL_TURNS:
raise NoToolCallsError(
f"The model ended {turns} consecutive turns without a tool call, after being "
"told each time that every turn must end with one. Giving up rather than "
"cycling to the recursion limit."
)
return {"messages": [HumanMessage(
content="Every AI turn must end with at least one tool call. Double check your "
"initial prompt for what tools you should be using. In particular, if you "
Expand Down
72 changes: 72 additions & 0 deletions tests/test_no_tool_turns.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
import pytest

from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, ToolMessage

from graphcore.graph import (
MAX_CONSECUTIVE_NO_TOOL_TURNS,
NoToolCallsError,
_consecutive_no_tool_turns,
_scolding_node,
)


def _no_tool_turn() -> AIMessage:
return AIMessage(content="I am done, thanks.")


def _tool_turn() -> AIMessage:
return AIMessage(
content="",
tool_calls=[{"name": "result", "args": {}, "id": "call-1"}],
)


def _scolding() -> HumanMessage:
return HumanMessage(content="call a tool", display_tag="scolding")


def test_counts_only_the_tail():
messages: list[AnyMessage] = [_tool_turn(), ToolMessage(content="ok", tool_call_id="call-1"), _no_tool_turn()]
assert _consecutive_no_tool_turns(messages) == 1


def test_scoldings_do_not_break_the_run():
messages: list[AnyMessage] = [_no_tool_turn(), _scolding(), _no_tool_turn(), _scolding(), _no_tool_turn()]
assert _consecutive_no_tool_turns(messages) == 3


def test_a_tool_call_resets_the_count():
messages: list[AnyMessage] = [
_no_tool_turn(), _scolding(),
_tool_turn(), ToolMessage(content="ok", tool_call_id="call-1"),
_no_tool_turn(),
]
assert _consecutive_no_tool_turns(messages) == 1


def test_scolds_while_under_the_limit():
messages: list[AnyMessage] = [_no_tool_turn()]
out = _scolding_node({"messages": messages})
assert len(out["messages"]) == 1
assert getattr(out["messages"][0], "display_tag", None) == "scolding"


def _run_of_no_tool_turns(n: int) -> list[AnyMessage]:
"""n no-tool turns, each scolded except the last — the state the node sees."""
messages: list[AnyMessage] = []
for _ in range(n - 1):
messages += [_no_tool_turn(), _scolding()]
messages.append(_no_tool_turn())
return messages


def test_last_turn_before_the_limit_is_still_scolded():
messages = _run_of_no_tool_turns(MAX_CONSECUTIVE_NO_TOOL_TURNS - 1)
out = _scolding_node({"messages": messages})
assert getattr(out["messages"][0], "display_tag", None) == "scolding"


def test_raises_once_scolding_stops_working():
messages = _run_of_no_tool_turns(MAX_CONSECUTIVE_NO_TOOL_TURNS)
with pytest.raises(NoToolCallsError):
_scolding_node({"messages": messages})
Loading