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:
@@ -33,6 +33,7 @@ REDIS_USERNAME=
|
||||
REDIS_PASSWORD=difyai123456
|
||||
REDIS_USE_SSL=false
|
||||
REDIS_DB=0
|
||||
REDIS_KEY_PREFIX=
|
||||
|
||||
# PostgreSQL database configuration
|
||||
DB_USERNAME=postgres
|
||||
|
||||
@@ -236,6 +236,41 @@ def test_pubsub_redis_url_required_when_default_unavailable(monkeypatch: pytest.
|
||||
_ = DifyConfig().normalized_pubsub_redis_url
|
||||
|
||||
|
||||
def test_dify_config_exposes_redis_key_prefix_default(monkeypatch: pytest.MonkeyPatch):
|
||||
os.environ.clear()
|
||||
|
||||
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
|
||||
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
|
||||
monkeypatch.setenv("DB_TYPE", "postgresql")
|
||||
monkeypatch.setenv("DB_USERNAME", "postgres")
|
||||
monkeypatch.setenv("DB_PASSWORD", "postgres")
|
||||
monkeypatch.setenv("DB_HOST", "localhost")
|
||||
monkeypatch.setenv("DB_PORT", "5432")
|
||||
monkeypatch.setenv("DB_DATABASE", "dify")
|
||||
|
||||
config = DifyConfig(_env_file=None)
|
||||
|
||||
assert config.REDIS_KEY_PREFIX == ""
|
||||
|
||||
|
||||
def test_dify_config_reads_redis_key_prefix_from_env(monkeypatch: pytest.MonkeyPatch):
|
||||
os.environ.clear()
|
||||
|
||||
monkeypatch.setenv("CONSOLE_API_URL", "https://example.com")
|
||||
monkeypatch.setenv("CONSOLE_WEB_URL", "https://example.com")
|
||||
monkeypatch.setenv("DB_TYPE", "postgresql")
|
||||
monkeypatch.setenv("DB_USERNAME", "postgres")
|
||||
monkeypatch.setenv("DB_PASSWORD", "postgres")
|
||||
monkeypatch.setenv("DB_HOST", "localhost")
|
||||
monkeypatch.setenv("DB_PORT", "5432")
|
||||
monkeypatch.setenv("DB_DATABASE", "dify")
|
||||
monkeypatch.setenv("REDIS_KEY_PREFIX", "enterprise-a")
|
||||
|
||||
config = DifyConfig(_env_file=None)
|
||||
|
||||
assert config.REDIS_KEY_PREFIX == "enterprise-a"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("broker_url", "expected_host", "expected_port", "expected_username", "expected_password", "expected_db"),
|
||||
[
|
||||
|
||||
@@ -7,6 +7,47 @@ from unittest.mock import MagicMock, patch
|
||||
class TestCelerySSLConfiguration:
|
||||
"""Test suite for Celery SSL configuration."""
|
||||
|
||||
def test_get_celery_broker_transport_options_includes_global_keyprefix_for_redis(self):
|
||||
mock_config = MagicMock()
|
||||
mock_config.CELERY_USE_SENTINEL = False
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
with patch("extensions.ext_celery.dify_config", mock_config):
|
||||
from extensions.ext_celery import get_celery_broker_transport_options
|
||||
|
||||
result = get_celery_broker_transport_options()
|
||||
|
||||
assert result["global_keyprefix"] == "enterprise-a:"
|
||||
|
||||
def test_get_celery_broker_transport_options_omits_global_keyprefix_when_prefix_empty(self):
|
||||
mock_config = MagicMock()
|
||||
mock_config.CELERY_USE_SENTINEL = False
|
||||
mock_config.REDIS_KEY_PREFIX = " "
|
||||
|
||||
with patch("extensions.ext_celery.dify_config", mock_config):
|
||||
from extensions.ext_celery import get_celery_broker_transport_options
|
||||
|
||||
result = get_celery_broker_transport_options()
|
||||
|
||||
assert "global_keyprefix" not in result
|
||||
|
||||
def test_get_celery_broker_transport_options_keeps_sentinel_and_adds_global_keyprefix(self):
|
||||
mock_config = MagicMock()
|
||||
mock_config.CELERY_USE_SENTINEL = True
|
||||
mock_config.CELERY_SENTINEL_MASTER_NAME = "mymaster"
|
||||
mock_config.CELERY_SENTINEL_SOCKET_TIMEOUT = 0.1
|
||||
mock_config.CELERY_SENTINEL_PASSWORD = "secret"
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
with patch("extensions.ext_celery.dify_config", mock_config):
|
||||
from extensions.ext_celery import get_celery_broker_transport_options
|
||||
|
||||
result = get_celery_broker_transport_options()
|
||||
|
||||
assert result["master_name"] == "mymaster"
|
||||
assert result["sentinel_kwargs"]["password"] == "secret"
|
||||
assert result["global_keyprefix"] == "enterprise-a:"
|
||||
|
||||
def test_get_celery_ssl_options_when_ssl_disabled(self):
|
||||
"""Test SSL options when BROKER_USE_SSL is False."""
|
||||
from configs import DifyConfig
|
||||
@@ -151,3 +192,49 @@ class TestCelerySSLConfiguration:
|
||||
# Check that SSL is also applied to Redis backend
|
||||
assert "redis_backend_use_ssl" in celery_app.conf
|
||||
assert celery_app.conf["redis_backend_use_ssl"] is not None
|
||||
|
||||
def test_celery_init_applies_global_keyprefix_to_broker_and_backend_transport(self):
|
||||
mock_config = MagicMock()
|
||||
mock_config.BROKER_USE_SSL = False
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
mock_config.HUMAN_INPUT_TIMEOUT_TASK_INTERVAL = 1
|
||||
mock_config.CELERY_BROKER_URL = "redis://localhost:6379/0"
|
||||
mock_config.CELERY_BACKEND = "redis"
|
||||
mock_config.CELERY_RESULT_BACKEND = "redis://localhost:6379/0"
|
||||
mock_config.CELERY_USE_SENTINEL = False
|
||||
mock_config.LOG_FORMAT = "%(message)s"
|
||||
mock_config.LOG_TZ = "UTC"
|
||||
mock_config.LOG_FILE = None
|
||||
mock_config.CELERY_TASK_ANNOTATIONS = {}
|
||||
|
||||
mock_config.CELERY_BEAT_SCHEDULER_TIME = 1
|
||||
mock_config.ENABLE_CLEAN_EMBEDDING_CACHE_TASK = False
|
||||
mock_config.ENABLE_CLEAN_UNUSED_DATASETS_TASK = False
|
||||
mock_config.ENABLE_CREATE_TIDB_SERVERLESS_TASK = False
|
||||
mock_config.ENABLE_UPDATE_TIDB_SERVERLESS_STATUS_TASK = False
|
||||
mock_config.ENABLE_CLEAN_MESSAGES = False
|
||||
mock_config.ENABLE_MAIL_CLEAN_DOCUMENT_NOTIFY_TASK = False
|
||||
mock_config.ENABLE_DATASETS_QUEUE_MONITOR = False
|
||||
mock_config.ENABLE_HUMAN_INPUT_TIMEOUT_TASK = False
|
||||
mock_config.ENABLE_CHECK_UPGRADABLE_PLUGIN_TASK = False
|
||||
mock_config.MARKETPLACE_ENABLED = False
|
||||
mock_config.WORKFLOW_LOG_CLEANUP_ENABLED = False
|
||||
mock_config.ENABLE_WORKFLOW_RUN_CLEANUP_TASK = False
|
||||
mock_config.ENABLE_WORKFLOW_SCHEDULE_POLLER_TASK = False
|
||||
mock_config.WORKFLOW_SCHEDULE_POLLER_INTERVAL = 1
|
||||
mock_config.ENABLE_TRIGGER_PROVIDER_REFRESH_TASK = False
|
||||
mock_config.TRIGGER_PROVIDER_REFRESH_INTERVAL = 15
|
||||
mock_config.ENABLE_API_TOKEN_LAST_USED_UPDATE_TASK = False
|
||||
mock_config.API_TOKEN_LAST_USED_UPDATE_INTERVAL = 30
|
||||
mock_config.ENTERPRISE_ENABLED = False
|
||||
mock_config.ENTERPRISE_TELEMETRY_ENABLED = False
|
||||
|
||||
with patch("extensions.ext_celery.dify_config", mock_config):
|
||||
from dify_app import DifyApp
|
||||
from extensions.ext_celery import init_app
|
||||
|
||||
app = DifyApp(__name__)
|
||||
celery_app = init_app(app)
|
||||
|
||||
assert celery_app.conf["broker_transport_options"]["global_keyprefix"] == "enterprise-a:"
|
||||
assert celery_app.conf["result_backend_transport_options"]["global_keyprefix"] == "enterprise-a:"
|
||||
|
||||
@@ -6,6 +6,7 @@ from libs.broadcast_channel.redis.sharded_channel import ShardedRedisBroadcastCh
|
||||
|
||||
def test_get_pubsub_broadcast_channel_defaults_to_pubsub(monkeypatch):
|
||||
monkeypatch.setattr(dify_config, "PUBSUB_REDIS_CHANNEL_TYPE", "pubsub")
|
||||
monkeypatch.setattr(ext_redis, "_pubsub_redis_client", object())
|
||||
|
||||
channel = ext_redis.get_pubsub_broadcast_channel()
|
||||
|
||||
@@ -14,6 +15,7 @@ def test_get_pubsub_broadcast_channel_defaults_to_pubsub(monkeypatch):
|
||||
|
||||
def test_get_pubsub_broadcast_channel_sharded(monkeypatch):
|
||||
monkeypatch.setattr(dify_config, "PUBSUB_REDIS_CHANNEL_TYPE", "sharded")
|
||||
monkeypatch.setattr(ext_redis, "_pubsub_redis_client", object())
|
||||
|
||||
channel = ext_redis.get_pubsub_broadcast_channel()
|
||||
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from redis import RedisError
|
||||
from redis.retry import Retry
|
||||
|
||||
from extensions.ext_redis import (
|
||||
RedisClientWrapper,
|
||||
_get_base_redis_params,
|
||||
_get_cluster_connection_health_params,
|
||||
_get_connection_health_params,
|
||||
_normalize_redis_key_prefix,
|
||||
_serialize_redis_name,
|
||||
redis_fallback,
|
||||
)
|
||||
|
||||
@@ -123,3 +126,99 @@ class TestRedisFallback:
|
||||
|
||||
assert test_func.__name__ == "test_func"
|
||||
assert test_func.__doc__ == "Test function docstring"
|
||||
|
||||
|
||||
class TestRedisKeyPrefixHelpers:
|
||||
def test_normalize_redis_key_prefix_trims_whitespace(self):
|
||||
assert _normalize_redis_key_prefix(" enterprise-a ") == "enterprise-a"
|
||||
|
||||
def test_normalize_redis_key_prefix_treats_whitespace_only_as_empty(self):
|
||||
assert _normalize_redis_key_prefix(" ") == ""
|
||||
|
||||
def test_serialize_redis_name_returns_original_when_prefix_empty(self):
|
||||
assert _serialize_redis_name("model_lb_index:test", "") == "model_lb_index:test"
|
||||
|
||||
def test_serialize_redis_name_adds_single_colon_separator(self):
|
||||
assert _serialize_redis_name("model_lb_index:test", "enterprise-a") == "enterprise-a:model_lb_index:test"
|
||||
|
||||
|
||||
class TestRedisClientWrapperKeyPrefix:
|
||||
def test_wrapper_get_prefixes_string_keys(self):
|
||||
mock_client = MagicMock()
|
||||
wrapper = RedisClientWrapper()
|
||||
wrapper.initialize(mock_client)
|
||||
|
||||
with patch("extensions.ext_redis.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
wrapper.get("oauth_state:abc")
|
||||
|
||||
mock_client.get.assert_called_once_with("enterprise-a:oauth_state:abc")
|
||||
|
||||
def test_wrapper_delete_prefixes_multiple_keys(self):
|
||||
mock_client = MagicMock()
|
||||
wrapper = RedisClientWrapper()
|
||||
wrapper.initialize(mock_client)
|
||||
|
||||
with patch("extensions.ext_redis.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
wrapper.delete("key:a", "key:b")
|
||||
|
||||
mock_client.delete.assert_called_once_with("enterprise-a:key:a", "enterprise-a:key:b")
|
||||
|
||||
def test_wrapper_lock_prefixes_lock_name(self):
|
||||
mock_client = MagicMock()
|
||||
wrapper = RedisClientWrapper()
|
||||
wrapper.initialize(mock_client)
|
||||
|
||||
with patch("extensions.ext_redis.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
wrapper.lock("resource-lock", timeout=10)
|
||||
|
||||
mock_client.lock.assert_called_once()
|
||||
args, kwargs = mock_client.lock.call_args
|
||||
assert args == ("enterprise-a:resource-lock",)
|
||||
assert kwargs["timeout"] == 10
|
||||
|
||||
def test_wrapper_hash_operations_prefix_key_name(self):
|
||||
mock_client = MagicMock()
|
||||
wrapper = RedisClientWrapper()
|
||||
wrapper.initialize(mock_client)
|
||||
|
||||
with patch("extensions.ext_redis.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
wrapper.hset("hash:key", "field", "value")
|
||||
wrapper.hgetall("hash:key")
|
||||
|
||||
mock_client.hset.assert_called_once_with("enterprise-a:hash:key", "field", "value")
|
||||
mock_client.hgetall.assert_called_once_with("enterprise-a:hash:key")
|
||||
|
||||
def test_wrapper_zadd_prefixes_sorted_set_name(self):
|
||||
mock_client = MagicMock()
|
||||
wrapper = RedisClientWrapper()
|
||||
wrapper.initialize(mock_client)
|
||||
|
||||
with patch("extensions.ext_redis.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
wrapper.zadd("zset:key", {"member": 1})
|
||||
|
||||
mock_client.zadd.assert_called_once()
|
||||
args, kwargs = mock_client.zadd.call_args
|
||||
assert args == ("enterprise-a:zset:key", {"member": 1})
|
||||
assert kwargs["nx"] is False
|
||||
|
||||
def test_wrapper_preserves_keys_when_prefix_is_empty(self):
|
||||
mock_client = MagicMock()
|
||||
wrapper = RedisClientWrapper()
|
||||
wrapper.initialize(mock_client)
|
||||
|
||||
with patch("extensions.ext_redis.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = " "
|
||||
|
||||
wrapper.get("plain:key")
|
||||
|
||||
mock_client.get.assert_called_once_with("plain:key")
|
||||
|
||||
@@ -139,6 +139,28 @@ class TestTopic:
|
||||
|
||||
mock_redis_client.publish.assert_called_once_with("test-topic", payload)
|
||||
|
||||
def test_publish_prefixes_regular_topic(self, mock_redis_client: MagicMock):
|
||||
with patch("extensions.redis_names.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
topic = Topic(mock_redis_client, "test-topic")
|
||||
|
||||
topic.publish(b"test message")
|
||||
|
||||
mock_redis_client.publish.assert_called_once_with("enterprise-a:test-topic", b"test message")
|
||||
|
||||
def test_subscribe_prefixes_regular_topic(self, mock_redis_client: MagicMock):
|
||||
with patch("extensions.redis_names.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
topic = Topic(mock_redis_client, "test-topic")
|
||||
|
||||
subscription = topic.subscribe()
|
||||
try:
|
||||
subscription._start_if_needed()
|
||||
finally:
|
||||
subscription.close()
|
||||
|
||||
mock_redis_client.pubsub.return_value.subscribe.assert_called_once_with("enterprise-a:test-topic")
|
||||
|
||||
|
||||
class TestShardedTopic:
|
||||
"""Test cases for the ShardedTopic class."""
|
||||
@@ -176,6 +198,15 @@ class TestShardedTopic:
|
||||
|
||||
mock_redis_client.spublish.assert_called_once_with("test-sharded-topic", payload)
|
||||
|
||||
def test_publish_prefixes_sharded_topic(self, mock_redis_client: MagicMock):
|
||||
with patch("extensions.redis_names.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
sharded_topic = ShardedTopic(mock_redis_client, "test-sharded-topic")
|
||||
|
||||
sharded_topic.publish(b"test sharded message")
|
||||
|
||||
mock_redis_client.spublish.assert_called_once_with("enterprise-a:test-sharded-topic", b"test sharded message")
|
||||
|
||||
def test_subscribe_returns_sharded_subscription(self, sharded_topic: ShardedTopic, mock_redis_client: MagicMock):
|
||||
"""Test that subscribe() returns a _RedisShardedSubscription instance."""
|
||||
subscription = sharded_topic.subscribe()
|
||||
@@ -185,6 +216,19 @@ class TestShardedTopic:
|
||||
assert subscription._pubsub is mock_redis_client.pubsub.return_value
|
||||
assert subscription._topic == "test-sharded-topic"
|
||||
|
||||
def test_subscribe_prefixes_sharded_topic(self, mock_redis_client: MagicMock):
|
||||
with patch("extensions.redis_names.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
sharded_topic = ShardedTopic(mock_redis_client, "test-sharded-topic")
|
||||
|
||||
subscription = sharded_topic.subscribe()
|
||||
try:
|
||||
subscription._start_if_needed()
|
||||
finally:
|
||||
subscription.close()
|
||||
|
||||
mock_redis_client.pubsub.return_value.ssubscribe.assert_called_once_with("enterprise-a:test-sharded-topic")
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class SubscriptionTestCase:
|
||||
|
||||
@@ -2,6 +2,7 @@ import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import cast
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -150,6 +151,25 @@ class TestStreamsBroadcastChannel:
|
||||
# Expire called after publish
|
||||
assert fake_redis._expire_calls.get("stream:beta", 0) >= 1
|
||||
|
||||
def test_topic_uses_prefixed_stream_key(self, fake_redis: FakeStreamsRedis):
|
||||
with patch("extensions.redis_names.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
|
||||
topic = StreamsBroadcastChannel(fake_redis, retention_seconds=60).topic("alpha")
|
||||
|
||||
assert topic._topic == "alpha"
|
||||
assert topic._key == "enterprise-a:stream:alpha"
|
||||
|
||||
def test_publish_uses_prefixed_stream_key(self, fake_redis: FakeStreamsRedis):
|
||||
with patch("extensions.redis_names.dify_config") as mock_config:
|
||||
mock_config.REDIS_KEY_PREFIX = "enterprise-a"
|
||||
topic = StreamsBroadcastChannel(fake_redis, retention_seconds=60).topic("beta")
|
||||
|
||||
topic.publish(b"hello")
|
||||
|
||||
assert fake_redis._store["enterprise-a:stream:beta"][0][1] == {b"data": b"hello"}
|
||||
assert fake_redis._expire_calls.get("enterprise-a:stream:beta", 0) >= 1
|
||||
|
||||
def test_topic_exposes_self_as_producer_and_subscriber(self, streams_channel: StreamsBroadcastChannel):
|
||||
topic = streams_channel.topic("producer-subscriber")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user