Underground-Digital/Workflow-Engine
0
1from core.tools.entities.values import default_tool_label_name_list2from core.tools.provider.api_tool_provider import ApiToolProviderController3from core.tools.provider.builtin_tool_provider import BuiltinToolProviderController4from core.tools.provider.tool_provider import ToolProviderController5from core.tools.provider.workflow_tool_provider import WorkflowToolProviderController6from extensions.ext_database import db7from models.tools import ToolLabelBinding8 9 10class ToolLabelManager:11 @classmethod12 def filter_tool_labels(cls, tool_labels: list[str]) -> list[str]:13 """14 Filter tool labels15 """16 tool_labels = [label for label in tool_labels if label in default_tool_label_name_list]17 return list(set(tool_labels))18 19 @classmethod20 def update_tool_labels(cls, controller: ToolProviderController, labels: list[str]):21 """22 Update tool labels23 """24 labels = cls.filter_tool_labels(labels)25 26 if isinstance(controller, ApiToolProviderController | WorkflowToolProviderController):27 provider_id = controller.provider_id28 else:29 raise ValueError("Unsupported tool type")30 31 # delete old labels32 db.session.query(ToolLabelBinding).filter(ToolLabelBinding.tool_id == provider_id).delete()33 34 # insert new labels35 for label in labels:36 db.session.add(37 ToolLabelBinding(38 tool_id=provider_id,39 tool_type=controller.provider_type.value,40 label_name=label,41 )42 )43 44 db.session.commit()45 46 @classmethod47 def get_tool_labels(cls, controller: ToolProviderController) -> list[str]:48 """49 Get tool labels50 """51 if isinstance(controller, ApiToolProviderController | WorkflowToolProviderController):52 provider_id = controller.provider_id53 elif isinstance(controller, BuiltinToolProviderController):54 return controller.tool_labels55 else:56 raise ValueError("Unsupported tool type")57 58 labels: list[ToolLabelBinding] = (59 db.session.query(ToolLabelBinding.label_name)60 .filter(61 ToolLabelBinding.tool_id == provider_id,62 ToolLabelBinding.tool_type == controller.provider_type.value,63 )64 .all()65 )66 67 return [label.label_name for label in labels]68 69 @classmethod70 def get_tools_labels(cls, tool_providers: list[ToolProviderController]) -> dict[str, list[str]]:71 """72 Get tools labels73 74 :param tool_providers: list of tool providers75 76 :return: dict of tool labels77 :key: tool id78 :value: list of tool labels79 """80 if not tool_providers:81 return {}82 83 for controller in tool_providers:84 if not isinstance(controller, ApiToolProviderController | WorkflowToolProviderController):85 raise ValueError("Unsupported tool type")86 87 provider_ids = [controller.provider_id for controller in tool_providers]88 89 labels: list[ToolLabelBinding] = (90 db.session.query(ToolLabelBinding).filter(ToolLabelBinding.tool_id.in_(provider_ids)).all()91 )92 93 tool_labels = {label.tool_id: [] for label in labels}94 95 for label in labels:96 tool_labels[label.tool_id].append(label.label_name)97 98 return tool_labels99 