import os
import traceback
from flask import Flask, request, jsonify
import numpy as np
import joblib

app = Flask(__name__)

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
WEIGHTS_PATH = os.path.join(BASE_DIR, 'weights.pkl')
SCALER_PATH = os.path.join(BASE_DIR, 'scaler.pkl')

scaler = None
lstm1_w = None
lstm2_w = None
dense1_w = None
dense2_w = None

def sigmoid(x):
    return 1.0 / (1.0 + np.exp(-np.clip(x, -30, 30)))

def relu(x):
    return np.maximum(0, x)

def load_engine():
    global scaler, lstm1_w, lstm2_w, dense1_w, dense2_w
    if scaler is None and os.path.exists(SCALER_PATH):
        scaler = joblib.load(SCALER_PATH)

    if lstm1_w is None and os.path.exists(WEIGHTS_PATH):
        w = joblib.load(WEIGHTS_PATH)
        lstm1_w = w['lstm1']
        lstm2_w = w['lstm2']
        dense1_w = w['dense1']
        dense2_w = w['dense2']

    return scaler

def lstm_step(x, kernel, recurrent_kernel, bias, return_sequences=False):
    units = recurrent_kernel.shape[0]
    timesteps = x.shape[0]

    W_i, W_f, W_c, W_o = np.split(kernel, 4, axis=1)
    U_i, U_f, U_c, U_o = np.split(recurrent_kernel, 4, axis=1)
    b_i, b_f, b_c, b_o = np.split(bias, 4, axis=0)

    h_t = np.zeros(units, dtype=np.float32)
    c_t = np.zeros(units, dtype=np.float32)
    seq = []

    for t in range(timesteps):
        x_t = x[t]
        i = sigmoid(np.dot(x_t, W_i) + np.dot(h_t, U_i) + b_i)
        f = sigmoid(np.dot(x_t, W_f) + np.dot(h_t, U_f) + b_f)
        c_cand = np.tanh(np.dot(x_t, W_c) + np.dot(h_t, U_c) + b_c)
        c_t = f * c_t + i * c_cand
        o = sigmoid(np.dot(x_t, W_o) + np.dot(h_t, U_o) + b_o)
        h_t = o * np.tanh(c_t)
        if return_sequences:
            seq.append(h_t)

    return np.array(seq) if return_sequences else h_t

def run_inference(sequence_data):
    s = load_engine()
    scaled = s.transform(np.array(sequence_data, dtype=np.float32).reshape(-1, 1)).flatten()
    x = scaled.reshape(len(sequence_data), 1)

    # 1. LSTM 1 (return_sequences=True)
    out1 = lstm_step(x, lstm1_w[0], lstm1_w[1], lstm1_w[2], return_sequences=True)

    # 2. LSTM 2 (return_sequences=False)
    out2 = lstm_step(out1, lstm2_w[0], lstm2_w[1], lstm2_w[2], return_sequences=False)

    # 3. Dense 1 (ReLU)
    out3 = relu(np.dot(out2, dense1_w[0]) + dense1_w[1])

    # 4. Dense 2 (Linear)
    out4 = np.dot(out3, dense2_w[0]) + dense2_w[1]

    # Inverse Transform
    final_result = s.inverse_transform(np.array(out4).reshape(-1, 1))
    return final_result.flatten().tolist()

@app.route('/', methods=['GET'])
def index():
    return jsonify({
        "status": "online",
        "message": "Flask ML API Apotik Limas is running (NumPy Engine)"
    }), 200

@app.route('/predict', methods=['POST'])
def predict():
    try:
        load_engine()
        req_data = request.get_json(force=True)
        raw_data = req_data.get('data')

        if not raw_data or len(raw_data) != 6:
            return jsonify({
                "status": "error",
                "message": "Input 'data' harus berisi tepat 6 angka."
            }), 400

        pred = run_inference(raw_data)

        return jsonify({
            "status": "success",
            "prediction": pred
        }), 200

    except Exception as e:
        return jsonify({
            "status": "error",
            "error_type": type(e).__name__,
            "message": str(e),
            "traceback": traceback.format_exc()
        }), 500

if __name__ == '__main__':
    app.run(host='0.0.0.0', port=5000)