Add RTVIObserverParams to control what information is included in function call events
This commit is contained in:
@@ -1 +1 @@
|
|||||||
- Added native RTVI function call lifecycle messages (`llm-function-call-start`, `llm-function-call`, `llm-function-call-cancelled`, `llm-function-call-result`) for visual feedback when bot executes functions.
|
- Added native RTVI function call lifecycle messages (`llm-function-call-started`, `llm-function-call-in-progress`, `llm-function-call-stopped`) for visual feedback when bot executes functions.
|
||||||
@@ -14,7 +14,8 @@ and frame observation for the RTVI protocol.
|
|||||||
import asyncio
|
import asyncio
|
||||||
import base64
|
import base64
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, field
|
||||||
|
from enum import Enum
|
||||||
from typing import (
|
from typing import (
|
||||||
Any,
|
Any,
|
||||||
Awaitable,
|
Awaitable,
|
||||||
@@ -645,10 +646,11 @@ class RTVIAppendToContext(BaseModel):
|
|||||||
class RTVILLMFunctionCallStartMessageData(BaseModel):
|
class RTVILLMFunctionCallStartMessageData(BaseModel):
|
||||||
"""Data for LLM function call start notification.
|
"""Data for LLM function call start notification.
|
||||||
|
|
||||||
Contains the function name being called.
|
Contains the function name being called. Fields may be omitted based on
|
||||||
|
the configured function_call_report_level for security.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
function_name: str
|
function_name: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
class RTVILLMFunctionCallStartMessage(BaseModel):
|
class RTVILLMFunctionCallStartMessage(BaseModel):
|
||||||
@@ -678,11 +680,12 @@ class RTVILLMFunctionCallInProgressMessageData(BaseModel):
|
|||||||
"""Data for LLM function call in-progress notification.
|
"""Data for LLM function call in-progress notification.
|
||||||
|
|
||||||
Contains function call details including name, ID, and arguments.
|
Contains function call details including name, ID, and arguments.
|
||||||
|
Fields may be omitted based on the configured function_call_report_level for security.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
function_name: str
|
|
||||||
tool_call_id: str
|
tool_call_id: str
|
||||||
args: Mapping[str, Any]
|
function_name: Optional[str] = None
|
||||||
|
args: Optional[Mapping[str, Any]] = None
|
||||||
|
|
||||||
|
|
||||||
class RTVILLMFunctionCallInProgressMessage(BaseModel):
|
class RTVILLMFunctionCallInProgressMessage(BaseModel):
|
||||||
@@ -701,11 +704,12 @@ class RTVILLMFunctionCallStoppedMessageData(BaseModel):
|
|||||||
|
|
||||||
Contains details about the function call that stopped, including
|
Contains details about the function call that stopped, including
|
||||||
whether it was cancelled or completed with a result.
|
whether it was cancelled or completed with a result.
|
||||||
|
Fields may be omitted based on the configured function_call_report_level for security.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
function_name: str
|
|
||||||
tool_call_id: str
|
tool_call_id: str
|
||||||
cancelled: bool
|
cancelled: bool
|
||||||
|
function_name: Optional[str] = None
|
||||||
result: Optional[Any] = None
|
result: Optional[Any] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -964,6 +968,24 @@ class RTVIServerMessageFrame(SystemFrame):
|
|||||||
return f"{self.name}(data: {self.data})"
|
return f"{self.name}(data: {self.data})"
|
||||||
|
|
||||||
|
|
||||||
|
class RTVIFunctionCallReportLevel(str, Enum):
|
||||||
|
"""Level of detail to include in function call RTVI events.
|
||||||
|
|
||||||
|
Controls what information is exposed in function call events for security.
|
||||||
|
|
||||||
|
Values:
|
||||||
|
DISABLED: No events emitted for this function call.
|
||||||
|
NONE: Events only with tool_call_id, no function name or metadata (most secure).
|
||||||
|
NAME: Events with function name, no arguments or results.
|
||||||
|
FULL: Events with function name, arguments, and results.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DISABLED = "disabled"
|
||||||
|
NONE = "none"
|
||||||
|
NAME = "name"
|
||||||
|
FULL = "full"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class RTVIObserverParams:
|
class RTVIObserverParams:
|
||||||
"""Parameters for configuring RTVI Observer behavior.
|
"""Parameters for configuring RTVI Observer behavior.
|
||||||
@@ -992,6 +1014,22 @@ class RTVIObserverParams:
|
|||||||
transformed text. To register, provide a list of tuples of
|
transformed text. To register, provide a list of tuples of
|
||||||
(aggregation_type | '*', transform_function).
|
(aggregation_type | '*', transform_function).
|
||||||
audio_level_period_secs: How often audio levels should be sent if enabled.
|
audio_level_period_secs: How often audio levels should be sent if enabled.
|
||||||
|
function_call_report_level: Controls what information is exposed in function call
|
||||||
|
events for security. A dict mapping function names to levels, where ``"*"``
|
||||||
|
sets the default level for unlisted functions::
|
||||||
|
|
||||||
|
function_call_report_level={
|
||||||
|
"*": RTVIFunctionCallReportLevel.DISABLED, # Default: no events
|
||||||
|
"get_weather": RTVIFunctionCallReportLevel.FULL, # Expose everything
|
||||||
|
}
|
||||||
|
|
||||||
|
Levels:
|
||||||
|
- DISABLED: No events emitted for this function.
|
||||||
|
- NONE: Events with tool_call_id only (most secure when events needed).
|
||||||
|
- NAME: Adds function name to events.
|
||||||
|
- FULL: Adds function name, arguments, and results.
|
||||||
|
|
||||||
|
Defaults to ``{"*": RTVIFunctionCallReportLevel.NONE}``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
bot_output_enabled: bool = True
|
bot_output_enabled: bool = True
|
||||||
@@ -1016,6 +1054,9 @@ class RTVIObserverParams:
|
|||||||
]
|
]
|
||||||
] = None
|
] = None
|
||||||
audio_level_period_secs: float = 0.15
|
audio_level_period_secs: float = 0.15
|
||||||
|
function_call_report_level: Dict[str, RTVIFunctionCallReportLevel] = field(
|
||||||
|
default_factory=lambda: {"*": RTVIFunctionCallReportLevel.NONE}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class RTVIObserver(BaseObserver):
|
class RTVIObserver(BaseObserver):
|
||||||
@@ -1104,6 +1145,21 @@ class RTVIObserver(BaseObserver):
|
|||||||
if not (agg_type == aggregation_type and func == transform_function)
|
if not (agg_type == aggregation_type and func == transform_function)
|
||||||
]
|
]
|
||||||
|
|
||||||
|
def _get_function_call_report_level(self, function_name: str) -> RTVIFunctionCallReportLevel:
|
||||||
|
"""Get the report level for a specific function call.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
function_name: The name of the function to get the report level for.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The report level for the function. Looks up the function name first,
|
||||||
|
then falls back to "*" key, then NONE.
|
||||||
|
"""
|
||||||
|
levels = self._params.function_call_report_level
|
||||||
|
if function_name in levels:
|
||||||
|
return levels[function_name]
|
||||||
|
return levels.get("*", RTVIFunctionCallReportLevel.NONE)
|
||||||
|
|
||||||
async def _logger_sink(self, message):
|
async def _logger_sink(self, message):
|
||||||
"""Logger sink so we can send system logs to RTVI clients."""
|
"""Logger sink so we can send system logs to RTVI clients."""
|
||||||
message = RTVISystemLogMessage(data=RTVITextMessageData(text=message))
|
message = RTVISystemLogMessage(data=RTVITextMessageData(text=message))
|
||||||
@@ -1190,40 +1246,60 @@ class RTVIObserver(BaseObserver):
|
|||||||
await self._handle_metrics(frame)
|
await self._handle_metrics(frame)
|
||||||
elif isinstance(frame, FunctionCallsStartedFrame):
|
elif isinstance(frame, FunctionCallsStartedFrame):
|
||||||
for function_call in frame.function_calls:
|
for function_call in frame.function_calls:
|
||||||
message = RTVILLMFunctionCallStartMessage(
|
report_level = self._get_function_call_report_level(function_call.function_name)
|
||||||
data=RTVILLMFunctionCallStartMessageData(
|
if report_level == RTVIFunctionCallReportLevel.DISABLED:
|
||||||
function_name=function_call.function_name
|
continue
|
||||||
)
|
data = RTVILLMFunctionCallStartMessageData()
|
||||||
)
|
if report_level in (
|
||||||
|
RTVIFunctionCallReportLevel.NAME,
|
||||||
|
RTVIFunctionCallReportLevel.FULL,
|
||||||
|
):
|
||||||
|
data.function_name = function_call.function_name
|
||||||
|
message = RTVILLMFunctionCallStartMessage(data=data)
|
||||||
await self.send_rtvi_message(message)
|
await self.send_rtvi_message(message)
|
||||||
elif isinstance(frame, FunctionCallInProgressFrame):
|
elif isinstance(frame, FunctionCallInProgressFrame):
|
||||||
message = RTVILLMFunctionCallInProgressMessage(
|
report_level = self._get_function_call_report_level(frame.function_name)
|
||||||
data=RTVILLMFunctionCallInProgressMessageData(
|
if report_level != RTVIFunctionCallReportLevel.DISABLED:
|
||||||
function_name=frame.function_name,
|
data = RTVILLMFunctionCallInProgressMessageData(tool_call_id=frame.tool_call_id)
|
||||||
tool_call_id=frame.tool_call_id,
|
if report_level in (
|
||||||
args=frame.arguments,
|
RTVIFunctionCallReportLevel.NAME,
|
||||||
)
|
RTVIFunctionCallReportLevel.FULL,
|
||||||
)
|
):
|
||||||
await self.send_rtvi_message(message)
|
data.function_name = frame.function_name
|
||||||
|
if report_level == RTVIFunctionCallReportLevel.FULL:
|
||||||
|
data.args = frame.arguments
|
||||||
|
message = RTVILLMFunctionCallInProgressMessage(data=data)
|
||||||
|
await self.send_rtvi_message(message)
|
||||||
elif isinstance(frame, FunctionCallCancelFrame):
|
elif isinstance(frame, FunctionCallCancelFrame):
|
||||||
message = RTVILLMFunctionCallStoppedMessage(
|
report_level = self._get_function_call_report_level(frame.function_name)
|
||||||
data=RTVILLMFunctionCallStoppedMessageData(
|
if report_level != RTVIFunctionCallReportLevel.DISABLED:
|
||||||
function_name=frame.function_name,
|
data = RTVILLMFunctionCallStoppedMessageData(
|
||||||
tool_call_id=frame.tool_call_id,
|
tool_call_id=frame.tool_call_id,
|
||||||
cancelled=True,
|
cancelled=True,
|
||||||
)
|
)
|
||||||
)
|
if report_level in (
|
||||||
await self.send_rtvi_message(message)
|
RTVIFunctionCallReportLevel.NAME,
|
||||||
|
RTVIFunctionCallReportLevel.FULL,
|
||||||
|
):
|
||||||
|
data.function_name = frame.function_name
|
||||||
|
message = RTVILLMFunctionCallStoppedMessage(data=data)
|
||||||
|
await self.send_rtvi_message(message)
|
||||||
elif isinstance(frame, FunctionCallResultFrame):
|
elif isinstance(frame, FunctionCallResultFrame):
|
||||||
message = RTVILLMFunctionCallStoppedMessage(
|
report_level = self._get_function_call_report_level(frame.function_name)
|
||||||
data=RTVILLMFunctionCallStoppedMessageData(
|
if report_level != RTVIFunctionCallReportLevel.DISABLED:
|
||||||
function_name=frame.function_name,
|
data = RTVILLMFunctionCallStoppedMessageData(
|
||||||
tool_call_id=frame.tool_call_id,
|
tool_call_id=frame.tool_call_id,
|
||||||
cancelled=False,
|
cancelled=False,
|
||||||
result=frame.result if frame.result else None,
|
|
||||||
)
|
)
|
||||||
)
|
if report_level in (
|
||||||
await self.send_rtvi_message(message)
|
RTVIFunctionCallReportLevel.NAME,
|
||||||
|
RTVIFunctionCallReportLevel.FULL,
|
||||||
|
):
|
||||||
|
data.function_name = frame.function_name
|
||||||
|
if report_level == RTVIFunctionCallReportLevel.FULL:
|
||||||
|
data.result = frame.result if frame.result else None
|
||||||
|
message = RTVILLMFunctionCallStoppedMessage(data=data)
|
||||||
|
await self.send_rtvi_message(message)
|
||||||
elif isinstance(frame, RTVIServerMessageFrame):
|
elif isinstance(frame, RTVIServerMessageFrame):
|
||||||
message = RTVIServerMessage(data=frame.data)
|
message = RTVIServerMessage(data=frame.data)
|
||||||
await self.send_rtvi_message(message)
|
await self.send_rtvi_message(message)
|
||||||
|
|||||||
Reference in New Issue
Block a user