323 lines
13 KiB
Python
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()
|