import argparse
import json
import numpy as np
import pandas as pd
import plotly.express as px
from sklearn.feature_extraction.text import CountVectorizer
from sklearn.decomposition import LatentDirichletAllocation
from gp_common import Run
from gp_text_analysis import prepare

MODEL_ID = "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"

def topics(path="outputs/google_play_reviews.csv", root="outputs/gp_runs", use_bertopic=False, embeddings=None):
    df, settings = prepare(path)
    reasons = []
    for _, r in df.iterrows():
        reasons.append("empty_text" if not r.content_raw.strip() else "invalid_score" if pd.isna(r.score_valid) else "empty_tokens" if not r.doc else "included")
    df["selection"] = reasons
    selected = df[df["selection"].eq("included")].copy()
    docs = selected["doc"].tolist()
    settings.update(lda={"min_reviews": 5, "max_topics": 10, "random_state": 42, "max_iter": 20, "min_df": 1, "max_df": 1.0},
                    bertopic={"enabled": use_bertopic, "min_reviews": 10, "embedding_model": MODEL_ID if embeddings is None else "provided_embeddings_not_model_inference",
                              "min_df": 1, "max_df": 1.0, "ngram_range": [1, 2], "umap_neighbors": min(15, max(1, len(docs)-1)),
                              "umap_components": 5, "umap_init": "random", "umap_metric": "cosine", "random_state": 42,
                              "hdbscan_min_cluster_size": 5, "calculate_probabilities": False})
    run = Run("topics", path, settings, root)
    audit = df.copy()
    for col in ["tokens", "tokens_raw"]:
        audit[col] = audit[col].map(lambda v: json.dumps(v, ensure_ascii=False))
    run.csv("selection.csv", audit)
    statuses = {}

    def save_assignment(method, assignments, terms):
        if len(assignments) != len(selected):
            raise ValueError("結果行数と分析対象が不一致。切り詰めず停止します。")
        assignment = pd.DataFrame({"row_id": selected["row_id"].tolist(), "topic_id": assignments})
        output = selected.drop(columns=["tokens", "tokens_raw"]).merge(assignment, on="row_id", validate="one_to_one")
        output["run_id"] = run.meta["run_id"]
        output["input_sha256"] = run.meta["input_sha256"]
        run.csv(method + "_assignments.csv", output)
        run.csv(method + "_terms.csv", pd.DataFrame(terms, columns=["topic_id", "rank", "term", "weight"]))
        monthly = output[output["month"].ne("") & output["topic_id"].ne(-1)].groupby(["date_basis", "month", "topic_id"]).size().reset_index(name="count")
        monthly["denominator"] = monthly.groupby(["date_basis", "month"])["count"].transform("sum")
        monthly["ratio"] = monthly["count"] / monthly["denominator"]
        run.csv(method + "_monthly.csv", monthly)
        if not monthly.empty:
            plot = monthly.copy(); plot["topic_id"] = plot["topic_id"].astype(str)
            run.figure(method + "_monthly.html", px.bar(plot, x="month", y="count", color="topic_id", facet_row="date_basis", title=method + "：有効日付・外れ値以外の割当件数"))
        else:
            run.meta["skipped"][method + "_monthly.html"] = "有効日付を持つ非外れ値0件"

    if len(docs) < 5 or len(set(docs)) < 2:
        statuses["lda"] = "SKIPPED_FEW_OR_IDENTICAL_DOCUMENTS"
    else:
        try:
            v = CountVectorizer(tokenizer=str.split, token_pattern=None, lowercase=False, min_df=1, max_df=1.0)
            x = v.fit_transform(docs)
            if x.shape[1] < 2:
                statuses["lda"] = "SKIPPED_FEW_TERMS"
            else:
                model = LatentDirichletAllocation(n_components=min(10, *x.shape), learning_method="batch", random_state=42, max_iter=20)
                dist = model.fit_transform(x); names = v.get_feature_names_out()
                terms = [(k, rank, names[i], float(weights[i])) for k, weights in enumerate(model.components_) for rank, i in enumerate(weights.argsort()[::-1][:10], 1)]
                save_assignment("lda", dist.argmax(axis=1), terms)
                statuses["lda"] = "OK"
        except ValueError as exc:
            statuses["lda"] = "FAILED: " + str(exc)
    if not use_bertopic:
        statuses["bertopic"] = "SKIPPED_NOT_REQUESTED"
    elif len(docs) < 10 or len(set(docs)) < 2:
        statuses["bertopic"] = "SKIPPED_FEW_OR_IDENTICAL_DOCUMENTS"
    else:
        try:
            from bertopic import BERTopic
            from umap import UMAP
            from hdbscan import HDBSCAN
            embedding_model = None
            if embeddings is None:
                from sentence_transformers import SentenceTransformer
                embedding_model = SentenceTransformer(MODEL_ID)
            elif len(embeddings) != len(docs):
                raise ValueError("embeddingsは分析対象と同じ行数・順序が必要です。")
            model = BERTopic(embedding_model=embedding_model, language="multilingual",
                vectorizer_model=CountVectorizer(tokenizer=str.split, token_pattern=None, lowercase=False, min_df=1, max_df=1.0, ngram_range=(1, 2)),
                umap_model=UMAP(n_neighbors=min(15, len(docs)-1), n_components=5, metric="cosine", init="random", random_state=42),
                hdbscan_model=HDBSCAN(min_cluster_size=5, prediction_data=True), calculate_probabilities=False)
            labels, _ = model.fit_transform(docs, embeddings=embeddings)
            terms = [(int(k), rank, term, float(weight)) for k in sorted(set(labels)) for rank, (term, weight) in enumerate(model.get_topic(k)[:10], 1)]
            save_assignment("bertopic", labels, terms)
            count = len(set(labels) - {-1})
            statuses["bertopic"] = "ALL_OUTLIERS" if count == 0 else "OK"
            jobs = []
            if count:
                jobs.append(("bertopic_barchart.html", lambda: model.visualize_barchart(top_n_topics=min(12, count))))
            if count >= 2:
                jobs.append(("bertopic_hierarchy.html", model.visualize_hierarchy))
            if count >= 4:
                jobs.append(("bertopic_distance.html", model.visualize_topics))
            for name, make in jobs:
                try:
                    run.figure(name, make())
                except (ValueError, TypeError, IndexError) as exc:
                    run.meta["skipped"][name] = type(exc).__name__ + ": " + str(exc)
            if count < 2:
                run.meta["skipped"]["bertopic_hierarchy.html"] = "非外れ値トピック2未満"
            if count < 4:
                run.meta["skipped"]["bertopic_distance.html"] = "非外れ値トピック4未満"
        except Exception as exc:
            statuses["bertopic"] = "FAILED: " + type(exc).__name__ + ": " + str(exc)
    return run.finish(statuses=statuses, input_rows=len(df), selected_rows=len(selected),
                      selected_row_ids=selected["row_id"].tolist(), models_saved=False)

if __name__ == "__main__":
    p = argparse.ArgumentParser(); p.add_argument("input", nargs="?", default="outputs/google_play_reviews.csv"); p.add_argument("--bertopic", action="store_true")
    a = p.parse_args(); topics(a.input, use_bertopic=a.bertopic)
