288 lines
13 KiB
Python
288 lines
13 KiB
Python
# -*- 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}
|