mirror of
https://github.com/fastapi-practices/fastapi-best-architecture.git
synced 2026-10-04 09:11:03 +00:00
191 lines
6.7 KiB
Python
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()
|