feat(provider): 添加厂商预设与模型发现接口

This commit is contained in:
2026-08-30 10:43:55 +08:00
parent 11cb384115
commit f7d864bc4d
5 changed files with 112 additions and 3 deletions
+12
View File
@@ -438,6 +438,18 @@ class ProviderListResponse(Contract):
items: list[ProviderConfig] = Field(default_factory=list) 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): class ModelInfo(Contract):
model: str model: str
display_name: str display_name: str
+25 -1
View File
@@ -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.base import ModelProvider
from app.providers.credentials import CredentialResolver from app.providers.credentials import CredentialResolver
from app.providers.ollama import OllamaProvider from app.providers.ollama import OllamaProvider
@@ -27,6 +27,30 @@ class ProviderFactory:
return OllamaProvider(config.base_url or "http://127.0.0.1:11434") return OllamaProvider(config.base_url or "http://127.0.0.1:11434")
raise UnsupportedProviderError(config.provider_type.value) 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 @staticmethod
def capabilities(provider_type: ProviderType) -> list[ModelCapability]: def capabilities(provider_type: ProviderType) -> list[ModelCapability]:
if provider_type in { if provider_type in {
+27 -1
View File
@@ -31,6 +31,7 @@ from app.contracts import (
ProviderCreateRequest, ProviderCreateRequest,
ProviderListResponse, ProviderListResponse,
ProviderModelsResponse, ProviderModelsResponse,
ProviderPresetListResponse,
ProviderTestRequest, ProviderTestRequest,
ProviderTestResponse, ProviderTestResponse,
ProviderUpdateRequest, ProviderUpdateRequest,
@@ -52,6 +53,7 @@ from app.errors import ApiError
from app.extensions import ExtensionError from app.extensions import ExtensionError
from app.providers.registry import ProviderNotFoundError from app.providers.registry import ProviderNotFoundError
from app.providers.factory import UnsupportedProviderError from app.providers.factory import UnsupportedProviderError
from app.providers.base import ProviderError
from app.retrieval.engine import engine from app.retrieval.engine import engine
from app.services import index_service, note_service, task_service, transcription_service 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()) 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( @router.get(
"/providers/{provider_id}", "/providers/{provider_id}",
response_model=ProviderConfig, response_model=ProviderConfig,
@@ -502,9 +513,24 @@ async def delete_provider(provider_id: str) -> OperationResponse:
) )
async def list_provider_models(provider_id: str) -> ProviderModelsResponse: async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
provider_or_404(provider_id) 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( return ProviderModelsResponse(
provider_id=provider_id, provider_id=provider_id,
items=await container.providers.list_models(provider_id), items=models,
) )
+26 -1
View File
@@ -1,7 +1,14 @@
import asyncio import asyncio
from app.main import health, service_status 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 ( from app.routes import (
create_provider, create_provider,
create_task, create_task,
@@ -53,6 +60,23 @@ def test_core_collections_are_typed() -> None:
assert index.status == "idle" 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: def test_openapi_contains_documented_frontend_interfaces() -> None:
from app.main import app 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}/enable",
"/api/plugins/{plugin_id}/disable", "/api/plugins/{plugin_id}/disable",
"/api/providers/test", "/api/providers/test",
"/api/providers/presets",
"/api/index/rebuild", "/api/index/rebuild",
} }
+22
View File
@@ -134,6 +134,28 @@ def test_openai_compatible_preserves_tool_call_context() -> None:
assert turn.text == "done" 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 test_ollama_maps_models_and_completion() -> None:
def handler(request: httpx.Request) -> httpx.Response: def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/api/tags": if request.url.path == "/api/tags":