feat: add session updates and message dialogs
This commit is contained in:
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()
|
||||
Reference in New Issue
Block a user