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()