mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-13 01:02:25 -06:00
refac
This commit is contained in:
+809
-1906
File diff suppressed because it is too large
Load Diff
@@ -1,265 +0,0 @@
|
||||
"""Database-backed configuration with environment variable defaults."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from functools import reduce
|
||||
from typing import Any, Optional, Union
|
||||
|
||||
import redis
|
||||
from open_webui.internal.db import Base, get_async_db, get_db
|
||||
from open_webui.utils.redis import get_redis_connection
|
||||
from sqlalchemy import JSON, Column, DateTime, Integer, func, select
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── Model ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ConfigTable(Base):
|
||||
__tablename__ = 'config'
|
||||
|
||||
id = Column(Integer, primary_key=True)
|
||||
data = Column(JSON, nullable=False)
|
||||
version = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, server_default=func.now())
|
||||
updated_at = Column(DateTime, nullable=True, onupdate=func.now())
|
||||
|
||||
|
||||
# ── Blob ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ConfigState:
|
||||
"""In-memory mirror of the single-row config JSON blob."""
|
||||
|
||||
__slots__ = ('_data',)
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._data: dict[str, Any] = {}
|
||||
|
||||
@property
|
||||
def snapshot(self) -> dict:
|
||||
return self._data
|
||||
|
||||
def read(self, path: str) -> Any:
|
||||
return reduce(
|
||||
lambda n, k: n.get(k) if isinstance(n, dict) else None,
|
||||
path.split('.'),
|
||||
self._data,
|
||||
)
|
||||
|
||||
def write(self, path: str, value: Any) -> None:
|
||||
keys = path.split('.')
|
||||
reduce(lambda d, k: d.setdefault(k, {}), keys[:-1], self._data)[keys[-1]] = value
|
||||
|
||||
def replace(self, data: dict) -> None:
|
||||
self._data = data
|
||||
|
||||
def load(self) -> dict:
|
||||
with get_db() as db:
|
||||
row = db.query(ConfigTable).order_by(ConfigTable.id.desc()).first()
|
||||
self._data = row.data if row else {'version': 0, 'ui': {}}
|
||||
return self._data
|
||||
|
||||
def persist(self, data: dict | None = None) -> None:
|
||||
if data is not None:
|
||||
self._data = data
|
||||
with get_db() as db:
|
||||
row = db.query(ConfigTable).first()
|
||||
if row is None:
|
||||
db.add(ConfigTable(data=self._data, version=0))
|
||||
else:
|
||||
row.data, row.updated_at = self._data, datetime.now()
|
||||
db.add(row)
|
||||
db.commit()
|
||||
|
||||
async def persist_async(self, data: dict | None = None) -> None:
|
||||
if data is not None:
|
||||
self._data = data
|
||||
async with get_async_db() as db:
|
||||
result = await db.execute(select(ConfigTable).limit(1))
|
||||
row = result.scalars().first()
|
||||
if row is None:
|
||||
db.add(ConfigTable(data=self._data, version=0))
|
||||
else:
|
||||
row.data, row.updated_at = self._data, datetime.now()
|
||||
db.add(row)
|
||||
await db.commit()
|
||||
|
||||
def clear(self) -> None:
|
||||
with get_db() as db:
|
||||
db.query(ConfigTable).delete()
|
||||
db.commit()
|
||||
|
||||
async def clear_async(self) -> None:
|
||||
from sqlalchemy import delete as sa_delete
|
||||
|
||||
async with get_async_db() as db:
|
||||
await db.execute(sa_delete(ConfigTable))
|
||||
await db.commit()
|
||||
|
||||
|
||||
STATE = ConfigState()
|
||||
|
||||
|
||||
# ── ConfigVar ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
_persist_enabled: bool = True
|
||||
_oauth_persist_enabled: bool = False
|
||||
_all_configs: list[ConfigVar] = []
|
||||
|
||||
|
||||
def initialize(*, enable_persistent: bool = True, enable_oauth_persistent: bool = False) -> dict:
|
||||
global _persist_enabled, _oauth_persist_enabled
|
||||
_persist_enabled = enable_persistent
|
||||
_oauth_persist_enabled = enable_oauth_persistent
|
||||
return STATE.load()
|
||||
|
||||
|
||||
class ConfigVar:
|
||||
__slots__ = ('env_name', 'config_path', 'env_value', 'config_value', 'value')
|
||||
|
||||
def __init__(self, env_name: str, config_path: str, env_value: Any) -> None:
|
||||
self.env_name = env_name
|
||||
self.config_path = config_path
|
||||
self.env_value = env_value
|
||||
self.config_value = STATE.read(config_path)
|
||||
|
||||
if self.config_value is not None and _persist_enabled:
|
||||
if config_path.startswith('oauth.') and not _oauth_persist_enabled:
|
||||
log.info("Skipping DB value for '%s' (OAuth persistence disabled)", env_name)
|
||||
self.value = env_value
|
||||
else:
|
||||
log.info("'%s' loaded from database", env_name)
|
||||
self.value = self.config_value
|
||||
else:
|
||||
self.value = env_value
|
||||
|
||||
_all_configs.append(self)
|
||||
|
||||
def __str__(self) -> str:
|
||||
return str(self.value)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f'<ConfigVar {self.env_name}={self.value!r}>'
|
||||
|
||||
@property
|
||||
def __dict__(self): # type: ignore[override]
|
||||
raise TypeError(f"ConfigVar('{self.env_name}') cannot be cast to dict; use .value")
|
||||
|
||||
def __getattribute__(self, item: str):
|
||||
if item == '__dict__':
|
||||
raise TypeError('ConfigVar cannot be cast to dict; use .value')
|
||||
return super().__getattribute__(item)
|
||||
|
||||
def refresh(self) -> None:
|
||||
current = STATE.read(self.config_path)
|
||||
if current is not None:
|
||||
self.value = current
|
||||
log.info('Refreshed %s → %s', self.env_name, self.value)
|
||||
|
||||
def commit(self) -> None:
|
||||
log.info("Persisting '%s'", self.env_name)
|
||||
STATE.write(self.config_path, self.value)
|
||||
self.config_value = self.value
|
||||
STATE.persist()
|
||||
|
||||
async def commit_async(self) -> None:
|
||||
log.info("Persisting '%s'", self.env_name)
|
||||
STATE.write(self.config_path, self.value)
|
||||
self.config_value = self.value
|
||||
await STATE.persist_async()
|
||||
|
||||
|
||||
# ── AppConfig ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AppConfig:
|
||||
"""Attribute-style container for ConfigVars with optional Redis sync."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
redis_url: Optional[str] = None,
|
||||
redis_sentinels: Optional[list] = None,
|
||||
redis_cluster: bool = False,
|
||||
redis_key_prefix: str = 'open-webui',
|
||||
) -> None:
|
||||
super().__setattr__('_entries', {})
|
||||
super().__setattr__('_key_prefix', redis_key_prefix)
|
||||
|
||||
# If sentinels weren't explicitly provided, read from env.
|
||||
if redis_sentinels is None:
|
||||
from open_webui.env import REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT
|
||||
from open_webui.utils.redis import get_sentinels_from_env
|
||||
|
||||
redis_sentinels = get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT)
|
||||
|
||||
rc: Union[redis.Redis, redis.cluster.RedisCluster, None] = None
|
||||
if redis_url:
|
||||
rc = get_redis_connection(redis_url, redis_sentinels or [], redis_cluster, decode_responses=True)
|
||||
super().__setattr__('_rc', rc)
|
||||
|
||||
def __setattr__(self, name: str, value: Any) -> None:
|
||||
entries: dict = super().__getattribute__('_entries')
|
||||
|
||||
if isinstance(value, ConfigVar):
|
||||
entries[name] = value
|
||||
return
|
||||
|
||||
entries[name].value = value
|
||||
|
||||
try:
|
||||
asyncio.get_running_loop().create_task(self._write_async(name))
|
||||
except RuntimeError:
|
||||
entries[name].commit()
|
||||
|
||||
rc = super().__getattribute__('_rc')
|
||||
if rc and _persist_enabled:
|
||||
prefix = super().__getattribute__('_key_prefix')
|
||||
try:
|
||||
rc.set(f'{prefix}:config:{name}', json.dumps(entries[name].value))
|
||||
except Exception as exc:
|
||||
log.error("Redis write failed for '%s': %s", name, exc)
|
||||
|
||||
async def _write_async(self, name: str) -> None:
|
||||
try:
|
||||
await self._entries[name].commit_async()
|
||||
except Exception as exc:
|
||||
log.error("Async persist failed for '%s': %s", name, exc)
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
entries = super().__getattribute__('_entries')
|
||||
if name not in entries:
|
||||
raise AttributeError(f"No config key '{name}'")
|
||||
|
||||
rc = super().__getattribute__('_rc')
|
||||
if rc and _persist_enabled:
|
||||
prefix = super().__getattribute__('_key_prefix')
|
||||
try:
|
||||
raw = rc.get(f'{prefix}:config:{name}')
|
||||
if raw is not None:
|
||||
decoded = json.loads(raw)
|
||||
if entries[name].value != decoded:
|
||||
entries[name].value = decoded
|
||||
log.info("Updated '%s' from Redis", name)
|
||||
except Exception as exc:
|
||||
log.error("Redis read failed for '%s': %s", name, exc)
|
||||
|
||||
return entries[name].value
|
||||
|
||||
def _sync_to_redis(self) -> None:
|
||||
rc = super().__getattribute__('_rc')
|
||||
if not rc or not _persist_enabled:
|
||||
return
|
||||
prefix = super().__getattribute__('_key_prefix')
|
||||
for name, s in super().__getattribute__('_entries').items():
|
||||
try:
|
||||
rc.set(f'{prefix}:config:{name}', json.dumps(s.value))
|
||||
except Exception as exc:
|
||||
log.error("Redis sync failed for '%s': %s", name, exc)
|
||||
+276
-914
File diff suppressed because it is too large
Load Diff
+580
@@ -0,0 +1,580 @@
|
||||
"""reshape config to per key rows
|
||||
|
||||
Revision ID: 3ff2c63645b8
|
||||
Revises: 461111b60977
|
||||
Create Date: 2026-06-17 00:50:51.477073
|
||||
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '3ff2c63645b8'
|
||||
down_revision: Union[str, None] = '461111b60977'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
# Maps every dot-notation blob path to its legacy env/config key name.
|
||||
# Built from the legacy persistent config declarations in config.py.
|
||||
BLOB_PATH_TO_KEY = {
|
||||
"audio.stt.allowed_extensions": "AUDIO_STT_ALLOWED_EXTENSIONS",
|
||||
"audio.stt.azure.api_key": "AUDIO_STT_AZURE_API_KEY",
|
||||
"audio.stt.azure.base_url": "AUDIO_STT_AZURE_BASE_URL",
|
||||
"audio.stt.azure.locales": "AUDIO_STT_AZURE_LOCALES",
|
||||
"audio.stt.azure.max_speakers": "AUDIO_STT_AZURE_MAX_SPEAKERS",
|
||||
"audio.stt.azure.region": "AUDIO_STT_AZURE_REGION",
|
||||
"audio.stt.deepgram.api_key": "DEEPGRAM_API_KEY",
|
||||
"audio.stt.engine": "AUDIO_STT_ENGINE",
|
||||
"audio.stt.mistral.api_base_url": "AUDIO_STT_MISTRAL_API_BASE_URL",
|
||||
"audio.stt.mistral.api_key": "AUDIO_STT_MISTRAL_API_KEY",
|
||||
"audio.stt.mistral.use_chat_completions": "AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS",
|
||||
"audio.stt.model": "AUDIO_STT_MODEL",
|
||||
"audio.stt.openai.api_base_url": "AUDIO_STT_OPENAI_API_BASE_URL",
|
||||
"audio.stt.openai.api_key": "AUDIO_STT_OPENAI_API_KEY",
|
||||
"audio.stt.supported_content_types": "AUDIO_STT_SUPPORTED_CONTENT_TYPES",
|
||||
"audio.stt.whisper_model": "WHISPER_MODEL",
|
||||
"audio.tts.api_key": "AUDIO_TTS_API_KEY",
|
||||
"audio.tts.azure.speech_base_url": "AUDIO_TTS_AZURE_SPEECH_BASE_URL",
|
||||
"audio.tts.azure.speech_output_format": "AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT",
|
||||
"audio.tts.azure.speech_region": "AUDIO_TTS_AZURE_SPEECH_REGION",
|
||||
"audio.tts.engine": "AUDIO_TTS_ENGINE",
|
||||
"audio.tts.mistral.api_base_url": "AUDIO_TTS_MISTRAL_API_BASE_URL",
|
||||
"audio.tts.mistral.api_key": "AUDIO_TTS_MISTRAL_API_KEY",
|
||||
"audio.tts.model": "AUDIO_TTS_MODEL",
|
||||
"audio.tts.openai.api_base_url": "AUDIO_TTS_OPENAI_API_BASE_URL",
|
||||
"audio.tts.openai.api_key": "AUDIO_TTS_OPENAI_API_KEY",
|
||||
"audio.tts.openai.params": "AUDIO_TTS_OPENAI_PARAMS",
|
||||
"audio.tts.split_on": "AUDIO_TTS_SPLIT_ON",
|
||||
"audio.tts.voice": "AUDIO_TTS_VOICE",
|
||||
"auth.admin.email": "ADMIN_EMAIL",
|
||||
"auth.admin.show": "SHOW_ADMIN_DETAILS",
|
||||
"auth.api_key.allowed_endpoints": "API_KEYS_ALLOWED_ENDPOINTS",
|
||||
"auth.api_key.endpoint_restrictions": "ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS",
|
||||
"auth.enable_api_keys": "ENABLE_API_KEYS",
|
||||
"auth.jwt_expiry": "JWT_EXPIRES_IN",
|
||||
"automations.enable": "ENABLE_AUTOMATIONS",
|
||||
"automations.max_count": "AUTOMATION_MAX_COUNT",
|
||||
"automations.min_interval": "AUTOMATION_MIN_INTERVAL",
|
||||
"calendar.enable": "ENABLE_CALENDAR",
|
||||
"channels.enable": "ENABLE_CHANNELS",
|
||||
"code_execution.enable": "ENABLE_CODE_EXECUTION",
|
||||
"code_execution.engine": "CODE_EXECUTION_ENGINE",
|
||||
"code_execution.jupyter.auth": "CODE_EXECUTION_JUPYTER_AUTH",
|
||||
"code_execution.jupyter.auth_password": "CODE_EXECUTION_JUPYTER_AUTH_PASSWORD",
|
||||
"code_execution.jupyter.auth_token": "CODE_EXECUTION_JUPYTER_AUTH_TOKEN",
|
||||
"code_execution.jupyter.timeout": "CODE_EXECUTION_JUPYTER_TIMEOUT",
|
||||
"code_execution.jupyter.url": "CODE_EXECUTION_JUPYTER_URL",
|
||||
"code_interpreter.enable": "ENABLE_CODE_INTERPRETER",
|
||||
"code_interpreter.engine": "CODE_INTERPRETER_ENGINE",
|
||||
"code_interpreter.jupyter.auth": "CODE_INTERPRETER_JUPYTER_AUTH",
|
||||
"code_interpreter.jupyter.auth_password": "CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD",
|
||||
"code_interpreter.jupyter.auth_token": "CODE_INTERPRETER_JUPYTER_AUTH_TOKEN",
|
||||
"code_interpreter.jupyter.timeout": "CODE_INTERPRETER_JUPYTER_TIMEOUT",
|
||||
"code_interpreter.jupyter.url": "CODE_INTERPRETER_JUPYTER_URL",
|
||||
"code_interpreter.prompt_template": "CODE_INTERPRETER_PROMPT_TEMPLATE",
|
||||
"direct.enable": "ENABLE_DIRECT_CONNECTIONS",
|
||||
"evaluation.arena.enable": "ENABLE_EVALUATION_ARENA_MODELS",
|
||||
"evaluation.arena.models": "EVALUATION_ARENA_MODELS",
|
||||
"file.image_compression_height": "FILE_IMAGE_COMPRESSION_HEIGHT",
|
||||
"file.image_compression_width": "FILE_IMAGE_COMPRESSION_WIDTH",
|
||||
"folders.enable": "ENABLE_FOLDERS",
|
||||
"folders.max_file_count": "FOLDER_MAX_FILE_COUNT",
|
||||
"google_drive.api_key": "GOOGLE_DRIVE_API_KEY",
|
||||
"google_drive.client_id": "GOOGLE_DRIVE_CLIENT_ID",
|
||||
"google_drive.enable": "ENABLE_GOOGLE_DRIVE_INTEGRATION",
|
||||
"image_generation.automatic1111.api_auth": "AUTOMATIC1111_API_AUTH",
|
||||
"image_generation.automatic1111.api_params": "AUTOMATIC1111_PARAMS",
|
||||
"image_generation.automatic1111.base_url": "AUTOMATIC1111_BASE_URL",
|
||||
"image_generation.comfyui.api_key": "COMFYUI_API_KEY",
|
||||
"image_generation.comfyui.base_url": "COMFYUI_BASE_URL",
|
||||
"image_generation.comfyui.nodes": "COMFYUI_WORKFLOW_NODES",
|
||||
"image_generation.comfyui.workflow": "COMFYUI_WORKFLOW",
|
||||
"image_generation.enable": "ENABLE_IMAGE_GENERATION",
|
||||
"image_generation.engine": "IMAGE_GENERATION_ENGINE",
|
||||
"image_generation.gemini.api_base_url": "IMAGES_GEMINI_API_BASE_URL",
|
||||
"image_generation.gemini.api_key": "IMAGES_GEMINI_API_KEY",
|
||||
"image_generation.gemini.endpoint_method": "IMAGES_GEMINI_ENDPOINT_METHOD",
|
||||
"image_generation.model": "IMAGE_GENERATION_MODEL",
|
||||
"image_generation.openai.api_base_url": "IMAGES_OPENAI_API_BASE_URL",
|
||||
"image_generation.openai.api_key": "IMAGES_OPENAI_API_KEY",
|
||||
"image_generation.openai.api_version": "IMAGES_OPENAI_API_VERSION",
|
||||
"image_generation.openai.params": "IMAGES_OPENAI_API_PARAMS",
|
||||
"image_generation.prompt.enable": "ENABLE_IMAGE_PROMPT_GENERATION",
|
||||
"image_generation.size": "IMAGE_SIZE",
|
||||
"image_generation.steps": "IMAGE_STEPS",
|
||||
"images.edit.comfyui.api_key": "IMAGES_EDIT_COMFYUI_API_KEY",
|
||||
"images.edit.comfyui.base_url": "IMAGES_EDIT_COMFYUI_BASE_URL",
|
||||
"images.edit.comfyui.nodes": "IMAGES_EDIT_COMFYUI_WORKFLOW_NODES",
|
||||
"images.edit.comfyui.workflow": "IMAGES_EDIT_COMFYUI_WORKFLOW",
|
||||
"images.edit.enable": "ENABLE_IMAGE_EDIT",
|
||||
"images.edit.engine": "IMAGE_EDIT_ENGINE",
|
||||
"images.edit.gemini.api_base_url": "IMAGES_EDIT_GEMINI_API_BASE_URL",
|
||||
"images.edit.gemini.api_key": "IMAGES_EDIT_GEMINI_API_KEY",
|
||||
"images.edit.model": "IMAGE_EDIT_MODEL",
|
||||
"images.edit.openai.api_base_url": "IMAGES_EDIT_OPENAI_API_BASE_URL",
|
||||
"images.edit.openai.api_key": "IMAGES_EDIT_OPENAI_API_KEY",
|
||||
"images.edit.openai.api_version": "IMAGES_EDIT_OPENAI_API_VERSION",
|
||||
"images.edit.size": "IMAGE_EDIT_SIZE",
|
||||
"ldap.enable": "ENABLE_LDAP",
|
||||
"ldap.group.enable_creation": "ENABLE_LDAP_GROUP_CREATION",
|
||||
"ldap.group.enable_management": "ENABLE_LDAP_GROUP_MANAGEMENT",
|
||||
"ldap.server.app_dn": "LDAP_APP_DN",
|
||||
"ldap.server.app_password": "LDAP_APP_PASSWORD",
|
||||
"ldap.server.attribute_for_groups": "LDAP_ATTRIBUTE_FOR_GROUPS",
|
||||
"ldap.server.attribute_for_mail": "LDAP_ATTRIBUTE_FOR_MAIL",
|
||||
"ldap.server.attribute_for_username": "LDAP_ATTRIBUTE_FOR_USERNAME",
|
||||
"ldap.server.ca_cert_file": "LDAP_CA_CERT_FILE",
|
||||
"ldap.server.ciphers": "LDAP_CIPHERS",
|
||||
"ldap.server.host": "LDAP_SERVER_HOST",
|
||||
"ldap.server.label": "LDAP_SERVER_LABEL",
|
||||
"ldap.server.port": "LDAP_SERVER_PORT",
|
||||
"ldap.server.search_filter": "LDAP_SEARCH_FILTER",
|
||||
"ldap.server.use_tls": "LDAP_USE_TLS",
|
||||
"ldap.server.users_dn": "LDAP_SEARCH_BASE",
|
||||
"ldap.server.validate_cert": "LDAP_VALIDATE_CERT",
|
||||
"memories.enable": "ENABLE_MEMORIES",
|
||||
"models.base_models_cache": "ENABLE_BASE_MODELS_CACHE",
|
||||
"models.default_metadata": "DEFAULT_MODEL_METADATA",
|
||||
"models.default_params": "DEFAULT_MODEL_PARAMS",
|
||||
"notes.enable": "ENABLE_NOTES",
|
||||
# OAuth — direct paths
|
||||
"oauth.admin_roles": "OAUTH_ADMIN_ROLES",
|
||||
"oauth.allowed_domains": "OAUTH_ALLOWED_DOMAINS",
|
||||
"oauth.allowed_roles": "OAUTH_ALLOWED_ROLES",
|
||||
"oauth.audience": "OAUTH_AUDIENCE",
|
||||
"oauth.auto_redirect": "OAUTH_AUTO_REDIRECT",
|
||||
"oauth.blocked_groups": "OAUTH_BLOCKED_GROUPS",
|
||||
"oauth.client.timeout": "OAUTH_CLIENT_TIMEOUT",
|
||||
"oauth.enable_group_creation": "ENABLE_OAUTH_GROUP_CREATION",
|
||||
"oauth.enable_group_mapping": "ENABLE_OAUTH_GROUP_MANAGEMENT",
|
||||
"oauth.enable_role_mapping": "ENABLE_OAUTH_ROLE_MANAGEMENT",
|
||||
"oauth.enable_signup": "ENABLE_OAUTH_SIGNUP",
|
||||
"oauth.group_default_share": "OAUTH_GROUP_DEFAULT_SHARE",
|
||||
"oauth.merge_accounts_by_email": "OAUTH_MERGE_ACCOUNTS_BY_EMAIL",
|
||||
"oauth.refresh_token_include_scope": "OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE",
|
||||
"oauth.roles_claim": "OAUTH_ROLES_CLAIM",
|
||||
"oauth.update_email_on_login": "OAUTH_UPDATE_EMAIL_ON_LOGIN",
|
||||
"oauth.update_name_on_login": "OAUTH_UPDATE_NAME_ON_LOGIN",
|
||||
"oauth.update_picture_on_login": "OAUTH_UPDATE_PICTURE_ON_LOGIN",
|
||||
# OAuth — generic provider paths
|
||||
"oauth.client_id": "OAUTH_CLIENT_ID",
|
||||
"oauth.client_secret": "OAUTH_CLIENT_SECRET",
|
||||
"oauth.code_challenge_method": "OAUTH_CODE_CHALLENGE_METHOD",
|
||||
"oauth.email_claim": "OAUTH_EMAIL_CLAIM",
|
||||
"oauth.end_session_endpoint": "OPENID_END_SESSION_ENDPOINT",
|
||||
"oauth.group_claim": "OAUTH_GROUP_CLAIM",
|
||||
"oauth.picture_claim": "OAUTH_PICTURE_CLAIM",
|
||||
"oauth.provider_name": "OAUTH_PROVIDER_NAME",
|
||||
"oauth.provider_url": "OPENID_PROVIDER_URL",
|
||||
"oauth.redirect_uri": "OPENID_REDIRECT_URI",
|
||||
"oauth.scopes": "OAUTH_SCOPES",
|
||||
"oauth.sub_claim": "OAUTH_SUB_CLAIM",
|
||||
"oauth.timeout": "OAUTH_TIMEOUT",
|
||||
"oauth.token_endpoint_auth_method": "OAUTH_TOKEN_ENDPOINT_AUTH_METHOD",
|
||||
"oauth.username_claim": "OAUTH_USERNAME_CLAIM",
|
||||
# OAuth — OIDC nested paths (flattened)
|
||||
"oauth.oidc.avatar_claim": "OAUTH_PICTURE_CLAIM",
|
||||
"oauth.oidc.client_id": "OAUTH_CLIENT_ID",
|
||||
"oauth.oidc.client_secret": "OAUTH_CLIENT_SECRET",
|
||||
"oauth.oidc.code_challenge_method": "OAUTH_CODE_CHALLENGE_METHOD",
|
||||
"oauth.oidc.email_claim": "OAUTH_EMAIL_CLAIM",
|
||||
"oauth.oidc.end_session_endpoint": "OPENID_END_SESSION_ENDPOINT",
|
||||
"oauth.oidc.group_claim": "OAUTH_GROUP_CLAIM", # renamed from OAUTH_GROUPS_CLAIM
|
||||
"oauth.oidc.oauth_timeout": "OAUTH_TIMEOUT",
|
||||
"oauth.oidc.provider_name": "OAUTH_PROVIDER_NAME",
|
||||
"oauth.oidc.provider_url": "OPENID_PROVIDER_URL",
|
||||
"oauth.oidc.redirect_uri": "OPENID_REDIRECT_URI",
|
||||
"oauth.oidc.scopes": "OAUTH_SCOPES",
|
||||
"oauth.oidc.sub_claim": "OAUTH_SUB_CLAIM",
|
||||
"oauth.oidc.token_endpoint_auth_method": "OAUTH_TOKEN_ENDPOINT_AUTH_METHOD",
|
||||
"oauth.oidc.username_claim": "OAUTH_USERNAME_CLAIM",
|
||||
# OAuth — provider-specific
|
||||
"oauth.feishu.client_id": "FEISHU_CLIENT_ID",
|
||||
"oauth.feishu.client_secret": "FEISHU_CLIENT_SECRET",
|
||||
"oauth.feishu.redirect_uri": "FEISHU_REDIRECT_URI",
|
||||
"oauth.feishu.scope": "FEISHU_OAUTH_SCOPE",
|
||||
"oauth.github.client_id": "GITHUB_CLIENT_ID",
|
||||
"oauth.github.client_secret": "GITHUB_CLIENT_SECRET",
|
||||
"oauth.github.redirect_uri": "GITHUB_CLIENT_REDIRECT_URI",
|
||||
"oauth.github.scope": "GITHUB_CLIENT_SCOPE",
|
||||
"oauth.google.client_id": "GOOGLE_CLIENT_ID",
|
||||
"oauth.google.client_secret": "GOOGLE_CLIENT_SECRET",
|
||||
"oauth.google.redirect_uri": "GOOGLE_REDIRECT_URI",
|
||||
"oauth.google.scope": "GOOGLE_OAUTH_SCOPE",
|
||||
"oauth.microsoft.client_id": "MICROSOFT_CLIENT_ID",
|
||||
"oauth.microsoft.client_secret": "MICROSOFT_CLIENT_SECRET",
|
||||
"oauth.microsoft.login_base_url": "MICROSOFT_CLIENT_LOGIN_BASE_URL",
|
||||
"oauth.microsoft.picture_url": "MICROSOFT_CLIENT_PICTURE_URL",
|
||||
"oauth.microsoft.redirect_uri": "MICROSOFT_REDIRECT_URI",
|
||||
"oauth.microsoft.scope": "MICROSOFT_OAUTH_SCOPE",
|
||||
"oauth.microsoft.tenant_id": "MICROSOFT_CLIENT_TENANT_ID",
|
||||
# Ollama / OpenAI
|
||||
"ollama.api_configs": "OLLAMA_API_CONFIGS",
|
||||
"ollama.base_urls": "OLLAMA_BASE_URLS",
|
||||
"ollama.enable": "ENABLE_OLLAMA_API",
|
||||
"onedrive.enable": "ENABLE_ONEDRIVE_INTEGRATION",
|
||||
"onedrive.sharepoint_tenant_id": "ONEDRIVE_SHAREPOINT_TENANT_ID",
|
||||
"onedrive.sharepoint_url": "ONEDRIVE_SHAREPOINT_URL",
|
||||
"openai.api_base_urls": "OPENAI_API_BASE_URLS",
|
||||
"openai.api_configs": "OPENAI_API_CONFIGS",
|
||||
"openai.api_keys": "OPENAI_API_KEYS",
|
||||
"openai.enable": "ENABLE_OPENAI_API",
|
||||
# RAG
|
||||
"rag.content_extraction_engine": "CONTENT_EXTRACTION_ENGINE",
|
||||
"rag.datalab_marker_use_llm": "DATALAB_MARKER_USE_LLM",
|
||||
"rag.mistral_ocr_api_base_url": "MISTRAL_OCR_API_BASE_URL",
|
||||
"rag.azure_openai.api_key": "RAG_AZURE_OPENAI_API_KEY",
|
||||
"rag.azure_openai.api_version": "RAG_AZURE_OPENAI_API_VERSION",
|
||||
"rag.azure_openai.base_url": "RAG_AZURE_OPENAI_BASE_URL",
|
||||
"rag.bypass_embedding_and_retrieval": "BYPASS_EMBEDDING_AND_RETRIEVAL",
|
||||
"rag.chunk_min_size_target": "CHUNK_MIN_SIZE_TARGET",
|
||||
"rag.chunk_overlap": "CHUNK_OVERLAP",
|
||||
"rag.chunk_size": "CHUNK_SIZE",
|
||||
"rag.datalab_marker_additional_config": "DATALAB_MARKER_ADDITIONAL_CONFIG",
|
||||
"rag.datalab_marker_api_base_url": "DATALAB_MARKER_API_BASE_URL",
|
||||
"rag.datalab_marker_api_key": "DATALAB_MARKER_API_KEY",
|
||||
"rag.datalab_marker_disable_image_extraction": "DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION",
|
||||
"rag.datalab_marker_force_ocr": "DATALAB_MARKER_FORCE_OCR",
|
||||
"rag.datalab_marker_format_lines": "DATALAB_MARKER_FORMAT_LINES",
|
||||
"rag.datalab_marker_output_format": "DATALAB_MARKER_OUTPUT_FORMAT",
|
||||
"rag.datalab_marker_paginate": "DATALAB_MARKER_PAGINATE",
|
||||
"rag.datalab_marker_skip_cache": "DATALAB_MARKER_SKIP_CACHE",
|
||||
"rag.datalab_marker_strip_existing_ocr": "DATALAB_MARKER_STRIP_EXISTING_OCR",
|
||||
"rag.docling_api_key": "DOCLING_API_KEY",
|
||||
"rag.docling_params": "DOCLING_PARAMS",
|
||||
"rag.docling_server_url": "DOCLING_SERVER_URL",
|
||||
"rag.document_intelligence_endpoint": "DOCUMENT_INTELLIGENCE_ENDPOINT",
|
||||
"rag.document_intelligence_key": "DOCUMENT_INTELLIGENCE_KEY",
|
||||
"rag.document_intelligence_model": "DOCUMENT_INTELLIGENCE_MODEL",
|
||||
"rag.embedding_batch_size": "RAG_EMBEDDING_BATCH_SIZE",
|
||||
"rag.embedding_concurrent_requests": "RAG_EMBEDDING_CONCURRENT_REQUESTS",
|
||||
"rag.embedding_engine": "RAG_EMBEDDING_ENGINE",
|
||||
"rag.embedding_model": "RAG_EMBEDDING_MODEL",
|
||||
"rag.enable_async_embedding": "ENABLE_ASYNC_EMBEDDING",
|
||||
"rag.enable_hybrid_search": "ENABLE_RAG_HYBRID_SEARCH",
|
||||
"rag.enable_hybrid_search_enriched_texts": "ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS",
|
||||
"rag.enable_markdown_header_text_splitter": "ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER",
|
||||
"rag.external_document_loader_api_key": "EXTERNAL_DOCUMENT_LOADER_API_KEY",
|
||||
"rag.external_document_loader_url": "EXTERNAL_DOCUMENT_LOADER_URL",
|
||||
"rag.external_reranker_api_key": "RAG_EXTERNAL_RERANKER_API_KEY",
|
||||
"rag.external_reranker_timeout": "RAG_EXTERNAL_RERANKER_TIMEOUT",
|
||||
"rag.external_reranker_url": "RAG_EXTERNAL_RERANKER_URL",
|
||||
"rag.file.allowed_extensions": "RAG_ALLOWED_FILE_EXTENSIONS",
|
||||
"rag.file.max_count": "RAG_FILE_MAX_COUNT",
|
||||
"rag.file.max_size": "RAG_FILE_MAX_SIZE",
|
||||
"rag.full_context": "RAG_FULL_CONTEXT",
|
||||
"rag.hybrid_bm25_weight": "RAG_HYBRID_BM25_WEIGHT",
|
||||
"rag.mineru_api_key": "MINERU_API_KEY",
|
||||
"rag.mineru_api_mode": "MINERU_API_MODE",
|
||||
"rag.mineru_api_timeout": "MINERU_API_TIMEOUT",
|
||||
"rag.mineru_api_url": "MINERU_API_URL",
|
||||
"rag.mineru_file_extensions": "MINERU_FILE_EXTENSIONS",
|
||||
"rag.mineru_params": "MINERU_PARAMS",
|
||||
"rag.mistral_ocr_api_key": "MISTRAL_OCR_API_KEY",
|
||||
"rag.ollama.key": "RAG_OLLAMA_API_KEY",
|
||||
"rag.ollama.url": "RAG_OLLAMA_BASE_URL",
|
||||
"rag.openai_api_base_url": "RAG_OPENAI_API_BASE_URL",
|
||||
"rag.openai_api_key": "RAG_OPENAI_API_KEY",
|
||||
"rag.paddleocr_vl_base_url": "PADDLEOCR_VL_BASE_URL",
|
||||
"rag.paddleocr_vl_token": "PADDLEOCR_VL_TOKEN",
|
||||
"rag.pdf_extract_images": "PDF_EXTRACT_IMAGES",
|
||||
"rag.pdf_loader_mode": "PDF_LOADER_MODE",
|
||||
"rag.relevance_threshold": "RAG_RELEVANCE_THRESHOLD",
|
||||
"rag.reranking_batch_size": "RAG_RERANKING_BATCH_SIZE",
|
||||
"rag.reranking_engine": "RAG_RERANKING_ENGINE",
|
||||
"rag.reranking_model": "RAG_RERANKING_MODEL",
|
||||
"rag.template": "RAG_TEMPLATE",
|
||||
"rag.text_splitter": "RAG_TEXT_SPLITTER",
|
||||
"rag.tika_server_url": "TIKA_SERVER_URL",
|
||||
"rag.tiktoken_encoding_name": "TIKTOKEN_ENCODING_NAME",
|
||||
"rag.top_k": "RAG_TOP_K",
|
||||
"rag.top_k_reranker": "RAG_TOP_K_RERANKER",
|
||||
# RAG — Web
|
||||
"rag.web.fetch.max_content_length": "WEB_FETCH_MAX_CONTENT_LENGTH",
|
||||
"rag.web.loader.concurrent_requests": "WEB_LOADER_CONCURRENT_REQUESTS",
|
||||
"rag.web.loader.engine": "WEB_LOADER_ENGINE",
|
||||
"rag.web.loader.external_web_loader_api_key": "EXTERNAL_WEB_LOADER_API_KEY",
|
||||
"rag.web.loader.external_web_loader_url": "EXTERNAL_WEB_LOADER_URL",
|
||||
"rag.web.loader.firecrawl_api_key": "FIRECRAWL_API_KEY",
|
||||
"rag.web.loader.firecrawl_api_url": "FIRECRAWL_API_BASE_URL",
|
||||
"rag.web.loader.firecrawl_timeout": "FIRECRAWL_TIMEOUT",
|
||||
"rag.web.loader.playwright_timeout": "PLAYWRIGHT_TIMEOUT",
|
||||
"rag.web.loader.playwright_ws_url": "PLAYWRIGHT_WS_URL",
|
||||
"rag.web.loader.ssl_verification": "ENABLE_WEB_LOADER_SSL_VERIFICATION",
|
||||
"rag.web.loader.timeout": "WEB_LOADER_TIMEOUT",
|
||||
"rag.web.search.azure_ai_search_api_key": "AZURE_AI_SEARCH_API_KEY",
|
||||
"rag.web.search.azure_ai_search_endpoint": "AZURE_AI_SEARCH_ENDPOINT",
|
||||
"rag.web.search.azure_ai_search_index_name": "AZURE_AI_SEARCH_INDEX_NAME",
|
||||
"rag.web.search.bing_search_v7_endpoint": "BING_SEARCH_V7_ENDPOINT",
|
||||
"rag.web.search.bing_search_v7_subscription_key": "BING_SEARCH_V7_SUBSCRIPTION_KEY",
|
||||
"rag.web.search.bocha_search_api_key": "BOCHA_SEARCH_API_KEY",
|
||||
"rag.web.search.brave_search_api_key": "BRAVE_SEARCH_API_KEY",
|
||||
"rag.web.search.brave_search_context_tokens": "BRAVE_SEARCH_CONTEXT_TOKENS",
|
||||
"rag.web.search.bypass_embedding_and_retrieval": "BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL",
|
||||
"rag.web.search.bypass_web_loader": "BYPASS_WEB_SEARCH_WEB_LOADER",
|
||||
"rag.web.search.concurrent_requests": "WEB_SEARCH_CONCURRENT_REQUESTS",
|
||||
"rag.web.search.ddgs_backend": "DDGS_BACKEND",
|
||||
"rag.web.search.domain.filter_list": "WEB_SEARCH_DOMAIN_FILTER_LIST",
|
||||
"rag.web.search.enable": "ENABLE_WEB_SEARCH",
|
||||
"rag.web.search.engine": "WEB_SEARCH_ENGINE",
|
||||
"rag.web.search.exa_api_key": "EXA_API_KEY",
|
||||
"rag.web.search.external_web_search_api_key": "EXTERNAL_WEB_SEARCH_API_KEY",
|
||||
"rag.web.search.external_web_search_url": "EXTERNAL_WEB_SEARCH_URL",
|
||||
"rag.web.search.google_pse_api_key": "GOOGLE_PSE_API_KEY",
|
||||
"rag.web.search.google_pse_engine_id": "GOOGLE_PSE_ENGINE_ID",
|
||||
"rag.web.search.jina_api_base_url": "JINA_API_BASE_URL",
|
||||
"rag.web.search.jina_api_key": "JINA_API_KEY",
|
||||
"rag.web.search.kagi_search_api_key": "KAGI_SEARCH_API_KEY",
|
||||
"rag.web.search.linkup_api_key": "LINKUP_API_KEY",
|
||||
"rag.web.search.linkup_search_params": "LINKUP_SEARCH_PARAMS",
|
||||
"rag.web.search.mojeek_search_api_key": "MOJEEK_SEARCH_API_KEY",
|
||||
"rag.web.search.ollama_cloud_api_key": "OLLAMA_CLOUD_WEB_SEARCH_API_KEY",
|
||||
"rag.web.search.perplexity_api_key": "PERPLEXITY_API_KEY",
|
||||
"rag.web.search.perplexity_model": "PERPLEXITY_MODEL",
|
||||
"rag.web.search.perplexity_search_api_url": "PERPLEXITY_SEARCH_API_URL",
|
||||
"rag.web.search.perplexity_search_context_usage": "PERPLEXITY_SEARCH_CONTEXT_USAGE",
|
||||
"rag.web.search.result_count": "WEB_SEARCH_RESULT_COUNT",
|
||||
"rag.web.search.searchapi_api_key": "SEARCHAPI_API_KEY",
|
||||
"rag.web.search.searchapi_engine": "SEARCHAPI_ENGINE",
|
||||
"rag.web.search.searxng_language": "SEARXNG_LANGUAGE",
|
||||
"rag.web.search.searxng_query_url": "SEARXNG_QUERY_URL",
|
||||
"rag.web.search.serpapi_api_key": "SERPAPI_API_KEY",
|
||||
"rag.web.search.serpapi_engine": "SERPAPI_ENGINE",
|
||||
"rag.web.search.serper_api_key": "SERPER_API_KEY",
|
||||
"rag.web.search.serply_api_key": "SERPLY_API_KEY",
|
||||
"rag.web.search.serpstack_api_key": "SERPSTACK_API_KEY",
|
||||
"rag.web.search.serpstack_https": "SERPSTACK_HTTPS",
|
||||
"rag.web.search.sougou_api_sid": "SOUGOU_API_SID",
|
||||
"rag.web.search.sougou_api_sk": "SOUGOU_API_SK",
|
||||
"rag.web.search.tavily_api_key": "TAVILY_API_KEY",
|
||||
"rag.web.search.tavily_extract_depth": "TAVILY_EXTRACT_DEPTH",
|
||||
"rag.web.search.trust_env": "WEB_SEARCH_TRUST_ENV",
|
||||
"rag.web.search.yacy_password": "YACY_PASSWORD",
|
||||
"rag.web.search.yacy_query_url": "YACY_QUERY_URL",
|
||||
"rag.web.search.yacy_username": "YACY_USERNAME",
|
||||
"rag.web.search.yandex_web_search_api_key": "YANDEX_WEB_SEARCH_API_KEY",
|
||||
"rag.web.search.yandex_web_search_config": "YANDEX_WEB_SEARCH_CONFIG",
|
||||
"rag.web.search.yandex_web_search_url": "YANDEX_WEB_SEARCH_URL",
|
||||
"rag.web.search.youcom_api_key": "YOUCOM_API_KEY",
|
||||
"rag.youtube_loader_language": "YOUTUBE_LOADER_LANGUAGE",
|
||||
"rag.youtube_loader_proxy_url": "YOUTUBE_LOADER_PROXY_URL",
|
||||
# Tasks
|
||||
"task.autocomplete.enable": "ENABLE_AUTOCOMPLETE_GENERATION",
|
||||
"task.autocomplete.input_max_length": "AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH",
|
||||
"task.autocomplete.prompt_template": "AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE",
|
||||
"task.follow_up.enable": "ENABLE_FOLLOW_UP_GENERATION",
|
||||
"task.follow_up.prompt_template": "FOLLOW_UP_GENERATION_PROMPT_TEMPLATE",
|
||||
"task.image.prompt_template": "IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE",
|
||||
"task.model.default": "TASK_MODEL",
|
||||
"task.model.external": "TASK_MODEL_EXTERNAL",
|
||||
"task.query.prompt_template": "QUERY_GENERATION_PROMPT_TEMPLATE",
|
||||
"task.query.retrieval.enable": "ENABLE_RETRIEVAL_QUERY_GENERATION",
|
||||
"task.query.search.enable": "ENABLE_SEARCH_QUERY_GENERATION",
|
||||
"task.tags.enable": "ENABLE_TAGS_GENERATION",
|
||||
"task.tags.prompt_template": "TAGS_GENERATION_PROMPT_TEMPLATE",
|
||||
"task.title.enable": "ENABLE_TITLE_GENERATION",
|
||||
"task.title.prompt_template": "TITLE_GENERATION_PROMPT_TEMPLATE",
|
||||
"task.tools.prompt_template": "TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE",
|
||||
"task.voice.prompt.enable": "ENABLE_VOICE_MODE_PROMPT",
|
||||
"task.voice.prompt_template": "VOICE_MODE_PROMPT_TEMPLATE",
|
||||
# Misc
|
||||
"terminal_server.connections": "TERMINAL_SERVER_CONNECTIONS",
|
||||
"tool_server.connections": "TOOL_SERVER_CONNECTIONS",
|
||||
"ui.banners": "WEBUI_BANNERS",
|
||||
"ui.default_group_id": "DEFAULT_GROUP_ID",
|
||||
"ui.default_locale": "DEFAULT_LOCALE",
|
||||
"ui.default_models": "DEFAULT_MODELS",
|
||||
"ui.default_pinned_models": "DEFAULT_PINNED_MODELS",
|
||||
"ui.default_user_role": "DEFAULT_USER_ROLE",
|
||||
"ui.enable_community_sharing": "ENABLE_COMMUNITY_SHARING",
|
||||
"ui.enable_login_form": "ENABLE_LOGIN_FORM",
|
||||
"ui.enable_message_rating": "ENABLE_MESSAGE_RATING",
|
||||
"ui.enable_password_change_form": "ENABLE_PASSWORD_CHANGE_FORM",
|
||||
"ui.enable_signup": "ENABLE_SIGNUP",
|
||||
"ui.enable_user_webhooks": "ENABLE_USER_WEBHOOKS",
|
||||
"ui.model_order_list": "MODEL_ORDER_LIST",
|
||||
"ui.pending_user_overlay_content": "PENDING_USER_OVERLAY_CONTENT",
|
||||
"ui.pending_user_overlay_title": "PENDING_USER_OVERLAY_TITLE",
|
||||
"ui.prompt_suggestions": "DEFAULT_PROMPT_SUGGESTIONS",
|
||||
"ui.watermark": "RESPONSE_WATERMARK",
|
||||
"user.permissions": "USER_PERMISSIONS",
|
||||
"users.enable_status": "ENABLE_USER_STATUS",
|
||||
"webhook_url": "WEBHOOK_URL",
|
||||
"webui.url": "WEBUI_URL",
|
||||
}
|
||||
|
||||
|
||||
STORAGE_KEY_REWRITES = {
|
||||
"oauth.refresh_token_include_scope": "oauth.refresh_token.include_scope",
|
||||
|
||||
"rag.openai_api_base_url": "rag.openai.api_base_url",
|
||||
"rag.openai_api_key": "rag.openai.api_key",
|
||||
"rag.ollama.url": "rag.ollama.base_url",
|
||||
"rag.ollama.key": "rag.ollama.api_key",
|
||||
"oauth.oidc.avatar_claim": "oauth.picture_claim",
|
||||
"oauth.oidc.client_id": "oauth.client_id",
|
||||
"oauth.oidc.client_secret": "oauth.client_secret",
|
||||
"oauth.oidc.code_challenge_method": "oauth.code_challenge_method",
|
||||
"oauth.oidc.email_claim": "oauth.email_claim",
|
||||
"oauth.oidc.end_session_endpoint": "oauth.end_session_endpoint",
|
||||
"oauth.oidc.group_claim": "oauth.group_claim",
|
||||
"oauth.oidc.oauth_timeout": "oauth.timeout",
|
||||
"oauth.oidc.provider_name": "oauth.provider_name",
|
||||
"oauth.oidc.provider_url": "oauth.provider_url",
|
||||
"oauth.oidc.redirect_uri": "oauth.redirect_uri",
|
||||
"oauth.oidc.scopes": "oauth.scopes",
|
||||
"oauth.oidc.sub_claim": "oauth.sub_claim",
|
||||
"oauth.oidc.token_endpoint_auth_method": "oauth.token_endpoint_auth_method",
|
||||
"oauth.oidc.username_claim": "oauth.username_claim",
|
||||
}
|
||||
|
||||
|
||||
LEGACY_KEY_TO_STORAGE_KEY = {
|
||||
legacy_key: STORAGE_KEY_REWRITES.get(blob_path, blob_path)
|
||||
for blob_path, legacy_key in BLOB_PATH_TO_KEY.items()
|
||||
}
|
||||
|
||||
|
||||
def _walk_blob(data: dict, prefix: str = '') -> dict:
|
||||
"""Recursively walk a nested dict, yielding (dot.path, value) for leaf nodes."""
|
||||
result = {}
|
||||
for key, value in data.items():
|
||||
path = f'{prefix}{key}' if not prefix else f'{prefix}.{key}'
|
||||
if isinstance(value, dict):
|
||||
result.update(_walk_blob(value, path))
|
||||
else:
|
||||
result[path] = value
|
||||
return result
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Reshape config from single-row JSON blob to per-key rows."""
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
table_names = set(inspector.get_table_names())
|
||||
config_columns = {column['name'] for column in inspector.get_columns('config')} if 'config' in table_names else set()
|
||||
has_old_config = {'id', 'data'}.issubset(config_columns)
|
||||
has_new_config = {'key', 'value'}.issubset(config_columns)
|
||||
|
||||
# Ad-hoc table reference for reading the old schema
|
||||
old_config = sa.table(
|
||||
'config',
|
||||
sa.column('id', sa.Integer),
|
||||
sa.column('data', sa.JSON),
|
||||
)
|
||||
|
||||
# 1. Read existing blob
|
||||
blob_data = {}
|
||||
if has_old_config:
|
||||
try:
|
||||
result = conn.execute(
|
||||
sa.select(old_config.c.data).order_by(old_config.c.id.desc()).limit(1)
|
||||
)
|
||||
row = result.fetchone()
|
||||
if row and row[0]:
|
||||
raw = row[0]
|
||||
blob_data = json.loads(raw) if isinstance(raw, str) else raw
|
||||
except Exception:
|
||||
pass # Table might be partially migrated or empty
|
||||
|
||||
# 2. Preserve old blob table for rollback/inspection, then create per-key table.
|
||||
if has_old_config:
|
||||
if 'config_old' in table_names:
|
||||
op.drop_table('config_old')
|
||||
op.rename_table('config', 'config_old')
|
||||
|
||||
# 3. Create new per-key table
|
||||
new_config = (
|
||||
sa.table(
|
||||
'config',
|
||||
sa.column('key', sa.Text),
|
||||
sa.column('value', sa.JSON()),
|
||||
sa.column('updated_at', sa.BigInteger),
|
||||
)
|
||||
if has_new_config
|
||||
else op.create_table(
|
||||
'config',
|
||||
sa.Column('key', sa.Text(), primary_key=True),
|
||||
sa.Column('value', sa.JSON(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
)
|
||||
|
||||
# 4. Flatten blob and insert per-key rows
|
||||
if blob_data:
|
||||
flat = _walk_blob(blob_data)
|
||||
|
||||
# Keep stable dot-notation paths as the database keys.
|
||||
# Known legacy env-style keys are rewritten to their dotted keys; unknown
|
||||
# keys are still copied so custom/future config is not silently lost.
|
||||
rows = {}
|
||||
for blob_path, value in flat.items():
|
||||
if blob_path in BLOB_PATH_TO_KEY:
|
||||
storage_key = STORAGE_KEY_REWRITES.get(blob_path, blob_path)
|
||||
elif blob_path in LEGACY_KEY_TO_STORAGE_KEY:
|
||||
storage_key = LEGACY_KEY_TO_STORAGE_KEY[blob_path]
|
||||
else:
|
||||
storage_key = STORAGE_KEY_REWRITES.get(blob_path, blob_path)
|
||||
|
||||
if storage_key not in rows:
|
||||
rows[storage_key] = value
|
||||
|
||||
# Batch insert via SQLAlchemy table reference
|
||||
if rows:
|
||||
now = int(time.time())
|
||||
op.bulk_insert(
|
||||
new_config,
|
||||
[
|
||||
{'key': k, 'value': v, 'updated_at': now}
|
||||
for k, v in rows.items()
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Restore preserved old single-row config table when available."""
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
table_names = set(inspector.get_table_names())
|
||||
|
||||
if 'config_old' in table_names:
|
||||
if 'config' in table_names:
|
||||
op.drop_table('config')
|
||||
op.rename_table('config_old', 'config')
|
||||
return
|
||||
|
||||
config_columns = {column['name'] for column in inspector.get_columns('config')} if 'config' in table_names else set()
|
||||
has_per_key_config = {'key', 'value'}.issubset(config_columns)
|
||||
|
||||
blob_data = {}
|
||||
if has_per_key_config:
|
||||
config = sa.table(
|
||||
'config',
|
||||
sa.column('key', sa.Text),
|
||||
sa.column('value', sa.JSON),
|
||||
)
|
||||
for key, value in conn.execute(sa.select(config.c.key, config.c.value)):
|
||||
blob_data[key] = json.loads(value) if isinstance(value, str) else value
|
||||
op.drop_table('config')
|
||||
|
||||
if 'config' in table_names and not has_per_key_config:
|
||||
return
|
||||
|
||||
old_config = op.create_table(
|
||||
'config',
|
||||
sa.Column('id', sa.Integer(), primary_key=True),
|
||||
sa.Column('data', sa.JSON(), nullable=False),
|
||||
sa.Column('version', sa.Integer(), nullable=False, server_default='0'),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column('updated_at', sa.DateTime(), nullable=True),
|
||||
)
|
||||
|
||||
if blob_data:
|
||||
op.bulk_insert(old_config, [{'data': blob_data, 'version': 0}])
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Database-backed configuration with per-key storage.
|
||||
|
||||
Replaces the old single-row JSON blob machinery with a simple per-key model
|
||||
mirroring cptr's Config.
|
||||
|
||||
Each config key is stored as its own row: key TEXT PK, value JSON.
|
||||
Reads are direct DB lookups. Writes are explicit awaited upserts that raise on
|
||||
failure (no more fire-and-forget create_task).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from open_webui.internal.db import Base, get_async_db
|
||||
from sqlalchemy import JSON, BigInteger, Column, Text, select
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── Model ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class Config(Base):
|
||||
"""Per-key config storage. Each row is one config key."""
|
||||
|
||||
__tablename__ = 'config'
|
||||
|
||||
key = Column(Text, primary_key=True)
|
||||
value = Column(JSON, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=True)
|
||||
|
||||
DEFAULTS: ClassVar[dict[str, Any]] = {}
|
||||
PERSISTENT_ENABLED: ClassVar[bool] = True
|
||||
OAUTH_PERSISTENT_ENABLED: ClassVar[bool] = False
|
||||
|
||||
# ── Class methods ────────────────────────────────────────
|
||||
|
||||
@classmethod
|
||||
def configure(
|
||||
cls,
|
||||
*,
|
||||
defaults: dict[str, Any] | None = None,
|
||||
enable_persistent: bool = True,
|
||||
enable_oauth_persistent: bool = False,
|
||||
) -> None:
|
||||
cls.DEFAULTS = defaults or {}
|
||||
cls.PERSISTENT_ENABLED = enable_persistent
|
||||
cls.OAUTH_PERSISTENT_ENABLED = enable_oauth_persistent
|
||||
|
||||
@classmethod
|
||||
def default_value(cls, key: str, default: Any = None) -> Any:
|
||||
return cls.DEFAULTS.get(key, default)
|
||||
|
||||
@classmethod
|
||||
def persistent_enabled_for(cls, key: str) -> bool:
|
||||
if not cls.PERSISTENT_ENABLED:
|
||||
return False
|
||||
if key.startswith('oauth.') and not cls.OAUTH_PERSISTENT_ENABLED:
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
async def get(key: str, default: Any = None) -> Any:
|
||||
"""Get a config value by key. Returns default if not set."""
|
||||
if not Config.persistent_enabled_for(key):
|
||||
return Config.default_value(key, default)
|
||||
async with get_async_db() as db:
|
||||
row = await db.get(Config, key)
|
||||
return row.value if row else Config.default_value(key, default)
|
||||
|
||||
@staticmethod
|
||||
async def get_many(*keys: str) -> dict:
|
||||
"""Get multiple config values. Returns {key: value} for keys that exist."""
|
||||
disabled_values = {
|
||||
key: Config.default_value(key)
|
||||
for key in keys
|
||||
if not Config.persistent_enabled_for(key) and key in Config.DEFAULTS
|
||||
}
|
||||
enabled_keys = {key for key in keys if Config.persistent_enabled_for(key)}
|
||||
if not enabled_keys:
|
||||
return disabled_values
|
||||
async with get_async_db() as db:
|
||||
result = await db.execute(select(Config).where(Config.key.in_(enabled_keys)))
|
||||
values = {row.key: row.value for row in result.scalars().all()}
|
||||
return {
|
||||
key: values.get(key, Config.default_value(key))
|
||||
for key in keys
|
||||
if key in values or key in Config.DEFAULTS or key in disabled_values
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
async def get_namespace(namespace: str) -> dict:
|
||||
"""Get all config keys under a dotted namespace."""
|
||||
default_values = {
|
||||
key: value
|
||||
for key, value in Config.DEFAULTS.items()
|
||||
if key.startswith(f'{namespace}.') and not Config.persistent_enabled_for(key)
|
||||
}
|
||||
if not Config.PERSISTENT_ENABLED:
|
||||
return default_values
|
||||
async with get_async_db() as db:
|
||||
result = await db.execute(select(Config).where(Config.key.like(f'{namespace}.%')))
|
||||
values = {row.key: row.value for row in result.scalars().all()}
|
||||
values.update(default_values)
|
||||
return values
|
||||
|
||||
@staticmethod
|
||||
async def get_all() -> dict:
|
||||
"""Get all config as {key: value}."""
|
||||
if not Config.PERSISTENT_ENABLED:
|
||||
return dict(Config.DEFAULTS)
|
||||
async with get_async_db() as db:
|
||||
result = await db.execute(select(Config))
|
||||
values = {row.key: row.value for row in result.scalars().all()}
|
||||
if not Config.OAUTH_PERSISTENT_ENABLED:
|
||||
values.update({key: value for key, value in Config.DEFAULTS.items() if key.startswith('oauth.')})
|
||||
return values
|
||||
|
||||
@staticmethod
|
||||
async def upsert(updates: dict) -> None:
|
||||
"""Upsert multiple config key-value pairs. Raises on failure."""
|
||||
async with get_async_db() as db:
|
||||
now = int(time.time())
|
||||
for key, value in updates.items():
|
||||
existing = await db.get(Config, key)
|
||||
if existing:
|
||||
existing.value = value
|
||||
existing.updated_at = now
|
||||
else:
|
||||
db.add(Config(key=key, value=value, updated_at=now))
|
||||
await db.commit()
|
||||
|
||||
@staticmethod
|
||||
async def delete(key: str) -> bool:
|
||||
"""Delete a config key. Returns True if it existed."""
|
||||
async with get_async_db() as db:
|
||||
row = await db.get(Config, key)
|
||||
if row:
|
||||
await db.delete(row)
|
||||
await db.commit()
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def clear() -> None:
|
||||
"""Delete all config rows."""
|
||||
from sqlalchemy import delete as sa_delete
|
||||
|
||||
async with get_async_db() as db:
|
||||
await db.execute(sa_delete(Config))
|
||||
await db.commit()
|
||||
|
||||
@staticmethod
|
||||
async def seed_defaults(defaults: dict) -> None:
|
||||
"""Insert keys that don't yet exist in the DB.
|
||||
|
||||
Called at startup to ensure all known config keys have values.
|
||||
Existing DB values take precedence over defaults.
|
||||
"""
|
||||
async with get_async_db() as db:
|
||||
result = await db.execute(select(Config.key))
|
||||
existing_keys = {row[0] for row in result.all()}
|
||||
|
||||
now = int(time.time())
|
||||
new_count = 0
|
||||
for key, value in defaults.items():
|
||||
if key not in existing_keys:
|
||||
db.add(Config(key=key, value=value, updated_at=now))
|
||||
existing_keys.add(key)
|
||||
new_count += 1
|
||||
|
||||
if new_count:
|
||||
await db.commit()
|
||||
log.info('Seeded %d new config defaults', new_count)
|
||||
@@ -39,6 +39,7 @@ from open_webui.models.chats import Chats
|
||||
from open_webui.models.files import Files
|
||||
from open_webui.models.knowledge import Knowledges
|
||||
from open_webui.models.notes import Notes
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.users import UserModel
|
||||
from open_webui.retrieval.loaders.youtube import YoutubeLoader
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
@@ -63,65 +64,84 @@ def is_youtube_url(url: str) -> bool:
|
||||
return re.match(youtube_regex, url) is not None
|
||||
|
||||
|
||||
def get_loader(request, url: str):
|
||||
LOADER_CONFIG_KEYS = {
|
||||
'youtube_language': 'rag.youtube_loader_language',
|
||||
'youtube_proxy_url': 'rag.youtube_loader_proxy_url',
|
||||
'web_loader_ssl_verification': 'rag.web.loader.ssl_verification',
|
||||
'web_loader_concurrent_requests': 'rag.web.loader.concurrent_requests',
|
||||
'web_search_trust_env': 'rag.web.search.trust_env',
|
||||
'CONTENT_EXTRACTION_ENGINE': 'rag.content_extraction_engine',
|
||||
'DATALAB_MARKER_API_KEY': 'rag.datalab_marker_api_key',
|
||||
'DATALAB_MARKER_API_BASE_URL': 'rag.datalab_marker_api_base_url',
|
||||
'DATALAB_MARKER_ADDITIONAL_CONFIG': 'rag.datalab_marker_additional_config',
|
||||
'DATALAB_MARKER_SKIP_CACHE': 'rag.datalab_marker_skip_cache',
|
||||
'DATALAB_MARKER_FORCE_OCR': 'rag.datalab_marker_force_ocr',
|
||||
'DATALAB_MARKER_PAGINATE': 'rag.datalab_marker_paginate',
|
||||
'DATALAB_MARKER_STRIP_EXISTING_OCR': 'rag.datalab_marker_strip_existing_ocr',
|
||||
'DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION': 'rag.datalab_marker_disable_image_extraction',
|
||||
'DATALAB_MARKER_FORMAT_LINES': 'rag.datalab_marker_format_lines',
|
||||
'DATALAB_MARKER_USE_LLM': 'rag.datalab_marker_use_llm',
|
||||
'DATALAB_MARKER_OUTPUT_FORMAT': 'rag.datalab_marker_output_format',
|
||||
'EXTERNAL_DOCUMENT_LOADER_URL': 'rag.external_document_loader_url',
|
||||
'EXTERNAL_DOCUMENT_LOADER_API_KEY': 'rag.external_document_loader_api_key',
|
||||
'TIKA_SERVER_URL': 'rag.tika_server_url',
|
||||
'DOCLING_SERVER_URL': 'rag.docling_server_url',
|
||||
'DOCLING_API_KEY': 'rag.docling_api_key',
|
||||
'DOCLING_PARAMS': 'rag.docling_params',
|
||||
'PDF_EXTRACT_IMAGES': 'rag.pdf_extract_images',
|
||||
'PDF_LOADER_MODE': 'rag.pdf_loader_mode',
|
||||
'DOCUMENT_INTELLIGENCE_ENDPOINT': 'rag.document_intelligence_endpoint',
|
||||
'DOCUMENT_INTELLIGENCE_KEY': 'rag.document_intelligence_key',
|
||||
'DOCUMENT_INTELLIGENCE_MODEL': 'rag.document_intelligence_model',
|
||||
'MISTRAL_OCR_API_BASE_URL': 'rag.mistral_ocr_api_base_url',
|
||||
'MISTRAL_OCR_API_KEY': 'rag.mistral_ocr_api_key',
|
||||
'PADDLEOCR_VL_BASE_URL': 'rag.paddleocr_vl_base_url',
|
||||
'PADDLEOCR_VL_TOKEN': 'rag.paddleocr_vl_token',
|
||||
'MINERU_API_MODE': 'rag.mineru_api_mode',
|
||||
'MINERU_API_URL': 'rag.mineru_api_url',
|
||||
'MINERU_API_KEY': 'rag.mineru_api_key',
|
||||
'MINERU_API_TIMEOUT': 'rag.mineru_api_timeout',
|
||||
'MINERU_PARAMS': 'rag.mineru_params',
|
||||
'MINERU_FILE_EXTENSIONS': 'rag.mineru_file_extensions',
|
||||
}
|
||||
|
||||
|
||||
async def get_loader_config():
|
||||
values = await Config.get_many(*LOADER_CONFIG_KEYS.values())
|
||||
return {name: values.get(key) for name, key in LOADER_CONFIG_KEYS.items()}
|
||||
|
||||
|
||||
def get_loader(request, url: str, config: dict):
|
||||
if is_youtube_url(url):
|
||||
return YoutubeLoader(
|
||||
url,
|
||||
language=request.app.state.config.YOUTUBE_LOADER_LANGUAGE,
|
||||
proxy_url=request.app.state.config.YOUTUBE_LOADER_PROXY_URL,
|
||||
language=config.get('youtube_language'),
|
||||
proxy_url=config.get('youtube_proxy_url'),
|
||||
)
|
||||
else:
|
||||
return get_web_loader(
|
||||
url,
|
||||
verify_ssl=request.app.state.config.ENABLE_WEB_LOADER_SSL_VERIFICATION,
|
||||
requests_per_second=request.app.state.config.WEB_LOADER_CONCURRENT_REQUESTS,
|
||||
trust_env=request.app.state.config.WEB_SEARCH_TRUST_ENV,
|
||||
)
|
||||
|
||||
|
||||
def build_loader_from_config(request):
|
||||
"""Build a Loader instance with the admin's configured extraction engine settings."""
|
||||
from open_webui.retrieval.loaders.main import Loader
|
||||
|
||||
config = request.app.state.config
|
||||
return Loader(
|
||||
engine=config.CONTENT_EXTRACTION_ENGINE,
|
||||
DATALAB_MARKER_API_KEY=config.DATALAB_MARKER_API_KEY,
|
||||
DATALAB_MARKER_API_BASE_URL=config.DATALAB_MARKER_API_BASE_URL,
|
||||
DATALAB_MARKER_ADDITIONAL_CONFIG=config.DATALAB_MARKER_ADDITIONAL_CONFIG,
|
||||
DATALAB_MARKER_SKIP_CACHE=config.DATALAB_MARKER_SKIP_CACHE,
|
||||
DATALAB_MARKER_FORCE_OCR=config.DATALAB_MARKER_FORCE_OCR,
|
||||
DATALAB_MARKER_PAGINATE=config.DATALAB_MARKER_PAGINATE,
|
||||
DATALAB_MARKER_STRIP_EXISTING_OCR=config.DATALAB_MARKER_STRIP_EXISTING_OCR,
|
||||
DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION=config.DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION,
|
||||
DATALAB_MARKER_FORMAT_LINES=config.DATALAB_MARKER_FORMAT_LINES,
|
||||
DATALAB_MARKER_USE_LLM=config.DATALAB_MARKER_USE_LLM,
|
||||
DATALAB_MARKER_OUTPUT_FORMAT=config.DATALAB_MARKER_OUTPUT_FORMAT,
|
||||
EXTERNAL_DOCUMENT_LOADER_URL=config.EXTERNAL_DOCUMENT_LOADER_URL,
|
||||
EXTERNAL_DOCUMENT_LOADER_API_KEY=config.EXTERNAL_DOCUMENT_LOADER_API_KEY,
|
||||
TIKA_SERVER_URL=config.TIKA_SERVER_URL,
|
||||
DOCLING_SERVER_URL=config.DOCLING_SERVER_URL,
|
||||
DOCLING_API_KEY=config.DOCLING_API_KEY,
|
||||
DOCLING_PARAMS=config.DOCLING_PARAMS,
|
||||
PDF_EXTRACT_IMAGES=config.PDF_EXTRACT_IMAGES,
|
||||
PDF_LOADER_MODE=config.PDF_LOADER_MODE,
|
||||
DOCUMENT_INTELLIGENCE_ENDPOINT=config.DOCUMENT_INTELLIGENCE_ENDPOINT,
|
||||
DOCUMENT_INTELLIGENCE_KEY=config.DOCUMENT_INTELLIGENCE_KEY,
|
||||
DOCUMENT_INTELLIGENCE_MODEL=config.DOCUMENT_INTELLIGENCE_MODEL,
|
||||
MISTRAL_OCR_API_BASE_URL=config.MISTRAL_OCR_API_BASE_URL,
|
||||
MISTRAL_OCR_API_KEY=config.MISTRAL_OCR_API_KEY,
|
||||
PADDLEOCR_VL_BASE_URL=config.PADDLEOCR_VL_BASE_URL,
|
||||
PADDLEOCR_VL_TOKEN=config.PADDLEOCR_VL_TOKEN,
|
||||
MINERU_API_MODE=config.MINERU_API_MODE,
|
||||
MINERU_API_URL=config.MINERU_API_URL,
|
||||
MINERU_API_KEY=config.MINERU_API_KEY,
|
||||
MINERU_API_TIMEOUT=config.MINERU_API_TIMEOUT,
|
||||
MINERU_PARAMS=config.MINERU_PARAMS,
|
||||
MINERU_FILE_EXTENSIONS=config.MINERU_FILE_EXTENSIONS,
|
||||
return get_web_loader(
|
||||
url,
|
||||
verify_ssl=config.get('web_loader_ssl_verification'),
|
||||
requests_per_second=config.get('web_loader_concurrent_requests'),
|
||||
trust_env=config.get('web_search_trust_env'),
|
||||
)
|
||||
|
||||
|
||||
def _extract_text_from_binary_response(request, response: requests.Response, url: str) -> tuple[str, list]:
|
||||
def build_loader_from_config(request, config: dict):
|
||||
"""Build a Loader instance with the admin's configured extraction engine settings."""
|
||||
from open_webui.retrieval.loaders.main import Loader
|
||||
|
||||
loader_config = {
|
||||
key: config.get(key)
|
||||
for key in LOADER_CONFIG_KEYS
|
||||
if key.isupper()
|
||||
}
|
||||
return Loader(
|
||||
engine=loader_config['CONTENT_EXTRACTION_ENGINE'],
|
||||
**{key: value for key, value in loader_config.items() if key != 'CONTENT_EXTRACTION_ENGINE'},
|
||||
)
|
||||
|
||||
|
||||
def _extract_text_from_binary_response(request, response: requests.Response, url: str, loader_config: dict) -> tuple[str, list]:
|
||||
"""Download response body to a temp file and extract text using the Loader pipeline."""
|
||||
import mimetypes
|
||||
import tempfile
|
||||
@@ -150,7 +170,7 @@ def _extract_text_from_binary_response(request, response: requests.Response, url
|
||||
tmp_path = tmp.name
|
||||
|
||||
try:
|
||||
loader = build_loader_from_config(request)
|
||||
loader = build_loader_from_config(request, loader_config)
|
||||
docs = loader.load(filename, content_type, tmp_path)
|
||||
for doc in docs:
|
||||
doc.metadata['source'] = url
|
||||
@@ -170,9 +190,11 @@ def _is_text_content_type(content_type: str) -> bool:
|
||||
return not ct # empty / missing → assume HTML
|
||||
|
||||
|
||||
def get_content_from_url(request, url: str) -> str:
|
||||
async def get_content_from_url(request, url: str) -> str:
|
||||
from open_webui.retrieval.web.utils import validate_url
|
||||
|
||||
loader_config = await get_loader_config()
|
||||
|
||||
# Validate URL before making any request (blocks private IPs, non-HTTP, filter list)
|
||||
validate_url(url)
|
||||
|
||||
@@ -183,7 +205,7 @@ def get_content_from_url(request, url: str) -> str:
|
||||
# when allow_redirects=False, causing the binary-content path to run
|
||||
# and produce empty docs → HTTP 400.
|
||||
if is_youtube_url(url):
|
||||
loader = get_loader(request, url)
|
||||
loader = get_loader(request, url, loader_config)
|
||||
docs = loader.load()
|
||||
content = ' '.join([doc.page_content for doc in docs])
|
||||
return content, docs
|
||||
@@ -205,14 +227,14 @@ def get_content_from_url(request, url: str) -> str:
|
||||
if response is None or _is_text_content_type(content_type):
|
||||
if response is not None:
|
||||
response.close()
|
||||
loader = get_loader(request, url)
|
||||
loader = get_loader(request, url, loader_config)
|
||||
docs = loader.load()
|
||||
content = ' '.join([doc.page_content for doc in docs])
|
||||
return content, docs
|
||||
|
||||
# Binary content (PDF, DOCX, XLSX, PPTX, etc.) — download and extract
|
||||
try:
|
||||
return _extract_text_from_binary_response(request, response, url)
|
||||
return _extract_text_from_binary_response(request, response, url, loader_config)
|
||||
finally:
|
||||
response.close()
|
||||
|
||||
@@ -539,8 +561,15 @@ async def query_collection(
|
||||
embedding_function,
|
||||
k: int,
|
||||
) -> dict:
|
||||
config = await Config.get_many(
|
||||
'rag.enable_hybrid_search',
|
||||
'rag.top_k_reranker',
|
||||
'rag.relevance_threshold',
|
||||
'rag.hybrid_bm25_weight',
|
||||
'rag.enable_hybrid_search_enriched_texts',
|
||||
)
|
||||
# When request is provided, try hybrid search + reranking if enabled
|
||||
if request and request.app.state.config.ENABLE_RAG_HYBRID_SEARCH:
|
||||
if request and config.get('rag.enable_hybrid_search'):
|
||||
try:
|
||||
reranking_function = (
|
||||
(lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents))
|
||||
@@ -553,10 +582,10 @@ async def query_collection(
|
||||
embedding_function=embedding_function,
|
||||
k=k,
|
||||
reranking_function=reranking_function,
|
||||
k_reranker=request.app.state.config.TOP_K_RERANKER,
|
||||
r=request.app.state.config.RELEVANCE_THRESHOLD,
|
||||
hybrid_bm25_weight=request.app.state.config.HYBRID_BM25_WEIGHT,
|
||||
enable_enriched_texts=request.app.state.config.ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS,
|
||||
k_reranker=config.get('rag.top_k_reranker'),
|
||||
r=config.get('rag.relevance_threshold'),
|
||||
hybrid_bm25_weight=config.get('rag.hybrid_bm25_weight'),
|
||||
enable_enriched_texts=config.get('rag.enable_hybrid_search_enriched_texts'),
|
||||
)
|
||||
except Exception as e:
|
||||
log.debug(f'Hybrid search failed, falling back to vector search: {e}')
|
||||
@@ -1165,6 +1194,7 @@ async def get_sources_from_items(
|
||||
):
|
||||
log.debug(f'items: {items} {queries} {embedding_function} {reranking_function} {full_context}')
|
||||
|
||||
bypass_embedding_and_retrieval = await Config.get('rag.bypass_embedding_and_retrieval')
|
||||
extracted_collections = []
|
||||
query_results = []
|
||||
|
||||
@@ -1244,14 +1274,14 @@ async def get_sources_from_items(
|
||||
}
|
||||
|
||||
elif item.get('type') == 'url':
|
||||
content, docs = get_content_from_url(request, item.get('url'))
|
||||
content, docs = await get_content_from_url(request, item.get('url'))
|
||||
if docs:
|
||||
query_result = {
|
||||
'documents': [[content]],
|
||||
'metadatas': [[{'url': item.get('url'), 'name': item.get('url')}]],
|
||||
}
|
||||
elif item.get('type') == 'file':
|
||||
if item.get('context') == 'full' or request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL:
|
||||
if item.get('context') == 'full' or bypass_embedding_and_retrieval:
|
||||
if item.get('file', {}).get('data', {}).get('content', ''):
|
||||
# Manual Full Mode Toggle
|
||||
# Used from chat file modal, we can assume that the file content will be available from item.get("file").get("data", {}).get("content")
|
||||
@@ -1323,7 +1353,7 @@ async def get_sources_from_items(
|
||||
permission='read',
|
||||
)
|
||||
):
|
||||
if item.get('context') == 'full' or request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL:
|
||||
if item.get('context') == 'full' or bypass_embedding_and_retrieval:
|
||||
if knowledge_base and (
|
||||
user.role == 'admin'
|
||||
or knowledge_base.user_id == user.id
|
||||
|
||||
@@ -38,9 +38,7 @@ def search_perplexity(
|
||||
|
||||
"""
|
||||
|
||||
# Handle ConfigVar object
|
||||
if hasattr(api_key, '__str__'):
|
||||
api_key = str(api_key)
|
||||
api_key = str(api_key)
|
||||
|
||||
try:
|
||||
url = 'https://api.perplexity.ai/chat/completions'
|
||||
|
||||
@@ -29,12 +29,8 @@ def search_perplexity_search(
|
||||
|
||||
"""
|
||||
|
||||
# Handle ConfigVar object
|
||||
if hasattr(api_key, '__str__'):
|
||||
api_key = str(api_key)
|
||||
|
||||
if hasattr(api_url, '__str__'):
|
||||
api_url = str(api_url)
|
||||
api_key = str(api_key)
|
||||
api_url = str(api_url)
|
||||
|
||||
try:
|
||||
url = api_url
|
||||
|
||||
@@ -21,7 +21,6 @@ from typing import (
|
||||
import aiohttp
|
||||
import aiohttp.resolver
|
||||
import certifi
|
||||
import requests
|
||||
import urllib3.connection
|
||||
import urllib3.connectionpool
|
||||
import validators
|
||||
@@ -777,13 +776,13 @@ def get_web_loader(
|
||||
'trust_env': trust_env,
|
||||
}
|
||||
|
||||
if WEB_LOADER_ENGINE.value == '' or WEB_LOADER_ENGINE.value == 'safe_web':
|
||||
if WEB_LOADER_ENGINE == '' or WEB_LOADER_ENGINE == 'safe_web':
|
||||
WebLoaderClass = SafeWebBaseLoader
|
||||
|
||||
request_kwargs = {}
|
||||
if WEB_LOADER_TIMEOUT.value:
|
||||
if WEB_LOADER_TIMEOUT:
|
||||
try:
|
||||
timeout_value = float(WEB_LOADER_TIMEOUT.value)
|
||||
timeout_value = float(WEB_LOADER_TIMEOUT)
|
||||
except ValueError:
|
||||
timeout_value = None
|
||||
|
||||
@@ -793,31 +792,31 @@ def get_web_loader(
|
||||
if request_kwargs:
|
||||
web_loader_args['requests_kwargs'] = request_kwargs
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'playwright':
|
||||
if WEB_LOADER_ENGINE == 'playwright':
|
||||
WebLoaderClass = SafePlaywrightURLLoader
|
||||
web_loader_args['playwright_timeout'] = PLAYWRIGHT_TIMEOUT.value
|
||||
if PLAYWRIGHT_WS_URL.value:
|
||||
web_loader_args['playwright_ws_url'] = PLAYWRIGHT_WS_URL.value
|
||||
web_loader_args['playwright_timeout'] = PLAYWRIGHT_TIMEOUT
|
||||
if PLAYWRIGHT_WS_URL:
|
||||
web_loader_args['playwright_ws_url'] = PLAYWRIGHT_WS_URL
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'firecrawl':
|
||||
if WEB_LOADER_ENGINE == 'firecrawl':
|
||||
WebLoaderClass = SafeFireCrawlLoader
|
||||
web_loader_args['api_key'] = FIRECRAWL_API_KEY.value
|
||||
web_loader_args['api_url'] = FIRECRAWL_API_BASE_URL.value
|
||||
if FIRECRAWL_TIMEOUT.value:
|
||||
web_loader_args['api_key'] = FIRECRAWL_API_KEY
|
||||
web_loader_args['api_url'] = FIRECRAWL_API_BASE_URL
|
||||
if FIRECRAWL_TIMEOUT:
|
||||
try:
|
||||
web_loader_args['timeout'] = int(FIRECRAWL_TIMEOUT.value)
|
||||
web_loader_args['timeout'] = int(FIRECRAWL_TIMEOUT)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'tavily':
|
||||
if WEB_LOADER_ENGINE == 'tavily':
|
||||
WebLoaderClass = SafeTavilyLoader
|
||||
web_loader_args['api_key'] = TAVILY_API_KEY.value
|
||||
web_loader_args['extract_depth'] = TAVILY_EXTRACT_DEPTH.value
|
||||
web_loader_args['api_key'] = TAVILY_API_KEY
|
||||
web_loader_args['extract_depth'] = TAVILY_EXTRACT_DEPTH
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'external':
|
||||
if WEB_LOADER_ENGINE == 'external':
|
||||
WebLoaderClass = ExternalWebLoader
|
||||
web_loader_args['external_url'] = EXTERNAL_WEB_LOADER_URL.value
|
||||
web_loader_args['external_api_key'] = EXTERNAL_WEB_LOADER_API_KEY.value
|
||||
web_loader_args['external_url'] = EXTERNAL_WEB_LOADER_URL
|
||||
web_loader_args['external_api_key'] = EXTERNAL_WEB_LOADER_API_KEY
|
||||
|
||||
if WebLoaderClass:
|
||||
web_loader = WebLoaderClass(**web_loader_args)
|
||||
@@ -831,6 +830,6 @@ def get_web_loader(
|
||||
return web_loader
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Invalid WEB_LOADER_ENGINE: {WEB_LOADER_ENGINE.value}. '
|
||||
f'Invalid WEB_LOADER_ENGINE: {WEB_LOADER_ENGINE}. '
|
||||
"Please set it to 'safe_web', 'playwright', 'firecrawl', or 'tavily'."
|
||||
)
|
||||
|
||||
+112
-154
@@ -52,6 +52,7 @@ from open_webui.env import (
|
||||
ENABLE_FORWARD_USER_INFO_HEADERS,
|
||||
ENV,
|
||||
)
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
@@ -71,6 +72,50 @@ AZURE_MAX_FILE_SIZE: int = AZURE_MAX_FILE_SIZE_MB * 1024 * 1024
|
||||
SPEECH_CACHE_DIR = CACHE_DIR / 'audio' / 'speech'
|
||||
SPEECH_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
TTS_CONFIG_KEYS = {
|
||||
'OPENAI_API_BASE_URL': 'audio.tts.openai.api_base_url',
|
||||
'OPENAI_API_KEY': 'audio.tts.openai.api_key',
|
||||
'OPENAI_PARAMS': 'audio.tts.openai.params',
|
||||
'API_KEY': 'audio.tts.api_key',
|
||||
'ENGINE': 'audio.tts.engine',
|
||||
'MODEL': 'audio.tts.model',
|
||||
'VOICE': 'audio.tts.voice',
|
||||
'SPLIT_ON': 'audio.tts.split_on',
|
||||
'AZURE_SPEECH_REGION': 'audio.tts.azure.speech_region',
|
||||
'AZURE_SPEECH_BASE_URL': 'audio.tts.azure.speech_base_url',
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': 'audio.tts.azure.speech_output_format',
|
||||
'MISTRAL_API_KEY': 'audio.tts.mistral.api_key',
|
||||
'MISTRAL_API_BASE_URL': 'audio.tts.mistral.api_base_url',
|
||||
}
|
||||
|
||||
STT_CONFIG_KEYS = {
|
||||
'OPENAI_API_BASE_URL': 'audio.stt.openai.api_base_url',
|
||||
'OPENAI_API_KEY': 'audio.stt.openai.api_key',
|
||||
'ENGINE': 'audio.stt.engine',
|
||||
'MODEL': 'audio.stt.model',
|
||||
'SUPPORTED_CONTENT_TYPES': 'audio.stt.supported_content_types',
|
||||
'ALLOWED_EXTENSIONS': 'audio.stt.allowed_extensions',
|
||||
'WHISPER_MODEL': 'audio.stt.whisper_model',
|
||||
'DEEPGRAM_API_KEY': 'audio.stt.deepgram.api_key',
|
||||
'AZURE_API_KEY': 'audio.stt.azure.api_key',
|
||||
'AZURE_REGION': 'audio.stt.azure.region',
|
||||
'AZURE_LOCALES': 'audio.stt.azure.locales',
|
||||
'AZURE_BASE_URL': 'audio.stt.azure.base_url',
|
||||
'AZURE_MAX_SPEAKERS': 'audio.stt.azure.max_speakers',
|
||||
'MISTRAL_API_KEY': 'audio.stt.mistral.api_key',
|
||||
'MISTRAL_API_BASE_URL': 'audio.stt.mistral.api_base_url',
|
||||
'MISTRAL_USE_CHAT_COMPLETIONS': 'audio.stt.mistral.use_chat_completions',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
def is_audio_conversion_required(file_path):
|
||||
"""
|
||||
@@ -228,119 +273,28 @@ class AudioConfigUpdateForm(BaseModel):
|
||||
@router.get('/config')
|
||||
async def get_audio_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'tts': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.TTS_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.TTS_OPENAI_API_KEY,
|
||||
'OPENAI_PARAMS': request.app.state.config.TTS_OPENAI_PARAMS,
|
||||
'API_KEY': request.app.state.config.TTS_API_KEY,
|
||||
'ENGINE': request.app.state.config.TTS_ENGINE,
|
||||
'MODEL': request.app.state.config.TTS_MODEL,
|
||||
'VOICE': request.app.state.config.TTS_VOICE,
|
||||
'SPLIT_ON': request.app.state.config.TTS_SPLIT_ON,
|
||||
'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION,
|
||||
'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL,
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL,
|
||||
},
|
||||
'stt': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.STT_OPENAI_API_KEY,
|
||||
'ENGINE': request.app.state.config.STT_ENGINE,
|
||||
'MODEL': request.app.state.config.STT_MODEL,
|
||||
'SUPPORTED_CONTENT_TYPES': request.app.state.config.STT_SUPPORTED_CONTENT_TYPES,
|
||||
'ALLOWED_EXTENSIONS': request.app.state.config.STT_ALLOWED_EXTENSIONS,
|
||||
'WHISPER_MODEL': request.app.state.config.WHISPER_MODEL,
|
||||
'DEEPGRAM_API_KEY': request.app.state.config.DEEPGRAM_API_KEY,
|
||||
'AZURE_API_KEY': request.app.state.config.AUDIO_STT_AZURE_API_KEY,
|
||||
'AZURE_REGION': request.app.state.config.AUDIO_STT_AZURE_REGION,
|
||||
'AZURE_LOCALES': request.app.state.config.AUDIO_STT_AZURE_LOCALES,
|
||||
'AZURE_BASE_URL': request.app.state.config.AUDIO_STT_AZURE_BASE_URL,
|
||||
'AZURE_MAX_SPEAKERS': request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS,
|
||||
'MISTRAL_API_KEY': request.app.state.config.AUDIO_STT_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL,
|
||||
'MISTRAL_USE_CHAT_COMPLETIONS': request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS,
|
||||
},
|
||||
'tts': await get_config_values(TTS_CONFIG_KEYS),
|
||||
'stt': await get_config_values(STT_CONFIG_KEYS),
|
||||
}
|
||||
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm, user=Depends(get_admin_user)):
|
||||
# TTS settings
|
||||
request.app.state.config.TTS_OPENAI_API_BASE_URL = form_data.tts.OPENAI_API_BASE_URL
|
||||
request.app.state.config.TTS_OPENAI_API_KEY = form_data.tts.OPENAI_API_KEY
|
||||
request.app.state.config.TTS_OPENAI_PARAMS = form_data.tts.OPENAI_PARAMS
|
||||
request.app.state.config.TTS_API_KEY = form_data.tts.API_KEY
|
||||
request.app.state.config.TTS_ENGINE = form_data.tts.ENGINE
|
||||
request.app.state.config.TTS_MODEL = form_data.tts.MODEL
|
||||
request.app.state.config.TTS_VOICE = form_data.tts.VOICE
|
||||
request.app.state.config.TTS_SPLIT_ON = form_data.tts.SPLIT_ON
|
||||
request.app.state.config.TTS_AZURE_SPEECH_REGION = form_data.tts.AZURE_SPEECH_REGION
|
||||
request.app.state.config.TTS_AZURE_SPEECH_BASE_URL = form_data.tts.AZURE_SPEECH_BASE_URL
|
||||
request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = form_data.tts.AZURE_SPEECH_OUTPUT_FORMAT
|
||||
request.app.state.config.TTS_MISTRAL_API_KEY = form_data.tts.MISTRAL_API_KEY
|
||||
request.app.state.config.TTS_MISTRAL_API_BASE_URL = form_data.tts.MISTRAL_API_BASE_URL
|
||||
await Config.upsert(
|
||||
{
|
||||
**config_updates(form_data.tts.model_dump(), TTS_CONFIG_KEYS),
|
||||
**config_updates(form_data.stt.model_dump(), STT_CONFIG_KEYS),
|
||||
}
|
||||
)
|
||||
|
||||
# STT settings
|
||||
request.app.state.config.STT_OPENAI_API_BASE_URL = form_data.stt.OPENAI_API_BASE_URL
|
||||
request.app.state.config.STT_OPENAI_API_KEY = form_data.stt.OPENAI_API_KEY
|
||||
request.app.state.config.STT_ENGINE = form_data.stt.ENGINE
|
||||
request.app.state.config.STT_MODEL = form_data.stt.MODEL
|
||||
request.app.state.config.STT_SUPPORTED_CONTENT_TYPES = form_data.stt.SUPPORTED_CONTENT_TYPES
|
||||
request.app.state.config.STT_ALLOWED_EXTENSIONS = form_data.stt.ALLOWED_EXTENSIONS
|
||||
request.app.state.config.WHISPER_MODEL = form_data.stt.WHISPER_MODEL
|
||||
request.app.state.config.DEEPGRAM_API_KEY = form_data.stt.DEEPGRAM_API_KEY
|
||||
request.app.state.config.AUDIO_STT_AZURE_API_KEY = form_data.stt.AZURE_API_KEY
|
||||
request.app.state.config.AUDIO_STT_AZURE_REGION = form_data.stt.AZURE_REGION
|
||||
request.app.state.config.AUDIO_STT_AZURE_LOCALES = form_data.stt.AZURE_LOCALES
|
||||
request.app.state.config.AUDIO_STT_AZURE_BASE_URL = form_data.stt.AZURE_BASE_URL
|
||||
request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS = form_data.stt.AZURE_MAX_SPEAKERS
|
||||
request.app.state.config.AUDIO_STT_MISTRAL_API_KEY = form_data.stt.MISTRAL_API_KEY
|
||||
request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL = form_data.stt.MISTRAL_API_BASE_URL
|
||||
request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS = form_data.stt.MISTRAL_USE_CHAT_COMPLETIONS
|
||||
|
||||
if request.app.state.config.STT_ENGINE == '':
|
||||
if form_data.stt.ENGINE == '':
|
||||
request.app.state.faster_whisper_model = set_faster_whisper_model(
|
||||
form_data.stt.WHISPER_MODEL, WHISPER_MODEL_AUTO_UPDATE
|
||||
)
|
||||
else:
|
||||
request.app.state.faster_whisper_model = None
|
||||
|
||||
return {
|
||||
'tts': {
|
||||
'ENGINE': request.app.state.config.TTS_ENGINE,
|
||||
'MODEL': request.app.state.config.TTS_MODEL,
|
||||
'VOICE': request.app.state.config.TTS_VOICE,
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.TTS_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.TTS_OPENAI_API_KEY,
|
||||
'OPENAI_PARAMS': request.app.state.config.TTS_OPENAI_PARAMS,
|
||||
'API_KEY': request.app.state.config.TTS_API_KEY,
|
||||
'SPLIT_ON': request.app.state.config.TTS_SPLIT_ON,
|
||||
'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION,
|
||||
'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL,
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL,
|
||||
},
|
||||
'stt': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.STT_OPENAI_API_KEY,
|
||||
'ENGINE': request.app.state.config.STT_ENGINE,
|
||||
'MODEL': request.app.state.config.STT_MODEL,
|
||||
'SUPPORTED_CONTENT_TYPES': request.app.state.config.STT_SUPPORTED_CONTENT_TYPES,
|
||||
'ALLOWED_EXTENSIONS': request.app.state.config.STT_ALLOWED_EXTENSIONS,
|
||||
'WHISPER_MODEL': request.app.state.config.WHISPER_MODEL,
|
||||
'DEEPGRAM_API_KEY': request.app.state.config.DEEPGRAM_API_KEY,
|
||||
'AZURE_API_KEY': request.app.state.config.AUDIO_STT_AZURE_API_KEY,
|
||||
'AZURE_REGION': request.app.state.config.AUDIO_STT_AZURE_REGION,
|
||||
'AZURE_LOCALES': request.app.state.config.AUDIO_STT_AZURE_LOCALES,
|
||||
'AZURE_BASE_URL': request.app.state.config.AUDIO_STT_AZURE_BASE_URL,
|
||||
'AZURE_MAX_SPEAKERS': request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS,
|
||||
'MISTRAL_API_KEY': request.app.state.config.AUDIO_STT_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL,
|
||||
'MISTRAL_USE_CHAT_COMPLETIONS': request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS,
|
||||
},
|
||||
}
|
||||
return await get_audio_config(request, user)
|
||||
|
||||
|
||||
def load_speech_pipeline(request):
|
||||
@@ -388,14 +342,16 @@ async def _write_tts_cache(
|
||||
|
||||
async def _tts_openai(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via an OpenAI-compatible TTS endpoint."""
|
||||
payload['model'] = request.app.state.config.TTS_MODEL
|
||||
payload['model'] = await Config.get('audio.tts.model')
|
||||
if not payload.get('voice'):
|
||||
payload['voice'] = request.app.state.config.TTS_VOICE
|
||||
payload = {**payload, **(request.app.state.config.TTS_OPENAI_PARAMS or {})}
|
||||
payload['voice'] = await Config.get('audio.tts.voice')
|
||||
payload = {**payload, **(await Config.get('audio.tts.openai.params') or {})}
|
||||
api_key = await Config.get('audio.tts.openai.api_key')
|
||||
api_base_url = await Config.get('audio.tts.openai.api_base_url')
|
||||
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': f'Bearer {request.app.state.config.TTS_OPENAI_API_KEY}',
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
}
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
@@ -404,7 +360,7 @@ async def _tts_openai(request, payload, file_path, file_body_path, user):
|
||||
try:
|
||||
session = await get_session()
|
||||
r = await session.post(
|
||||
url=f'{request.app.state.config.TTS_OPENAI_API_BASE_URL}/audio/speech',
|
||||
url=f'{api_base_url}/audio/speech',
|
||||
json=payload,
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -444,13 +400,13 @@ async def _tts_elevenlabs(request, payload, file_path, file_body_path, user):
|
||||
f'{ELEVENLABS_API_BASE_URL}/v1/text-to-speech/{voice_id}',
|
||||
json={
|
||||
'text': payload['input'],
|
||||
'model_id': request.app.state.config.TTS_MODEL,
|
||||
'model_id': await Config.get('audio.tts.model'),
|
||||
'voice_settings': {'stability': 0.5, 'similarity_boost': 0.5},
|
||||
},
|
||||
headers={
|
||||
'Accept': 'audio/mpeg',
|
||||
'Content-Type': 'application/json',
|
||||
'xi-api-key': request.app.state.config.TTS_API_KEY,
|
||||
'xi-api-key': await Config.get('audio.tts.api_key'),
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
@@ -464,11 +420,11 @@ async def _tts_elevenlabs(request, payload, file_path, file_body_path, user):
|
||||
|
||||
async def _tts_azure(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via Azure Cognitive Services TTS."""
|
||||
az_region = request.app.state.config.TTS_AZURE_SPEECH_REGION or 'eastus'
|
||||
az_base = request.app.state.config.TTS_AZURE_SPEECH_BASE_URL
|
||||
language = payload.get('voice') or request.app.state.config.TTS_VOICE
|
||||
az_region = await Config.get('audio.tts.azure.speech_region') or 'eastus'
|
||||
az_base = await Config.get('audio.tts.azure.speech_base_url')
|
||||
language = payload.get('voice') or await Config.get('audio.tts.voice')
|
||||
locale = '-'.join(language.split('-')[:2])
|
||||
output_format = request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT
|
||||
output_format = await Config.get('audio.tts.azure.speech_output_format')
|
||||
|
||||
ssml = (
|
||||
f'<speak version="1.0" xmlns="http://www.w3.org/2001/10/synthesis" xml:lang="{locale}">'
|
||||
@@ -482,7 +438,7 @@ async def _tts_azure(request, payload, file_path, file_body_path, user):
|
||||
async with session.post(
|
||||
(az_base or f'https://{az_region}.tts.speech.microsoft.com') + '/cognitiveservices/v1',
|
||||
headers={
|
||||
'Ocp-Apim-Subscription-Key': request.app.state.config.TTS_API_KEY,
|
||||
'Ocp-Apim-Subscription-Key': await Config.get('audio.tts.api_key'),
|
||||
'Content-Type': 'application/ssml+xml',
|
||||
'X-Microsoft-OutputFormat': output_format,
|
||||
},
|
||||
@@ -505,7 +461,7 @@ async def _tts_transformers(request, payload, file_path, file_body_path, user):
|
||||
load_speech_pipeline(request)
|
||||
|
||||
embeddings = request.app.state.speech_speaker_embeddings_dataset
|
||||
model_name = request.app.state.config.TTS_MODEL
|
||||
model_name = await Config.get('audio.tts.model')
|
||||
|
||||
idx = 6799
|
||||
try:
|
||||
@@ -533,8 +489,8 @@ async def _tts_transformers(request, payload, file_path, file_body_path, user):
|
||||
|
||||
async def _tts_mistral(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via the Mistral TTS API."""
|
||||
api_key = request.app.state.config.TTS_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
api_key = await Config.get('audio.tts.mistral.api_key')
|
||||
api_base_url = await Config.get('audio.tts.mistral.api_base_url') or 'https://api.mistral.ai/v1'
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=400, detail='Mistral API key is required for Mistral TTS')
|
||||
@@ -546,7 +502,7 @@ async def _tts_mistral(request, payload, file_path, file_body_path, user):
|
||||
url=f'{api_base_url}/audio/speech',
|
||||
json={
|
||||
'input': payload.get('input', ''), # text to synthesize
|
||||
'model': request.app.state.config.TTS_MODEL or 'voxtral-mini-tts-2603',
|
||||
'model': await Config.get('audio.tts.model') or 'voxtral-mini-tts-2603',
|
||||
'voice_id': payload.get('voice', ''),
|
||||
'response_format': 'mp3',
|
||||
},
|
||||
@@ -582,7 +538,7 @@ _TTS_ENGINES = {
|
||||
|
||||
@router.post('/speech')
|
||||
async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
engine = request.app.state.config.TTS_ENGINE
|
||||
engine = await Config.get('audio.tts.engine')
|
||||
if engine == '':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
@@ -590,7 +546,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'chat.tts', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -599,7 +555,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
body = await request.body()
|
||||
name = hashlib.sha256(
|
||||
body + str(engine).encode('utf-8') + str(request.app.state.config.TTS_MODEL).encode('utf-8')
|
||||
body + str(engine).encode('utf-8') + str(await Config.get('audio.tts.model')).encode('utf-8')
|
||||
).hexdigest()
|
||||
|
||||
file_path = SPEECH_CACHE_DIR.joinpath(f'{name}.mp3')
|
||||
@@ -624,7 +580,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
async def _transcribe_whisper(request, file_path, languages, file_dir, id):
|
||||
if request.app.state.faster_whisper_model is None:
|
||||
request.app.state.faster_whisper_model = set_faster_whisper_model(request.app.state.config.WHISPER_MODEL)
|
||||
request.app.state.faster_whisper_model = set_faster_whisper_model(await Config.get('audio.stt.whisper_model'))
|
||||
|
||||
model = request.app.state.faster_whisper_model
|
||||
|
||||
@@ -655,11 +611,13 @@ async def _transcribe_openai(request, file_path, filename, languages, file_dir,
|
||||
try:
|
||||
session = await get_session()
|
||||
for language in languages:
|
||||
payload = {'model': request.app.state.config.STT_MODEL}
|
||||
payload = {'model': await Config.get('audio.stt.model')}
|
||||
if language:
|
||||
payload['language'] = language
|
||||
api_key = await Config.get('audio.stt.openai.api_key')
|
||||
api_base_url = await Config.get('audio.stt.openai.api_base_url')
|
||||
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.STT_OPENAI_API_KEY}'}
|
||||
headers = {'Authorization': f'Bearer {api_key}'}
|
||||
if user and ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
@@ -669,7 +627,7 @@ async def _transcribe_openai(request, file_path, filename, languages, file_dir,
|
||||
form_data.add_field('file', open(file_path, 'rb'), filename=filename)
|
||||
|
||||
r = await session.post(
|
||||
url=f'{request.app.state.config.STT_OPENAI_API_BASE_URL}/audio/transcriptions',
|
||||
url=f'{api_base_url}/audio/transcriptions',
|
||||
headers=headers,
|
||||
data=form_data,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -703,8 +661,8 @@ async def _transcribe_deepgram(request, file_path, languages, file_dir, id):
|
||||
async with aiofiles.open(file_path, 'rb') as f:
|
||||
audio_bytes = await f.read()
|
||||
|
||||
api_key = request.app.state.config.DEEPGRAM_API_KEY
|
||||
stt_model = request.app.state.config.STT_MODEL
|
||||
api_key = await Config.get('audio.stt.deepgram.api_key')
|
||||
stt_model = await Config.get('audio.stt.model')
|
||||
|
||||
r = None
|
||||
try:
|
||||
@@ -771,11 +729,11 @@ async def _transcribe_azure(request, file_path, filename, file_dir, id):
|
||||
detail=f'File size ({audio_size // (1024 * 1024)}MB) exceeds Azure limit of {AZURE_MAX_FILE_SIZE_MB}MB',
|
||||
)
|
||||
|
||||
api_key = request.app.state.config.AUDIO_STT_AZURE_API_KEY
|
||||
region = request.app.state.config.AUDIO_STT_AZURE_REGION or 'eastus'
|
||||
locale_str = request.app.state.config.AUDIO_STT_AZURE_LOCALES
|
||||
base_url = request.app.state.config.AUDIO_STT_AZURE_BASE_URL
|
||||
max_speakers = request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS or 3
|
||||
api_key = await Config.get('audio.stt.azure.api_key')
|
||||
region = await Config.get('audio.stt.azure.region') or 'eastus'
|
||||
locale_str = await Config.get('audio.stt.azure.locales')
|
||||
base_url = await Config.get('audio.stt.azure.base_url')
|
||||
max_speakers = await Config.get('audio.stt.azure.max_speakers') or 3
|
||||
|
||||
# Default to a broad set of locales when none are configured
|
||||
if len(locale_str) < 2:
|
||||
@@ -885,16 +843,16 @@ async def transcription_handler(request, file_path, metadata, user=None):
|
||||
None, # Always fallback to None in case transcription fails
|
||||
]
|
||||
|
||||
if request.app.state.config.STT_ENGINE == '':
|
||||
if await Config.get('audio.stt.engine') == '':
|
||||
return await _transcribe_whisper(request, file_path, languages, file_dir, id)
|
||||
elif request.app.state.config.STT_ENGINE == 'openai':
|
||||
elif await Config.get('audio.stt.engine') == 'openai':
|
||||
return await _transcribe_openai(request, file_path, filename, languages, file_dir, id, user)
|
||||
elif request.app.state.config.STT_ENGINE == 'deepgram':
|
||||
elif await Config.get('audio.stt.engine') == 'deepgram':
|
||||
return await _transcribe_deepgram(request, file_path, languages, file_dir, id)
|
||||
elif request.app.state.config.STT_ENGINE == 'azure':
|
||||
elif await Config.get('audio.stt.engine') == 'azure':
|
||||
return await _transcribe_azure(request, file_path, filename, file_dir, id)
|
||||
|
||||
elif request.app.state.config.STT_ENGINE == 'mistral':
|
||||
elif await Config.get('audio.stt.engine') == 'mistral':
|
||||
return await _transcribe_mistral(request, file_path, filename, metadata, file_dir, id)
|
||||
|
||||
|
||||
@@ -907,16 +865,16 @@ async def _transcribe_mistral(request, file_path, filename, metadata, file_dir,
|
||||
if file_size > MAX_FILE_SIZE:
|
||||
raise HTTPException(status_code=400, detail=f'File size exceeds limit of {MAX_FILE_SIZE_MB}MB')
|
||||
|
||||
api_key = request.app.state.config.AUDIO_STT_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
use_chat_completions = request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS
|
||||
api_key = await Config.get('audio.stt.mistral.api_key')
|
||||
api_base_url = await Config.get('audio.stt.mistral.api_base_url') or 'https://api.mistral.ai/v1'
|
||||
use_chat_completions = await Config.get('audio.stt.mistral.use_chat_completions')
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=400, detail='Mistral API key is required for Mistral STT')
|
||||
|
||||
r = None
|
||||
try:
|
||||
model = request.app.state.config.STT_MODEL or 'voxtral-mini-latest'
|
||||
model = await Config.get('audio.stt.model') or 'voxtral-mini-latest'
|
||||
log.info(
|
||||
f'Mistral STT - model: {model}, method: {"chat_completions" if use_chat_completions else "transcriptions"}'
|
||||
)
|
||||
@@ -1158,14 +1116,14 @@ async def transcription(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'chat.stt', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
log.info(f'file.content_type: {file.content_type}')
|
||||
stt_supported_content_types = getattr(request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', [])
|
||||
stt_supported_content_types = await Config.get('audio.stt.supported_content_types', [])
|
||||
|
||||
if not strict_match_mime_type(stt_supported_content_types, file.content_type):
|
||||
raise HTTPException(
|
||||
@@ -1177,7 +1135,7 @@ async def transcription(
|
||||
safe_name = os.path.basename(file.filename) if file.filename else ''
|
||||
ext = safe_name.rsplit('.', 1)[-1].lower() if '.' in safe_name else ''
|
||||
|
||||
allowed_extensions = getattr(request.app.state.config, 'STT_ALLOWED_EXTENSIONS', [])
|
||||
allowed_extensions = await Config.get('audio.stt.allowed_extensions', [])
|
||||
if allowed_extensions and ext not in allowed_extensions:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
@@ -1237,11 +1195,11 @@ async def transcription(
|
||||
async def get_available_models(request: Request) -> list[dict]:
|
||||
"""Return the list of available TTS models for the configured engine."""
|
||||
available_models = []
|
||||
engine = request.app.state.config.TTS_ENGINE
|
||||
engine = await Config.get('audio.tts.engine')
|
||||
_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST)
|
||||
|
||||
if engine == 'openai':
|
||||
base_url = request.app.state.config.TTS_OPENAI_API_BASE_URL
|
||||
base_url = await Config.get('audio.tts.openai.api_base_url')
|
||||
if not base_url.startswith('https://api.openai.com'):
|
||||
session = await get_session()
|
||||
try:
|
||||
@@ -1276,7 +1234,7 @@ async def get_available_models(request: Request) -> list[dict]:
|
||||
async with session.get(
|
||||
f'{ELEVENLABS_API_BASE_URL}/v1/models',
|
||||
headers={
|
||||
'xi-api-key': request.app.state.config.TTS_API_KEY,
|
||||
'xi-api-key': await Config.get('audio.tts.api_key'),
|
||||
'Content-Type': 'application/json',
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -1311,11 +1269,11 @@ _OPENAI_DEFAULT_VOICES = {
|
||||
|
||||
async def get_available_voices(request) -> dict:
|
||||
"""Return ``{voice_id: voice_name}`` for the configured TTS engine."""
|
||||
engine = request.app.state.config.TTS_ENGINE
|
||||
engine = await Config.get('audio.tts.engine')
|
||||
_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST)
|
||||
|
||||
if engine == 'openai':
|
||||
base_url = request.app.state.config.TTS_OPENAI_API_BASE_URL
|
||||
base_url = await Config.get('audio.tts.openai.api_base_url')
|
||||
if not base_url.startswith('https://api.openai.com'):
|
||||
try:
|
||||
session = await get_session()
|
||||
@@ -1338,7 +1296,7 @@ async def get_available_voices(request) -> dict:
|
||||
async with session.get(
|
||||
f'{ELEVENLABS_API_BASE_URL}/v1/voices',
|
||||
headers={
|
||||
'xi-api-key': request.app.state.config.TTS_API_KEY,
|
||||
'xi-api-key': await Config.get('audio.tts.api_key'),
|
||||
'Content-Type': 'application/json',
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -1353,14 +1311,14 @@ async def get_available_voices(request) -> dict:
|
||||
|
||||
if engine == 'azure':
|
||||
try:
|
||||
region = request.app.state.config.TTS_AZURE_SPEECH_REGION
|
||||
base_url = request.app.state.config.TTS_AZURE_SPEECH_BASE_URL
|
||||
region = await Config.get('audio.tts.azure.speech_region')
|
||||
base_url = await Config.get('audio.tts.azure.speech_base_url')
|
||||
url = (base_url or f'https://{region}.tts.speech.microsoft.com') + '/cognitiveservices/voices/list'
|
||||
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url,
|
||||
headers={'Ocp-Apim-Subscription-Key': request.app.state.config.TTS_API_KEY},
|
||||
headers={'Ocp-Apim-Subscription-Key': await Config.get('audio.tts.api_key')},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
timeout=_timeout,
|
||||
) as resp:
|
||||
@@ -1372,8 +1330,8 @@ async def get_available_voices(request) -> dict:
|
||||
return {}
|
||||
|
||||
if engine == 'mistral':
|
||||
api_key = request.app.state.config.TTS_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
api_key = await Config.get('audio.tts.mistral.api_key')
|
||||
api_base_url = await Config.get('audio.tts.mistral.api_base_url') or 'https://api.mistral.ai/v1'
|
||||
if api_key:
|
||||
try:
|
||||
session = await get_session()
|
||||
|
||||
+263
-198
@@ -8,21 +8,15 @@ import time
|
||||
import urllib
|
||||
import uuid
|
||||
from ssl import CERT_NONE, CERT_REQUIRED, PROTOCOL_TLS
|
||||
from typing import List, Optional
|
||||
|
||||
from aiohttp import ClientSession
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import JSONResponse, RedirectResponse, Response
|
||||
from fastapi.responses import JSONResponse, Response
|
||||
from ldap3 import NONE, Connection, Server, Tls
|
||||
from ldap3.utils.conv import escape_filter_chars
|
||||
from open_webui.config import (
|
||||
ENABLE_LDAP,
|
||||
ENABLE_OAUTH_SIGNUP,
|
||||
ENABLE_PASSWORD_AUTH,
|
||||
OAUTH_MERGE_ACCOUNTS_BY_EMAIL,
|
||||
OAUTH_PROVIDERS,
|
||||
OPENID_END_SESSION_ENDPOINT,
|
||||
OPENID_PROVIDER_URL,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES, WEBHOOK_MESSAGES
|
||||
from open_webui.env import (
|
||||
@@ -50,6 +44,7 @@ from open_webui.models.auths import (
|
||||
Token,
|
||||
UpdatePasswordForm,
|
||||
)
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.models.users import (
|
||||
@@ -75,7 +70,6 @@ from open_webui.utils.auth import (
|
||||
)
|
||||
from open_webui.utils.groups import apply_default_group_assignment
|
||||
from open_webui.utils.misc import parse_duration, validate_email_format
|
||||
from open_webui.utils.oauth import auth_manager_config
|
||||
from open_webui.utils.rate_limit import RateLimiter
|
||||
from open_webui.utils.redis import get_redis_client
|
||||
from open_webui.utils.webhook import post_webhook
|
||||
@@ -90,6 +84,60 @@ log = logging.getLogger(__name__)
|
||||
# who exceed their allotted rate against this gate.
|
||||
signin_rate_limiter = RateLimiter(redis_client=get_redis_client(), limit=5 * 3, window=60 * 3)
|
||||
|
||||
ADMIN_CONFIG_KEYS = {
|
||||
'SHOW_ADMIN_DETAILS': 'auth.admin.show',
|
||||
'ADMIN_EMAIL': 'auth.admin.email',
|
||||
'WEBUI_URL': 'webui.url',
|
||||
'ENABLE_SIGNUP': 'ui.enable_signup',
|
||||
'ENABLE_API_KEYS': 'auth.enable_api_keys',
|
||||
'ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS': 'auth.api_key.endpoint_restrictions',
|
||||
'API_KEYS_ALLOWED_ENDPOINTS': 'auth.api_key.allowed_endpoints',
|
||||
'DEFAULT_USER_ROLE': 'ui.default_user_role',
|
||||
'DEFAULT_GROUP_ID': 'ui.default_group_id',
|
||||
'JWT_EXPIRES_IN': 'auth.jwt_expiry',
|
||||
'ENABLE_COMMUNITY_SHARING': 'ui.enable_community_sharing',
|
||||
'ENABLE_MESSAGE_RATING': 'ui.enable_message_rating',
|
||||
'ENABLE_FOLDERS': 'folders.enable',
|
||||
'FOLDER_MAX_FILE_COUNT': 'folders.max_file_count',
|
||||
'AUTOMATION_MAX_COUNT': 'automations.max_count',
|
||||
'AUTOMATION_MIN_INTERVAL': 'automations.min_interval',
|
||||
'ENABLE_AUTOMATIONS': 'automations.enable',
|
||||
'ENABLE_CHANNELS': 'channels.enable',
|
||||
'ENABLE_CALENDAR': 'calendar.enable',
|
||||
'ENABLE_MEMORIES': 'memories.enable',
|
||||
'ENABLE_NOTES': 'notes.enable',
|
||||
'ENABLE_USER_WEBHOOKS': 'ui.enable_user_webhooks',
|
||||
'ENABLE_USER_STATUS': 'users.enable_status',
|
||||
'PENDING_USER_OVERLAY_TITLE': 'ui.pending_user_overlay_title',
|
||||
'PENDING_USER_OVERLAY_CONTENT': 'ui.pending_user_overlay_content',
|
||||
'RESPONSE_WATERMARK': 'ui.watermark',
|
||||
}
|
||||
|
||||
LDAP_SERVER_CONFIG_KEYS = {
|
||||
'label': 'ldap.server.label',
|
||||
'host': 'ldap.server.host',
|
||||
'port': 'ldap.server.port',
|
||||
'attribute_for_mail': 'ldap.server.attribute_for_mail',
|
||||
'attribute_for_username': 'ldap.server.attribute_for_username',
|
||||
'app_dn': 'ldap.server.app_dn',
|
||||
'app_dn_password': 'ldap.server.app_password',
|
||||
'search_base': 'ldap.server.users_dn',
|
||||
'search_filters': 'ldap.server.search_filter',
|
||||
'use_tls': 'ldap.server.use_tls',
|
||||
'certificate_path': 'ldap.server.ca_cert_file',
|
||||
'validate_cert': 'ldap.server.validate_cert',
|
||||
'ciphers': 'ldap.server.ciphers',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
async def create_session_response(
|
||||
request: Request, user, db, response: Response = None, set_cookie: bool = False
|
||||
@@ -105,7 +153,7 @@ async def create_session_response(
|
||||
response: FastAPI response object (required if set_cookie is True)
|
||||
set_cookie: Whether to set the auth cookie on the response
|
||||
"""
|
||||
expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN)
|
||||
expires_delta = parse_duration(await Config.get('auth.jwt_expiry'))
|
||||
expires_at = None
|
||||
if expires_delta:
|
||||
expires_at = int(time.time()) + int(expires_delta.total_seconds())
|
||||
@@ -128,7 +176,7 @@ async def create_session_response(
|
||||
**({'max_age': max_age} if max_age is not None else {}),
|
||||
)
|
||||
|
||||
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
user_permissions = await get_permissions(user.id, await Config.get('user.permissions'), db=db)
|
||||
|
||||
return {
|
||||
'token': token,
|
||||
@@ -201,7 +249,7 @@ async def get_session_user(
|
||||
**({'max_age': max_age} if max_age is not None else {}),
|
||||
)
|
||||
|
||||
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
user_permissions = await get_permissions(user.id, await Config.get('user.permissions'), db=db)
|
||||
|
||||
response_data = {
|
||||
'token': token,
|
||||
@@ -320,7 +368,7 @@ async def ldap_auth(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
# Security checks FIRST - before loading any config
|
||||
if not request.app.state.config.ENABLE_LDAP:
|
||||
if not await Config.get('ldap.enable'):
|
||||
raise HTTPException(400, detail='LDAP authentication is not enabled')
|
||||
|
||||
if not ENABLE_PASSWORD_AUTH:
|
||||
@@ -338,19 +386,19 @@ async def ldap_auth(
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
|
||||
|
||||
# NOW load LDAP config variables
|
||||
LDAP_SERVER_LABEL = request.app.state.config.LDAP_SERVER_LABEL
|
||||
LDAP_SERVER_HOST = request.app.state.config.LDAP_SERVER_HOST
|
||||
LDAP_SERVER_PORT = request.app.state.config.LDAP_SERVER_PORT
|
||||
LDAP_ATTRIBUTE_FOR_MAIL = request.app.state.config.LDAP_ATTRIBUTE_FOR_MAIL
|
||||
LDAP_ATTRIBUTE_FOR_USERNAME = request.app.state.config.LDAP_ATTRIBUTE_FOR_USERNAME
|
||||
LDAP_SEARCH_BASE = request.app.state.config.LDAP_SEARCH_BASE
|
||||
LDAP_SEARCH_FILTERS = request.app.state.config.LDAP_SEARCH_FILTERS
|
||||
LDAP_APP_DN = request.app.state.config.LDAP_APP_DN
|
||||
LDAP_APP_PASSWORD = request.app.state.config.LDAP_APP_PASSWORD
|
||||
LDAP_USE_TLS = request.app.state.config.LDAP_USE_TLS
|
||||
LDAP_CA_CERT_FILE = request.app.state.config.LDAP_CA_CERT_FILE
|
||||
LDAP_VALIDATE_CERT = CERT_REQUIRED if request.app.state.config.LDAP_VALIDATE_CERT else CERT_NONE
|
||||
LDAP_CIPHERS = request.app.state.config.LDAP_CIPHERS if request.app.state.config.LDAP_CIPHERS else 'ALL'
|
||||
LDAP_SERVER_LABEL = await Config.get('ldap.server.label')
|
||||
LDAP_SERVER_HOST = await Config.get('ldap.server.host')
|
||||
LDAP_SERVER_PORT = await Config.get('ldap.server.port')
|
||||
LDAP_ATTRIBUTE_FOR_MAIL = await Config.get('ldap.server.attribute_for_mail')
|
||||
LDAP_ATTRIBUTE_FOR_USERNAME = await Config.get('ldap.server.attribute_for_username')
|
||||
LDAP_SEARCH_BASE = await Config.get('ldap.server.users_dn')
|
||||
LDAP_SEARCH_FILTERS = await Config.get('ldap.server.search_filter')
|
||||
LDAP_APP_DN = await Config.get('ldap.server.app_dn')
|
||||
LDAP_APP_PASSWORD = await Config.get('ldap.server.app_password')
|
||||
LDAP_USE_TLS = await Config.get('ldap.server.use_tls')
|
||||
LDAP_CA_CERT_FILE = await Config.get('ldap.server.ca_cert_file')
|
||||
LDAP_VALIDATE_CERT = CERT_REQUIRED if await Config.get('ldap.server.validate_cert') else CERT_NONE
|
||||
LDAP_CIPHERS = await Config.get('ldap.server.ciphers') if await Config.get('ldap.server.ciphers') else 'ALL'
|
||||
|
||||
try:
|
||||
tls = Tls(
|
||||
@@ -381,9 +429,9 @@ async def ldap_auth(
|
||||
if not await asyncio.to_thread(connection_app.bind):
|
||||
raise HTTPException(400, detail='Application account bind failed')
|
||||
|
||||
ENABLE_LDAP_GROUP_MANAGEMENT = request.app.state.config.ENABLE_LDAP_GROUP_MANAGEMENT
|
||||
ENABLE_LDAP_GROUP_CREATION = request.app.state.config.ENABLE_LDAP_GROUP_CREATION
|
||||
LDAP_ATTRIBUTE_FOR_GROUPS = request.app.state.config.LDAP_ATTRIBUTE_FOR_GROUPS
|
||||
ENABLE_LDAP_GROUP_MANAGEMENT = await Config.get('ldap.group.enable_management')
|
||||
ENABLE_LDAP_GROUP_CREATION = await Config.get('ldap.group.enable_creation')
|
||||
LDAP_ATTRIBUTE_FOR_GROUPS = await Config.get('ldap.server.attribute_for_groups')
|
||||
|
||||
search_attributes = [
|
||||
f'{LDAP_ATTRIBUTE_FOR_USERNAME}',
|
||||
@@ -500,7 +548,7 @@ async def ldap_auth(
|
||||
email=email,
|
||||
password=str(uuid.uuid4()),
|
||||
name=cn,
|
||||
role=request.app.state.config.DEFAULT_USER_ROLE,
|
||||
role=await Config.get('ui.default_user_role'),
|
||||
db=db,
|
||||
)
|
||||
|
||||
@@ -514,15 +562,15 @@ async def ldap_auth(
|
||||
user = await Users.get_user_by_id(user.id, db=db)
|
||||
|
||||
await apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
await Config.get('ui.default_group_id'),
|
||||
user.id,
|
||||
db=db,
|
||||
)
|
||||
|
||||
if request.app.state.config.WEBHOOK_URL:
|
||||
if await Config.get('webhook_url'):
|
||||
await post_webhook(
|
||||
request.app.state.WEBUI_NAME,
|
||||
request.app.state.config.WEBHOOK_URL,
|
||||
await Config.get('webhook_url'),
|
||||
WEBHOOK_MESSAGES.USER_SIGNUP(user.name),
|
||||
{
|
||||
'action': 'signup',
|
||||
@@ -703,7 +751,7 @@ async def signup_handler(
|
||||
password=hashed,
|
||||
name=name,
|
||||
profile_image_url=profile_image_url,
|
||||
role=request.app.state.config.DEFAULT_USER_ROLE,
|
||||
role=await Config.get('ui.default_user_role'),
|
||||
db=db,
|
||||
)
|
||||
if not user:
|
||||
@@ -714,12 +762,12 @@ async def signup_handler(
|
||||
if await Users.get_num_users(db=db) == 1:
|
||||
await Users.update_user_role_by_id(user.id, 'admin', db=db)
|
||||
user = await Users.get_user_by_id(user.id, db=db)
|
||||
request.app.state.config.ENABLE_SIGNUP = False
|
||||
await Config.upsert({'ui.enable_signup': False})
|
||||
|
||||
if request.app.state.config.WEBHOOK_URL:
|
||||
if await Config.get('webhook_url'):
|
||||
await post_webhook(
|
||||
request.app.state.WEBUI_NAME,
|
||||
request.app.state.config.WEBHOOK_URL,
|
||||
await Config.get('webhook_url'),
|
||||
WEBHOOK_MESSAGES.USER_SIGNUP(user.name),
|
||||
{
|
||||
'action': 'signup',
|
||||
@@ -729,7 +777,7 @@ async def signup_handler(
|
||||
)
|
||||
|
||||
await apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
await Config.get('ui.default_group_id'),
|
||||
user.id,
|
||||
db=db,
|
||||
)
|
||||
@@ -748,10 +796,10 @@ async def signup(
|
||||
|
||||
if WEBUI_AUTH:
|
||||
if has_users:
|
||||
if not request.app.state.config.ENABLE_SIGNUP or not request.app.state.config.ENABLE_LOGIN_FORM:
|
||||
if not await Config.get('ui.enable_signup') or not await Config.get('ui.enable_login_form'):
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
# Don't gate the first admin on ENABLE_SIGNUP: it auto-disables and can persist stale across a DB reset.
|
||||
elif not request.app.state.config.ENABLE_LOGIN_FORM and not ENABLE_INITIAL_ADMIN_SIGNUP:
|
||||
elif not await Config.get('ui.enable_login_form') and not ENABLE_INITIAL_ADMIN_SIGNUP:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
else:
|
||||
if has_users:
|
||||
@@ -812,19 +860,21 @@ async def signout(request: Request, response: Response, db: AsyncSession = Depen
|
||||
|
||||
# If a custom end_session_endpoint is configured (e.g. AWS Cognito), redirect
|
||||
# there directly instead of attempting OIDC discovery.
|
||||
if OPENID_END_SESSION_ENDPOINT.value:
|
||||
openid_end_session_endpoint = await Config.get('oauth.end_session_endpoint')
|
||||
if openid_end_session_endpoint:
|
||||
return JSONResponse(
|
||||
status_code=200,
|
||||
content={
|
||||
'status': True,
|
||||
'redirect_url': OPENID_END_SESSION_ENDPOINT.value,
|
||||
'redirect_url': openid_end_session_endpoint,
|
||||
},
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
openid_provider_url = await Config.get('oauth.provider_url')
|
||||
oauth_server_metadata_url = (
|
||||
request.app.state.oauth_manager.get_server_metadata_url(session.provider) if session else None
|
||||
) or OPENID_PROVIDER_URL.value
|
||||
) or openid_provider_url
|
||||
|
||||
if session and oauth_server_metadata_url:
|
||||
oauth_id_token = session.token.get('id_token')
|
||||
@@ -934,12 +984,12 @@ async def add_user(
|
||||
|
||||
if user:
|
||||
await apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
await Config.get('ui.default_group_id'),
|
||||
user.id,
|
||||
db=db,
|
||||
)
|
||||
|
||||
expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN)
|
||||
expires_delta = parse_duration(await Config.get('auth.jwt_expiry'))
|
||||
token = create_token(data={'id': user.id}, expires_delta=expires_delta)
|
||||
return {
|
||||
'token': token,
|
||||
@@ -968,8 +1018,8 @@ async def add_user(
|
||||
async def get_admin_details(
|
||||
request: Request, user=Depends(get_current_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
if request.app.state.config.SHOW_ADMIN_DETAILS:
|
||||
admin_email = request.app.state.config.ADMIN_EMAIL
|
||||
if await Config.get('auth.admin.show'):
|
||||
admin_email = await Config.get('auth.admin.email')
|
||||
admin_name = None
|
||||
|
||||
log.info(f'Admin details - Email: {admin_email}, Name: {admin_name}')
|
||||
@@ -999,34 +1049,7 @@ async def get_admin_details(
|
||||
|
||||
@router.get('/admin/config')
|
||||
async def get_admin_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'SHOW_ADMIN_DETAILS': request.app.state.config.SHOW_ADMIN_DETAILS,
|
||||
'ADMIN_EMAIL': request.app.state.config.ADMIN_EMAIL,
|
||||
'WEBUI_URL': request.app.state.config.WEBUI_URL,
|
||||
'ENABLE_SIGNUP': request.app.state.config.ENABLE_SIGNUP,
|
||||
'ENABLE_API_KEYS': request.app.state.config.ENABLE_API_KEYS,
|
||||
'ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS': request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS,
|
||||
'API_KEYS_ALLOWED_ENDPOINTS': request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS,
|
||||
'DEFAULT_USER_ROLE': request.app.state.config.DEFAULT_USER_ROLE,
|
||||
'DEFAULT_GROUP_ID': request.app.state.config.DEFAULT_GROUP_ID,
|
||||
'JWT_EXPIRES_IN': request.app.state.config.JWT_EXPIRES_IN,
|
||||
'ENABLE_COMMUNITY_SHARING': request.app.state.config.ENABLE_COMMUNITY_SHARING,
|
||||
'ENABLE_MESSAGE_RATING': request.app.state.config.ENABLE_MESSAGE_RATING,
|
||||
'ENABLE_FOLDERS': request.app.state.config.ENABLE_FOLDERS,
|
||||
'FOLDER_MAX_FILE_COUNT': request.app.state.config.FOLDER_MAX_FILE_COUNT,
|
||||
'AUTOMATION_MAX_COUNT': request.app.state.config.AUTOMATION_MAX_COUNT,
|
||||
'AUTOMATION_MIN_INTERVAL': request.app.state.config.AUTOMATION_MIN_INTERVAL,
|
||||
'ENABLE_AUTOMATIONS': request.app.state.config.ENABLE_AUTOMATIONS,
|
||||
'ENABLE_CHANNELS': request.app.state.config.ENABLE_CHANNELS,
|
||||
'ENABLE_CALENDAR': request.app.state.config.ENABLE_CALENDAR,
|
||||
'ENABLE_MEMORIES': request.app.state.config.ENABLE_MEMORIES,
|
||||
'ENABLE_NOTES': request.app.state.config.ENABLE_NOTES,
|
||||
'ENABLE_USER_WEBHOOKS': request.app.state.config.ENABLE_USER_WEBHOOKS,
|
||||
'ENABLE_USER_STATUS': request.app.state.config.ENABLE_USER_STATUS,
|
||||
'PENDING_USER_OVERLAY_TITLE': request.app.state.config.PENDING_USER_OVERLAY_TITLE,
|
||||
'PENDING_USER_OVERLAY_CONTENT': request.app.state.config.PENDING_USER_OVERLAY_CONTENT,
|
||||
'RESPONSE_WATERMARK': request.app.state.config.RESPONSE_WATERMARK,
|
||||
}
|
||||
return await get_config_values(ADMIN_CONFIG_KEYS)
|
||||
|
||||
|
||||
class AdminConfig(BaseModel):
|
||||
@@ -1060,81 +1083,24 @@ class AdminConfig(BaseModel):
|
||||
|
||||
@router.post('/admin/config')
|
||||
async def update_admin_config(request: Request, form_data: AdminConfig, user=Depends(get_admin_user)):
|
||||
request.app.state.config.SHOW_ADMIN_DETAILS = form_data.SHOW_ADMIN_DETAILS
|
||||
request.app.state.config.ADMIN_EMAIL = form_data.ADMIN_EMAIL
|
||||
request.app.state.config.WEBUI_URL = form_data.WEBUI_URL
|
||||
request.app.state.config.ENABLE_SIGNUP = form_data.ENABLE_SIGNUP
|
||||
|
||||
request.app.state.config.ENABLE_API_KEYS = form_data.ENABLE_API_KEYS
|
||||
request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS = form_data.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS
|
||||
request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS = form_data.API_KEYS_ALLOWED_ENDPOINTS
|
||||
|
||||
request.app.state.config.ENABLE_FOLDERS = form_data.ENABLE_FOLDERS
|
||||
request.app.state.config.FOLDER_MAX_FILE_COUNT = (
|
||||
int(form_data.FOLDER_MAX_FILE_COUNT) if form_data.FOLDER_MAX_FILE_COUNT else ''
|
||||
)
|
||||
request.app.state.config.AUTOMATION_MAX_COUNT = (
|
||||
int(form_data.AUTOMATION_MAX_COUNT) if form_data.AUTOMATION_MAX_COUNT else ''
|
||||
)
|
||||
request.app.state.config.AUTOMATION_MIN_INTERVAL = (
|
||||
updates = config_updates(form_data.model_dump(), ADMIN_CONFIG_KEYS)
|
||||
updates['folders.max_file_count'] = int(form_data.FOLDER_MAX_FILE_COUNT) if form_data.FOLDER_MAX_FILE_COUNT else ''
|
||||
updates['automations.max_count'] = int(form_data.AUTOMATION_MAX_COUNT) if form_data.AUTOMATION_MAX_COUNT else ''
|
||||
updates['automations.min_interval'] = (
|
||||
int(form_data.AUTOMATION_MIN_INTERVAL) if form_data.AUTOMATION_MIN_INTERVAL else ''
|
||||
)
|
||||
request.app.state.config.ENABLE_AUTOMATIONS = form_data.ENABLE_AUTOMATIONS
|
||||
request.app.state.config.ENABLE_CHANNELS = form_data.ENABLE_CHANNELS
|
||||
request.app.state.config.ENABLE_CALENDAR = form_data.ENABLE_CALENDAR
|
||||
request.app.state.config.ENABLE_MEMORIES = form_data.ENABLE_MEMORIES
|
||||
request.app.state.config.ENABLE_NOTES = form_data.ENABLE_NOTES
|
||||
|
||||
if form_data.DEFAULT_USER_ROLE in ['pending', 'user', 'admin']:
|
||||
request.app.state.config.DEFAULT_USER_ROLE = form_data.DEFAULT_USER_ROLE
|
||||
|
||||
request.app.state.config.DEFAULT_GROUP_ID = form_data.DEFAULT_GROUP_ID
|
||||
if form_data.DEFAULT_USER_ROLE not in ['pending', 'user', 'admin']:
|
||||
updates.pop('ui.default_user_role', None)
|
||||
|
||||
pattern = r'^(-1|0|(-?\d+(\.\d+)?)(ms|s|m|h|d|w))$'
|
||||
|
||||
# Check if the input string matches the pattern
|
||||
if re.match(pattern, form_data.JWT_EXPIRES_IN):
|
||||
request.app.state.config.JWT_EXPIRES_IN = form_data.JWT_EXPIRES_IN
|
||||
if not re.match(pattern, form_data.JWT_EXPIRES_IN):
|
||||
updates.pop('auth.jwt_expiry', None)
|
||||
|
||||
request.app.state.config.ENABLE_COMMUNITY_SHARING = form_data.ENABLE_COMMUNITY_SHARING
|
||||
request.app.state.config.ENABLE_MESSAGE_RATING = form_data.ENABLE_MESSAGE_RATING
|
||||
|
||||
request.app.state.config.ENABLE_USER_WEBHOOKS = form_data.ENABLE_USER_WEBHOOKS
|
||||
request.app.state.config.ENABLE_USER_STATUS = form_data.ENABLE_USER_STATUS
|
||||
|
||||
request.app.state.config.PENDING_USER_OVERLAY_TITLE = form_data.PENDING_USER_OVERLAY_TITLE
|
||||
request.app.state.config.PENDING_USER_OVERLAY_CONTENT = form_data.PENDING_USER_OVERLAY_CONTENT
|
||||
|
||||
request.app.state.config.RESPONSE_WATERMARK = form_data.RESPONSE_WATERMARK
|
||||
|
||||
return {
|
||||
'SHOW_ADMIN_DETAILS': request.app.state.config.SHOW_ADMIN_DETAILS,
|
||||
'ADMIN_EMAIL': request.app.state.config.ADMIN_EMAIL,
|
||||
'WEBUI_URL': request.app.state.config.WEBUI_URL,
|
||||
'ENABLE_SIGNUP': request.app.state.config.ENABLE_SIGNUP,
|
||||
'ENABLE_API_KEYS': request.app.state.config.ENABLE_API_KEYS,
|
||||
'ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS': request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS,
|
||||
'API_KEYS_ALLOWED_ENDPOINTS': request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS,
|
||||
'DEFAULT_USER_ROLE': request.app.state.config.DEFAULT_USER_ROLE,
|
||||
'DEFAULT_GROUP_ID': request.app.state.config.DEFAULT_GROUP_ID,
|
||||
'JWT_EXPIRES_IN': request.app.state.config.JWT_EXPIRES_IN,
|
||||
'ENABLE_COMMUNITY_SHARING': request.app.state.config.ENABLE_COMMUNITY_SHARING,
|
||||
'ENABLE_MESSAGE_RATING': request.app.state.config.ENABLE_MESSAGE_RATING,
|
||||
'ENABLE_FOLDERS': request.app.state.config.ENABLE_FOLDERS,
|
||||
'FOLDER_MAX_FILE_COUNT': request.app.state.config.FOLDER_MAX_FILE_COUNT,
|
||||
'AUTOMATION_MAX_COUNT': request.app.state.config.AUTOMATION_MAX_COUNT,
|
||||
'AUTOMATION_MIN_INTERVAL': request.app.state.config.AUTOMATION_MIN_INTERVAL,
|
||||
'ENABLE_AUTOMATIONS': request.app.state.config.ENABLE_AUTOMATIONS,
|
||||
'ENABLE_CHANNELS': request.app.state.config.ENABLE_CHANNELS,
|
||||
'ENABLE_CALENDAR': request.app.state.config.ENABLE_CALENDAR,
|
||||
'ENABLE_MEMORIES': request.app.state.config.ENABLE_MEMORIES,
|
||||
'ENABLE_NOTES': request.app.state.config.ENABLE_NOTES,
|
||||
'ENABLE_USER_WEBHOOKS': request.app.state.config.ENABLE_USER_WEBHOOKS,
|
||||
'ENABLE_USER_STATUS': request.app.state.config.ENABLE_USER_STATUS,
|
||||
'PENDING_USER_OVERLAY_TITLE': request.app.state.config.PENDING_USER_OVERLAY_TITLE,
|
||||
'PENDING_USER_OVERLAY_CONTENT': request.app.state.config.PENDING_USER_OVERLAY_CONTENT,
|
||||
'RESPONSE_WATERMARK': request.app.state.config.RESPONSE_WATERMARK,
|
||||
}
|
||||
await Config.upsert(updates)
|
||||
return await get_config_values(ADMIN_CONFIG_KEYS)
|
||||
|
||||
|
||||
class LdapServerConfig(BaseModel):
|
||||
@@ -1155,21 +1121,7 @@ class LdapServerConfig(BaseModel):
|
||||
|
||||
@router.get('/admin/config/ldap/server', response_model=LdapServerConfig)
|
||||
async def get_ldap_server(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'label': request.app.state.config.LDAP_SERVER_LABEL,
|
||||
'host': request.app.state.config.LDAP_SERVER_HOST,
|
||||
'port': request.app.state.config.LDAP_SERVER_PORT,
|
||||
'attribute_for_mail': request.app.state.config.LDAP_ATTRIBUTE_FOR_MAIL,
|
||||
'attribute_for_username': request.app.state.config.LDAP_ATTRIBUTE_FOR_USERNAME,
|
||||
'app_dn': request.app.state.config.LDAP_APP_DN,
|
||||
'app_dn_password': request.app.state.config.LDAP_APP_PASSWORD,
|
||||
'search_base': request.app.state.config.LDAP_SEARCH_BASE,
|
||||
'search_filters': request.app.state.config.LDAP_SEARCH_FILTERS,
|
||||
'use_tls': request.app.state.config.LDAP_USE_TLS,
|
||||
'certificate_path': request.app.state.config.LDAP_CA_CERT_FILE,
|
||||
'validate_cert': request.app.state.config.LDAP_VALIDATE_CERT,
|
||||
'ciphers': request.app.state.config.LDAP_CIPHERS,
|
||||
}
|
||||
return await get_config_values(LDAP_SERVER_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/admin/config/ldap/server')
|
||||
@@ -1186,40 +1138,16 @@ async def update_ldap_server(request: Request, form_data: LdapServerConfig, user
|
||||
if not value:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.REQUIRED_FIELD_EMPTY(key))
|
||||
|
||||
request.app.state.config.LDAP_SERVER_LABEL = form_data.label
|
||||
request.app.state.config.LDAP_SERVER_HOST = form_data.host
|
||||
request.app.state.config.LDAP_SERVER_PORT = form_data.port
|
||||
request.app.state.config.LDAP_ATTRIBUTE_FOR_MAIL = form_data.attribute_for_mail
|
||||
request.app.state.config.LDAP_ATTRIBUTE_FOR_USERNAME = form_data.attribute_for_username
|
||||
request.app.state.config.LDAP_APP_DN = form_data.app_dn or ''
|
||||
request.app.state.config.LDAP_APP_PASSWORD = form_data.app_dn_password or ''
|
||||
request.app.state.config.LDAP_SEARCH_BASE = form_data.search_base
|
||||
request.app.state.config.LDAP_SEARCH_FILTERS = form_data.search_filters
|
||||
request.app.state.config.LDAP_USE_TLS = form_data.use_tls
|
||||
request.app.state.config.LDAP_CA_CERT_FILE = form_data.certificate_path
|
||||
request.app.state.config.LDAP_VALIDATE_CERT = form_data.validate_cert
|
||||
request.app.state.config.LDAP_CIPHERS = form_data.ciphers
|
||||
|
||||
return {
|
||||
'label': request.app.state.config.LDAP_SERVER_LABEL,
|
||||
'host': request.app.state.config.LDAP_SERVER_HOST,
|
||||
'port': request.app.state.config.LDAP_SERVER_PORT,
|
||||
'attribute_for_mail': request.app.state.config.LDAP_ATTRIBUTE_FOR_MAIL,
|
||||
'attribute_for_username': request.app.state.config.LDAP_ATTRIBUTE_FOR_USERNAME,
|
||||
'app_dn': request.app.state.config.LDAP_APP_DN,
|
||||
'app_dn_password': request.app.state.config.LDAP_APP_PASSWORD,
|
||||
'search_base': request.app.state.config.LDAP_SEARCH_BASE,
|
||||
'search_filters': request.app.state.config.LDAP_SEARCH_FILTERS,
|
||||
'use_tls': request.app.state.config.LDAP_USE_TLS,
|
||||
'certificate_path': request.app.state.config.LDAP_CA_CERT_FILE,
|
||||
'validate_cert': request.app.state.config.LDAP_VALIDATE_CERT,
|
||||
'ciphers': request.app.state.config.LDAP_CIPHERS,
|
||||
}
|
||||
updates = config_updates(form_data.model_dump(), LDAP_SERVER_CONFIG_KEYS)
|
||||
updates['ldap.server.app_dn'] = form_data.app_dn or ''
|
||||
updates['ldap.server.app_password'] = form_data.app_dn_password or ''
|
||||
await Config.upsert(updates)
|
||||
return await get_config_values(LDAP_SERVER_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.get('/admin/config/ldap')
|
||||
async def get_ldap_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {'ENABLE_LDAP': request.app.state.config.ENABLE_LDAP}
|
||||
return {'ENABLE_LDAP': await Config.get('ldap.enable')}
|
||||
|
||||
|
||||
class LdapConfigForm(BaseModel):
|
||||
@@ -1228,8 +1156,8 @@ class LdapConfigForm(BaseModel):
|
||||
|
||||
@router.post('/admin/config/ldap')
|
||||
async def update_ldap_config(request: Request, form_data: LdapConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.ENABLE_LDAP = form_data.enable_ldap
|
||||
return {'ENABLE_LDAP': request.app.state.config.ENABLE_LDAP}
|
||||
await Config.upsert({'ldap.enable': form_data.enable_ldap})
|
||||
return {'ENABLE_LDAP': await Config.get('ldap.enable')}
|
||||
|
||||
|
||||
############################
|
||||
@@ -1237,11 +1165,148 @@ async def update_ldap_config(request: Request, form_data: LdapConfigForm, user=D
|
||||
############################
|
||||
|
||||
|
||||
class OAuthConfigForm(BaseModel):
|
||||
"""All OAuth/OIDC settings exposed to the admin panel."""
|
||||
|
||||
# General OAuth
|
||||
ENABLE_OAUTH_SIGNUP: bool | None = None
|
||||
OAUTH_MERGE_ACCOUNTS_BY_EMAIL: bool | None = None
|
||||
OAUTH_AUTO_REDIRECT: bool | None = None
|
||||
OAUTH_ALLOWED_DOMAINS: str | None = None
|
||||
OAUTH_BLOCKED_GROUPS: str | None = None
|
||||
|
||||
# Role management
|
||||
ENABLE_OAUTH_ROLE_MANAGEMENT: bool | None = None
|
||||
OAUTH_ROLES_CLAIM: str | None = None
|
||||
OAUTH_ADMIN_ROLES: str | None = None
|
||||
OAUTH_ALLOWED_ROLES: str | None = None
|
||||
|
||||
# Group management
|
||||
ENABLE_OAUTH_GROUP_MANAGEMENT: bool | None = None
|
||||
ENABLE_OAUTH_GROUP_CREATION: bool | None = None
|
||||
OAUTH_GROUP_CLAIM: str | None = None
|
||||
OAUTH_GROUP_DEFAULT_SHARE: bool | str | None = None
|
||||
|
||||
# OIDC provider settings
|
||||
OAUTH_PROVIDER_NAME: str | None = None
|
||||
OPENID_PROVIDER_URL: str | None = None
|
||||
OAUTH_CLIENT_ID: str | None = None
|
||||
OAUTH_CLIENT_SECRET: str | None = None
|
||||
OPENID_REDIRECT_URI: str | None = None
|
||||
OAUTH_SCOPES: str | None = None
|
||||
OAUTH_CODE_CHALLENGE_METHOD: str | None = None
|
||||
OAUTH_TOKEN_ENDPOINT_AUTH_METHOD: str | None = None
|
||||
OPENID_END_SESSION_ENDPOINT: str | None = None
|
||||
OAUTH_TIMEOUT: int | str | None = None
|
||||
OAUTH_CLIENT_TIMEOUT: int | str | None = None
|
||||
|
||||
# Claims
|
||||
OAUTH_EMAIL_CLAIM: str | None = None
|
||||
OAUTH_USERNAME_CLAIM: str | None = None
|
||||
OAUTH_PICTURE_CLAIM: str | None = None
|
||||
OAUTH_SUB_CLAIM: str | None = None
|
||||
OAUTH_AUDIENCE: str | None = None
|
||||
|
||||
# Profile update toggles
|
||||
OAUTH_UPDATE_EMAIL_ON_LOGIN: bool | None = None
|
||||
OAUTH_UPDATE_NAME_ON_LOGIN: bool | None = None
|
||||
OAUTH_UPDATE_PICTURE_ON_LOGIN: bool | None = None
|
||||
|
||||
# Token
|
||||
OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE: bool | None = None
|
||||
|
||||
|
||||
OAUTH_COMMA_LIST_FIELDS = {
|
||||
'OAUTH_ALLOWED_DOMAINS',
|
||||
'OAUTH_ADMIN_ROLES',
|
||||
'OAUTH_ALLOWED_ROLES',
|
||||
}
|
||||
|
||||
|
||||
OAUTH_CONFIG_KEYS = {
|
||||
'ENABLE_OAUTH_SIGNUP': 'oauth.enable_signup',
|
||||
'OAUTH_MERGE_ACCOUNTS_BY_EMAIL': 'oauth.merge_accounts_by_email',
|
||||
'OAUTH_AUTO_REDIRECT': 'oauth.auto_redirect',
|
||||
'OAUTH_ALLOWED_DOMAINS': 'oauth.allowed_domains',
|
||||
'OAUTH_BLOCKED_GROUPS': 'oauth.blocked_groups',
|
||||
'ENABLE_OAUTH_ROLE_MANAGEMENT': 'oauth.enable_role_mapping',
|
||||
'OAUTH_ROLES_CLAIM': 'oauth.roles_claim',
|
||||
'OAUTH_ADMIN_ROLES': 'oauth.admin_roles',
|
||||
'OAUTH_ALLOWED_ROLES': 'oauth.allowed_roles',
|
||||
'ENABLE_OAUTH_GROUP_MANAGEMENT': 'oauth.enable_group_mapping',
|
||||
'ENABLE_OAUTH_GROUP_CREATION': 'oauth.enable_group_creation',
|
||||
'OAUTH_GROUP_CLAIM': 'oauth.group_claim',
|
||||
'OAUTH_GROUP_DEFAULT_SHARE': 'oauth.group_default_share',
|
||||
'OAUTH_PROVIDER_NAME': 'oauth.provider_name',
|
||||
'OPENID_PROVIDER_URL': 'oauth.provider_url',
|
||||
'OAUTH_CLIENT_ID': 'oauth.client_id',
|
||||
'OAUTH_CLIENT_SECRET': 'oauth.client_secret',
|
||||
'OPENID_REDIRECT_URI': 'oauth.redirect_uri',
|
||||
'OAUTH_SCOPES': 'oauth.scopes',
|
||||
'OAUTH_CODE_CHALLENGE_METHOD': 'oauth.code_challenge_method',
|
||||
'OAUTH_TOKEN_ENDPOINT_AUTH_METHOD': 'oauth.token_endpoint_auth_method',
|
||||
'OPENID_END_SESSION_ENDPOINT': 'oauth.end_session_endpoint',
|
||||
'OAUTH_TIMEOUT': 'oauth.timeout',
|
||||
'OAUTH_CLIENT_TIMEOUT': 'oauth.client.timeout',
|
||||
'OAUTH_EMAIL_CLAIM': 'oauth.email_claim',
|
||||
'OAUTH_USERNAME_CLAIM': 'oauth.username_claim',
|
||||
'OAUTH_PICTURE_CLAIM': 'oauth.picture_claim',
|
||||
'OAUTH_SUB_CLAIM': 'oauth.sub_claim',
|
||||
'OAUTH_AUDIENCE': 'oauth.audience',
|
||||
'OAUTH_UPDATE_EMAIL_ON_LOGIN': 'oauth.update_email_on_login',
|
||||
'OAUTH_UPDATE_NAME_ON_LOGIN': 'oauth.update_name_on_login',
|
||||
'OAUTH_UPDATE_PICTURE_ON_LOGIN': 'oauth.update_picture_on_login',
|
||||
'OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE': 'oauth.refresh_token.include_scope',
|
||||
}
|
||||
|
||||
|
||||
def _format_oauth_form_value(field: str, value):
|
||||
if field in OAUTH_COMMA_LIST_FIELDS and isinstance(value, list):
|
||||
return ','.join(str(item) for item in value)
|
||||
return value
|
||||
|
||||
|
||||
def _parse_oauth_update_value(field: str, value):
|
||||
if field in OAUTH_COMMA_LIST_FIELDS and isinstance(value, str):
|
||||
return [item.strip() for item in value.split(',') if item.strip()]
|
||||
if field in {'OAUTH_TIMEOUT', 'OAUTH_CLIENT_TIMEOUT'} and value == '':
|
||||
return ''
|
||||
return value
|
||||
|
||||
|
||||
async def get_oauth_config_values() -> dict:
|
||||
values = await Config.get_many(*OAUTH_CONFIG_KEYS.values())
|
||||
return {
|
||||
field: _format_oauth_form_value(field, values[storage_key])
|
||||
for field, storage_key in OAUTH_CONFIG_KEYS.items()
|
||||
if storage_key in values
|
||||
}
|
||||
|
||||
|
||||
def oauth_config_updates(data: dict) -> dict:
|
||||
return {
|
||||
OAUTH_CONFIG_KEYS[field]: _parse_oauth_update_value(field, value)
|
||||
for field, value in data.items()
|
||||
if field in OAUTH_CONFIG_KEYS
|
||||
}
|
||||
|
||||
|
||||
@router.get('/admin/config/oauth', response_model=OAuthConfigForm)
|
||||
async def get_oauth_config(request: Request, user=Depends(get_admin_user)):
|
||||
return await get_oauth_config_values()
|
||||
|
||||
|
||||
@router.post('/admin/config/oauth', response_model=OAuthConfigForm)
|
||||
async def update_oauth_config(request: Request, form_data: OAuthConfigForm, user=Depends(get_admin_user)):
|
||||
await Config.upsert(oauth_config_updates(form_data.model_dump(exclude_none=True)))
|
||||
return await get_oauth_config_values()
|
||||
|
||||
|
||||
async def _check_api_key_permission(request: Request, user, db: AsyncSession):
|
||||
if not request.app.state.config.ENABLE_API_KEYS or (
|
||||
if not await Config.get('auth.enable_api_keys') or (
|
||||
user.role != 'admin'
|
||||
and not await has_permission(
|
||||
user.id, 'features.api_keys', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.api_keys', await Config.get('user.permissions'), db=db
|
||||
)
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -1354,11 +1419,11 @@ async def token_exchange(
|
||||
)
|
||||
|
||||
# Extract user information from the token claims
|
||||
email_claim = request.app.state.config.OAUTH_EMAIL_CLAIM
|
||||
username_claim = request.app.state.config.OAUTH_USERNAME_CLAIM
|
||||
email_claim = await Config.get('oauth.email_claim', 'email')
|
||||
|
||||
# Get sub claim
|
||||
sub = user_data.get(request.app.state.config.OAUTH_SUB_CLAIM or OAUTH_PROVIDERS[provider].get('sub_claim', 'sub'))
|
||||
sub_claim = await Config.get('oauth.sub_claim')
|
||||
sub = user_data.get(sub_claim or OAUTH_PROVIDERS[provider].get('sub_claim', 'sub'))
|
||||
if not sub:
|
||||
log.warning(f'Token exchange failed: sub claim missing from user data')
|
||||
raise HTTPException(
|
||||
@@ -1376,10 +1441,10 @@ async def token_exchange(
|
||||
email = email.lower()
|
||||
|
||||
# Enforce domain allowlist — same check as the normal OAuth callback
|
||||
if (
|
||||
'*' not in auth_manager_config.OAUTH_ALLOWED_DOMAINS
|
||||
and email.split('@')[-1] not in auth_manager_config.OAUTH_ALLOWED_DOMAINS
|
||||
):
|
||||
oauth_allowed_domains = await Config.get('oauth.allowed_domains', [])
|
||||
if isinstance(oauth_allowed_domains, str):
|
||||
oauth_allowed_domains = [domain.strip() for domain in oauth_allowed_domains.split(',') if domain.strip()]
|
||||
if '*' not in oauth_allowed_domains and email.split('@')[-1] not in oauth_allowed_domains:
|
||||
log.warning(f'Token exchange denied: email domain not in allowed domains list')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -1389,7 +1454,7 @@ async def token_exchange(
|
||||
# Try to find the user by OAuth sub
|
||||
user = await Users.get_user_by_oauth_sub(provider, sub, db=db)
|
||||
|
||||
if not user and OAUTH_MERGE_ACCOUNTS_BY_EMAIL.value:
|
||||
if not user and await Config.get('oauth.merge_accounts_by_email'):
|
||||
# Try to find by email if merge is enabled
|
||||
user = await Users.get_user_by_email(email, db=db)
|
||||
if user:
|
||||
|
||||
@@ -14,6 +14,7 @@ from open_webui.models.automations import (
|
||||
AutomationRuns,
|
||||
Automations,
|
||||
)
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.automations import (
|
||||
@@ -38,13 +39,14 @@ PAGE_ITEM_COUNT = 30
|
||||
|
||||
|
||||
async def check_automations_permission(request, user):
|
||||
if not request.app.state.config.ENABLE_AUTOMATIONS:
|
||||
config = await Config.get_many('automations.enable', 'user.permissions')
|
||||
if not config.get('automations.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.automations', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.automations', config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -72,7 +74,7 @@ async def check_automation_limits(request, user, rrule_str: str, db, is_create:
|
||||
|
||||
# Max count (create only)
|
||||
if is_create:
|
||||
max_count = request.app.state.config.AUTOMATION_MAX_COUNT
|
||||
max_count = await Config.get('automations.max_count')
|
||||
if max_count:
|
||||
max_count = int(max_count)
|
||||
if max_count > 0 and await Automations.count_by_user(user.id, db=db) >= max_count:
|
||||
@@ -82,7 +84,7 @@ async def check_automation_limits(request, user, rrule_str: str, db, is_create:
|
||||
)
|
||||
|
||||
# Min interval (create + update)
|
||||
min_interval = request.app.state.config.AUTOMATION_MIN_INTERVAL
|
||||
min_interval = await Config.get('automations.min_interval')
|
||||
if min_interval:
|
||||
min_interval = int(min_interval)
|
||||
if min_interval > 0:
|
||||
|
||||
@@ -19,6 +19,7 @@ from open_webui.models.calendar import (
|
||||
CalendarUpdateForm,
|
||||
RSVPForm,
|
||||
)
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import UserModel
|
||||
from open_webui.utils.access_control import filter_allowed_access_grants, has_permission
|
||||
@@ -34,13 +35,14 @@ SCHEDULED_TASKS_CALENDAR_ID = '__scheduled_tasks__'
|
||||
|
||||
async def check_calendar_permission(request: Request, user):
|
||||
"""Check global feature flag AND per-user permission for calendar access."""
|
||||
if not request.app.state.config.ENABLE_CALENDAR:
|
||||
config = await Config.get_many('calendar.enable', 'user.permissions')
|
||||
if not config.get('calendar.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.calendar', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.calendar', config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -50,11 +52,12 @@ async def check_calendar_permission(request: Request, user):
|
||||
|
||||
async def _user_has_automations(request: Request, user) -> bool:
|
||||
"""Check if automations feature is available to this user."""
|
||||
if not getattr(request.app.state.config, 'ENABLE_AUTOMATIONS', False):
|
||||
config = await Config.get_many('automations.enable', 'user.permissions')
|
||||
if not config.get('automations.enable', False):
|
||||
return False
|
||||
if user.role == 'admin':
|
||||
return True
|
||||
return await has_permission(user.id, 'features.automations', request.app.state.config.USER_PERMISSIONS)
|
||||
return await has_permission(user.id, 'features.automations', config.get('user.permissions'))
|
||||
|
||||
|
||||
async def _check_calendar_access(calendar_id: str, user: UserModel, permission: str = 'write') -> CalendarModel:
|
||||
@@ -116,7 +119,7 @@ async def create_calendar(request: Request, form_data: CalendarForm, user: UserM
|
||||
# could create a calendar with `principal_id='*' permission='read'|'write'`,
|
||||
# making their events readable or writable by any other verified user.
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -373,7 +376,7 @@ async def update_calendar(
|
||||
# publicly readable/writable without the corresponding sharing permission.
|
||||
if form_data.access_grants is not None:
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
@@ -11,6 +11,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import STATIC_DIR
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants, has_public_read_access_grant, has_public_write_access_grant
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.channels import (
|
||||
ChannelForm,
|
||||
ChannelModel,
|
||||
@@ -123,7 +124,7 @@ def get_channel_permitted_group_and_user_ids(
|
||||
|
||||
async def check_channels_access(request: Request, user: Optional[UserModel] = None):
|
||||
"""Dependency to ensure channels are globally enabled."""
|
||||
if not request.app.state.config.ENABLE_CHANNELS:
|
||||
if not await Config.get('channels.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Channels'),
|
||||
@@ -131,7 +132,7 @@ async def check_channels_access(request: Request, user: Optional[UserModel] = No
|
||||
|
||||
if user:
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.channels', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.channels', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -292,7 +293,7 @@ async def create_new_channel(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -663,7 +664,7 @@ async def update_channel_by_id(
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -857,8 +858,8 @@ async def get_pinned_channel_messages(
|
||||
|
||||
async def send_notification(request, channel, message, active_user_ids, db=None):
|
||||
name = request.app.state.WEBUI_NAME
|
||||
webui_url = request.app.state.config.WEBUI_URL
|
||||
enable_user_webhooks = request.app.state.config.ENABLE_USER_WEBHOOKS
|
||||
webui_url = await Config.get('webui.url')
|
||||
enable_user_webhooks = await Config.get('ui.enable_user_webhooks')
|
||||
|
||||
users = await get_channel_users_with_access(channel, 'read', db=db)
|
||||
|
||||
@@ -1009,7 +1010,7 @@ async def model_response_handler(request, channel, message, user, db=None):
|
||||
)
|
||||
|
||||
tool_ids = _resolve_model_tool_ids(request.app, model_id)
|
||||
features = _resolve_model_features(request.app, model_id)
|
||||
features = await _resolve_model_features(request.app, model_id)
|
||||
filter_ids = _resolve_model_filter_ids(request.app, model_id)
|
||||
|
||||
# Build full form_data — same shape as frontend POST.
|
||||
|
||||
@@ -12,6 +12,7 @@ from open_webui.config import ENABLE_ADMIN_CHAT_ACCESS, ENABLE_ADMIN_EXPORT
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.chats import (
|
||||
AggregateChatStats,
|
||||
ChatBody,
|
||||
@@ -46,7 +47,7 @@ router = APIRouter()
|
||||
|
||||
async def require_chat_import_permission(request: Request, user, db: AsyncSession):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.import', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'chat.import', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -412,7 +413,7 @@ async def export_chat_stats(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
# Check if the user has permission to share/export chats
|
||||
if (user.role != 'admin') and (not request.app.state.config.ENABLE_COMMUNITY_SHARING):
|
||||
if (user.role != 'admin') and (not await Config.get('ui.enable_community_sharing')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -461,7 +462,7 @@ async def export_single_chat_stats(
|
||||
Returns ChatStatsExport for the specified chat.
|
||||
"""
|
||||
# Check if the user has permission to share/export chats
|
||||
if (user.role != 'admin') and (not request.app.state.config.ENABLE_COMMUNITY_SHARING):
|
||||
if (user.role != 'admin') and (not await Config.get('ui.enable_community_sharing')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -508,7 +509,7 @@ async def delete_all_user_chats(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role == 'user' and not await has_permission(
|
||||
user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'chat.delete', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -1164,7 +1165,7 @@ async def delete_chat_by_id(
|
||||
|
||||
return result
|
||||
else:
|
||||
if not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
|
||||
if not await has_permission(user.id, 'chat.delete', await Config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -1384,7 +1385,7 @@ async def share_chat_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'chat.share', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
|
||||
@@ -1460,7 +1461,7 @@ async def update_shared_chat_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
@@ -7,7 +7,8 @@ from typing import Optional
|
||||
import aiohttp
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from mcp.shared.auth import OAuthMetadata
|
||||
from open_webui.config import BannerModel, async_save_config, get_config, save_config
|
||||
from open_webui.config import BannerModel
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
@@ -34,6 +35,44 @@ router = APIRouter()
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
CONNECTIONS_CONFIG_KEYS = {
|
||||
'ENABLE_DIRECT_CONNECTIONS': 'direct.enable',
|
||||
'ENABLE_BASE_MODELS_CACHE': 'models.base_models_cache',
|
||||
}
|
||||
CODE_EXECUTION_CONFIG_KEYS = {
|
||||
'ENABLE_CODE_EXECUTION': 'code_execution.enable',
|
||||
'CODE_EXECUTION_ENGINE': 'code_execution.engine',
|
||||
'CODE_EXECUTION_JUPYTER_URL': 'code_execution.jupyter.url',
|
||||
'CODE_EXECUTION_JUPYTER_AUTH': 'code_execution.jupyter.auth',
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_TOKEN': 'code_execution.jupyter.auth_token',
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD': 'code_execution.jupyter.auth_password',
|
||||
'CODE_EXECUTION_JUPYTER_TIMEOUT': 'code_execution.jupyter.timeout',
|
||||
'ENABLE_CODE_INTERPRETER': 'code_interpreter.enable',
|
||||
'CODE_INTERPRETER_ENGINE': 'code_interpreter.engine',
|
||||
'CODE_INTERPRETER_PROMPT_TEMPLATE': 'code_interpreter.prompt_template',
|
||||
'CODE_INTERPRETER_JUPYTER_URL': 'code_interpreter.jupyter.url',
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH': 'code_interpreter.jupyter.auth',
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN': 'code_interpreter.jupyter.auth_token',
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD': 'code_interpreter.jupyter.auth_password',
|
||||
'CODE_INTERPRETER_JUPYTER_TIMEOUT': 'code_interpreter.jupyter.timeout',
|
||||
}
|
||||
MODELS_CONFIG_KEYS = {
|
||||
'DEFAULT_MODELS': 'ui.default_models',
|
||||
'DEFAULT_PINNED_MODELS': 'ui.default_pinned_models',
|
||||
'MODEL_ORDER_LIST': 'ui.model_order_list',
|
||||
'DEFAULT_MODEL_METADATA': 'models.default_metadata',
|
||||
'DEFAULT_MODEL_PARAMS': 'models.default_params',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
############################
|
||||
# ImportConfig
|
||||
@@ -48,9 +87,8 @@ class ImportConfigForm(BaseModel):
|
||||
|
||||
@router.post('/import', response_model=dict)
|
||||
async def import_config(request: Request, form_data: ImportConfigForm, user=Depends(get_admin_user)):
|
||||
await async_save_config(form_data.config)
|
||||
request.app.state.config._sync_to_redis()
|
||||
return get_config()
|
||||
await Config.upsert(form_data.config)
|
||||
return await Config.get_all()
|
||||
|
||||
|
||||
############################
|
||||
@@ -60,7 +98,12 @@ async def import_config(request: Request, form_data: ImportConfigForm, user=Depe
|
||||
|
||||
@router.get('/export', response_model=dict)
|
||||
async def export_config(user=Depends(get_admin_user)):
|
||||
return get_config()
|
||||
return await Config.get_all()
|
||||
|
||||
|
||||
@router.get('/namespace/{namespace}', response_model=dict)
|
||||
async def get_config_namespace(namespace: str, user=Depends(get_admin_user)):
|
||||
return await Config.get_namespace(namespace)
|
||||
|
||||
|
||||
############################
|
||||
@@ -75,10 +118,7 @@ class ConnectionsConfigForm(BaseModel):
|
||||
|
||||
@router.get('/connections', response_model=ConnectionsConfigForm)
|
||||
async def get_connections_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_DIRECT_CONNECTIONS': request.app.state.config.ENABLE_DIRECT_CONNECTIONS,
|
||||
'ENABLE_BASE_MODELS_CACHE': request.app.state.config.ENABLE_BASE_MODELS_CACHE,
|
||||
}
|
||||
return await get_config_values(CONNECTIONS_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/connections', response_model=ConnectionsConfigForm)
|
||||
@@ -87,13 +127,8 @@ async def set_connections_config(
|
||||
form_data: ConnectionsConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
request.app.state.config.ENABLE_DIRECT_CONNECTIONS = form_data.ENABLE_DIRECT_CONNECTIONS
|
||||
request.app.state.config.ENABLE_BASE_MODELS_CACHE = form_data.ENABLE_BASE_MODELS_CACHE
|
||||
|
||||
return {
|
||||
'ENABLE_DIRECT_CONNECTIONS': request.app.state.config.ENABLE_DIRECT_CONNECTIONS,
|
||||
'ENABLE_BASE_MODELS_CACHE': request.app.state.config.ENABLE_BASE_MODELS_CACHE,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), CONNECTIONS_CONFIG_KEYS))
|
||||
return await get_config_values(CONNECTIONS_CONFIG_KEYS)
|
||||
|
||||
|
||||
class OAuthClientRegistrationForm(BaseModel):
|
||||
@@ -167,9 +202,7 @@ class ToolServersConfigForm(BaseModel):
|
||||
|
||||
@router.get('/tool_servers', response_model=ToolServersConfigForm)
|
||||
async def get_tool_servers_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'TOOL_SERVER_CONNECTIONS': request.app.state.config.TOOL_SERVER_CONNECTIONS,
|
||||
}
|
||||
return {'TOOL_SERVER_CONNECTIONS': await Config.get('tool_server.connections')}
|
||||
|
||||
|
||||
@router.post('/tool_servers', response_model=ToolServersConfigForm)
|
||||
@@ -178,7 +211,8 @@ async def set_tool_servers_config(
|
||||
form_data: ToolServersConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
for connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
existing_connections = await Config.get('tool_server.connections', []) or []
|
||||
for connection in existing_connections:
|
||||
server_type = connection.get('type', 'openapi')
|
||||
auth_type = connection.get('auth_type', 'none')
|
||||
|
||||
@@ -193,13 +227,12 @@ async def set_tool_servers_config(
|
||||
pass
|
||||
|
||||
# Set new tool server connections
|
||||
request.app.state.config.TOOL_SERVER_CONNECTIONS = [
|
||||
connection.model_dump() for connection in form_data.TOOL_SERVER_CONNECTIONS
|
||||
]
|
||||
connections = [connection.model_dump() for connection in form_data.TOOL_SERVER_CONNECTIONS]
|
||||
await Config.upsert({'tool_server.connections': connections})
|
||||
|
||||
await set_tool_servers(request)
|
||||
|
||||
for connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
for connection in connections:
|
||||
server_type = connection.get('type', 'openapi')
|
||||
if server_type == 'mcp':
|
||||
server_id = connection.get('info', {}).get('id')
|
||||
@@ -216,9 +249,7 @@ async def set_tool_servers_config(
|
||||
log.debug(f'Failed to add OAuth client for MCP tool server: {e}')
|
||||
continue
|
||||
|
||||
return {
|
||||
'TOOL_SERVER_CONNECTIONS': request.app.state.config.TOOL_SERVER_CONNECTIONS,
|
||||
}
|
||||
return {'TOOL_SERVER_CONNECTIONS': connections}
|
||||
|
||||
|
||||
class TerminalServerConnection(BaseModel):
|
||||
@@ -249,9 +280,7 @@ class TerminalServersConfigForm(BaseModel):
|
||||
|
||||
@router.get('/terminal_servers')
|
||||
async def get_terminal_servers_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'TERMINAL_SERVER_CONNECTIONS': request.app.state.config.TERMINAL_SERVER_CONNECTIONS,
|
||||
}
|
||||
return {'TERMINAL_SERVER_CONNECTIONS': await Config.get('terminal_server.connections')}
|
||||
|
||||
|
||||
@router.post('/terminal_servers')
|
||||
@@ -260,15 +289,12 @@ async def set_terminal_servers_config(
|
||||
form_data: TerminalServersConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
request.app.state.config.TERMINAL_SERVER_CONNECTIONS = [
|
||||
connection.model_dump() for connection in form_data.TERMINAL_SERVER_CONNECTIONS
|
||||
]
|
||||
connections = [connection.model_dump() for connection in form_data.TERMINAL_SERVER_CONNECTIONS]
|
||||
await Config.upsert({'terminal_server.connections': connections})
|
||||
|
||||
await set_terminal_servers(request)
|
||||
|
||||
return {
|
||||
'TERMINAL_SERVER_CONNECTIONS': request.app.state.config.TERMINAL_SERVER_CONNECTIONS,
|
||||
}
|
||||
return {'TERMINAL_SERVER_CONNECTIONS': connections}
|
||||
|
||||
|
||||
@router.post('/terminal_servers/verify')
|
||||
@@ -518,67 +544,15 @@ class CodeInterpreterConfigForm(BaseModel):
|
||||
|
||||
@router.get('/code_execution', response_model=CodeInterpreterConfigForm)
|
||||
async def get_code_execution_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_CODE_EXECUTION': request.app.state.config.ENABLE_CODE_EXECUTION,
|
||||
'CODE_EXECUTION_ENGINE': request.app.state.config.CODE_EXECUTION_ENGINE,
|
||||
'CODE_EXECUTION_JUPYTER_URL': request.app.state.config.CODE_EXECUTION_JUPYTER_URL,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_EXECUTION_JUPYTER_TIMEOUT': request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT,
|
||||
'ENABLE_CODE_INTERPRETER': request.app.state.config.ENABLE_CODE_INTERPRETER,
|
||||
'CODE_INTERPRETER_ENGINE': request.app.state.config.CODE_INTERPRETER_ENGINE,
|
||||
'CODE_INTERPRETER_PROMPT_TEMPLATE': request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE,
|
||||
'CODE_INTERPRETER_JUPYTER_URL': request.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_INTERPRETER_JUPYTER_TIMEOUT': request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
}
|
||||
return await get_config_values(CODE_EXECUTION_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/code_execution', response_model=CodeInterpreterConfigForm)
|
||||
async def set_code_execution_config(
|
||||
request: Request, form_data: CodeInterpreterConfigForm, user=Depends(get_admin_user)
|
||||
):
|
||||
request.app.state.config.ENABLE_CODE_EXECUTION = form_data.ENABLE_CODE_EXECUTION
|
||||
|
||||
request.app.state.config.CODE_EXECUTION_ENGINE = form_data.CODE_EXECUTION_ENGINE
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_URL = form_data.CODE_EXECUTION_JUPYTER_URL
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH = form_data.CODE_EXECUTION_JUPYTER_AUTH
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN = form_data.CODE_EXECUTION_JUPYTER_AUTH_TOKEN
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD = form_data.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT = form_data.CODE_EXECUTION_JUPYTER_TIMEOUT
|
||||
|
||||
request.app.state.config.ENABLE_CODE_INTERPRETER = form_data.ENABLE_CODE_INTERPRETER
|
||||
request.app.state.config.CODE_INTERPRETER_ENGINE = form_data.CODE_INTERPRETER_ENGINE
|
||||
request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE = form_data.CODE_INTERPRETER_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_URL = form_data.CODE_INTERPRETER_JUPYTER_URL
|
||||
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH = form_data.CODE_INTERPRETER_JUPYTER_AUTH
|
||||
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN = form_data.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD = form_data.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT = form_data.CODE_INTERPRETER_JUPYTER_TIMEOUT
|
||||
|
||||
return {
|
||||
'ENABLE_CODE_EXECUTION': request.app.state.config.ENABLE_CODE_EXECUTION,
|
||||
'CODE_EXECUTION_ENGINE': request.app.state.config.CODE_EXECUTION_ENGINE,
|
||||
'CODE_EXECUTION_JUPYTER_URL': request.app.state.config.CODE_EXECUTION_JUPYTER_URL,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_EXECUTION_JUPYTER_TIMEOUT': request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT,
|
||||
'ENABLE_CODE_INTERPRETER': request.app.state.config.ENABLE_CODE_INTERPRETER,
|
||||
'CODE_INTERPRETER_ENGINE': request.app.state.config.CODE_INTERPRETER_ENGINE,
|
||||
'CODE_INTERPRETER_PROMPT_TEMPLATE': request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE,
|
||||
'CODE_INTERPRETER_JUPYTER_URL': request.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_INTERPRETER_JUPYTER_TIMEOUT': request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), CODE_EXECUTION_CONFIG_KEYS))
|
||||
return await get_config_values(CODE_EXECUTION_CONFIG_KEYS)
|
||||
|
||||
|
||||
############################
|
||||
@@ -595,35 +569,19 @@ class ModelsConfigForm(BaseModel):
|
||||
@router.get('/models/defaults')
|
||||
async def get_models_defaults(request: Request, user=Depends(get_verified_user)):
|
||||
return {
|
||||
'DEFAULT_MODEL_METADATA': request.app.state.config.DEFAULT_MODEL_METADATA,
|
||||
'DEFAULT_MODEL_METADATA': await Config.get('models.default_metadata'),
|
||||
}
|
||||
|
||||
|
||||
@router.get('/models', response_model=ModelsConfigForm)
|
||||
async def get_models_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'DEFAULT_MODELS': request.app.state.config.DEFAULT_MODELS,
|
||||
'DEFAULT_PINNED_MODELS': request.app.state.config.DEFAULT_PINNED_MODELS,
|
||||
'MODEL_ORDER_LIST': request.app.state.config.MODEL_ORDER_LIST,
|
||||
'DEFAULT_MODEL_METADATA': request.app.state.config.DEFAULT_MODEL_METADATA,
|
||||
'DEFAULT_MODEL_PARAMS': request.app.state.config.DEFAULT_MODEL_PARAMS,
|
||||
}
|
||||
return await get_config_values(MODELS_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/models', response_model=ModelsConfigForm)
|
||||
async def set_models_config(request: Request, form_data: ModelsConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.DEFAULT_MODELS = form_data.DEFAULT_MODELS
|
||||
request.app.state.config.DEFAULT_PINNED_MODELS = form_data.DEFAULT_PINNED_MODELS
|
||||
request.app.state.config.MODEL_ORDER_LIST = form_data.MODEL_ORDER_LIST
|
||||
request.app.state.config.DEFAULT_MODEL_METADATA = form_data.DEFAULT_MODEL_METADATA
|
||||
request.app.state.config.DEFAULT_MODEL_PARAMS = form_data.DEFAULT_MODEL_PARAMS
|
||||
return {
|
||||
'DEFAULT_MODELS': request.app.state.config.DEFAULT_MODELS,
|
||||
'DEFAULT_PINNED_MODELS': request.app.state.config.DEFAULT_PINNED_MODELS,
|
||||
'MODEL_ORDER_LIST': request.app.state.config.MODEL_ORDER_LIST,
|
||||
'DEFAULT_MODEL_METADATA': request.app.state.config.DEFAULT_MODEL_METADATA,
|
||||
'DEFAULT_MODEL_PARAMS': request.app.state.config.DEFAULT_MODEL_PARAMS,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), MODELS_CONFIG_KEYS))
|
||||
return await get_config_values(MODELS_CONFIG_KEYS)
|
||||
|
||||
|
||||
class PromptSuggestion(BaseModel):
|
||||
@@ -642,8 +600,8 @@ async def set_default_suggestions(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
data = form_data.model_dump()
|
||||
request.app.state.config.DEFAULT_PROMPT_SUGGESTIONS = data['suggestions']
|
||||
return request.app.state.config.DEFAULT_PROMPT_SUGGESTIONS
|
||||
await Config.upsert({'ui.prompt_suggestions': data['suggestions']})
|
||||
return await Config.get('ui.prompt_suggestions')
|
||||
|
||||
|
||||
############################
|
||||
@@ -662,8 +620,8 @@ async def set_banners(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
data = form_data.model_dump()
|
||||
request.app.state.config.BANNERS = data['banners']
|
||||
return request.app.state.config.BANNERS
|
||||
await Config.upsert({'ui.banners': data['banners']})
|
||||
return await Config.get('ui.banners')
|
||||
|
||||
|
||||
@router.get('/banners', response_model=list[BannerModel])
|
||||
@@ -671,4 +629,4 @@ async def get_banners(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
return request.app.state.config.BANNERS
|
||||
return await Config.get('ui.banners')
|
||||
|
||||
@@ -5,6 +5,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.concurrency import run_in_threadpool
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.feedbacks import (
|
||||
FeedbackForm,
|
||||
FeedbackIdResponse,
|
||||
@@ -25,6 +26,16 @@ log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
EVALUATION_CONFIG_KEYS = {
|
||||
'ENABLE_EVALUATION_ARENA_MODELS': 'evaluation.arena.enable',
|
||||
'EVALUATION_ARENA_MODELS': 'evaluation.arena.models',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
# Leaderboard Elo Rating Computation
|
||||
# The judgment has already been rendered with grace;
|
||||
@@ -255,10 +266,7 @@ async def get_model_history(
|
||||
|
||||
@router.get('/config')
|
||||
async def get_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_EVALUATION_ARENA_MODELS': request.app.state.config.ENABLE_EVALUATION_ARENA_MODELS,
|
||||
'EVALUATION_ARENA_MODELS': request.app.state.config.EVALUATION_ARENA_MODELS,
|
||||
}
|
||||
return await get_config_values(EVALUATION_CONFIG_KEYS)
|
||||
|
||||
|
||||
############################
|
||||
@@ -277,15 +285,13 @@ async def update_config(
|
||||
form_data: UpdateConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
config = request.app.state.config
|
||||
updates = {}
|
||||
if form_data.ENABLE_EVALUATION_ARENA_MODELS is not None:
|
||||
config.ENABLE_EVALUATION_ARENA_MODELS = form_data.ENABLE_EVALUATION_ARENA_MODELS
|
||||
updates['evaluation.arena.enable'] = form_data.ENABLE_EVALUATION_ARENA_MODELS
|
||||
if form_data.EVALUATION_ARENA_MODELS is not None:
|
||||
config.EVALUATION_ARENA_MODELS = form_data.EVALUATION_ARENA_MODELS
|
||||
return {
|
||||
'ENABLE_EVALUATION_ARENA_MODELS': config.ENABLE_EVALUATION_ARENA_MODELS,
|
||||
'EVALUATION_ARENA_MODELS': config.EVALUATION_ARENA_MODELS,
|
||||
}
|
||||
updates['evaluation.arena.models'] = form_data.EVALUATION_ARENA_MODELS
|
||||
await Config.upsert(updates)
|
||||
return await get_config_values(EVALUATION_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.get('/feedbacks/models', response_model=list[str])
|
||||
|
||||
@@ -26,6 +26,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_db_context, get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.channels import Channels
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.files import (
|
||||
FileForm,
|
||||
@@ -123,7 +124,7 @@ async def process_uploaded_file(
|
||||
if _is_text_file(file_path):
|
||||
content_type = 'text/plain'
|
||||
|
||||
stt_supported = getattr(request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', [])
|
||||
stt_supported = await Config.get('audio.stt.supported_content_types', [])
|
||||
|
||||
if content_type and strict_match_mime_type(stt_supported, content_type):
|
||||
# Audio / STT-supported files → transcribe then index
|
||||
@@ -144,7 +145,7 @@ async def process_uploaded_file(
|
||||
elif (
|
||||
content_type
|
||||
and content_type.startswith(('image/', 'video/'))
|
||||
and request.app.state.config.CONTENT_EXTRACTION_ENGINE != 'external'
|
||||
and await Config.get('rag.content_extraction_engine') != 'external'
|
||||
):
|
||||
# Media files without an external extraction engine
|
||||
if content_type.startswith('video/'):
|
||||
@@ -288,12 +289,11 @@ async def upload_file_handler(
|
||||
# Remove the leading dot from the file extension and lowercase it
|
||||
file_extension = file_extension[1:].lower() if file_extension else ''
|
||||
|
||||
if process and request.app.state.config.ALLOWED_FILE_EXTENSIONS:
|
||||
request.app.state.config.ALLOWED_FILE_EXTENSIONS = [
|
||||
ext for ext in request.app.state.config.ALLOWED_FILE_EXTENSIONS if ext
|
||||
]
|
||||
allowed_file_extensions = await Config.get('rag.file.allowed_extensions')
|
||||
if process and allowed_file_extensions:
|
||||
allowed_file_extensions = [ext for ext in allowed_file_extensions if ext]
|
||||
|
||||
if file_extension not in request.app.state.config.ALLOWED_FILE_EXTENSIONS:
|
||||
if file_extension not in allowed_file_extensions:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(f'File type {file_extension} is not allowed'),
|
||||
|
||||
@@ -11,6 +11,7 @@ from fastapi.responses import FileResponse, StreamingResponse
|
||||
from open_webui.config import UPLOAD_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.folders import (
|
||||
FolderForm,
|
||||
@@ -42,7 +43,8 @@ from open_webui.utils.access_control.folders import has_folder_access as _has_fo
|
||||
|
||||
async def check_folders_permission(request: Request, user, db=None):
|
||||
"""Verify the folders feature is enabled and the user has permission."""
|
||||
if request.app.state.config.ENABLE_FOLDERS is False:
|
||||
config = await Config.get_many('folders.enable', 'user.permissions')
|
||||
if config.get('folders.enable') is False:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -50,7 +52,7 @@ async def check_folders_permission(request: Request, user, db=None):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'features.folders',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -411,7 +413,7 @@ async def update_folder_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id, user.role,
|
||||
form_data.access_grants,
|
||||
None,
|
||||
@@ -522,7 +524,7 @@ async def delete_folder_by_id(
|
||||
folder_ids = await Folders.get_folder_ids_by_id_and_user_id_in_subtree(id, folder_owner_id, db=db)
|
||||
if await Chats.count_chats_by_folder_ids_and_user_id(folder_ids, folder_owner_id, db=db):
|
||||
chat_delete_permission = await has_permission(
|
||||
user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'chat.delete', await Config.get('user.permissions'), db=db
|
||||
)
|
||||
if user.role != 'admin' and not chat_delete_permission:
|
||||
raise HTTPException(
|
||||
|
||||
@@ -9,6 +9,7 @@ import mimetypes
|
||||
import re
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Optional
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
@@ -24,6 +25,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_ALLOW_REDIRECTS, AIOHTTP_CLIENT_SESSION_SSL, ENABLE_FORWARD_USER_INFO_HEADERS
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.retrieval.web.utils import get_ssrf_safe_session, validate_url
|
||||
from open_webui.routers.files import get_file_content_by_id, upload_file_handler
|
||||
from open_webui.utils.access_control import has_permission
|
||||
@@ -50,17 +52,68 @@ IMAGE_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
IMAGE_CONFIG_KEYS = {
|
||||
'ENABLE_IMAGE_GENERATION': 'image_generation.enable',
|
||||
'ENABLE_IMAGE_PROMPT_GENERATION': 'image_generation.prompt.enable',
|
||||
'IMAGE_GENERATION_ENGINE': 'image_generation.engine',
|
||||
'IMAGE_GENERATION_MODEL': 'image_generation.model',
|
||||
'IMAGE_SIZE': 'image_generation.size',
|
||||
'IMAGE_STEPS': 'image_generation.steps',
|
||||
'IMAGES_OPENAI_API_BASE_URL': 'image_generation.openai.api_base_url',
|
||||
'IMAGES_OPENAI_API_KEY': 'image_generation.openai.api_key',
|
||||
'IMAGES_OPENAI_API_VERSION': 'image_generation.openai.api_version',
|
||||
'IMAGES_OPENAI_API_PARAMS': 'image_generation.openai.params',
|
||||
'AUTOMATIC1111_BASE_URL': 'image_generation.automatic1111.base_url',
|
||||
'AUTOMATIC1111_API_AUTH': 'image_generation.automatic1111.api_auth',
|
||||
'AUTOMATIC1111_PARAMS': 'image_generation.automatic1111.api_params',
|
||||
'COMFYUI_BASE_URL': 'image_generation.comfyui.base_url',
|
||||
'COMFYUI_API_KEY': 'image_generation.comfyui.api_key',
|
||||
'COMFYUI_WORKFLOW': 'image_generation.comfyui.workflow',
|
||||
'COMFYUI_WORKFLOW_NODES': 'image_generation.comfyui.nodes',
|
||||
'IMAGES_GEMINI_API_BASE_URL': 'image_generation.gemini.api_base_url',
|
||||
'IMAGES_GEMINI_API_KEY': 'image_generation.gemini.api_key',
|
||||
'IMAGES_GEMINI_ENDPOINT_METHOD': 'image_generation.gemini.endpoint_method',
|
||||
'ENABLE_IMAGE_EDIT': 'images.edit.enable',
|
||||
'IMAGE_EDIT_ENGINE': 'images.edit.engine',
|
||||
'IMAGE_EDIT_MODEL': 'images.edit.model',
|
||||
'IMAGE_EDIT_SIZE': 'images.edit.size',
|
||||
'IMAGES_EDIT_OPENAI_API_BASE_URL': 'images.edit.openai.api_base_url',
|
||||
'IMAGES_EDIT_OPENAI_API_KEY': 'images.edit.openai.api_key',
|
||||
'IMAGES_EDIT_OPENAI_API_VERSION': 'images.edit.openai.api_version',
|
||||
'IMAGES_EDIT_GEMINI_API_BASE_URL': 'images.edit.gemini.api_base_url',
|
||||
'IMAGES_EDIT_GEMINI_API_KEY': 'images.edit.gemini.api_key',
|
||||
'IMAGES_EDIT_COMFYUI_BASE_URL': 'images.edit.comfyui.base_url',
|
||||
'IMAGES_EDIT_COMFYUI_API_KEY': 'images.edit.comfyui.api_key',
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW': 'images.edit.comfyui.workflow',
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES': 'images.edit.comfyui.nodes',
|
||||
'USER_PERMISSIONS': 'user.permissions',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
async def get_image_config() -> SimpleNamespace:
|
||||
return SimpleNamespace(**await get_config_values(IMAGE_CONFIG_KEYS))
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
async def set_image_model(request: Request, model: str):
|
||||
log.info(f'Setting image model to {model}')
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL = model
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE in ['', 'automatic1111']:
|
||||
api_auth = get_automatic1111_api_auth(request)
|
||||
await Config.upsert({'image_generation.model': model})
|
||||
image_config = await get_image_config()
|
||||
if image_config.IMAGE_GENERATION_ENGINE in ['', 'automatic1111']:
|
||||
api_auth = get_automatic1111_api_auth(image_config)
|
||||
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': api_auth},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
@@ -68,7 +121,7 @@ async def set_image_model(request: Request, model: str):
|
||||
if model != options['sd_model_checkpoint']:
|
||||
options['sd_model_checkpoint'] = model
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
json=options,
|
||||
headers={'authorization': api_auth},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -77,35 +130,36 @@ async def set_image_model(request: Request, model: str):
|
||||
except Exception as e:
|
||||
log.debug(f'{e}')
|
||||
|
||||
return request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
return image_config.IMAGE_GENERATION_MODEL
|
||||
|
||||
|
||||
async def get_image_model(request):
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
image_config = await get_image_config()
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
return (
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
if request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
image_config.IMAGE_GENERATION_MODEL
|
||||
if image_config.IMAGE_GENERATION_MODEL
|
||||
else 'dall-e-2'
|
||||
)
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
return (
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
if request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
image_config.IMAGE_GENERATION_MODEL
|
||||
if image_config.IMAGE_GENERATION_MODEL
|
||||
else 'imagen-3.0-generate-002'
|
||||
)
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
return (
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL if request.app.state.config.IMAGE_GENERATION_MODEL else ''
|
||||
image_config.IMAGE_GENERATION_MODEL if image_config.IMAGE_GENERATION_MODEL else ''
|
||||
)
|
||||
elif (
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
|
||||
image_config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or image_config.IMAGE_GENERATION_ENGINE == ''
|
||||
):
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
options = await r.json()
|
||||
@@ -159,52 +213,11 @@ class ImagesConfig(BaseModel):
|
||||
|
||||
@router.get('/config', response_model=ImagesConfig)
|
||||
async def get_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_IMAGE_GENERATION': request.app.state.config.ENABLE_IMAGE_GENERATION,
|
||||
'ENABLE_IMAGE_PROMPT_GENERATION': request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION,
|
||||
'IMAGE_GENERATION_ENGINE': request.app.state.config.IMAGE_GENERATION_ENGINE,
|
||||
'IMAGE_GENERATION_MODEL': request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
'IMAGE_SIZE': request.app.state.config.IMAGE_SIZE,
|
||||
'IMAGE_STEPS': request.app.state.config.IMAGE_STEPS,
|
||||
'IMAGES_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_OPENAI_API_BASE_URL,
|
||||
'IMAGES_OPENAI_API_KEY': request.app.state.config.IMAGES_OPENAI_API_KEY,
|
||||
'IMAGES_OPENAI_API_VERSION': request.app.state.config.IMAGES_OPENAI_API_VERSION,
|
||||
'IMAGES_OPENAI_API_PARAMS': request.app.state.config.IMAGES_OPENAI_API_PARAMS,
|
||||
'AUTOMATIC1111_BASE_URL': request.app.state.config.AUTOMATIC1111_BASE_URL,
|
||||
'AUTOMATIC1111_API_AUTH': request.app.state.config.AUTOMATIC1111_API_AUTH,
|
||||
'AUTOMATIC1111_PARAMS': request.app.state.config.AUTOMATIC1111_PARAMS,
|
||||
'COMFYUI_BASE_URL': request.app.state.config.COMFYUI_BASE_URL,
|
||||
'COMFYUI_API_KEY': request.app.state.config.COMFYUI_API_KEY,
|
||||
'COMFYUI_WORKFLOW': request.app.state.config.COMFYUI_WORKFLOW,
|
||||
'COMFYUI_WORKFLOW_NODES': request.app.state.config.COMFYUI_WORKFLOW_NODES,
|
||||
'IMAGES_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_GEMINI_API_BASE_URL,
|
||||
'IMAGES_GEMINI_API_KEY': request.app.state.config.IMAGES_GEMINI_API_KEY,
|
||||
'IMAGES_GEMINI_ENDPOINT_METHOD': request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD,
|
||||
'ENABLE_IMAGE_EDIT': request.app.state.config.ENABLE_IMAGE_EDIT,
|
||||
'IMAGE_EDIT_ENGINE': request.app.state.config.IMAGE_EDIT_ENGINE,
|
||||
'IMAGE_EDIT_MODEL': request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
'IMAGE_EDIT_SIZE': request.app.state.config.IMAGE_EDIT_SIZE,
|
||||
'IMAGES_EDIT_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL,
|
||||
'IMAGES_EDIT_OPENAI_API_KEY': request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY,
|
||||
'IMAGES_EDIT_OPENAI_API_VERSION': request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION,
|
||||
'IMAGES_EDIT_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL,
|
||||
'IMAGES_EDIT_GEMINI_API_KEY': request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_BASE_URL': request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
'IMAGES_EDIT_COMFYUI_API_KEY': request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
return await get_config_values(IMAGE_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_config(request: Request, form_data: ImagesConfig, user=Depends(get_admin_user)):
|
||||
request.app.state.config.ENABLE_IMAGE_GENERATION = form_data.ENABLE_IMAGE_GENERATION
|
||||
|
||||
# Create Image
|
||||
request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION = form_data.ENABLE_IMAGE_PROMPT_GENERATION
|
||||
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE = form_data.IMAGE_GENERATION_ENGINE
|
||||
await set_image_model(request, form_data.IMAGE_GENERATION_MODEL)
|
||||
if form_data.IMAGE_SIZE == 'auto' and not re.match(
|
||||
IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN, form_data.IMAGE_GENERATION_MODEL
|
||||
):
|
||||
@@ -216,100 +229,31 @@ async def update_config(request: Request, form_data: ImagesConfig, user=Depends(
|
||||
)
|
||||
|
||||
pattern = r'^\d+x\d+$'
|
||||
if form_data.IMAGE_SIZE == 'auto' or form_data.IMAGE_SIZE == '' or re.match(pattern, form_data.IMAGE_SIZE):
|
||||
request.app.state.config.IMAGE_SIZE = form_data.IMAGE_SIZE
|
||||
else:
|
||||
if not (form_data.IMAGE_SIZE == 'auto' or form_data.IMAGE_SIZE == '' or re.match(pattern, form_data.IMAGE_SIZE)):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.INCORRECT_FORMAT(' (e.g., 512x512).'),
|
||||
)
|
||||
|
||||
if form_data.IMAGE_STEPS >= 0:
|
||||
request.app.state.config.IMAGE_STEPS = form_data.IMAGE_STEPS
|
||||
else:
|
||||
if form_data.IMAGE_STEPS < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.INCORRECT_FORMAT(' (e.g., 50).'),
|
||||
)
|
||||
|
||||
request.app.state.config.IMAGES_OPENAI_API_BASE_URL = form_data.IMAGES_OPENAI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_OPENAI_API_KEY = form_data.IMAGES_OPENAI_API_KEY
|
||||
request.app.state.config.IMAGES_OPENAI_API_VERSION = form_data.IMAGES_OPENAI_API_VERSION
|
||||
request.app.state.config.IMAGES_OPENAI_API_PARAMS = form_data.IMAGES_OPENAI_API_PARAMS
|
||||
|
||||
request.app.state.config.AUTOMATIC1111_BASE_URL = form_data.AUTOMATIC1111_BASE_URL
|
||||
request.app.state.config.AUTOMATIC1111_API_AUTH = form_data.AUTOMATIC1111_API_AUTH
|
||||
request.app.state.config.AUTOMATIC1111_PARAMS = form_data.AUTOMATIC1111_PARAMS
|
||||
|
||||
request.app.state.config.COMFYUI_BASE_URL = form_data.COMFYUI_BASE_URL.strip('/')
|
||||
request.app.state.config.COMFYUI_API_KEY = form_data.COMFYUI_API_KEY
|
||||
request.app.state.config.COMFYUI_WORKFLOW = form_data.COMFYUI_WORKFLOW
|
||||
request.app.state.config.COMFYUI_WORKFLOW_NODES = form_data.COMFYUI_WORKFLOW_NODES
|
||||
|
||||
request.app.state.config.IMAGES_GEMINI_API_BASE_URL = form_data.IMAGES_GEMINI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_GEMINI_API_KEY = form_data.IMAGES_GEMINI_API_KEY
|
||||
request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD = form_data.IMAGES_GEMINI_ENDPOINT_METHOD
|
||||
|
||||
# Edit Image
|
||||
request.app.state.config.ENABLE_IMAGE_EDIT = form_data.ENABLE_IMAGE_EDIT
|
||||
request.app.state.config.IMAGE_EDIT_ENGINE = form_data.IMAGE_EDIT_ENGINE
|
||||
request.app.state.config.IMAGE_EDIT_MODEL = form_data.IMAGE_EDIT_MODEL
|
||||
request.app.state.config.IMAGE_EDIT_SIZE = form_data.IMAGE_EDIT_SIZE
|
||||
|
||||
request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL = form_data.IMAGES_EDIT_OPENAI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY = form_data.IMAGES_EDIT_OPENAI_API_KEY
|
||||
request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION = form_data.IMAGES_EDIT_OPENAI_API_VERSION
|
||||
|
||||
request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL = form_data.IMAGES_EDIT_GEMINI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY = form_data.IMAGES_EDIT_GEMINI_API_KEY
|
||||
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL = form_data.IMAGES_EDIT_COMFYUI_BASE_URL.strip('/')
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY = form_data.IMAGES_EDIT_COMFYUI_API_KEY
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW = form_data.IMAGES_EDIT_COMFYUI_WORKFLOW
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES = form_data.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES
|
||||
|
||||
return {
|
||||
'ENABLE_IMAGE_GENERATION': request.app.state.config.ENABLE_IMAGE_GENERATION,
|
||||
'ENABLE_IMAGE_PROMPT_GENERATION': request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION,
|
||||
'IMAGE_GENERATION_ENGINE': request.app.state.config.IMAGE_GENERATION_ENGINE,
|
||||
'IMAGE_GENERATION_MODEL': request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
'IMAGE_SIZE': request.app.state.config.IMAGE_SIZE,
|
||||
'IMAGE_STEPS': request.app.state.config.IMAGE_STEPS,
|
||||
'IMAGES_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_OPENAI_API_BASE_URL,
|
||||
'IMAGES_OPENAI_API_KEY': request.app.state.config.IMAGES_OPENAI_API_KEY,
|
||||
'IMAGES_OPENAI_API_VERSION': request.app.state.config.IMAGES_OPENAI_API_VERSION,
|
||||
'IMAGES_OPENAI_API_PARAMS': request.app.state.config.IMAGES_OPENAI_API_PARAMS,
|
||||
'AUTOMATIC1111_BASE_URL': request.app.state.config.AUTOMATIC1111_BASE_URL,
|
||||
'AUTOMATIC1111_API_AUTH': request.app.state.config.AUTOMATIC1111_API_AUTH,
|
||||
'AUTOMATIC1111_PARAMS': request.app.state.config.AUTOMATIC1111_PARAMS,
|
||||
'COMFYUI_BASE_URL': request.app.state.config.COMFYUI_BASE_URL,
|
||||
'COMFYUI_API_KEY': request.app.state.config.COMFYUI_API_KEY,
|
||||
'COMFYUI_WORKFLOW': request.app.state.config.COMFYUI_WORKFLOW,
|
||||
'COMFYUI_WORKFLOW_NODES': request.app.state.config.COMFYUI_WORKFLOW_NODES,
|
||||
'IMAGES_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_GEMINI_API_BASE_URL,
|
||||
'IMAGES_GEMINI_API_KEY': request.app.state.config.IMAGES_GEMINI_API_KEY,
|
||||
'IMAGES_GEMINI_ENDPOINT_METHOD': request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD,
|
||||
'ENABLE_IMAGE_EDIT': request.app.state.config.ENABLE_IMAGE_EDIT,
|
||||
'IMAGE_EDIT_ENGINE': request.app.state.config.IMAGE_EDIT_ENGINE,
|
||||
'IMAGE_EDIT_MODEL': request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
'IMAGE_EDIT_SIZE': request.app.state.config.IMAGE_EDIT_SIZE,
|
||||
'IMAGES_EDIT_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL,
|
||||
'IMAGES_EDIT_OPENAI_API_KEY': request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY,
|
||||
'IMAGES_EDIT_OPENAI_API_VERSION': request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION,
|
||||
'IMAGES_EDIT_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL,
|
||||
'IMAGES_EDIT_GEMINI_API_KEY': request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_BASE_URL': request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
'IMAGES_EDIT_COMFYUI_API_KEY': request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
updates = config_updates(form_data.model_dump(), IMAGE_CONFIG_KEYS)
|
||||
updates['image_generation.comfyui.base_url'] = form_data.COMFYUI_BASE_URL.strip('/')
|
||||
updates['images.edit.comfyui.base_url'] = form_data.IMAGES_EDIT_COMFYUI_BASE_URL.strip('/')
|
||||
await Config.upsert(updates)
|
||||
await set_image_model(request, form_data.IMAGE_GENERATION_MODEL)
|
||||
return await get_config_values(IMAGE_CONFIG_KEYS)
|
||||
|
||||
|
||||
def get_automatic1111_api_auth(request: Request):
|
||||
if request.app.state.config.AUTOMATIC1111_API_AUTH is None:
|
||||
def get_automatic1111_api_auth(image_config):
|
||||
if image_config.AUTOMATIC1111_API_AUTH is None:
|
||||
return ''
|
||||
else:
|
||||
auth1111_byte_string = request.app.state.config.AUTOMATIC1111_API_AUTH.encode('utf-8')
|
||||
auth1111_byte_string = image_config.AUTOMATIC1111_API_AUTH.encode('utf-8')
|
||||
auth1111_base64_encoded_bytes = base64.b64encode(auth1111_byte_string)
|
||||
auth1111_base64_encoded_string = auth1111_base64_encoded_bytes.decode('utf-8')
|
||||
return f'Basic {auth1111_base64_encoded_string}'
|
||||
@@ -317,26 +261,27 @@ def get_automatic1111_api_auth(request: Request):
|
||||
|
||||
@router.get('/config/url/verify')
|
||||
async def verify_url(request: Request, user=Depends(get_admin_user)):
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111':
|
||||
image_config = await get_image_config()
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'automatic1111':
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
return True
|
||||
except Exception:
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.INVALID_URL)
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
headers = None
|
||||
if request.app.state.config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
|
||||
if image_config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {image_config.COMFYUI_API_KEY}'}
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.COMFYUI_BASE_URL}/object_info',
|
||||
url=f'{image_config.COMFYUI_BASE_URL}/object_info',
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
@@ -350,33 +295,34 @@ async def verify_url(request: Request, user=Depends(get_admin_user)):
|
||||
|
||||
@router.get('/models')
|
||||
async def get_models(request: Request, user=Depends(get_verified_user)):
|
||||
image_config = await get_image_config()
|
||||
try:
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
return [
|
||||
{'id': 'dall-e-2', 'name': 'DALL·E 2'},
|
||||
{'id': 'dall-e-3', 'name': 'DALL·E 3'},
|
||||
{'id': 'gpt-image-1', 'name': 'GPT-IMAGE 1'},
|
||||
{'id': 'gpt-image-1.5', 'name': 'GPT-IMAGE 1.5'},
|
||||
]
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
return [
|
||||
{'id': 'imagen-3.0-generate-002', 'name': 'imagen-3.0 generate-002'},
|
||||
]
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
# TODO - get models from comfyui
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
|
||||
headers = {'Authorization': f'Bearer {image_config.COMFYUI_API_KEY}'}
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.COMFYUI_BASE_URL}/object_info',
|
||||
url=f'{image_config.COMFYUI_BASE_URL}/object_info',
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
info = await r.json()
|
||||
|
||||
workflow = json.loads(request.app.state.config.COMFYUI_WORKFLOW)
|
||||
workflow = json.loads(image_config.COMFYUI_WORKFLOW)
|
||||
model_node_id = None
|
||||
|
||||
for node in request.app.state.config.COMFYUI_WORKFLOW_NODES:
|
||||
for node in image_config.COMFYUI_WORKFLOW_NODES:
|
||||
if node['type'] == 'model':
|
||||
if node['node_ids']:
|
||||
model_node_id = node['node_ids'][0]
|
||||
@@ -406,13 +352,13 @@ async def get_models(request: Request, user=Depends(get_verified_user)):
|
||||
)
|
||||
)
|
||||
elif (
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
|
||||
image_config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or image_config.IMAGE_GENERATION_ENGINE == ''
|
||||
):
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/sd-models',
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/sd-models',
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
models = await r.json()
|
||||
@@ -540,14 +486,15 @@ async def upload_image(request, image_data, content_type, metadata, user, db=Non
|
||||
|
||||
@router.post('/generations')
|
||||
async def generate_images(request: Request, form_data: CreateImageForm, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_IMAGE_GENERATION:
|
||||
image_config = await get_image_config()
|
||||
if not image_config.ENABLE_IMAGE_GENERATION:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.image_generation', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.image_generation', image_config.USER_PERMISSIONS
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
@@ -563,13 +510,14 @@ async def image_generations(
|
||||
metadata: dict | None = None,
|
||||
user=None,
|
||||
):
|
||||
image_config = await get_image_config()
|
||||
# if IMAGE_SIZE = 'auto', default WidthxHeight to the 512x512 default
|
||||
# This is only relevant when the user has set IMAGE_SIZE to 'auto' with an
|
||||
# image model other than gpt-image-1, which is warned about on settings save
|
||||
|
||||
size = '512x512'
|
||||
if request.app.state.config.IMAGE_SIZE and 'x' in request.app.state.config.IMAGE_SIZE:
|
||||
size = request.app.state.config.IMAGE_SIZE
|
||||
if image_config.IMAGE_SIZE and 'x' in image_config.IMAGE_SIZE:
|
||||
size = image_config.IMAGE_SIZE
|
||||
|
||||
if form_data.size and 'x' in form_data.size:
|
||||
size = form_data.size
|
||||
@@ -581,40 +529,40 @@ async def image_generations(
|
||||
model = await get_image_model(request)
|
||||
|
||||
try:
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
headers = {
|
||||
'Authorization': f'Bearer {request.app.state.config.IMAGES_OPENAI_API_KEY}',
|
||||
'Authorization': f'Bearer {image_config.IMAGES_OPENAI_API_KEY}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
url = f'{request.app.state.config.IMAGES_OPENAI_API_BASE_URL}/images/generations'
|
||||
if request.app.state.config.IMAGES_OPENAI_API_VERSION:
|
||||
url = f'{url}?api-version={request.app.state.config.IMAGES_OPENAI_API_VERSION}'
|
||||
url = f'{image_config.IMAGES_OPENAI_API_BASE_URL}/images/generations'
|
||||
if image_config.IMAGES_OPENAI_API_VERSION:
|
||||
url = f'{url}?api-version={image_config.IMAGES_OPENAI_API_VERSION}'
|
||||
|
||||
data = {
|
||||
'model': model,
|
||||
'prompt': form_data.prompt,
|
||||
'n': form_data.n,
|
||||
**(
|
||||
{'size': form_data.size or request.app.state.config.IMAGE_SIZE}
|
||||
if (form_data.size or request.app.state.config.IMAGE_SIZE)
|
||||
{'size': form_data.size or image_config.IMAGE_SIZE}
|
||||
if (form_data.size or image_config.IMAGE_SIZE)
|
||||
else {}
|
||||
),
|
||||
**(
|
||||
{}
|
||||
if re.match(
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
image_config.IMAGE_GENERATION_MODEL,
|
||||
)
|
||||
else {'response_format': 'b64_json'}
|
||||
),
|
||||
**(
|
||||
{}
|
||||
if not request.app.state.config.IMAGES_OPENAI_API_PARAMS
|
||||
else request.app.state.config.IMAGES_OPENAI_API_PARAMS
|
||||
if not image_config.IMAGES_OPENAI_API_PARAMS
|
||||
else image_config.IMAGES_OPENAI_API_PARAMS
|
||||
),
|
||||
}
|
||||
|
||||
@@ -643,17 +591,17 @@ async def image_generations(
|
||||
images.append({'url': url})
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'x-goog-api-key': request.app.state.config.IMAGES_GEMINI_API_KEY,
|
||||
'x-goog-api-key': image_config.IMAGES_GEMINI_API_KEY,
|
||||
}
|
||||
|
||||
data = {}
|
||||
|
||||
if (
|
||||
request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD == ''
|
||||
or request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD == 'predict'
|
||||
image_config.IMAGES_GEMINI_ENDPOINT_METHOD == ''
|
||||
or image_config.IMAGES_GEMINI_ENDPOINT_METHOD == 'predict'
|
||||
):
|
||||
model = f'{model}:predict'
|
||||
data = {
|
||||
@@ -664,13 +612,13 @@ async def image_generations(
|
||||
},
|
||||
}
|
||||
|
||||
elif request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD == 'generateContent':
|
||||
elif image_config.IMAGES_GEMINI_ENDPOINT_METHOD == 'generateContent':
|
||||
model = f'{model}:generateContent'
|
||||
data = {'contents': [{'parts': [{'text': form_data.prompt}]}]}
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.IMAGES_GEMINI_API_BASE_URL}/models/{model}',
|
||||
url=f'{image_config.IMAGES_GEMINI_API_BASE_URL}/models/{model}',
|
||||
json=data,
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -701,7 +649,7 @@ async def image_generations(
|
||||
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
data = {
|
||||
'prompt': form_data.prompt,
|
||||
'width': width,
|
||||
@@ -709,8 +657,8 @@ async def image_generations(
|
||||
'n': form_data.n,
|
||||
}
|
||||
|
||||
if request.app.state.config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else request.app.state.config.IMAGE_STEPS
|
||||
if image_config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else image_config.IMAGE_STEPS
|
||||
|
||||
if form_data.negative_prompt is not None:
|
||||
data['negative_prompt'] = form_data.negative_prompt
|
||||
@@ -719,8 +667,8 @@ async def image_generations(
|
||||
**{
|
||||
'workflow': ComfyUIWorkflow(
|
||||
**{
|
||||
'workflow': request.app.state.config.COMFYUI_WORKFLOW,
|
||||
'nodes': request.app.state.config.COMFYUI_WORKFLOW_NODES,
|
||||
'workflow': image_config.COMFYUI_WORKFLOW,
|
||||
'nodes': image_config.COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
),
|
||||
**data,
|
||||
@@ -730,8 +678,8 @@ async def image_generations(
|
||||
model,
|
||||
form_data,
|
||||
str(uuid.uuid4()),
|
||||
request.app.state.config.COMFYUI_BASE_URL,
|
||||
request.app.state.config.COMFYUI_API_KEY,
|
||||
image_config.COMFYUI_BASE_URL,
|
||||
image_config.COMFYUI_API_KEY,
|
||||
)
|
||||
log.debug(f'res: {res}')
|
||||
|
||||
@@ -739,13 +687,13 @@ async def image_generations(
|
||||
|
||||
for image in res['data']:
|
||||
headers = None
|
||||
if request.app.state.config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
|
||||
if image_config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {image_config.COMFYUI_API_KEY}'}
|
||||
|
||||
image_data, content_type = await get_image_data(
|
||||
image['url'],
|
||||
headers,
|
||||
trusted_base_url=request.app.state.config.COMFYUI_BASE_URL,
|
||||
trusted_base_url=image_config.COMFYUI_BASE_URL,
|
||||
)
|
||||
_, url = await upload_image(
|
||||
request,
|
||||
@@ -757,8 +705,8 @@ async def image_generations(
|
||||
images.append({'url': url})
|
||||
return images
|
||||
elif (
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
|
||||
image_config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or image_config.IMAGE_GENERATION_ENGINE == ''
|
||||
):
|
||||
if form_data.model:
|
||||
await set_image_model(request, form_data.model)
|
||||
@@ -770,20 +718,20 @@ async def image_generations(
|
||||
'height': height,
|
||||
}
|
||||
|
||||
if request.app.state.config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else request.app.state.config.IMAGE_STEPS
|
||||
if image_config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else image_config.IMAGE_STEPS
|
||||
|
||||
if form_data.negative_prompt is not None:
|
||||
data['negative_prompt'] = form_data.negative_prompt
|
||||
|
||||
if request.app.state.config.AUTOMATIC1111_PARAMS:
|
||||
data = {**data, **request.app.state.config.AUTOMATIC1111_PARAMS}
|
||||
if image_config.AUTOMATIC1111_PARAMS:
|
||||
data = {**data, **image_config.AUTOMATIC1111_PARAMS}
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/txt2img',
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/txt2img',
|
||||
json=data,
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
res = await r.json(content_type=None)
|
||||
@@ -826,17 +774,18 @@ async def image_edits(
|
||||
metadata: dict | None = None,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
image_config = await get_image_config()
|
||||
size = None
|
||||
width, height = None, None
|
||||
metadata = metadata or {}
|
||||
|
||||
if (request.app.state.config.IMAGE_EDIT_SIZE and 'x' in request.app.state.config.IMAGE_EDIT_SIZE) or (
|
||||
if (image_config.IMAGE_EDIT_SIZE and 'x' in image_config.IMAGE_EDIT_SIZE) or (
|
||||
form_data.size and 'x' in form_data.size
|
||||
):
|
||||
size = form_data.size if form_data.size else request.app.state.config.IMAGE_EDIT_SIZE
|
||||
size = form_data.size if form_data.size else image_config.IMAGE_EDIT_SIZE
|
||||
width, height = tuple(map(int, size.split('x')))
|
||||
|
||||
model = request.app.state.config.IMAGE_EDIT_MODEL if form_data.model is None else form_data.model
|
||||
model = image_config.IMAGE_EDIT_MODEL if form_data.model is None else form_data.model
|
||||
|
||||
try:
|
||||
|
||||
@@ -905,9 +854,9 @@ async def image_edits(
|
||||
)
|
||||
|
||||
try:
|
||||
if request.app.state.config.IMAGE_EDIT_ENGINE == 'openai':
|
||||
if image_config.IMAGE_EDIT_ENGINE == 'openai':
|
||||
headers = {
|
||||
'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY}',
|
||||
'Authorization': f'Bearer {image_config.IMAGES_EDIT_OPENAI_API_KEY}',
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
@@ -923,7 +872,7 @@ async def image_edits(
|
||||
{}
|
||||
if re.match(
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
image_config.IMAGE_EDIT_MODEL,
|
||||
)
|
||||
else {'response_format': 'b64_json'}
|
||||
),
|
||||
@@ -937,8 +886,8 @@ async def image_edits(
|
||||
files.append(get_image_file_item(img, 'image[]'))
|
||||
|
||||
url_search_params = ''
|
||||
if request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION:
|
||||
url_search_params += f'?api-version={request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION}'
|
||||
if image_config.IMAGES_EDIT_OPENAI_API_VERSION:
|
||||
url_search_params += f'?api-version={image_config.IMAGES_EDIT_OPENAI_API_VERSION}'
|
||||
|
||||
# Build multipart form data for aiohttp
|
||||
form = aiohttp.FormData()
|
||||
@@ -957,7 +906,7 @@ async def image_edits(
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL}/images/edits{url_search_params}',
|
||||
url=f'{image_config.IMAGES_EDIT_OPENAI_API_BASE_URL}/images/edits{url_search_params}',
|
||||
headers=headers,
|
||||
data=form,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -979,10 +928,10 @@ async def image_edits(
|
||||
images.append({'url': url})
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_EDIT_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_EDIT_ENGINE == 'gemini':
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'x-goog-api-key': request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
'x-goog-api-key': image_config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
}
|
||||
|
||||
model = f'{model}:generateContent'
|
||||
@@ -1012,7 +961,7 @@ async def image_edits(
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL}/models/{model}',
|
||||
url=f'{image_config.IMAGES_EDIT_GEMINI_API_BASE_URL}/models/{model}',
|
||||
json=data,
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -1036,7 +985,7 @@ async def image_edits(
|
||||
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_EDIT_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_EDIT_ENGINE == 'comfyui':
|
||||
try:
|
||||
files = []
|
||||
if isinstance(form_data.image, str):
|
||||
@@ -1050,8 +999,8 @@ async def image_edits(
|
||||
for file_item in files:
|
||||
res = await comfyui_upload_image(
|
||||
file_item,
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
image_config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
image_config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
)
|
||||
comfyui_images.append(res.get('name', file_item[1][0]))
|
||||
except Exception as e:
|
||||
@@ -1070,8 +1019,8 @@ async def image_edits(
|
||||
**{
|
||||
'workflow': ComfyUIWorkflow(
|
||||
**{
|
||||
'workflow': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'nodes': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
'workflow': image_config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'nodes': image_config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
),
|
||||
**data,
|
||||
@@ -1081,8 +1030,8 @@ async def image_edits(
|
||||
model,
|
||||
form_data,
|
||||
str(uuid.uuid4()),
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
image_config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
image_config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
)
|
||||
log.debug(f'res: {res}')
|
||||
|
||||
@@ -1101,13 +1050,13 @@ async def image_edits(
|
||||
|
||||
for image_url in image_urls:
|
||||
headers = None
|
||||
if request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY}'}
|
||||
if image_config.IMAGES_EDIT_COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {image_config.IMAGES_EDIT_COMFYUI_API_KEY}'}
|
||||
|
||||
image_data, content_type = await get_image_data(
|
||||
image_url,
|
||||
headers,
|
||||
trusted_base_url=request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
trusted_base_url=image_config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
)
|
||||
_, url = await upload_image(
|
||||
request,
|
||||
|
||||
@@ -13,6 +13,7 @@ from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.files import FileMetadataResponse, FileModel, FileModelResponse, Files
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.knowledge import (
|
||||
@@ -257,7 +258,7 @@ async def create_new_knowledge(
|
||||
# This prevents holding a connection during embed_knowledge_base_metadata()
|
||||
# which makes external embedding API calls (1-5+ seconds).
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'workspace.knowledge', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'workspace.knowledge', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -265,7 +266,7 @@ async def create_new_knowledge(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -469,7 +470,7 @@ async def update_knowledge_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -537,7 +538,7 @@ async def update_knowledge_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.memories import Memories, MemoryModel
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
|
||||
@@ -20,6 +21,23 @@ log = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
async def check_memories_permission(user):
|
||||
config = await Config.get_many('memories.enable', 'user.permissions')
|
||||
if not config.get('memories.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.memories', config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# GetMemories
|
||||
# Let what is remembered here spare someone the cost
|
||||
@@ -33,17 +51,7 @@ async def get_memories(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
return await Memories.get_memories_by_user_id(user.id, db=db)
|
||||
|
||||
@@ -73,17 +81,7 @@ async def add_memory(
|
||||
own short-lived sessions so a connection is not held during the external
|
||||
embedding API call (``EMBEDDING_FUNCTION``), which can take 1-5+ seconds.
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
memory = await Memories.insert_new_memory(user.id, form_data.content)
|
||||
|
||||
@@ -124,17 +122,7 @@ async def query_memory(
|
||||
# Database operations (get_memories_by_user_id) manage their own short-lived sessions.
|
||||
# This prevents holding a connection during EMBEDDING_FUNCTION()
|
||||
# which makes external embedding API calls (1-5+ seconds).
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
if not memories:
|
||||
@@ -154,7 +142,7 @@ async def query_memory(
|
||||
# same RELEVANCE_THRESHOLD used by RAG ensures only genuinely matching
|
||||
# memories are surfaced (distances are normalised to 0→1, higher is
|
||||
# better).
|
||||
relevance_threshold = getattr(request.app.state.config, 'RELEVANCE_THRESHOLD', 0.0)
|
||||
relevance_threshold = await Config.get('rag.relevance_threshold', 0.0)
|
||||
if results and relevance_threshold > 0.0 and results.distances and results.distances[0]:
|
||||
from open_webui.retrieval.vector.main import SearchResult
|
||||
|
||||
@@ -199,17 +187,7 @@ async def reset_memory_from_vector_db(
|
||||
calls simultaneously. With a session held, this could block a connection
|
||||
for MINUTES, completely exhausting the connection pool.
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
|
||||
@@ -250,17 +228,7 @@ async def delete_memory_by_user_id(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
result = await Memories.delete_memories_by_user_id(user.id, db=db)
|
||||
|
||||
@@ -290,17 +258,7 @@ async def update_memory_by_id(
|
||||
# Database operations (update_memory_by_id_and_user_id) manage their own
|
||||
# short-lived sessions. This prevents holding a connection during
|
||||
# EMBEDDING_FUNCTION() which makes external API calls (1-5+ seconds).
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content)
|
||||
if memory is None:
|
||||
@@ -339,17 +297,7 @@ async def delete_memory_by_id(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id, db=db)
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING, PROFILE_IMAGE_ALLOWED_MIME_TYPES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import (
|
||||
ModelAccessListResponse,
|
||||
@@ -230,7 +231,7 @@ async def create_new_model(
|
||||
):
|
||||
"""Create a new workspace model entry."""
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'workspace.models', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'workspace.models', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -258,7 +259,7 @@ async def create_new_model(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -289,7 +290,7 @@ async def export_models(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.models_export',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -322,7 +323,7 @@ async def import_models(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.models_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -403,7 +404,7 @@ async def import_models(
|
||||
# metadata-only imports.
|
||||
if 'access_grants' in model_data:
|
||||
updated_model.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
updated_model.access_grants,
|
||||
@@ -416,7 +417,7 @@ async def import_models(
|
||||
model_data['params'] = model_data.get('params', {})
|
||||
new_model = ModelForm(**model_data)
|
||||
new_model.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
new_model.access_grants,
|
||||
@@ -677,7 +678,7 @@ async def update_model_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -749,7 +750,7 @@ async def update_model_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
@@ -11,6 +11,7 @@ from open_webui.config import (
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.notes import (
|
||||
NoteForm,
|
||||
@@ -66,7 +67,7 @@ async def get_notes(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -114,7 +115,7 @@ async def get_pinned_notes(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -155,7 +156,7 @@ async def search_notes(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -208,7 +209,7 @@ async def create_new_note(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -216,7 +217,7 @@ async def create_new_note(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -249,7 +250,7 @@ async def get_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -308,7 +309,7 @@ async def update_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -332,7 +333,7 @@ async def update_note_by_id(
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -375,7 +376,7 @@ async def update_note_access_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -399,7 +400,7 @@ async def update_note_access_by_id(
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -427,7 +428,7 @@ async def pin_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -469,7 +470,7 @@ async def delete_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
|
||||
@@ -31,6 +31,7 @@ from open_webui.env import (
|
||||
)
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import Models
|
||||
from open_webui.models.users import UserModel
|
||||
@@ -181,6 +182,32 @@ def get_api_key(idx, url, configs):
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
OLLAMA_CONFIG_KEYS = {
|
||||
'ENABLE_OLLAMA_API': 'ollama.enable',
|
||||
'OLLAMA_BASE_URLS': 'ollama.base_urls',
|
||||
'OLLAMA_API_CONFIGS': 'ollama.api_configs',
|
||||
}
|
||||
|
||||
|
||||
async def get_ollama_config_values() -> dict:
|
||||
values = await Config.get_many(*OLLAMA_CONFIG_KEYS.values())
|
||||
return {field: values[storage_key] for field, storage_key in OLLAMA_CONFIG_KEYS.items() if storage_key in values}
|
||||
|
||||
|
||||
async def get_ollama_runtime_config() -> tuple[bool, list[str], dict]:
|
||||
values = await Config.get_many('ollama.enable', 'ollama.base_urls', 'ollama.api_configs')
|
||||
return (
|
||||
values.get('ollama.enable'),
|
||||
values.get('ollama.base_urls') or [],
|
||||
values.get('ollama.api_configs') or {},
|
||||
)
|
||||
|
||||
|
||||
async def get_ollama_connection(idx: int) -> tuple[str, dict, str | None]:
|
||||
_, base_urls, api_configs = await get_ollama_runtime_config()
|
||||
url = base_urls[idx]
|
||||
return url, resolve_api_config(api_configs, idx, url), get_api_key(idx, url, api_configs)
|
||||
|
||||
|
||||
@router.head('/')
|
||||
@router.get('/')
|
||||
@@ -236,11 +263,7 @@ async def get_config(
|
||||
user=Depends(get_admin_user),
|
||||
) -> dict:
|
||||
"""Return the current Ollama connection configuration."""
|
||||
return {
|
||||
'ENABLE_OLLAMA_API': request.app.state.config.ENABLE_OLLAMA_API,
|
||||
'OLLAMA_BASE_URLS': request.app.state.config.OLLAMA_BASE_URLS,
|
||||
'OLLAMA_API_CONFIGS': request.app.state.config.OLLAMA_API_CONFIGS,
|
||||
}
|
||||
return await get_ollama_config_values()
|
||||
|
||||
|
||||
class OllamaConfigForm(BaseModel):
|
||||
@@ -258,20 +281,20 @@ async def update_config(
|
||||
user=Depends(get_admin_user),
|
||||
) -> dict:
|
||||
"""Persist updated Ollama connection settings."""
|
||||
request.app.state.config.ENABLE_OLLAMA_API = form_data.ENABLE_OLLAMA_API
|
||||
request.app.state.config.OLLAMA_BASE_URLS = form_data.OLLAMA_BASE_URLS
|
||||
request.app.state.config.OLLAMA_API_CONFIGS = form_data.OLLAMA_API_CONFIGS
|
||||
|
||||
# Prune stale config entries that no longer map to a URL index
|
||||
valid_keys = {str(i) for i in range(len(request.app.state.config.OLLAMA_BASE_URLS))}
|
||||
request.app.state.config.OLLAMA_API_CONFIGS = {
|
||||
k: v for k, v in request.app.state.config.OLLAMA_API_CONFIGS.items() if k in valid_keys
|
||||
}
|
||||
valid_keys = {str(i) for i in range(len(form_data.OLLAMA_BASE_URLS))}
|
||||
api_configs = {k: v for k, v in form_data.OLLAMA_API_CONFIGS.items() if k in valid_keys}
|
||||
|
||||
await Config.upsert(
|
||||
{
|
||||
'ollama.enable': form_data.ENABLE_OLLAMA_API,
|
||||
'ollama.base_urls': form_data.OLLAMA_BASE_URLS,
|
||||
'ollama.api_configs': api_configs,
|
||||
}
|
||||
)
|
||||
return {
|
||||
'ENABLE_OLLAMA_API': request.app.state.config.ENABLE_OLLAMA_API,
|
||||
'OLLAMA_BASE_URLS': request.app.state.config.OLLAMA_BASE_URLS,
|
||||
'OLLAMA_API_CONFIGS': request.app.state.config.OLLAMA_API_CONFIGS,
|
||||
'ENABLE_OLLAMA_API': form_data.ENABLE_OLLAMA_API,
|
||||
'OLLAMA_BASE_URLS': form_data.OLLAMA_BASE_URLS,
|
||||
'OLLAMA_API_CONFIGS': api_configs,
|
||||
}
|
||||
|
||||
|
||||
@@ -293,9 +316,8 @@ def merge_models_lists(model_lists) -> list[dict]:
|
||||
return list(merged.values())
|
||||
|
||||
|
||||
def _resolve_api_config(request: Request, idx: int, url: str) -> dict:
|
||||
def resolve_api_config(api_configs: dict, idx: int, url: str) -> dict:
|
||||
"""Look up the API config for a backend by numeric index, falling back to URL key (legacy)."""
|
||||
api_configs = request.app.state.config.OLLAMA_API_CONFIGS
|
||||
return api_configs.get(str(idx), api_configs.get(url, {}))
|
||||
|
||||
|
||||
@@ -307,15 +329,15 @@ async def get_all_models(request: Request, user: UserModel | None = None):
|
||||
"""Aggregate model tags from every enabled Ollama backend."""
|
||||
log.info('get_all_models()')
|
||||
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
models_dict: dict = {'models': []}
|
||||
request.app.state.OLLAMA_MODELS = {}
|
||||
return models_dict
|
||||
|
||||
# Fan-out tag requests to every backend
|
||||
tasks = []
|
||||
for idx, url in enumerate(request.app.state.config.OLLAMA_BASE_URLS):
|
||||
api_config = _resolve_api_config(request, idx, url)
|
||||
for idx, url in enumerate(await Config.get('ollama.base_urls', [])):
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), idx, url)
|
||||
if not api_config:
|
||||
tasks.append(send_get_request(f'{url}/api/tags', user=user))
|
||||
elif api_config.get('enable', True):
|
||||
@@ -329,8 +351,8 @@ async def get_all_models(request: Request, user: UserModel | None = None):
|
||||
for idx, response in enumerate(responses):
|
||||
if not response:
|
||||
continue
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[idx]
|
||||
api_config = _resolve_api_config(request, idx, url)
|
||||
url = (await Config.get('ollama.base_urls', []))[idx]
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), idx, url)
|
||||
|
||||
connection_type = api_config.get('connection_type', 'local')
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
@@ -394,14 +416,14 @@ async def get_ollama_tags(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""List Ollama model tags, optionally from a specific backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
if url_idx is None:
|
||||
result = await get_all_models(request, user=user)
|
||||
else:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
result = await send_request(f'{url}/api/tags', 'GET', key=key, user=user)
|
||||
|
||||
if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL:
|
||||
@@ -416,12 +438,12 @@ async def get_ollama_loaded_models(
|
||||
user=Depends(get_admin_user),
|
||||
) -> dict:
|
||||
"""List models currently loaded in Ollama memory across all backends."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
return {'models': []}
|
||||
|
||||
tasks = []
|
||||
for idx, url in enumerate(request.app.state.config.OLLAMA_BASE_URLS):
|
||||
api_config = _resolve_api_config(request, idx, url)
|
||||
for idx, url in enumerate(await Config.get('ollama.base_urls', [])):
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), idx, url)
|
||||
if not api_config:
|
||||
tasks.append(send_get_request(f'{url}/api/ps', user=user))
|
||||
elif api_config.get('enable', True):
|
||||
@@ -434,7 +456,7 @@ async def get_ollama_loaded_models(
|
||||
for idx, response in enumerate(responses):
|
||||
if not response:
|
||||
continue
|
||||
api_config = _resolve_api_config(request.app.state.config, idx, request.app.state.config.OLLAMA_BASE_URLS[idx])
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), idx, (await Config.get('ollama.base_urls', []))[idx])
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
for m in response.get('models', []):
|
||||
@@ -450,19 +472,19 @@ async def get_ollama_versions(
|
||||
url_idx: int | None = None,
|
||||
):
|
||||
"""Return the lowest Ollama version across all configured backends."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
return {'version': False}
|
||||
|
||||
if url_idx is not None:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
return await send_request(f'{url}/api/version', 'GET')
|
||||
|
||||
# Fan-out to every enabled backend
|
||||
tasks = []
|
||||
for idx, url in enumerate(request.app.state.config.OLLAMA_BASE_URLS):
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
for idx, url in enumerate(await Config.get('ollama.base_urls', [])):
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
if api_config.get('enable', True):
|
||||
tasks.append(send_get_request(f'{url}/api/version', api_config.get('key')))
|
||||
@@ -511,11 +533,11 @@ async def unload_model(
|
||||
results = []
|
||||
errors = []
|
||||
for idx in url_indices:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
str(idx), request.app.state.config.OLLAMA_API_CONFIGS.get(url, {})
|
||||
url = (await Config.get('ollama.base_urls', []))[idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(idx), (await Config.get('ollama.api_configs', {})).get(url, {})
|
||||
)
|
||||
key = get_api_key(idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
key = get_api_key(idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
if prefix_id and model.startswith(f'{prefix_id}.'):
|
||||
@@ -552,20 +574,20 @@ async def pull_model(
|
||||
url_idx: int = 0,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
form_data = form_data.model_dump(exclude_none=True)
|
||||
form_data['model'] = form_data.get('model', form_data.get('name'))
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
log.info(f'url: {url}')
|
||||
|
||||
# Admins may pull from any registry
|
||||
return await send_request(
|
||||
f'{url}/api/pull',
|
||||
payload=json.dumps({**form_data, 'insecure': True}),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -588,7 +610,7 @@ async def push_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Push a local model to a remote registry."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
if url_idx is None:
|
||||
@@ -598,13 +620,13 @@ async def push_model(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = models[form_data.model]['urls'][0]
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
log.debug(f'url: {url}')
|
||||
|
||||
return await send_request(
|
||||
f'{url}/api/push',
|
||||
payload=form_data.model_dump_json(exclude_none=True).encode(),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -627,16 +649,16 @@ async def create_model(
|
||||
url_idx: int = 0,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
log.debug(f'form_data: {form_data}')
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
|
||||
return await send_request(
|
||||
f'{url}/api/create',
|
||||
payload=form_data.model_dump_json(exclude_none=True).encode(),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -658,7 +680,7 @@ async def copy_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Duplicate an existing model under a new name."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
if url_idx is None:
|
||||
@@ -668,8 +690,8 @@ async def copy_model(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.source))
|
||||
url_idx = models[form_data.source]['urls'][0]
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
await send_request(
|
||||
f'{url}/api/copy',
|
||||
@@ -689,7 +711,7 @@ async def delete_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Remove a model from an Ollama backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = form_data.model_dump(exclude_none=True)
|
||||
@@ -703,8 +725,8 @@ async def delete_model(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(model))
|
||||
url_idx = models[model]['urls'][0]
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
await send_request(
|
||||
f'{url}/api/delete',
|
||||
@@ -723,7 +745,7 @@ async def show_model_info(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Retrieve model metadata from the Ollama backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = form_data.model_dump(exclude_none=True)
|
||||
@@ -739,8 +761,8 @@ async def show_model_info(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(model))
|
||||
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
return await send_request(
|
||||
f'{url}/api/show',
|
||||
@@ -770,7 +792,7 @@ async def embed(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Generate embeddings via the Ollama /api/embed endpoint."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
log.info(f'generate_ollama_batch_embeddings {form_data}')
|
||||
@@ -787,12 +809,12 @@ async def embed(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -824,7 +846,7 @@ async def embeddings(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Generate embeddings via the legacy Ollama /api/embeddings endpoint."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
log.info(f'generate_ollama_embeddings {form_data}')
|
||||
@@ -841,12 +863,12 @@ async def embeddings(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -886,7 +908,7 @@ async def generate_completion(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Run text completion via Ollama /api/generate."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
await check_model_access(user, await Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL)
|
||||
@@ -900,10 +922,10 @@ async def generate_completion(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
@@ -913,7 +935,7 @@ async def generate_completion(
|
||||
return await send_request(
|
||||
f'{url}/api/generate',
|
||||
payload=form_data.model_dump_json(exclude_none=True).encode(),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -973,7 +995,7 @@ async def get_ollama_url(request: Request, model: str, url_idx: int | None = Non
|
||||
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(model),
|
||||
)
|
||||
url_idx = random.choice(models[model].get('urls', []))
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
return url, url_idx
|
||||
|
||||
|
||||
@@ -986,7 +1008,7 @@ async def generate_chat_completion(
|
||||
user=Depends(get_verified_user), # noqa: B008
|
||||
):
|
||||
"""Forward a chat completion request to an Ollama backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
|
||||
@@ -1035,7 +1057,7 @@ async def generate_chat_completion(
|
||||
await check_model_access(user, None, bypass_filter)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = _resolve_api_config(request, url_idx, url)
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -1044,7 +1066,7 @@ async def generate_chat_completion(
|
||||
return await send_request(
|
||||
f'{url}/api/chat',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=form_data.stream,
|
||||
content_type='application/x-ndjson',
|
||||
@@ -1121,7 +1143,7 @@ async def generate_openai_completion(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = _resolve_api_config(request, url_idx, url)
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -1130,7 +1152,7 @@ async def generate_openai_completion(
|
||||
return await send_request(
|
||||
f'{url}/v1/completions',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
metadata=metadata,
|
||||
@@ -1178,7 +1200,7 @@ async def generate_openai_chat_completion(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = _resolve_api_config(request, url_idx, url)
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -1187,7 +1209,7 @@ async def generate_openai_chat_completion(
|
||||
return await send_request(
|
||||
f'{url}/v1/chat/completions',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
metadata=metadata,
|
||||
@@ -1211,7 +1233,7 @@ async def generate_anthropic_messages(
|
||||
|
||||
See https://docs.ollama.com/api/anthropic-compatibility
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = {**form_data}
|
||||
@@ -1227,9 +1249,9 @@ async def generate_anthropic_messages(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}), # Legacy support
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}), # Legacy support
|
||||
)
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
@@ -1239,7 +1261,7 @@ async def generate_anthropic_messages(
|
||||
return await send_request(
|
||||
f'{url}/v1/messages',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
content_type='text/event-stream' if payload.get('stream', False) else None,
|
||||
@@ -1269,7 +1291,7 @@ async def generate_responses(
|
||||
|
||||
See https://ollama.com/blog/responses-api
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = form_data.model_dump()
|
||||
@@ -1285,9 +1307,9 @@ async def generate_responses(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}), # Legacy support
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}), # Legacy support
|
||||
)
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
@@ -1297,7 +1319,7 @@ async def generate_responses(
|
||||
return await send_request(
|
||||
f'{url}/v1/responses',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
content_type='text/event-stream' if payload.get('stream', False) else None,
|
||||
@@ -1317,7 +1339,7 @@ async def get_openai_models(
|
||||
model_list = await get_all_models(request, user=user)
|
||||
raw_models = model_list['models']
|
||||
else:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
model_list = await send_request(f'{url}/api/tags', 'GET')
|
||||
raw_models = model_list.get('models', [])
|
||||
|
||||
@@ -1429,7 +1451,7 @@ async def download_model(
|
||||
detail='Invalid file_url. Only URLs from allowed hosts are permitted.',
|
||||
)
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx if url_idx is not None else 0]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx if url_idx is not None else 0]
|
||||
file_name = parse_huggingface_url(form_data.url)
|
||||
|
||||
if not file_name:
|
||||
@@ -1450,7 +1472,7 @@ async def upload_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Upload a local model file, push it as a blob, and create the model in Ollama."""
|
||||
ollama_url = request.app.state.config.OLLAMA_BASE_URLS[url_idx if url_idx is not None else 0]
|
||||
ollama_url = (await Config.get('ollama.base_urls', []))[url_idx if url_idx is not None else 0]
|
||||
|
||||
filename = os.path.basename(file.filename)
|
||||
file_path = os.path.join(UPLOAD_DIR, filename)
|
||||
|
||||
@@ -34,6 +34,7 @@ from open_webui.env import (
|
||||
)
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import Models
|
||||
from open_webui.models.users import UserModel
|
||||
@@ -236,15 +237,50 @@ def get_microsoft_entra_id_access_token():
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
OPENAI_CONFIG_KEYS = {
|
||||
'ENABLE_OPENAI_API': 'openai.enable',
|
||||
'OPENAI_API_BASE_URLS': 'openai.api_base_urls',
|
||||
'OPENAI_API_KEYS': 'openai.api_keys',
|
||||
'OPENAI_API_CONFIGS': 'openai.api_configs',
|
||||
}
|
||||
|
||||
|
||||
async def get_openai_config() -> dict:
|
||||
values = await Config.get_many(*OPENAI_CONFIG_KEYS.values())
|
||||
return {field: values[storage_key] for field, storage_key in OPENAI_CONFIG_KEYS.items() if storage_key in values}
|
||||
|
||||
|
||||
async def get_openai_runtime_config() -> tuple[bool, list[str], list[str], dict]:
|
||||
values = await Config.get_many('openai.enable', 'openai.api_base_urls', 'openai.api_keys', 'openai.api_configs')
|
||||
return (
|
||||
values.get('openai.enable'),
|
||||
values.get('openai.api_base_urls') or [],
|
||||
values.get('openai.api_keys') or [],
|
||||
values.get('openai.api_configs') or {},
|
||||
)
|
||||
|
||||
|
||||
async def normalize_openai_api_keys(api_base_urls: list[str], api_keys: list[str]) -> list[str]:
|
||||
if len(api_keys) > len(api_base_urls):
|
||||
api_keys = api_keys[: len(api_base_urls)]
|
||||
elif len(api_keys) < len(api_base_urls):
|
||||
api_keys = [*api_keys, *([''] * (len(api_base_urls) - len(api_keys)))]
|
||||
|
||||
await Config.upsert({'openai.api_keys': api_keys})
|
||||
return api_keys
|
||||
|
||||
|
||||
async def get_openai_connection(idx: int) -> tuple[str, str, dict]:
|
||||
_, api_base_urls, api_keys, api_configs = await get_openai_runtime_config()
|
||||
url = api_base_urls[idx]
|
||||
key = api_keys[idx]
|
||||
api_config = api_configs.get(str(idx), api_configs.get(url, {}))
|
||||
return url, key, api_config
|
||||
|
||||
|
||||
@router.get('/config')
|
||||
async def get_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_OPENAI_API': request.app.state.config.ENABLE_OPENAI_API,
|
||||
'OPENAI_API_BASE_URLS': request.app.state.config.OPENAI_API_BASE_URLS,
|
||||
'OPENAI_API_KEYS': request.app.state.config.OPENAI_API_KEYS,
|
||||
'OPENAI_API_CONFIGS': request.app.state.config.OPENAI_API_CONFIGS,
|
||||
}
|
||||
return await get_openai_config()
|
||||
|
||||
|
||||
class OpenAIConfigForm(BaseModel):
|
||||
@@ -256,41 +292,37 @@ class OpenAIConfigForm(BaseModel):
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_config(request: Request, form_data: OpenAIConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.ENABLE_OPENAI_API = form_data.ENABLE_OPENAI_API
|
||||
request.app.state.config.OPENAI_API_BASE_URLS = form_data.OPENAI_API_BASE_URLS
|
||||
request.app.state.config.OPENAI_API_KEYS = form_data.OPENAI_API_KEYS
|
||||
api_keys = form_data.OPENAI_API_KEYS
|
||||
|
||||
# Check if API KEYS length is same than API URLS length
|
||||
if len(request.app.state.config.OPENAI_API_KEYS) != len(request.app.state.config.OPENAI_API_BASE_URLS):
|
||||
if len(request.app.state.config.OPENAI_API_KEYS) > len(request.app.state.config.OPENAI_API_BASE_URLS):
|
||||
request.app.state.config.OPENAI_API_KEYS = request.app.state.config.OPENAI_API_KEYS[
|
||||
: len(request.app.state.config.OPENAI_API_BASE_URLS)
|
||||
]
|
||||
else:
|
||||
request.app.state.config.OPENAI_API_KEYS += [''] * (
|
||||
len(request.app.state.config.OPENAI_API_BASE_URLS) - len(request.app.state.config.OPENAI_API_KEYS)
|
||||
)
|
||||
if len(api_keys) > len(form_data.OPENAI_API_BASE_URLS):
|
||||
api_keys = api_keys[: len(form_data.OPENAI_API_BASE_URLS)]
|
||||
elif len(api_keys) < len(form_data.OPENAI_API_BASE_URLS):
|
||||
api_keys = [*api_keys, *([''] * (len(form_data.OPENAI_API_BASE_URLS) - len(api_keys)))]
|
||||
|
||||
request.app.state.config.OPENAI_API_CONFIGS = form_data.OPENAI_API_CONFIGS
|
||||
valid_keys = set(map(str, range(len(form_data.OPENAI_API_BASE_URLS))))
|
||||
api_configs = {key: value for key, value in form_data.OPENAI_API_CONFIGS.items() if key in valid_keys}
|
||||
|
||||
# Remove the API configs that are not in the API URLS
|
||||
keys = list(map(str, range(len(request.app.state.config.OPENAI_API_BASE_URLS))))
|
||||
request.app.state.config.OPENAI_API_CONFIGS = {
|
||||
key: value for key, value in request.app.state.config.OPENAI_API_CONFIGS.items() if key in keys
|
||||
}
|
||||
await Config.upsert(
|
||||
{
|
||||
'openai.enable': form_data.ENABLE_OPENAI_API,
|
||||
'openai.api_base_urls': form_data.OPENAI_API_BASE_URLS,
|
||||
'openai.api_keys': api_keys,
|
||||
'openai.api_configs': api_configs,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
'ENABLE_OPENAI_API': request.app.state.config.ENABLE_OPENAI_API,
|
||||
'OPENAI_API_BASE_URLS': request.app.state.config.OPENAI_API_BASE_URLS,
|
||||
'OPENAI_API_KEYS': request.app.state.config.OPENAI_API_KEYS,
|
||||
'OPENAI_API_CONFIGS': request.app.state.config.OPENAI_API_CONFIGS,
|
||||
'ENABLE_OPENAI_API': form_data.ENABLE_OPENAI_API,
|
||||
'OPENAI_API_BASE_URLS': form_data.OPENAI_API_BASE_URLS,
|
||||
'OPENAI_API_KEYS': api_keys,
|
||||
'OPENAI_API_CONFIGS': api_configs,
|
||||
}
|
||||
|
||||
|
||||
@router.post('/audio/speech')
|
||||
async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'chat.tts', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -299,7 +331,8 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
idx = None
|
||||
try:
|
||||
idx = request.app.state.config.OPENAI_API_BASE_URLS.index('https://api.openai.com/v1')
|
||||
_, api_base_urls, _, _ = await get_openai_runtime_config()
|
||||
idx = api_base_urls.index('https://api.openai.com/v1')
|
||||
|
||||
body = await request.body()
|
||||
name = hashlib.sha256(body).hexdigest()
|
||||
@@ -313,12 +346,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
if file_path.is_file():
|
||||
return FileResponse(file_path)
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
|
||||
|
||||
@@ -368,29 +396,15 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
async def get_all_models_responses(request: Request, user: UserModel) -> list:
|
||||
if not request.app.state.config.ENABLE_OPENAI_API:
|
||||
enable_openai_api, api_base_urls, api_keys, api_configs = await get_openai_runtime_config()
|
||||
if not enable_openai_api:
|
||||
return []
|
||||
|
||||
# Cache config values locally to avoid repeated Redis lookups.
|
||||
# Each access to request.app.state.config.<KEY> triggers a Redis GET;
|
||||
# caching here avoids hundreds of redundant round-trips.
|
||||
api_base_urls = request.app.state.config.OPENAI_API_BASE_URLS
|
||||
api_keys = list(request.app.state.config.OPENAI_API_KEYS)
|
||||
api_configs = request.app.state.config.OPENAI_API_CONFIGS
|
||||
|
||||
# Check if API KEYS length is same than API URLS length
|
||||
num_urls = len(api_base_urls)
|
||||
num_keys = len(api_keys)
|
||||
|
||||
if num_keys != num_urls:
|
||||
# if there are more keys than urls, remove the extra keys
|
||||
if num_keys > num_urls:
|
||||
api_keys = api_keys[:num_urls]
|
||||
request.app.state.config.OPENAI_API_KEYS = api_keys
|
||||
# if there are more urls than keys, add empty keys
|
||||
else:
|
||||
api_keys += [''] * (num_urls - num_keys)
|
||||
request.app.state.config.OPENAI_API_KEYS = api_keys
|
||||
api_keys = await normalize_openai_api_keys(api_base_urls, api_keys)
|
||||
|
||||
request_tasks = []
|
||||
for idx, url in enumerate(api_base_urls):
|
||||
@@ -500,13 +514,10 @@ async def get_filtered_models(models, user, db=None):
|
||||
async def get_all_models(request: Request, user: UserModel) -> dict[str, list]:
|
||||
log.info('get_all_models()')
|
||||
|
||||
if not request.app.state.config.ENABLE_OPENAI_API:
|
||||
enable_openai_api, api_base_urls, _, _ = await get_openai_runtime_config()
|
||||
if not enable_openai_api:
|
||||
return {'data': []}
|
||||
|
||||
# Cache config value locally to avoid repeated Redis lookups inside
|
||||
# the nested loop in get_merged_models (one GET per model otherwise).
|
||||
api_base_urls = request.app.state.config.OPENAI_API_BASE_URLS
|
||||
|
||||
responses = await get_all_models_responses(request, user=user)
|
||||
|
||||
def extract_data(response):
|
||||
@@ -577,7 +588,7 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]:
|
||||
@router.get('/models')
|
||||
@router.get('/models/{url_idx}')
|
||||
async def get_models(request: Request, url_idx: int | None = None, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_OPENAI_API:
|
||||
if not await Config.get('openai.enable'):
|
||||
raise HTTPException(status_code=503, detail='OpenAI API is disabled')
|
||||
|
||||
models = {
|
||||
@@ -587,13 +598,7 @@ async def get_models(request: Request, url_idx: int | None = None, user=Depends(
|
||||
if url_idx is None:
|
||||
models = await get_all_models(request, user=user)
|
||||
else:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[url_idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[url_idx]
|
||||
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(url_idx)
|
||||
|
||||
r = None
|
||||
async with aiohttp.ClientSession(
|
||||
@@ -1122,13 +1127,7 @@ async def generate_chat_completion(
|
||||
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
|
||||
)
|
||||
|
||||
# Get the API config for the model
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
request.app.state.config.OPENAI_API_BASE_URLS[idx], {}
|
||||
), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
if prefix_id:
|
||||
@@ -1143,9 +1142,6 @@ async def generate_chat_completion(
|
||||
'role': user.role,
|
||||
}
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
|
||||
# Check if model is a reasoning model that needs special handling
|
||||
if is_openai_new_model(payload['model']):
|
||||
payload = openai_reasoning_model_handler(payload)
|
||||
@@ -1311,12 +1307,7 @@ async def embeddings(request: Request, form_data: dict, user):
|
||||
if model_id in models:
|
||||
idx = models[model_id]['urlIdx']
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
r = None
|
||||
streaming = False
|
||||
@@ -1434,12 +1425,7 @@ async def responses(
|
||||
if model_id in models:
|
||||
idx = models[model_id]['urlIdx']
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
r = None
|
||||
streaming = False
|
||||
@@ -1543,14 +1529,7 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
|
||||
if model_id in models:
|
||||
idx = models[model_id]['urlIdx']
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
request.app.state.config.OPENAI_API_BASE_URLS[idx], {}
|
||||
), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
r = None
|
||||
streaming = False
|
||||
|
||||
@@ -18,6 +18,7 @@ from fastapi import (
|
||||
from open_webui.config import CACHE_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.routers.openai import get_all_models_responses
|
||||
from open_webui.utils.auth import get_admin_user
|
||||
from pydantic import BaseModel
|
||||
@@ -51,6 +52,12 @@ def get_sorted_filters(model_id, models):
|
||||
return sorted_filters
|
||||
|
||||
|
||||
async def get_openai_connection(url_idx: int) -> tuple[str, str]:
|
||||
base_urls = await Config.get('openai.api_base_urls', [])
|
||||
api_keys = await Config.get('openai.api_keys', [])
|
||||
return base_urls[url_idx], api_keys[url_idx]
|
||||
|
||||
|
||||
async def process_pipeline_inlet_filter(request, payload, user, models):
|
||||
user = {'id': user.id, 'email': user.email, 'name': user.name, 'role': user.role}
|
||||
model_id = payload['model']
|
||||
@@ -69,8 +76,7 @@ async def process_pipeline_inlet_filter(request, payload, user, models):
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
if not key:
|
||||
continue
|
||||
@@ -133,8 +139,7 @@ async def process_pipeline_outlet_filter(request, payload, user, models):
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
if not key:
|
||||
continue
|
||||
@@ -194,11 +199,12 @@ async def get_pipelines_list(request: Request, user=Depends(get_admin_user)):
|
||||
log.debug(f'get_pipelines_list: get_openai_models_responses returned {responses}')
|
||||
|
||||
urlIdxs = [idx for idx, response in enumerate(responses) if response is not None and 'pipelines' in response]
|
||||
base_urls = await Config.get('openai.api_base_urls', [])
|
||||
|
||||
return {
|
||||
'data': [
|
||||
{
|
||||
'url': request.app.state.config.OPENAI_API_BASE_URLS[urlIdx],
|
||||
'url': base_urls[urlIdx],
|
||||
'idx': urlIdx,
|
||||
}
|
||||
for urlIdx in urlIdxs
|
||||
@@ -233,8 +239,7 @@ async def upload_pipeline(
|
||||
with open(file_path, 'wb') as buffer:
|
||||
shutil.copyfileobj(file.file, buffer)
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
headers = {'Authorization': f'Bearer {key}'}
|
||||
|
||||
@@ -294,8 +299,7 @@ async def add_pipeline(request: Request, form_data: AddPipelineForm, user=Depend
|
||||
try:
|
||||
urlIdx = form_data.urlIdx
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.post(
|
||||
@@ -338,8 +342,7 @@ async def delete_pipeline(request: Request, form_data: DeletePipelineForm, user=
|
||||
try:
|
||||
urlIdx = form_data.urlIdx
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.delete(
|
||||
@@ -375,8 +378,7 @@ async def delete_pipeline(request: Request, form_data: DeletePipelineForm, user=
|
||||
async def get_pipelines(request: Request, urlIdx: Optional[int] = None, user=Depends(get_admin_user)):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(
|
||||
@@ -416,8 +418,7 @@ async def get_pipeline_valves(
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(
|
||||
@@ -457,8 +458,7 @@ async def get_pipeline_valves_spec(
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(
|
||||
@@ -499,8 +499,7 @@ async def update_pipeline_valves(
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.post(
|
||||
|
||||
@@ -7,6 +7,7 @@ from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.prompt_history import (
|
||||
PromptHistories,
|
||||
@@ -149,13 +150,13 @@ async def create_new_prompt(
|
||||
await has_permission(
|
||||
user.id,
|
||||
'workspace.prompts',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
or await has_permission(
|
||||
user.id,
|
||||
'workspace.prompts_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -165,7 +166,7 @@ async def create_new_prompt(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -281,7 +282,7 @@ async def update_prompt_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -438,7 +439,7 @@ async def update_prompt_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -259,10 +259,6 @@ def get_scim_auth(request: Request, authorization: Optional[str] = Header(None))
|
||||
enable_scim = getattr(request.app.state, 'ENABLE_SCIM', False)
|
||||
log.info(f'SCIM auth check - raw ENABLE_SCIM: {enable_scim}, type: {type(enable_scim)}')
|
||||
|
||||
# Handle both ConfigVar and direct value
|
||||
if hasattr(enable_scim, 'value'):
|
||||
enable_scim = enable_scim.value
|
||||
|
||||
if not enable_scim:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -271,9 +267,6 @@ def get_scim_auth(request: Request, authorization: Optional[str] = Header(None))
|
||||
|
||||
# Verify the SCIM token
|
||||
scim_token = getattr(request.app.state, 'SCIM_TOKEN', None)
|
||||
# Handle both ConfigVar and direct value
|
||||
if hasattr(scim_token, 'value'):
|
||||
scim_token = scim_token.value
|
||||
log.debug(f'SCIM token configured: {bool(scim_token)}')
|
||||
if not scim_token or not hmac.compare_digest(token, scim_token):
|
||||
raise HTTPException(
|
||||
|
||||
@@ -6,6 +6,7 @@ from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.skills import (
|
||||
SkillAccessListResponse,
|
||||
@@ -130,7 +131,7 @@ async def export_skills(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.skills',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -157,7 +158,7 @@ async def create_new_skill(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'workspace.skills', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'workspace.skills', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -179,7 +180,7 @@ async def create_new_skill(
|
||||
# grants in the create payload, bypassing the sharing.public_skills gate
|
||||
# that the dedicated /access/update endpoint already enforces.
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -292,7 +293,7 @@ async def update_skill_by_id(
|
||||
# they may set, so a non-admin owner cannot make their own skill publicly
|
||||
# readable/writable without sharing.public_skills permission.
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -361,7 +362,7 @@ async def update_skill_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
@@ -16,6 +16,7 @@ from open_webui.config import (
|
||||
DEFAULT_VOICE_MODE_PROMPT_TEMPLATE,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES, TASKS
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.routers.pipelines import process_pipeline_inlet_filter
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.chat import generate_chat_completion
|
||||
@@ -36,6 +37,35 @@ log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
TASK_CONFIG_KEYS = {
|
||||
'TASK_MODEL': 'task.model.default',
|
||||
'TASK_MODEL_EXTERNAL': 'task.model.external',
|
||||
'TITLE_GENERATION_PROMPT_TEMPLATE': 'task.title.prompt_template',
|
||||
'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': 'task.image.prompt_template',
|
||||
'ENABLE_AUTOCOMPLETE_GENERATION': 'task.autocomplete.enable',
|
||||
'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': 'task.autocomplete.input_max_length',
|
||||
'TAGS_GENERATION_PROMPT_TEMPLATE': 'task.tags.prompt_template',
|
||||
'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': 'task.follow_up.prompt_template',
|
||||
'ENABLE_FOLLOW_UP_GENERATION': 'task.follow_up.enable',
|
||||
'ENABLE_TAGS_GENERATION': 'task.tags.enable',
|
||||
'ENABLE_TITLE_GENERATION': 'task.title.enable',
|
||||
'ENABLE_SEARCH_QUERY_GENERATION': 'task.query.search.enable',
|
||||
'ENABLE_RETRIEVAL_QUERY_GENERATION': 'task.query.retrieval.enable',
|
||||
'QUERY_GENERATION_PROMPT_TEMPLATE': 'task.query.prompt_template',
|
||||
'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': 'task.tools.prompt_template',
|
||||
'ENABLE_VOICE_MODE_PROMPT': 'task.voice.prompt.enable',
|
||||
'VOICE_MODE_PROMPT_TEMPLATE': 'task.voice.prompt_template',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
##################################
|
||||
#
|
||||
@@ -59,25 +89,7 @@ async def check_active_chats(request: Request, form_data: ActiveChatsForm, user=
|
||||
|
||||
@router.get('/config')
|
||||
async def get_task_config(request: Request, user=Depends(get_verified_user)):
|
||||
return {
|
||||
'TASK_MODEL': request.app.state.config.TASK_MODEL,
|
||||
'TASK_MODEL_EXTERNAL': request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
'TITLE_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE,
|
||||
'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_AUTOCOMPLETE_GENERATION': request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION,
|
||||
'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH,
|
||||
'TAGS_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE,
|
||||
'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_FOLLOW_UP_GENERATION': request.app.state.config.ENABLE_FOLLOW_UP_GENERATION,
|
||||
'ENABLE_TAGS_GENERATION': request.app.state.config.ENABLE_TAGS_GENERATION,
|
||||
'ENABLE_TITLE_GENERATION': request.app.state.config.ENABLE_TITLE_GENERATION,
|
||||
'ENABLE_SEARCH_QUERY_GENERATION': request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION,
|
||||
'ENABLE_RETRIEVAL_QUERY_GENERATION': request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION,
|
||||
'QUERY_GENERATION_PROMPT_TEMPLATE': request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE,
|
||||
'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE,
|
||||
'ENABLE_VOICE_MODE_PROMPT': request.app.state.config.ENABLE_VOICE_MODE_PROMPT,
|
||||
'VOICE_MODE_PROMPT_TEMPLATE': request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE,
|
||||
}
|
||||
return await get_config_values(TASK_CONFIG_KEYS)
|
||||
|
||||
|
||||
class TaskConfigForm(BaseModel):
|
||||
@@ -102,56 +114,13 @@ class TaskConfigForm(BaseModel):
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_task_config(request: Request, form_data: TaskConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.TASK_MODEL = form_data.TASK_MODEL
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL = form_data.TASK_MODEL_EXTERNAL
|
||||
request.app.state.config.ENABLE_TITLE_GENERATION = form_data.ENABLE_TITLE_GENERATION
|
||||
request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE = form_data.TITLE_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.ENABLE_FOLLOW_UP_GENERATION = form_data.ENABLE_FOLLOW_UP_GENERATION
|
||||
request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE = form_data.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE = form_data.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION = form_data.ENABLE_AUTOCOMPLETE_GENERATION
|
||||
request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH = (
|
||||
form_data.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH
|
||||
)
|
||||
|
||||
request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE = form_data.TAGS_GENERATION_PROMPT_TEMPLATE
|
||||
request.app.state.config.ENABLE_TAGS_GENERATION = form_data.ENABLE_TAGS_GENERATION
|
||||
request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION = form_data.ENABLE_SEARCH_QUERY_GENERATION
|
||||
request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION = form_data.ENABLE_RETRIEVAL_QUERY_GENERATION
|
||||
|
||||
request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE = form_data.QUERY_GENERATION_PROMPT_TEMPLATE
|
||||
request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE = form_data.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.ENABLE_VOICE_MODE_PROMPT = form_data.ENABLE_VOICE_MODE_PROMPT
|
||||
request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE = form_data.VOICE_MODE_PROMPT_TEMPLATE
|
||||
|
||||
return {
|
||||
'TASK_MODEL': request.app.state.config.TASK_MODEL,
|
||||
'TASK_MODEL_EXTERNAL': request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
'ENABLE_TITLE_GENERATION': request.app.state.config.ENABLE_TITLE_GENERATION,
|
||||
'TITLE_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE,
|
||||
'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_AUTOCOMPLETE_GENERATION': request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION,
|
||||
'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH,
|
||||
'TAGS_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_TAGS_GENERATION': request.app.state.config.ENABLE_TAGS_GENERATION,
|
||||
'ENABLE_FOLLOW_UP_GENERATION': request.app.state.config.ENABLE_FOLLOW_UP_GENERATION,
|
||||
'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_SEARCH_QUERY_GENERATION': request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION,
|
||||
'ENABLE_RETRIEVAL_QUERY_GENERATION': request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION,
|
||||
'QUERY_GENERATION_PROMPT_TEMPLATE': request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE,
|
||||
'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE,
|
||||
'ENABLE_VOICE_MODE_PROMPT': request.app.state.config.ENABLE_VOICE_MODE_PROMPT,
|
||||
'VOICE_MODE_PROMPT_TEMPLATE': request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), TASK_CONFIG_KEYS))
|
||||
return await get_config_values(TASK_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/title/completions')
|
||||
async def generate_title(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_TITLE_GENERATION:
|
||||
if not await Config.get('task.title.enable'):
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={'detail': 'Title generation is disabled'},
|
||||
@@ -181,15 +150,16 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating chat title using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE
|
||||
title_template = await Config.get('task.title.prompt_template')
|
||||
if title_template != '':
|
||||
template = title_template
|
||||
else:
|
||||
template = DEFAULT_TITLE_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -234,7 +204,7 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
|
||||
|
||||
@router.post('/follow_up/completions')
|
||||
async def generate_follow_ups(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_FOLLOW_UP_GENERATION:
|
||||
if not await Config.get('task.follow_up.enable'):
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={'detail': 'Follow-up generation is disabled'},
|
||||
@@ -259,15 +229,16 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating chat title using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE
|
||||
follow_up_template = await Config.get('task.follow_up.prompt_template')
|
||||
if follow_up_template != '':
|
||||
template = follow_up_template
|
||||
else:
|
||||
template = DEFAULT_FOLLOW_UP_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -303,7 +274,7 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
|
||||
|
||||
@router.post('/tags/completions')
|
||||
async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_TAGS_GENERATION:
|
||||
if not await Config.get('task.tags.enable'):
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={'detail': 'Tags generation is disabled'},
|
||||
@@ -328,15 +299,16 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating chat tags using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE
|
||||
tags_template = await Config.get('task.tags.prompt_template')
|
||||
if tags_template != '':
|
||||
template = tags_template
|
||||
else:
|
||||
template = DEFAULT_TAGS_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -391,15 +363,16 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends(
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating image prompt using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE
|
||||
image_prompt_template = await Config.get('task.image.prompt_template')
|
||||
if image_prompt_template != '':
|
||||
template = image_prompt_template
|
||||
else:
|
||||
template = DEFAULT_IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -437,13 +410,13 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends(
|
||||
async def generate_queries(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
type = form_data.get('type')
|
||||
if type == 'web_search':
|
||||
if not request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION:
|
||||
if not await Config.get('task.query.search.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Search query generation'),
|
||||
)
|
||||
elif type == 'retrieval':
|
||||
if not request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION:
|
||||
if not await Config.get('task.query.retrieval.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Query generation'),
|
||||
@@ -472,15 +445,16 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating {type} queries using model {task_model_id} for user {user.email}')
|
||||
|
||||
if (request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE).strip() != '':
|
||||
template = request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE
|
||||
query_template = await Config.get('task.query.prompt_template')
|
||||
if query_template.strip() != '':
|
||||
template = query_template
|
||||
else:
|
||||
template = DEFAULT_QUERY_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -515,7 +489,7 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
|
||||
|
||||
@router.post('/auto/completions')
|
||||
async def generate_autocompletion(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION:
|
||||
if not await Config.get('task.autocomplete.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Autocompletion generation'),
|
||||
@@ -525,11 +499,12 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
|
||||
prompt = form_data.get('prompt')
|
||||
messages = form_data.get('messages')
|
||||
|
||||
if request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH > 0:
|
||||
if len(prompt) > request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH:
|
||||
autocomplete_input_max_length = await Config.get('task.autocomplete.input_max_length')
|
||||
if autocomplete_input_max_length > 0:
|
||||
if len(prompt) > autocomplete_input_max_length:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.INPUT_TOO_LONG(request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH),
|
||||
detail=ERROR_MESSAGES.INPUT_TOO_LONG(autocomplete_input_max_length),
|
||||
)
|
||||
|
||||
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
|
||||
@@ -551,15 +526,16 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating autocompletion using model {task_model_id} for user {user.email}')
|
||||
|
||||
if (request.app.state.config.AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE).strip() != '':
|
||||
template = request.app.state.config.AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE
|
||||
autocomplete_template = await Config.get('task.autocomplete.prompt_template')
|
||||
if autocomplete_template.strip() != '':
|
||||
template = autocomplete_template
|
||||
else:
|
||||
template = DEFAULT_AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -614,8 +590,8 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ from fastapi import APIRouter, Depends, Request, Response, WebSocket
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from open_webui.config import TERMINAL_PROXY_HEADERS
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.access_control import has_connection_access
|
||||
@@ -62,7 +63,7 @@ def _sanitize_proxy_path(path: str) -> str | None:
|
||||
@router.get('/')
|
||||
async def list_terminal_servers(request: Request, user=Depends(get_verified_user)):
|
||||
"""Return terminal servers the authenticated user has access to."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
|
||||
|
||||
return [
|
||||
@@ -87,7 +88,7 @@ async def proxy_terminal(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Proxy a request to the admin terminal server identified by *server_id*."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
connection = next((c for c in connections if c.get('id') == server_id), None)
|
||||
|
||||
if connection is None:
|
||||
@@ -235,7 +236,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str):
|
||||
return None
|
||||
|
||||
# Resolve terminal server
|
||||
connections = ws.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
connection = next((c for c in connections if c.get('id') == server_id), None)
|
||||
|
||||
if connection is None:
|
||||
|
||||
@@ -13,6 +13,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.models.tools import (
|
||||
@@ -84,7 +85,7 @@ async def get_tools(
|
||||
server_access_grants = {}
|
||||
for server in await get_tool_servers(request):
|
||||
server_idx = server.get('idx', 0)
|
||||
connections = request.app.state.config.TOOL_SERVER_CONNECTIONS
|
||||
connections = await Config.get('tool_server.connections', [])
|
||||
if server_idx >= len(connections):
|
||||
log.warning(
|
||||
f'Tool server index {server_idx} out of range '
|
||||
@@ -113,7 +114,7 @@ async def get_tools(
|
||||
)
|
||||
|
||||
# MCP Tool Servers
|
||||
for server in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
for server in await Config.get('tool_server.connections', []):
|
||||
if server.get('type', 'openapi') == 'mcp' and server.get('config', {}).get('enable'):
|
||||
server_id = server.get('info', {}).get('id')
|
||||
auth_type = server.get('auth_type', 'none')
|
||||
@@ -303,7 +304,7 @@ async def export_tools(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.tools_export',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -331,11 +332,11 @@ async def create_new_tools(
|
||||
):
|
||||
"""Create a new tool from user-supplied Python source code."""
|
||||
if user.role != 'admin' and not (
|
||||
await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
await has_permission(user.id, 'workspace.tools', await Config.get('user.permissions'), db=db)
|
||||
or await has_permission(
|
||||
user.id,
|
||||
'workspace.tools_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -356,7 +357,7 @@ async def create_new_tools(
|
||||
if tools is None:
|
||||
try:
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -484,8 +485,8 @@ async def update_tools_by_id(
|
||||
# Content edits trigger exec on load — gate them behind workspace.tools (matches /create).
|
||||
if form_data.content != tools.content:
|
||||
if user.role != 'admin' and not (
|
||||
await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
or await has_permission(user.id, 'workspace.tools_import', request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
await has_permission(user.id, 'workspace.tools', await Config.get('user.permissions'), db=db)
|
||||
or await has_permission(user.id, 'workspace.tools_import', await Config.get('user.permissions'), db=db)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -503,7 +504,7 @@ async def update_tools_by_id(
|
||||
specs = get_tool_specs(TOOLS[id])
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -574,7 +575,7 @@ async def update_tool_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
|
||||
@@ -12,6 +12,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING, PROFILE_IMAGE_ALLOWED_MIME_TYPES, STATIC_DIR
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.auths import Auths
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.models.users import (
|
||||
@@ -157,7 +158,7 @@ async def get_user_permissisions(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
user_permissions = await get_permissions(user.id, await Config.get('user.permissions'), db=db)
|
||||
|
||||
return user_permissions
|
||||
|
||||
@@ -239,7 +240,7 @@ class FeaturesPermissions(BaseModel):
|
||||
memories: bool = True
|
||||
automations: bool = False
|
||||
calendar: bool = True
|
||||
user_webhooks: bool = False
|
||||
webhooks: bool = False
|
||||
|
||||
|
||||
class SettingsPermissions(BaseModel):
|
||||
@@ -257,20 +258,22 @@ class UserPermissions(BaseModel):
|
||||
|
||||
@router.get('/default/permissions', response_model=UserPermissions)
|
||||
async def get_default_user_permissions(request: Request, user=Depends(get_admin_user)):
|
||||
user_permissions = await Config.get('user.permissions')
|
||||
return {
|
||||
'workspace': WorkspacePermissions(**request.app.state.config.USER_PERMISSIONS.get('workspace', {})),
|
||||
'sharing': SharingPermissions(**request.app.state.config.USER_PERMISSIONS.get('sharing', {})),
|
||||
'access_grants': AccessGrantsPermissions(**request.app.state.config.USER_PERMISSIONS.get('access_grants', {})),
|
||||
'chat': ChatPermissions(**request.app.state.config.USER_PERMISSIONS.get('chat', {})),
|
||||
'features': FeaturesPermissions(**request.app.state.config.USER_PERMISSIONS.get('features', {})),
|
||||
'settings': SettingsPermissions(**request.app.state.config.USER_PERMISSIONS.get('settings', {})),
|
||||
'workspace': WorkspacePermissions(**user_permissions.get('workspace', {})),
|
||||
'sharing': SharingPermissions(**user_permissions.get('sharing', {})),
|
||||
'access_grants': AccessGrantsPermissions(**user_permissions.get('access_grants', {})),
|
||||
'chat': ChatPermissions(**user_permissions.get('chat', {})),
|
||||
'features': FeaturesPermissions(**user_permissions.get('features', {})),
|
||||
'settings': SettingsPermissions(**user_permissions.get('settings', {})),
|
||||
}
|
||||
|
||||
|
||||
@router.post('/default/permissions')
|
||||
async def update_default_user_permissions(request: Request, form_data: UserPermissions, user=Depends(get_admin_user)):
|
||||
request.app.state.config.USER_PERMISSIONS = form_data.model_dump(by_alias=True)
|
||||
return request.app.state.config.USER_PERMISSIONS
|
||||
user_permissions = form_data.model_dump(by_alias=True)
|
||||
await Config.upsert({'user.permissions': user_permissions})
|
||||
return user_permissions
|
||||
|
||||
|
||||
@router.get('/default/permissions/defaults', response_model=UserPermissions)
|
||||
@@ -321,7 +324,7 @@ async def update_user_settings_by_session_user(
|
||||
and not await has_permission(
|
||||
user.id,
|
||||
'features.direct_tool_servers',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
)
|
||||
):
|
||||
# If the user is not an admin and does not have permission to use tool servers, remove the key
|
||||
@@ -348,7 +351,7 @@ async def get_user_status_by_session_user(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_USER_STATUS:
|
||||
if not await Config.get('users.enable_status'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
|
||||
@@ -369,7 +372,7 @@ async def update_user_status_by_session_user(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_USER_STATUS:
|
||||
if not await Config.get('users.enable_status'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
|
||||
|
||||
@@ -7,6 +7,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from open_webui.config import DATA_DIR, ENABLE_ADMIN_EXPORT
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.models.chats import ChatTitleMessagesForm
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.code_interpreter import execute_code_jupyter
|
||||
from open_webui.utils.misc import get_gravatar_url
|
||||
@@ -41,27 +42,27 @@ async def format_code(form_data: CodeForm, user=Depends(get_admin_user)):
|
||||
|
||||
@router.post('/code/execute')
|
||||
async def execute_code(request: Request, form_data: CodeForm, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_CODE_EXECUTION:
|
||||
if not await Config.get('code_execution.enable'):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Code execution'),
|
||||
)
|
||||
|
||||
if request.app.state.config.CODE_EXECUTION_ENGINE == 'jupyter':
|
||||
if await Config.get('code_execution.engine') == 'jupyter':
|
||||
output = await execute_code_jupyter(
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_URL,
|
||||
await Config.get('code_execution.jupyter.url'),
|
||||
form_data.code,
|
||||
(
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN
|
||||
if request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH == 'token'
|
||||
await Config.get('code_execution.jupyter.auth_token')
|
||||
if await Config.get('code_execution.jupyter.auth') == 'token'
|
||||
else None
|
||||
),
|
||||
(
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD
|
||||
if request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH == 'password'
|
||||
await Config.get('code_execution.jupyter.auth_password')
|
||||
if await Config.get('code_execution.jupyter.auth') == 'password'
|
||||
else None
|
||||
),
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT,
|
||||
await Config.get('code_execution.jupyter.timeout'),
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@@ -18,6 +18,7 @@ from fastapi import Request
|
||||
|
||||
from open_webui.models.channels import Channel, ChannelMember, Channels
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.memories import Memories
|
||||
from open_webui.models.messages import Message, Messages
|
||||
@@ -225,10 +226,10 @@ async def search_web(
|
||||
return json.dumps({'error': 'Request context not available'})
|
||||
|
||||
try:
|
||||
engine = __request__.app.state.config.WEB_SEARCH_ENGINE
|
||||
engine = await Config.get('rag.web.search.engine')
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
|
||||
configured = __request__.app.state.config.WEB_SEARCH_RESULT_COUNT
|
||||
configured = await Config.get('rag.web.search.result_count')
|
||||
max_count = 5 if configured is None else configured
|
||||
count = max(1, min(count, max_count)) if count is not None else max_count
|
||||
|
||||
@@ -266,7 +267,7 @@ async def fetch_url(
|
||||
# Truncate if configured (WEB_FETCH_MAX_CONTENT_LENGTH)
|
||||
# Guard: content may be None if the web loader silently failed
|
||||
if content is not None:
|
||||
max_length = getattr(__request__.app.state.config, 'WEB_FETCH_MAX_CONTENT_LENGTH', None)
|
||||
max_length = await Config.get('rag.web.fetch.max_content_length')
|
||||
if max_length and max_length > 0 and len(content) > max_length:
|
||||
content = content[:max_length] + '\n\n[Content truncated...]'
|
||||
else:
|
||||
@@ -475,7 +476,7 @@ async def execute_code(
|
||||
)
|
||||
code = blocking_code + '\n' + code
|
||||
|
||||
engine = getattr(__request__.app.state.config, 'CODE_INTERPRETER_ENGINE', 'pyodide')
|
||||
engine = await Config.get('code_interpreter.engine', 'pyodide')
|
||||
if engine == 'pyodide':
|
||||
# Execute via frontend pyodide using bidirectional event call
|
||||
if __event_call__ is None:
|
||||
@@ -513,21 +514,22 @@ async def execute_code(
|
||||
|
||||
elif engine == 'jupyter':
|
||||
from open_webui.utils.code_interpreter import execute_code_jupyter
|
||||
jupyter_auth = await Config.get('code_interpreter.jupyter.auth')
|
||||
|
||||
output = await execute_code_jupyter(
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
await Config.get('code_interpreter.jupyter.url'),
|
||||
code,
|
||||
(
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN
|
||||
if __request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH == 'token'
|
||||
await Config.get('code_interpreter.jupyter.auth_token')
|
||||
if jupyter_auth == 'token'
|
||||
else None
|
||||
),
|
||||
(
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD
|
||||
if __request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH == 'password'
|
||||
await Config.get('code_interpreter.jupyter.auth_password')
|
||||
if jupyter_auth == 'password'
|
||||
else None
|
||||
),
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
await Config.get('code_interpreter.jupyter.timeout'),
|
||||
)
|
||||
|
||||
stdout = output.get('stdout', '')
|
||||
|
||||
@@ -114,7 +114,7 @@ async def has_access(
|
||||
Check if a user has the specified permission using an in-memory access_grants list.
|
||||
|
||||
Used for config-driven resources (arena models, tool servers) that store
|
||||
access control as JSON in ConfigVar rather than in the access_grant DB table.
|
||||
access control as JSON config rather than in the access_grant DB table.
|
||||
|
||||
Semantics:
|
||||
- None or [] → private (owner-only, deny all)
|
||||
|
||||
@@ -39,6 +39,7 @@ from fastapi.responses import JSONResponse, RedirectResponse
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
from open_webui.env import CUSTOM_API_KEY_HEADER
|
||||
from open_webui.internal.db import ScopedSession
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.auth import get_http_authorization_cred
|
||||
from starlette.datastructures import MutableHeaders
|
||||
from starlette.requests import Request
|
||||
@@ -165,7 +166,7 @@ class AuthTokenMiddleware:
|
||||
token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=api_key)
|
||||
|
||||
request.state.token = token
|
||||
request.state.enable_api_keys = self._fastapi_app.state.config.ENABLE_API_KEYS
|
||||
request.state.enable_api_keys = await Config.get('auth.enable_api_keys')
|
||||
|
||||
async def send_with_timing(message: Message) -> None:
|
||||
if message['type'] == 'http.response.start':
|
||||
|
||||
@@ -35,6 +35,7 @@ from open_webui.env import (
|
||||
pk,
|
||||
)
|
||||
from open_webui.models.auths import Auths
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from pytz import UTC
|
||||
@@ -409,12 +410,16 @@ async def get_current_user_by_api_key(request, api_key: str):
|
||||
detail=ERROR_MESSAGES.INVALID_TOKEN,
|
||||
)
|
||||
|
||||
user_permissions = await Config.get('user.permissions')
|
||||
enable_endpoint_restrictions = await Config.get('auth.api_key.endpoint_restrictions')
|
||||
allowed_endpoints = await Config.get('auth.api_key.allowed_endpoints', '')
|
||||
|
||||
if not request.state.enable_api_keys or (
|
||||
user.role != 'admin'
|
||||
and not await has_permission(
|
||||
user.id,
|
||||
'features.api_keys',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
user_permissions,
|
||||
)
|
||||
):
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED)
|
||||
@@ -422,10 +427,8 @@ async def get_current_user_by_api_key(request, api_key: str):
|
||||
# Enforce endpoint restrictions — checked here (not in middleware)
|
||||
# so it applies regardless of how the API key was transported
|
||||
# (Authorization header, cookie, x-api-key header, etc.).
|
||||
if request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS:
|
||||
allowed_paths = [
|
||||
path.strip() for path in str(request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',') if path.strip()
|
||||
]
|
||||
if enable_endpoint_restrictions:
|
||||
allowed_paths = [path.strip() for path in str(allowed_endpoints).split(',') if path.strip()]
|
||||
request_path = request.scope['path'] # Use raw ASGI path — not spoofable via Host header (CVE-2026-48710)
|
||||
is_allowed = any(request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths)
|
||||
if not is_allowed:
|
||||
|
||||
@@ -29,6 +29,7 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.internal.db import get_async_db
|
||||
from open_webui.models.automations import AutomationModel, AutomationRuns, Automations
|
||||
from open_webui.models.chats import ChatForm, Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.task import prompt_template
|
||||
from starlette.datastructures import Headers
|
||||
@@ -172,7 +173,7 @@ async def scheduler_worker_loop(app) -> None:
|
||||
while True:
|
||||
try:
|
||||
# ── Automations ──
|
||||
if getattr(app.state.config, 'ENABLE_AUTOMATIONS', False):
|
||||
if await Config.get('automations.enable'):
|
||||
try:
|
||||
async with get_async_db() as db:
|
||||
batch = await Automations.claim_due(int(time.time_ns()), limit=10, db=db)
|
||||
@@ -184,7 +185,7 @@ async def scheduler_worker_loop(app) -> None:
|
||||
log.exception('Scheduler: automation error')
|
||||
|
||||
# ── Calendar Alerts ──
|
||||
if getattr(app.state.config, 'ENABLE_CALENDAR', False):
|
||||
if await Config.get('calendar.enable'):
|
||||
try:
|
||||
await _check_calendar_alerts(app)
|
||||
except Exception:
|
||||
@@ -239,7 +240,7 @@ def _resolve_model_tool_ids(app, model_id: str) -> list[str]:
|
||||
return list(tool_ids) if tool_ids else []
|
||||
|
||||
|
||||
def _resolve_model_features(app, model_id: str) -> dict:
|
||||
async def _resolve_model_features(app, model_id: str) -> dict:
|
||||
"""Read model default features from model config.
|
||||
|
||||
The frontend does this in Chat.svelte (model.info.meta.defaultFeatureIds
|
||||
@@ -256,14 +257,13 @@ def _resolve_model_features(app, model_id: str) -> dict:
|
||||
return {}
|
||||
|
||||
capabilities = meta.get('capabilities', {})
|
||||
config = app.state.config
|
||||
features = {}
|
||||
|
||||
# code_interpreter is excluded: it requires the frontend event emitter
|
||||
# and does not work in headless backend execution.
|
||||
feature_checks = {
|
||||
'web_search': getattr(config, 'ENABLE_WEB_SEARCH', False),
|
||||
'image_generation': getattr(config, 'ENABLE_IMAGE_GENERATION', False),
|
||||
'web_search': await Config.get('rag.web.search.enable'),
|
||||
'image_generation': await Config.get('image_generation.enable'),
|
||||
}
|
||||
|
||||
for feature_id in default_feature_ids:
|
||||
@@ -364,7 +364,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
|
||||
if user.role not in ('user', 'admin') or (
|
||||
user.role != 'admin'
|
||||
and not await has_permission(user.id, 'features.automations', app.state.config.USER_PERMISSIONS)
|
||||
and not await has_permission(user.id, 'features.automations', await Config.get('user.permissions'))
|
||||
):
|
||||
await _record_run(automation.id, 'error', error='Owner no longer permitted to run automations')
|
||||
return
|
||||
@@ -436,7 +436,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
|
||||
# Resolve model defaults (frontend does this, backend doesn't)
|
||||
tool_ids = _resolve_model_tool_ids(app, model_id)
|
||||
features = _resolve_model_features(app, model_id)
|
||||
features = await _resolve_model_features(app, model_id)
|
||||
filter_ids = _resolve_model_filter_ids(app, model_id)
|
||||
|
||||
# Resolve terminal from model config
|
||||
@@ -561,7 +561,7 @@ async def _check_calendar_alerts(app) -> None:
|
||||
# Send webhook notification if user has one configured
|
||||
try:
|
||||
webui_name = getattr(app.state, 'WEBUI_NAME', 'Open WebUI')
|
||||
enable_user_webhooks = getattr(app.state.config, 'ENABLE_USER_WEBHOOKS', False)
|
||||
enable_user_webhooks = await Config.get('ui.enable_user_webhooks')
|
||||
|
||||
if enable_user_webhooks:
|
||||
user = await Users.get_user_by_id(event.user_id)
|
||||
|
||||
@@ -40,6 +40,7 @@ from open_webui.env import (
|
||||
RAG_SYSTEM_CONTEXT,
|
||||
)
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.folders import Folders
|
||||
from open_webui.models.functions import Functions
|
||||
from open_webui.models.models import Models
|
||||
@@ -980,13 +981,13 @@ async def apply_source_context_to_messages(
|
||||
|
||||
if RAG_SYSTEM_CONTEXT:
|
||||
return add_or_update_system_message(
|
||||
await rag_template(request.app.state.config.RAG_TEMPLATE, context, user_message),
|
||||
await rag_template(await Config.get('rag.template'), context, user_message),
|
||||
messages,
|
||||
append=True,
|
||||
)
|
||||
else:
|
||||
return add_or_update_user_message(
|
||||
await rag_template(request.app.state.config.RAG_TEMPLATE, context, user_message),
|
||||
await rag_template(await Config.get('rag.template'), context, user_message),
|
||||
messages,
|
||||
append=False,
|
||||
)
|
||||
@@ -1290,8 +1291,8 @@ async def chat_completion_tools_handler(
|
||||
|
||||
task_model_id = get_task_model_id(
|
||||
body['model'],
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
@@ -1301,8 +1302,8 @@ async def chat_completion_tools_handler(
|
||||
specs = [tool['spec'] for tool in tools.values()]
|
||||
tools_specs = json.dumps(specs, ensure_ascii=False)
|
||||
|
||||
if request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE
|
||||
if await Config.get('task.tools.prompt_template') != '':
|
||||
template = await Config.get('task.tools.prompt_template')
|
||||
else:
|
||||
template = DEFAULT_TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE
|
||||
|
||||
@@ -1792,7 +1793,7 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra
|
||||
|
||||
system_message_content = ''
|
||||
|
||||
if len(input_images) > 0 and request.app.state.config.ENABLE_IMAGE_EDIT:
|
||||
if len(input_images) > 0 and await Config.get('images.edit.enable'):
|
||||
# Edit image(s)
|
||||
try:
|
||||
images = await image_edits(
|
||||
@@ -1852,7 +1853,7 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra
|
||||
|
||||
else:
|
||||
# Create image(s)
|
||||
if request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION:
|
||||
if await Config.get('image_generation.prompt.enable'):
|
||||
try:
|
||||
res = await generate_image_prompt(
|
||||
request,
|
||||
@@ -2018,17 +2019,17 @@ async def chat_completion_files_handler(
|
||||
embedding_function=lambda query, prefix: request.app.state.EMBEDDING_FUNCTION(
|
||||
query, prefix=prefix, user=user
|
||||
),
|
||||
k=request.app.state.config.TOP_K,
|
||||
k=await Config.get('rag.top_k'),
|
||||
reranking_function=(
|
||||
(lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents, user=user))
|
||||
if request.app.state.RERANKING_FUNCTION
|
||||
else None
|
||||
),
|
||||
k_reranker=request.app.state.config.TOP_K_RERANKER,
|
||||
r=request.app.state.config.RELEVANCE_THRESHOLD,
|
||||
hybrid_bm25_weight=request.app.state.config.HYBRID_BM25_WEIGHT,
|
||||
hybrid_search=request.app.state.config.ENABLE_RAG_HYBRID_SEARCH,
|
||||
full_context=all_full_context or request.app.state.config.RAG_FULL_CONTEXT,
|
||||
k_reranker=await Config.get('rag.top_k_reranker'),
|
||||
r=await Config.get('rag.relevance_threshold'),
|
||||
hybrid_bm25_weight=await Config.get('rag.hybrid_bm25_weight'),
|
||||
hybrid_search=await Config.get('rag.enable_hybrid_search'),
|
||||
full_context=all_full_context or await Config.get('rag.full_context'),
|
||||
user=user,
|
||||
)
|
||||
except Exception as e:
|
||||
@@ -2269,7 +2270,7 @@ async def connect_mcp_server(
|
||||
Returns None if the server is not found or access is denied.
|
||||
"""
|
||||
mcp_server_connection = None
|
||||
for server_connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
for server_connection in await Config.get('tool_server.connections', []):
|
||||
if server_connection.get('type', '') == 'mcp' and server_connection.get('info', {}).get('id') == server_id:
|
||||
mcp_server_connection = server_connection
|
||||
break
|
||||
@@ -2442,8 +2443,8 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
|
||||
task_model_id = get_task_model_id(
|
||||
form_data['model'],
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
@@ -2550,9 +2551,9 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
extra_params['__features__'] = features
|
||||
if features:
|
||||
if 'voice' in features and features['voice']:
|
||||
if getattr(request.app.state.config, 'ENABLE_VOICE_MODE_PROMPT', True):
|
||||
if request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE:
|
||||
template = request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE
|
||||
if await Config.get('task.voice.prompt.enable'):
|
||||
if await Config.get('task.voice.prompt_template'):
|
||||
template = await Config.get('task.voice.prompt_template')
|
||||
else:
|
||||
template = DEFAULT_VOICE_MODE_PROMPT_TEMPLATE
|
||||
|
||||
@@ -2577,14 +2578,14 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
form_data = await chat_image_generation_handler(request, form_data, extra_params, user)
|
||||
|
||||
if 'code_interpreter' in features and features['code_interpreter']:
|
||||
engine = getattr(request.app.state.config, 'CODE_INTERPRETER_ENGINE', 'pyodide')
|
||||
engine = await Config.get('code_interpreter.engine', 'pyodide')
|
||||
|
||||
# Skip XML-tag prompt injection when native FC is enabled —
|
||||
# execute_code will be injected as a builtin tool instead
|
||||
if metadata.get('params', {}).get('function_calling') == 'legacy':
|
||||
prompt = (
|
||||
request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE
|
||||
if request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE != ''
|
||||
await Config.get('code_interpreter.prompt_template')
|
||||
if await Config.get('code_interpreter.prompt_template') != ''
|
||||
else DEFAULT_CODE_INTERPRETER_PROMPT
|
||||
)
|
||||
|
||||
@@ -3534,18 +3535,19 @@ async def non_streaming_chat_response_handler(response, ctx):
|
||||
)
|
||||
|
||||
# Send a webhook notification if the user is not active
|
||||
if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id):
|
||||
if await Config.get('ui.enable_user_webhooks') and not await Users.is_user_active(user.id):
|
||||
webhook_url = await Users.get_user_webhook_url_by_id(user.id)
|
||||
if webhook_url:
|
||||
webui_url = await Config.get('webui.url')
|
||||
await post_webhook(
|
||||
request.app.state.WEBUI_NAME,
|
||||
webhook_url,
|
||||
f'{content}\n\n{title} - {request.app.state.config.WEBUI_URL}/c/{metadata["chat_id"]}',
|
||||
f'{content}\n\n{title} - {webui_url}/c/{metadata["chat_id"]}',
|
||||
{
|
||||
'action': 'chat',
|
||||
'message': content,
|
||||
'title': title,
|
||||
'url': f'{request.app.state.config.WEBUI_URL}/c/{metadata["chat_id"]}',
|
||||
'url': f'{webui_url}/c/{metadata["chat_id"]}',
|
||||
},
|
||||
)
|
||||
|
||||
@@ -3883,14 +3885,14 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
DETECT_CODE_INTERPRETER = (
|
||||
bool(features.get('code_interpreter'))
|
||||
and builtin_tools_meta.get('code_interpreter', True)
|
||||
and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True)
|
||||
and await Config.get('code_interpreter.enable')
|
||||
and model_capabilities.get('code_interpreter', True)
|
||||
and (
|
||||
getattr(user, 'role', None) == 'admin'
|
||||
or await has_permission(
|
||||
getattr(user, 'id', ''),
|
||||
'features.code_interpreter',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
)
|
||||
)
|
||||
)
|
||||
@@ -4784,7 +4786,7 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
source_context = source_context.strip()
|
||||
if source_context:
|
||||
rag_content = await rag_template(
|
||||
request.app.state.config.RAG_TEMPLATE,
|
||||
await Config.get('rag.template'),
|
||||
source_context,
|
||||
user_message,
|
||||
)
|
||||
@@ -4978,7 +4980,7 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
""")
|
||||
code = blocking_code + '\n' + code
|
||||
|
||||
if request.app.state.config.CODE_INTERPRETER_ENGINE == 'pyodide':
|
||||
if await Config.get('code_interpreter.engine') == 'pyodide':
|
||||
ci_output = await event_caller(
|
||||
{
|
||||
'type': 'execute:python',
|
||||
@@ -4990,21 +4992,21 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
},
|
||||
}
|
||||
)
|
||||
elif request.app.state.config.CODE_INTERPRETER_ENGINE == 'jupyter':
|
||||
elif await Config.get('code_interpreter.engine') == 'jupyter':
|
||||
ci_output = await execute_code_jupyter(
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
await Config.get('code_interpreter.jupyter.url'),
|
||||
code,
|
||||
(
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN
|
||||
if request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH == 'token'
|
||||
await Config.get('code_interpreter.jupyter.auth_token')
|
||||
if await Config.get('code_interpreter.jupyter.auth') == 'token'
|
||||
else None
|
||||
),
|
||||
(
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD
|
||||
if request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH == 'password'
|
||||
await Config.get('code_interpreter.jupyter.auth_password')
|
||||
if await Config.get('code_interpreter.jupyter.auth') == 'password'
|
||||
else None
|
||||
),
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
await Config.get('code_interpreter.jupyter.timeout'),
|
||||
)
|
||||
else:
|
||||
ci_output = {'stdout': 'Code interpreter engine not configured.'}
|
||||
@@ -5148,18 +5150,19 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
)
|
||||
|
||||
# Send a webhook notification if the user is not active
|
||||
if request.app.state.config.ENABLE_USER_WEBHOOKS and not await Users.is_user_active(user.id):
|
||||
if await Config.get('ui.enable_user_webhooks') and not await Users.is_user_active(user.id):
|
||||
webhook_url = await Users.get_user_webhook_url_by_id(user.id)
|
||||
if webhook_url:
|
||||
webui_url = await Config.get('webui.url')
|
||||
await post_webhook(
|
||||
request.app.state.WEBUI_NAME,
|
||||
webhook_url,
|
||||
f'{content}\n\n{title} - {request.app.state.config.WEBUI_URL}/c/{metadata["chat_id"]}',
|
||||
f'{content}\n\n{title} - {webui_url}/c/{metadata["chat_id"]}',
|
||||
{
|
||||
'action': 'chat',
|
||||
'message': content,
|
||||
'title': title,
|
||||
'url': f'{request.app.state.config.WEBUI_URL}/c/{metadata["chat_id"]}',
|
||||
'url': f'{webui_url}/c/{metadata["chat_id"]}',
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from open_webui.config import (
|
||||
from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, GLOBAL_LOG_LEVEL
|
||||
from open_webui.functions import get_function_models
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.functions import Functions
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import Models
|
||||
@@ -52,14 +53,15 @@ async def fetch_openai_models(request: Request, user: UserModel = None):
|
||||
|
||||
|
||||
async def get_all_base_models(request: Request, user: UserModel = None):
|
||||
config = await Config.get_many('openai.enable', 'ollama.enable')
|
||||
openai_task = (
|
||||
fetch_openai_models(request, user)
|
||||
if request.app.state.config.ENABLE_OPENAI_API
|
||||
if config.get('openai.enable')
|
||||
else asyncio.sleep(0, result=[])
|
||||
)
|
||||
ollama_task = (
|
||||
fetch_ollama_models(request, user)
|
||||
if request.app.state.config.ENABLE_OLLAMA_API
|
||||
if config.get('ollama.enable')
|
||||
else asyncio.sleep(0, result=[])
|
||||
)
|
||||
function_task = get_function_models(request)
|
||||
@@ -70,10 +72,15 @@ async def get_all_base_models(request: Request, user: UserModel = None):
|
||||
|
||||
|
||||
async def get_all_models(request, refresh: bool = False, user: UserModel = None):
|
||||
config = await Config.get_many(
|
||||
'models.base_models_cache',
|
||||
'evaluation.arena.enable',
|
||||
'evaluation.arena.models',
|
||||
)
|
||||
if (
|
||||
request.app.state.MODELS
|
||||
and request.app.state.BASE_MODELS
|
||||
and (request.app.state.config.ENABLE_BASE_MODELS_CACHE and not refresh)
|
||||
and (config.get('models.base_models_cache') and not refresh)
|
||||
):
|
||||
base_models = request.app.state.BASE_MODELS
|
||||
else:
|
||||
@@ -88,9 +95,10 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None)
|
||||
return []
|
||||
|
||||
# Add arena models
|
||||
if request.app.state.config.ENABLE_EVALUATION_ARENA_MODELS:
|
||||
if config.get('evaluation.arena.enable'):
|
||||
arena_models = []
|
||||
if len(request.app.state.config.EVALUATION_ARENA_MODELS) > 0:
|
||||
arena_config = config.get('evaluation.arena.models') or []
|
||||
if len(arena_config) > 0:
|
||||
arena_models = [
|
||||
{
|
||||
'id': model['id'],
|
||||
@@ -103,7 +111,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None)
|
||||
'owned_by': 'arena',
|
||||
'arena': True,
|
||||
}
|
||||
for model in request.app.state.config.EVALUATION_ARENA_MODELS
|
||||
for model in arena_config
|
||||
]
|
||||
else:
|
||||
# Add default arena model
|
||||
@@ -289,7 +297,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None)
|
||||
|
||||
# Apply global model defaults to all models
|
||||
# Per-model overrides take precedence over global defaults
|
||||
default_metadata = getattr(request.app.state.config, 'DEFAULT_MODEL_METADATA', None) or {}
|
||||
default_metadata = await Config.get('models.default_metadata', {}) or {}
|
||||
|
||||
if default_metadata:
|
||||
for model in models:
|
||||
|
||||
@@ -1,18 +1,17 @@
|
||||
import base64
|
||||
import copy
|
||||
import fnmatch
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import re
|
||||
import secrets
|
||||
import sys
|
||||
import time
|
||||
import urllib
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import Literal, Optional
|
||||
|
||||
import aiohttp
|
||||
@@ -62,7 +61,6 @@ from open_webui.config import (
|
||||
OAUTH_UPDATE_PICTURE_ON_LOGIN,
|
||||
OAUTH_USERNAME_CLAIM,
|
||||
WEBHOOK_URL,
|
||||
AppConfig,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES, WEBHOOK_MESSAGES
|
||||
from open_webui.env import (
|
||||
@@ -78,6 +76,7 @@ from open_webui.env import (
|
||||
WEBUI_NAME,
|
||||
)
|
||||
from open_webui.models.auths import Auths
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import GroupForm, GroupModel, Groups, GroupUpdateForm
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.models.users import Users
|
||||
@@ -111,31 +110,73 @@ from open_webui.env import GLOBAL_LOG_LEVEL
|
||||
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
auth_manager_config = AppConfig()
|
||||
auth_manager_config.DEFAULT_USER_ROLE = DEFAULT_USER_ROLE
|
||||
auth_manager_config.ENABLE_OAUTH_SIGNUP = ENABLE_OAUTH_SIGNUP
|
||||
auth_manager_config.OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE = OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE
|
||||
auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL = OAUTH_MERGE_ACCOUNTS_BY_EMAIL
|
||||
auth_manager_config.ENABLE_OAUTH_ROLE_MANAGEMENT = ENABLE_OAUTH_ROLE_MANAGEMENT
|
||||
auth_manager_config.ENABLE_OAUTH_GROUP_MANAGEMENT = ENABLE_OAUTH_GROUP_MANAGEMENT
|
||||
auth_manager_config.ENABLE_OAUTH_GROUP_CREATION = ENABLE_OAUTH_GROUP_CREATION
|
||||
auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE = OAUTH_GROUP_DEFAULT_SHARE
|
||||
auth_manager_config.OAUTH_BLOCKED_GROUPS = OAUTH_BLOCKED_GROUPS
|
||||
auth_manager_config.OAUTH_ROLES_CLAIM = OAUTH_ROLES_CLAIM
|
||||
auth_manager_config.OAUTH_SUB_CLAIM = OAUTH_SUB_CLAIM
|
||||
auth_manager_config.OAUTH_GROUPS_CLAIM = OAUTH_GROUPS_CLAIM
|
||||
auth_manager_config.OAUTH_EMAIL_CLAIM = OAUTH_EMAIL_CLAIM
|
||||
auth_manager_config.OAUTH_PICTURE_CLAIM = OAUTH_PICTURE_CLAIM
|
||||
auth_manager_config.OAUTH_USERNAME_CLAIM = OAUTH_USERNAME_CLAIM
|
||||
auth_manager_config.OAUTH_ALLOWED_ROLES = OAUTH_ALLOWED_ROLES
|
||||
auth_manager_config.OAUTH_ADMIN_ROLES = OAUTH_ADMIN_ROLES
|
||||
auth_manager_config.OAUTH_ALLOWED_DOMAINS = OAUTH_ALLOWED_DOMAINS
|
||||
auth_manager_config.WEBHOOK_URL = WEBHOOK_URL
|
||||
auth_manager_config.JWT_EXPIRES_IN = JWT_EXPIRES_IN
|
||||
auth_manager_config.OAUTH_UPDATE_PICTURE_ON_LOGIN = OAUTH_UPDATE_PICTURE_ON_LOGIN
|
||||
auth_manager_config.OAUTH_UPDATE_NAME_ON_LOGIN = OAUTH_UPDATE_NAME_ON_LOGIN
|
||||
auth_manager_config.OAUTH_UPDATE_EMAIL_ON_LOGIN = OAUTH_UPDATE_EMAIL_ON_LOGIN
|
||||
auth_manager_config.OAUTH_AUDIENCE = OAUTH_AUDIENCE
|
||||
OAUTH_RUNTIME_CONFIG = {
|
||||
'DEFAULT_USER_ROLE': ('ui.default_user_role', DEFAULT_USER_ROLE),
|
||||
'ENABLE_OAUTH_SIGNUP': ('oauth.enable_signup', ENABLE_OAUTH_SIGNUP),
|
||||
'OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE': (
|
||||
'oauth.refresh_token.include_scope',
|
||||
OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE,
|
||||
),
|
||||
'OAUTH_MERGE_ACCOUNTS_BY_EMAIL': (
|
||||
'oauth.merge_accounts_by_email',
|
||||
OAUTH_MERGE_ACCOUNTS_BY_EMAIL,
|
||||
),
|
||||
'ENABLE_OAUTH_ROLE_MANAGEMENT': (
|
||||
'oauth.enable_role_mapping',
|
||||
ENABLE_OAUTH_ROLE_MANAGEMENT,
|
||||
),
|
||||
'ENABLE_OAUTH_GROUP_MANAGEMENT': (
|
||||
'oauth.enable_group_mapping',
|
||||
ENABLE_OAUTH_GROUP_MANAGEMENT,
|
||||
),
|
||||
'ENABLE_OAUTH_GROUP_CREATION': (
|
||||
'oauth.enable_group_creation',
|
||||
ENABLE_OAUTH_GROUP_CREATION,
|
||||
),
|
||||
'OAUTH_GROUP_DEFAULT_SHARE': (
|
||||
'oauth.group_default_share',
|
||||
OAUTH_GROUP_DEFAULT_SHARE,
|
||||
),
|
||||
'OAUTH_BLOCKED_GROUPS': ('oauth.blocked_groups', OAUTH_BLOCKED_GROUPS),
|
||||
'OAUTH_ROLES_CLAIM': ('oauth.roles_claim', OAUTH_ROLES_CLAIM),
|
||||
'OAUTH_SUB_CLAIM': ('oauth.sub_claim', OAUTH_SUB_CLAIM),
|
||||
'OAUTH_GROUPS_CLAIM': ('oauth.group_claim', OAUTH_GROUPS_CLAIM),
|
||||
'OAUTH_EMAIL_CLAIM': ('oauth.email_claim', OAUTH_EMAIL_CLAIM),
|
||||
'OAUTH_PICTURE_CLAIM': ('oauth.picture_claim', OAUTH_PICTURE_CLAIM),
|
||||
'OAUTH_USERNAME_CLAIM': ('oauth.username_claim', OAUTH_USERNAME_CLAIM),
|
||||
'OAUTH_ALLOWED_ROLES': ('oauth.allowed_roles', OAUTH_ALLOWED_ROLES),
|
||||
'OAUTH_ADMIN_ROLES': ('oauth.admin_roles', OAUTH_ADMIN_ROLES),
|
||||
'OAUTH_ALLOWED_DOMAINS': ('oauth.allowed_domains', OAUTH_ALLOWED_DOMAINS),
|
||||
'WEBHOOK_URL': ('webhook_url', WEBHOOK_URL),
|
||||
'JWT_EXPIRES_IN': ('auth.jwt_expiry', JWT_EXPIRES_IN),
|
||||
'OAUTH_UPDATE_PICTURE_ON_LOGIN': (
|
||||
'oauth.update_picture_on_login',
|
||||
OAUTH_UPDATE_PICTURE_ON_LOGIN,
|
||||
),
|
||||
'OAUTH_UPDATE_NAME_ON_LOGIN': (
|
||||
'oauth.update_name_on_login',
|
||||
OAUTH_UPDATE_NAME_ON_LOGIN,
|
||||
),
|
||||
'OAUTH_UPDATE_EMAIL_ON_LOGIN': (
|
||||
'oauth.update_email_on_login',
|
||||
OAUTH_UPDATE_EMAIL_ON_LOGIN,
|
||||
),
|
||||
'OAUTH_AUDIENCE': ('oauth.audience', OAUTH_AUDIENCE),
|
||||
}
|
||||
|
||||
|
||||
def _default_value(value):
|
||||
return getattr(value, 'value', value)
|
||||
|
||||
|
||||
async def get_oauth_runtime_config() -> SimpleNamespace:
|
||||
keys = [key for key, _default in OAUTH_RUNTIME_CONFIG.values()]
|
||||
stored = await Config.get_many(*keys)
|
||||
values = {
|
||||
name: stored.get(key, _default_value(default))
|
||||
for name, (key, default) in OAUTH_RUNTIME_CONFIG.items()
|
||||
}
|
||||
return SimpleNamespace(**values)
|
||||
|
||||
|
||||
# Conservative default when the provider omits both expires_in and expires_at.
|
||||
@@ -423,7 +464,8 @@ async def get_oauth_client_info_with_dynamic_client_registration(
|
||||
oauth_server_metadata = None
|
||||
oauth_server_metadata_url = None
|
||||
|
||||
redirect_base_url = (str(request.app.state.config.WEBUI_URL or request.base_url)).rstrip('/')
|
||||
webui_url = await Config.get('webui.url')
|
||||
redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
|
||||
|
||||
oauth_client_metadata = OAuthClientMetadata(
|
||||
client_name='Open WebUI',
|
||||
@@ -549,7 +591,8 @@ async def get_oauth_client_info_with_static_credentials(
|
||||
oauth_server_metadata = None
|
||||
oauth_server_metadata_url = None
|
||||
|
||||
redirect_base_url = (str(request.app.state.config.WEBUI_URL or request.base_url)).rstrip('/')
|
||||
webui_url = await Config.get('webui.url')
|
||||
redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
|
||||
redirect_uri = f'{redirect_base_url}/oauth/clients/{client_id}/callback'
|
||||
|
||||
# Discover server metadata (authorization endpoint, token endpoint, scopes, etc.)
|
||||
@@ -636,7 +679,7 @@ class OAuthClientManager:
|
||||
'client_secret': oauth_client_info.client_secret,
|
||||
'client_kwargs': {
|
||||
'follow_redirects': True,
|
||||
**({'timeout': int(OAUTH_CLIENT_TIMEOUT.value)} if OAUTH_CLIENT_TIMEOUT.value else {}),
|
||||
**({'timeout': int(OAUTH_CLIENT_TIMEOUT)} if OAUTH_CLIENT_TIMEOUT else {}),
|
||||
**({'scope': oauth_client_info.scope} if oauth_client_info.scope else {}),
|
||||
**(
|
||||
{'token_endpoint_auth_method': oauth_client_info.token_endpoint_auth_method}
|
||||
@@ -668,7 +711,7 @@ class OAuthClientManager:
|
||||
}
|
||||
return self.clients[client_id]
|
||||
|
||||
def ensure_client_from_config(self, client_id):
|
||||
async def ensure_client_from_config(self, client_id):
|
||||
"""
|
||||
Lazy-load an OAuth client from the current TOOL_SERVER_CONNECTIONS
|
||||
config if it hasn't been registered on this node yet.
|
||||
@@ -677,7 +720,7 @@ class OAuthClientManager:
|
||||
return self.clients[client_id]['client']
|
||||
|
||||
try:
|
||||
connections = getattr(self.app.state.config, 'TOOL_SERVER_CONNECTIONS', [])
|
||||
connections = await Config.get('tool_server.connections', [])
|
||||
except Exception:
|
||||
connections = []
|
||||
|
||||
@@ -783,22 +826,22 @@ class OAuthClientManager:
|
||||
|
||||
return True
|
||||
|
||||
def get_client(self, client_id):
|
||||
async def get_client(self, client_id):
|
||||
if client_id not in self.clients:
|
||||
self.ensure_client_from_config(client_id)
|
||||
await self.ensure_client_from_config(client_id)
|
||||
|
||||
client = self.clients.get(client_id)
|
||||
return client['client'] if client else None
|
||||
|
||||
def get_client_info(self, client_id):
|
||||
async def get_client_info(self, client_id):
|
||||
if client_id not in self.clients:
|
||||
self.ensure_client_from_config(client_id)
|
||||
await self.ensure_client_from_config(client_id)
|
||||
|
||||
client = self.clients.get(client_id)
|
||||
return client['client_info'] if client else None
|
||||
|
||||
def get_server_metadata_url(self, client_id):
|
||||
client = self.get_client(client_id)
|
||||
async def get_server_metadata_url(self, client_id):
|
||||
client = await self.get_client(client_id)
|
||||
if not client:
|
||||
return None
|
||||
|
||||
@@ -881,6 +924,7 @@ class OAuthClientManager:
|
||||
Returns:
|
||||
dict: New token data, or None if refresh failed
|
||||
"""
|
||||
auth_config = await get_oauth_runtime_config()
|
||||
client_id = session.provider
|
||||
token_data = session.token
|
||||
|
||||
@@ -889,14 +933,14 @@ class OAuthClientManager:
|
||||
return None
|
||||
|
||||
try:
|
||||
client = self.get_client(client_id)
|
||||
client = await self.get_client(client_id)
|
||||
if not client:
|
||||
log.error(f'No OAuth client found for provider {client_id}')
|
||||
return None
|
||||
|
||||
token_endpoint = None
|
||||
async with aiohttp.ClientSession(trust_env=True) as session_http:
|
||||
async with session_http.get(self.get_server_metadata_url(client_id)) as r:
|
||||
async with session_http.get(await self.get_server_metadata_url(client_id)) as r:
|
||||
if r.status == 200:
|
||||
openid_data = await r.json()
|
||||
token_endpoint = openid_data.get('token_endpoint')
|
||||
@@ -913,7 +957,7 @@ class OAuthClientManager:
|
||||
'client_id': client.client_id,
|
||||
}
|
||||
# RFC 8707: include resource indicator so refreshed tokens retain correct audience
|
||||
client_info = self.get_client_info(client_id)
|
||||
client_info = await self.get_client_info(client_id)
|
||||
if client_info and client_info.resource:
|
||||
refresh_data['resource'] = client_info.resource
|
||||
|
||||
@@ -924,7 +968,7 @@ class OAuthClientManager:
|
||||
if (
|
||||
hasattr(client, 'client_kwargs')
|
||||
and client.client_kwargs.get('scope')
|
||||
and getattr(self.app.state.config, 'OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE', False)
|
||||
and auth_config.OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE
|
||||
):
|
||||
refresh_data['scope'] = client.client_kwargs['scope']
|
||||
|
||||
@@ -957,13 +1001,13 @@ class OAuthClientManager:
|
||||
return None
|
||||
|
||||
async def handle_authorize(self, request, client_id: str) -> RedirectResponse:
|
||||
client = self.get_client(client_id) or self.ensure_client_from_config(client_id)
|
||||
client = await self.get_client(client_id)
|
||||
if client is None:
|
||||
raise HTTPException(404)
|
||||
client_info = self.get_client_info(client_id)
|
||||
client_info = await self.get_client_info(client_id)
|
||||
if client_info is None:
|
||||
# ensure_client_from_config registers client_info too
|
||||
client_info = self.get_client_info(client_id)
|
||||
# get_client registers client_info too
|
||||
client_info = await self.get_client_info(client_id)
|
||||
if client_info is None:
|
||||
raise HTTPException(404)
|
||||
|
||||
@@ -976,13 +1020,13 @@ class OAuthClientManager:
|
||||
return await client.authorize_redirect(request, redirect_uri_str, **kwargs)
|
||||
|
||||
async def handle_callback(self, request, client_id: str, user_id: str, response):
|
||||
client = self.get_client(client_id) or self.ensure_client_from_config(client_id)
|
||||
client = await self.get_client(client_id)
|
||||
if client is None:
|
||||
raise HTTPException(404)
|
||||
|
||||
error_message = None
|
||||
try:
|
||||
client_info = self.get_client_info(client_id)
|
||||
client_info = await self.get_client_info(client_id)
|
||||
|
||||
# Note: Do NOT pass client_id/client_secret explicitly here.
|
||||
# The Authlib client already has these configured during add_client().
|
||||
@@ -1035,7 +1079,8 @@ class OAuthClientManager:
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
redirect_url = (str(request.app.state.config.WEBUI_URL or request.base_url)).rstrip('/')
|
||||
webui_url = await Config.get('webui.url')
|
||||
redirect_url = (str(webui_url or request.base_url)).rstrip('/')
|
||||
|
||||
if error_message:
|
||||
log.debug(error_message)
|
||||
@@ -1202,7 +1247,7 @@ class OAuthManager:
|
||||
if (
|
||||
hasattr(client, 'client_kwargs')
|
||||
and client.client_kwargs.get('scope')
|
||||
and auth_manager_config.OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE
|
||||
and auth_config.OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE
|
||||
):
|
||||
refresh_data['scope'] = client.client_kwargs['scope']
|
||||
|
||||
@@ -1235,6 +1280,7 @@ class OAuthManager:
|
||||
return None
|
||||
|
||||
async def get_user_role(self, user, user_data):
|
||||
auth_config = await get_oauth_runtime_config()
|
||||
user_count = await Users.get_num_users()
|
||||
if user and user_count == 1:
|
||||
# If the user is the only user, assign the role "admin" - actually repairs role for single user on login
|
||||
@@ -1246,16 +1292,16 @@ class OAuthManager:
|
||||
# default role here (not 'admin') — admin promotion happens
|
||||
# race-safely *after* insert via get_num_users() == 1.
|
||||
log.debug('First user bootstrap: using default role (admin promotion deferred to post-insert)')
|
||||
return auth_manager_config.DEFAULT_USER_ROLE
|
||||
return auth_config.DEFAULT_USER_ROLE
|
||||
|
||||
if auth_manager_config.ENABLE_OAUTH_ROLE_MANAGEMENT:
|
||||
if auth_config.ENABLE_OAUTH_ROLE_MANAGEMENT:
|
||||
log.debug('Running OAUTH Role management')
|
||||
oauth_claim = auth_manager_config.OAUTH_ROLES_CLAIM
|
||||
oauth_allowed_roles = auth_manager_config.OAUTH_ALLOWED_ROLES
|
||||
oauth_admin_roles = auth_manager_config.OAUTH_ADMIN_ROLES
|
||||
oauth_claim = auth_config.OAUTH_ROLES_CLAIM
|
||||
oauth_allowed_roles = auth_config.OAUTH_ALLOWED_ROLES
|
||||
oauth_admin_roles = auth_config.OAUTH_ADMIN_ROLES
|
||||
oauth_roles = []
|
||||
# Default/fallback role if no matching roles are found
|
||||
role = auth_manager_config.DEFAULT_USER_ROLE
|
||||
role = auth_config.DEFAULT_USER_ROLE
|
||||
|
||||
# Next block extracts the roles from the user data, accepting nested claims of any depth
|
||||
if oauth_claim and oauth_allowed_roles and oauth_admin_roles:
|
||||
@@ -1313,7 +1359,7 @@ class OAuthManager:
|
||||
else:
|
||||
if not user:
|
||||
# If role management is disabled, use the default role for new users
|
||||
role = auth_manager_config.DEFAULT_USER_ROLE
|
||||
role = auth_config.DEFAULT_USER_ROLE
|
||||
else:
|
||||
# If role management is disabled, use the existing role for existing users
|
||||
role = user.role
|
||||
@@ -1321,11 +1367,12 @@ class OAuthManager:
|
||||
return role
|
||||
|
||||
async def update_user_groups(self, user, user_data, default_permissions, db=None):
|
||||
auth_config = await get_oauth_runtime_config()
|
||||
log.debug('Running OAUTH Group management')
|
||||
oauth_claim = auth_manager_config.OAUTH_GROUPS_CLAIM
|
||||
oauth_claim = auth_config.OAUTH_GROUPS_CLAIM
|
||||
|
||||
try:
|
||||
blocked_groups = json.loads(auth_manager_config.OAUTH_BLOCKED_GROUPS)
|
||||
blocked_groups = json.loads(auth_config.OAUTH_BLOCKED_GROUPS)
|
||||
except Exception as e:
|
||||
log.exception(f'Error loading OAUTH_BLOCKED_GROUPS: {e}')
|
||||
blocked_groups = []
|
||||
@@ -1353,7 +1400,7 @@ class OAuthManager:
|
||||
all_available_groups: list[GroupModel] = await Groups.get_all_groups(db=db)
|
||||
|
||||
# Create groups if they don't exist and creation is enabled
|
||||
if auth_manager_config.ENABLE_OAUTH_GROUP_CREATION:
|
||||
if auth_config.ENABLE_OAUTH_GROUP_CREATION:
|
||||
log.debug('Checking for missing groups to create...')
|
||||
all_group_names = {g.name for g in all_available_groups}
|
||||
groups_created = False
|
||||
@@ -1370,7 +1417,7 @@ class OAuthManager:
|
||||
name=group_name,
|
||||
description=f"Group '{group_name}' created automatically via OAuth.",
|
||||
permissions=default_permissions, # Use default permissions from function args
|
||||
data={'config': {'share': auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE}},
|
||||
data={'config': {'share': auth_config.OAUTH_GROUP_DEFAULT_SHARE}},
|
||||
)
|
||||
# Use determined creator ID (admin or fallback to current user)
|
||||
created_group = await Groups.insert_new_group(creator_id, new_group_form, db=db)
|
||||
@@ -1496,6 +1543,7 @@ class OAuthManager:
|
||||
return '/user.png'
|
||||
|
||||
async def handle_login(self, request, provider):
|
||||
auth_config = await get_oauth_runtime_config()
|
||||
if provider not in OAUTH_PROVIDERS:
|
||||
raise HTTPException(404)
|
||||
# If the provider has a custom redirect URL, use that, otherwise automatically generate one
|
||||
@@ -1507,14 +1555,15 @@ class OAuthManager:
|
||||
)
|
||||
|
||||
kwargs = {}
|
||||
if auth_manager_config.OAUTH_AUDIENCE:
|
||||
kwargs['audience'] = auth_manager_config.OAUTH_AUDIENCE
|
||||
if auth_config.OAUTH_AUDIENCE:
|
||||
kwargs['audience'] = auth_config.OAUTH_AUDIENCE
|
||||
if OAUTH_AUTHORIZE_PARAMS:
|
||||
kwargs.update(OAUTH_AUTHORIZE_PARAMS)
|
||||
|
||||
return await client.authorize_redirect(request, redirect_uri, **kwargs)
|
||||
|
||||
async def handle_callback(self, request, provider, response, db=None):
|
||||
auth_config = await get_oauth_runtime_config()
|
||||
if provider not in OAUTH_PROVIDERS:
|
||||
raise HTTPException(404)
|
||||
|
||||
@@ -1568,8 +1617,8 @@ class OAuthManager:
|
||||
id_token_claims = dict(user_data) if user_data else {}
|
||||
if (
|
||||
(not user_data)
|
||||
or (auth_manager_config.OAUTH_EMAIL_CLAIM not in user_data)
|
||||
or (auth_manager_config.OAUTH_USERNAME_CLAIM not in user_data)
|
||||
or (auth_config.OAUTH_EMAIL_CLAIM not in user_data)
|
||||
or (auth_config.OAUTH_USERNAME_CLAIM not in user_data)
|
||||
):
|
||||
user_data: UserInfo = await client.userinfo(token=token)
|
||||
# Merge back ID token claims that the userinfo endpoint doesn't
|
||||
@@ -1585,8 +1634,8 @@ class OAuthManager:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
|
||||
|
||||
# Extract the "sub" claim, using custom claim if configured
|
||||
if auth_manager_config.OAUTH_SUB_CLAIM:
|
||||
sub = user_data.get(auth_manager_config.OAUTH_SUB_CLAIM)
|
||||
if auth_config.OAUTH_SUB_CLAIM:
|
||||
sub = user_data.get(auth_config.OAUTH_SUB_CLAIM)
|
||||
else:
|
||||
# Fallback to the default sub claim if not configured
|
||||
sub = user_data.get(OAUTH_PROVIDERS[provider].get('sub_claim', 'sub'))
|
||||
@@ -1600,7 +1649,7 @@ class OAuthManager:
|
||||
}
|
||||
|
||||
# Email extraction
|
||||
email_claim = auth_manager_config.OAUTH_EMAIL_CLAIM
|
||||
email_claim = auth_config.OAUTH_EMAIL_CLAIM
|
||||
email = user_data.get(email_claim, '')
|
||||
# We currently mandate that email addresses are provided
|
||||
if not email:
|
||||
@@ -1642,8 +1691,8 @@ class OAuthManager:
|
||||
email = email.lower()
|
||||
# If allowed domains are configured, check if the email domain is in the list
|
||||
if (
|
||||
'*' not in auth_manager_config.OAUTH_ALLOWED_DOMAINS
|
||||
and email.split('@')[-1] not in auth_manager_config.OAUTH_ALLOWED_DOMAINS
|
||||
'*' not in auth_config.OAUTH_ALLOWED_DOMAINS
|
||||
and email.split('@')[-1] not in auth_config.OAUTH_ALLOWED_DOMAINS
|
||||
):
|
||||
log.warning(f'OAuth callback failed, e-mail domain is not in the list of allowed domains: {user_data}')
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
|
||||
@@ -1652,7 +1701,7 @@ class OAuthManager:
|
||||
user = await Users.get_user_by_oauth_sub(provider, sub, db=db)
|
||||
if not user:
|
||||
# If the user does not exist, check if merging is enabled
|
||||
if auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL:
|
||||
if auth_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL:
|
||||
# Check if the user exists by email
|
||||
user = await Users.get_user_by_email(email, db=db)
|
||||
if user:
|
||||
@@ -1667,8 +1716,8 @@ class OAuthManager:
|
||||
# to avoid problems with the ENABLE_OAUTH_GROUP_MANAGEMENT check below
|
||||
user.role = determined_role
|
||||
|
||||
if auth_manager_config.OAUTH_UPDATE_NAME_ON_LOGIN:
|
||||
username_claim = auth_manager_config.OAUTH_USERNAME_CLAIM
|
||||
if auth_config.OAUTH_UPDATE_NAME_ON_LOGIN:
|
||||
username_claim = auth_config.OAUTH_USERNAME_CLAIM
|
||||
if username_claim:
|
||||
new_name = user_data.get(username_claim)
|
||||
if new_name and new_name != user.name:
|
||||
@@ -1676,8 +1725,8 @@ class OAuthManager:
|
||||
user.name = new_name
|
||||
log.debug(f'Updated name for user {user.email}')
|
||||
|
||||
if auth_manager_config.OAUTH_UPDATE_EMAIL_ON_LOGIN:
|
||||
email_claim = auth_manager_config.OAUTH_EMAIL_CLAIM
|
||||
if auth_config.OAUTH_UPDATE_EMAIL_ON_LOGIN:
|
||||
email_claim = auth_config.OAUTH_EMAIL_CLAIM
|
||||
if email_claim:
|
||||
new_email = user_data.get(email_claim)
|
||||
if new_email and new_email.lower() != user.email.lower():
|
||||
@@ -1692,8 +1741,8 @@ class OAuthManager:
|
||||
log.debug(f'Updated email for user {user.id}')
|
||||
|
||||
# Update profile picture if enabled and different from current
|
||||
if auth_manager_config.OAUTH_UPDATE_PICTURE_ON_LOGIN:
|
||||
picture_claim = auth_manager_config.OAUTH_PICTURE_CLAIM
|
||||
if auth_config.OAUTH_UPDATE_PICTURE_ON_LOGIN:
|
||||
picture_claim = auth_config.OAUTH_PICTURE_CLAIM
|
||||
if picture_claim:
|
||||
new_picture_url = user_data.get(
|
||||
picture_claim,
|
||||
@@ -1707,13 +1756,13 @@ class OAuthManager:
|
||||
log.debug(f'Updated profile picture for user {user.email}')
|
||||
else:
|
||||
# If the user does not exist, check if signups are enabled
|
||||
if auth_manager_config.ENABLE_OAUTH_SIGNUP:
|
||||
if auth_config.ENABLE_OAUTH_SIGNUP:
|
||||
# Check if an existing user with the same email already exists
|
||||
existing_user = await Users.get_user_by_email(email, db=db)
|
||||
if existing_user:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
|
||||
|
||||
picture_claim = auth_manager_config.OAUTH_PICTURE_CLAIM
|
||||
picture_claim = auth_config.OAUTH_PICTURE_CLAIM
|
||||
if picture_claim:
|
||||
picture_url = user_data.get(
|
||||
picture_claim,
|
||||
@@ -1722,7 +1771,7 @@ class OAuthManager:
|
||||
picture_url = await self._process_picture_url(picture_url, token.get('access_token'))
|
||||
else:
|
||||
picture_url = '/user.png'
|
||||
username_claim = auth_manager_config.OAUTH_USERNAME_CLAIM
|
||||
username_claim = auth_config.OAUTH_USERNAME_CLAIM
|
||||
|
||||
name = user_data.get(username_claim)
|
||||
if not name:
|
||||
@@ -1749,10 +1798,10 @@ class OAuthManager:
|
||||
await Users.update_user_role_by_id(user.id, 'admin', db=db)
|
||||
user = await Users.get_user_by_id(user.id, db=db)
|
||||
|
||||
if auth_manager_config.WEBHOOK_URL:
|
||||
if auth_config.WEBHOOK_URL:
|
||||
await post_webhook(
|
||||
WEBUI_NAME,
|
||||
auth_manager_config.WEBHOOK_URL,
|
||||
auth_config.WEBHOOK_URL,
|
||||
WEBHOOK_MESSAGES.USER_SIGNUP(user.name),
|
||||
{
|
||||
'action': 'signup',
|
||||
@@ -1761,7 +1810,8 @@ class OAuthManager:
|
||||
},
|
||||
)
|
||||
|
||||
await apply_default_group_assignment(request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db)
|
||||
default_group_id = await Config.get('ui.default_group_id')
|
||||
await apply_default_group_assignment(default_group_id, user.id, db=db)
|
||||
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -1771,13 +1821,13 @@ class OAuthManager:
|
||||
|
||||
jwt_token = create_token(
|
||||
data={'id': user.id},
|
||||
expires_delta=parse_duration(auth_manager_config.JWT_EXPIRES_IN),
|
||||
expires_delta=parse_duration(auth_config.JWT_EXPIRES_IN),
|
||||
)
|
||||
if auth_manager_config.ENABLE_OAUTH_GROUP_MANAGEMENT:
|
||||
if auth_config.ENABLE_OAUTH_GROUP_MANAGEMENT:
|
||||
await self.update_user_groups(
|
||||
user=user,
|
||||
user_data=user_data,
|
||||
default_permissions=request.app.state.config.USER_PERMISSIONS,
|
||||
default_permissions=await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
|
||||
@@ -1789,7 +1839,8 @@ class OAuthManager:
|
||||
else ERROR_MESSAGES.DEFAULT('Error during OAuth process')
|
||||
)
|
||||
|
||||
redirect_base_url = (str(request.app.state.config.WEBUI_URL or request.base_url)).rstrip('/')
|
||||
webui_url = await Config.get('webui.url')
|
||||
redirect_base_url = (str(webui_url or request.base_url)).rstrip('/')
|
||||
redirect_url = f'{redirect_base_url}/auth'
|
||||
|
||||
if error_message:
|
||||
@@ -1799,7 +1850,7 @@ class OAuthManager:
|
||||
response = RedirectResponse(url=redirect_url, headers=response.headers)
|
||||
|
||||
# Compute cookie expiry from JWT lifetime
|
||||
expires_delta = parse_duration(auth_manager_config.JWT_EXPIRES_IN)
|
||||
expires_delta = parse_duration(auth_config.JWT_EXPIRES_IN)
|
||||
cookie_max_age = int(expires_delta.total_seconds()) if expires_delta else None
|
||||
|
||||
# Set the cookie token
|
||||
|
||||
@@ -42,6 +42,7 @@ from open_webui.env import (
|
||||
REDIS_KEY_PREFIX,
|
||||
)
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.tools import Tools
|
||||
from open_webui.models.users import UserModel
|
||||
@@ -345,7 +346,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
|
||||
continue
|
||||
|
||||
tool_server_idx = tool_server_data.get('idx', 0)
|
||||
connections = request.app.state.config.TOOL_SERVER_CONNECTIONS
|
||||
connections = await Config.get('tool_server.connections', [])
|
||||
if tool_server_idx >= len(connections):
|
||||
log.warning(
|
||||
f'Tool server index {tool_server_idx} out of range '
|
||||
@@ -451,6 +452,16 @@ async def get_builtin_tools(
|
||||
|
||||
# Helper to check user-level feature permission (admins always pass)
|
||||
user = extra_params.get('__user__', {})
|
||||
config = await Config.get_many(
|
||||
'rag.web.search.enable',
|
||||
'image_generation.enable',
|
||||
'images.edit.enable',
|
||||
'code_interpreter.enable',
|
||||
'notes.enable',
|
||||
'channels.enable',
|
||||
'automations.enable',
|
||||
'calendar.enable',
|
||||
)
|
||||
|
||||
async def has_user_permission(feature_key: str) -> bool:
|
||||
if user.get('role') == 'admin':
|
||||
@@ -458,7 +469,7 @@ async def get_builtin_tools(
|
||||
return await has_permission(
|
||||
user.get('id', ''),
|
||||
f'features.{feature_key}',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
)
|
||||
|
||||
# Time utilities - available for date calculations
|
||||
@@ -533,7 +544,7 @@ async def get_builtin_tools(
|
||||
# Add web search tools if builtin category enabled AND enabled globally AND model has web_search capability
|
||||
if (
|
||||
is_builtin_tool_enabled('web_search')
|
||||
and getattr(request.app.state.config, 'ENABLE_WEB_SEARCH', False)
|
||||
and config.get('rag.web.search.enable')
|
||||
and get_model_capability('web_search')
|
||||
and features.get('web_search')
|
||||
and await has_user_permission('web_search')
|
||||
@@ -543,7 +554,7 @@ async def get_builtin_tools(
|
||||
# Add image generation/edit tools if builtin category enabled AND enabled globally AND model has image_generation capability
|
||||
if (
|
||||
is_builtin_tool_enabled('image_generation')
|
||||
and getattr(request.app.state.config, 'ENABLE_IMAGE_GENERATION', False)
|
||||
and config.get('image_generation.enable')
|
||||
and get_model_capability('image_generation')
|
||||
and features.get('image_generation')
|
||||
and await has_user_permission('image_generation')
|
||||
@@ -551,7 +562,7 @@ async def get_builtin_tools(
|
||||
builtin_functions.append(generate_image)
|
||||
if (
|
||||
is_builtin_tool_enabled('image_generation')
|
||||
and getattr(request.app.state.config, 'ENABLE_IMAGE_EDIT', False)
|
||||
and config.get('images.edit.enable')
|
||||
and get_model_capability('image_generation')
|
||||
and features.get('image_generation')
|
||||
and await has_user_permission('image_generation')
|
||||
@@ -561,7 +572,7 @@ async def get_builtin_tools(
|
||||
# Add code interpreter tool if builtin category enabled AND enabled globally AND model has code_interpreter capability
|
||||
if (
|
||||
is_builtin_tool_enabled('code_interpreter')
|
||||
and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True)
|
||||
and config.get('code_interpreter.enable')
|
||||
and get_model_capability('code_interpreter')
|
||||
and features.get('code_interpreter')
|
||||
and await has_user_permission('code_interpreter')
|
||||
@@ -571,7 +582,7 @@ async def get_builtin_tools(
|
||||
# Notes tools - search, view, create, and update user's notes
|
||||
if (
|
||||
is_builtin_tool_enabled('notes')
|
||||
and getattr(request.app.state.config, 'ENABLE_NOTES', False)
|
||||
and config.get('notes.enable')
|
||||
and await has_user_permission('notes')
|
||||
):
|
||||
builtin_functions.extend([search_notes, view_note, write_note, replace_note_content])
|
||||
@@ -579,7 +590,7 @@ async def get_builtin_tools(
|
||||
# Channels tools - search channels and messages
|
||||
if (
|
||||
is_builtin_tool_enabled('channels')
|
||||
and getattr(request.app.state.config, 'ENABLE_CHANNELS', False)
|
||||
and config.get('channels.enable')
|
||||
and await has_user_permission('channels')
|
||||
):
|
||||
builtin_functions.extend(
|
||||
@@ -602,7 +613,7 @@ async def get_builtin_tools(
|
||||
# Automation tools - create and manage scheduled automations from chat
|
||||
if (
|
||||
is_builtin_tool_enabled('automations')
|
||||
and getattr(request.app.state.config, 'ENABLE_AUTOMATIONS', False)
|
||||
and config.get('automations.enable')
|
||||
and await has_user_permission('automations')
|
||||
):
|
||||
builtin_functions.extend(
|
||||
@@ -612,7 +623,7 @@ async def get_builtin_tools(
|
||||
# Calendar tools - search/create/update/delete events
|
||||
if (
|
||||
is_builtin_tool_enabled('calendar')
|
||||
and getattr(request.app.state.config, 'ENABLE_CALENDAR', False)
|
||||
and config.get('calendar.enable')
|
||||
and await has_user_permission('calendar')
|
||||
):
|
||||
builtin_functions.extend(
|
||||
@@ -958,7 +969,7 @@ def convert_openapi_to_tool_payload(openapi_spec):
|
||||
|
||||
async def set_tool_servers(request: Request):
|
||||
try:
|
||||
request.app.state.TOOL_SERVERS = await get_tool_servers_data(request.app.state.config.TOOL_SERVER_CONNECTIONS)
|
||||
request.app.state.TOOL_SERVERS = await get_tool_servers_data(await Config.get('tool_server.connections', []))
|
||||
except Exception as e:
|
||||
log.error(f'Error fetching tool server data: {e}')
|
||||
request.app.state.TOOL_SERVERS = getattr(request.app.state, 'TOOL_SERVERS', None) or []
|
||||
@@ -1055,7 +1066,7 @@ async def get_terminal_system_prompt(
|
||||
|
||||
async def set_terminal_servers(request: Request):
|
||||
"""Load and cache OpenAPI specs from all TERMINAL_SERVER_CONNECTIONS."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
|
||||
# Build server configs compatible with get_tool_servers_data
|
||||
# Terminal connections store id/name at top level; translate to info dict
|
||||
@@ -1148,7 +1159,7 @@ async def get_terminal_tools(
|
||||
- Loads specs from cache
|
||||
- Builds callables that route through the terminal proxy
|
||||
"""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
connection = next((c for c in connections if c.get('id') == terminal_id), None)
|
||||
if connection is None:
|
||||
log.warning(f'Terminal server not found: {terminal_id}')
|
||||
|
||||
@@ -254,6 +254,61 @@ export const updateLdapServer = async (token: string = '', body: object) => {
|
||||
return res;
|
||||
};
|
||||
|
||||
export const getOAuthConfig = async (token: string) => {
|
||||
let error = null;
|
||||
|
||||
const res = await fetch(`${WEBUI_API_BASE_URL}/auths/admin/config/oauth`, {
|
||||
method: 'GET',
|
||||
headers: {
|
||||
'Content-Type': 'application/json',
|
||||
Authorization: `Bearer ${token}`
|
||||
}
|
||||
})
|
||||
.then(async (res) => {
|
||||
if (!res.ok) throw await res.json();
|
||||
return res.json();
|
||||
})
|
||||
.catch((err) => {
|
||||
console.error(err);
|
||||
error = err.detail;
|
||||
return null;
|
||||
});
|
||||
|
||||
if (error) {
|
||||
throw error;
|
||||
}
|
||||
|
||||
return res;
|
||||
};
|
||||
|
||||
export const updateOAuthConfig = async (token: string, body: object) => {
|
||||
let error = null;
|
||||
|
||||
const res = await fetch(`${WEBUI_API_BASE_URL}/auths/admin/config/oauth`, {
|
||||
method: 'POST',
|
||||
headers: {
|
||||
'Content-Type': 'application/json',
|
||||
Authorization: `Bearer ${token}`
|
||||
},
|
||||
body: JSON.stringify(body)
|
||||
})
|
||||
.then(async (res) => {
|
||||
if (!res.ok) throw await res.json();
|
||||
return res.json();
|
||||
})
|
||||
.catch((err) => {
|
||||
console.error(err);
|
||||
error = err.detail;
|
||||
return null;
|
||||
});
|
||||
|
||||
if (error) {
|
||||
throw error;
|
||||
}
|
||||
|
||||
return res;
|
||||
};
|
||||
|
||||
export const userSignIn = async (email: string, password: string) => {
|
||||
let error = null;
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
import { getBackendConfig } from '$lib/apis';
|
||||
import Database from './Settings/Database.svelte';
|
||||
|
||||
import Authentication from './Settings/Authentication.svelte';
|
||||
import General from './Settings/General.svelte';
|
||||
import Pipelines from './Settings/Pipelines.svelte';
|
||||
import Audio from './Settings/Audio.svelte';
|
||||
@@ -37,6 +38,7 @@
|
||||
const tabFromPath = pathParts[pathParts.length - 1];
|
||||
selectedTab = [
|
||||
'general',
|
||||
'authentication',
|
||||
'connections',
|
||||
'models',
|
||||
'evaluations',
|
||||
@@ -94,6 +96,24 @@
|
||||
'channels'
|
||||
]
|
||||
},
|
||||
{
|
||||
id: 'authentication',
|
||||
title: 'Authentication',
|
||||
route: '/admin/settings/authentication',
|
||||
keywords: [
|
||||
'authentication',
|
||||
'auth',
|
||||
'login',
|
||||
'signup',
|
||||
'ldap',
|
||||
'oauth',
|
||||
'oidc',
|
||||
'sso',
|
||||
'roles',
|
||||
'groups',
|
||||
'identity'
|
||||
]
|
||||
},
|
||||
{
|
||||
id: 'connections',
|
||||
title: 'Connections',
|
||||
@@ -308,6 +328,7 @@
|
||||
</div>
|
||||
|
||||
<!-- {$i18n.t('General')} -->
|
||||
<!-- {$i18n.t('Authentication')} -->
|
||||
<!-- {$i18n.t('Connections')} -->
|
||||
<!-- {$i18n.t('Models')} -->
|
||||
<!-- {$i18n.t('Evaluations')} -->
|
||||
@@ -344,6 +365,19 @@
|
||||
clip-rule="evenodd"
|
||||
/>
|
||||
</svg>
|
||||
{:else if tab.id === 'authentication'}
|
||||
<svg
|
||||
xmlns="http://www.w3.org/2000/svg"
|
||||
viewBox="0 0 16 16"
|
||||
fill="currentColor"
|
||||
class="w-4 h-4"
|
||||
>
|
||||
<path
|
||||
fill-rule="evenodd"
|
||||
d="M8 1.5A3.5 3.5 0 0 0 4.5 5v1H4a2 2 0 0 0-2 2v4.5a2 2 0 0 0 2 2h8a2 2 0 0 0 2-2V8a2 2 0 0 0-2-2h-.5V5A3.5 3.5 0 0 0 8 1.5ZM6 6V5a2 2 0 1 1 4 0v1H6Zm2 3a.75.75 0 0 0-.75.75v1a.75.75 0 0 0 1.5 0v-1A.75.75 0 0 0 8 9Z"
|
||||
clip-rule="evenodd"
|
||||
/>
|
||||
</svg>
|
||||
{:else if tab.id === 'connections'}
|
||||
<svg
|
||||
xmlns="http://www.w3.org/2000/svg"
|
||||
@@ -515,6 +549,8 @@
|
||||
await config.set(await getBackendConfig());
|
||||
}}
|
||||
/>
|
||||
{:else if selectedTab === 'authentication'}
|
||||
<Authentication />
|
||||
{:else if selectedTab === 'connections'}
|
||||
<Connections
|
||||
on:save={() => {
|
||||
|
||||
@@ -0,0 +1,776 @@
|
||||
<script lang="ts">
|
||||
import { getBackendConfig } from '$lib/apis';
|
||||
import {
|
||||
getAdminConfig,
|
||||
getLdapConfig,
|
||||
getLdapServer,
|
||||
getOAuthConfig,
|
||||
updateLdapConfig,
|
||||
updateLdapServer,
|
||||
updateOAuthConfig,
|
||||
updateAdminConfig
|
||||
} from '$lib/apis/auths';
|
||||
import { getGroups } from '$lib/apis/groups';
|
||||
import SensitiveInput from '$lib/components/common/SensitiveInput.svelte';
|
||||
import Switch from '$lib/components/common/Switch.svelte';
|
||||
import Textarea from '$lib/components/common/Textarea.svelte';
|
||||
import Tooltip from '$lib/components/common/Tooltip.svelte';
|
||||
import { config } from '$lib/stores';
|
||||
import { getContext, onMount } from 'svelte';
|
||||
import { toast } from 'svelte-sonner';
|
||||
|
||||
const i18n = getContext('i18n');
|
||||
|
||||
let adminConfig = null;
|
||||
let groups = [];
|
||||
|
||||
let ENABLE_LDAP = false;
|
||||
let LDAP_SERVER = {
|
||||
label: '',
|
||||
host: '',
|
||||
port: '',
|
||||
attribute_for_mail: 'mail',
|
||||
attribute_for_username: 'uid',
|
||||
app_dn: '',
|
||||
app_dn_password: '',
|
||||
search_base: '',
|
||||
search_filters: '',
|
||||
use_tls: false,
|
||||
certificate_path: '',
|
||||
ciphers: ''
|
||||
};
|
||||
|
||||
let oauthConfig: any = null;
|
||||
|
||||
const updateLdapServerHandler = async () => {
|
||||
await updateLdapConfig(localStorage.token, ENABLE_LDAP);
|
||||
if (!ENABLE_LDAP) return true;
|
||||
|
||||
const res = await updateLdapServer(localStorage.token, LDAP_SERVER).catch((error) => {
|
||||
toast.error(`${error}`);
|
||||
return null;
|
||||
});
|
||||
|
||||
return !!res;
|
||||
};
|
||||
|
||||
const updateOAuthHandler = async () => {
|
||||
if (!oauthConfig) return true;
|
||||
const res = await updateOAuthConfig(localStorage.token, oauthConfig).catch((error) => {
|
||||
toast.error(`${error}`);
|
||||
return null;
|
||||
});
|
||||
if (res) {
|
||||
oauthConfig = res;
|
||||
}
|
||||
return !!res;
|
||||
};
|
||||
|
||||
const updateAdminHandler = async () => {
|
||||
if (!adminConfig) return true;
|
||||
const res = await updateAdminConfig(localStorage.token, adminConfig).catch((error) => {
|
||||
toast.error(`${error}`);
|
||||
return null;
|
||||
});
|
||||
return !!res;
|
||||
};
|
||||
|
||||
const submitHandler = async () => {
|
||||
const adminSaved = await updateAdminHandler();
|
||||
const ldapSaved = await updateLdapServerHandler();
|
||||
const oauthSaved = await updateOAuthHandler();
|
||||
|
||||
if (adminSaved && ldapSaved && oauthSaved) {
|
||||
toast.success($i18n.t('Settings saved successfully!'));
|
||||
await config.set(await getBackendConfig());
|
||||
}
|
||||
};
|
||||
|
||||
onMount(async () => {
|
||||
await Promise.all([
|
||||
(async () => {
|
||||
adminConfig = await getAdminConfig(localStorage.token);
|
||||
})(),
|
||||
(async () => {
|
||||
groups = await getGroups(localStorage.token);
|
||||
})(),
|
||||
(async () => {
|
||||
LDAP_SERVER = await getLdapServer(localStorage.token);
|
||||
})(),
|
||||
(async () => {
|
||||
oauthConfig = await getOAuthConfig(localStorage.token).catch(() => null);
|
||||
})()
|
||||
]);
|
||||
|
||||
const ldapConfig = await getLdapConfig(localStorage.token);
|
||||
ENABLE_LDAP = ldapConfig.ENABLE_LDAP;
|
||||
});
|
||||
</script>
|
||||
|
||||
<form
|
||||
class="flex flex-col h-full justify-between space-y-3 text-sm"
|
||||
on:submit|preventDefault={submitHandler}
|
||||
>
|
||||
<div class="space-y-3 overflow-y-scroll scrollbar-hidden h-full">
|
||||
{#if adminConfig !== null}
|
||||
<div class="mb-3">
|
||||
<div class="mt-0.5 mb-2.5 text-base font-medium">{$i18n.t('User Access')}</div>
|
||||
|
||||
<hr class="border-gray-100/30 dark:border-gray-850/30 my-2" />
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Default User Role')}</div>
|
||||
<div class="flex items-center relative">
|
||||
<select
|
||||
class="w-fit pr-8 rounded-sm px-2 text-xs bg-transparent outline-hidden text-right"
|
||||
bind:value={adminConfig.DEFAULT_USER_ROLE}
|
||||
placeholder={$i18n.t('Select a role')}
|
||||
>
|
||||
<option value="pending">{$i18n.t('pending')}</option>
|
||||
<option value="user">{$i18n.t('user')}</option>
|
||||
<option value="admin">{$i18n.t('admin')}</option>
|
||||
</select>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Default Group')}</div>
|
||||
<div class="flex items-center relative">
|
||||
<select
|
||||
class="w-fit pr-8 rounded-sm px-2 text-xs bg-transparent outline-hidden text-right"
|
||||
bind:value={adminConfig.DEFAULT_GROUP_ID}
|
||||
placeholder={$i18n.t('Select a group')}
|
||||
>
|
||||
<option value={''}>None</option>
|
||||
{#each groups as group}
|
||||
<option value={group.id}>{group.name}</option>
|
||||
{/each}
|
||||
</select>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Enable New Sign Ups')}</div>
|
||||
|
||||
<Switch bind:state={adminConfig.ENABLE_SIGNUP} />
|
||||
</div>
|
||||
|
||||
<div class="mb-2.5 flex w-full justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Enable API Keys')}</div>
|
||||
|
||||
<Switch bind:state={adminConfig.ENABLE_API_KEYS} />
|
||||
</div>
|
||||
|
||||
{#if adminConfig?.ENABLE_API_KEYS}
|
||||
<div class="mb-2.5 flex w-full justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">
|
||||
{$i18n.t('API Key Endpoint Restrictions')}
|
||||
</div>
|
||||
|
||||
<Switch bind:state={adminConfig.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS} />
|
||||
</div>
|
||||
|
||||
{#if adminConfig?.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS}
|
||||
<div class=" flex w-full flex-col pr-2 mb-2.5">
|
||||
<div class=" text-xs font-medium">
|
||||
{$i18n.t('Allowed Endpoints')}
|
||||
</div>
|
||||
|
||||
<input
|
||||
class="w-full mt-1 text-sm dark:text-gray-300 bg-transparent outline-hidden"
|
||||
type="text"
|
||||
placeholder={`e.g.) /api/v1/messages, /api/v1/channels`}
|
||||
bind:value={adminConfig.API_KEYS_ALLOWED_ENDPOINTS}
|
||||
/>
|
||||
|
||||
<div class="mt-2 text-xs text-gray-400 dark:text-gray-500">
|
||||
<a
|
||||
href="https://docs.openwebui.com/reference/api-endpoints"
|
||||
target="_blank"
|
||||
class=" text-gray-300 font-medium underline"
|
||||
>
|
||||
{$i18n.t('To learn more about available endpoints, visit our documentation.')}
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
{/if}
|
||||
|
||||
<div class=" mb-2.5 w-full justify-between">
|
||||
<div class="flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('JWT Expiration')}</div>
|
||||
</div>
|
||||
|
||||
<div class="flex mt-2 space-x-2">
|
||||
<input
|
||||
class="w-full rounded-lg py-2 px-4 text-sm bg-gray-50 dark:text-gray-300 dark:bg-gray-850 outline-hidden"
|
||||
type="text"
|
||||
placeholder={`e.g.) "30m","1h", "10d". `}
|
||||
bind:value={adminConfig.JWT_EXPIRES_IN}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="mt-2 text-xs text-gray-400 dark:text-gray-500">
|
||||
{$i18n.t('Valid time units:')}
|
||||
<span class=" text-gray-300 font-medium"
|
||||
>{$i18n.t("'s', 'm', 'h', 'd', 'w' or '-1' for no expiration.")}</span
|
||||
>
|
||||
</div>
|
||||
|
||||
{#if adminConfig.JWT_EXPIRES_IN === '-1'}
|
||||
<div class="mt-2 text-xs">
|
||||
<div
|
||||
class=" bg-yellow-500/20 text-yellow-700 dark:text-yellow-200 rounded-lg px-3 py-2"
|
||||
>
|
||||
<div>
|
||||
<span class=" font-medium">{$i18n.t('Warning')}:</span>
|
||||
<span
|
||||
><a
|
||||
href="https://docs.openwebui.com/reference/env-configuration#jwt_expires_in"
|
||||
target="_blank"
|
||||
class=" underline"
|
||||
>{$i18n.t('No expiration can pose security risks.')}
|
||||
</a></span
|
||||
>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
<div class=" mt-0.5 mb-2.5 text-base font-medium">{$i18n.t('Pending Accounts')}</div>
|
||||
|
||||
<hr class=" border-gray-100/30 dark:border-gray-850/30 my-2" />
|
||||
|
||||
<div class="mb-2.5 flex w-full items-center justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">
|
||||
{$i18n.t('Show Admin Details in Account Pending Overlay')}
|
||||
</div>
|
||||
|
||||
<Switch bind:state={adminConfig.SHOW_ADMIN_DETAILS} />
|
||||
</div>
|
||||
|
||||
{#if adminConfig.SHOW_ADMIN_DETAILS}
|
||||
<div class="mb-2.5 w-full justify-between">
|
||||
<div class="flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Admin Contact Email')}</div>
|
||||
</div>
|
||||
|
||||
<div class="flex mt-2 space-x-2">
|
||||
<input
|
||||
class="w-full rounded-lg py-2 px-4 text-sm bg-gray-50 dark:text-gray-300 dark:bg-gray-850 outline-hidden"
|
||||
type="email"
|
||||
placeholder={$i18n.t('Leave empty to use first admin user')}
|
||||
bind:value={adminConfig.ADMIN_EMAIL}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
|
||||
<div class="mb-2.5">
|
||||
<div class=" self-center text-xs font-medium mb-2">
|
||||
{$i18n.t('Pending User Overlay Title')}
|
||||
</div>
|
||||
<Textarea
|
||||
placeholder={$i18n.t(
|
||||
'Enter a title for the pending user info overlay. Leave empty for default.'
|
||||
)}
|
||||
bind:value={adminConfig.PENDING_USER_OVERLAY_TITLE}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="mb-2.5">
|
||||
<div class=" self-center text-xs font-medium mb-2">
|
||||
{$i18n.t('Pending User Overlay Content')}
|
||||
</div>
|
||||
<Textarea
|
||||
placeholder={$i18n.t(
|
||||
'Enter content for the pending user info overlay. Leave empty for default.'
|
||||
)}
|
||||
bind:value={adminConfig.PENDING_USER_OVERLAY_CONTENT}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
|
||||
<div class=" space-y-3">
|
||||
<div class="mt-2 space-y-2 pr-1.5">
|
||||
<div class="flex justify-between items-center text-sm">
|
||||
<div class=" font-medium">{$i18n.t('LDAP')}</div>
|
||||
|
||||
<div class="mt-1">
|
||||
<Switch bind:state={ENABLE_LDAP} />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{#if ENABLE_LDAP}
|
||||
<div class="flex flex-col gap-1">
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Label')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Enter server label')}
|
||||
bind:value={LDAP_SERVER.label}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full"></div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Host')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Enter server host')}
|
||||
bind:value={LDAP_SERVER.host}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Port')}
|
||||
</div>
|
||||
<Tooltip
|
||||
placement="top-start"
|
||||
content={$i18n.t('Default to 389 or 636 if TLS is enabled')}
|
||||
className="w-full"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
type="number"
|
||||
placeholder={$i18n.t('Enter server port')}
|
||||
bind:value={LDAP_SERVER.port}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Application DN')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t('The Application Account DN you bind with for search')}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Enter Application DN')}
|
||||
bind:value={LDAP_SERVER.app_dn}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Application DN Password')}
|
||||
</div>
|
||||
<SensitiveInput
|
||||
placeholder={$i18n.t('Enter Application DN Password')}
|
||||
required={false}
|
||||
bind:value={LDAP_SERVER.app_dn_password}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Attribute for Mail')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t(
|
||||
'The LDAP attribute that maps to the mail that users use to sign in.'
|
||||
)}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Example: mail')}
|
||||
bind:value={LDAP_SERVER.attribute_for_mail}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Attribute for Username')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t(
|
||||
'The LDAP attribute that maps to the username that users use to sign in.'
|
||||
)}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Example: sAMAccountName or uid or userPrincipalName')}
|
||||
bind:value={LDAP_SERVER.attribute_for_username}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Search Base')}
|
||||
</div>
|
||||
<Tooltip content={$i18n.t('The base to search for users')} placement="top-start">
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Example: ou=users,dc=foo,dc=example')}
|
||||
bind:value={LDAP_SERVER.search_base}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Search Filters')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Example: (&(objectClass=inetOrgPerson)(uid=%s))')}
|
||||
bind:value={LDAP_SERVER.search_filters}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="text-xs text-gray-400 dark:text-gray-500">
|
||||
<a
|
||||
class=" text-gray-300 font-medium underline"
|
||||
href="https://ldap.com/ldap-filters/"
|
||||
target="_blank"
|
||||
>
|
||||
{$i18n.t('Click here for filter guides.')}
|
||||
</a>
|
||||
</div>
|
||||
<div>
|
||||
<div class="flex justify-between items-center text-sm">
|
||||
<div class=" font-medium">{$i18n.t('TLS')}</div>
|
||||
|
||||
<div class="mt-1">
|
||||
<Switch bind:state={LDAP_SERVER.use_tls} />
|
||||
</div>
|
||||
</div>
|
||||
{#if LDAP_SERVER.use_tls}
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1 mt-1">
|
||||
{$i18n.t('Certificate Path')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Enter certificate path')}
|
||||
bind:value={LDAP_SERVER.certificate_path}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex justify-between items-center text-xs">
|
||||
<div class=" font-medium">{$i18n.t('Validate certificate')}</div>
|
||||
|
||||
<div class="mt-1">
|
||||
<Switch bind:state={LDAP_SERVER.validate_cert} />
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Ciphers')}
|
||||
</div>
|
||||
<Tooltip content={$i18n.t('Default to ALL')} placement="top-start">
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Example: ALL')}
|
||||
bind:value={LDAP_SERVER.ciphers}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
<div class="w-full"></div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
</div>
|
||||
{#if oauthConfig}
|
||||
<div class="mb-3">
|
||||
<div class="mt-0.5 mb-2.5 text-base font-medium">{$i18n.t('OAuth / OIDC')}</div>
|
||||
|
||||
<hr class="border-gray-100/30 dark:border-gray-850/30 my-2" />
|
||||
|
||||
<div class="pr-1.5">
|
||||
<div class="space-y-3">
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Provider Name')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="SSO"
|
||||
bind:value={oauthConfig.OAUTH_PROVIDER_NAME}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Provider URL')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="https://accounts.google.com/.well-known/openid-configuration"
|
||||
bind:value={oauthConfig.OPENID_PROVIDER_URL}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Client ID')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Enter Client ID')}
|
||||
bind:value={oauthConfig.OAUTH_CLIENT_ID}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Client Secret')}
|
||||
</div>
|
||||
<SensitiveInput
|
||||
placeholder={$i18n.t('Enter Client Secret')}
|
||||
required={false}
|
||||
outerClassName="flex flex-1 bg-transparent"
|
||||
inputClassName="w-full text-sm py-0.5 bg-transparent"
|
||||
bind:value={oauthConfig.OAUTH_CLIENT_SECRET}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Redirect URI')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Enter Redirect URI')}
|
||||
bind:value={oauthConfig.OPENID_REDIRECT_URI}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Scopes')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="openid email profile"
|
||||
bind:value={oauthConfig.OAUTH_SCOPES}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Email Claim')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="email"
|
||||
bind:value={oauthConfig.OAUTH_EMAIL_CLAIM}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Username Claim')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="name"
|
||||
bind:value={oauthConfig.OAUTH_USERNAME_CLAIM}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Picture Claim')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="picture"
|
||||
bind:value={oauthConfig.OAUTH_PICTURE_CLAIM}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Sub Claim')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="sub"
|
||||
bind:value={oauthConfig.OAUTH_SUB_CLAIM}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Enable OAuth Signup')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.ENABLE_OAUTH_SIGNUP} />
|
||||
</div>
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Merge Accounts by Email')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.OAUTH_MERGE_ACCOUNTS_BY_EMAIL} />
|
||||
</div>
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Auto Redirect')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.OAUTH_AUTO_REDIRECT} />
|
||||
</div>
|
||||
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Allowed Domains')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="* (all domains)"
|
||||
bind:value={oauthConfig.OAUTH_ALLOWED_DOMAINS}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Enable Role Mapping')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.ENABLE_OAUTH_ROLE_MANAGEMENT} />
|
||||
</div>
|
||||
|
||||
{#if oauthConfig.ENABLE_OAUTH_ROLE_MANAGEMENT}
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Roles Claim')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="roles"
|
||||
bind:value={oauthConfig.OAUTH_ROLES_CLAIM}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Admin Roles')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="admin"
|
||||
bind:value={oauthConfig.OAUTH_ADMIN_ROLES}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Allowed Roles')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="*"
|
||||
bind:value={oauthConfig.OAUTH_ALLOWED_ROLES}
|
||||
/>
|
||||
</div>
|
||||
{/if}
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Enable Group Mapping')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.ENABLE_OAUTH_GROUP_MANAGEMENT} />
|
||||
</div>
|
||||
|
||||
{#if oauthConfig.ENABLE_OAUTH_GROUP_MANAGEMENT}
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Auto-Create Groups')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.ENABLE_OAUTH_GROUP_CREATION} />
|
||||
</div>
|
||||
|
||||
<div class="grid grid-cols-1 sm:grid-cols-2 gap-2">
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Group Claim')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder="groups"
|
||||
bind:value={oauthConfig.OAUTH_GROUP_CLAIM}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class="self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Blocked Groups')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Comma-separated group names')}
|
||||
bind:value={oauthConfig.OAUTH_BLOCKED_GROUPS}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Update Email')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.OAUTH_UPDATE_EMAIL_ON_LOGIN} />
|
||||
</div>
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Update Name')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.OAUTH_UPDATE_NAME_ON_LOGIN} />
|
||||
</div>
|
||||
|
||||
<div class="flex w-full justify-between pr-2">
|
||||
<div class="self-center text-xs font-medium">
|
||||
{$i18n.t('Update Picture')}
|
||||
</div>
|
||||
<Switch bind:state={oauthConfig.OAUTH_UPDATE_PICTURE_ON_LOGIN} />
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
|
||||
<div class="flex justify-end pt-3 text-sm font-medium">
|
||||
<button
|
||||
class="px-3.5 py-1.5 text-sm font-medium bg-black hover:bg-gray-900 text-white dark:bg-white dark:text-black dark:hover:bg-gray-100 transition rounded-full"
|
||||
type="submit"
|
||||
>
|
||||
{$i18n.t('Save')}
|
||||
</button>
|
||||
</div>
|
||||
</form>
|
||||
@@ -3,17 +3,8 @@
|
||||
import { v4 as uuidv4 } from 'uuid';
|
||||
|
||||
import { getBackendConfig, getVersionUpdates, getWebhookUrl, updateWebhookUrl } from '$lib/apis';
|
||||
import {
|
||||
getAdminConfig,
|
||||
getLdapConfig,
|
||||
getLdapServer,
|
||||
updateAdminConfig,
|
||||
updateLdapConfig,
|
||||
updateLdapServer
|
||||
} from '$lib/apis/auths';
|
||||
import { getAdminConfig, updateAdminConfig } from '$lib/apis/auths';
|
||||
import { getBanners, setBanners } from '$lib/apis/configs';
|
||||
import { getGroups } from '$lib/apis/groups';
|
||||
import SensitiveInput from '$lib/components/common/SensitiveInput.svelte';
|
||||
import Switch from '$lib/components/common/Switch.svelte';
|
||||
import Tooltip from '$lib/components/common/Tooltip.svelte';
|
||||
import { WEBUI_BUILD_HASH, WEBUI_VERSION } from '$lib/constants';
|
||||
@@ -37,27 +28,9 @@
|
||||
|
||||
let adminConfig = null;
|
||||
let webhookUrl = '';
|
||||
let groups = [];
|
||||
|
||||
let banners: Banner[] = [];
|
||||
|
||||
// LDAP
|
||||
let ENABLE_LDAP = false;
|
||||
let LDAP_SERVER = {
|
||||
label: '',
|
||||
host: '',
|
||||
port: '',
|
||||
attribute_for_mail: 'mail',
|
||||
attribute_for_username: 'uid',
|
||||
app_dn: '',
|
||||
app_dn_password: '',
|
||||
search_base: '',
|
||||
search_filters: '',
|
||||
use_tls: false,
|
||||
certificate_path: '',
|
||||
ciphers: ''
|
||||
};
|
||||
|
||||
const checkForVersionUpdates = async () => {
|
||||
updateAvailable = null;
|
||||
version = await getVersionUpdates(localStorage.token).catch((error) => {
|
||||
@@ -73,17 +46,6 @@
|
||||
console.info(updateAvailable);
|
||||
};
|
||||
|
||||
const updateLdapServerHandler = async () => {
|
||||
if (!ENABLE_LDAP) return;
|
||||
const res = await updateLdapServer(localStorage.token, LDAP_SERVER).catch((error) => {
|
||||
toast.error(`${error}`);
|
||||
return null;
|
||||
});
|
||||
if (res) {
|
||||
toast.success($i18n.t('LDAP server updated'));
|
||||
}
|
||||
};
|
||||
|
||||
const updateBanners = async () => {
|
||||
_banners.set(await setBanners(localStorage.token, banners));
|
||||
};
|
||||
@@ -91,8 +53,6 @@
|
||||
const updateHandler = async () => {
|
||||
webhookUrl = await updateWebhookUrl(localStorage.token, webhookUrl);
|
||||
const res = await updateAdminConfig(localStorage.token, adminConfig);
|
||||
await updateLdapConfig(localStorage.token, ENABLE_LDAP);
|
||||
await updateLdapServerHandler();
|
||||
|
||||
await updateBanners();
|
||||
|
||||
@@ -113,18 +73,9 @@
|
||||
|
||||
(async () => {
|
||||
webhookUrl = await getWebhookUrl(localStorage.token);
|
||||
})(),
|
||||
(async () => {
|
||||
LDAP_SERVER = await getLdapServer(localStorage.token);
|
||||
})(),
|
||||
(async () => {
|
||||
groups = await getGroups(localStorage.token);
|
||||
})()
|
||||
]);
|
||||
|
||||
const ldapConfig = await getLdapConfig(localStorage.token);
|
||||
ENABLE_LDAP = ldapConfig.ENABLE_LDAP;
|
||||
|
||||
banners = [...$_banners];
|
||||
});
|
||||
</script>
|
||||
@@ -296,395 +247,6 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="mb-3">
|
||||
<div class=" mt-0.5 mb-2.5 text-base font-medium">{$i18n.t('Authentication')}</div>
|
||||
|
||||
<hr class=" border-gray-100/30 dark:border-gray-850/30 my-2" />
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Default User Role')}</div>
|
||||
<div class="flex items-center relative">
|
||||
<select
|
||||
class="w-fit pr-8 rounded-sm px-2 text-xs bg-transparent outline-hidden text-right"
|
||||
bind:value={adminConfig.DEFAULT_USER_ROLE}
|
||||
placeholder={$i18n.t('Select a role')}
|
||||
>
|
||||
<option value="pending">{$i18n.t('pending')}</option>
|
||||
<option value="user">{$i18n.t('user')}</option>
|
||||
<option value="admin">{$i18n.t('admin')}</option>
|
||||
</select>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Default Group')}</div>
|
||||
<div class="flex items-center relative">
|
||||
<select
|
||||
class="w-fit pr-8 rounded-sm px-2 text-xs bg-transparent outline-hidden text-right"
|
||||
bind:value={adminConfig.DEFAULT_GROUP_ID}
|
||||
placeholder={$i18n.t('Select a group')}
|
||||
>
|
||||
<option value={''}>None</option>
|
||||
{#each groups as group}
|
||||
<option value={group.id}>{group.name}</option>
|
||||
{/each}
|
||||
</select>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Enable New Sign Ups')}</div>
|
||||
|
||||
<Switch bind:state={adminConfig.ENABLE_SIGNUP} />
|
||||
</div>
|
||||
|
||||
<div class="mb-2.5 flex w-full items-center justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">
|
||||
{$i18n.t('Show Admin Details in Account Pending Overlay')}
|
||||
</div>
|
||||
|
||||
<Switch bind:state={adminConfig.SHOW_ADMIN_DETAILS} />
|
||||
</div>
|
||||
|
||||
{#if adminConfig.SHOW_ADMIN_DETAILS}
|
||||
<div class="mb-2.5 w-full justify-between">
|
||||
<div class="flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Admin Contact Email')}</div>
|
||||
</div>
|
||||
|
||||
<div class="flex mt-2 space-x-2">
|
||||
<input
|
||||
class="w-full rounded-lg py-2 px-4 text-sm bg-gray-50 dark:text-gray-300 dark:bg-gray-850 outline-hidden"
|
||||
type="email"
|
||||
placeholder={$i18n.t('Leave empty to use first admin user')}
|
||||
bind:value={adminConfig.ADMIN_EMAIL}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
|
||||
<div class="mb-2.5">
|
||||
<div class=" self-center text-xs font-medium mb-2">
|
||||
{$i18n.t('Pending User Overlay Title')}
|
||||
</div>
|
||||
<Textarea
|
||||
placeholder={$i18n.t(
|
||||
'Enter a title for the pending user info overlay. Leave empty for default.'
|
||||
)}
|
||||
bind:value={adminConfig.PENDING_USER_OVERLAY_TITLE}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="mb-2.5">
|
||||
<div class=" self-center text-xs font-medium mb-2">
|
||||
{$i18n.t('Pending User Overlay Content')}
|
||||
</div>
|
||||
<Textarea
|
||||
placeholder={$i18n.t(
|
||||
'Enter content for the pending user info overlay. Leave empty for default.'
|
||||
)}
|
||||
bind:value={adminConfig.PENDING_USER_OVERLAY_CONTENT}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="mb-2.5 flex w-full justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Enable API Keys')}</div>
|
||||
|
||||
<Switch bind:state={adminConfig.ENABLE_API_KEYS} />
|
||||
</div>
|
||||
|
||||
{#if adminConfig?.ENABLE_API_KEYS}
|
||||
<div class="mb-2.5 flex w-full justify-between pr-2">
|
||||
<div class=" self-center text-xs font-medium">
|
||||
{$i18n.t('API Key Endpoint Restrictions')}
|
||||
</div>
|
||||
|
||||
<Switch bind:state={adminConfig.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS} />
|
||||
</div>
|
||||
|
||||
{#if adminConfig?.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS}
|
||||
<div class=" flex w-full flex-col pr-2 mb-2.5">
|
||||
<div class=" text-xs font-medium">
|
||||
{$i18n.t('Allowed Endpoints')}
|
||||
</div>
|
||||
|
||||
<input
|
||||
class="w-full mt-1 text-sm dark:text-gray-300 bg-transparent outline-hidden"
|
||||
type="text"
|
||||
placeholder={`e.g.) /api/v1/messages, /api/v1/channels`}
|
||||
bind:value={adminConfig.API_KEYS_ALLOWED_ENDPOINTS}
|
||||
/>
|
||||
|
||||
<div class="mt-2 text-xs text-gray-400 dark:text-gray-500">
|
||||
<a
|
||||
href="https://docs.openwebui.com/reference/api-endpoints"
|
||||
target="_blank"
|
||||
class=" text-gray-300 font-medium underline"
|
||||
>
|
||||
{$i18n.t('To learn more about available endpoints, visit our documentation.')}
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
{/if}
|
||||
|
||||
<div class=" mb-2.5 w-full justify-between">
|
||||
<div class="flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('JWT Expiration')}</div>
|
||||
</div>
|
||||
|
||||
<div class="flex mt-2 space-x-2">
|
||||
<input
|
||||
class="w-full rounded-lg py-2 px-4 text-sm bg-gray-50 dark:text-gray-300 dark:bg-gray-850 outline-hidden"
|
||||
type="text"
|
||||
placeholder={`e.g.) "30m","1h", "10d". `}
|
||||
bind:value={adminConfig.JWT_EXPIRES_IN}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div class="mt-2 text-xs text-gray-400 dark:text-gray-500">
|
||||
{$i18n.t('Valid time units:')}
|
||||
<span class=" text-gray-300 font-medium"
|
||||
>{$i18n.t("'s', 'm', 'h', 'd', 'w' or '-1' for no expiration.")}</span
|
||||
>
|
||||
</div>
|
||||
|
||||
{#if adminConfig.JWT_EXPIRES_IN === '-1'}
|
||||
<div class="mt-2 text-xs">
|
||||
<div
|
||||
class=" bg-yellow-500/20 text-yellow-700 dark:text-yellow-200 rounded-lg px-3 py-2"
|
||||
>
|
||||
<div>
|
||||
<span class=" font-medium">{$i18n.t('Warning')}:</span>
|
||||
<span
|
||||
><a
|
||||
href="https://docs.openwebui.com/reference/env-configuration#jwt_expires_in"
|
||||
target="_blank"
|
||||
class=" underline"
|
||||
>{$i18n.t('No expiration can pose security risks.')}
|
||||
</a></span
|
||||
>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
|
||||
<div class=" space-y-3">
|
||||
<div class="mt-2 space-y-2 pr-1.5">
|
||||
<div class="flex justify-between items-center text-sm">
|
||||
<div class=" font-medium">{$i18n.t('LDAP')}</div>
|
||||
|
||||
<div class="mt-1">
|
||||
<Switch bind:state={ENABLE_LDAP} />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{#if ENABLE_LDAP}
|
||||
<div class="flex flex-col gap-1">
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Label')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Enter server label')}
|
||||
bind:value={LDAP_SERVER.label}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full"></div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Host')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Enter server host')}
|
||||
bind:value={LDAP_SERVER.host}
|
||||
/>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Port')}
|
||||
</div>
|
||||
<Tooltip
|
||||
placement="top-start"
|
||||
content={$i18n.t('Default to 389 or 636 if TLS is enabled')}
|
||||
className="w-full"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
type="number"
|
||||
placeholder={$i18n.t('Enter server port')}
|
||||
bind:value={LDAP_SERVER.port}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Application DN')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t('The Application Account DN you bind with for search')}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Enter Application DN')}
|
||||
bind:value={LDAP_SERVER.app_dn}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Application DN Password')}
|
||||
</div>
|
||||
<SensitiveInput
|
||||
placeholder={$i18n.t('Enter Application DN Password')}
|
||||
required={false}
|
||||
bind:value={LDAP_SERVER.app_dn_password}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Attribute for Mail')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t(
|
||||
'The LDAP attribute that maps to the mail that users use to sign in.'
|
||||
)}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Example: mail')}
|
||||
bind:value={LDAP_SERVER.attribute_for_mail}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Attribute for Username')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t(
|
||||
'The LDAP attribute that maps to the username that users use to sign in.'
|
||||
)}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t(
|
||||
'Example: sAMAccountName or uid or userPrincipalName'
|
||||
)}
|
||||
bind:value={LDAP_SERVER.attribute_for_username}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Search Base')}
|
||||
</div>
|
||||
<Tooltip
|
||||
content={$i18n.t('The base to search for users')}
|
||||
placement="top-start"
|
||||
>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
required
|
||||
placeholder={$i18n.t('Example: ou=users,dc=foo,dc=example')}
|
||||
bind:value={LDAP_SERVER.search_base}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Search Filters')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Example: (&(objectClass=inetOrgPerson)(uid=%s))')}
|
||||
bind:value={LDAP_SERVER.search_filters}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="text-xs text-gray-400 dark:text-gray-500">
|
||||
<a
|
||||
class=" text-gray-300 font-medium underline"
|
||||
href="https://ldap.com/ldap-filters/"
|
||||
target="_blank"
|
||||
>
|
||||
{$i18n.t('Click here for filter guides.')}
|
||||
</a>
|
||||
</div>
|
||||
<div>
|
||||
<div class="flex justify-between items-center text-sm">
|
||||
<div class=" font-medium">{$i18n.t('TLS')}</div>
|
||||
|
||||
<div class="mt-1">
|
||||
<Switch bind:state={LDAP_SERVER.use_tls} />
|
||||
</div>
|
||||
</div>
|
||||
{#if LDAP_SERVER.use_tls}
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1 mt-1">
|
||||
{$i18n.t('Certificate Path')}
|
||||
</div>
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Enter certificate path')}
|
||||
bind:value={LDAP_SERVER.certificate_path}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex justify-between items-center text-xs">
|
||||
<div class=" font-medium">{$i18n.t('Validate certificate')}</div>
|
||||
|
||||
<div class="mt-1">
|
||||
<Switch bind:state={LDAP_SERVER.validate_cert} />
|
||||
</div>
|
||||
</div>
|
||||
<div class="flex w-full gap-2">
|
||||
<div class="w-full">
|
||||
<div class=" self-center text-xs font-medium min-w-fit mb-1">
|
||||
{$i18n.t('Ciphers')}
|
||||
</div>
|
||||
<Tooltip content={$i18n.t('Default to ALL')} placement="top-start">
|
||||
<input
|
||||
class="w-full bg-transparent outline-hidden py-0.5"
|
||||
placeholder={$i18n.t('Example: ALL')}
|
||||
bind:value={LDAP_SERVER.ciphers}
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
<div class="w-full"></div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="mb-3">
|
||||
<div class=" mt-0.5 mb-2.5 text-base font-medium">{$i18n.t('Features')}</div>
|
||||
|
||||
|
||||
@@ -1007,9 +1007,9 @@
|
||||
<div class=" self-center text-xs font-medium">
|
||||
{$i18n.t('User Webhooks')}
|
||||
</div>
|
||||
<Switch bind:state={permissions.features.user_webhooks} />
|
||||
<Switch bind:state={permissions.features.webhooks} />
|
||||
</div>
|
||||
{#if defaultPermissions?.features?.user_webhooks && !permissions.features.user_webhooks}
|
||||
{#if defaultPermissions?.features?.webhooks && !permissions.features.webhooks}
|
||||
<div>
|
||||
<div class="text-xs text-gray-500">
|
||||
{$i18n.t('This is a default user permission and will remain enabled.')}
|
||||
|
||||
@@ -227,7 +227,7 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{#if $config?.features?.enable_user_webhooks && ($user?.role === 'admin' || ($user?.permissions?.features?.user_webhooks ?? false))}
|
||||
{#if $config?.features?.enable_user_webhooks && ($user?.role === 'admin' || ($user?.permissions?.features?.webhooks ?? false))}
|
||||
<div class="mt-2">
|
||||
<div class="flex flex-col w-full">
|
||||
<div class=" mb-1 text-xs font-medium">{$i18n.t('Notification Webhook')}</div>
|
||||
|
||||
@@ -67,7 +67,7 @@ export const DEFAULT_PERMISSIONS = {
|
||||
memories: true,
|
||||
automations: false,
|
||||
calendar: true,
|
||||
user_webhooks: false
|
||||
webhooks: false
|
||||
},
|
||||
settings: {
|
||||
interface: true
|
||||
|
||||
Reference in New Issue
Block a user