mem0 / server /server_state.py
keithproject's picture
fix: retry DB load on startup — wait up to 60s for Postgres crash recovery
d9cf3d7
Raw
History Blame Contribute Delete
4.13 kB
import json
import logging
import threading
import time
from copy import deepcopy
from typing import Any, Callable, Dict
from mem0 import Memory
_state_lock = threading.RLock()
_current_config: Dict[str, Any] = {}
_memory_instance: Memory | None = None
_session_factory: Callable | None = None
def set_session_factory(factory: Callable) -> None:
global _session_factory
_session_factory = factory
def _load_overrides() -> Dict[str, Any]:
try:
if _session_factory is None:
return {}
from models import Settings
session = _session_factory()
try:
row = session.get(Settings, "config_overrides")
if row is None:
return {}
return json.loads(row.value)
finally:
session.close()
except Exception:
return {}
def _save_overrides(overrides: Dict[str, Any]) -> None:
try:
if _session_factory is None:
return
from models import Settings
from sqlalchemy.dialects.postgresql import insert
session = _session_factory()
try:
serialized = json.dumps(overrides)
stmt = (
insert(Settings)
.values(key="config_overrides", value=serialized)
.on_conflict_do_update(index_elements=[Settings.key], set_={"value": serialized})
)
session.execute(stmt)
session.commit()
finally:
session.close()
except Exception:
logging.warning("Failed to persist config overrides to database", exc_info=True)
def _merge_config(base: Dict[str, Any], updates: Dict[str, Any]) -> Dict[str, Any]:
merged = deepcopy(base)
for key, value in updates.items():
if isinstance(value, dict) and isinstance(merged.get(key), dict):
merged[key] = _merge_config(merged[key], value)
else:
merged[key] = value
return merged
def initialize_state(default_config: Dict[str, Any]) -> None:
global _current_config, _memory_instance
with _state_lock:
_current_config = deepcopy(default_config)
# ponytail: retry DB load — Postgres needs ~90s crash recovery on HF FUSE
overrides: Dict[str, Any] = {}
for attempt in range(1, 13): # max 60s (12 x 5s)
overrides = _load_overrides()
if overrides:
logging.info("Config overrides loaded from DB on attempt %d", attempt)
break
if _session_factory is not None:
# DB reachable but no overrides yet — stop retrying
try:
from models import Settings
session = _session_factory()
try:
session.execute(__import__("sqlalchemy").text("SELECT 1"))
logging.info("DB ready, no config overrides stored yet")
break
finally:
session.close()
except Exception:
pass
logging.warning("DB not ready yet (attempt %d/12), retrying in 5s...", attempt)
time.sleep(5)
if overrides:
_current_config = _merge_config(_current_config, overrides)
_memory_instance = Memory.from_config(_current_config)
def update_config(updates: Dict[str, Any]) -> Dict[str, Any]:
global _current_config, _memory_instance
with _state_lock:
next_config = _merge_config(_current_config, updates)
_current_config = next_config
_memory_instance = Memory.from_config(next_config)
overrides = _load_overrides()
overrides = _merge_config(overrides, updates)
_save_overrides(overrides)
return deepcopy(_current_config)
def get_current_config() -> Dict[str, Any]:
with _state_lock:
return deepcopy(_current_config)
def get_memory_instance() -> Memory:
with _state_lock:
if _memory_instance is None:
raise RuntimeError("Mem0 runtime has not been initialized.")
return _memory_instance