Add RTVIObserverParams to control what information is included in function call events

This commit is contained in:
Mark Backman
2026-02-03 10:11:36 -05:00
parent 40c84faff5
commit e0e3b5250b
2 changed files with 107 additions and 31 deletions

View File

@@ -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.

View File

@@ -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)