预计阅读时间: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