1. Python函数基础与核心语法解析
函数作为Python编程中最基础也最重要的代码组织单元,每个Python开发者都需要深入掌握其运作机制。我们先从最基础的函数定义开始:
python复制def calculate_area(width, height):
"""计算矩形面积"""
return width * height
这个简单示例包含了函数定义的几个关键要素:def关键字、函数名、参数列表、文档字符串和return语句。但实际开发中我们遇到的场景要复杂得多。
1.1 函数参数的高级用法
Python的函数参数处理非常灵活,这既是优势也可能成为陷阱。以下是几种典型参数传递方式:
位置参数与关键字参数:
python复制def connect(host, port, timeout=10):
print(f"Connecting to {host}:{port} with timeout {timeout}s")
# 位置参数调用
connect("example.com", 8080)
# 关键字参数调用
connect(port=3306, host="db.example.com")
# 混合调用(位置参数必须在关键字参数前)
connect("cache.example.com", timeout=5, port=6379)
可变参数:
python复制def log_messages(*messages):
for msg in messages:
print(f"[LOG] {msg}")
log_messages("System started", "Loading modules", "Ready for connections")
关键字可变参数:
python复制def build_url(base, **query_params):
url = base + "?"
for key, value in query_params.items():
url += f"{key}={value}&"
return url.rstrip("&")
print(build_url("https://api.example.com", page=1, limit=20, sort="desc"))
重要提示:参数定义的顺序必须遵循:位置参数 -> 默认参数 -> *args -> **kwargs,否则会引发语法错误。
1.2 返回值处理技巧
Python函数可以返回任意类型的值,包括多个值(实际上是返回元组)。一些实用的返回值模式:
多值返回与解包:
python复制def analyze_data(data):
min_val = min(data)
max_val = max(data)
avg_val = sum(data)/len(data)
return min_val, max_val, avg_val
minimum, maximum, average = analyze_data([10, 20, 30, 40])
返回函数(闭包):
python复制def make_multiplier(factor):
def multiplier(x):
return x * factor
return multiplier
double = make_multiplier(2)
print(double(5)) # 输出10
返回None的常见情况:
- 函数没有显式return语句
- 只写了return没有跟值
- 实际需要表示"无结果"的场景
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 下划线在Python函数中的妙用
Python中的下划线(_)在不同位置有不同的语义含义,这是Python特有的语法糖。
2.1 单下划线的多种角色
作为临时变量:
python复制for _ in range(10):
print("Hello") # _表示我们不关心循环变量
在解释器中存储结果:
python复制>>> 3 + 4
7
>>> _ + 2 # 使用上一个结果
9
数字分隔符(Python 3.6+):
python复制big_number = 1_000_000 # 更易读,实际值为1000000
2.2 前置单下划线:约定私有
python复制class MyClass:
def __init__(self):
self._internal_cache = {} # 提示这是内部使用
def _helper_method(self): # 提示这是内部方法
pass
虽然Python没有真正的私有变量,但前置单下划线是一种约定,表示"这是内部实现细节,外部不要直接访问"。
2.3 前置双下划线:名称改写
python复制class Test:
def __init__(self):
self.__secret = 42 # 会被改写为_Test__secret
t = Test()
print(t.__dict__) # 可以看到'_Test__secret': 42
这种改写是为了避免子类中的命名冲突,实际开发中应谨慎使用。
2.4 前后双下划线:魔法方法
python复制class Vector:
def __init__(self, x, y):
self.x = x
self.y = y
def __add__(self, other): # 实现+运算符
return Vector(self.x + other.x, self.y + other.y)
Python中有大量这样的特殊方法,如__init__、__str__、__len__等,它们定义了类的各种行为。
3. 变量作用域:global与nonlocal详解
理解变量作用域是写出可靠Python代码的关键。Python有四种作用域:
- 局部(Local) - 函数内部
- 闭包(Enclosing) - 嵌套函数的外层函数
- 全局(Global) - 模块级别
- 内置(Built-in) - Python内置名称
3.1 global关键字实战
python复制count = 0 # 全局变量
def increment():
global count # 声明使用全局count
count += 1
increment()
print(count) # 输出1
常见使用场景:
- 在函数内修改模块级配置
- 维护全局状态(如计数器)
- 跨函数共享数据
注意:过度使用global通常意味着设计有问题,应考虑使用类或返回值替代。
3.2 nonlocal解决闭包问题
python复制def outer():
x = 10
def inner():
nonlocal x # 引用外层函数的x
x += 1
return x
return inner
f = outer()
print(f()) # 11
print(f()) # 12
nonlocal的典型应用:
- 实现闭包函数记住状态
- 在嵌套函数中修改外层变量
- 实现装饰器时保持状态
3.3 作用域查找规则
Python按照LEGB规则查找变量:
- 先在局部(Local)作用域查找
- 然后在闭包(Enclosing)作用域查找
- 接着在全局(Global)作用域查找
- 最后在内置(Built-in)作用域查找
python复制len = "global len" # 覆盖内置len
def test():
len = "local len"
print(len) # 输出"local len"
test()
print(len) # 输出"global len"
print(len([1,2,3])) # TypeError: 'str'对象不可调用
4. 函数最佳实践与常见陷阱
4.1 函数设计原则
-
单一职责原则:一个函数只做一件事
- 反例:
process_data_and_save_to_file() - 正例:
clean_data(),validate_data(),save_to_file()
- 反例:
-
合理参数数量:一般不超过5个
- 过多参数考虑使用对象或字典封装
- 示例重构:
python复制# 重构前 def create_user(name, email, password, age, address, phone): pass # 重构后 def create_user(user_data: dict): pass
-
避免副作用:函数应该只通过返回值与外界通信
- 反例:修改全局变量、修改输入参数
- 例外:类方法可以修改实例状态
4.2 常见错误排查
可变默认参数陷阱:
python复制def add_item(item, items=[]): # 默认值在定义时计算一次
items.append(item)
return items
print(add_item(1)) # [1]
print(add_item(2)) # [1, 2] 不是预期的[2]
正确做法:
python复制def add_item(item, items=None):
if items is None:
items = []
items.append(item)
return items
作用域混淆:
python复制x = 10
def confuse():
print(x) # 这里会报UnboundLocalError
x = 20
confuse()
修正方法:
python复制x = 10
def dont_confuse():
global x
print(x) # 正常输出10
x = 20
文档字符串规范:
python复制def calculate_stats(data):
"""计算数据的统计指标
Args:
data (list): 包含数值的列表
Returns:
dict: 包含min, max, avg的字典
Raises:
ValueError: 如果输入数据为空
"""
if not data:
raise ValueError("数据不能为空")
return {
"min": min(data),
"max": max(data),
"avg": sum(data)/len(data)
}
4.3 性能优化技巧
-
使用局部变量:访问局部变量比全局变量快
python复制def slow(): global big_data for item in big_data: process(item) def fast(): local_data = big_data # 复制到局部变量 for item in local_data: process(item) -
避免不必要的函数调用:
python复制# 不佳 for i in range(len(data)): process(data[i]) # 更好 for item in data: process(item) -
考虑使用生成器处理大数据集:
python复制def read_large_file(file_path): with open(file_path) as f: for line in f: yield process_line(line) for result in read_large_file("huge.log"): save_result(result)
5. 函数式编程在Python中的应用
虽然Python不是纯函数式语言,但它提供了一些函数式编程特性:
5.1 高阶函数
map/filter/reduce:
python复制numbers = [1, 2, 3, 4, 5]
# map示例
squares = list(map(lambda x: x**2, numbers))
# filter示例
evens = list(filter(lambda x: x % 2 == 0, numbers))
# reduce示例
from functools import reduce
product = reduce(lambda x, y: x * y, numbers)
sorted的key参数:
python复制users = [{"name": "Alice", "age": 25}, {"name": "Bob", "age": 30}]
sorted_users = sorted(users, key=lambda u: u["age"])
5.2 闭包与装饰器
计时装饰器:
python复制import time
def timer(func):
def wrapper(*args, **kwargs):
start = time.time()
result = func(*args, **kwargs)
end = time.time()
print(f"{func.__name__} took {end-start:.2f} seconds")
return result
return wrapper
@timer
def expensive_operation():
time.sleep(2)
expensive_operation()
带参数的装饰器:
python复制def retry(max_attempts=3):
def decorator(func):
def wrapper(*args, **kwargs):
attempts = 0
while attempts < max_attempts:
try:
return func(*args, **kwargs)
except Exception as e:
attempts += 1
print(f"Attempt {attempts} failed: {e}")
if attempts == max_attempts:
raise
return wrapper
return decorator
@retry(max_attempts=5)
def unreliable_api_call():
import random
if random.random() < 0.7:
raise ValueError("API error")
return "Success"
5.3 偏函数应用
python复制from functools import partial
def power(base, exponent):
return base ** exponent
square = partial(power, exponent=2)
cube = partial(power, exponent=3)
print(square(5)) # 25
print(cube(3)) # 27
6. 类型提示与函数注解
Python 3.5+引入了类型提示,可以显著提高代码可读性和IDE支持:
6.1 基本类型注解
python复制def greet(name: str) -> str:
return f"Hello, {name}"
def calculate(a: float, b: float) -> dict[str, float]:
return {
"sum": a + b,
"difference": a - b,
"product": a * b
}
6.2 复杂类型注解
python复制from typing import List, Tuple, Dict, Optional, Union
def process_items(
items: List[Union[int, str]],
settings: Dict[str, float]
) -> Tuple[bool, Optional[str]]:
if not items:
return (False, "No items provided")
return (True, None)
6.3 类型别名
python复制from typing import TypedDict
class Point(TypedDict):
x: float
y: float
def distance(p1: Point, p2: Point) -> float:
return ((p1["x"] - p2["x"])**2 + (p1["y"] - p2["y"])**2)**0.5
注意:Python运行时不会强制检查类型提示,需要使用mypy等工具进行静态检查。
7. 异步函数与协程
Python 3.5+引入了async/await语法支持协程:
7.1 基本异步函数
python复制import asyncio
async def fetch_data(url):
print(f"开始获取 {url}")
await asyncio.sleep(2) # 模拟IO操作
print(f"完成获取 {url}")
return f"{url}的数据"
async def main():
task1 = asyncio.create_task(fetch_data("url1"))
task2 = asyncio.create_task(fetch_data("url2"))
data1 = await task1
data2 = await task2
print(f"获取到数据: {data1}, {data2}")
asyncio.run(main())
7.2 异步生成器
python复制async def async_counter(max):
for i in range(max):
yield i
await asyncio.sleep(0.1)
async def main():
async for num in async_counter(5):
print(num)
asyncio.run(main())
7.3 常见异步模式
超时控制:
python复制async def might_timeout():
try:
await asyncio.wait_for(long_running_task(), timeout=1.0)
except asyncio.TimeoutError:
print("操作超时")
并行执行:
python复制async def gather_example():
results = await asyncio.gather(
fetch_data("url1"),
fetch_data("url2"),
fetch_data("url3"),
return_exceptions=True
)
for result in results:
if isinstance(result, Exception):
print(f"错误: {result}")
else:
print(f"结果: {result}")
8. 函数调试与测试技巧
8.1 调试技巧
使用pdb调试:
python复制import pdb
def problematic_function(data):
pdb.set_trace() # 在这里暂停进入调试器
result = []
for item in data:
processed = complex_processing(item)
result.append(processed)
return result
调试器常用命令:
- n(ext):执行下一行
- s(tep):进入函数调用
- c(ontinue):继续执行直到下一个断点
- l(ist):显示当前代码
- p(rint):打印变量值
- q(uit):退出调试器
8.2 单元测试
使用unittest测试函数:
python复制import unittest
def add(a, b):
return a + b
class TestAddFunction(unittest.TestCase):
def test_add_integers(self):
self.assertEqual(add(1, 2), 3)
def test_add_floats(self):
self.assertAlmostEqual(add(0.1, 0.2), 0.3, places=7)
def test_add_strings(self):
self.assertEqual(add("hello", " world"), "hello world")
if __name__ == "__main__":
unittest.main()
使用pytest的更简洁测试:
python复制# test_functions.py
import pytest
def test_add_integers():
assert add(1, 2) == 3
def test_add_floats():
assert add(0.1, 0.2) == pytest.approx(0.3)
@pytest.mark.parametrize("a,b,expected", [
(1, 1, 2),
(2, 3, 5),
(-1, 1, 0)
])
def test_add_with_params(a, b, expected):
assert add(a, b) == expected
8.3 性能分析
使用cProfile分析函数性能:
python复制import cProfile
def slow_function():
total = 0
for i in range(100000):
total += i**2
return total
profiler = cProfile.Profile()
profiler.enable()
slow_function()
profiler.disable()
profiler.print_stats(sort="cumulative")
使用timeit测量小段代码:
python复制import timeit
setup = "from math import sqrt"
stmt = "sqrt(100)"
time = timeit.timeit(stmt, setup, number=1000000)
print(f"执行100万次耗时: {time:.2f}秒")
9. 函数设计模式实战
9.1 策略模式
python复制class PaymentStrategy:
def pay(self, amount):
pass
class CreditCardPayment(PaymentStrategy):
def pay(self, amount):
print(f"信用卡支付 {amount}元")
class AlipayPayment(PaymentStrategy):
def pay(self, amount):
print(f"支付宝支付 {amount}元")
class PaymentContext:
def __init__(self, strategy):
self._strategy = strategy
def execute_payment(self, amount):
self._strategy.pay(amount)
# 使用示例
payment = PaymentContext(CreditCardPayment())
payment.execute_payment(100)
9.2 工厂模式
python复制class Logger:
def log(self, message):
pass
class FileLogger(Logger):
def log(self, message):
print(f"写入文件: {message}")
class DatabaseLogger(Logger):
def log(self, message):
print(f"存入数据库: {message}")
def get_logger(logger_type):
loggers = {
"file": FileLogger,
"db": DatabaseLogger
}
return loggers[logger_type]()
# 使用示例
logger = get_logger("file")
logger.log("系统启动")
9.3 观察者模式
python复制class EventObserver:
def update(self, event_data):
pass
class EmailNotifier(EventObserver):
def update(self, event_data):
print(f"发送邮件通知: {event_data}")
class LogRecorder(EventObserver):
def update(self, event_data):
print(f"记录日志: {event_data}")
class EventManager:
def __init__(self):
self._observers = []
def attach(self, observer):
self._observers.append(observer)
def notify(self, event_data):
for observer in self._observers:
observer.update(event_data)
# 使用示例
manager = EventManager()
manager.attach(EmailNotifier())
manager.attach(LogRecorder())
manager.notify("用户登录")
10. 函数优化与高级技巧
10.1 使用functools.lru_cache缓存结果
python复制from functools import lru_cache
@lru_cache(maxsize=128)
def fibonacci(n):
if n < 2:
return n
return fibonacci(n-1) + fibonacci(n-2)
print(fibonacci(50)) # 快速计算,因为有缓存
10.2 上下文管理器与with语句
python复制from contextlib import contextmanager
@contextmanager
def managed_resource(path):
print("获取资源")
resource = open(path, "r")
try:
yield resource
finally:
print("释放资源")
resource.close()
# 使用示例
with managed_resource("data.txt") as f:
content = f.read()
print(content)
10.3 动态函数创建
python复制def create_operation(op):
if op == "add":
return lambda a, b: a + b
elif op == "sub":
return lambda a, b: a - b
else:
return lambda a, b: a * b
adder = create_operation("add")
print(adder(3, 5)) # 8
10.4 函数签名检查
python复制from inspect import signature
def validate_args(func, *args, **kwargs):
sig = signature(func)
try:
sig.bind(*args, **kwargs)
return True
except TypeError as e:
print(f"参数错误: {e}")
return False
def example(a, b, c=10):
pass
print(validate_args(example, 1, 2)) # True
print(validate_args(example, 1)) # False
11. 函数与面向对象编程的结合
11.1 类方法与静态方法
python复制class DateUtil:
@staticmethod
def is_valid_date(date_str):
try:
from datetime import datetime
datetime.strptime(date_str, "%Y-%m-%d")
return True
except ValueError:
return False
@classmethod
def today_string(cls):
from datetime import datetime
return datetime.now().strftime("%Y-%m-%d")
# 使用示例
print(DateUtil.is_valid_date("2023-01-01")) # True
print(DateUtil.today_string()) # 当前日期
11.2 使用__call__使对象可调用
python复制class Adder:
def __init__(self, increment):
self.increment = increment
def __call__(self, x):
return x + self.increment
add5 = Adder(5)
print(add5(10)) # 15
11.3 函数与描述符协议
python复制class CachedProperty:
def __init__(self, func):
self.func = func
self.name = func.__name__
def __get__(self, obj, objtype=None):
if obj is None:
return self
value = self.func(obj)
obj.__dict__[self.name] = value
return value
class DataSet:
def __init__(self, data):
self.data = data
@CachedProperty
def stats(self):
print("计算统计数据...")
return {
"mean": sum(self.data)/len(self.data),
"max": max(self.data),
"min": min(self.data)
}
ds = DataSet([1, 2, 3, 4, 5])
print(ds.stats) # 第一次计算
print(ds.stats) # 直接从缓存获取
12. Python函数生态与工具链
12.1 常用函数工具库
operator模块:
python复制from operator import itemgetter, attrgetter, methodcaller
data = [{"name": "Alice", "age": 25}, {"name": "Bob", "age": 30}]
sorted_by_age = sorted(data, key=itemgetter("age"))
class Person:
def __init__(self, name):
self.name = name
people = [Person("Alice"), Person("Bob")]
names = list(map(attrgetter("name"), people))
strs = ["hello", "world"]
uppers = list(map(methodcaller("upper"), strs))
itertools模块:
python复制from itertools import chain, cycle, islice, groupby
# 合并多个迭代器
combined = chain([1, 2], ["a", "b"], (True, False))
# 无限循环
colors = cycle(["red", "green", "blue"])
# 分组
data = sorted([("a", 1), ("b", 2), ("a", 3)], key=itemgetter(0))
for key, group in groupby(data, key=itemgetter(0)):
print(f"{key}: {list(group)}")
12.2 函数文档与帮助
查看函数信息:
python复制import math
print(help(math.sqrt)) # 查看帮助文档
print(math.sqrt.__doc__) # 查看文档字符串
print(dir(math)) # 查看模块所有函数
生成API文档:
使用Sphinx可以自动从文档字符串生成漂亮的HTML文档:
code复制pip install sphinx
sphinx-quickstart
然后在.py文件中编写规范的文档字符串,运行:
code复制sphinx-apidoc -o docs .
cd docs && make html
12.3 函数可视化工具
使用PyCallGraph分析调用关系:
python复制from pycallgraph import PyCallGraph
from pycallgraph.output import GraphvizOutput
def complex_function():
# 复杂调用关系
pass
with PyCallGraph(output=GraphvizOutput()):
complex_function()
使用snakeviz可视化性能分析:
python复制import cProfile
import io
import pstats
from snakeviz.cli import main
def profile_me():
pr = cProfile.Profile()
pr.enable()
# 要分析的代码
pr.disable()
s = io.StringIO()
ps = pstats.Stats(pr, stream=s).sort_stats("cumulative")
ps.print_stats()
with open("profile.txt", "w") as f:
f.write(s.getvalue())
profile_me()
main(["profile.txt"]) # 启动可视化界面
13. 函数安全与防御性编程
13.1 输入验证
python复制def safe_divide(dividend, divisor):
if not isinstance(dividend, (int, float)):
raise TypeError("被除数必须是数字")
if not isinstance(divisor, (int, float)):
raise TypeError("除数必须是数字")
if divisor == 0:
raise ValueError("除数不能为零")
return dividend / divisor
13.2 参数消毒
python复制import html
def sanitize_input(text):
if not isinstance(text, str):
raise TypeError("输入必须是字符串")
return html.escape(text).strip()
user_input = "<script>alert('xss')</script>"
print(sanitize_input(user_input)) # <script>alert('xss')</script>
13.3 权限检查
python复制def requires_permission(permission):
def decorator(func):
def wrapper(user, *args, **kwargs):
if permission not in user.permissions:
raise PermissionError(f"需要{permission}权限")
return func(user, *args, **kwargs)
return wrapper
return decorator
class User:
def __init__(self, permissions):
self.permissions = permissions
@requires_permission("admin")
def delete_user(admin_user, target_user):
print(f"{target_user}已被删除")
admin = User(["admin"])
regular = User([])
delete_user(admin, "hacker") # 正常执行
delete_user(regular, "hacker") # 抛出PermissionError
14. 函数性能优化进阶
14.1 使用__slots__减少内存
python复制class Point:
__slots__ = ("x", "y") # 固定属性列表
def __init__(self, x, y):
self.x = x
self.y = y
def create_many_points():
return [Point(i, i+1) for i in range(100000)]
14.2 使用numpy向量化操作
python复制import numpy as np
def slow_sum(numbers):
total = 0
for num in numbers:
total += num
return total
def fast_sum(numbers):
return np.sum(numbers)
large_array = np.random.rand(1000000)
%timeit slow_sum(large_array) # 约100ms
%timeit fast_sum(large_array) # 约1ms
14.3 使用Cython加速关键函数
python复制# 保存为fast.pyx
def cython_sum(numbers):
cdef long total = 0
cdef int num
for num in numbers:
total += num
return total
# setup.py
from setuptools import setup
from Cython.Build import cythonize
setup(ext_modules=cythonize("fast.pyx"))
# 编译安装
# python setup.py build_ext --inplace
15. 函数在数据科学中的应用
15.1 Pandas中的函数应用
python复制import pandas as pd
df = pd.DataFrame({
"A": [1, 2, 3],
"B": [4, 5, 6]
})
# 应用函数到列
df["A_squared"] = df["A"].apply(lambda x: x**2)
# 应用函数到行
df["sum"] = df.apply(lambda row: row["A"] + row["B"], axis=1)
# 向量化操作
df["product"] = df["A"] * df["B"]
15.2 使用闭包创建特征工程管道
python复制def create_pipeline(*steps):
def pipeline(df):
for step in steps:
df = step(df)
return df
return pipeline
def add_sum(df):
df["sum"] = df["A"] + df["B"]
return df
def add_product(df):
df["product"] = df["A"] * df["B"]
return df
process_data = create_pipeline(add_sum, add_product)
result = process_data(df.copy())
15.3 使用生成器处理大数据
python复制def read_large_csv(file_path, chunk_size=10000):
for chunk in pd.read_csv(file_path, chunksize=chunk_size):
yield process_chunk(chunk)
def process_chunk(chunk):
chunk["new_column"] = chunk["existing"] * 2
return chunk
for processed in read_large_csv("huge_file.csv"):
save_to_database(processed)
16. 函数在Web开发中的应用
16.1 Flask路由函数
python复制from flask import Flask
app = Flask(__name__)
@app.route("/")
def home():
return "Welcome"
@app.route("/user/<username>")
def show_user(username):
return f"User: {username}"
if __name__ == "__main__":
app.run()
16.2 Django视图函数
python复制from django.http import HttpResponse
def hello(request):
name = request.GET.get("name", "World")
return HttpResponse(f"Hello, {name}!")
# urls.py
from django.urls import path
from . import views
urlpatterns = [
path("hello/", views.hello),
]
16.3 FastAPI异步端点
python复制from fastapi import FastAPI
app = FastAPI()
@app.get("/items/{item_id}")
async def read_item(item_id: int, q: str = None):
return {"item_id": item_id, "q": q}
17. 函数式编程库实践
17.1 使用toolz函数式工具
python复制from toolz import compose, pipe, curry
@curry
def add(a, b):
return a + b
add5 = add(5)
print(add5(10)) # 15
# 函数组合
double = lambda x: x * 2
square = lambda x: x ** 2
transform = compose(double, square)
print(transform(5)) # 50
# 管道操作
result = pipe(
5,
lambda x: x + 1,
lambda x: x * 2,
str
)
print(result) # "12"
17.2 使用fn.py函数式扩展
python复制from fn import F, _
# F对象支持管道操作
result = (F(range(10))
.map(_ * 2)
.filter(_ < 10)
.reduce(_ + _))
print(result) # 20
# 模式匹配
from fn.monad import Either, Left, Right
def divide(a, b):
return Right(a / b) if b != 0 else Left("Division by zero")
result = divide(10, 2).match(
right=lambda x: f"结果是 {x}",
left=lambda e: f"错误: {e}"
)
print(result) # "结果是 5.0"
18. 元编程与动态函数
18.1 动态创建函数
python复制def create_function(name, args, body):
code = f"def {name}({','.join(args)}):\n {body}"
locals_dict = {}
exec(code, globals(), locals_dict)
return locals_dict[name]
dynamic_func = create_function(
"say_hello",
["name"],
"print(f'Hello, {name}!')"
)
dynamic_func("World") # 输出 "Hello, World!"
18.2 函数装饰器工厂
python复制def validate_args(*validators):
def decorator(func):
def wrapper(*args, **kwargs):
for val, arg in zip(validators, args):
if not val(arg):
raise ValueError(f"参数 {arg} 无效")
return func(*args, **kwargs)
return wrapper
return decorator
@validate_args(lambda x: x > 0, lambda y: isinstance(y, str))
def process_data(num, text):
return text * num
print(process_data(3, "Hi")) # "HiHiHi"
print(process_data(-1, "Hi")) # ValueError
18.3 使用inspect修改函数行为
python复制import inspect
def log_calls(func):
def wrapper(*args, **kwargs):
print(f"调用 {func.__name__} 参数: {args}, {kwargs}")
return func(*args, **kwargs)
# 复制原始函数的元数据
wrapper.__name__ = func.__name__
wrapper.__doc__ = func.__doc__
wrapper.__module__ = func.__module__
wrapper.__annotations__ = func.__annotations__
wrapper.__signature__ = inspect.signature(func)
return wrapper
@log_calls
def add(a: int, b: int) -> int:
"""两数相加"""
return a + b
print(add(2, 3)) # 输出调用日志和结果5
print(help(add)) # 仍然显示正确的帮助信息
19. 函数测试与Mock技术
19.1 使用unittest.mock
python复制from unittest.mock import Mock, patch
def call_api(url):
# 实际会调用外部API
pass
def test_with_mock():
# 创建mock对象替代真实API调用
mock_api = Mock(return_value="mocked response")
# 替换函数
with patch("__main__.call_api", mock_api):
result = call_api("http://example.com")
assert result == "mocked response"
mock_api.assert_c
