mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
chore(api): Fix several typing errors (#37248)
This commit is contained in:
@@ -11,8 +11,10 @@ from core.tools.entities.tool_entities import (
|
||||
from core.tools.errors import ToolProviderCredentialValidationError
|
||||
|
||||
|
||||
class ToolProviderController(ABC):
|
||||
def __init__(self, entity: ToolProviderEntity):
|
||||
class ToolProviderController[ToolProviderEntityT: ToolProviderEntity, ToolProviderToolT: Tool | None](ABC):
|
||||
entity: ToolProviderEntityT
|
||||
|
||||
def __init__(self, entity: ToolProviderEntityT):
|
||||
self.entity = entity
|
||||
|
||||
def get_credentials_schema(self) -> list[ProviderConfig]:
|
||||
@@ -24,7 +26,7 @@ class ToolProviderController(ABC):
|
||||
return deepcopy(self.entity.credentials_schema)
|
||||
|
||||
@abstractmethod
|
||||
def get_tool(self, tool_name: str) -> Tool:
|
||||
def get_tool(self, tool_name: str) -> ToolProviderToolT:
|
||||
"""
|
||||
returns a tool that the provider can provide
|
||||
|
||||
|
||||
@@ -21,7 +21,7 @@ from core.tools.errors import (
|
||||
from core.tools.utils.yaml_utils import load_yaml_file_cached
|
||||
|
||||
|
||||
class BuiltinToolProviderController(ToolProviderController):
|
||||
class BuiltinToolProviderController(ToolProviderController[ToolProviderEntity, BuiltinTool | None]):
|
||||
tools: list[BuiltinTool]
|
||||
|
||||
def __init__(self, **data: Any):
|
||||
@@ -163,7 +163,8 @@ class BuiltinToolProviderController(ToolProviderController):
|
||||
"""
|
||||
return self._get_builtin_tools()
|
||||
|
||||
def get_tool(self, tool_name: str) -> BuiltinTool | None: # type: ignore
|
||||
@override
|
||||
def get_tool(self, tool_name: str) -> BuiltinTool | None:
|
||||
"""
|
||||
returns the tool that the provider can provide
|
||||
"""
|
||||
|
||||
@@ -24,7 +24,7 @@ from extensions.ext_database import db
|
||||
from models.tools import ApiToolProvider
|
||||
|
||||
|
||||
class ApiToolProviderController(ToolProviderController):
|
||||
class ApiToolProviderController(ToolProviderController[ToolProviderEntity, ApiTool]):
|
||||
provider_id: str
|
||||
tenant_id: str
|
||||
tools: list[ApiTool] = Field(default_factory=list)
|
||||
|
||||
@@ -18,7 +18,7 @@ from models.tools import MCPToolProvider
|
||||
from services.tools.tools_transform_service import ToolTransformService
|
||||
|
||||
|
||||
class MCPToolProviderController(ToolProviderController):
|
||||
class MCPToolProviderController(ToolProviderController[ToolProviderEntityWithPlugin, MCPTool]):
|
||||
def __init__(
|
||||
self,
|
||||
entity: ToolProviderEntityWithPlugin,
|
||||
|
||||
@@ -9,7 +9,9 @@ from core.tools.plugin_tool.tool import PluginTool
|
||||
|
||||
|
||||
class PluginToolProviderController(BuiltinToolProviderController):
|
||||
entity: ToolProviderEntityWithPlugin
|
||||
# TODO: Split the credential/schema helpers from BuiltinToolProviderController
|
||||
# so plugin providers do not need to inherit builtin tool-loading behavior.
|
||||
entity: ToolProviderEntityWithPlugin # pyrefly: ignore[bad-override-mutable-attribute]
|
||||
tenant_id: str
|
||||
plugin_id: str
|
||||
plugin_unique_identifier: str
|
||||
@@ -46,7 +48,8 @@ class PluginToolProviderController(BuiltinToolProviderController):
|
||||
):
|
||||
raise ToolProviderCredentialValidationError("Invalid credentials")
|
||||
|
||||
def get_tool(self, tool_name: str) -> PluginTool: # type: ignore
|
||||
@override
|
||||
def get_tool(self, tool_name: str) -> PluginTool: # type: ignore[override] # pyrefly: ignore[bad-override]
|
||||
"""
|
||||
return tool with given name
|
||||
"""
|
||||
@@ -65,7 +68,8 @@ class PluginToolProviderController(BuiltinToolProviderController):
|
||||
plugin_unique_identifier=self.plugin_unique_identifier,
|
||||
)
|
||||
|
||||
def get_tools(self) -> list[PluginTool]: # type: ignore
|
||||
@override
|
||||
def get_tools(self) -> list[PluginTool]: # type: ignore[override] # pyrefly: ignore[bad-override]
|
||||
"""
|
||||
get all tools
|
||||
"""
|
||||
|
||||
@@ -3,7 +3,6 @@ from __future__ import annotations
|
||||
from collections.abc import Mapping
|
||||
from typing import override
|
||||
|
||||
from pydantic import Field
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -43,13 +42,14 @@ VARIABLE_TO_PARAMETER_TYPE_MAPPING = {
|
||||
}
|
||||
|
||||
|
||||
class WorkflowToolProviderController(ToolProviderController):
|
||||
class WorkflowToolProviderController(ToolProviderController[ToolProviderEntity, WorkflowTool | None]):
|
||||
provider_id: str
|
||||
tools: list[WorkflowTool] = Field(default_factory=list)
|
||||
tools: list[WorkflowTool] | None
|
||||
|
||||
def __init__(self, entity: ToolProviderEntity, provider_id: str):
|
||||
super().__init__(entity=entity)
|
||||
self.provider_id = provider_id
|
||||
self.tools = None
|
||||
|
||||
@classmethod
|
||||
def from_db(cls, db_provider: WorkflowToolProvider) -> WorkflowToolProviderController:
|
||||
@@ -241,7 +241,8 @@ class WorkflowToolProviderController(ToolProviderController):
|
||||
|
||||
return self.tools
|
||||
|
||||
def get_tool(self, tool_name: str) -> WorkflowTool | None: # type: ignore
|
||||
@override
|
||||
def get_tool(self, tool_name: str) -> WorkflowTool | None:
|
||||
"""
|
||||
get tool by name
|
||||
|
||||
|
||||
Reference in New Issue
Block a user