Python重构工具案例如何封装代码重构

wen python案例 28

本文目录导读:

Python重构工具案例如何封装代码重构

  1. 函数提取与封装
  2. 数据验证封装
  3. 策略模式重构
  4. 数据管道封装
  5. 重构工具建议

我来分享几个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}")

重构工具建议

常用重构工具

  1. IDE内置工具

    • PyCharm/IntelliJ IDEA: Extract Method, Extract Class, Rename
    • VS Code + Python扩展: Code Actions
  2. 专门的重构库

    # 使用 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
  1. 单一职责原则: 每个类/函数只做一件事
  2. 开闭原则: 对扩展开放,对修改关闭
  3. 依赖倒置: 依赖抽象,不依赖具体实现
  4. 接口隔离: 最小化接口依赖
  5. DRY原则: 不要重复自己

通过这些案例,你可以看到好的封装可以:

  • 提高代码可读性
  • 增强可维护性
  • 提升测试性
  • 便于复用
  • 降低耦合度

选择哪种封装方式取决于具体场景,关键是要保持代码的清晰和可维护性。

抱歉,评论功能暂时关闭!