- SQLite 启用 WAL,允许读写并发 - 新增 6 个索引(eval_runs.status/campaign_id、eval_results.run_id、 turns.run_id、intelligent_evals.status、task_queue.assigned_at) - 幂等 Alembic 迁移(列/索引存在性检查) - domain.py 计数改 func.count 聚合,get_attention_reason 单次加载 sessions - scenario list_all 批量加载 bindings(1+N → 2 查询) - mark_orphans_failed 批量加载 campaigns(N → 1 IN 查询)
127 lines
5.0 KiB
Python
127 lines
5.0 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 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 get_all_for_scenarios(self, scenario_ids: list[str]) -> dict[str, dict[str, str]]:
|
|
"""Batch load bindings for multiple scenarios to avoid N+1 queries."""
|
|
if not scenario_ids:
|
|
return {}
|
|
statement = select(ScenarioModelBindingDB).where(
|
|
ScenarioModelBindingDB.scenario_id.in_(scenario_ids),
|
|
)
|
|
result: dict[str, dict[str, str]] = {sid: {} for sid in scenario_ids}
|
|
for item in self.session.exec(statement).all():
|
|
result[item.scenario_id][item.purpose] = item.model_config_id
|
|
return result
|
|
|
|
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, {})
|