# -*- coding:utf-8 -*- """挖掘批次编排:状态机驱动的全流程(pulling→embedding→clustering→naming→done/failed)。 入口 start_mining() 同步建批次记录后立即返回 batch_id,全流程 asyncio.create_task 后台跑(对齐 executor.py 范式);进度/错误实时落 opp_mining_batches(stats_json/error_msg), 前端轮询 opp_mining_status 可见——失败必须是可行动报错(用户铁律:禁前端挂起干等)。 """ import asyncio import json import logging import time from .opp_common import get_db, new_id, rows_to_dicts from .opp_mining import (get_mining_params, pull_demands, embed_snapshots, build_batch_collection, knn_graph, cluster_vectors, name_clusters, can_batch_transition, _vdb_base, vdb_post) logger = logging.getLogger("pipeline.opp_mining") _running = {} # batch_id -> asyncio.Task(防重复起跑,进程内) async def _set_status(sor, batch_id, cur, nxt, **extra): if not can_batch_transition(cur, nxt): raise RuntimeError("非法状态迁移 %s→%s" % (cur, nxt)) sets = ["status=${nxt}$", "updated_at=CURRENT_TIMESTAMP"] args = {"nxt": nxt, "bid": batch_id} for k, v in extra.items(): sets.append("%s=${%s}$" % (k, k)) args[k] = v await sor.sqlExe( "UPDATE opp_mining_batches SET " + ", ".join(sets) + " WHERE id=${bid}$", args) await sor.sqlExe("COMMIT", {}) async def start_mining(sor, ctx, scope="top", days=365, sources="", keyword="", category="", created_by=""): """创建挖掘批次并后台执行。返回 (batch_id, err)。 scope=top:全量需求 → 聚类 → TopX 排名。 scope=targeted:用户指定类型(keyword/category 过滤)→ 范围内聚类找子型。 """ org_id = ctx.get("org_id") or "" if scope not in ("top", "targeted"): return "", "scope 必须是 top 或 targeted" if scope == "targeted" and not (keyword or category): return "", "targeted 模式必须指定 keyword 或 category(如:合同管理)" bid = new_id() params_json = json.dumps({"scope": scope, "days": days, "sources": sources, "keyword": keyword, "category": category}, ensure_ascii=False) await sor.C("opp_mining_batches", { "id": bid, "org_id": org_id, "project_id": ctx.get("project_id") or "", # 项目归属(空=平台级挖掘,对齐报告口径) "scope": scope, "params_json": params_json, "status": "pulling", "vdb_col": "", "stats_json": "", "error_msg": "", "created_by": created_by or ctx.get("user_id") or "", }) await sor.sqlExe("COMMIT", {}) if bid in _running and not _running[bid].done(): return bid, "" task = asyncio.create_task(_run_mining(bid, ctx)) _running[bid] = task return bid, "" async def _run_mining(batch_id, ctx): db, DBNAME = get_db() try: async with db.sqlorContext(DBNAME) as sor: await _run_mining_inner(sor, batch_id, ctx) except Exception as e: logger.exception("mining batch %s crashed", batch_id) try: async with db.sqlorContext(DBNAME) as sor: recs = await sor.sqlExe( "SELECT status FROM opp_mining_batches WHERE id=${b}$", {"b": batch_id}) await sor.sqlExe("COMMIT", {}) cur = recs[0].status if recs else "pulling" await _set_status(sor, batch_id, cur, "failed", error_msg="批次异常: %s" % str(e)[:500]) except Exception: pass async def _run_mining_inner(sor, batch_id, ctx): t0 = time.time() recs = await sor.sqlExe( "SELECT * FROM opp_mining_batches WHERE id=${b}$", {"b": batch_id}) await sor.sqlExe("COMMIT", {}) if not recs: raise RuntimeError("批次不存在: %s" % batch_id) batch = rows_to_dicts(recs, limit=1)[0] params = json.loads(batch.get("params_json") or "{}") cfg = await get_mining_params(sor) def stats(**kw): kw["elapsed_s"] = round(time.time() - t0, 1) return json.dumps(kw, ensure_ascii=False) # ── 1. pulling ── n, err = await pull_demands( sor, batch, ctx, cfg, params.get("scope", "top"), keyword=params.get("keyword", ""), category=params.get("category", ""), days=int(params.get("days") or 365), sources=params.get("sources", "")) if err: return await _set_status(sor, batch_id, "pulling", "failed", error_msg=err, stats_json=stats(demands=n)) if n == 0: return await _set_status(sor, batch_id, "pulling", "done", stats_json=stats(demands=0, clusters=0, note="范围内无需求数据")) await _set_status(sor, batch_id, "pulling", "embedding", stats_json=stats(demands=n)) # ── 2. embedding(增量缓存,失败可行动报错)── project_id = ctx.get("project_id") or "" embed_res, err = await embed_snapshots(sor, batch, project_id, cfg) if err: return await _set_status(sor, batch_id, "embedding", "failed", error_msg=err, stats_json=stats(demands=n)) vecs = embed_res["vecs"] text_map = embed_res["text_map"] if not vecs: return await _set_status(sor, batch_id, "embedding", "failed", error_msg="全部需求向量化失败", stats_json=stats(demands=n)) await _set_status(sor, batch_id, "embedding", "clustering", stats_json=stats(demands=n, embedded=len(vecs))) # ── 3. clustering(工作集 + kNN + 并查集)── base = await _vdb_base(sor) col, err = await build_batch_collection(sor, base, batch, cfg, embed_res) if err: return await _set_status(sor, batch_id, "clustering", "failed", error_msg=err, stats_json=stats(demands=n, embedded=len(vecs))) await sor.sqlExe("UPDATE opp_mining_batches SET vdb_col=${c}$ WHERE id=${b}$", {"c": col, "b": batch_id}) await sor.sqlExe("COMMIT", {}) nbr, err = await knn_graph(sor, base, col, vecs, cfg["opp_mine_topk"]) if err: return await _set_status(sor, batch_id, "clustering", "failed", error_msg=err, stats_json=stats(demands=n, embedded=len(vecs))) clusters, other = cluster_vectors(vecs, nbr, cfg["opp_mine_tau"], cfg["opp_mine_min_cluster"], cfg["opp_mine_big_split"]) await _set_status(sor, batch_id, "clustering", "naming", stats_json=stats(demands=n, embedded=len(vecs), clusters=len(clusters), other=len(other), tau=cfg["opp_mine_tau"])) # ── 4. naming(LLM utility + 规则兜底)── named = await name_clusters(sor, batch, clusters, vecs, text_map, cfg, ctx) org_id = ctx.get("org_id") or "" proj_id = str(batch.get("project_id") or "") # 簇归属跟随批次(写入时已落库) # vid → snap_id 反查 vid2snap = {v: s for s, v in embed_res["id_map"].items()} for cl in named: cid = new_id() cent_snap = vid2snap.get(cl["centroid_vid"], "") await sor.C("opp_clusters", { "id": cid, "batch_id": batch_id, "org_id": org_id, "project_id": proj_id, "name": cl["name"][:128], "doc_count": cl["size"], "share": cl["share"], "heat_rank": cl["rank"], "naming_evidence": json.dumps( {"samples": cl["samples"], "by": cl["evidence"]}, ensure_ascii=False), "centroid_snap_id": cent_snap, }) await sor.sqlExe("COMMIT", {}) # 回填快照 cluster_id(分批 UPDATE,参数化防注入) snaps = [vid2snap[v] for v in cl["members"] if v in vid2snap] for i in range(0, len(snaps), 200): chunk = snaps[i:i + 200] ph = ",".join("${s%d}$" % j for j in range(len(chunk))) args = {"cid": cid} for j, sv in enumerate(chunk): args["s%d" % j] = sv await sor.sqlExe( "UPDATE opp_demand_snap SET cluster_id=${cid}$ WHERE id IN (%s)" % ph, args) await sor.sqlExe("COMMIT", {}) await _set_status(sor, batch_id, "naming", "done", stats_json=stats(demands=n, embedded=len(vecs), clusters=len(named), other=len(other), tau=cfg["opp_mine_tau"], top1=(named[0]["name"] if named else ""))) # ── 5. 清理:保留最近 N 批工作集,旧批 drop(缓存 collection 永不清)── await _cleanup_old_collections(sor, base, org_id, cfg) async def _cleanup_old_collections(sor, base, org_id, cfg): keep = cfg["opp_mine_keep_batches"] recs = await sor.sqlExe( "SELECT id, vdb_col FROM opp_mining_batches WHERE org_id=${o}$ " "ORDER BY created_at DESC LIMIT 200", {"o": org_id}) await sor.sqlExe("COMMIT", {}) cols = [r.vdb_col for r in (recs or []) if getattr(r, "vdb_col", "")] for col in cols[keep:]: await vdb_post(sor, base, "/v1/dropcollection", {"colname": col}) await sor.sqlExe("UPDATE opp_mining_batches SET vdb_col='' WHERE vdb_col=${c}$", {"c": col}) await sor.sqlExe("COMMIT", {}) logger.debug("dropped old batch collection %s", col) async def mining_status(sor, ctx, batch_id=""): """查批次状态(隔离:本机构 + 本项目/平台级)。batch_id 空 = 可见批次列表。 可见性口径与 opp_reports 对齐(2026-09-12 用户确认挖掘也分项目): - 挂项目的批次:仅当前会话项目内可见(project_id 相等) - 平台级批次(project_id 空):本机构登录可见 - 会话无当前项目时:只见平台级批次 """ org_id = ctx.get("org_id") or "" proj_id = ctx.get("project_id") or "" if batch_id: recs = await sor.sqlExe( "SELECT id, scope, status, project_id, stats_json, error_msg, vdb_col, created_at " "FROM opp_mining_batches WHERE id=${b}$ AND org_id=${o}$", {"b": batch_id, "o": org_id}) await sor.sqlExe("COMMIT", {}) if not recs: return False, "批次不存在或不属于当前机构" row = rows_to_dicts(recs, limit=1)[0] bproj = str(row.get("project_id") or "") if bproj and bproj != proj_id: return False, "批次属于其他项目,当前会话不可见" return True, row # 列表:本项目批次 + 平台级批次 recs = await sor.sqlExe( "SELECT id, scope, status, project_id, stats_json, error_msg, created_at " "FROM opp_mining_batches WHERE org_id=${o}$ " "AND (project_id='' OR project_id=${p}$) " "ORDER BY created_at DESC LIMIT 10", {"o": org_id, "p": proj_id}) await sor.sqlExe("COMMIT", {}) return True, rows_to_dicts(recs, limit=10) async def _check_batch_visible(sor, ctx, batch_id): """批次可见性单一规则源(status/list_clusters 复用,语义不漂移)。 返回 (ok, err)。""" org_id = ctx.get("org_id") or "" proj_id = ctx.get("project_id") or "" recs = await sor.sqlExe( "SELECT id, project_id FROM opp_mining_batches WHERE id=${b}$ AND org_id=${o}$", {"b": batch_id, "o": org_id}) await sor.sqlExe("COMMIT", {}) if not recs: return False, "批次不存在或不属于当前机构" bproj = str(getattr(recs[0], "project_id", "") or "") if bproj and bproj != proj_id: return False, "批次属于其他项目,当前会话不可见" return True, "" async def list_clusters(sor, ctx, batch_id): """类别排名表(隔离:批次必须本机构+本项目/平台级可见)。""" ok, err = await _check_batch_visible(sor, ctx, batch_id) if not ok: return False, err org_id = ctx.get("org_id") or "" recs = await sor.sqlExe( "SELECT id, name, doc_count, share, heat_rank, naming_evidence, centroid_snap_id " "FROM opp_clusters WHERE batch_id=${b}$ AND org_id=${o}$ ORDER BY heat_rank ASC", {"b": batch_id, "o": org_id}) await sor.sqlExe("COMMIT", {}) return True, rows_to_dicts(recs, limit=200) async def cluster_detail(sor, ctx, cluster_id, limit=20): """类内需求明细(样例,带来源 URL 可查证;经批次可见性校验)。""" org_id = ctx.get("org_id") or "" recs = await sor.sqlExe( "SELECT c.id, c.name, c.doc_count, c.share, c.heat_rank, c.naming_evidence, c.batch_id " "FROM opp_clusters c WHERE c.id=${c}$ AND c.org_id=${o}$", {"c": cluster_id, "o": org_id}) await sor.sqlExe("COMMIT", {}) if not recs: return False, "类别不存在或不属于当前机构" cl = rows_to_dicts(recs, limit=1)[0] ok, err = await _check_batch_visible(sor, ctx, cl.get("batch_id") or "") if not ok: return False, err items = rows_to_dicts(await sor.sqlExe( "SELECT id, src_id, title, source, budget_wan, url, publish_time FROM opp_demand_snap " "WHERE cluster_id=${c}$ ORDER BY budget_wan DESC LIMIT ${n}$", {"c": cluster_id, "n": int(limit)}), limit=int(limit)) await sor.sqlExe("COMMIT", {}) return True, {"cluster": cl, "samples": items}