AgentEvalTool/backend/agenteval/storage/model_config_repository.py
sinohqb ca208232c7 perf(db): WAL 模式 + 性能索引 + N+1 查询消除
- 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 查询)
2026-08-24 23:17:51 +08:00

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, {})