""" 批量保存单维度需求产生结果到 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}"