# exports/export_xlsx.py
from pathlib import Path
from typing import Dict, Any, List

import pandas as pd
from openpyxl.styles import Font, Alignment, NamedStyle, PatternFill
from openpyxl.utils import get_column_letter
from openpyxl.formatting.rule import Rule
from openpyxl.styles.differential import DifferentialStyle

def export_xlsx(
    out_path: Path,
    lv_meta: Any,
    offer_positions: Dict[str, Dict[str, Dict[str, Any]]],
) -> None:
    if not offer_positions:
        raise ValueError("Keine analysierten Offert-Positionen zum Exportieren vorhanden.")

    vendors_with_data = {v: p for v, p in offer_positions.items() if any(d.get('ep') or d.get('gp') for d in p.values())}
    if not vendors_with_data:
        raise ValueError("Die analysierten Offerten enthalten keine gültigen Preisdaten.")

    vendors = sorted(list(vendors_with_data.keys()))
    df = _build_comparison_df(lv_meta, vendors_with_data, vendors)

    with pd.ExcelWriter(str(out_path), engine="openpyxl") as writer:
        df.to_excel(writer, sheet_name="Preisvergleich", index=False)
        ws = writer.book["Preisvergleich"]
        _format_sheet(ws, len(vendors))

def _build_comparison_df(lv_meta: Any, offer_positions: Dict[str, Dict], vendors: List[str]) -> pd.DataFrame:
    rows = []
    for pos_code in sorted(lv_meta.positions.keys()):
        lv_pos = lv_meta.positions[pos_code]
        row = {
            "Pos": pos_code,
            "Kurztext": lv_pos.get("text", ""),
            "Menge": lv_pos.get("qty"),
            "Einheit": lv_pos.get("unit", ""),
        }
        for vendor in vendors:
            prices = offer_positions.get(vendor, {}).get(pos_code, {})
            ep, gp = prices.get("ep"), prices.get("gp")
            if gp is None and ep is not None and row["Menge"] is not None:
                gp = ep * row["Menge"]
            row[f"{vendor} (EP)"] = ep
            row[f"{vendor} (GP)"] = gp
        rows.append(row)
    return pd.DataFrame(rows)

def _format_sheet(ws, num_vendors: int):
    _bold_header_and_freeze(ws)
    chf_style = NamedStyle(name='chf_currency', number_format='#,##0.00 "CHF"')
    
    price_col_indices = [i for i in range(5, 5 + 2 * num_vendors)]
    for col_idx in price_col_indices:
        for row in range(2, ws.max_row + 2):
            ws.cell(row=row, column=col_idx).style = chf_style
    
    total_row_idx = ws.max_row + 1
    ws.cell(row=total_row_idx, column=2, value="Gesamtsumme").font = Font(bold=True)
    for col_idx in price_col_indices:
        header_val = ws.cell(row=1, column=col_idx).value
        if header_val and "(GP)" in header_val:
            col_letter = get_column_letter(col_idx)
            formula = f"=SUM({col_letter}2:{col_letter}{ws.max_row-1})"
            cell = ws.cell(row=total_row_idx, column=col_idx, value=formula)
            cell.style = chf_style
            cell.font = Font(bold=True)

    green_fill = PatternFill(start_color="C6EFCE", end_color="C6EFCE", fill_type="solid")
    green_font = Font(color="006100")
    dxf = DifferentialStyle(font=green_font, fill=green_fill)
    
    ep_col_indices = [i for i in range(5, 5 + 2 * num_vendors, 2)]
    
    for row in range(2, ws.max_row):
        ep_cells = [f"${get_column_letter(i)}${row}" for i in ep_col_indices]
        ep_cells_str = ",".join(ep_cells)
        for col_idx in ep_col_indices:
            col_letter = get_column_letter(col_idx)
            rule_formula = f"AND(ISNUMBER(${col_letter}{row}), ${col_letter}{row}=MIN({ep_cells_str}))"
            rule = Rule(type="expression", dxf=dxf, formula=[rule_formula])
            ws.conditional_formatting.add(f"{col_letter}{row}", rule)

    _auto_width(ws)

def _bold_header_and_freeze(ws):
    for cell in ws[1]:
        cell.font = Font(bold=True)
        cell.alignment = Alignment(horizontal='center', vertical='center', wrap_text=True)
    ws.freeze_panes = "A2"

def _auto_width(ws):
    for col in ws.columns:
        max_length = 0
        column = get_column_letter(col[0].column)
        for cell in col:
            try:
                if len(str(cell.value)) > max_length: max_length = len(str(cell.value))
            except: pass
        adjusted_width = (max_length + 2) if max_length < 60 else 60
        ws.column_dimensions[column].width = adjusted_width