121 lines
6.2 KiB
Python
121 lines
6.2 KiB
Python
"""Regression coverage for Ark video/chat confusion and activation errors."""
|
|
import asyncio
|
|
import json
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest import mock
|
|
|
|
import httpx
|
|
import admin_backend as backend
|
|
import model_gateway as gateway
|
|
import model_protocol as protocol
|
|
|
|
ROOT = "https://ark.cn-beijing.volces.com/api/v3"
|
|
VIDEO = ROOT + "/contents/generations/tasks"
|
|
SEEDANCE = "doubao-seedance-2-5-260628"
|
|
|
|
|
|
def config(base=ROOT, model="ep-chat-test", mode="auto"):
|
|
return backend.model_test_config({
|
|
"AI_PROVIDER_TYPE": "openai", "AI_API_BASE": base,
|
|
"AI_API_KEY": "test-private-key", "AI_MODEL": model,
|
|
"AI_ENDPOINT_MODE": mode,
|
|
}, {})
|
|
|
|
|
|
class ArkChatConfigTest(unittest.TestCase):
|
|
def test_video_configs_are_rejected_without_network_calls(self):
|
|
for base, model, mode in (
|
|
(VIDEO, SEEDANCE, "exact"),
|
|
(VIDEO + "/", "ep-video-test", "exact"),
|
|
(VIDEO, "ep-video-test", "auto"),
|
|
(ROOT, SEEDANCE, "auto"),
|
|
(ROOT, "volcengine/DOUBAO-SEEDANCE-2-5-260628", "auto"),
|
|
):
|
|
with self.subTest(base=base, model=model, mode=mode):
|
|
with mock.patch.object(backend, "_perform_http_request") as request:
|
|
result = backend.test_model_connection(config(base, model, mode))
|
|
request.assert_not_called()
|
|
self.assertFalse(result["ok"])
|
|
self.assertIsNone(result["http_status"])
|
|
self.assertIn("不能用于客服聊天", result["message"])
|
|
self.assertIn("/chat/completions", result["message"])
|
|
self.assertNotIn("test-private-key", json.dumps(result))
|
|
|
|
def test_valid_ark_chat_request_has_expected_url_and_payload(self):
|
|
for base, mode in ((ROOT, "auto"), (ROOT + "/chat/completions", "exact")):
|
|
with self.subTest(mode=mode):
|
|
response = json.dumps({"choices": [{"message": {"content": "OK"}}]}).encode()
|
|
with mock.patch.object(backend, "_perform_http_request", return_value=(200, response)) as request:
|
|
result = backend.test_model_connection(config(base, mode=mode))
|
|
self.assertTrue(result["ok"])
|
|
self.assertEqual(request.call_args.args[0], ROOT + "/chat/completions")
|
|
self.assertEqual(request.call_args.kwargs["payload"]["model"], "ep-chat-test")
|
|
self.assertIn("messages", request.call_args.kwargs["payload"])
|
|
self.assertEqual(protocol.endpoint_url("openai", base, mode), ROOT + "/chat/completions")
|
|
|
|
def test_activation_error_is_not_reported_as_missing_endpoint(self):
|
|
response = json.dumps({"error": {
|
|
"code": "ModelNotOpen",
|
|
"message": "Your account has not activated the model ep-chat-test. test-private-key",
|
|
}}).encode()
|
|
with mock.patch.object(backend, "_perform_http_request", return_value=(404, response)):
|
|
result = backend.test_model_connection(config())
|
|
self.assertFalse(result["ok"])
|
|
self.assertEqual(result["http_status"], 404)
|
|
self.assertIn("尚未开通", result["message"])
|
|
self.assertNotIn("接口地址或模型名称不存在", result["message"])
|
|
self.assertNotIn("test-private-key", json.dumps(result))
|
|
|
|
def test_regular_404_still_reports_missing_endpoint(self):
|
|
with mock.patch.object(backend, "_perform_http_request", return_value=(404, b'{"error":{"message":"not found"}}')):
|
|
result = backend.test_model_connection(config())
|
|
self.assertIn("接口地址或模型名称不存在", result["message"])
|
|
|
|
def test_dify_app_model_label_is_not_treated_as_a_video_model(self):
|
|
self.assertEqual(protocol.chat_config_error("dify", "https://dify.example/v1", SEEDANCE), "")
|
|
payload = protocol.chat_payload("openai", model="ep-chat-test", base_url=ROOT,
|
|
messages=[{"role": "user", "content": "hello"}], max_tokens=100, temperature=0.3)
|
|
self.assertEqual(payload["model"], "ep-chat-test")
|
|
self.assertEqual(protocol.chat_config_error("openai", "https://proxy.example/custom/chat", "chat-model"), "")
|
|
|
|
def test_existing_bad_gateway_outlet_never_calls_upstream(self):
|
|
async def run():
|
|
calls = []
|
|
def handler(request):
|
|
calls.append(request)
|
|
return httpx.Response(200, json={"id": "video-task"})
|
|
outlet = gateway.Outlet(config={"id": "bad", "kind": "openai", "model": SEEDANCE,
|
|
"base_url": VIDEO, "endpoint_mode": "exact"}, gate=asyncio.Semaphore(1))
|
|
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
|
|
result = await gateway.call_outlet(client, outlet, [{"role": "user", "content": "hi"}])
|
|
self.assertEqual(calls, [])
|
|
self.assertEqual(result["text"], "")
|
|
self.assertIn("不能用于客服聊天", result["error"])
|
|
asyncio.run(run())
|
|
|
|
def test_catalog_rejects_enabled_video_but_allows_disabling_old_config(self):
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
db = backend.Database(Path(directory) / "test.db")
|
|
db.initialize("InitialAdmin123")
|
|
uid = db.authenticate("admin", "InitialAdmin123")["id"]
|
|
item = {"id": "bad", "kind": "openai", "base_url": VIDEO,
|
|
"model": SEEDANCE, "endpoint_mode": "exact", "api_key": "test-private-key"}
|
|
with self.assertRaisesRegex(ValueError, "不能用于客服聊天"):
|
|
db.save_model_provider(item, uid, "127.0.0.1")
|
|
self.assertEqual(db.model_providers(), [])
|
|
db.save_model_provider({**item, "enabled": False}, uid, "127.0.0.1")
|
|
saved = db.model_providers()[0]
|
|
self.assertFalse(saved["enabled"])
|
|
self.assertNotIn("test-private-key", json.dumps(saved))
|
|
db.save_model_provider({**item, "model": "ep-chat-test", "base_url": ROOT,
|
|
"endpoint_mode": "auto", "api_key": "", "enabled": True}, uid, "127.0.0.1")
|
|
fixed = db.model_providers(include_secrets=True)[0]
|
|
self.assertEqual(fixed["api_key"], "test-private-key")
|
|
self.assertEqual(fixed["endpoint"], ROOT + "/chat/completions")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|