Files
2026-07-15 11:37:57 +04:00

268 lines
10 KiB
Python
Raw Permalink 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.
"""
Classify service — LLM-based document classification.
Архитектурное решение (Opus):
- Отдельный сервис, не встроен в upload. Upload быстрый (0.5с), classify — медленный (2-10с/файл).
- ThreadPoolExecutor(max_workers=4) — параллельная классификация с ограничением конкурентности,
чтобы не положить api.aillm.ru при 2000 файлах.
- Умная выжимка (_smart_extract): header ~1500 симв + regex-хиты по маркерам (договор/№/соглашение)
из всего документа. Экономия токенов в 5-10 раз при сохранении точности.
- Двухпроходная архитектура: LLM извлекает строки (тип/номер/дата/контрагент),
Python в grouping.py нормализует и группирует детерминированно.
"""
import json, re, os
from concurrent.futures import ThreadPoolExecutor, as_completed
import httpx
from db import documents as db_docs
from llm_prompt import build_classify_prompt
log = __import__("logging").getLogger(__name__)
# Лимит одновременных запросов к LLM API
# Увеличивать осторожно — api.aillm.ru может троттлить
MAX_WORKERS = 4
LLM_URL = "https://api.aillm.ru/v1/chat/completions"
LLM_KEY = os.environ.get("LLM_KEY") or os.environ.get("LLM_API_KEY", "")
LLM_MODEL = "gpt-oss-120b"
# Ленивый singleton — обратная совместимость
_classify_llm = None
def _get_classify_client():
global _classify_llm
if _classify_llm is None:
from services.llm_client import HttpxLLMClient
_classify_llm = HttpxLLMClient(url=LLM_URL, key=LLM_KEY, model=LLM_MODEL, max_tokens=1000, timeout=60)
return _classify_llm
# ── Garbage filter (Stage 1: filename regex) ────────────────────
_GARBAGE_FILENAME_RE = re.compile(
r'(сч[её]т|акт|плат[её]ж|УПД|сверк|инвойс|invoice|payment|act)',
re.IGNORECASE,
)
# ── Garbage filter (Stage 2: header keywords) ───────────────────
_GARBAGE_HEADER_MARKERS = [
'СЧЕТ-ФАКТУРА', 'СЧЕТ НА ОПЛАТУ', 'АКТ СВЕРКИ',
'АКТ ОКАЗАННЫХ УСЛУГ', 'АКТ ВЫПОЛНЕННЫХ РАБОТ',
'ПЛАТЁЖНОЕ ПОРУЧЕНИЕ', 'УНИВЕРСАЛЬНЫЙ ПЕРЕДАТОЧНЫЙ',
'УПД', 'ПЛАТЕЖНОЕ ПОРУЧЕНИЕ',
]
def _is_garbage_by_filename(filename: str) -> bool:
"""Stage 1: regex по имени файла — быстро, 0 токенов."""
return bool(_GARBAGE_FILENAME_RE.search(filename))
def _is_garbage_by_header(text: str) -> bool:
"""Stage 2: ключевые слова в первых 2KB текста — быстро, 0 токенов."""
header = text[:2000].upper()
return any(marker in header for marker in _GARBAGE_HEADER_MARKERS)
def _call_llm_classify(header_text, llm_client=None):
"""
Прямой вызов LLM для классификации ОДНОГО документа.
Возвращает (parsed_dict, raw_text, needed_fix).
llm_client: LLMClient (optional). Default — HttpxLLMClient.
"""
if llm_client is None:
llm_client = _get_classify_client()
prompt, _ = build_classify_prompt(header_text)
raw_text = llm_client.complete(prompt)
parsed, needed_fix = _safe_json_parse(raw_text)
return parsed, raw_text, needed_fix
def classify_batch(batch_id, llm_client=None, repo=None):
"""
Классифицировать все документы в batch.
Сбрасывает статус на 'pending' для всех перед началом.
Параллельно (ThreadPoolExecutor) обрабатывает до MAX_WORKERS документов.
Возвращает {ok, total, done, failed, garbage, json_fix_rate}.
llm_client: LLMClient (optional, default — HttpxLLMClient)
repo: Repository (optional, default — direct db.* calls)
"""
_db = repo if repo else db_docs
_llm = llm_client if llm_client else _get_classify_client()
# Сбросить статус — allow re-classify after file changes
_db.reset_classify_status(batch_id)
pending = _db.list_pending(batch_id)
if not pending:
return {"ok": False, "error": "no pending documents"}
total = len(pending)
done = 0
failed = 0
garbage = 0
json_fixes = 0
json_total = 0
type_counts = {} # doc_type → count for batch summary
def _classify_one(doc):
"""Классифицировать один документ: фильтр → выжимка → LLM → сохранить."""
nonlocal garbage, json_fixes, json_total, type_counts
try:
# Stage 1: garbage by filename (0 tokens)
if _is_garbage_by_filename(doc["filename"]):
_db.set_classify_garbage(doc["id"], "filename_regex")
garbage += 1
return True
# Stage 2: garbage by header keywords (0 tokens)
text = _smart_extract(doc["elements_json"])
if _is_garbage_by_header(text):
_db.set_classify_garbage(doc["id"], "header_keywords")
garbage += 1
return True
# Stage 3: LLM classification (only for remaining)
_db.set_classify_processing(doc["id"]) # crash recovery marker
result, raw, needed_fix = _call_llm_classify(text, _llm)
json_total += 1
if needed_fix:
json_fixes += 1
dtype = result.get("doc_type", "other")
type_counts[dtype] = type_counts.get(dtype, 0) + 1
_db.set_classification(
doc["id"],
result.get("doc_type", "other"),
result.get("own_number"),
result.get("parent_number"),
result.get("doc_date"),
result.get("counterparty"),
classify_raw=raw,
classify_input=text,
)
return True
except Exception as e:
_db.set_classify_failed(doc["id"], str(e))
return False
with ThreadPoolExecutor(max_workers=MAX_WORKERS) as pool:
futures = {pool.submit(_classify_one, d): d for d in pending}
for f in as_completed(futures):
if f.result():
done += 1
else:
failed += 1
return {"ok": True, "total": total, "done": done, "failed": failed,
"garbage": garbage, "json_fix_rate": round(json_fixes / max(json_total, 1), 3),
"types": type_counts, "summary": f"{total} total, {done} classified, {failed} failed, {garbage} garbage"}
def _safe_json_parse(raw):
"""Parse LLM response, fixing common JSON errors.
Returns (parsed_dict, needed_fix: bool)."""
if not raw:
raise ValueError("empty LLM response")
text = raw.strip()
# Strip markdown
if "```json" in text:
text = text.split("```json")[1].split("```")[0].strip()
elif "```" in text:
text = text.split("```")[1].split("```")[0].strip()
# Remove non-JSON prefix/suffix (LLM chatter)
brace_start = text.find("{")
brace_end = text.rfind("}")
if brace_start >= 0 and brace_end > brace_start:
text = text[brace_start:brace_end + 1]
# Try strict parse
try:
return json.loads(text), False
except json.JSONDecodeError:
pass
import re as _re
# Collapse multiline
text = _re.sub(r"\n\s*", " ", text)
# Remove trailing commas
text = _re.sub(r",\s*}", "}", text)
text = _re.sub(r",\s*]", "]", text)
# Try again
try:
return json.loads(text), True
except json.JSONDecodeError:
pass
# Aggressive: try adding missing closing quotes/braces
text = text.rstrip()
if not text.endswith("}"):
# Count unclosed quotes
in_string = False
for i, ch in enumerate(text):
if ch == '"' and (i == 0 or text[i-1] != "\\"):
in_string = not in_string
if in_string:
text += '"'
text += "}"
return json.loads(text), True
def _smart_extract(elements_json):
"""
Умная выжимка текста для классификации (решение Q3 от Opus).
Вместо отправки всего документа (дорого) или только header (теряет зарытые номера),
используется гибрид:
1. Первые ~1500 симв (титул, преамбула, стороны)
2. Regex-хиты по маркерам «договор|№|соглашение|приложение|спецификация»
из ВСЕГО документа
3. Дедупликация, лимит 10 строк, склейка → ~3000 симв на вход LLM
Это покрывает и титульную зону, и зарытые ссылки в середине документа.
"""
if not elements_json:
return ""
if isinstance(elements_json, str):
try:
elements = json.loads(elements_json)
except json.JSONDecodeError:
return elements_json[:2000]
elif isinstance(elements_json, list):
elements = elements_json
else:
return str(elements_json)[:2000]
# Build full text
lines = []
for el in elements:
if isinstance(el, dict):
t = el.get("type") or el.get("TYPE", "")
if t == "paragraph":
txt = el.get("text") or el.get("TEXT", "")
if txt:
lines.append(txt)
elif t == "table":
rows = el.get("rows") or el.get("ROWS", [])
for row in rows:
lines.append(" | ".join(str(c) for c in row))
full_text = "\n".join(lines)
# Header: first ~1500 chars
header = full_text[:1500]
# Marker lines: grep for key patterns
markers = re.findall(
r'.{0,200}(?:договор|№|соглашен|приложен|специф|контрагент|заказчик|арендатор).{0,200}',
full_text, re.IGNORECASE,
)
unique_markers = list(dict.fromkeys(markers))[:10]
combined = header + "\n---\n" + "\n".join(unique_markers)
return combined[:3000]