126 lines
4.4 KiB
Python
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))
|