refactor(storage): converge CRUD skeleton into BaseRepository
Target/Scenario/Run/Campaign repositories repeated the same __init__/list_all/get/create/delete skeleton (~120 lines). Extract a generic BaseRepository[M, DB]: subclasses declare the table + ordering column and implement instance-method _to_db/_from_db converters (so Scenario's _from_db can reach self.session for model bindings). Bespoke paths (Scenario create/delete, all update) stay per-subclass. Adds test_repository.py locking the CRUD + JSON round-trip contract.
This commit is contained in:
parent
9c01afa79b
commit
d411572607
@ -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:
|
||||
|
||||
113
tests/unit/test_repository.py
Normal file
113
tests/unit/test_repository.py
Normal file
@ -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
|
||||
Loading…
Reference in New Issue
Block a user