from __future__ import annotations
import json, math, sqlite3, statistics, shutil, uuid
from pathlib import Path
from datetime import datetime

from fastapi import FastAPI, UploadFile, File, Form, HTTPException
from fastapi.responses import FileResponse
from fastapi.staticfiles import StaticFiles
from openpyxl import load_workbook

BASE = Path(__file__).resolve().parent
STATIC = BASE / 'static'
STORAGE = BASE / 'storage'
UPLOADS = BASE / 'uploads'
DB = BASE / 'spylogger.db'
for p in (STATIC, STORAGE, UPLOADS): p.mkdir(exist_ok=True)

app = FastAPI(title='Spy Logger Línea A · Multi Tren')
app.mount('/static', StaticFiles(directory=STATIC), name='static')

FIELDS = ['TIME','TAC OK','VIGILANCIA','CLT2','TKV OK','CKD OK','PA','CMC','MAV15','MAV35','PAHS','AUPA','CMHS','DISPA','SVIT','NO SALIDA','ARRET','OG','OD','PODF','POGF','VIT TREN','VIT MAX','VIT BUT','DIST BUT','ABSCISA','SECTOR','BALIZA ID','CANAL','IdLog','TREN TIPO','MTYP']
STATIONS = ['Pantitlán','Agrícola Oriental','Canal de San Juan','Tepalcates','Guelatao','Peñón Viejo','Acatitla','Santa Marta','Los Reyes','La Paz']


def db():
    con = sqlite3.connect(DB)
    con.row_factory = sqlite3.Row
    return con


def init_db():
    con = db()
    con.executescript('''
    CREATE TABLE IF NOT EXISTS trains(
      id INTEGER PRIMARY KEY AUTOINCREMENT,
      code TEXT UNIQUE NOT NULL,
      display_name TEXT NOT NULL,
      created_at TEXT NOT NULL
    );
    CREATE TABLE IF NOT EXISTS recordings(
      id TEXT PRIMARY KEY,
      train_code TEXT NOT NULL,
      original_name TEXT NOT NULL,
      stored_xlsx TEXT NOT NULL,
      rows_count INTEGER NOT NULL DEFAULT 0,
      start_time TEXT,
      end_time TEXT,
      status TEXT NOT NULL DEFAULT 'processing',
      created_at TEXT NOT NULL,
      meta_path TEXT,
      error_text TEXT
    );
    ''')
    con.commit(); con.close()

init_db()


def num(v, d=0.0):
    try:
        x = float(v)
        return x if math.isfinite(x) else d
    except Exception:
        return d


def clean(v):
    if isinstance(v, datetime): return v.isoformat(timespec='milliseconds')
    if v is None: return None
    if isinstance(v, float):
        if not math.isfinite(v): return None
        return round(v, 6)
    if isinstance(v, (int, bool)): return int(v)
    return str(v)


def dt_seconds(a, b):
    if isinstance(a, datetime) and isinstance(b, datetime):
        return max(.01, min(2.0, (b-a).total_seconds()))
    return .315


def ensure_train(code: str):
    code = code.strip()
    if not code: raise HTTPException(400, 'Falta el número/código de tren.')
    con = db()
    con.execute('INSERT OR IGNORE INTO trains(code,display_name,created_at) VALUES(?,?,?)', (code, f'Tren {code}', datetime.now().isoformat(timespec='seconds')))
    con.commit(); con.close()


def process_excel(path: Path, recording_id: str, train_code: str, original_name: str):
    wb = load_workbook(path, read_only=True, data_only=True)
    ws = wb.active
    it = ws.iter_rows(values_only=True)
    headers = [str(v).strip() if v is not None else '' for v in next(it)]
    idx = {h:i for i,h in enumerate(headers)}
    required = ['TIME','VIT TREN','SECTOR']
    missing = [x for x in required if x not in idx]
    if missing:
        wb.close(); raise ValueError('Faltan columnas requeridas: ' + ', '.join(missing))

    raw = []
    for r in it:
        raw.append({f:(r[idx[f]] if f in idx and idx[f] < len(r) else None) for f in FIELDS})
    wb.close()
    if not raw: raise ValueError('El Excel no contiene registros.')

    candidate=[]
    for i,r in enumerate(raw):
        speed=abs(num(r['VIT TREN']))
        door=any(num(r[k]) != 0 for k in ['OG','OD','POGF','PODF'])
        s=int(round(num(r['SECTOR'])))
        if speed<=.5 and door and 1<=s<=10: candidate.append((i,s))

    episodes=[]
    if candidate:
        start_i=prev_i=candidate[0][0]; sectors=[candidate[0][1]]
        for i,s in candidate[1:]:
            if i-prev_i<=25:
                prev_i=i; sectors.append(s)
            else:
                episodes.append([start_i,prev_i,round(statistics.median(sectors))])
                start_i=prev_i=i; sectors=[s]
        episodes.append([start_i,prev_i,round(statistics.median(sectors))])

    merged=[]
    for ep in episodes:
        if merged:
            last=merged[-1]
            t1=raw[last[1]]['TIME']; t2=raw[ep[0]]['TIME']
            gap=(t2-t1).total_seconds() if isinstance(t1,datetime) and isinstance(t2,datetime) else 999
            if ep[2]==last[2] and gap<120:
                last[1]=ep[1]; continue
        merged.append(ep)
    episodes=merged

    route_pos=[None]*len(raw); at_station=[False]*len(raw)
    for a,b,s in episodes:
        anchor=(s-1)/9
        for i in range(a,b+1): route_pos[i]=anchor; at_station[i]=True

    for e1,e2 in zip(episodes,episodes[1:]):
        _,b1,s1=e1; a2,_,s2=e2; start,end=b1,a2
        if end<=start: continue
        p1=(s1-1)/9; p2=(s2-1)/9
        if s1==s2:
            for i in range(start,end+1): route_pos[i]=p1
            continue
        cum=[0.0]; total=0.0
        for i in range(start+1,end+1):
            total += (max(0,abs(num(raw[i]['VIT TREN'])))/3.6)*dt_seconds(raw[i-1]['TIME'],raw[i]['TIME'])
            cum.append(total)
        for j,i in enumerate(range(start,end+1)):
            f=cum[j]/total if total>0 else j/max(1,end-start)
            f=f*f*(3-2*f)
            route_pos[i]=p1+(p2-p1)*f

    if episodes:
        first_a,_,first_s=episodes[0]
        start_s=int(round(num(raw[0]['SECTOR'])))
        start_anchor=(max(1,min(10,start_s))-1)/9; first_anchor=(first_s-1)/9
        for i in range(first_a):
            f=i/max(1,first_a-1); route_pos[i]=start_anchor+(first_anchor-start_anchor)*f
        _,last_b,last_s=episodes[-1]
        last=(last_s-1)/9; last_sector=last_s; direction=0
        for i in range(last_b+1,len(raw)):
            s=int(round(num(raw[i]['SECTOR'])))
            if 1<=s<=10 and s!=last_sector:
                direction=1 if s>last_sector else -1; last_sector=s
            target=(max(1,min(10,s))-1)/9 if 1<=s<=10 else last
            cand=last+(target-last)*.04
            if direction>0: cand=max(last,cand)
            elif direction<0: cand=min(last,cand)
            last=max(0,min(1,cand)); route_pos[i]=last
    else:
        last=0.0
        for i,r in enumerate(raw):
            s=int(round(num(r['SECTOR'])))
            target=(max(1,min(10,s))-1)/9 if 1<=s<=10 else last
            last += (target-last)*.04; route_pos[i]=last

    last=0.0
    for i,p in enumerate(route_pos):
        if p is None: route_pos[i]=last
        else: last=p

    direction_arr=[]; curdir=-1; prev=route_pos[0]
    for p in route_pos:
        d=p-prev
        if d>1e-5: curdir=1
        elif d<-1e-5: curdir=-1
        direction_arr.append(curdir); prev=p

    out_fields=FIELDS+['ROUTE_POS','AT_STATION','DIRECTION']
    recdir=STORAGE/recording_id; recdir.mkdir(parents=True,exist_ok=True)
    chunks=[]; chunk_size=4000
    for cstart in range(0,len(raw),chunk_size):
        arr=[]
        for i in range(cstart,min(len(raw),cstart+chunk_size)):
            r=raw[i]
            row=[clean(r[f]) for f in FIELDS] + [round(route_pos[i],6),1 if at_station[i] else 0,direction_arr[i]]
            arr.append(row)
        name=f'chunk_{cstart//chunk_size:03d}.json'
        (recdir/name).write_text(json.dumps(arr,ensure_ascii=False,separators=(',',':')),encoding='utf-8')
        chunks.append(name)

    start_time=clean(raw[0]['TIME']); end_time=clean(raw[-1]['TIME'])
    meta={'id':recording_id,'train_code':train_code,'source':original_name,'rows':len(raw),'chunk_size':chunk_size,'chunks':chunks,'fields':out_fields,'stations':STATIONS,'start_time':start_time,'end_time':end_time,'station_episodes':[{'start':a,'end':b,'station_index':s-1,'station':STATIONS[s-1],'start_time':clean(raw[a]['TIME']),'end_time':clean(raw[b]['TIME'])} for a,b,s in episodes]}
    (recdir/'meta.json').write_text(json.dumps(meta,ensure_ascii=False,separators=(',',':')),encoding='utf-8')

    con=db(); con.execute("UPDATE recordings SET rows_count=?,start_time=?,end_time=?,status='ready',meta_path=?,error_text=NULL WHERE id=?",(len(raw),start_time,end_time,str(recdir/'meta.json'),recording_id)); con.commit(); con.close()
    return meta

@app.get('/')
def home(): return FileResponse(STATIC/'index.html')

@app.get('/api/trains')
def trains():
    con=db(); rows=con.execute('''SELECT t.code,t.display_name,COUNT(r.id) recordings,MAX(r.end_time) last_time,SUM(CASE WHEN r.status='ready' THEN 1 ELSE 0 END) ready_count FROM trains t LEFT JOIN recordings r ON r.train_code=t.code GROUP BY t.code,t.display_name ORDER BY CAST(t.code AS INTEGER),t.code''').fetchall(); con.close(); return [dict(r) for r in rows]

@app.get('/api/trains/{train_code}/recordings')
def recordings(train_code:str):
    con=db(); rows=con.execute('SELECT id,train_code,original_name,rows_count,start_time,end_time,status,created_at,error_text FROM recordings WHERE train_code=? ORDER BY COALESCE(start_time,created_at) DESC',(train_code,)).fetchall(); con.close(); return [dict(r) for r in rows]

@app.post('/api/upload')
async def upload(train_code:str=Form(...), file:UploadFile=File(...)):
    if not file.filename.lower().endswith('.xlsx'): raise HTTPException(400,'Solo se aceptan archivos .xlsx')
    train_code=train_code.strip(); ensure_train(train_code)
    rec_id=uuid.uuid4().hex[:12]; stored=UPLOADS/f'{rec_id}_{Path(file.filename).name}'
    with stored.open('wb') as out: shutil.copyfileobj(file.file,out)
    con=db(); con.execute('INSERT INTO recordings(id,train_code,original_name,stored_xlsx,status,created_at) VALUES(?,?,?,?,?,?)',(rec_id,train_code,file.filename,str(stored),'processing',datetime.now().isoformat(timespec='seconds'))); con.commit(); con.close()
    try: return {'ok':True,'recording':process_excel(stored,rec_id,train_code,file.filename)}
    except Exception as e:
        con=db(); con.execute("UPDATE recordings SET status='error',error_text=? WHERE id=?",(str(e),rec_id)); con.commit(); con.close(); raise HTTPException(400,str(e))

@app.get('/api/recordings/{recording_id}/meta')
def recording_meta(recording_id:str):
    p=STORAGE/recording_id/'meta.json'
    if not p.exists(): raise HTTPException(404,'Registro no encontrado')
    return json.loads(p.read_text(encoding='utf-8'))

@app.get('/api/recordings/{recording_id}/chunk/{chunk_name}')
def recording_chunk(recording_id:str,chunk_name:str):
    if '/' in chunk_name or '\\' in chunk_name: raise HTTPException(400,'Nombre inválido')
    p=STORAGE/recording_id/chunk_name
    if not p.exists(): raise HTTPException(404,'Bloque no encontrado')
    return FileResponse(p,media_type='application/json')
