feat: add session updates and message dialogs
This commit is contained in:
@@ -44,6 +44,14 @@ class BrainSpec:
|
|||||||
owns_context: bool
|
owns_context: bool
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SessionVariableUpdate:
|
||||||
|
"""Result of one silent update to declared session variables."""
|
||||||
|
|
||||||
|
changed: list[str]
|
||||||
|
dynamic_variables: dict[str, str | int | float | bool]
|
||||||
|
|
||||||
|
|
||||||
class CallEndPort(Protocol):
|
class CallEndPort(Protocol):
|
||||||
"""Small call-lifecycle surface available to a brain."""
|
"""Small call-lifecycle surface available to a brain."""
|
||||||
|
|
||||||
@@ -144,6 +152,13 @@ class BaseBrain:
|
|||||||
async def on_client_ready(self) -> None:
|
async def on_client_ready(self) -> None:
|
||||||
"""Replay client-visible state after its app message channel is ready."""
|
"""Replay client-visible state after its app message channel is ready."""
|
||||||
|
|
||||||
|
async def on_session_update(
|
||||||
|
self,
|
||||||
|
dynamic_variables: dict[str, Any],
|
||||||
|
) -> SessionVariableUpdate:
|
||||||
|
"""Apply a silent client state update without starting an LLM turn."""
|
||||||
|
raise ValueError(f"助手类型 {self.spec.type} 不支持会话动态变量更新")
|
||||||
|
|
||||||
def record_user_message(self, content: str) -> None:
|
def record_user_message(self, content: str) -> None:
|
||||||
"""Observe a committed user message for brain-owned routing state."""
|
"""Observe a committed user message for brain-owned routing state."""
|
||||||
|
|
||||||
@@ -194,6 +209,11 @@ class Brain(Protocol):
|
|||||||
|
|
||||||
async def on_client_ready(self) -> None: ...
|
async def on_client_ready(self) -> None: ...
|
||||||
|
|
||||||
|
async def on_session_update(
|
||||||
|
self,
|
||||||
|
dynamic_variables: dict[str, Any],
|
||||||
|
) -> SessionVariableUpdate: ...
|
||||||
|
|
||||||
def record_user_message(self, content: str) -> None: ...
|
def record_user_message(self, content: str) -> None: ...
|
||||||
|
|
||||||
async def on_user_turn_end(self, content: str) -> bool: ...
|
async def on_user_turn_end(self, content: str) -> bool: ...
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
from models import AssistantConfig
|
from models import AssistantConfig
|
||||||
@@ -15,10 +16,13 @@ from pipecat.services.llm_service import (
|
|||||||
)
|
)
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
from services.brains.base import BaseBrain, BrainRuntime, BrainSpec
|
from services.brains.base import (
|
||||||
from services.runtime_variables import (
|
BaseBrain,
|
||||||
DynamicVariableStore,
|
BrainRuntime,
|
||||||
|
BrainSpec,
|
||||||
|
SessionVariableUpdate,
|
||||||
)
|
)
|
||||||
|
from services.runtime_variables import DynamicVariableStore
|
||||||
from services.tool_executor import ToolExecutionError, ToolExecutor
|
from services.tool_executor import ToolExecutionError, ToolExecutor
|
||||||
from services.tool_policy import policy_for_tool
|
from services.tool_policy import policy_for_tool
|
||||||
|
|
||||||
@@ -76,6 +80,18 @@ class PromptBrain(BaseBrain):
|
|||||||
self._store.record("user", content)
|
self._store.record("user", content)
|
||||||
self._refresh_prompt()
|
self._refresh_prompt()
|
||||||
|
|
||||||
|
async def on_session_update(
|
||||||
|
self,
|
||||||
|
dynamic_variables: dict[str, Any],
|
||||||
|
) -> SessionVariableUpdate:
|
||||||
|
changed = self._store.assign_declared_many(dynamic_variables)
|
||||||
|
if changed:
|
||||||
|
self._refresh_prompt()
|
||||||
|
return SessionVariableUpdate(
|
||||||
|
changed=changed,
|
||||||
|
dynamic_variables=self._store.public_values(),
|
||||||
|
)
|
||||||
|
|
||||||
async def on_assistant_text_start(self, _turn_id: str) -> None:
|
async def on_assistant_text_start(self, _turn_id: str) -> None:
|
||||||
if self._runtime is not None:
|
if self._runtime is not None:
|
||||||
self._runtime.call_end.begin_response()
|
self._runtime.call_end.begin_response()
|
||||||
|
|||||||
@@ -28,7 +28,12 @@ from pipecat.services.llm_service import (
|
|||||||
FunctionCallParams,
|
FunctionCallParams,
|
||||||
FunctionCallResultProperties,
|
FunctionCallResultProperties,
|
||||||
)
|
)
|
||||||
from services.brains.base import BaseBrain, BrainRuntime, BrainSpec
|
from services.brains.base import (
|
||||||
|
BaseBrain,
|
||||||
|
BrainRuntime,
|
||||||
|
BrainSpec,
|
||||||
|
SessionVariableUpdate,
|
||||||
|
)
|
||||||
from services.knowledge import search as search_knowledge
|
from services.knowledge import search as search_knowledge
|
||||||
from services.runtime_variables import DynamicVariableStore
|
from services.runtime_variables import DynamicVariableStore
|
||||||
from services.tool_executor import ToolExecutionError, ToolExecutor
|
from services.tool_executor import ToolExecutionError, ToolExecutor
|
||||||
@@ -258,6 +263,27 @@ class WorkflowBrain(BaseBrain):
|
|||||||
node_id=current_node,
|
node_id=current_node,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def on_session_update(
|
||||||
|
self,
|
||||||
|
dynamic_variables: dict[str, Any],
|
||||||
|
) -> SessionVariableUpdate:
|
||||||
|
if self._ended:
|
||||||
|
raise ValueError("工作流会话已经结束")
|
||||||
|
changed = self._store.assign_declared_many(dynamic_variables)
|
||||||
|
current = self._state.current_node_id
|
||||||
|
if changed and current and self._engine.node_type(current) == "agent":
|
||||||
|
await self._refresh_agent_prompt(current)
|
||||||
|
if changed:
|
||||||
|
await self._emit_variables(
|
||||||
|
reason="session_update",
|
||||||
|
node_id=current or None,
|
||||||
|
changed=changed,
|
||||||
|
)
|
||||||
|
return SessionVariableUpdate(
|
||||||
|
changed=changed,
|
||||||
|
dynamic_variables=self._store.public_values(),
|
||||||
|
)
|
||||||
|
|
||||||
def record_user_message(self, content: str) -> None:
|
def record_user_message(self, content: str) -> None:
|
||||||
if content and not self._ended:
|
if content and not self._ended:
|
||||||
self._store.record("user", content)
|
self._store.record("user", content)
|
||||||
|
|||||||
@@ -74,6 +74,7 @@ from services.pipecat.processors import (
|
|||||||
PassthroughLLMAssistantAggregator,
|
PassthroughLLMAssistantAggregator,
|
||||||
RealtimeDynamicVariableProcessor,
|
RealtimeDynamicVariableProcessor,
|
||||||
RealtimeUserInputProcessor,
|
RealtimeUserInputProcessor,
|
||||||
|
SessionUpdateProcessor,
|
||||||
UserInput,
|
UserInput,
|
||||||
UserInputProcessor,
|
UserInputProcessor,
|
||||||
UserTurnRoutingProcessor,
|
UserTurnRoutingProcessor,
|
||||||
@@ -365,6 +366,10 @@ async def run_pipeline(
|
|||||||
or tool_interruption_strategy.is_muted
|
or tool_interruption_strategy.is_muted
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
session_update = SessionUpdateProcessor(
|
||||||
|
brain.on_session_update,
|
||||||
|
should_ignore_update=lambda: call_end.ending,
|
||||||
|
)
|
||||||
client_tools = ClientToolBroker()
|
client_tools = ClientToolBroker()
|
||||||
vision_capture = VisionCaptureProcessor()
|
vision_capture = VisionCaptureProcessor()
|
||||||
knowledge_retrieval = KnowledgeRetrievalProcessor(
|
knowledge_retrieval = KnowledgeRetrievalProcessor(
|
||||||
@@ -578,6 +583,7 @@ async def run_pipeline(
|
|||||||
[
|
[
|
||||||
transport.input(),
|
transport.input(),
|
||||||
client_tools,
|
client_tools,
|
||||||
|
session_update,
|
||||||
vision_capture,
|
vision_capture,
|
||||||
user_input,
|
user_input,
|
||||||
stt_processor,
|
stt_processor,
|
||||||
@@ -778,6 +784,16 @@ async def run_realtime_pipeline(
|
|||||||
input_sample_rate, output_sample_rate = realtime_audio_sample_rates(cfg)
|
input_sample_rate, output_sample_rate = realtime_audio_sample_rates(cfg)
|
||||||
user_input = RealtimeUserInputProcessor()
|
user_input = RealtimeUserInputProcessor()
|
||||||
dynamic_variables = RealtimeDynamicVariableProcessor(brain, cfg, realtime)
|
dynamic_variables = RealtimeDynamicVariableProcessor(brain, cfg, realtime)
|
||||||
|
|
||||||
|
async def refresh_realtime_instructions() -> None:
|
||||||
|
update = getattr(realtime, "update_instructions", None)
|
||||||
|
if callable(update):
|
||||||
|
await update(brain.system_prompt(cfg))
|
||||||
|
|
||||||
|
session_update = SessionUpdateProcessor(
|
||||||
|
brain.on_session_update,
|
||||||
|
after_update=refresh_realtime_instructions,
|
||||||
|
)
|
||||||
greeting = await brain.greeting(cfg)
|
greeting = await brain.greeting(cfg)
|
||||||
|
|
||||||
recorder = await ConversationRecorder.start(
|
recorder = await ConversationRecorder.start(
|
||||||
@@ -790,6 +806,7 @@ async def run_realtime_pipeline(
|
|||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(),
|
transport.input(),
|
||||||
|
session_update,
|
||||||
user_input,
|
user_input,
|
||||||
realtime,
|
realtime,
|
||||||
dynamic_variables,
|
dynamic_variables,
|
||||||
|
|||||||
@@ -1,13 +1,14 @@
|
|||||||
"""Reusable frame processors shared by cascade and realtime pipelines."""
|
"""Reusable frame processors shared by cascade and realtime pipelines."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from collections.abc import Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from models import AssistantConfig
|
from models import AssistantConfig
|
||||||
from services.brains import Brain
|
from services.brains import Brain
|
||||||
|
from services.brains.base import SessionVariableUpdate
|
||||||
from services.conversation_history import ConversationRecorder
|
from services.conversation_history import ConversationRecorder
|
||||||
from services.knowledge import search as search_knowledge
|
from services.knowledge import search as search_knowledge
|
||||||
from db.session import SessionLocal
|
from db.session import SessionLocal
|
||||||
@@ -38,6 +39,7 @@ from pipecat.utils.time import time_now_iso8601
|
|||||||
|
|
||||||
|
|
||||||
KNOWLEDGE_CONTEXT_MARKER = "<!-- knowledge-context -->"
|
KNOWLEDGE_CONTEXT_MARKER = "<!-- knowledge-context -->"
|
||||||
|
MAX_SESSION_UPDATE_VARIABLES = 20
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -69,6 +71,69 @@ class UserInputError(ValueError):
|
|||||||
self.input_id = input_id
|
self.input_id = input_id
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SessionUpdate:
|
||||||
|
"""Validated silent update to declared session variables."""
|
||||||
|
|
||||||
|
update_id: str
|
||||||
|
dynamic_variables: dict[str, str | int | float | bool]
|
||||||
|
|
||||||
|
|
||||||
|
class SessionUpdateError(ValueError):
|
||||||
|
def __init__(self, message: str, *, update_id: str = "") -> None:
|
||||||
|
super().__init__(message)
|
||||||
|
self.update_id = update_id
|
||||||
|
|
||||||
|
|
||||||
|
def parse_session_update(message) -> SessionUpdate | None:
|
||||||
|
"""Parse the silent session-update wire format."""
|
||||||
|
if not isinstance(message, dict) or message.get("type") != "session-update":
|
||||||
|
return None
|
||||||
|
|
||||||
|
update_id = str(message.get("update_id") or "").strip()
|
||||||
|
if not update_id or len(update_id) > 128:
|
||||||
|
raise SessionUpdateError(
|
||||||
|
"session-update 缺少有效的 update_id",
|
||||||
|
update_id=update_id,
|
||||||
|
)
|
||||||
|
if message.get("schema_version") != 1:
|
||||||
|
raise SessionUpdateError(
|
||||||
|
"不支持的 session-update schema_version",
|
||||||
|
update_id=update_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
raw_variables = message.get("dynamic_variables")
|
||||||
|
if not isinstance(raw_variables, dict) or not raw_variables:
|
||||||
|
raise SessionUpdateError(
|
||||||
|
"session-update dynamic_variables 不能为空",
|
||||||
|
update_id=update_id,
|
||||||
|
)
|
||||||
|
if len(raw_variables) > MAX_SESSION_UPDATE_VARIABLES:
|
||||||
|
raise SessionUpdateError(
|
||||||
|
f"session-update 一次最多更新 {MAX_SESSION_UPDATE_VARIABLES} 个变量",
|
||||||
|
update_id=update_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
dynamic_variables: dict[str, str | int | float | bool] = {}
|
||||||
|
for raw_name, value in raw_variables.items():
|
||||||
|
if not isinstance(raw_name, str):
|
||||||
|
raise SessionUpdateError(
|
||||||
|
"session-update 动态变量名必须是字符串",
|
||||||
|
update_id=update_id,
|
||||||
|
)
|
||||||
|
if not isinstance(value, (str, int, float, bool)) or value is None:
|
||||||
|
raise SessionUpdateError(
|
||||||
|
f"动态变量 {raw_name} 仅支持字符串、数字或布尔值",
|
||||||
|
update_id=update_id,
|
||||||
|
)
|
||||||
|
dynamic_variables[raw_name] = value
|
||||||
|
|
||||||
|
return SessionUpdate(
|
||||||
|
update_id=update_id,
|
||||||
|
dynamic_variables=dynamic_variables,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def parse_user_input(message) -> UserInput | None:
|
def parse_user_input(message) -> UserInput | None:
|
||||||
"""Parse the sole public user-input wire format."""
|
"""Parse the sole public user-input wire format."""
|
||||||
if not isinstance(message, dict):
|
if not isinstance(message, dict):
|
||||||
@@ -198,6 +263,94 @@ class UserInputProcessor(FrameProcessor):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SessionUpdateProcessor(FrameProcessor):
|
||||||
|
"""Apply silent dynamic-variable updates without creating a user turn."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
apply_update: Callable[
|
||||||
|
[dict[str, str | int | float | bool]],
|
||||||
|
Awaitable[SessionVariableUpdate],
|
||||||
|
],
|
||||||
|
*,
|
||||||
|
after_update: Callable[[], Awaitable[None]] | None = None,
|
||||||
|
should_ignore_update: Callable[[], bool] | None = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self._apply_update = apply_update
|
||||||
|
self._after_update = after_update
|
||||||
|
self._should_ignore_update = should_ignore_update or (lambda: False)
|
||||||
|
|
||||||
|
async def process_frame(self, frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
if not isinstance(frame, InputTransportMessageFrame):
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
update = parse_session_update(frame.message)
|
||||||
|
except SessionUpdateError as exc:
|
||||||
|
await self._emit_result(exc.update_id, "error", message=str(exc))
|
||||||
|
return
|
||||||
|
if update is None:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
return
|
||||||
|
if self._should_ignore_update():
|
||||||
|
await self._emit_result(
|
||||||
|
update.update_id,
|
||||||
|
"error",
|
||||||
|
message="当前不能更新会话状态",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await self._apply_update(update.dynamic_variables)
|
||||||
|
if result.changed and self._after_update is not None:
|
||||||
|
await self._after_update()
|
||||||
|
except Exception as exc: # noqa: BLE001 - protocol failures need a reply
|
||||||
|
logger.warning(f"会话状态更新失败: {exc}")
|
||||||
|
await self._emit_result(
|
||||||
|
update.update_id,
|
||||||
|
"error",
|
||||||
|
message=str(exc),
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
await self._emit_result(
|
||||||
|
update.update_id,
|
||||||
|
"accepted",
|
||||||
|
changed=result.changed,
|
||||||
|
dynamic_variables=result.dynamic_variables,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _emit_result(
|
||||||
|
self,
|
||||||
|
update_id: str,
|
||||||
|
status: str,
|
||||||
|
*,
|
||||||
|
message: str = "",
|
||||||
|
changed: list[str] | None = None,
|
||||||
|
dynamic_variables: dict[str, str | int | float | bool] | None = None,
|
||||||
|
) -> None:
|
||||||
|
await self.push_frame(
|
||||||
|
OutputTransportMessageUrgentFrame(
|
||||||
|
message={
|
||||||
|
"type": "session-update-result",
|
||||||
|
"update_id": update_id,
|
||||||
|
"status": status,
|
||||||
|
**({"message": message} if message else {}),
|
||||||
|
**({"changed": changed} if changed is not None else {}),
|
||||||
|
**(
|
||||||
|
{"dynamic_variables": dynamic_variables}
|
||||||
|
if dynamic_variables is not None
|
||||||
|
else {}
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class CallEndingUserMuteStrategy(BaseUserMuteStrategy):
|
class CallEndingUserMuteStrategy(BaseUserMuteStrategy):
|
||||||
"""Keep user media muted after an end-call tool starts terminating a call."""
|
"""Keep user media muted after an end-call tool starts terminating a call."""
|
||||||
|
|
||||||
|
|||||||
@@ -210,6 +210,39 @@ class DynamicVariableStore:
|
|||||||
raise DynamicVariableError(f"动态变量 {name} 类型应为 {expected}")
|
raise DynamicVariableError(f"动态变量 {name} 类型应为 {expected}")
|
||||||
self.values[name] = primitive
|
self.values[name] = primitive
|
||||||
|
|
||||||
|
def assign_declared_many(self, values: dict[str, Any]) -> list[str]:
|
||||||
|
"""Atomically update public variables declared by the assistant."""
|
||||||
|
staged: dict[str, Primitive] = {}
|
||||||
|
for name, value in values.items():
|
||||||
|
if (
|
||||||
|
name.startswith(("system__", "secret__"))
|
||||||
|
or not VARIABLE_NAME.fullmatch(name)
|
||||||
|
):
|
||||||
|
raise DynamicVariableError(f"会话不能更新保留变量: {name}")
|
||||||
|
if name not in self.variable_types:
|
||||||
|
raise DynamicVariableError(f"动态变量未声明: {name}")
|
||||||
|
primitive = _primitive(value, name)
|
||||||
|
expected = self.variable_types[name]
|
||||||
|
if not _type_matches(primitive, expected):
|
||||||
|
raise DynamicVariableError(f"动态变量 {name} 类型应为 {expected}")
|
||||||
|
staged[name] = primitive
|
||||||
|
|
||||||
|
changed = [
|
||||||
|
name
|
||||||
|
for name, value in staged.items()
|
||||||
|
if name not in self.values or self.values[name] != value
|
||||||
|
]
|
||||||
|
self.values.update(staged)
|
||||||
|
return changed
|
||||||
|
|
||||||
|
def public_values(self) -> dict[str, Primitive]:
|
||||||
|
"""Return the client-visible dynamic-variable snapshot."""
|
||||||
|
return {
|
||||||
|
name: value
|
||||||
|
for name, value in self.values.items()
|
||||||
|
if not name.startswith(("system__", "secret__"))
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def prepare_dynamic_config(
|
def prepare_dynamic_config(
|
||||||
cfg: AssistantConfig,
|
cfg: AssistantConfig,
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from pipecat.frames.frames import (
|
|||||||
LLMMessagesUpdateFrame,
|
LLMMessagesUpdateFrame,
|
||||||
LLMRunFrame,
|
LLMRunFrame,
|
||||||
LLMTextFrame,
|
LLMTextFrame,
|
||||||
|
LLMUpdateSettingsFrame,
|
||||||
OutputTransportMessageUrgentFrame,
|
OutputTransportMessageUrgentFrame,
|
||||||
TTSSpeakFrame,
|
TTSSpeakFrame,
|
||||||
)
|
)
|
||||||
@@ -28,7 +29,7 @@ from services.brains.dify_llm import (
|
|||||||
)
|
)
|
||||||
from services.brains.workflow_brain import WorkflowBrain
|
from services.brains.workflow_brain import WorkflowBrain
|
||||||
from services.runtime_variables import prepare_dynamic_config
|
from services.runtime_variables import prepare_dynamic_config
|
||||||
from services.workflow.models import LLMRouteResult, RouteStatus
|
from services.workflow.models import LLMRouteResult, RouteStatus, WorkflowStatus
|
||||||
|
|
||||||
|
|
||||||
class FakeLLM:
|
class FakeLLM:
|
||||||
@@ -267,6 +268,47 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
|||||||
f"{GREETING_CONTEXT_MARKER}\n您好,王先生",
|
f"{GREETING_CONTEXT_MARKER}\n您好,王先生",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def test_session_update_refreshes_prompt_without_running_llm(self):
|
||||||
|
cfg = prepare_dynamic_config(
|
||||||
|
AssistantConfig(
|
||||||
|
type="prompt",
|
||||||
|
runtimeMode="pipeline",
|
||||||
|
prompt="面板状态:{{panel_open}}",
|
||||||
|
dynamic_variable_definitions={
|
||||||
|
"panel_open": {"type": "boolean", "default": False}
|
||||||
|
},
|
||||||
|
),
|
||||||
|
{},
|
||||||
|
assistant_id="asst_session_update",
|
||||||
|
)
|
||||||
|
brain = build_brain(cfg)
|
||||||
|
prompts = []
|
||||||
|
queued_frames = []
|
||||||
|
|
||||||
|
async def queue_frame(frame):
|
||||||
|
queued_frames.append(frame)
|
||||||
|
|
||||||
|
await brain.setup(
|
||||||
|
cfg,
|
||||||
|
BrainRuntime(
|
||||||
|
context=LLMContext(
|
||||||
|
messages=[{"role": "system", "content": brain.system_prompt(cfg)}]
|
||||||
|
),
|
||||||
|
llm=FakeLLM(),
|
||||||
|
queue_frame=queue_frame,
|
||||||
|
set_system_prompt=prompts.append,
|
||||||
|
set_tools=lambda _tools: None,
|
||||||
|
call_end=FakeCallEnd(),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await brain.on_session_update({"panel_open": True})
|
||||||
|
|
||||||
|
self.assertEqual(result.changed, ["panel_open"])
|
||||||
|
self.assertEqual(result.dynamic_variables, {"panel_open": True})
|
||||||
|
self.assertEqual(prompts, ["面板状态:true"])
|
||||||
|
self.assertEqual(queued_frames, [])
|
||||||
|
|
||||||
async def test_end_call_tool_is_owned_by_prompt_brain(self):
|
async def test_end_call_tool_is_owned_by_prompt_brain(self):
|
||||||
brain = build_brain(
|
brain = build_brain(
|
||||||
AssistantConfig(
|
AssistantConfig(
|
||||||
@@ -472,6 +514,71 @@ class PromptBrainTests(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
|
|
||||||
class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
class WorkflowBrainTests(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def test_session_update_refreshes_current_agent_without_routing(self):
|
||||||
|
cfg = prepare_dynamic_config(
|
||||||
|
AssistantConfig(
|
||||||
|
type="workflow",
|
||||||
|
graph={
|
||||||
|
"specVersion": 3,
|
||||||
|
"settings": {"globalPrompt": "面板 {{panel_open}}"},
|
||||||
|
"nodes": [
|
||||||
|
{"id": "start", "type": "start", "data": {}},
|
||||||
|
{
|
||||||
|
"id": "agent",
|
||||||
|
"type": "agent",
|
||||||
|
"data": {"prompt": "等待用户操作"},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
"edges": [
|
||||||
|
{
|
||||||
|
"id": "begin",
|
||||||
|
"source": "start",
|
||||||
|
"target": "agent",
|
||||||
|
"data": {"mode": "always", "priority": 0},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
dynamic_variable_definitions={
|
||||||
|
"panel_open": {"type": "boolean", "default": False}
|
||||||
|
},
|
||||||
|
),
|
||||||
|
{},
|
||||||
|
assistant_id="asst_workflow_session_update",
|
||||||
|
)
|
||||||
|
brain = WorkflowBrain(cfg)
|
||||||
|
queued = []
|
||||||
|
|
||||||
|
async def queue_frame(frame):
|
||||||
|
queued.append(frame)
|
||||||
|
|
||||||
|
brain._runtime = BrainRuntime(
|
||||||
|
context=LLMContext(messages=[]),
|
||||||
|
llm=FakeLLM(),
|
||||||
|
queue_frame=queue_frame,
|
||||||
|
set_system_prompt=lambda _prompt: None,
|
||||||
|
set_tools=lambda _tools: None,
|
||||||
|
call_end=FakeCallEnd(),
|
||||||
|
)
|
||||||
|
brain._state.enter("agent", WorkflowStatus.WAITING_USER)
|
||||||
|
|
||||||
|
result = await brain.on_session_update({"panel_open": True})
|
||||||
|
|
||||||
|
self.assertEqual(result.changed, ["panel_open"])
|
||||||
|
prompt_updates = [
|
||||||
|
frame for frame in queued if isinstance(frame, LLMUpdateSettingsFrame)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(prompt_updates), 1)
|
||||||
|
self.assertIn("面板 true", prompt_updates[0].delta.system_instruction)
|
||||||
|
self.assertFalse(any(isinstance(frame, LLMRunFrame) for frame in queued))
|
||||||
|
variable_event = next(
|
||||||
|
frame.message
|
||||||
|
for frame in queued
|
||||||
|
if isinstance(frame, OutputTransportMessageUrgentFrame)
|
||||||
|
and frame.message.get("type") == "workflow-variables"
|
||||||
|
)
|
||||||
|
self.assertEqual(variable_event["reason"], "session_update")
|
||||||
|
self.assertEqual(variable_event["changed"], ["panel_open"])
|
||||||
|
|
||||||
async def test_agent_vision_tool_is_scoped_to_effective_stage(self):
|
async def test_agent_vision_tool_is_scoped_to_effective_stage(self):
|
||||||
brain = WorkflowBrain(
|
brain = WorkflowBrain(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -146,6 +146,34 @@ class DynamicVariableTests(unittest.TestCase):
|
|||||||
self.assertEqual(store.values["system__agent_turns"], 1)
|
self.assertEqual(store.values["system__agent_turns"], 1)
|
||||||
self.assertIn("查订单", store.values["system__conversation_history"])
|
self.assertIn("查订单", store.values["system__conversation_history"])
|
||||||
|
|
||||||
|
def test_session_assignment_is_declared_typed_and_atomic(self):
|
||||||
|
store = DynamicVariableStore(
|
||||||
|
{"panel_open": False, "selection": "none"},
|
||||||
|
variable_types={
|
||||||
|
"panel_open": "boolean",
|
||||||
|
"selection": "string",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
changed = store.assign_declared_many(
|
||||||
|
{"panel_open": True, "selection": "confirmed"}
|
||||||
|
)
|
||||||
|
self.assertEqual(changed, ["panel_open", "selection"])
|
||||||
|
self.assertEqual(
|
||||||
|
store.public_values(),
|
||||||
|
{"panel_open": True, "selection": "confirmed"},
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(DynamicVariableError, "类型应为 boolean"):
|
||||||
|
store.assign_declared_many(
|
||||||
|
{"selection": "cancelled", "panel_open": "yes"}
|
||||||
|
)
|
||||||
|
self.assertEqual(store.values["selection"], "confirmed")
|
||||||
|
self.assertIs(store.values["panel_open"], True)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(DynamicVariableError, "未声明"):
|
||||||
|
store.assign_declared_many({"unknown_state": True})
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
138
backend/tests/test_session_update.py
Normal file
138
backend/tests/test_session_update.py
Normal file
@@ -0,0 +1,138 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from pipecat.frames.frames import (
|
||||||
|
InputTransportMessageFrame,
|
||||||
|
OutputTransportMessageUrgentFrame,
|
||||||
|
)
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
|
from services.brains.base import SessionVariableUpdate
|
||||||
|
from services.pipecat.processors import (
|
||||||
|
SessionUpdateError,
|
||||||
|
SessionUpdateProcessor,
|
||||||
|
parse_session_update,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SessionUpdateParserTests(unittest.TestCase):
|
||||||
|
def test_parses_declared_primitive_values(self):
|
||||||
|
value = parse_session_update(
|
||||||
|
{
|
||||||
|
"type": "session-update",
|
||||||
|
"schema_version": 1,
|
||||||
|
"update_id": "update_1",
|
||||||
|
"dynamic_variables": {
|
||||||
|
"panel_open": True,
|
||||||
|
"selected_option": "approve",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIsNotNone(value)
|
||||||
|
self.assertEqual(value.update_id, "update_1")
|
||||||
|
self.assertEqual(
|
||||||
|
value.dynamic_variables,
|
||||||
|
{"panel_open": True, "selected_option": "approve"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_rejects_empty_or_nested_updates(self):
|
||||||
|
with self.assertRaisesRegex(SessionUpdateError, "不能为空"):
|
||||||
|
parse_session_update(
|
||||||
|
{
|
||||||
|
"type": "session-update",
|
||||||
|
"schema_version": 1,
|
||||||
|
"update_id": "update_empty",
|
||||||
|
"dynamic_variables": {},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
with self.assertRaisesRegex(SessionUpdateError, "仅支持"):
|
||||||
|
parse_session_update(
|
||||||
|
{
|
||||||
|
"type": "session-update",
|
||||||
|
"schema_version": 1,
|
||||||
|
"update_id": "update_nested",
|
||||||
|
"dynamic_variables": {"dialog": {"open": True}},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SessionUpdateProcessorTests(unittest.IsolatedAsyncioTestCase):
|
||||||
|
async def test_returns_snapshot_without_starting_a_turn(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
async def apply_update(dynamic_variables):
|
||||||
|
calls.append(dynamic_variables)
|
||||||
|
return SessionVariableUpdate(
|
||||||
|
changed=["panel_open"],
|
||||||
|
dynamic_variables={
|
||||||
|
"panel_open": True,
|
||||||
|
"selected_option": "none",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
processor = SessionUpdateProcessor(apply_update)
|
||||||
|
outbound = []
|
||||||
|
|
||||||
|
async def push_frame(frame, direction=FrameDirection.DOWNSTREAM):
|
||||||
|
outbound.append((frame, direction))
|
||||||
|
|
||||||
|
processor.push_frame = push_frame
|
||||||
|
await processor.process_frame(
|
||||||
|
InputTransportMessageFrame(
|
||||||
|
message={
|
||||||
|
"type": "session-update",
|
||||||
|
"schema_version": 1,
|
||||||
|
"update_id": "update_accepted",
|
||||||
|
"dynamic_variables": {"panel_open": True},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
FrameDirection.DOWNSTREAM,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(calls, [{"panel_open": True}])
|
||||||
|
self.assertEqual(len(outbound), 1)
|
||||||
|
self.assertIsInstance(outbound[0][0], OutputTransportMessageUrgentFrame)
|
||||||
|
self.assertEqual(
|
||||||
|
outbound[0][0].message,
|
||||||
|
{
|
||||||
|
"type": "session-update-result",
|
||||||
|
"update_id": "update_accepted",
|
||||||
|
"status": "accepted",
|
||||||
|
"changed": ["panel_open"],
|
||||||
|
"dynamic_variables": {
|
||||||
|
"panel_open": True,
|
||||||
|
"selected_option": "none",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_validation_failure_returns_error(self):
|
||||||
|
async def reject_update(_dynamic_variables):
|
||||||
|
raise ValueError("动态变量未声明: panel_open")
|
||||||
|
|
||||||
|
processor = SessionUpdateProcessor(reject_update)
|
||||||
|
outbound = []
|
||||||
|
|
||||||
|
async def push_frame(frame, direction=FrameDirection.DOWNSTREAM):
|
||||||
|
outbound.append((frame, direction))
|
||||||
|
|
||||||
|
processor.push_frame = push_frame
|
||||||
|
await processor.process_frame(
|
||||||
|
InputTransportMessageFrame(
|
||||||
|
message={
|
||||||
|
"type": "session-update",
|
||||||
|
"schema_version": 1,
|
||||||
|
"update_id": "update_rejected",
|
||||||
|
"dynamic_variables": {"panel_open": True},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
FrameDirection.DOWNSTREAM,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(outbound[0][0].message["status"], "error")
|
||||||
|
self.assertIn("未声明", outbound[0][0].message["message"])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -22,6 +22,7 @@ import {
|
|||||||
} from "lucide-react";
|
} from "lucide-react";
|
||||||
|
|
||||||
import { NetworkQualityIndicator } from "@/components/network-quality-indicator";
|
import { NetworkQualityIndicator } from "@/components/network-quality-indicator";
|
||||||
|
import { ClientMessageDialog } from "@/components/client-message-dialog";
|
||||||
import { AuraVisualizer } from "@/components/ui/aura-visualizer";
|
import { AuraVisualizer } from "@/components/ui/aura-visualizer";
|
||||||
import { Badge } from "@/components/ui/badge";
|
import { Badge } from "@/components/ui/badge";
|
||||||
import { Button } from "@/components/ui/button";
|
import { Button } from "@/components/ui/button";
|
||||||
@@ -684,6 +685,7 @@ function DebugVoicePanel({
|
|||||||
<div className="flex min-h-0 flex-1 flex-col">
|
<div className="flex min-h-0 flex-1 flex-col">
|
||||||
{/* 后端 TTS 音频经 WebRTC 媒体流过来,挂这里播放 */}
|
{/* 后端 TTS 音频经 WebRTC 媒体流过来,挂这里播放 */}
|
||||||
<audio ref={audioRef} autoPlay playsInline className="hidden" />
|
<audio ref={audioRef} autoPlay playsInline className="hidden" />
|
||||||
|
<ClientMessageDialog preview={preview} />
|
||||||
{vision && !showIdleHub ? (
|
{vision && !showIdleHub ? (
|
||||||
<DebugVisionWorkspace
|
<DebugVisionWorkspace
|
||||||
view={view}
|
view={view}
|
||||||
|
|||||||
202
frontend/src/components/client-message-dialog.tsx
Normal file
202
frontend/src/components/client-message-dialog.tsx
Normal file
@@ -0,0 +1,202 @@
|
|||||||
|
"use client";
|
||||||
|
|
||||||
|
import { useCallback, useEffect, useRef, useState } from "react";
|
||||||
|
|
||||||
|
import { Button } from "@/components/ui/button";
|
||||||
|
import {
|
||||||
|
Dialog,
|
||||||
|
DialogContent,
|
||||||
|
DialogDescription,
|
||||||
|
DialogFooter,
|
||||||
|
DialogHeader,
|
||||||
|
DialogTitle,
|
||||||
|
} from "@/components/ui/dialog";
|
||||||
|
import type { VoicePreview } from "@/hooks/use-voice-preview";
|
||||||
|
|
||||||
|
const SHOW_MESSAGE_TOOL = "show_message";
|
||||||
|
const MAX_ACTIONS = 5;
|
||||||
|
|
||||||
|
type MessageActionStyle = "primary" | "secondary" | "danger";
|
||||||
|
|
||||||
|
type MessageAction = {
|
||||||
|
id: string;
|
||||||
|
label: string;
|
||||||
|
style: MessageActionStyle;
|
||||||
|
};
|
||||||
|
|
||||||
|
type MessageState = {
|
||||||
|
title: string;
|
||||||
|
message: string;
|
||||||
|
actions: MessageAction[];
|
||||||
|
dismissible: boolean;
|
||||||
|
};
|
||||||
|
|
||||||
|
type PendingMessage = {
|
||||||
|
resolve: (result: { action: string }) => void;
|
||||||
|
reject: (error: Error) => void;
|
||||||
|
};
|
||||||
|
|
||||||
|
function requiredString(
|
||||||
|
value: unknown,
|
||||||
|
field: string,
|
||||||
|
maxLength: number,
|
||||||
|
): string {
|
||||||
|
if (typeof value !== "string" || !value.trim()) {
|
||||||
|
throw new Error(`${field} 必须是非空字符串`);
|
||||||
|
}
|
||||||
|
const normalized = value.trim();
|
||||||
|
if (normalized.length > maxLength) {
|
||||||
|
throw new Error(`${field} 最多允许 ${maxLength} 个字符`);
|
||||||
|
}
|
||||||
|
return normalized;
|
||||||
|
}
|
||||||
|
|
||||||
|
function normalizeAction(value: unknown, index: number): MessageAction {
|
||||||
|
if (!value || typeof value !== "object" || Array.isArray(value)) {
|
||||||
|
throw new Error(`actions[${index}] 格式不正确`);
|
||||||
|
}
|
||||||
|
const action = value as Record<string, unknown>;
|
||||||
|
const rawStyle = action.style ?? "secondary";
|
||||||
|
if (
|
||||||
|
rawStyle !== "primary" &&
|
||||||
|
rawStyle !== "secondary" &&
|
||||||
|
rawStyle !== "danger"
|
||||||
|
) {
|
||||||
|
throw new Error(`actions[${index}].style 不受支持`);
|
||||||
|
}
|
||||||
|
return {
|
||||||
|
id: requiredString(action.id, `actions[${index}].id`, 64),
|
||||||
|
label: requiredString(action.label, `actions[${index}].label`, 40),
|
||||||
|
style: rawStyle,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
function normalizeMessage(argumentsValue: Record<string, unknown>): MessageState {
|
||||||
|
const rawActions = argumentsValue.actions;
|
||||||
|
if (
|
||||||
|
rawActions !== undefined &&
|
||||||
|
(!Array.isArray(rawActions) || rawActions.length > MAX_ACTIONS)
|
||||||
|
) {
|
||||||
|
throw new Error(`actions 必须是数组且最多包含 ${MAX_ACTIONS} 项`);
|
||||||
|
}
|
||||||
|
const actions =
|
||||||
|
Array.isArray(rawActions) && rawActions.length > 0
|
||||||
|
? rawActions.map(normalizeAction)
|
||||||
|
: [
|
||||||
|
{
|
||||||
|
id: "acknowledged",
|
||||||
|
label: "知道了",
|
||||||
|
style: "primary" as const,
|
||||||
|
},
|
||||||
|
];
|
||||||
|
if (new Set(actions.map((action) => action.id)).size !== actions.length) {
|
||||||
|
throw new Error("actions.id 不能重复");
|
||||||
|
}
|
||||||
|
if (
|
||||||
|
argumentsValue.dismissible !== undefined &&
|
||||||
|
typeof argumentsValue.dismissible !== "boolean"
|
||||||
|
) {
|
||||||
|
throw new Error("dismissible 必须是布尔值");
|
||||||
|
}
|
||||||
|
|
||||||
|
return {
|
||||||
|
title:
|
||||||
|
argumentsValue.title === undefined
|
||||||
|
? "提示"
|
||||||
|
: requiredString(argumentsValue.title, "title", 120),
|
||||||
|
message: requiredString(argumentsValue.message, "message", 2_000),
|
||||||
|
actions,
|
||||||
|
dismissible: argumentsValue.dismissible !== false,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
function actionVariant(style: MessageActionStyle) {
|
||||||
|
if (style === "primary") return "default" as const;
|
||||||
|
if (style === "danger") return "destructive" as const;
|
||||||
|
return "outline" as const;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Renders messages requested by the agent through the generic Client Tool
|
||||||
|
* channel. The tool result is held until the user chooses an action.
|
||||||
|
*/
|
||||||
|
export function ClientMessageDialog({ preview }: { preview: VoicePreview }) {
|
||||||
|
const { registerClientTool, status } = preview;
|
||||||
|
const [message, setMessage] = useState<MessageState | null>(null);
|
||||||
|
const pendingRef = useRef<PendingMessage | null>(null);
|
||||||
|
|
||||||
|
const complete = useCallback((action: string) => {
|
||||||
|
const pending = pendingRef.current;
|
||||||
|
if (!pending) return;
|
||||||
|
pendingRef.current = null;
|
||||||
|
setMessage(null);
|
||||||
|
pending.resolve({ action });
|
||||||
|
}, []);
|
||||||
|
|
||||||
|
useEffect(() => {
|
||||||
|
const unregister = registerClientTool(SHOW_MESSAGE_TOOL, (argumentsValue) => {
|
||||||
|
if (pendingRef.current) {
|
||||||
|
throw new Error("已有消息弹窗正在等待用户操作");
|
||||||
|
}
|
||||||
|
const nextMessage = normalizeMessage(argumentsValue);
|
||||||
|
return new Promise<{ action: string }>((resolve, reject) => {
|
||||||
|
pendingRef.current = { resolve, reject };
|
||||||
|
setMessage(nextMessage);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
return () => {
|
||||||
|
unregister();
|
||||||
|
const pending = pendingRef.current;
|
||||||
|
pendingRef.current = null;
|
||||||
|
pending?.reject(new Error("消息弹窗已关闭"));
|
||||||
|
};
|
||||||
|
}, [registerClientTool]);
|
||||||
|
|
||||||
|
useEffect(() => {
|
||||||
|
if (status === "connected" || !pendingRef.current) return;
|
||||||
|
const pending = pendingRef.current;
|
||||||
|
pendingRef.current = null;
|
||||||
|
pending.reject(new Error("会话已结束"));
|
||||||
|
setMessage(null);
|
||||||
|
}, [status]);
|
||||||
|
|
||||||
|
return (
|
||||||
|
<Dialog
|
||||||
|
open={message !== null}
|
||||||
|
onOpenChange={(open) => {
|
||||||
|
if (!open && message?.dismissible) complete("dismissed");
|
||||||
|
}}
|
||||||
|
>
|
||||||
|
<DialogContent
|
||||||
|
className="gap-5 sm:max-w-md"
|
||||||
|
showCloseButton={message?.dismissible ?? true}
|
||||||
|
onEscapeKeyDown={(event) => {
|
||||||
|
if (!message?.dismissible) event.preventDefault();
|
||||||
|
}}
|
||||||
|
onInteractOutside={(event) => {
|
||||||
|
if (!message?.dismissible) event.preventDefault();
|
||||||
|
}}
|
||||||
|
>
|
||||||
|
<DialogHeader className="text-left">
|
||||||
|
<DialogTitle>{message?.title}</DialogTitle>
|
||||||
|
<DialogDescription className="whitespace-pre-wrap leading-6 text-body">
|
||||||
|
{message?.message}
|
||||||
|
</DialogDescription>
|
||||||
|
</DialogHeader>
|
||||||
|
<DialogFooter>
|
||||||
|
{message?.actions.map((action) => (
|
||||||
|
<Button
|
||||||
|
key={action.id}
|
||||||
|
type="button"
|
||||||
|
variant={actionVariant(action.style)}
|
||||||
|
onClick={() => complete(action.id)}
|
||||||
|
>
|
||||||
|
{action.label}
|
||||||
|
</Button>
|
||||||
|
))}
|
||||||
|
</DialogFooter>
|
||||||
|
</DialogContent>
|
||||||
|
</Dialog>
|
||||||
|
);
|
||||||
|
}
|
||||||
@@ -980,7 +980,7 @@ function ClientToolFields({
|
|||||||
)}
|
)}
|
||||||
<p className="text-xs leading-5 text-muted-foreground">
|
<p className="text-xs leading-5 text-muted-foreground">
|
||||||
函数名必须与客户端注册的 handler 名称一致,例如
|
函数名必须与客户端注册的 handler 名称一致,例如
|
||||||
set_photo_button_visible。
|
set_photo_button_visible 或 show_message。
|
||||||
</p>
|
</p>
|
||||||
</div>
|
</div>
|
||||||
);
|
);
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ import {
|
|||||||
SelectTrigger,
|
SelectTrigger,
|
||||||
SelectValue,
|
SelectValue,
|
||||||
} from "@/components/ui/select";
|
} from "@/components/ui/select";
|
||||||
|
import { ClientMessageDialog } from "@/components/client-message-dialog";
|
||||||
import { useCameraPreview } from "@/hooks/use-camera-preview";
|
import { useCameraPreview } from "@/hooks/use-camera-preview";
|
||||||
import { usePhotoCaptureTool } from "@/hooks/use-photo-capture-tool";
|
import { usePhotoCaptureTool } from "@/hooks/use-photo-capture-tool";
|
||||||
import {
|
import {
|
||||||
@@ -483,6 +484,7 @@ export function MobileCallPage({ assistantId }: { assistantId: string }) {
|
|||||||
return (
|
return (
|
||||||
<main className="flex h-dvh w-full items-center justify-center overflow-hidden bg-black lg:p-4">
|
<main className="flex h-dvh w-full items-center justify-center overflow-hidden bg-black lg:p-4">
|
||||||
<audio ref={audioRef} autoPlay playsInline className="hidden" />
|
<audio ref={audioRef} autoPlay playsInline className="hidden" />
|
||||||
|
<ClientMessageDialog preview={preview} />
|
||||||
<div
|
<div
|
||||||
data-testid="mobile-call-viewport"
|
data-testid="mobile-call-viewport"
|
||||||
className="relative isolate h-dvh min-h-80 w-full overflow-hidden bg-[#07101a] text-white lg:h-[calc(100dvh-2rem)] lg:min-h-0 lg:w-auto lg:aspect-[9/16] lg:rounded-[2rem] lg:ring-1 lg:ring-white/15 lg:shadow-2xl"
|
className="relative isolate h-dvh min-h-80 w-full overflow-hidden bg-[#07101a] text-white lg:h-[calc(100dvh-2rem)] lg:min-h-0 lg:w-auto lg:aspect-[9/16] lg:rounded-[2rem] lg:ring-1 lg:ring-white/15 lg:shadow-2xl"
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ export type VoicePreviewStatus = "idle" | "connecting" | "connected" | "failed";
|
|||||||
export type NetworkQuality = "unknown" | "good" | "fair" | "poor";
|
export type NetworkQuality = "unknown" | "good" | "fair" | "poor";
|
||||||
|
|
||||||
const NETWORK_SAMPLE_INTERVAL_MS = 2_000;
|
const NETWORK_SAMPLE_INTERVAL_MS = 2_000;
|
||||||
|
const SESSION_UPDATE_TIMEOUT_MS = 10_000;
|
||||||
|
|
||||||
type NetworkStatsSample = {
|
type NetworkStatsSample = {
|
||||||
packetsSent: number;
|
packetsSent: number;
|
||||||
@@ -46,7 +47,7 @@ export type ChatMessage = {
|
|||||||
};
|
};
|
||||||
|
|
||||||
type AppMessage = Record<string, unknown> & { type?: string };
|
type AppMessage = Record<string, unknown> & { type?: string };
|
||||||
type DynamicVariableValue = string | number | boolean;
|
export type DynamicVariableValue = string | number | boolean;
|
||||||
export type UserInputPart =
|
export type UserInputPart =
|
||||||
| { type: "input_text"; text: string }
|
| { type: "input_text"; text: string }
|
||||||
| {
|
| {
|
||||||
@@ -57,6 +58,15 @@ export type UserInputResult = {
|
|||||||
inputId: string;
|
inputId: string;
|
||||||
status: "accepted";
|
status: "accepted";
|
||||||
};
|
};
|
||||||
|
export type SessionUpdateRequest = {
|
||||||
|
dynamicVariables: Record<string, DynamicVariableValue>;
|
||||||
|
};
|
||||||
|
export type SessionUpdateResult = {
|
||||||
|
updateId: string;
|
||||||
|
status: "accepted";
|
||||||
|
changed: string[];
|
||||||
|
dynamicVariables: Record<string, DynamicVariableValue>;
|
||||||
|
};
|
||||||
export type ClientToolHandler = (
|
export type ClientToolHandler = (
|
||||||
argumentsValue: Record<string, unknown>,
|
argumentsValue: Record<string, unknown>,
|
||||||
) => unknown | Promise<unknown>;
|
) => unknown | Promise<unknown>;
|
||||||
@@ -67,6 +77,12 @@ type PendingUserInput = {
|
|||||||
timeout: number;
|
timeout: number;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
type PendingSessionUpdate = {
|
||||||
|
resolve: (result: SessionUpdateResult) => void;
|
||||||
|
reject: (error: Error) => void;
|
||||||
|
timeout: number;
|
||||||
|
};
|
||||||
|
|
||||||
function publicVariableSnapshot(
|
function publicVariableSnapshot(
|
||||||
value: unknown,
|
value: unknown,
|
||||||
): Record<string, DynamicVariableValue> {
|
): Record<string, DynamicVariableValue> {
|
||||||
@@ -142,6 +158,13 @@ function newInputId(): string {
|
|||||||
return `input_${Date.now()}_${Math.random().toString(16).slice(2)}`;
|
return `input_${Date.now()}_${Math.random().toString(16).slice(2)}`;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function newUpdateId(): string {
|
||||||
|
if (typeof crypto !== "undefined" && "randomUUID" in crypto) {
|
||||||
|
return `update_${crypto.randomUUID()}`;
|
||||||
|
}
|
||||||
|
return `update_${Date.now()}_${Math.random().toString(16).slice(2)}`;
|
||||||
|
}
|
||||||
|
|
||||||
function readNetworkMetrics(report: RTCStatsReport): NetworkMetrics | null {
|
function readNetworkMetrics(report: RTCStatsReport): NetworkMetrics | null {
|
||||||
const outbound: Partial<Record<"audio" | "video", number>> = {};
|
const outbound: Partial<Record<"audio" | "video", number>> = {};
|
||||||
const lost: Partial<Record<"audio" | "video", number>> = {};
|
const lost: Partial<Record<"audio" | "video", number>> = {};
|
||||||
@@ -246,6 +269,9 @@ export function useVoicePreview(
|
|||||||
new Map<string, { timestamp: string }>(),
|
new Map<string, { timestamp: string }>(),
|
||||||
);
|
);
|
||||||
const pendingUserInputsRef = useRef(new Map<string, PendingUserInput>());
|
const pendingUserInputsRef = useRef(new Map<string, PendingUserInput>());
|
||||||
|
const pendingSessionUpdatesRef = useRef(
|
||||||
|
new Map<string, PendingSessionUpdate>(),
|
||||||
|
);
|
||||||
const clientToolHandlersRef = useRef(new Map<string, ClientToolHandler>());
|
const clientToolHandlersRef = useRef(new Map<string, ClientToolHandler>());
|
||||||
const endedByServerRef = useRef(false);
|
const endedByServerRef = useRef(false);
|
||||||
const selectedDeviceIdRef = useRef("");
|
const selectedDeviceIdRef = useRef("");
|
||||||
@@ -306,6 +332,11 @@ export function useVoicePreview(
|
|||||||
reject(new Error("连接已关闭,用户输入未完成"));
|
reject(new Error("连接已关闭,用户输入未完成"));
|
||||||
});
|
});
|
||||||
pendingUserInputsRef.current.clear();
|
pendingUserInputsRef.current.clear();
|
||||||
|
pendingSessionUpdatesRef.current.forEach(({ reject, timeout }) => {
|
||||||
|
window.clearTimeout(timeout);
|
||||||
|
reject(new Error("连接已关闭,会话状态更新未完成"));
|
||||||
|
});
|
||||||
|
pendingSessionUpdatesRef.current.clear();
|
||||||
networkStatsRef.current = null;
|
networkStatsRef.current = null;
|
||||||
onNodeActiveRef.current?.(null);
|
onNodeActiveRef.current?.(null);
|
||||||
}, []);
|
}, []);
|
||||||
@@ -405,6 +436,33 @@ export function useVoicePreview(
|
|||||||
),
|
),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
} else if (
|
||||||
|
msg.type === "session-update-result" &&
|
||||||
|
typeof msg.update_id === "string"
|
||||||
|
) {
|
||||||
|
const pending = pendingSessionUpdatesRef.current.get(msg.update_id);
|
||||||
|
if (!pending) return;
|
||||||
|
pendingSessionUpdatesRef.current.delete(msg.update_id);
|
||||||
|
window.clearTimeout(pending.timeout);
|
||||||
|
if (msg.status === "accepted") {
|
||||||
|
const dynamicVariables = publicVariableSnapshot(msg.dynamic_variables);
|
||||||
|
const changed = Array.isArray(msg.changed)
|
||||||
|
? msg.changed.filter((name): name is string => typeof name === "string")
|
||||||
|
: [];
|
||||||
|
setSessionVariables(dynamicVariables);
|
||||||
|
pending.resolve({
|
||||||
|
updateId: msg.update_id,
|
||||||
|
status: "accepted",
|
||||||
|
changed,
|
||||||
|
dynamicVariables,
|
||||||
|
});
|
||||||
|
} else {
|
||||||
|
pending.reject(
|
||||||
|
new Error(
|
||||||
|
typeof msg.message === "string" ? msg.message : "会话状态更新失败",
|
||||||
|
),
|
||||||
|
);
|
||||||
|
}
|
||||||
} else if (
|
} else if (
|
||||||
msg.type === "assistant-text-start" &&
|
msg.type === "assistant-text-start" &&
|
||||||
typeof msg.turn_id === "string"
|
typeof msg.turn_id === "string"
|
||||||
@@ -754,6 +812,45 @@ export function useVoicePreview(
|
|||||||
[],
|
[],
|
||||||
);
|
);
|
||||||
|
|
||||||
|
const updateSession = useCallback(
|
||||||
|
({
|
||||||
|
dynamicVariables,
|
||||||
|
}: SessionUpdateRequest): Promise<SessionUpdateResult> => {
|
||||||
|
if (Object.keys(dynamicVariables).length === 0) {
|
||||||
|
return Promise.reject(new Error("至少需要更新一个动态变量"));
|
||||||
|
}
|
||||||
|
const transport = transportRef.current;
|
||||||
|
if (!transport || transport.state !== "connected") {
|
||||||
|
return Promise.reject(new Error("当前未连接语音服务"));
|
||||||
|
}
|
||||||
|
const updateId = newUpdateId();
|
||||||
|
return new Promise<SessionUpdateResult>((resolve, reject) => {
|
||||||
|
const timeout = window.setTimeout(() => {
|
||||||
|
pendingSessionUpdatesRef.current.delete(updateId);
|
||||||
|
reject(new Error("等待会话状态更新结果超时"));
|
||||||
|
}, SESSION_UPDATE_TIMEOUT_MS);
|
||||||
|
pendingSessionUpdatesRef.current.set(updateId, {
|
||||||
|
resolve,
|
||||||
|
reject,
|
||||||
|
timeout,
|
||||||
|
});
|
||||||
|
try {
|
||||||
|
transport.sendAppMessage({
|
||||||
|
type: "session-update",
|
||||||
|
schema_version: 1,
|
||||||
|
update_id: updateId,
|
||||||
|
dynamic_variables: dynamicVariables,
|
||||||
|
});
|
||||||
|
} catch (sendError) {
|
||||||
|
window.clearTimeout(timeout);
|
||||||
|
pendingSessionUpdatesRef.current.delete(updateId);
|
||||||
|
reject(new Error(errorMessage(sendError, "发送会话状态更新失败")));
|
||||||
|
}
|
||||||
|
});
|
||||||
|
},
|
||||||
|
[],
|
||||||
|
);
|
||||||
|
|
||||||
const registerClientTool = useCallback(
|
const registerClientTool = useCallback(
|
||||||
(functionName: string, handler: ClientToolHandler): (() => void) => {
|
(functionName: string, handler: ClientToolHandler): (() => void) => {
|
||||||
clientToolHandlersRef.current.set(functionName, handler);
|
clientToolHandlersRef.current.set(functionName, handler);
|
||||||
@@ -789,6 +886,7 @@ export function useVoicePreview(
|
|||||||
supportsOutputSelection,
|
supportsOutputSelection,
|
||||||
sendText,
|
sendText,
|
||||||
sendUserInput,
|
sendUserInput,
|
||||||
|
updateSession,
|
||||||
registerClientTool,
|
registerClientTool,
|
||||||
connect,
|
connect,
|
||||||
replaceVideoStream,
|
replaceVideoStream,
|
||||||
|
|||||||
Reference in New Issue
Block a user