| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286 |
- """
- 批量保存单维度需求产生结果到 generated_demand 表。
- """
- from __future__ import annotations
- import logging
- import uuid
- from decimal import Decimal
- from typing import Any
- from agents.generate_demand_agent.tools.dim_constants import (
- DIM_KEYS,
- DIM_LABEL,
- build_category_path,
- resolve_biz_dt,
- )
- from supply_agent.tools import tool
- from supply_infra.db.repositories.demand_belong_category_repo import (
- DemandBelongCategoryRepository,
- )
- from supply_infra.db.repositories.demand_popularity_stats_repo import (
- DemandPopularityStatsRepository,
- )
- from supply_infra.db.repositories.generated_demand_repo import GeneratedDemandRepository
- from supply_infra.db.repositories.global_tree_category_repo import (
- GlobalTreeCategoryRepository,
- )
- from supply_infra.db.session import get_session
- logger = logging.getLogger(__name__)
- def _optional_int(value: Any) -> int | None:
- if value is None or value == "":
- return None
- return int(value)
- def _optional_str(value: Any) -> str | None:
- if value is None:
- return None
- text = str(value).strip()
- return text or None
- def _normalize_items(
- items: list[dict[str, Any]],
- *,
- default_run_id: str,
- ) -> tuple[list[dict[str, Any]], list[str]]:
- """校验并规范化待插入行,返回 (rows, errors)。"""
- rows: list[dict[str, Any]] = []
- errors: list[str] = []
- seen_keys: set[tuple[str, str]] = set()
- for idx, item in enumerate(items):
- if not isinstance(item, dict):
- errors.append(f"第 {idx} 项不是对象")
- continue
- source_dim = _optional_str(item.get("source_dim"))
- overall_direction = _optional_str(item.get("overall_direction"))
- summary_event = _optional_str(item.get("summary_event"))
- demand_name = _optional_str(item.get("demand_name"))
- if not source_dim:
- errors.append(f"第 {idx} 项缺少 source_dim")
- continue
- if source_dim not in DIM_KEYS:
- allowed = "、".join(f"{k}({DIM_LABEL[k]})" for k in DIM_KEYS)
- errors.append(f"第 {idx} 项 source_dim 无效(只能是:{allowed}): {source_dim}")
- continue
- if not overall_direction:
- errors.append(f"第 {idx} 项缺少 overall_direction")
- continue
- if not summary_event:
- errors.append(f"第 {idx} 项缺少 summary_event")
- continue
- if not demand_name:
- errors.append(f"第 {idx} 项缺少 demand_name")
- continue
- dedupe_key = (source_dim, demand_name)
- if dedupe_key in seen_keys:
- errors.append(
- f"第 {idx} 项在本次请求中重复: source_dim={source_dim}, demand_name={demand_name}"
- )
- continue
- seen_keys.add(dedupe_key)
- try:
- category_id = _optional_int(item.get("category_id"))
- except (TypeError, ValueError):
- errors.append(f"第 {idx} 项 category_id 无效: {item.get('category_id')!r}")
- continue
- try:
- demand_belong_id = _optional_int(item.get("demand_belong_id"))
- except (TypeError, ValueError):
- errors.append(
- f"第 {idx} 项 demand_belong_id 无效: {item.get('demand_belong_id')!r}"
- )
- continue
- run_id = _optional_str(item.get("run_id")) or default_run_id
- reason = _optional_str(item.get("reason"))
- category_path = _optional_str(item.get("category_path"))
- biz_dt = _optional_str(item.get("biz_dt"))
- dim_avg: Decimal | None = None
- dim_count: int | None = None
- if item.get("dim_avg") is not None and item.get("dim_avg") != "":
- try:
- dim_avg = Decimal(str(item.get("dim_avg")))
- except Exception:
- errors.append(f"第 {idx} 项 dim_avg 无效: {item.get('dim_avg')!r}")
- continue
- if item.get("dim_count") is not None and item.get("dim_count") != "":
- try:
- dim_count = int(item.get("dim_count"))
- except (TypeError, ValueError):
- errors.append(f"第 {idx} 项 dim_count 无效: {item.get('dim_count')!r}")
- continue
- rows.append(
- {
- "source_dim": source_dim,
- "overall_direction": overall_direction,
- "summary_event": summary_event,
- "demand_name": demand_name,
- "demand_belong_id": demand_belong_id,
- "category_id": category_id,
- "category_path": category_path,
- "dim_avg": dim_avg,
- "dim_count": dim_count,
- "reason": reason,
- "biz_dt": biz_dt,
- "run_id": run_id,
- "is_delete": 0,
- }
- )
- return rows, errors
- @tool
- def batch_save_generated_demands(
- items: list[dict[str, Any]],
- biz_dt: str | None = None,
- ) -> str:
- """
- 批量保存单维度需求产生结果。
- 写入 generated_demand 表。demand_name 必须已存在于 demand_belong_category;
- 不存在的词会被跳过。同一 run_id 内 (source_dim, demand_name) 去重。
- Args:
- items: 待保存列表。每项必填:
- - source_dim: 四维之一(ext_pop / plat_sust_pop / plat_ly_pop / recent_pop)
- - overall_direction: 整体方向
- - summary_event: 汇总事件
- - demand_name: 需求名(须来自 demand_belong_category.name)
- 选填:
- - category_id, category_path, demand_belong_id, reason,
- dim_avg, dim_count, biz_dt, run_id
- 若省略 run_id,本次调用自动生成同一 run_id。
- 若省略 demand_belong_id / category_id / 热度快照,将尽量从库中补全。
- biz_dt: 业务日 YYYYMMDD;省略则使用 demand_popularity_stats 最新业务日。
- 会作为本次落库的默认 biz_dt,并用于补全热度快照。
- Returns:
- 保存结果摘要。
- """
- if not items:
- return "items 不能为空"
- default_run_id = uuid.uuid4().hex
- rows, errors = _normalize_items(items, default_run_id=default_run_id)
- if not rows:
- detail = ";".join(errors) if errors else "无有效数据"
- return f"没有可保存的数据: {detail}"
- try:
- with get_session() as session:
- stats_repo = DemandPopularityStatsRepository(session)
- resolved_biz_dt, err = resolve_biz_dt(
- biz_dt,
- get_latest=stats_repo.get_latest_biz_dt,
- has_data=stats_repo.has_biz_dt,
- table_label="demand_popularity_stats 数据",
- )
- if err:
- return err
- belong_repo = DemandBelongCategoryRepository(session)
- names = [r["demand_name"] for r in rows]
- belong_by_name = belong_repo.get_by_names(names)
- valid_rows: list[dict[str, Any]] = []
- skipped_names: list[str] = []
- for row in rows:
- belong = belong_by_name.get(row["demand_name"])
- if belong is None:
- skipped_names.append(row["demand_name"])
- continue
- if row["demand_belong_id"] is None:
- row["demand_belong_id"] = int(belong.id)
- if row["category_id"] is None:
- row["category_id"] = int(belong.category_id)
- valid_rows.append(row)
- if not valid_rows:
- detail = "、".join(skipped_names)
- extra = f";校验失败: {';'.join(errors)}" if errors else ""
- return f"没有可保存的数据: demand_name 不存在于 demand_belong_category: {detail}{extra}"
- # 补全路径
- need_path_ids = [
- int(r["category_id"])
- for r in valid_rows
- if r["category_id"] is not None and not r.get("category_path")
- ]
- if need_path_ids:
- categories = GlobalTreeCategoryRepository(session).list_active_categories()
- by_id = {int(c.id): c for c in categories}
- for row in valid_rows:
- if row.get("category_path") or row["category_id"] is None:
- continue
- path = build_category_path(int(row["category_id"]), by_id)
- if path:
- row["category_path"] = path
- # 补全 biz_dt 与热度快照
- need_stats_ids = [
- int(r["demand_belong_id"])
- for r in valid_rows
- if r["demand_belong_id"] is not None
- and (r.get("dim_avg") is None or r.get("dim_count") is None)
- ]
- stats_by_belong: dict[int, Any] = {}
- if need_stats_ids:
- for stats in stats_repo.list_by_biz_dt_and_belong_ids(
- resolved_biz_dt, need_stats_ids
- ):
- stats_by_belong[int(stats.demand_category_id)] = stats
- for row in valid_rows:
- if not row.get("biz_dt"):
- row["biz_dt"] = resolved_biz_dt
- belong_id = row.get("demand_belong_id")
- if belong_id is None:
- continue
- stats = stats_by_belong.get(int(belong_id))
- if stats is None:
- continue
- dim = row["source_dim"]
- if row.get("dim_count") is None:
- row["dim_count"] = int(getattr(stats, f"{dim}_count", 0) or 0)
- if row.get("dim_avg") is None:
- avg = getattr(stats, f"{dim}_avg", None)
- row["dim_avg"] = Decimal(str(avg)) if avg is not None else None
- inserted = GeneratedDemandRepository(session).bulk_insert(valid_rows)
- run_ids = sorted({str(r["run_id"]) for r in valid_rows})
- parts = [
- f"提交有效 {len(valid_rows)} 条,成功插入 {inserted} 条",
- f"run_id={','.join(run_ids)}",
- f"biz_dt={resolved_biz_dt}",
- ]
- if skipped_names:
- parts.append(
- f"因 demand_name 不存在跳过 {len(skipped_names)} 条: "
- + "、".join(skipped_names[:20])
- + ("…" if len(skipped_names) > 20 else "")
- )
- if errors:
- parts.append(f"校验失败 {len(errors)} 条: " + ";".join(errors[:20]))
- message = "。".join(parts)
- logger.info("batch_save_generated_demands completed: %s", message)
- return message
- except Exception as e:
- logger.error("batch_save_generated_demands failed: %s", e, exc_info=True)
- return f"批量保存生成需求失败: {e}"
|