import pygame
import random
import math

BASE_TILE = 50

class BossProjectile(pygame.sprite.Sprite):
    SPEED = 4
    SIZE = (20, 10)
    fire_sound = None  # None = not loaded, False = permanent fail

    @classmethod
    def _ensure_sound(cls):
        if cls.fire_sound is False:
            return
        if not pygame.mixer.get_init():
            cls.fire_sound = False
            return

        if isinstance(cls.fire_sound, pygame.mixer.Sound):
            return
        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

    @classmethod
    def _play_fire_sound_safe(cls):
        """Play boss projectile sound safely, without importing other modules."""
        if not pygame.mixer.get_init():
            return
        cls._ensure_sound()
        if not isinstance(cls.fire_sound, pygame.mixer.Sound):
            return
        try:
            cls.fire_sound.play()
        except Exception:
            pass

    @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
        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

    @staticmethod
    def load_enemy_laser_surface(size):
        """Shared loader for laser_enemies.png with robust fallbacks (Windows/macOS)."""
        paths = [
            # project root Textures
            r"Textures/laser_enemies.png",
            r"Textures\laser_enemies.png",
            # from inside Bosses/ (run from project root or IDE)
            r"../Textures/laser_enemies.png",
            r"..\Textures\laser_enemies.png",
        ]
        img = None
        for p in paths:
            try:
                img = pygame.image.load(p).convert_alpha()
                break
            except Exception:
                continue
        if img is None:
            img = pygame.Surface(size, pygame.SRCALPHA)
            img.fill((255, 80, 20))
        return img

    def __init__(self, x, y, vx, vy, tile_size=None, group=None):
        super().__init__()
        img = self.load_enemy_laser_surface(self.SIZE)
        self.base_image = pygame.transform.scale(img, self.SIZE)
        # compute rotation visually but keep base image for scaling cache
        angle = math.degrees(math.atan2(-vy, vx))
        self.image = pygame.transform.rotate(self.base_image, angle)
        self.rect = self.image.get_rect(center=(int(x), int(y)))
        # use float positions to avoid rounding drift when moving/scaling
        self.x = float(x)
        self.y = float(y)
        # normalize direction then apply pixel speed scaled by tile_size
        length = math.hypot(vx, vy)
        speed_scale = 1.0
        if tile_size:
            speed_scale = float(tile_size) / float(BASE_TILE)
        base_speed = float(self.SPEED) * speed_scale
        if length == 0:
            self.vx, self.vy = 0.0, 0.0
        else:
            self.vx = (vx / length) * base_speed
            self.vy = (vy / length) * base_speed
        # per-instance draw cache
        self._draw_cache = {"scale": None, "image": None}
        # auto-add to group if provided
        if group is not None:
            group.add(self)

        BossProjectile._play_fire_sound_safe()

    def update(self, world):
        # Move in logic space using floats then sync rect
        self.x += self.vx
        self.y += self.vy
        self.rect.center = (int(round(self.x)), int(round(self.y)))

        # Collide with tiles -> destroy (use collision_rects if present)
        tile_source = getattr(world, "collision_rects", None)
        if tile_source is None:
            # fallback to visual tile_list rects
            for _, tile_rect in world.tile_list:
                if tile_rect.colliderect(self.rect):
                    self.kill()
                    return
        else:
            for tile_rect in tile_source:
                if tile_rect.colliderect(self.rect):
                    self.kill()
                    return

        # Remove if outside world
        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))


class Boss3(pygame.sprite.Sprite):
    def __init__(self, x, y, hp=20, projectile_group=None):
        super().__init__()
        # Load boss image or fallback (Windows/macOS friendly)
        img = None
        for p in (
            r"Textures/Boss3.png",
            r"Textures\Boss3.png",
            r"../Textures/Boss3.png",
            r"..\Textures\Boss3.png",
        ):
            try:
                tmp = pygame.image.load(p).convert_alpha()
                img = pygame.transform.scale(tmp, (70, 70))
                break
            except Exception:
                continue
        if img is None:
            img = pygame.Surface((70, 70), pygame.SRCALPHA)
            pygame.draw.circle(img, (160, 0, 180), (35, 35), 35)
            pygame.draw.circle(img, (255, 255, 255), (35, 35), 20, 3)
        self.image = img
        self.rect = self.image.get_rect()
        self.rect.center = (x, y)

        # Stats
        self.max_hp = hp
        self.hp = hp

        # Movement (store base speeds; actual pixel speeds are computed using world.tile_size/BASE_TILE)
        self.roam_speed = 0.8     # was 1.2 -> boss roams slower
        self.dash_speed = 2.4     # was 3.8 -> slower dash
        self.target = None

        # Shooting
        self.bullets = pygame.sprite.Group()
        # global projectiles group from RARL.py (used only for drawing)
        self.projectile_group = projectile_group
        self.last_shot_ms = 0
        self.roam_shot_interval_ms = 1100
        self.special_shot_interval_ms = 130

        # State
        self.state = "roam"  # roam | dash_to_center | special
        now = pygame.time.get_ticks()
        self.next_special_ms = now + random.randint(6000, 10000)
        self.special_end_ms = 0

        # Drawing cache
        self._draw_cache = {"scale": None, "image": None}

    def _pick_random_target(self, world):
        pad = 20
        tx = random.randint(pad, max(pad, world.pixel_width - pad))
        ty = random.randint(pad, max(pad, world.pixel_height - pad))
        self.target = (tx, ty)

    def _move_towards(self, tx, ty, speed_px):
        dx = tx - self.rect.centerx
        dy = ty - self.rect.centery
        dist = math.hypot(dx, dy)
        if dist <= speed_px or dist == 0:
            self.rect.center = (tx, ty)
            return True
        self.rect.centerx += dx / dist * speed_px
        self.rect.centery += dy / dist * speed_px
        return False

    def _shoot_8_directions(self, world):
        cx, cy = self.rect.center
        dirs = [
            (1, 0), (-1, 0), (0, 1), (0, -1),
            (1, 1), (1, -1), (-1, 1), (-1, -1),
        ]
        for vx, vy in dirs:
            # Always add to self.bullets for damage handling
            p = BossProjectile(
                cx,
                cy,
                vx,
                vy,
                tile_size=getattr(world, "tile_size", BASE_TILE),
                group=self.bullets,
            )
            # Additionally mirror into global drawing group if available
            if self.projectile_group is not None:
                self.projectile_group.add(p)

    def _handle_player_hits(self, player_projectiles):
        # Take damage from player's projectiles
        for proj in list(player_projectiles):
            if self.rect.colliderect(proj.rect):
                self.hp = max(0, self.hp - 1)
                proj.kill()
                if self.hp <= 0:
                    self.kill()
                    # Clear remaining boss bullets on death
                    for b in list(self.bullets):
                        b.kill()
                    return

    def _handle_bullet_hits_player(self, player):
        # Damage player when boss bullets collide
        for b in list(self.bullets):
            if b.rect.colliderect(player.rect):
                try:
                    player.take_damage(1)
                except Exception:
                    # simple fallback if player has no take_damage
                    player.health = max(0, getattr(player, "health", 0) - 1)
                b.kill()

    def update(self, player, player_projectiles, world):
        now = pygame.time.get_ticks()

        # Process hits
        self._handle_player_hits(player_projectiles)

        # Compute scaled speeds in pixels
        tile_size = getattr(world, "tile_size", BASE_TILE)
        s = float(tile_size) / float(BASE_TILE)
        roam_px = max(0.1, float(self.roam_speed) * s)
        dash_px = max(0.1, float(self.dash_speed) * s)

        # State transitions & movement
        if self.state == "roam":
            if self.target is None:
                self._pick_random_target(world)
            reached = self._move_towards(self.target[0], self.target[1], roam_px)
            if reached:
                self._pick_random_target(world)
            # periodic shooting
            if now - self.last_shot_ms >= self.roam_shot_interval_ms:
                self._shoot_8_directions(world)
                self.last_shot_ms = now
            # randomly trigger special
            if now >= self.next_special_ms:
                self.state = "dash_to_center"
                self.target = (world.pixel_width // 2, world.pixel_height // 2)

        elif self.state == "dash_to_center":
            reached = self._move_towards(self.target[0], self.target[1], dash_px)
            if reached:
                self.state = "special"
                self.special_end_ms = now + random.randint(2000, 3500)
                self.last_shot_ms = 0  # fire immediately
        elif self.state == "special":
            # rapid fire 8-dir
            if now - self.last_shot_ms >= self.special_shot_interval_ms:
                self._shoot_8_directions(world)
                self.last_shot_ms = now
            if now >= self.special_end_ms:
                self.state = "roam"
                self.target = None
                self.next_special_ms = now + random.randint(6000, 10000)

        # Clamp to world bounds
        self.rect.left = max(0, self.rect.left)
        self.rect.top = max(0, self.rect.top)
        self.rect.right = min(world.pixel_width, self.rect.right)
        self.rect.bottom = min(world.pixel_height, self.rect.bottom)

        # Update own bullets
        self.bullets.update(world)

        # Damage player on bullet hit
        self._handle_bullet_hits_player(player)

    def draw(self, screen, scale, off_x=0, off_y=0):
        # Draw boss
        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))
        sx = int(round(self.rect.x * scale)) + off_x
        sy = int(round(self.rect.y * scale)) + off_y
        screen.blit(self._draw_cache["image"], (sx, sy))

        # Optional: small HP bar above boss
        bw = max(40, self._draw_cache["image"].get_width())
        bh = max(3, int(round(6 * scale)))
        bx = sx + self._draw_cache["image"].get_width() // 2 - bw // 2
        by = sy - bh - max(2, int(round(6 * scale)))
        pygame.draw.rect(screen, (40, 40, 40), (bx, by, bw, bh))
        ratio = 0 if self.max_hp == 0 else self.hp / self.max_hp
        pygame.draw.rect(screen, (220, 40, 40), (bx, by, int(round(bw * ratio)), bh))
        pygame.draw.rect(screen, (0, 0, 0), (bx, by, bw, bh), max(1, int(round(2 * scale))))

        # Draw bullets
        for b in self.bullets:
            b.draw(screen, scale, off_x, off_y)
