import csv, gzip, hashlib, json, os, tempfile
from datetime import datetime, timezone, timedelta
from pathlib import Path
from sqlalchemy import select
from app.models.entities import (FederatedQueryJob,FederatedJoinDefinition,FederatedWorkerRun,QuerySpillSegment,QueryExport,SemanticModel,QueryUsageLedger)
from app.services.semantic_execution import execute_semantic

WORK_ROOT=Path(os.getenv('FEDERATED_WORK_ROOT','/tmp/cedp-federated'))
EXPORT_ROOT=Path(os.getenv('FEDERATED_EXPORT_ROOT','/tmp/cedp-exports'))

def _now(): return datetime.now(timezone.utc)
def _safe_root(root): root.mkdir(parents=True,exist_ok=True); return root
def _row_dict(columns,row): return row if isinstance(row,dict) else dict(zip(columns,row))
def _query_side(db,user,model_id,spec,limit):
    model=db.get(SemanticModel,model_id)
    if not model: raise ValueError('Modelo semântico do join não encontrado')
    result=execute_semantic(db,user,model,spec.get('metrics',[]),spec.get('dimensions',[]),spec.get('filters',{}),limit,spec.get('use_cache',True))
    cols=result.get('columns',[])
    return cols,[_row_dict(cols,r) for r in result.get('rows',[])]

def _spill(db,run,side,rows,partition_size=10000):
    base=_safe_root(WORK_ROOT/run.id); segments=[]
    for no,start in enumerate(range(0,len(rows),partition_size)):
        chunk=rows[start:start+partition_size]; path=base/f'{side}-{no:06d}.jsonl.gz'
        h=hashlib.sha256()
        with gzip.open(path,'wt',encoding='utf-8') as fh:
            for row in chunk:
                line=json.dumps(row,ensure_ascii=False,separators=(',',':'))+'\n'; fh.write(line); h.update(line.encode())
        seg=QuerySpillSegment(worker_run_id=run.id,side=side,partition_no=no,object_key=str(path),checksum=h.hexdigest(),row_count=len(chunk),size_bytes=path.stat().st_size,encrypted=False,expires_at=_now()+timedelta(hours=24))
        db.add(seg); segments.append(seg); run.spilled_bytes+=seg.size_bytes
    db.commit(); return segments

def _load_segments(segments):
    rows=[]
    for seg in segments:
        path=Path(seg.object_key); h=hashlib.sha256()
        with gzip.open(path,'rt',encoding='utf-8') as fh:
            for line in fh:
                h.update(line.encode()); rows.append(json.loads(line))
        if h.hexdigest()!=seg.checksum: raise ValueError('Checksum inválido em segmento temporário')
    return rows

def _hash_join(left,right,left_key,right_key,join_type,max_output):
    index={}
    for r in right: index.setdefault(str(r.get(right_key)),[]).append(r)
    output=[]; matched_right=set()
    for l in left:
        matches=index.get(str(l.get(left_key)),[])
        if matches:
            for r in matches:
                merged={f'left.{k}':v for k,v in l.items()}; merged.update({f'right.{k}':v for k,v in r.items()}); output.append(merged); matched_right.add(id(r))
                if len(output)>=max_output: return output
        elif join_type in {'left','full'}:
            output.append({f'left.{k}':v for k,v in l.items()})
    if join_type in {'right','full'}:
        for r in right:
            if id(r) not in matched_right: output.append({f'right.{k}':v for k,v in r.items()})
            if len(output)>=max_output: break
    return output

def execute_federated_job(db,user,job,memory_limit_mb=256,worker_name='federated-worker'):
    if not job.request_json.get('join_id'): raise ValueError('Job não contém join federado')
    join=db.get(FederatedJoinDefinition,job.request_json['join_id'])
    if not join or not join.enabled: raise ValueError('Join governado não encontrado')
    run=FederatedWorkerRun(job_id=job.id,worker_name=worker_name,status='running',phase='left_source',memory_limit_mb=max(32,memory_limit_mb),heartbeat_at=_now(),started_at=_now())
    db.add(run); job.status='running'; job.progress_percent=5; db.commit(); db.refresh(run)
    try:
        req=job.request_json; left_spec=req.get('left',{}); right_spec=req.get('right',{}); side_limit=min(join.max_rows_per_side,int(req.get('max_rows_per_side',join.max_rows_per_side)))
        lcols,left=_query_side(db,user,join.left_model_id,left_spec,side_limit); run.checkpoint={'phase':'left_complete','rows':len(left)}; run.heartbeat_at=_now(); job.progress_percent=30; db.commit()
        run.phase='right_source'; rcols,right=_query_side(db,user,join.right_model_id,right_spec,side_limit); run.checkpoint={'phase':'right_complete','left_rows':len(left),'right_rows':len(right)}; job.progress_percent=55; db.commit()
        estimated=sum(len(json.dumps(r,default=str)) for r in left+right); run.peak_memory_mb=max(1,estimated//(1024*1024))
        if estimated>run.memory_limit_mb*1024*1024:
            run.phase='spill'; ls=_spill(db,run,'left',left); rs=_spill(db,run,'right',right); left=_load_segments(ls); right=_load_segments(rs)
        if job.cancel_requested: raise InterruptedError('Consulta cancelada')
        run.phase='join'; job.progress_percent=75; db.commit()
        rows=_hash_join(left,right,join.left_key,join.right_key,join.join_type,int(req.get('max_output_rows',100000)))
        columns=sorted({k for row in rows for k in row})
        job.execution_plan={**(job.execution_plan or {}),'strategy':'application_hash_join','worker_run_id':run.id,'result':{'columns':columns,'rows':rows},'left_rows':len(left),'right_rows':len(right),'spilled_bytes':run.spilled_bytes}
        job.row_count=len(rows); job.progress_percent=100; job.status='success'; job.finished_at=_now(); run.status='success'; run.phase='complete'; run.finished_at=_now(); run.checkpoint={'phase':'complete','output_rows':len(rows)}
        db.add(QueryUsageLedger(user_id=user.id,job_id=job.id,rows_processed=len(left)+len(right)+len(rows),cost_units=max(1,len(left)+len(right)),operation='federated_join'))
        db.commit(); db.refresh(job); return job
    except InterruptedError as exc:
        job.status='cancelled'; job.error=str(exc); job.finished_at=_now(); run.status='cancelled'; run.error=str(exc); run.finished_at=_now(); db.commit(); return job
    except Exception as exc:
        job.status='failed'; job.error=str(exc)[:2000]; job.finished_at=_now(); run.status='failed'; run.error=str(exc)[:2000]; run.finished_at=_now(); db.commit(); raise

def resume_worker_run(db,user,run):
    job=db.get(FederatedQueryJob,run.job_id)
    if not job: raise ValueError('Job do checkpoint não encontrado')
    if job.status=='success': return job
    return execute_federated_job(db,user,job,run.memory_limit_mb,run.worker_name)

def create_export(db,user,job,fmt='csv',ttl_hours=24):
    if job.status!='success': raise ValueError('Somente jobs concluídos podem ser exportados')
    fmt=fmt.lower();
    if fmt not in {'csv','jsonl'}: raise ValueError('Formato de exportação inválido')
    export=QueryExport(job_id=job.id,user_id=user.id,format=fmt,status='running',expires_at=_now()+timedelta(hours=ttl_hours)); db.add(export); db.commit(); db.refresh(export)
    try:
        result=(job.execution_plan or {}).get('result',{}); rows=result.get('rows',[]); cols=result.get('columns',[]) or sorted({k for r in rows for k in r})
        base=_safe_root(EXPORT_ROOT/user.id); path=base/f'{export.id}.{fmt}.gz'; h=hashlib.sha256()
        with gzip.open(path,'wt',encoding='utf-8',newline='') as fh:
            if fmt=='csv':
                writer=csv.DictWriter(fh,fieldnames=cols,extrasaction='ignore'); writer.writeheader(); writer.writerows(rows)
            else:
                for row in rows: fh.write(json.dumps(row,ensure_ascii=False,default=str)+'\n')
        with open(path,'rb') as raw:
            for chunk in iter(lambda:raw.read(1024*1024),b''): h.update(chunk)
        export.object_key=str(path); export.checksum=h.hexdigest(); export.row_count=len(rows); export.size_bytes=path.stat().st_size; export.status='success'; export.finished_at=_now()
        db.add(QueryUsageLedger(user_id=user.id,job_id=job.id,rows_processed=len(rows),cost_units=max(1,len(rows)//10),operation='export'))
        db.commit(); db.refresh(export); return export
    except Exception as exc:
        export.status='failed'; export.error=str(exc)[:2000]; export.finished_at=_now(); db.commit(); raise
