feat(provider): 添加厂商预设与模型发现接口
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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":
|
||||||
|
|||||||
Reference in New Issue
Block a user