Files
fastapi-best-architecture/backend/database/redis.py
T

191 lines
6.7 KiB
Python

import sys
from redis.asyncio import BlockingConnectionPool, Redis
from redis.exceptions import AuthenticationError, TimeoutError
from backend.common.log import log
from backend.core.conf import settings
class RedisCli(Redis):
"""Redis 客户端"""
def __init__(
self,
host: str = settings.REDIS_HOST,
port: int = settings.REDIS_PORT,
password: str = settings.REDIS_PASSWORD,
db: int = settings.REDIS_DATABASE,
socket_timeout: int | None = settings.REDIS_TIMEOUT,
socket_connect_timeout: int = settings.REDIS_TIMEOUT,
*,
socket_keepalive: bool = True,
health_check_interval: int = 30,
decode_responses: bool = True,
max_connections: int = settings.REDIS_MAX_CONNECTIONS,
pool_timeout: int = settings.REDIS_POOL_TIMEOUT,
) -> None:
"""
初始化 Redis 客户端
:param host: Redis 服务器的主机地址
:param port: Redis 服务器的端口号
:param password: Redis 认证密码
:param db: 使用的 Redis 逻辑数据库索引
:param socket_timeout: Socket 读写操作的超时时间
:param socket_connect_timeout: 建立 TCP 连接时的超时时间
:param socket_keepalive: 是否开启 TCP Keepalive 探测
:param health_check_interval: 健康检查间隔时间(秒)
:param decode_responses: 是否自动将 Redis 返回的字节流(bytes)解码为字符串(utf-8)
:param max_connections: 连接池最大连接数,超出后排队等待而不是无限新建
:param pool_timeout: 等待空闲连接的超时时间(秒)
"""
pool = BlockingConnectionPool(
max_connections=max_connections,
timeout=pool_timeout,
host=host,
port=port,
password=password,
db=db,
socket_timeout=socket_timeout,
socket_connect_timeout=socket_connect_timeout,
socket_keepalive=socket_keepalive,
health_check_interval=health_check_interval,
decode_responses=decode_responses,
)
super().__init__(connection_pool=pool)
# 连接池由客户端独占,aclose 时一并释放
self.auto_close_connection_pool = True
async def init(self) -> None:
"""初始化 Redis 服务器"""
try:
await self.ping()
except TimeoutError:
log.error('Redis 服务器连接超时')
sys.exit()
except AuthenticationError:
log.error('Redis 服务器连接认证失败')
sys.exit()
except Exception as e:
log.error('Redis 服务器连接异常 {}', e)
sys.exit()
async def delete_by_prefix(
self,
key_prefix: str,
exclude_keys: str | list[str] | None = None,
batch_size: int = 1000,
count: int = 1000,
) -> None:
"""
删除指定前缀的所有 key
:param key_prefix: 要删除的键前缀
:param exclude_keys: 要排除的键或键列表
:param batch_size: 批量删除的大小
:param count: 每次扫描批次的数量
:return:
"""
exclude_set = (
set(exclude_keys)
if isinstance(exclude_keys, list)
else {exclude_keys}
if isinstance(exclude_keys, str)
else set()
)
batch_keys = []
if key_prefix not in exclude_set and await self.exists(key_prefix):
batch_keys.append(key_prefix)
async for key in self.scan_iter(match=f'{key_prefix}:*', count=count):
if key not in exclude_set:
batch_keys.append(key)
if len(batch_keys) >= batch_size:
await self.delete(*batch_keys)
batch_keys.clear()
if batch_keys:
await self.delete(*batch_keys)
async def get_by_prefix(self, key_prefix: str, count: int = 1000) -> list[str]:
"""
获取指定前缀的所有 key
:param key_prefix: 要搜索的键前缀
:param count: 每次扫描批次的数量,值越大扫描速度越快,但会占用更多服务器资源
:return:
"""
return [key async for key in self.scan_iter(match=f'{key_prefix}:*', count=count)]
async def mget_batched(self, keys: list[str], batch_size: int = 1000) -> list[str | None]:
"""
分批获取多个 key 的值
:param keys: 键列表
:param batch_size: 每批数量
:return:
"""
if batch_size <= 0:
raise ValueError('batch_size 必须大于 0')
if not keys:
return []
values: list[str | None] = []
for index in range(0, len(keys), batch_size):
values.extend(await self.mget(keys[index : index + batch_size]))
return values
async def exists_batched(self, keys: list[str], batch_size: int = 1000) -> list[bool]:
"""
分批判断多个 key 是否存在
:param keys: 键列表
:param batch_size: 每批数量
:return:
"""
if batch_size <= 0:
raise ValueError('batch_size 必须大于 0')
return [value is not None for value in await self.mget_batched(keys, batch_size=batch_size)]
async def smembers_many(self, keys: list[str], batch_size: int = 100) -> list[set[str]]:
"""
分批获取多个集合的成员
:param keys: 键列表
:param batch_size: 每批数量
:return:
"""
if batch_size <= 0:
raise ValueError('batch_size 必须大于 0')
if not keys:
return []
members: list[set[str]] = []
for index in range(0, len(keys), batch_size):
batch = keys[index : index + batch_size]
async with self.pipeline(transaction=False) as pipe:
for key in batch:
pipe.smembers(key)
results = await pipe.execute()
members.extend(set(result) if result else set() for result in results)
return members
async def delete_batched(self, keys: list[str], batch_size: int = 1000) -> int:
"""
分批删除多个 key
:param keys: 键列表
:param batch_size: 每批数量
:return:
"""
if batch_size <= 0:
raise ValueError('batch_size 必须大于 0')
if not keys:
return 0
deleted = 0
for index in range(0, len(keys), batch_size):
batch = keys[index : index + batch_size]
deleted += await self.delete(*batch)
return deleted
# 创建 redis 客户端单例
redis_client: RedisCli = RedisCli()