diff --git a/backend/agenteval/storage/repository.py b/backend/agenteval/storage/repository.py index bae0404..183d1cc 100644 --- a/backend/agenteval/storage/repository.py +++ b/backend/agenteval/storage/repository.py @@ -1,6 +1,6 @@ """Repository layer for database access.""" -from typing import Optional +from typing import Generic, Optional, TypeVar from sqlmodel import Session, select @@ -18,131 +18,56 @@ from agenteval.storage.db import ( ) from agenteval.storage.model_config_repository import ScenarioModelBindingRepository - -def _target_to_db(target: EvalTarget) -> EvalTargetDB: - db = EvalTargetDB( - id=target.id, - name=target.name, - description=target.description, - platform=target.platform.value, - channel_type=target.channel_type.value, - status=target.status.value, - created_at=target.created_at, - updated_at=target.updated_at or utc_now(), - ) - db.set_config(target.channel_config) - return db +M = TypeVar("M") # domain model +DB = TypeVar("DB") # persisted table row -def _target_from_db(db: EvalTargetDB) -> EvalTarget: - return EvalTarget( - id=db.id, - name=db.name, - description=db.description, - platform=db.platform, - channel_type=db.channel_type, - channel_config=db.get_config(), - status=db.status, - created_at=db.created_at, - updated_at=db.updated_at, - ) +class BaseRepository(Generic[M, DB]): + """Shared CRUD skeleton for id-keyed entity repositories. + Subclasses declare the table (``_table``) and the ``list_all`` ordering + column name (``_order_by``, newest-first), and implement the ``_to_db`` / + ``_from_db`` converter pair. The converters are instance methods so a + subclass whose ``_from_db`` needs cross-table reads (e.g. Scenario's model + bindings) can reach ``self.session``. Entities with bespoke create/update + (binding validation, versioning) override just those methods. + """ -def _scenario_to_db(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 + _table: type + _order_by: str + def __init__(self, session: Optional[Session] = None): + self.session = session or get_session() -def _scenario_from_db(db: ScenarioDB, session: Session) -> Scenario: - bindings = ScenarioModelBindingRepository(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 _to_db(self, obj: M) -> DB: + raise NotImplementedError + def _from_db(self, db: DB) -> M: + raise NotImplementedError -def _run_to_db(run: EvalRun) -> EvalRunDB: - db = EvalRunDB( - id=run.id, - target_id=run.target_id, - scenario_id=run.scenario_id, - scenario_version=run.scenario_version, - campaign_id=run.campaign_id, - status=run.status.value, - triggered_by=run.triggered_by.value, - started_at=run.started_at, - completed_at=run.completed_at, - ) - if run.summary is not None: - db.set_summary(run.summary.model_dump(mode="json")) - return db + def list_all(self) -> list[M]: + column = getattr(self._table, self._order_by) + statement = select(self._table).order_by(column.desc()) + return [self._from_db(r) for r in self.session.exec(statement).all()] + def get(self, entity_id: str) -> Optional[M]: + db = self.session.get(self._table, entity_id) + return self._from_db(db) if db else None -def _run_from_db(db: EvalRunDB) -> EvalRun: - return EvalRun( - id=db.id, - target_id=db.target_id, - scenario_id=db.scenario_id, - scenario_version=db.scenario_version or 1, - campaign_id=db.campaign_id, - status=db.status, - triggered_by=db.triggered_by or "manual", - started_at=db.started_at, - completed_at=db.completed_at, - summary=db.get_summary(), - ) + def create(self, obj: M) -> M: + db = self._to_db(obj) + self.session.add(db) + self.session.commit() + self.session.refresh(db) + return self._from_db(db) - -def _campaign_to_db(campaign: Campaign) -> CampaignDB: - db = CampaignDB( - id=campaign.id, - name=campaign.name, - target_id=campaign.target_id, - window_seconds=campaign.window_seconds, - time_scale=campaign.time_scale, - status=campaign.status.value, - started_at=campaign.started_at, - completed_at=campaign.completed_at, - created_at=campaign.created_at, - ) - db.set_plan([entry.model_dump(mode="json") for entry in campaign.plan]) - if campaign.summary: - db.set_summary(campaign.summary) - return db - - -def _campaign_from_db(db: CampaignDB) -> Campaign: - return Campaign( - id=db.id, - name=db.name, - target_id=db.target_id, - window_seconds=db.window_seconds, - time_scale=db.time_scale, - plan=db.get_plan(), - status=db.status, - started_at=db.started_at, - completed_at=db.completed_at, - created_at=db.created_at, - summary=db.get_summary(), - ) + def delete(self, entity_id: str) -> bool: + db = self.session.get(self._table, entity_id) + if not db: + return False + self.session.delete(db) + self.session.commit() + return True def _result_to_db(result: EvalResult) -> EvalResultDB: @@ -171,26 +96,38 @@ def _result_from_db(db: EvalResultDB) -> EvalResult: ) -class TargetRepository: +class TargetRepository(BaseRepository[EvalTarget, EvalTargetDB]): """Repository for evaluation targets.""" - def __init__(self, session: Optional[Session] = None): - self.session = session or get_session() + _table = EvalTargetDB + _order_by = "created_at" - def list_all(self) -> list[EvalTarget]: - statement = select(EvalTargetDB).order_by(EvalTargetDB.created_at.desc()) - return [_target_from_db(r) for r in self.session.exec(statement).all()] + def _to_db(self, target: EvalTarget) -> EvalTargetDB: + db = EvalTargetDB( + id=target.id, + name=target.name, + description=target.description, + platform=target.platform.value, + channel_type=target.channel_type.value, + status=target.status.value, + created_at=target.created_at, + updated_at=target.updated_at or utc_now(), + ) + db.set_config(target.channel_config) + return db - def get(self, target_id: str) -> Optional[EvalTarget]: - db = self.session.get(EvalTargetDB, target_id) - return _target_from_db(db) if db else None - - def create(self, target: EvalTarget) -> EvalTarget: - db = _target_to_db(target) - self.session.add(db) - self.session.commit() - self.session.refresh(db) - return _target_from_db(db) + def _from_db(self, db: EvalTargetDB) -> EvalTarget: + return EvalTarget( + id=db.id, + name=db.name, + description=db.description, + platform=db.platform, + channel_type=db.channel_type, + channel_config=db.get_config(), + status=db.status, + created_at=db.created_at, + updated_at=db.updated_at, + ) def update(self, target: EvalTarget) -> Optional[EvalTarget]: existing = self.session.get(EvalTargetDB, target.id) @@ -206,33 +143,46 @@ class TargetRepository: self.session.add(existing) self.session.commit() self.session.refresh(existing) - return _target_from_db(existing) - - def delete(self, target_id: str) -> bool: - db = self.session.get(EvalTargetDB, target_id) - if not db: - return False - self.session.delete(db) - self.session.commit() - return True + return self._from_db(existing) -class ScenarioRepository: +class ScenarioRepository(BaseRepository[Scenario, ScenarioDB]): """Repository for evaluation scenarios.""" - def __init__(self, session: Optional[Session] = None): - self.session = session or get_session() + _table = ScenarioDB + _order_by = "created_at" - def list_all(self) -> list[Scenario]: - statement = select(ScenarioDB).order_by(ScenarioDB.created_at.desc()) - return [_scenario_from_db(r, self.session) for r in self.session.exec(statement).all()] + 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 get(self, scenario_id: str) -> Optional[Scenario]: - db = self.session.get(ScenarioDB, scenario_id) - return _scenario_from_db(db, self.session) if db else None + 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 = _scenario_to_db(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) @@ -244,7 +194,7 @@ class ScenarioRepository: except Exception: self.session.rollback() raise - return _scenario_from_db(db, self.session) + return self._from_db(db) def update(self, scenario: Scenario) -> Optional[Scenario]: existing = self.session.get(ScenarioDB, scenario.id) @@ -277,7 +227,7 @@ class ScenarioRepository: except Exception: self.session.rollback() raise - return _scenario_from_db(existing, self.session) + return self._from_db(existing) def delete(self, scenario_id: str) -> bool: db = self.session.get(ScenarioDB, scenario_id) @@ -293,15 +243,41 @@ class ScenarioRepository: return True -class RunRepository: +class RunRepository(BaseRepository[EvalRun, EvalRunDB]): """Repository for evaluation runs.""" - def __init__(self, session: Optional[Session] = None): - self.session = session or get_session() + _table = EvalRunDB + _order_by = "started_at" - def list_all(self) -> list[EvalRun]: - statement = select(EvalRunDB).order_by(EvalRunDB.started_at.desc()) - return [_run_from_db(r) for r in self.session.exec(statement).all()] + def _to_db(self, run: EvalRun) -> EvalRunDB: + db = EvalRunDB( + id=run.id, + target_id=run.target_id, + scenario_id=run.scenario_id, + scenario_version=run.scenario_version, + campaign_id=run.campaign_id, + status=run.status.value, + triggered_by=run.triggered_by.value, + started_at=run.started_at, + completed_at=run.completed_at, + ) + if run.summary is not None: + db.set_summary(run.summary.model_dump(mode="json")) + return db + + def _from_db(self, db: EvalRunDB) -> EvalRun: + return EvalRun( + id=db.id, + target_id=db.target_id, + scenario_id=db.scenario_id, + scenario_version=db.scenario_version or 1, + campaign_id=db.campaign_id, + status=db.status, + triggered_by=db.triggered_by or "manual", + started_at=db.started_at, + completed_at=db.completed_at, + summary=db.get_summary(), + ) def list_by_campaign(self, campaign_id: str) -> list[EvalRun]: statement = ( @@ -309,18 +285,7 @@ class RunRepository: .where(EvalRunDB.campaign_id == campaign_id) .order_by(EvalRunDB.started_at) ) - return [_run_from_db(r) for r in self.session.exec(statement).all()] - - def get(self, run_id: str) -> Optional[EvalRun]: - db = self.session.get(EvalRunDB, run_id) - return _run_from_db(db) if db else None - - def create(self, run: EvalRun) -> EvalRun: - db = _run_to_db(run) - self.session.add(db) - self.session.commit() - self.session.refresh(db) - return _run_from_db(db) + return [self._from_db(r) for r in self.session.exec(statement).all()] def mark_orphans_failed(self) -> int: """服务启动时清理:把遗留的 running/pending 运行标记为 failed。 @@ -357,7 +322,7 @@ class RunRepository: self.session.add(existing) self.session.commit() self.session.refresh(existing) - return _run_from_db(existing) + return self._from_db(existing) def get_turns(self, run_id: str) -> list[TurnDB]: statement = select(TurnDB).where(TurnDB.run_id == run_id).order_by(TurnDB.sent_at) @@ -367,36 +332,44 @@ class RunRepository: statement = select(EvalResultDB).where(EvalResultDB.run_id == run_id) return [_result_from_db(r) for r in self.session.exec(statement).all()] - def delete(self, run_id: str) -> bool: - """Delete a run. ORM-level cascade removes associated turns/results.""" - db = self.session.get(EvalRunDB, run_id) - if not db: - return False - self.session.delete(db) - self.session.commit() - return True - -class CampaignRepository: +class CampaignRepository(BaseRepository[Campaign, CampaignDB]): """Repository for evaluation campaigns (评估活动).""" - def __init__(self, session: Optional[Session] = None): - self.session = session or get_session() + _table = CampaignDB + _order_by = "created_at" - def list_all(self) -> list[Campaign]: - statement = select(CampaignDB).order_by(CampaignDB.created_at.desc()) - return [_campaign_from_db(r) for r in self.session.exec(statement).all()] + def _to_db(self, campaign: Campaign) -> CampaignDB: + db = CampaignDB( + id=campaign.id, + name=campaign.name, + target_id=campaign.target_id, + window_seconds=campaign.window_seconds, + time_scale=campaign.time_scale, + status=campaign.status.value, + started_at=campaign.started_at, + completed_at=campaign.completed_at, + created_at=campaign.created_at, + ) + db.set_plan([entry.model_dump(mode="json") for entry in campaign.plan]) + if campaign.summary: + db.set_summary(campaign.summary) + return db - def get(self, campaign_id: str) -> Optional[Campaign]: - db = self.session.get(CampaignDB, campaign_id) - return _campaign_from_db(db) if db else None - - def create(self, campaign: Campaign) -> Campaign: - db = _campaign_to_db(campaign) - self.session.add(db) - self.session.commit() - self.session.refresh(db) - return _campaign_from_db(db) + def _from_db(self, db: CampaignDB) -> Campaign: + return Campaign( + id=db.id, + name=db.name, + target_id=db.target_id, + window_seconds=db.window_seconds, + time_scale=db.time_scale, + plan=db.get_plan(), + status=db.status, + started_at=db.started_at, + completed_at=db.completed_at, + created_at=db.created_at, + summary=db.get_summary(), + ) def update(self, campaign: Campaign) -> Optional[Campaign]: existing = self.session.get(CampaignDB, campaign.id) @@ -415,7 +388,7 @@ class CampaignRepository: self.session.add(existing) self.session.commit() self.session.refresh(existing) - return _campaign_from_db(existing) + return self._from_db(existing) class ResultRepository: diff --git a/tests/unit/test_repository.py b/tests/unit/test_repository.py new file mode 100644 index 0000000..cc6abef --- /dev/null +++ b/tests/unit/test_repository.py @@ -0,0 +1,113 @@ +"""CRUD + serialization round-trip characterization for the repository layer. + +These lock the create → get → list_all → delete contract shared by the +Target / Run / Campaign repositories (the BaseRepository skeleton), plus the +summary/plan JSON round-trip, so the generic-base refactor stays behaviour- +preserving. Scenario's bespoke create/update (binding validation, versioning) +is covered by its own tests. +""" + +from datetime import datetime, timezone + +from agenteval.models import ( + Campaign, + CampaignPlanEntry, + ChannelType, + EvalRun, + EvalTarget, + PlatformType, + RunStatus, + RunSummary, + TargetStatus, +) +from agenteval.storage.repository import ( + CampaignRepository, + RunRepository, + TargetRepository, +) + + +def _make_target(tid: str = "t-1") -> EvalTarget: + return EvalTarget( + id=tid, + name="target", + platform=PlatformType.AI_DIGITAL_EMPLOYEE, + channel_type=ChannelType.TUTU_API, + channel_config={ + "base_url": "x", + "token": "x", + "tenant": "x", + "chat_channel_id": "x", + "chat_contact_id": "x", + }, + ) + + +def test_target_crud_round_trip(db_session): + repo = TargetRepository(db_session) + repo.create(_make_target()) + + fetched = repo.get("t-1") + assert fetched is not None + assert fetched.name == "target" + assert fetched.channel_config["base_url"] == "x" + assert fetched.status == TargetStatus.ACTIVE or fetched.status is not None + + assert [t.id for t in repo.list_all()] == ["t-1"] + + assert repo.delete("t-1") is True + assert repo.get("t-1") is None + assert repo.delete("t-1") is False + + +def test_target_list_all_newest_first(db_session): + repo = TargetRepository(db_session) + older = _make_target("t-old") + older.created_at = datetime(2026, 1, 1, tzinfo=timezone.utc) + newer = _make_target("t-new") + newer.created_at = datetime(2026, 6, 1, tzinfo=timezone.utc) + repo.create(older) + repo.create(newer) + assert [t.id for t in repo.list_all()] == ["t-new", "t-old"] + + +def test_run_summary_json_round_trip(db_session): + TargetRepository(db_session).create(_make_target()) + repo = RunRepository(db_session) + repo.create( + EvalRun( + id="r-1", + target_id="t-1", + scenario_id="s-1", + status=RunStatus.COMPLETED, + summary=RunSummary(total_cases=2, passed_cases=1, pass_rate=0.5), + ) + ) + fetched = repo.get("r-1") + assert fetched is not None + assert fetched.summary is not None + assert fetched.summary.total_cases == 2 + assert fetched.summary.pass_rate == 0.5 + + +def test_campaign_crud_and_plan_round_trip(db_session): + TargetRepository(db_session).create(_make_target()) + repo = CampaignRepository(db_session) + repo.create( + Campaign( + id="cp-1", + name="campaign", + target_id="t-1", + window_seconds=3600, + plan=[CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=2)], + ) + ) + fetched = repo.get("cp-1") + assert fetched is not None + assert fetched.plan[0].scenario_id == "s-1" + assert fetched.plan[0].count == 2 + assert [c.id for c in repo.list_all()] == ["cp-1"] + + # Campaign inherits the shared delete() from the base repository. + assert repo.delete("cp-1") is True + assert repo.get("cp-1") is None