import os
import json
import base64
import secrets
from datetime import datetime
from decimal import Decimal, InvalidOperation
from io import BytesIO
import logging

import mysql.connector
import qrcode
import requests
from flask import Flask, jsonify, render_template, request, send_from_directory, url_for
from mysql.connector import Error, pooling
from PIL import Image

# Configure logging for cPanel
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)

# Load .env only if it exists (local development)
if os.path.exists(os.path.join(os.path.dirname(__file__), '.env')):
    from dotenv import load_dotenv
    load_dotenv()
    logger.info("Loaded .env file for local development")
else:
    logger.info("Running in production mode - using cPanel environment variables")

app = Flask(__name__)
app.config["SECRET_KEY"] = os.getenv("SECRET_KEY") or secrets.token_hex(32)
app.config["QR_FOLDER"] = os.path.join(app.root_path, "static", "qrcodes")
os.makedirs(app.config["QR_FOLDER"], exist_ok=True)

# -------------------------------------------------------------------
# Configuration
# -------------------------------------------------------------------

MPESA_ENV = os.getenv("MPESA_ENV", "sandbox").lower()
ALLOW_SIMULATION_QR = os.getenv("ALLOW_SIMULATION_QR", "false").lower() == "true"

if MPESA_ENV == "production":
    MPESA_BASE_URL = "https://api.safaricom.co.ke"
else:
    MPESA_BASE_URL = "https://sandbox.safaricom.co.ke"

MPESA_CONSUMER_KEY = os.getenv("MPESA_CONSUMER_KEY", "").strip()
MPESA_CONSUMER_SECRET = os.getenv("MPESA_CONSUMER_SECRET", "").strip()
MPESA_SHORTCODE = os.getenv("MPESA_SHORTCODE", "174379")
MPESA_PASSKEY = os.getenv("MPESA_PASSKEY", "").strip()
MPESA_CALLBACK_URL = os.getenv("MPESA_CALLBACK_URL", "https://easternlight.co.ke/jayjaympesa/api/mpesa/stk-callback")
MPESA_MERCHANT_NAME = os.getenv("MPESA_MERCHANT_NAME", "My POS Shop")
MPESA_QR_TRX_CODE = os.getenv("MPESA_QR_TRX_CODE", "BG")

# -------------------------------------------------------------------
# Daraja authentication and API helpers (MUST BE DEFINED BEFORE USE)
# -------------------------------------------------------------------

def daraja_is_configured():
    """Check whether live/sandbox Daraja credentials are configured."""
    if not MPESA_CONSUMER_KEY or len(MPESA_CONSUMER_KEY) < 20:
        logger.warning(f"Consumer Key invalid or too short: {MPESA_CONSUMER_KEY[:10] if MPESA_CONSUMER_KEY else 'EMPTY'}...")
        return False
    
    if not MPESA_CONSUMER_SECRET or len(MPESA_CONSUMER_SECRET) < 20:
        logger.warning(f"Consumer Secret invalid or too short: {MPESA_CONSUMER_SECRET[:10] if MPESA_CONSUMER_SECRET else 'EMPTY'}...")
        return False
    
    logger.info("Daraja credentials appear valid")
    return True


def get_daraja_access_token():
    """Request an OAuth access token from Daraja."""
    if not daraja_is_configured():
        raise ValueError(
            "Daraja credentials are not configured. "
            "Please check your Consumer Key and Consumer Secret in cPanel environment variables."
        )

    url = f"{MPESA_BASE_URL}/oauth/v1/generate?grant_type=client_credentials"

    try:
        logger.info(f"Requesting Daraja token from: {url}")
        response = requests.get(
            url,
            auth=(MPESA_CONSUMER_KEY, MPESA_CONSUMER_SECRET),
            timeout=30
        )
        
        logger.info(f"Daraja token response status: {response.status_code}")
        logger.info(f"Daraja token response: {response.text[:200]}")
        
        response.raise_for_status()

        response_json = response.json()
        access_token = response_json.get("access_token")
        
        if not access_token:
            logger.error(f"Daraja did not return an access token. Full response: {response_json}")
            raise ValueError("Daraja did not return an access token.")

        logger.info("Daraja access token obtained successfully")
        return access_token
    
    except requests.exceptions.HTTPError as err:
        logger.error(f"Daraja HTTP error: {err}")
        logger.error(f"Response: {err.response.text if hasattr(err, 'response') else 'No response'}")
        raise ValueError(f"Daraja authentication failed: {err.response.text if hasattr(err, 'response') else str(err)}")
    
    except requests.RequestException as err:
        logger.error(f"Daraja token request failed: {err}")
        raise ValueError(f"Daraja connection failed: {str(err)}")


def normalise_phone_number(phone):
    """Convert Kenyan numbers to 2547XXXXXXXX."""
    phone = str(phone).strip().replace(" ", "").replace("-", "")

    if phone.startswith("+"):
        phone = phone[1:]

    if phone.startswith("0") and len(phone) == 10:
        phone = "254" + phone[1:]

    elif phone.startswith("7") and len(phone) == 9:
        phone = "254" + phone

    if not (phone.startswith("2547") and len(phone) == 12 and phone.isdigit()):
        raise ValueError("Use a valid Kenyan phone number, e.g. 0712345678.")

    return phone


# -------------------------------------------------------------------
# Log configuration AFTER helper functions are defined
# -------------------------------------------------------------------

logger.info(f"MPESA_ENV: {MPESA_ENV}")
logger.info(f"MPESA_BASE_URL: {MPESA_BASE_URL}")
logger.info(f"MPESA_SHORTCODE: {MPESA_SHORTCODE}")
logger.info(f"MPESA_CALLBACK_URL: {MPESA_CALLBACK_URL}")
logger.info(f"Daraja configured: {daraja_is_configured()}")

# -------------------------------------------------------------------
# Database Connection Pool (Optimized for cPanel)
# -------------------------------------------------------------------

db_pool = None

def init_db_pool():
    """Initialize database connection pool on application startup."""
    global db_pool
    
    try:
        db_pool = pooling.MySQLConnectionPool(
            pool_name=os.getenv("MYSQL_POOL_NAME", "elpool"),
            pool_size=int(os.getenv("MYSQL_POOL_SIZE", 5)),
            pool_reset_session=True,
            host=os.getenv("MYSQL_HOST", "localhost"),
            port=int(os.getenv("MYSQL_PORT", 3306)),
            user=os.getenv("MYSQL_USER"),
            password=os.getenv("MYSQL_PASSWORD"),
            database=os.getenv("MYSQL_DATABASE"),
            autocommit=False,
            connection_timeout=30,
            use_pure=True
        )
        logger.info("Database pool initialized successfully")
    except Error as err:
        logger.error(f"Database pool initialization failed: {err}")
        raise

def get_db_connection():
    """Get a connection from the pool."""
    global db_pool
    if db_pool is None:
        init_db_pool()
    
    try:
        return db_pool.get_connection()
    except Error as err:
        logger.error(f"Failed to get database connection: {err}")
        raise

# Initialize pool when app starts
with app.app_context():
    init_db_pool()

# -------------------------------------------------------------------
# Database helpers
# -------------------------------------------------------------------

def fetch_one(query, params=None):
    """Execute a SELECT query and return one dictionary row."""
    connection = None
    cursor = None
    
    try:
        connection = get_db_connection()
        cursor = connection.cursor(dictionary=True)
        cursor.execute(query, params or ())
        return cursor.fetchone()
    except Error as err:
        logger.error(f"Database fetch_one error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


def fetch_all(query, params=None):
    """Execute a SELECT query and return all dictionary rows."""
    connection = None
    cursor = None
    
    try:
        connection = get_db_connection()
        cursor = connection.cursor(dictionary=True)
        cursor.execute(query, params or ())
        return cursor.fetchall()
    except Error as err:
        logger.error(f"Database fetch_all error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


def create_order(items):
    """Create an order and its items. Returns the newly inserted order dictionary."""
    order_reference = (
        f"POS-{datetime.now().strftime('%Y%m%d-%H%M%S')}-"
        f"{secrets.token_hex(3).upper()}"
    )

    total_amount = sum(
        Decimal(str(item["price"])) * int(item.get("quantity", 1))
        for item in items
    )

    connection = None
    cursor = None

    try:
        connection = get_db_connection()
        cursor = connection.cursor(dictionary=True)

        cursor.execute(
            """
            INSERT INTO orders (order_reference, total_amount, status)
            VALUES (%s, %s, 'PENDING')
            """,
            (order_reference, total_amount)
        )
        order_id = cursor.lastrowid

        item_sql = """
            INSERT INTO order_items
            (order_id, item_name, quantity, unit_price, line_total)
            VALUES (%s, %s, %s, %s, %s)
        """

        for index, item in enumerate(items, start=1):
            item_name = str(item.get("name") or f"Item {index}").strip()
            quantity = int(item.get("quantity", 1))
            price = Decimal(str(item["price"]))
            line_total = price * quantity

            cursor.execute(
                item_sql,
                (order_id, item_name, quantity, price, line_total)
            )

        connection.commit()

        cursor.execute(
            "SELECT * FROM orders WHERE id = %s",
            (order_id,)
        )
        return cursor.fetchone()

    except Error as err:
        if connection:
            connection.rollback()
        logger.error(f"Database create_order error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


def update_order_qr(order_id, filename):
    """Save the generated QR file name for the order."""
    connection = None
    cursor = None

    try:
        connection = get_db_connection()
        cursor = connection.cursor()

        cursor.execute(
            "UPDATE orders SET qr_filename = %s WHERE id = %s",
            (filename, order_id)
        )
        connection.commit()
    except Error as err:
        if connection:
            connection.rollback()
        logger.error(f"Database update_order_qr error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


def update_order_stk_ids(order_id, merchant_request_id, checkout_request_id):
    """Store Daraja identifiers after an accepted STK request."""
    connection = None
    cursor = None

    try:
        connection = get_db_connection()
        cursor = connection.cursor()

        cursor.execute(
            """
            UPDATE orders
            SET merchant_request_id = %s,
                checkout_request_id = %s,
                status = 'PROCESSING'
            WHERE id = %s
            """,
            (merchant_request_id, checkout_request_id, order_id)
        )
        connection.commit()
    except Error as err:
        if connection:
            connection.rollback()
        logger.error(f"Database update_order_stk_ids error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


def insert_callback_log(callback_type, checkout_request_id, payload):
    """Store every received callback, even invalid/unknown callbacks."""
    connection = None
    cursor = None

    try:
        connection = get_db_connection()
        cursor = connection.cursor()

        cursor.execute(
            """
            INSERT INTO mpesa_callback_logs
            (callback_type, checkout_request_id, payload)
            VALUES (%s, %s, %s)
            """,
            (callback_type, checkout_request_id, json.dumps(payload))
        )
        connection.commit()
    except Error as err:
        if connection:
            connection.rollback()
        logger.error(f"Database insert_callback_log error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


# -------------------------------------------------------------------
# QR Generation Functions
# -------------------------------------------------------------------

def generate_local_qr(order):
    """Development-only QR fallback."""
    payload = (
        f"SIMULATION|ORDER={order['order_reference']}|"
        f"AMOUNT={order['total_amount']}|"
        f"SHORTCODE={MPESA_SHORTCODE}"
    )

    qr = qrcode.QRCode(
        version=None,
        error_correction=qrcode.constants.ERROR_CORRECT_M,
        box_size=10,
        border=4
    )
    qr.add_data(payload)
    qr.make(fit=True)

    image = qr.make_image(fill_color="black", back_color="white")

    filename = f"sim_{order['order_reference']}.png"
    filepath = os.path.join(app.config["QR_FOLDER"], filename)
    image.save(filepath)

    update_order_qr(order["id"], filename)

    return {
        "success": True,
        "simulation": True,
        "qr_url": f"/jayjaympesa/static/qrcodes/{filename}",
        "order_reference": order["order_reference"],
        "amount": str(order["total_amount"]),
        "message": "Simulation QR created. It cannot receive real M-PESA payments."
    }
    

def generate_daraja_qr(order):
    """Generate a real Dynamic QR through Daraja."""
    if not daraja_is_configured():
        if ALLOW_SIMULATION_QR and MPESA_ENV != "production":
            logger.warning("Daraja not configured, using simulation QR")
            return generate_local_qr(order)
        
        raise ValueError(
            "Daraja credentials are not configured. "
            "Configure Consumer Key, Consumer Secret, and verify QR API access."
        )

    try:
        logger.info("Getting Daraja access token...")
        token = get_daraja_access_token()
        logger.info(f"Token obtained: {token[:20]}...")

        payload = {
            "MerchantName": MPESA_MERCHANT_NAME,
            "RefNo": order["order_reference"],
            "Amount": int(Decimal(str(order["total_amount"]))),
            "TrxCode": MPESA_QR_TRX_CODE,
            "CPI": MPESA_SHORTCODE,
            "Size": "300"
        }

        logger.info(f"Generating QR with payload: {payload}")
        
        headers = {
            "Authorization": f"Bearer {token}",
            "Content-Type": "application/json"
        }
        
        logger.info(f"Request headers: Authorization=Bearer {token[:20]}...")
        
        response = requests.post(
            f"{MPESA_BASE_URL}/mpesa/qrcode/v1/generate",
            json=payload,
            headers=headers,
            timeout=30
        )

        logger.info(f"QR API Response Status: {response.status_code}")
        logger.info(f"QR API Response: {response.text[:500]}")

        if response.status_code != 200:
            error_msg = f"Daraja QR API returned status {response.status_code}"
            try:
                error_data = response.json()
                error_msg = error_data.get("errorMessage", error_data.get("ResponseDescription", error_msg))
            except:
                error_msg = f"{error_msg}: {response.text[:200]}"
            
            logger.error(f"QR generation failed: {error_msg}")
            raise ValueError(error_msg)

        response_data = response.json()

        if "QRCode" not in response_data:
            raise ValueError(
                response_data.get(
                    "errorMessage",
                    response_data.get("ResponseDescription", "QR generation failed - no QRCode in response.")
                )
            )

        image_data = base64.b64decode(response_data["QRCode"])
        image = Image.open(BytesIO(image_data))

        filename = f"daraja_{order['order_reference']}.png"
        filepath = os.path.join(app.config["QR_FOLDER"], filename)
        image.save(filepath)

        update_order_qr(order["id"], filename)

        return {
            "success": True,
            "simulation": False,
            "qr_url": f"/jayjaympesa/static/qrcodes/{filename}",
            "order_reference": order["order_reference"],
            "amount": str(order["total_amount"]),
            "message": "Dynamic M-PESA QR code generated."
        }
    
    except ValueError as err:
        logger.error(f"Daraja QR generation failed: {err}", exc_info=True)
        raise
    
    except Exception as err:
        logger.error(f"Unexpected error in QR generation: {err}", exc_info=True)
        raise ValueError(f"QR generation failed: {str(err)}")


def initiate_stk_push(order, phone_number):
    """Send an STK push."""
    if not daraja_is_configured():
        raise ValueError(
            "Daraja credentials are missing. Add Consumer Key and Secret to cPanel environment variables."
        )

    if not MPESA_PASSKEY or MPESA_PASSKEY.startswith("PASTE_"):
        raise ValueError(
            "MPESA_PASSKEY is missing. Add the Daraja passkey to cPanel environment variables."
        )

    if not MPESA_CALLBACK_URL or not MPESA_CALLBACK_URL.startswith("https://"):
        raise ValueError(
            "MPESA_CALLBACK_URL must be a publicly reachable HTTPS URL."
        )

    phone_number = normalise_phone_number(phone_number)
    token = get_daraja_access_token()

    timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
    password_string = f"{MPESA_SHORTCODE}{MPESA_PASSKEY}{timestamp}"
    password = base64.b64encode(password_string.encode()).decode()

    payload = {
        "BusinessShortCode": MPESA_SHORTCODE,
        "Password": password,
        "Timestamp": timestamp,
        "TransactionType": "CustomerPayBillOnline",
        "Amount": int(Decimal(str(order["total_amount"]))),
        "PartyA": phone_number,
        "PartyB": MPESA_SHORTCODE,
        "PhoneNumber": phone_number,
        "CallBackURL": MPESA_CALLBACK_URL,
        "AccountReference": order["order_reference"],
        "TransactionDesc": f"Payment for {order['order_reference']}"
    }

    logger.info(f"STK Push payload: {payload}")

    try:
        response = requests.post(
            f"{MPESA_BASE_URL}/mpesa/stkpush/v1/processrequest",
            json=payload,
            headers={
                "Authorization": f"Bearer {token}",
                "Content-Type": "application/json"
            },
            timeout=30
        )

        logger.info(f"STK Push Response Status: {response.status_code}")
        logger.info(f"STK Push Response: {response.text[:500]}")

        if response.status_code != 200:
            error_msg = f"STK API returned status {response.status_code}"
            try:
                error_data = response.json()
                error_msg = error_data.get("errorMessage", error_data.get("ResponseDescription", error_msg))
            except:
                error_msg = f"{error_msg}: {response.text[:200]}"
            raise ValueError(error_msg)

        response_data = response.json()

        checkout_request_id = response_data.get("CheckoutRequestID")
        merchant_request_id = response_data.get("MerchantRequestID")

        if not checkout_request_id:
            raise ValueError(
                response_data.get(
                    "ResponseDescription",
                    "Daraja did not return CheckoutRequestID."
                )
            )

        update_order_stk_ids(
            order["id"],
            merchant_request_id,
            checkout_request_id
        )

        return {
            "success": True,
            "message": response_data.get(
                "CustomerMessage",
                "STK payment prompt sent. Enter your M-PESA PIN."
            ),
            "checkout_request_id": checkout_request_id,
            "order_reference": order["order_reference"]
        }
    
    except requests.RequestException as err:
        logger.error(f"STK push request failed: {err}", exc_info=True)
        raise ValueError(f"STK push failed: {str(err)}")


# -------------------------------------------------------------------
# Callback processing
# -------------------------------------------------------------------

def callback_items_to_dict(callback_metadata):
    """Convert Daraja CallbackMetadata Item array into a Python dictionary."""
    values = {}

    if not callback_metadata:
        return values

    for item in callback_metadata.get("Item", []):
        name = item.get("Name")
        value = item.get("Value")

        if name:
            values[name] = value

    return values


def process_stk_callback(payload):
    """Validate and save an STK callback."""
    stk_callback = payload.get("Body", {}).get("stkCallback", {})
    checkout_request_id = stk_callback.get("CheckoutRequestID")
    result_code = stk_callback.get("ResultCode")
    result_desc = stk_callback.get("ResultDesc", "No result description.")

    insert_callback_log("STK_CALLBACK", checkout_request_id, payload)

    if not checkout_request_id:
        return False, "Missing CheckoutRequestID in callback."

    order = fetch_one(
        "SELECT * FROM orders WHERE checkout_request_id = %s",
        (checkout_request_id,)
    )

    if not order:
        return False, "Callback received for an unknown order."

    metadata = callback_items_to_dict(stk_callback.get("CallbackMetadata"))
    callback_json = json.dumps(payload)

    connection = None
    cursor = None

    try:
        connection = get_db_connection()
        cursor = connection.cursor()

        if result_code != 0:
            status = "CANCELLED" if result_code == 1032 else "FAILED"

            cursor.execute(
                """
                UPDATE orders
                SET status = %s,
                    result_code = %s,
                    result_description = %s,
                    raw_callback = %s
                WHERE id = %s
                """,
                (status, result_code, result_desc, callback_json, order["id"])
            )
            connection.commit()
            return True, f"Order marked {status}."

        amount_received = metadata.get("Amount")
        receipt_number = metadata.get("MpesaReceiptNumber")
        transaction_date = metadata.get("TransactionDate")
        phone_number = metadata.get("PhoneNumber")

        if amount_received is None:
            raise ValueError("Successful callback has no Amount.")

        expected_amount = Decimal(str(order["total_amount"]))
        received_amount = Decimal(str(amount_received))

        if received_amount != expected_amount:
            cursor.execute(
                """
                UPDATE orders
                SET status = 'FAILED',
                    result_code = %s,
                    result_description = %s,
                    raw_callback = %s
                WHERE id = %s
                """,
                (result_code, f"Amount mismatch. Expected {expected_amount}; received {received_amount}.", callback_json, order["id"])
            )
            connection.commit()
            return False, "Amount mismatch; order was not marked paid."

        cursor.execute(
            """
            UPDATE orders
            SET status = 'PAID',
                customer_phone = %s,
                mpesa_receipt_number = %s,
                mpesa_transaction_date = %s,
                result_code = %s,
                result_description = %s,
                raw_callback = %s,
                paid_at = NOW()
            WHERE id = %s
            """,
            (str(phone_number) if phone_number else None, receipt_number, str(transaction_date) if transaction_date else None, result_code, result_desc, callback_json, order["id"])
        )

        connection.commit()
        return True, "Payment saved successfully."

    except Error as err:
        if connection:
            connection.rollback()
        logger.error(f"Database process_stk_callback error: {err}")
        raise
    finally:
        if cursor:
            cursor.close()
        if connection:
            connection.close()


# -------------------------------------------------------------------
# Web pages
# -------------------------------------------------------------------

@app.route("/")
def index():
    return render_template("index.html")


@app.route("/transactions")
def transactions_page():
    """View recent transactions in the browser."""
    orders = fetch_all(
        """
        SELECT id, order_reference, total_amount, status,
               customer_phone, mpesa_receipt_number,
               created_at, paid_at
        FROM orders
        ORDER BY id DESC
        LIMIT 100
        """
    )
    return render_template("transactions.html", orders=orders)


# -------------------------------------------------------------------
# Frontend API routes
# -------------------------------------------------------------------

@app.route("/api/orders", methods=["POST"])
def api_create_order():
    """Create a new order."""
    data = request.get_json(silent=True) or {}
    items = data.get("items", [])

    if not isinstance(items, list) or len(items) == 0:
        return jsonify({
            "success": False,
            "error": "Add at least one item before creating an order."
        }), 400

    cleaned_items = []

    try:
        for index, item in enumerate(items, start=1):
            name = str(item.get("name") or f"Item {index}").strip()
            quantity = int(item.get("quantity", 1))
            price = Decimal(str(item.get("price")))

            if not name:
                raise ValueError("Every item must have a name.")

            if quantity < 1:
                raise ValueError("Quantity must be at least 1.")

            if price <= 0:
                raise ValueError("Item price must be greater than zero.")

            cleaned_items.append({
                "name": name,
                "quantity": quantity,
                "price": str(price)
            })

        order = create_order(cleaned_items)

        return jsonify({
            "success": True,
            "order_id": order["id"],
            "order_reference": order["order_reference"],
            "amount": str(order["total_amount"]),
            "status": order["status"]
        })

    except (ValueError, InvalidOperation) as error:
        logger.error(f"Order creation error: {error}")
        return jsonify({
            "success": False,
            "error": str(error)
        }), 400

    except Error as error:
        logger.error(f"Database error in api_create_order: {error}")
        return jsonify({
            "success": False,
            "error": f"Database error: {error}"
        }), 500


@app.route("/api/orders/<int:order_id>/qr", methods=["POST"])
def api_generate_order_qr(order_id):
    """Generate and save a QR image for a previously created order."""
    logger.info("=" * 60)
    logger.info(f"QR GENERATION REQUEST - Order ID: {order_id}")
    logger.info("=" * 60)
    
    try:
        order = fetch_one("SELECT * FROM orders WHERE id = %s", (order_id,))

        if not order:
            logger.warning(f"Order {order_id} not found")
            return jsonify({"success": False, "error": "Order not found."}), 404

        if order["status"] == "PAID":
            logger.warning(f"Order {order_id} already paid")
            return jsonify({"success": False, "error": "This order is already paid."}), 400

        logger.info(f"Order status: {order['status']}")
        result = generate_daraja_qr(order)
        
        logger.info(f"QR generated successfully: {result}")
        logger.info("=" * 60)
        return jsonify(result)

    except Exception as error:
        logger.error(f"QR generation error for order {order_id}: {error}", exc_info=True)
        logger.error(f"Error type: {type(error).__name__}")
        logger.info("=" * 60)
        return jsonify({
            "success": False,
            "error": str(error),
            "error_type": type(error).__name__
        }), 500


@app.route("/api/orders/<int:order_id>/stk-push", methods=["POST"])
def api_stk_push(order_id):
    """Send STK push for an order."""
    data = request.get_json(silent=True) or {}
    phone_number = data.get("phone_number", "")

    logger.info(f"STK Push request for order {order_id}, phone: {phone_number}")

    try:
        order = fetch_one("SELECT * FROM orders WHERE id = %s", (order_id,))

        if not order:
            return jsonify({"success": False, "error": "Order not found."}), 404

        if order["status"] == "PAID":
            return jsonify({"success": False, "error": "This order has already been paid."}), 400

        result = initiate_stk_push(order, phone_number)
        return jsonify(result)

    except (ValueError, requests.RequestException) as error:
        logger.error(f"STK push error: {error}")
        return jsonify({"success": False, "error": str(error)}), 400

    except Error as error:
        logger.error(f"Database error in api_stk_push: {error}")
        return jsonify({"success": False, "error": f"Database error: {error}"}), 500


@app.route("/api/orders/<int:order_id>", methods=["GET"])
def api_order_status(order_id):
    """Check order payment status."""
    order = fetch_one(
        """
        SELECT id, order_reference, total_amount, status,
               customer_phone, mpesa_receipt_number,
               result_code, result_description,
               created_at, paid_at
        FROM orders
        WHERE id = %s
        """,
        (order_id,)
    )
    
    if not order:
        return jsonify({"success": False, "error": "Order not found."}), 404

    for key in ("total_amount", "created_at", "paid_at"):
        if order.get(key) is not None:
            order[key] = str(order[key])

    return jsonify({"success": True, "order": order})


# -------------------------------------------------------------------
# M-PESA callback endpoint
# -------------------------------------------------------------------

@app.route("/api/mpesa/stk-callback", methods=["POST"])
def mpesa_stk_callback():
    """Safaricom posts the STK result to this route."""
    logger.info(f"Received STK callback")
    logger.info(f"Request data: {request.get_data()}")
    
    payload = request.get_json(silent=True)

    if not payload:
        logger.error("Invalid JSON callback payload")
        return jsonify({
            "ResultCode": 1,
            "ResultDesc": "Invalid JSON callback payload."
        }), 400

    try:
        success, message = process_stk_callback(payload)
        logger.info(f"Callback processed: {success}, {message}")

        return jsonify({
            "ResultCode": 0,
            "ResultDesc": "Callback received."
        }), 200

    except Exception as error:
        logger.exception("STK callback processing failed: %s", error)

        return jsonify({
            "ResultCode": 1,
            "ResultDesc": "Callback received but processing failed."
        }), 500


# -------------------------------------------------------------------
# Debug routes
# -------------------------------------------------------------------

@app.route("/test")
def test_route():
    """Simple test route to verify Flask is working"""
    return jsonify({
        "status": "ok",
        "message": "Flask is running!",
        "timestamp": datetime.now().isoformat()
    })

@app.route("/debug/routes")
def debug_routes():
    """Show all registered routes"""
    routes = []
    for rule in app.url_map.iter_rules():
        methods = ','.join(sorted(rule.methods - {'HEAD', 'OPTIONS'}))
        if methods:
            routes.append({
                "route": str(rule),
                "methods": methods,
                "endpoint": rule.endpoint
            })
    
    return jsonify({"total_routes": len(routes), "routes": routes})

@app.route("/debug/db-test")
def debug_db_test():
    """Test database connection"""
    try:
        conn = get_db_connection()
        cursor = conn.cursor()
        cursor.execute("SELECT 1 as test")
        result = cursor.fetchone()
        cursor.close()
        conn.close()
        
        return jsonify({
            "success": True,
            "message": "Database connection OK",
            "test_result": result
        })
    except Exception as e:
        return jsonify({
            "success": False,
            "error": str(e),
            "type": type(e).__name__
        }), 500

@app.route("/debug/check-credentials")
def debug_check_credentials():
    """Check if Daraja credentials are loaded"""
    return jsonify({
        "MPESA_ENV": os.getenv("MPESA_ENV", "NOT SET"),
        "MPESA_BASE_URL": MPESA_BASE_URL,
        "MPESA_CONSUMER_KEY": "SET" if MPESA_CONSUMER_KEY and len(MPESA_CONSUMER_KEY) > 20 else "NOT SET or INVALID",
        "MPESA_CONSUMER_SECRET": "SET" if MPESA_CONSUMER_SECRET and len(MPESA_CONSUMER_SECRET) > 20 else "NOT SET or INVALID",
        "MPESA_SHORTCODE": MPESA_SHORTCODE,
        "MPESA_PASSKEY": "SET" if MPESA_PASSKEY and not MPESA_PASSKEY.startswith("PASTE_") else "NOT SET or INVALID",
        "MPESA_CALLBACK_URL": MPESA_CALLBACK_URL,
        "daraja_configured": daraja_is_configured()
    })


# -------------------------------------------------------------------
# cPanel/Production WSGI entry point
# -------------------------------------------------------------------

if __name__ == "__main__":
    print("=" * 65)
    print("M-PESA QR POS SYSTEM WITH MYSQL AND CALLBACK SUPPORT")
    print("=" * 65)
    print("Open: http://127.0.0.1:5000/jayjaympesa/")
    print("Transactions: http://127.0.0.1:5000/jayjaympesa/transactions")
    print("Callback URL:", MPESA_CALLBACK_URL or "Not configured")
    print("=" * 65)

    app.run(debug=False, host="0.0.0.0", port=5000)