Underground-Digital/Workflow-Engine
0
1import logging2from typing import Optional3 4from core.app.app_config.entities import AppConfig5from core.moderation.base import ModerationAction, ModerationError6from core.moderation.factory import ModerationFactory7from core.ops.entities.trace_entity import TraceTaskName8from core.ops.ops_trace_manager import TraceQueueManager, TraceTask9from core.ops.utils import measure_time10 11logger = logging.getLogger(__name__)12 13 14class InputModeration:15 def check(16 self,17 app_id: str,18 tenant_id: str,19 app_config: AppConfig,20 inputs: dict,21 query: str,22 message_id: str,23 trace_manager: Optional[TraceQueueManager] = None,24 ) -> tuple[bool, dict, str]:25 """26 Process sensitive_word_avoidance.27 :param app_id: app id28 :param tenant_id: tenant id29 :param app_config: app config30 :param inputs: inputs31 :param query: query32 :param message_id: message id33 :param trace_manager: trace manager34 :return:35 """36 if not app_config.sensitive_word_avoidance:37 return False, inputs, query38 39 sensitive_word_avoidance_config = app_config.sensitive_word_avoidance40 moderation_type = sensitive_word_avoidance_config.type41 42 moderation_factory = ModerationFactory(43 name=moderation_type, app_id=app_id, tenant_id=tenant_id, config=sensitive_word_avoidance_config.config44 )45 46 with measure_time() as timer:47 moderation_result = moderation_factory.moderation_for_inputs(inputs, query)48 49 if trace_manager:50 trace_manager.add_trace_task(51 TraceTask(52 TraceTaskName.MODERATION_TRACE,53 message_id=message_id,54 moderation_result=moderation_result,55 inputs=inputs,56 timer=timer,57 )58 )59 60 if not moderation_result.flagged:61 return False, inputs, query62 63 if moderation_result.action == ModerationAction.DIRECT_OUTPUT:64 raise ModerationError(moderation_result.preset_response)65 elif moderation_result.action == ModerationAction.OVERRIDDEN:66 inputs = moderation_result.inputs67 query = moderation_result.query68 69 return True, inputs, query70 