You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Python中如何让类型检查器自动推断工厂方法返回的具体类型?

解决方案:使用函数重载 + Literal 类型

要让类型检查器(mypy/IntelliJ)根据传入的枚举值自动推断返回的具体连接类型,无需手动给conn变量加类型注解,可以通过函数重载(@overload)结合Literal类型实现。

修改后的代码示例

import abc
import enum
import typing
from typing import Literal, overload


class BaseConnection(abc.ABC):
    @abc.abstractmethod
    def sql(self, query: str) -> typing.List[typing.Any]:
        ...


class PostgresConnection(BaseConnection):
    def sql(self, query: str) -> typing.List[typing.Any]:
        return "This is a postgres result".split()

    def only_postgres_things(self):
        pass


class MySQLConnection(BaseConnection):
    def sql(self, query: str) -> typing.List[typing.Any]:
        return "This is a mysql result".split()

    def only_mysql_things(self):
        pass


class ConnectionType(enum.Enum):
    POSTGRES = 1
    MYSQL = 2


@overload
def connect(conn_type: Literal[ConnectionType.POSTGRES]) -> PostgresConnection:
    ...


@overload
def connect(conn_type: Literal[ConnectionType.MYSQL]) -> MySQLConnection:
    ...


def connect(conn_type: ConnectionType) -> typing.Union[PostgresConnection, MySQLConnection]:
    if conn_type is ConnectionType.POSTGRES:
        return PostgresConnection()
    if conn_type is ConnectionType.MYSQL:
        return MySQLConnection()


conn = connect(ConnectionType.POSTGRES)
conn.only_postgres_things()  # 类型检查器会准确推断出conn是PostgresConnection,仅提示该类的方法

原理说明

  1. @overload装饰器:为connect函数定义多个类型签名,分别对应不同的枚举输入和返回类型。
  2. Literal类型:明确指定参数必须是枚举的某个具体实例(如ConnectionType.POSTGRES),而不仅仅是ConnectionType类型。
  3. 类型检查器会根据传入的具体枚举值,匹配对应的重载签名,从而自动推断出返回对象的精确类型,无需手动给conn变量添加注解。

注意事项

  • 如果使用Python 3.7及更早版本,需要安装typing_extensions库,并从typing_extensions导入Literal和overload。

内容的提问来源于stack exchange,提问作者mnowotka

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.03 17:51:30