#include <fstream>
#include <iostream>
#include <vector>
#include <algorithm>

#include <cassert>
#define __GXX_EXPERIMENTAL_CXX0X__ // Fixed-width integer types in C++ are just too radical for GCC to stay silent.
#include <cstdint>

#include <jdkmidi/filewritemultitrack.h>

// These meta event defs are mixed up in jdkmidi...
#define META_TRACK_NAME      0x03
#define META_INSTRUMENT_NAME 0x04

#define CLOCK_RESOLUTION 24
#define MIN_SIZE 0x10
#define BUF_SIZE 0xFFFF

#define VCE_ARG 1
#define KMS_ARG 2
#define OUT_ARG 3

#define DEFAULT_SEQ_POS 0

#define Str2Id(a,b,c,d)  (a | b << 8 | c << 16 | d << 24)
#define Id2Str(a)        (char)a << (char)(a >> 8) << (char)(a >> 16) << (char)(a >> 24)

#pragma pack(push, 1)
struct resource_t {
  uint32_t size;
  uint16_t count;
  uint32_t ids[];
};
#pragma pack(pop)

struct voice_t {
  uint32_t id;
  bool     enabled;
  int8_t   channel;
  int8_t   program;
  int8_t   volume;
  int8_t   pan;
  int8_t   transpose;
  bool     hasHits;
};

struct track_t {
  uint32_t id;
  bool     enabled;
  voice_t  voice;

  jdkmidi::MIDITrack midi;
};

struct sequence_t {
  uint32_t id;
  uint8_t  numVoices;
  voice_t  voices[255];
  uint8_t  numTracks;
  track_t  tracks[255];

  jdkmidi::MIDITrack metaTrack;
};

enum {
  EVENT_END_OF_TRK = 0x00,
  EVENT_VOICE      = 0x03,
  EVENT_TEMPO      = 0x04,
  EVENT_CONTROLLER = 0x06,
  EVENT_PITCH_BEND = 0x0C,
  EVENT_TRACK_NAME = 0x0E,
  EVENT_SYS_EX     = 0x0F
};

bool readFile(const char* filename, char* buf)
{
  std::ifstream fs;
  fs.open(filename, std::ios::in | std::ios::binary | std::ios::ate);
  if (!fs) {
    std::cerr << "Couldn't open file \"" << filename << "\" for reading." << std::endl;
    return false;
  }

  int size = fs.tellg();
  if (size < MIN_SIZE || size > BUF_SIZE) {
    std::cerr << "Invalid size of file \"" << filename << "\": " << size << " bytes" << std::endl;
    fs.close();
    return false;
  }

  fs.seekg(0, std::ios_base::beg);

  fs.read(buf, size);

  if (!fs.good()) {
    std::cerr << "Error reading file \"" << filename << "\"" << std::endl;
    return false;
  }

  fs.close();

  return true;
}

char* resourceAt(char* buf, uint16_t pos)
{
  resource_t* res = (resource_t*)buf;

  if (pos > res->count) {
    return NULL;
  }

  return buf + 6 + res->count * 8 + res->ids[res->count + pos];
}

char* findResource(char* buf, uint32_t id)
{
  resource_t* res = (resource_t*)buf;

  for (int i = 0; i < res->count; ++i) {
    if (res->ids[i] == id) {
      return resourceAt(buf, i);
    }
  }

  return NULL;
}

bool midiTimeSorter(const jdkmidi::MIDITimedBigMessage& a, const jdkmidi::MIDITimedBigMessage& b)
{
  return a.GetTime() < b.GetTime();
}

inline unsigned int readVarLenValue(unsigned char** src)
{
  unsigned int val = 0;

  for (; **src & 0x80; ++*src) {
    val = (val << 14) | (**src & 0x7F) << 7;
  }
  val |= **src;
  ++*src;

  return val;
}

bool parseKmsTrack(char* trkBuf, track_t* trk, sequence_t* seq)
{
  char txtBuf[256];

  int lastNote = -1;
  unsigned int lastNoteOffAt = -1;

  std::vector<jdkmidi::MIDITimedBigMessage> messages;
  jdkmidi::MIDISystemExclusive sysEx;

  unsigned char* trkPos = (unsigned char*)trkBuf;
  unsigned char* trkEnd = trkPos + *(uint32_t*)(trkPos);
  trkPos += 4;

  unsigned int timestamp = 0;

  // Process track events.
  while (trkPos < trkEnd) {
    jdkmidi::MIDITimedBigMessage m;

    unsigned int delta = readVarLenValue(&trkPos);

    timestamp += delta;

    uint8_t c = *trkPos - 0xD9;
    if (c < 0x12) {
      ++trkPos;

      switch (c) {
        case EVENT_END_OF_TRK:
          m.SetTime(timestamp);
          m.SetType(jdkmidi::META_EVENT);
          m.SetMetaType(jdkmidi::META_END_OF_TRACK);
          messages.push_back(m);
          break;

        // Code 2 and 10 always occur at the same time on every track at the
        // end of the song. Probably related to looping.
        case 1:
        case 2:
        case 10:
          std::cout << "  Time " << timestamp << ": code = " << (int)c << std::endl;
          break;

        case EVENT_VOICE:
          trk->voice = seq->voices[*trkPos++];

          // Track has no voice definition in current set. Disable track, but
          // keep parsing in case it contains global meta events.
          if (!trk->voice.enabled) {
            std::cout << "Skipping track " << Id2Str(trk->id) << std::endl;
            trk->enabled = false;
            continue;
          }

          std::cout << "Setting voice " << Id2Str(trk->voice.id) << " for track " << Id2Str(trk->id) << std::endl;

          m.SetTime(timestamp);
          m.SetProgramChange(trk->voice.channel, trk->voice.program);
          messages.push_back(m);

          m.SetControlChange(trk->voice.channel, jdkmidi::C_MAIN_VOLUME, trk->voice.volume);
          messages.push_back(m);

          m.SetControlChange(trk->voice.channel, jdkmidi::C_PAN, trk->voice.pan);
          messages.push_back(m);

          // Use voice id as instrument name.
          sysEx = jdkmidi::MIDISystemExclusive((unsigned char*)&trk->voice.id, 4, 4, false);
          m.CopySysEx(&sysEx);
          m.SetType(jdkmidi::META_EVENT);
          m.SetMetaType(META_INSTRUMENT_NAME);
          messages.push_back(m);
          break;

        case EVENT_TEMPO:
          std::cout << "  Tempo = " << (int)*trkPos << " BPM" << std::endl;
          m.SetTime(0);
          m.SetTempo32((int)*(trkPos++) * 32);
          seq->metaTrack.PutEvent(m);
          break;

        case 5:
        case 7:
        case 8:
        case 9:
        case 11:
        case 16:
        case 17:
          std::cout << "  Time " << timestamp << ": code " << (int)c << " = " << (int)*(trkPos++) << std::endl;
          break;

        case EVENT_CONTROLLER:
          m.SetTime(timestamp);
          m.SetControlChange(trk->voice.channel, (int)*trkPos, (int)*(trkPos+1));
          messages.push_back(m);
          trkPos += 2;
          break;

        case EVENT_PITCH_BEND:
          m.SetTime(timestamp);
          m.SetPitchBend(trk->voice.channel, (int)*trkPos, (int)*(trkPos+1));
          messages.push_back(m);
          trkPos += 2;
          break;

        // This does-a-lot ... but isnt called in regular songs .. maybe some engine stuff?
        case 13:
          assert(false);
          break;

        case EVENT_TRACK_NAME:
        case EVENT_SYS_EX:
          sysEx = jdkmidi::MIDISystemExclusive(trkPos+1, trkPos[0], trkPos[0], false);
          m.SetTime(timestamp);
          m.CopySysEx(&sysEx);
          if (c == EVENT_TRACK_NAME) {
            std::cout << "Track " << Id2Str(trk->id) << ": \"" << (trkPos+1) << "\"" << std::endl;
            m.SetType(jdkmidi::META_EVENT);
            m.SetMetaType(META_TRACK_NAME);
          }
          else {
            m.SetSysEx(); // TODO: Trim start and end marker byte?
          }
          messages.push_back(m);
          trkPos += *trkPos + 1;
          break;

        // Shouldn't happen...
        default:
          assert(false);
          break;
      }
    }
    // Regular note event.
    else {
      int note     = -1;
      int velocity = 0;
      int duration = 0;
      int noteDuration;

      if (*trkPos & 0x80) {
        note = (*trkPos & 0x7F) + trk->voice.transpose;
        ++trkPos;
      }

      // TODO: Was the previous block just a sanity check, or are we missing
      // more possible conditions?
      assert(note != -1);

      velocity = *trkPos++;

      // TODO: How do we handle MSb set on velocity?
      assert((velocity & 0x80) != 0x80);

      duration = readVarLenValue(&trkPos);

      // Lame attempt at detecting "hits".
      if (trk->voice.hasHits && note == lastNote && timestamp == lastNoteOffAt && velocity > 63) {
        noteDuration = 1;
      }
      else {
        noteDuration = duration;
      }

      lastNote = note;
      lastNoteOffAt = timestamp + duration;

      m.SetTime(timestamp);
      m.SetNoteOn(trk->voice.channel, note, 127);
      messages.push_back(m);

      m.SetTime(timestamp + noteDuration);
      m.SetNoteOff(trk->voice.channel, note, 0);
      messages.push_back(m);
    }
  }

  // Sort by time and add to MIDI track.
  std::sort(messages.begin(), messages.end(), midiTimeSorter);
  for (std::vector<jdkmidi::MIDITimedBigMessage>::iterator i = messages.begin(); i != messages.end(); ++i) {
    //std::cout << "  " << (*i).GetTime() << "\t" << (*i).MsgToText(txtBuf) << std::endl;
    trk->midi.PutEvent(*i);
  }

  return true;
}

bool parseKmsHeader(char* seqBuf, char* vceBuf, sequence_t* seq)
{
  char* hdr = findResource(seqBuf, Str2Id('H','D','R','1'));

  if (hdr == NULL) {
    std::cerr << "Couldn't find sequence header." << std::endl;
    return false;
  }

  hdr += 6; // Skip boring counters.

  // Read voice list.
  seq->numVoices = *hdr++;
  for (int i = 0; i < seq->numVoices; ++i, hdr += 4) {
    seq->voices[i].id = *(int32_t*)hdr;

    // Look-up voice properties.
    char *vce = findResource(vceBuf, seq->voices[i].id);

    if (vce == NULL) {
      seq->voices[i].enabled = false;
    }
    else {
      seq->voices[i].enabled   = true;
      seq->voices[i].channel   = vce[0x43];
      seq->voices[i].program   = vce[0x44];
      seq->voices[i].volume    = vce[0x45];
      seq->voices[i].pan       = vce[0x46];
      seq->voices[i].transpose = vce[0x10];
      seq->voices[i].hasHits   = vce[0x25];
    }
  }

  // Read track list.
  seq->numTracks = *hdr++;
  for (int i = 0; i < seq->numTracks; ++i, hdr += 5) {
    seq->tracks[i].id = *(int32_t*)hdr;
  }

  return true;
}

// Insert sequence id as text event in meta track.
void setSequenceTitle(char* kmsBuf, sequence_t* seq)
{
  resource_t* kmsRes = (resource_t*)kmsBuf;
  seq->id = kmsRes->ids[DEFAULT_SEQ_POS];

  jdkmidi::MIDISystemExclusive sysEx((unsigned char*)&seq->id, 4, 4, false);
  jdkmidi::MIDITimedBigMessage m;

  m.CopySysEx(&sysEx);
  m.SetType(jdkmidi::META_EVENT);
  m.SetMetaType(META_TRACK_NAME);

  seq->metaTrack.PutEvent(m);
}

bool parseKmsSequence(char* seqBuf, sequence_t* seq)
{
  for (int i = 0; i < seq->numTracks; ++i) {
    char* trk = findResource(seqBuf, seq->tracks[i].id);
    if (trk == NULL) {
      std::cerr << "Couldn't find track \"" << seq->tracks[i].id << "\"." << std::endl;
      return false;
    }
    else {
      seq->tracks[i].enabled = true;
      parseKmsTrack(trk, &seq->tracks[i], seq);
    }
  }

  return true;
}

bool writeMidiFile(const char* filename, sequence_t* seq)
{
  int numMidiTracks = 0;
  jdkmidi::MIDIMultiTrack midiTracks(seq->numTracks + 1, false); // Don't take ownership of the tracks.

  midiTracks.SetTrack(numMidiTracks++, &seq->metaTrack);

  for (int i = 0; i < seq->numTracks; ++i) {
    if (seq->tracks[i].enabled) {
      midiTracks.SetTrack(numMidiTracks++, &seq->tracks[i].midi);
    }
  }

  jdkmidi::MIDIFileWriteStreamFileName outStream(filename);

  if (!outStream.IsValid()) {
    std::cerr << "Error: Can't open file \"" << filename << "\" for writing." << std::endl;
    return false;
  }

  jdkmidi::MIDIFileWriteMultiTrack writer(&midiTracks, &outStream);

  if (!writer.Write(numMidiTracks, CLOCK_RESOLUTION)) {
    std::cerr << "Error: Writing to file \"" << filename << "\" failed." << std::endl;
  }

  return true;
}

int main(int argc, char** argv)
{
  if (argc != 4) {
    std::cerr << "usage: " << argv[0] << " voices.vce song.kms output.mid" << std::endl;
    return 1;
  }

  char vceBuf[BUF_SIZE];
  char kmsBuf[BUF_SIZE];

  if (!readFile(argv[VCE_ARG], vceBuf)) {
    return 2;
  }

  if (!readFile(argv[KMS_ARG], kmsBuf)) {
    return 2;
  }

  sequence_t seq;
  char* seqRes = resourceAt(kmsBuf, DEFAULT_SEQ_POS);

  if (seqRes == NULL) {
    std::cerr << "Error: No sequence found in KMS file \"" << argv[KMS_ARG] << "\"." << std::endl;
    return 3;
  }

  setSequenceTitle(kmsBuf, &seq);

  if (!parseKmsHeader(seqRes, vceBuf, &seq)) {
    return 3;
  }

  if (!parseKmsSequence(seqRes, &seq)) {
    return 3;
  }

  if (!writeMidiFile(argv[OUT_ARG], &seq)) {
    return 4;
  }

  return 0;
}
