from math import atan2, sqrt, copysign
from controller import Robot

MAX_SPEED = 6.28
BASE_FWD = 0.65 * MAX_SPEED
TURN_GAIN = 3.2
SLOW_TURN_GAIN = 2.0
AVOID_GAIN = 0.9
EDGE_BACK_TIME = 10
BALL_LOST_SPIN = 0.45 * MAX_SPEED
BALL_CLOSE_DIST = 0.25
BALL_SEEN_HOLD = 30
FIELD_X_HALF = 1.0
FIELD_Z_HALF = 0.7
SAFE_MARGIN = 0.08

LEFT_MOTOR_NAMES  = ['left wheel motor','left motor','left_wheel','left_wheel_motor','motor_left','left']
RIGHT_MOTOR_NAMES = ['right wheel motor','right motor','right_wheel','right_wheel_motor','motor_right','right']
SONAR_NAME_CANDIDATES = [[f'us{i}' for i in range(16)],[f'us{i}' for i in range(8)],[f'ps{i}' for i in range(8)]]
CAMERA_NAMES = ['camera','cam','head_camera','ball_cam']
COMPASS_NAMES = ['compass']
GPS_NAMES = ['gps']

def try_get_device(robot, names):
    for n in names:
        try:
            d = robot.getDevice(n)
            if d is not None:
                return d
        except Exception:
            pass
    return None

robot = Robot()
TIME_STEP = int(robot.getBasicTimeStep()) if robot.getBasicTimeStep() > 0 else 32

left_motor = try_get_device(robot, LEFT_MOTOR_NAMES)
right_motor = try_get_device(robot, RIGHT_MOTOR_NAMES)
if left_motor is None or right_motor is None:
    raise RuntimeError("Wheel motors not found.")

left_motor.setPosition(float('inf'))
right_motor.setPosition(float('inf'))
left_motor.setVelocity(0.0)
right_motor.setVelocity(0.0)

sonars = []
for candidate in SONAR_NAME_CANDIDATES:
    tmp = []
    for name in candidate:
        try:
            s = robot.getDevice(name)
            if s:
                s.enable(TIME_STEP)
                tmp.append(s)
        except Exception:
            pass
    if len(tmp) >= 2:
        sonars = tmp
        break

compass = try_get_device(robot, COMPASS_NAMES)
if compass:
    compass.enable(TIME_STEP)

gps = try_get_device(robot, GPS_NAMES)
if gps:
    gps.enable(TIME_STEP)

camera = try_get_device(robot, CAMERA_NAMES)
if camera:
    camera.enable(TIME_STEP)
    try:
        camera.recognitionEnable(TIME_STEP)
    except Exception:
        pass

initial_x = None
opponent_goal = [0.0, 0.0]
last_ball_angle = None
last_ball_seen_ticks = 0
edge_timer = 0

def heading_yaw():
    if not compass:
        return None
    v = compass.getValues()
    return atan2(v[0], v[2])

def gps_xz():
    if not gps:
        return (0.0, 0.0)
    p = gps.getValues()
    return (p[0], p[2])

def ball_from_camera():
    if not camera:
        return (None, None)
    try:
        objs = camera.getRecognitionObjects()
    except Exception:
        return (None, None)
    if not objs:
        return (None, None)
    best = None
    best_area = -1
    for o in objs:
        pos = o.get_position() if hasattr(o, 'get_position') else None
        size = o.get_size() if hasattr(o, 'get_size') else None
        area = 0.0 if not size else size[0] * size[1]
        if pos is not None and area >= best_area:
            best = (pos, area)
    if best is None:
        return (None, None)
    bx, by, bz = best[0]
    if bz <= 0:
        return (None, None)
    ang = atan2(bx, bz)
    dist = sqrt(bx*bx + bz*bz)
    return (ang, dist)

def wall_avoid_vector(x, z):
    b = 0.0
    if abs(x) > FIELD_X_HALF - SAFE_MARGIN:
        b += copysign(1.0, -x)
    if abs(z) > FIELD_Z_HALF - SAFE_MARGIN:
        b += copysign(0.7, -z)
    return b

def obstacle_avoid_bias():
    if not sonars:
        return 0.0
    n = len(sonars)
    if n < 2:
        return 0.0
    vals = [s.getValue() for s in sonars]
    mid = n // 2
    left = sum(vals[:mid]) / max(1, mid)
    right = sum(vals[mid:]) / max(1, n - mid)
    return (right - left)

def set_wheels(vl, vr):
    left_motor.setVelocity(max(-MAX_SPEED, min(MAX_SPEED, vl)))
    right_motor.setVelocity(max(-MAX_SPEED, min(MAX_SPEED, vr)))

while robot.step(TIME_STEP) != -1:
    x, z = gps_xz()
    if initial_x is None and gps:
        initial_x = x
        opponent_goal = [-copysign(FIELD_X_HALF - 0.05, initial_x), 0.0]

    near_edge = (abs(x) > FIELD_X_HALF - SAFE_MARGIN) or (abs(z) > FIELD_Z_HALF - SAFE_MARGIN)
    if near_edge and edge_timer == 0:
        edge_timer = EDGE_BACK_TIME

    if edge_timer > 0:
        to_center_ang = 0.0
        if compass:
            yaw = heading_yaw()
            vecx = -x
            vecz = -z
            desired = atan2(vecx, vecz)
            err = desired - yaw
            while err > 3.14159: err -= 2*3.14159
            while err < -3.14159: err += 2*3.14159
            to_center_ang = err
        turn = TURN_GAIN * to_center_ang
        set_wheels(-0.4*MAX_SPEED - turn, -0.4*MAX_SPEED + turn)
        edge_timer -= 1
        continue

    avoid = obstacle_avoid_bias() * AVOID_GAIN

    ball_ang, ball_dist = ball_from_camera()
    if ball_ang is not None:
        last_ball_angle = ball_ang
        last_ball_seen_ticks = BALL_SEEN_HOLD
    else:
        last_ball_seen_ticks = max(0, last_ball_seen_ticks - 1)

    if last_ball_angle is None:
        set_wheels(BALL_LOST_SPIN, -BALL_LOST_SPIN)
        continue

    target_ang = last_ball_angle

    if ball_dist is not None and ball_dist < BALL_CLOSE_DIST and gps and compass:
        yaw = heading_yaw()
        vecx = opponent_goal[0] - x
        vecz = opponent_goal[1] - z
        desired = atan2(vecx, vecz)
        err_goal = desired - yaw
        while err_goal > 3.14159: err_goal -= 2*3.14159
        while err_goal < -3.14159: err_goal += 2*3.14159
        target_ang = 0.65 * target_ang + 0.35 * err_goal

    target_ang += wall_avoid_vector(x, z) * 0.4
    target_ang += avoid * 0.002

    gain = SLOW_TURN_GAIN if (ball_dist is not None and ball_dist < BALL_CLOSE_DIST) else TURN_GAIN
    turn = gain * target_ang
    fwd = BASE_FWD * (0.8 if abs(target_ang) > 0.6 else 1.0)

    vl = fwd - turn
    vr = fwd + turn
    set_wheels(vl, vr)
