Django基于用户的数据库路由实现方案优化咨询
Django基于用户的数据库路由优化方案咨询
我正尝试在Django中实现基于用户的数据库路由(PER USER database routing),但推进过程并不顺利。所有数据库均为预定义状态,且结构完全一致。我目前已实现了一套可行方案(代码如下),但想咨询是否存在更优的实现方式:
class DatabaseMiddleware: def __init__(self, get_response): self.get_response = get_response def __call__(self, request): if (str(request.company) == 'company1'): request.database = 'COMPANY_X' else: request.database = 'default' response = self.get_response(request) return response # 当前方案存在安全风险,但暂未找到更好的替代方案 from threadlocals.threadlocals import get_current_request class UserDatabaseRouter: def db_for_read(self, model,user_id=None, **hints): request = get_current_request() if(not (request is None)): return request.database else: return None def db_for_write(self, model,user_id=None, **hints): request = get_current_request() if(not (request is None)): return request.database else: return None
优化方向建议
1. 移除第三方threadlocals依赖,改用原生线程存储
Django生态有原生的线程/异步上下文存储方案,无需依赖第三方库,同时能避免潜在的兼容性问题:
# 自定义上下文存储(兼容WSGI/ASGI) from asgiref.local import Local local_storage = Local() class DatabaseMiddleware: def __init__(self, get_response): self.get_response = get_response def __call__(self, request): # 用映射表替代硬编码,便于扩展 db_mapping = { 'company1': 'COMPANY_X', # 可添加更多公司-数据库映射 } # 优先从映射表取,默认用default local_storage.database = db_mapping.get(str(request.company), 'default') response = self.get_response(request) # 请求结束后清理上下文,避免线程复用导致脏数据 if hasattr(local_storage, 'database'): del local_storage.database return response class UserDatabaseRouter: def db_for_read(self, model, **hints): return getattr(local_storage, 'database', 'default') def db_for_write(self, model, **hints): return getattr(local_storage, 'database', 'default')
2. 强化安全与可维护性
- 把映射关系移到配置文件:避免硬编码,在
settings.py中定义映射表,便于统一管理:
中间件中读取配置:# settings.py COMPANY_DB_MAPPING = { 'company1': 'COMPANY_X', 'company2': 'COMPANY_Y', } DEFAULT_DATABASE = 'default'from django.conf import settings # 中间件内替换映射逻辑 local_storage.database = settings.COMPANY_DB_MAPPING.get(str(request.company), settings.DEFAULT_DATABASE) - 增加参数校验:对
request.company做合法性校验,比如判断是否在映射表的键集合内,防止非法值导致错误路由:# 中间件内 company_str = str(request.company) if company_str not in settings.COMPANY_DB_MAPPING: company_str = 'default' local_storage.database = settings.COMPANY_DB_MAPPING.get(company_str, settings.DEFAULT_DATABASE)
3. 兼容无请求上下文场景
比如Celery异步任务、管理命令这类没有request对象的场景,可以通过路由的hints参数传递标识:
class UserDatabaseRouter: def db_for_read(self, model, **hints): # 优先从hints获取用户/公司信息 user = hints.get('user') if user: return settings.COMPANY_DB_MAPPING.get(str(user.company), settings.DEFAULT_DATABASE) # 再尝试从上下文获取 return getattr(local_storage, 'database', settings.DEFAULT_DATABASE) def db_for_write(self, model, **hints): return self.db_for_read(model, **hints)
异步任务中调用示例:
# Celery任务内 MyModel.objects.db_manager(hints={'user': current_user}).all()
内容的提问来源于stack exchange,提问作者Lerex
相关产品推荐
相关产品推荐

