本文目录导读:

我来分享几个Python代码重构的实用案例,重点介绍如何封装重构逻辑。
函数提取与封装
原始代码
# 混乱的业务逻辑
def process_order(order_data):
# 验证订单
if not order_data.get('user_id'):
raise ValueError("用户ID不能为空")
if not order_data.get('items'):
raise ValueError("订单项不能为空")
if len(order_data['items']) == 0:
raise ValueError("订单项不能为空")
# 计算价格
total = 0
for item in order_data['items']:
price = item['price'] * item['quantity']
total += price
# 应用折扣
if total > 1000:
total = total * 0.9
elif total > 500:
total = total * 0.95
# 记录日志
print(f"订单处理完成,总价: {total}")
return total
重构后的代码
class OrderProcessor:
"""订单处理器,封装订单处理逻辑"""
def __init__(self):
self.logger = Logger()
def process_order(self, order_data: dict) -> float:
"""处理订单"""
self._validate_order(order_data)
total = self._calculate_total(order_data['items'])
total = self._apply_discount(total)
self._log_result(total)
return total
def _validate_order(self, order_data: dict):
"""验证订单数据"""
if not order_data.get('user_id'):
raise ValueError("用户ID不能为空")
items = order_data.get('items', [])
if not items:
raise ValueError("订单项不能为空")
def _calculate_total(self, items: list) -> float:
"""计算订单总价"""
return sum(item['price'] * item['quantity'] for item in items)
def _apply_discount(self, total: float) -> float:
"""应用折扣策略"""
discount_strategy = DiscountStrategy()
return discount_strategy.apply_discount(total)
def _log_result(self, total: float):
"""记录处理结果"""
self.logger.info(f"订单处理完成,总价: {total}")
class DiscountStrategy:
"""折扣策略封装"""
def apply_discount(self, total: float) -> float:
"""根据总价应用不同折扣"""
if total > 1000:
return total * 0.9
elif total > 500:
return total * 0.95
return total
class Logger:
"""日志记录器封装"""
@staticmethod
def info(message: str):
"""记录信息日志"""
print(f"[INFO] {message}")
数据验证封装
原始代码
# 散落的验证逻辑
def register_user(username, email, age, password):
if len(username) < 3:
return "用户名太短"
if len(username) > 20:
return "用户名太长"
if '@' not in email:
return "邮箱格式错误"
if age < 18:
return "年龄不符合要求"
if len(password) < 8:
return "密码太短"
# ... 更多业务逻辑
重构后的代码
from dataclasses import dataclass
from typing import List, Optional
@dataclass
class ValidationResult:
"""验证结果封装"""
is_valid: bool
errors: List[str]
@classmethod
def success(cls):
return cls(is_valid=True, errors=[])
@classmethod
def failure(cls, errors: List[str]):
return cls(is_valid=False, errors=errors)
class UserValidator:
"""用户数据验证器"""
def validate_registration(self, user_data: dict) -> ValidationResult:
"""验证用户注册数据"""
errors = []
errors.extend(self._validate_username(user_data.get('username', '')))
errors.extend(self._validate_email(user_data.get('email', '')))
errors.extend(self._validate_age(user_data.get('age', 0)))
errors.extend(self._validate_password(user_data.get('password', '')))
if errors:
return ValidationResult.failure(errors)
return ValidationResult.success()
def _validate_username(self, username: str) -> List[str]:
"""验证用户名"""
errors = []
if len(username) < 3:
errors.append("用户名至少需要3个字符")
if len(username) > 20:
errors.append("用户名不能超过20个字符")
if not username.isalnum():
errors.append("用户名只能包含字母和数字")
return errors
def _validate_email(self, email: str) -> List[str]:
"""验证邮箱"""
errors = []
if '@' not in email:
errors.append("邮箱必须包含@符号")
if '.' not in email.split('@')[-1]:
errors.append("邮箱域名格式不正确")
return errors
def _validate_age(self, age: int) -> List[str]:
"""验证年龄"""
errors = []
if age < 18:
errors.append("年龄必须大于18岁")
if age > 120:
errors.append("年龄似乎不太现实")
return errors
def _validate_password(self, password: str) -> List[str]:
"""验证密码"""
errors = []
if len(password) < 8:
errors.append("密码至少需要8个字符")
if not any(c.isupper() for c in password):
errors.append("密码需要包含大写字母")
if not any(c.islower() for c in password):
errors.append("密码需要包含小写字母")
if not any(c.isdigit() for c in password):
errors.append("密码需要包含数字")
return errors
# 使用示例
class UserRegistration:
"""用户注册处理器"""
def __init__(self):
self.validator = UserValidator()
def register(self, user_data: dict) -> dict:
"""处理用户注册"""
validation_result = self.validator.validate_registration(user_data)
if not validation_result.is_valid:
return {
'success': False,
'errors': validation_result.errors
}
# 实际的注册逻辑
return {
'success': True,
'message': '注册成功'
}
策略模式重构
原始代码
# 大量的条件判断
def calculate_shipping_cost(order, shipping_type):
if shipping_type == 'standard':
cost = order.weight * 0.5
if order.weight > 10:
cost += 5
elif shipping_type == 'express':
cost = order.weight * 1.0
cost += 10
elif shipping_type == 'overnight':
cost = order.weight * 2.0
cost += 20
if order.urgent:
cost += 15
return cost
重构后的代码
from abc import ABC, abstractmethod
from dataclasses import dataclass
@dataclass
class Order:
"""订单数据封装"""
weight: float
urgent: bool = False
class ShippingStrategy(ABC):
"""运费策略基类"""
@abstractmethod
def calculate(self, order: Order) -> float:
"""计算运费"""
pass
class StandardShipping(ShippingStrategy):
"""标准配送策略"""
def calculate(self, order: Order) -> float:
cost = order.weight * 0.5
if order.weight > 10:
cost += 5
return cost
class ExpressShipping(ShippingStrategy):
"""快递配送策略"""
def calculate(self, order: Order) -> float:
return order.weight * 1.0 + 10
class OvernightShipping(ShippingStrategy):
"""次日达配送策略"""
def calculate(self, order: Order) -> float:
cost = order.weight * 2.0 + 20
if order.urgent:
cost += 15
return cost
class ShippingCalculator:
"""运费计算器"""
def __init__(self):
self._strategies = {
'standard': StandardShipping(),
'express': ExpressShipping(),
'overnight': OvernightShipping()
}
def calculate(self, order: Order, shipping_type: str) -> float:
"""计算运费"""
strategy = self._strategies.get(shipping_type)
if not strategy:
raise ValueError(f"不支持的配送类型: {shipping_type}")
return strategy.calculate(order)
def add_strategy(self, name: str, strategy: ShippingStrategy):
"""动态添加策略"""
self._strategies[name] = strategy
# 使用示例
calculator = ShippingCalculator()
order = Order(weight=15, urgent=True)
cost = calculator.calculate(order, 'overnight')
print(f"运费: {cost}元")
数据管道封装
原始代码
# 混杂的数据处理逻辑
def process_data(data):
# 清洗
cleaned = []
for item in data:
if item is not None and item != '':
if isinstance(item, str):
item = item.strip()
cleaned.append(item)
# 转换
transformed = []
for item in cleaned:
if isinstance(item, str):
item = item.lower()
transformed.append(item)
# 过滤
filtered = []
for item in transformed:
if isinstance(item, str) and len(item) > 3:
filtered.append(item)
elif isinstance(item, (int, float)) and item > 0:
filtered.append(item)
return filtered
重构后的代码
from typing import Any, List, Callable
from functools import reduce
class DataPipeline:
"""数据处理管道"""
def __init__(self):
self._stages: List[Callable] = []
def add_stage(self, stage: Callable) -> 'DataPipeline':
"""添加处理阶段"""
self._stages.append(stage)
return self
def process(self, data: Any) -> Any:
"""执行管道处理"""
return reduce(lambda d, stage: stage(d), self._stages, data)
class DataProcessor:
"""数据处理器,提供各种处理函数"""
@staticmethod
def clean_data(data: List) -> List:
"""数据清洗"""
return [
item.strip() if isinstance(item, str) and item.strip()
else item
for item in data
if item is not None
]
@staticmethod
def transform_data(data: List) -> List:
"""数据转换"""
return [
item.lower() if isinstance(item, str)
else item
for item in data
]
@staticmethod
def filter_data(data: List) -> List:
"""数据过滤"""
result = []
for item in data:
if isinstance(item, str) and len(item) > 3:
result.append(item)
elif isinstance(item, (int, float)) and item > 0:
result.append(item)
return result
@staticmethod
def deduplicate_data(data: List) -> List:
"""数据去重"""
seen = set()
result = []
for item in data:
key = item if isinstance(item, (str, int, float)) else str(item)
if key not in seen:
seen.add(key)
result.append(item)
return result
@staticmethod
def sort_data(data: List, reverse: bool = False) -> List:
"""数据排序"""
return sorted(data, reverse=reverse)
# 使用示例
pipeline = DataPipeline()
processor = DataProcessor()
pipeline.add_stage(processor.clean_data)
pipeline.add_stage(processor.transform_data)
pipeline.add_stage(processor.filter_data)
pipeline.add_stage(processor.deduplicate_data)
pipeline.add_stage(processor.sort_data)
# 特定场景的自定义管道
custom_pipeline = DataPipeline()
custom_pipeline.add_stage(processor.clean_data)
custom_pipeline.add_stage(processor.filter_data)
# 处理数据
raw_data = [" Hello ", "WORLD", "", None, 123, " ", "hello", -5, 10]
result = pipeline.process(raw_data)
print(f"处理结果: {result}")
重构工具建议
常用重构工具
-
IDE内置工具
- PyCharm/IntelliJ IDEA: Extract Method, Extract Class, Rename
- VS Code + Python扩展: Code Actions
-
专门的重构库
# 使用 rope 进行自动化重构 import rope.base.project from rope.refactor.extract import ExtractMethod
project = rope.base.project.Project('.') resource = project.find_module('my_module') extracted = ExtractMethod(project, resource, start_offset, end_offset) changes = extracted.get_changes('new_method_name') project.do(changes)
3. **代码质量检查**
```bash
# 安装 pylint 和 autopep8
pip install pylint autopep8
# 检查代码质量
pylint my_module.py
# 自动格式化
autopep8 --in-place --aggressive my_module.py
- 单一职责原则: 每个类/函数只做一件事
- 开闭原则: 对扩展开放,对修改关闭
- 依赖倒置: 依赖抽象,不依赖具体实现
- 接口隔离: 最小化接口依赖
- DRY原则: 不要重复自己
通过这些案例,你可以看到好的封装可以:
- 提高代码可读性
- 增强可维护性
- 提升测试性
- 便于复用
- 降低耦合度
选择哪种封装方式取决于具体场景,关键是要保持代码的清晰和可维护性。