254 lines
12 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, "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 ""
# 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,
"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=""):
"""查批次状态(org 隔离)。batch_id 空 = 本机构最近批次列表。"""
org_id = ctx.get("org_id") or ""
if batch_id:
recs = await sor.sqlExe(
"SELECT id, scope, status, 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, "批次不存在或不属于当前机构"
return True, rows_to_dicts(recs, limit=1)[0]
recs = await sor.sqlExe(
"SELECT id, scope, status, stats_json, error_msg, created_at "
"FROM opp_mining_batches WHERE org_id=${o}$ ORDER BY created_at DESC LIMIT 10",
{"o": org_id})
await sor.sqlExe("COMMIT", {})
return True, rows_to_dicts(recs, limit=10)
async def list_clusters(sor, ctx, batch_id):
"""类别排名表(org 隔离:批次必须属于本机构)。"""
org_id = ctx.get("org_id") or ""
recs = await sor.sqlExe(
"SELECT 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, "批次不存在或不属于当前机构"
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 id, name, doc_count, share, heat_rank, naming_evidence "
"FROM opp_clusters WHERE id=${c}$ AND 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]
items = rows_to_dicts(await sor.sqlExe(
"SELECT 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}