//
//  MusicSynthesizer.swift
//  DashSmash
//
//  Renders a procedural drum + bass track for a Level into a 16-bit WAV blob
//  and plays it via AVAudioPlayer. Using AVAudioPlayer (rather than
//  AVAudioEngine + AVAudioPlayerNode) is far more tolerant of the device's
//  hardware audio format, so it works reliably on real devices, not just the
//  simulator.
//

import Foundation
import AVFoundation

final class MusicSynthesizer {

    private var player: AVAudioPlayer?
    private let sampleRate: Double = 44100
    private(set) var bpm: Double = 120
    private(set) var isRunning = false
    private var startHost: TimeInterval = 0
    private var rate: Float = 1.0

    deinit { stop() }

    func prepare(level: Level) {
        self.bpm = level.bpm
        guard let wav = Self.renderWAV(level: level, sampleRate: sampleRate) else { return }
        do {
            try AVAudioSession.sharedInstance().setCategory(.ambient, options: [.mixWithOthers])
            try AVAudioSession.sharedInstance().setActive(true)
            let p = try AVAudioPlayer(data: wav)
            p.enableRate = true
            p.rate = rate
            p.prepareToPlay()
            player = p
        } catch {
            player = nil
        }
    }

    func setRate(_ rate: Float) {
        self.rate = rate
        player?.enableRate = true
        player?.rate = rate
    }

    func start() {
        player?.rate = rate
        player?.play()
        isRunning = (player?.isPlaying ?? false)
        startHost = CACurrentMediaTime()
    }

    func stop() {
        player?.stop()
        isRunning = false
    }

    func currentTime() -> TimeInterval {
        player?.currentTime ?? max(0, CACurrentMediaTime() - startHost)
    }

    // MARK: - WAV building

    private static func renderWAV(level: Level, sampleRate: Double) -> Data? {
        let duration = level.durationSeconds + 1.0
        let totalSamples = Int(duration * sampleRate)
        guard totalSamples > 0 else { return nil }

        var samples = [Float](repeating: 0, count: totalSamples)
        samples.withUnsafeMutableBufferPointer { buf in
            guard let base = buf.baseAddress else { return }
            renderSamples(base: base, totalSamples: totalSamples, level: level, sampleRate: sampleRate)
        }

        // Soft master limiter
        var peak: Float = 0
        for s in samples { let a = abs(s); if a > peak { peak = a } }
        if peak > 0.95 {
            let scale: Float = 0.9 / peak
            for i in 0..<samples.count { samples[i] *= scale }
        }

        // Convert to Int16 little-endian
        var int16: [Int16] = []
        int16.reserveCapacity(totalSamples)
        for s in samples {
            let clamped = max(-1.0, min(1.0, s))
            int16.append(Int16(clamped * 32767))
        }

        var data = Data()
        data.reserveCapacity(44 + totalSamples * 2)
        appendWAVHeader(into: &data, sampleRate: Int(sampleRate), channelCount: 1, sampleCount: totalSamples)
        int16.withUnsafeBufferPointer { ptr in
            data.append(UnsafeBufferPointer(start: ptr.baseAddress, count: ptr.count))
        }
        return data
    }

    private static func appendWAVHeader(into data: inout Data, sampleRate: Int, channelCount: Int, sampleCount: Int) {
        let bitsPerSample = 16
        let byteRate = sampleRate * channelCount * (bitsPerSample / 8)
        let blockAlign = channelCount * (bitsPerSample / 8)
        let subChunk2Size = sampleCount * channelCount * (bitsPerSample / 8)
        let chunkSize = 36 + subChunk2Size

        func appendAscii(_ s: String) {
            data.append(contentsOf: Array(s.utf8))
        }
        func append<T: FixedWidthInteger>(_ value: T) {
            var v = value.littleEndian
            withUnsafeBytes(of: &v) { data.append(contentsOf: $0) }
        }

        appendAscii("RIFF")
        append(UInt32(chunkSize))
        appendAscii("WAVE")

        appendAscii("fmt ")
        append(UInt32(16))                    // PCM chunk size
        append(UInt16(1))                     // PCM format
        append(UInt16(channelCount))
        append(UInt32(sampleRate))
        append(UInt32(byteRate))
        append(UInt16(blockAlign))
        append(UInt16(bitsPerSample))

        appendAscii("data")
        append(UInt32(subChunk2Size))
    }

    // MARK: - Synthesis (writes mono samples into pre-allocated buffer)

    private static func renderSamples(base: UnsafeMutablePointer<Float>, totalSamples: Int, level: Level, sampleRate: Double) {
        let secondsPerBeat = 60.0 / level.bpm
        let samplesPerBeat = secondsPerBeat * sampleRate

        var rng = SeededRNG(string: level.songName + "|" + level.bandName + "|root")
        let rootSemitones = rng.intInRange(-3...8)
        let rootFreq = 110.0 * pow(2.0, Double(rootSemitones) / 12.0)
        let scale = [0, 3, 5, 7, 10]

        let totalBeats = Int(Double(totalSamples) / samplesPerBeat)
        for beatIndex in 0..<totalBeats {
            let beatStartSample = Int(Double(beatIndex) * samplesPerBeat)
            let section = currentSection(atBeat: Double(beatIndex), in: level)

            let beatInBar = beatIndex % 4
            if beatInBar == 0 || beatInBar == 2 {
                renderKick(into: base, start: beatStartSample, totalSamples: totalSamples, sampleRate: sampleRate)
            }
            if beatInBar == 1 || beatInBar == 3 {
                renderSnare(into: base, start: beatStartSample, totalSamples: totalSamples, sampleRate: sampleRate)
            }
            if beatIndex > 0 {
                let halfBeat = beatStartSample + Int(samplesPerBeat / 2)
                renderHat(into: base, start: halfBeat, totalSamples: totalSamples, sampleRate: sampleRate, gain: 0.18)
                renderHat(into: base, start: beatStartSample, totalSamples: totalSamples, sampleRate: sampleRate, gain: 0.12)
            }

            let modeBias: Int
            switch section?.mode ?? .cube {
            case .cube: modeBias = 0
            case .ship: modeBias = 1
            case .gravity: modeBias = 2
            case .mini: modeBias = 3
            }
            let degree = scale[(beatIndex + modeBias) % scale.count]
            let octaveShift: Int
            switch section?.mode ?? .cube {
            case .cube: octaveShift = 0
            case .ship: octaveShift = 0
            case .gravity: octaveShift = -12
            case .mini: octaveShift = 12
            }
            let bassFreq = rootFreq * pow(2.0, Double(degree + octaveShift) / 12.0)
            renderBass(
                into: base,
                start: beatStartSample,
                lengthSamples: Int(samplesPerBeat * 0.9),
                totalSamples: totalSamples,
                frequency: bassFreq,
                sampleRate: sampleRate
            )

            if let section = section, section.label.contains("Solo") || section.label == "Drop" {
                let melodyDegree = scale[(beatIndex * 3 + 1) % scale.count]
                let melodyFreq = rootFreq * pow(2.0, Double(melodyDegree + 24) / 12.0)
                renderLead(
                    into: base,
                    start: beatStartSample + Int(samplesPerBeat / 2),
                    lengthSamples: Int(samplesPerBeat * 0.4),
                    totalSamples: totalSamples,
                    frequency: melodyFreq,
                    sampleRate: sampleRate
                )
            }
        }
    }

    private static func currentSection(atBeat beat: Double, in level: Level) -> LevelSection? {
        var current: LevelSection?
        for s in level.sections {
            if s.startBeat <= beat { current = s } else { break }
        }
        return current
    }

    // MARK: - Voice rendering helpers

    private static func renderKick(into data: UnsafeMutablePointer<Float>, start: Int, totalSamples: Int, sampleRate: Double) {
        let length = Int(sampleRate * 0.18)
        for i in 0..<length {
            let idx = start + i
            if idx < 0 || idx >= totalSamples { return }
            let t = Double(i) / sampleRate
            let freq = 45.0 + 75.0 * exp(-t * 30.0)
            let env: Float = Float(exp(-t * 18.0))
            let sample = Float(sin(2.0 * .pi * freq * t)) * env * 0.9
            data[idx] += sample
        }
    }

    private static func renderSnare(into data: UnsafeMutablePointer<Float>, start: Int, totalSamples: Int, sampleRate: Double) {
        let length = Int(sampleRate * 0.18)
        var seed: UInt32 = 0xC0FFEE
        for i in 0..<length {
            let idx = start + i
            if idx < 0 || idx >= totalSamples { return }
            let t = Double(i) / sampleRate
            seed = seed &* 1664525 &+ 1013904223
            let noise = Float(Int32(bitPattern: seed)) / Float(Int32.max)
            let tone = Float(sin(2.0 * .pi * 200.0 * t))
            let env: Float = Float(exp(-t * 22.0))
            let sample = (noise * 0.6 + tone * 0.4) * env * 0.5
            data[idx] += sample
        }
    }

    private static func renderHat(into data: UnsafeMutablePointer<Float>, start: Int, totalSamples: Int, sampleRate: Double, gain: Float) {
        let length = Int(sampleRate * 0.06)
        var seed: UInt32 = 0xBEEF1234
        for i in 0..<length {
            let idx = start + i
            if idx < 0 || idx >= totalSamples { return }
            let t = Double(i) / sampleRate
            seed = seed &* 1664525 &+ 1013904223
            let noise = Float(Int32(bitPattern: seed)) / Float(Int32.max)
            let env: Float = Float(exp(-t * 80.0))
            data[idx] += noise * env * gain
        }
    }

    private static func renderBass(into data: UnsafeMutablePointer<Float>, start: Int, lengthSamples: Int, totalSamples: Int, frequency: Double, sampleRate: Double) {
        for i in 0..<lengthSamples {
            let idx = start + i
            if idx < 0 || idx >= totalSamples { return }
            let t = Double(i) / sampleRate
            let attack = min(1.0, t * 40.0)
            let env: Float = Float(attack * exp(-t * 3.5))
            let raw = sin(2.0 * .pi * frequency * t)
            let shaped = Float(tanh(raw * 2.5))
            data[idx] += shaped * env * 0.25
        }
    }

    private static func renderLead(into data: UnsafeMutablePointer<Float>, start: Int, lengthSamples: Int, totalSamples: Int, frequency: Double, sampleRate: Double) {
        for i in 0..<lengthSamples {
            let idx = start + i
            if idx < 0 || idx >= totalSamples { return }
            let t = Double(i) / sampleRate
            let attack = min(1.0, t * 60.0)
            let env: Float = Float(attack * exp(-t * 6.0))
            let saw1 = Float(2.0 * (frequency * t - floor(frequency * t + 0.5)))
            let saw2 = Float(2.0 * (frequency * 1.005 * t - floor(frequency * 1.005 * t + 0.5)))
            data[idx] += (saw1 + saw2) * 0.18 * env
        }
    }
}
