import Metal

public typealias ConfigureAction = (MTLRenderCommandEncoder?) -> Void
let ConfigureActionEmpty: ConfigureAction = {_ in }

open class QuadRenderable: Renderable {
    private var pipelineState: MTLRenderPipelineState!
    var fragmentShaderName: String {
        return QuadRenderConstants.defaultFragmentShaderName
    }
    
    open var fragmentNameOverride: String? { return nil }
    open var configureAction: ConfigureAction { return ConfigureActionEmpty }
    
    public init() {
        createRenderPipelineState()
    }
    
    private func createRenderPipelineState() {
        let pipelineDescriptor = MTLRenderPipelineDescriptor()
        let fragmentShader = fragmentNameOverride ?? fragmentShaderName
        guard
            let device = Adamantium.sharedDevice,
            let library = Adamantium.sharedLibrary
        else { return }
        pipelineDescriptor.vertexFunction = library.makeFunction(name: QuadRenderConstants.vertexShaderName)
        pipelineDescriptor.fragmentFunction = library.makeFunction(name: fragmentShader)
        pipelineDescriptor.colorAttachments[0].pixelFormat = .bgra8Unorm
        guard let pipelineState = try? device.makeRenderPipelineState(descriptor: pipelineDescriptor) else { return }
        self.pipelineState = pipelineState
    }
    
    private func setPipeline(encoder: MTLRenderCommandEncoder?) {
        guard let encoder = encoder else { return }
        encoder.setRenderPipelineState(pipelineState)
    }
    
    private func draw(encoder: MTLRenderCommandEncoder?) {
        guard let encoder = encoder else { return }
        encoder.drawPrimitives(
            type: .triangle,
            vertexStart: QuadRenderConstants.quadVertexStart,
            vertexCount: QuadRenderConstants.quadVertexCount)
    }
    
    public func render(encoder: MTLRenderCommandEncoder?) {
        setPipeline(encoder: encoder)
        configureAction(encoder)
        draw(encoder: encoder)
    }
    
    public var usesPreviouslyRenderedColors: Bool {
        return false
    }
}

open class ConfiguredQuadRenderable<T: QuadRenderConfiguration>: Renderable {
    public var config = T.defaultConfiguration
    private var pipelineState: MTLRenderPipelineState?
    var fragmentShaderName: String {
        return QuadRenderConstants.defaultFragmentShaderName
    }
    
    open var fragmentNameOverride: String? { return nil }
    open var configureAction: ConfigureAction { return ConfigureActionEmpty }
    
    public init() {
        createRenderPipelineState()
    }
    
    private func createRenderPipelineState() {
        let pipelineDescriptor = MTLRenderPipelineDescriptor()
        let fragmentShader = fragmentNameOverride ?? fragmentShaderName
        
        guard
            let device = Adamantium.sharedDevice,
            let library = Adamantium.sharedLibrary
        else { return }
        
        pipelineDescriptor.vertexFunction = library
            .makeFunction(name: QuadRenderConstants.vertexShaderName)
        pipelineDescriptor.fragmentFunction = library.makeFunction(name: fragmentShader)
        pipelineDescriptor.colorAttachments[0].pixelFormat = .bgra8Unorm
        
        do {
            let pipelineState = try device
                .makeRenderPipelineState(descriptor: pipelineDescriptor)
            self.pipelineState = pipelineState
        } catch {
            print(error)
        }
    }
    
    private func setPipeline(encoder: MTLRenderCommandEncoder?) {
        guard let encoder = encoder, let pipelineState = pipelineState else { return }
        encoder.setRenderPipelineState(pipelineState)
    }
    
    private func setConfigurationAndDraw(encoder: MTLRenderCommandEncoder?) {
        guard let encoder = encoder else { return }
        
        withUnsafePointer(to: &config) { pointer in
            encoder.setFragmentBytes(
                pointer,
                length: MemoryLayout<T>.stride,
                index: QuadRenderConstants.localUniformsIndex)
        }
        
        encoder.drawPrimitives(
            type: .triangle,
            vertexStart: QuadRenderConstants.quadVertexStart,
            vertexCount: QuadRenderConstants.quadVertexCount
        )
    }
    
    public func render(encoder: MTLRenderCommandEncoder?) {
        setPipeline(encoder: encoder)
        configureAction(encoder)
        setConfigurationAndDraw(encoder: encoder)
    }
}

enum QuadRenderConstants {
    static let defaultFragmentShaderName = "white"
    static let vertexShaderName = "quad_render"
    static let quadVertexCount = 6
    static let quadVertexStart = 0
    static let localUniformsIndex = 0
}
