import pygame
velkost_plochy = 50
BASE_TILE = 50

class JumpingEnemy(pygame.sprite.Sprite):
    def __init__(self, x, y, speed=2, patrol_blocks=5, image_path=None, tile_size=50):
        super().__init__()
        self.base_speed = speed
        base_img = self._load_image(image_path)
        base_img = pygame.transform.scale(base_img, (tile_size, tile_size))
        self.image_right = base_img
        self.image_left = pygame.transform.flip(base_img, True, False)
        self.image = self.image_right
        # hitbox smaller than image and bottom-aligned
        hit_w = max(8, int(round(tile_size * 0.6)))
        hit_h = max(8, int(round(tile_size * 0.9)))
        hit_x = x + (tile_size - hit_w) // 2
        hit_y = y + tile_size - hit_h
        self.rect = pygame.Rect(hit_x, hit_y, hit_w, hit_h)  # hitbox
        self.image_offset = ((self.image_right.get_width() - hit_w) // 2, self.image_right.get_height() - hit_h)
        self.start_x = self.rect.x
        self.patrol_blocks = patrol_blocks
        self.patrol_distance = patrol_blocks * tile_size
        self.direction = 1
        self.health = 2
        self.vel_y = 0
        self.on_ground = False
        self.jump_force_base = 15  # interpret relative to BASE_TILE
        self.jump_interval_ms = 900
        self.last_jump_ms = 0
        self._scale_cache = {"scale": None, "right": None, "left": None}

    def _load_image(self, image_path):
        # Try provided path, then common default, then fallback colored box
        candidates = [p for p in [image_path, r"Textures\Player.png"] if p]
        for p in candidates:
            if not p:
                continue
            try:
                return pygame.image.load(p).convert_alpha()
            except Exception:
                try:
                    return pygame.image.load(p.replace("\\", "/")).convert_alpha()
                except Exception:
                    continue
        surf = pygame.Surface((50, 50), pygame.SRCALPHA)
        surf.fill((220, 40, 40, 255))
        return surf

    def update(self, player=None, projectiles=None, world=None):
        # compute scaled horizontal speed
        tile_size = getattr(world, "tile_size", velkost_plochy)
        speed_px = max(1, int(round(self.base_speed * (float(tile_size) / float(BASE_TILE)))))

        # Horizontal patrol (operate on self.rect hitbox)
        intended_dx = int(speed_px * self.direction)
        if intended_dx != 0:
            step_x = 1 if intended_dx > 0 else -1
            for _ in range(abs(intended_dx)):
                self.rect.x += step_x
                hit_wall = False
                if world:
                    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
                            hit_wall = True
                            break
                if hit_wall:
                    self.direction *= -1
                    break
                if self.rect.x > self.start_x + self.patrol_distance:
                    self.rect.x = self.start_x + self.patrol_distance
                    self.direction = -1
                    break
                elif self.rect.x < self.start_x:
                    self.rect.x = self.start_x
                    self.direction = 1
                    break

        # NEW: trigger jump when on ground and cooldown elapsed
        now_ms = pygame.time.get_ticks()
        if self.on_ground and (now_ms - self.last_jump_ms) >= self.jump_interval_ms:
            # scale jump force
            tile_size = getattr(world, "tile_size", velkost_plochy)
            self.vel_y = - (self.jump_force_base * (float(tile_size) / float(BASE_TILE)))
            self.on_ground = False
            self.last_jump_ms = now_ms

        # Gravity + vertical collision with tiles (hitbox)
        gravity = 1.0 * (float(tile_size) / float(BASE_TILE))
        max_fall = 10.0 * (float(tile_size) / float(BASE_TILE))
        self.vel_y = min(self.vel_y + gravity, max_fall)
        dy = self.vel_y
        if dy != 0:
            step_y = 1 if dy > 0 else -1
            for _ in range(int(abs(round(dy)))):
                self.rect.y += step_y
                collided = False
                if world:
                    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.on_ground = True
                            collided = True
                            break
                if collided:
                    break
        if dy > 0 and self.vel_y != 0:
            self.on_ground = False

        # Take damage from projectiles (use projectile.damage if present)
        if projectiles and isinstance(projectiles, pygame.sprite.Group):
            total_damage = 0
            for p in list(projectiles):
                try:
                    owner = getattr(p, 'owner', None)
                    if owner not in (None, 'player'):
                        continue
                    if p.rect.colliderect(self.rect):
                        try:
                            p.kill()
                        except Exception:
                            pass
                        dmg = getattr(p, "damage", 1)
                        total_damage += max(0, int(dmg))
                except Exception:
                    continue

            if total_damage > 0:
                self.health -= total_damage
                if self.health <= 0:
                    self.kill()
                    return

    # NEW: scaled draw with cache
    def draw(self, screen, scale, off_x=0, off_y=0):
        if self._scale_cache["scale"] != scale:
            self._scale_cache["scale"] = scale
            sw = max(1, int(round(self.image_right.get_width() * scale)))
            sh = max(1, int(round(self.image_right.get_height() * scale)))
            self._scale_cache["right"] = pygame.transform.smoothscale(self.image_right, (sw, sh))
            self._scale_cache["left"] = pygame.transform.smoothscale(self.image_left, (sw, sh))
        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))
        img = self._scale_cache["right"] if self.direction == 1 else self._scale_cache["left"]
        screen.blit(img, (sx - ox, sy - oy))
