"""
Test Environments for AGI Evaluation

Various tasks for testing the recursive optimization system:
    - Classification tasks (Two Moons, Circles, Spiral)
    - Grid World navigation
    - Continuous control
"""

import numpy as np
from typing import Tuple, Optional, List
from dataclasses import dataclass
from sklearn.datasets import make_moons, make_circles


@dataclass
class TaskConfig:
    """Base configuration for tasks."""
    n_samples: int = 500
    noise: float = 0.1
    random_state: int = 42


class BaseTask:
    """Base class for all tasks."""
    
    def __init__(self):
        self.current_idx = 0
        self.episode_steps = 0
        self.max_steps = 200
        
    def reset(self) -> np.ndarray:
        raise NotImplementedError
        
    def step(self, action: int) -> Tuple[np.ndarray, float, bool]:
        raise NotImplementedError
        
    def get_goal(self) -> Optional[np.ndarray]:
        return None
    
    def get_observation(self) -> np.ndarray:
        raise NotImplementedError


class TwoMoonsTask(BaseTask):
    """
    Two interlocking moons classification.
    
    Non-linear decision boundary requiring hidden representations.
    """
    
    def __init__(self, config: Optional[TaskConfig] = None):
        super().__init__()
        self.config = config or TaskConfig()
        
        self.X, self.y = make_moons(
            n_samples=self.config.n_samples,
            noise=self.config.noise,
            random_state=self.config.random_state
        )
        self.y = self.y.reshape(-1, 1)
        
        self.current_obs = None
        self.current_target = None
        
    def reset(self) -> np.ndarray:
        self.current_idx = np.random.randint(len(self.X))
        self.current_obs = self.X[self.current_idx]
        self.current_target = self.y[self.current_idx]
        self.episode_steps = 0
        return self.current_obs
    
    def step(self, action: int) -> Tuple[np.ndarray, float, bool]:
        # For classification, action is the predicted class
        predicted = 1 if action > 0 else 0
        actual = self.current_target[0]
        
        reward = 1.0 if predicted == actual else -0.1
        
        self.episode_steps += 1
        done = self.episode_steps >= self.max_steps
        
        # Move to next sample
        self.current_idx = (self.current_idx + 1) % len(self.X)
        self.current_obs = self.X[self.current_idx]
        self.current_target = self.y[self.current_idx]
        
        return self.current_obs, reward, done
    
    def get_observation(self) -> np.ndarray:
        return self.current_obs if self.current_obs is not None else self.reset()
    
    def get_data(self) -> Tuple[np.ndarray, np.ndarray]:
        """Get full dataset for batch training."""
        return self.X, self.y


class ConcentricCirclesTask(BaseTask):
    """
    Concentric circles classification.
    
    Requires learning radial decision boundary.
    """
    
    def __init__(self, config: Optional[TaskConfig] = None):
        super().__init__()
        self.config = config or TaskConfig()
        
        self.X, self.y = make_circles(
            n_samples=self.config.n_samples,
            noise=self.config.noise,
            factor=0.5,
            random_state=self.config.random_state
        )
        self.y = self.y.reshape(-1, 1)
        
        self.current_obs = None
        self.current_target = None
        
    def reset(self) -> np.ndarray:
        self.current_idx = np.random.randint(len(self.X))
        self.current_obs = self.X[self.current_idx]
        self.current_target = self.y[self.current_idx]
        self.episode_steps = 0
        return self.current_obs
    
    def step(self, action: int) -> Tuple[np.ndarray, float, bool]:
        predicted = 1 if action > 0 else 0
        actual = self.current_target[0]
        
        reward = 1.0 if predicted == actual else -0.1
        
        self.episode_steps += 1
        done = self.episode_steps >= self.max_steps
        
        self.current_idx = (self.current_idx + 1) % len(self.X)
        self.current_obs = self.X[self.current_idx]
        self.current_target = self.y[self.current_idx]
        
        return self.current_obs, reward, done
    
    def get_observation(self) -> np.ndarray:
        return self.current_obs if self.current_obs is not None else self.reset()
    
    def get_data(self) -> Tuple[np.ndarray, np.ndarray]:
        return self.X, self.y


class SpiralTask(BaseTask):
    """
    Two interleaved spirals classification.
    
    Highly non-linear - requires complex decision boundary.
    """
    
    def __init__(self, config: Optional[TaskConfig] = None):
        super().__init__()
        self.config = config or TaskConfig()
        
        n = self.config.n_samples // 2
        
        # Generate spirals
        theta = np.sqrt(np.random.rand(n)) * 2 * np.pi
        r_a = 2 * theta + np.pi
        r_b = -2 * theta - np.pi
        
        data_a = np.column_stack([
            r_a * np.cos(theta) + np.random.randn(n) * self.config.noise,
            r_a * np.sin(theta) + np.random.randn(n) * self.config.noise
        ])
        data_b = np.column_stack([
            r_b * np.cos(theta) + np.random.randn(n) * self.config.noise,
            r_b * np.sin(theta) + np.random.randn(n) * self.config.noise
        ])
        
        self.X = np.vstack([data_a, data_b])
        self.y = np.array([0] * n + [1] * n).reshape(-1, 1)
        
        # Normalize
        self.X = (self.X - self.X.mean(axis=0)) / (self.X.std(axis=0) + 1e-8)
        
        self.current_obs = None
        self.current_target = None
        
    def reset(self) -> np.ndarray:
        self.current_idx = np.random.randint(len(self.X))
        self.current_obs = self.X[self.current_idx]
        self.current_target = self.y[self.current_idx]
        self.episode_steps = 0
        return self.current_obs
    
    def step(self, action: int) -> Tuple[np.ndarray, float, bool]:
        predicted = 1 if action > 0 else 0
        actual = self.current_target[0]
        
        reward = 1.0 if predicted == actual else -0.1
        
        self.episode_steps += 1
        done = self.episode_steps >= self.max_steps
        
        self.current_idx = (self.current_idx + 1) % len(self.X)
        self.current_obs = self.X[self.current_idx]
        self.current_target = self.y[self.current_idx]
        
        return self.current_obs, reward, done
    
    def get_observation(self) -> np.ndarray:
        return self.current_obs if self.current_obs is not None else self.reset()
    
    def get_data(self) -> Tuple[np.ndarray, np.ndarray]:
        return self.X, self.y


class GridWorldTask(BaseTask):
    """
    Simple grid world navigation.
    
    Agent must reach goal while avoiding obstacles.
    """
    
    def __init__(self, size: int = 10):
        super().__init__()
        self.size = size
        self.grid = np.zeros((size, size))
        
        # Actions: 0=up, 1=right, 2=down, 3=left
        self.action_map = {
            0: (-1, 0),
            1: (0, 1),
            2: (1, 0),
            3: (0, -1)
        }
        
        self.agent_pos = None
        self.goal_pos = None
        self.max_steps = size * 4
        
    def reset(self) -> np.ndarray:
        # Random start and goal
        self.agent_pos = np.array([0, 0])
        self.goal_pos = np.array([self.size - 1, self.size - 1])
        self.episode_steps = 0
        return self._get_obs()
    
    def _get_obs(self) -> np.ndarray:
        """Observation: normalized position + goal direction."""
        pos_norm = self.agent_pos / self.size
        goal_dir = (self.goal_pos - self.agent_pos) / self.size
        return np.concatenate([pos_norm, goal_dir])
    
    def step(self, action: int) -> Tuple[np.ndarray, float, bool]:
        action = action % 4  # Ensure valid action
        
        # Move agent
        delta = np.array(self.action_map[action])
        new_pos = self.agent_pos + delta
        
        # Clip to grid bounds
        new_pos = np.clip(new_pos, 0, self.size - 1)
        self.agent_pos = new_pos
        
        # Compute reward
        dist_to_goal = np.linalg.norm(self.agent_pos - self.goal_pos)
        reached_goal = dist_to_goal < 0.5
        
        reward = -0.01  # Step cost
        if reached_goal:
            reward = 1.0
        else:
            # Reward for getting closer
            reward += 0.1 * (1 - dist_to_goal / (self.size * np.sqrt(2)))
        
        self.episode_steps += 1
        done = reached_goal or self.episode_steps >= self.max_steps
        
        return self._get_obs(), reward, done
    
    def get_observation(self) -> np.ndarray:
        if self.agent_pos is None:
            return self.reset()
        return self._get_obs()
    
    def get_goal(self) -> np.ndarray:
        if self.goal_pos is None:
            self.reset()
        return np.concatenate([self.goal_pos / self.size, np.zeros(2)])


class ContinuousControlTask(BaseTask):
    """
    Continuous control task (simplified pendulum-like).
    
    Agent controls velocity to reach target.
    """
    
    def __init__(self):
        super().__init__()
        self.pos = None
        self.vel = None
        self.target = None
        self.dt = 0.1
        self.max_steps = 100
        
    def reset(self) -> np.ndarray:
        self.pos = np.random.randn(2) * 0.5
        self.vel = np.zeros(2)
        self.target = np.random.randn(2)
        self.episode_steps = 0
        return self._get_obs()
    
    def _get_obs(self) -> np.ndarray:
        return np.concatenate([self.pos, self.vel, self.target - self.pos])
    
    def step(self, action: int) -> Tuple[np.ndarray, float, bool]:
        # Map discrete action to acceleration
        acc_map = {
            0: np.array([0, 0.1]),   # up
            1: np.array([0.1, 0]),   # right
            2: np.array([0, -0.1]),  # down
            3: np.array([-0.1, 0])   # left
        }
        
        acc = acc_map.get(action % 4, np.zeros(2))
        
        # Physics update
        self.vel = self.vel * 0.95 + acc  # Damping
        self.pos = self.pos + self.vel * self.dt
        
        # Reward
        dist = np.linalg.norm(self.pos - self.target)
        reward = -dist * 0.1  # Distance penalty
        
        if dist < 0.2:
            reward = 1.0  # Bonus for reaching target
        
        self.episode_steps += 1
        done = dist < 0.1 or self.episode_steps >= self.max_steps
        
        return self._get_obs(), reward, done
    
    def get_observation(self) -> np.ndarray:
        if self.pos is None:
            return self.reset()
        return self._get_obs()
    
    def get_goal(self) -> np.ndarray:
        if self.target is None:
            self.reset()
        return np.concatenate([self.target, np.zeros(4)])
