AgentEvalTool/backend/agenteval/web/routers/model_configs.py

126 lines
4.4 KiB
Python

"""API routes for centralized model configurations."""
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlmodel import Session
from agenteval.model_gateway import ModelGateway, ModelGatewayError
from agenteval.models import ModelCapability
from agenteval.services.model_configs import (
ModelConfigError,
ModelConfigInUseError,
ModelConfigNotFoundError,
ModelConfigService,
)
from agenteval.storage.db import ScenarioDB
from agenteval.web.deps import get_db
from agenteval.web.model_config_schemas import (
ModelConfigCreate,
ModelConfigReference,
ModelConfigResponse,
ModelConfigUpdate,
ModelConnectionTestResponse,
)
router = APIRouter()
def _http_error(exc: ModelConfigError) -> HTTPException:
if isinstance(exc, ModelConfigNotFoundError):
return HTTPException(status_code=404, detail=str(exc))
if isinstance(exc, ModelConfigInUseError):
return HTTPException(status_code=409, detail=str(exc))
return HTTPException(status_code=400, detail=str(exc))
@router.get("", response_model=list[ModelConfigResponse])
def list_model_configs(
capability: ModelCapability | None = Query(default=None),
enabled: bool | None = Query(default=None),
session: Session = Depends(get_db),
) -> list[ModelConfigResponse]:
configs = ModelConfigService(session).repo.list_all(
capability=capability.value if capability else None,
enabled=enabled,
)
return [ModelConfigResponse.from_db(config) for config in configs]
@router.post("", response_model=ModelConfigResponse)
def create_model_config(payload: ModelConfigCreate, session: Session = Depends(get_db)) -> ModelConfigResponse:
try:
config = ModelConfigService(session).create(**payload.model_dump(mode="json"))
except ModelConfigError as exc:
raise _http_error(exc) from exc
return ModelConfigResponse.from_db(config)
@router.get("/{config_id}", response_model=ModelConfigResponse)
def get_model_config(config_id: str, session: Session = Depends(get_db)) -> ModelConfigResponse:
try:
config = ModelConfigService(session).require(config_id)
except ModelConfigError as exc:
raise _http_error(exc) from exc
return ModelConfigResponse.from_db(config)
@router.put("/{config_id}", response_model=ModelConfigResponse)
def update_model_config(
config_id: str,
payload: ModelConfigUpdate,
session: Session = Depends(get_db),
) -> ModelConfigResponse:
try:
config = ModelConfigService(session).update(config_id, **payload.model_dump(mode="json"))
except ModelConfigError as exc:
raise _http_error(exc) from exc
return ModelConfigResponse.from_db(config)
@router.delete("/{config_id}")
def delete_model_config(config_id: str, session: Session = Depends(get_db)) -> dict[str, bool]:
try:
ModelConfigService(session).delete(config_id)
except ModelConfigError as exc:
raise _http_error(exc) from exc
return {"ok": True}
@router.get("/{config_id}/references", response_model=list[ModelConfigReference])
def list_model_config_references(
config_id: str,
session: Session = Depends(get_db),
) -> list[ModelConfigReference]:
service = ModelConfigService(session)
try:
service.require(config_id)
except ModelConfigError as exc:
raise _http_error(exc) from exc
references = []
for binding in service.repo.list_references(config_id):
scenario = session.get(ScenarioDB, binding.scenario_id)
references.append(
ModelConfigReference(
scenario_id=binding.scenario_id,
scenario_name=scenario.name if scenario else binding.scenario_id,
purpose=binding.purpose,
)
)
return references
@router.post("/{config_id}/test", response_model=ModelConnectionTestResponse)
async def test_model_config(
config_id: str,
session: Session = Depends(get_db),
) -> ModelConnectionTestResponse:
try:
runtime = ModelConfigService(session).resolve(config_id)
message = await ModelGateway().test_connection(runtime)
except ModelConfigError as exc:
raise _http_error(exc) from exc
except ModelGatewayError as exc:
return ModelConnectionTestResponse(ok=False, message=str(exc), tested_at=datetime.now(timezone.utc))
return ModelConnectionTestResponse(ok=True, message=message, tested_at=datetime.now(timezone.utc))