Files
drhider/site/routes/api_bp.py
T

471 lines
22 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.
"""
Blueprint: API DrHider.
Five endpoints:
- POST /api/upload — upload one file -> {session_id}
- GET /api/process_stream/<sid> — SSE: process files, per-file progress
- POST /api/process/<sid> — process all session files -> {status:"done"} (legacy)
- GET /api/download/<sid> — download ZIP (with timestamp name)
- GET /api/csv/<sid> — download CSV separately
"""
import io
import json
import queue
import threading
import time
import zipfile
import traceback
import logging
import httpx
from datetime import datetime, timedelta
from flask import Blueprint, request, send_file, jsonify, Response, stream_with_context
from drhider import obfuscate_files, LLMClient
from session import (create_session, add_file, get_files, store_result,
get_result, store_csv, get_csv, cleanup, file_count,
MAX_FILE_BYTES, pause_ttl, resume_ttl,
request_cancel, get_cancel_event)
api_bp = Blueprint("api", __name__, url_prefix="/api")
log = logging.getLogger("routes.api_bp")
# Ретраи pull из ВМ-буфера: защита от разовых DNS/сетевых сбоев (gaierror -5 и т.п.)
PULL_RETRIES = 3
PULL_RETRY_DELAY = 2 # секунды между попытками
# Доверенный префикс ВМ-буфера — валидация URL при pull (защита от SSRF)
VM_UPLOAD_PREFIX = "https://contracts.kube5s.ru/drhider-upload/"
def _safe_name(name: str) -> str:
"""Санитизировать имя файла: защита от path traversal, сохраняя подпапки.
Запрещает '..' и абсолютные пути; нормализует слэши. Возвращает "" если
имя пустое или небезопасное.
"""
if not name:
return ""
name = name.replace("\\", "/")
parts = [p for p in name.split("/") if p and p != "."]
if not parts or any(p == ".." for p in parts):
return ""
return "/".join(parts)
def _disconnect_exceptions():
"""Исключения, означающие отключение клиента SSE."""
return (GeneratorExit, BrokenPipeError, ConnectionResetError)
@api_bp.route("/upload", methods=["POST"])
def upload():
"""Upload files to session (один или несколько)."""
sid = request.form.get("session", "")
if not sid:
sid = create_session()
uploaded = request.files.getlist("files")
if not uploaded:
log.warning("upload: no files, sid=%s", sid)
return jsonify({"ok": False, "error": "No file"}), 400
added = 0
had_unnamed = False
for f in uploaded:
name = _safe_name(f.filename)
if not name:
had_unnamed = True
continue
data = f.read()
log.info("upload: sid=%s file=%r size=%d", sid, name, len(data))
if len(data) > MAX_FILE_BYTES:
log.warning("upload: file exceeds %dMB, skipped sid=%s file=%r size=%d",
MAX_FILE_BYTES // (1024 * 1024), sid, name, len(data))
continue
if not add_file(sid, name, data):
log.warning("upload: session not found/limit, sid=%s file=%r", sid, f.filename)
return jsonify({"ok": False, "error": "Session not found"}), 404
added += 1
if added == 0:
err = "No filename" if had_unnamed else "No file"
log.warning("upload: %s, sid=%s", err, sid)
return jsonify({"ok": False, "error": err}), 400
log.info("upload: done sid=%s added=%d total=%d", sid, added, file_count(sid))
return jsonify({"ok": True, "session": sid, "count": file_count(sid)})
@api_bp.route("/upload_refs", methods=["POST"])
def upload_refs():
"""Принять ссылки на файлы (загружены на ВМ-буфер), забрать по egress.
Вход: JSON {"session": "...", "files": [{"name": str, "size": int, "url": str}]}.
Каждый файл тянется ИСХОДЯЩИМ GET'ом с ВМ (egress не ограничен шлюзом),
читается по частям (stream), кладётся в сессию. После успешного pull файл
удаляется с ВМ (best-effort; TTL-чистка на ВМ тоже есть).
"""
data = request.get_json(silent=True) or {}
sid = data.get("session") or create_session()
refs = data.get("files") or []
if not refs:
log.warning("upload_refs: no files, sid=%s", sid)
return jsonify({"ok": False, "error": "No files"}), 400
added = 0
try:
with httpx.Client(timeout=120, follow_redirects=True) as client:
for ref in refs:
name = _safe_name(ref.get("name") or "")
url = ref.get("url")
if not name or not url:
continue
# SSRF-защита: тянуть можно ТОЛЬКО с доверенного ВМ-буфера
if not url.startswith(VM_UPLOAD_PREFIX):
log.warning("upload_refs: unsafe URL, skip sid=%s url=%r", sid, url)
continue
# Лимит на один файл (50 МБ): сверх лимита — пропускаем (не участвует)
if (ref.get("size") or 0) > MAX_FILE_BYTES:
log.warning("upload_refs: file exceeds %dMB, skip sid=%s file=%r size=%s",
MAX_FILE_BYTES // (1024 * 1024), sid, name, ref.get("size"))
try:
client.delete(url)
except Exception:
pass
continue
# Pull с ретраями: разовые DNS/сетевые сбои не роняют всю загрузку
content = None
last_err = None
for attempt in range(PULL_RETRIES):
try:
with client.stream("GET", url) as resp:
resp.raise_for_status()
content = b"".join(resp.iter_bytes())
last_err = None
break
except Exception as e:
last_err = e
log.warning("upload_refs: pull attempt %d/%d failed sid=%s file=%r: %r",
attempt + 1, PULL_RETRIES, sid, name, e)
time.sleep(PULL_RETRY_DELAY)
if content is None:
raise last_err if last_err else RuntimeError("pull failed")
log.info("upload_refs: pulled sid=%s file=%r size=%d", sid, name, len(content))
if len(content) > MAX_FILE_BYTES:
log.warning("upload_refs: pulled file exceeds %dMB, skip sid=%s file=%r size=%d",
MAX_FILE_BYTES // (1024 * 1024), sid, name, len(content))
try:
client.delete(url)
except Exception:
pass
continue
if not add_file(sid, name, content):
# Различить: сессия исчезла vs превышен суммарный лимит сессии
if get_files(sid) is None:
log.warning("upload_refs: session not found, sid=%s file=%r", sid, name)
return jsonify({"ok": False, "error": "Session not found"}), 404
log.warning("upload_refs: session limit exceeded, skip sid=%s file=%r", sid, name)
try:
client.delete(url)
except Exception:
pass
continue
try:
client.delete(url) # убрать файл с ВМ после загрузки
except Exception:
pass
added += 1
except Exception as e:
log.error("upload_refs: pull error sid=%s: %r", sid, e)
return jsonify({"ok": False, "error": "Pull failed: %s" % e}), 502
log.info("upload_refs: done sid=%s added=%d total=%d", sid, added, file_count(sid))
return jsonify({"ok": True, "session": sid, "count": file_count(sid)})
@api_bp.route("/session_files/<sid>", methods=["GET"])
def session_files(sid):
"""Отладка: список файлов сессии с размерами (для теста загрузки)."""
files = get_files(sid)
if files is None:
return jsonify({"ok": False, "error": "Session not found"}), 404
return jsonify({
"ok": True, "count": len(files),
"files": [{"name": n, "size": len(b)} for n, b in files],
})
@api_bp.route("/cancel/<sid>", methods=["POST"])
def cancel(sid):
"""Запросить мягкое прерывание обработки сессии.
Воркер останавливается на ближайшей границе файла (текущий добирается),
собирает частичный результат (готовые файлы + mapping) и шлёт SSE-событие `cancelled`.
"""
if request_cancel(sid):
log.info("cancel: requested sid=%s", sid)
return jsonify({"ok": True}), 200
log.warning("cancel: session not found sid=%s", sid)
return jsonify({"ok": False, "error": "Session not found"}), 404
@api_bp.route("/process_stream/<sid>", methods=["GET"])
def process_stream(sid):
"""SSE: process all session files, streaming per-file progress.
Все файлы обрабатываются ЕДИНЫМ вызовом obfuscate_files (общий mapping,
согласованные токены). Обработка идёт в отдельном потоке; прогресс
передаётся через очередь. Разрыв соединения клиента корректно
перехватывается и останавливает генератор (и воркер — через cancel_event).
"""
files = get_files(sid)
if files is None:
return jsonify({"ok": False, "error": "Session not found"}), 404
if not files:
return jsonify({"ok": False, "error": "No files"}), 400
all_files = [(fname, content, "") for fname, content in files]
log.info("process_stream: start sid=%s files=%d", sid, len(all_files))
# Сессия живёт, пока идёт обработка (TTL возобновляется в finally генератора)
pause_ttl(sid)
def generate():
llm = LLMClient()
q = queue.Queue()
cancel = threading.Event() # локальный: разрыв клиента (стоп heartbeat)
cancel_event = get_cancel_event(sid) # из сессии: мягкая отмена (кнопка «Прервать»)
# Состояние для глобальной ETA (символы)
eta = {"total_chars": 0, "done_chars": 0, "cur_chars": 0, "cur_total": 0, "cur_done": 0}
_TOKENS_PER_CHAR = 8.0 # эмпирический коэффициент символов -> токенов LLM
def progress(phase, idx, total_, name, elapsed):
q.put(("progress", phase, idx, name, total_, elapsed))
# Состояние для per-file ETA (чанки)
fstate = {"t0": 0.0, "total": 0, "done": 0}
def file_progress(event, fname, **fields):
if event == "file_start":
fstate["t0"] = time.time()
fstate["total"] = fields.get("chunks", 0)
fstate["done"] = 0
elif event == "file_chunk":
fstate["done"] = fields.get("chunks_done", 0)
elapsed = time.time() - fstate["t0"]
rate = fstate["done"] / elapsed if elapsed > 0 else 0
rem = fstate["total"] - fstate["done"]
if rate > 0:
fields["eta_sec"] = max(0, round(rem / rate))
q.put(("file", event, fname, fields))
def worker():
log.info("worker: start sid=%s files=%d", sid, len(all_files))
t0 = datetime.utcnow()
try:
zip_data, csv_str, meta = obfuscate_files(
all_files, llm_client=llm, progress_cb=progress,
cancel_event=cancel_event, file_progress=file_progress,
)
stats = {
"tokens": llm.tokens_total,
"llm_sec": round(llm.llm_sec, 1),
}
dt = (datetime.utcnow() - t0).total_seconds()
if meta and meta.get("cancelled"):
log.info("worker: cancelled sid=%s in %.1fs processed=%d/%d",
sid, dt, meta.get("processed", 0), meta.get("total", 0))
q.put(("cancelled", zip_data, csv_str, stats, meta))
else:
log.info("worker: done sid=%s in %.1fs tokens=%d llm_sec=%.1f zip_len=%d",
sid, dt, llm.tokens_total, llm.llm_sec, len(zip_data))
q.put(("result", zip_data, csv_str, stats))
except Exception as e:
log.error("worker: exception sid=%s: %r\n%s", sid, e, traceback.format_exc())
q.put(("error", repr(e)))
threading.Thread(target=worker, daemon=True).start()
log.debug("process_stream: worker thread started sid=%s", sid)
try:
while True:
try:
evt = q.get(timeout=1)
except queue.Empty:
if cancel.is_set():
log.info("process_stream: cancelled sid=%s (generator exit)", sid)
return
# Heartbeat: живая статистика LLM + глобальная ETA
tokens = llm.tokens_total
elapsed = llm.llm_elapsed_now()
eta_sec = None
if eta["total_chars"] > 0 and tokens > 0 and elapsed > 0:
est_total_tokens = eta["total_chars"] * _TOKENS_PER_CHAR
rate = tokens / elapsed
if rate > 0:
eta_sec = max(0, int((est_total_tokens - tokens) / rate))
done_chars = eta["done_chars"]
if eta["cur_total"] > 0:
done_chars += int(eta["cur_chars"] * (eta["cur_done"] / eta["cur_total"]))
try:
yield (
f"event: llm\n"
f"data: {json.dumps({'active': llm.llm_active, 'elapsed': round(elapsed, 1), 'tokens': tokens, 'eta_sec': eta_sec, 'done_chars': done_chars, 'total_chars': eta['total_chars']})}\n\n"
)
except _disconnect_exceptions() as e:
log.warning("process_stream: disconnect during heartbeat sid=%s err=%r", sid, e)
cancel.set()
if cancel_event:
cancel_event.set()
return
continue
kind = evt[0]
if kind == "progress":
_, phase, idx, name, total_, elapsed = evt
log.debug("process_stream: event=%s idx=%d name=%r elapsed=%s sid=%s",
phase, idx, name, elapsed, sid)
try:
yield (
f"event: {phase}\n"
f"data: {json.dumps({'idx': idx, 'name': name, 'total': total_, 'elapsed': elapsed})}\n\n"
)
except _disconnect_exceptions() as e:
log.warning("process_stream: disconnect on progress sid=%s phase=%s err=%r", sid, phase, e)
cancel.set()
if cancel_event:
cancel_event.set()
return
elif kind == "file":
_, event, fname, fields = evt
# Обновляем состояние для глобальной ETA
if event == "extract_done":
eta["total_chars"] = fields.get("total_chars", 0)
elif event == "file_start":
eta["cur_chars"] = fields.get("chars", 0)
eta["cur_total"] = fields.get("chunks", 0)
eta["cur_done"] = 0
elif event == "file_chunk":
eta["cur_done"] = fields.get("chunks_done", 0)
elif event == "file_done":
eta["done_chars"] += eta["cur_chars"]
eta["cur_chars"] = 0
eta["cur_total"] = 0
eta["cur_done"] = 0
try:
yield (
f"event: {event}\n"
f"data: {json.dumps({'name': fname, **fields})}\n\n"
)
except _disconnect_exceptions() as e:
log.warning("process_stream: disconnect on file event sid=%s err=%r", sid, e)
cancel.set()
if cancel_event:
cancel_event.set()
return
elif kind == "result":
_, zip_data, csv_str, stats = evt
log.info("process_stream: result sid=%s, storing result", sid)
store_result(sid, zip_data)
if csv_str:
store_csv(sid, csv_str)
count = 0
with zipfile.ZipFile(io.BytesIO(zip_data)) as zf:
count = len([n for n in zf.namelist() if n != "mapping.csv"])
log.info("process_stream: complete sid=%s count=%d stats=%r", sid, count, stats)
try:
yield (
f"event: complete\n"
f"data: {json.dumps({'total': count, **stats})}\n\n"
)
except _disconnect_exceptions() as e:
log.warning("process_stream: disconnect on complete sid=%s err=%r", sid, e)
return
return
elif kind == "cancelled":
_, zip_data, csv_str, stats, meta = evt
log.info("process_stream: cancelled sid=%s, storing partial result", sid)
store_result(sid, zip_data)
if csv_str:
store_csv(sid, csv_str)
try:
yield (
f"event: cancelled\n"
f"data: {json.dumps({'saved': meta.get('processed', 0), 'total': meta.get('total', 0), **stats})}\n\n"
)
except _disconnect_exceptions() as e:
log.warning("process_stream: disconnect on cancelled sid=%s err=%r", sid, e)
return
return
elif kind == "error":
_, msg = evt
log.error("process_stream: error event sid=%s msg=%r", sid, msg)
try:
yield f"event: proc_error\ndata: {json.dumps({'error': msg})}\n\n"
except _disconnect_exceptions() as e:
log.warning("process_stream: disconnect on error sid=%s err=%r", sid, e)
return
return
finally:
resume_ttl(sid)
log.debug("process_stream: generator exit sid=%s, TTL resumed", sid)
return Response(
stream_with_context(generate()),
content_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}
)
@api_bp.route("/process/<sid>", methods=["POST"])
def process(sid):
"""Process all session files -> ZIP (единым вызовом, общий mapping)."""
files = get_files(sid)
if files is None:
return jsonify({"ok": False, "error": "Session not found"}), 404
if not files:
return jsonify({"ok": False, "error": "No files"}), 400
try:
llm = LLMClient()
all_files = [(fname, content, "") for fname, content in files]
pause_ttl(sid)
try:
zip_data, csv_str, _ = obfuscate_files(all_files, llm_client=llm)
finally:
resume_ttl(sid)
store_result(sid, zip_data)
if csv_str:
store_csv(sid, csv_str)
return jsonify({"ok": True, "status": "done"})
except Exception as e:
traceback.print_exc()
return jsonify({"ok": False, "error": str(e)}), 500
@api_bp.route("/download/<sid>", methods=["GET"])
def download(sid):
"""Download result and cleanup session."""
zip_data = get_result(sid)
if zip_data is None:
return jsonify({"ok": False, "error": "Not found"}), 404
ts = (datetime.now() + timedelta(hours=3)).strftime("%Y-%m-%d_%H-%M-%S")
return send_file(io.BytesIO(zip_data), mimetype="application/zip",
as_attachment=True, download_name=f"drhider_{ts}.zip")
@api_bp.route("/csv/<sid>", methods=["GET"])
def csv_download(sid):
"""Download CSV separately."""
csv_str = get_csv(sid)
if csv_str is None:
return jsonify({"ok": False, "error": "Not found"}), 404
ts = (datetime.now() + timedelta(hours=3)).strftime("%Y-%m-%d_%H-%M-%S")
buf = io.BytesIO()
buf.write('\ufeff'.encode('utf-8') + csv_str.encode('utf-8'))
buf.seek(0)
return send_file(buf, mimetype="text/csv",
as_attachment=True, download_name=f"mapping_{ts}.csv")