# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.

import json
from typing import Any, Dict, Generator, List, Union

import dashscope
from dashscope.aigc.generation import Generation
from dashscope.api_entities.chat_completion_types import (
    ChatCompletion,
    ChatCompletionChunk,
)
from dashscope.api_entities.dashscope_response import (
    GenerationResponse,
    Message,
)
from dashscope.client.base_api import BaseAioApi, CreateMixin
from dashscope.common.error import InputRequired, ModelRequired
from dashscope.common.utils import _get_task_group_and_task


class Completions(CreateMixin):
    """Support openai compatible chat completion interface."""

    SUB_PATH = ""

    @classmethod
    def create(
        cls,
        *,
        model: str,
        messages: List[Message],
        stream: bool = False,
        temperature: float = None,
        top_p: float = None,
        top_k: int = None,
        stop: Union[List[str], List[List[int]]] = None,
        max_tokens: int = None,
        repetition_penalty: float = None,
        api_key: str = None,
        workspace: str = None,
        extra_headers: Dict = None,
        extra_body: Dict = None,
        **kwargs,
    ) -> Union[ChatCompletion, Generator[ChatCompletionChunk, None, None]]:
        """Call openai compatible chat completion model service.

        Args:
            model (str): The requested model, such as qwen-long
            messages (list): The generation messages.
                examples:
                    [{'role': 'user',
                      'content': 'The weather is fine today.'},
                      {'role': 'assistant', 'content': 'Suitable for outings'}]
            stream(bool, `optional`): Enable server-sent events
                (ref: https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events)  # noqa E501  # pylint: disable=line-too-long
                the result will back partially[qwen-turbo,bailian-v1].
            temperature(float, `optional`): Used to control the degree
                of randomness and diversity. Specifically, the temperature
                value controls the degree to which the probability distribution
                of each candidate word is smoothed when generating text.
                A higher temperature value will reduce the peak value of
                the probability, allowing more low-probability words to be
                selected, and the generated results will be more diverse;
                while a lower temperature value will enhance the peak value
                of the probability, making it easier for high-probability
                words to be selected, the generated results are more
                deterministic.
            top_p(float, `optional`): A sampling strategy, called nucleus
                sampling, where the model considers the results of the
                tokens with top_p probability mass. So 0.1 means only
                the tokens comprising the top 10% probability mass are
                considered.
            top_k(int, `optional`): The size of the sample candidate set when generated.  # noqa E501  # pylint: disable=line-too-long
                For example, when the value is 50, only the 50 highest-scoring tokens  # noqa E501
                in a single generation form a randomly sampled candidate set. # noqa E501
                The larger the value, the higher the randomness generated;  # noqa E501
                the smaller the value, the higher the certainty generated. # noqa E501
                The default value is 0, which means the top_k policy is  # noqa E501
                not enabled. At this time, only the top_p policy takes effect. # noqa E501
            stop(list[str] or list[list[int]], `optional`): Used to control the generation to stop  # noqa E501  # pylint: disable=line-too-long
                when encountering setting str or token ids, the result will not include # noqa E501
                stop words or tokens.
            max_tokens(int, `optional`): The maximum token num expected to be output. It should be # noqa E501  # pylint: disable=line-too-long
                noted that the length generated by the model will only be less than max_tokens,  # noqa E501  # pylint: disable=line-too-long
                not necessarily equal to it. If max_tokens is set too large, the service will # noqa E501  # pylint: disable=line-too-long
                directly prompt that the length exceeds the limit. It is generally # noqa E501
                not recommended to set this value.
            repetition_penalty(float, `optional`): Used to control the repeatability when generating models.  # noqa E501  # pylint: disable=line-too-long
                Increasing repetition_penalty can reduce the duplication of model generation.  # noqa E501  # pylint: disable=line-too-long
                1.0 means no punishment.
            api_key (str, optional): The api api_key, can be None,
                if None, will get by default rule.
            workspace (str, optional): The bailian workspace id.
            **kwargs:
                timeout: set request timeout.
        Raises:
            InvalidInput: The history and auto_history are mutually exclusive.

        Returns:
            Union[ChatCompletion,
                  Generator[ChatCompletionChunk, None, None]]: If
            stream is True, return Generator, otherwise ChatCompletion.
        """
        if messages is None or not messages:
            raise InputRequired("Messages is required!")
        if model is None or not model:
            raise ModelRequired("Model is required!")
        data = {}
        data["model"] = model
        data["messages"] = messages
        if temperature is not None:
            data["temperature"] = temperature
        if top_p is not None:
            data["top_p"] = top_p
        if top_k is not None:
            data["top_k"] = top_k
        if stop is not None:
            data["stop"] = stop
        if max_tokens is not None:
            data["max_tokens"] = max_tokens
        if repetition_penalty is not None:
            data["repetition_penalty"] = repetition_penalty
        if extra_body is not None and extra_body:
            data = {**data, **extra_body}

        if extra_headers is not None and extra_headers:
            kwargs["headers"] = {
                **kwargs.get("headers", {}),
                **extra_headers,
            }

        response = super().call(
            data=data,
            path="chat/completions",
            base_address=dashscope.base_compatible_api_url,
            api_key=api_key,
            flattened_output=True,
            stream=stream,
            workspace=workspace,
            **kwargs,
        )
        if stream:
            return (ChatCompletionChunk(**item) for _, item in response)
        else:
            return ChatCompletion(**response)


class AioGeneration(BaseAioApi):
    task = "text-generation"
    """API for AI-Generated Content(AIGC) models.

    """

    class Models:
        """@deprecated, use qwen_turbo instead"""

        qwen_v1 = "qwen-v1"
        """@deprecated, use qwen_plus instead"""
        qwen_plus_v1 = "qwen-plus-v1"

        bailian_v1 = "bailian-v1"
        dolly_12b_v2 = "dolly-12b-v2"
        qwen_turbo = "qwen-turbo"
        qwen_plus = "qwen-plus"
        qwen_max = "qwen-max"

    @classmethod
    # type: ignore[override]
    async def call(  # pylint: disable=arguments-renamed  # type: ignore[override] # noqa: E501
        cls,
        model: str,
        prompt: Any = None,
        history: list = None,
        api_key: str = None,
        messages: List[Message] = None,
        plugins: Union[str, Dict[str, Any]] = None,
        workspace: str = None,
        **kwargs,
    ) -> Union[GenerationResponse, Generator[GenerationResponse, None, None]]:
        """Call generation model service.

        Args:
            model (str): The requested model, such as qwen-turbo
            prompt (Any): The input prompt.
            history (list):The user provided history, deprecated
                examples:
                    [{'user':'The weather is fine today.',
                    'bot': 'Suitable for outings'}].
                Defaults to None.
            api_key (str, optional): The api api_key, can be None,
                if None, will get by default rule(TODO: api key doc).
            messages (list): The generation messages.
                examples:
                    [{'role': 'user',
                      'content': 'The weather is fine today.'},
                      {'role': 'assistant', 'content': 'Suitable for outings'}]
            plugins (Any): The plugin config. Can be plugins config str, or dict.
            **kwargs:
                stream(bool, `optional`): Enable server-sent events
                    (ref: https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events)  # noqa E501  # pylint: disable=line-too-long
                    the result will back partially[qwen-turbo,bailian-v1].
                temperature(float, `optional`): Used to control the degree
                    of randomness and diversity. Specifically, the temperature
                    value controls the degree to which the probability distribution
                    of each candidate word is smoothed when generating text.
                    A higher temperature value will reduce the peak value of
                    the probability, allowing more low-probability words to be
                    selected, and the generated results will be more diverse;
                    while a lower temperature value will enhance the peak value
                    of the probability, making it easier for high-probability
                    words to be selected, the generated results are more
                    deterministic, range(0, 2) .[qwen-turbo,qwen-plus].
                top_p(float, `optional`): A sampling strategy, called nucleus
                    sampling, where the model considers the results of the
                    tokens with top_p probability mass. So 0.1 means only
                    the tokens comprising the top 10% probability mass are
                    considered[qwen-turbo,bailian-v1].
                top_k(int, `optional`): The size of the sample candidate set when generated.  # noqa E501  # pylint: disable=line-too-long
                    For example, when the value is 50, only the 50 highest-scoring tokens  # noqa E501  # pylint: disable=line-too-long
                    in a single generation form a randomly sampled candidate set. # noqa E501
                    The larger the value, the higher the randomness generated;  # noqa E501
                    the smaller the value, the higher the certainty generated. # noqa E501
                    The default value is 0, which means the top_k policy is  # noqa E501
                    not enabled. At this time, only the top_p policy takes effect. # noqa E501
                enable_search(bool, `optional`): Whether to enable web search(quark).  # noqa E501
                    Currently works best only on the first round of conversation.
                    Default to False, support model: [qwen-turbo].
                customized_model_id(str, required) The enterprise-specific
                    large model id, which needs to be generated from the
                    operation background of the enterprise-specific
                    large model product, support model: [bailian-v1].
                result_format(str, `optional`): [message|text] Set result result format. # noqa E501
                    Default result is text
                incremental_output(bool, `optional`): Used to control the streaming output mode. # noqa E501  # pylint: disable=line-too-long
                    If true, the subsequent output will include the previously input content. # noqa E501  # pylint: disable=line-too-long
                    Otherwise, the subsequent output will not include the previously output # noqa E501  # pylint: disable=line-too-long
                    content. Default false.
                stop(list[str] or list[list[int]], `optional`): Used to control the generation to stop  # noqa E501  # pylint: disable=line-too-long
                    when encountering setting str or token ids, the result will not include # noqa E501  # pylint: disable=line-too-long
                    stop words or tokens.
                max_tokens(int, `optional`): The maximum token num expected to be output. It should be # noqa E501  # pylint: disable=line-too-long
                    noted that the length generated by the model will only be less than max_tokens,  # noqa E501  # pylint: disable=line-too-long
                    not necessarily equal to it. If max_tokens is set too large, the service will # noqa E501  # pylint: disable=line-too-long
                    directly prompt that the length exceeds the limit. It is generally # noqa E501
                    not recommended to set this value.
                repetition_penalty(float, `optional`): Used to control the repeatability when generating models.  # noqa E501  # pylint: disable=line-too-long
                    Increasing repetition_penalty can reduce the duplication of model generation.  # noqa E501  # pylint: disable=line-too-long
                    1.0 means no punishment.
            workspace (str): The dashscope workspace id.
        Raises:
            InvalidInput: The history and auto_history are mutually exclusive.

        Returns:
            Union[GenerationResponse,
                  Generator[GenerationResponse, None, None]]: If
            stream is True, return Generator, otherwise GenerationResponse.
        """
        if (prompt is None or not prompt) and (
            messages is None or not messages
        ):
            raise InputRequired("prompt or messages is required!")
        if model is None or not model:
            raise ModelRequired("Model is required!")
        task_group, function = _get_task_group_and_task(__name__)
        if plugins is not None:
            headers = kwargs.pop("headers", {})
            if isinstance(plugins, str):
                headers["X-DashScope-Plugin"] = plugins
            else:
                headers["X-DashScope-Plugin"] = json.dumps(plugins)
            kwargs["headers"] = headers
        # pylint: disable=protected-access
        (
            input,  # pylint: disable=redefined-builtin
            parameters,
        ) = Generation._build_input_parameters(
            model,
            prompt,
            history,
            messages,
            **kwargs,
        )
        response = await super().call(
            model=model,
            task_group=task_group,
            task=Generation.task,
            function=function,
            api_key=api_key,
            input=input,
            workspace=workspace,
            **parameters,
        )
        is_stream = kwargs.get("stream", False)
        if is_stream:
            return (  # type: ignore[return-value]
                GenerationResponse.from_api_response(rsp)
                async for rsp in response
            )
        else:
            return GenerationResponse.from_api_response(response)
