package com.lemenzo.gallery.editor;

import android.graphics.Bitmap;
import android.graphics.Canvas;
import android.graphics.Color;
import android.graphics.Gainmap;
import android.graphics.Matrix;
import android.graphics.Paint;
import android.graphics.RadialGradient;
import android.graphics.RectF;
import android.graphics.Shader;

import java.util.List;

/** Replays non-destructive edit state identically for preview and export. */
public final class EditPipeline {
    public static final int ROTATE_LEFT = 1;
    public static final int ROTATE_RIGHT = 2;
    public static final int FLIP_HORIZONTAL = 3;

    private EditPipeline() {}

    public static Bitmap render(Bitmap source, List<Integer> operations, float cropAspectRatio,
            RectF freeCrop, float straighten, float perspectiveX, float perspectiveY,
            float brightness, float contrast, float saturation, float warmth,
            float exposure, float tint, float highlights, float shadows, float sharpness,
            float vignette, int filter) {
        if (source == null) throw new IllegalArgumentException("source == null");

        Matrix matrix = new Matrix();
        int width = source.getWidth();
        int height = source.getHeight();

        if (operations != null) {
            for (Integer operation : operations) {
                if (operation == null) continue;
                Matrix op = new Matrix();
                if (operation == ROTATE_LEFT || operation == ROTATE_RIGHT) {
                    float degrees = operation == ROTATE_LEFT ? -90f : 90f;
                    op.setRotate(degrees);
                    RectF bounds = new RectF(0f, 0f, width, height);
                    op.mapRect(bounds);
                    op.postTranslate(-bounds.left, -bounds.top);
                    matrix.postConcat(op);
                    width = Math.max(1, Math.round(bounds.width()));
                    height = Math.max(1, Math.round(bounds.height()));
                } else if (operation == FLIP_HORIZONTAL) {
                    op.setScale(-1f, 1f);
                    op.postTranslate(width, 0f);
                    matrix.postConcat(op);
                }
            }
        }

        if (Math.abs(straighten) > .01f) {
            Matrix op = new Matrix();
            op.setRotate(straighten, width / 2f, height / 2f);
            RectF bounds = new RectF(0f, 0f, width, height);
            op.mapRect(bounds);
            op.postTranslate(-bounds.left, -bounds.top);
            matrix.postConcat(op);
            width = Math.max(1, Math.round(bounds.width()));
            height = Math.max(1, Math.round(bounds.height()));
        }

        if (Math.abs(perspectiveX) > .001f || Math.abs(perspectiveY) > .001f) {
            float px = clamp(perspectiveX, -.2f, .2f);
            float py = clamp(perspectiveY, -.2f, .2f);
            float dx = Math.abs(px) * width * .22f;
            float dy = Math.abs(py) * height * .22f;
            float[] src = {0,0, width,0, width,height, 0,height};
            float[] dst = {0,0, width,0, width,height, 0,height};
            // Expand the corrected edge beyond the output bounds rather than shrinking it inward;
            // the normal output rectangle then crops the overscan and avoids transparent wedges.
            if (px > 0f) { dst[0] -= dx; dst[2] += dx; }
            else if (px < 0f) { dst[6] -= dx; dst[4] += dx; }
            if (py > 0f) { dst[1] -= dy; dst[7] += dy; }
            else if (py < 0f) { dst[3] -= dy; dst[5] += dy; }
            Matrix perspective = new Matrix();
            if (perspective.setPolyToPoly(src, 0, dst, 0, 4)) matrix.postConcat(perspective);
        }

        int cropLeft = 0, cropTop = 0, cropRight = width, cropBottom = height;
        if (freeCrop != null && !isFullCrop(freeCrop)) {
            cropLeft = Math.max(0, Math.min(width - 1, Math.round(clamp(freeCrop.left,0f,1f) * width)));
            cropTop = Math.max(0, Math.min(height - 1, Math.round(clamp(freeCrop.top,0f,1f) * height)));
            cropRight = Math.max(cropLeft + 1, Math.min(width, Math.round(clamp(freeCrop.right,0f,1f) * width)));
            cropBottom = Math.max(cropTop + 1, Math.min(height, Math.round(clamp(freeCrop.bottom,0f,1f) * height)));
        } else if (cropAspectRatio > 0f) {
            float currentRatio = width / (float) height;
            if (currentRatio > cropAspectRatio) {
                int outWidth = Math.max(1, Math.round(height * cropAspectRatio));
                cropLeft = (width - outWidth) / 2;
                cropRight = cropLeft + outWidth;
            } else if (currentRatio < cropAspectRatio) {
                int outHeight = Math.max(1, Math.round(width / cropAspectRatio));
                cropTop = (height - outHeight) / 2;
                cropBottom = cropTop + outHeight;
            }
        }
        int outputWidth = Math.max(1, cropRight - cropLeft);
        int outputHeight = Math.max(1, cropBottom - cropTop);
        matrix.postTranslate(-cropLeft, -cropTop);

        Bitmap result = Bitmap.createBitmap(outputWidth, outputHeight, Bitmap.Config.ARGB_8888);
        result.setHasAlpha(source.hasAlpha());
        Canvas canvas = new Canvas(result);
        Paint paint = new Paint(Paint.ANTI_ALIAS_FLAG | Paint.FILTER_BITMAP_FLAG | Paint.DITHER_FLAG);
        if (!neutralColor(brightness, contrast, saturation, warmth, exposure, tint, filter)) {
            paint.setColorFilter(ImageAdjustments.createColorFilter(
                    brightness, contrast, saturation, warmth, exposure, tint, filter));
        }
        canvas.drawBitmap(source, matrix, paint);

        applyHighlightsShadowsInPlace(result, highlights, shadows);
        applySharpenInPlace(result, sharpness);

        float v = clamp(vignette, 0f, 1f);
        if (v > .001f) {
            float cx = outputWidth / 2f, cy = outputHeight / 2f;
            float radius = (float) Math.hypot(cx, cy);
            int edgeAlpha = Math.min(210, Math.round(190f * v));
            RadialGradient gradient = new RadialGradient(cx, cy, radius,
                    new int[] {Color.TRANSPARENT, Color.TRANSPARENT, Color.argb(edgeAlpha, 0, 0, 0)},
                    new float[] {0f, .52f, 1f}, Shader.TileMode.CLAMP);
            Paint vignettePaint = new Paint(Paint.ANTI_ALIAS_FLAG);
            vignettePaint.setShader(gradient);
            canvas.drawRect(0f, 0f, outputWidth, outputHeight, vignettePaint);
        }

        // Gainmap is preserved only when the edit is geometrical. Tone/color/markup operations
        // intentionally produce SDR rather than carrying a stale gainmap.
        boolean toneNeutral = neutralColor(brightness, contrast, saturation, warmth, exposure, tint, filter)
                && Math.abs(highlights) < .001f && Math.abs(shadows) < .001f
                && Math.abs(sharpness) < .001f && v <= .001f;
        if (source.hasGainmap() && toneNeutral) {
            try {
                Bitmap gainmapSource = source.getGainmap().getGainmapContents();
                Bitmap gainmapOutput = render(gainmapSource, operations, cropAspectRatio, freeCrop,
                        straighten, perspectiveX, perspectiveY,
                        0f, 0f, 1f, 0f, 0f, 0f, 0f, 0f, 0f, 0f, ImageAdjustments.FILTER_NONE);
                result.setGainmap(new Gainmap(source.getGainmap(), gainmapOutput));
            } catch (Exception ignored) {
                // Safe fallback is SDR output; never attach a stale/untransformed gainmap.
            }
        }
        return result;
    }

    private static void applyHighlightsShadowsInPlace(Bitmap bitmap, float highlights, float shadows) {
        float hi = clamp(highlights, -1f, 1f);
        float sh = clamp(shadows, -1f, 1f);
        if (Math.abs(hi) < .001f && Math.abs(sh) < .001f) return;
        int width = bitmap.getWidth(), height = bitmap.getHeight();
        int[] row = new int[width];
        for (int y = 0; y < height; y++) {
            bitmap.getPixels(row, 0, width, 0, y, width, 1);
            for (int x = 0; x < width; x++) {
                int c = row[x]; int a = Color.alpha(c), r = Color.red(c), g = Color.green(c), b = Color.blue(c);
                float lum = (0.2126f*r + 0.7152f*g + 0.0722f*b) / 255f;
                float shadowWeight = (1f - lum) * (1f - lum);
                float highlightWeight = lum * lum;
                float delta = sh * shadowWeight * 95f + hi * highlightWeight * 80f;
                // Negative highlights should recover brightness instead of crushing all channels equally.
                float scale = delta >= 0 ? 1f : .82f;
                r = clamp255(Math.round(r + delta * scale));
                g = clamp255(Math.round(g + delta * scale));
                b = clamp255(Math.round(b + delta * scale));
                row[x] = Color.argb(a,r,g,b);
            }
            bitmap.setPixels(row, 0, width, 0, y, width, 1);
        }
    }

    private static void applySharpenInPlace(Bitmap bitmap, float sharpness) {
        float amount = clamp(sharpness, 0f, 1f) * .65f;
        if (amount < .001f || bitmap.getWidth() < 3 || bitmap.getHeight() < 3) return;
        int width = bitmap.getWidth(), height = bitmap.getHeight();
        int[] prev = new int[width], cur = new int[width], next = new int[width], out = new int[width];
        bitmap.getPixels(prev,0,width,0,0,width,1);
        bitmap.getPixels(cur,0,width,0,1,width,1);
        for (int y=1; y<height-1; y++) {
            bitmap.getPixels(next,0,width,0,y+1,width,1);
            out[0]=cur[0]; out[width-1]=cur[width-1];
            for (int x=1; x<width-1; x++) {
                int c=cur[x], l=cur[x-1], r=cur[x+1], u=prev[x], d=next[x];
                int rr=unsharp(Color.red(c),Color.red(l),Color.red(r),Color.red(u),Color.red(d),amount);
                int gg=unsharp(Color.green(c),Color.green(l),Color.green(r),Color.green(u),Color.green(d),amount);
                int bb=unsharp(Color.blue(c),Color.blue(l),Color.blue(r),Color.blue(u),Color.blue(d),amount);
                out[x]=Color.argb(Color.alpha(c),rr,gg,bb);
            }
            bitmap.setPixels(out,0,width,0,y,width,1);
            int[] tmp=prev; prev=cur; cur=next; next=tmp;
        }
    }

    private static int unsharp(int center,int left,int right,int up,int down,float amount) {
        float blur=(left+right+up+down)/4f;
        return clamp255(Math.round(center + (center-blur)*amount));
    }

    private static boolean neutralColor(float brightness, float contrast, float saturation, float warmth,
            float exposure, float tint, int filter) {
        return Math.abs(brightness) < .001f && Math.abs(contrast) < .001f
                && Math.abs(saturation - 1f) < .001f && Math.abs(warmth) < .001f
                && Math.abs(exposure) < .001f && Math.abs(tint) < .001f
                && filter == ImageAdjustments.FILTER_NONE;
    }

    private static boolean isFullCrop(RectF crop) {
        return crop == null || (crop.left <= .001f && crop.top <= .001f && crop.right >= .999f && crop.bottom >= .999f);
    }
    private static float clamp(float value,float min,float max){return Math.max(min,Math.min(max,value));}
    private static int clamp255(int value){return Math.max(0,Math.min(255,value));}
}
