//
//  MIDIFileParser.swift
//  DashSmash
//
//  Small Standard MIDI File parser that extracts tempo, program changes, and
//  note events. It intentionally ignores controller/pitch data because level
//  generation only needs instrument timing and register.
//

import Foundation

enum MIDIFileParserError: Error, LocalizedError {
    case invalidHeader
    case unsupportedTimeDivision
    case malformedTrack

    var errorDescription: String? {
        switch self {
        case .invalidHeader: return "The file is not a valid MIDI file."
        case .unsupportedTimeDivision: return "This MIDI file uses SMPTE timing, which is not supported yet."
        case .malformedTrack: return "The MIDI track data is malformed."
        }
    }
}

enum MIDIFileParser {
    static func parse(_ data: Data) throws -> MIDIParsedFile {
        var reader = MIDIByteReader(data: data)
        guard try reader.readString(count: 4) == "MThd" else { throw MIDIFileParserError.invalidHeader }
        let headerLength = try reader.readUInt32()
        guard headerLength >= 6 else { throw MIDIFileParserError.invalidHeader }
        _ = try reader.readUInt16() // format
        let trackCount = Int(try reader.readUInt16())
        let division = Int(try reader.readUInt16())
        if headerLength > 6 { try reader.skip(Int(headerLength) - 6) }
        guard division & 0x8000 == 0 else { throw MIDIFileParserError.unsupportedTimeDivision }
        let ticksPerQuarter = max(1, division)

        var allNotes: [MIDINote] = []
        var tempoBPM: Double?

        for _ in 0..<trackCount where !reader.isAtEnd {
            guard try reader.readString(count: 4) == "MTrk" else { throw MIDIFileParserError.malformedTrack }
            let length = Int(try reader.readUInt32())
            let trackEnd = reader.offset + length
            let trackResult = try parseTrack(
                reader: &reader,
                trackEnd: min(trackEnd, data.count),
                ticksPerQuarter: ticksPerQuarter
            )
            allNotes.append(contentsOf: trackResult.notes)
            if tempoBPM == nil { tempoBPM = trackResult.firstTempoBPM }
            reader.offset = min(trackEnd, data.count)
        }

        let sorted = allNotes
            .filter { $0.durationBeats > 0.02 && $0.velocity > 0 }
            .sorted { $0.beat == $1.beat ? $0.note < $1.note : $0.beat < $1.beat }
        return MIDIParsedFile(
            ticksPerQuarter: ticksPerQuarter,
            bpm: max(50, min(220, (tempoBPM ?? 120).rounded())),
            notes: sorted
        )
    }

    private static func parseTrack(
        reader: inout MIDIByteReader,
        trackEnd: Int,
        ticksPerQuarter: Int
    ) throws -> (notes: [MIDINote], firstTempoBPM: Double?) {
        var tick = 0
        var runningStatus: UInt8?
        var channelPrograms = Array<Int?>(repeating: nil, count: 16)
        var activeNotes: [ActiveNoteKey: ActiveNote] = [:]
        var notes: [MIDINote] = []
        var firstTempoBPM: Double?

        while reader.offset < trackEnd {
            tick += try reader.readVariableLengthQuantity(limit: trackEnd)
            var status = try reader.readUInt8()
            var firstDataByte: UInt8?

            if status < 0x80 {
                guard let running = runningStatus else { throw MIDIFileParserError.malformedTrack }
                firstDataByte = status
                status = running
            } else if status < 0xF0 {
                runningStatus = status
            }

            if status == 0xFF {
                let metaType = try reader.readUInt8()
                let length = try reader.readVariableLengthQuantity(limit: trackEnd)
                if metaType == 0x51 && length == 3 {
                    let b1 = UInt32(try reader.readUInt8())
                    let b2 = UInt32(try reader.readUInt8())
                    let b3 = UInt32(try reader.readUInt8())
                    let microsPerQuarter = Double((b1 << 16) | (b2 << 8) | b3)
                    if firstTempoBPM == nil && microsPerQuarter > 0 {
                        firstTempoBPM = 60_000_000.0 / microsPerQuarter
                    }
                } else {
                    try reader.skip(length)
                }
                runningStatus = nil
                continue
            }

            if status == 0xF0 || status == 0xF7 {
                let length = try reader.readVariableLengthQuantity(limit: trackEnd)
                try reader.skip(length)
                runningStatus = nil
                continue
            }

            let eventType = status & 0xF0
            let channel = Int(status & 0x0F)
            let data1 = try firstDataByte ?? reader.readUInt8()

            switch eventType {
            case 0x80:
                let velocity = try reader.readUInt8()
                closeNote(note: Int(data1), velocity: Int(velocity), channel: channel, tick: tick, activeNotes: &activeNotes, output: &notes, ticksPerQuarter: ticksPerQuarter)
            case 0x90:
                let velocity = Int(try reader.readUInt8())
                let key = ActiveNoteKey(channel: channel, note: Int(data1))
                if velocity == 0 {
                    closeNote(note: Int(data1), velocity: 0, channel: channel, tick: tick, activeNotes: &activeNotes, output: &notes, ticksPerQuarter: ticksPerQuarter)
                } else {
                    activeNotes[key] = ActiveNote(startTick: tick, velocity: velocity, program: channelPrograms[channel])
                }
            case 0xA0, 0xB0, 0xE0:
                _ = try reader.readUInt8()
            case 0xC0:
                channelPrograms[channel] = Int(data1)
            case 0xD0:
                break
            default:
                throw MIDIFileParserError.malformedTrack
            }
        }

        for (key, active) in activeNotes {
            appendNote(key: key, active: active, endTick: tick, output: &notes, ticksPerQuarter: ticksPerQuarter)
        }
        return (notes, firstTempoBPM)
    }

    private static func closeNote(
        note: Int,
        velocity: Int,
        channel: Int,
        tick: Int,
        activeNotes: inout [ActiveNoteKey: ActiveNote],
        output: inout [MIDINote],
        ticksPerQuarter: Int
    ) {
        let key = ActiveNoteKey(channel: channel, note: note)
        guard let active = activeNotes.removeValue(forKey: key) else { return }
        appendNote(key: key, active: active, endTick: tick, output: &output, ticksPerQuarter: ticksPerQuarter)
    }

    private static func appendNote(
        key: ActiveNoteKey,
        active: ActiveNote,
        endTick: Int,
        output: inout [MIDINote],
        ticksPerQuarter: Int
    ) {
        let durationTicks = max(1, endTick - active.startTick)
        output.append(MIDINote(
            beat: Double(active.startTick) / Double(ticksPerQuarter),
            durationBeats: Double(durationTicks) / Double(ticksPerQuarter),
            channel: key.channel,
            note: key.note,
            velocity: active.velocity,
            program: active.program
        ))
    }
}

private struct ActiveNoteKey: Hashable {
    let channel: Int
    let note: Int
}

private struct ActiveNote {
    let startTick: Int
    let velocity: Int
    let program: Int?
}

struct MIDIByteReader {
    let data: Data
    var offset = 0

    var isAtEnd: Bool { offset >= data.count }

    mutating func readUInt8() throws -> UInt8 {
        guard offset < data.count else { throw MIDIFileParserError.malformedTrack }
        defer { offset += 1 }
        return data[offset]
    }

    mutating func readUInt16() throws -> UInt16 {
        let b1 = UInt16(try readUInt8())
        let b2 = UInt16(try readUInt8())
        return (b1 << 8) | b2
    }

    mutating func readUInt32() throws -> UInt32 {
        let b1 = UInt32(try readUInt8())
        let b2 = UInt32(try readUInt8())
        let b3 = UInt32(try readUInt8())
        let b4 = UInt32(try readUInt8())
        return (b1 << 24) | (b2 << 16) | (b3 << 8) | b4
    }

    mutating func readString(count: Int) throws -> String {
        guard offset + count <= data.count else { throw MIDIFileParserError.malformedTrack }
        let slice = data[offset..<(offset + count)]
        offset += count
        return String(decoding: slice, as: UTF8.self)
    }

    mutating func readVariableLengthQuantity(limit: Int) throws -> Int {
        var value = 0
        for _ in 0..<4 {
            guard offset < min(limit, data.count) else { throw MIDIFileParserError.malformedTrack }
            let byte = try readUInt8()
            value = (value << 7) | Int(byte & 0x7F)
            if byte & 0x80 == 0 { return value }
        }
        throw MIDIFileParserError.malformedTrack
    }

    mutating func skip(_ count: Int) throws {
        guard count >= 0, offset + count <= data.count else { throw MIDIFileParserError.malformedTrack }
        offset += count
    }
}
