services(anthropic): allow setting enable prompt caching via frame
This commit is contained in:
@@ -186,6 +186,13 @@ class LLMSetToolsFrame(DataFrame):
|
|||||||
tools: List[dict]
|
tools: List[dict]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMEnablePromptCachingFrame(DataFrame):
|
||||||
|
"""A frame to enable/disable prompt caching in certain LLMs.
|
||||||
|
"""
|
||||||
|
enable: bool
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TTSSpeakFrame(DataFrame):
|
class TTSSpeakFrame(DataFrame):
|
||||||
"""A frame that contains a text that should be spoken by the TTS in the
|
"""A frame that contains a text that should be spoken by the TTS in the
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import re
|
|||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
Frame,
|
Frame,
|
||||||
|
LLMEnablePromptCachingFrame,
|
||||||
LLMModelUpdateFrame,
|
LLMModelUpdateFrame,
|
||||||
TextFrame,
|
TextFrame,
|
||||||
VisionImageRawFrame,
|
VisionImageRawFrame,
|
||||||
@@ -62,10 +63,10 @@ class AnthropicContextAggregatorPair:
|
|||||||
_user: 'AnthropicUserContextAggregator'
|
_user: 'AnthropicUserContextAggregator'
|
||||||
_assistant: 'AnthropicAssistantContextAggregator'
|
_assistant: 'AnthropicAssistantContextAggregator'
|
||||||
|
|
||||||
def user(self) -> str:
|
def user(self) -> 'AnthropicUserContextAggregator':
|
||||||
return self._user
|
return self._user
|
||||||
|
|
||||||
def assistant(self) -> str:
|
def assistant(self) -> 'AnthropicAssistantContextAggregator':
|
||||||
return self._assistant
|
return self._assistant
|
||||||
|
|
||||||
|
|
||||||
@@ -227,6 +228,9 @@ class AnthropicLLMService(LLMService):
|
|||||||
elif isinstance(frame, LLMModelUpdateFrame):
|
elif isinstance(frame, LLMModelUpdateFrame):
|
||||||
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
||||||
self._model = frame.model
|
self._model = frame.model
|
||||||
|
elif isinstance(frame, LLMEnablePromptCachingFrame):
|
||||||
|
logger.debug(f"Setting enable prompt caching to: [{frame.enable}]")
|
||||||
|
self._enable_prompt_caching_beta = frame.enable
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|||||||
@@ -229,10 +229,10 @@ class OpenAIContextAggregatorPair:
|
|||||||
_user: 'OpenAIUserContextAggregator'
|
_user: 'OpenAIUserContextAggregator'
|
||||||
_assistant: 'OpenAIAssistantContextAggregator'
|
_assistant: 'OpenAIAssistantContextAggregator'
|
||||||
|
|
||||||
def user(self) -> str:
|
def user(self) -> 'OpenAIUserContextAggregator':
|
||||||
return self._user
|
return self._user
|
||||||
|
|
||||||
def assistant(self) -> str:
|
def assistant(self) -> 'OpenAIAssistantContextAggregator':
|
||||||
return self._assistant
|
return self._assistant
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -49,10 +49,10 @@ class TogetherContextAggregatorPair:
|
|||||||
_user: 'TogetherUserContextAggregator'
|
_user: 'TogetherUserContextAggregator'
|
||||||
_assistant: 'TogetherAssistantContextAggregator'
|
_assistant: 'TogetherAssistantContextAggregator'
|
||||||
|
|
||||||
def user(self) -> str:
|
def user(self) -> 'TogetherUserContextAggregator':
|
||||||
return self._user
|
return self._user
|
||||||
|
|
||||||
def assistant(self) -> str:
|
def assistant(self) -> 'TogetherAssistantContextAggregator':
|
||||||
return self._assistant
|
return self._assistant
|
||||||
|
|
||||||
|
|
||||||
@@ -75,7 +75,7 @@ class TogetherLLMService(LLMService):
|
|||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@ staticmethod
|
@staticmethod
|
||||||
def create_context_aggregator(context: OpenAILLMContext) -> TogetherContextAggregatorPair:
|
def create_context_aggregator(context: OpenAILLMContext) -> TogetherContextAggregatorPair:
|
||||||
user = TogetherUserContextAggregator(context)
|
user = TogetherUserContextAggregator(context)
|
||||||
assistant = TogetherAssistantContextAggregator(user)
|
assistant = TogetherAssistantContextAggregator(user)
|
||||||
@@ -191,14 +191,14 @@ class TogetherLLMContext(OpenAILLMContext):
|
|||||||
):
|
):
|
||||||
super().__init__(messages=messages)
|
super().__init__(messages=messages)
|
||||||
|
|
||||||
@ classmethod
|
@classmethod
|
||||||
def from_openai_context(cls, openai_context: OpenAILLMContext):
|
def from_openai_context(cls, openai_context: OpenAILLMContext):
|
||||||
self = cls(
|
self = cls(
|
||||||
messages=openai_context.messages,
|
messages=openai_context.messages,
|
||||||
)
|
)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
@ classmethod
|
@classmethod
|
||||||
def from_messages(cls, messages: List[dict]) -> "TogetherLLMContext":
|
def from_messages(cls, messages: List[dict]) -> "TogetherLLMContext":
|
||||||
return cls(messages=messages)
|
return cls(messages=messages)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user