import io
from datetime import date, datetime, timedelta

import pandas as pd
import streamlit as st
from rdflib import Graph, Namespace, Literal
from rdflib.namespace import RDF, RDFS, SKOS, XSD, DCTERMS, SH
from pyshacl import validate

# ------------------------------------------------------------------ CONFIG
ONTOLOGY_FILE = "fdic_ontology_fixed.ttl"
SHACL_FILE    = "fdic_shacl_constraints_optimized.ttl"

FDIC = Namespace("https://example.org/fdic#")
EX   = Namespace("https://example.org/data#")
DCAT = Namespace("http://www.w3.org/ns/dcat#")
VOID = Namespace("http://rdfs.org/ns/void#")

DATE_COLUMNS = {"ACQDATE", "ESTYMD", "RUNDATE"}

# ------------------------------------------------------------------ HELPERS
@st.cache_resource(show_spinner=False)
def load_ontology(path: str):
    g = Graph().parse(path, format="turtle")
    mapping = {str(o): s for s, _, o in g.triples((None, SKOS.altLabel, None))}
    return g, mapping

@st.cache_resource(show_spinner=False)
def load_shacl_graph(path: str):
    return Graph().parse(path, format="turtle")

def excel_serial_to_iso(value: str) -> str | None:
    v = str(value).strip()
    if not v:
        return None
    if v.isdigit():
        try:
            return (datetime(1899, 12, 30) + timedelta(days=int(v))).strftime("%Y-%m-%d")
        except Exception:
            return None
    for fmt in ("%m/%d/%Y", "%d/%m/%Y", "%Y-%m-%d"):
        try:
            return datetime.strptime(v, fmt).strftime("%Y-%m-%d")
        except ValueError:
            continue
    return None

def transform(csv_bytes: bytes, fname: str) -> tuple[str, int]:
    onto, col2pred = load_ontology(ONTOLOGY_FILE)
    df = pd.read_csv(io.BytesIO(csv_bytes), dtype=str).fillna("")

    g = Graph()
    for p, ns in [("fdic", FDIC), ("skos", SKOS), ("rdfs", RDFS), ("xsd", XSD),
                  ("dct", DCTERMS), ("dcat", DCAT), ("void", VOID)]:
        g.bind(p, ns)

    for idx, row in df.iterrows():
        subj = EX[f"bankbranch-{idx}"]
        g.add((subj, RDF.type, FDIC.BankBranch))
        for col, val in row.items():
            if col not in col2pred:
                continue
            pred  = col2pred[col]
            rtype = onto.value(subject=pred, predicate=RDFS.range)
            if col in DATE_COLUMNS:
                iso = excel_serial_to_iso(val)
                if iso:
                    g.add((subj, pred, Literal(iso, datatype=XSD.date)))
                continue
            if rtype == SKOS.Concept:
                g.add((subj, pred, FDIC[val]))
            else:
                g.add((subj, pred, Literal(val)))

    # DCAT/VoID metadata
    dataset, dist = EX["fdic-dataset"], EX["fdic-dist"]
    g.add((dataset, RDF.type, DCAT.Dataset))
    g.add((dataset, DCTERMS.title, Literal("FDIC Insured Banks Data", lang="en")))
    g.add((dataset, DCTERMS.issued, Literal(date.today().isoformat(), datatype=XSD.date)))
    g.add((dataset, DCAT.distribution, dist))
    g.add((dist, RDF.type, DCAT.Distribution))
    g.add((dist, DCTERMS.format, Literal("text/turtle")))
    g.add((dataset, RDF.type, VOID.Dataset))
    g.add((dataset, VOID.triples, Literal(len(g), datatype=XSD.integer)))

    ttl_bytes = g.serialize(format="turtle").encode("utf-8")
    return ttl_bytes.decode("utf-8"), len(g)

# ------------------------------------------------------------------ NEW: robust RDF subset extractor
def extract_first_n_subjects_rdf(data_ttl: str, n: int) -> str:
    """
    Safely create Turtle containing only triples for the first N fdic:BankBranch
    subjects by using rdflib, so syntax stays valid.
    """
    full_g = Graph().parse(data=data_ttl, format="turtle")
    subset = Graph()
    for p, ns in full_g.namespaces():
        subset.bind(p, ns)

    subjects = list(full_g.subjects(RDF.type, FDIC.BankBranch))[:n]
    for s in subjects:
        for p, o in full_g.predicate_objects(s):
            subset.add((s, p, o))

    ttl = subset.serialize(format="turtle")
    return ttl.decode("utf-8") if isinstance(ttl, bytes) else ttl

def run_shacl_validation(data_ttl: str, shacl_path: str, record_limit):
    sample_ttl = (
        extract_first_n_subjects_rdf(data_ttl, record_limit)
        if isinstance(record_limit, int)
        else data_ttl
    )

    data_g  = Graph().parse(data=sample_ttl, format="turtle")
    shacl_g = load_shacl_graph(shacl_path)

    conforms, results_g, results_txt = validate(
        data_graph=data_g,
        shacl_graph=shacl_g,
        inference="none",
        serialize_report_graph=True
    )

    issues = [{
        "Focus Node":  str(results_g.value(r, SH.focusNode)),
        "Path":        str(results_g.value(r, SH.resultPath)),
        "Message":     str(results_g.value(r, SH.resultMessage))
    } for r in results_g.subjects(RDF.type, SH.ValidationResult)]

    return conforms, pd.DataFrame(issues), results_txt

# ------------------------------------------------------------------ UI
st.set_page_config(page_title="FDIC CSV → RDF Transformer", layout="centered")
st.title("🚀 FDIC Semantic Transformer")

uploaded = st.file_uploader("Upload FDIC_Insured_Banks CSV", type=["csv"])
if uploaded:
    st.success(f"Loaded **{uploaded.name}** – {uploaded.size/1_048_576:.2f} MB")

    if st.button("🔄 Transform to RDF"):
        with st.spinner("Processing …"):
            ttl_text, triple_count = transform(uploaded.getvalue(), uploaded.name)
            st.session_state.ttl_text     = ttl_text
            st.session_state.triple_count = triple_count
            st.session_state.df_full      = pd.read_csv(uploaded)

    # Show results only after transform
    if "ttl_text" in st.session_state:
        ttl_text      = st.session_state.ttl_text
        triple_count  = st.session_state.triple_count
        df_full       = st.session_state.df_full

        st.subheader("📊 First Record from Source CSV")
        st.dataframe(df_full.head(1))

        first_subject = "<" + ttl_text.split("\n<")[1].split("\n<")[0].strip()
        st.subheader("🐢 Turtle for First Record Only")
        st.code(first_subject, language="turtle")

        st.info(f"Generated **{triple_count:,} triples** from **{len(df_full)} records**")

        st.download_button("⬇️ Download full Turtle",
                           data=ttl_text,
                           file_name="fdic_data_transformed.ttl",
                           mime="text/turtle",
                           use_container_width=True)

        st.markdown("---")
        num_to_validate = st.selectbox(
            "🔎 Number of records to validate with SHACL",
            options=[1, 2, 5, 10, "All"],
            index=1
        )

        if st.button("✅ Validate with SHACL", type="primary"):
            with st.spinner("Running SHACL validation …"):
                conforms, issues_df, report_txt = run_shacl_validation(
                    ttl_text,
                    SHACL_FILE,
                    num_to_validate if num_to_validate != "All" else "All"
                )

            if conforms:
                st.success("🎉  Data conforms to all SHACL constraints!")
            else:
                st.error(f"⚠️  Data does NOT conform – {len(issues_df)} issue(s) found.")
                st.dataframe(issues_df, use_container_width=True)
                with st.expander("Raw SHACL report"):
                    st.code(report_txt, language="turtle")
                st.download_button("⬇️ Download SHACL report",
                                   data=report_txt,
                                   file_name="fdic_shacl_report.ttl",
                                   mime="text/turtle",
                                   use_container_width=True)
else:
    st.info("Upload the FDIC CSV to begin.")