Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 48 additions & 0 deletions src/memos/api/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,11 @@ def is_scheduler_enabled() -> bool:
"""Check if scheduler is enabled via environment variable."""
return os.getenv("MOS_ENABLE_SCHEDULER", "false").lower() == "true"

@staticmethod
def is_default_cube_config_enabled() -> bool:
"""Check if default cube config is enabled via environment variable."""
return os.getenv("MOS_ENABLE_DEFAULT_CUBE_CONFIG", "false").lower() == "true"

@staticmethod
def get_product_default_config() -> dict[str, Any]:
"""Get default configuration for Product API."""
Expand Down Expand Up @@ -321,3 +326,46 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General

default_mem_cube = GeneralMemCube(default_cube_config)
return default_config, default_mem_cube

@staticmethod
def get_default_cube_config() -> GeneralMemCubeConfig | None:
"""Get default cube configuration for product initialization.

Returns:
GeneralMemCubeConfig | None: Default cube configuration if enabled, None otherwise.
"""
if not APIConfig.is_default_cube_config_enabled():
return None

openai_config = APIConfig.get_openai_config()
neo4j_config = APIConfig.get_neo4j_config()

return GeneralMemCubeConfig.model_validate(
{
"user_id": "default",
"cube_id": "default_cube",
"text_mem": {
"backend": "tree_text",
"config": {
"extractor_llm": {"backend": "openai", "config": openai_config},
"dispatcher_llm": {"backend": "openai", "config": openai_config},
"graph_db": {
"backend": "neo4j",
"config": neo4j_config,
},
"embedder": {
"backend": "ollama",
"config": {
"model_name_or_path": "nomic-embed-text:latest",
"api_base": os.getenv("OLLAMA_API_BASE", "http://localhost:11434"),
},
},
"reorganize": os.getenv("MOS_ENABLE_REORGANIZE", "false").lower() == "true",
},
},
"act_mem": {}
if os.getenv("ENABLE_ACTIVATION_MEMORY", "false").lower() == "false"
else APIConfig.get_activation_vllm_config(),
"para_mem": {},
}
)
9 changes: 8 additions & 1 deletion src/memos/api/routers/product_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,13 @@ def get_mos_product_instance():
from memos.configs.mem_os import MOSConfig

mos_config = MOSConfig(**default_config)
MOS_PRODUCT_INSTANCE = MOSProduct(default_config=mos_config)

# Get default cube config from APIConfig (may be None if disabled)
default_cube_config = APIConfig.get_default_cube_config()
print("*********default_cube_config*********", default_cube_config)
MOS_PRODUCT_INSTANCE = MOSProduct(
default_config=mos_config, default_cube_config=default_cube_config
)
logger.info("MOSProduct instance created successfully with inheritance architecture")
return MOS_PRODUCT_INSTANCE

Expand All @@ -68,6 +74,7 @@ async def register_user(user_req: UserRegisterRequest):
logger.info(f"user_config: {user_config.model_dump(mode='json')}")
logger.info(f"default_mem_cube: {default_mem_cube.config.model_dump(mode='json')}")
mos_product = get_mos_product_instance()

# Register user with default config and mem cube
result = mos_product.user_register(
user_id=user_req.user_id,
Expand Down
18 changes: 15 additions & 3 deletions src/memos/mem_cube/general.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from memos.exceptions import ConfigurationError, MemCubeError
from memos.log import get_logger
from memos.mem_cube.base import BaseMemCube
from memos.mem_cube.utils import download_repo
from memos.mem_cube.utils import download_repo, merge_config_with_default
from memos.memories.activation.base import BaseActMemory
from memos.memories.factory import MemoryFactory
from memos.memories.parametric.base import BaseParaMemory
Expand Down Expand Up @@ -114,20 +114,30 @@ def dump(

@staticmethod
def init_from_dir(
dir: str, memory_types: list[Literal["text_mem", "act_mem", "para_mem"]] | None = None
dir: str,
memory_types: list[Literal["text_mem", "act_mem", "para_mem"]] | None = None,
default_config: GeneralMemCubeConfig | None = None,
) -> "GeneralMemCube":
"""Create a MemCube instance from a MemCube directory.

Args:
dir (str): The directory containing the memory files.
memory_types (list[str], optional): List of memory types to load.
If None, loads all available memory types.
default_config (GeneralMemCubeConfig, optional): Default configuration to merge with existing config.
If provided, will merge general settings while preserving critical user-specific fields.

Returns:
MemCube: An instance of MemCube loaded with memories from the specified directory.
"""
config_path = os.path.join(dir, "config.json")
config = GeneralMemCubeConfig.from_json_file(config_path)

# Merge with default config if provided
if default_config is not None:
config = merge_config_with_default(config, default_config)
logger.info(f"Applied default config to cube {config.cube_id}")

mem_cube = GeneralMemCube(config)
mem_cube.load(dir, memory_types)
return mem_cube
Expand All @@ -137,6 +147,7 @@ def init_from_remote_repo(
cube_id: str,
base_url: str = "https://huggingface.co/datasets",
memory_types: list[Literal["text_mem", "act_mem", "para_mem"]] | None = None,
default_config: GeneralMemCubeConfig | None = None,
) -> "GeneralMemCube":
"""Create a MemCube instance from a remote repository.

Expand All @@ -145,12 +156,13 @@ def init_from_remote_repo(
base_url (str): The base URL of the remote repository.
memory_types (list[str], optional): List of memory types to load.
If None, loads all available memory types.
default_config (GeneralMemCubeConfig, optional): Default configuration to merge with existing config.

Returns:
MemCube: An instance of MemCube loaded with memories from the specified remote repository.
"""
dir = download_repo(cube_id, base_url)
return GeneralMemCube.init_from_dir(dir, memory_types)
return GeneralMemCube.init_from_dir(dir, memory_types, default_config)

@property
def text_mem(self) -> "BaseTextMemory | None":
Expand Down
102 changes: 102 additions & 0 deletions src/memos/mem_cube/utils.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,15 @@
import copy
import logging
import subprocess
import tempfile

from typing import Any

from memos.configs.mem_cube import GeneralMemCubeConfig


logger = logging.getLogger(__name__)


def download_repo(repo: str, base_url: str, dir: str | None = None) -> str:
"""Download a repository from a remote source.
Expand All @@ -22,3 +31,96 @@ def download_repo(repo: str, base_url: str, dir: str | None = None) -> str:
subprocess.run(["git", "clone", repo_url, dir], check=True)

return dir


def merge_config_with_default(
existing_config: GeneralMemCubeConfig, default_config: GeneralMemCubeConfig
) -> GeneralMemCubeConfig:
"""
Merge existing cube config with default config, preserving critical fields.

This method updates general configuration fields (like API keys, model parameters)
while preserving critical user-specific fields (like user_id, cube_id, graph_db settings).

Args:
existing_config (GeneralMemCubeConfig): The existing cube configuration loaded from file
default_config (GeneralMemCubeConfig): The default configuration to merge from

Returns:
GeneralMemCubeConfig: Merged configuration
"""

def deep_merge_dicts(
existing: dict[str, Any], default: dict[str, Any], preserve_keys: set[str] | None = None
) -> dict[str, Any]:
"""Recursively merge dictionaries, preserving specified keys from existing dict."""
if preserve_keys is None:
preserve_keys = set()

result = copy.deepcopy(existing)

for key, default_value in default.items():
if key in preserve_keys:
# Preserve existing value for critical keys
continue

if key in result and isinstance(result[key], dict) and isinstance(default_value, dict):
# Recursively merge nested dictionaries
result[key] = deep_merge_dicts(result[key], default_value, preserve_keys)
elif key not in result or result[key] is None:
# Use default value if key doesn't exist or is None
result[key] = copy.deepcopy(default_value)
# For non-dict values, keep existing value unless it's None

return result

# Convert configs to dictionaries
existing_dict = existing_config.model_dump(mode="json")
default_dict = default_config.model_dump(mode="json")

# Merge text_mem config
if "text_mem" in existing_dict and "text_mem" in default_dict:
existing_text_config = existing_dict["text_mem"].get("config", {})
default_text_config = default_dict["text_mem"].get("config", {})

# Handle nested graph_db config specially
if "graph_db" in existing_text_config and "graph_db" in default_text_config:
existing_graph_config = existing_text_config["graph_db"].get("config", {})
default_graph_config = default_text_config["graph_db"].get("config", {})

# Merge graph_db config, preserving critical keys
merged_graph_config = deep_merge_dicts(
existing_graph_config,
default_graph_config,
preserve_keys={"uri", "user", "password", "db_name", "auto_create"},
)

# Update the configs
existing_text_config["graph_db"]["config"] = merged_graph_config
default_text_config["graph_db"]["config"] = merged_graph_config

# Merge other text_mem config fields
merged_text_config = deep_merge_dicts(existing_text_config, default_text_config)
existing_dict["text_mem"]["config"] = merged_text_config

# Merge act_mem config
if "act_mem" in existing_dict and "act_mem" in default_dict:
existing_act_config = existing_dict["act_mem"].get("config", {})
default_act_config = default_dict["act_mem"].get("config", {})
merged_act_config = deep_merge_dicts(existing_act_config, default_act_config)
existing_dict["act_mem"]["config"] = merged_act_config

# Merge para_mem config
if "para_mem" in existing_dict and "para_mem" in default_dict:
existing_para_config = existing_dict["para_mem"].get("config", {})
default_para_config = default_dict["para_mem"].get("config", {})
merged_para_config = deep_merge_dicts(existing_para_config, default_para_config)
existing_dict["para_mem"]["config"] = merged_para_config

# Create new config from merged dictionary
merged_config = GeneralMemCubeConfig.model_validate(existing_dict)
logger.info(
f"Merged cube config for user {merged_config.user_id}, cube {merged_config.cube_id}"
)

return merged_config
Loading