Python 装饰器

Python基础 2026-04-20 11
预计阅读时间:40 分钟

Python 装饰器深度解析

一、理解装饰器:从函数说起

装饰器是 Python 最强大、最优雅的特性之一。但在深入装饰器之前,我们必须先理解一个关键概念:在 Python 中,函数是一等公民。

1.1 函数是对象

def greet(name):
    return f"Hello, {name}!"

# 函数可以赋值给变量
my_function = greet
print(my_function("Alice"))  # Hello, Alice!

# 函数可以存储在数据结构中
func_list = [greet, str.upper, len]
print(func_list[0]("Bob"))   # Hello, Bob!

# 函数可以作为参数传递
def call_function(func, arg):
    print(f"调用 {func.__name__} 函数")
    return func(arg)

result = call_function(greet, "Charlie")
print(result)  # Hello, Charlie!

# 函数可以在内部定义并返回
def make_multiplier(n):
    def multiplier(x):
        return x * n
    return multiplier  # 返回函数对象

times_3 = make_multiplier(3)
print(times_3(10))  # 30
print(times_3.__name__)  # multiplier

# 函数有元数据
print(f"函数名: {greet.__name__}")
print(f"文档字符串: {greet.__doc__}")
print(f"模块: {greet.__module__}")

二、装饰器基础

2.1 什么是装饰器?

装饰器本质上是一个接受函数作为参数并返回新函数的可调用对象。它允许在不修改原函数代码的情况下,增加额外的功能。

# 最简单的装饰器形式
def simple_decorator(func):
    """一个什么都不做的装饰器"""
    def wrapper(*args, **kwargs):
        print(f"调用函数: {func.__name__}")
        return func(*args, **kwargs)
    return wrapper

# 使用装饰器
@simple_decorator
def say_hello(name):
    """问候函数"""
    return f"Hello, {name}!"

# 等价于:say_hello = simple_decorator(say_hello)

print(say_hello("World"))
# 输出:
# 调用函数: say_hello
# Hello, World!

# 查看被装饰后的函数
print(f"函数名: {say_hello.__name__}")  # wrapper (丢失了原函数名!)
print(f"文档: {say_hello.__doc__}")     # None (丢失了文档!)

2.2 保留函数元数据

from functools import wraps

def better_decorator(func):
    """使用 wraps 保留原函数的元数据"""
    @wraps(func)  # 这个装饰器会复制原函数的元数据
    def wrapper(*args, **kwargs):
        print(f"调用函数: {func.__name__}")
        return func(*args, **kwargs)
    return wrapper

@better_decorator
def say_hello(name):
    """问候函数"""
    return f"Hello, {name}!"

print(f"函数名: {say_hello.__name__}")  # say_hello (保留了!)
print(f"文档: {say_hello.__doc__}")     # 问候函数 (保留了!)

2.3 装饰器的执行时机

def decorator_with_side_effect(func):
    """带副作用的装饰器 - 演示执行时机"""
    print(f"装饰器正在装饰 {func.__name__}")  # 这会在定义时执行

    @wraps(func)
    def wrapper(*args, **kwargs):
        print(f"包装器调用 {func.__name__}")  # 这会在调用时执行
        return func(*args, **kwargs)
    return wrapper

print("开始定义函数...")

@decorator_with_side_effect
def my_function():
    print("原函数执行")

print("函数定义完成")
print("开始调用函数...")
my_function()

# 输出顺序:
# 开始定义函数...
# 装饰器正在装饰 my_function  <-- 注意:定义时就执行了
# 函数定义完成
# 开始调用函数...
# 包装器调用 my_function      <-- 调用时才执行
# 原函数执行

三、装饰器的各种形式

3.1 带参数的装饰器

# 装饰器工厂:返回真正的装饰器
def repeat(times):
    """重复执行函数的装饰器工厂"""
    def decorator(func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            result = None
            for i in range(times):
                print(f"第 {i+1} 次执行")
                result = func(*args, **kwargs)
            return result
        return wrapper
    return decorator

@repeat(3)
def greet(name):
    return f"Hello, {name}!"

print(greet("Python"))
# 输出:
# 第 1 次执行
# 第 2 次执行
# 第 3 次执行
# Hello, Python!

# 理解嵌套:repeat(3) 返回 decorator,然后 decorator(greet) 返回 wrapper
# greet = repeat(3)(greet)

# 实战:带参数的计时器
import time

def timer(unit='ms'):
    """计时装饰器,可选择时间单位"""
    def decorator(func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            start = time.perf_counter()
            result = func(*args, **kwargs)
            elapsed = time.perf_counter() - start

            if unit == 'ms':
                elapsed_time = elapsed * 1000
                unit_str = 'ms'
            elif unit == 'us':
                elapsed_time = elapsed * 1_000_000
                unit_str = 'μs'
            else:  # seconds
                elapsed_time = elapsed
                unit_str = 's'

            print(f"{func.__name__} 执行时间: {elapsed_time:.2f}{unit_str}")
            return result
        return wrapper
    return decorator

@timer(unit='ms')
def slow_function():
    time.sleep(0.1)
    return "完成"

slow_function()

3.2 类作为装饰器

# 方法1:使用 __call__ 方法
class CountCalls:
    """统计函数调用次数的装饰器"""

    def __init__(self, func):
        self.func = func
        self.count = 0
        wraps(func)(self)  # 手动应用 wraps

    def __call__(self, *args, **kwargs):
        self.count += 1
        print(f"{self.func.__name__} 已被调用 {self.count} 次")
        return self.func(*args, **kwargs)

@CountCalls
def say_hello():
    print("Hello!")

say_hello()
say_hello()
say_hello()
print(f"总调用次数: {say_hello.count}")  # 可以访问状态

# 方法2:带参数的类装饰器
class Retry:
    """重试装饰器类"""

    def __init__(self, max_attempts=3, delay=1):
        self.max_attempts = max_attempts
        self.delay = delay

    def __call__(self, func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            import time
            last_exception = None

            for attempt in range(self.max_attempts):
                try:
                    return func(*args, **kwargs)
                except Exception as e:
                    last_exception = e
                    if attempt < self.max_attempts - 1:
                        print(f"尝试 {attempt + 1} 失败,{self.delay}秒后重试...")
                        time.sleep(self.delay)

            raise last_exception

        return wrapper

@Retry(max_attempts=3, delay=0.5)
def unstable_network_call():
    """模拟不稳定的网络调用"""
    import random
    if random.random() < 0.7:
        raise ConnectionError("网络错误")
    return "成功获取数据"

# 测试
for i in range(3):
    try:
        result = unstable_network_call()
        print(f"结果: {result}")
    except ConnectionError:
        print("最终失败")

3.3 多个装饰器的叠加

def decorator_a(func):
    """装饰器 A"""
    @wraps(func)
    def wrapper(*args, **kwargs):
        print("装饰器 A 开始")
        result = func(*args, **kwargs)
        print("装饰器 A 结束")
        return result
    return wrapper

def decorator_b(func):
    """装饰器 B"""
    @wraps(func)
    def wrapper(*args, **kwargs):
        print("装饰器 B 开始")
        result = func(*args, **kwargs)
        print("装饰器 B 结束")
        return result
    return wrapper

@decorator_a
@decorator_b
def my_func():
    print("原函数执行")

print("执行顺序:")
my_func()
# 输出:
# 装饰器 A 开始
# 装饰器 B 开始
# 原函数执行
# 装饰器 B 结束
# 装饰器 A 结束

# 等价于:my_func = decorator_a(decorator_b(my_func))

# 理解装饰器的数学性质
def add_exclamation(func):
    @wraps(func)
    def wrapper(*args, **kwargs):
        return func(*args, **kwargs) + "!"
    return wrapper

def add_question(func):
    @wraps(func)
    def wrapper(*args, **kwargs):
        return func(*args, **kwargs) + "?"
    return wrapper

@add_exclamation
@add_question
def get_message():
    return "Hello"

print(get_message())  # Hello?!

@add_question
@add_exclamation
def get_message():
    return "Hello"

print(get_message())  # Hello!?
# 注意顺序的不同导致结果不同

四、装饰器的实用模式

4.1 缓存与记忆化

from functools import lru_cache
import time

# 1. 使用 functools.lru_cache(最常用)
@lru_cache(maxsize=128)
def fibonacci(n):
    """计算斐波那契数(带缓存)"""
    if n < 2:
        return n
    return fibonacci(n-1) + fibonacci(n-2)

# 测试性能
def test_fibonacci():
    start = time.perf_counter()
    result = fibonacci(35)
    elapsed = time.perf_counter() - start
    print(f"fibonacci(35) = {result}, 耗时: {elapsed:.4f}s")

    # 再次计算(从缓存获取)
    start = time.perf_counter()
    result = fibonacci(35)
    elapsed = time.perf_counter() - start
    print(f"第二次调用耗时: {elapsed:.6f}s")

test_fibonacci()

# 2. 自定义缓存装饰器(支持过期时间)
from datetime import datetime, timedelta

def cache_with_ttl(ttl_seconds=60):
    """带过期时间的缓存装饰器"""
    def decorator(func):
        cache = {}

        @wraps(func)
        def wrapper(*args, **kwargs):
            # 创建缓存键
            key = (args, tuple(sorted(kwargs.items())))

            # 检查缓存
            if key in cache:
                result, timestamp = cache[key]
                if datetime.now() - timestamp < timedelta(seconds=ttl_seconds):
                    print(f"从缓存返回: {key}")
                    return result
                else:
                    print(f"缓存已过期: {key}")
                    del cache[key]

            # 计算新值
            result = func(*args, **kwargs)
            cache[key] = (result, datetime.now())
            return result

        def clear_cache():
            """清空缓存"""
            cache.clear()
            print("缓存已清空")

        wrapper.clear_cache = clear_cache  # 添加辅助方法
        return wrapper
    return decorator

@cache_with_ttl(ttl_seconds=2)
def expensive_computation(x, y):
    """模拟耗时计算"""
    print(f"计算 {x} + {y}...")
    time.sleep(1)
    return x + y

# 测试
print(expensive_computation(3, 5))
print(expensive_computation(3, 5))  # 从缓存获取
time.sleep(3)
print(expensive_computation(3, 5))  # 缓存过期,重新计算
expensive_computation.clear_cache()

4.2 参数验证与类型检查

def validate_args(**validators):
    """参数验证装饰器"""
    def decorator(func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            # 获取函数参数名
            import inspect
            sig = inspect.signature(func)
            bound = sig.bind(*args, **kwargs)
            bound.apply_defaults()

            # 验证每个参数
            for param_name, validator in validators.items():
                if param_name in bound.arguments:
                    value = bound.arguments[param_name]
                    if not validator(value):
                        raise ValueError(
                            f"参数 '{param_name}' 的值 {value} 验证失败"
                        )

            return func(*args, **kwargs)
        return wrapper
    return decorator

# 使用示例
def is_positive(x):
    return x > 0

def is_valid_email(email):
    return '@' in email and '.' in email

def is_not_empty(s):
    return bool(s and s.strip())

@validate_args(
    age=is_positive,
    email=is_valid_email,
    name=is_not_empty
)
def register_user(name, email, age):
    """用户注册"""
    return f"用户 {name} ({email}) 年龄 {age} 注册成功"

# 测试
print(register_user("Alice", "alice@example.com", 25))

try:
    print(register_user("", "invalid-email", -5))
except ValueError as e:
    print(f"验证失败: {e}")

# 类型检查装饰器
def type_check(**expected_types):
    """类型检查装饰器"""
    def decorator(func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            import inspect
            sig = inspect.signature(func)
            bound = sig.bind(*args, **kwargs)

            for param_name, expected_type in expected_types.items():
                if param_name in bound.arguments:
                    value = bound.arguments[param_name]
                    if not isinstance(value, expected_type):
                        raise TypeError(
                            f"参数 '{param_name}' 应为 {expected_type.__name__},"
                            f"实际为 {type(value).__name__}"
                        )

            result = func(*args, **kwargs)

            # 检查返回值类型
            if 'return' in expected_types:
                if not isinstance(result, expected_types['return']):
                    raise TypeError(
                        f"返回值应为 {expected_types['return'].__name__},"
                        f"实际为 {type(result).__name__}"
                    )

            return result
        return wrapper
    return decorator

@type_check(a=int, b=int, return=int)
def add(a, b):
    return a + b

print(add(10, 20))  # 30
# add(10, "20")  # TypeError

4.3 单例模式装饰器

from functools import wraps

def singleton(cls):
    """单例模式装饰器"""
    instances = {}

    @wraps(cls)
    def get_instance(*args, **kwargs):
        if cls not in instances:
            instances[cls] = cls(*args, **kwargs)
            print(f"创建 {cls.__name__} 的新实例")
        else:
            print(f"返回 {cls.__name__} 的现有实例")
        return instances[cls]

    return get_instance

@singleton
class DatabaseConnection:
    """数据库连接类"""

    def __init__(self, host='localhost', port=3306):
        self.host = host
        self.port = port
        self.connected = False
        print(f"初始化连接: {host}:{port}")

    def connect(self):
        self.connected = True
        return f"已连接到 {self.host}:{self.port}"

    def query(self, sql):
        return f"执行查询: {sql}"

# 测试
db1 = DatabaseConnection('db.example.com', 5432)
db2 = DatabaseConnection('another.com', 9999)  # 参数被忽略
print(db1 is db2)  # True
print(db1.connect())

# 带参数的线程安全单例
import threading

def thread_safe_singleton(cls):
    """线程安全的单例装饰器"""
    instances = {}
    lock = threading.Lock()

    @wraps(cls)
    def get_instance(*args, **kwargs):
        if cls not in instances:
            with lock:
                if cls not in instances:  # 双重检查
                    instances[cls] = cls(*args, **kwargs)
        return instances[cls]

    return get_instance

4.4 路由注册装饰器

class Router:
    """简单的路由系统"""

    def __init__(self):
        self.routes = {}

    def route(self, path, methods=None):
        """路由注册装饰器"""
        if methods is None:
            methods = ['GET']

        def decorator(func):
            @wraps(func)
            def wrapper(*args, **kwargs):
                return func(*args, **kwargs)

            # 注册路由
            self.routes[path] = {
                'handler': wrapper,
                'methods': methods,
                'func_name': func.__name__
            }
            print(f"注册路由: {methods} {path} -> {func.__name__}")
            return wrapper

        return decorator

    def handle_request(self, path, method='GET'):
        """处理请求"""
        if path in self.routes:
            route_info = self.routes[path]
            if method in route_info['methods']:
                handler = route_info['handler']
                return handler()
            else:
                return f"405 Method Not Allowed: {method}"
        return "404 Not Found"

# 创建路由实例
app = Router()

@app.route('/')
def index():
    return "首页"

@app.route('/users', methods=['GET', 'POST'])
def users():
    return "用户列表"

@app.route('/about')
def about():
    return "关于我们"

# 测试路由
print("\n路由测试:")
print(f"GET / -> {app.handle_request('/')}")
print(f"GET /users -> {app.handle_request('/users')}")
print(f"POST /users -> {app.handle_request('/users', 'POST')}")
print(f"DELETE /users -> {app.handle_request('/users', 'DELETE')}")
print(f"GET /nonexistent -> {app.handle_request('/nonexistent')}")

# 查看所有注册的路由
print(f"\n已注册的路由:")
for path, info in app.routes.items():
    print(f"  {info['methods']} {path}")

五、高级装饰器技巧

5.1 装饰器的可选参数

def flexible_decorator(func=None, *, option1=None, option2=None):
    """
    既可以作为普通装饰器,也可以作为装饰器工厂
    使用方式:
    @flexible_decorator
    def func1(): ...

    @flexible_decorator(option1='value')
    def func2(): ...
    """
    def decorator(f):
        @wraps(f)
        def wrapper(*args, **kwargs):
            if option1:
                print(f"选项1: {option1}")
            if option2:
                print(f"选项2: {option2}")
            return f(*args, **kwargs)
        return wrapper

    # 如果直接传入函数,就作为装饰器
    if func is not None:
        return decorator(func)
    # 否则返回装饰器工厂
    return decorator

# 两种使用方式
@flexible_decorator
def hello1():
    return "Hello 1"

@flexible_decorator(option1="enabled", option2=42)
def hello2():
    return "Hello 2"

print(hello1())
print(hello2())

# 更优雅的实现方式
import functools

def elegant_decorator(func=None, *, param1='default1', param2='default2'):
    """更优雅的可选参数装饰器"""
    def decorator(func):
        @functools.wraps(func)
        def wrapper(*args, **kwargs):
            print(f"装饰器参数: {param1}, {param2}")
            return func(*args, **kwargs)
        return wrapper

    if func is None:
        return decorator
    return decorator(func)

5.2 类方法装饰器

# 装饰类方法需要注意 self 参数
def method_decorator(func):
    """装饰实例方法"""
    @wraps(func)
    def wrapper(self, *args, **kwargs):
        print(f"调用 {self.__class__.__name__}.{func.__name__}")
        print(f"  实例属性: {getattr(self, 'name', 'N/A')}")
        return func(self, *args, **kwargs)
    return wrapper

class MyClass:
    def __init__(self, name):
        self.name = name

    @method_decorator
    def greet(self, message):
        return f"{self.name} says: {message}"

obj = MyClass("Alice")
print(obj.greet("Hello!"))

# 类方法装饰器
def classmethod_decorator(func):
    """装饰类方法"""
    @wraps(func)
    def wrapper(cls, *args, **kwargs):
        print(f"调用类方法 {cls.__name__}.{func.__name__}")
        return func(cls, *args, **kwargs)
    return wrapper

# 静态方法装饰器
def staticmethod_decorator(func):
    """装饰静态方法"""
    @wraps(func)
    def wrapper(*args, **kwargs):
        print(f"调用静态方法 {func.__name__}")
        return func(*args, **kwargs)
    return wrapper

class Example:
    class_var = "类变量"

    @classmethod
    @classmethod_decorator
    def class_method(cls):
        return cls.class_var

    @staticmethod
    @staticmethod_decorator
    def static_method(x, y):
        return x + y

print(Example.class_method())
print(Example.static_method(10, 20))

# 注意:装饰器的顺序很重要
# @classmethod 必须放在最外层(最靠近函数定义)
# 因为 classmethod 返回的是 descriptor

5.3 带状态的装饰器

class StatefulDecorator:
    """有状态的装饰器"""

    def __init__(self, func):
        self.func = func
        self.call_count = 0
        self.results = []
        wraps(func)(self)

    def __call__(self, *args, **kwargs):
        self.call_count += 1
        result = self.func(*args, **kwargs)
        self.results.append(result)
        return result

    def get_stats(self):
        """获取统计信息"""
        return {
            'call_count': self.call_count,
            'results': self.results,
            'last_result': self.results[-1] if self.results else None
        }

    def reset(self):
        """重置状态"""
        self.call_count = 0
        self.results.clear()

@StatefulDecorator
def random_number():
    """生成随机数"""
    import random
    return random.randint(1, 100)

# 使用
for _ in range(5):
    print(f"生成: {random_number()}")

print(f"统计信息: {random_number.get_stats()}")
random_number.reset()
print(f"重置后: {random_number.get_stats()}")

# 函数式带状态装饰器
def stateful(initial_state=None):
    """创建带状态的装饰器"""
    def decorator(func):
        state = initial_state.copy() if initial_state else {}

        @wraps(func)
        def wrapper(*args, **kwargs):
            nonlocal state
            result = func(*args, **kwargs, state=state)
            return result

        wrapper.get_state = lambda: state.copy()
        wrapper.update_state = lambda updates: state.update(updates)
        wrapper.reset_state = lambda: state.clear()

        return wrapper
    return decorator

@stateful(initial_state={'total': 0, 'count': 0})
def add_number(x, state):
    """累加数字"""
    state['total'] += x
    state['count'] += 1
    return state['total']

print(add_number(10))
print(add_number(20))
print(f"状态: {add_number.get_state()}")

5.4 装饰器实现依赖注入

class Container:
    """简单的依赖注入容器"""

    def __init__(self):
        self._dependencies = {}
        self._instances = {}

    def register(self, name, factory, singleton=False):
        """注册依赖"""
        self._dependencies[name] = {
            'factory': factory,
            'singleton': singleton
        }

    def resolve(self, name):
        """解析依赖"""
        if name not in self._dependencies:
            raise KeyError(f"未注册的依赖: {name}")

        dep_info = self._dependencies[name]

        # 单例模式
        if dep_info['singleton']:
            if name not in self._instances:
                self._instances[name] = dep_info['factory']()
            return self._instances[name]

        # 每次创建新实例
        return dep_info['factory']()

    def inject(self, *dep_names):
        """依赖注入装饰器"""
        def decorator(func):
            @wraps(func)
            def wrapper(*args, **kwargs):
                # 解析依赖
                deps = [self.resolve(name) for name in dep_names]
                # 注入到函数
                return func(*deps, *args, **kwargs)
            return wrapper
        return decorator

# 使用示例
container = Container()

# 注册依赖
class Database:
    def query(self, sql):
        return f"执行查询: {sql}"

class EmailService:
    def send(self, to, subject, body):
        return f"发送邮件到 {to}: {subject}"

container.register('db', Database, singleton=True)
container.register('email', EmailService, singleton=True)

# 使用注入
@container.inject('db', 'email')
def create_user(db, email, username, user_email):
    """创建用户(自动注入 db 和 email)"""
    result = db.query(f"INSERT INTO users VALUES ('{username}')")
    email_result = email.send(user_email, "欢迎", f"欢迎 {username}")
    return f"{result}\n{email_result}"

print(create_user("alice", "alice@example.com"))

六、内置装饰器详解

6.1 @property 及其相关装饰器

class Temperature:
    """温度类,演示 property 装饰器"""

    def __init__(self, celsius=0):
        self._celsius = celsius

    @property
    def celsius(self):
        """获取摄氏度"""
        print("获取摄氏度")
        return self._celsius

    @celsius.setter
    def celsius(self, value):
        """设置摄氏度"""
        print(f"设置摄氏度为 {value}")
        if value < -273.15:
            raise ValueError("温度不能低于绝对零度")
        self._celsius = value

    @property
    def fahrenheit(self):
        """华氏度(只读属性)"""
        return self._celsius * 9/5 + 32

    @property
    def kelvin(self):
        """开尔文(只读属性)"""
        return self._celsius + 273.15

# 使用
temp = Temperature(25)
print(f"摄氏度: {temp.celsius}")
print(f"华氏度: {temp.fahrenheit}")
print(f"开尔文: {temp.kelvin}")

temp.celsius = 30
print(f"新的华氏度: {temp.fahrenheit}")

# 缓存属性值
class Circle:
    def __init__(self, radius):
        self.radius = radius
        self._area = None  # 缓存
        self._circumference = None

    @property
    def area(self):
        """计算面积(带缓存)"""
        if self._area is None:
            print("计算面积...")
            import math
            self._area = math.pi * self.radius ** 2
        return self._area

    @property
    def circumference(self):
        """计算周长(带缓存)"""
        if self._circumference is None:
            print("计算周长...")
            import math
            self._circumference = 2 * math.pi * self.radius
        return self._circumference

    @property
    def radius(self):
        return self._radius

    @radius.setter
    def radius(self, value):
        """修改半径时清空缓存"""
        self._radius = value
        self._area = None
        self._circumference = None

circle = Circle(5)
print(f"面积: {circle.area:.2f}")      # 计算
print(f"面积: {circle.area:.2f}")      # 从缓存获取
circle.radius = 10
print(f"新面积: {circle.area:.2f}")    # 重新计算

6.2 @staticmethod 和 @classmethod

class Date:
    """日期类,演示静态方法和类方法"""

    def __init__(self, year, month, day):
        self.year = year
        self.month = month
        self.day = day

    @classmethod
    def from_string(cls, date_string):
        """从字符串创建日期对象"""
        year, month, day = map(int, date_string.split('-'))
        return cls(year, month, day)

    @classmethod
    def today(cls):
        """获取今天日期"""
        import datetime
        today = datetime.date.today()
        return cls(today.year, today.month, today.day)

    @staticmethod
    def is_valid_date(year, month, day):
        """验证日期是否有效"""
        if not (1 <= month <= 12):
            return False
        if month in (4, 6, 9, 11):
            return 1 <= day <= 30
        elif month == 2:
            # 简化:不考虑闰年
            return 1 <= day <= 28
        else:
            return 1 <= day <= 31

    @staticmethod
    def days_in_month(year, month):
        """返回指定月份的天数"""
        if month in (4, 6, 9, 11):
            return 30
        elif month == 2:
            # 判断闰年
            if (year % 4 == 0 and year % 100 != 0) or (year % 400 == 0):
                return 29
            return 28
        else:
            return 31

    def __str__(self):
        return f"{self.year:04d}-{self.month:02d}-{self.day:02d}"

# 使用类方法创建实例
date1 = Date(2024, 1, 15)
date2 = Date.from_string("2024-12-25")
date3 = Date.today()

print(f"date1: {date1}")
print(f"date2: {date2}")
print(f"date3: {date3}")

# 使用静态方法
print(f"2024-02-30 有效吗? {Date.is_valid_date(2024, 2, 30)}")
print(f"2024年2月有 {Date.days_in_month(2024, 2)} 天")

# 子类继承类方法
class DateTime(Date):
    def __init__(self, year, month, day, hour=0, minute=0):
        super().__init__(year, month, day)
        self.hour = hour
        self.minute = minute

    def __str__(self):
        return f"{super().__str__()} {self.hour:02d}:{self.minute:02d}"

# 类方法会自动使用子类
dt = DateTime.from_string("2024-01-15")  # 返回 DateTime 实例
print(f"DateTime: {dt}")

6.3 @dataclass 装饰器

from dataclasses import dataclass, field, asdict, astuple
from typing import List, Optional

# 基础用法
@dataclass
class Point:
    x: float
    y: float

    def distance_from_origin(self):
        return (self.x ** 2 + self.y ** 2) ** 0.5

p1 = Point(3, 4)
p2 = Point(3, 4)
print(f"p1: {p1}")
print(f"p1 == p2: {p1 == p2}")  # True(自动实现 __eq__)
print(f"距离: {p1.distance_from_origin()}")

# 高级选项
@dataclass(order=True, frozen=False)
class Person:
    sort_index: int = field(init=False, repr=False)
    name: str
    age: int
    email: Optional[str] = None
    tags: List[str] = field(default_factory=list)

    def __post_init__(self):
        """初始化后处理"""
        self.sort_index = self.age
        if self.email is None:
            self.email = f"{self.name.lower()}@example.com"

# 测试
p1 = Person("Alice", 30, tags=["developer"])
p2 = Person("Bob", 25, "bob@company.com")
p3 = Person("Charlie", 35)

print(f"\n排序测试:")
people = sorted([p1, p2, p3])
for person in people:
    print(f"  {person}")

print(f"\n转换为字典: {asdict(p1)}")
print(f"转换为元组: {astuple(p1)}")

# 继承
@dataclass
class Employee(Person):
    department: str
    salary: float = field(default=0.0, metadata={"unit": "USD"})

    def __post_init__(self):
        super().__post_init__()
        if self.salary < 0:
            raise ValueError("工资不能为负")

emp = Employee("David", 40, "david@company.com", ["manager"], "IT", 75000)
print(f"\nEmployee: {emp}")

6.4 @contextmanager 装饰器

from contextlib import contextmanager
import time

@contextmanager
def timer(description="操作"):
    """计时上下文管理器"""
    print(f"{description} 开始...")
    start = time.perf_counter()
    try:
        yield  # 执行 with 块中的代码
    finally:
        elapsed = time.perf_counter() - start
        print(f"{description} 完成,耗时: {elapsed:.4f}秒")

# 使用
with timer("数据处理"):
    time.sleep(0.5)
    result = sum(range(1000000))
    print(f"计算结果: {result}")

@contextmanager
def temporary_attribute(obj, name, value):
    """临时修改对象属性"""
    old_value = getattr(obj, name, None)
    setattr(obj, name, value)
    try:
        yield
    finally:
        if old_value is None:
            delattr(obj, name)
        else:
            setattr(obj, name, old_value)

class Config:
    debug = False
    log_level = "INFO"

config = Config()
print(f"原始: debug={config.debug}, log_level={config.log_level}")

with temporary_attribute(config, 'debug', True):
    with temporary_attribute(config, 'log_level', 'DEBUG'):
        print(f"临时: debug={config.debug}, log_level={config.log_level}")

print(f"恢复: debug={config.debug}, log_level={config.log_level}")

@contextmanager
def ignored(*exceptions):
    """忽略指定的异常"""
    try:
        yield
    except exceptions:
        pass

# 使用
with ignored(ValueError, TypeError):
    int("not a number")  # 不会抛出异常
    print("继续执行")

# 嵌套上下文管理器
@contextmanager
def nested_context(*managers):
    """合并多个上下文管理器"""
    with ExitStack() as stack:
        yield [stack.enter_context(m) for m in managers]

七、实战案例

7.1 API 限流装饰器

import time
from collections import defaultdict, deque
from functools import wraps

class RateLimiter:
    """API 限流器"""

    def __init__(self, max_calls, time_window):
        self.max_calls = max_calls
        self.time_window = time_window
        self.calls = defaultdict(deque)

    def __call__(self, func):
        @wraps(func)
        def wrapper(*args, **kwargs):
            # 使用函数名作为标识,也可以使用 IP 或其他标识
            key = func.__name__
            now = time.time()

            # 清理过期的调用记录
            while self.calls[key] and self.calls[key][0] < now - self.time_window:
                self.calls[key].popleft()

            # 检查是否超过限制
            if len(self.calls[key]) >= self.max_calls:
                wait_time = self.calls[key][0] + self.time_window - now
                raise Exception(f"API 限流,请在 {wait_time:.1f} 秒后重试")

            # 记录本次调用
            self.calls[key].append(now)

            return func(*args, **kwargs)

        return wrapper

# 使用限流器
@RateLimiter(max_calls=3, time_window=10)
def api_endpoint(user_id):
    """模拟 API 端点"""
    return f"用户 {user_id} 的数据"

# 测试
print("API 限流测试:")
for i in range(5):
    try:
        result = api_endpoint(123)
        print(f"  请求 {i+1}: {result}")
    except Exception as e:
        print(f"  请求 {i+1}: {e}")
    time.sleep(1)

# 基于 IP 的限流
class IPRateLimiter:
    """基于 IP 的 API 限流"""

    def __init__(self, max_calls_per_minute=60):
        self.max_calls = max_calls_per_minute
        self.window = 60  # 1分钟
        self.ip_calls = defaultdict(deque)

    def __call__(self, func):
        @wraps(func)
        def wrapper(request, *args, **kwargs):
            # 假设 request 对象有 ip 属性
            ip = getattr(request, 'ip', 'unknown')
            now = time.time()

            # 清理过期记录
            while self.ip_calls[ip] and self.ip_calls[ip][0] < now - self.window:
                self.ip_calls[ip].popleft()

            if len(self.ip_calls[ip]) >= self.max_calls:
                raise Exception(f"IP {ip} 请求过于频繁")

            self.ip_calls[ip].append(now)
            return func(request, *args, **kwargs)

        return wrapper

7.2 性能分析与监控装饰器

import time
import functools
from collections import defaultdict
import statistics

class PerformanceMonitor:
    """性能监控装饰器"""

    def __init__(self):
        self.stats = defaultdict(lambda: {
            'calls': 0,
            'total_time': 0,
            'min_time': float('inf'),
            'max_time': 0,
            'times': []
        })

    def monitor(self, func=None, *, sample_size=100):
        """监控函数性能"""
        def decorator(func):
            @functools.wraps(func)
            def wrapper(*args, **kwargs):
                start = time.perf_counter()
                try:
                    return func(*args, **kwargs)
                finally:
                    elapsed = time.perf_counter() - start

                    stat = self.stats[func.__name__]
                    stat['calls'] += 1
                    stat['total_time'] += elapsed
                    stat['min_time'] = min(stat['min_time'], elapsed)
                    stat['max_time'] = max(stat['max_time'], elapsed)

                    # 保留最近的样本
                    stat['times'].append(elapsed)
                    if len(stat['times']) > sample_size:
                        stat['times'].pop(0)

            return wrapper

        if func is None:
            return decorator
        return decorator(func)

    def report(self):
        """生成性能报告"""
        print("\n=== 性能监控报告 ===")
        for func_name, stat in self.stats.items():
            print(f"\n函数: {func_name}")
            print(f"  调用次数: {stat['calls']}")
            print(f"  总耗时: {stat['total_time']:.4f}s")
            print(f"  平均耗时: {stat['total_time']/stat['calls']:.6f}s")
            print(f"  最小耗时: {stat['min_time']:.6f}s")
            print(f"  最大耗时: {stat['max_time']:.6f}s")

            if len(stat['times']) > 1:
                print(f"  标准差: {statistics.stdev(stat['times']):.6f}s")
                print(f"  中位数: {statistics.median(stat['times']):.6f}s")

    def reset(self):
        """重置统计"""
        self.stats.clear()

# 使用示例
monitor = PerformanceMonitor()

@monitor.monitor
def slow_operation(n):
    """模拟耗时操作"""
    time.sleep(n * 0.1)
    return n * n

@monitor.monitor(sample_size=50)
def quick_operation(x, y):
    """快速操作"""
    return x + y

# 模拟调用
import random
for _ in range(10):
    slow_operation(random.random())
    for _ in range(100):
        quick_operation(random.randint(1, 10), random.randint(1, 10))

monitor.report()

7.3 事务装饰器

class TransactionManager:
    """事务管理器"""

    def __init__(self):
        self.operations = []
        self.committed = False

    def add_operation(self, do_func, undo_func):
        """添加操作及其回滚函数"""
        self.operations.append((do_func, undo_func))

    def commit(self):
        """提交事务"""
        self.committed = True
        self.operations.clear()

    def rollback(self):
        """回滚事务"""
        # 逆序执行回滚操作
        for do_func, undo_func in reversed(self.operations):
            try:
                undo_func()
            except Exception as e:
                print(f"回滚失败: {e}")
        self.operations.clear()

def transactional(func):
    """事务装饰器"""
    @functools.wraps(func)
    def wrapper(*args, **kwargs):
        tm = TransactionManager()

        # 将事务管理器注入到函数中
        kwargs['transaction'] = tm

        try:
            result = func(*args, **kwargs)
            tm.commit()
            return result
        except Exception as e:
            print(f"事务失败,执行回滚: {e}")
            tm.rollback()
            raise

    return wrapper

# 模拟数据库操作
class BankAccount:
    def __init__(self, owner, balance=0):
        self.owner = owner
        self.balance = balance
        self.history = []

    def deposit(self, amount, transaction=None):
        """存款"""
        old_balance = self.balance
        self.balance += amount
        self.history.append(f"存款: +{amount}")

        if transaction:
            transaction.add_operation(
                lambda: None,  # 不需要 do,已经执行了
                lambda: self._undo_deposit(old_balance)  # 回滚操作
            )

    def withdraw(self, amount, transaction=None):
        """取款"""
        if self.balance < amount:
            raise ValueError("余额不足")

        old_balance = self.balance
        self.balance -= amount
        self.history.append(f"取款: -{amount}")

        if transaction:
            transaction.add_operation(
                lambda: None,
                lambda: self._undo_withdraw(old_balance)
            )

    def _undo_deposit(self, old_balance):
        """回滚存款"""
        self.balance = old_balance
        self.history.append("回滚: 取消存款")

    def _undo_withdraw(self, old_balance):
        """回滚取款"""
        self.balance = old_balance
        self.history.append("回滚: 取消取款")

@transactional
def transfer_money(from_account, to_account, amount, transaction=None):
    """转账(事务性)"""
    print(f"开始转账: {from_account.owner} -> {to_account.owner}, 金额: {amount}")

    from_account.withdraw(amount, transaction)
    to_account.deposit(amount, transaction)

    # 模拟可能发生的错误
    if amount > 1000:
        raise ValueError("转账金额超过限额")

    print("转账成功")

# 测试
alice = BankAccount("Alice", 2000)
bob = BankAccount("Bob", 1000)

print(f"转账前: Alice={alice.balance}, Bob={bob.balance}")

# 成功的事务
try:
    transfer_money(alice, bob, 500)
except ValueError:
    pass

print(f"第一次转账后: Alice={alice.balance}, Bob={bob.balance}")

# 失败的事务(会回滚)
try:
    transfer_money(alice, bob, 1500)
except ValueError as e:
    print(f"转账失败: {e}")

print(f"第二次转账后: Alice={alice.balance}, Bob={bob.balance}")
print(f"Alice 的历史: {alice.history}")
print(f"Bob 的历史: {bob.history}")

八、装饰器的最佳实践与陷阱

8.1 常见陷阱及解决方案

# 陷阱1:装饰器在模块加载时执行
print("\n=== 陷阱1:执行时机 ===")

def eager_decorator(func):
    print(f"装饰器立即执行: {func.__name__}")  # 导入时就会打印
    return func

@eager_decorator
def some_function():
    pass
# 即使从不调用 some_function,print 也会执行

# 解决:延迟执行
def lazy_decorator(func):
    @functools.wraps(func)
    def wrapper(*args, **kwargs):
        print(f"延迟执行: {func.__name__}")  # 只在调用时打印
        return func(*args, **kwargs)
    return wrapper

# 陷阱2:装饰器改变函数签名
def bad_decorator(func):
    def wrapper(x, y):  # 改变了原函数签名
        return func(x, y)
    return wrapper

@bad_decorator
def add(a, b, c=0):  # 原函数有3个参数
    return a + b + c

# print(add(1, 2, 3))  # TypeError: wrapper() takes 2 positional arguments

# 解决:使用 *args, **kwargs
def good_decorator(func):
    @functools.wraps(func)
    def wrapper(*args, **kwargs):
        return func(*args, **kwargs)
    return wrapper

# 陷阱3:装饰器中的闭包变量
def counter_decorator_bad():
    count = 0
    def decorator(func):
        @functools.wraps(func)
        def wrapper(*args, **kwargs):
            nonlocal count  # Python 3 需要 nonlocal
            count += 1
            print(f"调用次数: {count}")
            return func(*args, **kwargs)
        return wrapper
    return decorator

# 每个被装饰的函数会共享同一个 count
counter = counter_decorator_bad()

@counter
def func1():
    pass

@counter
def func2():
    pass

func1()  # 1
func2()  # 2(共享了)
func1()  # 3

# 解决:每个函数独立的计数器
def counter_decorator_good():
    def decorator(func):
        count = 0  # 在 decorator 内定义
        @functools.wraps(func)
        def wrapper(*args, **kwargs):
            nonlocal count
            count += 1
            print(f"{func.__name__} 调用次数: {count}")
            return func(*args, **kwargs)
        return wrapper
    return decorator

# 陷阱4:装饰器与描述符(如 property)
class MyClass:
    def __init__(self):
        self._value = 0

    @property
    @good_decorator  # 注意顺序
    def value(self):
        return self._value

# 陷阱5:递归函数被装饰后的问题
@lru_cache(maxsize=None)
def recursive_func(n):
    if n < 2:
        return n
    return recursive_func(n-1) + recursive_func(n-2)
# 注意:lru_cache 会缓存每次调用的结果,包括递归调用
# 这对于斐波那契数列是有益的优化

8.2 最佳实践总结

"""
装饰器最佳实践清单:

1. ✓ 始终使用 @functools.wraps
   - 保留原函数的元数据(__name__, __doc__ 等)
   - 使调试更容易

2. ✓ 使用 *args, **kwargs
   - 保持装饰器的通用性
   - 不改变原函数的签名

3. ✓ 明确装饰器的职责
   - 单一职责原则
   - 每个装饰器只做一件事

4. ✓ 考虑装饰器的可组合性
   - 装饰器应该可以与其他装饰器组合使用
   - 注意装饰器的顺序

5. ✓ 提供合理的默认参数
   - 使装饰器易于使用
   - 同时保持灵活性

6. ✓ 文档化装饰器
   - 说明装饰器的功能
   - 说明参数的含义
   - 提供使用示例

7. ✓ 考虑性能影响
   - 装饰器会增加函数调用开销
   - 在性能关键路径谨慎使用

8. ✓ 测试装饰器
   - 测试装饰器本身
   - 测试被装饰的函数
   - 测试边界情况

9. ✓ 避免副作用
   - 装饰器应该是纯函数
   - 避免修改全局状态

10. ✓ 使用类型提示
    - 帮助 IDE 提供更好的支持
    - 提高代码可读性
"""

# 示例:遵循最佳实践的装饰器
from typing import Callable, Any, Optional
import logging

def retry(
    max_attempts: int = 3,
    delay: float = 1.0,
    exceptions: tuple = (Exception,),
    logger: Optional[logging.Logger] = None
) -> Callable:
    """
    重试装饰器

    Args:
        max_attempts: 最大尝试次数
        delay: 重试延迟(秒)
        exceptions: 需要重试的异常类型
        logger: 日志记录器

    Returns:
        装饰后的函数

    Example:
        @retry(max_attempts=3, delay=0.5, exceptions=(ConnectionError,))
        def fetch_data():
            ...
    """
    if logger is None:
        logger = logging.getLogger(__name__)

    def decorator(func: Callable) -> Callable:
        @functools.wraps(func)
        def wrapper(*args: Any, **kwargs: Any) -> Any:
            last_exception = None

            for attempt in range(max_attempts):
                try:
                    return func(*args, **kwargs)
                except exceptions as e:
                    last_exception = e
                    if attempt < max_attempts - 1:
                        logger.warning(
                            f"{func.__name__} 第 {attempt + 1} 次尝试失败: {e},"
                            f"{delay}秒后重试..."
                        )
                        time.sleep(delay)
                    else:
                        logger.error(
                            f"{func.__name__} 所有 {max_attempts} 次尝试均失败"
                        )

            raise last_exception

        return wrapper
    return decorator

九、总结

装饰器是 Python 最具表现力的特性之一,它体现了 Python "显式优于隐式" 和 "可读性很重要" 的设计哲学。

核心概念回顾: - 一等公民函数:函数可以作为参数、返回值 - 闭包:内部函数可以访问外部函数的变量 - @ 语法糖:简化装饰器的使用 - functools.wraps:保留原函数的元数据

装饰器的应用场景: - 日志记录和性能监控 - 缓存和记忆化 - 权限验证和访问控制 - 参数验证和类型检查 - 事务管理 - 路由注册 - 依赖注入

选择建议: - 简单功能 → 函数装饰器 - 需要状态 → 类装饰器 - 可配置 → 装饰器工厂 - 复杂逻辑 → 拆分为多个简单装饰器

装饰器让 Python 代码更加优雅、模块化和可重用。掌握装饰器,你就能写出更具 Python 风格的代码,更好地理解许多流行框架(如 Flask、Django)的设计思想。


本文由 尚先生 原创,转载请注明出处。

📖相关推荐

评论

0
暂无评论,来发表第一条评论吧

发表评论

登录 后发表评论