//
//  MaskedAuraStillRender.swift
//  Adamantium
//
//  Created by Sunny Uppal on 11/25/24.
//

import Metal
import Foundation

public class MaskedAuraStillRender {
    
    public enum MaskingType {
        case noMask
        case topLeft
        case bottomRight
    }
    
    public enum AuraShape: String, Identifiable {
        
        public var id: String {
            return rawValue
        }
        
        case standard           // Done
        case cornerRipple       // Done
        case rightArrow         // Done
        case bottomRipple       // Done
        case curveEcho          // Done
        case circleTile         // Done
        case diamond            // Done
        case sunburst           // Done
        case circleGlass
        case circleFlower       // Done
        case star               // Done
        case circleRipple       // Done
        case curveRotation      // Done
        case circle             // Done
        case waves              // Done
        case gradientDown       // Done
        case customMap          // -- TODO
        
        public var displayName: String {
            switch self {
            case .standard:
                return "Standard"
            case .cornerRipple:
                return "Corner Ripple"
            case .rightArrow:
                return "Right Arrow"
            case .bottomRipple:
                return "Bottom Ripple"
            case .curveEcho:
                return "Curve Echo"
            case .circleTile:
                return "Circle Tile"
            case .diamond:
                return "Diamond"
            case .sunburst:
                return "Sunburst"
            case .circleGlass:
                return "Circle Glass"
            case .circleFlower:
                return "Circle Flower"
            case .star:
                return "Star"
            case .circleRipple:
                return "Circle Ripple"
            case .curveRotation:
                return "Curve Rotation"
            case .circle:
                return "Circle"
            case .waves:
                return "Waves"
            case .gradientDown:
                return "Gradient Down"
            case .customMap:
                return "Custom Map"
            }
        }
    }
    
    public struct RenderConfig {
        
        public let width: Int
        public let height: Int
        public let colorSet: SunoTriColorSet
        public let overlayColor: SunoColorPalette
        public let maskingType: MaskingType
        public let auraShape: AuraShape
        public let auraShapeCustomMap: Data?
        public let auraShapeMapOffsetScalar: Float
        public let scale: Float
        public let seed: Float
        
        public var sizeVec2: ADVector2 {
            return ADVector2(ADFloat(width), ADFloat(height))
        }
        
        public init(
            width: Int,
            height: Int,
            colorSet: SunoTriColorSet,
            overlayColor: SunoColorPalette,
            maskingType: MaskingType,
            auraShape: AuraShape,
            auraShapeCustomMap: Data? = nil,
            auraShapeMapOffsetScalar: Float,
            scale: Float,
            seed: Float
        ) {
            self.width = width
            self.height = height
            self.colorSet = colorSet
            self.overlayColor = overlayColor
            self.maskingType = maskingType
            self.auraShape = auraShape
            self.auraShapeCustomMap = auraShapeCustomMap
            self.auraShapeMapOffsetScalar = auraShapeMapOffsetScalar
            self.scale = scale
            self.seed = seed
        }
    }
    
    public var lastRender: MTLTexture?
    
    public init(lastRender: MTLTexture? = nil) {
        self.lastRender = lastRender
    }
    
    public func render(_ config: RenderConfig) -> MTLTexture? {
        let renderHelper = OfflineRenderHelper(width: config.width, height: config.height)
        let offsetMapTexture = getOffsetMap(config: config, renderHelper: renderHelper)
        let overlayMapTexture = getOverlayColorFadeMap(config: config, renderHelper: renderHelper)
        
        let finalAuraTexture = renderAura(
            config: config,
            offsetMap: offsetMapTexture,
            overlayMap: overlayMapTexture,
            renderHelper: renderHelper
        )
        
        return finalAuraTexture
    }
    
    public func renderPNG(_ config: RenderConfig) throws -> URL? {
        guard let texture = render(config) else { return nil }
        return try StillCaptureUtility.png(texture: texture)
    }
}

private extension MaskedAuraStillRender {
    
    func renderAura(
        config: RenderConfig,
        offsetMap: MTLTexture?,
        overlayMap: MTLTexture?,
        renderHelper: OfflineRenderHelper) -> MTLTexture? {
            
        var size: ADVector2 {
            return ADVector2(
                x: Float(config.width),
                y: Float(config.height)
            )
        }
        
        let gradientSeed = config.seed
        let gradientScale = config.scale
        let mixPointStart: ADFloat = -0.2
        let mixPointBackground: ADFloat = -0.1
        let mixPointBackgroundA: ADFloat = 0.0
        let mixPointAB: ADFloat = 0.2
        let mixPointBC: ADFloat = 0.4
        let mixPointEnd: ADFloat = 2.0
        
        let noiseRenderable = TextureMaskedSimplex3DAuraRenderable()
        noiseRenderable.inputTexture = offsetMap
        noiseRenderable.colorFadeMap = overlayMap
        noiseRenderable.config = .init(
            displaySize: size,
            fractalNoiseX: .zero,
            fractalNoiseY: .zero,
            fractalNoiseZ: gradientSeed,
            fractalUniformScale: 0.5 * gradientScale,
            fractalMapOffsetScalar: config.auraShapeMapOffsetScalar,
            grainNoiseZ: gradientSeed,
            grainNoiseScale: 5.0 * gradientScale,
            circleMaskRadius: 100.0,
            backgroundColor: SunoColorPalette.backgroundDotCom.floatThree,
            colorA: config.colorSet.floatColorA,
            colorB: config.colorSet.floatColorB,
            colorC: config.colorSet.floatColorC,
            colorOverlay: config.overlayColor.floatThree,
            mixPointStart: mixPointStart,
            mixPointBackground: mixPointBackground,
            mixPointBackgroundA: mixPointBackgroundA,
            mixPointAB: mixPointAB,
            mixPointBC: mixPointBC,
            mixPointEnd: mixPointEnd
        )
        
        renderHelper.clearLastRender()
        renderHelper.commitRenderableLayers(layers: [noiseRenderable])
        lastRender = renderHelper.lastRender
        return renderHelper.lastRender
    }
    
    func getEchoBoxMaskRenderable() -> MaskingEchoBox {
        let maskedEchoBox = MaskingEchoBox()
        maskedEchoBox.config = .init(
            boxWidth: 1.0,
            boxHeight: 1.0,
            boxRotation: 0.125,
            boxRoundness: 0.05,
        
            repeatCount: 5.0,
            repeatXSpacing: 0.3,
            repeatYSpacing: 0.0,
            repeatXOffset: 1.7,
            repeatYOffset: 0.0,
        
            fadeMaskMix: 0.0,
            fadeMaskRotation: -0.125,
            fadeMaskOffset: 1.1,
            fadeMaskScaling: 2.0,
            invertMix: 1.0)
        return maskedEchoBox
    }
    
    func getOffsetMap(config: RenderConfig, renderHelper: OfflineRenderHelper) -> MTLTexture? {
        
        switch config.auraShape {
        case .standard, .gradientDown, .diamond:
            let fillRenderable = ColorFillRenderable()
            fillRenderable.config = .init(color: .black)
            renderHelper.commitRenderableLayers(layers: [ fillRenderable ])
        case .cornerRipple:
            let offsetMaskRenderable = MaskedEchoCircleRenderable()
            offsetMaskRenderable.config = .init(
                circleRadius: 0.1,
                spaceScale: 7,
                repeatValueScalar: 0.125,
                xOffset: 0.5,
                yOffset: -0.5,
                invertMix: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .rightArrow:
            let offsetMaskRenderable = getEchoBoxMaskRenderable()
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .bottomRipple:
            let offsetMaskRenderable = MaskedEchoCircleRenderable()
            offsetMaskRenderable.config = .init(
                circleRadius: 0.1,
                spaceScale: 3,
                repeatValueScalar: 0.1,
                xOffset: 0.0,
                yOffset: -0.5,
                invertMix: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .curveEcho:
            let offsetMaskRenderable = MaskedRepeatCircleCurveRenderable()
            let curveScale: ADFloat = 0.55
            offsetMaskRenderable.config = .init(
                canvasSize: config.sizeVec2,
                lineWidth: 0.1,
                xSpaceScale: curveScale,
                ySpaceScale: curveScale,
                xOffset: 1.68,
                yOffset: 0.2,
                pinchAmount: 0.6,
                repeatCount: 3,
                repeatXSpacing: -0.2,
                repeatYSpacing: 0.0,
                repeatRotation: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .sunburst:
            let offsetMaskRenderable = MaskedBurstRenderable()
            let curveScale: ADFloat = 0.5
            let rayCount: ADFloat = 6.0
            offsetMaskRenderable.config = .init(
                rotation: ((Float.pi * 2.0) / (rayCount * 2.0)) * 0.5,
                rayCount: rayCount,
                xOffset: 0.0,
                yOffset: 0.5
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .star:
            let offsetMaskRenderable = MaskedStarRenderable()
            let curveScale: ADFloat = 0.5
            offsetMaskRenderable.config = .init(
                starRadius: 0.45,
                armRatio: 0.65,
                rounding: 0.06,
                spaceScale: 1.0,
                xOffset: 0.0,
                yOffset: 0.05
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .curveRotation:
            let offsetMaskRenderable = MaskedRepeatCircleCurveRenderable()
            let curveScale: ADFloat = 0.5
            offsetMaskRenderable.config = .init(
                canvasSize: config.sizeVec2,
                lineWidth: 0.09,
                xSpaceScale: curveScale,
                ySpaceScale: curveScale,
                xOffset: 1.0,
                yOffset: -0.125,
                pinchAmount: 0.25,
                repeatCount: 4,
                repeatXSpacing: 0.0,
                repeatYSpacing: 0.0,
                repeatRotation: 0.8
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .circle:
            let offsetMaskRenderable = MaskedRepeatCircleRenderable()
            offsetMaskRenderable.config = .init(
                circleRadius: 0.5,
                spaceScale: 1.0,
                spaceRotation: 0.0,
                patternOneAndOneMix: 0.0,
                xOffset: 0.0,
                yOffset: 0.0,
                invertMix: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .circleRipple:
            let offsetMaskRenderable = MaskedEchoCircleRenderable()
            offsetMaskRenderable.config = .init(
                circleRadius: 1.0,
                spaceScale: 6,
                repeatValueScalar: 0.1,
                xOffset: 0.0,
                yOffset: 0.0,
                invertMix: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .circleGlass:
            let offsetMaskRenderable = MaskedCircleGlassRenderable()
            offsetMaskRenderable.config = .init(
                scale: 1.65,
                xOffset: 0.0,
                yOffset: 0.0,
                invertMix: 1.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
            
        case .circleFlower:
            let offsetMaskRenderable = MaskedRepeatCircleRenderable()
            offsetMaskRenderable.config = .init(
                circleRadius: 0.5,
                spaceScale: 1.0,
                spaceRotation: 0.125,
                patternOneAndOneMix: 1.0,
                xOffset: 0.71,
                yOffset: 0.0,
                invertMix: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .circleTile:
            let offsetMaskRenderable = MaskedRepeatCircleRenderable()
            offsetMaskRenderable.config = .init(
                circleRadius: 0.4,
                spaceScale: 1.5,
                spaceRotation: 0.0, // Rotation by circle ratio
                patternOneAndOneMix: 0.0,
                xOffset: 0.0,
                yOffset: 0.0,
                invertMix: 0.0)
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .waves:
            let offsetMaskRenderable = MaskedRepeatCircleCurveRenderable()
            let curveScale: ADFloat = 0.8
            offsetMaskRenderable.config = .init(
                canvasSize: config.sizeVec2,
                lineWidth: 0.173,
                xSpaceScale: curveScale,
                ySpaceScale: curveScale,
                xOffset: 2.0,
                yOffset: 0.55,
                pinchAmount: 0.2,
                repeatCount: 3,
                repeatXSpacing: 0.35,
                repeatYSpacing: -0.28,
                repeatRotation: 0.0
            )
            renderHelper.commitRenderableLayers(layers: [
                offsetMaskRenderable
            ])
        case .customMap:
            guard let pngData = config.auraShapeCustomMap else { return nil }
            let mapTexture = try? Adamantium.sharedTextureLoader.newTexture(data: pngData, options: [:])
            let uvTexturePass = UVTexturePassRenderable()
            uvTexturePass.inputTexture = mapTexture
            renderHelper.commitRenderableLayers(layers: [
                uvTexturePass
            ])
        }
        
        return renderHelper.lastRender
    }
    
    func getOverlayColorFadeMap(
        config: RenderConfig,
        renderHelper: OfflineRenderHelper) -> MTLTexture? {
        
        switch config.auraShape {
        case .cornerRipple:
            let renderable = MaskedCircleFadeRenderable()
            renderable.config = .init(
                radius: 0.7,
                scale: 1.0,
                fadeStart: -0.1,
                fadeEnd: 0.9,
                xOffset: -0.5,
                yOffset: 0.6,
                invertMix: 1.0
            )
            renderHelper.commitRenderableLayers(layers: [
                renderable
            ])
            
        case .diamond:
            let renderable = MaskedDiamondGradientRenderable()
            renderable.config = .init(scale: 3.0, centerOffset: 2.0 / 3.0)
            renderHelper.commitRenderableLayers(layers: [
                renderable
            ])
            
        case .rightArrow:
            let renderable = MaskingEchoBox()
            renderable.config = .init(
                boxWidth: 1.0,
                boxHeight: 1.0,
                boxRotation: 0.125,
                boxRoundness: 0.05,
            
                repeatCount: 5.0,
                repeatXSpacing: 0.3,
                repeatYSpacing: 0.0,
                repeatXOffset: 1.7,
                repeatYOffset: 0.0,
            
                fadeMaskMix: 1.0,
                fadeMaskRotation: -0.125,
                fadeMaskOffset: 0.9,
                fadeMaskScaling: 2.0,
                invertMix: 0.0)
            renderHelper.commitRenderableLayers(layers: [
                renderable
            ])
            
        case .gradientDown:
            let renderable = GrayscaleLinearGradientRenderable()
            renderable.config = .init(rotationRatio: 0.25, start: .zero, stop: 0.6)
            renderHelper.commitRenderableLayers(layers: [
                renderable
            ])
            
        case .standard, .bottomRipple, .curveEcho, .circleTile, .sunburst,
            .circleGlass, .circleFlower, .star, .circleRipple, .curveRotation,
            .circle, .waves, .customMap:
            let renderable = ColorFillRenderable()
            renderable.config = .init(color: .black)
            renderHelper.commitRenderableLayers(layers: [
                renderable
            ])
        }
        
        return renderHelper.lastRender
    }
}
