skipping provider-specific messages during summarization
This commit is contained in:
@@ -15,7 +15,7 @@ from typing import List, Optional
|
|||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from pipecat.processors.aggregators.llm_context import LLMContext
|
from pipecat.processors.aggregators.llm_context import LLMContext, LLMSpecificMessage
|
||||||
|
|
||||||
# Token estimation constants
|
# Token estimation constants
|
||||||
CHARS_PER_TOKEN = 4 # Industry-standard heuristic: 1 token ≈ 4 characters
|
CHARS_PER_TOKEN = 4 # Industry-standard heuristic: 1 token ≈ 4 characters
|
||||||
@@ -188,6 +188,9 @@ class LLMContextSummarizationUtil:
|
|||||||
total = 0
|
total = 0
|
||||||
|
|
||||||
for message in context.messages:
|
for message in context.messages:
|
||||||
|
if isinstance(message, LLMSpecificMessage):
|
||||||
|
continue
|
||||||
|
|
||||||
# Role and structure overhead
|
# Role and structure overhead
|
||||||
total += TOKEN_OVERHEAD_PER_MESSAGE
|
total += TOKEN_OVERHEAD_PER_MESSAGE
|
||||||
|
|
||||||
@@ -248,6 +251,9 @@ class LLMContextSummarizationUtil:
|
|||||||
|
|
||||||
for i in range(start_idx, len(messages)):
|
for i in range(start_idx, len(messages)):
|
||||||
msg = messages[i]
|
msg = messages[i]
|
||||||
|
if isinstance(msg, LLMSpecificMessage):
|
||||||
|
continue
|
||||||
|
|
||||||
role = msg.get("role")
|
role = msg.get("role")
|
||||||
|
|
||||||
# Check for tool calls in assistant messages
|
# Check for tool calls in assistant messages
|
||||||
@@ -298,7 +304,12 @@ class LLMContextSummarizationUtil:
|
|||||||
|
|
||||||
# Find first system message index
|
# Find first system message index
|
||||||
first_system_index = next(
|
first_system_index = next(
|
||||||
(i for i, msg in enumerate(messages) if msg.get("role") == "system"), -1
|
(
|
||||||
|
i
|
||||||
|
for i, msg in enumerate(messages)
|
||||||
|
if not isinstance(msg, LLMSpecificMessage) and msg.get("role") == "system"
|
||||||
|
),
|
||||||
|
-1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Messages to summarize are between first system and recent messages
|
# Messages to summarize are between first system and recent messages
|
||||||
@@ -356,6 +367,9 @@ class LLMContextSummarizationUtil:
|
|||||||
transcript_parts = []
|
transcript_parts = []
|
||||||
|
|
||||||
for msg in messages:
|
for msg in messages:
|
||||||
|
if isinstance(msg, LLMSpecificMessage):
|
||||||
|
continue
|
||||||
|
|
||||||
role = msg.get("role", "unknown")
|
role = msg.get("role", "unknown")
|
||||||
content = msg.get("content", "")
|
content = msg.get("content", "")
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ import unittest
|
|||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
from pipecat.frames.frames import LLMContextSummaryRequestFrame
|
from pipecat.frames.frames import LLMContextSummaryRequestFrame
|
||||||
from pipecat.processors.aggregators.llm_context import LLMContext
|
from pipecat.processors.aggregators.llm_context import LLMContext, LLMSpecificMessage
|
||||||
from pipecat.services.llm_service import LLMService
|
from pipecat.services.llm_service import LLMService
|
||||||
from pipecat.utils.context.llm_context_summarization import (
|
from pipecat.utils.context.llm_context_summarization import (
|
||||||
LLMContextSummarizationConfig,
|
LLMContextSummarizationConfig,
|
||||||
@@ -602,5 +602,77 @@ class TestSummaryGenerationExceptions(unittest.IsolatedAsyncioTestCase):
|
|||||||
self.assertEqual(last_index, 1) # Should be the index of the last summarized message
|
self.assertEqual(last_index, 1) # Should be the index of the last summarized message
|
||||||
|
|
||||||
|
|
||||||
|
class TestLLMSpecificMessageHandling(unittest.TestCase):
|
||||||
|
"""Tests that LLMSpecificMessage objects are correctly skipped in summarization."""
|
||||||
|
|
||||||
|
def test_estimate_context_tokens_skips_specific_messages(self):
|
||||||
|
"""Test that estimate_context_tokens skips LLMSpecificMessage objects."""
|
||||||
|
context = LLMContext()
|
||||||
|
context.add_message({"role": "user", "content": "Hello"})
|
||||||
|
context.add_message(LLMSpecificMessage(llm="google", message={}))
|
||||||
|
context.add_message({"role": "assistant", "content": "Hi there"})
|
||||||
|
|
||||||
|
tokens_with_specific = LLMContextSummarizationUtil.estimate_context_tokens(context)
|
||||||
|
|
||||||
|
context_without = LLMContext()
|
||||||
|
context_without.add_message({"role": "user", "content": "Hello"})
|
||||||
|
context_without.add_message({"role": "assistant", "content": "Hi there"})
|
||||||
|
tokens_without = LLMContextSummarizationUtil.estimate_context_tokens(context_without)
|
||||||
|
|
||||||
|
self.assertEqual(tokens_with_specific, tokens_without)
|
||||||
|
|
||||||
|
def test_get_messages_to_summarize_with_specific_messages(self):
|
||||||
|
"""Test that get_messages_to_summarize handles LLMSpecificMessage objects."""
|
||||||
|
context = LLMContext()
|
||||||
|
context.add_message({"role": "system", "content": "System prompt"})
|
||||||
|
context.add_message(LLMSpecificMessage(llm="google", message={}))
|
||||||
|
context.add_message({"role": "user", "content": "Message 1"})
|
||||||
|
context.add_message({"role": "assistant", "content": "Response 1"})
|
||||||
|
context.add_message(LLMSpecificMessage(llm="google", message={}))
|
||||||
|
context.add_message({"role": "user", "content": "Message 2"})
|
||||||
|
context.add_message({"role": "assistant", "content": "Response 2"})
|
||||||
|
|
||||||
|
result = LLMContextSummarizationUtil.get_messages_to_summarize(context, 2)
|
||||||
|
|
||||||
|
self.assertGreater(len(result.messages), 0)
|
||||||
|
self.assertGreater(result.last_summarized_index, 0)
|
||||||
|
|
||||||
|
def test_format_messages_skips_specific_messages(self):
|
||||||
|
"""Test that format_messages_for_summary skips LLMSpecificMessage objects."""
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "Hello"},
|
||||||
|
LLMSpecificMessage(llm="google", message={}),
|
||||||
|
{"role": "assistant", "content": "Hi there"},
|
||||||
|
]
|
||||||
|
|
||||||
|
transcript = LLMContextSummarizationUtil.format_messages_for_summary(messages)
|
||||||
|
|
||||||
|
self.assertIn("USER: Hello", transcript)
|
||||||
|
self.assertIn("ASSISTANT: Hi there", transcript)
|
||||||
|
|
||||||
|
def test_function_call_tracking_skips_specific_messages(self):
|
||||||
|
"""Test that _get_function_calls_in_progress_index skips LLMSpecificMessage."""
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "What time is it?"},
|
||||||
|
LLMSpecificMessage(llm="google", message={}),
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": "",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": "call_123",
|
||||||
|
"type": "function",
|
||||||
|
"function": {"name": "get_time", "arguments": "{}"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
LLMSpecificMessage(llm="google", message={}),
|
||||||
|
{"role": "tool", "tool_call_id": "call_123", "content": '{"time": "10:30 AM"}'},
|
||||||
|
]
|
||||||
|
|
||||||
|
result = LLMContextSummarizationUtil._get_function_calls_in_progress_index(messages, 0)
|
||||||
|
self.assertEqual(result, -1)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user