From 6438f8b7c2e490ee18cb8c21d400e9dc7dde2cb6 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 27 May 2026 19:56:26 +0800 Subject: [PATCH] =?UTF-8?q?feat(models):=20=E5=A2=9E=E5=BC=BA=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E9=85=8D=E7=BD=AE=E8=83=BD=E5=8A=9B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- REFACTOR.md | 2 +- .../package/yuxi/config/builtin_providers.py | 16 ++++ backend/package/yuxi/models/embed.py | 7 +- backend/package/yuxi/models/rerank.py | 13 +++ backend/package/yuxi/services/model_cache.py | 2 +- .../yuxi/services/model_provider_service.py | 8 +- .../server/routers/model_provider_router.py | 3 + .../unit/server/test_model_provider_router.py | 39 +++++++++ .../test/unit/services/test_model_cache.py | 32 +++++++ .../services/test_model_provider_service.py | 14 +++ .../unit/services/test_model_selectors.py | 34 ++++++++ .../ModelProviderManagePanel.vue | 87 +++++++++++++++++-- 12 files changed, 246 insertions(+), 11 deletions(-) create mode 100644 backend/test/unit/services/test_model_cache.py diff --git a/REFACTOR.md b/REFACTOR.md index 91070364..2c4eff51 100644 --- a/REFACTOR.md +++ b/REFACTOR.md @@ -26,7 +26,7 @@ - [ ] add model retry times to agent context config - [ ] 添加用户级别的 Skills 的安装 - [ ] 子智能体的优化,参考 PR 的方案。 -- [ ] 附件上传能够支持转换为 PDF +- [ ] 附件上传能够支持转换为 PDF,待办:查看 OCR 模型的状态,样式优化,保存的文件名不对 - [ ] 参考 PR,实现内置 Dashscope 的 Embedding 和 rerank 的方法 - [ ] 优化知识库的 API 接口设计,使用 /{db_id}/xxx 的形式,整合 mindmap / eval 接口 - [x] allow multi-hop qa generate diff --git a/backend/package/yuxi/config/builtin_providers.py b/backend/package/yuxi/config/builtin_providers.py index 311d234c..9b390118 100644 --- a/backend/package/yuxi/config/builtin_providers.py +++ b/backend/package/yuxi/config/builtin_providers.py @@ -41,8 +41,24 @@ BUILTIN_PROVIDERS: list[dict[str, Any]] = [ "provider_id": "alibaba", "display_name": "DashScope", "base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1", + "embedding_base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings", + "rerank_base_url": "https://dashscope.aliyuncs.com/compatible-api/v1/reranks", "api_key_env": "DASHSCOPE_API_KEY", + "capabilities": ["chat", "embedding", "rerank"], "models_endpoint": "https://dashscope.aliyuncs.com/compatible-mode/v1/models", + "enabled_models": [ + { + "id": "text-embedding-v4", + "type": "embedding", + "display_name": "text-embedding-v4", + "dimension": 1024, + }, + { + "id": "qwen3-rerank", + "type": "rerank", + "display_name": "qwen3-rerank", + }, + ], }, { "provider_id": "alibaba-coding-plan-cn", diff --git a/backend/package/yuxi/models/embed.py b/backend/package/yuxi/models/embed.py index ba2b888c..b9f31a5a 100644 --- a/backend/package/yuxi/models/embed.py +++ b/backend/package/yuxi/models/embed.py @@ -91,7 +91,12 @@ class BaseEmbeddingModel(ABC): async def test_connection(self) -> tuple[bool, str]: try: - await self.aencode(["Hello world"]) + embeddings = await self.aencode(["Hello world"]) + if self.dimension not in (None, ""): + actual_dimension = len(embeddings[0]) if embeddings else 0 + expected_dimension = int(self.dimension) + if actual_dimension != expected_dimension: + return False, f"Embedding 维度不一致:配置 {expected_dimension},实际 {actual_dimension}" return True, "连接正常" except Exception as e: error_msg = str(e) diff --git a/backend/package/yuxi/models/rerank.py b/backend/package/yuxi/models/rerank.py index 941452aa..1b2c8275 100644 --- a/backend/package/yuxi/models/rerank.py +++ b/backend/package/yuxi/models/rerank.py @@ -105,6 +105,19 @@ class BaseReranker(ABC): return asyncio.run(self.acompute_score(sentence_pairs, batch_size, max_length, normalize)) raise RuntimeError("compute_score cannot be used while an event loop is running. Use acompute_score instead.") + async def test_connection(self) -> tuple[bool, str]: + try: + scores = await self._batch_rerank("test query", ["test document"], max_length=128) + if scores: + return True, "连接正常" + return False, "响应无效" + except Exception as e: + error_msg = str(e) + logger.error(f"Rerank connection test failed: {error_msg}") + return False, error_msg + finally: + await self.aclose() + async def aclose(self) -> None: if self.session and not self.session.closed: await self.session.close() diff --git a/backend/package/yuxi/services/model_cache.py b/backend/package/yuxi/services/model_cache.py index 660a0023..b7f3cb9b 100644 --- a/backend/package/yuxi/services/model_cache.py +++ b/backend/package/yuxi/services/model_cache.py @@ -161,7 +161,7 @@ class ModelCache: for model in provider.enabled_models or []: model_type = model.get("type", "chat") - base_url = self._get_base_url_for_type(provider, model_type) + base_url = model.get("base_url_override") or self._get_base_url_for_type(provider, model_type) info = ModelInfo( provider_id=provider.provider_id, diff --git a/backend/package/yuxi/services/model_provider_service.py b/backend/package/yuxi/services/model_provider_service.py index a764e302..92a45f43 100644 --- a/backend/package/yuxi/services/model_provider_service.py +++ b/backend/package/yuxi/services/model_provider_service.py @@ -387,10 +387,14 @@ async def test_model_status_by_spec(spec: str) -> dict: "model_type": "embedding", } if info.model_type == "rerank": + from yuxi.models.rerank import get_reranker + + model = get_reranker(spec) + success, message = await model.test_connection() return { "spec": spec, - "status": "unsupported", - "message": "暂不支持 rerank 模型在线连接测试", + "status": "available" if success else "unavailable", + "message": "连接正常" if success else message, "model_type": "rerank", } diff --git a/backend/server/routers/model_provider_router.py b/backend/server/routers/model_provider_router.py index 71dd01bf..c06ef1bd 100644 --- a/backend/server/routers/model_provider_router.py +++ b/backend/server/routers/model_provider_router.py @@ -87,6 +87,7 @@ async def create_provider( payload.model_dump(exclude_none=True), current_user.username, ) + await db.commit() await _refresh_model_cache() return {"success": True, "data": provider.to_dict()} except ValueError as e: @@ -138,6 +139,7 @@ async def update_provider( provider = await update_provider_config(db, provider_id, data, current_user.username) if provider is None: raise HTTPException(status_code=404, detail=f"供应商 {provider_id} 不存在") + await db.commit() await _refresh_model_cache() return {"success": True, "data": provider.to_dict()} except HTTPException: @@ -159,6 +161,7 @@ async def delete_provider( deleted = await delete_provider_config(db, provider_id) if not deleted: raise HTTPException(status_code=404, detail=f"供应商 {provider_id} 不存在") + await db.commit() await _refresh_model_cache() return {"success": True} diff --git a/backend/test/unit/server/test_model_provider_router.py b/backend/test/unit/server/test_model_provider_router.py index 89941a1f..23251efe 100644 --- a/backend/test/unit/server/test_model_provider_router.py +++ b/backend/test/unit/server/test_model_provider_router.py @@ -1,3 +1,6 @@ +import pytest + +from server.routers import model_provider_router from server.routers.model_provider_router import ModelProviderPayload @@ -15,3 +18,39 @@ def test_model_provider_payload_accepts_embedding_and_rerank_urls(): assert data["embedding_base_url"] == "https://api.example.com/v1/embeddings" assert data["rerank_base_url"] == "https://api.example.com/v1/rerank" + + +@pytest.mark.asyncio +async def test_update_provider_commits_before_refreshing_cache(monkeypatch): + calls = [] + + class Db: + async def commit(self): + calls.append("commit") + + class User: + username = "admin" + + class Provider: + def to_dict(self): + return {"provider_id": "alibaba"} + + async def fake_update_provider_config(db, provider_id, data, username): + calls.append("update") + return Provider() + + async def fake_refresh_model_cache(): + calls.append("refresh") + + monkeypatch.setattr(model_provider_router, "update_provider_config", fake_update_provider_config) + monkeypatch.setattr(model_provider_router, "_refresh_model_cache", fake_refresh_model_cache) + + result = await model_provider_router.update_provider( + "alibaba", + ModelProviderPayload(enabled_models=[]), + current_user=User(), + db=Db(), + ) + + assert result == {"success": True, "data": {"provider_id": "alibaba"}} + assert calls == ["update", "commit", "refresh"] diff --git a/backend/test/unit/services/test_model_cache.py b/backend/test/unit/services/test_model_cache.py new file mode 100644 index 00000000..2a6762b0 --- /dev/null +++ b/backend/test/unit/services/test_model_cache.py @@ -0,0 +1,32 @@ +from yuxi.services.model_cache import ModelCache + + +def test_model_cache_prefers_model_base_url_override(monkeypatch): + saved_cache = {} + + class Provider: + is_enabled = True + provider_id = "alibaba" + api_key = "sk-test" + api_key_env = None + provider_type = "openai" + base_url = "https://dashscope.aliyuncs.com/compatible-mode/v1" + embedding_base_url = "https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings" + rerank_base_url = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" + headers_json = {} + extra_json = {} + enabled_models = [ + { + "id": "qwen3-rerank", + "type": "rerank", + "display_name": "Qwen3 Rerank", + "base_url_override": "https://invalid.example/rerank", + } + ] + + cache = ModelCache() + monkeypatch.setattr(cache, "_save_cache", lambda data: saved_cache.update(data)) + + cache.rebuild([Provider()]) + + assert saved_cache["alibaba:qwen3-rerank"].base_url == "https://invalid.example/rerank" diff --git a/backend/test/unit/services/test_model_provider_service.py b/backend/test/unit/services/test_model_provider_service.py index c822210a..0cda2179 100644 --- a/backend/test/unit/services/test_model_provider_service.py +++ b/backend/test/unit/services/test_model_provider_service.py @@ -157,6 +157,20 @@ def test_builtin_siliconflow_provider_includes_default_runnable_models(): assert "base_url_override" not in models["Pro/BAAI/bge-reranker-v2-m3"] +def test_builtin_dashscope_provider_includes_default_embedding_and_rerank_models(): + provider = next(item for item in BUILTIN_PROVIDERS if item["provider_id"] == "alibaba") + models = {model["id"]: model for model in provider["enabled_models"]} + + assert provider["capabilities"] == ["chat", "embedding", "rerank"] + assert provider["embedding_base_url"] == "https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings" + assert provider["rerank_base_url"] == "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" + assert "embedding_models_endpoint" not in provider + assert "rerank_models_endpoint" not in provider + assert models["text-embedding-v4"]["type"] == "embedding" + assert models["text-embedding-v4"]["dimension"] == 2048 + assert models["qwen3-rerank"]["type"] == "rerank" + + def testcheck_credential_status_disabled_provider_always_ok(): """未启用的 provider 无论凭证如何配置,状态始终为 ok。""" class Provider: diff --git a/backend/test/unit/services/test_model_selectors.py b/backend/test/unit/services/test_model_selectors.py index d978b66b..d2d05915 100644 --- a/backend/test/unit/services/test_model_selectors.py +++ b/backend/test/unit/services/test_model_selectors.py @@ -47,6 +47,40 @@ def test_select_embedding_model_loads_model_from_cache(monkeypatch): assert model.dimension == 1024 +@pytest.mark.asyncio +async def test_embedding_connection_checks_configured_dimension(monkeypatch): + model = OtherEmbedding( + model="namespace/embedding-model", + base_url="https://example.com/v1/embeddings", + api_key="test-key", + dimension=3, + ) + + async def fake_aencode(_messages): + return [[0.1, 0.2, 0.3]] + + monkeypatch.setattr(model, "aencode", fake_aencode) + + assert await model.test_connection() == (True, "连接正常") + + +@pytest.mark.asyncio +async def test_embedding_connection_reports_dimension_mismatch(monkeypatch): + model = OtherEmbedding( + model="namespace/embedding-model", + base_url="https://example.com/v1/embeddings", + api_key="test-key", + dimension=4, + ) + + async def fake_aencode(_messages): + return [[0.1, 0.2, 0.3]] + + monkeypatch.setattr(model, "aencode", fake_aencode) + + assert await model.test_connection() == (False, "Embedding 维度不一致:配置 4,实际 3") + + def test_get_reranker_loads_model_from_cache(monkeypatch): monkeypatch.setattr( "yuxi.models.rerank.model_cache.get_model_info", diff --git a/web/src/components/model-management/ModelProviderManagePanel.vue b/web/src/components/model-management/ModelProviderManagePanel.vue index f2848e79..5342774a 100644 --- a/web/src/components/model-management/ModelProviderManagePanel.vue +++ b/web/src/components/model-management/ModelProviderManagePanel.vue @@ -1,7 +1,17 @@