Files
agent-call/tests/test_real_cell.py
rogee 758c0a29e6
management-images / build-and-publish (push) Successful in 34s
fix RTP packet pacing
2026-09-17 18:46:51 +08:00

583 lines
21 KiB
Python

from __future__ import annotations
import io
import json
import socket
import struct
import tempfile
import unittest
import wave
from pathlib import Path
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import patch
from agent_call.real_cell import (
CellCallConfig,
CellCallError,
CellExecutionLedger,
CellRoute,
RealCellCall,
RealCellWorker,
RTPMedia,
alaw_to_pcm16,
load_cell_routes,
voice_level,
)
class FakeBroker:
def __init__(self, messages: list[dict]) -> None:
self.messages = list(messages)
self.acked: list[dict] = []
self.requeued: list[dict] = []
self.rejected: list[dict] = []
self.published: list[dict] = []
def declare_tenant(self, tenant_key: str) -> None:
self.tenant_key = tenant_key
def consume(self, _queue: str) -> dict | None:
return self.messages.pop(0) if self.messages else None
def ack(self, message: dict) -> None:
self.acked.append(message)
def requeue(self, message: dict) -> None:
self.requeued.append(message)
def reject(self, message: dict) -> None:
self.rejected.append(message)
def publish(
self, _exchange: str, _route: str, body: dict, **_kwargs: object
) -> None:
self.published.append(body)
def wav_bytes(pcm: bytes) -> bytes:
output = io.BytesIO()
with wave.open(output, "wb") as source:
source.setnchannels(1)
source.setsampwidth(2)
source.setframerate(8000)
source.writeframes(pcm)
return output.getvalue()
class RealCellTests(unittest.TestCase):
def _build_lifecycle_call(
self, order: list[str], answer_error: CellCallError | None = None
) -> tuple[RealCellCall, Any]:
config = CellCallConfig(
"http://127.0.0.1:8088", "u", "p", event_timeout_s=0.1
)
engine = SimpleNamespace(
config={
"agent_version_id": "agent_v1",
"conversation": {
"opening": "hello",
"max_turns": 1,
"allow_interrupt": False,
},
},
history=[],
)
engine.speak = lambda _text, _turn_id: (
order.append("opening") or {"audio": b""}
)
call = RealCellCall(config, cast(Any, engine), cast(Any, SimpleNamespace()))
class FakeARI:
def __init__(self) -> None:
self.record_requests = 0
def events(self, _app: str) -> Any:
return SimpleNamespace(close=lambda: None)
def request(
self, method: str, resource: str, _params=None, **_kwargs: object
) -> dict[str, str] | dict:
if method == "POST" and resource.endswith("/record"):
self.record_requests += 1
order.append("record")
if resource == "channels/externalMedia":
return {"id": "external"}
if resource == "channels":
return {"id": "target"}
return {}
fake_ari = FakeARI()
call_any = cast(Any, call)
call_any.ari = fake_ari
call_any._start_event_reader = lambda: None
call_any._try_set_external_media_peer = lambda: False
call_any._add_when_ready = lambda _channel_id: True
def wait_answer(_result: Any) -> None:
order.append("answer")
if answer_error is not None:
raise answer_error
call_any._wait_answer = wait_answer
call_any._play = lambda _audio: order.append("play")
call_any._capture_turn = lambda _timeout: (
order.append("capture") or b""
)
call_any._finish_recording = lambda: None
call_any._cleanup = lambda: None
return call, fake_ari
def test_recording_starts_after_answer_before_opening_tts(self) -> None:
order: list[str] = []
call, fake_ari = self._build_lifecycle_call(order)
try:
result = call.start_authorized_call("123")
finally:
if call.media is not None:
call.media.close()
self.assertTrue(result.connected)
self.assertEqual(fake_ari.record_requests, 1)
self.assertLess(order.index("answer"), order.index("record"))
self.assertLess(order.index("record"), order.index("opening"))
def test_answer_failure_does_not_start_recording(self) -> None:
order: list[str] = []
call, fake_ari = self._build_lifecycle_call(
order, CellCallError("CALL_NOT_ANSWERED", "not answered")
)
try:
result = call.start_authorized_call("123")
finally:
if call.media is not None:
call.media.close()
self.assertFalse(result.connected)
self.assertEqual(fake_ari.record_requests, 0)
self.assertNotIn("opening", order)
def test_rtp_send_keeps_packet_clock_across_frames(self) -> None:
media = RTPMedia("127.0.0.1", 0)
receiver = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
receiver.bind(("127.0.0.1", 0))
media.peer = ("127.0.0.1", receiver.getsockname()[1])
clock = [0.0]
sleeps: list[float] = []
def monotonic() -> float:
return clock[0]
def sleep(seconds: float) -> None:
sleeps.append(seconds)
clock[0] += seconds
try:
with (
patch("agent_call.real_cell.time.monotonic", monotonic),
patch("agent_call.real_cell.time.sleep", sleep),
):
media.send_pcm16(bytes([0, 0]) * 160, 8000)
media.send_pcm16(bytes([0, 0]) * 160, 8000)
self.assertEqual(sleeps, [0.02])
self.assertEqual(len(receiver.recvfrom(512)[0]), 172)
self.assertEqual(len(receiver.recvfrom(512)[0]), 172)
finally:
media.close()
receiver.close()
def test_pcma_decode_and_rtp_payload(self) -> None:
media = RTPMedia("127.0.0.1", 0)
receiver = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
receiver.bind(("127.0.0.1", 0))
try:
media.peer = ("127.0.0.1", receiver.getsockname()[1])
sent = media.send_pcm16(b"\x00\x00" * 160, 8000)
packet, _ = receiver.recvfrom(512)
self.assertEqual(sent, 160)
self.assertEqual(len(packet), 172)
self.assertEqual(packet[1] & 0x7F, 8)
self.assertEqual(packet[12:], b"\xd5" * 160)
self.assertEqual(len(alaw_to_pcm16(packet[12:])), 320)
self.assertEqual(voice_level(alaw_to_pcm16(packet[12:])), 8)
finally:
media.close()
receiver.close()
def test_rtp_payload_handles_extension_and_padding(self) -> None:
payload = b"abc"
packet = struct.pack("!BBHII", 0xB0, 8, 1, 2, 3)
packet += struct.pack("!HH", 0, 1) + b"xxxx" + payload + b"\x00\x00\x03"
self.assertEqual(RTPMedia._payload(packet), payload)
self.assertIsNone(RTPMedia._payload(b"bad"))
def test_nonblocking_receive_returns_no_packet(self) -> None:
media = RTPMedia("127.0.0.1", 0)
try:
self.assertIsNone(media.receive(0.0))
finally:
media.close()
def test_rtp_receive_filters_payload_peer_and_ssrc(self) -> None:
media = RTPMedia("127.0.0.1", 0)
sender = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
sender.bind(("127.0.0.1", 0))
media.peer = ("127.0.0.1", sender.getsockname()[1])
def packet(payload_type: int, ssrc: int, payload: bytes) -> bytes:
return struct.pack("!BBHII", 0x80, payload_type, 1, 0, ssrc) + payload
try:
sender.sendto(packet(0, 1, b"wrong-pt"), media.address)
self.assertIsNone(media.receive(0.2))
sender.sendto(packet(8, 1, b"voice"), media.address)
self.assertEqual(media.receive(0.2), b"voice")
sender.sendto(packet(8, 2, b"wrong-ssrc"), media.address)
self.assertIsNone(media.receive(0.2))
finally:
sender.close()
media.close()
def test_external_media_peer_comes_from_ari(self) -> None:
config = CellCallConfig("http://127.0.0.1:8088", "u", "p")
call = cast(
Any,
RealCellCall(
config,
cast(Any, SimpleNamespace(config={})),
cast(Any, SimpleNamespace()),
),
)
media = RTPMedia("127.0.0.1", 0)
calls: list[tuple[str, str, dict | None]] = []
class FakeARI:
def request(self, method: str, resource: str, params=None, **_kwargs):
calls.append((method, resource, params))
if params and params.get("variable") == "UNICASTRTP_LOCAL_ADDRESS":
return {"value": "127.0.0.1"}
return {"value": "12345"}
call.media = media
call.ari = FakeARI()
try:
call._set_external_media_peer("external-1")
self.assertEqual(media.peer, ("127.0.0.1", 12345))
self.assertEqual(len(calls), 2)
finally:
media.close()
def test_add_when_ready_marks_only_after_success(self) -> None:
config = CellCallConfig("http://127.0.0.1:8088", "u", "p")
call = cast(
Any,
RealCellCall(
config,
cast(Any, SimpleNamespace(config={})),
cast(Any, SimpleNamespace()),
),
)
class FakeARI:
def __init__(self) -> None:
self.attempts = 0
def request(self, method: str, _resource: str, _params=None, **_kwargs):
if method == "POST":
self.attempts += 1
if self.attempts == 1:
raise CellCallError("ARI_HTTP_422", "not in Stasis", True)
return {}
fake_ari = FakeARI()
call.ari = fake_ari
self.assertTrue(call._add_when_ready("target"))
self.assertEqual(fake_ari.attempts, 2)
self.assertIn("target", call._known_channels)
def test_wait_answer_requires_bridge_membership(self) -> None:
config = CellCallConfig("http://127.0.0.1:8088", "u", "p", event_timeout_s=1.0)
call = cast(
Any,
RealCellCall(
config,
cast(Any, SimpleNamespace(config={})),
cast(Any, SimpleNamespace()),
),
)
call.target_channel_id = "target"
call.external_channel_id = "external"
call.media = RTPMedia("127.0.0.1", 0)
call.media.peer = ("127.0.0.1", 12345)
bridge_members = [{"external"}, {"external", "target"}]
class FakeARI:
def request(self, method: str, resource: str, _params=None, **_kwargs):
if method == "GET" and resource == f"bridges/{call.bridge_id}":
return {"channels": list(bridge_members.pop(0))}
return {}
call.ari = FakeARI()
call.events.put(
{"type": "ChannelStateChange", "channel": {"id": "target", "state": "Up"}}
)
call.events.put({"type": "StasisStart", "channel": {"id": "target"}})
try:
call._wait_answer(SimpleNamespace())
self.assertEqual(bridge_members, [])
finally:
call.media.close()
def test_finish_recording_stops_then_reads_stored_file(self) -> None:
class FakeARI:
def __init__(self) -> None:
self.calls: list[tuple[str, str]] = []
def request(self, method: str, resource: str, **_kwargs):
self.calls.append((method, resource))
return {} if method == "POST" else wav_bytes(b"\x00\x00")
with tempfile.TemporaryDirectory() as directory:
config = CellCallConfig(
"http://127.0.0.1:8088", "u", "p", recording_dir=directory
)
call = cast(
Any,
RealCellCall(
config,
cast(Any, SimpleNamespace(config={})),
cast(Any, SimpleNamespace()),
),
)
fake_ari = FakeARI()
call.ari = fake_ari
path = call._finish_recording()
self.assertIsNotNone(path)
self.assertEqual(fake_ari.calls[0][0], "POST")
self.assertIn("/stop", fake_ari.calls[0][1])
self.assertNotIn(
("DELETE", fake_ari.calls[0][1].rsplit("/stop", 1)[0]), fake_ari.calls
)
def test_finish_recording_rejects_zero_frame_wav(self) -> None:
class FakeARI:
def request(self, method: str, resource: str, **_kwargs):
return {} if method == "POST" else wav_bytes(b"")
with tempfile.TemporaryDirectory() as directory:
config = CellCallConfig(
"http://127.0.0.1:8088", "u", "p", recording_dir=directory
)
call = cast(
Any,
RealCellCall(
config,
cast(Any, SimpleNamespace(config={})),
cast(Any, SimpleNamespace()),
),
)
call.ari = FakeARI()
self.assertIsNone(call._finish_recording())
self.assertFalse(
(Path(directory) / f"{call.recording_name}.wav").exists()
)
def test_cell_config_rejects_non_pcma_bad_port_and_prefix(self) -> None:
with self.assertRaises(CellCallError):
CellCallConfig("http://127.0.0.1:8088", "u", "p", rtp_format="ulaw")
with self.assertRaises(CellCallError):
CellCallConfig("http://127.0.0.1:8088", "u", "p", rtp_bind_port=65536)
with self.assertRaises(CellCallError):
CellCallConfig("http://127.0.0.1:8088", "u", "p", dial_prefix="7089+")
self.assertEqual(
CellCallConfig(
"http://127.0.0.1:8088", "u", "p", dial_prefix=""
).dial_prefix,
"",
)
self.assertEqual(
CellCallConfig(
"http://127.0.0.1:8088", "u", "p", dial_prefix="mka755"
).dial_prefix,
"mka755",
)
def test_route_map_binds_task_policy_to_trusted_trunk(self) -> None:
routes = load_cell_routes(
json.dumps(
{
"route-a": {
"caller_profile_id": "caller-a",
"trunk_id": "provider-second",
"caller_id": "mbkq",
"dial_prefix": "",
}
}
)
)
self.assertEqual(routes["route-a"].trunk_id, "provider-second")
self.assertEqual(routes["route-a"].caller_id, "mbkq")
with self.assertRaises(CellCallError):
load_cell_routes(
json.dumps(
{
"route-a": {
"caller_profile_id": "caller-a",
"trunk_id": "provider-second",
"caller_id": "mbkq\nspoof",
"dial_prefix": "",
}
}
)
)
def test_worker_selects_installed_route_without_payload_sip_values(self) -> None:
command = {
"body": {
"command_type": "call.execute",
"tenant_id": "tenant-demo",
"tenant_key": "tenant-key",
"payload": {
"execution_id": "exec-route-1",
"agent_version_id": "agent_v1",
"callee": "18625770806",
"route_policy_id": "route-a",
"caller_profile_id": "caller-a",
},
}
}
broker = FakeBroker([command])
with tempfile.TemporaryDirectory() as directory:
ledger = CellExecutionLedger(Path(directory) / "ledger.sqlite3")
fake_result = SimpleNamespace(
as_dict=lambda: {
"call_id": "call-route",
"status": "failed",
"turns": [],
}
)
class FakeExecutor:
def __init__(self) -> None:
self.engine = SimpleNamespace(
config={"agent_version_id": "agent_v1"}
)
self.calls: list[tuple[str, CellRoute | None]] = []
def start_authorized_call(
self, callee: str, route: CellRoute | None = None
) -> SimpleNamespace:
self.calls.append((callee, route))
return fake_result
fake_executor = FakeExecutor()
worker = RealCellWorker(
broker,
"tenant-key",
ledger,
cast(RealCellCall, fake_executor),
routes={
"route-a": CellRoute(
"route-a", "caller-a", "provider-second", "mbkq", ""
)
},
)
event = worker.process_once()
if event is None:
self.fail("worker did not publish a call.finished event")
self.assertEqual(fake_executor.calls[0][0], "18625770806")
selected_route = fake_executor.calls[0][1]
assert selected_route is not None
self.assertEqual(selected_route.trunk_id, "provider-second")
def test_ledger_marks_in_progress_as_in_doubt(self) -> None:
with tempfile.TemporaryDirectory() as directory:
ledger = CellExecutionLedger(Path(directory) / "ledger.sqlite3")
first = ledger.claim("exec-1", "15003164745")
second = ledger.claim("exec-1", "15003164745")
self.assertTrue(first["claimed"])
self.assertTrue(second["in_doubt"])
def test_worker_uses_queue_and_never_redials_terminal_execution(self) -> None:
command = {
"body": {
"command_type": "call.execute",
"tenant_id": "tenant-demo",
"tenant_key": "tenant-key",
"payload": {
"execution_id": "exec-1",
"agent_version_id": "agent_v1",
"callee": "15003164745",
},
}
}
broker = FakeBroker([command, command])
with tempfile.TemporaryDirectory() as directory:
ledger = CellExecutionLedger(Path(directory) / "ledger.sqlite3")
fake_result = SimpleNamespace(
as_dict=lambda: {
"call_id": "call-1",
"status": "failed",
"reason_code": "CUSTOMER_SILENT",
"recording_path": "/private/recording.wav",
}
)
class FakeExecutor:
def __init__(self) -> None:
self.engine = SimpleNamespace(
config={"agent_version_id": "agent_v1"}
)
self.calls: list[str] = []
def start_authorized_call(self, callee: str) -> SimpleNamespace:
self.calls.append(callee)
return fake_result
fake_executor = FakeExecutor()
executor = cast(RealCellCall, fake_executor)
worker = RealCellWorker(broker, "tenant-key", ledger, executor)
first = worker.process_once()
second = worker.process_once()
if first is None or second is None:
self.fail("worker did not publish a call.finished event")
self.assertEqual(first["event_type"], "call.finished")
self.assertEqual(second["event_type"], "call.finished")
self.assertEqual(fake_executor.calls, ["15003164745"])
self.assertEqual(len(broker.acked), 2)
self.assertNotIn("recording_path", broker.published[0]["payload"])
def test_worker_rejects_wrong_tenant_route(self) -> None:
broker = FakeBroker(
[
{
"body": {
"command_type": "call.execute",
"tenant_id": "tenant-demo",
"tenant_key": "other-tenant",
"payload": {
"execution_id": "exec-1",
"agent_version_id": "agent_v1",
"callee": "15003164745",
},
}
}
]
)
with tempfile.TemporaryDirectory() as directory:
ledger = CellExecutionLedger(Path(directory) / "ledger.sqlite3")
executor = cast(
RealCellCall,
SimpleNamespace(
engine=SimpleNamespace(config={"agent_version_id": "agent_v1"})
),
)
worker = RealCellWorker(broker, "tenant-key", ledger, executor)
result = worker.process_once()
if result is None:
self.fail("worker returned no command result")
self.assertEqual(result["reason_code"], "COMMAND_INVALID")
self.assertEqual(len(broker.rejected), 1)
if __name__ == "__main__":
unittest.main()