Files
agent-call/tests/test_ai_runtime.py
T

323 lines
13 KiB
Python

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