From f7d864bc4d08058b2aed22e7cb7f0b9ccdef20ed Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Sun, 30 Aug 2026 10:43:55 +0800 Subject: [PATCH] =?UTF-8?q?feat(provider):=20=E6=B7=BB=E5=8A=A0=E5=8E=82?= =?UTF-8?q?=E5=95=86=E9=A2=84=E8=AE=BE=E4=B8=8E=E6=A8=A1=E5=9E=8B=E5=8F=91?= =?UTF-8?q?=E7=8E=B0=E6=8E=A5=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/contracts.py | 12 +++++++++++ backend/app/providers/factory.py | 26 ++++++++++++++++++++++- backend/app/routes.py | 28 ++++++++++++++++++++++++- backend/tests/test_api.py | 27 +++++++++++++++++++++++- backend/tests/test_provider_adapters.py | 22 +++++++++++++++++++ 5 files changed, 112 insertions(+), 3 deletions(-) diff --git a/backend/app/contracts.py b/backend/app/contracts.py index e335a70..c44e885 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -438,6 +438,18 @@ class ProviderListResponse(Contract): items: list[ProviderConfig] = Field(default_factory=list) +class ProviderPreset(Contract): + preset_id: str + name: str + provider_type: ProviderType + base_url: str + requires_credential: bool = True + + +class ProviderPresetListResponse(Contract): + items: list[ProviderPreset] = Field(default_factory=list) + + class ModelInfo(Contract): model: str display_name: str diff --git a/backend/app/providers/factory.py b/backend/app/providers/factory.py index 8236882..e9f1947 100644 --- a/backend/app/providers/factory.py +++ b/backend/app/providers/factory.py @@ -1,4 +1,4 @@ -from app.contracts import ModelCapability, ProviderConfig, ProviderType +from app.contracts import ModelCapability, ProviderConfig, ProviderPreset, ProviderType from app.providers.base import ModelProvider from app.providers.credentials import CredentialResolver from app.providers.ollama import OllamaProvider @@ -27,6 +27,30 @@ class ProviderFactory: return OllamaProvider(config.base_url or "http://127.0.0.1:11434") raise UnsupportedProviderError(config.provider_type.value) + @staticmethod + def presets() -> list[ProviderPreset]: + return [ + ProviderPreset( + preset_id="openai", + name="OpenAI", + provider_type=ProviderType.openai_chat, + base_url="https://api.openai.com/v1", + ), + ProviderPreset( + preset_id="deepseek", + name="DeepSeek", + provider_type=ProviderType.openai_compatible, + base_url="https://api.deepseek.com", + ), + ProviderPreset( + preset_id="ollama", + name="Ollama", + provider_type=ProviderType.ollama, + base_url="http://127.0.0.1:11434", + requires_credential=False, + ), + ] + @staticmethod def capabilities(provider_type: ProviderType) -> list[ModelCapability]: if provider_type in { diff --git a/backend/app/routes.py b/backend/app/routes.py index ecdf374..1fc1994 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -31,6 +31,7 @@ from app.contracts import ( ProviderCreateRequest, ProviderListResponse, ProviderModelsResponse, + ProviderPresetListResponse, ProviderTestRequest, ProviderTestResponse, ProviderUpdateRequest, @@ -52,6 +53,7 @@ from app.errors import ApiError from app.extensions import ExtensionError from app.providers.registry import ProviderNotFoundError from app.providers.factory import UnsupportedProviderError +from app.providers.base import ProviderError from app.retrieval.engine import engine from app.services import index_service, note_service, task_service, transcription_service @@ -416,6 +418,15 @@ async def list_providers() -> ProviderListResponse: return ProviderListResponse(items=container.providers.list_configs()) +@router.get( + "/providers/presets", + response_model=ProviderPresetListResponse, + tags=["Providers"], +) +async def list_provider_presets() -> ProviderPresetListResponse: + return ProviderPresetListResponse(items=container.provider_factory.presets()) + + @router.get( "/providers/{provider_id}", response_model=ProviderConfig, @@ -502,9 +513,24 @@ async def delete_provider(provider_id: str) -> OperationResponse: ) async def list_provider_models(provider_id: str) -> ProviderModelsResponse: provider_or_404(provider_id) + try: + models = await container.providers.list_models(provider_id) + except ProviderError as exc: + status_code = { + "PROVIDER_AUTH_FAILED": 401, + "MODEL_NOT_FOUND": 404, + "PROVIDER_RATE_LIMITED": 429, + "PROVIDER_TIMEOUT": 504, + }.get(exc.code, 502) + raise ApiError( + status_code, + exc.code, + exc.message, + {"provider_id": provider_id}, + ) from exc return ProviderModelsResponse( provider_id=provider_id, - items=await container.providers.list_models(provider_id), + items=models, ) diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index 8b57a08..f107842 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -1,7 +1,14 @@ import asyncio from app.main import health, service_status -from app.routes import get_index_status, list_notes, list_plugins, list_providers, list_skills +from app.routes import ( + get_index_status, + list_notes, + list_plugins, + list_provider_presets, + list_providers, + list_skills, +) from app.routes import ( create_provider, create_task, @@ -53,6 +60,23 @@ def test_core_collections_are_typed() -> None: assert index.status == "idle" +def test_provider_presets_include_openai_and_deepseek() -> None: + presets = asyncio.run(list_provider_presets()) + by_id = {item.preset_id: item for item in presets.items} + + assert by_id["openai"].base_url == "https://api.openai.com/v1" + assert by_id["deepseek"].base_url == "https://api.deepseek.com" + assert by_id["deepseek"].provider_type == ProviderType.openai_compatible + + +def test_provider_presets_static_route_precedes_provider_id_route() -> None: + from app.routes import router + + get_paths = [route.path for route in router.routes if "GET" in getattr(route, "methods", set())] + + assert get_paths.index("/api/providers/presets") < get_paths.index("/api/providers/{provider_id}") + + def test_openapi_contains_documented_frontend_interfaces() -> None: from app.main import app @@ -70,6 +94,7 @@ def test_openapi_contains_documented_frontend_interfaces() -> None: "/api/plugins/{plugin_id}/enable", "/api/plugins/{plugin_id}/disable", "/api/providers/test", + "/api/providers/presets", "/api/index/rebuild", } diff --git a/backend/tests/test_provider_adapters.py b/backend/tests/test_provider_adapters.py index dba0151..f6c5f45 100644 --- a/backend/tests/test_provider_adapters.py +++ b/backend/tests/test_provider_adapters.py @@ -134,6 +134,28 @@ def test_openai_compatible_preserves_tool_call_context() -> None: assert turn.text == "done" +def test_openai_compatible_fetches_and_maps_model_list() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert request.method == "GET" + assert request.url.path == "/v1/models" + assert request.headers["Authorization"] == "Bearer secret-test-key" + return httpx.Response( + 200, + json={"data": [{"id": "model-b"}, {"id": "model-a"}]}, + ) + + provider = OpenAICompatibleProvider( + base_url="https://provider.test/v1", + credential_id="provider-test", + credentials=StaticCredentials(), + transport=httpx.MockTransport(handler), + ) + + models = run(provider.list_models()) + + assert [item.model for item in models] == ["model-b", "model-a"] + + def test_ollama_maps_models_and_completion() -> None: def handler(request: httpx.Request) -> httpx.Response: if request.url.path == "/api/tags":