"""Orchestrator for the weekly sales upload.

Pipeline (one HTTP request):
    1. Parse each uploaded Excel into a normalised DataFrame
    2. Discover new SKUs and insert them into dim_products
    3. Discover new partners / stores and insert them into dim_partners / dim_stores
    4. Resolve SKU → product_id, partner code → partner_id, unit → store_id,
       tip → channel_map_id
    5. Bulk-insert rows into erp_transactions (batched 1000 at a time) with the
       FULL per-row detail (partner, store, sales rep, document, all financials)
    6. REFRESH MATERIALIZED VIEW v_sales_weekly_full (concurrent if possible)
    7. Return summary counts + warnings

The Rekapitulacija carries per-row partner ('Partner' code + 'Naziv partnera'),
store ('Jedinica'), sales rep ('Komercijalist'), document and full financials —
all persisted so the app can slice by buyer/store/rep, not just by channel.
Partner codes are matched zero-stripped (Excel depads '05770' → '5770'); new
partners/stores auto-discover like db/migrate_sales.py.
"""
from __future__ import annotations

from typing import Optional

import pandas as pd
from sqlalchemy.orm import Session

from backend.repositories.upload_repo import (
    UploadRepository, parse_excel_file, _partner_key,
)


class UploadService:
    def __init__(self, db: Session):
        self.db = db
        self.repo = UploadRepository(db)

    def _discover(self, df: pd.DataFrame, code_col: str, name_col: str,
                  current_map: dict, create_fn, *,
                  code_field: str, name_field: str) -> int:
        """Find codes in `code_col` not already in `current_map` (matched
        zero-stripped), insert them via `create_fn`, and merge the result
        back into `current_map`. Returns the count newly created."""
        if code_col not in df.columns:
            return 0
        sub = df[[code_col, name_col]].copy()
        sub["country"] = df["country"] if "country" in df.columns else None
        sub["_code"] = sub[code_col].astype(str).str.strip()
        sub = sub[(sub["_code"] != "") & (sub["_code"].str.lower() != "nan")]
        if sub.empty:
            return 0
        sub["_key"] = sub["_code"].map(_partner_key)
        unknown = [k for k in sub["_key"].dropna().unique() if k not in current_map]
        if not unknown:
            return 0
        info = (sub.groupby("_key")
                   .agg(code=("_code", "first"), name=(name_col, "first"),
                        country=("country", "first"))
                   .to_dict("index"))

        def _clean(v):
            return None if (v is None or pd.isna(v)) else str(v).strip() or None

        recs = [{code_field: info[k]["code"],
                 name_field: _clean(info[k]["name"]),
                 "country":  _clean(info[k]["country"])}
                for k in unknown]
        added = create_fn(recs)
        current_map.update(added)
        return len(added)

    def process_weekly_update(
        self,
        files: list[tuple[str, bytes]],
    ) -> dict:
        """Parse, discover new entities, insert transactions, refresh view.

        files: list of (filename, bytes) tuples.
        Returns: dict with rows_added, n_files, new_products, new_partners,
                 new_stores, warnings, view_refreshed, date_range.
        """
        if not files:
            return {
                "rows_added":      0,
                "n_files":         0,
                "new_products":    0,
                "new_partners":    0,
                "new_stores":      0,
                "warnings":        ["No files supplied"],
                "view_refreshed":  False,
                "date_range":      None,
            }

        warnings: list[str] = []
        all_dfs: list[pd.DataFrame] = []
        for fname, data in files:
            try:
                parsed = parse_excel_file(data, filename=fname)
                all_dfs.append(parsed["df"])
                warnings.extend(parsed["warnings"])
            except Exception as exc:
                warnings.append(f"{fname}: {exc}")

        if not all_dfs:
            return {
                "rows_added":     0,
                "n_files":        len(files),
                "new_products":   0,
                "new_partners":   0,
                "new_stores":     0,
                "warnings":       warnings or ["No usable rows in any file"],
                "view_refreshed": False,
                "date_range":     None,
            }

        df = pd.concat(all_dfs, ignore_index=True)
        if df.empty:
            return {
                "rows_added":     0,
                "n_files":        len(files),
                "new_products":   0,
                "new_partners":   0,
                "new_stores":     0,
                "warnings":       warnings or ["No valid transactions after filtering"],
                "view_refreshed": False,
                "date_range":     None,
            }

        date_min = df["date"].min()
        date_max = df["date"].max()
        date_range = {
            "from": date_min.date().isoformat() if pd.notna(date_min) else None,
            "to":   date_max.date().isoformat() if pd.notna(date_max) else None,
            "n_weeks": int(df["date"].dt.isocalendar().week.nunique()),
        }

        # ----- Discover new SKUs -----
        existing_skus = self.repo.get_existing_skus()
        df_skus = set(df["sku"].unique())
        new_sku_codes = df_skus - existing_skus.keys()
        new_products: list[dict] = []
        if new_sku_codes:
            # Pick first non-empty name/cat seen for each new SKU
            for sku in new_sku_codes:
                slice_df = df[df["sku"] == sku]
                name = ""
                cat = ""
                for _, r in slice_df.iterrows():
                    if not name and isinstance(r.get("name"), str) and r["name"].strip():
                        name = r["name"].strip()
                    if not cat and isinstance(r.get("cat"), str) and r["cat"].strip():
                        cat = r["cat"].strip()
                    if name and cat:
                        break
                new_products.append({"sku": sku, "name": name, "category": cat})

        created_map = self.repo.create_products_batch(new_skus=new_products) if new_products else {}
        sku_to_pid = {**existing_skus, **created_map}

        # ----- Resolve channel map -----
        channel_map = self.repo.get_channel_map()

        # ----- Discover + resolve partners and stores -----
        # Codes are matched zero-stripped (Excel depads). Unknown codes
        # auto-discover into dim_partners / dim_stores, mirroring migrate_sales.
        partner_map = self.repo.get_partner_map()
        store_map = self.repo.get_store_map()
        new_partners = self._discover(df, "partner", "partner_name",
                                      partner_map, self.repo.create_partners_batch,
                                      code_field="code", name_field="name")
        new_stores = self._discover(df, "store", "store_name",
                                    store_map, self.repo.create_stores_batch,
                                    code_field="unit_code", name_field="name")

        def _id_for(code, mp):
            if code is None or (isinstance(code, float) and pd.isna(code)):
                return None
            s = str(code).strip()
            return mp.get(_partner_key(s)) if s and s.lower() != "nan" else None

        # ----- Build insert rows (full per-row detail) -----
        rows: list[dict] = []
        n_skipped_missing_pid = 0
        n_skipped_missing_chan = 0

        def _num(v) -> Optional[float]:
            if v is None or (pd.isna(v) if not isinstance(v, str) else v.strip() == ""):
                return None
            try:
                return float(v)
            except (TypeError, ValueError):
                return None

        def _str(v) -> Optional[str]:
            if v is None or (isinstance(v, float) and pd.isna(v)):
                return None
            s = str(v).strip()
            return s if s and s.lower() != "nan" else None

        for _, r in df.iterrows():
            sku = str(r["sku"])
            pid = sku_to_pid.get(sku)
            if pid is None:
                n_skipped_missing_pid += 1
                continue
            chan_id = channel_map.get(str(r["tip"]))
            if chan_id is None:
                n_skipped_missing_chan += 1
                continue
            rows.append({
                "transaction_date":  r["date"].date(),
                "product_id":        int(pid),
                "channel_map_id":    int(chan_id),
                "quantity":          float(r["qty"]),
                "total_value":       float(r["value"]),
                "ruc_eur":           _num(r.get("ruc")),
                "partner_id":        _id_for(r.get("partner"), partner_map),
                "store_id":          _id_for(r.get("store"), store_map),
                "document":          _str(r.get("document")),
                "sales_rep":         _str(r.get("sales_rep")),
                "purchase_value":    _num(r.get("purchase_value")),
                "ruc_pct":           _num(r.get("ruc_pct")),
                "tax_base":          _num(r.get("tax_base")),
                "vat":               _num(r.get("vat")),
                "approved_discount": _num(r.get("discount")),
                "has_loyalty":       bool(r.get("has_loyalty")) if pd.notna(r.get("has_loyalty")) else False,
            })

        if n_skipped_missing_pid > 0:
            warnings.append(f"Skipped {n_skipped_missing_pid} rows — could not resolve product_id")
        if n_skipped_missing_chan > 0:
            warnings.append(f"Skipped {n_skipped_missing_chan} rows — unknown tip/doc-type")

        n_inserted = self.repo.insert_transactions(rows)
        view_refreshed = self.repo.refresh_sales_weekly() if n_inserted > 0 else False

        return {
            "rows_added":     n_inserted,
            "n_files":        len(files),
            "new_products":   len(created_map),
            "new_partners":   new_partners,
            "new_stores":     new_stores,
            "warnings":       warnings,
            "view_refreshed": view_refreshed,
            "date_range":     date_range,
        }
