from __future__ import annotations
import math, re
from datetime import datetime, timezone
from typing import Any
import pandas as pd
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.models.entities import (Dataset, EnterpriseQualityRule, QualityIncident, QualityResult,
 QualityRun, QuarantineRecord, User)
from app.services.analytics import latest_frame

def now(): return datetime.now(timezone.utc)

def profile_frame(df:pd.DataFrame)->dict[str,Any]:
    columns={}
    for name in df.columns:
        s=df[name]; non_null=int(s.notna().sum()); unique=int(s.nunique(dropna=True)); item={
            'dtype':str(s.dtype),'null_count':int(s.isna().sum()),'null_rate':round(float(s.isna().mean())*100,2),
            'non_null_count':non_null,'unique_count':unique,'unique_rate':round((unique/non_null*100) if non_null else 0,2)
        }
        if pd.api.types.is_numeric_dtype(s):
            clean=pd.to_numeric(s,errors='coerce').dropna()
            if len(clean): item.update({'min':float(clean.min()),'max':float(clean.max()),'mean':float(clean.mean()),'std':float(clean.std(ddof=0)) if len(clean)>1 else 0})
        else:
            lengths=s.dropna().astype(str).str.len()
            if len(lengths): item.update({'min_length':int(lengths.min()),'max_length':int(lengths.max()),'avg_length':round(float(lengths.mean()),2)})
            item['top_values']={str(k):int(v) for k,v in s.fillna('[NULL]').astype(str).value_counts().head(10).items()}
        columns[str(name)]=item
    return {'row_count':len(df),'column_count':len(df.columns),'columns':columns}

def detect_anomalies(df:pd.DataFrame)->dict[str,Any]:
    findings=[]
    for col in df.select_dtypes(include='number').columns:
        s=pd.to_numeric(df[col],errors='coerce').dropna()
        if len(s)<5: continue
        q1,q3=s.quantile(.25),s.quantile(.75); iqr=q3-q1; low,high=q1-1.5*iqr,q3+1.5*iqr
        count=int(((s<low)|(s>high)).sum())
        if count: findings.append({'column':str(col),'method':'iqr','count':count,'lower':float(low),'upper':float(high)})
    return {'findings':findings,'total_anomalies':sum(x['count'] for x in findings)}

def _failure_mask(df:pd.DataFrame,rule:EnterpriseQualityRule):
    cfg=rule.config or {}; col=rule.column_name; t=rule.rule_type
    if col and col not in df.columns: return pd.Series([True]*len(df),index=df.index),{'error':'column_missing'}
    if t=='not_null': return df[col].isna(),{}
    if t=='unique': return df[col].duplicated(keep=False)&df[col].notna(),{}
    if t=='accepted_values': return ~df[col].isin(cfg.get('values',[])),{}
    if t=='regex': return ~df[col].fillna('').astype(str).str.match(cfg.get('pattern','.*')) ,{}
    if t=='range':
        s=pd.to_numeric(df[col],errors='coerce'); mask=s.isna()
        if cfg.get('min') is not None: mask=mask|(s<float(cfg['min']))
        if cfg.get('max') is not None: mask=mask|(s>float(cfg['max']))
        return mask,{}
    if t=='length':
        lengths=df[col].fillna('').astype(str).str.len(); mask=pd.Series(False,index=df.index)
        if cfg.get('min') is not None: mask=mask|(lengths<int(cfg['min']))
        if cfg.get('max') is not None: mask=mask|(lengths>int(cfg['max']))
        return mask,{}
    if t=='freshness':
        s=pd.to_datetime(df[col],errors='coerce',utc=True); days=int(cfg.get('max_age_days',1)); threshold=pd.Timestamp.now(tz='UTC')-pd.Timedelta(days=days)
        return s.isna()|(s<threshold),{'threshold':threshold.isoformat()}
    if t=='row_count':
        ok=len(df)>=int(cfg.get('min',0)) and (cfg.get('max') is None or len(df)<=int(cfg['max']))
        return pd.Series([not ok]*len(df),index=df.index),{'row_count':len(df)}
    if t=='completeness':
        columns=[x for x in cfg.get('columns',[]) if x in df.columns]; mask=df[columns].isna().any(axis=1) if columns else pd.Series([False]*len(df),index=df.index)
        return mask,{'columns':columns}
    raise ValueError(f'Unsupported quality rule: {t}')

def execute_quality_run(db:Session,dataset:Dataset,user:User,trigger_type='manual')->QualityRun:
    run=QualityRun(dataset_id=dataset.id,trigger_type=trigger_type,status='running',initiated_by=user.id); db.add(run); db.flush()
    try:
        df=latest_frame(db,dataset); rules=db.scalars(select(EnterpriseQualityRule).where(EnterpriseQualityRule.dataset_id==dataset.id,EnterpriseQualityRule.enabled.is_(True))).all()
        dimension_points={}; dimension_weights={}; passed=0
        for rule in rules:
            mask,metrics=_failure_mask(df,rule); failed=int(mask.sum()); checked=max(len(df),1); score=max(0,round((1-failed/checked)*100)); status='passed' if score>=int((rule.config or {}).get('minimum_score',100)) else 'failed'
            result=QualityResult(quality_run_id=run.id,rule_id=rule.id,status=status,score=score,checked_rows=len(df),failed_rows=failed,failure_rate=round(failed/checked*100),sample_failures=df.loc[mask].head(20).where(pd.notnull(df.loc[mask].head(20)),None).to_dict('records'),metrics=metrics)
            db.add(result); rule.last_status=status; rule.last_score=score; rule.last_run_at=now(); w=max(rule.weight,1); dimension_points[rule.dimension]=dimension_points.get(rule.dimension,0)+score*w; dimension_weights[rule.dimension]=dimension_weights.get(rule.dimension,0)+w
            if status=='passed': passed+=1
            else:
                incident=QualityIncident(dataset_id=dataset.id,quality_run_id=run.id,rule_id=rule.id,title=f'Falha: {rule.name}',description=f'{failed} de {len(df)} registros falharam.',severity=rule.severity,owner_id=rule.owner_id); db.add(incident)
                if rule.quarantine_on_failure:
                    for idx,row in df.loc[mask].head(100).iterrows(): db.add(QuarantineRecord(dataset_id=dataset.id,quality_run_id=run.id,rule_id=rule.id,row_reference=str(idx),row_data={k:(None if pd.isna(v) else v) for k,v in row.to_dict().items()},reason=rule.name))
        dims={k:round(dimension_points[k]/dimension_weights[k]) for k in dimension_points}; run.total_rules=len(rules); run.passed_rules=passed; run.failed_rules=len(rules)-passed; run.dimension_scores=dims; run.overall_score=round(sum(dimension_points.values())/sum(dimension_weights.values())) if dimension_weights else 100; run.status='completed'; run.finished_at=now()
    except Exception as exc:
        run.status='failed'; run.error=str(exc); run.finished_at=now()
    db.commit(); db.refresh(run); return run
