import board
import analogio
import digitalio
import time

# -- Pin definitions (Matching your diagram exactly) --
# Columns (Vertical traces - Bottom layer)
COL_PINS = [board.D10, board.D9, board.D8, board.D7]
# Rows (Horizontal traces - Top layer)
ROW_PINS = [board.D6, board.D5, board.D4, board.D3]

# Initialize Drive Pins as Outputs
cols = [digitalio.DigitalInOut(p) for p in COL_PINS]
rows = [digitalio.DigitalInOut(p) for p in ROW_PINS]

for p in cols + rows:
    p.direction = digitalio.Direction.OUTPUT
    p.value = False

# Pen ADC (A0 / D0)
pen_adc = analogio.AnalogIn(board.A0)

# -- Tuning --
SAMPLES = 20
SETTLE_S = 0.000050  # 50 microseconds (CircuitPython uses seconds)
THRESHOLD = 1000     # CircuitPython ADC is 16-bit (0-65535)
                     # Start high and calibrate down.

# -- ADC helper --
def read_adc():
    """Average SAMPLES from the 16-bit ADC."""
    total = 0
    for _ in range(SAMPLES):
        total += pen_adc.value
    return total // SAMPLES

# -- Single electrode scan step --
def read_electrode(pin_obj):
    """Energize one pin, wait, read, and discharge."""
    pin_obj.value = True
    time.sleep(SETTLE_S)
    val = read_adc()
    pin_obj.value = False
    return val
    
# -- Full scan --
def scan():
    col_vals = [read_electrode(p) for p in cols]
    row_vals = [read_electrode(p) for p in rows]
    return col_vals, row_vals

# -- Position decode (Weighted Centroid) --
def weighted_centroid(vals):
    total, wsum = 0, 0
    for i, v in enumerate(vals):
        if v > THRESHOLD:
            wsum += i * v
            total += v
    return (wsum / total) if total > 0 else None

# -- Main Loop --
print("RAND tablet — XIAO nRF52840 CircuitPython")
print("Calibration Mode: RAW_MODE = True")

# Set to False once you determine your THRESHOLD
RAW_MODE = False


# -- New Calibration Logic --
BASELINES_COL = [0, 0, 0, 0]
BASELINES_ROW = [0, 0, 0, 0]
SENSITIVITY = 1500  # We look for a signal that is 1500 ABOVE the baseline

def calibrate():
    global BASELINES_COL, BASELINES_ROW
    print("Calibrating... KEEP PEN AWAY")
    time.sleep(1.0)
    c, r = scan()
    BASELINES_COL, BASELINES_ROW = c, r
    print("Calibration Done.")

def weighted_centroid(vals, baselines):
    total, wsum = 0, 0
    for i, v in enumerate(vals):
        # Calculate how much the signal changed from the idle state
        diff = v - baselines[i]
        
        if diff > SENSITIVITY:
            wsum += i * diff
            total += diff
    return (wsum / total) if total > 0 else None

# -- In your main loop --
calibrate() # Run this once at start

while True:
    col_vals, row_vals = scan()
    x = weighted_centroid(col_vals, BASELINES_COL)
    y = weighted_centroid(row_vals, BASELINES_ROW)
    if RAW_MODE:
        # Print raw values for calibration
        c_str = " ".join([f"C{i}:{v:5d}" for i, v in enumerate(col_vals)])
        r_str = " ".join([f"R{i}:{v:5d}" for i, v in enumerate(row_vals)])
        print(f"{c_str} | {r_str}")
    else:    
        if x is not None and y is not None:
            print(f"Pos: ({x:.2f}, {y:.2f})")
        else:
            print("-- no pen --")
        time.sleep(0.1)
    
# 
# while True:
# 
#     col_vals, row_vals = scan()
#     x = weighted_centroid(col_vals)
#     y = weighted_centroid(row_vals)
#     
# 
#     if RAW_MODE:
#         # Print raw values for calibration
#         c_str = " ".join([f"C{i}:{v:5d}" for i, v in enumerate(col_vals)])
#         r_str = " ".join([f"R{i}:{v:5d}" for i, v in enumerate(row_vals)])
#         print(f"{c_str} | {r_str}")
#     else:
#         # Print calculated position
#         if x is not None and y is not None:
#             print(f"Pos: ({x:.2f}, {y:.2f})")
#         else:
#             print("-- no pen --")
# 
#     time.sleep(0.2) # ~50Hz