import java.io.*;
import java.lang.*;
import java.awt.*;
import java.awt.image.*;

public class BmpToImage
{

  /* file layout variables */

  char letterB;
  char letterM;
  int fileSize;
  int reserved1;
  int reserved2;
  int pixelArrayOffset;
  int reserved3;
  int structSize;
  int imageWidth;
  int imageHeight;
  int nPlanes;
  int bitCount;
  int compression;
  int imageSize;
  int xPixelsPerMeter;
  int yPixelsPerMeter;
  int colorsUsed;
  int colorsImportant;

  int scanLinePadSize;

  int[] pixels;

  byte[] binary;
  int binaryIdx = 0;

  Image image;

  public BmpToImage(byte[] binary, Image image) {
    this.binary = binary;
    this.image = image;
    loadBinary();
  }

  private int loadByte() {
    int value = (int)binary[binaryIdx];
    if (value < 0)
      value += 256;
    binaryIdx++;
    return value;
  }

  private int loadShort() {
    return
      loadByte() |
      loadByte() << 8;
  }

  private int loadInteger() {
    return
      loadByte() |
      loadByte() << 8 |
      loadByte() << 16 |
      loadByte() << 24;
  }

  private int getScanPadSize (int bitCount, int imageWidth) {
    if (bitCount == 1)
      return (3 - (((imageWidth - 1) % 32) / 8));
    else /* assume 24 bit for now */
      return (imageWidth % 4);
  }

  private void plotPixel(int i, int j) {
    int r, g, b;
    r = loadByte();
    g = loadByte();
    b = loadByte();

    Graphics graphics = image.getGraphics();
    graphics.setColor(new Color(r, g, b));
    graphics.drawLine(i, j, i, j);
  }

  private void loadBinary() {
    int pixelSize = 3;

    letterB = 'B';
    letterM = 'M';
    if ((char)loadByte() != letterB) {
      System.out.println("source is not a .bmp file\n");
      return;
    }
    if ((char)loadByte() != letterM) {
      System.out.println("source is not a .bmp file\n");
      return;
    }
    fileSize = loadInteger();
    reserved1 = loadShort();
    reserved2 = loadShort();
    pixelArrayOffset = loadShort();
    reserved3 = loadShort();
    structSize = loadInteger();
    imageWidth = loadInteger();
    imageHeight = loadInteger();
    nPlanes = loadShort();
    bitCount = loadShort();
    scanLinePadSize = getScanPadSize(bitCount, imageWidth);
    compression = loadInteger();
    imageSize = loadInteger();
    xPixelsPerMeter = loadInteger();
    yPixelsPerMeter = loadInteger();
    colorsUsed = loadInteger();
    colorsImportant = loadInteger();

    if (bitCount != 24) {
      System.out.println("so far only support 24 bit\n");
      return;
    }

    int i, j;
    int padIdx;
    for (j = imageHeight - 1; j >= 0; j--)
      for (i = 0; i < imageWidth; i++) {
        plotPixel(i, j);
        for (padIdx = 0; padIdx < scanLinePadSize; padIdx++) // skip pad
          loadByte(); 
      }
  }
}
