mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
feat: support configurable redis key prefix (#35139)
This commit is contained in:
@@ -9,6 +9,7 @@ from typing_extensions import TypedDict
|
||||
|
||||
from configs import dify_config
|
||||
from dify_app import DifyApp
|
||||
from extensions.redis_names import normalize_redis_key_prefix
|
||||
|
||||
|
||||
class _CelerySentinelKwargsDict(TypedDict):
|
||||
@@ -16,9 +17,10 @@ class _CelerySentinelKwargsDict(TypedDict):
|
||||
password: str | None
|
||||
|
||||
|
||||
class CelerySentinelTransportDict(TypedDict):
|
||||
class CelerySentinelTransportDict(TypedDict, total=False):
|
||||
master_name: str | None
|
||||
sentinel_kwargs: _CelerySentinelKwargsDict
|
||||
global_keyprefix: str
|
||||
|
||||
|
||||
class CelerySSLOptionsDict(TypedDict):
|
||||
@@ -61,15 +63,31 @@ def get_celery_ssl_options() -> CelerySSLOptionsDict | None:
|
||||
|
||||
def get_celery_broker_transport_options() -> CelerySentinelTransportDict | dict[str, Any]:
|
||||
"""Get broker transport options (e.g. Redis Sentinel) for Celery connections."""
|
||||
transport_options: CelerySentinelTransportDict | dict[str, Any]
|
||||
if dify_config.CELERY_USE_SENTINEL:
|
||||
return CelerySentinelTransportDict(
|
||||
transport_options = CelerySentinelTransportDict(
|
||||
master_name=dify_config.CELERY_SENTINEL_MASTER_NAME,
|
||||
sentinel_kwargs=_CelerySentinelKwargsDict(
|
||||
socket_timeout=dify_config.CELERY_SENTINEL_SOCKET_TIMEOUT,
|
||||
password=dify_config.CELERY_SENTINEL_PASSWORD,
|
||||
),
|
||||
)
|
||||
return {}
|
||||
else:
|
||||
transport_options = {}
|
||||
|
||||
global_keyprefix = get_celery_redis_global_keyprefix()
|
||||
if global_keyprefix:
|
||||
transport_options["global_keyprefix"] = global_keyprefix
|
||||
|
||||
return transport_options
|
||||
|
||||
|
||||
def get_celery_redis_global_keyprefix() -> str | None:
|
||||
"""Return the Redis transport prefix for Celery when namespace isolation is enabled."""
|
||||
normalized_prefix = normalize_redis_key_prefix(dify_config.REDIS_KEY_PREFIX)
|
||||
if not normalized_prefix:
|
||||
return None
|
||||
return f"{normalized_prefix}:"
|
||||
|
||||
|
||||
def init_app(app: DifyApp) -> Celery:
|
||||
|
||||
+153
-64
@@ -3,7 +3,7 @@ import logging
|
||||
import ssl
|
||||
from collections.abc import Callable
|
||||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import Any, Union, cast
|
||||
|
||||
import redis
|
||||
from redis import RedisError
|
||||
@@ -18,17 +18,26 @@ from typing_extensions import TypedDict
|
||||
|
||||
from configs import dify_config
|
||||
from dify_app import DifyApp
|
||||
from extensions.redis_names import (
|
||||
normalize_redis_key_prefix,
|
||||
serialize_redis_name,
|
||||
serialize_redis_name_arg,
|
||||
serialize_redis_name_args,
|
||||
)
|
||||
from libs.broadcast_channel.channel import BroadcastChannel as BroadcastChannelProtocol
|
||||
from libs.broadcast_channel.redis.channel import BroadcastChannel as RedisBroadcastChannel
|
||||
from libs.broadcast_channel.redis.sharded_channel import ShardedRedisBroadcastChannel
|
||||
from libs.broadcast_channel.redis.streams_channel import StreamsBroadcastChannel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from redis.lock import Lock
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_normalize_redis_key_prefix = normalize_redis_key_prefix
|
||||
_serialize_redis_name = serialize_redis_name
|
||||
_serialize_redis_name_arg = serialize_redis_name_arg
|
||||
_serialize_redis_name_args = serialize_redis_name_args
|
||||
|
||||
|
||||
class RedisClientWrapper:
|
||||
"""
|
||||
A wrapper class for the Redis client that addresses the issue where the global
|
||||
@@ -59,68 +68,148 @@ class RedisClientWrapper:
|
||||
if self._client is None:
|
||||
self._client = client
|
||||
|
||||
if TYPE_CHECKING:
|
||||
# Type hints for IDE support and static analysis
|
||||
# These are not executed at runtime but provide type information
|
||||
def get(self, name: str | bytes) -> Any: ...
|
||||
|
||||
def set(
|
||||
self,
|
||||
name: str | bytes,
|
||||
value: Any,
|
||||
ex: int | None = None,
|
||||
px: int | None = None,
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
keepttl: bool = False,
|
||||
get: bool = False,
|
||||
exat: int | None = None,
|
||||
pxat: int | None = None,
|
||||
) -> Any: ...
|
||||
|
||||
def setex(self, name: str | bytes, time: int | timedelta, value: Any) -> Any: ...
|
||||
def setnx(self, name: str | bytes, value: Any) -> Any: ...
|
||||
def delete(self, *names: str | bytes) -> Any: ...
|
||||
def incr(self, name: str | bytes, amount: int = 1) -> Any: ...
|
||||
def expire(
|
||||
self,
|
||||
name: str | bytes,
|
||||
time: int | timedelta,
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
gt: bool = False,
|
||||
lt: bool = False,
|
||||
) -> Any: ...
|
||||
def lock(
|
||||
self,
|
||||
name: str,
|
||||
timeout: float | None = None,
|
||||
sleep: float = 0.1,
|
||||
blocking: bool = True,
|
||||
blocking_timeout: float | None = None,
|
||||
thread_local: bool = True,
|
||||
) -> Lock: ...
|
||||
def zadd(
|
||||
self,
|
||||
name: str | bytes,
|
||||
mapping: dict[str | bytes | int | float, float | int | str | bytes],
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
ch: bool = False,
|
||||
incr: bool = False,
|
||||
gt: bool = False,
|
||||
lt: bool = False,
|
||||
) -> Any: ...
|
||||
def zremrangebyscore(self, name: str | bytes, min: float | str, max: float | str) -> Any: ...
|
||||
def zcard(self, name: str | bytes) -> Any: ...
|
||||
def getdel(self, name: str | bytes) -> Any: ...
|
||||
def pubsub(self) -> PubSub: ...
|
||||
def pipeline(self, transaction: bool = True, shard_hint: str | None = None) -> Any: ...
|
||||
|
||||
def __getattr__(self, item: str) -> Any:
|
||||
def _require_client(self) -> redis.Redis | RedisCluster:
|
||||
if self._client is None:
|
||||
raise RuntimeError("Redis client is not initialized. Call init_app first.")
|
||||
return getattr(self._client, item)
|
||||
return self._client
|
||||
|
||||
def _get_prefix(self) -> str:
|
||||
return dify_config.REDIS_KEY_PREFIX
|
||||
|
||||
def get(self, name: str | bytes) -> Any:
|
||||
return self._require_client().get(_serialize_redis_name_arg(name, self._get_prefix()))
|
||||
|
||||
def set(
|
||||
self,
|
||||
name: str | bytes,
|
||||
value: Any,
|
||||
ex: int | None = None,
|
||||
px: int | None = None,
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
keepttl: bool = False,
|
||||
get: bool = False,
|
||||
exat: int | None = None,
|
||||
pxat: int | None = None,
|
||||
) -> Any:
|
||||
return self._require_client().set(
|
||||
_serialize_redis_name_arg(name, self._get_prefix()),
|
||||
value,
|
||||
ex=ex,
|
||||
px=px,
|
||||
nx=nx,
|
||||
xx=xx,
|
||||
keepttl=keepttl,
|
||||
get=get,
|
||||
exat=exat,
|
||||
pxat=pxat,
|
||||
)
|
||||
|
||||
def setex(self, name: str | bytes, time: int | timedelta, value: Any) -> Any:
|
||||
return self._require_client().setex(_serialize_redis_name_arg(name, self._get_prefix()), time, value)
|
||||
|
||||
def setnx(self, name: str | bytes, value: Any) -> Any:
|
||||
return self._require_client().setnx(_serialize_redis_name_arg(name, self._get_prefix()), value)
|
||||
|
||||
def delete(self, *names: str | bytes) -> Any:
|
||||
return self._require_client().delete(*_serialize_redis_name_args(names, self._get_prefix()))
|
||||
|
||||
def incr(self, name: str | bytes, amount: int = 1) -> Any:
|
||||
return self._require_client().incr(_serialize_redis_name_arg(name, self._get_prefix()), amount)
|
||||
|
||||
def expire(
|
||||
self,
|
||||
name: str | bytes,
|
||||
time: int | timedelta,
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
gt: bool = False,
|
||||
lt: bool = False,
|
||||
) -> Any:
|
||||
return self._require_client().expire(
|
||||
_serialize_redis_name_arg(name, self._get_prefix()),
|
||||
time,
|
||||
nx=nx,
|
||||
xx=xx,
|
||||
gt=gt,
|
||||
lt=lt,
|
||||
)
|
||||
|
||||
def exists(self, *names: str | bytes) -> Any:
|
||||
return self._require_client().exists(*_serialize_redis_name_args(names, self._get_prefix()))
|
||||
|
||||
def ttl(self, name: str | bytes) -> Any:
|
||||
return self._require_client().ttl(_serialize_redis_name_arg(name, self._get_prefix()))
|
||||
|
||||
def getdel(self, name: str | bytes) -> Any:
|
||||
return self._require_client().getdel(_serialize_redis_name_arg(name, self._get_prefix()))
|
||||
|
||||
def lock(
|
||||
self,
|
||||
name: str,
|
||||
timeout: float | None = None,
|
||||
sleep: float = 0.1,
|
||||
blocking: bool = True,
|
||||
blocking_timeout: float | None = None,
|
||||
thread_local: bool = True,
|
||||
) -> Any:
|
||||
return self._require_client().lock(
|
||||
_serialize_redis_name(name, self._get_prefix()),
|
||||
timeout=timeout,
|
||||
sleep=sleep,
|
||||
blocking=blocking,
|
||||
blocking_timeout=blocking_timeout,
|
||||
thread_local=thread_local,
|
||||
)
|
||||
|
||||
def hset(self, name: str | bytes, *args: Any, **kwargs: Any) -> Any:
|
||||
return self._require_client().hset(_serialize_redis_name_arg(name, self._get_prefix()), *args, **kwargs)
|
||||
|
||||
def hgetall(self, name: str | bytes) -> Any:
|
||||
return self._require_client().hgetall(_serialize_redis_name_arg(name, self._get_prefix()))
|
||||
|
||||
def hdel(self, name: str | bytes, *keys: str | bytes) -> Any:
|
||||
return self._require_client().hdel(_serialize_redis_name_arg(name, self._get_prefix()), *keys)
|
||||
|
||||
def hlen(self, name: str | bytes) -> Any:
|
||||
return self._require_client().hlen(_serialize_redis_name_arg(name, self._get_prefix()))
|
||||
|
||||
def zadd(
|
||||
self,
|
||||
name: str | bytes,
|
||||
mapping: dict[str | bytes | int | float, float | int | str | bytes],
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
ch: bool = False,
|
||||
incr: bool = False,
|
||||
gt: bool = False,
|
||||
lt: bool = False,
|
||||
) -> Any:
|
||||
return self._require_client().zadd(
|
||||
_serialize_redis_name_arg(name, self._get_prefix()),
|
||||
cast(Any, mapping),
|
||||
nx=nx,
|
||||
xx=xx,
|
||||
ch=ch,
|
||||
incr=incr,
|
||||
gt=gt,
|
||||
lt=lt,
|
||||
)
|
||||
|
||||
def zremrangebyscore(self, name: str | bytes, min: float | str, max: float | str) -> Any:
|
||||
return self._require_client().zremrangebyscore(_serialize_redis_name_arg(name, self._get_prefix()), min, max)
|
||||
|
||||
def zcard(self, name: str | bytes) -> Any:
|
||||
return self._require_client().zcard(_serialize_redis_name_arg(name, self._get_prefix()))
|
||||
|
||||
def pubsub(self) -> PubSub:
|
||||
return self._require_client().pubsub()
|
||||
|
||||
def pipeline(self, transaction: bool = True, shard_hint: str | None = None) -> Any:
|
||||
return self._require_client().pipeline(transaction=transaction, shard_hint=shard_hint)
|
||||
|
||||
def __getattr__(self, item: str) -> Any:
|
||||
return getattr(self._require_client(), item)
|
||||
|
||||
|
||||
redis_client: RedisClientWrapper = RedisClientWrapper()
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
from configs import dify_config
|
||||
|
||||
|
||||
def normalize_redis_key_prefix(prefix: str | None) -> str:
|
||||
"""Normalize the configured Redis key prefix for consistent runtime use."""
|
||||
if prefix is None:
|
||||
return ""
|
||||
return prefix.strip()
|
||||
|
||||
|
||||
def get_redis_key_prefix() -> str:
|
||||
"""Read and normalize the current Redis key prefix from config."""
|
||||
return normalize_redis_key_prefix(dify_config.REDIS_KEY_PREFIX)
|
||||
|
||||
|
||||
def serialize_redis_name(name: str, prefix: str | None = None) -> str:
|
||||
"""Convert a logical Redis name into the physical name used in Redis."""
|
||||
normalized_prefix = get_redis_key_prefix() if prefix is None else normalize_redis_key_prefix(prefix)
|
||||
if not normalized_prefix:
|
||||
return name
|
||||
return f"{normalized_prefix}:{name}"
|
||||
|
||||
|
||||
def serialize_redis_name_arg(name: str | bytes, prefix: str | None = None) -> str | bytes:
|
||||
"""Prefix string Redis names while preserving bytes inputs unchanged."""
|
||||
if isinstance(name, bytes):
|
||||
return name
|
||||
return serialize_redis_name(name, prefix)
|
||||
|
||||
|
||||
def serialize_redis_name_args(names: tuple[str | bytes, ...], prefix: str | None = None) -> tuple[str | bytes, ...]:
|
||||
return tuple(serialize_redis_name_arg(name, prefix) for name in names)
|
||||
Reference in New Issue
Block a user