98 lines
3.6 KiB
Python
98 lines
3.6 KiB
Python
"""Persistence for reusable model configurations and scenario bindings."""
|
|
|
|
from sqlmodel import Session, select
|
|
|
|
from agenteval.storage.db import ModelConfigDB, ScenarioModelBindingDB, utc_now
|
|
|
|
|
|
class ModelConfigRepository:
|
|
def __init__(self, session: Session):
|
|
self.session = session
|
|
|
|
def list_all(
|
|
self,
|
|
capability: str | None = None,
|
|
enabled: bool | None = None,
|
|
) -> list[ModelConfigDB]:
|
|
statement = select(ModelConfigDB)
|
|
if capability is not None:
|
|
statement = statement.where(ModelConfigDB.capability == capability)
|
|
if enabled is not None:
|
|
statement = statement.where(ModelConfigDB.enabled == enabled)
|
|
statement = statement.order_by(ModelConfigDB.created_at.desc())
|
|
return list(self.session.exec(statement).all())
|
|
|
|
def get(self, config_id: str) -> ModelConfigDB | None:
|
|
return self.session.get(ModelConfigDB, config_id)
|
|
|
|
def get_by_name(self, name: str) -> ModelConfigDB | None:
|
|
return self.session.exec(select(ModelConfigDB).where(ModelConfigDB.name == name)).first()
|
|
|
|
def create(self, config: ModelConfigDB) -> ModelConfigDB:
|
|
if config.is_default:
|
|
self.clear_default(config.capability)
|
|
self.session.add(config)
|
|
self.session.commit()
|
|
self.session.refresh(config)
|
|
return config
|
|
|
|
def update(self, config: ModelConfigDB) -> ModelConfigDB:
|
|
if config.is_default:
|
|
self.clear_default(config.capability, exclude_id=config.id)
|
|
config.updated_at = utc_now()
|
|
self.session.add(config)
|
|
self.session.commit()
|
|
self.session.refresh(config)
|
|
return config
|
|
|
|
def delete(self, config: ModelConfigDB) -> None:
|
|
self.session.delete(config)
|
|
self.session.commit()
|
|
|
|
def clear_default(self, capability: str, exclude_id: str | None = None) -> None:
|
|
statement = select(ModelConfigDB).where(
|
|
ModelConfigDB.capability == capability,
|
|
ModelConfigDB.is_default.is_(True),
|
|
)
|
|
for item in self.session.exec(statement).all():
|
|
if item.id == exclude_id:
|
|
continue
|
|
item.is_default = False
|
|
item.updated_at = utc_now()
|
|
self.session.add(item)
|
|
|
|
def list_references(self, config_id: str) -> list[ScenarioModelBindingDB]:
|
|
statement = select(ScenarioModelBindingDB).where(
|
|
ScenarioModelBindingDB.model_config_id == config_id,
|
|
)
|
|
return list(self.session.exec(statement).all())
|
|
|
|
|
|
class ScenarioModelBindingRepository:
|
|
def __init__(self, session: Session):
|
|
self.session = session
|
|
|
|
def get_for_scenario(self, scenario_id: str) -> dict[str, str]:
|
|
statement = select(ScenarioModelBindingDB).where(
|
|
ScenarioModelBindingDB.scenario_id == scenario_id,
|
|
)
|
|
return {item.purpose: item.model_config_id for item in self.session.exec(statement).all()}
|
|
|
|
def replace_for_scenario(self, scenario_id: str, bindings: dict[str, str]) -> None:
|
|
statement = select(ScenarioModelBindingDB).where(
|
|
ScenarioModelBindingDB.scenario_id == scenario_id,
|
|
)
|
|
for item in self.session.exec(statement).all():
|
|
self.session.delete(item)
|
|
for purpose, model_config_id in bindings.items():
|
|
self.session.add(
|
|
ScenarioModelBindingDB(
|
|
scenario_id=scenario_id,
|
|
purpose=purpose,
|
|
model_config_id=model_config_id,
|
|
)
|
|
)
|
|
|
|
def delete_for_scenario(self, scenario_id: str) -> None:
|
|
self.replace_for_scenario(scenario_id, {})
|