如何在Repository子类中定义兼容的create方法签名?
问题背景
我在练习用类型提示实现仓储模式,有多个SQLAlchemy模型:
from sqlalchemy.orm import MappedAsDataclass, DeclarativeBase, mapped_column from typing import Mapped, Annotated from sqlalchemy import String, primary_key class Base(MappedAsDataclass, DeclarativeBase): id: Mapped[primary_key] = mapped_column(init=False) default_string = Annotated[str, mapped_column(String(100))] class User(Base): __tablename__ = "sample" name: Mapped[default_string] class Sample(Base): __tablename__ = "sample" value: Mapped[default_string] location: Mapped[default_string]
为统一查询方式,定义了泛型仓储基类:
from abc import ABC, abstractmethod from typing import TypeVar, Generic, list from sqlalchemy.orm import Session T = TypeVar("T", bound=Base) class Repository(ABC, Generic[T]): def __init__(self, model: type[T]): self.db = SessionLocal() self.model = model def list(self, skip: int = 0, limit: int = 100): return self.db.query(self.model).offset(skip).limit(limit).all()
希望每个模型的仓储子类能重载create方法,使用符合模型的参数签名,但重载时报错“create方法签名与父类Repository不兼容”:
# 父类新增的create方法 class Repository(ABC, Generic[T]): ... def create(self, **data) -> T: instance = self.model(**data) self.db.add(instance) self.db.commit() self.db.refresh(instance) return instance # 子类重载(注:Sample模型无name字段,此处为示例笔误) class SampleRepository(Repository[Sample]): def create(self, name: str) -> Sample: return super().create(name=name)
问题原因
类型检查器报错是因为违反了里氏替换原则:父类Repository的create方法接受任意关键字参数,而子类的create只接受固定参数。如果用父类类型引用子类实例(比如repo: Repository[Sample] = SampleRepository(Sample)),调用repo.create(foo="bar")时子类方法无法处理,所以类型检查器判定签名不兼容。
解决方案
有两种可行的实现方式,既能复用父类的通用创建逻辑,又能满足子类的参数类型规范:
方案一:抽象父类create方法,抽离通用逻辑为私有方法
将父类的create定义为抽象方法,把通用的数据库操作逻辑抽成私有方法,子类实现具体参数签名并调用私有方法:
class Repository(ABC, Generic[T]): def __init__(self, model: type[T]): self.db = SessionLocal() self.model = model def list(self, skip: int = 0, limit: int = 100) -> list[T]: return self.db.query(self.model).offset(skip).limit(limit).all() # 私有通用创建逻辑,封装数据库操作 def _create_and_save(self, **data) -> T: instance = self.model(**data) self.db.add(instance) self.db.commit() self.db.refresh(instance) return instance @abstractmethod def create(self, *args, **kwargs) -> T: """子类必须实现符合模型参数的create方法""" pass # Sample仓储子类 class SampleRepository(Repository[Sample]): def create(self, value: str, location: str) -> Sample: return self._create_and_save(value=value, location=location) # User仓储子类 class UserRepository(Repository[User]): def create(self, name: str) -> User: return self._create_and_save(name=name)
这种方式完全符合类型系统要求,类型检查器不会报错,同时保证了每个子类的create参数严格对应模型的构造参数。
方案二:用TypedDict+Unpack绑定模型参数(Python3.11+)
利用Python3.11引入的Unpack类型,结合TypedDict定义每个模型的创建参数,让父类的create方法参数与模型绑定:
from typing import TypeVar, Generic, TypedDict, Unpack T = TypeVar("T", bound=Base) CreateData = TypeVar("CreateData", bound=TypedDict) class Repository(ABC, Generic[T, CreateData]): def __init__(self, model: type[T]): self.db = SessionLocal() self.model = model def list(self, skip: int = 0, limit: int = 100) -> list[T]: return self.db.query(self.model).offset(skip).limit(limit).all() def create(self, **data: Unpack[CreateData]) -> T: instance = self.model(**data) self.db.add(instance) self.db.commit() self.db.refresh(instance) return instance # 定义Sample的创建参数TypedDict class SampleCreateData(TypedDict): value: str location: str # 定义User的创建参数TypedDict class UserCreateData(TypedDict): name: str # 子类继承时指定模型和对应的TypedDict class SampleRepository(Repository[Sample, SampleCreateData]): pass class UserRepository(Repository[User, UserCreateData]): pass # 使用示例 sample_repo = SampleRepository(Sample) sample_repo.create(value="test", location="beijing") # 类型检查会自动提示正确参数 user_repo = UserRepository(User) user_repo.create(name="Alice")
这种方式无需子类重载create方法,直接通过泛型参数绑定参数类型,既能复用父类逻辑,又能获得严格的类型提示。
总结
这两种方案都能解决类型兼容问题,并非类型系统要求过高,而是需要遵循类型检查的核心原则(如里氏替换)来设计代码结构。根据Python版本和项目需求选择合适的方案即可。
内容的提问来源于stack exchange,提问作者tutuca

