本文目录导读:

我来通过几个实际案例来讲解令牌桶算法的应用。
案例1:API网关限流(最典型)
import time
import threading
from collections import deque
class TokenBucket:
def __init__(self, capacity, refill_rate):
"""
capacity: 桶容量(最大令牌数)
refill_rate: 每秒补充令牌数
"""
self.capacity = capacity
self.refill_rate = refill_rate
self.tokens = capacity # 初始状态桶是满的
self.last_refill_time = time.time()
self.lock = threading.Lock()
def try_acquire(self, tokens=1):
with self.lock:
# 计算自上次补充以来应添加的令牌数
now = time.time()
delta = now - self.last_refill_time
self.tokens = min(self.capacity,
self.tokens + delta * self.refill_rate)
self.last_refill_time = now
# 尝试获取令牌
if self.tokens >= tokens:
self.tokens -= tokens
return True
return False
# 使用示例:API限流
api_bucket = TokenBucket(capacity=10, refill_rate=2) # 突发10个请求,每秒恢复2个
def api_request(user_id):
if api_bucket.try_acquire():
print(f"时间 {time.time():.2f}: 用户{user_id} 请求成功")
# 处理请求
return "success"
else:
print(f"时间 {time.time():.2f}: 用户{user_id} 请求被限流")
return "rate_limited"
# 模拟并发请求
import random
for i in range(20):
time.sleep(0.1) # 每100ms一个请求
api_request(random.randint(1, 5))
案例2:流量整形(平滑突发流量)
class TrafficShaper:
def __init__(self, peak_rate, sustained_rate, bucket_size):
"""
peak_rate: 峰值速率(请求/秒)
sustained_rate: 持续速率(请求/秒)
bucket_size: 令牌桶大小
"""
self.bucket = TokenBucket(capacity=bucket_size,
refill_rate=sustained_rate)
self.peak_rate = peak_rate
self.min_interval = 1.0 / peak_rate
self.last_request_time = 0
def wait_and_send(self, data):
"""等待令牌并发送数据"""
# 控制峰值速率
now = time.time()
if now - self.last_request_time < self.min_interval:
wait_time = self.min_interval - (now - self.last_request_time)
time.sleep(wait_time)
# 获取令牌
if self.bucket.try_acquire():
self.last_request_time = time.time()
print(f"发送数据包大小: {len(data)} bytes")
return True
return False
# 使用示例:网络流量整形
shaper = TrafficShaper(peak_rate=10, sustained_rate=3, bucket_size=5)
# 模拟数据发送
import random
for i in range(15):
data = "x" * random.randint(100, 1000)
success = shaper.wait_and_send(data)
if not success:
print(f"等待:数据包 {i} 被缓冲")
案例3:分布式系统中的单机限流
import asyncio
from datetime import datetime
class ServiceRateLimiter:
def __init__(self, service_name, qps, burst_percentage=0.5):
"""
service_name: 服务名称
qps: 每秒钟的查询量
burst_percentage: 允许的突发比例
"""
self.service_name = service_name
self.base_qps = qps
# 允许的突发量 = 基础QPS * 突发百分比
burst_capacity = int(qps * (1 + burst_percentage))
self.bucket = TokenBucket(
capacity=burst_capacity, # 100 * 1.5 = 150
refill_rate=qps # 正常速率
)
self.total_requests = 0
self.rejected_requests = 0
async def validate_request(self, user_id, request_type):
"""验证请求是否允许处理"""
self.total_requests += 1
if self.bucket.try_acquire():
# 记录成功请求
print(f"[{datetime.now().strftime('%H:%M:%S')}] "
f"服务 {self.service_name} 接受请求: "
f"用户 {user_id}, 类型 {request_type}")
return True
else:
self.rejected_requests += 1
print(f"[{datetime.now().strftime('%H:%M:%S')}] "
f"服务 {self.service_name} 拒绝请求: "
f"用户 {user_id}, 原因: 限流")
return False
def get_stats(self):
"""获取限流统计信息"""
reject_rate = (self.rejected_requests / self.total_requests * 100
if self.total_requests > 0 else 0)
return {
'total_requests': self.total_requests,
'rejected_requests': self.rejected_requests,
'reject_rate': f"{reject_rate:.2f}%"
}
# 使用示例
import asyncio
async def simulate_service_load():
# 创建一个支持100 QPS的服务,允许50%的突发
rate_limiter = ServiceRateLimiter("订单服务", qps=100, burst_percentage=0.5)
# 模拟突发流量
tasks = []
for i in range(200):
if i < 50: # 模拟初始突发
await asyncio.sleep(0.01) # 10ms一个请求
else:
await asyncio.sleep(0.02) # 20ms一个请求
task = asyncio.create_task(
rate_limiter.validate_request(
user_id=i % 10,
request_type="create_order"
)
)
tasks.append(task)
await asyncio.gather(*tasks)
# 输出统计信息
stats = rate_limiter.get_stats()
print(f"\n限流统计: {stats}")
# 运行模拟
asyncio.run(simulate_service_load())
案例4:多级限流(网关 + 服务)
class MultiLevelRateLimiter:
"""多级限流:全局 + 用户级 + IP级"""
def __init__(self, global_qps=1000, user_qps=10, ip_qps=5):
self.global_bucket = TokenBucket(capacity=global_qps,
refill_rate=global_qps)
self.user_buckets = {} # {user_id: TokenBucket}
self.ip_buckets = {} # {ip_address: TokenBucket}
self.user_qps = user_qps
self.ip_qps = ip_qps
def check_all_levels(self, user_id, ip_address):
"""检查所有级别的限流"""
# 1. 全局限流
if not self.global_bucket.try_acquire():
print(f"全局限流生效")
return False
# 2. 用户级限流
if user_id not in self.user_buckets:
self.user_buckets[user_id] = TokenBucket(
capacity=self.user_qps,
refill_rate=self.user_qps
)
if not self.user_buckets[user_id].try_acquire():
print(f"用户 {user_id} 限流生效")
return False
# 3. IP级限流
if ip_address not in self.ip_buckets:
self.ip_buckets[ip_address] = TokenBucket(
capacity=self.ip_qps,
refill_rate=self.ip_qps
)
if not self.ip_buckets[ip_address].try_acquire():
print(f"IP {ip_address} 限流生效")
return False
return True
# 使用示例
limiter = MultiLevelRateLimiter(global_qps=1000, user_qps=3, ip_qps=2)
def process_request(request):
user_id = request['user_id']
ip = request['ip']
if limiter.check_all_levels(user_id, ip):
print(f"请求处理成功: 用户{user_id}, IP:{ip}")
return 200
else:
print(f"请求被拒绝: 用户{user_id}, IP:{ip}")
return 429 # Too Many Requests
# 模拟测试
requests = [
{'user_id': 'A', 'ip': '192.168.1.1'},
{'user_id': 'A', 'ip': '192.168.1.1'},
{'user_id': 'A', 'ip': '192.168.1.1'},
{'user_id': 'A', 'ip': '192.168.1.1'}, # 应该被拒绝(用户限流)
{'user_id': 'B', 'ip': '192.168.1.1'}, # 应该被拒绝(IP限流)
{'user_id': 'B', 'ip': '192.168.1.2'}, # 应该成功
]
for req in requests:
process_request(req)
关键设计要点
参数选择建议
# 不同场景的参数配置示例
# API网关:高吞吐、低延迟
api_config = {
'capacity': 1000,
'refill_rate': 500 # 每秒500个请求
}
# 数据库保护:防止过载
db_config = {
'capacity': 20,
'refill_rate': 5 # 每秒5个查询
}
# 短信服务:严格限制
sms_config = {
'capacity': 10,
'refill_rate': 1/60 # 每分钟1条
}
# 文件上传:允许突发但限制持续速率
upload_config = {
'capacity': 100, # 允许突发100个
'refill_rate': 10 # 每钞10个
}
性能优化技巧
class OptimizedTokenBucket:
"""高性能令牌桶实现"""
def __init__(self, tokens_per_sec, bucket_size):
self.tokens_per_sec = tokens_per_sec
self.bucket_size = bucket_size
self.tokens = bucket_size
self.last_refill = time.monotonic()
self.lock = threading.Lock()
def acquire(self, tokens=1):
"""原子操作,减少锁竞争"""
with self.lock:
now = time.monotonic()
# 使用单调时钟避免时间跳跃
self.tokens = min(self.bucket_size,
self.tokens + (now - self.last_refill) * self.tokens_per_sec)
self.last_refill = now
if self.tokens >= tokens:
self.tokens -= tokens
return True
return False
# 批量版本
def acquire_batch(self, num_tokens):
"""批量获取令牌,用于批次操作"""
acquired = []
with self.lock:
for _ in range(num_tokens):
if self.acquire():
acquired.append(time.monotonic())
else:
break
return acquired
实际应用案例总结
| 场景 | 桶容量 | 补充速率 | 特点 |
|---|---|---|---|
| API网关 | 1000 | 500/s | 允许突发,平稳持续流量 |
| 用户级别 | 10 | 2/s | 防止单个用户滥用 |
| 数据库保护 | 100 | 30/s | 防止连接池过载 |
| 短信服务 | 5 | 1/min | 严格控制频率 |
| 文件上传 | 50 | 10/s | 应对突发上传 |
这些案例展示了令牌桶算法的核心优势:
- 允许突发流量:桶容量支持突发
- 平滑限流:持续限制速率
- 灵活配置:可针对不同需求调整参数
- 简单高效:实现简单,性能好
选择令牌桶而不是漏桶的原因通常是需要支持突发流量,这是Web应用、API服务等场景的常见需求。