from __future__ import annotations

import asyncio
import random
from abc import abstractmethod
import json
from inspect import signature, Parameter
from typing import Optional, _GenericAlias, AsyncIterator
from pathlib import Path
from aiohttp import ClientSession

try:
    from types import NoneType
except ImportError:
    NoneType = type(None)

from ..typing import AsyncResult, Messages
from .types import BaseProvider
from .asyncio import to_sync_generator, to_async_iterator
from .response import BaseConversation, AuthResult
from ..cookies import get_cookies_dir
from ..requests import raise_for_status
from ..errors import (
    ResponseError,
    MissingAuthError,
    NoValidHarFileError,
    PaymentRequiredError,
    CloudflareError,
    BadRequestError,
    RateLimitError,
)
from .. import debug

SAFE_PARAMETERS = [
    "model",
    "messages",
    "stream",
    "timeout",
    "media",
    "response_format",
    "prompt",
    "negative_prompt",
    "tools",
    "conversation",
    "history_disabled",
    "temperature",
    "top_k",
    "top_p",
    "frequency_penalty",
    "presence_penalty",
    "max_tokens",
    "stop",
    "api_key",
    "seed",
    "width",
    "height",
    "max_retries",
    "web_search",
    "cache",
    "guidance_scale",
    "num_inference_steps",
    "randomize_seed",
    "safe",
    "enhance",
    "private",
    "aspect_ratio",
    "n",
    "transparent",
    "reasoning_effort",
]

BASIC_PARAMETERS = {
    "provider": None,
    "model": "",
    "messages": [],
    "stream": False,
    "timeout": 0,
    "response_format": None,
    "max_tokens": 4096,
    "stop": ["stop1", "stop2"],
}

PARAMETER_EXAMPLES = {
    "proxy": "http://user:password@127.0.0.1:3128",
    "temperature": 1,
    "top_k": 1,
    "top_p": 1,
    "frequency_penalty": 1,
    "presence_penalty": 1,
    "messages": [{"role": "system", "content": ""}, {"role": "user", "content": ""}],
    "media": [["data:image/jpeg;base64,...", "filename.jpg"]],
    "response_format": {"type": "json_object"},
    "conversation": {
        "conversation_id": "550e8400-e29b-11d4-a716-...",
        "message_id": "550e8400-e29b-11d4-a716-...",
    },
    "seed": 42,
    "tools": [],
    "width": 1024,
    "height": 1024,
    "reasoning_effort": "medium",
    "aspect_ratio": "1:1",
}


async def wait_for(response: AsyncIterator, timeout: int = None) -> AsyncIterator:
    if timeout is not None:
        while True:
            try:
                async def wait_for_next():
                    try:
                        return await response.__anext__()
                    except TimeoutError as e:
                        raise TimeoutError(str(e) or "The operation timed out") from e
                yield await asyncio.wait_for(wait_for_next(), timeout=timeout)
            except TimeoutError as e:
                if str(e):
                    raise TimeoutError(str(e)) from e
                raise TimeoutError(
                    "The operation timed out after {} seconds".format(timeout)
                ) from e
            except StopAsyncIteration:
                break
    else:
        async for chunk in response:
            yield chunk


def get_async_provider_method(provider: type) -> Optional[callable]:
    if hasattr(provider, "create_async_generator"):
        return provider.create_async_generator
    if hasattr(provider, "create_async"):

        async def wrapper(*args, **kwargs):
            yield await provider.create_async(*args, **kwargs)

        return wrapper
    if hasattr(provider, "create_completion"):

        async def wrapper(*args, **kwargs):
            for chunk in provider.create_completion(*args, **kwargs):
                yield chunk

        return wrapper
    raise NotImplementedError(f"{provider.__name__} does not implement an async method")


def get_provider_method(provider: type) -> Optional[callable]:
    if hasattr(provider, "create_completion"):
        return provider.create_completion
    if hasattr(provider, "create_async_generator"):

        def wrapper(*args, **kwargs):
            return to_sync_generator(
                provider.create_async_generator(*args, **kwargs),
                stream=provider.supports_stream,
            )

        return wrapper
    if hasattr(provider, "create_async"):

        def wrapper(*args, **kwargs):
            yield asyncio.run(provider.create_async(*args, **kwargs))

        return wrapper
    raise NotImplementedError(f"{provider.__name__} does not implement a create method")


class AbstractProvider(BaseProvider):
    @classmethod
    def get_parameters(cls, as_json: bool = False) -> dict[str, Parameter]:
        params = {
            name: parameter
            for name, parameter in signature(
                cls.create_async_generator
                if issubclass(cls, AsyncGeneratorProvider)
                else cls.create_async
                if issubclass(cls, AsyncProvider)
                else cls.create_completion
            ).parameters.items()
            if name in SAFE_PARAMETERS and (name != "stream" or cls.supports_stream)
        }
        if as_json:

            def get_type_as_var(annotation: type, key: str, default):
                if key in PARAMETER_EXAMPLES:
                    if key == "messages" and not cls.supports_system_message:
                        return [PARAMETER_EXAMPLES[key][-1]]
                    return PARAMETER_EXAMPLES[key]
                if isinstance(annotation, type):
                    if issubclass(annotation, int):
                        return 0
                    elif issubclass(annotation, float):
                        return 0.0
                    elif issubclass(annotation, bool):
                        return False
                    elif issubclass(annotation, str):
                        return ""
                    elif issubclass(annotation, dict):
                        return {}
                    elif issubclass(annotation, list):
                        return []
                    elif issubclass(annotation, BaseConversation):
                        return {}
                    elif issubclass(annotation, NoneType):
                        return {}
                elif annotation is None:
                    return None
                elif annotation == "str" or annotation == "list[str]":
                    return default
                elif isinstance(annotation, _GenericAlias):
                    if annotation.__origin__ is Optional:
                        return get_type_as_var(annotation.__args__[0])
                else:
                    return str(annotation)

            return {
                name: (
                    param.default
                    if isinstance(param, Parameter)
                    and param.default is not Parameter.empty
                    and param.default is not None
                    else get_type_as_var(param.annotation, name, param.default)
                    if isinstance(param, Parameter)
                    else param
                )
                for name, param in {
                    **BASIC_PARAMETERS,
                    **params,
                    **{
                        "provider": cls.__name__,
                        "model": getattr(cls, "default_model", ""),
                        "stream": cls.supports_stream,
                    },
                }.items()
            }
        return params

    @classmethod
    @property
    def params(cls) -> str:
        """
        Returns the parameters supported by the provider.

        Args:
            cls (type): The class on which this property is called.

        Returns:
            str: A string listing the supported parameters.
        """

        def get_type_name(annotation: type) -> str:
            return (
                getattr(annotation, "__name__", str(annotation))
                if annotation is not Parameter.empty
                else ""
            )

        args = ""
        for name, param in cls.get_parameters().items():
            args += f"\n    {name}"
            args += f": {get_type_name(param.annotation)}"
            default_value = (
                getattr(cls, "default_model", "") if name == "model" else param.default
            )
            default_value = (
                f'"{default_value}"'
                if isinstance(default_value, str)
                else default_value
            )
            args += (
                f" = {default_value}" if param.default is not Parameter.empty else ""
            )
            args += ","

        return f"g4f.Provider.{cls.__name__} supports: ({args}\n)"


class AsyncProvider(AbstractProvider):
    """
    Provides asynchronous functionality for creating completions.
    """

    @staticmethod
    @abstractmethod
    async def create_async(model: str, messages: Messages, **kwargs) -> str:
        """
        Abstract method for creating asynchronous results.

        Args:
            model (str): The model to use for creation.
            messages (Messages): The messages to process.
            **kwargs: Additional keyword arguments.

        Raises:
            NotImplementedError: If this method is not overridden in derived classes.

        Returns:
            str: The created result as a string.
        """
        raise NotImplementedError()


class AsyncGeneratorProvider(AbstractProvider):
    """
    Provides asynchronous generator functionality for streaming results.
    """

    supports_stream = True
    use_stream_timeout = True
    quota_url = None

    @classmethod
    async def get_quota(cls, api_key: Optional[str] = None, **kwargs) -> dict:
        """Get the quota information for the API key."""
        if cls.quota_url is None:
            raise NotImplementedError(
                f"{cls.__name__} does not implement get_quota method"
            )
        if not api_key and cls.needs_auth:
            raise MissingAuthError("API key is required.")
        headers = {"authorization": f"Bearer {api_key}"} if api_key else {}
        async with ClientSession() as session:
            async with session.get(cls.quota_url, headers=headers) as response:
                await raise_for_status(response)
                return await response.json()

    @staticmethod
    @abstractmethod
    async def create_async_generator(
        model: str, messages: Messages, **kwargs
    ) -> AsyncResult:
        """
        Abstract method for creating an asynchronous generator.

        Args:
            model (str): The model to use for creation.
            messages (Messages): The messages to process.
            **kwargs: Additional keyword arguments.

        Raises:
            NotImplementedError: If this method is not overridden in derived classes.

        Returns:
            AsyncResult: An asynchronous generator yielding results.
        """
        raise NotImplementedError()


class ProviderModelMixin:
    default_model: str = None
    models: list[str] = []
    model_aliases: Optional[dict[str, str]] = None
    models_count: dict = {}
    image_models: list = []
    vision_models: list = []
    video_models: list = []
    audio_models: dict = {}
    last_model: str = None
    models_loaded: bool = False
    models_tags: dict[str, list[str]] = None

    @classmethod
    def get_models(cls, api_key: str = None, **kwargs) -> list[str]:
        if not cls.models and cls.default_model is not None:
            cls.models = [cls.default_model]
        if not cls.models_loaded and hasattr(cls, "get_cache_file"):
            if cls.get_cache_file().exists():
                cls.live += 1
        cls.models_loaded = True
        return cls.models

    @classmethod
    def get_model(cls, model: str, **kwargs) -> str:
        if (not model or model == "auto") and cls.default_model is not None:
            return cls.default_model
        if cls.model_aliases is not None and model in cls.model_aliases:
            alias = cls.model_aliases[model]
            if isinstance(alias, list):
                selected_model = random.choice(alias)
                debug.log(
                    f"{cls.__name__}: Selected model '{selected_model}' from alias '{model}'"
                )
                return selected_model
            debug.log(f"{cls.__name__}: Using model '{alias}' for alias '{model}'")
            return alias
        for m in cls.get_models(**kwargs):
            if not model or model == "auto":
                if m == "auto" or m.endswith("/free"):
                    cls.default_model = m
                    return m
            elif m.lower().endswith(model.lower()):
                debug.log(f"{cls.__name__}: Using model '{m}'")
                return m
        return model


class RaiseErrorMixin:
    @staticmethod
    def raise_error(data: dict, status: int = None):
        message = None
        if "error_message" in data:
            message = data["error_message"]
        elif "error" in data:
            if isinstance(data["error"], str):
                message = data["error"]
            elif isinstance(data["error"], bool):
                message = str(data)
            elif "code" in data["error"] and data["error"]["code"]:
                message = "\n".join(
                    [
                        e
                        for e in [
                            f'Error {data["error"]["code"]}: {data["error"]["message"]}',
                            data["error"].get("failed_generation"),
                        ]
                        if e is not None
                    ]
                )
            elif "message" in data["error"]:
                message = data["error"]["message"]
            else:
                message = str(data["error"])
        if message is not None:
            if status is not None:
                if status == 401:
                    raise MissingAuthError(f"Error {status}: {message}")
                elif status == 402:
                    raise PaymentRequiredError(f"Error {status}: {message}")
                elif status == 400:
                    raise BadRequestError(f"Error {status}: {message}")
                elif status == 429:
                    raise RateLimitError(f"Error {status}: {message}")
                raise ResponseError(f"Error {status}: {message}")
            raise ResponseError(f"Error: {message}")


class AuthFileMixin:
    @classmethod
    def get_cache_file(cls) -> Path:
        return (
            Path(get_cookies_dir())
            / f"auth_{cls.parent if hasattr(cls, 'parent') else cls.__name__}.json"
        )

    @classmethod
    def delete_cache_file(cls):
        cache_file = cls.get_cache_file()
        if cache_file.exists():
            try:
                cache_file.unlink()
            except OSError:
                pass


class AsyncAuthedProvider(AsyncGeneratorProvider, AuthFileMixin):
    @classmethod
    async def on_auth_async(cls, **kwargs) -> AuthResult:
        if "api_key" not in kwargs:
            raise MissingAuthError(f"API key is required for {cls.__name__}")
        return AuthResult()

    @classmethod
    def reset_auth(cls):
        """
        Invalidate cached authentication after an auth failure.

        Removes the persisted auth cache file. Providers that keep auth
        state in class attributes (e.g. an access token) should override
        this to clear that in-memory state as well, so the following
        login performs a fresh authentication instead of reusing the
        rejected credentials.
        """
        cls.delete_cache_file()

    @classmethod
    def write_cache_file(cls, cache_file: Path, auth_result: AuthResult = None):
        if auth_result is not None:
            cache_file.parent.mkdir(parents=True, exist_ok=True)
            try:

                def toJSON(obj):
                    if hasattr(obj, "get_dict"):
                        return obj.get_dict()
                    return str(obj)

                with cache_file.open("w") as cache_file:
                    json.dump(auth_result, cache_file, default=toJSON)
            except TypeError as e:
                raise RuntimeError(
                    f"Failed to save: {auth_result.get_dict()}\n{type(e).__name__}: {e}"
                )
        # elif cache_file.exists():
        #    cache_file.unlink()

    @classmethod
    def get_auth_result(cls) -> AuthResult:
        """
        Retrieves the authentication result from cache.
        """
        cache_file = cls.get_cache_file()
        if cache_file.exists():
            try:
                with cache_file.open("r") as f:
                    return AuthResult(**json.load(f))
            except json.JSONDecodeError:
                cache_file.unlink()
                raise MissingAuthError(f"Invalid auth file: {cache_file}")
        else:
            raise MissingAuthError

    @classmethod
    async def create_async_generator(
        cls, model: str, messages: Messages, **kwargs
    ) -> AsyncResult:
        auth_result: AuthResult = None
        cache_file = cls.get_cache_file()
        try:
            auth_result = cls.get_auth_result()
            response = to_async_iterator(
                cls.create_authed(model, messages, **kwargs, auth_result=auth_result)
            )
            if "stream_timeout" in kwargs or "timeout" in kwargs:
                timeout = (
                    kwargs.get("stream_timeout")
                    if cls.use_stream_timeout
                    else kwargs.get("timeout")
                )
                while True:
                    try:
                        yield await asyncio.wait_for(
                            response.__anext__(), timeout=timeout
                        )
                    except TimeoutError as e:
                        raise TimeoutError(
                            "The operation timed out after {} seconds in {}".format(
                                timeout, cls.__name__
                            )
                        ) from e
                    except StopAsyncIteration:
                        break
            else:
                async for chunk in response:
                    yield chunk
        except (MissingAuthError, NoValidHarFileError, CloudflareError):
            # The cached auth is no longer valid (e.g. a revoked or expired
            # access token). Drop the persisted cache file and any in-memory
            # auth state so the re-login below fetches fresh credentials
            # instead of reusing the rejected ones.
            cls.reset_auth()
            response = cls.on_auth_async(**kwargs)
            async for chunk in response:
                if isinstance(chunk, AuthResult):
                    auth_result = chunk
                else:
                    yield chunk
            response = to_async_iterator(
                cls.create_authed(model, messages, **kwargs, auth_result=auth_result)
            )
            async for chunk in response:
                if cache_file is not None:
                    cls.write_cache_file(cache_file, auth_result)
                    cache_file = None
                yield chunk
