#!/usr/bin/env python3
"""Keyword clustering for Arabic and English keyword lists (standard library only).

Usage:
    python keyword_cluster.py keywords.csv                    # writes keywords_clustered.csv and clusters.csv
    python keyword_cluster.py keywords.csv --threshold 0.4    # looser clusters (default 0.5)
    python keyword_cluster.py keywords.csv --out results      # output folder
    python keyword_cluster.py keywords.csv --xlsx             # also writes clusters.xlsx (needs: pip install openpyxl)

Input: a CSV or TSV with a header. Recognised columns (any order):
    keyword (or keywords, query, term), volume (or search volume), kd (or keyword difficulty), cpc.
A plain text file with one keyword per line also works.

How it works (the same algorithm runs in the browser tool at /en/tools/keyword-clustering/):
    1. Normalize: lowercase, strip Arabic diacritics and tatweel, unify alef, ya and ta marbuta forms, drop punctuation.
    2. Tokenize, drop filler words (the, how, for, في, من, كيف ...), strip the Arabic "ال" prefix and simple English plurals.
    3. Order keywords by search volume (highest first). The top keyword that is not yet in a cluster becomes the seed.
    4. Every remaining keyword joins the seed's cluster if its token overlap (Jaccard: shared words divided by all words
       in both) is at least the threshold.
    5. Score each cluster: total volume x (1 - average KD / 100); with no volume column at all, the keyword count is used. Top 20% of clusters = High, next 30% = Medium, rest = Low.
    6. Label intent (transactional, commercial, informational, unclear) from modifier words, and suggest a page type.

Limits: it matches words, not meaning. "refund" and "money back" stay in different clusters, and Arabic and English
versions of the same idea are not merged. Always review the clusters by hand before building pages.
"""
import argparse
import csv
import math
import os
import re
import sys
import unicodedata

DIACRITICS = re.compile(r"[ً-ٰٟـ]")
ARABIC = re.compile(r"[؀-ۿ]")


def normalize(s):
    s = DIACRITICS.sub("", s.lower())
    s = re.sub("[أإآٱ]", "ا", s).replace("ى", "ي").replace("ة", "ه")
    s = "".join(ch if unicodedata.category(ch)[0] in "LN" else " " for ch in s)
    return re.sub(" +", " ", s).strip()


def stem(t):
    if ARABIC.search(t):
        return t[2:] if t.startswith("ال") and len(t) >= 5 else t
    if len(t) > 3 and t.endswith("ies"):
        return t[:-3] + "y"
    if len(t) > 3 and t.endswith("s") and not t.endswith("ss"):
        return t[:-1]
    return t


def prep(words):
    return {stem(normalize(w)) for w in words.split()}


STOP = prep("a an the of for to in on at and or with is are do does how what why when can i my your me we you it by from "
            "في من على عن الى او ما ماهو ماهي هو هي كيف كيفية لماذا متى هل مع")
TRANSACTIONAL = prep("buy price prices cheap cheapest discount coupon deal deals "
                     "شراء سعر اسعار اشتري رخيص ارخص خصم كوبون تخفيض")
COMMERCIAL = prep("best top review reviews compare comparison vs alternative alternatives software tool service services "
                  "company companies provider platform system pricing cost افضل مقارنة تقييم مراجعة بديل شركات شركة خدمات برنامج منصة نظام تكلفة")
INFORMATIONAL = prep("how what why when guide tutorial tips meaning definition example "
                     "كيف كيفية ما ماهو ماهي لماذا متى شرح دليل طريقة معنى تعريف نصائح")
PAGE = {"transactional": "Product or category page", "commercial": "Landing or comparison page",
        "informational": "Blog article", "unclear": "Review by hand"}
ORDER = ["transactional", "commercial", "informational", "unclear"]


def round_half_up(x, digits):
    f = 10 ** digits
    return math.floor(x * f + 0.5) / f


def intent_of(tokens):
    if any(t in TRANSACTIONAL for t in tokens):
        return "transactional"
    if any(t in COMMERCIAL for t in tokens):
        return "commercial"
    if any(t in INFORMATIONAL for t in tokens):
        return "informational"
    return "unclear"


def prepare(row):
    norm = normalize(row["keyword"])
    allt = [stem(t) for t in norm.split() if t]
    core = [t for t in allt if t not in STOP]
    if not core:
        core = allt if allt else [norm or "#"]
    return {**row, "norm": norm, "lang": "ar" if ARABIC.search(row["keyword"]) else "en",
            "intent": intent_of(allt), "tokens": sorted(set(core))}


def similar(a, b, threshold):
    sb = set(b)
    inter = sum(1 for t in a if t in sb)
    if not inter:
        return False
    union = len(a) + len(b) - inter
    return inter / union >= threshold


def mean(xs):
    return sum(xs) / len(xs) if xs else None


def cluster(rows, threshold=0.5):
    seen, kws = set(), []
    for r in rows:
        k = prepare(r)
        if not k["norm"] or k["norm"] in seen:
            continue
        seen.add(k["norm"])
        kws.append(k)
    kws.sort(key=lambda k: (-(k.get("volume") or 0), len(k["tokens"]), k["norm"]))
    any_volume = any(k.get("volume") is not None for k in kws)  # if any keyword has volume, clusters without volume score 0
    taken = [False] * len(kws)
    out = []
    for i, seed in enumerate(kws):
        if taken[i]:
            continue
        taken[i] = True
        members = [seed]
        for j in range(i + 1, len(kws)):
            if not taken[j] and similar(seed["tokens"], kws[j]["tokens"], threshold):
                taken[j] = True
                members.append(kws[j])
        has_volume = any(m.get("volume") is not None for m in members)
        volume = sum(m.get("volume") or 0 for m in members)
        kd = mean([m["kd"] for m in members if m.get("kd") is not None])
        kd = None if kd is None else round_half_up(kd, 1)
        cpc = mean([m["cpc"] for m in members if m.get("cpc") is not None])
        cpc = None if cpc is None else round_half_up(cpc, 2)
        weight = {}
        for m in members:
            weight[m["intent"]] = weight.get(m["intent"], 0) + (m.get("volume") or 0) + 1e-6
        intent = "unclear"
        for x in ORDER:
            if weight.get(x, 0) > weight.get(intent, 0):
                intent = x
        langs = {m["lang"] for m in members}
        score = volume * (1 - min(kd or 0, 100) / 100) if any_volume else len(members)
        out.append({"name": seed["keyword"], "keywords": members, "count": len(members), "volume": volume, "kd": kd,
                    "cpc": cpc, "score": round_half_up(score, 2), "intent": intent,
                    "lang": next(iter(langs)) if len(langs) == 1 else "mixed", "page": PAGE[intent], "has_volume": has_volume})
    out.sort(key=lambda c: (-c["score"], c["name"]))
    n = len(out)
    hi, mid = math.ceil(n * 0.2), math.ceil(n * 0.5)
    for i, c in enumerate(out):
        c["id"] = i + 1
        c["priority"] = "High" if i < hi else "Medium" if i < mid else "Low"
    return out


ALIAS = {
    "keyword": ["keyword", "keywords", "query", "term", "phrase", "الكلمة", "الكلمه"],
    "volume": ["volume", "search volume", "sv", "vol", "monthly volume", "volume sa"],
    "kd": ["kd", "keyword difficulty", "difficulty", "kd %", "kd%"],
    "cpc": ["cpc", "cpc usd", "cost per click"],
}


def to_num(s):
    t = re.sub(r"[,%$\s]", "", s or "")
    try:
        return float(t) if t != "" else None
    except ValueError:
        return None


def read_rows(path):
    with open(path, encoding="utf-8-sig", newline="") as f:
        lines = [l.strip() for l in f.read().replace("\r", "").split("\n") if l.strip()]
    if not lines:
        return []
    delim = next((d for d in ("\t", ";", ",") if d in lines[0]), None)
    if delim is None:
        return [{"keyword": l} for l in lines]
    table = list(csv.reader(lines, delimiter=delim))
    head = [normalize(h) for h in table[0]]

    def col(key):
        names = [normalize(a) for a in ALIAS[key]]
        return next((i for i, h in enumerate(head) if h in names), -1)

    ik, iv, idf, ic = col("keyword"), col("volume"), col("kd"), col("cpc")
    if ik < 0:
        return [{"keyword": r[0].strip()} for r in table if r and r[0].strip()]
    rows = []
    for r in table[1:]:
        kw = r[ik].strip() if ik < len(r) else ""
        if not kw:
            continue
        get = lambda i: to_num(r[i]) if 0 <= i < len(r) else None
        row = {"keyword": kw}
        for key, i in (("volume", iv), ("kd", idf), ("cpc", ic)):
            v = get(i)
            if v is not None:
                row[key] = v
        rows.append(row)
    return rows


def fmt(v):
    return "" if v is None else (int(v) if float(v).is_integer() else v)


def write_outputs(clusters, out_dir, xlsx):
    os.makedirs(out_dir, exist_ok=True)
    kw_path = os.path.join(out_dir, "keywords_clustered.csv")
    cl_path = os.path.join(out_dir, "clusters.csv")
    with open(kw_path, "w", encoding="utf-8-sig", newline="") as f:
        w = csv.writer(f)
        w.writerow(["cluster_id", "cluster", "keyword", "language", "intent", "volume", "kd", "cpc", "priority", "suggested_page"])
        for c in clusters:
            for k in c["keywords"]:
                w.writerow([c["id"], c["name"], k["keyword"], k["lang"], k["intent"], fmt(k.get("volume")), fmt(k.get("kd")),
                            fmt(k.get("cpc")), c["priority"], c["page"]])
    with open(cl_path, "w", encoding="utf-8-sig", newline="") as f:
        w = csv.writer(f)
        w.writerow(["cluster_id", "cluster", "keywords", "total_volume", "avg_kd", "avg_cpc", "score", "priority", "intent", "language", "suggested_page"])
        for c in clusters:
            w.writerow([c["id"], c["name"], c["count"], fmt(c["volume"]) if c["has_volume"] else "", fmt(c["kd"]), fmt(c["cpc"]),
                        fmt(c["score"]), c["priority"], c["intent"], c["lang"], c["page"]])
    print(f"wrote {kw_path}\nwrote {cl_path}")
    if xlsx:
        try:
            from openpyxl import Workbook
        except ImportError:
            print("openpyxl is not installed (pip install openpyxl), skipping the xlsx file", file=sys.stderr)
            return
        wb = Workbook()
        for title, path in (("Clusters", cl_path), ("Keywords", kw_path)):
            ws = wb.active if title == "Clusters" else wb.create_sheet(title)
            ws.title = title
            with open(path, encoding="utf-8-sig", newline="") as f:
                for row in csv.reader(f):
                    ws.append([float(x) if re.fullmatch(r"-?\d+(\.\d+)?", x) else x for x in row])
            ws.freeze_panes = "A2"
            ws.auto_filter.ref = ws.dimensions
        xp = os.path.join(out_dir, "clusters.xlsx")
        wb.save(xp)
        print(f"wrote {xp}")


def main():
    ap = argparse.ArgumentParser(description="Cluster Arabic and English keyword lists.")
    ap.add_argument("input")
    ap.add_argument("--threshold", type=float, default=0.5, help="token overlap needed to join a cluster, 0 to 1 (default 0.5)")
    ap.add_argument("--out", default=".", help="output folder (default: current folder)")
    ap.add_argument("--xlsx", action="store_true", help="also write clusters.xlsx (needs openpyxl)")
    a = ap.parse_args()
    rows = read_rows(a.input)
    if not rows:
        sys.exit("no keywords found in the input file")
    clusters = cluster(rows, a.threshold)
    write_outputs(clusters, a.out, a.xlsx)
    total = sum(c["count"] for c in clusters)
    print(f"{total} keywords in {len(clusters)} clusters; top 5 by score:")
    for c in clusters[:5]:
        print(f"  {c['id']:>3}. {c['name']}  ({c['count']} kw, {c['priority']}, {c['intent']})")


if __name__ == "__main__":
    main()
