import os
from flask import Flask, request, jsonify
from flask_cors import CORS
import tensorflow as tf
from tensorflow.keras.models import load_model
from PIL import Image
import numpy as np

app = Flask(__name__)
CORS(app)

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
MODEL_PATH = os.path.join(BASE_DIR, "skin_model.keras")

print("=" * 50)
print("Sedang memuat model...")
model = load_model(MODEL_PATH)
print("Model berhasil dimuat!")
print("INPUT MODEL :", model.input_shape)
print("OUTPUT MODEL:", model.output_shape)
print("=" * 50)

# Keras ImageDataGenerator/image_dataset_from_directory secara default
# mengurutkan kelas berdasarkan abjad (A-Z) saat training:
# Index 0: Acanthosis Nigricans
# Index 1: Normal
class_names = ["Acanthosis Nigricans", "Normal"]

@app.route("/", methods=["GET"])
def home():
    return jsonify({
        "success": True,
        "message": "Skin Prediction API berjalan",
        "model": "MobileNetV2",
        "classes": class_names
    })

@app.route("/predict", methods=["POST"])
def predict():
    try:
        if "image" not in request.files:
            return jsonify({"success": False, "error": "File gambar tidak ditemukan."}), 400

        file = request.files["image"]
        if file.filename == "":
            return jsonify({"success": False, "error": "Tidak ada file yang dipilih."}), 400

        # 1. Buka dan konversi gambar
        image = Image.open(file).convert("RGB")
        image = image.resize((224, 224))

        # 2. Convert ke Numpy Array (piksel mentah 0-255)
        image_array = np.array(image, dtype=np.float32)

        # 3. Tanpa normalisasi manual/preprocess_input karena model 
        #    sudah memiliki TrueDivide & Subtract secara internal!

        # 4. Tambahkan batch dimension (1, 224, 224, 3)
        image_array = np.expand_dims(image_array, axis=0)

        # 5. Prediksi
        prediction = model.predict(image_array, verbose=0)

        print("\n" + "=" * 50)
        print("RAW PREDICTION :", prediction)
        print("=" * 50 + "\n")

        # Handle jika output 1 neuron (Binary Classification / Sigmoid)
        if prediction.shape[-1] == 1:
            raw_val = float(prediction[0][0])
            
            # Jika raw_val >= 0.5 maka Normal (Index 1), jika < 0.5 maka Acanthosis (Index 0)
            if raw_val >= 0.5:
                predicted_class = "Normal"
                confidence = raw_val * 100
            else:
                predicted_class = "Acanthosis Nigricans"
                confidence = (1.0 - raw_val) * 100

            probabilities = {
                "Acanthosis Nigricans": round((1.0 - raw_val) * 100, 2),
                "Normal": round(raw_val * 100, 2)
            }

        # Handle jika output 2 neuron (Categorical / Softmax)
        else:
            predicted_index = int(np.argmax(prediction[0]))
            predicted_class = class_names[predicted_index]
            confidence = float(prediction[0][predicted_index] * 100)

            probabilities = {}
            for idx, name in enumerate(class_names):
                probabilities[name] = round(float(prediction[0][idx]) * 100, 2)

        description = (
            f"Citra kulit diklasifikasikan sebagai {predicted_class} "
            f"oleh model CNN MobileNetV2."
        )

        return jsonify({
            "success": True,
            "prediction": predicted_class,
            "label": predicted_class,
            "class": predicted_class,
            "confidence": round(confidence, 2),
            "model": "MobileNetV2",
            "description": description,
            "probabilities": probabilities
        })

    except Exception as e:
        print("ERROR PREDIKSI:", str(e))
        return jsonify({"success": False, "error": str(e)}), 500

if __name__ == "__main__":
    app.run(host="127.0.0.1", port=5000, debug=True)