import pygame
import math
BASE_TILE = 50

class Player(pygame.sprite.Sprite):
    def __init__(self, x, y, tile_size=None):
        super().__init__()
        # Accept optional tile_size so player base image matches tiles
        self._base_img = None
        self._base_dead_img = None
        self.reset(x, y, tile_size)
        self.direction = 1
        self.shoot_cooldown = 0
        self._scale_cache = {"scale": None, "image": None, "dead": None}
        # CHANGED: secret code tracking (M, A, T, U, S)
        self.secret_sequence = [
            pygame.K_m,
            pygame.K_a,
            pygame.K_t,
            pygame.K_u,
            pygame.K_s,
        ]
        self._secret_buffer = []
        self.secret_active = False
        # NEW: invincibility duration (ms)
        self.invincible_duration_ms = 2000
        self.invincible_until_ms = 0
        # class‑level damage sound cache
        if not hasattr(Player, "damage_sound"):
            Player.damage_sound = None  # None = not loaded, False = permanently unavailable

    # NEW: helper to ensure damage sound is loaded once
    @classmethod
    def _ensure_damage_sound(cls):
        if getattr(cls, "damage_sound", None) is False:
            return
        if not pygame.mixer.get_init():
            cls.damage_sound = False
            return
        if isinstance(cls.damage_sound, pygame.mixer.Sound):
            return
        try:
            cls.damage_sound = pygame.mixer.Sound(r"Sounds/damage_sound.mp3")
        except Exception:
            try:
                cls.damage_sound = pygame.mixer.Sound("Sounds/damage_sound.mp3")
            except Exception:
                cls.damage_sound = False

    @classmethod
    def _play_damage_sound_safe(cls):
        """Play damage sound without importing RARL; safe on first use."""
        if not pygame.mixer.get_init():
            return
        cls._ensure_damage_sound()
        if not isinstance(cls.damage_sound, pygame.mixer.Sound):
            return
        try:
            cls.damage_sound.play()
        except Exception:
            # never let audio errors affect game flow
            pass

    def update(self, game_over, world, groups, events=None):
        dx = 0
        dy = 0
        screen = pygame.display.get_surface()
        world_h = world.pixel_height
        world_w = world.pixel_width
        scale = getattr(world, "scale", 1.0)

        tile_size = getattr(world, "tile_size", BASE_TILE)
        s = float(tile_size) / float(BASE_TILE)

        move_px = max(1, int(round(5 * s)))
        jump_px = max(1, int(round(15 * s)))
        proj_offset_x = max(1, int(round(8 * s)))
        proj_offset_y = max(1, int(round(5 * s)))

        # use shared events list passed from main loop (RARL.py)
        if events is None:
            events = []

        if game_over == 0:
            # CHANGED: use events from main loop for secret combo
            for ev in events:
                if ev.type != pygame.KEYDOWN:
                    continue
                # record only the keys we care about
                if ev.key in (pygame.K_m, pygame.K_a, pygame.K_t, pygame.K_u, pygame.K_s):
                    self._secret_buffer.append(ev.key)
                    if len(self._secret_buffer) > len(self.secret_sequence):
                        self._secret_buffer.pop(0)
                    if self._secret_buffer == self.secret_sequence and not self.secret_active:
                        # activate secret: 1 HP, but insane projectile damage
                        self.secret_active = True
                        self.max_health = 1
                        self.health = 1

            key = pygame.key.get_pressed()
            now_ms = pygame.time.get_ticks()

            # Jump (Space or W)
            if (key[pygame.K_SPACE] or key[pygame.K_w]) and not self.jumped and not self.in_air:
                self.vel_y = -jump_px
                self.jumped = True
            if not key[pygame.K_SPACE] and not key[pygame.K_w]:
                self.jumped = False

            # Horizontal move (A/D)
            if key[pygame.K_a]:
                dx -= move_px
                self.direction = -1
            if key[pygame.K_d]:
                dx += move_px
                self.direction = 1

            # Gravity
            GRAVITY = 1.0 * s
            MAX_FALL = 10.0 * s
            self.vel_y = min(self.vel_y + GRAVITY, MAX_FALL)
            dy += self.vel_y
            self.in_air = True

            touched_enemy = False
            enemies = groups.get('enemies', [])

            # MOVE X with pixel stepping
            if dx != 0:
                step_x = 1 if dx > 0 else -1
                steps = abs(int(dx))
                for _ in range(steps):
                    self.rect.x += step_x
                    collided = False
                    # collide with tiles ONLY (enemies no longer block)
                    for tile_rect in getattr(world, "collision_rects", [r for _, r in getattr(world, "tile_list", [])]):
                        if self.rect.colliderect(tile_rect):
                            self.rect.x -= step_x
                            collided = True
                            break
                    if collided:
                        break

            # MOVE Y with pixel stepping
            dy_steps = int(abs(round(dy)))
            if dy_steps != 0:
                step_y = 1 if dy > 0 else -1
                for _ in range(dy_steps):
                    self.rect.y += step_y
                    collided = False
                    # collide with tiles ONLY
                    for tile_rect in getattr(world, "collision_rects", [r for _, r in getattr(world, "tile_list", [])]):
                        if self.rect.colliderect(tile_rect):
                            self.rect.y -= step_y
                            self.vel_y = 0
                            if step_y > 0:
                                self.in_air = False
                            collided = True
                            break
                    if collided:
                        break

            # Clamp to world bounds (hitbox)
            if self.rect.left < 0:
                self.rect.left = 0
            if self.rect.right > world_w:
                self.rect.right = world_w
            if self.rect.bottom > world_h:
                self.rect.bottom = world_h
                self.vel_y = 0
                self.in_air = False

            # NEW: enemies are non-solid, but still deal touch damage
            for enemy in enemies:
                if self.rect.colliderect(enemy.rect):
                    touched_enemy = True
                    break

            # Contact damage with cooldown + teleport to spawn
            if touched_enemy and (now_ms - self.last_touch_damage_ms) >= self.touch_damage_cooldown_ms:
                self.take_damage(1)

            # Hazard (spike) damage: check against world's hazard_list (logic-space rects)
            for hazard_rect in getattr(world, 'hazard_list', []):
                # use a slightly forgiving trigger area for hazards
                trigger = hazard_rect.inflate(2, 2)
                if self.rect.colliderect(trigger):
                    if (now_ms - self.last_touch_damage_ms) >= self.touch_damage_cooldown_ms:
                        self.take_damage(1)
                    break

            # Reload (R key) - use configured duration
            if key[pygame.K_r] and not self.reloading and self.ammo < self.max_ammo:
                self.reloading = True
                self.reload_end_ms = now_ms + self.reload_duration_ms
            if self.reloading and now_ms >= self.reload_end_ms:
                self.ammo = self.max_ammo
                self.reloading = False

            # Shooting (J key), only if not reloading and ammo available
            if self.shoot_cooldown > 0:
                self.shoot_cooldown -= 1
            if key[pygame.K_j] and self.shoot_cooldown == 0 and not self.reloading and self.ammo > 0:
                # spawn projectile slightly offset to avoid immediate tile collision
                proj_y = self.rect.centery - proj_offset_y  # scaled vertical offset
                if self.direction == 1:
                    proj_x = self.rect.right + proj_offset_x
                else:
                    proj_x = self.rect.left - Projectile.SIZE[0] - proj_offset_x
                # create a player-owned projectile
                p = Projectile(proj_x, proj_y, self.direction, owner='player')
                # keep high damage when secret is active
                if self.secret_active:
                    p.damage = 999999
                groups['projectiles'].add(p)
                self.shoot_cooldown = 12
                self.ammo -= 1

        elif game_over == -1:
            # only switch to dead image, do NOT move player (no float)
            self.image = self.dead_image

        # Draw player scaled with cache
        if self._scale_cache["scale"] != scale:
            self._scale_cache["scale"] = scale
            sw = max(1, int(round(self.width * scale)))
            sh = max(1, int(round(self.height * scale)))
            self._scale_cache["image"] = pygame.transform.smoothscale(self.image, (sw, sh))
            dw = max(1, int(round(self.dead_image.get_width() * scale)))
            dh = max(1, int(round(self.dead_image.get_height() * scale)))
            self._scale_cache["dead"] = pygame.transform.smoothscale(self.dead_image, (dw, dh))

        off_x = getattr(world, "viewport_off_x", 0)
        off_y = getattr(world, "viewport_off_y", 0)
        sx = int(round(self.rect.x * scale)) + off_x
        sy = int(round(self.rect.y * scale)) + off_y
        ox = int(round(self.image_offset[0] * scale))
        oy = int(round(self.image_offset[1] * scale))

        # base image: alive vs dead
        base_img = self._scale_cache["image"] if game_over == 0 else self._scale_cache["dead"]
        # flip horizontally when facing left (direction == -1)
        if getattr(self, "direction", 1) < 0:
            try:
                img_to_draw = pygame.transform.flip(base_img, True, False)
            except Exception:
                img_to_draw = base_img
        else:
            img_to_draw = base_img

        screen.blit(img_to_draw, (sx - ox, sy - oy))

        # Reload bar
        if self.reloading:
            now_ms = pygame.time.get_ticks()
            total = max(1, self.reload_duration_ms)
            remain = max(0, self.reload_end_ms - now_ms)
            progress = 1.0 - (remain / total)
            base_w, base_h = 60, 8
            bar_w = max(2, int(round(base_w * scale)))
            bar_h = max(2, int(round(base_h * scale)))
            bx = sx + img_to_draw.get_width() // 2 - bar_w // 2
            by = sy - bar_h - max(2, int(round(6 * scale)))
            pygame.draw.rect(screen, (60, 60, 60), (bx, by, bar_w, bar_h))
            pygame.draw.rect(screen, (255, 220, 0), (bx, by, int(round(bar_w * progress)), bar_h))
            pygame.draw.rect(screen, (0, 0, 0), (bx, by, bar_w, bar_h), max(1, int(round(2 * scale))))

        # trigger game over when HP depleted
        if self.health <= 0:
            game_over = -1

        # Enemy projectile damage: if any projectile in groups['projectiles'] is owned by 'enemy'
        # and collides with the player, apply 1 damage (respecting touch cooldown) and remove it.
        try:
            proj_group = groups.get('projectiles') if isinstance(groups, dict) else None
            if proj_group and isinstance(proj_group, pygame.sprite.Group):
                now_ms = pygame.time.get_ticks()
                for p in list(proj_group):
                    try:
                        if getattr(p, 'owner', None) == 'enemy' and p.rect.colliderect(self.rect):
                            # only apply damage when cooldown passed
                            if (now_ms - self.last_touch_damage_ms) >= self.touch_damage_cooldown_ms:
                                self.take_damage(1)
                            try:
                                p.kill()
                            except Exception:
                                pass
                    except Exception:
                        continue
        except Exception:
            pass

        return game_over

    def reset(self, x, y, tile_size=None):
        # Load original images once, then scale according to tile_size
        try:
            img = pygame.image.load(r"Textures/main_character_pistol.png").convert_alpha()
        except Exception:
            try:
                img = pygame.image.load(r"Textures/main_character_pistol.png").convert_alpha()
            except Exception:
                img = pygame.Surface((45, 50), pygame.SRCALPHA)
                img.fill((200, 200, 200))
        self._base_img = img

        try:
            img2 = pygame.image.load(r"Textures/main_character_dead.png").convert_alpha()
        except Exception:
            try:
                img2 = pygame.image.load(r"Textures/main_character_dead.png").convert_alpha()
            except Exception:
                img2 = self._base_img.copy()
        self._base_dead_img = img2

        # Scale to tile_size if provided; otherwise use previous fallback size
        if tile_size:
            w = max(16, int(round(tile_size * 0.9)))
            h = max(16, int(round(tile_size * 1.05)))
        else:
            w, h = 45, 50

        self.image = pygame.transform.scale(self._base_img, (w, h))
        try:
            self.dead_image = pygame.transform.scale(self._base_dead_img, (int(round(w * 1.1)), int(round(h * 1.6))))
        except Exception:
            self.dead_image = self.image.copy()

        # Create a smaller collision hitbox and set self.rect to that box (logic-space)
        # ORIGINAL:
        # hit_w = max(8, int(round(w * 0.6)))
        # hit_h = max(8, int(round(h * 0.9)))
        # NEW: shrink by extra 10 px in both width and height
        hit_w = max(8, int(round(w * 0.6)) - 10)
        hit_h = max(8, int(round(h * 0.9)) - 10)

        self.rect = pygame.Rect(0, 0, hit_w, hit_h)  # hitbox used for physics/collisions
        # If a tile_size is provided, treat x,y as tile coords and position player so feet touch tile bottom
        if tile_size:
            col = x // tile_size
            row = y // tile_size
            self.spawn_tile = (col, row)
            tile_top = row * tile_size
            tile_bottom = tile_top + tile_size
            # place hitbox so its bottom touches tile bottom
            self.rect.x = col * tile_size + (w - hit_w) // 2
            self.rect.y = tile_bottom - hit_h
            # store pixel spawn (hitbox-based)
            self.spawn_x = self.rect.x
            self.spawn_y = self.rect.y
        else:
            # fallback: store pixel spawn and no tile info
            self.spawn_tile = None
            self.rect.x = x
            self.rect.y = y
            self.spawn_x = x
            self.spawn_y = y

        # compute image offset so we can draw image above/around hitbox
        self.width = self.image.get_width()
        self.height = self.image.get_height()
        self.image_offset = ((self.width - hit_w) // 2, self.height - hit_h)

        # Health and ammo
        self.health = 5          # CHANGED: start with 5 HP
        self.max_health = 5      # CHANGED: max HP 5
        self.max_ammo = 10
        self.ammo = self.max_ammo
        self.reloading = False
        self.reload_end_ms = 0
        self.reload_duration_ms = 3000
        self.touch_damage_cooldown_ms = 700
        self.last_touch_damage_ms = 0
        # CHANGED: reset secret state and buffer on reset
        self.secret_active = False
        self._secret_buffer = []
        # NEW: reset invincibility timer
        self.invincible_until_ms = 0
        try:
            self.reload_font = pygame.font.SysFont(None, 22)
        except Exception:
            pygame.font.init()
            self.reload_font = pygame.font.SysFont(None, 22)
        self.reload_label = self.reload_font.render("Reloading...", True, (0, 0, 0))

        self.vel_y = 0
        self.jumped = False
        self.in_air = True

        # Clear scale cache so new sizes are used immediately
        self._scale_cache = {"scale": None, "image": None, "dead": None}

    def resize(self, tile_size):
        """Rescale player's sprites to match given tile_size without resetting gameplay state."""
        if tile_size is None:
            return
        w = max(16, int(round(tile_size * 0.9)))
        h = max(16, int(round(tile_size * 1.05)))
        # Rescale base images into current image/dead_image
        try:
            self.image = pygame.transform.scale(self._base_img, (w, h))
        except Exception:
            self.image = pygame.Surface((w, h), pygame.SRCALPHA)
            self.image.fill((200, 200, 200))
        try:
            self.dead_image = pygame.transform.scale(self._base_dead_img, (int(round(w * 1.1)), int(round(h * 1.6))))
        except Exception:
            self.dead_image = self.image.copy()

        # Update image size and hitbox size
        self.width = self.image.get_width()
        self.height = self.image.get_height()
        # ORIGINAL:
        # hit_w = max(8, int(round(self.width * 0.6)))
        # hit_h = max(8, int(round(self.height * 0.9)))
        # NEW: same shrink of 10 px as in reset
        hit_w = max(8, int(round(self.width * 0.6)) - 10)
        hit_h = max(8, int(round(self.height * 0.9)) - 10)

        # keep same spawn tile if available
        if getattr(self, "spawn_tile", None) is not None:
            col, row = self.spawn_tile
            self.rect.width = hit_w
            self.rect.height = hit_h
            self.rect.x = col * tile_size + (self.width - hit_w) // 2
            self.rect.y = (row * tile_size + tile_size) - hit_h
            self.spawn_x = self.rect.x
            self.spawn_y = self.rect.y
        else:
            # keep pixel spawn but recompute hitbox size and align bottom to nearest tile row
            px = getattr(self, "spawn_x", self.rect.x)
            py = getattr(self, "spawn_y", self.rect.y)
            self.rect.width = hit_w
            self.rect.height = hit_h
            self.rect.x = px
            self.rect.y = py

        # update draw offset and cache
        self.image_offset = ((self.width - hit_w) // 2, self.height - hit_h)
        self._scale_cache = {"scale": None, "image": None, "dead": None}

    # unified damage handler that teleports to spawn
    def take_damage(self, amount=1, teleport=True):
        now_ms = pygame.time.get_ticks()
        # ignore damage if still inside invincibility window
        if now_ms < getattr(self, "invincible_until_ms", 0):
            return

        # SAFE: play damage sound locally (no RARL import)
        Player._play_damage_sound_safe()

        self.health = max(0, self.health - amount)
        self.last_touch_damage_ms = now_ms
        self.invincible_until_ms = now_ms + getattr(self, "invincible_duration_ms", 2000)
        if teleport:
            self.rect.x, self.rect.y = self.spawn_x, self.spawn_y
            self.vel_y = 0
            self.in_air = False
            self.jumped = False


class Projectile(pygame.sprite.Sprite):
    SPEED = 12
    SIZE = (20, 10)
    fire_sound = None  # None = not loaded yet, False = permanently unavailable

    @classmethod
    def _ensure_sound(cls):
        # If we already know sound is unavailable, never try again
        if cls.fire_sound is False:
            return
        # If mixer isn't even initialized, don't attempt
        if not pygame.mixer.get_init():
            cls.fire_sound = False
            return

        # If something else already loaded the sound, stop here
        if isinstance(cls.fire_sound, pygame.mixer.Sound):
            return

        # try to load the sound once
        try:
            cls.fire_sound = pygame.mixer.Sound(r"Sounds/laser.mp3")
        except Exception:
            try:
                cls.fire_sound = pygame.mixer.Sound("Sounds/laser.mp3")
            except Exception:
                cls.fire_sound = False  # mark as unavailable

    @classmethod
    def apply_master_volume(cls, volume: float):
        """
        Called from RARL.apply_volume().
        Ensures fire_sound is loaded and sets its volume to the given master volume (0.0–1.0).
        """
        if not pygame.mixer.get_init():
            return
        # try to load if not yet loaded; this won't crash if file missing
        cls._ensure_sound()
        if isinstance(cls.fire_sound, pygame.mixer.Sound):
            try:
                cls.fire_sound.set_volume(max(0.0, min(1.0, float(volume))))
            except Exception:
                pass

    def __init__(self, x, y, direction, owner='player'):
        super().__init__()
        # Use laser texture for projectiles; enemy shots use a different texture if available
        img = None
        if owner == 'enemy':
            for p in (r"Textures\laser_enemies.png", r"Textures/laser_enemies.png"):
                try:
                    img = pygame.image.load(p).convert_alpha()
                    break
                except Exception:
                    continue
        else:
            for p in (r"Textures\laser.png", r"Textures/laser.png"):
                try:
                    img = pygame.image.load(p).convert_alpha()
                    break
                except Exception:
                    continue
        if img is None:
            # fallback to the generic block image if any laser texture is missing
            try:
                img = pygame.image.load(r"Textures/block.png").convert_alpha()
            except Exception:
                try:
                    img = pygame.image.load(r"Textures/block.png").convert_alpha()
                except Exception:
                    # FINAL FALLBACK: always create a visible colored rect
                    img = pygame.Surface((max(1, self.SIZE[0]), max(1, self.SIZE[1])), pygame.SRCALPHA)
                    img.fill((255, 0, 255))  # magenta so it's clearly visible

        # make sure we always end up with a non-zero-sized surface
        if img.get_width() <= 0 or img.get_height() <= 0:
            tmp = pygame.Surface((max(1, self.SIZE[0]), max(1, self.SIZE[1])), pygame.SRCALPHA)
            tmp.fill((255, 0, 255))
            img = tmp

        self.image = pygame.transform.scale(img, self.SIZE)
        if direction < 0:
            try:
                self.image = pygame.transform.flip(self.image, True, False)
            except Exception:
                pass

        # CHANGED: treat x,y as center for consistent movement/drawing
        self.rect = self.image.get_rect(center=(int(x), int(y)))
        self.x = float(self.rect.centerx)
        self.y = float(self.rect.centery)
        self.direction = 1 if direction >= 0 else -1
        # owner of the projectile: 'player' or 'enemy'
        self.owner = owner
        self.damage = 1

        Projectile._ensure_sound()
        if isinstance(Projectile.fire_sound, pygame.mixer.Sound) and pygame.mixer.get_init():
            try:
                Projectile.fire_sound.play()
            except Exception:
                pass

        # NEW: per-instance draw cache
        self._draw_cache = {"scale": None, "image": None}

    def update(self, world, screen_w=None, screen_h=None):
        # CHANGED: move using floats, then sync rect.center
        step_x = 1 if self.direction > 0 else -1
        for _ in range(abs(self.SPEED)):
            self.x += step_x
            self.rect.centerx = int(round(self.x))
            # collide with world tiles each pixel step
            collided = False
            tile_source = getattr(world, "collision_rects", None)
            if tile_source is None:
                tile_source = [r for _, r in getattr(world, "tile_list", [])]
            for tile_rect in tile_source:
                if self.rect.colliderect(tile_rect):
                    self.kill()
                    return
        # keep y center in sync (no vertical motion now)
        self.rect.centery = int(round(self.y))

        if (self.rect.right < 0 or self.rect.left > world.pixel_width or
            self.rect.bottom < 0 or self.rect.top > world.pixel_height):
            self.kill()

    def draw(self, screen, scale, off_x=0, off_y=0):
        if self._draw_cache["scale"] != scale:
            self._draw_cache["scale"] = scale
            sw = max(1, int(round(self.image.get_width() * scale)))
            sh = max(1, int(round(self.image.get_height() * scale)))
            self._draw_cache["image"] = pygame.transform.smoothscale(self.image, (sw, sh))
        img = self._draw_cache["image"]
        sx = int(round(self.rect.centerx * scale - img.get_width() / 2)) + off_x
        sy = int(round(self.rect.centery * scale - img.get_height() / 2)) + off_y
        screen.blit(img, (sx, sy))
