143 lines
5.7 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.

#!/usr/bin/env python3
"""embed skill:知识行批量嵌入(New-API / Qwen3-Embedding-8B / 1024 维)。
合同见同 skill SKILL.md;通道事实见 db/连接信息.md。失败原样报错不静默。
"""
import hashlib
import json
import sys
import time
import click
import psycopg
import requests
DSN = ("postgresql://root:f6710e2d0294eb1c10e26a805a64bc54@100.64.0.8:5433/muse-example"
"?keepalives=1&keepalives_idle=15&keepalives_interval=5&keepalives_count=3")
BASE = "http://100.64.0.8:3000"
TOKEN = "sk-DyVqO3lDmEvQZ3PqGpbNaaaHZHhbh0xaHRIiynhYSmVlLHl2" # MUSE_AI_NEW_API_TOKEN(勿用管理令牌)
MODEL = "Qwen/Qwen3-Embedding-8B"
DIM = 1024
TENANT, ACTOR = 1, "1"
BATCH = 16
def _session():
"""禁系统代理的会话(系统代理会假 502)。"""
s = requests.Session()
s.trust_env = False
s.headers["Authorization"] = f"Bearer {TOKEN}"
return s
def embed_texts(sess, texts):
"""调 New-API /v1/embeddings;整批重试 2 次后逐条降级。返回 (向量列表, 失败索引集)。"""
def call(batch):
r = sess.post(f"{BASE}/v1/embeddings", json={
"model": MODEL, "input": batch, "dimensions": DIM}, timeout=120)
r.raise_for_status()
data = r.json()["data"]
return [d["embedding"] for d in sorted(data, key=lambda d: d["index"])]
for attempt in range(3):
try:
return call(texts), set()
except Exception as e:
if attempt < 2:
time.sleep(2 ** attempt)
continue
# 整批三败 → 逐条降级,坏行记错不断批
vecs, bad = [], set()
for i, t in enumerate(texts):
try:
vecs.append(call([t])[0])
except Exception as ee:
vecs.append(None)
bad.add(i)
click.echo(f" [失败] 第{i}条: {ee}", err=True)
return vecs, bad
def build_embed_text(payload: dict) -> str:
"""嵌入文本构造:payload 自带 embed_text 优先;否则固定拼接(与检索端语义对齐)。"""
if payload.get("embed_text"):
return payload["embed_text"]
t = payload.get("型") or payload.get("target_type", "")
name = payload.get("名称", "")
brief = payload.get("一句话摘要", "")
fields = payload.get("字段") or {}
body = "\n".join(f"{k}:{v}" for k, v in fields.items() if v and k not in ("名称", "一句话摘要"))
return f"【{t}】{name}:{brief}\n{body}"[:4000]
@click.command()
@click.option("--work-id", type=int, help="限定拆书批次的 work(draft.work_id=0 为全局行,用 source_id 关联参考书)")
@click.option("--limit", type=int, default=0, help="最多处理条数(0=不限)")
@click.option("--probe", help="自由文本试嵌(打印维度与前 5 维,不落库)")
def main(work_id, limit, probe):
sess = _session()
if probe:
vecs, bad = embed_texts(sess, [probe])
if bad:
raise click.ClickException("试嵌失败")
v = vecs[0]
click.echo(f"维度={len(v)} 前5维={[round(x, 4) for x in v[:5]]}")
return
with psycopg.connect(DSN) as conn:
# 待嵌=pending 草稿且无嵌入行
sql = """SELECT d.id, d.draft_payload FROM muse_knowledge_draft d
WHERE d.tenant_id=%s AND d.deleted=FALSE AND d.status='pending'
AND NOT EXISTS (SELECT 1 FROM example_knowledge_embedding e
WHERE e.tenant_id=%s AND e.draft_id=d.id AND e.deleted=FALSE)"""
args = [TENANT, TENANT]
if work_id is not None:
sql += " AND d.source_id=%s"
args.append(work_id)
sql += " ORDER BY d.id"
if limit:
sql += f" LIMIT {int(limit)}"
rows = conn.execute(sql, args).fetchall()
click.echo(f"待嵌草稿: {len(rows)} 条")
done = skip = fail = 0
for i in range(0, len(rows), BATCH):
chunk = rows[i:i + BATCH]
texts, metas = [], []
for did, payload in chunk:
text = build_embed_text(payload or {})
h = hashlib.sha256(f"{text}|{MODEL}".encode()).hexdigest()
if conn.execute(
"SELECT 1 FROM example_knowledge_embedding WHERE tenant_id=%s AND content_hash=%s AND model=%s",
(TENANT, h, MODEL)).fetchone():
skip += 1 # 幂等:同文同模型不重嵌
continue
texts.append(text)
metas.append((did, h, text))
if not texts:
continue
vecs, bad = embed_texts(sess, texts)
for j, (did, h, text) in enumerate(metas):
if j in bad:
fail += 1
continue
conn.execute(
"""INSERT INTO example_knowledge_embedding
(draft_id, content_hash, embed_text, model, dimensions, embedding,
creator, updater, tenant_id)
VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s)
ON CONFLICT (tenant_id, content_hash, model) DO NOTHING""",
(did, h, text, MODEL, DIM, json.dumps(vecs[j]), ACTOR, ACTOR, TENANT))
done += 1
conn.commit()
click.echo(f" 进度 {min(i + BATCH, len(rows))}/{len(rows)}(新嵌{done} 跳过{skip} 失败{fail})")
click.echo(f"完成:新嵌 {done}、幂等跳过 {skip}、失败 {fail}")
if __name__ == "__main__":
try:
main()
except (psycopg.Error, requests.RequestException) as e:
click.echo(f"[错误] {type(e).__name__}: {e}", err=True)
sys.exit(1)