Introduce ModelPurpose.ANALYSIS and a globally-unique is_analysis_default marker on chat model configs so campaign analysis can resolve its model. Service rejects disabled or non-chat configs; repo clears the previous holder on set. Documented the analysis role in CONTEXT.md.
111 lines
4.2 KiB
Python
111 lines
4.2 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)
|
|
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, {})
|