diff --git a/backend/open_webui/internal/migrations/019_encrypt_user_valves.py b/backend/open_webui/internal/migrations/019_encrypt_user_valves.py index 69c259884e..009a3cc2a9 100644 --- a/backend/open_webui/internal/migrations/019_encrypt_user_valves.py +++ b/backend/open_webui/internal/migrations/019_encrypt_user_valves.py @@ -29,6 +29,7 @@ Some examples (model - class or model name):: from contextlib import suppress import json +import logging import peewee as pw from peewee_migrate import Migrator @@ -38,20 +39,31 @@ from open_webui.utils.valve_encryption import encrypt_user_valves with suppress(ImportError): import playhouse.postgres_ext as pw_pext +log = logging.getLogger(__name__) + def migrate(migrator: Migrator, database: pw.Database, *, fake=False): """Encrypt existing plaintext user valves in user.settings.""" - cursor = database.execute_sql('SELECT "id", "settings" FROM "user" WHERE "settings" IS NOT NULL') + User = migrator.orm["user"] - for user_id, settings_raw in cursor.fetchall(): + for user in User.select().where(User.settings.is_null(False)): + settings_raw = user.settings if not settings_raw: continue - settings = ( - json.loads(settings_raw) - if isinstance(settings_raw, str) - else settings_raw - ) + try: + settings = ( + json.loads(settings_raw) + if isinstance(settings_raw, str) + else settings_raw + ) + except (json.JSONDecodeError, TypeError): + log.warning("Skipping user %s: malformed settings JSON", user.id) + continue + + if not isinstance(settings, dict): + log.warning("Skipping user %s: settings is not a dict", user.id) + continue changed = False @@ -69,27 +81,35 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False): changed = True if changed: - database.execute_sql( - 'UPDATE "user" SET "settings" = %s WHERE "id" = %s', - (json.dumps(settings), user_id), - ) + User.update(settings=json.dumps(settings)).where( + User.id == user.id + ).execute() def rollback(migrator: Migrator, database: pw.Database, *, fake=False): """Decrypt user valves back to plaintext.""" from open_webui.utils.valve_encryption import decrypt_user_valves - cursor = database.execute_sql('SELECT "id", "settings" FROM "user" WHERE "settings" IS NOT NULL') + User = migrator.orm["user"] - for user_id, settings_raw in cursor.fetchall(): + for user in User.select().where(User.settings.is_null(False)): + settings_raw = user.settings if not settings_raw: continue - settings = ( - json.loads(settings_raw) - if isinstance(settings_raw, str) - else settings_raw - ) + try: + settings = ( + json.loads(settings_raw) + if isinstance(settings_raw, str) + else settings_raw + ) + except (json.JSONDecodeError, TypeError): + log.warning("Rollback skipping user %s: malformed settings JSON", user.id) + continue + + if not isinstance(settings, dict): + log.warning("Rollback skipping user %s: settings is not a dict", user.id) + continue changed = False @@ -107,10 +127,12 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False): ) changed = True except Exception: - pass + log.warning( + "Rollback failed to decrypt valve %s/%s for user %s", + category, valve_id, user.id, + ) if changed: - database.execute_sql( - 'UPDATE "user" SET "settings" = %s WHERE "id" = %s', - (json.dumps(settings), user_id), - ) + User.update(settings=json.dumps(settings)).where( + User.id == user.id + ).execute() diff --git a/backend/open_webui/utils/valve_encryption.py b/backend/open_webui/utils/valve_encryption.py index 28a856d45c..45389d5881 100644 --- a/backend/open_webui/utils/valve_encryption.py +++ b/backend/open_webui/utils/valve_encryption.py @@ -3,23 +3,32 @@ import hashlib import json import logging -from cryptography.fernet import Fernet +from cryptography.fernet import Fernet, InvalidToken from open_webui.env import WEBUI_SECRET_KEY log = logging.getLogger(__name__) +_DEFAULT_SECRET = "t0p-s3cr3t" + def _make_fernet(key: str) -> Fernet: """Derive a Fernet instance from an arbitrary string key.""" - if len(key) != 44: + try: + return Fernet(key.encode()) + except Exception: key_bytes = hashlib.sha256(key.encode()).digest() - key_encoded = base64.urlsafe_b64encode(key_bytes) - else: - key_encoded = key.encode() - return Fernet(key_encoded) + return Fernet(base64.urlsafe_b64encode(key_bytes)) +if WEBUI_SECRET_KEY == _DEFAULT_SECRET: + log.warning( + "WEBUI_SECRET_KEY is set to the default value '%s'. " + "Encrypted valve data is trivially decryptable. " + "Set a strong, unique WEBUI_SECRET_KEY in production.", + _DEFAULT_SECRET, + ) + _fernet = _make_fernet(WEBUI_SECRET_KEY) @@ -37,7 +46,19 @@ def decrypt_user_valves(stored) -> dict: try: decrypted = _fernet.decrypt(stored.encode()).decode() return json.loads(decrypted) + except InvalidToken: + log.error( + "Failed to decrypt user valves: key mismatch or corrupted data. " + "Returning empty valves. The original encrypted value is preserved in the database." + ) + return {} + except json.JSONDecodeError: + log.error( + "Decrypted user valves but got malformed JSON. " + "Returning empty valves. The original encrypted value is preserved in the database." + ) + return {} except Exception as e: - log.error(f"Error decrypting user valves: {type(e).__name__}: {e}") - raise + log.error("Unexpected error decrypting user valves: %s: %s", type(e).__name__, e) + return {} return {}