Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
tool_label_manager.py99 linesDownload Raw Back to tools
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