rebuild(llm): add audited model connectivity tests
This commit is contained in:
@@ -290,6 +290,13 @@ class ModelPoolService:
|
||||
raise BusinessError("model_not_configured", "智能解读服务尚未配置。")
|
||||
return ModelRuntimeConfig(primary=primary, fallback=fallback)
|
||||
|
||||
def runtime_model(self, model_id: int) -> ModelPoolRecord:
|
||||
with self._database.read() as connection:
|
||||
record = self._repository.get(connection, model_id)
|
||||
if record is None:
|
||||
raise BusinessError("model_not_found", "模型不存在。")
|
||||
return record
|
||||
|
||||
def decrypt_api_key(self, record: ModelPoolRecord) -> str:
|
||||
return self._cipher.decrypt(record.encrypted_api_key)
|
||||
|
||||
|
||||
@@ -35,9 +35,12 @@ from backend.features.accounts.schemas import (
|
||||
ModelInput,
|
||||
ModelPoolItemResponse,
|
||||
ModelSelectionInput,
|
||||
ModelTestResponse,
|
||||
ModelUpdateInput,
|
||||
PasswordChangeInput,
|
||||
)
|
||||
from backend.http.errors import AppError
|
||||
from backend.llm.gateway import LLMGatewayError
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -312,3 +315,16 @@ def delete_model(
|
||||
) -> MessageResponse:
|
||||
request.app.state.container.model_pool.delete(model_id)
|
||||
return MessageResponse(message="模型已删除。")
|
||||
|
||||
|
||||
@router.post("/admin/models/{model_id}/test", response_model=ModelTestResponse)
|
||||
def test_model(
|
||||
request: Request,
|
||||
model_id: Annotated[int, Path(ge=1)],
|
||||
principal: AdminWritePrincipal,
|
||||
) -> dict[str, object]:
|
||||
try:
|
||||
return request.app.state.container.llm.test_model(principal, model_id)
|
||||
except LLMGatewayError as exc:
|
||||
status = 404 if exc.code == "model_not_found" else 503
|
||||
raise AppError(exc.code, str(exc), status) from exc
|
||||
|
||||
@@ -119,3 +119,10 @@ class ModelPoolItemResponse(BaseModel):
|
||||
class ModelSelectionInput(BaseModel):
|
||||
primary_model_id: int = Field(ge=1)
|
||||
fallback_model_id: int | None = Field(default=None, ge=1)
|
||||
|
||||
|
||||
class ModelTestResponse(BaseModel):
|
||||
connected: bool
|
||||
message: str
|
||||
duration_ms: int = Field(ge=0)
|
||||
request_id: str
|
||||
|
||||
@@ -49,6 +49,8 @@ class ModelRuntimeAccess(Protocol):
|
||||
class ModelPoolAccess(Protocol):
|
||||
def runtime_config(self) -> ModelRuntimeAccess: ...
|
||||
|
||||
def runtime_model(self, model_id: int) -> ModelRecordAccess: ...
|
||||
|
||||
def decrypt_api_key(self, record: ModelRecordAccess) -> str: ...
|
||||
|
||||
|
||||
@@ -237,6 +239,57 @@ class LLMGateway:
|
||||
code, message = _safe_error(last_error)
|
||||
raise LLMGatewayError(code, message)
|
||||
|
||||
def test_model(self, principal: PrincipalAccess, model_id: int) -> dict[str, object]:
|
||||
try:
|
||||
record = self._model_pool.runtime_model(model_id)
|
||||
except BusinessError as exc:
|
||||
raise LLMGatewayError("model_not_found", "模型不存在。") from exc
|
||||
now = datetime.now(SHANGHAI)
|
||||
request_id = uuid.uuid4().hex
|
||||
call = LLMCall(
|
||||
request_id=request_id,
|
||||
user_id=principal.user.id,
|
||||
feature="model_connectivity",
|
||||
prompt_version="system:model-connectivity:v1",
|
||||
business_id=f"model:{model_id}",
|
||||
started_at=now,
|
||||
usage_date=now.date().isoformat(),
|
||||
input_chars=15,
|
||||
quota_exempt=True,
|
||||
profiles=(
|
||||
LLMProfile(
|
||||
model_id=record.id,
|
||||
role="primary",
|
||||
base_url=record.base_url,
|
||||
model_identifier=record.model_identifier,
|
||||
api_key=self._model_pool.decrypt_api_key(record),
|
||||
),
|
||||
),
|
||||
)
|
||||
with self._database.transaction() as connection:
|
||||
self._repository.reserve(
|
||||
connection,
|
||||
request_id=request_id,
|
||||
user_id=principal.user.id,
|
||||
feature=call.feature,
|
||||
business_id=call.business_id,
|
||||
prompt_version=call.prompt_version,
|
||||
started_at=now.isoformat(timespec="seconds"),
|
||||
input_chars=call.input_chars,
|
||||
)
|
||||
started = time.perf_counter()
|
||||
for _event in self.stream(
|
||||
call,
|
||||
[{"role": "user", "content": "连接测试,只回复:连接成功"}],
|
||||
):
|
||||
pass
|
||||
return {
|
||||
"connected": True,
|
||||
"message": "连接成功",
|
||||
"duration_ms": round((time.perf_counter() - started) * 1000),
|
||||
"request_id": request_id,
|
||||
}
|
||||
|
||||
def _finish_attempt(
|
||||
self,
|
||||
attempt_id: int,
|
||||
|
||||
Reference in New Issue
Block a user