This commit is contained in:
@@ -0,0 +1,301 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""推理后端抽象接口(backends,Issue #57,PRD 5.6)单元测试。
|
||||
|
||||
覆盖:
|
||||
- 抽象基类不可直接实例化(必须由子类实现四个生命周期方法);
|
||||
- 值对象 BackendCapabilities / BackendHealth / InferResult 的字段与序列化;
|
||||
- LocalBackend / CloudBackend 占位实现的生命周期(load/infer/health/unload)与幂等;
|
||||
- 向后兼容:``generate`` 转发到 ``infer`` 并返回 ``text``;
|
||||
- 能力声明差异(本地出厂内闭环 / 云端出厂外);
|
||||
- 注册表与 ``build_backend`` 的配置驱动构造 + 未知后端报错。
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from abc import ABC
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
import _bootstrap # noqa: F401
|
||||
|
||||
from llm_gateway.backends import ( # noqa: E402
|
||||
BackendCapabilities,
|
||||
BackendHealth,
|
||||
CloudBackend,
|
||||
InferResult,
|
||||
InferenceBackend,
|
||||
LocalBackend,
|
||||
_PlaceholderBackend,
|
||||
build_backend,
|
||||
default_registry,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 抽象基类契约
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class AbstractionContractTest(unittest.TestCase):
|
||||
"""PRD 5.6:InferenceBackend 是抽象接口,业务代码只依赖它。"""
|
||||
|
||||
def test_cannot_instantiate_abstract_base(self):
|
||||
# 缺少四个抽象方法 → 不能实例化
|
||||
with self.assertRaises(TypeError):
|
||||
InferenceBackend() # noqa: E721
|
||||
|
||||
def test_is_abc_subclass(self):
|
||||
self.assertTrue(issubclass(InferenceBackend, ABC))
|
||||
|
||||
def test_required_abstract_methods(self):
|
||||
# PRD 5.6 明列的生命周期动作
|
||||
abstract = InferenceBackend.__abstractmethods__
|
||||
for name in ("load_model", "infer", "health_check", "unload"):
|
||||
self.assertIn(name, abstract)
|
||||
|
||||
def test_concrete_backends_are_inference_backends(self):
|
||||
for cls in (LocalBackend, CloudBackend):
|
||||
self.assertTrue(issubclass(cls, InferenceBackend),
|
||||
f"{cls.__name__} 必须实现 InferenceBackend")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 值对象
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BackendCapabilitiesTest(unittest.TestCase):
|
||||
def test_defaults(self):
|
||||
cap = BackendCapabilities()
|
||||
self.assertFalse(cap.streaming)
|
||||
self.assertIsNone(cap.max_concurrency)
|
||||
self.assertFalse(cap.on_premises)
|
||||
self.assertEqual(cap.modalities, ("text",))
|
||||
|
||||
def test_supports_modality(self):
|
||||
cap = BackendCapabilities(modalities=("text", "image"))
|
||||
self.assertTrue(cap.supports("text"))
|
||||
self.assertTrue(cap.supports("image"))
|
||||
self.assertFalse(cap.supports("audio"))
|
||||
|
||||
def test_to_dict_roundtrip(self):
|
||||
cap = BackendCapabilities(streaming=True, max_concurrency=4,
|
||||
on_premises=False, modalities=("text",))
|
||||
d = cap.to_dict()
|
||||
self.assertEqual(d["streaming"], True)
|
||||
self.assertEqual(d["max_concurrency"], 4)
|
||||
self.assertEqual(d["modalities"], ["text"])
|
||||
|
||||
|
||||
class BackendHealthTest(unittest.TestCase):
|
||||
def test_fields(self):
|
||||
h = BackendHealth(healthy=True, detail="ok")
|
||||
self.assertTrue(h.healthy)
|
||||
self.assertEqual(h.detail, "ok")
|
||||
self.assertTrue(h.checked_at) # 自动生成时间戳
|
||||
|
||||
def test_to_dict(self):
|
||||
d = BackendHealth(healthy=False, detail="down").to_dict()
|
||||
self.assertEqual(d["healthy"], False)
|
||||
self.assertIn("checked_at", d)
|
||||
|
||||
|
||||
class InferResultTest(unittest.TestCase):
|
||||
def test_required_fields(self):
|
||||
r = InferResult(text="hello", backend_name="local-70b")
|
||||
self.assertEqual(r.text, "hello")
|
||||
self.assertEqual(r.backend_name, "local-70b")
|
||||
self.assertIsNone(r.prompt_tokens)
|
||||
|
||||
def test_to_dict(self):
|
||||
r = InferResult(text="a", backend_name="b", model_id="m",
|
||||
prompt_tokens=3, completion_tokens=5)
|
||||
d = r.to_dict()
|
||||
self.assertEqual(d["text"], "a")
|
||||
self.assertEqual(d["prompt_tokens"], 3)
|
||||
self.assertEqual(d["completion_tokens"], 5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 占位实现生命周期
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PlaceholderLifecycleTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.b = LocalBackend()
|
||||
|
||||
def test_health_reflects_load_state(self):
|
||||
# 未加载 → 不健康
|
||||
self.assertFalse(self.b.health_check().healthy)
|
||||
self.b.load_model("local-70b-base")
|
||||
self.assertTrue(self.b.health_check().healthy)
|
||||
|
||||
def test_load_is_idempotent(self):
|
||||
self.b.load_model("local-70b-base")
|
||||
# 重复加载同一 model_id 不报错
|
||||
self.b.load_model("local-70b-base")
|
||||
self.assertTrue(self.b.health_check().healthy)
|
||||
|
||||
def test_infer_lazy_loads_when_not_loaded(self):
|
||||
# 演示态:未显式 load_model 也能 infer(惰性自加载)
|
||||
r = self.b.infer("炉温是多少", context=["SOP-炉温"])
|
||||
self.assertIsInstance(r, InferResult)
|
||||
self.assertEqual(r.backend_name, "local-70b")
|
||||
self.assertIn("炉温是多少", r.text)
|
||||
self.assertIn("[来源: SOP-炉温]", r.text)
|
||||
|
||||
def test_infer_after_explicit_load(self):
|
||||
self.b.load_model("local-70b-base")
|
||||
r = self.b.infer("hello")
|
||||
self.assertEqual(r.model_id, "local-70b-base")
|
||||
self.assertIn("hello", r.text)
|
||||
|
||||
def test_unload_is_idempotent(self):
|
||||
self.b.load_model("local-70b-base")
|
||||
self.b.unload()
|
||||
self.assertFalse(self.b.health_check().healthy)
|
||||
# 未加载再 unload 也不报错
|
||||
self.b.unload()
|
||||
|
||||
def test_echo_context_disabled(self):
|
||||
b = LocalBackend(echo_context=False)
|
||||
b.load_model("m")
|
||||
r = b.infer("q", context=["src1", "src2"])
|
||||
self.assertNotIn("[来源:", r.text)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 向后兼容:generate 转发到 infer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BackwardCompatGenerateTest(unittest.TestCase):
|
||||
def test_generate_returns_text_of_infer(self):
|
||||
b = CloudBackend()
|
||||
b.load_model("cloud-qwen-plus")
|
||||
txt = b.generate("海绵钛是什么", context=["科普手册"])
|
||||
# 与 infer().text 一致
|
||||
self.assertEqual(txt, b.infer("海绵钛是什么", context=["科普手册"]).text)
|
||||
self.assertIn("云端API占位", txt)
|
||||
self.assertIn("[来源: 科普手册]", txt)
|
||||
|
||||
def test_gateway_still_works_with_new_backends(self):
|
||||
# 集成校验:LLMGateway.ask() 经 generate 路径仍正常(不导入失败)。
|
||||
# 复用 test_gateway.py 的模板配置加载 prompts,避免默认空注册表 KeyError。
|
||||
from llm_gateway.dlp import DlpEngine
|
||||
from llm_gateway.gateway import LLMGateway
|
||||
from llm_gateway.prompts import PromptRegistry
|
||||
from llm_gateway.router import SensitivityRouter
|
||||
cfg_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
prompts = PromptRegistry.from_template_config(
|
||||
os.path.join(cfg_dir, "config", "prompts.template.yaml"))
|
||||
router = SensitivityRouter.from_template_config(
|
||||
os.path.join(cfg_dir, "config", "router.template.yaml"))
|
||||
gw = LLMGateway(
|
||||
dlp=DlpEngine(), router=router, prompts=prompts,
|
||||
local=LocalBackend(), cloud=CloudBackend())
|
||||
result = gw.ask("海绵钛是什么", rag_context=["科普手册"])
|
||||
self.assertTrue(result.answer)
|
||||
# 后端占位回显特征仍在(证明走的是新 backends 的 generate 路径)
|
||||
self.assertIn("云端API占位", result.answer)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 能力声明差异(本地 vs 云端)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CapabilitiesDifferenceTest(unittest.TestCase):
|
||||
def test_local_is_on_premises(self):
|
||||
cap = LocalBackend().capabilities
|
||||
self.assertTrue(cap.on_premises)
|
||||
self.assertTrue(cap.streaming)
|
||||
self.assertGreater(cap.max_concurrency, 0)
|
||||
|
||||
def test_cloud_is_off_premises(self):
|
||||
cap = CloudBackend().capabilities
|
||||
self.assertFalse(cap.on_premises)
|
||||
self.assertTrue(cap.streaming)
|
||||
|
||||
def test_local_and_cloud_differ_on_premises(self):
|
||||
# 关键差异:本地出厂内闭环,云端数据出厂
|
||||
self.assertNotEqual(
|
||||
LocalBackend().capabilities.on_premises,
|
||||
CloudBackend().capabilities.on_premises,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 注册表与配置驱动构造
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RegistryTest(unittest.TestCase):
|
||||
def test_default_registry_has_known_backends(self):
|
||||
reg = default_registry()
|
||||
self.assertIn("local-70b", reg)
|
||||
self.assertIn("cloud-api", reg)
|
||||
self.assertIs(reg["local-70b"], LocalBackend)
|
||||
self.assertIs(reg["cloud-api"], CloudBackend)
|
||||
|
||||
def test_build_backend_by_name(self):
|
||||
b = build_backend("local-70b")
|
||||
self.assertIsInstance(b, LocalBackend)
|
||||
self.assertIsInstance(b, InferenceBackend)
|
||||
self.assertEqual(b.name, "local-70b")
|
||||
|
||||
def test_build_unknown_backend_raises_with_hint(self):
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
build_backend("npu-cann") # 尚未实现(#59 才接入)
|
||||
self.assertIn("npu-cann", str(ctx.exception))
|
||||
self.assertIn("local-70b", str(ctx.exception)) # 提示已知项
|
||||
|
||||
def test_build_passes_kwargs(self):
|
||||
b = build_backend("cloud-api", echo_context=False)
|
||||
self.assertIsInstance(b, CloudBackend)
|
||||
self.assertFalse(b.echo_context)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 自定义后端通过实现接口接入(证明「业务代码不感知硬件」)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CustomBackendImplementationTest(unittest.TestCase):
|
||||
"""模拟 #59 昇腾后端:只需实现四个方法即可被当作 InferenceBackend 使用。"""
|
||||
|
||||
def test_custom_backend_satisfies_interface(self):
|
||||
class NpuCannBackend(InferenceBackend):
|
||||
name = "npu-cann"
|
||||
|
||||
def __init__(self):
|
||||
self._loaded = False
|
||||
|
||||
def load_model(self, model_id):
|
||||
self._loaded = True
|
||||
|
||||
def infer(self, prompt, context=None):
|
||||
if not self._loaded:
|
||||
self.load_model("ascend-cann")
|
||||
return InferResult(text=f"[NPU] {prompt}", backend_name=self.name)
|
||||
|
||||
def health_check(self):
|
||||
return BackendHealth(healthy=self._loaded)
|
||||
|
||||
def unload(self):
|
||||
self._loaded = False
|
||||
|
||||
b = NpuCannBackend()
|
||||
self.assertIsInstance(b, InferenceBackend)
|
||||
self.assertFalse(b.health_check().healthy)
|
||||
b.load_model("ascend-cann")
|
||||
self.assertTrue(b.health_check().healthy)
|
||||
self.assertEqual(b.infer("q").text, "[NPU] q")
|
||||
# generate 兼容路径
|
||||
self.assertEqual(b.generate("q", context=[]), "[NPU] q")
|
||||
b.unload()
|
||||
self.assertFalse(b.health_check().healthy)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,239 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""NVIDIA GPU 推理后端(gpu_backend,Issue #58,PRD 5.6)单元测试。
|
||||
|
||||
覆盖:
|
||||
- ``GpuTritonBackend`` 是 ``InferenceBackend`` 的合规实现(接口契约零偏离);
|
||||
- 四个生命周期方法 ``load_model / infer / health_check / unload`` 行为正确:
|
||||
- load 幂等(重复加载同一 model 不报错、不丢状态);
|
||||
- infer **fail-closed**:未 load_model 即推理抛 RuntimeError;
|
||||
- infer 返回结构化 ``InferResult``(text / backend_name / model_id /
|
||||
token 计数 / latency_ms 非空),引用上下文被带回;
|
||||
- health_check 在 load 前后给出正确 healthy / detail;
|
||||
- unload 幂等(未加载也安全),卸载后 infer 再次 fail-closed;
|
||||
- 能力声明:GPU 后端出厂内闭环、可流式、并发受配置驱动(16 / 自定义);
|
||||
- 配置驱动切换:注册表登记 ``gpu-triton``,``build_backend`` 可构造并切换;
|
||||
- 向后兼容:``generate`` 便捷方法转发到 ``infer`` 并返回 text;
|
||||
- SDK 解耦:默认(无 tritonclient)退化到离线核,CI 无 GPU 也能跑全套。
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from abc import ABC
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
import _bootstrap # noqa: F401
|
||||
|
||||
from llm_gateway.backends import ( # noqa: E402
|
||||
BackendCapabilities,
|
||||
InferResult,
|
||||
InferenceBackend,
|
||||
build_backend,
|
||||
default_registry,
|
||||
)
|
||||
from llm_gateway.gpu_backend import ( # noqa: E402
|
||||
GpuTritonBackend,
|
||||
_OfflineKernel,
|
||||
_try_import_tritonclient,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 接口契约
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class GpuBackendContractTest(unittest.TestCase):
|
||||
"""PRD 5.6:GPU 后端必须落地 InferenceBackend 契约。"""
|
||||
|
||||
def test_is_inference_backend(self):
|
||||
self.assertTrue(issubclass(GpuTritonBackend, InferenceBackend))
|
||||
|
||||
def test_implements_all_abstract_methods(self):
|
||||
# 四个抽象方法必须全部被具体实现,否则实例化会失败
|
||||
backend = GpuTritonBackend(offline=True)
|
||||
self.assertIsInstance(backend, InferenceBackend)
|
||||
# 抽象方法集合在子类中应为空
|
||||
self.assertFalse(GpuTritonBackend.__abstractmethods__)
|
||||
|
||||
def test_default_name(self):
|
||||
self.assertEqual(GpuTritonBackend.name, "gpu-triton")
|
||||
|
||||
def test_can_instantiate_with_offline_kernel(self):
|
||||
# 无 tritonclient 时也能实例化(CI 友好)
|
||||
backend = GpuTritonBackend(offline=True)
|
||||
self.assertIsNotNone(backend)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 生命周期:load_model / infer / health_check / unload
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class LifecycleTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.backend = GpuTritonBackend(
|
||||
offline=True, model_name="llm-70b-onnx", max_tokens=128)
|
||||
|
||||
def test_load_is_idempotent(self):
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
self.assertTrue(self.backend._loaded)
|
||||
# 重复加载同一模型不报错、状态保持
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
self.assertTrue(self.backend._loaded)
|
||||
self.assertEqual(self.backend._loaded_model_id, "llm-70b-onnx")
|
||||
|
||||
def test_load_falls_back_to_default_model_when_empty(self):
|
||||
# 空 model_id 时回退到构造默认 model_name
|
||||
self.backend.load_model("")
|
||||
self.assertEqual(self.backend._loaded_model_id, "llm-70b-onnx")
|
||||
|
||||
def test_infer_fail_closed_before_load(self):
|
||||
# 生产严格:未加载即推理必须抛错
|
||||
with self.assertRaises(RuntimeError):
|
||||
self.backend.infer("ping")
|
||||
|
||||
def test_infer_returns_structured_result(self):
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
result = self.backend.infer("海绵钛还蒸能耗?", context=["SOP-A", "国标-B"])
|
||||
self.assertIsInstance(result, InferResult)
|
||||
self.assertEqual(result.backend_name, "gpu-triton")
|
||||
self.assertEqual(result.model_id, "llm-70b-onnx")
|
||||
self.assertIn("海绵钛还蒸能耗?", result.text)
|
||||
# 引用溯源:上下文被带回
|
||||
self.assertIn("[来源: SOP-A]", result.text)
|
||||
self.assertIn("[来源: 国标-B]", result.text)
|
||||
# 审计字段
|
||||
self.assertIsNotNone(result.prompt_tokens)
|
||||
self.assertGreater(result.prompt_tokens, 0)
|
||||
self.assertIsNotNone(result.completion_tokens)
|
||||
self.assertGreater(result.completion_tokens, 0)
|
||||
self.assertIsNotNone(result.latency_ms)
|
||||
self.assertGreaterEqual(result.latency_ms, 0.0)
|
||||
|
||||
def test_health_check_before_load(self):
|
||||
health = self.backend.health_check()
|
||||
self.assertFalse(health.healthy)
|
||||
self.assertIn("loaded=False", health.detail)
|
||||
|
||||
def test_health_check_after_load(self):
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
health = self.backend.health_check()
|
||||
# 离线核 load 后 server_live + model_ready 均为真
|
||||
self.assertTrue(health.healthy)
|
||||
self.assertIn("server_live=True", health.detail)
|
||||
self.assertIn("model_ready=True", health.detail)
|
||||
self.assertIn("loaded=True", health.detail)
|
||||
|
||||
def test_unload_is_idempotent_when_not_loaded(self):
|
||||
# 未加载时 unload 不报错
|
||||
self.backend.unload()
|
||||
self.assertFalse(self.backend._loaded)
|
||||
|
||||
def test_unload_disables_inference(self):
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
self.backend.infer("ok")
|
||||
self.backend.unload()
|
||||
self.assertFalse(self.backend._loaded)
|
||||
# 卸载后再次推理应 fail-closed
|
||||
with self.assertRaises(RuntimeError):
|
||||
self.backend.infer("ok")
|
||||
|
||||
def test_reload_after_unload(self):
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
self.backend.unload()
|
||||
# 可重新加载并推理
|
||||
self.backend.load_model("llm-70b-onnx")
|
||||
result = self.backend.infer("again")
|
||||
self.assertIn("again", result.text)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 能力声明
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CapabilitiesTest(unittest.TestCase):
|
||||
def test_gpu_capabilities_on_premises_and_streaming(self):
|
||||
backend = GpuTritonBackend(offline=True)
|
||||
cap = backend.capabilities
|
||||
self.assertIsInstance(cap, BackendCapabilities)
|
||||
# GPU 后端数据不出厂、支持流式
|
||||
self.assertTrue(cap.on_premises)
|
||||
self.assertTrue(cap.streaming)
|
||||
self.assertIn("text", cap.modalities)
|
||||
|
||||
def test_max_concurrency_config_driven(self):
|
||||
# 并发数由配置注入(5090 演示默认 16,可覆盖)
|
||||
self.assertEqual(
|
||||
GpuTritonBackend(offline=True).capabilities.max_concurrency, 16)
|
||||
self.assertEqual(
|
||||
GpuTritonBackend(offline=True, max_concurrency=32)
|
||||
.capabilities.max_concurrency, 32)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 向后兼容:generate 转发到 infer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BackwardCompatTest(unittest.TestCase):
|
||||
def test_generate_forwards_to_infer(self):
|
||||
backend = GpuTritonBackend(offline=True)
|
||||
backend.load_model("llm-70b-onnx")
|
||||
text = backend.generate("能耗预测", ["SOP-A"])
|
||||
self.assertIsInstance(text, str)
|
||||
self.assertIn("能耗预测", text)
|
||||
self.assertIn("[来源: SOP-A]", text)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 配置驱动切换(注册表 + build_backend)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RegistrySwitchTest(unittest.TestCase):
|
||||
def test_registered_in_default_registry(self):
|
||||
registry = default_registry()
|
||||
self.assertIn("gpu-triton", registry)
|
||||
self.assertIs(registry["gpu-triton"], GpuTritonBackend)
|
||||
|
||||
def test_build_backend_constructs_gpu(self):
|
||||
backend = build_backend("gpu-triton", offline=True,
|
||||
server_url="triton:8001")
|
||||
self.assertIsInstance(backend, GpuTritonBackend)
|
||||
self.assertEqual(backend.server_url, "triton:8001")
|
||||
self.assertEqual(backend.name, "gpu-triton")
|
||||
|
||||
def test_build_backend_unknown_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
build_backend("not-a-backend")
|
||||
|
||||
def test_switch_backend_by_config(self):
|
||||
# 切换后端 = 改 name + 配置,业务代码零改动
|
||||
gpu = build_backend("gpu-triton", offline=True, max_concurrency=32)
|
||||
local = build_backend("local-70b")
|
||||
self.assertNotEqual(gpu.name, local.name)
|
||||
self.assertEqual(gpu.capabilities.max_concurrency, 32)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SDK 解耦:无 tritonclient 时退化到离线核
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SdkDecouplingTest(unittest.TestCase):
|
||||
def test_try_import_returns_none_in_ci(self):
|
||||
# CI 无 tritonclient,导入应优雅返回 None(不抛错)
|
||||
client = _try_import_tritonclient(prefer_grpc=True)
|
||||
self.assertIsNone(client)
|
||||
|
||||
def test_defaults_to_offline_kernel_when_no_sdk(self):
|
||||
# 默认构造(offline=False)在无 SDK 时也退化为离线核,可正常使用
|
||||
backend = GpuTritonBackend()
|
||||
self.assertIsInstance(backend._kernel, _OfflineKernel)
|
||||
backend.load_model("llm-70b-onnx")
|
||||
self.assertTrue(backend.health_check().healthy)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user