Files
QMDSearch/tests/test_llm_factory.py
kplam 2ab8b56a01 feat: 完成全量功能开发,包括前端管理后台与后端服务优化
此提交实现了完整的知识库管理系统:
1. 新增Vue3 + Antd Vue前端管理后台,包含登录、文档管理、检索、类目设置等完整页面
2. 重构后端LLM调用抽象层,支持Ollama与OpenAI兼容服务动态切换
3. 调整默认嵌入模型配置为本地bge-m3模式
4. 优化入库任务去重逻辑与缓存清理机制
5. 完善Docker镜像构建与docker-compose部署配置
6. 修复多项测试用例与兼容性问题
7. 新增运行时配置API,支持动态调整系统参数
2026-07-31 12:05:25 +08:00

296 lines
12 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""LLM 工厂与客户端单元测试:mock httpx 验证 Ollama / OpenAICompatible 两条路径 + 缓存
通过 monkeypatch 替换 `httpx.AsyncClient` 为伪造的上下文管理器,避免真实网络请求。
"""
import httpx
import pytest
from unittest.mock import AsyncMock
from app.core import runtime_settings as rs
from app.services import llm as llm_mod
from app.services.llm import (
LLMClient,
OllamaLLMClient,
OpenAICompatibleLLMClient,
create_llm_client,
invalidate_llm_client_cache,
)
# ---------------------------------------------------------------------------- #
# httpx 伪造工具
# ---------------------------------------------------------------------------- #
class _FakeResponse:
def __init__(self, status_code: int = 200, json_data: dict | None = None) -> None:
self.status_code = status_code
self._json = json_data or {}
def raise_for_status(self) -> None:
if self.status_code >= 400:
raise httpx.HTTPStatusError(
f"HTTP {self.status_code}", request=httpx.Request("POST", "http://x"), response=self
)
def json(self) -> dict:
return self._json
class _FakeAsyncClient:
"""伪造 httpx.AsyncClient 上下文管理器
用法:把 _FakeAsyncClient.next_response 设为期望响应,构造后 post/get 返回该响应。
每次 __init__ 记录构造参数到 _FakeAsyncClient.last_kwargs。
"""
next_response: _FakeResponse | None = None
last_kwargs: dict | None = None
last_instances: list["_FakeAsyncClient"] = []
def __init__(self, *args, **kwargs) -> None:
self.post = AsyncMock()
self.get = AsyncMock()
if _FakeAsyncClient.next_response is not None:
self.post.return_value = _FakeAsyncClient.next_response
self.get.return_value = _FakeAsyncClient.next_response
_FakeAsyncClient.last_kwargs = kwargs
_FakeAsyncClient.last_instances.append(self)
async def __aenter__(self) -> "_FakeAsyncClient":
return self
async def __aexit__(self, *args) -> bool:
return False
@classmethod
def reset(cls) -> None:
cls.next_response = None
cls.last_kwargs = None
cls.last_instances = []
@pytest.fixture
def patch_httpx(monkeypatch: pytest.MonkeyPatch):
"""替换 httpx.AsyncClient 为伪造客户端"""
_FakeAsyncClient.reset()
monkeypatch.setattr(llm_mod.httpx, "AsyncClient", _FakeAsyncClient)
# ollama 模块内 OllamaClient 也用 httpx.AsyncClient
from app.services import ollama as ollama_mod
monkeypatch.setattr(ollama_mod.httpx, "AsyncClient", _FakeAsyncClient)
yield _FakeAsyncClient
_FakeAsyncClient.reset()
@pytest.fixture
def isolated_settings_path(tmp_path, monkeypatch: pytest.MonkeyPatch):
path = tmp_path / "runtime_settings.json"
monkeypatch.setenv("RUNTIME_SETTINGS_PATH", str(path))
rs._runtime_settings = None
invalidate_llm_client_cache()
yield path
rs._runtime_settings = None
invalidate_llm_client_cache()
# ---------------------------------------------------------------------------- #
# OllamaLLMClient
# ---------------------------------------------------------------------------- #
class TestOllamaLLMClient:
async def test_generate_returns_response_field(self, patch_httpx):
"""Ollama 响应:取 data['response'] 字段"""
patch_httpx.next_response = _FakeResponse(
json_data={"response": "你好世界", "model": "qwen2.5:1.5b"}
)
client = OllamaLLMClient(base_url="http://ollama:11434", model="qwen2.5:1.5b")
out = await client.generate("prompt")
assert out == "你好世界"
# 校验请求 url 与 payload
fake = patch_httpx.last_instances[-1]
fake.post.assert_awaited_once()
call_args = fake.post.await_args
assert call_args.args[0] == "http://ollama:11434/api/generate"
payload = call_args.kwargs["json"]
assert payload["model"] == "qwen2.5:1.5b"
assert payload["prompt"] == "prompt"
assert payload["stream"] is False
async def test_generate_json_mode_adds_format(self, patch_httpx):
patch_httpx.next_response = _FakeResponse(json_data={"response": "{}"})
client = OllamaLLMClient(base_url="http://o:11434", model="m")
await client.generate("p", json_mode=True)
payload = patch_httpx.last_instances[-1].post.await_args.kwargs["json"]
assert payload["format"] == "json"
async def test_is_available_true(self, patch_httpx):
patch_httpx.next_response = _FakeResponse(status_code=200)
client = OllamaLLMClient(base_url="http://o:11434", model="m")
assert await client.is_available() is True
async def test_is_available_false_on_http_error(self, monkeypatch):
"""httpx 抛 HTTPError 时 OllamaClient 内部捕获返回 FalseOllamaLLMClient 透传"""
class _RaisingClient:
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
async def get(self, *args, **kwargs):
raise httpx.HTTPError("conn refused")
from app.services import ollama as ollama_mod
monkeypatch.setattr(ollama_mod.httpx, "AsyncClient", lambda *a, **kw: _RaisingClient())
client = OllamaLLMClient(base_url="http://o:11434", model="m")
assert await client.is_available() is False
# ---------------------------------------------------------------------------- #
# OpenAICompatibleLLMClient
# ---------------------------------------------------------------------------- #
class TestOpenAICompatibleClient:
async def test_generate_returns_choices_content(self, patch_httpx):
patch_httpx.next_response = _FakeResponse(
json_data={"choices": [{"message": {"content": "答案"}}]}
)
client = OpenAICompatibleLLMClient(
base_url="https://api.openai.com/v1", api_key="sk-x", model="gpt-4o-mini"
)
out = await client.generate("hello")
assert out == "答案"
fake = patch_httpx.last_instances[-1]
call_args = fake.post.await_args
assert call_args.args[0] == "https://api.openai.com/v1/chat/completions"
payload = call_args.kwargs["json"]
assert payload["model"] == "gpt-4o-mini"
assert payload["messages"] == [{"role": "user", "content": "hello"}]
assert payload["stream"] is False
# Authorization header
headers = call_args.kwargs["headers"]
assert headers["Authorization"] == "Bearer sk-x"
async def test_generate_json_mode_adds_response_format(self, patch_httpx):
patch_httpx.next_response = _FakeResponse(
json_data={"choices": [{"message": {"content": "{}"}}]}
)
client = OpenAICompatibleLLMClient(
base_url="https://api.openai.com/v1", api_key="sk-x", model="m"
)
await client.generate("p", json_mode=True)
payload = patch_httpx.last_instances[-1].post.await_args.kwargs["json"]
assert payload["response_format"] == {"type": "json_object"}
async def test_generate_empty_choices_returns_empty(self, patch_httpx):
patch_httpx.next_response = _FakeResponse(json_data={"choices": []})
client = OpenAICompatibleLLMClient(
base_url="https://api.openai.com/v1", api_key="sk-x", model="m"
)
assert await client.generate("p") == ""
async def test_is_available_true(self, patch_httpx):
patch_httpx.next_response = _FakeResponse(status_code=200)
client = OpenAICompatibleLLMClient(
base_url="https://api.openai.com/v1", api_key="sk-x", model="m"
)
assert await client.is_available() is True
async def test_is_available_false_on_http_error(self, monkeypatch):
"""httpx 抛 HTTPError 时返回 False"""
client = OpenAICompatibleLLMClient(
base_url="https://api.openai.com/v1", api_key="sk-x", model="m"
)
class _Raiser:
async def __aenter__(self):
raise httpx.HTTPError("conn refused")
async def __aexit__(self, *args):
return False
monkeypatch.setattr(llm_mod.httpx, "AsyncClient", lambda *a, **kw: _Raiser())
assert await client.is_available() is False
# ---------------------------------------------------------------------------- #
# 工厂
# ---------------------------------------------------------------------------- #
class TestFactory:
def test_ollama_provider_returns_ollama_client(self, isolated_settings_path):
"""runtime_settings 默认 provider=ollama,工厂返回 OllamaLLMClient"""
# 确保配置走 ollama(默认即 ollama
client = create_llm_client("summarize", use_cache=False)
assert isinstance(client, OllamaLLMClient)
def test_openai_compatible_provider_returns_openai_client(
self, isolated_settings_path
):
rs.update_runtime_settings(
{"models": {"query": {"provider": "openai_compatible", "api_key": "sk-x", "model": "gpt-4o-mini"}}}
)
client = create_llm_client("query", use_cache=False)
assert isinstance(client, OpenAICompatibleLLMClient)
assert client.model == "gpt-4o-mini"
assert client.api_key == "sk-x"
def test_openai_compatible_falls_back_to_settings(self, isolated_settings_path):
"""openai_compatible 时 base_url/api_key 为空用 settings 默认;model 为空用 gpt-4o-mini"""
# 显式置空 model 以触发工厂默认值(env fallback 会预填 ollama_model
rs.update_runtime_settings(
{"models": {"classify": {"provider": "openai_compatible", "model": ""}}}
)
from app.config import settings
client = create_llm_client("classify", use_cache=False)
assert isinstance(client, OpenAICompatibleLLMClient)
assert client.api_key == settings.openai_api_key
assert client.model == "gpt-4o-mini" # openai_compatible 默认模型
def test_cache_returns_same_instance(self, isolated_settings_path):
"""use_cache=True 时同 purpose 复用单例"""
c1 = create_llm_client("summarize")
c2 = create_llm_client("summarize")
assert c1 is c2
def test_cache_different_purposes_different_instances(self, isolated_settings_path):
c1 = create_llm_client("summarize")
c2 = create_llm_client("query")
assert c1 is not c2
def test_no_cache_returns_new_instance_each_call(self, isolated_settings_path):
c1 = create_llm_client("summarize", use_cache=False)
c2 = create_llm_client("summarize", use_cache=False)
assert c1 is not c2
def test_invalidate_clears_cache(self, isolated_settings_path):
c1 = create_llm_client("summarize")
invalidate_llm_client_cache()
c2 = create_llm_client("summarize")
assert c1 is not c2
def test_invalidate_single_purpose(self, isolated_settings_path):
c_sum = create_llm_client("summarize")
c_qry = create_llm_client("query")
invalidate_llm_client_cache("summarize")
# summarize 已清,query 未清
assert create_llm_client("summarize") is not c_sum
assert create_llm_client("query") is c_qry
def test_client_implements_protocol(self, isolated_settings_path):
"""OllamaLLMClient 与 OpenAICompatibleLLMClient 都满足 LLMClient Protocol"""
c1 = create_llm_client("summarize", use_cache=False)
assert isinstance(c1, LLMClient)
rs.update_runtime_settings(
{"models": {"summarize": {"provider": "openai_compatible", "api_key": "k"}}}
)
c2 = create_llm_client("summarize", use_cache=False)
assert isinstance(c2, LLMClient)