FastAPI + SQLAlchemy全局事务管理

方案一:依赖注入 + 上下文变量(推荐)

基础设置

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都支持) ❌ 增加复杂度 ❌ 调试困难 批量处理,需要部分失败的复杂业务流程

最佳实践建议

  1. 推荐使用方案一:依赖注入 + 上下文变量的组合,它提供了最佳平衡

  2. 分层设计

    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
  3. 异常处理策略

    python 复制代码
    async 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
  4. 连接池优化

    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,
    )
  5. 监控和日志

    python 复制代码
    import 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                        │
└─────────────────────────────────────────────┘

这种架构结合了方案一的灵活性和方案四的结构性,既保持了代码的清晰性,又提供了强大的事务管理能力。对于大多数生产环境,这是最优的选择。

相关推荐
用户8356290780511 小时前
如何使用 Python 在 Excel 中添加、编辑和删除超链接
后端·python
AC赳赳老秦3 小时前
企业工商公开信息采集分析:OpenClaw 批量查询企业工商信息,生成企业画像报告
大数据·开发语言·python·自动化·php·deepseek·openclaw
FoldWinCard3 小时前
D6 Python 基础语法 --- 保留关键字
开发语言·python
暮暮祈安4 小时前
Celery 新手入门指南
java·数据库·python·flask·httpx
JustNow_Man4 小时前
每日高频场景英文口语
python
风吹心凉5 小时前
python3基础2026.7.22
开发语言·python
小猴子爱上树6 小时前
一站式解决图片视频翻译难题的高效AI工具
人工智能·python·音视频
weixin_538601976 小时前
智能体测开Day31
开发语言·python
lbb 小魔仙6 小时前
Git + Python 项目工作流最佳实践:pre-commit、CI、CHANGELOG 自动化
git·python·ci/cd