nex_docus/backend/tests/test_model_and_vector_confi...

222 lines
7.9 KiB
Python

import tempfile
import unittest
from pathlib import Path
from unittest.mock import AsyncMock, patch
from pydantic import ValidationError
from app.api.v1.llm_model_configs import (
LLMModelConfigUpsertRequest,
ensure_default_config,
normalize_payload,
to_storage_payload,
)
from app.models.llm_model_config import LLMModelConfig
from app.services.file_vector_sync_service import FileVectorSyncService
from app.services.llm_provider_service import LLMProviderService
from app.services.local_embedding_service import LocalEmbeddingService
from app.services.project_file_service import ProjectFileService
from app.services.zvec_service import ZVecService
class _ScalarResult:
def __init__(self, value):
self.value = value
def scalar_one_or_none(self):
return self.value
class _RecordingDB:
def __init__(self, values):
self.values = iter(values)
self.statements = []
async def execute(self, statement):
self.statements.append(statement)
return _ScalarResult(next(self.values))
class ModelConfigurationTest(unittest.IsolatedAsyncioTestCase):
def test_chat_type_config_contains_only_chat_specific_fields(self):
request = LLMModelConfigUpsertRequest(
model_type="chat",
provider="openai",
llm_model_name="gpt-test",
llm_temperature=0.4,
llm_top_p=0.8,
llm_max_tokens=2048,
llm_system_prompt="测试提示词",
)
payload = to_storage_payload(normalize_payload(request))
self.assertEqual(
payload["type_config"],
{
"temperature": 0.4,
"top_p": 0.8,
"max_tokens": 2048,
"system_prompt": "测试提示词",
},
)
self.assertNotIn("chunk_size", payload["type_config"])
self.assertNotIn("llm_temperature", payload)
def test_embedding_type_config_contains_only_embedding_specific_fields(self):
request = LLMModelConfigUpsertRequest(
model_type="embedding",
provider="local",
llm_model_name="m3e-small",
embedding_dimension=512,
chunk_size=600,
chunk_overlap=100,
)
payload = to_storage_payload(normalize_payload(request))
config = LLMModelConfig(
model_type="embedding",
type_config=payload["type_config"],
)
self.assertEqual(
payload["type_config"],
{"dimension": 512, "chunk_size": 600, "chunk_overlap": 100},
)
self.assertNotIn("temperature", payload["type_config"])
self.assertEqual(config.embedding_dimension, 512)
self.assertEqual(config.chunk_size, 600)
self.assertEqual(config.chunk_overlap, 100)
def test_chunk_overlap_must_be_smaller_than_chunk_size(self):
with self.assertRaises(ValidationError):
LLMModelConfigUpsertRequest(
model_type="embedding",
provider="local",
llm_model_name="m3e-small",
chunk_size=500,
chunk_overlap=500,
)
async def test_default_update_is_scoped_to_model_type(self):
db = _RecordingDB([None, 7, None, None])
await ensure_default_config(db, "embedding", preferred_config_id=7)
update_statement = str(db.statements[2])
self.assertIn("llm_model_config.model_type", update_statement)
self.assertIn("is_default", update_statement)
async def test_local_embedding_dimension_is_validated(self):
with patch.object(
LocalEmbeddingService,
"generate_embedding",
new=AsyncMock(return_value=[0.1] * 384),
):
vector = await LLMProviderService.generate_embedding(
provider="local",
endpoint_url="",
api_key="",
llm_model_name="paraphrase-multilingual-MiniLM-L12-v2",
text="测试文本",
dimension=384,
)
self.assertEqual(len(vector), 384)
with self.assertRaisesRegex(ValueError, "实际输出 384 维"):
await LLMProviderService.generate_embedding(
provider="local",
endpoint_url="",
api_key="",
llm_model_name="paraphrase-multilingual-MiniLM-L12-v2",
text="测试文本",
dimension=768,
)
def test_local_model_path_cannot_escape_models_directory(self):
with tempfile.TemporaryDirectory() as temp_dir:
models_dir = Path(temp_dir) / "models"
models_dir.mkdir()
(models_dir / "valid-model").mkdir()
with patch.object(LocalEmbeddingService, "MODELS_DIR", models_dir):
self.assertEqual(
LocalEmbeddingService.resolve_model_path("valid-model"),
(models_dir / "valid-model").resolve(),
)
with self.assertRaisesRegex(ValueError, "必须位于"):
LocalEmbeddingService.resolve_model_path("../outside")
def test_local_model_catalog_reports_dimension_and_readiness(self):
with tempfile.TemporaryDirectory() as temp_dir:
models_dir = Path(temp_dir) / "models"
model_dir = models_dir / "local-model"
pooling_dir = model_dir / "1_Pooling"
pooling_dir.mkdir(parents=True)
(pooling_dir / "config.json").write_text(
'{"word_embedding_dimension": 384}',
encoding="utf-8",
)
(model_dir / "model.safetensors").touch()
with patch.object(LocalEmbeddingService, "MODELS_DIR", models_dir):
self.assertEqual(
LocalEmbeddingService.list_available_models(),
[{"name": "local-model", "dimension": 384, "ready": True}],
)
def test_embedding_config_controls_chunking(self):
chunks = ZVecService._chunk_text("abcdefghij", chunk_size=6, overlap=2)
self.assertEqual([item["text"] for item in chunks], ["abcdef", "efghij", "ij"])
class FileVectorSyncTest(unittest.IsolatedAsyncioTestCase):
def test_directory_markdown_paths_and_moves(self):
with tempfile.TemporaryDirectory() as temp_dir:
root = Path(temp_dir) / "docs"
(root / "nested").mkdir(parents=True)
(root / "a.md").write_text("a", encoding="utf-8")
(root / "nested" / "b.md").write_text("b", encoding="utf-8")
(root / "ignored.txt").write_text("x", encoding="utf-8")
self.assertEqual(
sorted(ProjectFileService._markdown_paths(root, "docs")),
["docs/a.md", "docs/nested/b.md"],
)
self.assertEqual(
sorted(ProjectFileService._markdown_moves(root, "docs", "archive")),
[
("docs/a.md", "archive/a.md"),
("docs/nested/b.md", "archive/nested/b.md"),
],
)
async def test_background_sync_uses_its_own_database_session(self):
db = AsyncMock()
class SessionContext:
async def __aenter__(self):
return db
async def __aexit__(self, exc_type, exc, traceback):
return False
service = FileVectorSyncService()
with (
patch(
"app.services.file_vector_sync_service.AsyncSessionLocal",
return_value=SessionContext(),
),
patch.object(
ZVecService,
"revectorize_path",
new=AsyncMock(return_value=True),
) as revectorize,
):
await service._run(3, "revectorize", "guide.md", "")
revectorize.assert_awaited_once_with(db, 3, "guide.md")
if __name__ == "__main__":
unittest.main()