chore(api): Fix several typing errors (#37248)

This commit is contained in:
chariri
2026-06-12 14:02:09 +00:00
committed by GitHub
parent ad96501e09
commit 7cf75c3cc5
18 changed files with 148 additions and 93 deletions
+5 -3
View File
@@ -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
+3 -2
View File
@@ -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
"""
+1 -1
View File
@@ -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)
+1 -1
View File
@@ -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,
+7 -3
View File
@@ -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
"""
+5 -4
View File
@@ -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