rag/wwwroot/knowledge_bases_list/upload_file.dspy
ymq 1054dbb3a9 feat(embedding): 知识库支持文本(bge-m3 /txte)/多媒体(CLIP /mme)双向量引擎
- create_kb/new_kb_form 加向量引擎选择(文本 bge-m3 / 多媒体 clip-vith14)
- upload_file/batch_ingest/search_result/init.py 按 embedding_engine 路由到 /txte 或 /mme
- 文本知识库上传媒体文件友好拒绝(不再崩溃报错)
2026-08-25 13:15:11 +08:00

327 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 documents WHERE org_id=${org_id}$", {"org_id": userorgid})
if rec: used = int(rec[0].used)
lim = await sor.sqlExe("SELECT limit_bytes FROM 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()
# 读知识库向量引擎bge-m3=文本(走 /txte)clip-vith14=多媒体(走 /mme)
emb_engine = 'clip-vith14'
try:
async with db.sqlorContext('rag') as sor:
krecs = await sor.sqlExe("SELECT embedding_engine FROM 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('rag') as sor:
await sor.sqlExe(
"UPDATE 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('rag') as sor:
await sor.sqlExe(
"INSERT INTO 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('rag') as sor:
for i, chunk_text in enumerate(chunks):
vid = vector_ids[i] if i < len(vector_ids) else ''
await sor.sqlExe(
"INSERT INTO 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('rag') as sor:
await sor.sqlExe(
"UPDATE 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('rag') as sor:
await sor.sqlExe(
"UPDATE 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 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 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 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)