feat(provider): 完成阶段E协议适配、国内预设与模型路由

This commit is contained in:
2026-09-04 06:19:32 +08:00
parent 9b8b10cdb1
commit 1fe75e3fd2
57 changed files with 4208 additions and 592 deletions
+51 -10
View File
@@ -1,5 +1,6 @@
import asyncio
from collections.abc import AsyncIterator
from contextlib import aclosing
from datetime import datetime, timezone
from uuid import uuid4
@@ -33,6 +34,12 @@ from app.contracts import (
McpToolSummaryListResponse,
ModelEvent,
ModelEventType,
EmbeddingRequest,
EmbeddingResult,
ModelRoutingConfig,
ModelRoutingResponse,
SpeakerMatchRequest,
SpeakerMatchResult,
Note,
NoteCreateRequest,
NoteListResponse,
@@ -96,6 +103,7 @@ from app.services import (
transcription_service,
workspace_service,
)
from app.services.attachment_service import attachment_path
router = APIRouter(prefix="/api")
@@ -294,17 +302,22 @@ async def chat(request: ChatRequest) -> StreamingResponse:
provider = provider_or_404(request.provider_id)
async def stream() -> AsyncIterator[str]:
sequence = 0
try:
async for event in provider.adapter.stream(request):
yield as_sse(event.event.value, event.model_dump_json())
except Exception as exc:
async with aclosing(provider.adapter.stream(request)) as events:
async for event in events:
sequence = event.sequence + 1
yield as_sse(event.event.value, event.model_dump_json())
except Exception:
error = ModelEvent(
event=ModelEventType.error,
data={"code": "PROVIDER_ERROR", "message": str(exc)},
sequence=sequence,
data={"code": "PROVIDER_ERROR", "message": "Provider could not complete the request."},
timestamp=utc_now(),
)
done = ModelEvent(
event=ModelEventType.done, sequence=1, timestamp=utc_now()
event=ModelEventType.done, sequence=sequence + 1,
data={"status": "failed"}, timestamp=utc_now()
)
yield as_sse(error.event.value, error.model_dump_json())
yield as_sse(done.event.value, done.model_dump_json())
@@ -911,13 +924,13 @@ async def update_provider(
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified."
)
fields = request.model_fields_set
if ("name" in fields and request.name is None) or (
if ("provider_type" in fields and request.provider_type is None) or ("name" in fields and request.name is None) or (
"enabled" in fields and request.enabled is None
):
raise ApiError(
422,
"VALIDATION_ERROR",
"name and enabled cannot be null when explicitly provided.",
"provider_type, name and enabled cannot be null when explicitly provided.",
)
updates = {name: getattr(request, name) for name in fields}
if "credential_id" in fields:
@@ -925,7 +938,11 @@ async def update_provider(
config = ProviderConfig.model_validate(
{**current.model_dump(mode="python"), **updates}
)
adapter = container.provider_factory.build(config)
config.capabilities = container.provider_factory.capabilities(config.provider_type)
try:
adapter = container.provider_factory.build(config)
except UnsupportedProviderError as exc:
raise ApiError(422, "PROVIDER_TYPE_UNSUPPORTED", "Provider adapter is not supported.") from exc
container.providers.replace(config, adapter)
return config
@@ -941,6 +958,8 @@ async def delete_provider(provider_id: str) -> OperationResponse:
raise ApiError(
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted."
)
if container.model_routing.uses_provider(provider_id):
raise ApiError(409, "PROVIDER_IN_USE", "请先在索引与模型中解除该提供商的模型绑定。")
container.providers.unregister(provider_id)
return OperationResponse(status="completed", resource_id=provider_id)
@@ -1043,6 +1062,28 @@ async def delete_task(task_id: str) -> OperationResponse:
# Media and index
@router.get("/model-routing", response_model=ModelRoutingResponse, tags=["Providers"])
async def get_model_routing() -> ModelRoutingResponse:
return container.model_routing.describe()
@router.put("/model-routing", response_model=ModelRoutingResponse, tags=["Providers"])
async def update_model_routing(request: ModelRoutingConfig) -> ModelRoutingResponse:
return container.model_routing.update(request)
@router.post("/models/embeddings", response_model=EmbeddingResult, tags=["Providers"])
async def create_embeddings(request: EmbeddingRequest) -> EmbeddingResult:
return await container.model_routing.embed(request.texts)
@router.post("/media/speaker-matches", response_model=SpeakerMatchResult, tags=["Media"])
async def match_speakers(request: SpeakerMatchRequest) -> SpeakerMatchResult:
return await container.model_routing.match_speakers(
attachment_path(request.attachment_id), attachment_path(request.reference_attachment_id),
)
@router.post(
"/media/transcriptions",
response_model=TranscriptionJob,
@@ -1050,8 +1091,8 @@ async def delete_task(task_id: str) -> OperationResponse:
tags=["Media"],
)
async def create_transcription(request: TranscriptionRequest) -> TranscriptionJob:
return transcription_service.create_transcription(
request.attachment_id, request.language
return await transcription_service.create_transcription(
request.attachment_id, request.language, diarization=request.diarization
)