--- doc_type: short full_text: sources/02_Inheritance.md --- # 02_Inheritance 总结 ## 核心主题 本文介绍 Python 中的继承机制,以及如何用继承编写可扩展、可定制的程序。重点不仅在语法本身,还在于通过继承定义稳定接口,让应用代码与具体实现解耦,形成可插拔的设计。相关主题可连接到 python inheritance、object oriented programming、polymorphism、extensible design。 ## 继承基础 继承用于基于已有类创建更专门的新类: ```python class Parent: ... class Child(Parent): ... ``` 其中: - `Child` 是派生类或子类。 - `Parent` 是基类或父类。 - 父类写在类名后的括号中。 继承的核心用途是扩展已有代码,包括: - 添加新方法。 - 重定义已有方法。 - 为实例添加新属性。 例如已有 `Stock` 类: ```python class Stock: def __init__(self, name, shares, price): self.name = name self.shares = shares self.price = price def cost(self): return self.shares * self.price def sell(self, nshares): self.shares -= nshares ``` 可以通过继承创建 `MyStock`,添加新行为: ```python class MyStock(Stock): def panic(self): self.sell(self.shares) ``` 也可以重定义已有方法: ```python class MyStock(Stock): def cost(self): return 1.25 * self.shares * self.price ``` 重定义的方法会替代父类中的同名方法,而其他未重定义的方法仍然来自父类。 ## 方法覆盖与 `super()` 当子类希望扩展父类方法,而不是完全替换它时,应使用 `super()` 调用父类版本: ```python class MyStock(Stock): def cost(self): actual_cost = super().cost() return 1.25 * actual_cost ``` `super()` 表示“调用继承链中的上一个实现”。这让子类可以复用父类逻辑,并在其基础上添加额外行为。 在 Python 2 中写法更繁琐: ```python actual_cost = super(MyStock, self).cost() ``` ## `__init__` 与继承 如果子类重定义了 `__init__()`,通常必须显式调用父类的 `__init__()`,否则父类负责初始化的属性不会被创建: ```python class MyStock(Stock): def __init__(self, name, shares, price, factor): super().__init__(name, shares, price) self.factor = factor def cost(self): return self.factor * super().cost() ``` 这体现了继承中的一个常见模式: 1. 父类初始化通用状态。 2. 子类通过 `super().__init__()` 复用父类初始化。 3. 子类再初始化自身新增的状态。 ## 继承的用途 继承有两类常见用途。 ### 1. 表达类型层次结构 例如: ```python class Shape: ... class Circle(Shape): ... class Rectangle(Shape): ... ``` 这表达了“Circle 是一种 Shape”的关系,即典型的 is-a 关系。 可以使用 `isinstance()` 检查实例是否属于父类类型: ```python c = Circle(4.0) isinstance(c, Shape) # True ``` 重要原则:理想情况下,凡是能处理父类实例的代码,也应该能处理子类实例。这与 polymorphism 和面向对象替换原则相关。 ### 2. 编写可扩展代码 更实用的用途是框架式扩展。例如框架提供一个基类,用户继承它并重写部分方法: ```python class CustomHandler(TCPHandler): def handle_request(self): ... ``` 父类包含通用逻辑,子类只负责定制特定行为。这是很多库和框架使用继承的主要原因。 ## `object` 基类 Python 中所有类最终都继承自 `object`。 有时会看到: ```python class Shape(object): ... ``` 在现代 Python 中,即使不显式写 `object`,类也会隐式继承自 `object`。显式写法主要是 Python 2 时代遗留下来的习惯。 ## 多重继承 Python 允许一个类同时继承多个父类: ```python class Mother: ... class Father: ... class Child(Mother, Father): ... ``` `Child` 会继承两个父类的功能。但多重继承涉及复杂的方法解析顺序等细节,文中提醒:除非清楚自己在做什么,否则不要轻易使用。 ## 练习主题:用继承解决可扩展输出格式问题 练习部分围绕 `report.py` 中的 `print_report()` 函数展开。原始函数只能输出固定的纯文本表格: ```python def print_report(reportdata): headers = ('Name','Shares','Price','Change') print('%10s %10s %10s %10s' % headers) print(('-'*10 + ' ')*len(headers)) for row in reportdata: print('%10s %10d %10.2f %10.2f' % row) ``` 问题是:如果希望支持纯文本、HTML、CSV、XML 等多种输出格式,把所有逻辑都写进一个巨大函数会导致代码难以维护。继承提供了更好的可扩展方案。 ## 抽象基类:`TableFormatter` 练习首先要求创建 `tableformat.py`,定义一个表格格式化器基类: ```python class TableFormatter: def headings(self, headers): ''' Emit the table headings. ''' raise NotImplementedError() def row(self, rowdata): ''' Emit a single row of table data. ''' raise NotImplementedError() ``` 这个类本身不实现具体功能,而是规定接口: - `headings(headers)`:输出表头。 - `row(rowdata)`:输出一行数据。 它相当于一个“设计规范”或抽象基类。具体格式化器通过继承它并实现这些方法。 随后 `print_report()` 被改写为接收一个 formatter 对象: ```python def print_report(reportdata, formatter): formatter.headings(['Name','Shares','Price','Change']) for name, shares, price, change in reportdata: rowdata = [ name, str(shares), f'{price:0.2f}', f'{change:0.2f}' ] formatter.row(rowdata) ``` 这样,`print_report()` 不再关心输出格式,只依赖统一接口。这是 loose coupling 的体现。 ## 具体格式化器实现 ### 纯文本格式:`TextTableFormatter` ```python class TextTableFormatter(TableFormatter): ''' Emit a table in plain-text format ''' def headings(self, headers): for h in headers: print(f'{h:>10s}', end=' ') print() print(('-'*10 + ' ')*len(headers)) def row(self, rowdata): for d in rowdata: print(f'{d:>10s}', end=' ') print() ``` 它产生与原始程序相同的固定宽度表格输出。 ### CSV 格式:`CSVTableFormatter` ```python class CSVTableFormatter(TableFormatter): ''' Output portfolio data in CSV format. ''' def headings(self, headers): print(','.join(headers)) def row(self, rowdata): print(','.join(rowdata)) ``` 它通过逗号连接字段,输出 CSV 格式。 ### HTML 格式:`HTMLTableFormatter` 练习要求实现一个 HTML 表格行格式化器,输出形式类似: ```html