import json
import logging
from collections.abc import Callable, Generator, Mapping
from typing import Union, cast

from sqlalchemy import select
from sqlalchemy.orm import Session

from core.app.app_config.entities import EasyUIBasedAppConfig, EasyUIBasedAppModelConfigFrom
from core.app.apps.base_app_generator import BaseAppGenerator
from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.streaming_utils import stream_topic_events
from core.app.entities.app_invoke_entities import (
    AdvancedChatAppGenerateEntity,
    AgentChatAppGenerateEntity,
    AppGenerateEntity,
    ChatAppGenerateEntity,
    CompletionAppGenerateEntity,
    ConversationAppGenerateEntity,
    InvokeFrom,
)
from core.app.entities.task_entities import (
    ChatbotAppBlockingResponse,
    ChatbotAppStreamResponse,
    CompletionAppBlockingResponse,
    CompletionAppStreamResponse,
)
from core.app.task_pipeline.easy_ui_based_generate_task_pipeline import EasyUIBasedGenerateTaskPipeline
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
from core.workflow.file_reference import resolve_file_record_id
from extensions.ext_database import db
from extensions.ext_redis import get_pubsub_broadcast_channel
from libs.broadcast_channel.channel import Topic
from libs.datetime_utils import naive_utc_now
from models import Account
from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo
from models.model import App, AppMode, AppModelConfig, Conversation, EndUser, Message, MessageFile
from services.errors.app_model_config import AppModelConfigBrokenError
from services.errors.conversation import ConversationNotExistsError
from services.errors.message import MessageNotExistsError

logger = logging.getLogger(__name__)


class MessageBasedAppGenerator(BaseAppGenerator):
    def _handle_response(
        self,
        application_generate_entity: Union[
            ChatAppGenerateEntity,
            CompletionAppGenerateEntity,
            AgentChatAppGenerateEntity,
        ],
        queue_manager: AppQueueManager,
        conversation: Conversation,
        message: Message,
        user: Union[Account, EndUser],
        stream: bool = False,
    ) -> Union[
        ChatbotAppBlockingResponse,
        CompletionAppBlockingResponse,
        Generator[Union[ChatbotAppStreamResponse, CompletionAppStreamResponse], None, None],
    ]:
        """
        Handle response.
        :param application_generate_entity: application generate entity
        :param queue_manager: queue manager
        :param conversation: conversation
        :param message: message
        :param user: user
        :param stream: is stream
        :return:
        """
        # init generate task pipeline
        generate_task_pipeline = EasyUIBasedGenerateTaskPipeline(
            application_generate_entity=application_generate_entity,
            queue_manager=queue_manager,
            conversation=conversation,
            message=message,
            stream=stream,
        )

        try:
            return generate_task_pipeline.process()
        except ValueError as e:
            if len(e.args) > 0 and e.args[0] == "I/O operation on closed file.":  # ignore this error
                raise GenerateTaskStoppedError()
            else:
                logger.exception("Failed to handle response, conversation_id: %s", conversation.id)
                raise e

    def _get_app_model_config(self, app_model: App, conversation: Conversation | None = None) -> AppModelConfig:
        if conversation:
            stmt = select(AppModelConfig).where(
                AppModelConfig.id == conversation.app_model_config_id, AppModelConfig.app_id == app_model.id
            )
            app_model_config = db.session.scalar(stmt)

            if not app_model_config:
                raise AppModelConfigBrokenError()
        else:
            if app_model.app_model_config_id is None:
                raise AppModelConfigBrokenError()

            app_model_config = app_model.app_model_config

            if not app_model_config:
                raise AppModelConfigBrokenError()

        return app_model_config

    def _init_generate_records(
        self,
        application_generate_entity: Union[
            ChatAppGenerateEntity,
            CompletionAppGenerateEntity,
            AgentChatAppGenerateEntity,
            AdvancedChatAppGenerateEntity,
        ],
        conversation: Conversation | None = None,
    ) -> tuple[Conversation, Message]:
        """
        Initialize generate records
        :param application_generate_entity: application generate entity
        :conversation conversation
        :return:
        """
        app_config: EasyUIBasedAppConfig = cast(EasyUIBasedAppConfig, application_generate_entity.app_config)

        # get from source
        end_user_id = None
        account_id = None
        if application_generate_entity.invoke_from in {InvokeFrom.WEB_APP, InvokeFrom.SERVICE_API}:
            from_source = ConversationFromSource.API
            end_user_id = application_generate_entity.user_id
        else:
            from_source = ConversationFromSource.CONSOLE
            account_id = application_generate_entity.user_id

        if isinstance(application_generate_entity, AdvancedChatAppGenerateEntity):
            app_model_config_id = None
            override_model_configs = None
            model_provider = None
            model_id = None
        else:
            app_model_config_id = app_config.app_model_config_id
            model_provider = application_generate_entity.model_conf.provider
            model_id = application_generate_entity.model_conf.model
            override_model_configs = None
            if app_config.app_model_config_from == EasyUIBasedAppModelConfigFrom.ARGS and app_config.app_mode in {
                AppMode.AGENT_CHAT,
                AppMode.CHAT,
                AppMode.COMPLETION,
            }:
                override_model_configs = app_config.app_model_config_dict

        # get conversation introduction
        introduction = self._get_conversation_introduction(application_generate_entity)

        # get conversation name
        query = application_generate_entity.query or "New conversation"
        conversation_name = (query[:20] + "…") if len(query) > 20 else query

        created_new_conversation = conversation is None
        try:
            if not conversation:
                conversation = Conversation(
                    app_id=app_config.app_id,
                    app_model_config_id=app_model_config_id,
                    model_provider=model_provider,
                    model_id=model_id,
                    override_model_configs=json.dumps(override_model_configs) if override_model_configs else None,
                    mode=app_config.app_mode.value,
                    name=conversation_name,
                    inputs=application_generate_entity.inputs,
                    introduction=introduction,
                    system_instruction="",
                    system_instruction_tokens=0,
                    status="normal",
                    invoke_from=application_generate_entity.invoke_from.value,
                    from_source=from_source,
                    from_end_user_id=end_user_id,
                    from_account_id=account_id,
                )

                db.session.add(conversation)
                db.session.flush()
                db.session.refresh(conversation)
            else:
                conversation.updated_at = naive_utc_now()

            message = Message(
                app_id=app_config.app_id,
                model_provider=model_provider,
                model_id=model_id,
                override_model_configs=json.dumps(override_model_configs) if override_model_configs else None,
                conversation_id=conversation.id,
                inputs=application_generate_entity.inputs,
                query=application_generate_entity.query,
                message="",
                message_tokens=0,
                message_unit_price=0,
                message_price_unit=0,
                answer="",
                answer_tokens=0,
                answer_unit_price=0,
                answer_price_unit=0,
                parent_message_id=getattr(application_generate_entity, "parent_message_id", None),
                provider_response_latency=0,
                total_price=0,
                currency="USD",
                invoke_from=application_generate_entity.invoke_from.value,
                from_source=from_source,
                from_end_user_id=end_user_id,
                from_account_id=account_id,
                app_mode=app_config.app_mode,
            )

            db.session.add(message)
            db.session.flush()
            db.session.refresh(message)

            message_files = []
            for file in application_generate_entity.files:
                message_file = MessageFile(
                    message_id=message.id,
                    type=file.type,
                    transfer_method=file.transfer_method,
                    belongs_to=MessageFileBelongsTo.USER,
                    url=file.remote_url,
                    upload_file_id=resolve_file_record_id(file.reference),
                    created_by_role=(CreatorUserRole.ACCOUNT if account_id else CreatorUserRole.END_USER),
                    created_by=account_id or end_user_id or "",
                )
                message_files.append(message_file)

            if message_files:
                db.session.add_all(message_files)

            db.session.commit()

            if isinstance(application_generate_entity, ConversationAppGenerateEntity):
                application_generate_entity.conversation_id = conversation.id
                application_generate_entity.is_new_conversation = created_new_conversation
            return conversation, message
        except Exception:
            db.session.rollback()
            raise

    def _get_conversation_introduction(self, application_generate_entity: AppGenerateEntity) -> str:
        """
        Get conversation introduction
        :param application_generate_entity: application generate entity
        :return: conversation introduction
        """
        app_config = application_generate_entity.app_config
        introduction = app_config.additional_features.opening_statement

        if introduction:
            try:
                inputs = application_generate_entity.inputs
                prompt_template = PromptTemplateParser(template=introduction)
                prompt_inputs = {k: inputs[k] for k in prompt_template.variable_keys if k in inputs}
                introduction = prompt_template.format(prompt_inputs)
            except KeyError:
                pass

        return introduction or ""

    def _get_conversation(self, conversation_id: str) -> Conversation:
        """
        Get conversation by conversation id
        :param conversation_id: conversation id
        :return: conversation
        """
        with Session(db.engine, expire_on_commit=False) as session:
            conversation = session.scalar(select(Conversation).where(Conversation.id == conversation_id))

        if not conversation:
            raise ConversationNotExistsError("Conversation not exists")

        return conversation

    def _get_message(self, message_id: str) -> Message:
        """
        Get message by message id
        :param message_id: message id
        :return: message
        """
        with Session(db.engine, expire_on_commit=False) as session:
            message = session.scalar(select(Message).where(Message.id == message_id))

        if message is None:
            raise MessageNotExistsError("Message not exists")

        return message

    @staticmethod
    def _make_channel_key(app_mode: AppMode, workflow_run_id: str):
        return f"channel:{app_mode}:{workflow_run_id}"

    @classmethod
    def get_response_topic(cls, app_mode: AppMode, workflow_run_id: str) -> Topic:
        key = cls._make_channel_key(app_mode, workflow_run_id)
        channel = get_pubsub_broadcast_channel()
        topic = channel.topic(key)
        return topic

    @classmethod
    def retrieve_events(
        cls,
        app_mode: AppMode,
        workflow_run_id: str,
        idle_timeout=300,
        on_subscribe: Callable[[], None] | None = None,
    ) -> Generator[Mapping | str, None, None]:
        topic = cls.get_response_topic(app_mode, workflow_run_id)
        return stream_topic_events(
            topic=topic,
            idle_timeout=idle_timeout,
            on_subscribe=on_subscribe,
        )
