from typing import Literal, Optional, Union, Any, TextIO


class Colors:
    """ANSI color codes for terminal output"""

    HEADER = "\033[95m"
    OKBLUE = "\033[94m"
    OKCYAN = "\033[96m"
    OKGREEN = "\033[92m"
    WARNING = "\033[93m"
    FAIL = "\033[91m"
    ENDC = "\033[0m"
    BOLD = "\033[1m"
    UNDERLINE = "\033[4m"

    # Additional colors
    BLACK = "\033[30m"
    RED = "\033[31m"
    GREEN = "\033[32m"
    YELLOW = "\033[33m"
    BLUE = "\033[34m"
    MAGENTA = "\033[35m"
    CYAN = "\033[36m"
    WHITE = "\033[37m"

    # Bright colors
    BRIGHT_BLACK = "\033[90m"
    BRIGHT_RED = "\033[91m"
    BRIGHT_GREEN = "\033[92m"
    BRIGHT_YELLOW = "\033[93m"
    BRIGHT_BLUE = "\033[94m"
    BRIGHT_MAGENTA = "\033[95m"
    BRIGHT_CYAN = "\033[96m"
    BRIGHT_WHITE = "\033[97m"


# Rainbow colors in order
RAINBOW_COLORS = [
    Colors.RED,
    Colors.BRIGHT_RED,  # Orange-ish
    Colors.YELLOW,
    Colors.GREEN,
    Colors.CYAN,
    Colors.BLUE,
    Colors.MAGENTA,
]


ColorType = Union[
    Literal[
        "black",
        "red",
        "green",
        "yellow",
        "blue",
        "magenta",
        "cyan",
        "white",
        "bright_black",
        "bright_red",
        "bright_green",
        "bright_yellow",
        "bright_blue",
        "bright_magenta",
        "bright_cyan",
        "bright_white",
        "header",
        "okblue",
        "okcyan",
        "okgreen",
        "warning",
        "fail",
        "bold",
        "underline",
    ],
    str,  # Allow any string for direct ANSI codes
]
def cprint(
    *args: Any,
    color: Optional[ColorType] = None,
    rainbow: bool = False,
    sep: str = " ",
    end: str = "\n",
    file: Optional[TextIO] = None,
    flush: bool = False
) -> None:
    """
    Print text in color with optional rainbow effect.

    Args:
        text: The text to print
        color: Color name (str) or ANSI color code. Valid names:
               'red', 'green', 'blue', 'yellow', 'cyan', 'magenta', 'white', 'black',
               'header', 'okblue', 'okcyan', 'okgreen', 'warning', 'fail',
               'bold', 'underline', and bright variants (e.g., 'bright_red')
        end: String appended after the text (default: '\n')
        rainbow: If True, print each character in a different rainbow color
    """
    print_kwargs = {
        "end": end,
        "file": file,
        "flush": flush,
    }
    text = sep.join(str(arg) for arg in args)

    if rainbow:
        rainbowed = []
        for i, char in enumerate(text):
            color_code = RAINBOW_COLORS[i % len(RAINBOW_COLORS)]
            rainbowed.append(f"{color_code}{char}")
        rainbowed.append(Colors.ENDC)
        print("".join(rainbowed), **print_kwargs)
    else:
        if color is None:
            print(text, **print_kwargs)
        else:
            # Map color names to color codes
            color_map = {
                "black": Colors.BLACK,
                "red": Colors.RED,
                "green": Colors.GREEN,
                "yellow": Colors.YELLOW,
                "blue": Colors.BLUE,
                "magenta": Colors.MAGENTA,
                "cyan": Colors.CYAN,
                "white": Colors.WHITE,
                "bright_black": Colors.BRIGHT_BLACK,
                "bright_red": Colors.BRIGHT_RED,
                "bright_green": Colors.BRIGHT_GREEN,
                "bright_yellow": Colors.BRIGHT_YELLOW,
                "bright_blue": Colors.BRIGHT_BLUE,
                "bright_magenta": Colors.BRIGHT_MAGENTA,
                "bright_cyan": Colors.BRIGHT_CYAN,
                "bright_white": Colors.BRIGHT_WHITE,
                "header": Colors.HEADER,
                "okblue": Colors.OKBLUE,
                "okcyan": Colors.OKCYAN,
                "okgreen": Colors.OKGREEN,
                "warning": Colors.WARNING,
                "fail": Colors.FAIL,
                "bold": Colors.BOLD,
                "underline": Colors.UNDERLINE,
            }

            # Get color code from map or use directly if it's already a code
            if isinstance(color, str) and color.lower() in color_map:
                color_code = color_map[color.lower()]
            elif isinstance(color, str) and color in [getattr(Colors, attr) for attr in dir(Colors) if not attr.startswith('_')]:
                # Check if it's a valid color code from Colors class
                color_code = color
            else:
                raise ValueError(f"Invalid color: '{color}'. Must be a valid color name or color code from Colors class.")

            print(f"{color_code}{text}{Colors.ENDC}", **print_kwargs)
