288 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# -*- 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}