from flask import Blueprint, request, render_template, redirect, url_for, flash, current_app
from flask_login import login_required, current_user
from datetime import datetime, timedelta, date
from sqlalchemy.exc import IntegrityError
from extensions import db
from models import Employee, LeaveType, LeaveBalance, LeaveRequest, LeaveApprovalLog, RHConfig, LeavePolicy, Role, Department, Holiday, LeaveSettings, CompOffRequest, LeaveAllocationBatch, LeaveAllocationEntry
from utils.leave_utils import calculate_sandwich_days
from utils.attendance_utils import sync_leave_to_attendance
from flask import url_for
from utils.notify import notify_employee, notify_many
from typing import Optional
from sqlalchemy import and_, or_


# --- Notification helpers ---

from urllib.parse import urljoin

def _link(endpoint: str, **kwargs) -> str:
    """
    Build public, absolute links for notifications.
    Prefer EXTERNAL_BASE_URL if configured; else fall back to _external=True
    (which relies on ProxyFix + forwarded headers).
    """
    base = current_app.config.get("EXTERNAL_BASE_URL")
    if base:
        path = url_for(endpoint, _external=False, **kwargs)
        return urljoin(base, path)
    return url_for(endpoint, _external=True, **kwargs)

def _format_range(d1, d2):
    if d1 == d2:
        return d1.strftime('%d-%b-%Y')
    return f"{d1.strftime('%d-%b-%Y')} → {d2.strftime('%d-%b-%Y')}"

def _hr_ids():
    """
    Try to find HR recipients.
    - First by role.name == 'HR'
    - Fallback by department.name contains 'HR'
    - If none, return [] (no HR ping)
    """
    ids = []
    try:
        # Role-based
        ids = [e.id for e in Employee.query.join(Role).filter(Role.name.ilike('hr')).all()]
        if ids:
            return ids
    except Exception:
        pass

    try:
        # Department-based fallback
        ids = [e.id for e in Employee.query.join(Department).filter(Department.name.ilike('%hr%')).all()]
        return ids
    except Exception:
        return []


leave_bp = Blueprint('leave', __name__, url_prefix='/leave')

from sqlalchemy import and_

def get_leave_year_range():
    settings = LeaveSettings.query.first()
    if settings and settings.leave_year_start and settings.leave_year_end:
        return settings.leave_year_start, settings.leave_year_end
    # fallback: financial year Apr-Mar if not configured
    today = datetime.today().date()
    fy_start = date(today.year if today.month >= 4 else today.year - 1, 4, 1)
    fy_end   = date(fy_start.year + 1, 3, 31)
    return fy_start, fy_end

def get_optional_holiday_dates(start_dt, end_dt, employee=None):
    # fetch all holidays in window once
    hols = (Holiday.query
        .filter(Holiday.date >= start_dt, Holiday.date <= end_dt)
        .all())

    by_date = {}
    for h in hols:
        by_date.setdefault(h.date, []).append(h)

    whitelist = set()
    emp_loc = getattr(employee, 'location', None)

    for dt, items in by_date.items():
        # Prefer specific location if it exists
        specific = next((x for x in items if emp_loc and x.location == emp_loc), None)
        if specific:
            # specific decides RH eligibility
            if (specific.type or '').lower() == 'optional':
                whitelist.add(dt)
            continue  # if specific is mandatory, NOT eligible; skip

        # No specific entry -> fall back to ALL
        all_row = next((x for x in items if x.location == 'ALL'), None)
        if all_row and (all_row.type or '').lower() == 'optional':
            whitelist.add(dt)

    return whitelist
    
def get_applicable_holiday_for(employee, dt):
    q = Holiday.query.filter_by(date=dt)
    emp_loc = getattr(employee, 'location', None)
    if emp_loc:
        h = q.filter_by(location=emp_loc).first()
        if h:
            return h
    return q.filter_by(location='ALL').first()

from typing import Optional

def _normalize_holiday_type(raw: str) -> Optional[str]:
    """
    Map various inputs to canonical {'mandatory','optional'}.
    Treat 'National Holiday' as mandatory (non-RH).
    """
    if not raw:
        return None
    t = raw.strip().lower()

    mandatory_aliases = {
        'mandatory', 'national', 'national holiday', 'gazetted',
        'public', 'public holiday', 'closed'
    }
    optional_aliases = {'optional', 'restricted', 'restricted holiday', 'rh'}

    if t in mandatory_aliases:
        return 'mandatory'
    if t in optional_aliases:
        return 'optional'
    return None

# -----------------------------
# Allocation helpers (freeze + delta-based)
# -----------------------------
from datetime import datetime, date
from sqlalchemy import func
from typing import Optional
from models import LeaveAllocationBatch, LeaveAllocationEntry

def _make_batch_key(fy_label:int, leave_type_code:str, strategy:str) -> str:
    ts = datetime.now().strftime('%Y%m%d_%H%M%S')
    return f"FY{fy_label}_{leave_type_code}_{strategy}_{ts}"

def _latest_non_revoked_batch(leave_type_id:int, fy_label:int, strategy:str):
    return (LeaveAllocationBatch.query
            .filter_by(leave_type_id=leave_type_id, fy_label=fy_label, strategy=strategy, is_revoked=False)
            .order_by(LeaveAllocationBatch.run_at.desc())
            .first())

def fy_bounds_for_today(today: date):
    """Return (fy_start_date, fy_end_date) where FY is Apr–Mar containing 'today'."""
    fy_start_year = today.year if today.month >= 4 else today.year - 1
    fy_start = date(fy_start_year, 4, 1)
    fy_end   = date(fy_start_year + 1, 3, 31)
    return fy_start, fy_end

def months_in_fy_until(today: date):
    """List of (year, month) from Apr..today’s month within the same FY."""
    fy_start, _ = fy_bounds_for_today(today)
    months = [(fy_start.year, m) for m in range(4, 13)]           # Apr..Dec
    months += [(fy_start.year + 1, m) for m in range(1, 4)]       # Jan..Mar
    cut = []
    for (y, m) in months:
        cut.append((y, m))
        if y == today.year and m == today.month:
            break
    return cut

def months_after_freeze_until(today: date, freeze_year:int, freeze_month:int):
    """Months strictly AFTER (freeze_year, freeze_month) up to today's month within same FY."""
    all_months = months_in_fy_until(today)
    return [(y, m) for (y, m) in all_months if (y, m) > (freeze_year, freeze_month)]

def calc_chl_earned_leave_total_after_freeze(doj: Optional[date], today: date, freeze_year:int, freeze_month:int) -> float:
    """
    CHL EL rule, but only counting months strictly AFTER the freeze month.
    - If DOJ <= FY start: Apr=3, May=2, Jun–current=1 each (but only those months > freeze cut)
    - If DOJ > FY start: 1 per month from max(DOJ month, freeze+1) up to current
    """
    fy_start, _ = fy_bounds_for_today(today)
    months = months_after_freeze_until(today, freeze_year, freeze_month)

    if doj is None:
        doj = date(1900, 1, 1)

    total = 0.0
    if doj <= fy_start:
        for (_, m) in months:
            if m == 4:   total += 3
            elif m == 5: total += 2
            else:        total += 1
        return total

    # New in FY: start from max(DOJ, freeze+1)
    for (y, m) in months:
        if (y, m) >= (doj.year, doj.month):
            total += 1
    return total

def _sum_credited_since(emp_id:int, lt_id:int, since_dt:datetime) -> float:
    """Sum of non-revoked credit deltas for this employee/type credited on/after since_dt."""
    q = (db.session.query(db.func.coalesce(db.func.sum(LeaveAllocationEntry.credit_delta), 0.0))
         .join(LeaveAllocationBatch, LeaveAllocationBatch.id == LeaveAllocationEntry.batch_id)
         .filter(
             LeaveAllocationEntry.employee_id == emp_id,
             LeaveAllocationEntry.leave_type_id == lt_id,
             LeaveAllocationBatch.is_revoked == False,
             LeaveAllocationBatch.run_at >= since_dt
         ))
    return float(q.scalar() or 0.0)
# --------------------------------------------------------------------------------------------------------------------

# -----------------------------
# LEAVE TYPE MASTER
# -----------------------------

@leave_bp.route('/leave-types')
@login_required
def view_leave_types():
    leave_types = LeaveType.query.order_by(LeaveType.id).all()
    return render_template('leave/leave_types.html', leave_types=leave_types)

@leave_bp.route('/leave-types/add', methods=['POST'])
@login_required
def add_leave_type():
    form = request.form
    lt = LeaveType(
        name=form['name'],
        code=form['code'].upper(),
        default_quota=float(form.get('default_quota', 0)),
        allow_half_day='allow_half_day' in form,
        requires_approval='requires_approval' in form,
        carry_forward='carry_forward' in form,
        encashable='encashable' in form,
        sandwich_applicable='sandwich_applicable' in form,
        probation_restricted='probation_restricted' in form,
        remarks=form.get('remarks', '')
    )
    db.session.add(lt)
    db.session.commit()
    flash('✅ Leave type added.', 'success')
    return redirect(url_for('leave.view_leave_types'))

@leave_bp.route('/leave-types/edit/<int:type_id>', methods=['GET', 'POST'])
@login_required
def edit_leave_type(type_id):
    lt = LeaveType.query.get_or_404(type_id)
    if request.method == 'POST':
        form = request.form
        lt.name = form['name']
        lt.code = form['code'].upper()
        lt.default_quota = float(form.get('default_quota', 0))
        lt.allow_half_day = 'allow_half_day' in form
        lt.requires_approval = 'requires_approval' in form
        lt.carry_forward = 'carry_forward' in form
        lt.encashable = 'encashable' in form
        lt.sandwich_applicable = 'sandwich_applicable' in form
        lt.probation_restricted = 'probation_restricted' in form
        lt.remarks = form.get('remarks', '')
        db.session.commit()
        flash('✅ Leave type updated.', 'success')
        return redirect(url_for('leave.view_leave_types'))
    return render_template('leave/edit_leave_type.html', lt=lt)

@leave_bp.route('/leave-types/delete/<int:type_id>', methods=['POST'])
@login_required
def delete_leave_type(type_id):
    lt = LeaveType.query.get_or_404(type_id)
    db.session.delete(lt)
    db.session.commit()
    flash('✅ Leave type deleted.', 'success')
    return redirect(url_for('leave.view_leave_types'))



@leave_bp.route('/settings', methods=['GET', 'POST'])
@login_required
def leave_settings():
    settings = LeaveSettings.query.first()
    if not settings:
        settings = LeaveSettings(sandwich_policy_enabled=True, rh_max_per_year=5)
        db.session.add(settings)
        db.session.commit()

    if request.method == 'POST':
        settings.sandwich_policy_enabled = 'sandwich_policy_enabled' in request.form
        settings.rh_max_per_year = int(request.form['rh_max_per_year'])
        db.session.commit()
        flash("✅ Settings updated successfully.", "success")
        return redirect(url_for('leave.leave_settings'))

    return render_template('leave/leave_settings.html', settings=settings)

# -----------------------------
# LEAVE POLICY CONFIGURATION
# -----------------------------

@leave_bp.route('/leave-policies')
@login_required
def view_leave_policies():
    leave_types = LeaveType.query.order_by(LeaveType.name).all()
    policies = LeavePolicy.query.all()
    roles = Role.query.order_by(Role.name).all()
    departments = Department.query.order_by(Department.name).all()

    return render_template('leave/leave_policies.html',
        policies=policies,
        leave_types=leave_types,
        roles=roles,
        departments=departments
    )

@leave_bp.route('/leave-policies/add', methods=['POST'])
@login_required
def add_leave_policy():
    form = request.form
    lp = LeavePolicy(
        company_id=None,
        leave_type_id=form['leave_type_id'],
        role_id=form.get('role_id') or None,
        department_id=form.get('department_id') or None,
        annual_quota=form.get('annual_quota') or 0,
        monthly_credit='monthly_credit' in form,
        max_consecutive_days=form.get('max_consecutive_days') or 30,
        join_after_cutoff_month=form.get('join_after_cutoff_month') or None,
        rh_consecutive_limit=form.get('rh_consecutive_limit') or None
    )
    db.session.add(lp)
    db.session.commit()
    flash('✅ Leave policy added.', 'success')
    return redirect(url_for('leave.view_leave_policies'))

@leave_bp.route('/leave-policies/delete/<int:policy_id>', methods=['POST'])
@login_required
def delete_leave_policy(policy_id):
    policy = LeavePolicy.query.get_or_404(policy_id)
    db.session.delete(policy)
    db.session.commit()
    flash('✅ Policy deleted.', 'success')
    return redirect(url_for('leave.view_leave_policies'))


# -----------------------------
# HOLIDAY MASTER + RH CONFIG
# -----------------------------

@leave_bp.route('/holidays')
@login_required
def view_holidays():
    holidays = Holiday.query.order_by(Holiday.date).all()
    settings = LeaveSettings.query.all()
    return render_template('leave/holidays.html', holidays=holidays, settings=settings)

@leave_bp.route('/holidays/add', methods=['POST'])
@login_required
def add_holiday():
    form = request.form
    # Validate and normalize inputs
    try:
        date_ = datetime.strptime(form['date'], '%Y-%m-%d').date()
    except Exception:
        flash('❌ Invalid date format.', 'danger')
        return redirect(url_for('leave.view_holidays'))

    name = (form.get('name') or '').strip()
    raw_type = (form.get('type') or '').strip()
    canonical_type = _normalize_holiday_type(raw_type)   # <— normalize here

    # Location handling (keep your existing logic)
    raw_loc = (form.get('location') or '').strip()
    if not raw_loc or raw_loc.lower() in {'other', 'others', 'other location', 'others location', '-', 'all'}:
        location = 'ALL'
    else:
        location = raw_loc

    if not name:
        flash('❌ Holiday name is required.', 'danger')
        return redirect(url_for('leave.view_holidays'))

    if not canonical_type:
        flash('❌ Invalid holiday type.', 'danger')
        return redirect(url_for('leave.view_holidays'))

    # Friendly pre-check (matches composite unique key)
    exists = Holiday.query.filter(
        and_(Holiday.date == date_, Holiday.location == location)
    ).first()
    if exists:
        flash(f"⚠️ A holiday already exists on {date_.isoformat()} for '{location}'.", 'warning')
        return redirect(url_for('leave.view_holidays'))

    # Save canonical type: 'mandatory' or 'optional'
    h = Holiday(date=date_, name=name, type=canonical_type, location=location)
    db.session.add(h)
    try:
        db.session.commit()
        flash('✅ Holiday added.', 'success')
    except IntegrityError:
        db.session.rollback()
        clash_locs = [x.location for x in Holiday.query.filter(Holiday.date == date_).all()]
        where = ", ".join(sorted(set(clash_locs))) or "unknown"
        flash(f"❌ Conflict: a holiday already exists on {date_.isoformat()} for [{where}].", 'danger')

    return redirect(url_for('leave.view_holidays'))

@leave_bp.route('/holidays/delete/<int:holiday_id>', methods=['POST'])
@login_required
def delete_holiday(holiday_id):
    h = Holiday.query.get_or_404(holiday_id)
    db.session.delete(h)
    db.session.commit()
    flash('✅ Holiday deleted.', 'success')
    return redirect(url_for('leave.view_holidays'))

@leave_bp.route('/rh-config/update', methods=['POST'])
@login_required
def update_rh_config():
    rh_max = int(request.form['rh_max_per_year'])
    sandwich_enabled = 'sandwich_policy_enabled' in request.form

    setting = LeaveSettings.query.filter_by(company_id=None).first()
    if not setting:
        setting = LeaveSettings(company_id=None)

    setting.rh_max_per_year = rh_max
    setting.sandwich_policy_enabled = sandwich_enabled

    db.session.add(setting)
    db.session.commit()
    flash('✅ RH config and sandwich policy updated!', 'success')
    return redirect(url_for('leave.view_holidays'))

@leave_bp.route('/apply', methods=['GET', 'POST'])
@login_required
def apply_leave():
    today = datetime.today()
    year = today.year

    leave_types = LeaveType.query.order_by(LeaveType.name).all()
    balances = LeaveBalance.query.filter_by(employee_id=current_user.id, year=year).all()
    balance_map = {b.leave_type_id: b.balance for b in balances}

    # holidays for sandwich warning only (doesn't affect final calc)
    holidays = Holiday.query.with_entities(Holiday.date).all()
    holiday_dates = {h.date for h in holidays}

    settings = LeaveSettings.query.first()
    max_rh = settings.rh_max_per_year if settings else 5

    # Use leave-year window for RH counters
    ly_start, ly_end = get_leave_year_range()

    # Count approved RH (final) as DAYS (0.5 / 1.0); treat NULL total_days as 1.0 for legacy rows
    rh_approved_days = (
        db.session.query(func.coalesce(func.sum(func.coalesce(LeaveRequest.total_days, 1.0)), 0.0))
        .filter(
            LeaveRequest.employee_id == current_user.id,
            LeaveRequest.leave_type.has(code='RH'),
            LeaveRequest.start_date >= ly_start,
            LeaveRequest.end_date <= ly_end,
            LeaveRequest.status == 'final_approved'
        )
        .scalar()
    )

    # In-flight RH (pending + manager_approved + final_approved) as DAYS
    in_flight_statuses = ['pending', 'manager_approved', 'final_approved']
    rh_in_flight_days = (
        db.session.query(func.coalesce(func.sum(func.coalesce(LeaveRequest.total_days, 1.0)), 0.0))
        .filter(
            LeaveRequest.employee_id == current_user.id,
            LeaveRequest.leave_type.has(code='RH'),
            LeaveRequest.start_date >= ly_start,
            LeaveRequest.end_date <= ly_end,
            LeaveRequest.status.in_(in_flight_statuses)
        )
        .scalar()
    )

    # Conservative remaining, from the in-flight tally
    rh_remaining = max(0.0, float(max_rh) - float(rh_in_flight_days or 0.0))

    if request.method == 'POST':
        form = request.form
        try:
            leave_type_id = int(form['leave_type_id'])
            start_date = datetime.strptime(form['start_date'], '%Y-%m-%d').date()
            end_date = datetime.strptime(form['end_date'], '%Y-%m-%d').date()
        except Exception:
            flash("Invalid form input. Please check dates and leave type.", "danger")
            return redirect(url_for('leave.apply_leave'))

        if end_date < start_date:
            flash("End date cannot be before start date.", "danger")
            return redirect(url_for('leave.apply_leave'))

        reason = form['reason'].strip()
        half_day = 'half_day' in form

        leave_type = LeaveType.query.get(leave_type_id)
        if not leave_type:
            flash("Invalid leave type selected.", "danger")
            return redirect(url_for('leave.apply_leave'))

        # ---------- RH VALIDATION (STRICT) ----------

        if leave_type.code == 'RH':
            # RH must be a single calendar date (but can be 0.5 or 1.0 day)
            if end_date != start_date:
                flash("❌ RH must be a single-day leave. Please select exactly one Optional Holiday date.", "danger")
                return redirect(url_for('leave.apply_leave'))

            # Eligible date must be an Optional Holiday for the employee's location
            rh_whitelist = get_optional_holiday_dates(ly_start, ly_end, employee=current_user)
            if start_date not in rh_whitelist:
                flash("❌ Only Optional (RH) holidays can be selected for RH.", "danger")
                return redirect(url_for('leave.apply_leave'))

            # Requested units for cap (0.5 if half-day, else 1.0)
            requested_units = 0.5 if half_day else 1.0

            # Cap check in DAYS: in-flight-days + requested_units <= max_rh
            if (float(rh_in_flight_days or 0.0) + requested_units) > float(max_rh):
                remaining = max(0.0, float(max_rh) - float(rh_in_flight_days or 0.0))
                flash(f"❌ RH limit exceeded. Remaining RH this year: {remaining} day(s).", "danger")
                return redirect(url_for('leave.apply_leave'))

            # For RH, set total_days based on half_day, then continue to common flow
            total_days = requested_units
        else:
            # Non-RH: compute total_days as before
            if half_day:
                if end_date != start_date:
                    flash("❌ Half-day leave can only be applied for a single day.", "danger")
                    return redirect(url_for('leave.apply_leave'))
                total_days = 0.5
            else:
                total_days = (end_date - start_date).days + 1

        # ---------- END RH VALIDATION ----------

        # Compute total_days
        if half_day:
            if start_date != end_date:
                flash("❌ Half-day leave can only be applied for a single day.", "danger")
                return redirect(url_for('leave.apply_leave'))
            total_days = 0.5
        else:
            # inclusive day count
            total_days = (end_date - start_date).days + 1

        # Optional: sandwich warning (display only)
        span_days = 1 if total_days == 0.5 else (end_date - start_date).days + 1
        date_set = {start_date + timedelta(days=i) for i in range(span_days)}
        sandwich_days = {d for d in date_set if d.weekday() in (5, 6) or d in holiday_dates}
        if sandwich_days:
            flash(f"⚠️ Sandwich leave: includes {len(sandwich_days)} weekend/holiday day(s).", "warning")

        # ------------------ BALANCE CHECK + SPLIT INTO LWP IF NEEDED ------------------
        def _create_lr(lt_id, d1, d2, days, half=False):
            lr = LeaveRequest(
                employee_id=current_user.id,
                leave_type_id=lt_id,
                start_date=d1,
                end_date=d2,
                total_days=days,
                reason=reason,
                half_day=half,
                status='pending',
            )
            db.session.add(lr)
            return lr

        created_reqs = []

        if leave_type.code in ('LWP', 'RH'):
            # RH already validated above; LWP never needs balance.
            lr = _create_lr(leave_type_id, start_date, end_date, total_days, (total_days == 0.5))
            created_reqs.append(lr)

        else:
            current_balance = float(balance_map.get(leave_type_id, 0) or 0.0)

            # HALF-DAY request
            if total_days == 0.5:
                if current_balance >= 0.5:
                    lr = _create_lr(leave_type_id, start_date, end_date, 0.5, True)
                    created_reqs.append(lr)
                else:
                    lwp = (LeaveType.query.filter(db.func.upper(LeaveType.code).in_(['LWP','LOP','UL'])).first())
                    if not lwp:
                        flash(f"Insufficient balance: {current_balance} day(s) left and no LWP fallback configured.", "danger")
                        return redirect(url_for('leave.apply_leave'))
                    lr = _create_lr(lwp.id, start_date, end_date, 0.5, True)
                    created_reqs.append(lr)
                    flash("⚠️ No balance left. Applied as LWP (0.5 day).", "warning")

            else:
                # MULTI-DAY (whole-day) request
                requested = int(total_days)  # inclusive whole days
                if current_balance >= requested:
                    lr = _create_lr(leave_type_id, start_date, end_date, requested, False)
                    created_reqs.append(lr)
                else:
                    # Split into paid portion + LWP portion
                    lwp = LeaveType.query.filter_by(code='LWP').first()
                    if not lwp:
                        flash(f"Insufficient balance: {current_balance} day(s) left and no LWP fallback configured.", "danger")
                        return redirect(url_for('leave.apply_leave'))

                    paid_days = int(max(0, min(requested, current_balance)))
                    lwp_days  = int(requested - paid_days)

                    # Part A: consume balance (if any)
                    if paid_days > 0:
                        partA_start = start_date
                        partA_end   = start_date + timedelta(days=paid_days - 1)
                        lrA = _create_lr(leave_type_id, partA_start, partA_end, paid_days, False)
                        created_reqs.append(lrA)

                    # Part B: leftover as LWP (if any)
                    if lwp_days > 0:
                        partB_start = (start_date + timedelta(days=paid_days)) if paid_days > 0 else start_date
                        partB_end   = end_date
                        lrB = _create_lr(lwp.id, partB_start, partB_end, lwp_days, False)
                        created_reqs.append(lrB)

                    flash(f"⚠️ Not enough balance. Split into {paid_days} day(s) {leave_type.code} + {lwp_days} day(s) LWP.", "warning")

        # Persist all created requests
        db.session.commit()

        # ------------------ NOTIFICATIONS ------------------
        def _rng(d1, d2):
            return d1.strftime('%d-%b-%Y') if d1 == d2 else f"{d1:%d-%b-%Y} → {d2:%d-%b-%Y}"

        if len(created_reqs) == 1:
            lr = created_reqs[0]
            rng = _rng(lr.start_date, lr.end_date)
            lt  = LeaveType.query.get(lr.leave_type_id)
            notify_employee(
                employee_id=current_user.id,
                title="Leave request submitted",
                message=f"Your {lt.name} request ({rng}) has been submitted and is pending manager approval.",
                link=_link('leave.view_history')
            )
            mgr_id = getattr(current_user, 'reporting_manager_id', None)
            if mgr_id:
                notify_employee(
                    employee_id=mgr_id,
                    title="Leave approval needed",
                    message=f"{current_user.full_name} applied for {lt.name} ({rng}). Please review.",
                    link=_link('leave.manager_pending')
                )
            hr_list = _hr_ids()
            if hr_list:
                notify_many(
                    employee_ids=hr_list,
                    title="New leave request (FYI)",
                    message=f"{current_user.full_name} applied for {lt.name} ({rng}). Awaiting manager action.",
                    link=_link('leave.hr_pending')
                )
        else:
            parts = []
            for r in created_reqs:
                lt = LeaveType.query.get(r.leave_type_id)
                parts.append(f"{lt.code} ({_rng(r.start_date, r.end_date)})")
            summary = " + ".join(parts)

            notify_employee(
                employee_id=current_user.id,
                title="Leave requests submitted (split)",
                message=f"Submitted: {summary}. Pending manager approval.",
                link=_link('leave.view_history')
            )
            mgr_id = getattr(current_user, 'reporting_manager_id', None)
            if mgr_id:
                notify_employee(
                    employee_id=mgr_id,
                    title="Leave approval needed",
                    message=f"{current_user.full_name} applied: {summary}. Please review.",
                    link=_link('leave.manager_pending')
                )
            hr_list = _hr_ids()
            if hr_list:
                notify_many(
                    employee_ids=hr_list,
                    title="New leave requests (FYI)",
                    message=f"{current_user.full_name} applied: {summary}. Awaiting manager action.",
                    link=_link('leave.hr_pending')
                )

        flash("✅ Leave request submitted.", "success")
        return redirect(url_for('leave.view_history'))

    rh_type = LeaveType.query.filter_by(code='RH').first()
    rh_type_id = rh_type.id if rh_type else None

    # GET
    return render_template(
        'leave/apply_leave.html',
        leave_types=leave_types,
        balances=balances,
        max_rh=max_rh,
        rh_remaining=rh_remaining,
        current_year=year,  # if your template shows this
    )


@leave_bp.route('/history')
@login_required
def view_history():
    leave_requests = LeaveRequest.query.filter_by(employee_id=current_user.id).order_by(LeaveRequest.applied_on.desc()).all()
    return render_template('leave/view_history.html', leave_requests=leave_requests)

# -----------------------------
# MANAGER PENDING + APPROVE
# -----------------------------
from calendar import monthrange
@leave_bp.route('/manager/pending')
@login_required
def manager_pending():
    # ---- filters from query params ----
    emp_id     = request.args.get('employee_id', type=int)
    month_from = request.args.get('month_from', type=str)  # 'YYYY-MM'
    month_to   = request.args.get('month_to', type=str)    # 'YYYY-MM'
    remark_q   = (request.args.get('remark', '') or '').strip()

    def month_bounds(yyyy_mm: str):
        y, m = map(int, yyyy_mm.split('-'))
        start = date(y, m, 1)
        end = date(y, m, monthrange(y, m)[1])
        return start, end

    # restrict requests to manager’s team
    team_ids = [e.id for e in Employee.query.filter_by(reporting_manager_id=current_user.id).all()]

    q = (LeaveRequest.query
         .join(Employee, LeaveRequest.employee_id == Employee.id)
         .filter(LeaveRequest.employee_id.in_(team_ids))
         .order_by(LeaveRequest.applied_on.desc()))

    # Employee filter
    if emp_id:
        q = q.filter(LeaveRequest.employee_id == emp_id)

    # Month range filter (on start_date)
    if month_from and month_to:
        d1 = month_bounds(month_from)[0]
        d2 = month_bounds(month_to)[1]
        q = q.filter(LeaveRequest.start_date.between(d1, d2))
    elif month_from:
        d1, d2 = month_bounds(month_from)
        q = q.filter(LeaveRequest.start_date.between(d1, d2))
    elif month_to:
        d2 = month_bounds(month_to)[1]
        q = q.filter(LeaveRequest.start_date <= d2)

    team_requests = q.all()

    # Split into pending vs processed
    pending_requests, processed_requests = [], []
    for r in team_requests:
        manager_logs = [log for log in getattr(r, 'approval_logs', []) if log.approver_role == 'manager']
        if r.status == 'pending' and not manager_logs:
            pending_requests.append(r)
        else:
            processed_requests.append(r)

    # Manager remark filter applies to processed (logs exist there)
    if remark_q:
        rq = remark_q.lower()
        processed_requests = [
            r for r in processed_requests
            if any(
                (log.approver_role == 'manager') and (log.remark or '').lower().find(rq) >= 0
                for log in getattr(r, 'approval_logs', [])
            )
        ]
        # OPTIONAL: also filter pending by employee "reason" text:
        # pending_requests = [r for r in pending_requests if (r.reason or '').lower().find(rq) >= 0]

    # Team employees for the filter dropdown
    team_employees = Employee.query.filter(Employee.id.in_(team_ids)).order_by(Employee.full_name).all()

    return render_template(
        'leave/manager_pending.html',
        pending_requests=pending_requests,
        processed_requests=processed_requests,
        employees=team_employees,
        # echo filter values so the template can keep selections
        emp_id=emp_id,
        month_from=month_from,
        month_to=month_to,
        remark_q=remark_q
    )

@leave_bp.route('/manager/approve/<int:req_id>', methods=['POST'])
@login_required
def manager_approve(req_id):
    req = LeaveRequest.query.get_or_404(req_id)
    if req.status != 'pending':
        flash('⚠️ Already processed.', 'warning')
        return redirect(url_for('leave.manager_pending'))

    req.status = 'manager_approved'
    db.session.add(LeaveApprovalLog(
        leave_request_id=req.id,
        approver_id=current_user.id,
        approver_role='manager',
        status='approved',
        remark=request.form['remark']
    ))
    db.session.commit()

    # Notifications
    rng = _format_range(req.start_date, req.end_date)
    emp = Employee.query.get(req.employee_id)
    lt  = LeaveType.query.get(req.leave_type_id)

    # 4.1 Notify employee
    notify_employee(
        employee_id=req.employee_id,
        title="Leave approved by Manager",
        message=f"Your {lt.name} request ({rng}) was approved by your manager and sent to HR.",
        link=_link('leave.view_history')
    )

    # 4.2 Notify HR to act
    hr_list = _hr_ids()
    if hr_list:
        notify_many(
            employee_ids=hr_list,
            title="Leave request awaiting HR action",
            message=f"{emp.full_name}'s {lt.name} ({rng}) needs HR approval.",
            link=_link('leave.hr_pending')
        )

    flash('✅ Approved and forwarded to HR.', 'success')
    return redirect(url_for('leave.manager_pending'))

# -----------------------------
# HR PENDING + APPROVE
# -----------------------------

from calendar import monthrange
from datetime import date
from flask import request, render_template

@leave_bp.route('/hr/pending')
@login_required
def hr_pending():
    # ---- filters ----
    emp_id     = request.args.get('employee_id', type=int)
    month_from = request.args.get('month_from', type=str)  # 'YYYY-MM'
    month_to   = request.args.get('month_to', type=str)    # 'YYYY-MM'
    remark_q   = (request.args.get('remark', '') or '').strip()  # HR remark contains

    def month_bounds(yyyy_mm: str):
        y, m = map(int, yyyy_mm.split('-'))
        start = date(y, m, 1)
        end = date(y, m, monthrange(y, m)[1])
        return start, end

    # Base queries
    q_pending = (LeaveRequest.query
                 .filter(LeaveRequest.status == 'manager_approved')
                 .order_by(LeaveRequest.applied_on.desc()))

    q_processed = (LeaveRequest.query
                   .filter(LeaveRequest.id.in_(
                       db.session.query(LeaveApprovalLog.leave_request_id)
                       .filter(LeaveApprovalLog.approver_role == 'hr')
                   ))
                   .order_by(LeaveRequest.applied_on.desc()))

    # Apply Employee filter
    if emp_id:
        q_pending = q_pending.filter(LeaveRequest.employee_id == emp_id)
        q_processed = q_processed.filter(LeaveRequest.employee_id == emp_id)

    # Apply Month range on start_date
    if month_from and month_to:
        d1 = month_bounds(month_from)[0]
        d2 = month_bounds(month_to)[1]
        q_pending   = q_pending.filter(LeaveRequest.start_date.between(d1, d2))
        q_processed = q_processed.filter(LeaveRequest.start_date.between(d1, d2))
    elif month_from:
        d1, d2 = month_bounds(month_from)
        q_pending   = q_pending.filter(LeaveRequest.start_date.between(d1, d2))
        q_processed = q_processed.filter(LeaveRequest.start_date.between(d1, d2))
    elif month_to:
        d2 = month_bounds(month_to)[1]
        q_pending   = q_pending.filter(LeaveRequest.start_date <= d2)
        q_processed = q_processed.filter(LeaveRequest.start_date <= d2)

    pending_requests = q_pending.all()
    processed_requests = q_processed.all()

    # HR remark filter applies to processed (based on HR logs)
    if remark_q:
        rq = remark_q.lower()
        processed_requests = [
            r for r in processed_requests
            if any(
                (log.approver_role == 'hr') and (log.remark or '').lower().find(rq) >= 0
                for log in getattr(r, 'approval_logs', [])
            )
        ]
        # If you also want to filter PENDING by the employee's reason, uncomment:
        # pending_requests = [r for r in pending_requests if (r.reason or '').lower().find(rq) >= 0]

    # Dropdown data (all employees)
    employees = Employee.query.order_by(Employee.full_name).all()

    return render_template(
        'leave/hr_pending.html',
        pending_requests=pending_requests,
        processed_requests=processed_requests,
        employees=employees,
        emp_id=emp_id,
        month_from=month_from,
        month_to=month_to,
        remark_q=remark_q
    )


@leave_bp.route('/hr/approve/<int:req_id>', methods=['POST'])
@login_required
def hr_approve(req_id):
    req = LeaveRequest.query.get_or_404(req_id)
    leave_type = LeaveType.query.get(req.leave_type_id)
    settings = LeaveSettings.query.first()
    year = datetime.now().year

    if req.status != 'manager_approved':
        flash('⚠️ Not eligible for HR approval.', 'warning')
        return redirect(url_for('leave.hr_pending'))

    # Deduct balance only for non-RH, non-LWP
    if leave_type.code not in ('RH', 'LWP'):
        balance = LeaveBalance.query.filter_by(
            employee_id=req.employee_id,
            leave_type_id=req.leave_type_id,
            year=year
        ).first()
        if not balance:
            balance = LeaveBalance(
                employee_id=req.employee_id,
                leave_type_id=req.leave_type_id,
                year=year,
                balance=0
            )
            db.session.add(balance)

        if balance.balance < req.total_days:
            flash(f'❌ Insufficient balance. Available: {balance.balance}', 'danger')
            return redirect(url_for('leave.hr_pending'))

        balance.balance -= req.total_days

    # Finalise
    req.status = 'final_approved'
    db.session.add(LeaveApprovalLog(
        leave_request_id=req.id,
        approver_id=current_user.id,
        approver_role='hr',
        status='approved',
        remark=request.form.get('remark', '').strip()
    ))
    db.session.commit()

    # Sync to attendance
    sync_leave_to_attendance(db, req.employee_id, req.start_date, req.end_date, leave_type.name)
    db.session.commit()

    # Notifications
    rng = _format_range(req.start_date, req.end_date)
    emp = Employee.query.get(req.employee_id)
    lt  = LeaveType.query.get(req.leave_type_id)

    # 6.1 Notify employee
    notify_employee(
        employee_id=req.employee_id,
        title="Leave fully approved",
        message=f"Your {lt.name} request ({rng}) has been approved by HR.",
        link=_link('leave.view_history')
    )

    # 6.2 Notify manager (FYI)
    if emp and emp.reporting_manager_id:
        notify_employee(
            employee_id=emp.reporting_manager_id,
            title="Leave approved (HR)",
            message=f"{emp.full_name}'s {lt.name} ({rng}) is fully approved.",
            link=_link('leave.manager_pending')
        )

    flash('✅ Final approval done & attendance updated.', 'success')
    return redirect(url_for('leave.hr_pending'))

# -----------------------------
# Comp-Off Application
# -----------------------------
@leave_bp.route('/comp-off/apply', methods=['GET', 'POST'])
@login_required
def apply_comp_off():
    if request.method == 'POST':
        work_date = datetime.strptime(request.form['work_date'], '%Y-%m-%d').date()
        reason = (request.form['reason'] or '').strip()
        valid_until = CompOffRequest.default_valid_until(work_date, days=30)

        if CompOffRequest.query.filter_by(employee_id=current_user.id, work_date=work_date).first():
            flash('⚠️ Comp-off already requested for this date.', 'warning')
            return redirect(url_for('leave.apply_comp_off'))

        db.session.add(CompOffRequest(
            employee_id=current_user.id,
            work_date=work_date,
            reason=reason,
            valid_until=valid_until,
            status='pending'
        ))
        db.session.commit()
        flash('✅ Comp-off request submitted.', 'success')
        return redirect(url_for('leave.apply_comp_off'))

    # --- Filters for "my requests" ---
    d_from   = request.args.get('date_from')
    d_to     = request.args.get('date_to')
    status_q = (request.args.get('status') or '').lower()  # pending|approved|rejected|all
    q_text   = (request.args.get('q') or '').strip().lower()

    q = (CompOffRequest.query
         .filter(CompOffRequest.employee_id == current_user.id)
         .order_by(CompOffRequest.work_date.desc()))

    if d_from:
        try:
            q = q.filter(CompOffRequest.work_date >= datetime.strptime(d_from, "%Y-%m-%d").date())
        except Exception:
            pass
    if d_to:
        try:
            q = q.filter(CompOffRequest.work_date <= datetime.strptime(d_to, "%Y-%m-%d").date())
        except Exception:
            pass
    if status_q and status_q != 'all':
        q = q.filter(CompOffRequest.status == status_q)
    if q_text:
        q = q.filter(db.func.lower(CompOffRequest.reason).contains(q_text))

    existing = q.all()

    return render_template(
        'leave/comp_off_apply.html',
        requests=existing,
        date_from=d_from, date_to=d_to, status_q=status_q, q_text=q_text
    )
# -----------------------------
# Comp-Off Approval (Admin/HR)
# -----------------------------
from calendar import monthrange
from datetime import datetime, date

@leave_bp.route('/comp-off/pending')
@login_required
def pending_comp_off():
    """
    Filters (all optional, via query params):
      employee_id=<int>
      status=pending|approved|rejected|all   (default: pending view still shows pending; processed section shows approved+rejected)
      date_from=YYYY-MM-DD
      date_to=YYYY-MM-DD
      q=<free text in reason>
    """
    emp_id   = request.args.get('employee_id', type=int)
    status_q = (request.args.get('status') or '').lower()
    d_from   = request.args.get('date_from')
    d_to     = request.args.get('date_to')
    q_text   = (request.args.get('q') or '').strip().lower()

    # base queries
    qp = CompOffRequest.query.order_by(CompOffRequest.work_date.desc())
    qr = CompOffRequest.query.order_by(CompOffRequest.work_date.desc())

    # employee filter
    if emp_id:
        qp = qp.filter(CompOffRequest.employee_id == emp_id)
        qr = qr.filter(CompOffRequest.employee_id == emp_id)

    # date range (work_date)
    if d_from:
        try:
            d1 = datetime.strptime(d_from, "%Y-%m-%d").date()
            qp = qp.filter(CompOffRequest.work_date >= d1)
            qr = qr.filter(CompOffRequest.work_date >= d1)
        except Exception:
            pass
    if d_to:
        try:
            d2 = datetime.strptime(d_to, "%Y-%m-%d").date()
            qp = qp.filter(CompOffRequest.work_date <= d2)
            qr = qr.filter(CompOffRequest.work_date <= d2)
        except Exception:
            pass

    # text search (reason)
    if q_text:
        qp = qp.filter(db.func.lower(CompOffRequest.reason).contains(q_text))
        qr = qr.filter(db.func.lower(CompOffRequest.reason).contains(q_text))

    # status
    # left table = pending; right/bottom table = processed (approved/rejected)
    pending_requests   = qp.filter(CompOffRequest.status == 'pending').all()
    processed_requests = qr.filter(CompOffRequest.status.in_(['approved', 'rejected'])).all()

    # dropdown data
    employees = Employee.query.order_by(Employee.full_name).all()

    return render_template(
        'leave/comp_off_pending.html',
        pending_requests=pending_requests,
        processed_requests=processed_requests,
        employees=employees,
        # echo filters back to template
        emp_id=emp_id,
        status_q=status_q,
        date_from=d_from,
        date_to=d_to,
        q_text=q_text,
    )

@leave_bp.route('/comp-off/approve/<int:req_id>', methods=['POST'])
@login_required
def approve_comp_off(req_id):
    req = CompOffRequest.query.get_or_404(req_id)

    if req.status != 'pending':
        flash('⚠️ Already processed.', 'warning')
        return redirect(url_for('leave.pending_comp_off'))

    # 1) Mark approved
    req.status = 'approved'
    db.session.flush()

    # 2) Find Comp Off leave type (support common codes)
    co_type = (LeaveType.query
               .filter(db.func.upper(LeaveType.code).in_(['CO', 'COMP_OFF']))
               .first())
    if not co_type:
        flash('❌ Comp Off Leave Type missing (code CO or COMP_OFF). Create it first.', 'danger')
        db.session.rollback()
        return redirect(url_for('leave.pending_comp_off'))

    # 3) Credit balance (+1 day per approved comp-off)
    #    Keep it consistent with your other flows that use the *calendar* year for balances.
    credit_year = datetime.now().year

    bal = (LeaveBalance.query
           .filter_by(employee_id=req.employee_id, leave_type_id=co_type.id, year=credit_year)
           .first())
    if not bal:
        bal = LeaveBalance(
            employee_id=req.employee_id,
            leave_type_id=co_type.id,
            year=credit_year,
            balance=0.0
        )
        db.session.add(bal)
        db.session.flush()

    bal.balance = float(bal.balance or 0.0) + 1.0

    db.session.commit()
    flash('✅ Comp-off approved and 1 day credited to CO balance.', 'success')
    return redirect(url_for('leave.pending_comp_off'))

@leave_bp.route('/comp-off/reject/<int:req_id>', methods=['POST'])
@login_required
def reject_comp_off(req_id):
    req = CompOffRequest.query.get_or_404(req_id)
    req.status = 'rejected'
    db.session.commit()
    flash('❌ Comp-off rejected.', 'warning')
    return redirect(url_for('leave.pending_comp_off'))

@leave_bp.route('/allocation', methods=['GET', 'POST'])
@login_required
def leave_allocation():
    leave_types = LeaveType.query.order_by(LeaveType.name).all()
    results = []

    today = date.today()
    fy_start, fy_end = fy_bounds_for_today(today)
    selected_year = fy_start.year  # FY label (your LeaveBalance.year)

    if request.method == 'POST':
        form = request.form
        leave_type_id = int(form['leave_type_id'])
        strategy = form['strategy']
        force = ('force' in form)  # allow override of guard
        leave_type = LeaveType.query.get_or_404(leave_type_id)

        # ----- Freeze cutoff (from form or default to Aug of this FY) -----
        # Input comes as YYYY-MM (type="month")
        freeze_till = (form.get('freeze_till') or '').strip()  # e.g., "2025-08"
        if freeze_till:
            f_year, f_month = map(int, freeze_till.split('-'))
        else:
            # Default: freeze till August of the FY label (Apr-Mar)
            f_year, f_month = selected_year, 8

        # ‘since’ moment is the first day of the month after the freeze
        # e.g., freeze 2025-08 -> since_dt = 2025-09-01 00:00:00
        if f_month == 12:
            since_dt = datetime(f_year + 1, 1, 1)
        else:
            since_dt = datetime(f_year, f_month + 1, 1)

        # Guard against accidental re-run for same FY + LT + strategy (unless forced)
        existing_batch = _latest_non_revoked_batch(leave_type_id, selected_year, strategy)
        if existing_batch and not force:
            flash(f"⚠️ An allocation batch already exists for FY {selected_year}, "
                  f"{leave_type.code}, strategy {strategy}. Tick 'Force apply' to proceed or Revoke the batch.", "warning")
            return render_template(
                'leave/leave_allocation.html',
                leave_types=leave_types,
                selected_year=selected_year,
                batch_result=None,
                guard_exists=True,
                guard_batch=existing_batch,
                freeze_default=f"{f_year:04d}-{f_month:02d}"
            )

        # ---- Build a new batch audit shell ----
        batch_key = _make_batch_key(selected_year, leave_type.code, strategy)
        batch = LeaveAllocationBatch(
            batch_key=batch_key,
            leave_type_id=leave_type.id,
            fy_label=selected_year,
            strategy=strategy,
            run_by_id=getattr(current_user, 'id', None),
            note=f"Allocation post-freeze (frozen till {f_year}-{f_month:02d})"
        )
        db.session.add(batch)
        db.session.flush()  # get batch.id

        employees = Employee.query.filter(db.func.lower(Employee.status) == 'active').all()

        for emp in employees:
            doj = getattr(emp, 'date_of_joining', None)

            # If DOJ is after FY end, skip
            if doj and doj > fy_end:
                results.append({"employee_name": emp.full_name, "allocated": 0, "remark": "Joined after FY, skipped"})
                continue

            # Ensure there is a balance row for (emp, lt, FY)
            bal = (LeaveBalance.query
                   .filter_by(employee_id=emp.id, leave_type_id=leave_type.id, year=selected_year)
                   .with_for_update(read=True)
                   .first())
            if not bal:
                bal = LeaveBalance(employee_id=emp.id, leave_type_id=leave_type.id, year=selected_year, balance=0.0)
                db.session.add(bal)
                db.session.flush()
            prev = float(bal.balance or 0.0)

            # ---- Compute target_total for months strictly AFTER freeze ----
            if strategy == 'monthly_credit':
                months = months_after_freeze_until(today, f_year, f_month)
                monthly = round((leave_type.default_quota or 0) / 12.0, 2)
                # cap at default_quota, but only for counted months
                target_total = round(min(len(months) * monthly, (leave_type.default_quota or 0)), 2)
                remark_core = f"{len(months)} month(s) × {monthly} post-freeze"
            elif strategy == 'chl_earned_leave':
                target_total = calc_chl_earned_leave_total_after_freeze(doj, today, f_year, f_month)
                remark_core = "CHL EL rule (post-freeze)"
            elif strategy == 'annual':
                # If annual and freeze is set, annual post-freeze is ambiguous.
                # Common approach: allocate 0 until after freeze, then allocate full once (if you want).
                # Here we credit nothing pre-freeze and full once after freeze month:
                months = months_after_freeze_until(today, f_year, f_month)
                target_total = float(leave_type.default_quota or 0) if months else 0.0
                remark_core = "Annual post-freeze"
            else:
                results.append({"employee_name": emp.full_name, "allocated": 0, "remark": "❌ Unknown strategy"})
                continue

            # ---- Delta = target_total - already_credited_since_freeze ----
            already_since = _sum_credited_since(emp.id, leave_type.id, since_dt)
            delta = max(0.0, round(target_total - already_since, 2))

            # Audit entry BEFORE changing balance
            entry = LeaveAllocationEntry(
                batch_id=batch.id,
                employee_id=emp.id,
                leave_type_id=leave_type.id,
                fy_label=selected_year,
                previous_balance=prev,
                new_balance=round(prev + delta, 2),
                credit_delta=delta
            )
            db.session.add(entry)

            # Apply only the delta; never overwrite usage
            if delta > 0:
                bal.balance = round(prev + delta, 2)

            results.append({
                "employee_name": emp.full_name,
                "allocated": delta,
                "remark": f"{strategy.replace('_',' ').title()} ({remark_core}); "
                          f"freeze_till={f_year}-{f_month:02d}, target_post_freeze={target_total}, already_post_freeze={already_since}, was={prev}"
            })

        db.session.commit()
        flash(f"✅ Allocation (post-freeze) applied in batch {batch.batch_key} for {len(results)} active employee(s).", "success")

        return render_template(
            'leave/leave_allocation.html',
            leave_types=leave_types,
            selected_year=selected_year,
            batch_result={"leave_type": leave_type.name, "entries": results, "batch_key": batch.batch_key},
            freeze_default=f"{f_year:04d}-{f_month:02d}"
        )

    # GET
    # default freeze shown as Aug of current FY (editable in UI)
    return render_template(
        'leave/leave_allocation.html',
        leave_types=leave_types,
        selected_year=selected_year,
        batch_result=None,
        freeze_default=f"{selected_year}-08"
    )

@leave_bp.route('/hr/reject/<int:req_id>', methods=['POST'])
@login_required
def hr_reject(req_id):
    req = LeaveRequest.query.get_or_404(req_id)

    # Check if already acted on
    existing_log = LeaveApprovalLog.query.filter_by(
        leave_request_id=req.id,
        approver_role='hr'
    ).first()
    if existing_log:
        flash('⚠️ HR has already acted on this request.', 'warning')
        return redirect(url_for('leave.hr_pending'))

    req.status = 'rejected'

    db.session.add(LeaveApprovalLog(
        leave_request_id=req.id,
        approver_id=current_user.id,
        approver_role='hr',
        status='rejected',
        remark=request.form['remark']
    ))
    db.session.commit()

    # Notifications
    rng = _format_range(req.start_date, req.end_date)
    emp = Employee.query.get(req.employee_id)
    lt  = LeaveType.query.get(req.leave_type_id)

    # 7.1 Employee
    notify_employee(
        employee_id=req.employee_id,
        title="Leave rejected by HR",
        message=f"Your {lt.name} request ({rng}) was rejected by HR.",
        link=_link('leave.view_history')
    )

    # 7.2 Manager (FYI)
    if emp and emp.reporting_manager_id:
        notify_employee(
            employee_id=emp.reporting_manager_id,
            title="Leave rejected (HR)",
            message=f"{emp.full_name}'s {lt.name} ({rng}) was rejected by HR.",
            link=_link('leave.manager_pending')
        )

    flash('❌ Leave rejected by HR.', 'danger')
    return redirect(url_for('leave.hr_pending'))

@leave_bp.route('/manager/reject/<int:req_id>', methods=['POST'])
@login_required
def manager_reject(req_id):
    req = LeaveRequest.query.get_or_404(req_id)
    
    if req.status != 'pending':
        flash('⚠️ Already processed.', 'warning')
        return redirect(url_for('leave.manager_pending'))

    # Update the request status
    req.status = 'rejected'

    # Log the manager's rejection
    db.session.add(LeaveApprovalLog(
        leave_request_id=req.id,
        approver_id=current_user.id,
        approver_role='manager',
        status='rejected',
        remark=request.form['remark']
    ))
    db.session.commit()

    # Notifications
    rng = _format_range(req.start_date, req.end_date)
    lt  = LeaveType.query.get(req.leave_type_id)
    notify_employee(
        employee_id=req.employee_id,
        title="Leave rejected by Manager",
        message=f"Your {lt.name} request ({rng}) was rejected by your manager.",
        link=_link('leave.view_history')
    )

    flash('❌ Leave rejected by Manager.', 'danger')
    return redirect(url_for('leave.manager_pending'))


@leave_bp.route('/cancel/<int:req_id>', methods=['POST'])
@login_required
def cancel_leave(req_id):
    req = LeaveRequest.query.get_or_404(req_id)

    if req.employee_id != current_user.id:
        flash('❌ You are not authorized to cancel this request.', 'danger')
        return redirect(url_for('leave.view_history'))

    if req.status != 'pending':
        flash('⚠️ Only pending requests can be cancelled.', 'warning')
        return redirect(url_for('leave.view_history'))

    req.status = 'cancelled'
    db.session.commit()
    # Notifications
    rng = _format_range(req.start_date, req.end_date)
    lt  = LeaveType.query.get(req.leave_type_id)

    # 8.1 Manager
    mgr_id = getattr(current_user, 'reporting_manager_id', None)
    if mgr_id:
        notify_employee(
            employee_id=mgr_id,
            title="Leave request cancelled",
            message=f"{current_user.full_name} cancelled their {lt.name} request ({rng}).",
            link=_link('leave.manager_pending')
        )

    # 8.2 HR
    hr_list = _hr_ids()
    if hr_list:
        notify_many(
            employee_ids=hr_list,
            title="Leave request cancelled",
            message=f"{current_user.full_name} cancelled {lt.name} ({rng}).",
            link=_link('leave.hr_pending')
        )

    flash('✅ Leave request cancelled successfully.', 'success')
    return redirect(url_for('leave.view_history'))

@leave_bp.route('/eligible-dates', methods=['GET'])
@login_required
def eligible_dates():
    q_type = request.args.get('type', '').lower()
    ly_start, ly_end = get_leave_year_range()

    if q_type == 'rh':
        all_opt = get_optional_holiday_dates(ly_start, ly_end, employee=current_user)

        # Exclude in-flight RH dates
        in_flight = set()
        in_flight_statuses = ['pending', 'manager_approved', 'final_approved']
        existing = LeaveRequest.query.filter(
            LeaveRequest.employee_id == current_user.id,
            LeaveRequest.leave_type.has(code='RH'),
            LeaveRequest.start_date >= ly_start,
            LeaveRequest.end_date <= ly_end,
            LeaveRequest.status.in_(in_flight_statuses)
        ).all()
        for r in existing:
            dt = r.start_date
            while dt <= r.end_date:
                in_flight.add(dt)
                dt += timedelta(days=1)

        allowed = sorted(d for d in all_opt - in_flight)

        # Build {date, name}
        items = []
        for d in allowed:
            h = get_applicable_holiday_for(current_user, d)  # uses location-aware fallback
            items.append({
                "date": d.strftime("%Y-%m-%d"),
                "name": (h.name if h else "Optional Holiday")
            })
        return {"dates": items}

    return {"dates": []}


# -----------------------------
# BULK ACTIONS: MANAGER + HR
# -----------------------------
from sqlalchemy import and_

def _ensure_manager_scope(qs, manager_id: int):
    """Limit queryset to the manager's team."""
    team_ids = [e.id for e in Employee.query.filter_by(reporting_manager_id=manager_id).all()]
    return [r for r in qs if r.employee_id in team_ids]

@leave_bp.post('/manager/bulk-action')
@login_required
def manager_bulk_action():
    """Approve/Reject multiple PENDING requests belonging to current manager's team."""
    action = request.form.get('action')  # 'approve' | 'reject'
    ids = request.form.getlist('req_ids')  # list of request IDs (strings)
    remark = (request.form.get('remark') or '').strip()

    if action not in ('approve', 'reject') or not ids:
        flash('❌ Select at least one request and an action.', 'danger')
        return redirect(url_for('leave.manager_pending'))

    # Fetch requests and scope to manager's team
    reqs = LeaveRequest.query.filter(
        LeaveRequest.id.in_(ids),
        LeaveRequest.status == 'pending'
    ).all()
    reqs = _ensure_manager_scope(reqs, current_user.id)

    approved, rejected, skipped = 0, 0, 0
    for r in reqs:
        if action == 'approve':
            # Already processed by someone?
            if r.status != 'pending':
                skipped += 1
                continue
            r.status = 'manager_approved'
            db.session.add(LeaveApprovalLog(
                leave_request_id=r.id,
                approver_id=current_user.id,
                approver_role='manager',
                status='approved',
                remark=remark
            ))
            approved += 1
            # Notify HR (batching later is OK; keep simple)
            hr_list = _hr_ids()
            if hr_list:
                emp = Employee.query.get(r.employee_id)
                lt  = LeaveType.query.get(r.leave_type_id)
                rng = _format_range(r.start_date, r.end_date)
                notify_many(
                    employee_ids=hr_list,
                    title="Leave request awaiting HR action",
                    message=f"{emp.full_name}'s {lt.name} ({rng}) needs HR approval.",
                    link=_link('leave.hr_pending')
                )
            # Notify employee
            lt = LeaveType.query.get(r.leave_type_id)
            notify_employee(
                employee_id=r.employee_id,
                title="Leave approved by Manager",
                message=f"Your {lt.name} request ({_format_range(r.start_date, r.end_date)}) was approved by your manager and sent to HR.",
                link=_link('leave.view_history')
            )
        else:
            if r.status != 'pending':
                skipped += 1
                continue
            r.status = 'rejected'
            db.session.add(LeaveApprovalLog(
                leave_request_id=r.id,
                approver_id=current_user.id,
                approver_role='manager',
                status='rejected',
                remark=remark
            ))
            rejected += 1
            # Notify employee
            lt = LeaveType.query.get(r.leave_type_id)
            notify_employee(
                employee_id=r.employee_id,
                title="Leave rejected by Manager",
                message=f"Your {lt.name} request ({_format_range(r.start_date, r.end_date)}) was rejected by your manager.",
                link=_link('leave.view_history')
            )

    db.session.commit()
    flash(f"✅ Manager bulk: approved={approved}, rejected={rejected}, skipped={skipped}.", "success")
    return redirect(url_for('leave.manager_pending'))


@leave_bp.post('/hr/bulk-action')
@login_required
def hr_bulk_action():
    """Approve/Reject multiple requests pending for HR (status='manager_approved')."""
    action = request.form.get('action')  # 'approve' | 'reject'
    ids = request.form.getlist('req_ids')
    remark = (request.form.get('remark') or '').strip()

    # Simple HR guard
    if not _hr_ids() or current_user.id not in _hr_ids():
        # If you have a better is_hr() function, use it.
        # Or replace with your existing is_hr() helper if available.
        pass  # Optional: enforce role

    if action not in ('approve', 'reject') or not ids:
        flash('❌ Select at least one request and an action.', 'danger')
        return redirect(url_for('leave.hr_pending'))

    reqs = LeaveRequest.query.filter(
        LeaveRequest.id.in_(ids),
        LeaveRequest.status == 'manager_approved'
    ).all()

    approved, rejected, skipped, balance_errors = 0, 0, 0, 0

    # Common
    year_now = datetime.now().year

    for r in reqs:
        # Re-fetch leave type
        lt = LeaveType.query.get(r.leave_type_id)
        if not lt:
            skipped += 1
            continue

        if action == 'approve':
            # Balance deduction (non-RH, non-LWP)
            if lt.code not in ('RH', 'LWP'):
                bal = LeaveBalance.query.filter_by(
                    employee_id=r.employee_id,
                    leave_type_id=r.leave_type_id,
                    year=year_now
                ).first()
                if not bal:
                    bal = LeaveBalance(
                        employee_id=r.employee_id,
                        leave_type_id=r.leave_type_id,
                        year=year_now,
                        balance=0
                    )
                    db.session.add(bal)
                    db.session.flush()

                need = 0.5 if r.half_day or (r.total_days == 0.5) else (r.total_days or 0)
                if (bal.balance or 0) < need:
                    # Not enough balance -> skip this item
                    balance_errors += 1
                    continue
                bal.balance = (bal.balance or 0) - need

            # Finalise approval
            r.status = 'final_approved'
            db.session.add(LeaveApprovalLog(
                leave_request_id=r.id,
                approver_id=current_user.id,
                approver_role='hr',
                status='approved',
                remark=remark
            ))
            approved += 1

            # Attendance sync
            try:
                sync_leave_to_attendance(db, r.employee_id, r.start_date, r.end_date, lt.name)
            except Exception:
                # Don’t fail the whole batch on sync error; you can add logging if needed
                pass

            # Notify employee + manager FYI
            emp = Employee.query.get(r.employee_id)
            rng = _format_range(r.start_date, r.end_date)
            notify_employee(
                employee_id=r.employee_id,
                title="Leave fully approved",
                message=f"Your {lt.name} request ({rng}) has been approved by HR.",
                link=_link('leave.view_history')
            )
            if emp and emp.reporting_manager_id:
                notify_employee(
                    employee_id=emp.reporting_manager_id,
                    title="Leave approved (HR)",
                    message=f"{emp.full_name}'s {lt.name} ({rng}) is fully approved.",
                    link=_link('leave.manager_pending')
                )

        else:  # reject
            r.status = 'rejected'
            db.session.add(LeaveApprovalLog(
                leave_request_id=r.id,
                approver_id=current_user.id,
                approver_role='hr',
                status='rejected',
                remark=remark
            ))
            rejected += 1

            # Notify employee + manager FYI
            emp = Employee.query.get(r.employee_id)
            rng = _format_range(r.start_date, r.end_date)
            notify_employee(
                employee_id=r.employee_id,
                title="Leave rejected by HR",
                message=f"Your {lt.name} request ({rng}) was rejected by HR.",
                link=_link('leave.view_history')
            )
            if emp and emp.reporting_manager_id:
                notify_employee(
                    employee_id=emp.reporting_manager_id,
                    title="Leave rejected (HR)",
                    message=f"{emp.full_name}'s {lt.name} ({rng}) was rejected by HR.",
                    link=_link('leave.manager_pending')
                )

    db.session.commit()
    msg = f"✅ HR bulk: approved={approved}, rejected={rejected}, skipped={skipped}"
    if balance_errors:
        msg += f", balance_insufficient={balance_errors}"
    flash(msg + ".", "success")
    return redirect(url_for('leave.hr_pending'))


# -----------------------------
# LEAVE BALANCE REPORT + EXPORT
# -----------------------------
from io import StringIO, BytesIO
import csv

def _fy_range_for_report():
    """Use the leave-year window (Apr–Mar) you already compute elsewhere."""
    start, end = get_leave_year_range()
    # Your LeaveBalance.year stores FY start year (as in your allocation screen)
    fy_label = start.year
    return start, end, fy_label

def _sum_used_days(employee_id: int, leave_type_id: int, d1: date, d2: date) -> float:
    """Sum approved (final_approved) leave days in the FY for a type."""
    q = (LeaveRequest.query
         .filter(
            LeaveRequest.employee_id == employee_id,
            LeaveRequest.leave_type_id == leave_type_id,
            LeaveRequest.status == 'final_approved',
            LeaveRequest.start_date >= d1,
            LeaveRequest.end_date <= d2
         )
    )
    used = 0.0
    for lr in q:
        if lr.total_days and lr.total_days > 0:
            used += float(lr.total_days)
        else:
            # Fallback if total_days was not set
            span = (lr.end_date - lr.start_date).days + 1
            used += 0.5 if lr.half_day else span
    return round(used, 2)

@leave_bp.get('/balance-report')
@login_required
def balance_report():
    # Filters
    dept_id = request.args.get('department_id', type=int)
    emp_id  = request.args.get('employee_id', type=int)
    lt_id   = request.args.get('leave_type_id', type=int)

    fy_start, fy_end, fy_label = _fy_range_for_report()

    # Base query of balances for FY
    qb = (db.session.query(LeaveBalance, Employee, LeaveType)
          .join(Employee, LeaveBalance.employee_id == Employee.id)
          .join(LeaveType, LeaveBalance.leave_type_id == LeaveType.id)
          .filter(LeaveBalance.year == fy_label))

    if dept_id:
        qb = qb.join(Department, Department.id == Employee.department_id).filter(Department.id == dept_id)
    if emp_id:
        qb = qb.filter(LeaveBalance.employee_id == emp_id)
    if lt_id:
        qb = qb.filter(LeaveBalance.leave_type_id == lt_id)

    rows = qb.order_by(Employee.full_name, LeaveType.name).all()

    # Build report rows
    report = []
    totals = {"credited": 0.0, "used": 0.0, "remaining": 0.0}
    for bal, emp, lt in rows:
        remaining = float(bal.balance or 0.0)
        used = _sum_used_days(emp.id, lt.id, fy_start, fy_end)
        credited = round(remaining + used, 2)   # since remaining = credited - used in your flow

        report.append({
            "emp_code": emp.employee_code,
            "employee": emp.full_name,
            "department": getattr(emp.department, "name", "-") if hasattr(emp, "department") else "-",
            "leave_type": lt.name,
            "leave_code": lt.code,
            "credited": credited,
            "used": used,
            "remaining": remaining,
        })
        totals["credited"] += credited
        totals["used"] += used
        totals["remaining"] += remaining

    # dropdown data
    departments = Department.query.order_by(Department.name).all()
    employees   = Employee.query.order_by(Employee.full_name).all()
    leave_types = LeaveType.query.order_by(LeaveType.name).all()

    return render_template(
        'leave/balance_report.html',
        report=report,
        totals=totals,
        fy_label=fy_label,
        fy_start=fy_start,
        fy_end=fy_end,
        departments=departments,
        employees=employees,
        leave_types=leave_types,
        dept_id=dept_id,
        emp_id=emp_id,
        lt_id=lt_id
    )

@leave_bp.post('/allocation/preview')
@login_required
def leave_allocation_preview():
    form = request.form
    leave_type_id = int(form['leave_type_id'])
    strategy = form['strategy']
    leave_type = LeaveType.query.get_or_404(leave_type_id)

    today = date.today()
    fy_start, fy_end = fy_bounds_for_today(today)
    fy_label = fy_start.year

    # Freeze cutoff
    freeze_till = (form.get('freeze_till') or '').strip()
    if freeze_till:
        f_year, f_month = map(int, freeze_till.split('-'))
    else:
        f_year, f_month = fy_label, 8

    if f_month == 12:
        since_dt = datetime(f_year + 1, 1, 1)
    else:
        since_dt = datetime(f_year, f_month + 1, 1)

    employees = Employee.query.filter(func.lower(Employee.status) == 'active').all()
    rows = []

    for emp in employees:
        bal = LeaveBalance.query.filter_by(employee_id=emp.id, leave_type_id=leave_type.id, year=fy_label).first()
        prev = float(bal.balance if bal else 0.0)
        doj = getattr(emp, 'date_of_joining', None)

        if doj and doj > fy_end:
            rows.append({
                "emp_code": emp.employee_code, "employee": emp.full_name,
                "previous": prev, "new": prev, "delta": 0.0,
                "remark": "Joined after FY, skipped"
            })
            continue

        if strategy == 'monthly_credit':
            months = months_after_freeze_until(today, f_year, f_month)
            monthly = round((leave_type.default_quota or 0) / 12.0, 2)
            target_total = round(min(len(months) * monthly, (leave_type.default_quota or 0)), 2)
            remark = f"{len(months)} × {monthly} post-freeze"
        elif strategy == 'chl_earned_leave':
            target_total = calc_chl_earned_leave_total_after_freeze(doj, today, f_year, f_month)
            remark = "CHL EL rule (post-freeze)"
        elif strategy == 'annual':
            months = months_after_freeze_until(today, f_year, f_month)
            target_total = float(leave_type.default_quota or 0) if months else 0.0
            remark = "Annual post-freeze"
        else:
            target_total = 0.0
            remark = "Unknown strategy"

        already_since = _sum_credited_since(emp.id, leave_type.id, since_dt)
        delta = max(0.0, round(target_total - already_since, 2))

        rows.append({
            "emp_code": emp.employee_code,
            "employee": emp.full_name,
            "previous": prev,
            "new": round(prev + delta, 2),
            "delta": delta,
            "remark": f"{remark}; freeze_till={f_year}-{f_month:02d}, already_post_freeze={already_since}, target_post_freeze={target_total}"
        })

    return render_template(
        'leave/leave_allocation_preview.html',
        leave_type=leave_type,
        fy_label=fy_label,
        strategy=strategy,
        freeze_till=f"{f_year:04d}-{f_month:02d}",
        rows=rows
    )

@leave_bp.post('/allocation/revoke/<int:batch_id>')
@login_required
def revoke_allocation_batch(batch_id:int):
    batch = LeaveAllocationBatch.query.get_or_404(batch_id)
    if batch.is_revoked:
        flash('⚠️ This batch is already revoked.', 'warning')
        return redirect(url_for('leave.leave_allocation'))

    adjusted = 0

    if (batch.strategy or '').lower() == 'manual_set_balance':
        # Use the previous_balance restoration path
        restored, created = _revoke_set_balance_batch(batch)
        adjusted = restored
        note_tail = f" (manual_set_balance: restored={restored}, created_missing_rows={created})"
    else:
        # Your original delta-based revoke flow
        entries = LeaveAllocationEntry.query.filter_by(batch_id=batch.id).all()
        for ent in entries:
            if not ent.credit_delta or ent.credit_delta <= 0:
                continue
            bal = (LeaveBalance.query
                   .filter_by(employee_id=ent.employee_id, leave_type_id=ent.leave_type_id, year=ent.fy_label)
                   .with_for_update()
                   .first())
            if not bal:
                continue
            bal.balance = round(max(0.0, (bal.balance or 0.0) - ent.credit_delta), 2)
            adjusted += 1
        note_tail = ""

    batch.is_revoked = True
    batch.revoked_at = datetime.utcnow()
    batch.revoked_by_id = getattr(current_user, 'id', None)

    db.session.commit()
    flash(f"🧹 Reverted batch {batch.batch_key}. Adjusted {adjusted} balances.{note_tail}", "success")
    return redirect(url_for('leave.leave_allocation'))


# --- Bulk correction via CSV (desired TOTAL credit-to-date) ---
# POST /leave/allocation/correct (multipart/form-data with file=csv)
@leave_bp.post('/allocation/correct')
@login_required
def leave_allocation_correct():
    from werkzeug.utils import secure_filename
    import csv, io
    file = request.files.get('file')
    if not file:
        flash('Upload a CSV with: emp_code, lt_code, fy_label, desired_total_credit', 'danger')
        return redirect(url_for('leave.leave_allocation'))

    # Read CSV into memory
    data = file.read().decode('utf-8', errors='ignore')
    rows = list(csv.DictReader(io.StringIO(data)))
    required = {'emp_code','lt_code','fy_label','desired_total_credit'}
    if not required.issubset({c.strip() for c in rows[0].keys()}):
        flash('CSV must have columns: emp_code, lt_code, fy_label, desired_total_credit', 'danger')
        return redirect(url_for('leave.leave_allocation'))

    # Build a correction batch
    batch_key = f"CORRECT_{datetime.utcnow().strftime('%Y%m%d_%H%M%S')}"
    batch = LeaveAllocationBatch(
        batch_key=batch_key,
        leave_type_id=0,      # not a single type; per-row below
        fy_label=0,           # per-row below
        strategy='manual_correction',
        run_by_id=getattr(current_user, 'id', None),
        note='CSV corrections to total credit-to-date'
    )
    db.session.add(batch)
    db.session.flush()

    # Maps for lookups
    lt_by_code = {lt.code.upper(): lt for lt in LeaveType.query.all()}
    emp_by_code = {e.employee_code: e for e in Employee.query.all()}
    adjusted, skipped = 0, 0

    for r in rows:
        emp_code = r['emp_code'].strip()
        lt_code  = r['lt_code'].strip().upper()
        try:
            fy_label = int(r['fy_label'])
            desired_total = float(r['desired_total_credit'])
        except Exception:
            skipped += 1
            continue

        emp = emp_by_code.get(emp_code)
        lt  = lt_by_code.get(lt_code)
        if not emp or not lt:
            skipped += 1
            continue

        # How much credit has already been granted (via batches) this FY?
        already = _sum_credited_since(emp.id, lt.id, fy_label)
        delta = round(desired_total - already, 2)  # can be negative

        # Get current balance row
        bal = (LeaveBalance.query
               .filter_by(employee_id=emp.id, leave_type_id=lt.id, year=fy_label)
               .with_for_update()
               .first())
        if not bal:
            bal = LeaveBalance(employee_id=emp.id, leave_type_id=lt.id, year=fy_label, balance=0.0)
            db.session.add(bal)
            db.session.flush()

        prev = float(bal.balance or 0.0)
        new_bal = round(max(0.0, prev + delta), 2)

        # Audit entry with delta (can be negative)
        db.session.add(LeaveAllocationEntry(
            batch_id=batch.id,
            employee_id=emp.id,
            leave_type_id=lt.id,
            fy_label=fy_label,
            previous_balance=prev,
            new_balance=new_bal,
            credit_delta=delta
        ))

        # Apply delta (preserve usage)
        bal.balance = new_bal
        adjusted += 1

    db.session.commit()
    flash(f'✅ Correction batch {batch.batch_key} applied: adjusted={adjusted}, skipped={skipped}', 'success')
    return redirect(url_for('leave.leave_allocation'))


# =============================
# Allocation Batches: list, detail, export
# =============================
from flask import Response
from sqlalchemy import desc

@leave_bp.get('/allocation/batches')
@login_required
def allocation_batches():
    """
    Filters via query params (all optional):
      status=active|revoked|all   (default: active)
      leave_type_id=<int>
      fy_label=<int>
      strategy=<str>
      q=<batch_key contains>
      page=<int>  (default 1)
      per_page=<int> (default 20)
    """
    status      = (request.args.get('status') or 'active').lower()
    leave_type_id = request.args.get('leave_type_id', type=int)
    fy_label    = request.args.get('fy_label', type=int)
    strategy    = (request.args.get('strategy') or '').strip()
    q           = (request.args.get('q') or '').strip()
    page        = request.args.get('page', default=1, type=int)
    per_page    = request.args.get('per_page', default=20, type=int)

    qb = LeaveAllocationBatch.query

    if status in ('active', 'revoked'):
        qb = qb.filter(LeaveAllocationBatch.is_revoked == (status == 'revoked'))

    if leave_type_id:
        qb = qb.filter(LeaveAllocationBatch.leave_type_id == leave_type_id)
    if fy_label:
        qb = qb.filter(LeaveAllocationBatch.fy_label == fy_label)
    if strategy:
        qb = qb.filter(LeaveAllocationBatch.strategy == strategy)
    if q:
        qb = qb.filter(LeaveAllocationBatch.batch_key.ilike(f"%{q}%"))

    qb = qb.order_by(desc(LeaveAllocationBatch.run_at))

    # Pagination
    pagination = qb.paginate(page=page, per_page=per_page, error_out=False)
    batches = pagination.items

    # Pre-aggregate totals for each batch (sum of credit_delta and count of entries)
    # Avoid n+1: pull ids and do group sums
    batch_ids = [b.id for b in batches]
    totals_by_batch = {}
    if batch_ids:
        sums = (db.session.query(
                    LeaveAllocationEntry.batch_id,
                    db.func.coalesce(db.func.sum(LeaveAllocationEntry.credit_delta), 0.0),
                    db.func.count(LeaveAllocationEntry.id)
                )
                .filter(LeaveAllocationEntry.batch_id.in_(batch_ids))
                .group_by(LeaveAllocationEntry.batch_id)
                .all())
        totals_by_batch = {bid: {"sum_delta": float(s), "count": int(c)} for (bid, s, c) in sums}

    # Dropdown data
    leave_types = LeaveType.query.order_by(LeaveType.name).all()
    strategies  = ['monthly_credit', 'chl_earned_leave', 'annual', 'manual_correction', 'manual_set_balance', 'manual_adjust_one']

    return render_template(
        'leave/allocation_batches.html',
        batches=batches,
        totals_by_batch=totals_by_batch,
        leave_types=leave_types,
        strategies=strategies,
        status=status,
        leave_type_id=leave_type_id,
        fy_label=fy_label,
        strategy=strategy,
        q=q,
        pagination=pagination
    )

@leave_bp.get('/allocation/batches/<int:batch_id>')
@login_required
def allocation_batch_detail(batch_id: int):
    """
    Show all entries for a batch with employee details.
    """
    batch = LeaveAllocationBatch.query.get_or_404(batch_id)

    # Join entries with employees & leave types for display
    entries = (db.session.query(
                    LeaveAllocationEntry,
                    Employee.employee_code,
                    Employee.full_name,
                    LeaveType.code.label('lt_code'),
                    LeaveType.name.label('lt_name')
                )
                .join(Employee, Employee.id == LeaveAllocationEntry.employee_id)
                .join(LeaveType, LeaveType.id == LeaveAllocationEntry.leave_type_id)
                .filter(LeaveAllocationEntry.batch_id == batch.id)
                .order_by(Employee.full_name.asc())
                .all())

    # Totals
    total_delta = sum(float(e.LeaveAllocationEntry.credit_delta) for e in entries) if entries else 0.0
    count = len(entries)

    return render_template(
        'leave/allocation_batch_detail.html',
        batch=batch,
        entries=entries,
        total_delta=round(total_delta, 2),
        count=count
    )

@leave_bp.get('/allocation/batches/<int:batch_id>/export')
@login_required
def allocation_batch_export(batch_id: int):
    """
    Export a batch to CSV: emp_code, employee, leave_type, fy_label, prev, delta, new
    """
    batch = LeaveAllocationBatch.query.get_or_404(batch_id)
    rows = (db.session.query(
                Employee.employee_code,
                Employee.full_name,
                LeaveType.code.label('lt_code'),
                LeaveAllocationEntry.fy_label,
                LeaveAllocationEntry.previous_balance,
                LeaveAllocationEntry.credit_delta,
                LeaveAllocationEntry.new_balance
            )
            .join(LeaveAllocationEntry, LeaveAllocationEntry.employee_id == Employee.id)
            .join(LeaveType, LeaveType.id == LeaveAllocationEntry.leave_type_id)
            .filter(LeaveAllocationEntry.batch_id == batch.id)
            .order_by(Employee.full_name.asc())
            .all())

    def _gen():
        yield "emp_code,employee,leave_type,fy_label,previous,delta,new\n"
        for r in rows:
            yield f"{r.employee_code},{r.full_name},{r.lt_code},{r.fy_label},{r.previous_balance},{r.credit_delta},{r.new_balance}\n"

    filename = f"allocation_batch_{batch.batch_key}.csv"
    return Response(_gen(), mimetype='text/csv',
                    headers={"Content-Disposition": f"attachment; filename={filename}"})
                    

# ============================
# MANUAL "SET BALANCE (AS-OF TODAY)" – overrides current balance
# ============================
from io import StringIO
import csv

@leave_bp.get('/balance/set')
@login_required
def balance_set_form():
    """
    UI to set balances:
      - Option A: choose a leave type and a single value -> apply to all active employees
      - Option B: upload CSV with columns: emp_code, lt_code, fy_label, set_balance
    """
    leave_types = LeaveType.query.order_by(LeaveType.name).all()

    # Default FY label = Apr–Mar containing today
    _, _, fy_label = _fy_range_for_report()

    return render_template(
        'leave/balance_set.html',
        leave_types=leave_types,
        fy_label=fy_label
    )

@leave_bp.post('/balance/set/apply')
@login_required
def balance_set_apply():
    """
    Apply a single value for one leave type to all active employees (override).
    form: leave_type_id, fy_label, set_value, include_notice, include_resigned
    """
    try:
        leave_type_id = int(request.form['leave_type_id'])
        fy_label = int(request.form['fy_label'])
        set_value = float(request.form['set_value'])
    except Exception:
        flash("Invalid inputs. Please verify leave type, FY, and value.", "danger")
        return redirect(url_for('leave.balance_set_form'))

    include_notice = ('include_notice' in request.form)
    include_resigned = ('include_resigned' in request.form)

    # Employee scope
    statuses = ['active']
    if include_notice: statuses.append('notice_period')
    if include_resigned: statuses.append('resigned')

    employees = (Employee.query
                 .filter(func.lower(Employee.status).in_([s.lower() for s in statuses]))
                 .order_by(Employee.full_name).all())

    if not employees:
        flash("No employees matched the selected statuses.", "warning")
        return redirect(url_for('leave.balance_set_form'))

    lt = LeaveType.query.get_or_404(leave_type_id)

    # Create an audit batch – strategy ensures we don't count this as a credit
    batch = LeaveAllocationBatch(
        batch_key=f"SETBAL_{datetime.utcnow().strftime('%Y%m%d_%H%M%S')}",
        leave_type_id=leave_type_id,
        fy_label=fy_label,
        strategy='manual_set_balance',
        run_by_id=getattr(current_user, 'id', None),
        note='Manual set balance (override) for all employees'
    )
    db.session.add(batch)
    db.session.flush()

    adjusted = 0
    for emp in employees:
        bal = (LeaveBalance.query
               .filter_by(employee_id=emp.id, leave_type_id=leave_type_id, year=fy_label)
               .with_for_update(read=True)
               .first())
        if not bal:
            bal = LeaveBalance(employee_id=emp.id, leave_type_id=leave_type_id, year=fy_label, balance=0.0)
            db.session.add(bal)
            db.session.flush()

        prev = float(bal.balance or 0.0)
        new_bal = round(max(0.0, set_value), 2)   # guard against negative

        # Audit entry – credit_delta=0 so future allocations don't treat this as "credited"
        db.session.add(LeaveAllocationEntry(
            batch_id=batch.id,
            employee_id=emp.id,
            leave_type_id=leave_type_id,
            fy_label=fy_label,
            previous_balance=prev,
            new_balance=new_bal,
            credit_delta=0.0
        ))

        bal.balance = new_bal
        adjusted += 1

    db.session.commit()
    flash(f"✅ Set balance to {set_value} for {adjusted} employee(s) [{lt.code}, FY {fy_label}].", "success")
    return redirect(url_for('leave.allocation_batches'))


@leave_bp.post('/balance/set/csv')
@login_required
def balance_set_csv():
    """
    CSV override (per-employee):
      Columns: emp_code, lt_code, fy_label, set_balance
      Behaviour: set LeaveBalance.balance = set_balance (absolute), audit with credit_delta=0
    """
    file = request.files.get('file')
    if not file:
        flash('Upload a CSV with: emp_code, lt_code, fy_label, set_balance', 'danger')
        return redirect(url_for('leave.balance_set_form'))

    data = file.read().decode('utf-8', errors='ignore')
    rows = list(csv.DictReader(StringIO(data)))
    required = {'emp_code','lt_code','fy_label','set_balance'}
    if not rows or not required.issubset({c.strip() for c in rows[0].keys()}):
        flash('CSV must have columns: emp_code, lt_code, fy_label, set_balance', 'danger')
        return redirect(url_for('leave.balance_set_form'))

    # Lookups
    lt_by_code = {lt.code.upper(): lt for lt in LeaveType.query.all()}
    emp_by_code = {e.employee_code: e for e in Employee.query.all()}

    # One batch for the whole file
    batch = LeaveAllocationBatch(
        batch_key=f"SETBALCSV_{datetime.utcnow().strftime('%Y%m%d_%H%M%S')}",
        leave_type_id=0,     # mixed types per row
        fy_label=0,          # mixed FY per row
        strategy='manual_set_balance',
        run_by_id=getattr(current_user, 'id', None),
        note='CSV manual set balance (override)'
    )
    db.session.add(batch)
    db.session.flush()

    adjusted = skipped = 0
    for r in rows:
        emp_code = (r.get('emp_code') or '').strip()
        lt_code  = (r.get('lt_code') or '').strip().upper()
        try:
            fy_label = int(r.get('fy_label'))
            set_value = float(r.get('set_balance'))
        except Exception:
            skipped += 1
            continue

        emp = emp_by_code.get(emp_code)
        lt  = lt_by_code.get(lt_code)
        if not emp or not lt:
            skipped += 1
            continue

        bal = (LeaveBalance.query
               .filter_by(employee_id=emp.id, leave_type_id=lt.id, year=fy_label)
               .with_for_update(read=True)
               .first())
        if not bal:
            bal = LeaveBalance(employee_id=emp.id, leave_type_id=lt.id, year=fy_label, balance=0.0)
            db.session.add(bal)
            db.session.flush()

        prev = float(bal.balance or 0.0)
        new_bal = round(max(0.0, set_value), 2)

        # Audit entry – credit_delta=0 so future allocations add on top
        db.session.add(LeaveAllocationEntry(
            batch_id=batch.id,
            employee_id=emp.id,
            leave_type_id=lt.id,
            fy_label=fy_label,
            previous_balance=prev,
            new_balance=new_bal,
            credit_delta=0.0
        ))

        bal.balance = new_bal
        adjusted += 1

    db.session.commit()
    flash(f'✅ Manual set via CSV done: adjusted={adjusted}, skipped={skipped}', 'success')
    return redirect(url_for('leave.allocation_batches'))

from sqlalchemy import and_

def _revoke_set_balance_batch(batch):
    """
    For strategy='manual_set_balance' batches:
    Restore LeaveBalance.balance to the Entry.previous_balance for every entry.
    Returns (restored_count, created_missing_rows).
    """
    entries = LeaveAllocationEntry.query.filter_by(batch_id=batch.id).all()
    restored = 0
    created = 0

    for ent in entries:
        # Fetch (or create) the balance row for this (emp, lt, FY)
        bal = (LeaveBalance.query
               .filter_by(employee_id=ent.employee_id,
                          leave_type_id=ent.leave_type_id,
                          year=ent.fy_label)
               .with_for_update()
               .first())
        if not bal:
            bal = LeaveBalance(
                employee_id=ent.employee_id,
                leave_type_id=ent.leave_type_id,
                year=ent.fy_label,
                balance=0.0
            )
            db.session.add(bal)
            db.session.flush()
            created += 1

        # Restore to the exact previous balance we captured when setting
        bal.balance = float(ent.previous_balance or 0.0)
        restored += 1

    return restored, created


@leave_bp.post('/allocation/revoke-set/<int:batch_id>')
@login_required
def revoke_set_balance_batch(batch_id: int):
    """
    Revoke a 'manual_set_balance' batch by restoring the balances
    to the per-entry previous_balance.
    """
    batch = LeaveAllocationBatch.query.get_or_404(batch_id)

    if batch.is_revoked:
        flash('⚠️ This batch is already revoked.', 'warning')
        return redirect(url_for('leave.allocation_batches'))

    if (batch.strategy or '').lower() != 'manual_set_balance':
        flash('❌ This revoke endpoint is only for "manual_set_balance" batches.', 'danger')
        return redirect(url_for('leave.allocation_batches'))

    restored, created = _revoke_set_balance_batch(batch)

    batch.is_revoked = True
    batch.revoked_at = datetime.utcnow()
    batch.revoked_by_id = getattr(current_user, 'id', None)

    db.session.commit()
    flash(f"🧹 Reverted manual set-balance batch {batch.batch_key}. "
          f"Restored={restored}, created_missing_rows={created}.", "success")
    return redirect(url_for('leave.allocation_batches'))

# put this near your other Set Balance routes
from flask import Response
from io import StringIO
import csv

@leave_bp.get('/balance/set/template')
@login_required
def balance_set_template():
    """
    Downloadable CSV template for manual set balance.
    Columns: emp_code, lt_code, fy_label, set_balance
    Includes a few sample rows for convenience.
    """
    # FY label (Apr–Mar) using your helper
    _, _, fy_label = _fy_range_for_report()

    si = StringIO()
    w = csv.writer(si)
    w.writerow(['emp_code', 'lt_code', 'fy_label', 'set_balance'])

    # Optional sample rows (first 5 employees)
    sample_emps = Employee.query.order_by(Employee.employee_code).limit(5).all()
    for e in sample_emps:
        w.writerow([e.employee_code or '', 'EL', fy_label, 0])

    out = si.getvalue()
    filename = f"set_balance_template_{datetime.utcnow().strftime('%Y%m%d')}.csv"
    return Response(out, mimetype='text/csv',
                    headers={"Content-Disposition": f"attachment; filename={filename}"})