#!/usr/bin/env python3
"""
Индексатор документов — строит SQLite базу с эмбедингами через aitunnel.ru.

Что хранит:
  - Все метаданные элемента (element_id, contractor, doc_types, ...)
  - file_ids для загрузки по требованию агентом
  - Текст анализа (PROPERTY_2989) если есть
  - Эмбединг для семантического поиска

Использование:
    python3 indexer.py                        # индексирует matched_records.json
    python3 indexer.py --all                  # забирает ВСЕ записи из API (~12к)
    python3 indexer.py --search "авансирование риск"
    python3 indexer.py --id 66991             # показать запись по ID
    python3 indexer.py --stats                # статистика базы
    python3 indexer.py --model gemini-embedding-001  # выбрать модель
"""

import json
import os
import re
import sqlite3
import argparse
import time
import requests
import numpy as np
from datetime import datetime

# ─── Конфиг ───────────────────────────────────────────────────────────────────

DB_PATH        = 'data/documents.db'
WEBHOOK        = 'https://engeocentr.bitrix24.ru/rest/4391/n8rq6nyzr9z8tvdz/'
AITUNNEL_URL   = 'https://api.aitunnel.ru/v1/embeddings'
AITUNNEL_KEY   = 'sk-aitunnel-0H1gwP4EewVgEnwkMKFE1bjm8ZtzkZ6j'

# Доступные модели на aitunnel.ru
MODELS = {
    'gemini2': 'gemini-embedding-2-preview',  # 3072 dim, лучшее качество для RU
    'gemini':  'gemini-embedding-001',         # 3072 dim, большой контекст 20k
    'small':   'text-embedding-3-small',       # 1536 dim, цена/качество
    'large':   'text-embedding-3-large',       # 3072 dim, OpenAI топ
    'gigachat':'gigachat-embeddings-giga-r',   # нативный русский, Sber
    'qwen':    'qwen3-embedding-8b',           # 32k контекст
}
DEFAULT_MODEL = 'gemini-embedding-2-preview'

BATCH_SIZE = 32   # запросов в батче к aitunnel
RETRY_MAX  = 3    # повторов при ошибке

# ─── Маппинги B24 ─────────────────────────────────────────────────────────────

PAYER_MAP  = {'373': 'БДК', '375': 'ТЗК', '377': 'ИНГЕО', '1059': 'ВЦПО'}
OBJECT_MAP = {
    '1071':  'Камергер (общий)',
    '1073':  'Дмитровка (общий)',
    '15881': 'БД-ПА',
    '15883': 'КА-НС',
    '15885': 'КА-ПА',
    '17363': 'Стр.7 Б. Дмитровка, дом.9',
    '39129': 'БД-НС',
}
STATUS_MAP = {'363': 'Отклонено', '365': 'Согласовано',
              '367': 'Согласовано и распечатано', '419': 'Не установлен'}
TYPE_MAP   = {'357': 'Договор/ДС', '361': 'Счет/Акт'}

DOC_PATTERNS = [
    ('КС-2/КС-3',    r'КС-2\s*/\s*КС-3'),
    ('УПД',          r'\bУПД\b'),
    ('Счет-фактура', r'[Сс]чет-фактура|[Сс]ч-ф\b'),
    ('Акт',          r'\bАкт\b'),
    ('Счет',         r'[Сс]чет\s*№'),
    ('КС-2',         r'КС-2\b'),
    ('КС-3',         r'КС-3\b'),
    ('Справка',      r'[Сс]правка'),
    ('Договор',      r'Дог(?:овор)?\.?\s*№'),
    ('ДС',           r'\bДС\s*№'),
]


# ─── Утилиты ──────────────────────────────────────────────────────────────────

def scalar(field):
    if not field: return ''
    if isinstance(field, dict):
        vals = list(field.values())
        return vals[0] if vals else ''
    return field

def multi_values(field):
    if not field: return []
    if isinstance(field, dict):
        vals = list(field.values())
        if vals and isinstance(vals[0], list): return vals[0]
        return [str(v) for v in vals if v]
    if isinstance(field, list): return field
    return [str(field)]

def parse_doc_types(text):
    found = []
    for name, pat in DOC_PATTERNS:
        if re.search(pat, text or ''): found.append(name)
    return list(dict.fromkeys(found))

def extract_contract_refs(text):
    refs = []
    for m in re.findall(r'к\s*Договору\s*№\s*([\w\d\.\-\/]+)', text or ''):
        refs.append('Договор №' + m)
    for m in re.findall(r'по\s*ДС\s*№\s*([\w\d\.]+)', text or ''):
        refs.append('ДС №' + m)
    return refs

def strip_html(text):
    if not text: return ''
    return re.sub(r'<[^>]+>', ' ', text).strip()

def build_embed_text(record: dict) -> str:
    """Текст для эмбединга — объединяем всё значимое поле."""
    parts = [
        record.get('name', ''),
        record.get('contractor', ''),
        record.get('payer', ''),
        ' '.join(record.get('doc_types', [])),
        record.get('num_date_str', ''),
        record.get('num_date_in', ''),
        record.get('amount', ''),
        record.get('status', ''),
        record.get('object', ''),
        record.get('expense_item', ''),
        record.get('comment', ''),
        ' '.join(record.get('contract_refs', [])),
        strip_html(record.get('analysis_text', ''))[:2000],
    ]
    return ' '.join(p for p in parts if p).strip()

def get_analysis_text(item):
    res = scalar(item.get('PROPERTY_2989'))
    if isinstance(res, dict): return res.get('TEXT', '')
    return res or ''

def build_record_from_api(item: dict) -> dict:
    nom_dat    = scalar(item.get('PROPERTY_391', ''))
    title      = item.get('NAME', '')
    file_ids   = []
    for prop in ['PROPERTY_415', 'PROPERTY_451', 'PROPERTY_503', 'PROPERTY_2185']:
        file_ids += multi_values(item.get(prop))
    file_ids   = [f for f in file_ids if f]
    payer_id   = str(scalar(item.get('PROPERTY_427', '')))
    status_id  = str(scalar(item.get('PROPERTY_401', '')))
    type_id    = str(scalar(item.get('PROPERTY_389', '')))
    doc_types  = parse_doc_types(nom_dat) or parse_doc_types(title)
    return {
        'element_id':    item['ID'],
        'name':          title,
        'date_create':   item.get('DATE_CREATE', '')[:10],
        'contractor':    scalar(item.get('PROPERTY_1367', '')),
        'company_id':    scalar(item.get('PROPERTY_423', '')),
        'payer':         PAYER_MAP.get(payer_id, payer_id),
        'doc_type_raw':  TYPE_MAP.get(type_id, type_id),
        'doc_types':     doc_types,
        'contract_refs': extract_contract_refs(nom_dat),
        'num_date_str':  nom_dat,
        'num_date_in':   scalar(item.get('PROPERTY_383', '')),
        'amount':        scalar(item.get('PROPERTY_421', '')),
        'status':        STATUS_MAP.get(status_id, status_id),
        'object':        OBJECT_MAP.get(str(scalar(item.get('PROPERTY_387', ''))), scalar(item.get('PROPERTY_387', ''))),
        'expense_item':  scalar(item.get('PROPERTY_397', '')),
        'comment':       scalar(item.get('PROPERTY_413', '')),
        'payment_due':   scalar(item.get('PROPERTY_1415', '')),
        'approved_date': scalar(item.get('PROPERTY_431', '')),
        'file_ids':      file_ids,
        'has_analysis':  bool(get_analysis_text(item)),
        'analysis_text': get_analysis_text(item),
    }


# ─── Aitunnel Embeddings ───────────────────────────────────────────────────────

def embed_batch(texts: list, model: str) -> list:
    """
    Отправить батч текстов в aitunnel, вернуть список numpy-векторов.
    Повторяет при ошибке до RETRY_MAX раз.
    """
    for attempt in range(RETRY_MAX):
        try:
            resp = requests.post(
                AITUNNEL_URL,
                headers={
                    'Authorization': f'Bearer {AITUNNEL_KEY}',
                    'Content-Type':  'application/json',
                },
                json={'model': model, 'input': texts},
                timeout=60,
            )
            resp.raise_for_status()
            data = resp.json()

            # Сортируем по индексу (API гарантирует порядок, но на всякий случай)
            items = sorted(data['data'], key=lambda x: x['index'])
            vectors = [np.array(item['embedding'], dtype='float32') for item in items]

            # Логируем стоимость если есть
            usage = data.get('usage', {})
            cost  = data.get('cost_rub', '')
            if cost:
                print(f'    tokens: {usage.get("total_tokens", "?")} | cost: {cost} руб.')

            return vectors

        except Exception as e:
            if attempt < RETRY_MAX - 1:
                wait = 2 ** attempt
                print(f'    ⚠️  Попытка {attempt+1} ошибка: {e} — повтор через {wait}с')
                time.sleep(wait)
            else:
                print(f'    ❌ Ошибка после {RETRY_MAX} попыток: {e}')
                return [None] * len(texts)


def cosine_search(query_vec: np.ndarray, matrix: np.ndarray) -> np.ndarray:
    """Косинусное сходство query против матрицы всех эмбедингов."""
    q = query_vec / (np.linalg.norm(query_vec) + 1e-9)
    m = matrix / (np.linalg.norm(matrix, axis=1, keepdims=True) + 1e-9)
    return m @ q


# ─── База данных ──────────────────────────────────────────────────────────────

def init_db(conn):
    conn.executescript("""
        CREATE TABLE IF NOT EXISTS documents (
            element_id      TEXT PRIMARY KEY,
            name            TEXT,
            date_create     TEXT,
            contractor      TEXT,
            company_id    TEXT,
            payer           TEXT,
            doc_type_raw    TEXT,
            doc_types       TEXT,
            contract_refs   TEXT,
            num_date_str    TEXT,
            num_date_in     TEXT,
            amount          TEXT,
            status          TEXT,
            object          TEXT,
            expense_item    TEXT,
            comment         TEXT,
            payment_due     TEXT,
            approved_date   TEXT,
            file_ids        TEXT,
            has_analysis    INTEGER DEFAULT 0,
            analysis_text   TEXT,
            embed_text      TEXT,
            embed_model     TEXT,
            embedding       BLOB,
            indexed_at      TEXT
        );
        CREATE INDEX IF NOT EXISTS idx_contractor ON documents(contractor);
        CREATE INDEX IF NOT EXISTS idx_date       ON documents(date_create);
        CREATE INDEX IF NOT EXISTS idx_status     ON documents(status);
        CREATE INDEX IF NOT EXISTS idx_payer      ON documents(payer);
    """)
    conn.commit()


def upsert(conn, record: dict, embedding: np.ndarray = None, model: str = None):
    embed_text = build_embed_text(record)
    emb_blob   = embedding.tobytes() if embedding is not None else None
    conn.execute("""
        INSERT INTO documents
            (element_id, name, date_create, contractor, company_id, payer, doc_type_raw,
             doc_types, contract_refs, num_date_str, num_date_in, amount, status,
             object, expense_item, comment, payment_due, approved_date,
             file_ids, has_analysis, analysis_text, embed_text,
             embed_model, embedding, indexed_at)
        VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
        ON CONFLICT(element_id) DO UPDATE SET
            name          = excluded.name,
            date_create   = excluded.date_create,
            contractor    = excluded.contractor,
            company_id  = excluded.company_id,
            payer         = excluded.payer,
            doc_type_raw  = excluded.doc_type_raw,
            doc_types     = excluded.doc_types,
            contract_refs = excluded.contract_refs,
            num_date_str  = excluded.num_date_str,
            num_date_in   = excluded.num_date_in,
            amount        = excluded.amount,
            status        = excluded.status,
            object        = excluded.object,
            expense_item  = excluded.expense_item,
            comment       = excluded.comment,
            payment_due   = excluded.payment_due,
            approved_date = excluded.approved_date,
            file_ids      = excluded.file_ids,
            has_analysis  = excluded.has_analysis,
            analysis_text = excluded.analysis_text,
            embed_text    = excluded.embed_text,
            embed_model   = CASE WHEN excluded.embedding IS NOT NULL
                            THEN excluded.embed_model ELSE embed_model END,
            embedding     = CASE WHEN excluded.embedding IS NOT NULL
                            THEN excluded.embedding ELSE embedding END,
            indexed_at    = excluded.indexed_at
    """, (
        record['element_id'], record.get('name', ''), record.get('date_create', ''),
        record.get('contractor', ''), record.get('company_id', ''), record.get('payer', ''),
        record.get('doc_type_raw', ''),
        json.dumps(record.get('doc_types', []),     ensure_ascii=False),
        json.dumps(record.get('contract_refs', []), ensure_ascii=False),
        record.get('num_date_str', ''), record.get('num_date_in', ''),
        record.get('amount', ''), record.get('status', ''),
        record.get('object', ''), record.get('expense_item', ''),
        record.get('comment', ''), record.get('payment_due', ''), record.get('approved_date', ''),
        json.dumps(record.get('file_ids', []), ensure_ascii=False),
        int(record.get('has_analysis', False)),
        record.get('analysis_text', ''),
        embed_text, model, emb_blob,
        datetime.now().isoformat(),
    ))


def get_by_id(conn, element_id: str) -> dict | None:
    row = conn.execute('SELECT * FROM documents WHERE element_id = ?', (str(element_id),)).fetchone()
    if not row: return None
    cols = [d[0] for d in conn.execute('SELECT * FROM documents LIMIT 0').description]
    rec  = dict(zip(cols, row))
    for f in ('file_ids', 'doc_types', 'contract_refs'):
        rec[f] = json.loads(rec[f] or '[]')
    return rec


# ─── Загрузка из API ──────────────────────────────────────────────────────────

def fetch_all_api() -> list:
    select = [
        'ID', 'NAME', 'DATE_CREATE', 'CREATED_BY',
        'PROPERTY_1367', 'PROPERTY_423',
        'PROPERTY_389', 'PROPERTY_391', 'PROPERTY_383',
        'PROPERTY_421', 'PROPERTY_401', 'PROPERTY_427',
        'PROPERTY_387', 'PROPERTY_397', 'PROPERTY_413',
        'PROPERTY_415', 'PROPERTY_451', 'PROPERTY_503', 'PROPERTY_2185',
        'PROPERTY_1415', 'PROPERTY_431',
        'PROPERTY_2989',
    ]
    all_items, start, page = [], 0, 1
    while True:
        resp  = requests.post(WEBHOOK + 'lists.element.get', json={
            'IBLOCK_TYPE_ID': 'bitrix_processes', 'IBLOCK_ID': '57',
            'FILTER': {'ACTIVE': 'Y'}, 'ELEMENT_ORDER': {'DATE_CREATE': 'DESC'},
            'SELECT': select, 'LIST_NUM': start,
        }, timeout=30)
        resp.raise_for_status()
        data  = resp.json()
        items = data.get('result', [])
        if not items: break
        all_items += items
        total = int(data.get('total', 0))
        if page % 10 == 0 or len(all_items) >= total:
            print(f'  Страница {page}: {len(all_items)}/{total}')
        if len(all_items) >= total: break
        start += 50
        page  += 1
    return all_items


# ─── Main ─────────────────────────────────────────────────────────────────────

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--all',    action='store_true', help='Загрузить все записи из API')
    ap.add_argument('--search', type=str,            help='Семантический поиск')
    ap.add_argument('--id',     type=str,            help='Показать запись по element_id')
    ap.add_argument('--stats',  action='store_true', help='Статистика базы')
    ap.add_argument('--model',  type=str,            default=DEFAULT_MODEL,
                    help=f'Модель эмбедингов (default: {DEFAULT_MODEL})\n' +
                         '\n'.join(f'  {k}: {v}' for k, v in MODELS.items()))
    args = ap.parse_args()

    # Разрешаем короткие алиасы (small, large, gemini...)
    model = MODELS.get(args.model, args.model)

    os.makedirs('data', exist_ok=True)
    conn = sqlite3.connect(DB_PATH)
    init_db(conn)

    # ── Показать запись ──────────────────────────────────────────────────────
    if args.id:
        rec = get_by_id(conn, args.id)
        if not rec:
            print(f'Запись {args.id} не найдена')
        else:
            print(f'element_id:    {rec["element_id"]}')
            print(f'name:          {rec["name"]}')
            print(f'date:          {rec["date_create"]}')
            print(f'contractor:    {rec["contractor"]}')
            print(f'payer:         {rec["payer"]}')
            print(f'doc_types:     {rec["doc_types"]}')
            print(f'contract_refs: {rec["contract_refs"]}')
            print(f'amount:        {rec["amount"]}')
            print(f'status:        {rec["status"]}')
            print(f'file_ids:      {rec["file_ids"]}')
            print(f'has_analysis:  {bool(rec["has_analysis"])}')
            print(f'embed_model:   {rec["embed_model"]}')
        conn.close()
        return

    # ── Статистика ───────────────────────────────────────────────────────────
    if args.stats:
        total = conn.execute('SELECT COUNT(*) FROM documents').fetchone()[0]
        emb   = conn.execute('SELECT COUNT(*) FROM documents WHERE embedding IS NOT NULL').fetchone()[0]
        anal  = conn.execute('SELECT COUNT(*) FROM documents WHERE has_analysis=1').fetchone()[0]
        models_used = conn.execute(
            'SELECT embed_model, COUNT(*) FROM documents WHERE embed_model IS NOT NULL GROUP BY embed_model'
        ).fetchall()
        print(f'Записей в базе:  {total}')
        print(f'С эмбедингами:   {emb}')
        print(f'С анализом:      {anal}')
        if models_used:
            print('Модели:')
            for m, c in models_used:
                print(f'  {c}x {m}')
        conn.close()
        return

    # ── Семантический поиск ──────────────────────────────────────────────────
    if args.search:
        rows = conn.execute(
            'SELECT element_id, name, contractor, doc_types, status, amount, '
            'file_ids, date_create, embedding FROM documents WHERE embedding IS NOT NULL'
        ).fetchall()

        if not rows:
            print('❌ Нет эмбедингов — сначала запусти индексацию')
            conn.close()
            return

        # Определяем размерность из первой записи
        first_vec = np.frombuffer(rows[0][8], dtype='float32')
        dim = len(first_vec)
        matrix = np.frombuffer(
            b''.join(r[8] for r in rows), dtype='float32'
        ).reshape(len(rows), dim)

        print(f'Запрашиваю эмбединг для поиска через {model}...')
        vecs = embed_batch([args.search], model)
        if not vecs or vecs[0] is None:
            print('❌ Не удалось получить эмбединг запроса')
            conn.close()
            return

        scores  = cosine_search(vecs[0], matrix)
        top_idx = scores.argsort()[::-1][:10]

        print(f'\nРезультаты для: "{args.search}"\n')
        for i in top_idx:
            r = rows[i]
            print(f'  [{scores[i]:.3f}] ID={r[0]} | {r[2]}')
            print(f'           {json.loads(r[3] or "[]")} | {r[7]} | {r[5] or "—"}')
            print(f'           file_ids: {json.loads(r[6] or "[]")}')
            print()
        conn.close()
        return

    # ── Индексация ───────────────────────────────────────────────────────────
    if args.all:
        # Приоритет: raw_all.json (уже скачан) → all_records.json → API
        raw_path = 'data/raw_all.json'
        all_path = 'data/all_records.json'

        if os.path.exists(raw_path):
            print(f'Читаю сырые данные из {raw_path}...')
            with open(raw_path, encoding='utf-8') as f:
                raw_items = json.load(f)
            records = [build_record_from_api(i) for i in raw_items]
            print(f'Загружено: {len(records)} записей\n')
        elif os.path.exists(all_path):
            print(f'Читаю из {all_path}...')
            with open(all_path, encoding='utf-8') as f:
                records = json.load(f)
            print(f'Загружено: {len(records)} записей\n')
        else:
            print('❌ Нет данных — сначала запусти: python3 fetch_api.py --all')
            conn.close()
            return

        # Пропускаем уже проиндексированные с эмбедингом
        existing = set(
            r[0] for r in conn.execute(
                'SELECT element_id FROM documents WHERE embedding IS NOT NULL'
            ).fetchall()
        )
        if existing:
            before = len(records)
            records = [r for r in records if r['element_id'] not in existing]
            print(f'Уже в базе: {len(existing)}, осталось проиндексировать: {len(records)} из {before}\n')

    else:
        matched_path = 'data/matched_records.json'
        if not os.path.exists(matched_path):
            print(f'❌ Не найден {matched_path} — запусти match.py')
            conn.close()
            return
        with open(matched_path, encoding='utf-8') as f:
            data = json.load(f)
        records = data.get('matched', [])
        print(f'Из matched_records: {len(records)} записей\n')

    print(f'Модель эмбедингов: {model}')
    print(f'Батч: {BATCH_SIZE} текстов за запрос\n')

    total_ok  = 0
    total_err = 0

    for i in range(0, len(records), BATCH_SIZE):
        batch  = records[i:i + BATCH_SIZE]
        texts  = [build_embed_text(r) for r in batch]
        n      = i + len(batch)

        print(f'[{n}/{len(records)}] Эмбедингую батч...')
        vectors = embed_batch(texts, model)

        for record, vec in zip(batch, vectors):
            upsert(conn, record, vec, model if vec is not None else None)
            if vec is not None: total_ok  += 1
            else:               total_err += 1

        conn.commit()

    total_db = conn.execute('SELECT COUNT(*) FROM documents').fetchone()[0]
    print(f'\n✅ Готово')
    print(f'   База:            {DB_PATH}')
    print(f'   Всего записей:   {total_db}')
    print(f'   Эмбединги OK:    {total_ok}')
    if total_err:
        print(f'   Ошибки:          {total_err}')

    conn.close()


if __name__ == '__main__':
    main()
