LLM 应用的依赖注入工程实践:解耦 Client、Prompt 和 Tool Registry,让 AI 系统真正可测试可替换

作者按:我们团队有一个 LLM 应用上线后跑了 3 个月,等到要切换模型供应商的时候,才发现 LLM 客户端实例被散落在 17 个文件里,单元测试覆盖率是 0%,因为没法 mock。这篇文章就是为了不让你重蹈覆辙。


为什么 LLM 应用特别容易写成"依赖地狱"

传统 Web 应用的依赖注入(DI)已经是标配------Spring、NestJS、FastAPI 都内置 DI 容器。但 LLM 应用开发者往往从"快速原型"出发,把大模型 SDK 的调用直接写进业务逻辑,或者把 prompt 字符串硬编码在函数里。几个月后,这些决定会让你付出代价:

典型症状 1:Client 紧耦合

python 复制代码
# 反例:client 到处散落
class DocumentSummarizer:
    def summarize(self, text: str) -> str:
        client = LLMClient(api_key=os.getenv("LLM_API_KEY"))  # ← 每次都创建
        response = client.chat.completions.create(
            model="deepseek-chat",
            messages=[{"role": "user", "content": f"Summarize: {text}"}]
        )
        return response.choices[0].message.content

问题:

  • 无法单元测试(必须真实调用 API,每次跑测试都花钱)
  • 切换模型供应商要改 17 处代码
  • 没有统一的 retry/timeout/tracing 配置入口

典型症状 2:Prompt 内嵌在业务代码里

python 复制代码
def classify_intent(user_message: str) -> str:
    prompt = f"""
    你是一个意图分类器。将以下消息分类为:search/book/cancel/other
    消息:{user_message}
    """  # ← prompt 和业务代码耦合
    ...

问题:

  • Prompt A/B 测试需要改代码重新部署
  • 同一个 prompt 模板在多处复制粘贴,改一处漏了其他处
  • 无法对 prompt 做版本管理和灰度

典型症状 3:Tool Registry 的混乱

在 Agent 应用里,tool/function 定义往往被写成全局常量或者在每个 Agent 初始化时重新定义,既不能按租户动态加载,也不能做细粒度的权限控制。


依赖注入在 LLM 应用里的三个核心维度

DI 的本质是"控制反转":谁使用依赖,谁不负责创建依赖。在 LLM 应用里,有三类依赖值得特别处理:

依赖类型 注入什么 不注入的代价
LLM Client 统一的 provider 抽象 切换模型、添加 fallback、做 tracing 都要改业务代码
Prompt Provider 外部可管理的 prompt 模板 Prompt 迭代必须走代码 + 部署流程
Tool Registry 动态可注册的工具集合 Agent 能力固定,无法按权限/场景动态裁剪

下面逐一拆解。


第一维:LLM Client 的依赖注入

定义 Provider 接口

python 复制代码
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import AsyncIterator

@dataclass
class LLMMessage:
    role: str  # "user" | "assistant" | "system"
    content: str

@dataclass
class LLMResponse:
    content: str
    model: str
    prompt_tokens: int
    completion_tokens: int
    finish_reason: str

class LLMProvider(ABC):
    """LLM Client 的统一抽象接口"""

    @abstractmethod
    async def complete(
        self,
        messages: list[LLMMessage],
        *,
        model: str | None = None,
        temperature: float = 0.7,
        max_tokens: int | None = None,
    ) -> LLMResponse:
        ...

    @abstractmethod
    async def stream(
        self,
        messages: list[LLMMessage],
        *,
        model: str | None = None,
    ) -> AsyncIterator[str]:
        ...

实现具体 Provider

python 复制代码
from deepseek_sdk import AsyncDeepSeek  # 或 qwen_sdk 等

class DeepSeekProvider(LLMProvider):  # 以 DeepSeek 为例
    def __init__(self, client: AsyncLLMClient, default_model: str = "deepseek-chat"):
        self._client = client
        self._default_model = default_model

    async def complete(self, messages, *, model=None, temperature=0.7, max_tokens=None):
        response = await self._client.chat.completions.create(
            model=model or self._default_model,
            messages=[{"role": m.role, "content": m.content} for m in messages],
            temperature=temperature,
            max_tokens=max_tokens,
        )
        choice = response.choices[0]
        return LLMResponse(
            content=choice.message.content,
            model=response.model,
            prompt_tokens=response.usage.prompt_tokens,
            completion_tokens=response.usage.completion_tokens,
            finish_reason=choice.finish_reason,
        )

    async def stream(self, messages, *, model=None):
        async for chunk in await self._client.chat.completions.create(
            model=model or self._default_model,
            messages=[{"role": m.role, "content": m.content} for m in messages],
            stream=True,
        ):
            if chunk.choices[0].delta.content:
                yield chunk.choices[0].delta.content

切换到其他国内模型只需实现同一接口:

python 复制代码
from qwen_sdk import AsyncQwen  # 千问 SDK

class QwenProvider(LLMProvider):
    def __init__(self, client: AsyncQwen, default_model: str = "qwen-max"):
        self._client = client
        self._default_model = default_model

    async def complete(self, messages, *, model=None, temperature=0.7, max_tokens=None):
        # 注意:千问 API 的 system message 格式与标准 标准兼容格式略有差异
        system_msgs = [m for m in messages if m.role == "system"]
        user_msgs = [m for m in messages if m.role != "system"]
        system_content = "\n".join(m.content for m in system_msgs) if system_msgs else None

        kwargs = {
            "model": model or self._default_model,
            "messages": [{"role": m.role, "content": m.content} for m in user_msgs],
            "max_tokens": max_tokens or 4096,
        }
        if system_content:
            kwargs["system"] = system_content

        response = await self._client.chat.create(**kwargs)
        return LLMResponse(
            content=response.choices[0].message.content,
            model=response.model,
            prompt_tokens=response.usage.prompt_tokens,
            completion_tokens=response.usage.completion_tokens,
            finish_reason=response.choices[0].finish_reason,
        )

Mock Provider:让单测变可能

python 复制代码
from collections import deque

class MockLLMProvider(LLMProvider):
    """测试专用:预设响应队列"""

    def __init__(self, responses: list[str] | None = None):
        self._queue: deque[str] = deque(responses or [])
        self.calls: list[list[LLMMessage]] = []  # 记录调用历史供断言

    def enqueue(self, *responses: str) -> "MockLLMProvider":
        self._queue.extend(responses)
        return self

    async def complete(self, messages, **kwargs) -> LLMResponse:
        self.calls.append(messages)
        content = self._queue.popleft() if self._queue else "mock response"
        return LLMResponse(
            content=content,
            model="mock-model",
            prompt_tokens=len(str(messages)),
            completion_tokens=len(content),
            finish_reason="stop",
        )

    async def stream(self, messages, **kwargs):
        content = self._queue.popleft() if self._queue else "mock response"
        for char in content:
            yield char

现在单测干净了:

python 复制代码
import pytest
from unittest.mock import MagicMock

@pytest.mark.asyncio
async def test_summarizer_calls_llm_with_correct_prompt():
    mock_provider = MockLLMProvider(["这是一个关于 AI 工程的总结。"])
    summarizer = DocumentSummarizer(llm=mock_provider)  # 注入!

    result = await summarizer.summarize("一篇关于 AI 工程实践的长文...")

    assert len(mock_provider.calls) == 1
    assert "Summarize" in mock_provider.calls[0][-1].content
    assert result == "这是一个关于 AI 工程的总结。"
    # 整个测试:0 API 调用,0 花费,< 5ms

第二维:Prompt Provider 的依赖注入

Prompt 是 LLM 应用的"配置",不应该是代码。把 prompt 作为依赖注入,有几个好处:

  • 支持 A/B 测试:同一份代码,运行时切换不同 prompt
  • 支持热更新:prompt 改动不需要部署
  • 支持多租户:不同客户使用定制化 prompt

Prompt Provider 接口

python 复制代码
from abc import ABC, abstractmethod
from string import Template
from dataclasses import dataclass

@dataclass
class PromptTemplate:
    name: str
    version: str
    template: str  # 使用 {variable} 占位符
    metadata: dict

    def render(self, **kwargs) -> str:
        return self.template.format(**kwargs)

class PromptProvider(ABC):
    @abstractmethod
    async def get(self, name: str, version: str | None = None) -> PromptTemplate:
        ...

    @abstractmethod
    async def list_versions(self, name: str) -> list[str]:
        ...

三种实现:从文件到远程服务

实现 1:YAML 文件(本地开发首选)

python 复制代码
import yaml
from pathlib import Path

class FilePromptProvider(PromptProvider):
    def __init__(self, prompts_dir: Path):
        self._dir = prompts_dir
        self._cache: dict[str, dict[str, PromptTemplate]] = {}

    async def get(self, name: str, version: str | None = None) -> PromptTemplate:
        if name not in self._cache:
            await self._load(name)
        versions = self._cache[name]
        ver = version or max(versions.keys())  # 默认最新版本
        return versions[ver]

    async def _load(self, name: str):
        path = self._dir / f"{name}.yaml"
        data = yaml.safe_load(path.read_text())
        self._cache[name] = {
            v["version"]: PromptTemplate(
                name=name,
                version=v["version"],
                template=v["template"],
                metadata=v.get("metadata", {}),
            )
            for v in data["versions"]
        }

对应的 YAML 文件 prompts/intent_classifier.yaml

yaml 复制代码
versions:
  - version: "v1.0"
    template: |
      将以下用户消息分类为:search/book/cancel/other
      消息:{user_message}
      只返回类别名称,不要解释。
    metadata:
      created_by: team-a
      notes: "初始版本"

  - version: "v1.1"
    template: |
      你是一个意图分类助手。分析用户消息并返回最匹配的意图类别。
      可选类别:search(搜索信息)、book(预订服务)、cancel(取消订单)、other(其他)
      用户消息:{user_message}
      返回格式:{{"intent": "<category>", "confidence": <0.0-1.0>}}
    metadata:
      created_by: team-b
      notes: "添加置信度输出,改善分类准确率 +8%"

实现 2:远程 Prompt 服务(生产多实例)

python 复制代码
import httpx
from functools import lru_cache

class RemotePromptProvider(PromptProvider):
    def __init__(self, base_url: str, api_key: str, ttl_seconds: int = 300):
        self._base_url = base_url
        self._headers = {"Authorization": f"Bearer {api_key}"}
        self._ttl = ttl_seconds
        self._client = httpx.AsyncClient()
        self._cache: dict[str, tuple[PromptTemplate, float]] = {}

    async def get(self, name: str, version: str | None = None) -> PromptTemplate:
        import time
        cache_key = f"{name}:{version or 'latest'}"
        if cache_key in self._cache:
            template, cached_at = self._cache[cache_key]
            if time.time() - cached_at < self._ttl:
                return template

        response = await self._client.get(
            f"{self._base_url}/prompts/{name}",
            params={"version": version} if version else {},
            headers=self._headers,
        )
        response.raise_for_status()
        data = response.json()
        template = PromptTemplate(**data)
        self._cache[cache_key] = (template, time.time())
        return template

在业务代码里使用

python 复制代码
class IntentClassifier:
    def __init__(self, llm: LLMProvider, prompts: PromptProvider):
        self._llm = llm
        self._prompts = prompts

    async def classify(self, user_message: str, *, prompt_version: str | None = None) -> str:
        template = await self._prompts.get("intent_classifier", version=prompt_version)
        rendered = template.render(user_message=user_message)

        response = await self._llm.complete([
            LLMMessage(role="user", content=rendered)
        ])
        return response.content.strip()

A/B 测试的切换只需要传不同的 prompt_version,业务代码零改动。


第三维:Tool Registry 的依赖注入

在 Agent 场景,tool/function 的注册和使用是紧耦合的重灾区。典型反例:

python 复制代码
# 反例:tools 硬编码,全局共享
TOOLS = [
    {
        "type": "function",
        "function": {
            "name": "search_web",
            "description": "Search the web",
            "parameters": {...}
        }
    },
    # ... 全局 30 个 tools,每个用户都能用
]

class Agent:
    async def run(self, message: str):
        response = await llm_client.chat.completions.create(
            model="deepseek-chat",
            tools=TOOLS,  # ← 全局 tools,无法按权限裁剪
            messages=[...]
        )

可注入的 Tool Registry

python 复制代码
from typing import Callable, Any
import inspect

@dataclass
class ToolDefinition:
    name: str
    description: str
    parameters: dict  # JSON Schema
    handler: Callable[..., Any]
    requires_permission: str | None = None  # 权限标签

class ToolRegistry:
    def __init__(self):
        self._tools: dict[str, ToolDefinition] = {}

    def register(
        self,
        name: str,
        description: str,
        parameters: dict,
        requires_permission: str | None = None,
    ):
        """装饰器工厂,用于注册 tool"""
        def decorator(func: Callable) -> Callable:
            self._tools[name] = ToolDefinition(
                name=name,
                description=description,
                parameters=parameters,
                handler=func,
                requires_permission=requires_permission,
            )
            return func
        return decorator

    def get_llm_tools(self, allowed_permissions: set[str] | None = None) -> list[dict]:
        """按权限过滤,返回标准 tools 格式"""
        tools = []
        for tool in self._tools.values():
            if allowed_permissions is not None and tool.requires_permission:
                if tool.requires_permission not in allowed_permissions:
                    continue
            tools.append({
                "type": "function",
                "function": {
                    "name": tool.name,
                    "description": tool.description,
                    "parameters": tool.parameters,
                },
            })
        return tools

    async def execute(self, name: str, arguments: dict) -> Any:
        if name not in self._tools:
            raise KeyError(f"Unknown tool: {name}")
        tool = self._tools[name]
        if inspect.iscoroutinefunction(tool.handler):
            return await tool.handler(**arguments)
        return tool.handler(**arguments)

注册和使用:

python 复制代码
# 注册 tools(通常在应用启动时)
registry = ToolRegistry()

@registry.register(
    name="search_web",
    description="Search the web for current information",
    parameters={
        "type": "object",
        "properties": {
            "query": {"type": "string", "description": "Search query"},
        },
        "required": ["query"],
    },
)
async def search_web(query: str) -> str:
    # 实际搜索实现
    return f"Results for: {query}"

@registry.register(
    name="execute_code",
    description="Execute Python code in a sandbox",
    parameters={...},
    requires_permission="code_execution",  # 高权限 tool
)
async def execute_code(code: str) -> str:
    ...


# Agent 使用时按权限注入
class Agent:
    def __init__(
        self,
        llm: LLMProvider,
        tools: ToolRegistry,
        user_permissions: set[str],
    ):
        self._llm = llm
        self._tools = tools
        self._permissions = user_permissions

    async def run(self, message: str) -> str:
        available_tools = self._tools.get_llm_tools(
            allowed_permissions=self._permissions
        )
        # ... 普通用户看不到 execute_code,无需在 Agent 内部做权限判断

把三个维度串起来:DI 容器

在小项目里,手动注入够用。在中大型项目里,考虑用 DI 容器统一管理生命周期。

用 Python dependency_injector

python 复制代码
from dependency_injector import containers, providers
from deepseek_sdk import AsyncDeepSeek  # 或 qwen_sdk 等
from pathlib import Path

class Container(containers.DeclarativeContainer):
    config = providers.Configuration()

    # LLM Client
    llm_client = providers.Singleton(
        AsyncLLMClient,
        api_key=config.llm.api_key,
    )

    llm_provider = providers.Singleton(
        DeepSeekProvider,
        client=llm_client,
        default_model=config.llm.default_model,
    )

    # Prompt Provider
    prompt_provider = providers.Singleton(
        FilePromptProvider,
        prompts_dir=providers.Factory(Path, config.prompts.dir),
    )

    # Tool Registry
    tool_registry = providers.Singleton(ToolRegistry)

    # 业务组件
    intent_classifier = providers.Factory(
        IntentClassifier,
        llm=llm_provider,
        prompts=prompt_provider,
    )

初始化和使用:

python 复制代码
container = Container()
container.config.from_yaml("config.yaml")

# FastAPI 集成
@app.get("/classify")
async def classify_endpoint(
    message: str,
    classifier: IntentClassifier = Depends(container.intent_classifier),
):
    return {"intent": await classifier.classify(message)}

测试时覆盖依赖:

python 复制代码
def test_with_mock_providers():
    container = Container()
    container.llm_provider.override(MockLLMProvider(["search"]))
    container.prompt_provider.override(MockPromptProvider())

    classifier = container.intent_classifier()
    # 所有依赖都是 mock,测试完全隔离

五个生产陷阱

陷阱 1:Provider 的生命周期错误

python 复制代码
# 反例:每次请求创建新 client(连接池无法复用)
class BadService:
    async def handle(self, text: str):
        provider = DeepSeekProvider(AsyncLLMClient())  # ← 每次请求都 new
        ...

# 正例:Singleton provider,连接池复用
class GoodService:
    def __init__(self, llm: LLMProvider):  # ← Singleton 注入
        self._llm = llm

实测对比:10 QPS 下,每次 new AsyncLLMClient() vs Singleton 复用,P99 延迟相差约 80ms(连接建立成本)。

陷阱 2:Mock Provider 不模拟延迟

python 复制代码
# 单测用 MockProvider 跑通,但上线后因超时崩溃
# 原因:Mock 是同步/即时的,没有测到超时处理逻辑

class RealisticMockProvider(LLMProvider):
    def __init__(self, responses: list[str], latency_ms: int = 0):
        self._responses = deque(responses)
        self._latency_ms = latency_ms

    async def complete(self, messages, **kwargs):
        if self._latency_ms:
            await asyncio.sleep(self._latency_ms / 1000)
        ...

建议在集成测试里使用 latency_ms=2000 来验证超时逻辑。

陷阱 3:Prompt 模板的并发渲染问题

python 复制代码
# 反例:用 string.Template,$ 符号在 prompt 里经常出现
template = Template("Answer in $language: $question")
# 如果 question 包含 "$" 就会出错

# 正例:用 str.format_map + 明确的占位符
template = "Answer in {language}: {question}"
result = template.format_map({"language": "Chinese", "question": question})
# 额外安全:使用 SafeDict 忽略多余/缺失 key
class SafeDict(dict):
    def __missing__(self, key):
        return "{" + key + "}"  # 保留未匹配的占位符

result = template.format_map(SafeDict(language="Chinese"))

陷阱 4:Tool Registry 的循环依赖

当 tool 的 handler 依赖注入了 LLMProvider(比如 search_and_summarize tool 内部调用 LLM),而 LLMProvider 又通过 DI 容器注入了 ToolRegistry,就会出现循环依赖。

解法:tool handler 使用 lazy provider

python 复制代码
from typing import TYPE_CHECKING
if TYPE_CHECKING:
    from myapp.container import Container

class SummarizeTool:
    def __init__(self, llm_factory: Callable[[], LLMProvider]):
        # 不直接持有 LLMProvider,而是持有工厂函数(延迟初始化)
        self._llm_factory = llm_factory
        self._llm: LLMProvider | None = None

    @property
    def _llm(self) -> LLMProvider:
        if self.__llm is None:
            self.__llm = self._llm_factory()
        return self.__llm

陷阱 5:在测试里误用真实 Provider

python 复制代码
# 危险:conftest.py 里没有明确声明测试使用 Mock
# 如果环境变量 OPENAI_API_KEY 存在,测试就会真实调用 API

# 正例:在 conftest.py 里强制 mock
@pytest.fixture(autouse=True)
def ensure_no_real_llm_calls(monkeypatch):
    """防止测试意外调用真实 LLM API"""
    if "OPENAI_API_KEY" in os.environ:
        monkeypatch.setenv("OPENAI_API_KEY", "test-key-should-not-be-used")
    # 或者直接覆盖 container
    container.llm_provider.override(
        MockLLMProvider(["default mock response"])
    )
    yield
    container.llm_provider.reset_override()

完整架构对比

维度 未使用 DI 使用 DI
切换 LLM 供应商 改 N 个文件 改 1 个 container 配置
单元测试 需要真实 API Key 完全 Mock,0 成本
Prompt 热更新 重新部署 远程 Provider 自动刷新
Tool 权限控制 每个 Agent 自己判断 Registry 统一过滤
A/B 测试 代码 if/else 注入不同 PromptProvider
生命周期管理 散乱,内存泄漏风险 Container 统一管理
代码可读性 业务逻辑混杂创建逻辑 业务代码只关心业务

小结:LLM 应用 DI 的最小可行规则

  1. LLM Client 必须是接口,不是具体类------任何直接在业务代码里 import 大模型 SDK 的地方都是技术债。
  2. Prompt 不是字符串常量,是可注入的配置------至少从独立的 YAML/JSON 文件加载,为后续 A/B 测试预留空间。
  3. Tool Registry 按权限注入,不暴露全量------每个 Agent 实例只拿到它能用的 tools。
  4. Mock 是一等公民------不能 mock 的 LLM 客户端,就是不可测试的系统。
  5. DI 容器做生命周期管理 ------Singleton 的 Provider 比 Factory 节省连接建立开销,不要每请求 new。

LLM 应用不是要重新发明一套软件工程哲学------依赖注入这个 30 年前的 pattern 依然成立,只是 AI 应用开发者往往急于跑通 demo 而跳过了它。等到要做生产切换、做测试、做 A/B 的时候,补这个债会很痛。


参考资料

相关推荐
腾渊信息科技公司1 小时前
Spring Boot集成TDengine实战:工业时序数据存储选型与迁移方案
spring boot·后端·tdengine
codeGoogle10 小时前
自研 IM 还是选择第三方 SDK?企业开发者应该如何权衡?
前端·后端·程序员
ltl11 小时前
PagedAttention 与 Continuous Batching
llm
DLYSB_13 小时前
API 网关流量洪峰与突发 CC 攻击:我用 Go 写了个“现场物理防御哨兵”,把故障响应压缩到秒级
开发语言·后端·golang·报警灯
markinmarkin13 小时前
Spring 中Bean 的作用域有哪些?
java·后端·spring
橙子家14 小时前
使用方法 ToDictionary() 来优化查询时间复杂度:O(N*M) -> O(1*M)【C# 基础】
后端
IT_陈寒15 小时前
我又被JavaScript的隐式类型转换坑了
前端·人工智能·后端
用户83562907805115 小时前
Python Word 转 PDF 和 PDF 转 Word 指南
后端·python
用户83562907805115 小时前
如何使用 Python 加密和保护 Word 文档
后端·python