#!/usr/bin/python

#cereal basic
def is_cereal(cereal):
	return (type(cereal) == list) and (type(cereal[0]) == list) and (type(cereal[0][0]) == tuple)

def get_width(cereal):
	return len(cereal[0])

def get_height(cereal):
	return len(cereal)

def get_nchannel(cereal):
	return len(cereal[0][0])

def create_cereal(width,height,fill=None):
	new_cereal = []
	for y in range(height):
		new_cereal.append([])
		for x in range(width):
			new_cereal[y].append(())
			if fill:
				new_cereal[y][x] += (fill,)
	return new_cereal

def create_cereal_from(cereal,fill=None):
	#does not dupe data , just dupe size
	if not is_cereal(cereal):
		print "not a cereal"
		return
	w = get_width(cereal)
	h = get_height(cereal)
	new_cereal = create_cereal(w,h,fill)
	return new_cereal
	

#WAVE file
def load_wave(filename):
	import wave, audioop, struct
	wave_read = wave.open(filename,"r")
	(channels,bitrate,frequency,n,comptype,compname) = wave_read.getparams()
	if comptype != 'NONE' : return
	x,y = 0,0
	cereal = []
	cereal.append([])
	for i in range(n):
		sample = wave_read.readframes(1)	#if mono 8bits , it give 1 bytes , if stereo 16bits , it give 4 bytes
		#stereo
		if channels == 2:
			lsample = audioop.tomono(sample,bitrate,1,0)
			rsample = audioop.tomono(sample,bitrate,0,1)
			#16bits stereo
			if bitrate == 2:
				lfloat = (float(struct.unpack('h',lsample)[0]) / 65535) + 0.5
				rfloat = (float(struct.unpack('h',rsample)[0]) / 65535) + 0.5
			#8bits stereo
			elif bitrate == 1:
				lfloat = float(struct.unpack('B',lsample)[0]) / 255
				rfloat = float(struct.unpack('B',rsample)[0]) / 255
			#fuck the rest
			else:
				print("error bitrate not supported")
				return
			sample = (lfloat, rfloat)
		#mono
		elif channels == 1:
			#16bits mono
			if bitrate == 2:
				lfloat = (float(struct.unpack('h',sample)[0]) / 65535) + 0.5
			#8bits mono
			elif bitrate == 1:
				lfloat = float(struct.unpack('B',sample)[0]) / 255
			#fuck the rest
			else:
				print ("error bitrate not supported")
			sample = (lfloat,)
		#quadraphonic ,2.1 , 4.1 , 5.1 ..... nah !
		else:
			print("error , nchannel not supported")
			return
		#a tuple of every channel in integer between 0 and 256
		cereal[y].append(sample)
		x += 1
		if x == frequency:
			y +=1
			x = 0
			cereal.append([])
	wave_read.close()
	return cereal

def save_wave(cereal,filename):
	import wave, struct
	wave_write = wave.open(filename,"w")
	wave_write.setparams((len(cereal[0][0]),1,len(cereal[0]),0,'NONE','not compressed'))	#(channels,bytesrate 1=8bits 2=16bits,frequency,size,compression,compression type)
	n = 0
	for row in cereal:
		for col in row:
			sample = ""
			for value in col:
				sample += struct.pack('B', int(value*255))	#8bits
			wave_write.writeframesraw(sample)
			n += 1
	wave_write.close()

#image FILE
def load_image(filename):
	import Image
	image_read = Image.open(filename,"r")
	(w,h) = image_read.size
	mode = image_read.mode  #1 L RGB CMYK
	if mode == "P":
		print ("8bits palette not supported yet")
		return
	cereal = []
	image = image_read.load()
	for y in range(h):
		tmp = []
		for x in range(w):
			color = image[x,y]
			if (mode == "1") or (mode == "L"):
				tmp.append((color/255.0,))
			else:
				tttmp = ()
				for comp in color:
					tttmp += (comp/255.0,)
				tmp.append(tttmp)
		cereal.append(tmp)
	return cereal


def save_image(cereal,filename):
	import Image
	size = len(cereal[0][0])
	w = len(cereal[0])
	h = len(cereal)
	if size == 1:
		mode = "L"
	elif size == 3:
		mode = "RGB"
	elif size == 4:
		mode = "RGBA"
	else:
		print "no image format for these"
		return
	image_write = Image.new(mode,(w,h))
	image = image_write.load()
	for y in range(h):
		for x in range(w):
			if (mode == "L") or (mode == "1"):
				image[x,y] = int(cereal[y][x][0]*255)
			else:
				color = ()
				for comp in cereal[y][x]:
					color += (int(comp*255),)
				image[x,y] = color
	image_write.save(filename)

#cereal permutation
def join_row(cereal):
	if type(cereal[0]) != list:
		print "no row to join"
		return
	new_cereal = []
	for i in cereal:
		new_cereal += i
	return [new_cereal]

def split_in_row(cereal,lenght):
	if (type(cereal[0]) != list) and (type(cereal[0][0]) != tuple):
		print "the array must contain tuple"
		return
	new_cereal = []
	tmp_row = []
	pivot = 0
	for i in cereal[0]:
		if pivot < lenght:
			tmp_row.append(i)
			pivot += 1
		else:
			new_cereal.append(tmp_row)
			tmp_row = [i]
			pivot = 0
	return new_cereal

def swap_xy(cereal):
	if (type(cereal[0]) != list) and (type(cereal[0][0]) != tuple):
		print "the array must contain tuple"
		return
	new_cereal = []
	for col in range(len(cereal[0])):
		new_row = []
		for row in cereal:
			new_row.append(row.pop(0))
		new_cereal.append(new_row)
	return new_cereal

def invert_row(cereal):
	if (type(cereal[0]) != list) and (type(cereal[0][0]) != tuple):
		print "the array must contain tuple"
		return
	new_cereal = []
	for row in cereal:
		new_row = []
		while row:
			new_row.append(row.pop(-1))
		new_cereal.append(new_row)
	return new_cereal

def invert_col(cereal):
	if (type(cereal[0]) != list) and (type(cereal[0][0]) != tuple):
		print "the array must contain tuple"
		return
	new_cereal = []
	while cereal:
		new_cereal.append(cereal.pop(-1))
	return new_cereal

def get_channel(cereal, index=0):
	if (type(cereal[0]) != list) and (type(cereal[0][0]) != tuple):
		print "the array must contain tuple"
		return
	if index > len(cereal[0][0]):
		print "index not available"
		return
	new_cereal = []
	for row in cereal:
		new_row = []
		for col in row:
			new_row.append((col[index],))
		new_cereal.append(new_row)
	return new_cereal

def join_channel(cereal1, cereal2):
	if (type(cereal1[0]) != list) and (type(cereal1[0][0]) != tuple) and (type(cereal2[0]) != list) and (type(cereal2[0][0]) != tuple):
		print "the array must contain tuple"
		return
	if (len(cereal1) != len(cereal2)) and (len(cereal1[0]) != len(cereal2[0])):
		print "the two array must be the same size column and row"
		return
	new_cereal = []
	for y in range(len(cereal1)):
		new_row = []
		for x in range(len(cereal1[y])):
			new_row.append(cereal1[y][x] + cereal2[y][x])
		new_cereal.append(new_row)
	return new_cereal

def dupe(cereal):
	if not is_cereal(cereal):
		print "not a cereal array"
		return
	new_cereal = []
	for row in cereal:
		new_row = []
		for col in row:
			new_row.append(col)
		new_cereal.append(new_row)
	return new_cereal

#test
if __name__ == "__main__":
	#cereal = load_wave("8bits.wav")
	#save_wave(cereal,"out.wav")		#8bits only
	cereal = load_image("test.bmp")
	cerealR = get_channel(cereal,0)
	cerealG = get_channel(cereal,1)
	cerealB = get_channel(cereal,2)
	cereal_ = create_cereal_from(cerealR,0.0)
	cereal = join_channel(cerealB,cereal_)
	cereal = join_channel(cereal,cereal_)
	save_image(cereal,"out.bmp")
