This commit is contained in:
Timothy Jaeryang Baek
2026-03-17 17:58:01 -05:00
parent fcf7208352
commit de3317e26b
220 changed files with 17200 additions and 22836 deletions
+20 -30
View File
@@ -56,17 +56,15 @@ def handle_peewee_migration(DATABASE_URL):
# db = None
try:
# Replace the postgresql:// with postgres:// to handle the peewee migration
db = register_connection(DATABASE_URL.replace("postgresql://", "postgres://"))
migrate_dir = OPEN_WEBUI_DIR / "internal" / "migrations"
db = register_connection(DATABASE_URL.replace('postgresql://', 'postgres://'))
migrate_dir = OPEN_WEBUI_DIR / 'internal' / 'migrations'
router = Router(db, logger=log, migrate_dir=migrate_dir)
router.run()
db.close()
except Exception as e:
log.error(f"Failed to initialize the database connection: {e}")
log.warning(
"Hint: If your database password contains special characters, you may need to URL-encode it."
)
log.error(f'Failed to initialize the database connection: {e}')
log.warning('Hint: If your database password contains special characters, you may need to URL-encode it.')
raise
finally:
# Properly closing the database connection
@@ -74,7 +72,7 @@ def handle_peewee_migration(DATABASE_URL):
db.close()
# Assert if db connection has been closed
assert db.is_closed(), "Database connection is still open."
assert db.is_closed(), 'Database connection is still open.'
if ENABLE_DB_MIGRATIONS:
@@ -84,15 +82,13 @@ if ENABLE_DB_MIGRATIONS:
SQLALCHEMY_DATABASE_URL = DATABASE_URL
# Handle SQLCipher URLs
if SQLALCHEMY_DATABASE_URL.startswith("sqlite+sqlcipher://"):
database_password = os.environ.get("DATABASE_PASSWORD")
if not database_password or database_password.strip() == "":
raise ValueError(
"DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs"
)
if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'):
database_password = os.environ.get('DATABASE_PASSWORD')
if not database_password or database_password.strip() == '':
raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
# Extract database path from SQLCipher URL
db_path = SQLALCHEMY_DATABASE_URL.replace("sqlite+sqlcipher://", "")
db_path = SQLALCHEMY_DATABASE_URL.replace('sqlite+sqlcipher://', '')
# Create a custom creator function that uses sqlcipher3
def create_sqlcipher_connection():
@@ -109,7 +105,7 @@ if SQLALCHEMY_DATABASE_URL.startswith("sqlite+sqlcipher://"):
# or QueuePool if DATABASE_POOL_SIZE is explicitly configured.
if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0:
engine = create_engine(
"sqlite://",
'sqlite://',
creator=create_sqlcipher_connection,
pool_size=DATABASE_POOL_SIZE,
max_overflow=DATABASE_POOL_MAX_OVERFLOW,
@@ -121,28 +117,26 @@ if SQLALCHEMY_DATABASE_URL.startswith("sqlite+sqlcipher://"):
)
else:
engine = create_engine(
"sqlite://",
'sqlite://',
creator=create_sqlcipher_connection,
poolclass=NullPool,
echo=False,
)
log.info("Connected to encrypted SQLite database using SQLCipher")
log.info('Connected to encrypted SQLite database using SQLCipher')
elif "sqlite" in SQLALCHEMY_DATABASE_URL:
engine = create_engine(
SQLALCHEMY_DATABASE_URL, connect_args={"check_same_thread": False}
)
elif 'sqlite' in SQLALCHEMY_DATABASE_URL:
engine = create_engine(SQLALCHEMY_DATABASE_URL, connect_args={'check_same_thread': False})
def on_connect(dbapi_connection, connection_record):
cursor = dbapi_connection.cursor()
if DATABASE_ENABLE_SQLITE_WAL:
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute('PRAGMA journal_mode=WAL')
else:
cursor.execute("PRAGMA journal_mode=DELETE")
cursor.execute('PRAGMA journal_mode=DELETE')
cursor.close()
event.listen(engine, "connect", on_connect)
event.listen(engine, 'connect', on_connect)
else:
if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0:
@@ -156,16 +150,12 @@ else:
poolclass=QueuePool,
)
else:
engine = create_engine(
SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool
)
engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool)
else:
engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True)
SessionLocal = sessionmaker(
autocommit=False, autoflush=False, bind=engine, expire_on_commit=False
)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
metadata_obj = MetaData(schema=DATABASE_SCHEMA)
Base = declarative_base(metadata=metadata_obj)
ScopedSession = scoped_session(SessionLocal)
@@ -56,7 +56,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
active = pw.BooleanField()
class Meta:
table_name = "auth"
table_name = 'auth'
@migrator.create_model
class Chat(pw.Model):
@@ -67,7 +67,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "chat"
table_name = 'chat'
@migrator.create_model
class ChatIdTag(pw.Model):
@@ -78,7 +78,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "chatidtag"
table_name = 'chatidtag'
@migrator.create_model
class Document(pw.Model):
@@ -92,7 +92,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "document"
table_name = 'document'
@migrator.create_model
class Modelfile(pw.Model):
@@ -103,7 +103,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "modelfile"
table_name = 'modelfile'
@migrator.create_model
class Prompt(pw.Model):
@@ -115,7 +115,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "prompt"
table_name = 'prompt'
@migrator.create_model
class Tag(pw.Model):
@@ -125,7 +125,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
data = pw.TextField(null=True)
class Meta:
table_name = "tag"
table_name = 'tag'
@migrator.create_model
class User(pw.Model):
@@ -137,7 +137,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "user"
table_name = 'user'
def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
@@ -149,7 +149,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
active = pw.BooleanField()
class Meta:
table_name = "auth"
table_name = 'auth'
@migrator.create_model
class Chat(pw.Model):
@@ -160,7 +160,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "chat"
table_name = 'chat'
@migrator.create_model
class ChatIdTag(pw.Model):
@@ -171,7 +171,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "chatidtag"
table_name = 'chatidtag'
@migrator.create_model
class Document(pw.Model):
@@ -185,7 +185,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "document"
table_name = 'document'
@migrator.create_model
class Modelfile(pw.Model):
@@ -196,7 +196,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "modelfile"
table_name = 'modelfile'
@migrator.create_model
class Prompt(pw.Model):
@@ -208,7 +208,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "prompt"
table_name = 'prompt'
@migrator.create_model
class Tag(pw.Model):
@@ -218,7 +218,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
data = pw.TextField(null=True)
class Meta:
table_name = "tag"
table_name = 'tag'
@migrator.create_model
class User(pw.Model):
@@ -230,24 +230,24 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
timestamp = pw.BigIntegerField()
class Meta:
table_name = "user"
table_name = 'user'
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_model("user")
migrator.remove_model('user')
migrator.remove_model("tag")
migrator.remove_model('tag')
migrator.remove_model("prompt")
migrator.remove_model('prompt')
migrator.remove_model("modelfile")
migrator.remove_model('modelfile')
migrator.remove_model("document")
migrator.remove_model('document')
migrator.remove_model("chatidtag")
migrator.remove_model('chatidtag')
migrator.remove_model("chat")
migrator.remove_model('chat')
migrator.remove_model("auth")
migrator.remove_model('auth')
@@ -36,12 +36,10 @@ with suppress(ImportError):
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
migrator.add_fields(
"chat", share_id=pw.CharField(max_length=255, null=True, unique=True)
)
migrator.add_fields('chat', share_id=pw.CharField(max_length=255, null=True, unique=True))
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_fields("chat", "share_id")
migrator.remove_fields('chat', 'share_id')
@@ -36,12 +36,10 @@ with suppress(ImportError):
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
migrator.add_fields(
"user", api_key=pw.CharField(max_length=255, null=True, unique=True)
)
migrator.add_fields('user', api_key=pw.CharField(max_length=255, null=True, unique=True))
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_fields("user", "api_key")
migrator.remove_fields('user', 'api_key')
@@ -36,10 +36,10 @@ with suppress(ImportError):
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
migrator.add_fields("chat", archived=pw.BooleanField(default=False))
migrator.add_fields('chat', archived=pw.BooleanField(default=False))
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_fields("chat", "archived")
migrator.remove_fields('chat', 'archived')
@@ -45,22 +45,20 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
# Adding fields created_at and updated_at to the 'chat' table
migrator.add_fields(
"chat",
'chat',
created_at=pw.DateTimeField(null=True), # Allow null for transition
updated_at=pw.DateTimeField(null=True), # Allow null for transition
)
# Populate the new fields from an existing 'timestamp' field
migrator.sql(
"UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL"
)
migrator.sql('UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL')
# Now that the data has been copied, remove the original 'timestamp' field
migrator.remove_fields("chat", "timestamp")
migrator.remove_fields('chat', 'timestamp')
# Update the fields to be not null now that they are populated
migrator.change_fields(
"chat",
'chat',
created_at=pw.DateTimeField(null=False),
updated_at=pw.DateTimeField(null=False),
)
@@ -69,22 +67,20 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
# Adding fields created_at and updated_at to the 'chat' table
migrator.add_fields(
"chat",
'chat',
created_at=pw.BigIntegerField(null=True), # Allow null for transition
updated_at=pw.BigIntegerField(null=True), # Allow null for transition
)
# Populate the new fields from an existing 'timestamp' field
migrator.sql(
"UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL"
)
migrator.sql('UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL')
# Now that the data has been copied, remove the original 'timestamp' field
migrator.remove_fields("chat", "timestamp")
migrator.remove_fields('chat', 'timestamp')
# Update the fields to be not null now that they are populated
migrator.change_fields(
"chat",
'chat',
created_at=pw.BigIntegerField(null=False),
updated_at=pw.BigIntegerField(null=False),
)
@@ -101,29 +97,29 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
def rollback_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
# Recreate the timestamp field initially allowing null values for safe transition
migrator.add_fields("chat", timestamp=pw.DateTimeField(null=True))
migrator.add_fields('chat', timestamp=pw.DateTimeField(null=True))
# Copy the earliest created_at date back into the new timestamp field
# This assumes created_at was originally a copy of timestamp
migrator.sql("UPDATE chat SET timestamp = created_at")
migrator.sql('UPDATE chat SET timestamp = created_at')
# Remove the created_at and updated_at fields
migrator.remove_fields("chat", "created_at", "updated_at")
migrator.remove_fields('chat', 'created_at', 'updated_at')
# Finally, alter the timestamp field to not allow nulls if that was the original setting
migrator.change_fields("chat", timestamp=pw.DateTimeField(null=False))
migrator.change_fields('chat', timestamp=pw.DateTimeField(null=False))
def rollback_external(migrator: Migrator, database: pw.Database, *, fake=False):
# Recreate the timestamp field initially allowing null values for safe transition
migrator.add_fields("chat", timestamp=pw.BigIntegerField(null=True))
migrator.add_fields('chat', timestamp=pw.BigIntegerField(null=True))
# Copy the earliest created_at date back into the new timestamp field
# This assumes created_at was originally a copy of timestamp
migrator.sql("UPDATE chat SET timestamp = created_at")
migrator.sql('UPDATE chat SET timestamp = created_at')
# Remove the created_at and updated_at fields
migrator.remove_fields("chat", "created_at", "updated_at")
migrator.remove_fields('chat', 'created_at', 'updated_at')
# Finally, alter the timestamp field to not allow nulls if that was the original setting
migrator.change_fields("chat", timestamp=pw.BigIntegerField(null=False))
migrator.change_fields('chat', timestamp=pw.BigIntegerField(null=False))
@@ -38,45 +38,45 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
# Alter the tables with timestamps
migrator.change_fields(
"chatidtag",
'chatidtag',
timestamp=pw.BigIntegerField(),
)
migrator.change_fields(
"document",
'document',
timestamp=pw.BigIntegerField(),
)
migrator.change_fields(
"modelfile",
'modelfile',
timestamp=pw.BigIntegerField(),
)
migrator.change_fields(
"prompt",
'prompt',
timestamp=pw.BigIntegerField(),
)
migrator.change_fields(
"user",
'user',
timestamp=pw.BigIntegerField(),
)
# Alter the tables with varchar to text where necessary
migrator.change_fields(
"auth",
'auth',
password=pw.TextField(),
)
migrator.change_fields(
"chat",
'chat',
title=pw.TextField(),
)
migrator.change_fields(
"document",
'document',
title=pw.TextField(),
filename=pw.TextField(),
)
migrator.change_fields(
"prompt",
'prompt',
title=pw.TextField(),
)
migrator.change_fields(
"user",
'user',
profile_image_url=pw.TextField(),
)
@@ -87,43 +87,43 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
if isinstance(database, pw.SqliteDatabase):
# Alter the tables with timestamps
migrator.change_fields(
"chatidtag",
'chatidtag',
timestamp=pw.DateField(),
)
migrator.change_fields(
"document",
'document',
timestamp=pw.DateField(),
)
migrator.change_fields(
"modelfile",
'modelfile',
timestamp=pw.DateField(),
)
migrator.change_fields(
"prompt",
'prompt',
timestamp=pw.DateField(),
)
migrator.change_fields(
"user",
'user',
timestamp=pw.DateField(),
)
migrator.change_fields(
"auth",
'auth',
password=pw.CharField(max_length=255),
)
migrator.change_fields(
"chat",
'chat',
title=pw.CharField(),
)
migrator.change_fields(
"document",
'document',
title=pw.CharField(),
filename=pw.CharField(),
)
migrator.change_fields(
"prompt",
'prompt',
title=pw.CharField(),
)
migrator.change_fields(
"user",
'user',
profile_image_url=pw.CharField(),
)
@@ -38,7 +38,7 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
# Adding fields created_at and updated_at to the 'user' table
migrator.add_fields(
"user",
'user',
created_at=pw.BigIntegerField(null=True), # Allow null for transition
updated_at=pw.BigIntegerField(null=True), # Allow null for transition
last_active_at=pw.BigIntegerField(null=True), # Allow null for transition
@@ -50,11 +50,11 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
)
# Now that the data has been copied, remove the original 'timestamp' field
migrator.remove_fields("user", "timestamp")
migrator.remove_fields('user', 'timestamp')
# Update the fields to be not null now that they are populated
migrator.change_fields(
"user",
'user',
created_at=pw.BigIntegerField(null=False),
updated_at=pw.BigIntegerField(null=False),
last_active_at=pw.BigIntegerField(null=False),
@@ -65,14 +65,14 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
# Recreate the timestamp field initially allowing null values for safe transition
migrator.add_fields("user", timestamp=pw.BigIntegerField(null=True))
migrator.add_fields('user', timestamp=pw.BigIntegerField(null=True))
# Copy the earliest created_at date back into the new timestamp field
# This assumes created_at was originally a copy of timestamp
migrator.sql('UPDATE "user" SET timestamp = created_at')
# Remove the created_at and updated_at fields
migrator.remove_fields("user", "created_at", "updated_at", "last_active_at")
migrator.remove_fields('user', 'created_at', 'updated_at', 'last_active_at')
# Finally, alter the timestamp field to not allow nulls if that was the original setting
migrator.change_fields("user", timestamp=pw.BigIntegerField(null=False))
migrator.change_fields('user', timestamp=pw.BigIntegerField(null=False))
@@ -43,10 +43,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
created_at = pw.BigIntegerField(null=False)
class Meta:
table_name = "memory"
table_name = 'memory'
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_model("memory")
migrator.remove_model('memory')
@@ -51,10 +51,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
updated_at = pw.BigIntegerField(null=False)
class Meta:
table_name = "model"
table_name = 'model'
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_model("model")
migrator.remove_model('model')
@@ -42,12 +42,12 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
# Fetch data from 'modelfile' table and insert into 'model' table
migrate_modelfile_to_model(migrator, database)
# Drop the 'modelfile' table
migrator.remove_model("modelfile")
migrator.remove_model('modelfile')
def migrate_modelfile_to_model(migrator: Migrator, database: pw.Database):
ModelFile = migrator.orm["modelfile"]
Model = migrator.orm["model"]
ModelFile = migrator.orm['modelfile']
Model = migrator.orm['model']
modelfiles = ModelFile.select()
@@ -57,25 +57,25 @@ def migrate_modelfile_to_model(migrator: Migrator, database: pw.Database):
modelfile.modelfile = json.loads(modelfile.modelfile)
meta = json.dumps(
{
"description": modelfile.modelfile.get("desc"),
"profile_image_url": modelfile.modelfile.get("imageUrl"),
"ollama": {"modelfile": modelfile.modelfile.get("content")},
"suggestion_prompts": modelfile.modelfile.get("suggestionPrompts"),
"categories": modelfile.modelfile.get("categories"),
"user": {**modelfile.modelfile.get("user", {}), "community": True},
'description': modelfile.modelfile.get('desc'),
'profile_image_url': modelfile.modelfile.get('imageUrl'),
'ollama': {'modelfile': modelfile.modelfile.get('content')},
'suggestion_prompts': modelfile.modelfile.get('suggestionPrompts'),
'categories': modelfile.modelfile.get('categories'),
'user': {**modelfile.modelfile.get('user', {}), 'community': True},
}
)
info = parse_ollama_modelfile(modelfile.modelfile.get("content"))
info = parse_ollama_modelfile(modelfile.modelfile.get('content'))
# Insert the processed data into the 'model' table
Model.create(
id=f"ollama-{modelfile.tag_name}",
id=f'ollama-{modelfile.tag_name}',
user_id=modelfile.user_id,
base_model_id=info.get("base_model_id"),
name=modelfile.modelfile.get("title"),
base_model_id=info.get('base_model_id'),
name=modelfile.modelfile.get('title'),
meta=meta,
params=json.dumps(info.get("params", {})),
params=json.dumps(info.get('params', {})),
created_at=modelfile.timestamp,
updated_at=modelfile.timestamp,
)
@@ -86,7 +86,7 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
recreate_modelfile_table(migrator, database)
move_data_back_to_modelfile(migrator, database)
migrator.remove_model("model")
migrator.remove_model('model')
def recreate_modelfile_table(migrator: Migrator, database: pw.Database):
@@ -102,8 +102,8 @@ def recreate_modelfile_table(migrator: Migrator, database: pw.Database):
def move_data_back_to_modelfile(migrator: Migrator, database: pw.Database):
Model = migrator.orm["model"]
Modelfile = migrator.orm["modelfile"]
Model = migrator.orm['model']
Modelfile = migrator.orm['modelfile']
models = Model.select()
@@ -112,13 +112,13 @@ def move_data_back_to_modelfile(migrator: Migrator, database: pw.Database):
meta = json.loads(model.meta)
modelfile_data = {
"title": model.name,
"desc": meta.get("description"),
"imageUrl": meta.get("profile_image_url"),
"content": meta.get("ollama", {}).get("modelfile"),
"suggestionPrompts": meta.get("suggestion_prompts"),
"categories": meta.get("categories"),
"user": {k: v for k, v in meta.get("user", {}).items() if k != "community"},
'title': model.name,
'desc': meta.get('description'),
'imageUrl': meta.get('profile_image_url'),
'content': meta.get('ollama', {}).get('modelfile'),
'suggestionPrompts': meta.get('suggestion_prompts'),
'categories': meta.get('categories'),
'user': {k: v for k, v in meta.get('user', {}).items() if k != 'community'},
}
# Insert the processed data back into the 'modelfile' table
@@ -37,11 +37,11 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
# Adding fields settings to the 'user' table
migrator.add_fields("user", settings=pw.TextField(null=True))
migrator.add_fields('user', settings=pw.TextField(null=True))
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
# Remove the settings field
migrator.remove_fields("user", "settings")
migrator.remove_fields('user', 'settings')
@@ -51,10 +51,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
updated_at = pw.BigIntegerField(null=False)
class Meta:
table_name = "tool"
table_name = 'tool'
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_model("tool")
migrator.remove_model('tool')
@@ -37,11 +37,11 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
# Adding fields info to the 'user' table
migrator.add_fields("user", info=pw.TextField(null=True))
migrator.add_fields('user', info=pw.TextField(null=True))
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
# Remove the settings field
migrator.remove_fields("user", "info")
migrator.remove_fields('user', 'info')
@@ -45,10 +45,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
created_at = pw.BigIntegerField(null=False)
class Meta:
table_name = "file"
table_name = 'file'
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_model("file")
migrator.remove_model('file')
@@ -51,10 +51,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
updated_at = pw.BigIntegerField(null=False)
class Meta:
table_name = "function"
table_name = 'function'
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_model("function")
migrator.remove_model('function')
@@ -36,14 +36,14 @@ with suppress(ImportError):
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
migrator.add_fields("tool", valves=pw.TextField(null=True))
migrator.add_fields("function", valves=pw.TextField(null=True))
migrator.add_fields("function", is_active=pw.BooleanField(default=False))
migrator.add_fields('tool', valves=pw.TextField(null=True))
migrator.add_fields('function', valves=pw.TextField(null=True))
migrator.add_fields('function', is_active=pw.BooleanField(default=False))
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_fields("tool", "valves")
migrator.remove_fields("function", "valves")
migrator.remove_fields("function", "is_active")
migrator.remove_fields('tool', 'valves')
migrator.remove_fields('function', 'valves')
migrator.remove_fields('function', 'is_active')
@@ -33,7 +33,7 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
migrator.add_fields(
"user",
'user',
oauth_sub=pw.TextField(null=True, unique=True),
)
@@ -41,4 +41,4 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_fields("user", "oauth_sub")
migrator.remove_fields('user', 'oauth_sub')
@@ -37,7 +37,7 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your migrations here."""
migrator.add_fields(
"function",
'function',
is_global=pw.BooleanField(default=False),
)
@@ -45,4 +45,4 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
"""Write your rollback migrations here."""
migrator.remove_fields("function", "is_global")
migrator.remove_fields('function', 'is_global')
+15 -17
View File
@@ -10,13 +10,13 @@ from playhouse.shortcuts import ReconnectMixin
log = logging.getLogger(__name__)
db_state_default = {"closed": None, "conn": None, "ctx": None, "transactions": None}
db_state = ContextVar("db_state", default=db_state_default.copy())
db_state_default = {'closed': None, 'conn': None, 'ctx': None, 'transactions': None}
db_state = ContextVar('db_state', default=db_state_default.copy())
class PeeweeConnectionState(object):
def __init__(self, **kwargs):
super().__setattr__("_state", db_state)
super().__setattr__('_state', db_state)
super().__init__(**kwargs)
def __setattr__(self, name, value):
@@ -30,10 +30,10 @@ class PeeweeConnectionState(object):
class CustomReconnectMixin(ReconnectMixin):
reconnect_errors = (
# psycopg2
(OperationalError, "termin"),
(InterfaceError, "closed"),
(OperationalError, 'termin'),
(InterfaceError, 'closed'),
# peewee
(PeeWeeInterfaceError, "closed"),
(PeeWeeInterfaceError, 'closed'),
)
@@ -43,23 +43,21 @@ class ReconnectingPostgresqlDatabase(CustomReconnectMixin, PostgresqlDatabase):
def register_connection(db_url):
# Check if using SQLCipher protocol
if db_url.startswith("sqlite+sqlcipher://"):
database_password = os.environ.get("DATABASE_PASSWORD")
if not database_password or database_password.strip() == "":
raise ValueError(
"DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs"
)
if db_url.startswith('sqlite+sqlcipher://'):
database_password = os.environ.get('DATABASE_PASSWORD')
if not database_password or database_password.strip() == '':
raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
from playhouse.sqlcipher_ext import SqlCipherDatabase
# Parse the database path from SQLCipher URL
# Convert sqlite+sqlcipher:///path/to/db.sqlite to /path/to/db.sqlite
db_path = db_url.replace("sqlite+sqlcipher://", "")
db_path = db_url.replace('sqlite+sqlcipher://', '')
# Use Peewee's native SqlCipherDatabase with encryption
db = SqlCipherDatabase(db_path, passphrase=database_password)
db.autoconnect = True
db.reuse_if_open = True
log.info("Connected to encrypted SQLite database using SQLCipher")
log.info('Connected to encrypted SQLite database using SQLCipher')
else:
# Standard database connection (existing logic)
@@ -68,7 +66,7 @@ def register_connection(db_url):
# Enable autoconnect for SQLite databases, managed by Peewee
db.autoconnect = True
db.reuse_if_open = True
log.info("Connected to PostgreSQL database")
log.info('Connected to PostgreSQL database')
# Get the connection details
connection = parse(db_url, unquote_user=True, unquote_password=True)
@@ -80,7 +78,7 @@ def register_connection(db_url):
# Enable autoconnect for SQLite databases, managed by Peewee
db.autoconnect = True
db.reuse_if_open = True
log.info("Connected to SQLite database")
log.info('Connected to SQLite database')
else:
raise ValueError("Unsupported database connection")
raise ValueError('Unsupported database connection')
return db