Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
conversation.py322 linesDownload Raw Back to app
1from datetime import datetime, timezone2 3import pytz4from flask_login import current_user5from flask_restful import Resource, marshal_with, reqparse6from flask_restful.inputs import int_range7from sqlalchemy import func, or_8from sqlalchemy.orm import joinedload9from werkzeug.exceptions import Forbidden, NotFound10 11from controllers.console import api12from controllers.console.app.wraps import get_app_model13from controllers.console.wraps import account_initialization_required, setup_required14from core.app.entities.app_invoke_entities import InvokeFrom15from extensions.ext_database import db16from fields.conversation_fields import (17    conversation_detail_fields,18    conversation_message_detail_fields,19    conversation_pagination_fields,20    conversation_with_summary_pagination_fields,21)22from libs.helper import DatetimeString23from libs.login import login_required24from models import Conversation, EndUser, Message, MessageAnnotation25from models.model import AppMode26 27 28class CompletionConversationApi(Resource):29    @setup_required30    @login_required31    @account_initialization_required32    @get_app_model(mode=AppMode.COMPLETION)33    @marshal_with(conversation_pagination_fields)34    def get(self, app_model):35        if not current_user.is_editor:36            raise Forbidden()37        parser = reqparse.RequestParser()38        parser.add_argument("keyword", type=str, location="args")39        parser.add_argument("start", type=DatetimeString("%Y-%m-%d %H:%M"), location="args")40        parser.add_argument("end", type=DatetimeString("%Y-%m-%d %H:%M"), location="args")41        parser.add_argument(42            "annotation_status", type=str, choices=["annotated", "not_annotated", "all"], default="all", location="args"43        )44        parser.add_argument("page", type=int_range(1, 99999), default=1, location="args")45        parser.add_argument("limit", type=int_range(1, 100), default=20, location="args")46        args = parser.parse_args()47 48        query = db.select(Conversation).where(Conversation.app_id == app_model.id, Conversation.mode == "completion")49 50        if args["keyword"]:51            query = query.join(Message, Message.conversation_id == Conversation.id).filter(52                or_(53                    Message.query.ilike("%{}%".format(args["keyword"])),54                    Message.answer.ilike("%{}%".format(args["keyword"])),55                )56            )57 58        account = current_user59        timezone = pytz.timezone(account.timezone)60        utc_timezone = pytz.utc61 62        if args["start"]:63            start_datetime = datetime.strptime(args["start"], "%Y-%m-%d %H:%M")64            start_datetime = start_datetime.replace(second=0)65 66            start_datetime_timezone = timezone.localize(start_datetime)67            start_datetime_utc = start_datetime_timezone.astimezone(utc_timezone)68 69            query = query.where(Conversation.created_at >= start_datetime_utc)70 71        if args["end"]:72            end_datetime = datetime.strptime(args["end"], "%Y-%m-%d %H:%M")73            end_datetime = end_datetime.replace(second=59)74 75            end_datetime_timezone = timezone.localize(end_datetime)76            end_datetime_utc = end_datetime_timezone.astimezone(utc_timezone)77 78            query = query.where(Conversation.created_at < end_datetime_utc)79 80        if args["annotation_status"] == "annotated":81            query = query.options(joinedload(Conversation.message_annotations)).join(82                MessageAnnotation, MessageAnnotation.conversation_id == Conversation.id83            )84        elif args["annotation_status"] == "not_annotated":85            query = (86                query.outerjoin(MessageAnnotation, MessageAnnotation.conversation_id == Conversation.id)87                .group_by(Conversation.id)88                .having(func.count(MessageAnnotation.id) == 0)89            )90 91        query = query.order_by(Conversation.created_at.desc())92 93        conversations = db.paginate(query, page=args["page"], per_page=args["limit"], error_out=False)94 95        return conversations96 97 98class CompletionConversationDetailApi(Resource):99    @setup_required100    @login_required101    @account_initialization_required102    @get_app_model(mode=AppMode.COMPLETION)103    @marshal_with(conversation_message_detail_fields)104    def get(self, app_model, conversation_id):105        if not current_user.is_editor:106            raise Forbidden()107        conversation_id = str(conversation_id)108 109        return _get_conversation(app_model, conversation_id)110 111    @setup_required112    @login_required113    @account_initialization_required114    @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT])115    def delete(self, app_model, conversation_id):116        if not current_user.is_editor:117            raise Forbidden()118        conversation_id = str(conversation_id)119 120        conversation = (121            db.session.query(Conversation)122            .filter(Conversation.id == conversation_id, Conversation.app_id == app_model.id)123            .first()124        )125 126        if not conversation:127            raise NotFound("Conversation Not Exists.")128 129        conversation.is_deleted = True130        db.session.commit()131 132        return {"result": "success"}, 204133 134 135class ChatConversationApi(Resource):136    @setup_required137    @login_required138    @account_initialization_required139    @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT])140    @marshal_with(conversation_with_summary_pagination_fields)141    def get(self, app_model):142        if not current_user.is_editor:143            raise Forbidden()144        parser = reqparse.RequestParser()145        parser.add_argument("keyword", type=str, location="args")146        parser.add_argument("start", type=DatetimeString("%Y-%m-%d %H:%M"), location="args")147        parser.add_argument("end", type=DatetimeString("%Y-%m-%d %H:%M"), location="args")148        parser.add_argument(149            "annotation_status", type=str, choices=["annotated", "not_annotated", "all"], default="all", location="args"150        )151        parser.add_argument("message_count_gte", type=int_range(1, 99999), required=False, location="args")152        parser.add_argument("page", type=int_range(1, 99999), required=False, default=1, location="args")153        parser.add_argument("limit", type=int_range(1, 100), required=False, default=20, location="args")154        parser.add_argument(155            "sort_by",156            type=str,157            choices=["created_at", "-created_at", "updated_at", "-updated_at"],158            required=False,159            default="-updated_at",160            location="args",161        )162        args = parser.parse_args()163 164        subquery = (165            db.session.query(166                Conversation.id.label("conversation_id"), EndUser.session_id.label("from_end_user_session_id")167            )168            .outerjoin(EndUser, Conversation.from_end_user_id == EndUser.id)169            .subquery()170        )171 172        query = db.select(Conversation).where(Conversation.app_id == app_model.id)173 174        if args["keyword"]:175            keyword_filter = "%{}%".format(args["keyword"])176            query = (177                query.join(178                    Message,179                    Message.conversation_id == Conversation.id,180                )181                .join(subquery, subquery.c.conversation_id == Conversation.id)182                .filter(183                    or_(184                        Message.query.ilike(keyword_filter),185                        Message.answer.ilike(keyword_filter),186                        Conversation.name.ilike(keyword_filter),187                        Conversation.introduction.ilike(keyword_filter),188                        subquery.c.from_end_user_session_id.ilike(keyword_filter),189                    ),190                )191                .group_by(Conversation.id)192            )193 194        account = current_user195        timezone = pytz.timezone(account.timezone)196        utc_timezone = pytz.utc197 198        if args["start"]:199            start_datetime = datetime.strptime(args["start"], "%Y-%m-%d %H:%M")200            start_datetime = start_datetime.replace(second=0)201 202            start_datetime_timezone = timezone.localize(start_datetime)203            start_datetime_utc = start_datetime_timezone.astimezone(utc_timezone)204 205            match args["sort_by"]:206                case "updated_at" | "-updated_at":207                    query = query.where(Conversation.updated_at >= start_datetime_utc)208                case "created_at" | "-created_at" | _:209                    query = query.where(Conversation.created_at >= start_datetime_utc)210 211        if args["end"]:212            end_datetime = datetime.strptime(args["end"], "%Y-%m-%d %H:%M")213            end_datetime = end_datetime.replace(second=59)214 215            end_datetime_timezone = timezone.localize(end_datetime)216            end_datetime_utc = end_datetime_timezone.astimezone(utc_timezone)217 218            match args["sort_by"]:219                case "updated_at" | "-updated_at":220                    query = query.where(Conversation.updated_at <= end_datetime_utc)221                case "created_at" | "-created_at" | _:222                    query = query.where(Conversation.created_at <= end_datetime_utc)223 224        if args["annotation_status"] == "annotated":225            query = query.options(joinedload(Conversation.message_annotations)).join(226                MessageAnnotation, MessageAnnotation.conversation_id == Conversation.id227            )228        elif args["annotation_status"] == "not_annotated":229            query = (230                query.outerjoin(MessageAnnotation, MessageAnnotation.conversation_id == Conversation.id)231                .group_by(Conversation.id)232                .having(func.count(MessageAnnotation.id) == 0)233            )234 235        if args["message_count_gte"] and args["message_count_gte"] >= 1:236            query = (237                query.options(joinedload(Conversation.messages))238                .join(Message, Message.conversation_id == Conversation.id)239                .group_by(Conversation.id)240                .having(func.count(Message.id) >= args["message_count_gte"])241            )242 243        if app_model.mode == AppMode.ADVANCED_CHAT.value:244            query = query.where(Conversation.invoke_from != InvokeFrom.DEBUGGER.value)245 246        match args["sort_by"]:247            case "created_at":248                query = query.order_by(Conversation.created_at.asc())249            case "-created_at":250                query = query.order_by(Conversation.created_at.desc())251            case "updated_at":252                query = query.order_by(Conversation.updated_at.asc())253            case "-updated_at":254                query = query.order_by(Conversation.updated_at.desc())255            case _:256                query = query.order_by(Conversation.created_at.desc())257 258        conversations = db.paginate(query, page=args["page"], per_page=args["limit"], error_out=False)259 260        return conversations261 262 263class ChatConversationDetailApi(Resource):264    @setup_required265    @login_required266    @account_initialization_required267    @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT])268    @marshal_with(conversation_detail_fields)269    def get(self, app_model, conversation_id):270        if not current_user.is_editor:271            raise Forbidden()272        conversation_id = str(conversation_id)273 274        return _get_conversation(app_model, conversation_id)275 276    @setup_required277    @login_required278    @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT])279    @account_initialization_required280    def delete(self, app_model, conversation_id):281        if not current_user.is_editor:282            raise Forbidden()283        conversation_id = str(conversation_id)284 285        conversation = (286            db.session.query(Conversation)287            .filter(Conversation.id == conversation_id, Conversation.app_id == app_model.id)288            .first()289        )290 291        if not conversation:292            raise NotFound("Conversation Not Exists.")293 294        conversation.is_deleted = True295        db.session.commit()296 297        return {"result": "success"}, 204298 299 300api.add_resource(CompletionConversationApi, "/apps/<uuid:app_id>/completion-conversations")301api.add_resource(CompletionConversationDetailApi, "/apps/<uuid:app_id>/completion-conversations/<uuid:conversation_id>")302api.add_resource(ChatConversationApi, "/apps/<uuid:app_id>/chat-conversations")303api.add_resource(ChatConversationDetailApi, "/apps/<uuid:app_id>/chat-conversations/<uuid:conversation_id>")304 305 306def _get_conversation(app_model, conversation_id):307    conversation = (308        db.session.query(Conversation)309        .filter(Conversation.id == conversation_id, Conversation.app_id == app_model.id)310        .first()311    )312 313    if not conversation:314        raise NotFound("Conversation Not Exists.")315 316    if not conversation.read_at:317        conversation.read_at = datetime.now(timezone.utc).replace(tzinfo=None)318        conversation.read_account_id = current_user.id319        db.session.commit()320 321    return conversation322