Underground-Digital/Workflow-Engine
0
1from datetime import datetime, timezone2from typing import Optional, Union3 4from sqlalchemy import asc, desc, or_5 6from core.app.entities.app_invoke_entities import InvokeFrom7from core.llm_generator.llm_generator import LLMGenerator8from extensions.ext_database import db9from libs.infinite_scroll_pagination import InfiniteScrollPagination10from models.account import Account11from models.model import App, Conversation, EndUser, Message12from services.errors.conversation import ConversationNotExistsError, LastConversationNotExistsError13from services.errors.message import MessageNotExistsError14 15 16class ConversationService:17 @classmethod18 def pagination_by_last_id(19 cls,20 app_model: App,21 user: Optional[Union[Account, EndUser]],22 last_id: Optional[str],23 limit: int,24 invoke_from: InvokeFrom,25 include_ids: Optional[list] = None,26 exclude_ids: Optional[list] = None,27 sort_by: str = "-updated_at",28 ) -> InfiniteScrollPagination:29 if not user:30 return InfiniteScrollPagination(data=[], limit=limit, has_more=False)31 32 base_query = db.session.query(Conversation).filter(33 Conversation.is_deleted == False,34 Conversation.app_id == app_model.id,35 Conversation.from_source == ("api" if isinstance(user, EndUser) else "console"),36 Conversation.from_end_user_id == (user.id if isinstance(user, EndUser) else None),37 Conversation.from_account_id == (user.id if isinstance(user, Account) else None),38 or_(Conversation.invoke_from.is_(None), Conversation.invoke_from == invoke_from.value),39 )40 41 if include_ids is not None:42 base_query = base_query.filter(Conversation.id.in_(include_ids))43 44 if exclude_ids is not None:45 base_query = base_query.filter(~Conversation.id.in_(exclude_ids))46 47 # define sort fields and directions48 sort_field, sort_direction = cls._get_sort_params(sort_by)49 50 if last_id:51 last_conversation = base_query.filter(Conversation.id == last_id).first()52 if not last_conversation:53 raise LastConversationNotExistsError()54 55 # build filters based on sorting56 filter_condition = cls._build_filter_condition(sort_field, sort_direction, last_conversation)57 base_query = base_query.filter(filter_condition)58 59 base_query = base_query.order_by(sort_direction(getattr(Conversation, sort_field)))60 61 conversations = base_query.limit(limit).all()62 63 has_more = False64 if len(conversations) == limit:65 current_page_last_conversation = conversations[-1]66 rest_filter_condition = cls._build_filter_condition(67 sort_field, sort_direction, current_page_last_conversation, is_next_page=True68 )69 rest_count = base_query.filter(rest_filter_condition).count()70 71 if rest_count > 0:72 has_more = True73 74 return InfiniteScrollPagination(data=conversations, limit=limit, has_more=has_more)75 76 @classmethod77 def _get_sort_params(cls, sort_by: str) -> tuple[str, callable]:78 if sort_by.startswith("-"):79 return sort_by[1:], desc80 return sort_by, asc81 82 @classmethod83 def _build_filter_condition(84 cls, sort_field: str, sort_direction: callable, reference_conversation: Conversation, is_next_page: bool = False85 ):86 field_value = getattr(reference_conversation, sort_field)87 if (sort_direction == desc and not is_next_page) or (sort_direction == asc and is_next_page):88 return getattr(Conversation, sort_field) < field_value89 else:90 return getattr(Conversation, sort_field) > field_value91 92 @classmethod93 def rename(94 cls,95 app_model: App,96 conversation_id: str,97 user: Optional[Union[Account, EndUser]],98 name: str,99 auto_generate: bool,100 ):101 conversation = cls.get_conversation(app_model, conversation_id, user)102 103 if auto_generate:104 return cls.auto_generate_name(app_model, conversation)105 else:106 conversation.name = name107 conversation.updated_at = datetime.now(timezone.utc).replace(tzinfo=None)108 db.session.commit()109 110 return conversation111 112 @classmethod113 def auto_generate_name(cls, app_model: App, conversation: Conversation):114 # get conversation first message115 message = (116 db.session.query(Message)117 .filter(Message.app_id == app_model.id, Message.conversation_id == conversation.id)118 .order_by(Message.created_at.asc())119 .first()120 )121 122 if not message:123 raise MessageNotExistsError()124 125 # generate conversation name126 try:127 name = LLMGenerator.generate_conversation_name(128 app_model.tenant_id, message.query, conversation.id, app_model.id129 )130 conversation.name = name131 except:132 pass133 134 db.session.commit()135 136 return conversation137 138 @classmethod139 def get_conversation(cls, app_model: App, conversation_id: str, user: Optional[Union[Account, EndUser]]):140 conversation = (141 db.session.query(Conversation)142 .filter(143 Conversation.id == conversation_id,144 Conversation.app_id == app_model.id,145 Conversation.from_source == ("api" if isinstance(user, EndUser) else "console"),146 Conversation.from_end_user_id == (user.id if isinstance(user, EndUser) else None),147 Conversation.from_account_id == (user.id if isinstance(user, Account) else None),148 Conversation.is_deleted == False,149 )150 .first()151 )152 153 if not conversation:154 raise ConversationNotExistsError()155 156 return conversation157 158 @classmethod159 def delete(cls, app_model: App, conversation_id: str, user: Optional[Union[Account, EndUser]]):160 conversation = cls.get_conversation(app_model, conversation_id, user)161 162 conversation.is_deleted = True163 db.session.commit()164 