feat: add configurable Bailian AI cell call runtime

This commit is contained in:
2026-09-14 23:10:39 +08:00
parent 9bade94ada
commit 5dd69b8779
36 changed files with 6455 additions and 30 deletions
+322
View File
@@ -0,0 +1,322 @@
from __future__ import annotations
import copy
import http.client
import json
import os
import tempfile
import threading
import time
import unittest
from pathlib import Path
from agent_call.ai_runtime import (
AIConfigError,
AIProviderError,
ConversationEngine,
MockLLM,
MockTTS,
build_mock_config,
config_digest,
load_prompt,
pcm16_to_alaw,
pcm_to_wav,
render_prompt,
validate_agent_config,
validate_wav,
)
from agent_call.core import (
AgentCallService,
ConflictError,
InMemoryBroker,
NotFoundError,
iso,
utcnow,
)
from agent_call.http import make_server
ROOT = Path(__file__).resolve().parents[1]
class FailingLLM:
def stream(self, messages, config, turn_id, cancelled):
raise AIProviderError(
"LLM_TIMEOUT", "synthetic provider timeout", retryable=True
)
class AIRuntimeTests(unittest.TestCase):
def config(self, version: str = "ai_test_v1") -> dict:
return build_mock_config(version, load_prompt(ROOT / "prompts/test-call.txt"))
def test_prompt_is_utf8_bounded_and_variables_are_explicit(self) -> None:
prompt = "你好,${customer_name}。"
self.assertEqual(
render_prompt(prompt, {"customer_name": "甲"}, ["customer_name"]),
"你好,甲。",
)
with self.assertRaisesRegex(AIConfigError, "not allowed"):
render_prompt(prompt, {"other": "甲"}, ["customer_name"])
with self.assertRaisesRegex(AIConfigError, "missing"):
render_prompt(prompt, {}, ["customer_name"])
with tempfile.NamedTemporaryFile() as handle:
handle.write("坏".encode("gbk"))
handle.flush()
with self.assertRaises(AIConfigError):
load_prompt(handle.name)
def test_immutable_config_rejects_secrets_and_model_voice_mismatch(self) -> None:
config = self.config()
config["llm"]["api_key"] = "must-not-be-stored"
with self.assertRaisesRegex(AIConfigError, "credential or provider URL"):
validate_agent_config(config)
config = self.config()
config["tts"]["voice"] = "vendor-voice"
with self.assertRaisesRegex(AIConfigError, "mock TTS"):
validate_agent_config(config)
def test_real_provider_never_falls_back_to_mock(self) -> None:
config = self.config()
config["llm"]["provider_ref"] = "real-llm"
with self.assertRaisesRegex(AIConfigError, "provider-specific adapter"):
ConversationEngine(config)
def test_opening_and_provider_failure_are_explicit(self) -> None:
config = self.config()
config["conversation"]["opening"] = "配置开场白。"
engine = ConversationEngine(config)
opening = engine.speak(config["conversation"]["opening"])
self.assertEqual(opening["playback_state"], "playback_confirmed")
self.assertGreater(len(opening["audio"]), 0)
failed = ConversationEngine(config, llm=FailingLLM()).run_text("触发错误")
self.assertEqual(failed["status"], "failed")
self.assertEqual(failed["reason_code"], "LLM_TIMEOUT")
def test_text_stream_produces_first_token_audio_and_phone_codec(self) -> None:
engine = ConversationEngine(self.config(), llm=MockLLM(), tts=MockTTS())
result = engine.run_text("请回声测试。")
self.assertEqual(result["status"], "completed")
self.assertTrue(result["text"])
self.assertGreater(
result["audio_bytes"] if "audio_bytes" in result else len(result["audio"]),
0,
)
self.assertIsNotNone(result["llm_first_token_ms"])
self.assertIsNotNone(result["tts_first_audio_ms"])
self.assertEqual(
result["tts_model_evidence"]["configured_model"],
self.config()["tts"]["model"],
)
self.assertFalse(result["tts_model_evidence"]["provider_echoed_model"])
wav = pcm_to_wav(result["audio"])
self.assertEqual(validate_wav(wav)["channels"], 1)
pcma = pcm16_to_alaw(result["audio"])
self.assertEqual(len(pcma), len(result["audio"]) // 4)
def test_audio_stream_uses_input_fingerprint_not_fixed_text(self) -> None:
first = pcm_to_wav(b"\x00\x00" * 1600)
second = pcm_to_wav(b"\x01\x00" * 1600)
engine = ConversationEngine(self.config())
first_result = engine.run_audio(first)
second_result = engine.run_audio(second)
self.assertNotEqual(
first_result["asr_segments"][-1]["text"],
second_result["asr_segments"][-1]["text"],
)
def test_cancel_discards_late_provider_chunks(self) -> None:
engine = ConversationEngine(
self.config(),
llm=MockLLM(chunk_delay_s=0.01, late_chunks_after_cancel=3),
tts=MockTTS(chunk_delay_s=0.01, late_chunks_after_cancel=3),
)
result: dict = {}
worker = threading.Thread(
target=lambda: result.update(engine.run_text("长文本"))
)
worker.start()
time.sleep(0.025)
self.assertIsNotNone(engine.interrupt())
worker.join(timeout=2)
self.assertFalse(worker.is_alive())
self.assertEqual(result.get("status"), "cancelled")
self.assertGreater(result.get("discarded_late_chunks", 0), 0)
def test_version_publish_is_idempotent_and_immutable(self) -> None:
with tempfile.TemporaryDirectory() as directory:
service = AgentCallService(
db_path=Path(directory) / "state.sqlite3",
object_dir=Path(directory) / "objects",
broker=InMemoryBroker(),
)
try:
config = self.config("ai_store_v1")
first = service.publish_agent_version(
"tenant-demo", "ai_store_v1", config
)
again = service.publish_agent_version(
"tenant-demo", "ai_store_v1", config
)
self.assertEqual(first["content_sha256"], again["content_sha256"])
changed = copy.deepcopy(config)
changed["prompt"]["text"] = "different"
with self.assertRaises(ConflictError):
service.publish_agent_version("tenant-demo", "ai_store_v1", changed)
with self.assertRaises(NotFoundError):
service.get_agent_version("tenant-demo", "missing_version")
with self.assertRaises(NotFoundError):
service.get_agent_version("tenant-b", "ai_store_v1")
snapshot = service.get_agent_version("tenant-demo", "ai_store_v1")
self.assertTrue(snapshot["immutable"])
self.assertEqual(
snapshot["content_sha256"], config_digest(snapshot["config"])
)
finally:
service.store.close()
def test_call_persists_selected_version_and_digest(self) -> None:
with tempfile.TemporaryDirectory() as directory:
service = AgentCallService(
db_path=Path(directory) / "state.sqlite3",
object_dir=Path(directory) / "objects",
broker=InMemoryBroker(),
)
try:
config = self.config("ai_call_v1")
service.publish_agent_version("tenant-demo", "ai_call_v1", config)
fixture = json.loads(
(ROOT / "docs/contracts/examples/call.execute.json").read_text(
encoding="utf-8"
)
)
fixture["command_id"] = "ai_call_cmd"
fixture["trace_id"] = "ai_call_trace"
fixture["issued_at"] = iso(utcnow())
fixture["payload"]["execution_id"] = "ai_call_exec"
fixture["payload"]["agent_version_id"] = "ai_call_v1"
service.publish_execute(fixture)
service.wait_for_idle()
call = service.store.one(
"SELECT * FROM calls WHERE execution_id=?", ("ai_call_exec",)
)
self.assertIsNotNone(call)
self.assertEqual(call["agent_version_id"], "ai_call_v1")
self.assertEqual(call["agent_config_sha256"], config_digest(config))
snapshot = service.get_call("tenant-demo", call["call_id"])
self.assertEqual(snapshot["ai"]["turns"], 2)
self.assertEqual(len(snapshot["ai"]["llm_first_token_ms"]), 2)
finally:
service.store.close()
class AIConfigHTTPTests(unittest.TestCase):
def setUp(self) -> None:
self.temp = tempfile.TemporaryDirectory()
self.previous_tokens = os.environ.get("HTTP_TOKENS")
self.previous_issuer = os.environ.get("AI_CONFIG_ISSUER")
self.previous_audience = os.environ.get("AI_CONFIG_AUDIENCE")
os.environ["AI_CONFIG_ISSUER"] = "issuer-ai-test"
os.environ["AI_CONFIG_AUDIENCE"] = "audience-ai-test"
os.environ["HTTP_TOKENS"] = json.dumps(
{
"publish": {
"tenant_ids": ["tenant-demo"],
"scopes": ["ai.config.publish"],
"auth_domain": "ai-config",
"issuer": "issuer-ai-test",
"audience": "audience-ai-test",
},
"read": {
"tenant_ids": ["tenant-demo"],
"scopes": ["ai.config.read"],
"auth_domain": "ai-config",
"issuer": "issuer-ai-test",
"audience": "audience-ai-test",
},
"normal": {"tenant_ids": ["tenant-demo"], "scopes": ["outbound.read"]},
"wrong-domain": {
"tenant_ids": ["tenant-demo"],
"scopes": ["ai.config.read"],
"auth_domain": "ai-config",
"issuer": "wrong-issuer",
"audience": "audience-ai-test",
},
}
)
self.service = AgentCallService(
db_path=Path(self.temp.name) / "state.sqlite3",
object_dir=Path(self.temp.name) / "objects",
broker=InMemoryBroker(),
)
self.server = make_server(self.service, "127.0.0.1", 0)
self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
self.thread.start()
self.port = self.server.server_address[1]
def tearDown(self) -> None:
self.server.shutdown()
self.server.server_close()
self.thread.join(timeout=2)
self.service.stop()
self.service.store.close()
if self.previous_tokens is None:
os.environ.pop("HTTP_TOKENS", None)
else:
os.environ["HTTP_TOKENS"] = self.previous_tokens
for name, previous in (
("AI_CONFIG_ISSUER", self.previous_issuer),
("AI_CONFIG_AUDIENCE", self.previous_audience),
):
if previous is None:
os.environ.pop(name, None)
else:
os.environ[name] = previous
self.temp.cleanup()
def request(
self, method: str, path: str, token: str, body: dict | None = None
) -> tuple[int, dict]:
headers = {
"Authorization": f"Bearer {token}",
"X-Tenant-ID": "tenant-demo",
"X-Request-ID": "ai-http-test",
}
encoded = None
if body is not None:
encoded = json.dumps(body, ensure_ascii=False).encode("utf-8")
headers["Content-Type"] = "application/json"
connection = http.client.HTTPConnection("127.0.0.1", self.port, timeout=3)
try:
connection.request(method, path, body=encoded, headers=headers)
response = connection.getresponse()
return response.status, json.loads(response.read().decode("utf-8"))
finally:
connection.close()
def test_config_scope_is_separate_from_normal_control_scope(self) -> None:
config = build_mock_config("http_ai_v1", "test")
body = {"agent_version_id": "http_ai_v1", "config": config}
status, receipt = self.request(
"POST", "/internal/v1/ai/agent-versions", "publish", body
)
self.assertEqual(status, 201)
self.assertTrue(receipt["immutable"])
status, snapshot = self.request(
"GET", "/internal/v1/ai/agent-versions/http_ai_v1", "read"
)
self.assertEqual(status, 200)
self.assertEqual(snapshot["config"]["agent_version_id"], "http_ai_v1")
status, denied = self.request(
"POST", "/internal/v1/ai/agent-versions", "normal", body
)
self.assertEqual(status, 403)
self.assertEqual(denied["code"], "FORBIDDEN")
status, denied = self.request(
"GET", "/internal/v1/ai/agent-versions/http_ai_v1", "wrong-domain"
)
self.assertEqual(status, 401)
self.assertEqual(denied["code"], "UNAUTHORIZED")
if __name__ == "__main__":
unittest.main()
+52
View File
@@ -0,0 +1,52 @@
import json
import os
import queue
import unittest
from unittest.mock import patch
from agent_call.bailian import (
BailianTTS,
_BailianCosyVoiceCallback,
build_bailian_config,
)
class BailianAdapterTests(unittest.TestCase):
def test_qwen_realtime_url_is_derived_from_asr_endpoint(self) -> None:
adapter = BailianTTS("wss://example.test/api-ws/v1/inference/", "secret")
self.assertEqual(
adapter._url_for_qwen_realtime(),
"wss://example.test/api-ws/v1/realtime",
)
def test_config_selects_cosyvoice_for_custom_voice(self) -> None:
with patch.dict(
os.environ,
{
"BAILIAN_TTS_VOICE": "cosyvoice-v3.5-plus-bailian-test",
"BAILIAN_TTS_MODEL": "",
},
clear=False,
):
config = build_bailian_config("agent_test", "回答收到。")
self.assertEqual(config["tts"]["model"], "cosyvoice-v3.5-plus")
self.assertEqual(config["asr"]["provider_ref"], "bailian")
self.assertEqual(config["asr"]["model"], "fun-asr-realtime")
def test_cosyvoice_callback_records_returned_model_when_present(self) -> None:
callback = _BailianCosyVoiceCallback(4)
callback.on_event(json.dumps({"header": {"model": "cosyvoice-test"}}))
self.assertEqual(callback.provider_model, "cosyvoice-test")
def test_cosyvoice_callback_is_bounded(self) -> None:
callback = _BailianCosyVoiceCallback(1)
callback.on_data(b"\x00\x00")
callback.on_complete()
self.assertTrue(callback.overflowed)
self.assertEqual(callback.events.get_nowait()["type"], "__audio__")
with self.assertRaises(queue.Empty):
callback.events.get_nowait()
if __name__ == "__main__":
unittest.main()
+13
View File
@@ -48,6 +48,19 @@ class ContractTests(unittest.TestCase):
"/internal/v1/outbound/recording-uploads/{upload_id}/complete:", saas
)
def test_ai_config_contract_and_example(self) -> None:
schema = read_json(ROOT / "docs/contracts/ai-config.schema.json")
example = read_json(ROOT / "docs/contracts/examples/agent-version.json")
self.assertEqual(list(Draft202012Validator(schema).iter_errors(example)), [])
openapi = (ROOT / "docs/contracts/ai-config.openapi.yaml").read_text(
encoding="utf-8"
)
self.assertIn("openapi: 3.1.0", openapi)
self.assertIn("ai.config.publish", openapi)
self.assertIn("ai.config.read", openapi)
self.assertIn("/internal/v1/ai/agent-versions:", openapi)
self.assertNotIn("api_key", json.dumps(example, ensure_ascii=False))
def test_command_fixture_and_invalid_version(self) -> None:
schema = read_json(ROOT / "docs/contracts/mq.schema.json")
fixture = read_json(ROOT / "docs/contracts/examples/call.execute.json")
+175
View File
@@ -0,0 +1,175 @@
from __future__ import annotations
import socket
import struct
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from typing import cast
from agent_call.real_cell import (
CellCallConfig,
CellCallError,
CellExecutionLedger,
RealCellCall,
RealCellWorker,
RTPMedia,
alaw_to_pcm16,
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)
class RealCellTests(unittest.TestCase):
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_cell_config_rejects_non_pcma_and_bad_port(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)
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()