222 lines
7.9 KiB
Python
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()
|