rag/wwwroot/knowledge_bases_list/upload_file.dspy
ymq 51c9968093 fix(host): 后台入库任务禁硬编码库名,改经 get_module_dbname 映射
ingest_doc 是 background_reco 后台任务(无 request/env),原用 db.sqlorContext('rag') 6 处。
在 ragserver 上库名恰为 rag 所以能跑,接入 pipeline-app(库名 pipeline)后
sqlorFactory 报 NoneType.get → 入库全程失败,status 永久 pending、向量丢失。
改为函数开头 _dbname = get_module_dbname('rag') 统一映射。
2026-08-25 15:45:44 +08:00

330 lines
17 KiB
Plaintext
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.

import base64, os, subprocess
ns = params_kw.copy()
kb_id = ns.get('kb_id', '')
folder_id = ns.get('folder', '')
file_name = ns.get('file_name', 'upload.bin')
if not kb_id:
return json.dumps({"status": "error", "error": "kb_id required"}, ensure_ascii=False)
file_data = await request.read()
if not file_data:
return json.dumps({"status": "error", "error": "no file data"}, ensure_ascii=False)
env = request._run_ns
userorgid = await env.get_userorgid()
file_size = len(file_data)
# ---- ORG STORAGE QUOTA CHECK (per-org limit, not global) ----
def fmt_bytes(n):
if n < 1024: return str(n) + 'B'
if n < 1048576: return str(round(n/1024, 1)) + 'KB'
return str(round(n/1048576, 1)) + 'MB'
quota_limit = 104857600
used = 0
async with get_sor_context(env, 'rag') as sor:
rec = await sor.sqlExe("SELECT COALESCE(SUM(file_size),0) AS used FROM rag_documents WHERE org_id=${org_id}$", {"org_id": userorgid})
if rec: used = int(rec[0].used)
lim = await sor.sqlExe("SELECT limit_bytes FROM rag_org_storage_limits WHERE org_id=${org_id}$", {"org_id": userorgid})
if lim: quota_limit = int(lim[0].limit_bytes)
if used + file_size > quota_limit:
return json.dumps({"status": "error", "error": "storage_quota_exceeded",
"message": "存储配额超限:机构已用 " + fmt_bytes(used) + ",限额 " + fmt_bytes(quota_limit) + ",本文件 " + fmt_bytes(file_size)}, ensure_ascii=False)
# Save file via FileStorage (returns web path e.g. /idfile/191/193/197/97/xxx.txt)
web_path = await env.save_file(file_data, file_name)
real_path = env.realpath(web_path)
doc_id = str(uuid()).replace('-', '')[:16]
ext = '.' + file_name.rsplit('.', 1)[1] if '.' in file_name else '.bin'
ext_l = ext.lower()
# ============================================================
# BACKGROUND INGESTION — runs in a separate asyncio task after
# the response is returned (background_reco = create_task wrapper).
# SELF-CONTAINED: never touches env/request (invalid after response),
# creates its own DBPools connection. All args are primitives.
# ============================================================
async def ingest_doc(doc_id, kb_id, file_name, ext_l, real_path):
db = DBPools()
# 库名必须经宿主的 get_module_dbname 映射ragserver→'rag'pipeline-app→'pipeline'
# 后台任务无 request/env用注入的全局函数解析禁硬编码库名
_dbname = get_module_dbname('rag')
# 读知识库向量引擎bge-m3=文本(走 /txte)clip-vith14=多媒体(走 /mme)
emb_engine = 'clip-vith14'
try:
async with db.sqlorContext(_dbname) as sor:
krecs = await sor.sqlExe("SELECT embedding_engine FROM rag_knowledge_bases WHERE id=${kb_id}$", {"kb_id": kb_id})
if krecs:
emb_engine = (getattr(krecs[0], 'embedding_engine', '') or 'clip-vith14').strip()
except: pass
is_text = emb_engine == 'bge-m3'
if is_text:
emb_url = 'https://embedding.opencomputing.net:10443/txte/api/embed'
emb_model = 'bge-m3'
else:
emb_url = 'https://embedding.opencomputing.net:10443/mme/api/embed'
emb_model = 'CLIP-ViT-H-14'
async def ensure_vdb_collection(client, kb_id):
payload = {"colname": kb_id, "fields": [{"name": "id", "type": "str", "is_primary": True, "max_length": 64}, {"name": "vector", "type": "fvector", "dim": 1024}, {"name": "text", "type": "str", "max_length": 65535}], "description": "RAG kb", "metric": "COSINE"}
await client.request('POST', 'https://vectordb.opencomputing.net:10443/v1/createcollection', json=payload)
text = ''
chunks_n = 0
face_count = 0
voice_speakers = 0
meta_parts = {}
try:
with open(real_path, 'rb') as f:
file_data = f.read()
text_exts = {'.txt', '.md', '.csv', '.json', '.xml', '.html', '.htm', '.py', '.js', '.css', '.yaml', '.yml', '.log', '.rst'}
image_exts = {'.jpg', '.jpeg', '.png', '.bmp', '.gif', '.webp'}
audio_exts = {'.mp3', '.wav', '.flac', '.ogg', '.m4a', '.aac'}
video_exts = {'.mp4', '.avi', '.mov', '.mkv', '.webm'}
# --- OFFICE DOCS: text extraction (PDF/DOCX/PPTX/XLSX) ---
if ext_l == '.pdf' and not text:
import io; from PyPDF2 import PdfReader
reader = PdfReader(io.BytesIO(file_data))
text = '\n'.join(p.extract_text() or '' for p in reader.pages)
elif ext_l == '.docx' and not text:
import io; from docx import Document
doc = Document(io.BytesIO(file_data))
text = '\n'.join(p.text for p in doc.paragraphs)
elif ext_l == '.pptx' and not text:
import io; from pptx import Presentation
prs = Presentation(io.BytesIO(file_data))
parts = []
for slide in prs.slides:
for shape in slide.shapes:
if hasattr(shape, 'text') and shape.text:
parts.append(shape.text)
text = '\n'.join(parts)
elif ext_l == '.xlsx' and not text:
import io; from openpyxl import load_workbook
wb = load_workbook(io.BytesIO(file_data), data_only=True)
parts = []
for sheet in wb.worksheets:
for row in sheet.iter_rows(values_only=True):
parts.append('\t'.join(str(c or '') for c in row))
text = '\n'.join(parts)
# --- TEXT EXTRACTION ---
if ext_l in text_exts:
text = file_data.decode('utf-8', errors='replace')
# --- 文本知识库不支持媒体文件 ---
if is_text and (ext_l in image_exts or ext_l in audio_exts or ext_l in video_exts):
try:
async with db.sqlorContext(_dbname) as sor:
await sor.sqlExe(
"UPDATE rag_documents SET status='failed', metadata=${meta}$, updated_at=NOW() WHERE id=${id}$",
{"id": doc_id, "meta": json.dumps({"error": "文本知识库不支持媒体文件,请上传文本类文件(txt/md/pdf/docx等)或改用多媒体知识库"}, ensure_ascii=False)})
except: pass
return
# --- IMAGE: face detection ---
if ext_l in image_exts:
img_b64 = base64.b64encode(file_data).decode()
try:
client = StreamHttpClient()
resp = await client.request('POST', 'https://media.opencomputing.net:10443/face/api/detect', json={"images": [img_b64]})
fd = json.loads(resp)
results = fd.get("results", [])
if results and isinstance(results[0], dict):
faces = results[0].get("faces", results[0].get("detections", []))
face_count = len(faces)
if faces and isinstance(faces[0], dict):
meta_parts['face_bboxes'] = [f.get("bbox", {}) for f in faces[:10]]
meta_parts['face'] = face_count
except: pass
# --- AUDIO: voiceprint ---
if ext_l in audio_exts:
try:
client = StreamHttpClient()
resp = await client.request('POST', 'https://media.opencomputing.net:10443/voiceprint/extract/submit',
files={'file': (file_name, file_data)})
vd = json.loads(resp)
voice_speakers = vd.get('speakers', 1) if vd.get('status') == 'SUCCEEDED' else (1 if vd.get('embedding') else 0)
meta_parts['voiceprint'] = voice_speakers
except: pass
# --- VIDEO: frame extraction + voiceprint ---
if ext_l in video_exts:
meta_parts['video'] = 'pending'
video_ok = False
try:
tmp_img = '/tmp/' + doc_id + '_frame.jpg'
subprocess.run(['ffmpeg', '-y', '-i', real_path, '-vframes', '1', '-q:v', '2', tmp_img],
capture_output=True, timeout=30)
if os.path.exists(tmp_img):
with open(tmp_img, 'rb') as fi:
frame_data = fi.read()
img_b64 = base64.b64encode(frame_data).decode()
frame_bboxes = []
try:
client = StreamHttpClient()
resp = await client.request('POST', 'https://media.opencomputing.net:10443/face/api/detect', json={"images": [img_b64]})
fd = json.loads(resp)
results = fd.get("results", [])
if results and isinstance(results[0], dict):
faces = results[0].get("faces", results[0].get("detections", []))
face_count = len(faces)
frame_bboxes = [f.get("bbox", {}) for f in faces[:10]] if faces else []
except: pass
# --- CLIP image embedding for video frame ---
try:
client2 = StreamHttpClient()
resp2 = await client2.request('POST', emb_url,
json={"images": [img_b64], "model": emb_model})
emb_data = json.loads(resp2)
img_embeddings = emb_data.get("image_embeddings", emb_data.get("embeddings", []))
except:
img_embeddings = []
if img_embeddings:
try:
client3 = StreamHttpClient()
await ensure_vdb_collection(client3, kb_id)
vdb_data = {"colname": kb_id, "data": [
{"id": doc_id + "_c0", "vector": img_embeddings[0], "text": file_name}
]}
await client3.request('POST', 'https://vectordb.opencomputing.net:10443/v1/upsert', json=vdb_data)
chunk_meta = {"start_time": 0}
if frame_bboxes:
chunk_meta["bboxes"] = frame_bboxes
async with db.sqlorContext(_dbname) as sor:
await sor.sqlExe(
"INSERT INTO rag_document_chunks (id, doc_id, kb_id, chunk_index, content, vector_id, metadata, created_at) "
"VALUES (${id}$, ${doc_id}$, ${kb_id}$, 0, ${content}$, ${vid}$, ${meta}$, NOW())",
{"id": doc_id + "_c0", "doc_id": doc_id, "kb_id": kb_id,
"content": file_name, "vid": doc_id + "_c0",
"meta": json.dumps(chunk_meta, ensure_ascii=False)})
except:
pass
os.remove(tmp_img)
video_ok = True
except: pass
# --- Voiceprint: extract audio from video ---
if video_ok:
try:
tmp_wav = '/tmp/' + doc_id + '_audio.wav'
subprocess.run(['ffmpeg', '-y', '-i', real_path, '-vn', '-acodec', 'pcm_s16le',
'-ar', '16000', '-ac', '1', tmp_wav],
capture_output=True, timeout=60)
if os.path.exists(tmp_wav) and os.path.getsize(tmp_wav) > 1000:
with open(tmp_wav, 'rb') as fa:
audio_data = fa.read()
try:
client4 = StreamHttpClient()
resp4 = await client4.request('POST',
'https://media.opencomputing.net:10443/voiceprint/extract/submit',
files={'file': (file_name.rsplit('.', 1)[0] + '.wav', audio_data)})
vd = json.loads(resp4)
voice_speakers = vd.get('speakers', 1) if vd.get('status') == 'SUCCEEDED' else (1 if vd.get('embedding') else 0)
meta_parts['voiceprint'] = voice_speakers
except: pass
if os.path.exists(tmp_wav):
os.remove(tmp_wav)
except: pass
if video_ok:
meta_parts['video'] = 'done'
# --- RAG INGEST for text ---
if text and len(text.strip()) > 10:
paragraphs = text.split('\n')
chunks = []
cur = ''
for p in paragraphs:
p = p.strip()
if not p:
if cur: chunks.append(cur); cur = ''
continue
if len(cur) + len(p) < 500:
cur = (cur + '\n' + p).strip()
else:
if cur: chunks.append(cur)
cur = p
if cur: chunks.append(cur)
if chunks:
try:
client = StreamHttpClient()
resp = await client.request('POST', emb_url,
json={"texts": chunks, "model": emb_model})
emb_data = json.loads(resp)
embeddings = emb_data.get("text_embeddings", emb_data.get("embeddings", []))
except:
embeddings = []
vector_ids = []
if embeddings:
try:
client2 = StreamHttpClient()
await ensure_vdb_collection(client2, kb_id)
vdb_data = {"colname": kb_id, "data": [
{"id": doc_id + "_c" + str(i), "vector": emb, "text": chunks[i]}
for i, emb in enumerate(embeddings)]}
resp3 = await client2.request('POST', 'https://vectordb.opencomputing.net:10443/v1/upsert', json=vdb_data)
if json.loads(resp3).get('status') == 'SUCCEEDED':
vector_ids = [doc_id + "_c" + str(i) for i in range(len(embeddings))]
except:
pass
async with db.sqlorContext(_dbname) as sor:
for i, chunk_text in enumerate(chunks):
vid = vector_ids[i] if i < len(vector_ids) else ''
await sor.sqlExe(
"INSERT INTO rag_document_chunks (id, doc_id, kb_id, chunk_index, content, vector_id, created_at) "
"VALUES (${id}$, ${doc_id}$, ${kb_id}$, ${idx}$, ${content}$, ${vid}$, NOW())",
{"id": doc_id + "_c" + str(i), "doc_id": doc_id, "kb_id": kb_id,
"idx": i, "content": chunk_text[:2000], "vid": vid})
chunks_n = len(chunks)
except Exception as e:
try:
async with db.sqlorContext(_dbname) as sor:
await sor.sqlExe(
"UPDATE rag_documents SET status='failed', metadata=${meta}$, updated_at=NOW() WHERE id=${id}$",
{"id": doc_id, "meta": json.dumps({"error": str(e)[:300]}, ensure_ascii=False)})
except: pass
return
# --- finalize: mark document done + update KB chunk counts ---
meta_json = json.dumps(meta_parts, ensure_ascii=False)
try:
async with db.sqlorContext(_dbname) as sor:
await sor.sqlExe(
"UPDATE rag_documents SET status='done', chunk_count=${chunks}$, metadata=${meta}$, updated_at=NOW() WHERE id=${id}$",
{"id": doc_id, "chunks": chunks_n, "meta": meta_json})
if chunks_n:
await sor.sqlExe(
"UPDATE rag_knowledge_bases SET chunk_count=chunk_count+${n}$ WHERE id=${kb_id}$",
{"n": chunks_n, "kb_id": kb_id})
except: pass
# ============================================================
# SYNC PART — record document as 'pending', update KB counts,
# fire background ingestion, return immediately.
# ============================================================
async with get_sor_context(env, 'rag') as sor:
await sor.sqlExe(
"INSERT INTO rag_documents (id, kb_id, folder_id, file_name, file_type, file_size, file_path, mime_type, status, chunk_count, metadata, org_id, created_at, updated_at) "
"VALUES (${id}$, ${kb_id}$, ${folder_id}$, ${file_name}$, 'other', ${file_size}$, ${file_path}$, 'application/octet-stream', 'pending', 0, '{}', ${org_id}$, NOW(), NOW())",
{"id": doc_id, "kb_id": kb_id, "folder_id": folder_id, "file_name": file_name,
"file_size": file_size, "file_path": web_path, "org_id": userorgid})
await sor.sqlExe(
"UPDATE rag_knowledge_bases SET doc_count=doc_count+1, total_size=total_size+${size}$ WHERE id=${kb_id}$",
{"size": file_size, "kb_id": kb_id})
# Fire background ingestion — pass primitives only (no env/request/proxy objects)
background_reco(ingest_doc, doc_id, kb_id, file_name, ext_l, real_path)
result = {"status": "SUCCEEDED", "doc_id": doc_id, "file_name": file_name,
"file_size": file_size, "folder_id": folder_id, "ingest": "pending"}
return json.dumps(result, ensure_ascii=False, default=str)