AgentEvalTool/backend/agenteval/storage/repository/scenario.py
sinohqb 3705945a7d test: 完整测试覆盖补全(+163 用例)
架构重构(候选 1-6):
- storage/repository.py 按域拆分为包(target/scenario/run/campaign/result)
- storage/db.py 按域拆分为包(eval/campaign/file/model_config/intelligent_eval)
- intelligent_eval/lifecycle.py 按状态机阶段拆分为包
- services/runs.py 编排逻辑下沉
- Campaigns.tsx 拆分为 campaigns/ 子组件

测试补全(候选 7):
前端(+125 用例,107→232):
- utils/ 纯函数:date/campaignTime/ruleLabels/fileTree/fileFormat/colors
- stores/tabStore 状态管理
- 核心组件:FormDrawer/PageWrapper/ChatBubble/GeneratedMessages/SectionHeader/StatCard/TurnList
- 业务组件:CaseBlock/CaseDetail/RuleOverview/WindowTimeline/RunList/TabBar/CampaignRunTimeline
- 文件管理:FileCategoryTree/FileTable
- hooks:sessionReducer/useFiles/useRunSession

后端(+38 用例,916→954):
- targets API CRUD + 404 路径
- WebSocket 连接管理器
- proxy 头部重写(CSP/X-Frame-Options)
- target 仓储 update 方法
- app 健康检查 + SPA 404
- scenarios 模板端点 + 404
- files API 边缘分支(404 场景 + 500 兜底)
- files service update_category
- 智能评估状态机迁移测试

门禁状态:
- 前端:tsc 干净 + 232 passed
- 后端:954 passed + ruff 全绿
2026-08-24 15:56:09 +08:00

113 lines
4.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Repository for evaluation scenarios."""
from typing import Optional
from sqlmodel import select
from agenteval.models import Case, Scenario
from agenteval.services.model_configs import ModelConfigService
from agenteval.storage.db import ScenarioDB, utc_now
from agenteval.storage.model_config_repository import ScenarioModelBindingRepository
from agenteval.storage.repository.base import BaseRepository
class ScenarioRepository(BaseRepository[Scenario, ScenarioDB]):
"""Repository for evaluation scenarios."""
_table = ScenarioDB
_order_by = "created_at"
def _to_db(self, scenario: Scenario) -> ScenarioDB:
db = ScenarioDB(
id=scenario.id,
name=scenario.name,
description=scenario.description,
created_at=scenario.created_at,
updated_at=scenario.updated_at or utc_now(),
)
db.set_tags(scenario.tags)
# mode="json" 与 update() 的考纲比较保持同一序列化形态,避免假升版
db.set_cases([case.model_dump(mode="json") for case in scenario.cases])
db.set_llm_config(scenario.llm_config)
return db
def _from_db(self, db: ScenarioDB) -> Scenario:
bindings = ScenarioModelBindingRepository(self.session).get_for_scenario(db.id or "")
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 create(self, scenario: Scenario) -> Scenario:
db = self._to_db(scenario)
bindings = {purpose.value: config_id for purpose, config_id in scenario.model_bindings.items()}
try:
ModelConfigService(self.session).validate_bindings(bindings)
self.session.add(db)
self.session.flush()
ScenarioModelBindingRepository(self.session).replace_for_scenario(db.id or "", bindings)
self.session.commit()
self.session.refresh(db)
except Exception:
self.session.rollback()
raise
return self._from_db(db)
def update(self, scenario: Scenario) -> Optional[Scenario]:
existing = self.session.get(ScenarioDB, scenario.id)
if not existing:
return None
bindings = {purpose.value: config_id for purpose, config_id in scenario.model_bindings.items()}
try:
ModelConfigService(self.session).validate_bindings(bindings)
# 考纲字段cases / model_bindings / llm_config变更才升版ADR-0001
# 版本由系统维护,忽略 scenario.version 的外部传入值。
new_cases = [case.model_dump(mode="json") for case in scenario.cases]
old_bindings = ScenarioModelBindingRepository(self.session).get_for_scenario(existing.id or "")
syllabus_changed = (
existing.get_cases() != new_cases
or existing.get_llm_config() != scenario.llm_config
or old_bindings != bindings
)
if syllabus_changed:
existing.version = (existing.version or 1) + 1
existing.name = scenario.name
existing.description = scenario.description
existing.set_tags(scenario.tags)
existing.set_cases(new_cases)
existing.set_llm_config(scenario.llm_config)
existing.updated_at = utc_now()
self.session.add(existing)
ScenarioModelBindingRepository(self.session).replace_for_scenario(existing.id or "", bindings)
self.session.commit()
self.session.refresh(existing)
except Exception:
self.session.rollback()
raise
return self._from_db(existing)
def delete(self, scenario_id: str) -> bool:
db = self.session.get(ScenarioDB, scenario_id)
if not db:
return False
try:
ScenarioModelBindingRepository(self.session).delete_for_scenario(scenario_id)
self.session.delete(db)
self.session.commit()
except Exception:
self.session.rollback()
raise
return True
def name_map(self) -> dict[str, str]:
"""scenario_id → 名称映射:报告 / 时间线 / 列表等读路径共用的场景名取法。"""
return {sid: name for sid, name in self.session.exec(select(ScenarioDB.id, ScenarioDB.name)).all()}