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,仅提示该类的方法
原理说明
- @overload装饰器:为
connect函数定义多个类型签名,分别对应不同的枚举输入和返回类型。 - Literal类型:明确指定参数必须是枚举的某个具体实例(如
ConnectionType.POSTGRES),而不仅仅是ConnectionType类型。 - 类型检查器会根据传入的具体枚举值,匹配对应的重载签名,从而自动推断出返回对象的精确类型,无需手动给
conn变量添加注解。
注意事项
- 如果使用Python 3.7及更早版本,需要安装
typing_extensions库,并从typing_extensions导入Literal和overload。
内容的提问来源于stack exchange,提问作者mnowotka
相关产品推荐
相关产品推荐

