import pygame

velkost_plochy = 50
BASE_TILE = 50

class Enemy(pygame.sprite.Sprite):
    def __init__(self, x, y, speed=2, patrol_blocks=2, image_path=None, tile_size=50):
        super().__init__()
        # store base speed (interpreted relative to BASE_TILE)
        self.base_speed = speed
        # load and scale image to tile_size
        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
        # compute 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
        # use hitbox as self.rect (so spritecollide and other code continue to work)
        self.rect = pygame.Rect(hit_x, hit_y, hit_w, hit_h)
        # store drawing offset (image sits above hitbox)
        self.image_offset = ((self.image_right.get_width() - hit_w) // 2, self.image_right.get_height() - hit_h)
        # patrol
        self.start_x = self.rect.x
        self.patrol_blocks = patrol_blocks
        self.patrol_distance = patrol_blocks * tile_size
        self.direction = 1  # 1 -> right, -1 -> left
        self.health = 5
        self.vel_y = 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 speed in pixels using world's tile_size (fallback to constructor size)
        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 with pixel stepping on hitbox (self.rect)
        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

        # Gravity + vertical collision with tiles (pixel stepping) on hitbox
        # scale gravity too so enemies fall in tile-consistent way
        tile_size = getattr(world, "tile_size", velkost_plochy)
        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
                            collided = True
                            break
                if collided:
                    break

        # Take damage from projectiles (use projectile.damage if present)
        if projectiles and isinstance(projectiles, pygame.sprite.Group):
            hits_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)
                        hits_damage += max(0, int(dmg))
                except Exception:
                    continue

            if hits_damage > 0:
                self.health -= hits_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))
