| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103 |
- from __future__ import annotations
- from collections.abc import Iterable
- from sqlalchemy import func, select
- from sqlalchemy.dialects.mysql import insert
- from supply_infra.db.models.demand_grade import DemandGrade
- from supply_infra.db.repositories.base import BaseRepository
- _BATCH_SIZE = 500
- _UPSERT_COLUMNS = (
- "category_ids",
- "grade",
- "score",
- "prior_total_score",
- "posterior_rov_avg",
- "posterior_rov_count",
- "has_posterior",
- "related_pool_ids",
- "video_list",
- "strategies",
- "reason",
- )
- class DemandGradeRepository(BaseRepository[DemandGrade]):
- """需求分级结果表 — 按 (biz_dt, demand_name) 增量/更新写入。"""
- model = DemandGrade
- def get_existing_demand_names(self, biz_dt: str, names: Iterable[str] | None = None) -> set[str]:
- """返回指定业务日已分级的需求名集合;传入 names 时只在其中查交集。"""
- stmt = select(DemandGrade.demand_name).where(DemandGrade.biz_dt == biz_dt)
- if names is not None:
- name_list = [n for n in names if n]
- if not name_list:
- return set()
- stmt = stmt.where(DemandGrade.demand_name.in_(name_list))
- return {n for n in self.session.scalars(stmt).all() if n}
- def count_by_biz_dt(self, biz_dt: str) -> int:
- """统计指定业务日已分级的需求数。"""
- stmt = select(func.count()).select_from(DemandGrade).where(DemandGrade.biz_dt == biz_dt)
- return int(self.session.scalar(stmt) or 0)
- def list_by_biz_dt(self, biz_dt: str) -> list[DemandGrade]:
- """返回指定业务日的全部分级结果,按等级、需求名排序。"""
- stmt = (
- select(DemandGrade)
- .where(DemandGrade.biz_dt == biz_dt)
- .order_by(DemandGrade.grade, DemandGrade.demand_name)
- )
- return list(self.session.scalars(stmt).all())
- def list_by_biz_dt_and_grades(
- self,
- biz_dt: str,
- grades: Iterable[str] = ("S", "A"),
- ) -> list[DemandGrade]:
- """返回指定业务日、指定等级的分级结果,按等级、需求名排序。"""
- grade_list = [g for g in grades if g]
- if not grade_list:
- return []
- stmt = (
- select(DemandGrade)
- .where(DemandGrade.biz_dt == biz_dt, DemandGrade.grade.in_(grade_list))
- .order_by(DemandGrade.grade, DemandGrade.demand_name)
- )
- return list(self.session.scalars(stmt).all())
- def get_latest_biz_dt(self) -> str | None:
- """返回 demand_grade 中最新的业务日期。"""
- stmt = select(func.max(DemandGrade.biz_dt))
- return self.session.scalar(stmt)
- def get_ids_by_names(self, biz_dt: str, names: Iterable[str]) -> dict[str, int]:
- """按 (biz_dt, demand_name) 反查 id,供 upsert 后写关联表使用。"""
- name_list = [n for n in names if n]
- if not name_list:
- return {}
- stmt = select(DemandGrade.demand_name, DemandGrade.id).where(
- DemandGrade.biz_dt == biz_dt,
- DemandGrade.demand_name.in_(name_list),
- )
- return {name: int(id_) for name, id_ in self.session.execute(stmt).all()}
- def bulk_upsert(self, rows: list[dict]) -> int:
- """按 (biz_dt, demand_name) 批量 upsert。"""
- if not rows:
- return 0
- affected = 0
- for i in range(0, len(rows), _BATCH_SIZE):
- batch = rows[i : i + _BATCH_SIZE]
- stmt = insert(DemandGrade).values(batch)
- stmt = stmt.on_duplicate_key_update(
- **{col: stmt.inserted[col] for col in _UPSERT_COLUMNS}
- )
- result = self.session.execute(stmt)
- affected += result.rowcount or 0
- return affected
|