mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-04 02:33:43 +00:00
fix: address review feedback for valve encryption
- Use Peewee ORM instead of raw SQL to fix SQLite placeholder compatibility - Validate Fernet key by attempting direct init before SHA256 fallback - Warn at startup when WEBUI_SECRET_KEY is the default value - Return empty dict on decryption failure instead of raising - Handle malformed settings JSON gracefully during migration - Log rollback decryption failures with user/valve context Co-Authored-By: Claude Opus 4 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
8c06fd5bb3
commit
a789201bfe
2 changed files with 74 additions and 31 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue