Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
based_generate_task_pipeline.py153 linesDownload Raw Back to task_pipeline
1import logging2import time3from typing import Optional, Union4 5from core.app.apps.base_app_queue_manager import AppQueueManager6from core.app.entities.app_invoke_entities import (7    AppGenerateEntity,8)9from core.app.entities.queue_entities import (10    QueueErrorEvent,11)12from core.app.entities.task_entities import (13    ErrorStreamResponse,14    PingStreamResponse,15    TaskState,16)17from core.errors.error import QuotaExceededError18from core.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError19from core.moderation.output_moderation import ModerationRule, OutputModeration20from extensions.ext_database import db21from models.account import Account22from models.model import EndUser, Message23 24logger = logging.getLogger(__name__)25 26 27class BasedGenerateTaskPipeline:28    """29    BasedGenerateTaskPipeline is a class that generate stream output and state management for Application.30    """31 32    _task_state: TaskState33    _application_generate_entity: AppGenerateEntity34 35    def __init__(36        self,37        application_generate_entity: AppGenerateEntity,38        queue_manager: AppQueueManager,39        user: Union[Account, EndUser],40        stream: bool,41    ) -> None:42        """43        Initialize GenerateTaskPipeline.44        :param application_generate_entity: application generate entity45        :param queue_manager: queue manager46        :param user: user47        :param stream: stream48        """49        self._application_generate_entity = application_generate_entity50        self._queue_manager = queue_manager51        self._user = user52        self._start_at = time.perf_counter()53        self._output_moderation_handler = self._init_output_moderation()54        self._stream = stream55 56    def _handle_error(self, event: QueueErrorEvent, message: Optional[Message] = None):57        """58        Handle error event.59        :param event: event60        :param message: message61        :return:62        """63        logger.debug("error: %s", event.error)64        e = event.error65 66        if isinstance(e, InvokeAuthorizationError):67            err = InvokeAuthorizationError("Incorrect API key provided")68        elif isinstance(e, InvokeError | ValueError):69            err = e70        else:71            err = Exception(e.description if getattr(e, "description", None) is not None else str(e))72 73        if message:74            refetch_message = db.session.query(Message).filter(Message.id == message.id).first()75 76            if refetch_message:77                err_desc = self._error_to_desc(err)78                refetch_message.status = "error"79                refetch_message.error = err_desc80 81                db.session.commit()82 83        return err84 85    def _error_to_desc(self, e: Exception) -> str:86        """87        Error to desc.88        :param e: exception89        :return:90        """91        if isinstance(e, QuotaExceededError):92            return (93                "Your quota for Dify Hosted Model Provider has been exhausted. "94                "Please go to Settings -> Model Provider to complete your own provider credentials."95            )96 97        message = getattr(e, "description", str(e))98        if not message:99            message = "Internal Server Error, please contact support."100 101        return message102 103    def _error_to_stream_response(self, e: Exception):104        """105        Error to stream response.106        :param e: exception107        :return:108        """109        return ErrorStreamResponse(task_id=self._application_generate_entity.task_id, err=e)110 111    def _ping_stream_response(self) -> PingStreamResponse:112        """113        Ping stream response.114        :return:115        """116        return PingStreamResponse(task_id=self._application_generate_entity.task_id)117 118    def _init_output_moderation(self) -> Optional[OutputModeration]:119        """120        Init output moderation.121        :return:122        """123        app_config = self._application_generate_entity.app_config124        sensitive_word_avoidance = app_config.sensitive_word_avoidance125 126        if sensitive_word_avoidance:127            return OutputModeration(128                tenant_id=app_config.tenant_id,129                app_id=app_config.app_id,130                rule=ModerationRule(type=sensitive_word_avoidance.type, config=sensitive_word_avoidance.config),131                queue_manager=self._queue_manager,132            )133 134    def _handle_output_moderation_when_task_finished(self, completion: str) -> Optional[str]:135        """136        Handle output moderation when task finished.137        :param completion: completion138        :return:139        """140        # response moderation141        if self._output_moderation_handler:142            self._output_moderation_handler.stop_thread()143 144            completion = self._output_moderation_handler.moderation_completion(145                completion=completion, public_event=False146            )147 148            self._output_moderation_handler = None149 150            return completion151 152        return None153