refac
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
@@ -35,7 +36,7 @@ from open_webui.env import (
|
||||
WEBUI_NAME,
|
||||
log,
|
||||
)
|
||||
from open_webui.internal.db import Base, get_db
|
||||
from open_webui.internal.db import Base, get_db, get_async_db
|
||||
from open_webui.utils.redis import get_redis_connection
|
||||
|
||||
|
||||
@@ -90,6 +91,7 @@ def load_json_config():
|
||||
|
||||
|
||||
def save_to_db(data):
|
||||
"""Sync save — used ONLY at startup/import time."""
|
||||
with get_db() as db:
|
||||
existing_config = db.query(Config).first()
|
||||
if not existing_config:
|
||||
@@ -102,12 +104,39 @@ def save_to_db(data):
|
||||
db.commit()
|
||||
|
||||
|
||||
async def async_save_to_db(data):
|
||||
"""Async save — used for ALL runtime config persistence."""
|
||||
from sqlalchemy import select
|
||||
|
||||
async with get_async_db() as db:
|
||||
result = await db.execute(select(Config).limit(1))
|
||||
existing_config = result.scalars().first()
|
||||
if not existing_config:
|
||||
new_config = Config(data=data, version=0)
|
||||
db.add(new_config)
|
||||
else:
|
||||
existing_config.data = data
|
||||
existing_config.updated_at = datetime.now()
|
||||
db.add(existing_config)
|
||||
await db.commit()
|
||||
|
||||
|
||||
def reset_config():
|
||||
"""Sync reset — used ONLY at startup."""
|
||||
with get_db() as db:
|
||||
db.query(Config).delete()
|
||||
db.commit()
|
||||
|
||||
|
||||
async def async_reset_config():
|
||||
"""Async reset — used at runtime."""
|
||||
from sqlalchemy import delete as sa_delete
|
||||
|
||||
async with get_async_db() as db:
|
||||
await db.execute(sa_delete(Config))
|
||||
await db.commit()
|
||||
|
||||
|
||||
# When initializing, check if config.json exists and migrate it to the database
|
||||
if os.path.exists(f'{DATA_DIR}/config.json'):
|
||||
data = load_json_config()
|
||||
@@ -144,6 +173,7 @@ PERSISTENT_CONFIG_REGISTRY = []
|
||||
|
||||
|
||||
def save_config(config):
|
||||
"""Sync save — used ONLY at startup/import time."""
|
||||
global CONFIG_DATA
|
||||
global PERSISTENT_CONFIG_REGISTRY
|
||||
try:
|
||||
@@ -159,6 +189,23 @@ def save_config(config):
|
||||
return True
|
||||
|
||||
|
||||
async def async_save_config(config):
|
||||
"""Async save — used for ALL runtime config persistence."""
|
||||
global CONFIG_DATA
|
||||
global PERSISTENT_CONFIG_REGISTRY
|
||||
try:
|
||||
await async_save_to_db(config)
|
||||
CONFIG_DATA = config
|
||||
|
||||
# Trigger updates on all registered PersistentConfig entries
|
||||
for config_item in PERSISTENT_CONFIG_REGISTRY:
|
||||
config_item.update()
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
T = TypeVar('T')
|
||||
|
||||
ENABLE_PERSISTENT_CONFIG = os.environ.get('ENABLE_PERSISTENT_CONFIG', 'True').lower() == 'true'
|
||||
@@ -202,6 +249,7 @@ class PersistentConfig(Generic[T]):
|
||||
log.info(f'Updated {self.env_name} to new value {self.value}')
|
||||
|
||||
def save(self):
|
||||
"""Sync save — used ONLY at startup/import time."""
|
||||
log.info(f"Saving '{self.env_name}' to the database")
|
||||
path_parts = self.config_path.split('.')
|
||||
sub_config = CONFIG_DATA
|
||||
@@ -213,6 +261,19 @@ class PersistentConfig(Generic[T]):
|
||||
save_to_db(CONFIG_DATA)
|
||||
self.config_value = self.value
|
||||
|
||||
async def async_save(self):
|
||||
"""Async save — used for ALL runtime config persistence."""
|
||||
log.info(f"Saving '{self.env_name}' to the database")
|
||||
path_parts = self.config_path.split('.')
|
||||
sub_config = CONFIG_DATA
|
||||
for key in path_parts[:-1]:
|
||||
if key not in sub_config:
|
||||
sub_config[key] = {}
|
||||
sub_config = sub_config[key]
|
||||
sub_config[path_parts[-1]] = self.value
|
||||
await async_save_to_db(CONFIG_DATA)
|
||||
self.config_value = self.value
|
||||
|
||||
|
||||
class AppConfig:
|
||||
_redis: Union[redis.Redis, redis.cluster.RedisCluster] = None
|
||||
@@ -246,12 +307,27 @@ class AppConfig:
|
||||
self._state[key] = value
|
||||
else:
|
||||
self._state[key].value = value
|
||||
self._state[key].save()
|
||||
|
||||
# At runtime (inside the event loop) persist via the async engine
|
||||
# to avoid blocking the loop and contending with the async DB pool.
|
||||
# At startup/import time, fall back to sync.
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
loop.create_task(self._async_persist(key))
|
||||
except RuntimeError:
|
||||
self._state[key].save()
|
||||
|
||||
if self._redis and ENABLE_PERSISTENT_CONFIG:
|
||||
redis_key = f'{self._redis_key_prefix}:config:{key}'
|
||||
self._redis.set(redis_key, json.dumps(self._state[key].value))
|
||||
|
||||
async def _async_persist(self, key):
|
||||
"""Persist a single config key via the async engine."""
|
||||
try:
|
||||
await self._state[key].async_save()
|
||||
except Exception as e:
|
||||
log.error(f'Failed to async-persist config key {key}: {e}')
|
||||
|
||||
def __getattr__(self, key):
|
||||
if key not in self._state:
|
||||
raise AttributeError(f"Config key '{key}' not found")
|
||||
|
||||
@@ -470,6 +470,7 @@ from open_webui.config import (
|
||||
AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH,
|
||||
AppConfig,
|
||||
reset_config,
|
||||
async_reset_config,
|
||||
)
|
||||
from open_webui.env import (
|
||||
ENABLE_CUSTOM_MODEL_FALLBACK,
|
||||
@@ -625,7 +626,7 @@ async def lifespan(app: FastAPI):
|
||||
start_logger()
|
||||
|
||||
if RESET_CONFIG_ON_START:
|
||||
reset_config()
|
||||
await async_reset_config()
|
||||
|
||||
if LICENSE_KEY:
|
||||
get_license_data(app, LICENSE_KEY)
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Optional
|
||||
|
||||
from open_webui.env import AIOHTTP_CLIENT_TIMEOUT
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.config import get_config, save_config
|
||||
from open_webui.config import get_config, save_config, async_save_config
|
||||
from open_webui.config import BannerModel
|
||||
|
||||
from open_webui.utils.tools import (
|
||||
@@ -49,7 +49,7 @@ class ImportConfigForm(BaseModel):
|
||||
|
||||
@router.post('/import', response_model=dict)
|
||||
async def import_config(form_data: ImportConfigForm, user=Depends(get_admin_user)):
|
||||
save_config(form_data.config)
|
||||
await async_save_config(form_data.config)
|
||||
return get_config()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user