"""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 get_analysis_default(self) -> ModelConfigDB | None: statement = select(ModelConfigDB).where(ModelConfigDB.is_analysis_default.is_(True)) return self.session.exec(statement).first() def create(self, config: ModelConfigDB) -> ModelConfigDB: if config.is_default: self.clear_default(config.capability) if config.is_analysis_default: self.clear_analysis_default() 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) if config.is_analysis_default: self.clear_analysis_default(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 clear_analysis_default(self, exclude_id: str | None = None) -> None: statement = select(ModelConfigDB).where(ModelConfigDB.is_analysis_default.is_(True)) for item in self.session.exec(statement).all(): if item.id == exclude_id: continue item.is_analysis_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, {})