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 查询)
This commit is contained in:
sinohqb 2026-08-24 23:17:51 +08:00
parent f8f87815c6
commit ca208232c7
8 changed files with 176 additions and 15 deletions

View File

@ -7,7 +7,7 @@ delegate here so a single source of truth governs "8-10h"-style semantics.
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Optional from typing import Optional
from sqlmodel import Session, select from sqlmodel import Session, func, select
from agenteval.storage.db import ( from agenteval.storage.db import (
IntelligentEvalDB, IntelligentEvalDB,
@ -82,21 +82,21 @@ def count_sessions_in_slot(
start_hour, end_hour = parsed start_hour, end_hour = parsed
slot_start = (as_utc(eval_started_at) + timedelta(hours=start_hour)).replace(tzinfo=None) slot_start = (as_utc(eval_started_at) + timedelta(hours=start_hour)).replace(tzinfo=None)
slot_end = (as_utc(eval_started_at) + timedelta(hours=end_hour)).replace(tzinfo=None) slot_end = (as_utc(eval_started_at) + timedelta(hours=end_hour)).replace(tzinfo=None)
rows = session.exec( return session.exec(
select(IntelligentEvalSessionDB).where( select(func.count(IntelligentEvalSessionDB.id)).where(
IntelligentEvalSessionDB.eval_id == eval_id, IntelligentEvalSessionDB.eval_id == eval_id,
IntelligentEvalSessionDB.created_at >= slot_start, IntelligentEvalSessionDB.created_at >= slot_start,
IntelligentEvalSessionDB.created_at < slot_end, IntelligentEvalSessionDB.created_at < slot_end,
) )
).all() ).one()
return len(rows)
def count_total_sessions(eval_id: str, session: Session) -> int: def count_total_sessions(eval_id: str, session: Session) -> int:
rows = session.exec( return session.exec(
select(IntelligentEvalSessionDB).where(IntelligentEvalSessionDB.eval_id == eval_id) select(func.count(IntelligentEvalSessionDB.id)).where(
).all() IntelligentEvalSessionDB.eval_id == eval_id
return len(rows) )
).one()
def calculate_session_deficit(eval_db: IntelligentEvalDB, session: Session) -> int: def calculate_session_deficit(eval_db: IntelligentEvalDB, session: Session) -> int:
@ -146,12 +146,21 @@ def get_attention_reason(eval_db: IntelligentEvalDB, session: Session) -> Option
plan = eval_db.get_plan() plan = eval_db.get_plan()
time_distribution = plan.get("time_distribution", []) time_distribution = plan.get("time_distribution", [])
current_offset = utc_now() - as_utc(eval_db.started_at) current_offset = utc_now() - as_utc(eval_db.started_at)
if any(is_slot_due(s, current_offset) for s in time_distribution):
if calculate_session_deficit(eval_db, session) > 0: # Load sessions once and reuse for both checks
return "slot_due"
sessions = session.exec( sessions = session.exec(
select(IntelligentEvalSessionDB).where(IntelligentEvalSessionDB.eval_id == eval_db.id) select(IntelligentEvalSessionDB).where(IntelligentEvalSessionDB.eval_id == eval_db.id)
).all() ).all()
# Check slot due with deficit
if any(is_slot_due(s, current_offset) for s in time_distribution):
should_have = sum(
slot.get("sessions", 0) for slot in time_distribution if is_slot_due(slot, current_offset)
)
actual = len(sessions)
if should_have > actual:
return "slot_due"
# ADR-0011expired/failed 同为会话终态——存在过期会话时也必须触发 analyst # ADR-0011expired/failed 同为会话终态——存在过期会话时也必须触发 analyst
# 否则评估永远等不到"全部完成"而卡在 executing # 否则评估永远等不到"全部完成"而卡在 executing
if sessions and all(s.status in ("completed", "failed", "expired") for s in sessions): if sessions and all(s.status in ("completed", "failed", "expired") for s in sessions):

View File

@ -10,6 +10,7 @@ from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
import sqlalchemy as sa
from sqlalchemy.pool import StaticPool from sqlalchemy.pool import StaticPool
from sqlmodel import Session, SQLModel, create_engine from sqlmodel import Session, SQLModel, create_engine
@ -26,6 +27,12 @@ engine = create_engine(
poolclass=StaticPool, poolclass=StaticPool,
) )
# WAL mode allows concurrent reads during writes — critical for the evaluation
# engine's high-frequency turn/result writes while the frontend polls.
with engine.connect() as _conn:
_conn.execute(sa.text("PRAGMA journal_mode=WAL"))
_conn.commit()
def utc_now() -> datetime: def utc_now() -> datetime:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)

View File

@ -87,6 +87,8 @@ class EvalRunDB(SQLModel, table=True):
"campaign_occurrence_index", "campaign_occurrence_index",
unique=True, unique=True,
), ),
sa.Index("ix_eval_runs_status", "status"),
sa.Index("ix_eval_runs_campaign_id", "campaign_id"),
) )
id: Optional[str] = Field(default_factory=new_uuid, primary_key=True) id: Optional[str] = Field(default_factory=new_uuid, primary_key=True)
@ -124,6 +126,7 @@ class TurnDB(SQLModel, table=True):
"""Database table for conversation turns.""" """Database table for conversation turns."""
__tablename__ = "turns" __tablename__ = "turns"
__table_args__ = (sa.Index("ix_turns_run_id", "run_id"),)
id: Optional[str] = Field(default_factory=new_uuid, primary_key=True) id: Optional[str] = Field(default_factory=new_uuid, primary_key=True)
run_id: Optional[str] = Field(default=None, foreign_key="eval_runs.id") run_id: Optional[str] = Field(default=None, foreign_key="eval_runs.id")
@ -155,6 +158,7 @@ class EvalResultDB(SQLModel, table=True):
"""Database table for evaluation results.""" """Database table for evaluation results."""
__tablename__ = "eval_results" __tablename__ = "eval_results"
__table_args__ = (sa.Index("ix_eval_results_run_id", "run_id"),)
id: Optional[str] = Field(default_factory=new_uuid, primary_key=True) id: Optional[str] = Field(default_factory=new_uuid, primary_key=True)
run_id: Optional[str] = Field(default=None, foreign_key="eval_runs.id") run_id: Optional[str] = Field(default=None, foreign_key="eval_runs.id")

View File

@ -13,6 +13,7 @@ class IntelligentEvalDB(SQLModel, table=True):
"""Intelligent evaluation (智能评估) — independent entity, peer to Campaign.""" """Intelligent evaluation (智能评估) — independent entity, peer to Campaign."""
__tablename__ = "intelligent_evals" __tablename__ = "intelligent_evals"
__table_args__ = (sa.Index("ix_intelligent_evals_status", "status"),)
id: Optional[str] = Field(default_factory=new_uuid, primary_key=True) id: Optional[str] = Field(default_factory=new_uuid, primary_key=True)
name: str name: str
@ -138,6 +139,7 @@ class IntelligentEvalTaskQueueDB(SQLModel, table=True):
__table_args__ = ( __table_args__ = (
sa.Index("idx_task_queue_status_priority", "status", "priority"), sa.Index("idx_task_queue_status_priority", "status", "priority"),
sa.Index("idx_task_queue_eval_status", "eval_id", "status"), sa.Index("idx_task_queue_eval_status", "eval_id", "status"),
sa.Index("idx_task_queue_assigned_at", "assigned_at"),
) )

View File

@ -95,6 +95,18 @@ class ScenarioModelBindingRepository:
) )
return {item.purpose: item.model_config_id for item in self.session.exec(statement).all()} 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: def replace_for_scenario(self, scenario_id: str, bindings: dict[str, str]) -> None:
statement = select(ScenarioModelBindingDB).where( statement = select(ScenarioModelBindingDB).where(
ScenarioModelBindingDB.scenario_id == scenario_id, ScenarioModelBindingDB.scenario_id == scenario_id,

View File

@ -199,9 +199,17 @@ class RunRepository(BaseRepository[EvalRun, EvalRunDB]):
""" """
statement = select(EvalRunDB).where(EvalRunDB.status.in_(["running", "pending"])) # type: ignore[attr-defined] statement = select(EvalRunDB).where(EvalRunDB.status.in_(["running", "pending"])) # type: ignore[attr-defined]
candidates = self.session.exec(statement).all() candidates = self.session.exec(statement).all()
# Batch load all campaigns for pending runs to avoid N+1 queries
campaign_ids = {db.campaign_id for db in candidates if db.campaign_id and db.status == RunStatus.PENDING.value}
campaigns = {}
if campaign_ids:
campaign_stmt = select(CampaignDB).where(CampaignDB.id.in_(campaign_ids)) # type: ignore[attr-defined]
campaigns = {c.id: c for c in self.session.exec(campaign_stmt).all()}
failed: list[EvalRunDB] = [] failed: list[EvalRunDB] = []
for db in candidates: for db in candidates:
if db.status == RunStatus.PENDING.value and self._is_recoverable_campaign_pending(db): if db.status == RunStatus.PENDING.value and self._is_recoverable_campaign_pending(db, campaigns):
continue continue
self._set_interrupted(db) self._set_interrupted(db)
failed.append(db) failed.append(db)
@ -209,10 +217,10 @@ class RunRepository(BaseRepository[EvalRun, EvalRunDB]):
self.session.commit() self.session.commit()
return len(failed) return len(failed)
def _is_recoverable_campaign_pending(self, db: EvalRunDB) -> bool: def _is_recoverable_campaign_pending(self, db: EvalRunDB, campaigns: dict[str, CampaignDB]) -> bool:
if db.campaign_id is None or db.campaign_plan_index is None or db.campaign_occurrence_index is None: if db.campaign_id is None or db.campaign_plan_index is None or db.campaign_occurrence_index is None:
return False return False
campaign = self.session.get(CampaignDB, db.campaign_id) campaign = campaigns.get(db.campaign_id)
return campaign is not None and campaign.status == CampaignStatus.RUNNING.value return campaign is not None and campaign.status == CampaignStatus.RUNNING.value
def _set_interrupted(self, db: EvalRunDB) -> None: def _set_interrupted(self, db: EvalRunDB) -> None:

View File

@ -107,6 +107,34 @@ class ScenarioRepository(BaseRepository[Scenario, ScenarioDB]):
raise raise
return True return True
def list_all(self) -> list[Scenario]:
"""Batch load scenarios with bindings to avoid N+1 queries."""
column = getattr(self._table, self._order_by)
statement = select(self._table).order_by(column.desc())
scenarios = self.session.exec(statement).all()
# Batch load all bindings at once
all_bindings = ScenarioModelBindingRepository(self.session).get_all_for_scenarios(
[s.id for s in scenarios if s.id]
)
return [self._from_db_with_bindings(s, all_bindings.get(s.id or "", [])) for s in scenarios]
def _from_db_with_bindings(self, db: ScenarioDB, bindings: list) -> Scenario:
"""Convert DB model to domain model with pre-loaded bindings."""
return Scenario(
id=db.id,
name=db.name,
description=db.description,
tags=db.get_tags(),
cases=[Case(**case) for case in db.get_cases()],
model_bindings=bindings,
llm_config=db.get_llm_config(),
version=db.version or 1,
created_at=db.created_at,
updated_at=db.updated_at,
)
def name_map(self) -> dict[str, str]: def name_map(self) -> dict[str, str]:
"""scenario_id → 名称映射:报告 / 时间线 / 列表等读路径共用的场景名取法。""" """scenario_id → 名称映射:报告 / 时间线 / 列表等读路径共用的场景名取法。"""
return {sid: name for sid, name in self.session.exec(select(ScenarioDB.id, ScenarioDB.name)).all()} return {sid: name for sid, name in self.session.exec(select(ScenarioDB.id, ScenarioDB.name)).all()}

View File

@ -0,0 +1,91 @@
"""add performance indexes for query optimization
Revision ID: d5f6193ba7c8
Revises: e1a2b3c4d5f6
Create Date: 2026-08-24 21:16:08.829576
"""
from typing import Sequence, Union
from alembic import op
from sqlalchemy import inspect
# revision identifiers, used by Alembic.
revision: str = 'd5f6193ba7c8'
down_revision: Union[str, Sequence[str], None] = 'e1a2b3c4d5f6'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def _index_exists(inspector, table_name: str, index_name: str) -> bool:
"""Check if an index exists on a table."""
indexes = inspector.get_indexes(table_name)
return any(idx['name'] == index_name for idx in indexes)
def _column_exists(inspector, table_name: str, column_name: str) -> bool:
"""Check if a column exists on a table."""
columns = inspector.get_columns(table_name)
return any(col['name'] == column_name for col in columns)
def upgrade() -> None:
"""Add indexes for query performance optimization."""
conn = op.get_bind()
inspector = inspect(conn)
# Only create indexes that don't already exist and whose columns exist
if _column_exists(inspector, 'eval_results', 'run_id') and not _index_exists(inspector, 'eval_results', 'ix_eval_results_run_id'):
with op.batch_alter_table('eval_results', schema=None) as batch_op:
batch_op.create_index('ix_eval_results_run_id', ['run_id'], unique=False)
if _column_exists(inspector, 'eval_runs', 'campaign_id') and not _index_exists(inspector, 'eval_runs', 'ix_eval_runs_campaign_id'):
with op.batch_alter_table('eval_runs', schema=None) as batch_op:
batch_op.create_index('ix_eval_runs_campaign_id', ['campaign_id'], unique=False)
if _column_exists(inspector, 'eval_runs', 'status') and not _index_exists(inspector, 'eval_runs', 'ix_eval_runs_status'):
with op.batch_alter_table('eval_runs', schema=None) as batch_op:
batch_op.create_index('ix_eval_runs_status', ['status'], unique=False)
if _column_exists(inspector, 'intelligent_eval_task_queue', 'assigned_at') and not _index_exists(inspector, 'intelligent_eval_task_queue', 'idx_task_queue_assigned_at'):
with op.batch_alter_table('intelligent_eval_task_queue', schema=None) as batch_op:
batch_op.create_index('idx_task_queue_assigned_at', ['assigned_at'], unique=False)
if _column_exists(inspector, 'intelligent_evals', 'status') and not _index_exists(inspector, 'intelligent_evals', 'ix_intelligent_evals_status'):
with op.batch_alter_table('intelligent_evals', schema=None) as batch_op:
batch_op.create_index('ix_intelligent_evals_status', ['status'], unique=False)
if _column_exists(inspector, 'turns', 'run_id') and not _index_exists(inspector, 'turns', 'ix_turns_run_id'):
with op.batch_alter_table('turns', schema=None) as batch_op:
batch_op.create_index('ix_turns_run_id', ['run_id'], unique=False)
def downgrade() -> None:
"""Remove performance indexes."""
conn = op.get_bind()
inspector = inspect(conn)
if _index_exists(inspector, 'turns', 'ix_turns_run_id'):
with op.batch_alter_table('turns', schema=None) as batch_op:
batch_op.drop_index('ix_turns_run_id')
if _index_exists(inspector, 'intelligent_evals', 'ix_intelligent_evals_status'):
with op.batch_alter_table('intelligent_evals', schema=None) as batch_op:
batch_op.drop_index('ix_intelligent_evals_status')
if _index_exists(inspector, 'intelligent_eval_task_queue', 'idx_task_queue_assigned_at'):
with op.batch_alter_table('intelligent_eval_task_queue', schema=None) as batch_op:
batch_op.drop_index('idx_task_queue_assigned_at')
if _index_exists(inspector, 'eval_runs', 'ix_eval_runs_status'):
with op.batch_alter_table('eval_runs', schema=None) as batch_op:
batch_op.drop_index('ix_eval_runs_status')
if _index_exists(inspector, 'eval_runs', 'ix_eval_runs_campaign_id'):
with op.batch_alter_table('eval_runs', schema=None) as batch_op:
batch_op.drop_index('ix_eval_runs_campaign_id')
if _index_exists(inspector, 'eval_results', 'ix_eval_results_run_id'):
with op.batch_alter_table('eval_results', schema=None) as batch_op:
batch_op.drop_index('ix_eval_results_run_id')