"""Parameter sweep for guide candidate clustering. Read-only."""
import numpy as np
from supporthub.app.main import app
from supporthub.app.db import session_scope
from supporthub.app.models import AddonsProducts, Tickets, ThreadEmbeddings, GeneratedGuide
from supporthub.app.services.embedding_service import embedding_index

SCHEMA = "tenant_internal"

with app.app_context():
    from flask import g
    g.tenant_schema = SCHEMA
    embedding_index.load_from_db(schema=SCHEMA)
    embedding_index.load_doc_pages_from_db(schema=SCHEMA)

    with session_scope(schema=SCHEMA) as session:
        prods = []
        for p in session.query(AddonsProducts).all():
            try:
                prov = int(p.provider_product_id)
            except (ValueError, TypeError):
                continue
            n = (session.query(ThreadEmbeddings)
                 .filter(ThreadEmbeddings.product_id == prov,
                         ThreadEmbeddings.embedding_quality == 'strong',
                         ThreadEmbeddings.summary_text.isnot(None)).count())
            prods.append((p.id, prov, p.name, n))
        prods.sort(key=lambda r: r[3], reverse=True)
        prods = prods[:6]

        strong_rows = (session.query(ThreadEmbeddings.thread_id, ThreadEmbeddings.product_id)
                       .filter(ThreadEmbeddings.embedding_quality == 'strong',
                               ThreadEmbeddings.summary_text.isnot(None)).all())
        done = set(r[0] for r in session.query(Tickets.thread_id)
                   .join(GeneratedGuide, GeneratedGuide.ticket_id == Tickets.id).all())

    snap = embedding_index.snapshot()
    t2v = snap['thread_to_vec']
    per = {}
    for tid, pid in strong_rows:
        per.setdefault(pid, []).append(tid)

    def cluster(tids, thr):
        seeds, sizes = [], []
        for tid in sorted(tids, reverse=True):
            v = t2v.get(tid)
            if v is None: continue
            vn = v/(float(np.linalg.norm(v))+1e-10)
            if seeds:
                sims = np.array([float(np.dot(vn,s)) for s in seeds])
                bi = int(sims.argmax())
                if sims[bi] >= thr:
                    sizes[bi]+=1; continue
            seeds.append(vn); sizes.append(1)
        return sizes

    print(f"{'product':40} {'strong':>6}", end="")
    combos = [(0.78,2),(0.82,3),(0.85,4),(0.86,5),(0.88,5),(0.88,6),(0.90,6)]
    for thr,mcs in combos:
        print(f" t{thr}/m{mcs:>1}", end="")
    print()
    for (pid, prov, name, n) in prods:
        tids = [t for t in per.get(prov,[]) if t not in done]
        print(f"{name[:40]:40} {n:>6}", end="")
        for thr,mcs in combos:
            sizes = cluster(tids, thr)
            cnt = sum(1 for s in sizes if s>=mcs)
            print(f" {cnt:>7}", end="")
        print()
