from datetime import date, datetime, timedelta
from calendar import monthrange
import sqlite3
from db import DB_PATH

ESTADOS_LIVRES = {'Cancelada', 'Faltou', 'Cancelado'}
ESTADOS_ATIVOS = {'Marcada', 'Marcado', 'Confirmada'}


def _conn():
    conn = sqlite3.connect(DB_PATH)
    conn.row_factory = sqlite3.Row
    return conn


def _row(row):
    return dict(row) if row else None


def _rows(rows):
    return [dict(r) for r in rows]


def _ensure_cols():
    with _conn() as conn:
        cols = {r['name'] for r in conn.execute('PRAGMA table_info(marcacoes)').fetchall()}
        if 'hora_fim' not in cols:
            conn.execute("ALTER TABLE marcacoes ADD COLUMN hora_fim TEXT")
        if 'remarcada_de' not in cols:
            conn.execute("ALTER TABLE marcacoes ADD COLUMN remarcada_de INTEGER")
        conn.commit()


def _add_minutes(hora, mins=30):
    try:
        d = datetime.strptime(hora, '%H:%M') + timedelta(minutes=mins)
        return d.strftime('%H:%M')
    except Exception:
        return hora


def _normalizar_estado(estado):
    mapa = {'Marcado': 'Marcada', 'Concluido': 'Realizada', 'Concluído': 'Realizada', 'Cancelado': 'Cancelada'}
    return mapa.get(estado or 'Marcada', estado or 'Marcada')


def _consulta_select(where='', params=()):
    _ensure_cols()
    sql = f'''
        SELECT
            m.id,
            m.utente_id,
            u.empresa_id,
            m.medico_id,
            u.nome AS utente_nome,
            COALESCE(e.nome,'') AS empresa_nome,
            COALESCE(med.nome,'') AS medico_nome,
            COALESCE(m.tipo_exame,'Medicina do Trabalho') AS tipo,
            m.data_marcacao AS data,
            m.hora_marcacao AS hora_inicio,
            COALESCE(NULLIF(m.hora_fim,''), time(m.hora_marcacao, '+30 minutes')) AS hora_fim,
            m.estado,
            COALESCE(m.observacoes,'') AS observacoes,
            m.criado_em AS created_at,
            m.remarcada_de
        FROM marcacoes m
        INNER JOIN utentes u ON u.id = m.utente_id
        LEFT JOIN empresas e ON e.id = u.empresa_id
        LEFT JOIN utilizadores med ON med.id = m.medico_id
        {where}
    '''
    with _conn() as conn:
        dados = _rows(conn.execute(sql, params).fetchall())
    for c in dados:
        c['estado'] = _normalizar_estado(c.get('estado'))
        if c.get('hora_fim') and len(c['hora_fim']) >= 5:
            c['hora_fim'] = c['hora_fim'][:5]
    return dados


def _add_month(data_iso):
    d = datetime.strptime(data_iso, '%Y-%m-%d').date()
    month = d.month + 1
    year = d.year
    if month == 13:
        month = 1
        year += 1
    day = min(d.day, monthrange(year, month)[1])
    return date(year, month, day).isoformat()


def _medico_por_id_ou_nome(medico_id=None, medico_nome=''):
    with _conn() as conn:
        if medico_id:
            row = conn.execute("SELECT id, nome FROM utilizadores WHERE id=? AND ativo=1 AND cargo='Medico'", (medico_id,)).fetchone()
            if row:
                return dict(row)
        nome = (medico_nome or '').strip()
        if nome:
            row = conn.execute("SELECT id, nome FROM utilizadores WHERE ativo=1 AND cargo='Medico' AND LOWER(nome)=LOWER(?)", (nome,)).fetchone()
            if row:
                return dict(row)
    return None


def _utente_por_id(utente_id):
    if not utente_id:
        return None
    with _conn() as conn:
        row = conn.execute('''
            SELECT u.id, u.nome, u.empresa_id, COALESCE(e.nome,'') AS empresa_nome
            FROM utentes u
            LEFT JOIN empresas e ON e.id = u.empresa_id
            WHERE u.id=? AND u.ativo=1
        ''', (utente_id,)).fetchone()
        return dict(row) if row else None


def listar_consultas():
    return _consulta_select('ORDER BY m.data_marcacao, m.hora_marcacao')


def obter_consulta(consulta_id):
    dados = _consulta_select('WHERE m.id=? LIMIT 1', (consulta_id,))
    return dados[0] if dados else None


def conflito_ativo(medico_id, data, hora_inicio, hora_fim, ignorar_id=None):
    _ensure_cols()
    sql = '''
        SELECT m.id
        FROM marcacoes m
        WHERE m.medico_id=?
          AND m.data_marcacao=?
          AND m.estado IN ('Marcada','Marcado','Confirmada')
          AND NOT (COALESCE(NULLIF(m.hora_fim,''), time(m.hora_marcacao, '+30 minutes')) <= ? OR m.hora_marcacao >= ?)
    '''
    params = [medico_id, data, hora_inicio, hora_fim]
    if ignorar_id:
        sql += ' AND m.id != ?'
        params.append(ignorar_id)
    sql += ' LIMIT 1'
    with _conn() as conn:
        row = conn.execute(sql, params).fetchone()
    return obter_consulta(row['id']) if row else None


def consulta_livre_no_horario(medico_id, data, hora_inicio, hora_fim, ignorar_id=None):
    _ensure_cols()
    sql = '''
        SELECT m.id
        FROM marcacoes m
        WHERE m.medico_id=?
          AND m.data_marcacao=?
          AND m.estado IN ('Cancelada','Cancelado','Faltou')
          AND NOT (COALESCE(NULLIF(m.hora_fim,''), time(m.hora_marcacao, '+30 minutes')) <= ? OR m.hora_marcacao >= ?)
    '''
    params = [medico_id, data, hora_inicio, hora_fim]
    if ignorar_id:
        sql += ' AND m.id != ?'
        params.append(ignorar_id)
    sql += ' ORDER BY m.id DESC LIMIT 1'
    with _conn() as conn:
        row = conn.execute(sql, params).fetchone()
    return obter_consulta(row['id']) if row else None


def _normalizar_data(valor):
    valor = (valor or '').strip()
    if not valor:
        return ''
    for fmt in ('%Y-%m-%d', '%d/%m/%Y'):
        try:
            return datetime.strptime(valor, fmt).date().isoformat()
        except Exception:
            pass
    return valor


def _normalizar_hora(valor):
    valor = (valor or '').strip()
    if len(valor) >= 5:
        return valor[:5]
    return valor


def _criar_ou_obter_utente(nome, telefone='', empresa_id=None):
    nome = (nome or '').strip()
    if not nome:
        return None
    telefone = (telefone or '').strip()
    empresa_id = empresa_id or None
    with _conn() as conn:
        row = conn.execute('''
            SELECT u.id, u.nome, u.empresa_id, COALESCE(e.nome,'') AS empresa_nome
            FROM utentes u
            LEFT JOIN empresas e ON e.id = u.empresa_id
            WHERE u.ativo=1
              AND LOWER(u.nome)=LOWER(?)
              AND (COALESCE(u.telefone,'')=? OR ?='')
            ORDER BY CASE WHEN u.empresa_id IS ? THEN 0 ELSE 1 END, u.id DESC
            LIMIT 1
        ''', (nome, telefone, telefone, empresa_id)).fetchone()
        if row:
            return dict(row)
        cur = conn.execute('''
            INSERT INTO utentes (nome, telefone, empresa_id, ativo)
            VALUES (?, ?, ?, 1)
        ''', (nome, telefone, empresa_id))
        conn.commit()
        row = conn.execute('''
            SELECT u.id, u.nome, u.empresa_id, COALESCE(e.nome,'') AS empresa_nome
            FROM utentes u
            LEFT JOIN empresas e ON e.id = u.empresa_id
            WHERE u.id=?
        ''', (cur.lastrowid,)).fetchone()
        return dict(row) if row else None


def _validar_data(data):
    clean = dict(data or {})
    clean['data'] = _normalizar_data(clean.get('data'))
    clean['hora_inicio'] = _normalizar_hora(clean.get('hora_inicio'))
    clean['hora_fim'] = _normalizar_hora(clean.get('hora_fim')) or _add_minutes(clean.get('hora_inicio','09:00'), 30)

    if not clean.get('data') or not clean.get('hora_inicio') or not clean.get('hora_fim'):
        return None, {'ok': False, 'tipo': 'dados_invalidos', 'mensagem': 'Data, hora de início e hora de fim são obrigatórias.'}
    if clean['hora_inicio'] >= clean['hora_fim']:
        return None, {'ok': False, 'tipo': 'hora_invalida', 'mensagem': 'A hora de fim tem de ser depois da hora de início.'}

    medico = _medico_por_id_ou_nome(clean.get('medico_id'), clean.get('medico_nome',''))
    if not medico:
        return None, {'ok': False, 'tipo': 'medico_invalido', 'mensagem': 'Escolha um médico criado ou escreva exatamente o nome de um médico já registado.'}

    utente = _utente_por_id(clean.get('utente_id'))
    if not utente:
        utente = _criar_ou_obter_utente(clean.get('utente_nome'), clean.get('telefone',''), clean.get('empresa_id'))
    if not utente:
        return None, {'ok': False, 'tipo': 'utente_invalido', 'mensagem': 'Escolha um utente guardado ou escreva o nome do utente.'}

    clean['medico_id'] = medico['id']
    clean['medico_nome'] = medico['nome']
    clean['utente_id'] = utente['id']
    clean['utente_nome'] = utente['nome']
    clean['empresa_id'] = utente.get('empresa_id')
    clean['empresa_nome'] = utente.get('empresa_nome') or clean.get('empresa_nome') or ''
    clean['estado'] = _normalizar_estado(clean.get('estado'))
    return clean, None


def criar_consulta(data, forcar=False):
    data, erro = _validar_data(data)
    if erro:
        return erro
    conflito = conflito_ativo(data['medico_id'], data['data'], data['hora_inicio'], data['hora_fim'])
    if conflito and not forcar:
        return {'ok': False, 'tipo': 'conflito_ativo', 'consulta': conflito}
    livre = consulta_livre_no_horario(data['medico_id'], data['data'], data['hora_inicio'], data['hora_fim'])
    if livre and not forcar:
        return {'ok': False, 'tipo': 'horario_livre_cancelado', 'consulta': livre}
    _ensure_cols()
    with _conn() as conn:
        cur = conn.execute('''
            INSERT INTO marcacoes (utente_id, medico_id, data_marcacao, hora_marcacao, hora_fim, tipo_exame, estado, observacoes, remarcada_de)
            VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
        ''', (
            data.get('utente_id'), data.get('medico_id'), data['data'], data['hora_inicio'], data['hora_fim'],
            data.get('tipo','Medicina do Trabalho'), data.get('estado','Marcada'), data.get('observacoes',''), data.get('remarcada_de')
        ))
        conn.commit()
        return {'ok': True, 'consulta': obter_consulta(cur.lastrowid)}


def atualizar_consulta(consulta_id, data):
    atual = obter_consulta(consulta_id)
    if not atual:
        return {'ok': False, 'tipo': 'nao_encontrada'}
    data, erro = _validar_data(data)
    if erro:
        return erro
    conflito = conflito_ativo(data['medico_id'], data['data'], data['hora_inicio'], data['hora_fim'], consulta_id)
    if conflito:
        return {'ok': False, 'tipo': 'conflito_ativo', 'consulta': conflito}
    _ensure_cols()
    with _conn() as conn:
        conn.execute('''
            UPDATE marcacoes
            SET utente_id=?, medico_id=?, data_marcacao=?, hora_marcacao=?, hora_fim=?, tipo_exame=?, estado=?, observacoes=?
            WHERE id=?
        ''', (
            data.get('utente_id'), data.get('medico_id'), data['data'], data['hora_inicio'], data['hora_fim'],
            data.get('tipo','Medicina do Trabalho'), data.get('estado','Marcada'), data.get('observacoes',''), consulta_id
        ))
        conn.commit()
    nova = obter_consulta(consulta_id)
    aviso = nova['estado'] in {'Cancelada','Faltou'} and atual['estado'] != nova['estado']
    return {'ok': True, 'consulta': nova, 'aviso_livre': aviso}


def apagar_consulta(consulta_id):
    with _conn() as conn:
        conn.execute('DELETE FROM marcacoes WHERE id=?', (consulta_id,))
        conn.commit()
    return {'ok': True}


def remarcar_proximo_mes(consulta_id):
    c = obter_consulta(consulta_id)
    if not c:
        return {'ok': False, 'tipo': 'nao_encontrada'}
    nova_data = dict(c)
    nova_data['data'] = _add_month(c['data'])
    nova_data['estado'] = 'Marcada'
    nova_data['observacoes'] = (c.get('observacoes') or '') + f"\nRemarcada automaticamente a partir da consulta #{consulta_id}."
    nova_data['remarcada_de'] = consulta_id
    return criar_consulta(nova_data, forcar=False)


def estatisticas():
    hoje = date.today().isoformat()
    inicio_mes = date.today().replace(day=1).isoformat()
    with _conn() as conn:
        total_mes = conn.execute('SELECT COUNT(*) FROM marcacoes WHERE data_marcacao >= ?', (inicio_mes,)).fetchone()[0]
        hoje_count = conn.execute('SELECT COUNT(*) FROM marcacoes WHERE data_marcacao=?', (hoje,)).fetchone()[0]
        confirmadas = conn.execute("SELECT COUNT(*) FROM marcacoes WHERE estado='Confirmada' AND data_marcacao >= ?", (inicio_mes,)).fetchone()[0]
        faltas = conn.execute("SELECT COUNT(*) FROM marcacoes WHERE estado='Faltou' AND data_marcacao >= ?", (inicio_mes,)).fetchone()[0]
        canceladas = conn.execute("SELECT COUNT(*) FROM marcacoes WHERE estado IN ('Cancelada','Cancelado') AND data_marcacao >= ?", (inicio_mes,)).fetchone()[0]
        realizadas = conn.execute("SELECT COUNT(*) FROM marcacoes WHERE estado IN ('Realizada','Concluido','Concluído') AND data_marcacao >= ?", (inicio_mes,)).fetchone()[0]
        medicos = conn.execute("SELECT COUNT(*) FROM utilizadores WHERE ativo=1 AND cargo='Medico'").fetchone()[0]
    comparencia_base = realizadas + faltas
    taxa = round((realizadas / comparencia_base) * 100) if comparencia_base else 100
    return {'total_mes': total_mes, 'hoje': hoje_count, 'confirmadas': confirmadas, 'faltas': faltas, 'canceladas': canceladas, 'realizadas': realizadas, 'medicos': medicos, 'taxa_comparencia': taxa}


def painel_hoje():
    hoje = date.today().isoformat()
    return _consulta_select('WHERE m.data_marcacao=? ORDER BY m.hora_marcacao', (hoje,))


def proximos_lembretes():
    alvo = (date.today() + timedelta(days=1)).isoformat()
    return _consulta_select("WHERE m.data_marcacao=? AND m.estado IN ('Marcada','Marcado','Confirmada') ORDER BY m.hora_marcacao", (alvo,))


def historico_utente(nome):
    return _consulta_select('WHERE u.nome LIKE ? ORDER BY m.data_marcacao DESC, m.hora_marcacao DESC', (f'%{nome}%',))


def consumos_da_consulta(consulta_id):
    from agenda_mod.database import get_connection, rows_to_list
    with get_connection() as conn:
        return rows_to_list(conn.execute('SELECT * FROM consumos_consulta WHERE consulta_id=? ORDER BY id DESC', (consulta_id,)).fetchall())


def adicionar_consumo(consulta_id, produto, quantidade):
    from agenda_mod.database import get_connection
    with get_connection() as conn:
        conn.execute('INSERT INTO consumos_consulta (consulta_id, produto, quantidade) VALUES (?,?,?)', (consulta_id, produto, quantidade))
        conn.commit()
    return {'ok': True, 'consumos': consumos_da_consulta(consulta_id)}
