"""Когортный анализ 1-1: удержание, сезонность входа, каналы.

Считает то, чего нет ни в AlfaCRM, ни в Финтабло: доживаемость когорт по месяцам,
зависимость удержания от месяца входа и от канала, LTV за сопоставимый срок.

    python cohorts.py             # посчитать из кэша (быстро)
    python cohorts.py --refresh   # заново выгрузить из AlfaCRM (~5 мин)

Разбор результатов и выводы — в Obsidian:
projects/mars/marketing/{mars} {research} когорты retention сезонность и каналы – 2026-08-01.md

Грабли, о которых помнит код:
  * Без истории с 2020 в «новых клиентов 2023» попадают возвращенцы из старой базы —
    когорты раздуваются, удержание завышается. Тянем всё.
  * Выручку на клиента НЕЛЬЗЯ сравнивать между когортами разного возраста: старая
    когорта всегда выглядит лучше. Сравниваем только за первые N месяцев жизни.
  * customer/index по умолчанию отдаёт только removed=0, is_study=1 — это меньше
    четверти базы. Нужны все четыре комбинации.
  * Месячные выборки — 10–40 человек. Разницу меньше 10 п.п. считать шумом.
"""
import collections
import datetime as dt
import json
import os
import statistics as st
import sys

import core

CACHE = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'exports', 'cohorts_cache.json')
BRANCHES = (2, 3)
SINCE = '2020-01-01'

# pay_item_id по группам (см. margins.py — там тот же справочник для выручки)
ABON = {2, 8, 14, 15, 11, 17, 7, 18}    # абонементы 1-1: с них начинается когорта
RAZ = {1, 9, 16, 13, 27}                # разовые 1-1
PROB = {4, 10, 12}                      # пробные
ACTIVE = ABON | RAZ                     # что считаем «клиент жив в этом месяце»

RU = ['', 'янв', 'фев', 'мар', 'апр', 'май', 'июн', 'июл', 'авг', 'сен', 'окт', 'ноя', 'дек']


# ── выгрузка ──
def fetch():
    """Платежи + карточки клиентов + справочник источников из AlfaCRM."""
    core.load_env()
    tok = core.alfa_login()

    pays, seen = [], set()
    for br in BRANCHES:
        rows = core.alfa_all(f'/v2api/{br}/pay/index',
                             {'pay_type_id': 1, 'date_from': SINCE,
                              'date_to': dt.date.today().isoformat()}, tok)
        for x in rows:
            if x['id'] in seen:
                continue
            seen.add(x['id'])
            pays.append({k: x.get(k) for k in
                         ('id', 'customer_id', 'pay_item_id', 'document_date', 'income')})
        print(f'  платежи, филиал {br}: {len(rows)}', file=sys.stderr)

    # customer/index молча фильтрует: нужны все комбинации removed × is_study,
    # иначе теряется архив (это большая часть истории).
    cust = {}
    for br in BRANCHES:
        for removed in (0, 1):
            for study in (0, 1):
                page, got = 0, 0
                while page < 60:
                    r = core.alfa(f'/v2api/{br}/customer/index',
                                  {'removed': removed, 'is_study': study,
                                   'page': page, 'pageSize': 200}, tok)
                    items = r.get('items', [])
                    for x in items:
                        cust.setdefault(str(x['id']), {'src': x.get('lead_source_id')})
                    got += len(items)
                    if len(items) < 200:
                        break
                    page += 1
                print(f'  клиенты, филиал {br} removed={removed} is_study={study}: {got}',
                      file=sys.stderr)

    src = {str(x['id']): x['name']
           for br in BRANCHES
           for x in core.alfa(f'/v2api/{br}/lead-source/index', {'page': 0}, tok).get('items', [])}

    os.makedirs(os.path.dirname(CACHE), exist_ok=True)
    json.dump({'pays': pays, 'cust': cust, 'src': src},
              open(CACHE, 'w'), ensure_ascii=False)
    return pays, cust, src


def load(refresh=False):
    if refresh or not os.path.exists(CACHE):
        return fetch()
    d = json.load(open(CACHE))
    return d['pays'], d['cust'], d['src']


# ── подготовка ──
def _date(s):
    dd, mm, yy = s.split('.')
    return dt.date(int(yy), int(mm), int(dd))


def _mkey(d):
    """Месяц как одно число — чтобы арифметика по месяцам была без возни с годами."""
    return d.year * 12 + d.month


def build(pays):
    """{customer_id: [(дата, pay_item, сумма)]} и когорты по первому абонементу.

    Последний месяц отрезаем целиком: он неполный и занижает свежие когорты.
    """
    today = dt.date.today()
    cut = today.replace(day=1) - dt.timedelta(days=1)

    by = collections.defaultdict(list)
    for p in pays:
        if not p['customer_id'] or not p['document_date']:
            continue
        d = _date(p['document_date'])
        if d > cut:
            continue
        by[p['customer_id']].append((d, p['pay_item_id'] or 0, float(p['income'] or 0)))
    for v in by.values():
        v.sort(key=lambda r: (r[0], r[1]))

    coh = {}
    for cid, rows in by.items():
        first_any = next((r for r in rows if r[1] in ACTIVE | PROB), None)
        first_ab = next((r for r in rows if r[1] in ABON), None)
        if not first_ab:
            continue
        # Возвращенец, а не новый клиент: первый контакт задолго до этого абонемента.
        if first_any and (first_ab[0] - first_any[0]).days > 60:
            continue
        coh[cid] = {'m': _mkey(first_ab[0]), 'date': first_ab[0], 'check': first_ab[2]}
    return by, coh, _mkey(cut)


# ── метрики ──
def retention(by, sel, cut_m, k):
    """Доля когорты, заплатившей в месяце `старт + k`. Незрелые когорты не считаем."""
    n = hit = 0
    for cid, v in sel.items():
        if v['m'] + k > cut_m:
            continue
        n += 1
        hit += any(_mkey(r[0]) == v['m'] + k and r[1] in ACTIVE for r in by[cid])
    return (hit / n * 100 if n else None), n


def revenue_first(by, sel, cut_m, months=6):
    """Выручка за первые `months` месяцев жизни — единственная сравнимая между когортами."""
    vals = [sum(r[2] for r in by[cid] if r[1] in ACTIVE and _mkey(r[0]) < v['m'] + months)
            for cid, v in sel.items() if v['m'] + months - 1 <= cut_m]
    return (st.mean(vals) if vals else 0), len(vals)


# ── отчёты ──
def report_curve(by, coh, cut_m):
    print('\n=== КРИВАЯ УДЕРЖАНИЯ (все когорты) ===')
    ks, vs, ns = [], [], []
    for k in range(13):
        r, n = retention(by, coh, cut_m, k)
        if r is None or n < 25:
            break
        ks.append(k + 1); vs.append(r); ns.append(n)
    print('  месяц ' + ' '.join(f'{k:>5}' for k in ks))
    print('      % ' + ' '.join(f'{v:>5.0f}' for v in vs))
    print('   база ' + ' '.join(f'{n:>5}' for n in ns))


def report_years(by, coh, cut_m):
    print('\n=== ПО ГОДАМ КОГОРТ ===')
    print(f'{"год":6}{"когорт":>8}{"2-й мес":>9}{"6-й мес":>9}{"выручка 6 мес":>16}{"ср. чек":>11}')
    years = sorted({v['date'].year for v in coh.values()})
    for y in years:
        sel = {c: v for c, v in coh.items() if v['date'].year == y}
        if len(sel) < 20:
            continue
        r2, _ = retention(by, sel, cut_m, 1)
        r6, _ = retention(by, sel, cut_m, 5)
        m6, _ = revenue_first(by, sel, cut_m)
        chk = [v['check'] for v in sel.values() if v['check'] > 0]
        f = lambda x: f'{x:.0f}%' if x is not None else '—'
        print(f'{y:<6}{len(sel):>8}{f(r2):>9}{f(r6):>9}{m6:>15,.0f}₽{st.mean(chk):>10,.0f}₽')


def report_season(by, coh, cut_m):
    print('\n=== СЕЗОННОСТЬ ВХОДА (удержание 2-го месяца по месяцу старта) ===')
    rows = []
    for mo in range(1, 13):
        sel = {c: v for c, v in coh.items() if v['date'].month == mo}
        r, n = retention(by, sel, cut_m, 1)
        if r is not None and n >= 30:
            rows.append((mo, r, n))
    for mo, r, n in sorted(rows, key=lambda x: -x[1]):
        print(f'  старт {RU[mo]:4} → 2-й мес {RU[mo % 12 + 1]:4}: {r:5.1f}%   n={n:>4}')
    if rows:
        best, worst = max(rows, key=lambda x: x[1]), min(rows, key=lambda x: x[1])
        print(f'  разрыв {RU[best[0]]} / {RU[worst[0]]}: {best[1] / worst[1]:.1f}×')


def report_channels(by, coh, cut_m, cust, src, since_year=2024):
    print(f'\n=== КАНАЛЫ (когорты {since_year}+, выручка нормирована на 6 месяцев) ===')
    recent = {c: v for c, v in coh.items() if v['date'].year >= since_year}
    grp = collections.defaultdict(dict)
    for cid, v in recent.items():
        info = cust.get(str(cid)) or {}
        grp[src.get(str(info.get('src'))) or '⚠ источник не указан'][cid] = v
    rows = []
    for name, sel in grp.items():
        m6, n6 = revenue_first(by, sel, cut_m)
        if n6 < 15:
            continue
        r2, _ = retention(by, sel, cut_m, 1)
        rows.append((name, len(sel), n6, r2, m6))
    rows.sort(key=lambda r: -r[4])
    print(f'{"канал":42}{"когорт":>8}{"n(6мес)":>9}{"2-й мес":>9}{"выручка 6 мес":>16}')
    for name, n, n6, r2, m6 in rows:
        print(f'{name[:41]:42}{n:>8}{n6:>9}{r2:>8.0f}%{m6:>15,.0f}₽')
    m6all, _ = revenue_first(by, recent, cut_m)
    print(f'{"— среднее по всем —":42}{len(recent):>8}{"":>9}{"":>9}{m6all:>15,.0f}₽')
    unknown = len(grp.get('⚠ источник не указан', {}))
    print(f'  без источника: {unknown} из {len(recent)} ({unknown / len(recent) * 100:.0f}%)')


def report_check(by, coh, cut_m, since_year=2025):
    print(f'\n=== ПЕРВЫЙ ЧЕК (когорты {since_year}+) ===')
    sel_all = {c: v for c, v in coh.items() if v['date'].year >= since_year and v['check'] > 0}
    vals = sorted(v['check'] for v in sel_all.values())
    q1, q3 = vals[len(vals) // 4], vals[3 * len(vals) // 4]
    bands = ((0, q1, f'ниже {q1:,.0f}'), (q1, q3, f'{q1:,.0f}–{q3:,.0f}'), (q3, 9e9, f'выше {q3:,.0f}'))
    for lo, hi, name in bands:
        sel = {c: v for c, v in sel_all.items() if lo <= v['check'] < hi}
        if len(sel) < 20:
            continue
        r2, _ = retention(by, sel, cut_m, 1)
        r6, _ = retention(by, sel, cut_m, 5)
        f = lambda x: f'{x:.0f}%' if x is not None else '—'
        print(f'  чек {name:>18}: 2-й мес {f(r2):>5}   6-й мес {f(r6):>5}   n={len(sel)}')


def main():
    pays, cust, src = load('--refresh' in sys.argv)
    by, coh, cut_m = build(pays)
    print(f'платежей: {len(pays)} | клиентов: {len(by)} | когорт: {len(coh)}')
    report_curve(by, coh, cut_m)
    report_years(by, coh, cut_m)
    report_season(by, coh, cut_m)
    report_channels(by, coh, cut_m, cust, src)
    report_check(by, coh, cut_m)


if __name__ == '__main__':
    main()
