fix: preserve http_auth in _safe_deepcopy_config for OpenSearch (#3580) (#4418)

Co-authored-by: utkarsh240799 <utkarsh240799@users.noreply.github.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Utkarsh
2026-03-19 16:37:56 +05:30
committed by GitHub
parent a029cc9d43
commit 46b4b2e9c8
3 changed files with 444 additions and 14 deletions
+67 -10
View File
@@ -51,18 +51,77 @@ warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*swigva
logger = logging.getLogger(__name__)
# Fields that hold runtime auth/connection objects and must be preserved.
# These are non-serializable objects (e.g. AWSV4SignerAuth, RequestsHttpConnection)
# needed by clients like OpenSearch — not sensitive strings to redact.
_RUNTIME_FIELDS = frozenset({
"http_auth",
"auth",
"connection_class",
"ssl_context",
"use_azure_credential",
})
# Fields that are known to contain sensitive secrets and must be redacted.
_SENSITIVE_FIELDS_EXACT = frozenset({
"api_key",
"secret_key",
"private_key",
"access_key",
"password",
"credentials",
"credential",
"secret",
"token",
"access_token",
"refresh_token",
"auth_token",
"session_token",
"client_secret",
"auth_client_secret",
"azure_client_secret",
"service_account_json",
"aws_session_token",
})
# Suffixes that indicate a field likely holds a secret value.
_SENSITIVE_SUFFIXES = (
"_password",
"_secret",
"_token",
"_credential",
"_credentials",
)
def _is_sensitive_field(field_name: str) -> bool:
"""Check if a field should be redacted for telemetry safety.
Uses a layered approach:
1. Runtime fields (allowlist) — always preserved, highest priority.
2. Exact deny list — known secret field names.
3. Suffix deny list — catches patterns like db_password, auth_secret, etc.
"""
name = field_name.lower().strip()
if name in _RUNTIME_FIELDS:
return False
if name in _SENSITIVE_FIELDS_EXACT:
return True
return any(name.endswith(suffix) for suffix in _SENSITIVE_SUFFIXES)
def _safe_deepcopy_config(config):
"""Safely deepcopy config, falling back to JSON serialization for non-serializable objects."""
"""Safely deepcopy config, falling back to dict-based cloning for non-serializable objects."""
try:
return deepcopy(config)
except Exception as e:
logger.debug(f"Deepcopy failed, using JSON serialization: {e}")
logger.debug(f"Deepcopy failed, using dict-based cloning: {e}")
config_class = type(config)
if hasattr(config, "model_dump"):
try:
clone_dict = config.model_dump(mode="json")
clone_dict = config.model_dump()
except Exception:
clone_dict = {k: v for k, v in config.__dict__.items()}
elif hasattr(config, "__dataclass_fields__"):
@@ -70,12 +129,11 @@ def _safe_deepcopy_config(config):
clone_dict = asdict(config)
else:
clone_dict = {k: v for k, v in config.__dict__.items()}
sensitive_tokens = ("auth", "credential", "password", "token", "secret", "key", "connection_class")
for field_name in list(clone_dict.keys()):
if any(token in field_name.lower() for token in sensitive_tokens):
if _is_sensitive_field(field_name):
clone_dict[field_name] = None
try:
return config_class(**clone_dict)
except Exception as reconstruction_error:
@@ -222,7 +280,6 @@ class Memory(MemoryBase):
telemetry_config_dict['collection_name'] = "mem0migrations"
# Set path for file-based vector stores
telemetry_config = _safe_deepcopy_config(self.config.vector_store.config)
if self.config.vector_store.provider in ["faiss", "qdrant"]:
provider_path = f"migrations_{self.config.vector_store.provider}"
telemetry_config_dict['path'] = os.path.join(mem0_dir, provider_path)
+371
View File
@@ -0,0 +1,371 @@
"""Tests for _safe_deepcopy_config and _is_sensitive_field (Issue #3580).
Validates that runtime auth objects (http_auth, connection_class, etc.) are
preserved while genuinely sensitive fields (password, api_key, etc.) are
redacted during config cloning for telemetry.
"""
import threading
from dataclasses import dataclass
import pytest
from mem0.memory.main import _is_sensitive_field, _safe_deepcopy_config
# ---------------------------------------------------------------------------
# _is_sensitive_field tests
# ---------------------------------------------------------------------------
class TestRuntimeFieldsPreserved:
"""Runtime/allowlist fields must NOT be considered sensitive."""
@pytest.mark.parametrize("field", [
"http_auth",
"auth",
"connection_class",
"ssl_context",
])
def test_runtime_fields_are_not_sensitive(self, field):
assert _is_sensitive_field(field) is False
def test_runtime_fields_case_insensitive(self):
assert _is_sensitive_field("HTTP_AUTH") is False
assert _is_sensitive_field("Connection_Class") is False
class TestExactDenyList:
"""Known secret field names must be redacted."""
@pytest.mark.parametrize("field", [
"api_key",
"secret_key",
"private_key",
"access_key",
"password",
"credentials",
"credential",
"secret",
"token",
"access_token",
"refresh_token",
"auth_token",
"session_token",
"client_secret",
"auth_client_secret",
"azure_client_secret",
"service_account_json",
"aws_session_token",
])
def test_exact_sensitive_fields(self, field):
assert _is_sensitive_field(field) is True
def test_exact_fields_case_insensitive(self):
assert _is_sensitive_field("API_KEY") is True
assert _is_sensitive_field("Password") is True
class TestSuffixDenyList:
"""Fields ending with sensitive suffixes must be redacted."""
@pytest.mark.parametrize("field", [
"db_password",
"user_password",
"redis_password",
"app_secret",
"client_secret",
"oauth_token",
"bearer_token",
"aws_credential",
"gcp_credentials",
])
def test_suffix_matches(self, field):
assert _is_sensitive_field(field) is True
class TestNonSensitiveFields:
"""Common config fields that must NOT be redacted."""
@pytest.mark.parametrize("field", [
"host",
"port",
"collection_name",
"embedding_model_dims",
"use_ssl",
"verify_certs",
"index_name",
"dimension",
"metric",
"path",
"url",
"timeout",
"pool_maxsize",
])
def test_common_config_fields(self, field):
assert _is_sensitive_field(field) is False
class TestOverMatchingPrevention:
"""Fields that previously matched due to broad substring matching
but should NOT be redacted."""
@pytest.mark.parametrize("field", [
"primary_key", # contains "key" but is a DB concept
"partition_key", # contains "key" but is a DB concept
"shard_key", # contains "key" but is a DB concept
"token_type", # contains "token" but is metadata
"token_count", # contains "token" but is a count
"tokenizer", # contains "token" but is a tool name
"key_space", # contains "key" but is a namespace
"keyboard", # contains "key" but is unrelated
"monkey", # contains "key" but is unrelated
"authenticate", # contains "auth" but is a verb
"authorization_url", # contains "auth" but is a URL
"credentials_path", # contains "credential" but is a file path
"secret_agent_name", # contains "secret" but is not a suffix match
])
def test_no_over_matching(self, field):
assert _is_sensitive_field(field) is False
class TestEdgeCases:
def test_empty_string(self):
assert _is_sensitive_field("") is False
def test_whitespace_stripped(self):
assert _is_sensitive_field(" api_key ") is True
assert _is_sensitive_field(" http_auth ") is False
class TestRealWorldFieldCoverage:
"""Verify behavior for actual field names from mem0 vector store configs."""
@pytest.mark.parametrize("field,expected", [
# OpenSearch
("password", True),
("api_key", True),
("http_auth", False),
("connection_class", False),
("host", False),
("port", False),
("verify_certs", False),
("use_ssl", False),
("pool_maxsize", False),
# Weaviate
("auth_client_secret", True),
# Databricks
("access_token", True),
("client_secret", True),
("azure_client_secret", True),
# Upstash / Milvus
("token", True),
# Vertex AI
("service_account_json", True),
("credentials_path", False),
# AWS
("aws_session_token", True),
# Azure MySQL
("use_azure_credential", False),
# General non-sensitive
("collection_name", False),
("embedding_model_dims", False),
("user", False),
("path", False),
("url", False),
("dimension", False),
("metric_type", False),
("batch_size", False),
("index_type", False),
])
def test_field_sensitivity(self, field, expected):
assert _is_sensitive_field(field) is expected
# ---------------------------------------------------------------------------
# _safe_deepcopy_config integration tests
# ---------------------------------------------------------------------------
class MockNonCopyableAuth:
"""Simulates AWSV4SignerAuth which cannot be deep-copied due to thread locks."""
def __init__(self):
self._lock = threading.Lock()
self.region = "us-east-1"
def __deepcopy__(self, memo):
raise TypeError("cannot pickle '_thread.lock' object")
class MockConnectionClass:
def __init__(self):
self._state = {"connected": False}
def __deepcopy__(self, memo):
raise TypeError("cannot pickle connection state")
class PlainConfig:
"""Config object using plain attributes (not Pydantic)."""
def __init__(self, **kwargs):
for k, v in kwargs.items():
setattr(self, k, v)
class TestSafeDeepcopyClonesNormally:
"""When deepcopy succeeds, config is returned as-is (no sanitization)."""
def test_deepcopy_success_returns_clone(self):
config = PlainConfig(host="localhost", port=9200, password="super_secret")
result = _safe_deepcopy_config(config)
assert result is not config
assert result.host == "localhost"
assert result.port == 9200
# deepcopy success path does not sanitize
assert result.password == "super_secret"
class TestSafeDeepcopyCopiesWithAuth:
"""When deepcopy fails (auth objects), fallback preserves auth and redacts secrets."""
def test_preserves_http_auth_and_connection_class(self):
auth = MockNonCopyableAuth()
conn = MockConnectionClass()
config = PlainConfig(
host="localhost",
port=9200,
http_auth=auth,
connection_class=conn,
api_key="secret123",
password="hunter2",
collection_name="test",
)
result = _safe_deepcopy_config(config)
# Runtime objects preserved (not None)
assert result.http_auth is not None
assert result.connection_class is not None
# Sensitive fields redacted
assert result.api_key is None
assert result.password is None
# Normal fields preserved
assert result.host == "localhost"
assert result.port == 9200
assert result.collection_name == "test"
def test_preserves_auth_field(self):
auth = MockNonCopyableAuth()
config = PlainConfig(
host="localhost",
auth=auth,
credentials={"key": "val"},
)
result = _safe_deepcopy_config(config)
assert result.auth is not None
assert result.credentials is None
class TestSafeDeepcopyWithPydantic:
"""Test fallback path with Pydantic-like model_dump objects."""
def test_pydantic_like_config(self):
class PydanticLikeConfig:
def __init__(self, **kwargs):
for k, v in kwargs.items():
setattr(self, k, v)
def model_dump(self, mode=None):
return {k: v for k, v in self.__dict__.items()
if not k.startswith("_")}
def __deepcopy__(self, memo):
raise TypeError("cannot deepcopy")
config = PydanticLikeConfig(
host="localhost",
api_key="secret",
http_auth="signer_obj",
)
result = _safe_deepcopy_config(config)
assert result.host == "localhost"
assert result.api_key is None
assert result.http_auth is not None
class TestSafeDeepcopyWithRealPydanticModel:
"""Test with real Pydantic BaseModel matching the OpenSearch config pattern.
This validates the model_dump() path (without mode='json') preserves
actual auth objects rather than losing them to JSON serialization.
"""
def test_real_pydantic_model_preserves_auth_objects(self):
from pydantic import BaseModel, Field
from typing import Optional
class OpenSearchLikeConfig(BaseModel):
host: str = "localhost"
port: int = 9200
collection_name: str = "test"
password: Optional[str] = None
api_key: Optional[str] = None
http_auth: Optional[object] = Field(None)
connection_class: Optional[object] = Field(None)
auth = MockNonCopyableAuth()
conn = MockConnectionClass()
config = OpenSearchLikeConfig(
host="myhost",
password="hunter2",
api_key="sk-secret",
http_auth=auth,
connection_class=conn,
)
result = _safe_deepcopy_config(config)
# Auth objects must be the actual objects, not string representations
assert result.http_auth is auth
assert result.connection_class is conn
# Sensitive fields must be redacted
assert result.password is None
assert result.api_key is None
# Normal fields preserved
assert result.host == "myhost"
assert result.port == 9200
class TestSafeDeepcopyWithDataclass:
"""Test fallback path with dataclasses."""
def test_dataclass_config(self):
@dataclass
class DCConfig:
host: str = "localhost"
api_key: str = None
db_password: str = None
http_auth: object = None
def __deepcopy__(self, memo):
raise TypeError("cannot deepcopy")
config = DCConfig(
host="myhost",
api_key="secret",
db_password="pass123",
http_auth="auth_obj",
)
result = _safe_deepcopy_config(config)
assert result.host == "myhost"
assert result.api_key is None
assert result.db_password is None
assert result.http_auth is not None
+6 -4
View File
@@ -299,11 +299,13 @@ def test_safe_deepcopy_config_handles_opensearch_auth(mock_sqlite, mock_llm_fact
config_with_auth = MockOpenSearchConfig(collection_name="opensearch_test", include_auth=True)
safe_config = _safe_deepcopy_config(config_with_auth)
assert safe_config.http_auth is None
assert safe_config.auth is None
# Runtime auth objects must be preserved (Issue #3580)
assert safe_config.http_auth is not None
assert safe_config.auth is not None
assert safe_config.connection_class is not None
# Credentials dict is a sensitive secret and should be redacted
assert safe_config.credentials is None
assert safe_config.connection_class is None
assert safe_config.collection_name == "opensearch_test"
assert safe_config.host == "localhost"