import pygame
import random
import math

from Bosses.Boss3 import BASE_TILE, BossProjectile  # import BossProjectile for shared loader


class Boss4Projectile(pygame.sprite.Sprite):
    SPEED = 3
    SIZE = (20, 10)

    def __init__(self, x, y, vx, vy, tile_size=None, group=None,
                 start_delay_ms=0, is_hazard=True, use_special_texture=False,
                 orientation=None):
        super().__init__()
        # texture: shared enemy laser loader + optional special texture
        if use_special_texture:
            img = None
            # Try project-root Textures, then ../Textures, both slash styles
            for p in (
                r"Textures/Special_laser.png",
                r"Textures\Special_laser.png",
                r"../Textures/Special_laser.png",
                r"..\Textures\Special_laser.png",
            ):
                try:
                    img = pygame.image.load(p).convert_alpha()
                    break
                except Exception:
                    continue
            if img is None:
                img = BossProjectile.load_enemy_laser_surface(self.SIZE)

            base = pygame.transform.scale(img, self.SIZE)
            if orientation == "vertical":
                try:
                    base = pygame.transform.rotate(base, 90)
                except Exception:
                    pass
            self.image = base
        else:
            img = BossProjectile.load_enemy_laser_surface(self.SIZE)
            base = pygame.transform.scale(img, self.SIZE)
            angle = math.degrees(math.atan2(-vy, vx))
            self.image = pygame.transform.rotate(base, angle)

        self.rect = self.image.get_rect(center=(int(x), int(y)))
        # ...existing movement / delay init...
        self.x = float(x)
        self.y = float(y)
        length = math.hypot(vx, vy)
        s = float(tile_size) / float(BASE_TILE) if tile_size else 1.0
        base_speed = float(self.SPEED) * s
        if use_special_texture:
            self.vx = 0.0
            self.vy = 0.0
        else:
            if length == 0:
                self.vx = self.vy = 0.0
            else:
                self.vx = vx / length * base_speed
                self.vy = vy / length * base_speed
        self.delay_ms = int(max(0, start_delay_ms))
        self.spawn_time_ms = pygame.time.get_ticks()
        self.is_hazard = bool(is_hazard)
        self._draw_cache = {"scale": None, "image": None}
        if group is not None:
            group.add(self)

    def update(self, world):
        # wait some time before starting to move & damage
        now = pygame.time.get_ticks()
        if now - self.spawn_time_ms < self.delay_ms:
            return
        # ...existing movement / collision...
        self.x += self.vx
        self.y += self.vy
        self.rect.center = (int(round(self.x)), int(round(self.y)))

        tiles = getattr(world, "collision_rects", None)
        if tiles is None:
            tiles = [r for _, r in getattr(world, "tile_list", [])]
        for r in tiles:
            if self.rect.colliderect(r):
                self.kill()
                return

        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):
        # ...existing draw...
        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)))
            try:
                self._draw_cache["image"] = pygame.transform.smoothscale(self.image, (sw, sh))
            except Exception:
                self._draw_cache["image"] = pygame.transform.scale(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 LaneSegment(pygame.sprite.Sprite):
    """
    Red warning/damage lane:
    - orientation: 'vertical' (top->bottom) or 'horizontal' (left->right)
    - thickness: 1 tile
    - grows over build_time_ms; only damages when fully built.
    - disappears automatically after lifetime_ms, with a fade-out at the end.
    """
    def __init__(self, x, y, orientation, world, build_time_ms=1500):
        super().__init__()
        self.orientation = orientation
        self.world = world
        self.build_time_ms = max(1, int(build_time_ms))
        self.spawn_ms = pygame.time.get_ticks()
        self.lifetime_ms = 3000
        self.fade_ms = 800
        self.fade_start_ms = self.spawn_ms + max(0, self.lifetime_ms - self.fade_ms)

        ts = getattr(world, "tile_size", BASE_TILE)

        # logical full size
        if orientation == "vertical":
            full_w = ts
            full_h = world.pixel_height
        else:  # horizontal
            full_w = world.pixel_width
            full_h = ts

        # base rect anchored from given (x,y) edge
        if orientation == "vertical":
            self.anchor = (x, 0)
            # NEW: remember logical lane index (column) for duplicate suppression
            self.lane_index = x // ts
        else:
            self.anchor = (0, y)
            # NEW: lane index is row for horizontal lanes
            self.lane_index = y // ts

        self.full_size = (full_w, full_h)
        # start with zero length
        self.image = pygame.Surface((full_w, full_h), pygame.SRCALPHA)
        self.image.fill((255, 0, 0))
        self.rect = pygame.Rect(self.anchor[0], self.anchor[1], full_w, full_h)
        # draw cache
        self._draw_cache = {"scale": None, "image": None}

        # --- use tiled Special_laser.png with cross‑platform paths ---
        tile_img = None
        for p in (
            r"Textures/Special_laser.png",
            r"Textures\Special_laser.png",
            r"../Textures/Special_laser.png",
            r"..\Textures\Special_laser.png",
        ):
            try:
                tile_img = pygame.image.load(p).convert_alpha()
                break
            except Exception:
                continue
        if tile_img is None:
            # fallback: simple red block (old behaviour)
            tile_img = pygame.Surface((ts, ts), pygame.SRCALPHA)
            tile_img.fill((255, 0, 0))

        # scale base tile to one logical tile
        try:
            tile_img = pygame.transform.smoothscale(tile_img, (ts, ts))
        except Exception:
            tile_img = pygame.transform.scale(tile_img, (ts, ts))

        # NEW: rotate vertical lanes 90 degrees, keep horizontal as-is
        if self.orientation == "vertical":
            try:
                tile_img = pygame.transform.rotate(tile_img, 90)
            except Exception:
                pass

        # create lane image and tile the segment texture across it
        self.image = pygame.Surface(self.full_size, pygame.SRCALPHA)
        full_w, full_h = self.full_size
        if self.orientation == "vertical":
            # tile along height (x stays 0, y increases)
            for yy in range(0, full_h, ts):
                self.image.blit(tile_img, (0, yy))
        else:
            # horizontal: tile along width
            for xx in range(0, full_w, ts):
                self.image.blit(tile_img, (xx, 0))

        self.rect = pygame.Rect(self.anchor[0], self.anchor[1], full_w, full_h)
        self._draw_cache = {"scale": None, "image": None}

    @property
    def progress(self):
        now = pygame.time.get_ticks()
        return max(0.0, min(1.0, (now - self.spawn_ms) / self.build_time_ms))

    @property
    def active(self):
        # only deal damage when fully built (during visible / fade phase)
        return self.progress >= 1.0

    def update(self, world):
        now = pygame.time.get_ticks()
        # NEW: auto-despawn after lifetime_ms
        if now - self.spawn_ms >= self.lifetime_ms:
            self.kill()
            return

        # keep rect within world; nothing else to move, just exists and grows visually
        ts = getattr(world, "tile_size", BASE_TILE)
        full_w, full_h = self.full_size
        p = self.progress

        if self.orientation == "vertical":
            cur_h = max(1, int(round(full_h * p)))
            self.rect.x = self.anchor[0]
            self.rect.y = 0
            self.rect.width = full_w
            self.rect.height = cur_h
        else:
            cur_w = max(1, int(round(full_w * p)))
            self.rect.x = 0
            self.rect.y = self.anchor[1]
            self.rect.width = cur_w
            self.rect.height = full_h

    def _current_alpha(self):
        """Compute current alpha (0–255) based on fade window."""
        now = pygame.time.get_ticks()
        if now < self.fade_start_ms:
            return 255
        if self.fade_ms <= 0:
            return 0
        t = now - self.fade_start_ms
        ratio = 1.0 - max(0.0, min(1.0, t / float(self.fade_ms)))
        return int(round(255 * ratio))

    def draw(self, screen, scale, off_x=0, off_y=0):
        # scale full lane texture once per scale
        if self._draw_cache["scale"] != scale:
            self._draw_cache["scale"] = scale
            full_w, full_h = self.full_size
            sw = max(1, int(round(full_w * scale)))
            sh = max(1, int(round(full_h * scale)))
            try:
                self._draw_cache["image"] = pygame.transform.smoothscale(self.image, (sw, sh))
            except Exception:
                self._draw_cache["image"] = pygame.transform.scale(self.image, (sw, sh))

        img = self._draw_cache["image"]
        if img is None:
            return

        # apply fading alpha
        alpha = self._current_alpha()
        if alpha <= 0:
            return
        img.set_alpha(alpha)

        full_w, full_h = self.full_size
        p = self.progress
        if self.orientation == "vertical":
            vis_h = max(1, int(round(full_h * p)))
            area = pygame.Rect(0, 0, img.get_width(), vis_h)
            sx = int(round(self.anchor[0] * scale)) + off_x
            sy = off_y
        else:
            vis_w = max(1, int(round(full_w * p)))
            area = pygame.Rect(0, 0, vis_w, img.get_height())
            sx = off_x
            sy = int(round(self.anchor[1] * scale)) + off_y

        screen.blit(img, (sx, sy), area)


class Boss4(pygame.sprite.Sprite):
    """
    Teleporting boss for level 20:
    - No walking, only teleports.
    - Fires projectile patterns over the whole screen.
    - Pattern phase: random teleports + shooting.
    - Stun phase: teleports to middle, 3s no shooting, stands still.
    """
    def __init__(self, x, y, hp=40, projectile_group=None):
        super().__init__()
        base_img = None
        # cross-platform + cwd-safe paths for Boss4.png
        for p in (
            r"Textures/Boss4.png",
            r"Textures\Boss4.png",
            r"../Textures/Boss4.png",
            r"..\Textures\Boss4.png",
        ):
            try:
                base_img = pygame.image.load(p).convert_alpha()
                break
            except Exception:
                continue
        if base_img is None:
            base_img = pygame.Surface((BASE_TILE, BASE_TILE), pygame.SRCALPHA)
            pygame.draw.circle(base_img, (160, 0, 180), (BASE_TILE // 2, BASE_TILE // 2), BASE_TILE // 2)
        self._base_img = base_img
        self.image = base_img.copy()
        self.rect = self.image.get_rect(center=(x, y))
        self.hit_rect = self.rect.copy()

        self.max_hp = 25
        self.hp = 25

        self.bullets = pygame.sprite.Group()
        self.projectile_group = projectile_group
        self.lanes = pygame.sprite.Group()
        # NEW: track which columns/rows already have an active lane, to avoid duplicates
        self.active_vertical_lane_cols = set()
        self.active_horizontal_lane_rows = set()

        self.state = "pattern"
        now = pygame.time.get_ticks()
        self.pattern_index = 0
        self.pattern_duration_ms = 3500
        self.stun_duration_ms = 3000
        self.state_end_ms = now + self.pattern_duration_ms
        self.last_shot_ms = 0
        self.pattern_shot_interval_ms = 500

        self._draw_cache = {"scale": None, "image": None}
        self._last_tile_size = None

        # NEW: fade state for teleporting
        self.fade_state = "idle"      # idle | fading_out | fading_in
        self.fade_start_ms = 0
        self.fade_duration_ms = 400   # ms for each fade phase
        self.current_alpha = 255
        self._pending_teleport = None  # ("random" or "center", target_args)

    def _ensure_scaled_size(self, world):
        """Scale boss image and derive smaller hitbox."""
        tile_size = getattr(world, "tile_size", BASE_TILE)
        if self._last_tile_size == tile_size:
            return
        self._last_tile_size = tile_size

        w = tile_size * 8
        h = tile_size * 5
        try:
            scaled = pygame.transform.smoothscale(self._base_img, (w, h))
        except Exception:
            scaled = pygame.transform.scale(self._base_img, (w, h))

        center = self.rect.center
        self.image = scaled
        self.rect = self.image.get_rect(center=center)

        shrink_w = tile_size * 4
        shrink_h = tile_size * 1
        hit_w = max(1, self.rect.width - shrink_w)
        hit_h = max(1, self.rect.height - shrink_h)
        hit_x = self.rect.centerx - hit_w // 2
        hit_y = self.rect.centery - hit_h // 2
        self.hit_rect = pygame.Rect(hit_x, hit_y, hit_w, hit_h)

        self._draw_cache["scale"] = None
        self._draw_cache["image"] = None

    # ---------------- utils / teleport ----------------

    def _safe_random_position(self, world, player, max_tries=10, min_dist_tiles=2):
        """Find a random position for teleporting that is not too close to the player."""
        ts = getattr(world, "tile_size", BASE_TILE)
        min_dist = ts * min_dist_tiles
        for _ in range(max_tries):
            x = random.randint(0, world.pixel_width)
            y = random.randint(0, world.pixel_height)
            if player is not None:
                dx = x - player.rect.centerx
                dy = y - player.rect.centery
                if math.hypot(dx, dy) < min_dist:
                    continue
            return x, y
        # fallback to center
        return world.pixel_width // 2, world.pixel_height // 2

    def _teleport_random(self, world, player=None):
        """Teleport instantly to new position (used after fade_out)."""
        if player is not None:
            x, y = self._safe_random_position(world, player)
        else:
            x = random.randint(0, world.pixel_width)
            y = random.randint(0, world.pixel_height)
        self.rect.center = (x, y)
        self.hit_rect.center = self.rect.center

    def _teleport_center(self, world, player=None):
        """Teleport instantly to center (used after fade_out)."""
        cx = world.pixel_width // 2
        cy = world.pixel_height // 2
        if player is not None:
            # if overlapping, nudge one tile up
            ts = getattr(world, "tile_size", BASE_TILE)
            tmp = self.hit_rect.copy()
            tmp.center = (cx, cy)
            if tmp.colliderect(player.rect):
                cy = max(0, cy - ts)
        self.rect.center = (cx, cy)
        self.hit_rect.center = self.rect.center

    # NEW: request a teleport with fade
    def _request_teleport(self, kind, world, player):
        """Schedule a teleport; actual move happens when fade-out completes."""
        if self.fade_state != "idle":
            return  # already in fade, ignore new request
        self.fade_state = "fading_out"
        self.fade_start_ms = pygame.time.get_ticks()
        # store what we want to do once invisible
        self._pending_teleport = (kind, world, player)

    def _update_fade(self):
        """Update alpha for fading in/out; return True if we should skip normal actions."""
        now = pygame.time.get_ticks()
        if self.fade_state == "idle":
            self.current_alpha = 255
            return False

        elapsed = now - self.fade_start_ms
        t = max(0.0, min(1.0, elapsed / float(self.fade_duration_ms)))

        if self.fade_state == "fading_out":
            self.current_alpha = int(round(255 * (1.0 - t)))
            if elapsed >= self.fade_duration_ms:
                # perform pending teleport when fully invisible
                if self._pending_teleport is not None:
                    kind, world, player = self._pending_teleport
                    if kind == "random":
                        self._teleport_random(world, player)
                    elif kind == "center":
                        self._teleport_center(world, player)
                # start fade in
                self.fade_state = "fading_in"
                self.fade_start_ms = now
                self.current_alpha = 0
        elif self.fade_state == "fading_in":
            self.current_alpha = int(round(255 * t))
            if elapsed >= self.fade_duration_ms:
                self.fade_state = "idle"
                self.current_alpha = 255
                self._pending_teleport = None

        return self.fade_state in ("fading_out", "fading_in")

    # ------------- damage / hits / patterns -------------
    def _handle_player_hits(self, player_projectiles):
        """Use smaller hit_rect for taking projectile damage."""
        for proj in list(player_projectiles):
            if self.hit_rect.colliderect(proj.rect):
                dmg = getattr(proj, "damage", 1)
                self.hp = max(0, self.hp - max(0, int(dmg)))
                proj.kill()
                if self.hp <= 0:
                    self.kill()
                    for b in list(self.bullets):
                        b.kill()
                    for ln in list(self.lanes):
                        ln.kill()
                    return

    def _handle_bullet_hits_player(self, player):
        """Use hit_rect for boss bullets to hit the player (unchanged)."""
        for b in list(self.bullets):
            if b.rect.colliderect(player.rect):
                try:
                    player.take_damage(1)
                except Exception:
                    player.health = max(0, getattr(player, "health", 0) - 1)
                b.kill()

    def _handle_lane_hits_player(self, player):
        """Damage player when fully-built lanes touch the player."""
        for ln in list(self.lanes):
            if not getattr(ln, "active", False):
                continue
            if ln.rect.colliderect(player.rect):
                try:
                    player.take_damage(1)
                except Exception:
                    player.health = max(0, getattr(player, "health", 0) - 1)

    def _shoot_circle(self, world, count=24):
        # LESS PROJECTILES: reduce count from 24/28 to 16
        cx, cy = self.rect.center
        ts = getattr(world, "tile_size", BASE_TILE)
        for i in range(16):
            ang = (2 * math.pi * i) / 16
            vx = math.cos(ang)
            vy = math.sin(ang)
            p = Boss4Projectile(cx, cy, vx, vy, tile_size=ts, group=self.bullets)
            if self.projectile_group is not None:
                self.projectile_group.add(p)

    def _shoot_cross_grid(self, world, count=10):
        # LESS PROJECTILES: shrink +/- range so fewer rays
        cx, cy = self.rect.center
        ts = getattr(world, "tile_size", BASE_TILE)
        for i in range(-6, 7):  # was -count..count with count=10/12
            if i == 0:
                continue
            vx = i * 0.2
            vy = 1 if i % 2 == 0 else -1
            p = Boss4Projectile(cx, cy, vx, vy, tile_size=ts, group=self.bullets)
            if self.projectile_group is not None:
                self.projectile_group.add(p)
        for i in range(-6, 7):
            if i == 0:
                continue
            vy = i * 0.2
            vx = 1 if i % 2 == 0 else -1
            p = Boss4Projectile(cx, cy, vx, vy, tile_size=ts, group=self.bullets)
            if self.projectile_group is not None:
                self.projectile_group.add(p)

    def _shoot_rain(self, world, rays=18):
        # LESS PROJECTILES: reduce number of falling bullets
        ts = getattr(world, "tile_size", BASE_TILE)
        y_top = -ts
        for _ in range(12):  # was rays=18/22
            x = random.randint(0, world.pixel_width)
            p = Boss4Projectile(x, y_top, 0, 1, tile_size=ts, group=self.bullets)
            if self.projectile_group is not None:
                self.projectile_group.add(p)

    def _shoot_net(self, world, step_tiles=3, delay_ms=700):
        """
        Special net attack:
        - Uses LaneSegment lines built from Special_laser.png tiles.
        - Behaviour (growth, fade, damage) matches current red lanes.
        - Avoids spawning duplicate lanes on the same column/row.
        """
        ts = getattr(world, "tile_size", BASE_TILE)

        # vertical lanes (like current lanes but spaced for net)
        for col_px in range(0, world.pixel_width, step_tiles * ts * 2):
            lane_col = col_px // ts
            # skip if this column already has an active/net lane
            if lane_col in self.active_vertical_lane_cols:
                continue
            lane = LaneSegment(col_px, 0, "vertical", world, build_time_ms=delay_ms)
            self.lanes.add(lane)
            self.active_vertical_lane_cols.add(lane.lane_index)

        # horizontal lanes
        for row_px in range(0, world.pixel_height, step_tiles * ts * 2):
            lane_row = row_px // ts
            if lane_row in self.active_horizontal_lane_rows:
                continue
            lane = LaneSegment(0, row_px, "horizontal", world, build_time_ms=delay_ms)
            self.lanes.add(lane)
            self.active_horizontal_lane_rows.add(lane.lane_index)

    def _shoot_lanes(self, world, lane_count=3):
        """
        red lane pattern:
        - spawn several 1-tile-thick vertical or horizontal lanes.
        - lanes slowly appear over 1.5s and only hurt when fully built.
        - avoid creating duplicates on the same column/row.
        """
        ts = getattr(world, "tile_size", BASE_TILE)
        lane_count = min(lane_count, 2)
        for _ in range(lane_count):
            orient = random.choice(["vertical", "horizontal"])
            if orient == "vertical":
                col = random.randint(0, max(0, world.cols - 1))
                if col in self.active_vertical_lane_cols:
                    continue
                x = col * ts
                y = 0
                lane = LaneSegment(x, y, orient, world, build_time_ms=1500)
                self.lanes.add(lane)
                self.active_vertical_lane_cols.add(lane.lane_index)
            else:
                row = random.randint(0, max(0, world.rows - 1))
                if row in self.active_horizontal_lane_rows:
                    continue
                x = 0
                y = row * ts
                lane = LaneSegment(x, y, orient, world, build_time_ms=1500)
                self.lanes.add(lane)
                self.active_horizontal_lane_rows.add(lane.lane_index)

    def _do_pattern_shoot(self, world):
        now = pygame.time.get_ticks()
        if now - self.last_shot_ms < self.pattern_shot_interval_ms:
            return
        self.last_shot_ms = now

        # CHANGED: 5 patterns (circle, cross, rain, net, lanes)
        idx = self.pattern_index % 5
        if idx == 0:
            self._shoot_circle(world)
        elif idx == 1:
            self._shoot_cross_grid(world)
        elif idx == 2:
            self._shoot_rain(world)
        elif idx == 3:
            self._shoot_net(world, step_tiles=3, delay_ms=800)
        else:
            self._shoot_lanes(world, lane_count=3)

    def update(self, player, player_projectiles, world):
        self._ensure_scaled_size(world)
        now = pygame.time.get_ticks()

        self._handle_player_hits(player_projectiles)
        if not self.alive():
            return

        # update fading (may also perform teleport)
        fading = self._update_fade()

        # Only do movement/teleport decisions when not mid-fade
        if not fading:
            if self.state == "pattern":
                self._do_pattern_shoot(world)
                # request teleport instead of instant teleport
                if random.random() < 0.012:
                    self.state = "stun"
                    self.state_end_ms = now + self.stun_duration_ms
                    self._request_teleport("center", world, player)
            elif self.state == "stun":
                if now >= self.state_end_ms:
                    self.state = "pattern"
                    self.pattern_index += 1
                    self.state_end_ms = now + self.pattern_duration_ms
                    self.last_shot_ms = 0
                    self._request_teleport("random", world, player)

        # keep inside world using the hitbox, then sync image
        if self.hit_rect.left < 0:
            self.hit_rect.left = 0
        if self.hit_rect.top < 0:
            self.hit_rect.top = 0
        if self.hit_rect.right > world.pixel_width:
            self.hit_rect.right = world.pixel_width
        if self.hit_rect.bottom > world.pixel_height:
            self.hit_rect.bottom = world.pixel_height
        self.rect.center = self.hit_rect.center

        # update projectiles and lanes + damage player
        self.bullets.update(world)
        self.lanes.update(world)

        # NEW: clean up lane index sets so finished lanes free their column/row
        alive_vertical = set()
        alive_horizontal = set()
        for ln in self.lanes:
            if ln.orientation == "vertical":
                alive_vertical.add(ln.lane_index)
            else:
                alive_horizontal.add(ln.lane_index)
        self.active_vertical_lane_cols = alive_vertical
        self.active_horizontal_lane_rows = alive_horizontal

        self._handle_bullet_hits_player(player)
        self._handle_lane_hits_player(player)

    def draw(self, screen, scale, off_x=0, off_y=0):
        # boss sprite
        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)))
            try:
                self._draw_cache["image"] = pygame.transform.smoothscale(self.image, (sw, sh))
            except Exception:
                self._draw_cache["image"] = pygame.transform.scale(self.image, (sw, sh))
        img = self._draw_cache["image"].copy()

        # apply current alpha from fade state
        try:
            img.set_alpha(self.current_alpha)
        except Exception:
            pass

        sx = int(round(self.rect.x * scale)) + off_x
        sy = int(round(self.rect.y * scale)) + off_y
        screen.blit(img, (sx, sy))

        # HP bar
        bw = max(50, img.get_width())
        bh = max(4, int(round(6 * scale)))
        bx = sx + img.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))))

        # REMOVED: green hitbox outline (hx, hy, hw, hh)

        # bullets
        for b in self.bullets:
            b.draw(screen, scale, off_x, off_y)

        # lanes
        for ln in self.lanes:
            ln.draw(screen, scale, off_x, off_y)
