# Import step-response and neopixel libraries
from steptime import STEPTIME
from ws2812 import WS2812
from ssd1306 import SSD1306_I2C

# Import native libraries
from machine import Pin, freq, I2C
import utime
from random import randint

# Set up clock frequency, needed for steptime lib
freq(250000000)

# Power up built-in XIAO NeoPixel LED
power = machine.Pin(11, machine.Pin.OUT)
power.value(1)

# Define colors, we need 6, one for each of the touch buttons
RED = (255, 0, 0)
YELLOW = (255, 150, 0)
GREEN = (0, 255, 0)
CYAN = (0, 255, 255)
BLUE = (0, 0, 255)
PURPLE = (180, 0, 255)
BLACK = (0, 0, 0) # And black, of course

# Init LED and set it to initial color
led = WS2812(12, 1, 0.5, 6) # args: pin_num, led_count, brightness, state_machine_id (6, since we take 0-5 with touch pads)

# Define function for setting neopixel color
def set_led_color(color):
    led.pixels_fill(color)
    led.pixels_show()

set_led_color(RED)

# Init OLED display
i2c = I2C(1, scl=Pin(7), sda=Pin(6), freq=200000)#Grove - OLED Display 0.96" (SSD1315)
oled = SSD1306_I2C(128, 64, i2c)

# Configure left side buttons
#   x    [26]
# x   x  [27, 1]
#   x    [2]
Pin(26,Pin.IN,Pin.PULL_UP)
Pin(27,Pin.IN,Pin.PULL_UP)
Pin(1,Pin.IN,Pin.PULL_UP)
Pin(2,Pin.IN,Pin.PULL_UP)

# Configure right side buttons
#    x   [4]
#  x     [3]
Pin(3,Pin.IN,Pin.PULL_UP)
Pin(4,Pin.IN,Pin.PULL_UP)

actions = [
    (set_led_color, RED),
    (set_led_color, YELLOW),
    (set_led_color, GREEN),
    (set_led_color, CYAN),
    (set_led_color, BLUE),
    (set_led_color, PURPLE),
]

# Create megastructure so we can loop through it comfortably
# StateMachine, MinVal, action, color, button state
channels = [
    [STEPTIME(0,26), [1e6], actions[0], [BLACK], [False]], # STEPTIME args: state_machine_id, pin_num
    [STEPTIME(1,27), [1e6], actions[1], [BLACK], [False]],
    [STEPTIME(2,1),  [1e6], actions[2], [BLACK], [False]],
    [STEPTIME(3,2),  [1e6], actions[3], [BLACK], [False]],
    [STEPTIME(4,3),  [1e6], actions[4], [BLACK], [False]],
    [STEPTIME(5,4),  [1e6], actions[5], [BLACK], [False]],
]

buttons = [0, 0, 0, 0, 0, 0]

loop = 200
settle = 20000
thresh = 10000

class Spaceship:
    def __init__ (self, height):
        self.height = height 
        self.counter = 0

    def get_height(self):
        return self.height

    def move(self, direction):
        if direction == "UP":
            if self.height > 0:
                self.height -= 2
        if direction == "DOWN": 
            if self.height < 58:
                self.height += 2

    def shoot(self):
        laser = 5
        while laser < 128:
            oled.text(".",laser, ship.get_height()-3)
            laser += 5

class Invader:
    def __init__(self):
        self.coordinates = [randint(128,168), randint(0,58)]
        self.speed = randint(1,3)

    def get_coordinates(self):
        return self.coordinates

    def explode(self):
        set_led_color(RED)
        
        self.coordinates = [150, 100]
        self.speed = 0

    def move(self):
        self.coordinates[0] -= self.speed
        if self.coordinates[0] < 5:
            lose()

class Fleet:
    def __init__(self):
        self.invaders = []
        
    def add_invader(self):
        self.invaders.append(Invader())

def process_touch():
    out_parts = []
    for ch in channels:
        sm, min_val, (fn, val), color, btn_state = ch
        sm.put(loop)
        sm.put(settle)
        result = 4294967296 - sm.get()
        if result < min_val[0]:
            min_val[0] = result    
        out_parts.append(str(result - min_val[0]))
        if result - min_val[0] > thresh:
            btn_state[0] = True
        else:
            btn_state[0] = False
    line = ",".join(out_parts)
    #print(f"7500,{line}") # 7500 for scale

def draw(ship, fleet):
    oled.fill(0)
    height = ship.get_height()
    oled.text(">",0, height)
    for invader in fleet.invaders: 
        x,y = invader.coordinates
        oled.text("#",x, y)
        invader.move()
    set_led_color(YELLOW)
    try:
        x,y = move_ship(ship, fleet)
        print(x, y)
        oled.text("*", x+4, y+4)
        oled.text("*", x-4, y+4)
        oled.text("*", x+4, y-4)
        oled.text("*", x-4, y-4)
    except:
        pass
    oled.show()

def move_ship(ship, fleet):
    if channels[0][4][0]:
        ship.move("UP")
    if channels[1][4][0]:
        ship.move("DOWN")
    if channels[5][4][0]:
        ship.shoot()
        for invader in fleet.invaders:
            if ship.height >= invader.coordinates[1]-2 and ship.height <= invader.coordinates[1]+2:
                x,y = invader.coordinates
                invader.explode()
                ship.counter += 1
                return x,y
    return False


def check_level(fleet):           # level tells the number of invader spawning on each level, counter tells how many killed cumulatively
    global level
    if ship.counter == 15:
        win()
    elif ship.counter < 1:          # level 1, counter 0
        if level < 1:             # this if-statement ensures invaders are spawned only once per level
            level += 1
            fleet.add_invader()
            print("Level: 1")        
    elif ship.counter == 1:          # level 2, counter 1
        if level < 2:
            level += 1
            for i in range(level): 
                fleet.add_invader() # same number of invaders spawned as the level
                print("Level: 2")        
    elif ship.counter == 3:          # level 3, counter 3
        if level < 3:
            level += 1
            for i in range(level):
                fleet.add_invader()
                print("Level: 3")        
    elif ship.counter == 6:          # level 4, counter 6
        if level < 4:
            level += 1
            for i in range(level): 
                fleet.add_invader()
                print("Level: 4")
    elif ship.counter == 10:           # level 5, counter 10
        if level < 5:
            level += 1
            for i in range(level): 
                fleet.add_invader()
                print("Level: 5")

def lose():
    global game_on
    game_on = False
    while not game_on:
        oled.fill(0)
        set_led_color(RED)
        oled.text("GAME OVER", 30, 18)
        oled.text("Press any button", 0, 44)
        oled.text("to restart", 20, 54)
        oled.show()
        utime.sleep(0.5)
        if restart():
            game()


def win():
    end = True 
    global game_on
    game_on = False
    while not game_on:
        oled.fill(0)
        set_led_color(CYAN)
        oled.text("VICTORY!!!", 25, 18)
        oled.text("Press any button", 0, 44)
        oled.text("to restart", 20, 54)
        oled.show()
        utime.sleep(0.5)
        if restart():
            game()


def restart():
    set_led_color(BLACK)
    global game_on
    process_touch()
    if any(ch[4][0] for ch in channels):
        ship = Spaceship(32)
        fleet = Fleet()
        level = 0
        return True
    utime.sleep(0.5)
    return False


def game():
    global game_on
    game_on = True
    global level
    level = 0
    ship.counter = 0
    ship.height = 32
    fleet = Fleet()
    while game_on:
        process_touch()
        check_level(fleet)
        draw(ship, fleet)


ship = Spaceship(32)
fleet = Fleet()
level = 0
game_on = True
game()
