diff --git a/backend/open_webui/internal/migrations/019_encrypt_user_valves.py b/backend/open_webui/internal/migrations/019_encrypt_user_valves.py new file mode 100644 index 0000000000..5a77ddff52 --- /dev/null +++ b/backend/open_webui/internal/migrations/019_encrypt_user_valves.py @@ -0,0 +1,116 @@ +"""Peewee migrations -- 019_encrypt_user_valves.py. + +Encrypts existing plaintext user valve data stored in user.settings JSON. + +Some examples (model - class or model name):: + + > Model = migrator.orm['table_name'] # Return model in current state by name + > Model = migrator.ModelClass # Return model in current state by name + + > migrator.sql(sql) # Run custom SQL + > migrator.run(func, *args, **kwargs) # Run python function with the given args + > migrator.create_model(Model) # Create a model (could be used as decorator) + > migrator.remove_model(model, cascade=True) # Remove a model + > migrator.add_fields(model, **fields) # Add fields to a model + > migrator.change_fields(model, **fields) # Change fields + > migrator.remove_fields(model, *field_names, cascade=True) + > migrator.rename_field(model, old_field_name, new_field_name) + > migrator.rename_table(model, new_table_name) + > migrator.add_index(model, *col_names, unique=False) + > migrator.add_not_null(model, *field_names) + > migrator.add_default(model, field_name, default) + > migrator.add_constraint(model, name, sql) + > migrator.drop_index(model, *col_names) + > migrator.drop_not_null(model, *field_names) + > migrator.drop_constraints(model, *constraints) + +""" + +from contextlib import suppress + +import json + +import peewee as pw +from peewee_migrate import Migrator + +from open_webui.utils.valve_encryption import encrypt_user_valves + +with suppress(ImportError): + import playhouse.postgres_ext as pw_pext + + +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') + + for user_id, settings_raw in cursor.fetchall(): + if not settings_raw: + continue + + settings = ( + json.loads(settings_raw) + if isinstance(settings_raw, str) + else settings_raw + ) + + changed = False + + for category in ("tools", "functions"): + valves = ( + settings.get(category, {}).get("valves", {}) + if isinstance(settings.get(category), dict) + else {} + ) + for valve_id, valve_data in valves.items(): + if isinstance(valve_data, dict) and valve_data: + settings[category]["valves"][valve_id] = encrypt_user_valves( + valve_data + ) + changed = True + + if changed: + database.execute_sql( + 'UPDATE "user" SET "settings" = %s WHERE "id" = %s', + (json.dumps(settings), user_id), + ) + + +def rollback(migrator: Migrator, database: pw.Database, *, fake=False): + """Decrypt user valves back to plaintext.""" + from open_webui.utils.crypto import decrypt_user_valves + + cursor = database.execute_sql('SELECT "id", "settings" FROM "user" WHERE "settings" IS NOT NULL') + + for user_id, settings_raw in cursor.fetchall(): + if not settings_raw: + continue + + settings = ( + json.loads(settings_raw) + if isinstance(settings_raw, str) + else settings_raw + ) + + changed = False + + for category in ("tools", "functions"): + valves = ( + settings.get(category, {}).get("valves", {}) + if isinstance(settings.get(category), dict) + else {} + ) + for valve_id, valve_data in valves.items(): + if isinstance(valve_data, str): + try: + settings[category]["valves"][valve_id] = decrypt_user_valves( + valve_data + ) + changed = True + except Exception: + pass + + if changed: + database.execute_sql( + 'UPDATE "user" SET "settings" = %s WHERE "id" = %s', + (json.dumps(settings), user_id), + ) diff --git a/backend/open_webui/models/functions.py b/backend/open_webui/models/functions.py index ddac317863..88d8d341cd 100644 --- a/backend/open_webui/models/functions.py +++ b/backend/open_webui/models/functions.py @@ -6,6 +6,7 @@ from sqlalchemy import select, delete, update from sqlalchemy.ext.asyncio import AsyncSession from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.users import Users, UserModel, UserResponse +from open_webui.utils.valve_encryption import decrypt_user_valves, encrypt_user_valves from pydantic import BaseModel, ConfigDict from sqlalchemy import BigInteger, Boolean, Column, String, Text, Index @@ -357,7 +358,8 @@ class FunctionsTable: if 'valves' not in user_settings['functions']: user_settings['functions']['valves'] = {} - return user_settings['functions']['valves'].get(id, {}) + stored = user_settings['functions']['valves'].get(id, {}) + return decrypt_user_valves(stored) except Exception as e: log.exception(f'Error getting user values by id {id} and user id {user_id}') return None @@ -375,12 +377,12 @@ class FunctionsTable: if 'valves' not in user_settings['functions']: user_settings['functions']['valves'] = {} - user_settings['functions']['valves'][id] = valves + user_settings['functions']['valves'][id] = encrypt_user_valves(valves) # Update the user settings in the database await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) - return user_settings['functions']['valves'][id] + return valves except Exception as e: log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}') return None diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index 70035121aa..84e2797b35 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -8,6 +8,7 @@ from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.users import Users, UserResponse from open_webui.models.groups import Groups from open_webui.models.access_grants import AccessGrantModel, AccessGrants +from open_webui.utils.valve_encryption import decrypt_user_valves, encrypt_user_valves from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import BigInteger, Column, String, Text @@ -244,7 +245,8 @@ class ToolsTable: if 'valves' not in user_settings['tools']: user_settings['tools']['valves'] = {} - return user_settings['tools']['valves'].get(id, {}) + stored = user_settings['tools']['valves'].get(id, {}) + return decrypt_user_valves(stored) except Exception as e: log.exception(f'Error getting user values by id {id} and user_id {user_id}: {e}') return None @@ -262,12 +264,12 @@ class ToolsTable: if 'valves' not in user_settings['tools']: user_settings['tools']['valves'] = {} - user_settings['tools']['valves'][id] = valves + user_settings['tools']['valves'][id] = encrypt_user_valves(valves) # Update the user settings in the database await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) - return user_settings['tools']['valves'][id] + return valves except Exception as e: log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}') return None diff --git a/backend/open_webui/utils/valve_encryption.py b/backend/open_webui/utils/valve_encryption.py new file mode 100644 index 0000000000..28a856d45c --- /dev/null +++ b/backend/open_webui/utils/valve_encryption.py @@ -0,0 +1,43 @@ +import base64 +import hashlib +import json +import logging + +from cryptography.fernet import Fernet + +from open_webui.env import WEBUI_SECRET_KEY + +log = logging.getLogger(__name__) + + +def _make_fernet(key: str) -> Fernet: + """Derive a Fernet instance from an arbitrary string key.""" + if len(key) != 44: + key_bytes = hashlib.sha256(key.encode()).digest() + key_encoded = base64.urlsafe_b64encode(key_bytes) + else: + key_encoded = key.encode() + return Fernet(key_encoded) + + +_fernet = _make_fernet(WEBUI_SECRET_KEY) + + +def encrypt_user_valves(valves: dict) -> str: + """Encrypt a UserValves dict to an opaque string for DB storage.""" + valves_json = json.dumps(valves) + return _fernet.encrypt(valves_json.encode()).decode() + + +def decrypt_user_valves(stored) -> dict: + """Decrypt UserValves from DB storage. Handles both encrypted (str) and legacy plaintext (dict).""" + if isinstance(stored, dict): + return stored + if isinstance(stored, str): + try: + decrypted = _fernet.decrypt(stored.encode()).decode() + return json.loads(decrypted) + except Exception as e: + log.error(f"Error decrypting user valves: {type(e).__name__}: {e}") + raise + return {}