from time import sleep, ticks_ms, ticks_diff
from machine import Pin, SPI
from ssd1309 import Display

CLK = Pin(10, Pin.IN, Pin.PULL_UP)
DT  = Pin(11, Pin.IN, Pin.PULL_UP)
SW  = Pin(12, Pin.IN, Pin.PULL_UP)

spi = SPI(0, baudrate=10000000, polarity=0, phase=0, sck=Pin(18), mosi=Pin(19))
display = Display(spi, dc=Pin(16), cs=Pin(17), rst=Pin(20))
display.flip()

last_encoded = 0
encoder_delta = 0

def handle_encoder(pin):
    global last_encoded, encoder_delta
    MSB = CLK.value()
    LSB = DT.value()
    encoded = (MSB << 1) | LSB
    sum = (last_encoded << 2) | encoded

    if sum in (0b1101, 0b0100, 0b0010, 0b1011):
        encoder_delta += 1
    elif sum in (0b1110, 0b0111, 0b0001, 0b1000):
        encoder_delta -= 1

    last_encoded = encoded

CLK.irq(trigger=Pin.IRQ_RISING | Pin.IRQ_FALLING, handler=handle_encoder)
DT.irq(trigger=Pin.IRQ_RISING | Pin.IRQ_FALLING,  handler=handle_encoder)

cx = display.width // 2
cy = display.height // 2
crosshair_index = 0
crosshair_size = 10
MIN_SIZE = 3
MAX_SIZE = 20

MODE_NORMAL  = 0
MODE_CALIB_Y = 1
MODE_CALIB_X = 2
mode = MODE_NORMAL

btn_down = False
btn_press_time = 0
LONG_PRESS_MS = 700
DEBOUNCE_MS   = 50
last_btn_change = 0

def draw_plus(cx, cy, size):
    display.draw_hline(cx - size, cy, size * 2)
    display.draw_vline(cx, cy - size, size * 2)
    display.draw_pixel(cx, cy)

def draw_dot(cx, cy, size):
    display.fill_circle(cx, cy, size)

def draw_heart(cx, cy, size):
    scale = size / 6.0
    for dy in range(-size, size + 1):
        for dx in range(-size, size + 1):
            x = dx / scale
            y = -dy / scale
            if (x*x + y*y - 1)**3 - x*x * y*y*y <= 0:
                px, py = cx + dx, cy + dy
                if 0 <= px < display.width and 0 <= py < display.height:
                    display.draw_pixel(px, py)

def draw_circle_crosshair(cx, cy, size):
    display.draw_circle(cx, cy, size)
    tick = 4
    display.draw_hline(cx - size - tick, cy, tick)
    display.draw_hline(cx + size + 1,    cy, tick)
    display.draw_vline(cx, cy - size - tick, tick)
    display.draw_vline(cx, cy + size + 1,    tick)

def draw_cross(cx, cy, size):
    """Multiplication sign / X crosshair."""
    for i in range(-size, size + 1):
        display.draw_pixel(cx + i, cy + i)
        display.draw_pixel(cx + i, cy - i)

def draw_bracket_dot(cx, cy, size):
    """Square brackets with dot in center."""
    arm = size // 2  #length of bracket crosshiar
    
    #Left
    display.draw_vline(cx - size, cy - size, size * 2)      
    display.draw_hline(cx - size, cy - size, arm)           
    display.draw_hline(cx - size, cy + size, arm)           

    #Right
    display.draw_vline(cx + size, cy - size, size * 2)      
    display.draw_hline(cx + size - arm, cy - size, arm)     
    display.draw_hline(cx + size - arm, cy + size, arm)     

    #dot
    display.fill_circle(cx, cy, max(2, size // 5))


crosshairs = [
    ("Plus",    draw_plus),
    ("Dot",     draw_dot),
    ("Circle",  draw_circle_crosshair),
    ("Heart",   draw_heart),
    ("Cross",   draw_cross),
    ("Bracket", draw_bracket_dot),
]

def render():
    display.clear_buffers()

    if mode != MODE_NORMAL:
        display.draw_rectangle(0, 0, display.width - 1, display.height - 1)

    name, draw_fn = crosshairs[crosshair_index]
    draw_fn(cx, cy, crosshair_size)

    if mode == MODE_CALIB_Y:
        display.draw_text8x8(2, 2, "Cal Y")
    elif mode == MODE_CALIB_X:
        display.draw_text8x8(2, 2, "Cal X")
    else:
        size_label = str(crosshair_size) + "x"
        x = display.width - len(size_label) * 8 - 2
        display.draw_text8x8(x, display.height - 10, size_label)

    display.present()

render()

while True:
    now = ticks_ms()

    if encoder_delta != 0:
        delta = encoder_delta
        encoder_delta = 0

        if mode == MODE_NORMAL:
            crosshair_size = max(MIN_SIZE, min(MAX_SIZE, crosshair_size + delta))
        elif mode == MODE_CALIB_Y:
            cy = max(0, min(display.height - 1, cy + delta))
        elif mode == MODE_CALIB_X:
            cx = max(0, min(display.width - 1, cx + delta))
        render()
    #debounce
    sw_val = SW.value()
    if ticks_diff(now, last_btn_change) > DEBOUNCE_MS:
        if sw_val == 0 and not btn_down:
            btn_down = True
            btn_press_time = now
            last_btn_change = now

        elif sw_val == 1 and btn_down:
            held = ticks_diff(now, btn_press_time)
            btn_down = False
            last_btn_change = now

            if held >= LONG_PRESS_MS:
                #calibration
                if mode == MODE_NORMAL:
                    mode = MODE_CALIB_Y   #enter calibration of Y-axis
                else:
                    mode = MODE_NORMAL    #exit
            else:
                if mode == MODE_NORMAL:
                    crosshair_index = (crosshair_index + 1) % len(crosshairs)
                elif mode == MODE_CALIB_Y:
                    mode = MODE_CALIB_X   # switch to X axis
                elif mode == MODE_CALIB_X:
                    mode = MODE_CALIB_Y   # switch back to Y axis

            render()

    sleep(0.001)
