Merge pull request #4301 from bnovik0v/fix-4300-missing-tool-lifecycle
Fail missing tool calls cleanly
This commit is contained in:
1
changelog/4301.fixed.md
Normal file
1
changelog/4301.fixed.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Fixed missing tool handlers so unregistered tool calls fail with a normal final tool result instead of leaving tool-call state hanging.
|
||||||
@@ -737,7 +737,7 @@ class LLMService(UserTurnCompletionLLMServiceMixin, AIService):
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
f"{self} is calling '{function_call.function_name}', but it's not registered."
|
f"{self} is calling '{function_call.function_name}', but it's not registered."
|
||||||
)
|
)
|
||||||
continue
|
item = self._build_missing_function_call_registry_item(function_call.function_name)
|
||||||
|
|
||||||
runner_items.append(
|
runner_items.append(
|
||||||
FunctionCallRunnerItem(
|
FunctionCallRunnerItem(
|
||||||
@@ -794,12 +794,21 @@ class LLMService(UserTurnCompletionLLMServiceMixin, AIService):
|
|||||||
await self._sequential_runner_queue.put(runner_item)
|
await self._sequential_runner_queue.put(runner_item)
|
||||||
|
|
||||||
async def _run_function_call(self, runner_item: FunctionCallRunnerItem):
|
async def _run_function_call(self, runner_item: FunctionCallRunnerItem):
|
||||||
|
# Re-resolve the registry item at execution time. The function may have
|
||||||
|
# been unregistered between queuing and execution, in which case we
|
||||||
|
# fall back to the missing-function handler so the call still terminates
|
||||||
|
# with a normal tool result.
|
||||||
if runner_item.function_name in self._functions.keys():
|
if runner_item.function_name in self._functions.keys():
|
||||||
item = self._functions[runner_item.function_name]
|
item = self._functions[runner_item.function_name]
|
||||||
elif None in self._functions.keys():
|
elif None in self._functions.keys():
|
||||||
item = self._functions[None]
|
item = self._functions[None]
|
||||||
|
elif runner_item.registry_item.handler == self._missing_function_call_handler:
|
||||||
|
item = runner_item.registry_item
|
||||||
else:
|
else:
|
||||||
return
|
logger.warning(
|
||||||
|
f"{self} is calling '{runner_item.function_name}', but it was just unregistered."
|
||||||
|
)
|
||||||
|
item = self._build_missing_function_call_registry_item(runner_item.function_name)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"{self} Calling function [{runner_item.function_name}:{runner_item.tool_call_id}] with arguments {runner_item.arguments}"
|
f"{self} Calling function [{runner_item.function_name}:{runner_item.tool_call_id}] with arguments {runner_item.arguments}"
|
||||||
@@ -908,6 +917,20 @@ class LLMService(UserTurnCompletionLLMServiceMixin, AIService):
|
|||||||
if timeout_task and not timeout_task.done():
|
if timeout_task and not timeout_task.done():
|
||||||
await self.cancel_task(timeout_task)
|
await self.cancel_task(timeout_task)
|
||||||
|
|
||||||
|
def _build_missing_function_call_registry_item(
|
||||||
|
self, function_name: str
|
||||||
|
) -> FunctionCallRegistryItem:
|
||||||
|
"""Build a registry item that routes to the missing-function handler."""
|
||||||
|
return FunctionCallRegistryItem(
|
||||||
|
function_name=function_name,
|
||||||
|
handler=self._missing_function_call_handler,
|
||||||
|
cancel_on_interruption=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _missing_function_call_handler(self, params: FunctionCallParams):
|
||||||
|
"""Return a terminal tool result when the LLM calls an unknown function."""
|
||||||
|
await params.result_callback(f"Error: function '{params.function_name}' is not registered.")
|
||||||
|
|
||||||
def _has_async_tools(self) -> bool:
|
def _has_async_tools(self) -> bool:
|
||||||
"""Return True if at least one non-builtin async tool is registered."""
|
"""Return True if at least one non-builtin async tool is registered."""
|
||||||
return any(
|
return any(
|
||||||
|
|||||||
173
tests/test_llm_service.py
Normal file
173
tests/test_llm_service.py
Normal file
@@ -0,0 +1,173 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024-2026, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
from pipecat.frames.frames import (
|
||||||
|
FunctionCallFromLLM,
|
||||||
|
FunctionCallInProgressFrame,
|
||||||
|
FunctionCallResultFrame,
|
||||||
|
FunctionCallsStartedFrame,
|
||||||
|
)
|
||||||
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
|
from pipecat.services.llm_service import LLMService
|
||||||
|
from pipecat.services.settings import LLMSettings
|
||||||
|
from pipecat.turns.user_mute.function_call_user_mute_strategy import FunctionCallUserMuteStrategy
|
||||||
|
|
||||||
|
|
||||||
|
class MockLLMService(LLMService):
|
||||||
|
"""Minimal LLM service for testing function call execution."""
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
settings = LLMSettings(
|
||||||
|
model="test-model",
|
||||||
|
system_instruction=None,
|
||||||
|
temperature=None,
|
||||||
|
max_tokens=None,
|
||||||
|
top_p=None,
|
||||||
|
top_k=None,
|
||||||
|
frequency_penalty=None,
|
||||||
|
presence_penalty=None,
|
||||||
|
seed=None,
|
||||||
|
filter_incomplete_user_turns=None,
|
||||||
|
user_turn_completion_config=None,
|
||||||
|
)
|
||||||
|
super().__init__(settings=settings, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class TestLLMService(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def _run_function_calls_inline(self, service: MockLLMService):
|
||||||
|
async def run_inline(runner_items):
|
||||||
|
for runner_item in runner_items:
|
||||||
|
await service._run_function_call(runner_item)
|
||||||
|
|
||||||
|
service._run_parallel_function_calls = run_inline
|
||||||
|
service._run_sequential_function_calls = run_inline
|
||||||
|
|
||||||
|
async def test_missing_function_call_emits_terminal_result(self):
|
||||||
|
service = MockLLMService()
|
||||||
|
service._call_event_handler = AsyncMock()
|
||||||
|
await self._run_function_calls_inline(service)
|
||||||
|
|
||||||
|
recorded_frames = []
|
||||||
|
|
||||||
|
async def mock_broadcast_frame(frame_cls, **kwargs):
|
||||||
|
recorded_frames.append(frame_cls(**kwargs))
|
||||||
|
|
||||||
|
service.broadcast_frame = mock_broadcast_frame
|
||||||
|
|
||||||
|
with patch("pipecat.services.llm_service.logger") as mock_logger:
|
||||||
|
await service.run_function_calls(
|
||||||
|
[
|
||||||
|
FunctionCallFromLLM(
|
||||||
|
function_name="missing_tool",
|
||||||
|
tool_call_id="call_1",
|
||||||
|
arguments={"query": "weather"},
|
||||||
|
context=LLMContext(),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[type(frame) for frame in recorded_frames],
|
||||||
|
[
|
||||||
|
FunctionCallsStartedFrame,
|
||||||
|
FunctionCallInProgressFrame,
|
||||||
|
FunctionCallResultFrame,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertEqual(recorded_frames[1].function_name, "missing_tool")
|
||||||
|
self.assertEqual(
|
||||||
|
recorded_frames[2].result,
|
||||||
|
"Error: function 'missing_tool' is not registered.",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Only the queue-time warning should fire; the execution-time
|
||||||
|
# "just unregistered" warning must not double-log.
|
||||||
|
warnings = [c.args[0] for c in mock_logger.warning.call_args_list]
|
||||||
|
self.assertTrue(any("not registered" in w for w in warnings))
|
||||||
|
self.assertFalse(any("just unregistered" in w for w in warnings))
|
||||||
|
|
||||||
|
async def test_function_unregistered_between_queue_and_execute(self):
|
||||||
|
"""Function unregistered between queuing and execution still terminates."""
|
||||||
|
service = MockLLMService()
|
||||||
|
service._call_event_handler = AsyncMock()
|
||||||
|
|
||||||
|
async def real_handler(params):
|
||||||
|
await params.result_callback("should not be called")
|
||||||
|
|
||||||
|
service.register_function("doomed_tool", real_handler)
|
||||||
|
|
||||||
|
recorded_frames = []
|
||||||
|
|
||||||
|
async def mock_broadcast_frame(frame_cls, **kwargs):
|
||||||
|
recorded_frames.append(frame_cls(**kwargs))
|
||||||
|
|
||||||
|
service.broadcast_frame = mock_broadcast_frame
|
||||||
|
|
||||||
|
async def run_inline(runner_items):
|
||||||
|
# Simulate the function being unregistered after queuing but before execution.
|
||||||
|
service.unregister_function("doomed_tool")
|
||||||
|
for runner_item in runner_items:
|
||||||
|
await service._run_function_call(runner_item)
|
||||||
|
|
||||||
|
service._run_parallel_function_calls = run_inline
|
||||||
|
service._run_sequential_function_calls = run_inline
|
||||||
|
|
||||||
|
await service.run_function_calls(
|
||||||
|
[
|
||||||
|
FunctionCallFromLLM(
|
||||||
|
function_name="doomed_tool",
|
||||||
|
tool_call_id="call_1",
|
||||||
|
arguments={},
|
||||||
|
context=LLMContext(),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[type(frame) for frame in recorded_frames],
|
||||||
|
[
|
||||||
|
FunctionCallsStartedFrame,
|
||||||
|
FunctionCallInProgressFrame,
|
||||||
|
FunctionCallResultFrame,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
recorded_frames[2].result,
|
||||||
|
"Error: function 'doomed_tool' is not registered.",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_missing_function_call_allows_user_mute_cleanup(self):
|
||||||
|
service = MockLLMService()
|
||||||
|
service._call_event_handler = AsyncMock()
|
||||||
|
await self._run_function_calls_inline(service)
|
||||||
|
|
||||||
|
recorded_frames = []
|
||||||
|
|
||||||
|
async def mock_broadcast_frame(frame_cls, **kwargs):
|
||||||
|
recorded_frames.append(frame_cls(**kwargs))
|
||||||
|
|
||||||
|
service.broadcast_frame = mock_broadcast_frame
|
||||||
|
|
||||||
|
await service.run_function_calls(
|
||||||
|
[
|
||||||
|
FunctionCallFromLLM(
|
||||||
|
function_name="missing_tool",
|
||||||
|
tool_call_id="call_1",
|
||||||
|
arguments={},
|
||||||
|
context=LLMContext(),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
strategy = FunctionCallUserMuteStrategy()
|
||||||
|
muted = False
|
||||||
|
for frame in recorded_frames:
|
||||||
|
muted = await strategy.process_frame(frame)
|
||||||
|
|
||||||
|
self.assertFalse(muted)
|
||||||
Reference in New Issue
Block a user