方案一:依赖注入 + 上下文变量(推荐)
基础设置
python
# database.py
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from sqlalchemy.orm import declarative_base
from typing import AsyncGenerator, Dict
from contextvars import ContextVar
from functools import lru_cache
import asyncio
# 声明基类
Base = declarative_base()
# 上下文变量存储当前事务会话
_db_context: ContextVar[Dict[str, AsyncSession]] = ContextVar("db_context", default={})
_transaction_context: ContextVar[bool] = ContextVar("transaction_context", default=False)
@lru_cache()
def get_engine(db_url: str):
"""缓存引擎创建"""
return create_async_engine(
db_url,
echo=True,
pool_size=20,
max_overflow=40,
pool_pre_ping=True,
pool_recycle=3600,
)
class DatabaseManager:
def __init__(self):
self.engines = {}
self.session_makers = {}
def init_db(self, db_configs: dict):
"""初始化多个数据库"""
for db_name, db_url in db_configs.items():
engine = get_engine(db_url)
self.engines[db_name] = engine
self.session_makers[db_name] = async_sessionmaker(
engine,
class_=AsyncSession,
expire_on_commit=False,
autoflush=False,
autocommit=False
)
async def close_all(self):
"""关闭所有数据库连接"""
for engine in self.engines.values():
await engine.dispose()
db_manager = DatabaseManager()
# 依赖注入 - 获取会话
async def get_db_session(db_name: str = "default") -> AsyncGenerator[AsyncSession, None]:
"""获取数据库会话的依赖"""
session_maker = db_manager.session_makers.get(db_name)
if not session_maker:
raise ValueError(f"Database {db_name} not configured")
# 检查是否已在事务上下文中
sessions = _db_context.get()
if db_name in sessions:
yield sessions[db_name]
return
async with session_maker() as session:
try:
yield session
finally:
await session.close()
# 事务管理依赖
async def transactional(
*db_names: str,
isolation_level: str = "READ_COMMITTED"
) -> AsyncGenerator[Dict[str, AsyncSession], None]:
"""
全局事务依赖,支持多数据库
Args:
db_names: 需要参与事务的数据库名称
isolation_level: 事务隔离级别
"""
sessions = {}
current_sessions = _db_context.get()
# 创建所有需要的会话
for db_name in db_names:
if db_name in current_sessions:
sessions[db_name] = current_sessions[db_name]
else:
session_maker = db_manager.session_makers.get(db_name)
if not session_maker:
raise ValueError(f"Database {db_name} not configured")
session = session_maker()
# 设置事务隔离级别
await session.execute(f"SET TRANSACTION ISOLATION LEVEL {isolation_level}")
sessions[db_name] = session
# 保存当前会话到上下文
new_sessions = {**current_sessions, **sessions}
token = _db_context.set(new_sessions)
trans_token = _transaction_context.set(True)
try:
yield sessions
# 提交所有事务
commit_tasks = []
for session in sessions.values():
if session.in_transaction():
commit_tasks.append(session.commit())
if commit_tasks:
await asyncio.gather(*commit_tasks)
except Exception as e:
# 回滚所有事务
rollback_tasks = []
for session in sessions.values():
if session.in_transaction():
rollback_tasks.append(session.rollback())
if rollback_tasks:
await asyncio.gather(*rollback_tasks)
raise e
finally:
# 清理资源
_db_context.reset(token)
_transaction_context.reset(trans_token)
# 关闭新创建的会话
for db_name, session in sessions.items():
if db_name not in current_sessions:
await session.close()
使用示例
python
# main.py
from fastapi import FastAPI, Depends
from sqlalchemy import select
from .database import get_db_session, transactional, db_manager
from .models import User, Order
from .schemas import UserCreate, OrderCreate
app = FastAPI()
@app.on_event("startup")
async def startup():
db_manager.init_db({
"default": "postgresql+asyncpg://user:pass@localhost/default_db",
"order_db": "postgresql+asyncpg://user:pass@localhost/order_db",
"user_db": "postgresql+asyncpg://user:pass@localhost/user_db",
})
@app.post("/users/", response_model=UserCreate)
async def create_user(
user: UserCreate,
sessions: dict = Depends(lambda: transactional("user_db"))
):
"""单数据库事务"""
session = sessions["user_db"]
db_user = User(**user.dict())
session.add(db_user)
await session.flush()
return db_user
@app.post("/orders/", response_model=OrderCreate)
async def create_order_with_user(
order_data: OrderCreate,
sessions: dict = Depends(lambda: transactional("user_db", "order_db"))
):
"""多数据库分布式事务"""
user_session = sessions["user_db"]
order_session = sessions["order_db"]
# 创建用户
user = User(name=order_data.user_name, email=order_data.email)
user_session.add(user)
await user_session.flush()
# 创建订单
order = Order(
user_id=user.id,
product_name=order_data.product_name,
amount=order_data.amount
)
order_session.add(order)
# 注意:这里不需要显式commit,依赖会自动处理
return order
方案二:中间件全局事务管理
python
# middleware.py
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from contextvars import ContextVar
from typing import Dict, Set
import asyncio
_db_sessions: ContextVar[Dict[str, AsyncSession]] = ContextVar("db_sessions", default={})
_active_dbs: ContextVar[Set[str]] = ContextVar("active_dbs", default=set())
class GlobalTransactionMiddleware(BaseHTTPMiddleware):
def __init__(
self,
app,
db_names: list,
auto_transaction: bool = True,
isolation_level: str = "READ_COMMITTED"
):
super().__init__(app)
self.db_names = db_names
self.auto_transaction = auto_transaction
self.isolation_level = isolation_level
async def dispatch(self, request: Request, call_next):
if not self.auto_transaction:
return await call_next(request)
# 初始化会话字典
sessions = {}
active_dbs = set()
try:
# 为每个数据库创建会话
for db_name in self.db_names:
session_maker = db_manager.session_makers.get(db_name)
if session_maker:
session = session_maker()
await session.execute(
f"SET TRANSACTION ISOLATION LEVEL {self.isolation_level}"
)
sessions[db_name] = session
active_dbs.add(db_name)
# 保存到上下文
_db_sessions.set(sessions)
_active_dbs.set(active_dbs)
# 执行请求
response = await call_next(request)
# 提交所有事务
commit_tasks = []
for session in sessions.values():
if session.in_transaction():
commit_tasks.append(session.commit())
if commit_tasks:
await asyncio.gather(*commit_tasks)
return response
except Exception as e:
# 回滚所有事务
rollback_tasks = []
for session in sessions.values():
if session.in_transaction():
rollback_tasks.append(session.rollback())
if rollback_tasks:
await asyncio.gather(*rollback_tasks)
raise e
finally:
# 清理会话
close_tasks = [session.close() for session in sessions.values()]
if close_tasks:
await asyncio.gather(*close_tasks)
_db_sessions.set({})
_active_dbs.set(set())
# 获取当前请求的会话
async def get_current_session(db_name: str = "default") -> AsyncSession:
"""获取当前请求上下文中的会话"""
sessions = _db_sessions.get()
if db_name not in sessions:
raise ValueError(f"No active session for database: {db_name}")
return sessions[db_name]
# 动态添加数据库到当前事务
async def add_db_to_transaction(db_name: str):
"""动态将数据库加入当前事务"""
sessions = _db_sessions.get()
active_dbs = _active_dbs.get()
if db_name not in sessions and db_name not in active_dbs:
session_maker = db_manager.session_makers.get(db_name)
if session_maker:
session = session_maker()
sessions[db_name] = session
active_dbs.add(db_name)
_db_sessions.set(sessions)
_active_dbs.set(active_dbs)
使用中间件
python
# main.py
from fastapi import FastAPI
from .middleware import GlobalTransactionMiddleware
from .database import db_manager
app = FastAPI()
# 添加全局事务中间件
app.add_middleware(
GlobalTransactionMiddleware,
db_names=["user_db", "order_db"],
auto_transaction=True,
isolation_level="READ_COMMITTED"
)
@app.post("/create-user-order/")
async def create_user_order():
"""自动参与全局事务,无需显式声明依赖"""
# 获取当前会话
user_session = await get_current_session("user_db")
order_session = await get_current_session("order_db")
# 执行业务逻辑...
方案三:装饰器模式
python
# decorators.py
from functools import wraps
from typing import Callable, List
import asyncio
import inspect
def transactional_decorator(
*db_names: str,
isolation_level: str = "READ_COMMITTED"
):
"""事务装饰器"""
def decorator(func: Callable):
@wraps(func)
async def wrapper(*args, **kwargs):
# 获取函数签名
sig = inspect.signature(func)
bound_args = sig.bind(*args, **kwargs)
bound_args.apply_defaults()
sessions = {}
created_sessions = []
try:
# 创建会话
for db_name in db_names:
session_maker = db_manager.session_makers.get(db_name)
if not session_maker:
raise ValueError(f"Database {db_name} not configured")
session = session_maker()
await session.execute(
f"SET TRANSACTION ISOLATION LEVEL {isolation_level}"
)
sessions[db_name] = session
created_sessions.append(session)
# 将会话注入函数参数
new_kwargs = kwargs.copy()
for db_name, session in sessions.items():
param_name = f"{db_name}_session"
if param_name in sig.parameters:
new_kwargs[param_name] = session
# 执行函数
result = await func(*args, **new_kwargs)
# 提交事务
commit_tasks = [
session.commit()
for session in created_sessions
if session.in_transaction()
]
if commit_tasks:
await asyncio.gather(*commit_tasks)
return result
except Exception as e:
# 回滚事务
rollback_tasks = [
session.rollback()
for session in created_sessions
if session.in_transaction()
]
if rollback_tasks:
await asyncio.gather(*rollback_tasks)
raise e
finally:
# 关闭会话
close_tasks = [session.close() for session in created_sessions]
if close_tasks:
await asyncio.gather(*close_tasks)
return wrapper
return decorator
# 使用装饰器
@app.post("/transfer/")
@transactional_decorator("user_db", "account_db")
async def transfer_funds(
from_user: int,
to_user: int,
amount: float,
user_db_session: AsyncSession,
account_db_session: AsyncSession
):
"""资金转账示例"""
# 扣减转出账户
await user_db_session.execute(
"UPDATE accounts SET balance = balance - :amount WHERE user_id = :user_id",
{"amount": amount, "user_id": from_user}
)
# 增加转入账户
await account_db_session.execute(
"UPDATE accounts SET balance = balance + :amount WHERE user_id = :user_id",
{"amount": amount, "user_id": to_user}
)
return {"status": "success"}
方案四:Unit of Work 模式
python
# unit_of_work.py
from abc import ABC, abstractmethod
from typing import Type, Dict, Any
from contextlib import asynccontextmanager
class UnitOfWork(ABC):
"""工作单元抽象基类"""
def __init__(self):
self.sessions: Dict[str, AsyncSession] = {}
self._committed = False
self._rolledback = False
@abstractmethod
async def __aenter__(self):
pass
@abstractmethod
async def __aexit__(self, exc_type, exc_val, exc_tb):
pass
async def commit(self):
"""提交所有事务"""
if not self._committed and not self._rolledback:
commit_tasks = [
session.commit()
for session in self.sessions.values()
if session.in_transaction()
]
if commit_tasks:
await asyncio.gather(*commit_tasks)
self._committed = True
async def rollback(self):
"""回滚所有事务"""
if not self._rolledback:
rollback_tasks = [
session.rollback()
for session in self.sessions.values()
if session.in_transaction()
]
if rollback_tasks:
await asyncio.gather(*rollback_tasks)
self._rolledback = True
async def close(self):
"""关闭所有会话"""
close_tasks = [session.close() for session in self.sessions.values()]
if close_tasks:
await asyncio.gather(*close_tasks)
class MultiDBUnitOfWork(UnitOfWork):
"""多数据库工作单元"""
def __init__(self, *db_names: str):
super().__init__()
self.db_names = db_names
async def __aenter__(self):
for db_name in self.db_names:
session_maker = db_manager.session_makers.get(db_name)
if session_maker:
session = session_maker()
self.sessions[db_name] = session
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
try:
if exc_type is not None:
await self.rollback()
else:
await self.commit()
finally:
await self.close()
def get_session(self, db_name: str) -> AsyncSession:
"""获取指定数据库的会话"""
if db_name not in self.sessions:
raise ValueError(f"No session for database: {db_name}")
return self.sessions[db_name]
# 依赖注入版本
async def get_uow(*db_names: str) -> AsyncGenerator[MultiDBUnitOfWork, None]:
"""获取工作单元的依赖"""
async with MultiDBUnitOfWork(*db_names) as uow:
yield uow
# 使用示例
@app.post("/complex-operation/")
async def complex_operation(
uow: MultiDBUnitOfWork = Depends(lambda: get_uow("user_db", "order_db", "log_db"))
):
"""复杂业务操作"""
# 获取各数据库会话
user_session = uow.get_session("user_db")
order_session = uow.get_session("order_db")
log_session = uow.get_session("log_db")
try:
# 业务逻辑
# ...
# 记录日志
log_entry = Log(action="create_order", status="success")
log_session.add(log_entry)
except Exception as e:
# 记录错误日志
error_log = Log(action="create_order", status="failed", error=str(e))
log_session.add(error_log)
raise
方案五:嵌套事务与保存点
python
# nested_transactions.py
from contextlib import asynccontextmanager
from typing import AsyncGenerator
class NestedTransactionManager:
"""嵌套事务管理器,支持保存点"""
def __init__(self):
self.savepoints = {}
@asynccontextmanager
async def nested_transaction(
self,
session: AsyncSession,
savepoint_name: str = None
) -> AsyncGenerator[None, None]:
"""创建嵌套事务(保存点)"""
if not savepoint_name:
import uuid
savepoint_name = f"sp_{uuid.uuid4().hex[:8]}"
# 创建保存点
await session.execute(f"SAVEPOINT {savepoint_name}")
self.savepoints[savepoint_name] = session
try:
yield
# 释放保存点
await session.execute(f"RELEASE SAVEPOINT {savepoint_name}")
except Exception:
# 回滚到保存点
await session.execute(f"ROLLBACK TO SAVEPOINT {savepoint_name}")
raise
finally:
self.savepoints.pop(savepoint_name, None)
# 增强版事务依赖,支持嵌套
async def transactional_with_nesting(
*db_names: str,
isolation_level: str = "READ_COMMITTED"
) -> AsyncGenerator[Dict[str, AsyncSession], None]:
"""支持嵌套事务的依赖"""
sessions = {}
transaction_manager = NestedTransactionManager()
# ... 初始化会话代码类似之前 ...
# 为每个会话添加嵌套事务方法
class SessionWithNesting:
def __init__(self, session: AsyncSession):
self.session = session
self.nested_manager = transaction_manager
@asynccontextmanager
async def nested(self, savepoint_name: str = None):
async with self.nested_manager.nested_transaction(
self.session, savepoint_name
):
yield
# 包装会话
wrapped_sessions = {
name: SessionWithNesting(session)
for name, session in sessions.items()
}
try:
yield wrapped_sessions
# 提交逻辑...
except Exception:
# 回滚逻辑...
raise
# 使用示例
@app.post("/batch-process/")
async def batch_process(
sessions: dict = Depends(lambda: transactional_with_nesting("user_db"))
):
"""批量处理,部分失败不影响整体"""
session_wrapper = sessions["user_db"]
for i in range(10):
try:
async with session_wrapper.nested(f"batch_item_{i}"):
# 处理单个项目
user = User(name=f"user_{i}")
session_wrapper.session.add(user)
# 模拟可能失败的操作
if i == 5:
raise ValueError("Processing failed")
except Exception:
# 单个项目失败,继续处理下一个
continue
配置和模型示例
python
# models.py
from sqlalchemy import Column, Integer, String, Float, ForeignKey, DateTime
from sqlalchemy.sql import func
from .database import Base
class User(Base):
__tablename__ = "users"
id = Column(Integer, primary_key=True, index=True)
name = Column(String(100), nullable=False)
email = Column(String(100), unique=True, nullable=False)
created_at = Column(DateTime(timezone=True), server_default=func.now())
class Order(Base):
__tablename__ = "orders"
id = Column(Integer, primary_key=True, index=True)
user_id = Column(Integer, ForeignKey("users.id"), nullable=False)
product_name = Column(String(200), nullable=False)
amount = Column(Float, nullable=False)
created_at = Column(DateTime(timezone=True), server_default=func.now())
class Log(Base):
__tablename__ = "logs"
id = Column(Integer, primary_key=True, index=True)
action = Column(String(100), nullable=False)
status = Column(String(50), nullable=False)
error = Column(String(500))
created_at = Column(DateTime(timezone=True), server_default=func.now())
# schemas.py
from pydantic import BaseModel, EmailStr
from datetime import datetime
from typing import Optional
class UserBase(BaseModel):
name: str
email: EmailStr
class UserCreate(UserBase):
pass
class UserResponse(UserBase):
id: int
created_at: datetime
class Config:
from_attributes = True
class OrderBase(BaseModel):
user_name: str
email: EmailStr
product_name: str
amount: float
class OrderCreate(OrderBase):
pass
class OrderResponse(OrderBase):
id: int
user_id: int
created_at: datetime
class Config:
from_attributes = True
方案比较汇总
| 方案 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 方案一:依赖注入 + 上下文变量 | ✅ 最灵活,精确控制事务边界 ✅ 支持动态添加数据库 ✅ 符合FastAPI设计哲学 ✅ 易于测试和维护 | ❌ 需要显式声明依赖 ❌ 代码侵入性稍强 | 大多数业务场景,特别是需要精细控制事务的场景 |
| 方案二:中间件全局事务 | ✅ 完全透明,无需修改业务代码 ✅ 统一管理所有请求 ✅ 适合CRUD密集型应用 | ❌ 不够灵活,所有请求都参与事务 ❌ 难以处理复杂事务逻辑 ❌ 性能开销较大 | 简单的CRUD应用,所有操作都需要事务保障 |
| 方案三:装饰器模式 | ✅ 代码简洁,声明式事务 ✅ 复用性强 ✅ 对业务代码侵入小 | ❌ 灵活性较差 ❌ 调试相对困难 ❌ 动态添加数据库不便 | 固定的事务边界,方法级别的原子性要求 |
| 方案四:Unit of Work模式 | ✅ 最符合DDD思想 ✅ 业务逻辑与数据访问分离 ✅ 易于实现复杂业务规则 ✅ 支持聚合根管理 | ❌ 学习曲线较陡 ❌ 代码量较多 ❌ 过度设计风险 | 复杂领域模型,需要严格业务规则的应用 |
| 方案五:嵌套事务与保存点 | ✅ 支持部分回滚 ✅ 适合批处理场景 ✅ 提高容错性 | ❌ 数据库支持限制(非所有DB都支持) ❌ 增加复杂度 ❌ 调试困难 | 批量处理,需要部分失败的复杂业务流程 |
最佳实践建议
-
推荐使用方案一:依赖注入 + 上下文变量的组合,它提供了最佳平衡
-
分层设计:
python# repository层只负责数据访问 class UserRepository: def __init__(self, session: AsyncSession): self.session = session async def create(self, user: User) -> User: self.session.add(user) await self.session.flush() return user # service层处理业务逻辑和事务 class UserService: async def create_user_with_order( self, user_data: dict, order_data: dict, sessions: dict ): async with sessions["user_db"].begin_nested(): # 业务逻辑 pass -
异常处理策略:
pythonasync def safe_transaction(func): @wraps(func) async def wrapper(*args, **kwargs): try: return await func(*args, **kwargs) except IntegrityError as e: # 处理数据库完整性错误 raise HTTPException(status_code=400, detail=str(e)) except Exception as e: # 记录日志并重新抛出 logger.error(f"Transaction failed: {e}") raise return wrapper -
连接池优化:
python# 针对不同数据库配置不同的连接池 engine = create_async_engine( db_url, pool_size=10 if db_name == "main" else 5, max_overflow=20, pool_timeout=30, pool_recycle=1800, ) -
监控和日志:
pythonimport logging import time async def log_transaction_time(func): @wraps(func) async def wrapper(*args, **kwargs): start = time.time() try: result = await func(*args, **kwargs) duration = time.time() - start logging.info(f"Transaction completed in {duration:.2f}s") return result except Exception as e: duration = time.time() - start logging.error(f"Transaction failed after {duration:.2f}s: {e}") raise return wrapper
最终推荐架构
┌─────────────────────────────────────────────┐
│ FastAPI App │
├─────────────────────────────────────────────┤
│ Middleware (可选,用于全局监控/日志) │
├─────────────────────────────────────────────┤
│ Dependencies (事务管理入口) │
│ ├── transactional() │
│ ├── get_uow() │
│ └── get_db_session() │
├─────────────────────────────────────────────┤
│ Services (业务逻辑层) │
│ ├── UserService │
│ ├── OrderService │
│ └── TransactionCoordinator │
├─────────────────────────────────────────────┤
│ Repositories (数据访问层) │
│ ├── UserRepository │
│ └── OrderRepository │
├─────────────────────────────────────────────┤
│ Models & Schemas │
├─────────────────────────────────────────────┤
│ Database Manager (多数据库连接管理) │
│ ├── Engines Pool │
│ └── Session Makers │
└─────────────────────────────────────────────┘
这种架构结合了方案一的灵活性和方案四的结构性,既保持了代码的清晰性,又提供了强大的事务管理能力。对于大多数生产环境,这是最优的选择。