diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 75b9359168..1ec947b4e0 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -1,3 +1,9 @@ + + # Pull Request Checklist ### Note to first-time contributors: Please open a discussion post in [Discussions](https://github.com/open-webui/open-webui/discussions) to discuss your idea/fix with the community before creating a pull request, and describe your changes before submitting a pull request. @@ -6,14 +12,16 @@ This is to ensure large feature PRs are discussed with the community first, befo **Before submitting, make sure you've checked the following:** -- [ ] **Target branch:** Verify that the pull request targets the `dev` branch. **Not targeting the `dev` branch will lead to immediate closure of the PR.** +- [ ] **Target branch:** Verify that the pull request targets the `dev` branch. **PRs targeting `main` will be immediately closed.** - [ ] **Description:** Provide a concise description of the changes made in this pull request down below. - [ ] **Changelog:** Ensure a changelog entry following the format of [Keep a Changelog](https://keepachangelog.com/) is added at the bottom of the PR description. -- [ ] **Documentation:** If necessary, update relevant documentation [Open WebUI Docs](https://github.com/open-webui/docs) like environment variables, the tutorials, or other documentation sources. -- [ ] **Dependencies:** Are there any new dependencies? Have you updated the dependency versions in the documentation? -- [ ] **Testing:** Perform manual tests to **verify the implemented fix/feature works as intended AND does not break any other functionality**. Take this as an opportunity to **make screenshots of the feature/fix and include it in the PR description**. +- [ ] **Documentation:** Add docs in [Open WebUI Docs Repository](https://github.com/open-webui/docs). Document user-facing behavior, environment variables, public APIs/interfaces, or deployment steps. +- [ ] **Dependencies:** Are there any new or upgraded dependencies? If so, explain why, update the changelog/docs, and include any compatibility notes. Actually run the code/function that uses updated library to ensure it doesn't crash. +- [ ] **Testing:** Perform manual tests to **verify the implemented fix/feature works as intended AND does not break any other functionality**. Include reproducible steps to demonstrate the issue before the fix. Test edge cases (URL encoding, HTML entities, types). Take this as an opportunity to **make screenshots of the feature/fix and include them in the PR description**. - [ ] **Agentic AI Code:** Confirm this Pull Request is **not written by any AI Agent** or has at least **gone through additional human review AND manual testing**. If any AI Agent is the co-author of this PR, it may lead to immediate closure of the PR. - [ ] **Code review:** Have you performed a self-review of your code, addressing any coding standard issues and ensuring adherence to the project's coding standards? +- [ ] **Design & Architecture:** Prefer smart defaults over adding new settings; use local state for ephemeral UI logic. Open a Discussion for major architectural or UX changes. +- [ ] **Git Hygiene:** Keep PRs atomic (one logical change). Clean up commits and rebase on `dev` to ensure no unrelated commits (e.g. from `main`) are included. Push updates to the existing PR branch instead of closing and reopening. - [ ] **Title Prefix:** To clearly categorize this pull request, prefix the pull request title using one of the following: - **BREAKING CHANGE**: Significant changes that may affect compatibility - **build**: Changes that affect the build system or external dependencies @@ -76,7 +84,13 @@ This is to ensure large feature PRs are discussed with the community first, befo ### Contributor License Agreement + + By submitting this pull request, I confirm that I have read and fully agree to the [Contributor License Agreement (CLA)](https://github.com/open-webui/open-webui/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT), and I am providing my contributions under its terms. > [!NOTE] -> Deleting the CLA section will lead to immediate closure of your PR and it will not be merged in. +> Deleting the CLA section will lead to immediate closure of your PR and it will not be merged in. \ No newline at end of file diff --git a/Dockerfile b/Dockerfile index ac81a3943c..aa0bbf6f04 100644 --- a/Dockerfile +++ b/Dockerfile @@ -128,7 +128,7 @@ RUN apt-get update && \ apt-get install -y --no-install-recommends \ git build-essential pandoc gcc netcat-openbsd curl jq \ python3-dev \ - ffmpeg libsm6 libxext6 \ + ffmpeg libsm6 libxext6 zstd \ && rm -rf /var/lib/apt/lists/* # install python dependencies diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 49fbb84dea..675176ee54 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -257,7 +257,7 @@ class AppConfig: self._state[key].value = value self._state[key].save() - if self._redis: + if self._redis and ENABLE_PERSISTENT_CONFIG: redis_key = f"{self._redis_key_prefix}:config:{key}" self._redis.set(redis_key, json.dumps(self._state[key].value)) @@ -265,8 +265,8 @@ class AppConfig: if key not in self._state: raise AttributeError(f"Config key '{key}' not found") - # If Redis is available, check for an updated value - if self._redis: + # If Redis is available and persistent config is enabled, check for an updated value + if self._redis and ENABLE_PERSISTENT_CONFIG: redis_key = f"{self._redis_key_prefix}:config:{key}" redis_value = self._redis.get(redis_key) @@ -2246,9 +2246,13 @@ ENABLE_QDRANT_MULTITENANCY_MODE = ( QDRANT_COLLECTION_PREFIX = os.environ.get("QDRANT_COLLECTION_PREFIX", "open-webui") WEAVIATE_HTTP_HOST = os.environ.get("WEAVIATE_HTTP_HOST", "") +WEAVIATE_GRPC_HOST = os.environ.get("WEAVIATE_GRPC_HOST", "") WEAVIATE_HTTP_PORT = int(os.environ.get("WEAVIATE_HTTP_PORT", "8080")) WEAVIATE_GRPC_PORT = int(os.environ.get("WEAVIATE_GRPC_PORT", "50051")) WEAVIATE_API_KEY = os.environ.get("WEAVIATE_API_KEY") +WEAVIATE_HTTP_SECURE = os.environ.get("WEAVIATE_HTTP_SECURE", "false").lower() == "true" +WEAVIATE_GRPC_SECURE = os.environ.get("WEAVIATE_GRPC_SECURE", "false").lower() == "true" +WEAVIATE_SKIP_INIT_CHECKS = os.environ.get("WEAVIATE_SKIP_INIT_CHECKS", "false").lower() == "true" # OpenSearch OPENSEARCH_URI = os.environ.get("OPENSEARCH_URI", "https://localhost:9200") @@ -2817,6 +2821,12 @@ PDF_EXTRACT_IMAGES = PersistentConfig( os.environ.get("PDF_EXTRACT_IMAGES", "False").lower() == "true", ) +PDF_LOADER_MODE = PersistentConfig( + "PDF_LOADER_MODE", + "rag.pdf_loader_mode", + os.environ.get("PDF_LOADER_MODE", "page"), +) + RAG_EMBEDDING_MODEL = PersistentConfig( "RAG_EMBEDDING_MODEL", "rag.embedding_model", @@ -3412,6 +3422,24 @@ EXTERNAL_WEB_LOADER_API_KEY = PersistentConfig( os.environ.get("EXTERNAL_WEB_LOADER_API_KEY", ""), ) +YANDEX_WEB_SEARCH_URL = PersistentConfig( + "YANDEX_WEB_SEARCH_URL", + "rag.web.search.yandex_web_search_url", + os.environ.get("YANDEX_WEB_SEARCH_URL", ""), +) + +YANDEX_WEB_SEARCH_API_KEY = PersistentConfig( + "YANDEX_WEB_SEARCH_API_KEY", + "rag.web.search.yandex_web_search_api_key", + os.environ.get("YANDEX_WEB_SEARCH_API_KEY", ""), +) + +YANDEX_WEB_SEARCH_CONFIG = PersistentConfig( + "YANDEX_WEB_SEARCH_CONFIG", + "rag.web.search.yandex_web_search_config", + os.environ.get("YANDEX_WEB_SEARCH_CONFIG", ""), +) + #################################### # Images #################################### diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index ff48a3abfe..3fc75c72bb 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -199,6 +199,8 @@ ENABLE_STAR_SESSIONS_MIDDLEWARE = ( os.environ.get("ENABLE_STAR_SESSIONS_MIDDLEWARE", "False").lower() == "true" ) +ENABLE_EASTER_EGGS = os.environ.get("ENABLE_EASTER_EGGS", "True").lower() == "true" + #################################### # WEBUI_BUILD_HASH #################################### @@ -390,6 +392,22 @@ try: REDIS_SOCKET_CONNECT_TIMEOUT = float(REDIS_SOCKET_CONNECT_TIMEOUT) except ValueError: REDIS_SOCKET_CONNECT_TIMEOUT = None + +REDIS_RECONNECT_DELAY = os.environ.get( + "REDIS_RECONNECT_DELAY", "" +) + +if REDIS_RECONNECT_DELAY == "": + REDIS_RECONNECT_DELAY = None +else: + try: + REDIS_RECONNECT_DELAY = float( + REDIS_RECONNECT_DELAY + ) + if REDIS_RECONNECT_DELAY < 0: + REDIS_RECONNECT_DELAY = None + except Exception: + REDIS_RECONNECT_DELAY = None #################################### # UVICORN WORKERS @@ -455,6 +473,8 @@ except Exception as e: r"^(?=.*[a-z])(?=.*[A-Z])(?=.*\d)(?=.*[^\w\s]).{8,}$" ) +PASSWORD_VALIDATION_HINT = os.environ.get("PASSWORD_VALIDATION_HINT", "") + BYPASS_MODEL_ACCESS_CONTROL = ( os.environ.get("BYPASS_MODEL_ACCESS_CONTROL", "False").lower() == "true" @@ -519,6 +539,12 @@ OAUTH_SESSION_TOKEN_ENCRYPTION_KEY = os.environ.get( "OAUTH_SESSION_TOKEN_ENCRYPTION_KEY", WEBUI_SECRET_KEY ) +# Token Exchange Configuration +# Allows external apps to exchange OAuth tokens for OpenWebUI tokens +ENABLE_OAUTH_TOKEN_EXCHANGE = ( + os.environ.get("ENABLE_OAUTH_TOKEN_EXCHANGE", "False").lower() == "true" +) + #################################### # SCIM Configuration #################################### @@ -673,7 +699,11 @@ WEBSOCKET_SERVER_LOGGING = ( os.environ.get("WEBSOCKET_SERVER_LOGGING", "False").lower() == "true" ) WEBSOCKET_SERVER_ENGINEIO_LOGGING = ( - os.environ.get("WEBSOCKET_SERVER_LOGGING", "False").lower() == "true" + os.environ.get( + "WEBSOCKET_SERVER_ENGINEIO_LOGGING", + os.environ.get("WEBSOCKET_SERVER_LOGGING", "False"), + ).lower() + == "true" ) WEBSOCKET_SERVER_PING_TIMEOUT = os.environ.get("WEBSOCKET_SERVER_PING_TIMEOUT", "20") try: diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index bc6867491c..747859db88 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -290,6 +290,7 @@ from open_webui.config import ( ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER, TIKTOKEN_ENCODING_NAME, PDF_EXTRACT_IMAGES, + PDF_LOADER_MODE, YOUTUBE_LOADER_LANGUAGE, YOUTUBE_LOADER_PROXY_URL, # Retrieval (Web Search) @@ -354,6 +355,9 @@ from open_webui.config import ( EXTERNAL_WEB_SEARCH_API_KEY, EXTERNAL_WEB_LOADER_URL, EXTERNAL_WEB_LOADER_API_KEY, + YANDEX_WEB_SEARCH_URL, + YANDEX_WEB_SEARCH_API_KEY, + YANDEX_WEB_SEARCH_CONFIG, # WebUI WEBUI_AUTH, WEBUI_NAME, @@ -492,6 +496,7 @@ from open_webui.env import ( WEBUI_ADMIN_EMAIL, WEBUI_ADMIN_PASSWORD, WEBUI_ADMIN_NAME, + ENABLE_EASTER_EGGS, ) @@ -948,6 +953,7 @@ app.state.config.RAG_OLLAMA_BASE_URL = RAG_OLLAMA_BASE_URL app.state.config.RAG_OLLAMA_API_KEY = RAG_OLLAMA_API_KEY app.state.config.PDF_EXTRACT_IMAGES = PDF_EXTRACT_IMAGES +app.state.config.PDF_LOADER_MODE = PDF_LOADER_MODE app.state.config.YOUTUBE_LOADER_LANGUAGE = YOUTUBE_LOADER_LANGUAGE app.state.config.YOUTUBE_LOADER_PROXY_URL = YOUTUBE_LOADER_PROXY_URL @@ -1009,6 +1015,9 @@ app.state.config.EXTERNAL_WEB_SEARCH_URL = EXTERNAL_WEB_SEARCH_URL app.state.config.EXTERNAL_WEB_SEARCH_API_KEY = EXTERNAL_WEB_SEARCH_API_KEY app.state.config.EXTERNAL_WEB_LOADER_URL = EXTERNAL_WEB_LOADER_URL app.state.config.EXTERNAL_WEB_LOADER_API_KEY = EXTERNAL_WEB_LOADER_API_KEY +app.state.config.YANDEX_WEB_SEARCH_URL = YANDEX_WEB_SEARCH_URL +app.state.config.YANDEX_WEB_SEARCH_API_KEY = YANDEX_WEB_SEARCH_API_KEY +app.state.config.YANDEX_WEB_SEARCH_CONFIG = YANDEX_WEB_SEARCH_CONFIG app.state.config.PLAYWRIGHT_WS_URL = PLAYWRIGHT_WS_URL @@ -1370,6 +1379,13 @@ async def check_url(request: Request, call_next): request.state.token = get_http_authorization_cred( request.headers.get("Authorization") ) + # Fallback to cookie token for browser sessions + if request.state.token is None and request.cookies.get("token"): + from fastapi.security import HTTPAuthorizationCredentials + request.state.token = HTTPAuthorizationCredentials( + scheme="Bearer", + credentials=request.cookies.get("token") + ) request.state.enable_api_keys = app.state.config.ENABLE_API_KEYS response = await call_next(request) @@ -1934,6 +1950,7 @@ async def get_app_config(request: Request): "enable_websocket": ENABLE_WEBSOCKET_SUPPORT, "enable_version_update_check": ENABLE_VERSION_UPDATE_CHECK, "enable_public_active_users_count": ENABLE_PUBLIC_ACTIVE_USERS_COUNT, + "enable_easter_eggs": ENABLE_EASTER_EGGS, **( { "enable_direct_connections": app.state.config.ENABLE_DIRECT_CONNECTIONS, diff --git a/backend/open_webui/migrations/versions/374d2f66af06_add_prompt_history_table.py b/backend/open_webui/migrations/versions/374d2f66af06_add_prompt_history_table.py new file mode 100644 index 0000000000..c61196fcb0 --- /dev/null +++ b/backend/open_webui/migrations/versions/374d2f66af06_add_prompt_history_table.py @@ -0,0 +1,248 @@ +"""Add prompt history table + +Revision ID: 374d2f66af06 +Revises: c440947495f3 +Create Date: 2026-01-23 17:15:00.000000 + +""" + +from typing import Sequence, Union +import uuid + +from alembic import op +import sqlalchemy as sa + + +revision: str = "374d2f66af06" +down_revision: Union[str, None] = "c440947495f3" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + conn = op.get_bind() + + # Step 1: Read existing data from OLD table (schema likely command as PK) + # We use batch_alter previously, but we want to move to new table. + # We need to assume the OLD structure. + + old_prompt_table = sa.table( + "prompt", + sa.column("command", sa.Text()), + sa.column("user_id", sa.Text()), + sa.column("title", sa.Text()), + sa.column("content", sa.Text()), + sa.column("timestamp", sa.BigInteger()), + sa.column("access_control", sa.JSON()), + ) + + # Check if table exists/read data + try: + existing_prompts = conn.execute( + sa.select( + old_prompt_table.c.command, + old_prompt_table.c.user_id, + old_prompt_table.c.title, + old_prompt_table.c.content, + old_prompt_table.c.timestamp, + old_prompt_table.c.access_control, + ) + ).fetchall() + except Exception: + # Fallback if table doesn't exist (new install) + existing_prompts = [] + + # Step 2: Create new prompt table with 'id' as PRIMARY KEY + op.create_table( + "prompt_new", + sa.Column("id", sa.Text(), primary_key=True), + sa.Column("command", sa.String(), unique=True, index=True), + sa.Column("user_id", sa.String(), nullable=False), + sa.Column("name", sa.Text(), nullable=False), + sa.Column("content", sa.Text(), nullable=False), + sa.Column("data", sa.JSON(), nullable=True), + sa.Column("meta", sa.JSON(), nullable=True), + sa.Column("access_control", sa.JSON(), nullable=True), + sa.Column("is_active", sa.Boolean(), nullable=False, server_default="1"), + sa.Column("version_id", sa.Text(), nullable=True), + sa.Column("tags", sa.JSON(), nullable=True), + sa.Column("created_at", sa.BigInteger(), nullable=False), + sa.Column("updated_at", sa.BigInteger(), nullable=False), + ) + + # Step 3: Create prompt_history table + op.create_table( + "prompt_history", + sa.Column("id", sa.Text(), primary_key=True), + sa.Column("prompt_id", sa.Text(), nullable=False, index=True), + sa.Column("parent_id", sa.Text(), nullable=True), + sa.Column("snapshot", sa.JSON(), nullable=False), + sa.Column("user_id", sa.Text(), nullable=False), + sa.Column("commit_message", sa.Text(), nullable=True), + sa.Column("created_at", sa.BigInteger(), nullable=False), + ) + + # Step 4: Migrate data + prompt_new_table = sa.table( + "prompt_new", + sa.column("id", sa.Text()), + sa.column("command", sa.String()), + sa.column("user_id", sa.String()), + sa.column("name", sa.Text()), + sa.column("content", sa.Text()), + sa.column("data", sa.JSON()), + sa.column("meta", sa.JSON()), + sa.column("access_control", sa.JSON()), + sa.column("is_active", sa.Boolean()), + sa.column("version_id", sa.Text()), + sa.column("tags", sa.JSON()), + sa.column("created_at", sa.BigInteger()), + sa.column("updated_at", sa.BigInteger()), + ) + + prompt_history_table = sa.table( + "prompt_history", + sa.column("id", sa.Text()), + sa.column("prompt_id", sa.Text()), + sa.column("parent_id", sa.Text()), + sa.column("snapshot", sa.JSON()), + sa.column("user_id", sa.Text()), + sa.column("commit_message", sa.Text()), + sa.column("created_at", sa.BigInteger()), + ) + + for row in existing_prompts: + command = row[0] + user_id = row[1] + title = row[2] + content = row[3] + timestamp = row[4] + access_control = row[5] + + new_uuid = str(uuid.uuid4()) + history_uuid = str(uuid.uuid4()) + clean_command = command[1:] if command and command.startswith("/") else command + + # Insert into prompt_new + conn.execute( + sa.insert(prompt_new_table).values( + id=new_uuid, + command=clean_command, + user_id=user_id, + name=title, + content=content, + data={}, + meta={}, + access_control=access_control, + is_active=True, + version_id=history_uuid, + tags=[], + created_at=timestamp, + updated_at=timestamp, + ) + ) + + # Create initial history entry + conn.execute( + sa.insert(prompt_history_table).values( + id=history_uuid, + prompt_id=new_uuid, + parent_id=None, + snapshot={ + "name": title, + "content": content, + "command": clean_command, + "data": {}, + "meta": {}, + "access_control": access_control, + }, + user_id=user_id, + commit_message=None, + created_at=timestamp, + ) + ) + + # Step 5: Replace old table with new one + op.drop_table("prompt") + op.rename_table("prompt_new", "prompt") + + +def downgrade() -> None: + conn = op.get_bind() + + # Step 1: Read new data + prompt_table = sa.table( + "prompt", + sa.column("command", sa.String()), + sa.column("name", sa.Text()), + sa.column("created_at", sa.BigInteger()), + sa.column("user_id", sa.Text()), + sa.column("content", sa.Text()), + sa.column("access_control", sa.JSON()), + ) + + try: + current_data = conn.execute( + sa.select( + prompt_table.c.command, + prompt_table.c.name, + prompt_table.c.created_at, + prompt_table.c.user_id, + prompt_table.c.content, + prompt_table.c.access_control, + ) + ).fetchall() + except Exception: + current_data = [] + + # Step 2: Drop history and table + op.drop_table("prompt_history") + op.drop_table("prompt") + + # Step 3: Recreate old table (command as PK?) + # Assuming old schema: + op.create_table( + "prompt", + sa.Column("command", sa.String(), primary_key=True), + sa.Column("user_id", sa.String()), + sa.Column("title", sa.Text()), + sa.Column("content", sa.Text()), + sa.Column("timestamp", sa.BigInteger()), + sa.Column("access_control", sa.JSON()), + sa.Column("id", sa.Integer(), nullable=True), + ) + + # Step 4: Restore data + old_prompt_table = sa.table( + "prompt", + sa.column("command", sa.String()), + sa.column("user_id", sa.String()), + sa.column("title", sa.Text()), + sa.column("content", sa.Text()), + sa.column("timestamp", sa.BigInteger()), + sa.column("access_control", sa.JSON()), + ) + + for row in current_data: + command = row[0] + name = row[1] + created_at = row[2] + user_id = row[3] + content = row[4] + access_control = row[5] + + # Restore leading / + old_command = ( + "/" + command if command and not command.startswith("/") else command + ) + + conn.execute( + sa.insert(old_prompt_table).values( + command=old_command, + user_id=user_id, + title=name, + content=content, + timestamp=created_at, + access_control=access_control, + ) + ) diff --git a/backend/open_webui/models/auths.py b/backend/open_webui/models/auths.py index 93f17dff11..2d795cd2b1 100644 --- a/backend/open_webui/models/auths.py +++ b/backend/open_webui/models/auths.py @@ -4,7 +4,7 @@ from typing import Optional from sqlalchemy.orm import Session from open_webui.internal.db import Base, JSONField, get_db, get_db_context -from open_webui.models.users import UserModel, UserProfileImageResponse, Users +from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users from pydantic import BaseModel from sqlalchemy import Boolean, Column, String, Text @@ -155,10 +155,17 @@ class AuthsTable: log.info(f"authenticate_user_by_email: {email}") try: with get_db_context(db) as db: - auth = db.query(Auth).filter_by(email=email, active=True).first() - if auth: - user = Users.get_user_by_id(auth.id, db=db) - return user + # Single JOIN query instead of two separate queries + result = ( + db.query(Auth, User) + .join(User, Auth.id == User.id) + .filter(Auth.email == email, Auth.active == True) + .first() + ) + if result: + _, user = result + return UserModel.model_validate(user) + return None except Exception: return None diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index 12359eec9f..eb0763048b 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -168,6 +168,14 @@ class ChatTitleIdResponse(BaseModel): created_at: int +class SharedChatResponse(BaseModel): + id: str + title: str + share_id: Optional[str] = None + updated_at: int + created_at: int + + class ChatListResponse(BaseModel): items: list[ChatModel] total: int @@ -675,6 +683,49 @@ class ChatTable: all_chats = query.all() return [ChatModel.model_validate(chat) for chat in all_chats] + def get_shared_chat_list_by_user_id( + self, + user_id: str, + filter: Optional[dict] = None, + skip: int = 0, + limit: int = 50, + db: Optional[Session] = None, + ) -> list[ChatModel]: + + with get_db_context(db) as db: + query = db.query(Chat).filter_by(user_id=user_id).filter( + Chat.share_id.isnot(None) + ) + + if filter: + query_key = filter.get("query") + if query_key: + query = query.filter(Chat.title.ilike(f"%{query_key}%")) + + order_by = filter.get("order_by") + direction = filter.get("direction") + + if order_by and direction: + if not getattr(Chat, order_by, None): + raise ValueError("Invalid order_by field") + + if direction.lower() == "asc": + query = query.order_by(getattr(Chat, order_by).asc()) + elif direction.lower() == "desc": + query = query.order_by(getattr(Chat, order_by).desc()) + else: + raise ValueError("Invalid direction for ordering") + else: + query = query.order_by(Chat.updated_at.desc()) + + if skip: + query = query.offset(skip) + if limit: + query = query.limit(limit) + + all_chats = query.all() + return [ChatModel.model_validate(chat) for chat in all_chats] + def get_chat_list_by_user_id( self, user_id: str, diff --git a/backend/open_webui/models/feedbacks.py b/backend/open_webui/models/feedbacks.py index 048c10f85c..2a4d752f59 100644 --- a/backend/open_webui/models/feedbacks.py +++ b/backend/open_webui/models/feedbacks.py @@ -460,23 +460,15 @@ class FeedbackTable: self, user_id: str, db: Optional[Session] = None ) -> bool: with get_db_context(db) as db: - feedbacks = db.query(Feedback).filter_by(user_id=user_id).all() - if not feedbacks: - return False - for feedback in feedbacks: - db.delete(feedback) + result = db.query(Feedback).filter_by(user_id=user_id).delete() db.commit() - return True + return result > 0 def delete_all_feedbacks(self, db: Optional[Session] = None) -> bool: with get_db_context(db) as db: - feedbacks = db.query(Feedback).all() - if not feedbacks: - return False - for feedback in feedbacks: - db.delete(feedback) + result = db.query(Feedback).delete() db.commit() - return True + return result > 0 Feedbacks = FeedbackTable() diff --git a/backend/open_webui/models/files.py b/backend/open_webui/models/files.py index 4097ae08e1..c24b242bd8 100644 --- a/backend/open_webui/models/files.py +++ b/backend/open_webui/models/files.py @@ -4,7 +4,7 @@ from typing import Optional from sqlalchemy.orm import Session from open_webui.internal.db import Base, JSONField, get_db, get_db_context -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, model_validator from sqlalchemy import BigInteger, Column, String, Text, JSON log = logging.getLogger(__name__) @@ -63,6 +63,25 @@ class FileMeta(BaseModel): model_config = ConfigDict(extra="allow") + @model_validator(mode="before") + @classmethod + def sanitize_meta(cls, data): + """Sanitize metadata fields to handle malformed legacy data.""" + if not isinstance(data, dict): + return data + + # Handle content_type that may be a list like ['application/pdf', None] + content_type = data.get("content_type") + if isinstance(content_type, list): + # Extract first non-None string value + data["content_type"] = next( + (item for item in content_type if isinstance(item, str)), None + ) + elif content_type is not None and not isinstance(content_type, str): + data["content_type"] = None + + return data + class FileModelResponse(BaseModel): id: str @@ -74,7 +93,7 @@ class FileModelResponse(BaseModel): meta: FileMeta created_at: int # timestamp in epoch - updated_at: int # timestamp in epoch + updated_at: Optional[int] = None # timestamp in epoch, optional for legacy files model_config = ConfigDict(extra="allow") diff --git a/backend/open_webui/models/functions.py b/backend/open_webui/models/functions.py index 8e23bac093..c41b328317 100644 --- a/backend/open_webui/models/functions.py +++ b/backend/open_webui/models/functions.py @@ -195,6 +195,26 @@ class FunctionsTable: except Exception: return None + def get_functions_by_ids( + self, ids: list[str], db: Optional[Session] = None + ) -> list[FunctionModel]: + """ + Batch fetch multiple functions by their IDs in a single query. + Returns functions in the same order as the input IDs (None entries filtered out). + """ + if not ids: + return [] + try: + with get_db_context(db) as db: + functions = db.query(Function).filter(Function.id.in_(ids)).all() + # Create a dict for O(1) lookup + func_dict = {f.id: FunctionModel.model_validate(f) for f in functions} + # Return in original order, filtering out any not found + return [func_dict[id] for id in ids if id in func_dict] + except Exception: + return [] + + def get_functions( self, active_only=False, include_valves=False, db: Optional[Session] = None ) -> list[FunctionModel | FunctionWithValvesModel]: @@ -299,7 +319,7 @@ class FunctionsTable: function.updated_at = int(time.time()) db.commit() db.refresh(function) - return self.get_function_by_id(id, db=db) + return FunctionModel.model_validate(function) except Exception: return None @@ -319,7 +339,7 @@ class FunctionsTable: function.updated_at = int(time.time()) db.commit() db.refresh(function) - return self.get_function_by_id(id, db=db) + return FunctionModel.model_validate(function) else: return None except Exception as e: @@ -381,7 +401,8 @@ class FunctionsTable: } ) db.commit() - return self.get_function_by_id(id, db=db) + function = db.get(Function, id) + return FunctionModel.model_validate(function) if function else None except Exception: return None diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index 31b34c45c8..eab817534b 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -589,11 +589,10 @@ class GroupTable: if not user_ids: return GroupModel.model_validate(group) - # Remove each user from group_member - for user_id in user_ids: - db.query(GroupMember).filter( - GroupMember.group_id == id, GroupMember.user_id == user_id - ).delete() + # Remove users from group_member in batch + db.query(GroupMember).filter( + GroupMember.group_id == id, GroupMember.user_id.in_(user_ids) + ).delete(synchronize_session=False) # Update group timestamp group.updated_at = int(time.time()) diff --git a/backend/open_webui/models/knowledge.py b/backend/open_webui/models/knowledge.py index 7f99f828c7..81aa4099d9 100644 --- a/backend/open_webui/models/knowledge.py +++ b/backend/open_webui/models/knowledge.py @@ -229,6 +229,9 @@ class KnowledgeTable: or_( Knowledge.name.ilike(f"%{query_key}%"), Knowledge.description.ilike(f"%{query_key}%"), + User.name.ilike(f"%{query_key}%"), + User.email.ilike(f"%{query_key}%"), + User.username.ilike(f"%{query_key}%"), ) ) @@ -240,7 +243,7 @@ class KnowledgeTable: query = has_permission(db, Knowledge, query, filter) - query = query.order_by(Knowledge.updated_at.desc()) + query = query.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc()) total = query.count() if skip: @@ -300,7 +303,7 @@ class KnowledgeTable: query = query.filter(File.filename.ilike(f"%{q}%")) # Order by file changes - query = query.order_by(File.updated_at.desc()) + query = query.order_by(File.updated_at.desc(), File.id.asc()) # Count before pagination total = query.count() @@ -427,6 +430,9 @@ class KnowledgeTable: .filter(KnowledgeFile.knowledge_id == knowledge_id) ) + # Default sort: updated_at descending + primary_sort = File.updated_at.desc() + if filter: query_key = filter.get("query") if query_key: @@ -440,27 +446,17 @@ class KnowledgeTable: order_by = filter.get("order_by") direction = filter.get("direction") + is_asc = direction == "asc" if order_by == "name": - if direction == "asc": - query = query.order_by(File.filename.asc()) - else: - query = query.order_by(File.filename.desc()) + primary_sort = File.filename.asc() if is_asc else File.filename.desc() elif order_by == "created_at": - if direction == "asc": - query = query.order_by(File.created_at.asc()) - else: - query = query.order_by(File.created_at.desc()) + primary_sort = File.created_at.asc() if is_asc else File.created_at.desc() elif order_by == "updated_at": - if direction == "asc": - query = query.order_by(File.updated_at.asc()) - else: - query = query.order_by(File.updated_at.desc()) - else: - query = query.order_by(File.updated_at.desc()) + primary_sort = File.updated_at.asc() if is_asc else File.updated_at.desc() - else: - query = query.order_by(File.updated_at.desc()) + # Apply sort with secondary key for deterministic pagination + query = query.order_by(primary_sort, File.id.asc()) # Count BEFORE pagination total = query.count() diff --git a/backend/open_webui/models/memories.py b/backend/open_webui/models/memories.py index 2dc9656856..e6b70a3020 100644 --- a/backend/open_webui/models/memories.py +++ b/backend/open_webui/models/memories.py @@ -82,7 +82,8 @@ class MemoriesTable: memory.updated_at = int(time.time()) db.commit() - return self.get_memory_by_id(id) + db.refresh(memory) + return MemoryModel.model_validate(memory) except Exception: return None diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index 5457413f0d..5a59861dd7 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -294,6 +294,9 @@ class ModelsTable: or_( Model.name.ilike(f"%{query_key}%"), Model.base_model_id.ilike(f"%{query_key}%"), + User.name.ilike(f"%{query_key}%"), + User.email.ilike(f"%{query_key}%"), + User.username.ilike(f"%{query_key}%"), ) ) @@ -391,17 +394,16 @@ class ModelsTable: ) -> Optional[ModelModel]: with get_db_context(db) as db: try: - is_active = db.query(Model).filter_by(id=id).first().is_active + model = db.query(Model).filter_by(id=id).first() + if not model: + return None - db.query(Model).filter_by(id=id).update( - { - "is_active": not is_active, - "updated_at": int(time.time()), - } - ) + model.is_active = not model.is_active + model.updated_at = int(time.time()) db.commit() + db.refresh(model) - return self.get_model_by_id(id, db=db) + return ModelModel.model_validate(model) except Exception: return None diff --git a/backend/open_webui/models/prompt_history.py b/backend/open_webui/models/prompt_history.py new file mode 100644 index 0000000000..ea7f566fb1 --- /dev/null +++ b/backend/open_webui/models/prompt_history.py @@ -0,0 +1,223 @@ +"""Prompt history model for version tracking.""" + +import time +import uuid +from typing import Optional +import json +import difflib + +from sqlalchemy.orm import Session +from open_webui.internal.db import Base, get_db_context +from open_webui.models.users import Users, UserResponse + +from pydantic import BaseModel, ConfigDict +from sqlalchemy import BigInteger, Column, Text, JSON, Index + + +#################### +# PromptHistory DB Schema +#################### + + +class PromptHistory(Base): + __tablename__ = "prompt_history" + + id = Column(Text, primary_key=True) + prompt_id = Column(Text, nullable=False, index=True) + parent_id = Column(Text, nullable=True) # Reference to parent commit + snapshot = Column(JSON, nullable=False) + user_id = Column(Text, nullable=False) + commit_message = Column(Text, nullable=True) + created_at = Column(BigInteger, nullable=False) + + +class PromptHistoryModel(BaseModel): + id: str + prompt_id: str + parent_id: Optional[str] = None + snapshot: dict + user_id: str + commit_message: Optional[str] = None + created_at: int + + model_config = ConfigDict(from_attributes=True) + + +class PromptHistoryResponse(PromptHistoryModel): + """Response model with user info.""" + user: Optional[UserResponse] = None + + +class PromptHistoryTable: + def create_history_entry( + self, + prompt_id: str, + snapshot: dict, + user_id: str, + parent_id: Optional[str] = None, + commit_message: Optional[str] = None, + db: Optional[Session] = None, + ) -> Optional[PromptHistoryModel]: + """Create a new history entry (commit) for a prompt.""" + with get_db_context(db) as db: + history = PromptHistory( + id=str(uuid.uuid4()), + prompt_id=prompt_id, + parent_id=parent_id, + snapshot=snapshot, + user_id=user_id, + commit_message=commit_message, + created_at=int(time.time()), + ) + db.add(history) + db.commit() + db.refresh(history) + return PromptHistoryModel.model_validate(history) + + def get_history_by_prompt_id( + self, + prompt_id: str, + limit: int = 50, + offset: int = 0, + db: Optional[Session] = None, + ) -> list[PromptHistoryResponse]: + """Get all history entries for a prompt, ordered by created_at desc.""" + with get_db_context(db) as db: + entries = ( + db.query(PromptHistory) + .filter(PromptHistory.prompt_id == prompt_id) + .order_by(PromptHistory.created_at.desc()) + .offset(offset) + .limit(limit) + .all() + ) + + # Get user info for each entry + user_ids = list(set(e.user_id for e in entries)) + users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] + users_dict = {user.id: user for user in users} + + return [ + PromptHistoryResponse( + **PromptHistoryModel.model_validate(entry).model_dump(), + user=users_dict.get(entry.user_id).model_dump() if users_dict.get(entry.user_id) else None, + ) + for entry in entries + ] + + def get_history_entry_by_id( + self, + history_id: str, + db: Optional[Session] = None, + ) -> Optional[PromptHistoryModel]: + """Get a specific history entry by ID.""" + with get_db_context(db) as db: + entry = db.query(PromptHistory).filter(PromptHistory.id == history_id).first() + if entry: + return PromptHistoryModel.model_validate(entry) + return None + + def get_latest_history_entry( + self, + prompt_id: str, + db: Optional[Session] = None, + ) -> Optional[PromptHistoryModel]: + """Get the most recent history entry for a prompt.""" + with get_db_context(db) as db: + entry = ( + db.query(PromptHistory) + .filter(PromptHistory.prompt_id == prompt_id) + .order_by(PromptHistory.created_at.desc()) + .first() + ) + if entry: + return PromptHistoryModel.model_validate(entry) + return None + + def get_history_count( + self, + prompt_id: str, + db: Optional[Session] = None, + ) -> int: + """Get the number of history entries for a prompt.""" + with get_db_context(db) as db: + return ( + db.query(PromptHistory) + .filter(PromptHistory.prompt_id == prompt_id) + .count() + ) + + def compute_diff( + self, + from_id: str, + to_id: str, + db: Optional[Session] = None, + ) -> Optional[dict]: + """Compute diff between two history entries.""" + with get_db_context(db) as db: + from_entry = db.query(PromptHistory).filter(PromptHistory.id == from_id).first() + to_entry = db.query(PromptHistory).filter(PromptHistory.id == to_id).first() + + if not from_entry or not to_entry: + return None + + from_snapshot = from_entry.snapshot + to_snapshot = to_entry.snapshot + + # Compute diff for content field + from_content = from_snapshot.get("content", "") + to_content = to_snapshot.get("content", "") + + diff_lines = list(difflib.unified_diff( + from_content.splitlines(keepends=True), + to_content.splitlines(keepends=True), + fromfile=f"v{from_id[:8]}", + tofile=f"v{to_id[:8]}", + lineterm="", + )) + + return { + "from_id": from_id, + "to_id": to_id, + "from_snapshot": from_snapshot, + "to_snapshot": to_snapshot, + "content_diff": diff_lines, + "name_changed": from_snapshot.get("name") != to_snapshot.get("name"), + "access_control_changed": from_snapshot.get("access_control") != to_snapshot.get("access_control"), + } + + def delete_history_by_prompt_id( + self, + prompt_id: str, + db: Optional[Session] = None, + ) -> bool: + """Delete all history entries for a prompt.""" + with get_db_context(db) as db: + db.query(PromptHistory).filter(PromptHistory.prompt_id == prompt_id).delete() + db.commit() + return True + + def delete_history_entry( + self, + history_id: str, + db: Optional[Session] = None, + ) -> bool: + """Delete a history entry and reparent its children to grandparent.""" + with get_db_context(db) as db: + entry = db.query(PromptHistory).filter_by(id=history_id).first() + if not entry: + return False + + # Find children that reference this entry as parent + children = db.query(PromptHistory).filter_by(parent_id=history_id).all() + + # Reparent children to grandparent + for child in children: + child.parent_id = entry.parent_id + + db.delete(entry) + db.commit() + return True + + +PromptHistories = PromptHistoryTable() diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index 847597bc65..4a85ba9029 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -1,16 +1,21 @@ import time +import uuid from typing import Optional from sqlalchemy.orm import Session from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.groups import Groups from open_webui.models.users import Users, UserResponse +from open_webui.models.prompt_history import PromptHistories + from pydantic import BaseModel, ConfigDict -from sqlalchemy import BigInteger, Column, String, Text, JSON +from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON, or_, func, cast + from open_webui.utils.access_control import has_access + #################### # Prompts DB Schema #################### @@ -19,11 +24,18 @@ from open_webui.utils.access_control import has_access class Prompt(Base): __tablename__ = "prompt" - command = Column(String, primary_key=True) + id = Column(Text, primary_key=True) + command = Column(String, unique=True, index=True) user_id = Column(String) - title = Column(Text) + name = Column(Text) content = Column(Text) - timestamp = Column(BigInteger) + data = Column(JSON, nullable=True) + meta = Column(JSON, nullable=True) + tags = Column(JSON, nullable=True) + is_active = Column(Boolean, default=True) + version_id = Column(Text, nullable=True) # Points to active history entry + created_at = Column(BigInteger, nullable=True) + updated_at = Column(BigInteger, nullable=True) access_control = Column(JSON, nullable=True) # Controls data access levels. # Defines access control rules for this entry. @@ -44,13 +56,20 @@ class Prompt(Base): class PromptModel(BaseModel): + id: Optional[str] = None command: str user_id: str - title: str + name: str content: str - timestamp: int # timestamp in epoch - + data: Optional[dict] = None + meta: Optional[dict] = None + tags: Optional[list[str]] = None + is_active: Optional[bool] = True + version_id: Optional[str] = None + created_at: Optional[int] = None + updated_at: Optional[int] = None access_control: Optional[dict] = None + model_config = ConfigDict(from_attributes=True) @@ -67,23 +86,50 @@ class PromptAccessResponse(PromptUserResponse): write_access: Optional[bool] = False +class PromptListResponse(BaseModel): + items: list[PromptUserResponse] + total: int + + +class PromptAccessListResponse(BaseModel): + items: list[PromptAccessResponse] + total: int + + class PromptForm(BaseModel): + command: str - title: str + name: str # Changed from title content: str + data: Optional[dict] = None + meta: Optional[dict] = None + tags: Optional[list[str]] = None access_control: Optional[dict] = None + version_id: Optional[str] = None # Active version + commit_message: Optional[str] = None # For history tracking + is_production: Optional[bool] = True # Whether to set new version as production class PromptsTable: def insert_new_prompt( self, user_id: str, form_data: PromptForm, db: Optional[Session] = None ) -> Optional[PromptModel]: + now = int(time.time()) + prompt_id = str(uuid.uuid4()) + prompt = PromptModel( - **{ - "user_id": user_id, - **form_data.model_dump(), - "timestamp": int(time.time()), - } + id=prompt_id, + user_id=user_id, + command=form_data.command, + name=form_data.name, + content=form_data.content, + data=form_data.data or {}, + meta=form_data.meta or {}, + tags=form_data.tags or [], + access_control=form_data.access_control, + is_active=True, + created_at=now, + updated_at=now, ) try: @@ -92,26 +138,72 @@ class PromptsTable: db.add(result) db.commit() db.refresh(result) + if result: + snapshot = { + "name": form_data.name, + "content": form_data.content, + "command": form_data.command, + "data": form_data.data or {}, + "meta": form_data.meta or {}, + "tags": form_data.tags or [], + "access_control": form_data.access_control, + } + + history_entry = PromptHistories.create_history_entry( + prompt_id=prompt_id, + snapshot=snapshot, + user_id=user_id, + parent_id=None, # Initial commit has no parent + commit_message=form_data.commit_message or "Initial version", + db=db, + ) + + # Set the initial version as the production version + if history_entry: + result.version_id = history_entry.id + db.commit() + db.refresh(result) + return PromptModel.model_validate(result) else: return None except Exception: return None + def get_prompt_by_id( + self, prompt_id: str, db: Optional[Session] = None + ) -> Optional[PromptModel]: + """Get prompt by UUID.""" + try: + with get_db_context(db) as db: + prompt = db.query(Prompt).filter_by(id=prompt_id).first() + if prompt: + return PromptModel.model_validate(prompt) + return None + except Exception: + return None + def get_prompt_by_command( self, command: str, db: Optional[Session] = None ) -> Optional[PromptModel]: try: with get_db_context(db) as db: prompt = db.query(Prompt).filter_by(command=command).first() - return PromptModel.model_validate(prompt) + if prompt: + return PromptModel.model_validate(prompt) + return None except Exception: return None def get_prompts(self, db: Optional[Session] = None) -> list[PromptUserResponse]: with get_db_context(db) as db: - all_prompts = db.query(Prompt).order_by(Prompt.timestamp.desc()).all() + all_prompts = ( + db.query(Prompt) + .filter(Prompt.is_active == True) + .order_by(Prompt.updated_at.desc()) + .all() + ) user_ids = list(set(prompt.user_id for prompt in all_prompts)) @@ -147,17 +239,306 @@ class PromptsTable: or has_access(user_id, permission, prompt.access_control, user_group_ids) ] + def search_prompts( + self, + user_id: str, + filter: dict = {}, + skip: int = 0, + limit: int = 30, + db: Optional[Session] = None, + ) -> PromptListResponse: + with get_db_context(db) as db: + from open_webui.models.users import User, UserModel + + # Join with User table for user filtering and sorting + query = db.query(Prompt, User).outerjoin(User, User.id == Prompt.user_id) + query = query.filter(Prompt.is_active == True) + + if filter: + query_key = filter.get("query") + if query_key: + query = query.filter( + or_( + Prompt.name.ilike(f"%{query_key}%"), + Prompt.command.ilike(f"%{query_key}%"), + Prompt.content.ilike(f"%{query_key}%"), + User.name.ilike(f"%{query_key}%"), + User.email.ilike(f"%{query_key}%"), + ) + ) + + view_option = filter.get("view_option") + if view_option == "created": + query = query.filter(Prompt.user_id == user_id) + elif view_option == "shared": + query = query.filter(Prompt.user_id != user_id) + + # Apply access control filtering + group_ids = filter.get("group_ids", []) + filter_user_id = filter.get("user_id") + + if filter_user_id: + # User must have access: owner OR public OR explicit access + access_conditions = [ + Prompt.user_id == filter_user_id, # Owner + Prompt.access_control == None, # Public + ] + query = query.filter(or_(*access_conditions)) + + tag = filter.get("tag") + if tag: + # Search for tag in JSON array field + like_pattern = f'%"{tag.lower()}"%' + tags_text = func.lower(cast(Prompt.tags, String)) + query = query.filter(tags_text.like(like_pattern)) + + order_by = filter.get("order_by") + direction = filter.get("direction") + + if order_by == "name": + if direction == "asc": + query = query.order_by(Prompt.name.asc()) + else: + query = query.order_by(Prompt.name.desc()) + elif order_by == "created_at": + if direction == "asc": + query = query.order_by(Prompt.created_at.asc()) + else: + query = query.order_by(Prompt.created_at.desc()) + elif order_by == "updated_at": + if direction == "asc": + query = query.order_by(Prompt.updated_at.asc()) + else: + query = query.order_by(Prompt.updated_at.desc()) + else: + query = query.order_by(Prompt.updated_at.desc()) + else: + query = query.order_by(Prompt.updated_at.desc()) + + # Count BEFORE pagination + total = query.count() + + if skip: + query = query.offset(skip) + if limit: + query = query.limit(limit) + + items = query.all() + + prompts = [] + for prompt, user in items: + prompts.append( + PromptUserResponse( + **PromptModel.model_validate(prompt).model_dump(), + user=( + UserResponse(**UserModel.model_validate(user).model_dump()) + if user + else None + ), + ) + ) + + return PromptListResponse(items=prompts, total=total) + def update_prompt_by_command( - self, command: str, form_data: PromptForm, db: Optional[Session] = None + + self, + command: str, + form_data: PromptForm, + user_id: str, + db: Optional[Session] = None, ) -> Optional[PromptModel]: try: with get_db_context(db) as db: prompt = db.query(Prompt).filter_by(command=command).first() - prompt.title = form_data.title + if not prompt: + return None + + latest_history = PromptHistories.get_latest_history_entry( + prompt.id, db=db + ) + parent_id = latest_history.id if latest_history else None + + # Check if content changed to decide on history creation + content_changed = ( + prompt.name != form_data.name + or prompt.content != form_data.content + or prompt.access_control != form_data.access_control + ) + + # Update prompt fields + prompt.name = form_data.name prompt.content = form_data.content + prompt.data = form_data.data or prompt.data + prompt.meta = form_data.meta or prompt.meta prompt.access_control = form_data.access_control - prompt.timestamp = int(time.time()) + prompt.updated_at = int(time.time()) + db.commit() + + # Create history entry only if content changed + if content_changed: + snapshot = { + "name": form_data.name, + "content": form_data.content, + "command": command, + "data": form_data.data or {}, + "meta": form_data.meta or {}, + "access_control": form_data.access_control, + } + + history_entry = PromptHistories.create_history_entry( + prompt_id=prompt.id, + snapshot=snapshot, + user_id=user_id, + parent_id=parent_id, + commit_message=form_data.commit_message, + db=db, + ) + + # Set as production if flag is True (default) + if form_data.is_production and history_entry: + prompt.version_id = history_entry.id + db.commit() + + return PromptModel.model_validate(prompt) + except Exception: + return None + + def update_prompt_by_id( + self, + prompt_id: str, + form_data: PromptForm, + user_id: str, + db: Optional[Session] = None, + ) -> Optional[PromptModel]: + try: + with get_db_context(db) as db: + prompt = db.query(Prompt).filter_by(id=prompt_id).first() + if not prompt: + return None + + latest_history = PromptHistories.get_latest_history_entry( + prompt.id, db=db + ) + parent_id = latest_history.id if latest_history else None + + # Check if content changed to decide on history creation + content_changed = ( + prompt.name != form_data.name + or prompt.command != form_data.command + or prompt.content != form_data.content + or prompt.access_control != form_data.access_control + or (form_data.tags is not None and prompt.tags != form_data.tags) + ) + + # Update prompt fields + prompt.name = form_data.name + prompt.command = form_data.command + prompt.content = form_data.content + prompt.data = form_data.data or prompt.data + prompt.meta = form_data.meta or prompt.meta + prompt.access_control = form_data.access_control + + if form_data.tags is not None: + prompt.tags = form_data.tags + + prompt.updated_at = int(time.time()) + + db.commit() + + # Create history entry only if content changed + if content_changed: + snapshot = { + "name": form_data.name, + "content": form_data.content, + "command": prompt.command, + "data": form_data.data or {}, + "meta": form_data.meta or {}, + "tags": prompt.tags or [], + "access_control": form_data.access_control, + } + + history_entry = PromptHistories.create_history_entry( + prompt_id=prompt.id, + snapshot=snapshot, + user_id=user_id, + parent_id=parent_id, + commit_message=form_data.commit_message, + db=db, + ) + + # Set as production if flag is True (default) + if form_data.is_production and history_entry: + prompt.version_id = history_entry.id + db.commit() + + return PromptModel.model_validate(prompt) + except Exception: + return None + + def update_prompt_metadata( + self, + prompt_id: str, + name: str, + command: str, + tags: Optional[list[str]] = None, + db: Optional[Session] = None, + ) -> Optional[PromptModel]: + """Update only name and command (no history created).""" + try: + with get_db_context(db) as db: + prompt = db.query(Prompt).filter_by(id=prompt_id).first() + if not prompt: + return None + + prompt.name = name + prompt.command = command + + if tags is not None: + prompt.tags = tags + + prompt.updated_at = int(time.time()) + db.commit() + + return PromptModel.model_validate(prompt) + except Exception: + return None + + def update_prompt_version( + self, + prompt_id: str, + version_id: str, + db: Optional[Session] = None, + ) -> Optional[PromptModel]: + """Set the active version of a prompt and restore content from that version's snapshot.""" + try: + with get_db_context(db) as db: + prompt = db.query(Prompt).filter_by(id=prompt_id).first() + if not prompt: + return None + + history_entry = PromptHistories.get_history_entry_by_id( + version_id, db=db + ) + + if not history_entry: + return None + + # Restore prompt content from the snapshot + snapshot = history_entry.snapshot + if snapshot: + prompt.name = snapshot.get("name", prompt.name) + prompt.content = snapshot.get("content", prompt.content) + prompt.data = snapshot.get("data", prompt.data) + prompt.meta = snapshot.get("meta", prompt.meta) + prompt.tags = snapshot.get("tags", prompt.tags) + # Note: command and access_control are not restored from snapshot + + prompt.version_id = version_id + prompt.updated_at = int(time.time()) + db.commit() + return PromptModel.model_validate(prompt) except Exception: return None @@ -165,14 +546,68 @@ class PromptsTable: def delete_prompt_by_command( self, command: str, db: Optional[Session] = None ) -> bool: + """Soft delete a prompt by setting is_active to False.""" try: with get_db_context(db) as db: - db.query(Prompt).filter_by(command=command).delete() - db.commit() + prompt = db.query(Prompt).filter_by(command=command).first() + if prompt: + PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) - return True + prompt.is_active = False + prompt.updated_at = int(time.time()) + db.commit() + return True + return False except Exception: return False + def delete_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> bool: + """Soft delete a prompt by setting is_active to False.""" + try: + with get_db_context(db) as db: + prompt = db.query(Prompt).filter_by(id=prompt_id).first() + if prompt: + PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + + prompt.is_active = False + prompt.updated_at = int(time.time()) + db.commit() + return True + return False + except Exception: + return False + + def hard_delete_prompt_by_command( + self, command: str, db: Optional[Session] = None + ) -> bool: + """Permanently delete a prompt and its history.""" + try: + with get_db_context(db) as db: + prompt = db.query(Prompt).filter_by(command=command).first() + if prompt: + PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + + # Delete prompt + db.query(Prompt).filter_by(command=command).delete() + db.commit() + return True + return False + except Exception: + return False + + def get_tags(self, db: Optional[Session] = None) -> list[str]: + try: + with get_db_context(db) as db: + prompts = db.query(Prompt).filter_by(is_active=True).all() + tags = set() + for prompt in prompts: + if prompt.tags: + for tag in prompt.tags: + if tag: + tags.add(tag) + return sorted(list(tags)) + except Exception: + return [] + Prompts = PromptsTable() diff --git a/backend/open_webui/models/users.py b/backend/open_webui/models/users.py index 0d36d94b8f..deeda29a85 100644 --- a/backend/open_webui/models/users.py +++ b/backend/open_webui/models/users.py @@ -530,9 +530,12 @@ class UsersTable: ) -> Optional[UserModel]: try: with get_db_context(db) as db: - db.query(User).filter_by(id=id).update({"role": role}) - db.commit() user = db.query(User).filter_by(id=id).first() + if not user: + return None + user.role = role + db.commit() + db.refresh(user) return UserModel.model_validate(user) except Exception: return None @@ -542,12 +545,13 @@ class UsersTable: ) -> Optional[UserModel]: try: with get_db_context(db) as db: - db.query(User).filter_by(id=id).update( - {**form_data.model_dump(exclude_none=True)} - ) - db.commit() - user = db.query(User).filter_by(id=id).first() + if not user: + return None + for key, value in form_data.model_dump(exclude_none=True).items(): + setattr(user, key, value) + db.commit() + db.refresh(user) return UserModel.model_validate(user) except Exception: return None @@ -557,12 +561,12 @@ class UsersTable: ) -> Optional[UserModel]: try: with get_db_context(db) as db: - db.query(User).filter_by(id=id).update( - {"profile_image_url": profile_image_url} - ) - db.commit() - user = db.query(User).filter_by(id=id).first() + if not user: + return None + user.profile_image_url = profile_image_url + db.commit() + db.refresh(user) return UserModel.model_validate(user) except Exception: return None @@ -573,12 +577,12 @@ class UsersTable: ) -> Optional[UserModel]: try: with get_db_context(db) as db: - db.query(User).filter_by(id=id).update( - {"last_active_at": int(time.time())} - ) - db.commit() - user = db.query(User).filter_by(id=id).first() + if not user: + return None + user.last_active_at = int(time.time()) + db.commit() + db.refresh(user) return UserModel.model_validate(user) except Exception: return None @@ -620,12 +624,14 @@ class UsersTable: ) -> Optional[UserModel]: try: with get_db_context(db) as db: - db.query(User).filter_by(id=id).update(updated) - db.commit() - user = db.query(User).filter_by(id=id).first() + if not user: + return None + for key, value in updated.items(): + setattr(user, key, value) + db.commit() + db.refresh(user) return UserModel.model_validate(user) - # return UserModel(**user.dict()) except Exception as e: print(e) return None diff --git a/backend/open_webui/retrieval/loaders/main.py b/backend/open_webui/retrieval/loaders/main.py index 64784ad61f..51ecd627ac 100644 --- a/backend/open_webui/retrieval/loaders/main.py +++ b/backend/open_webui/retrieval/loaders/main.py @@ -143,7 +143,7 @@ class DoclingLoader: with open(self.file_path, "rb") as f: headers = {} if self.api_key: - headers["X-Api-Key"] = f"Bearer {self.api_key}" + headers["X-Api-Key"] = f"{self.api_key}" r = requests.post( f"{self.url}/v1/convert/file", @@ -363,7 +363,9 @@ class Loader: else: if file_ext == "pdf": loader = PyPDFLoader( - file_path, extract_images=self.kwargs.get("PDF_EXTRACT_IMAGES") + file_path, + extract_images=self.kwargs.get("PDF_EXTRACT_IMAGES"), + mode=self.kwargs.get("PDF_LOADER_MODE", "page"), ) elif file_ext == "csv": loader = CSVLoader(file_path, autodetect_encoding=True) diff --git a/backend/open_webui/retrieval/vector/dbs/weaviate.py b/backend/open_webui/retrieval/vector/dbs/weaviate.py index d204e8293a..dcc648c788 100644 --- a/backend/open_webui/retrieval/vector/dbs/weaviate.py +++ b/backend/open_webui/retrieval/vector/dbs/weaviate.py @@ -12,9 +12,13 @@ from open_webui.retrieval.vector.main import ( from open_webui.retrieval.vector.utils import process_metadata from open_webui.config import ( WEAVIATE_HTTP_HOST, + WEAVIATE_GRPC_HOST, WEAVIATE_HTTP_PORT, WEAVIATE_GRPC_PORT, WEAVIATE_API_KEY, + WEAVIATE_HTTP_SECURE, + WEAVIATE_GRPC_SECURE, + WEAVIATE_SKIP_INIT_CHECKS, ) @@ -52,9 +56,13 @@ class WeaviateClient(VectorDBBase): try: # Build connection parameters connection_params = { - "host": WEAVIATE_HTTP_HOST, - "port": WEAVIATE_HTTP_PORT, + "http_host": WEAVIATE_HTTP_HOST, + "http_port": WEAVIATE_HTTP_PORT, + "http_secure": WEAVIATE_HTTP_SECURE, + "grpc_host": WEAVIATE_GRPC_HOST, "grpc_port": WEAVIATE_GRPC_PORT, + "grpc_secure": WEAVIATE_GRPC_SECURE, + "skip_init_checks": WEAVIATE_SKIP_INIT_CHECKS, } # Only add auth_credentials if WEAVIATE_API_KEY exists and is not empty @@ -63,7 +71,7 @@ class WeaviateClient(VectorDBBase): weaviate.classes.init.Auth.api_key(WEAVIATE_API_KEY) ) - self.client = weaviate.connect_to_local(**connection_params) + self.client = weaviate.connect_to_custom(**connection_params) self.client.connect() except Exception as e: raise ConnectionError(f"Failed to connect to Weaviate: {e}") from e diff --git a/backend/open_webui/retrieval/web/yandex.py b/backend/open_webui/retrieval/web/yandex.py new file mode 100644 index 0000000000..def134d996 --- /dev/null +++ b/backend/open_webui/retrieval/web/yandex.py @@ -0,0 +1,147 @@ +import base64 +import io +import json +import logging +import os +from typing import Optional, List + +import requests + +from fastapi import Request + +from open_webui.retrieval.web.main import SearchResult, get_filtered_results +from open_webui.utils.headers import include_user_info_headers + +from xml.etree import ElementTree as ET +from xml.etree.ElementTree import Element + +log = logging.getLogger(__name__) + + +def xml_element_contents_to_string(element: Element) -> str: + buffer = [element.text if element.text else ""] + + for child in element: + buffer.append(xml_element_contents_to_string(child)) + + buffer.append(element.tail if element.tail else "") + + return "".join(buffer) + + +def search_yandex( + request: Request, + yandex_search_url: str, + yandex_search_api_key: str, + yandex_search_config: str, + query: str, + count: int, + filter_list: Optional[List[str]] = None, + user=None, +) -> List[SearchResult]: + try: + headers = { + "User-Agent": "Open WebUI (https://github.com/open-webui/open-webui) RAG Bot", + "Authorization": f"Api-Key {yandex_search_api_key}", + } + + if user is not None: + headers = include_user_info_headers(headers, user) + + chat_id = getattr(request.state, "chat_id", None) + if chat_id: + headers["X-OpenWebUI-Chat-Id"] = str(chat_id) + + payload = {} if yandex_search_config == "" else json.loads(yandex_search_config) + + if type(payload.get("query", None)) != dict: + payload["query"] = {} + + if "searchType" not in payload["query"]: + payload["query"]["searchType"] = "SEARCH_TYPE_RU" + + payload["query"]["queryText"] = query + + if type(payload.get("groupSpec", None)) != dict: + payload["groupSpec"] = {} + + if "groupMode" not in payload["groupSpec"]: + payload["groupSpec"]["groupMode"] = "GROUP_MODE_DEEP" + + payload["groupSpec"]["groupsOnPage"] = count + payload["groupSpec"]["docsInGroup"] = 1 + + response = requests.post( + "https://searchapi.api.cloud.yandex.net/v2/web/search" if yandex_search_url == "" else yandex_search_url, + headers=headers, + json=payload, + ) + + response.raise_for_status() + + response_body = response.json() + if "rawData" not in response_body: + raise Exception(f"No `rawData` in response body: {response_body}") + + search_result_body_bytes = base64.decodebytes(bytes(response_body["rawData"], "utf-8")) + + doc_root = ET.parse(io.BytesIO(search_result_body_bytes)) + + results = [] + + for group in doc_root.findall("response/results/grouping/group"): + results.append({ + "url": xml_element_contents_to_string(group.find("doc/url")).strip("\n"), + "title": xml_element_contents_to_string(group.find("doc/title")).strip("\n"), + "snippet": xml_element_contents_to_string(group.find("doc/passages/passage")), + }) + + results = get_filtered_results(results, filter_list) + + results = [ + SearchResult( + link=result.get("url"), + title=result.get("title"), + snippet=result.get("snippet"), + ) + for result in results[:count] + ] + + log.info(f"Yandex search results: {results}") + + return results + except Exception as e: + log.error(f"Error in search: {e}") + + return [] + + +if __name__ == "__main__": + from starlette.datastructures import Headers + from fastapi import FastAPI + + result = search_yandex( + Request( + { + "type": "http", + "asgi.version": "3.0", + "asgi.spec_version": "2.0", + "method": "GET", + "path": "/internal", + "query_string": b"", + "headers": Headers({}).raw, + "client": ("127.0.0.1", 12345), + "server": ("127.0.0.1", 80), + "scheme": "http", + "app": FastAPI(), + }, + None, + ), + os.environ.get("YANDEX_WEB_SEARCH_URL", ""), + os.environ.get("YANDEX_WEB_SEARCH_API_KEY", ""), + os.environ.get("YANDEX_WEB_SEARCH_CONFIG", "{\"query\": {\"searchType\": \"SEARCH_TYPE_COM\"}}"), + "TOP movies of the past year", + 3, + ) + + print(result) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 30d4ebe4cc..586fc66ec2 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -37,6 +37,8 @@ from open_webui.env import ( WEBUI_AUTH_COOKIE_SECURE, WEBUI_AUTH_SIGNOUT_REDIRECT_URL, ENABLE_INITIAL_ADMIN_SIGNUP, + ENABLE_OAUTH_TOKEN_EXCHANGE, + AIOHTTP_CLIENT_SESSION_SSL, ) from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi.responses import RedirectResponse, Response, JSONResponse @@ -45,6 +47,8 @@ from open_webui.config import ( ENABLE_OAUTH_SIGNUP, ENABLE_LDAP, ENABLE_PASSWORD_AUTH, + OAUTH_PROVIDERS, + OAUTH_MERGE_ACCOUNTS_BY_EMAIL, ) from pydantic import BaseModel @@ -87,6 +91,63 @@ signin_rate_limiter = RateLimiter( redis_client=get_redis_client(), limit=5 * 3, window=60 * 3 ) + +def create_session_response( + request: Request, user, db, response: Response = None, set_cookie: bool = False +) -> dict: + """ + Create JWT token and build session response for a user. + Shared helper for signin, signup, ldap_auth, add_user, and token_exchange endpoints. + + Args: + request: FastAPI request object + user: User object + db: Database session + 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_at = None + if expires_delta: + expires_at = int(time.time()) + int(expires_delta.total_seconds()) + + token = create_token( + data={"id": user.id}, + expires_delta=expires_delta, + ) + + if set_cookie and response: + datetime_expires_at = ( + datetime.datetime.fromtimestamp(expires_at, datetime.timezone.utc) + if expires_at + else None + ) + response.set_cookie( + key="token", + value=token, + expires=datetime_expires_at, + httponly=True, + samesite=WEBUI_AUTH_COOKIE_SAME_SITE, + secure=WEBUI_AUTH_COOKIE_SECURE, + ) + + user_permissions = get_permissions( + user.id, request.app.state.config.USER_PERMISSIONS, db=db + ) + + return { + "token": token, + "token_type": "Bearer", + "expires_at": expires_at, + "id": user.id, + "email": user.email, + "name": user.name, + "role": user.role, + "profile_image_url": user.profile_image_url, + "permissions": user_permissions, + } + + ############################ # GetSessionUser ############################ @@ -482,36 +543,6 @@ async def ldap_auth( user = Auths.authenticate_user_by_email(email, db=db) if user: - expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN) - expires_at = None - if expires_delta: - expires_at = int(time.time()) + int(expires_delta.total_seconds()) - - token = create_token( - data={"id": user.id}, - expires_delta=expires_delta, - ) - - # Set the cookie token - response.set_cookie( - key="token", - value=token, - expires=( - datetime.datetime.fromtimestamp( - expires_at, datetime.timezone.utc - ) - if expires_at - else None - ), - httponly=True, # Ensures the cookie is not accessible via JavaScript - samesite=WEBUI_AUTH_COOKIE_SAME_SITE, - secure=WEBUI_AUTH_COOKIE_SECURE, - ) - - user_permissions = get_permissions( - user.id, request.app.state.config.USER_PERMISSIONS, db=db - ) - if ( user.role != "admin" and ENABLE_LDAP_GROUP_MANAGEMENT @@ -527,17 +558,7 @@ async def ldap_auth( except Exception as e: log.error(f"Failed to sync groups for user {user.id}: {e}") - return { - "token": token, - "token_type": "Bearer", - "expires_at": expires_at, - "id": user.id, - "email": user.email, - "name": user.name, - "role": user.role, - "profile_image_url": user.profile_image_url, - "permissions": user_permissions, - } + return create_session_response(request, user, db, response, set_cookie=True) else: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) else: @@ -646,48 +667,7 @@ async def signin( ) if user: - - expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN) - expires_at = None - if expires_delta: - expires_at = int(time.time()) + int(expires_delta.total_seconds()) - - token = create_token( - data={"id": user.id}, - expires_delta=expires_delta, - ) - - datetime_expires_at = ( - datetime.datetime.fromtimestamp(expires_at, datetime.timezone.utc) - if expires_at - else None - ) - - # Set the cookie token - response.set_cookie( - key="token", - value=token, - expires=datetime_expires_at, - httponly=True, # Ensures the cookie is not accessible via JavaScript - samesite=WEBUI_AUTH_COOKIE_SAME_SITE, - secure=WEBUI_AUTH_COOKIE_SECURE, - ) - - user_permissions = get_permissions( - user.id, request.app.state.config.USER_PERMISSIONS, db=db - ) - - return { - "token": token, - "token_type": "Bearer", - "expires_at": expires_at, - "id": user.id, - "email": user.email, - "name": user.name, - "role": user.role, - "profile_image_url": user.profile_image_url, - "permissions": user_permissions, - } + return create_session_response(request, user, db, response, set_cookie=True) else: raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) @@ -748,32 +728,6 @@ async def signup( ) if user: - expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN) - expires_at = None - if expires_delta: - expires_at = int(time.time()) + int(expires_delta.total_seconds()) - - token = create_token( - data={"id": user.id}, - expires_delta=expires_delta, - ) - - datetime_expires_at = ( - datetime.datetime.fromtimestamp(expires_at, datetime.timezone.utc) - if expires_at - else None - ) - - # Set the cookie token - response.set_cookie( - key="token", - value=token, - expires=datetime_expires_at, - httponly=True, # Ensures the cookie is not accessible via JavaScript - samesite=WEBUI_AUTH_COOKIE_SAME_SITE, - secure=WEBUI_AUTH_COOKIE_SECURE, - ) - if request.app.state.config.WEBHOOK_URL: await post_webhook( request.app.state.WEBUI_NAME, @@ -786,10 +740,6 @@ async def signup( }, ) - user_permissions = get_permissions( - user.id, request.app.state.config.USER_PERMISSIONS, db=db - ) - if not has_users: # Disable signup after the first user is created request.app.state.config.ENABLE_SIGNUP = False @@ -800,19 +750,11 @@ async def signup( db=db, ) - return { - "token": token, - "token_type": "Bearer", - "expires_at": expires_at, - "id": user.id, - "email": user.email, - "name": user.name, - "role": user.role, - "profile_image_url": user.profile_image_url, - "permissions": user_permissions, - } + return create_session_response(request, user, db, response, set_cookie=True) else: raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) + except HTTPException: + raise except Exception as err: log.error(f"Signup error: {str(err)}") raise HTTPException(500, detail="An internal error occurred during signup.") @@ -954,6 +896,8 @@ async def add_user( } else: raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) + except HTTPException: + raise except Exception as err: log.error(f"Add user error: {str(err)}") raise HTTPException( @@ -1283,3 +1227,108 @@ async def get_api_key( } else: raise HTTPException(404, detail=ERROR_MESSAGES.API_KEY_NOT_FOUND) + + +############################ +# Token Exchange +############################ + + +class TokenExchangeForm(BaseModel): + token: str # OAuth access token from external provider + + +@router.post("/oauth/{provider}/token/exchange", response_model=SessionUserResponse) +async def token_exchange( + request: Request, + response: Response, + provider: str, + form_data: TokenExchangeForm, + db: Session = Depends(get_session), +): + """ + Exchange an external OAuth provider token for an OpenWebUI JWT. + This endpoint is disabled by default. Set ENABLE_OAUTH_TOKEN_EXCHANGE=True to enable. + """ + if not ENABLE_OAUTH_TOKEN_EXCHANGE: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Token exchange is disabled", + ) + + provider = provider.lower() + + # Check if provider is configured + if provider not in OAUTH_PROVIDERS: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider}' is not configured", + ) + # Get the OAuth client for this provider + oauth_manager = request.app.state.oauth_manager + client = oauth_manager.get_client(provider) + if not client: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"OAuth client for '{provider}' not found", + ) + + # Validate the token by calling the userinfo endpoint + try: + token_data = {"access_token": form_data.token, "token_type": "Bearer"} + user_data = await client.userinfo(token=token_data) + + if not user_data: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Invalid token or unable to fetch user info", + ) + except Exception as e: + log.warning(f"Token exchange failed for provider {provider}: {e}") + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Invalid token or unable to validate with provider", + ) + + # 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 + + # Get sub claim + sub = user_data.get( + request.app.state.config.OAUTH_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( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Token missing required 'sub' claim", + ) + + email = user_data.get(email_claim, "") + if not email: + log.warning(f"Token exchange failed: email claim missing from user data") + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Token missing required email claim", + ) + email = email.lower() + + # Try to find the user by OAuth sub + user = Users.get_user_by_oauth_sub(provider, sub, db=db) + + if not user and OAUTH_MERGE_ACCOUNTS_BY_EMAIL.value: + # Try to find by email if merge is enabled + user = Users.get_user_by_email(email, db=db) + if user: + # Link the OAuth sub to this user + Users.update_user_oauth_by_id(user.id, provider, sub, db=db) + + if not user: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="User not found. Please sign in via the web interface first.", + ) + + return create_session_response(request, user, db) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 7f4b347bef..e463276589 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -22,6 +22,7 @@ from open_webui.models.users import ( UserListResponse, UserModelResponse, Users, + UserModel, UserNameResponse, ) @@ -80,7 +81,7 @@ router = APIRouter() ############################ -def check_channels_access(request: Request): +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: raise HTTPException( @@ -88,6 +89,15 @@ def check_channels_access(request: Request): detail="Channels are not enabled", ) + if user: + if user.role != "admin" and not has_permission( + user.id, "features.channels", request.app.state.config.USER_PERMISSIONS + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.UNAUTHORIZED, + ) + ############################ # GetChatList @@ -355,7 +365,7 @@ async def get_channel_by_id( user=Depends(get_verified_user), db: Session = Depends(get_session), ): - check_channels_access(request) + check_channels_access(request, user) channel = Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException( @@ -467,7 +477,7 @@ async def get_channel_members_by_id( user=Depends(get_verified_user), db: Session = Depends(get_session), ): - check_channels_access(request) + check_channels_access(request, user) channel = Channels.get_channel_by_id(id, db=db) if not channel: @@ -788,7 +798,7 @@ async def get_channel_messages( user=Depends(get_verified_user), db: Session = Depends(get_session), ): - check_channels_access(request) + check_channels_access(request, user) channel = Channels.get_channel_by_id(id, db=db) if not channel: raise HTTPException( diff --git a/backend/open_webui/routers/chats.py b/backend/open_webui/routers/chats.py index 9a43234aa6..e03cdc7ba9 100644 --- a/backend/open_webui/routers/chats.py +++ b/backend/open_webui/routers/chats.py @@ -16,6 +16,7 @@ from open_webui.models.chats import ( ChatResponse, Chats, ChatTitleIdResponse, + SharedChatResponse, ChatStatsExport, AggregateChatStats, ChatBody, @@ -357,9 +358,7 @@ def _process_chat_for_export(chat) -> Optional[ChatStatsExport]: return None -def calculate_chat_stats( - user_id, skip=0, limit=10, filter=None, db: Optional[Session] = None -): +def calculate_chat_stats(user_id, skip=0, limit=10, filter=None): if filter is None: filter = {} @@ -368,7 +367,6 @@ def calculate_chat_stats( skip=skip, limit=limit, filter=filter, - db=db, ) chat_stats_export_list = [] @@ -424,7 +422,6 @@ async def export_chat_stats( page: Optional[int] = 1, stream: bool = False, user=Depends(get_verified_user), - db: Session = Depends(get_session), ): # Check if the user has permission to share/export chats if (user.role != "admin") and ( @@ -455,7 +452,7 @@ async def export_chat_stats( skip = (page - 1) * limit chat_stats_export_list, total = await asyncio.to_thread( - calculate_chat_stats, user.id, skip, limit, filter, db=db + calculate_chat_stats, user.id, skip, limit, filter ) return ChatStatsExportList( @@ -862,6 +859,48 @@ async def unarchive_all_chats( return Chats.unarchive_all_chats_by_user_id(user.id, db=db) +############################ +# GetSharedChats +############################ + + +@router.get("/shared", response_model=list[SharedChatResponse]) +async def get_shared_session_user_chat_list( + page: Optional[int] = None, + query: Optional[str] = None, + order_by: Optional[str] = None, + direction: Optional[str] = None, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + if page is None: + page = 1 + + limit = 60 + skip = (page - 1) * limit + + filter = {} + if query: + filter["query"] = query + if order_by: + filter["order_by"] = order_by + if direction: + filter["direction"] = direction + + chat_list = [ + SharedChatResponse(**chat.model_dump()) + for chat in Chats.get_shared_chat_list_by_user_id( + user.id, + filter=filter, + skip=skip, + limit=limit, + db=db, + ) + ] + + return chat_list + + ############################ # GetSharedChatById ############################ diff --git a/backend/open_webui/routers/configs.py b/backend/open_webui/routers/configs.py index 152e2c3edc..4dee4488cf 100644 --- a/backend/open_webui/routers/configs.py +++ b/backend/open_webui/routers/configs.py @@ -224,7 +224,7 @@ async def verify_tool_servers_config( try: if form_data.type == "mcp": if form_data.auth_type == "oauth_2.1": - discovery_urls = get_discovery_urls(form_data.url) + discovery_urls = await get_discovery_urls(form_data.url) for discovery_url in discovery_urls: log.debug( f"Trying to fetch OAuth 2.1 discovery document from {discovery_url}" diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index e3dd63525a..8b17dc406b 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -282,7 +282,11 @@ def upload_file_handler( }, "meta": { "name": name, - "content_type": file.content_type, + "content_type": ( + file.content_type + if isinstance(file.content_type, str) + else None + ), "size": len(contents), "data": file_metadata, }, @@ -332,6 +336,8 @@ def upload_file_handler( detail=ERROR_MESSAGES.DEFAULT("Error uploading file"), ) + except HTTPException as e: + raise e except Exception as e: log.exception(e) raise HTTPException( @@ -575,7 +581,7 @@ class ContentForm(BaseModel): @router.post("/{id}/data/content/update") -async def update_file_data_content_by_id( +def update_file_data_content_by_id( request: Request, id: str, form_data: ContentForm, @@ -825,6 +831,23 @@ async def delete_file_by_id( or has_access_to_file(id, "write", user, db=db) ): + # Clean up KB associations and embeddings before deleting + knowledges = Knowledges.get_knowledges_by_file_id(id, db=db) + for knowledge in knowledges: + # Remove KB-file relationship + Knowledges.remove_file_from_knowledge_by_id(knowledge.id, id, db=db) + # Clean KB embeddings (same logic as /knowledge/{id}/file/remove) + try: + VECTOR_DB_CLIENT.delete( + collection_name=knowledge.id, filter={"file_id": id} + ) + if file.hash: + VECTOR_DB_CLIENT.delete( + collection_name=knowledge.id, filter={"hash": file.hash} + ) + except Exception as e: + log.debug(f"KB embedding cleanup for {knowledge.id}: {e}") + result = Files.delete_file_by_id(id, db=db) if result: try: diff --git a/backend/open_webui/routers/functions.py b/backend/open_webui/routers/functions.py index ad47318911..a31b958e21 100644 --- a/backend/open_webui/routers/functions.py +++ b/backend/open_webui/routers/functions.py @@ -19,6 +19,7 @@ from open_webui.utils.plugin import ( load_function_module_by_id, replace_imports, get_function_module_from_cache, + resolve_valves_schema_options, ) from open_webui.config import CACHE_DIR from open_webui.constants import ERROR_MESSAGES @@ -446,7 +447,10 @@ async def get_function_valves_spec_by_id( if hasattr(function_module, "Valves"): Valves = function_module.Valves - return Valves.schema() + schema = Valves.schema() + # Resolve dynamic options for select dropdowns + schema = resolve_valves_schema_options(Valves, schema, user) + return schema return None else: raise HTTPException( @@ -546,7 +550,10 @@ async def get_function_user_valves_spec_by_id( if hasattr(function_module, "UserValves"): UserValves = function_module.UserValves - return UserValves.schema() + schema = UserValves.schema() + # Resolve dynamic options for select dropdowns + schema = resolve_valves_schema_options(UserValves, schema, user) + return schema return None else: raise HTTPException( diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index 0fc6930b81..7fdd84b3fa 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -16,6 +16,7 @@ from fastapi.responses import FileResponse from open_webui.config import CACHE_DIR from open_webui.constants import ERROR_MESSAGES +from open_webui.retrieval.web.utils import validate_url from open_webui.env import ENABLE_FORWARD_USER_INFO_HEADERS from open_webui.models.chats import Chats @@ -881,6 +882,8 @@ async def image_edits( return data if data.startswith("http://") or data.startswith("https://"): + # Validate URL to prevent SSRF attacks against local/private networks + validate_url(data) r = await asyncio.to_thread(requests.get, data) r.raise_for_status() @@ -910,7 +913,8 @@ async def image_edits( if isinstance(form_data.image, str): form_data.image = await load_url_image(form_data.image) elif isinstance(form_data.image, list): - form_data.image = [await load_url_image(img) for img in form_data.image] + # Load all images in parallel for better performance + form_data.image = list(await asyncio.gather(*[load_url_image(img) for img in form_data.image])) except Exception as e: raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e)) diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py index 19d25685ad..fc24ccaf4b 100644 --- a/backend/open_webui/routers/prompts.py +++ b/backend/open_webui/routers/prompts.py @@ -5,18 +5,40 @@ from open_webui.models.prompts import ( PromptForm, PromptUserResponse, PromptAccessResponse, + PromptAccessListResponse, PromptModel, Prompts, ) +from open_webui.models.groups import Groups +from open_webui.models.prompt_history import ( + PromptHistories, + PromptHistoryModel, + PromptHistoryResponse, +) from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.access_control import has_access, has_permission from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL from open_webui.internal.db import get_session from sqlalchemy.orm import Session +from pydantic import BaseModel + + +class PromptVersionUpdateForm(BaseModel): + version_id: str + + +class PromptMetadataForm(BaseModel): + name: str + command: str + tags: Optional[list[str]] = None + router = APIRouter() +PAGE_ITEM_COUNT = 30 + + ############################ # GetPrompts ############################ @@ -34,26 +56,72 @@ async def get_prompts( return prompts -@router.get("/list", response_model=list[PromptAccessResponse]) -async def get_prompt_list( +@router.get("/tags", response_model=list[str]) +async def get_prompt_tags( user=Depends(get_verified_user), db: Session = Depends(get_session) ): if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL: - prompts = Prompts.get_prompts(db=db) + return Prompts.get_tags(db=db) else: prompts = Prompts.get_prompts_by_user_id(user.id, "read", db=db) + tags = set() + for prompt in prompts: + if prompt.tags: + tags.update(prompt.tags) + return sorted(list(tags)) - return [ - PromptAccessResponse( - **prompt.model_dump(), - write_access=( - (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) - or user.id == prompt.user_id - or has_access(user.id, "write", prompt.access_control, db=db) - ), - ) - for prompt in prompts - ] + +@router.get("/list", response_model=PromptAccessListResponse) +async def get_prompt_list( + query: Optional[str] = None, + view_option: Optional[str] = None, + tag: Optional[str] = None, + order_by: Optional[str] = None, + direction: Optional[str] = None, + page: Optional[int] = 1, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + limit = PAGE_ITEM_COUNT + + page = max(1, page) + skip = (page - 1) * limit + + filter = {} + if query: + filter["query"] = query + if view_option: + filter["view_option"] = view_option + if tag: + filter["tag"] = tag + if order_by: + filter["order_by"] = order_by + if direction: + filter["direction"] = direction + + if not (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL): + groups = Groups.get_groups_by_member_id(user.id, db=db) + if groups: + filter["group_ids"] = [group.id for group in groups] + + filter["user_id"] = user.id + + result = Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db) + + return PromptAccessListResponse( + items=[ + PromptAccessResponse( + **prompt.model_dump(), + write_access=( + (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) + or user.id == prompt.user_id + or has_access(user.id, "write", prompt.access_control, db=db) + ), + ) + for prompt in result.items + ], + total=result.total, + ) ############################ @@ -112,7 +180,7 @@ async def create_new_prompt( async def get_prompt_by_command( command: str, user=Depends(get_verified_user), db: Session = Depends(get_session) ): - prompt = Prompts.get_prompt_by_command(f"/{command}", db=db) + prompt = Prompts.get_prompt_by_command(command, db=db) if prompt: if ( @@ -128,29 +196,62 @@ async def get_prompt_by_command( or has_access(user.id, "write", prompt.access_control, db=db) ), ) - else: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=ERROR_MESSAGES.NOT_FOUND, - ) + + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) ############################ -# UpdatePromptByCommand +# GetPromptById ############################ -@router.post("/command/{command}/update", response_model=Optional[PromptModel]) -async def update_prompt_by_command( - command: str, +@router.get("/id/{prompt_id}", response_model=Optional[PromptAccessResponse]) +async def get_prompt_by_id( + prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session) +): + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + + if prompt: + if ( + user.role == "admin" + or prompt.user_id == user.id + or has_access(user.id, "read", prompt.access_control, db=db) + ): + return PromptAccessResponse( + **prompt.model_dump(), + write_access=( + (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) + or user.id == prompt.user_id + or has_access(user.id, "write", prompt.access_control, db=db) + ), + ) + + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + +############################ +# UpdatePromptById +############################ + + +@router.post("/id/{prompt_id}/update", response_model=Optional[PromptModel]) +async def update_prompt_by_id( + prompt_id: str, form_data: PromptForm, user=Depends(get_verified_user), db: Session = Depends(get_session), ): - prompt = Prompts.get_prompt_by_command(f"/{command}", db=db) + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + if not prompt: raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, + status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) @@ -165,29 +266,46 @@ async def update_prompt_by_command( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - prompt = Prompts.update_prompt_by_command(f"/{command}", form_data, db=db) - if prompt: - return prompt + # Check for command collision if command is being changed + if form_data.command != prompt.command: + existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db) + if existing_prompt and existing_prompt.id != prompt.id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Command '/{form_data.command}' is already in use by another prompt", + ) + + # Use the ID from the found prompt + updated_prompt = Prompts.update_prompt_by_id( + prompt.id, form_data, user.id, db=db + ) + if updated_prompt: + return updated_prompt else: raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT(), ) ############################ -# DeletePromptByCommand +# UpdatePromptMetadata ############################ -@router.delete("/command/{command}/delete", response_model=bool) -async def delete_prompt_by_command( - command: str, user=Depends(get_verified_user), db: Session = Depends(get_session) +@router.post("/id/{prompt_id}/update/meta", response_model=Optional[PromptModel]) +async def update_prompt_metadata( + prompt_id: str, + form_data: PromptMetadataForm, + user=Depends(get_verified_user), + db: Session = Depends(get_session), ): - prompt = Prompts.get_prompt_by_command(f"/{command}", db=db) + """Update prompt name and command only (no history created).""" + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + if not prompt: raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, + status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) @@ -201,5 +319,252 @@ async def delete_prompt_by_command( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - result = Prompts.delete_prompt_by_command(f"/{command}", db=db) + # Check for command collision if command is being changed + if form_data.command != prompt.command: + existing_prompt = Prompts.get_prompt_by_command(form_data.command, db=db) + if existing_prompt and existing_prompt.id != prompt.id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Command '/{form_data.command}' is already in use", + ) + + updated_prompt = Prompts.update_prompt_metadata( + prompt.id, form_data.name, form_data.command, form_data.tags, db=db + ) + if updated_prompt: + return updated_prompt + else: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT(), + ) + + +@router.post("/id/{prompt_id}/update/version", response_model=Optional[PromptModel]) +async def set_prompt_version( + prompt_id: str, + form_data: PromptVersionUpdateForm, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + if not prompt: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + if ( + prompt.user_id != user.id + and not has_access(user.id, "write", prompt.access_control, db=db) + and user.role != "admin" + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + updated_prompt = Prompts.update_prompt_version( + prompt.id, form_data.version_id, db=db + ) + if updated_prompt: + return updated_prompt + else: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT(), + ) + + +############################ +# DeletePromptById +############################ + + +@router.delete("/id/{prompt_id}/delete", response_model=bool) +async def delete_prompt_by_id( + prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session) +): + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + + if not prompt: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + if ( + prompt.user_id != user.id + and not has_access(user.id, "write", prompt.access_control, db=db) + and user.role != "admin" + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + result = Prompts.delete_prompt_by_id(prompt.id, db=db) return result + + +############################ +# Prompt History Endpoints +############################ + + +@router.get("/id/{prompt_id}/history", response_model=list[PromptHistoryResponse]) +async def get_prompt_history( + prompt_id: str, + page: int = 0, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + """Get version history for a prompt.""" + PAGE_SIZE = 20 + + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + + if not prompt: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + # Check read access + if not ( + user.role == "admin" + or prompt.user_id == user.id + or has_access(user.id, "read", prompt.access_control, db=db) + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + history = PromptHistories.get_history_by_prompt_id( + prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db + ) + return history + + +@router.get( + "/id/{prompt_id}/history/{history_id}", response_model=PromptHistoryModel +) +async def get_prompt_history_entry( + prompt_id: str, + history_id: str, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + """Get a specific version from history.""" + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + + if not prompt: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + # Check read access + if not ( + user.role == "admin" + or prompt.user_id == user.id + or has_access(user.id, "read", prompt.access_control, db=db) + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + history_entry = PromptHistories.get_history_entry_by_id(history_id, db=db) + if not history_entry or history_entry.prompt_id != prompt.id: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + return history_entry + + +@router.delete( + "/id/{prompt_id}/history/{history_id}", response_model=bool +) +async def delete_prompt_history_entry( + prompt_id: str, + history_id: str, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + """Delete a history entry. Cannot delete the active production version.""" + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + + if not prompt: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + # Check write access + if not ( + user.role == "admin" + or prompt.user_id == user.id + or has_access(user.id, "write", prompt.access_control, db=db) + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + # Cannot delete active production version + if prompt.version_id == history_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Cannot delete the active production version", + ) + + success = PromptHistories.delete_history_entry(history_id, db=db) + if not success: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + return success + + +@router.get("/id/{prompt_id}/history/diff") +async def get_prompt_diff( + prompt_id: str, + from_id: str, + to_id: str, + user=Depends(get_verified_user), + db: Session = Depends(get_session), +): + """Get diff between two versions.""" + prompt = Prompts.get_prompt_by_id(prompt_id, db=db) + + if not prompt: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + # Check read access + if not ( + user.role == "admin" + or prompt.user_id == user.id + or has_access(user.id, "read", prompt.access_control, db=db) + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + diff = PromptHistories.compute_diff(from_id, to_id, db=db) + if not diff: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="One or both history entries not found", + ) + + return diff diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 9e22ed77ad..72bd0a42b0 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -39,7 +39,7 @@ from langchain_core.documents import Document from open_webui.models.files import FileModel, FileUpdateForm, Files from open_webui.models.knowledge import Knowledges from open_webui.storage.provider import Storage -from open_webui.internal.db import get_session +from open_webui.internal.db import get_session, get_db from sqlalchemy.orm import Session @@ -76,6 +76,7 @@ from open_webui.retrieval.web.perplexity import search_perplexity from open_webui.retrieval.web.sougou import search_sougou from open_webui.retrieval.web.firecrawl import search_firecrawl from open_webui.retrieval.web.external import search_external +from open_webui.retrieval.web.yandex import search_yandex from open_webui.retrieval.utils import ( get_content_from_url, @@ -468,6 +469,7 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)): # Content extraction settings "CONTENT_EXTRACTION_ENGINE": request.app.state.config.CONTENT_EXTRACTION_ENGINE, "PDF_EXTRACT_IMAGES": request.app.state.config.PDF_EXTRACT_IMAGES, + "PDF_LOADER_MODE": request.app.state.config.PDF_LOADER_MODE, "DATALAB_MARKER_API_KEY": request.app.state.config.DATALAB_MARKER_API_KEY, "DATALAB_MARKER_API_BASE_URL": request.app.state.config.DATALAB_MARKER_API_BASE_URL, "DATALAB_MARKER_ADDITIONAL_CONFIG": request.app.state.config.DATALAB_MARKER_ADDITIONAL_CONFIG, @@ -579,6 +581,9 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)): "YOUTUBE_LOADER_LANGUAGE": request.app.state.config.YOUTUBE_LOADER_LANGUAGE, "YOUTUBE_LOADER_PROXY_URL": request.app.state.config.YOUTUBE_LOADER_PROXY_URL, "YOUTUBE_LOADER_TRANSLATION": request.app.state.YOUTUBE_LOADER_TRANSLATION, + "YANDEX_WEB_SEARCH_URL": request.app.state.config.YANDEX_WEB_SEARCH_URL, + "YANDEX_WEB_SEARCH_API_KEY": request.app.state.config.YANDEX_WEB_SEARCH_API_KEY, + "YANDEX_WEB_SEARCH_CONFIG": request.app.state.config.YANDEX_WEB_SEARCH_CONFIG, }, } @@ -642,6 +647,9 @@ class WebConfig(BaseModel): YOUTUBE_LOADER_LANGUAGE: Optional[List[str]] = None YOUTUBE_LOADER_PROXY_URL: Optional[str] = None YOUTUBE_LOADER_TRANSLATION: Optional[str] = None + YANDEX_WEB_SEARCH_URL: Optional[str] = None + YANDEX_WEB_SEARCH_API_KEY: Optional[str] = None + YANDEX_WEB_SEARCH_CONFIG: Optional[str] = None class ConfigForm(BaseModel): @@ -661,6 +669,7 @@ class ConfigForm(BaseModel): # Content extraction settings CONTENT_EXTRACTION_ENGINE: Optional[str] = None PDF_EXTRACT_IMAGES: Optional[bool] = None + PDF_LOADER_MODE: Optional[str] = None DATALAB_MARKER_API_KEY: Optional[str] = None DATALAB_MARKER_API_BASE_URL: Optional[str] = None @@ -790,6 +799,11 @@ async def update_rag_config( if form_data.PDF_EXTRACT_IMAGES is not None else request.app.state.config.PDF_EXTRACT_IMAGES ) + request.app.state.config.PDF_LOADER_MODE = ( + form_data.PDF_LOADER_MODE + if form_data.PDF_LOADER_MODE is not None + else request.app.state.config.PDF_LOADER_MODE + ) request.app.state.config.DATALAB_MARKER_API_KEY = ( form_data.DATALAB_MARKER_API_KEY if form_data.DATALAB_MARKER_API_KEY is not None @@ -1020,6 +1034,11 @@ async def update_rag_config( if form_data.TEXT_SPLITTER is not None else request.app.state.config.TEXT_SPLITTER ) + request.app.state.config.ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER = ( + form_data.ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER + if form_data.ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER is not None + else request.app.state.config.ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER + ) request.app.state.config.CHUNK_SIZE = ( form_data.CHUNK_SIZE if form_data.CHUNK_SIZE is not None @@ -1178,6 +1197,15 @@ async def update_rag_config( request.app.state.YOUTUBE_LOADER_TRANSLATION = ( form_data.web.YOUTUBE_LOADER_TRANSLATION ) + request.app.state.config.YANDEX_WEB_SEARCH_URL = ( + form_data.web.YANDEX_WEB_SEARCH_URL + ) + request.app.state.config.YANDEX_WEB_SEARCH_API_KEY = ( + form_data.web.YANDEX_WEB_SEARCH_API_KEY + ) + request.app.state.config.YANDEX_WEB_SEARCH_CONFIG = ( + form_data.web.YANDEX_WEB_SEARCH_CONFIG + ) return { "status": True, @@ -1194,6 +1222,7 @@ async def update_rag_config( # Content extraction settings "CONTENT_EXTRACTION_ENGINE": request.app.state.config.CONTENT_EXTRACTION_ENGINE, "PDF_EXTRACT_IMAGES": request.app.state.config.PDF_EXTRACT_IMAGES, + "PDF_LOADER_MODE": request.app.state.config.PDF_LOADER_MODE, "DATALAB_MARKER_API_KEY": request.app.state.config.DATALAB_MARKER_API_KEY, "DATALAB_MARKER_API_BASE_URL": request.app.state.config.DATALAB_MARKER_API_BASE_URL, "DATALAB_MARKER_ADDITIONAL_CONFIG": request.app.state.config.DATALAB_MARKER_ADDITIONAL_CONFIG, @@ -1303,6 +1332,9 @@ async def update_rag_config( "YOUTUBE_LOADER_LANGUAGE": request.app.state.config.YOUTUBE_LOADER_LANGUAGE, "YOUTUBE_LOADER_PROXY_URL": request.app.state.config.YOUTUBE_LOADER_PROXY_URL, "YOUTUBE_LOADER_TRANSLATION": request.app.state.YOUTUBE_LOADER_TRANSLATION, + "YANDEX_WEB_SEARCH_URL": request.app.state.config.YANDEX_WEB_SEARCH_URL, + "YANDEX_WEB_SEARCH_API_KEY": request.app.state.config.YANDEX_WEB_SEARCH_API_KEY, + "YANDEX_WEB_SEARCH_CONFIG": request.app.state.config.YANDEX_WEB_SEARCH_CONFIG, }, } @@ -1433,8 +1465,16 @@ def save_docs_to_vector_db( if result is not None and result.ids and len(result.ids) > 0: existing_doc_ids = result.ids[0] if existing_doc_ids: - log.info(f"Document with hash {metadata['hash']} already exists") - raise ValueError(ERROR_MESSAGES.DUPLICATE_CONTENT) + # Check if the existing document belongs to the same file + # If same file_id, this is a re-add/reindex - allow it + # If different file_id, this is a duplicate - block it + existing_file_id = None + if result.metadatas and result.metadatas[0]: + existing_file_id = result.metadatas[0][0].get("file_id") + + if existing_file_id != metadata.get("file_id"): + log.info(f"Document with hash {metadata['hash']} already exists") + raise ValueError(ERROR_MESSAGES.DUPLICATE_CONTENT) if split: if request.app.state.config.ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER: @@ -1602,6 +1642,9 @@ def process_file( ): """ Process a file and save its content to the vector database. + Process a file and save its content to the vector database. + Note: granular session management is used to prevent connection pool exhaustion. + The session is committed before external API calls, and updates use a fresh session. """ if user.role == "admin": file = Files.get_file_by_id(form_data.file_id, db=db) @@ -1765,6 +1808,12 @@ def process_file( } else: try: + # Commit any pending changes before the slow embedding step. + # Note: file is already a Pydantic model (not ORM), so no expunge needed. + db.commit() + + # External embedding API takes time (5-60s+). + # Subsequent updates use fresh sessions via get_db(). result = save_docs_to_vector_db( request, docs=docs, @@ -1780,27 +1829,29 @@ def process_file( log.info(f"added {len(docs)} items to collection {collection_name}") if result: - Files.update_file_metadata_by_id( - file.id, - { + # Fresh session for the final update. + with get_db() as session: + Files.update_file_metadata_by_id( + file.id, + { + "collection_name": collection_name, + }, + db=session, + ) + + Files.update_file_data_by_id( + file.id, + {"status": "completed"}, + db=session, + ) + Files.update_file_hash_by_id(file.id, hash, db=session) + + return { + "status": True, "collection_name": collection_name, - }, - db=db, - ) - - Files.update_file_data_by_id( - file.id, - {"status": "completed"}, - db=db, - ) - Files.update_file_hash_by_id(file.id, hash, db=db) - - return { - "status": True, - "collection_name": collection_name, - "filename": file.filename, - "content": text_content, - } + "filename": file.filename, + "content": text_content, + } else: raise Exception("Error saving document to vector database") except Exception as e: @@ -1808,13 +1859,15 @@ def process_file( except Exception as e: log.exception(e) - Files.update_file_data_by_id( - file.id, - {"status": "failed"}, - db=db, - ) - # Clear the hash so the file can be re-uploaded after fixing the issue - Files.update_file_hash_by_id(file.id, None, db=db) + # Fresh session for error status update. + with get_db() as session: + Files.update_file_data_by_id( + file.id, + {"status": "failed"}, + db=session, + ) + # Clear the hash so the file can be re-uploaded after fixing the issue + Files.update_file_hash_by_id(file.id, None, db=session) if "No pandoc was found" in str(e): raise HTTPException( @@ -2224,6 +2277,17 @@ def search_web( request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST, user=user, ) + elif engine == "yandex": + return search_yandex( + request, + request.app.state.config.YANDEX_WEB_SEARCH_URL, + request.app.state.config.YANDEX_WEB_SEARCH_API_KEY, + request.app.state.config.YANDEX_WEB_SEARCH_CONFIG, + query, + request.app.state.config.WEB_SEARCH_RESULT_COUNT, + request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST, + user=user, + ) else: raise Exception("No search engine API key found in environment variables") @@ -2711,9 +2775,7 @@ async def process_files_batch( # Update all files with collection name for file_update, file_result in zip(file_updates, file_results): - Files.update_file_by_id( - id=file_result.file_id, form_data=file_update - ) + Files.update_file_by_id(id=file_result.file_id, form_data=file_update) file_result.status = "completed" except Exception as e: @@ -2723,7 +2785,9 @@ async def process_files_batch( for file_result in file_results: file_result.status = "failed" file_errors.append( - BatchProcessFilesResult(file_id=file_result.file_id, error=str(e)) + BatchProcessFilesResult( + file_id=file_result.file_id, status="failed", error=str(e) + ) ) return BatchProcessFilesResponse(results=file_results, errors=file_errors) diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index 9070256770..681be3c7d2 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -352,18 +352,17 @@ def user_to_scim(user: UserModel, request: Request, db=None) -> SCIMUser: def group_to_scim(group: GroupModel, request: Request, db=None) -> SCIMGroup: """Convert internal Group model to SCIM Group""" member_ids = Groups.get_group_user_ids_by_id(group.id, db) or [] - members = [] - for user_id in member_ids: - user = Users.get_user_by_id(user_id, db=db) - if user: - members.append( - SCIMGroupMember( - value=user.id, - ref=f"{request.base_url}api/v1/scim/v2/Users/{user.id}", - display=user.name, - ) - ) + # Batch-fetch all users to avoid N+1 queries + users = Users.get_users_by_user_ids(member_ids, db=db) if member_ids else [] + members = [ + SCIMGroupMember( + value=user.id, + ref=f"{request.base_url}api/v1/scim/v2/Users/{user.id}", + display=user.name, + ) + for user in users + ] return SCIMGroup( id=group.id, diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 03018d24a1..7f9b23c7ce 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -25,6 +25,7 @@ from open_webui.utils.plugin import ( load_tool_module_by_id, replace_imports, get_tool_module_from_cache, + resolve_valves_schema_options, ) from open_webui.utils.tools import get_tool_specs from open_webui.utils.auth import get_admin_user, get_verified_user @@ -553,7 +554,10 @@ async def get_tools_valves_spec_by_id( if hasattr(tools_module, "Valves"): Valves = tools_module.Valves - return Valves.schema() + schema = Valves.schema() + # Resolve dynamic options for select dropdowns + schema = resolve_valves_schema_options(Valves, schema, user) + return schema return None else: raise HTTPException( @@ -662,7 +666,10 @@ async def get_tools_user_valves_spec_by_id( if hasattr(tools_module, "UserValves"): UserValves = tools_module.UserValves - return UserValves.schema() + schema = UserValves.schema() + # Resolve dynamic options for select dropdowns + schema = resolve_valves_schema_options(UserValves, schema, user) + return schema return None else: raise HTTPException( diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index d26e032727..3dca901c42 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -36,6 +36,7 @@ from open_webui.models.chats import Chats from open_webui.models.channels import Channels, ChannelMember, Channel from open_webui.models.messages import Messages, Message from open_webui.models.groups import Groups +from open_webui.utils.sanitize import strip_markdown_code_fences log = logging.getLogger(__name__) @@ -166,7 +167,7 @@ async def search_web( engine = __request__.app.state.config.WEB_SEARCH_ENGINE user = UserModel(**__user__) if __user__ else None - results = _search_web(__request__, engine, query, user) + results = await asyncio.to_thread(_search_web, __request__, engine, query, user) # Limit results results = results[:count] if results else [] @@ -370,6 +371,9 @@ async def execute_code( return json.dumps({"error": "Request context not available"}) try: + # Strip markdown fences if model included them + code = strip_markdown_code_fences(code) + # Import blocked modules from config (same as middleware) from open_webui.config import CODE_INTERPRETER_BLOCKED_MODULES @@ -892,6 +896,7 @@ async def search_chats( end_timestamp: Optional[int] = None, __request__: Request = None, __user__: dict = None, + __chat_id__: str = None, ) -> str: """ Search the user's previous chat conversations by title and message content. @@ -921,6 +926,10 @@ async def search_chats( results = [] for chat in chats: + # Skip the current chat to avoid showing it in search results + if __chat_id__ and chat.id == __chat_id__: + continue + # Apply date filters (updated_at is in seconds) if start_timestamp and chat.updated_at < start_timestamp: continue @@ -1599,6 +1608,25 @@ async def query_knowledge_files( if not __user__: return json.dumps({"error": "User context not available"}) + # Coerce parameters from LLM tool calls (may come as strings) + if isinstance(count, str): + try: + count = int(count) + except ValueError: + count = 5 # Default fallback + + # Handle knowledge_ids being string "None", "null", or empty + if isinstance(knowledge_ids, str): + if knowledge_ids.lower() in ("none", "null", ""): + knowledge_ids = None + else: + # Try to parse as JSON array if it looks like one + try: + knowledge_ids = json.loads(knowledge_ids) + except json.JSONDecodeError: + # Treat as single ID + knowledge_ids = [knowledge_ids] + try: from open_webui.models.knowledge import Knowledges from open_webui.models.files import Files diff --git a/backend/open_webui/utils/audit.py b/backend/open_webui/utils/audit.py index dc1226a080..73dc140de5 100644 --- a/backend/open_webui/utils/audit.py +++ b/backend/open_webui/utils/audit.py @@ -221,7 +221,8 @@ class AuditLoggingMiddleware: return False # Do NOT skip logging for auth endpoints # Skip logging if the request is not authenticated - if not request.headers.get("authorization"): + # Check both Authorization header (API keys) and token cookie (browser sessions) + if not request.headers.get("authorization") and not request.cookies.get("token"): return True # match either /api//...(for the endpoint /api/chat case) or /api/v1//... diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index c1f6910ddb..ef09a6004d 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -33,6 +33,7 @@ from open_webui.env import ( ENABLE_PASSWORD_VALIDATION, OFFLINE_MODE, LICENSE_BLOB, + PASSWORD_VALIDATION_HINT, PASSWORD_VALIDATION_REGEX_PATTERN, REDIS_KEY_PREFIX, pk, @@ -173,7 +174,7 @@ def validate_password(password: str) -> bool: if ENABLE_PASSWORD_VALIDATION: if not PASSWORD_VALIDATION_REGEX_PATTERN.match(password): - raise Exception(ERROR_MESSAGES.INVALID_PASSWORD()) + raise Exception(ERROR_MESSAGES.INVALID_PASSWORD(PASSWORD_VALIDATION_HINT)) return True diff --git a/backend/open_webui/utils/chat.py b/backend/open_webui/utils/chat.py index be700dda76..1ed34b6ca7 100644 --- a/backend/open_webui/utils/chat.py +++ b/backend/open_webui/utils/chat.py @@ -341,12 +341,8 @@ async def chat_completed(request: Request, form_data: dict, user: Any): } try: - filter_functions = [ - Functions.get_function_by_id(filter_id) - for filter_id in get_sorted_filter_ids( - request, model, metadata.get("filter_ids", []) - ) - ] + filter_ids = get_sorted_filter_ids(request, model, metadata.get("filter_ids", [])) + filter_functions = Functions.get_functions_by_ids(filter_ids) result, _ = await process_filter_functions( request=request, diff --git a/backend/open_webui/utils/files.py b/backend/open_webui/utils/files.py index a37ecf31c6..af8818d59b 100644 --- a/backend/open_webui/utils/files.py +++ b/backend/open_webui/utils/files.py @@ -18,6 +18,7 @@ from open_webui.storage.provider import Storage from open_webui.models.chats import Chats from open_webui.models.files import Files from open_webui.routers.files import upload_file_handler +from open_webui.retrieval.web.utils import validate_url import mimetypes import base64 @@ -33,6 +34,8 @@ MARKDOWN_IMAGE_URL_PATTERN = re.compile(r"!\[(.*?)\]\((.+?)\)", re.IGNORECASE) def get_image_base64_from_url(url: str) -> Optional[str]: try: if url.startswith("http"): + # Validate URL to prevent SSRF attacks against local/private networks + validate_url(url) # Download the image from the URL response = requests.get(url) response.raise_for_status() diff --git a/backend/open_webui/utils/mcp/client.py b/backend/open_webui/utils/mcp/client.py index 6edfca4f6c..33803648be 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -8,6 +8,15 @@ from mcp import ClientSession from mcp.client.auth import OAuthClientProvider, TokenStorage from mcp.client.streamable_http import streamablehttp_client from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken +import httpx +from mcp.shared._httpx_utils import create_mcp_http_client +from open_webui.env import AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL + + +def create_insecure_httpx_client(headers=None, timeout=None, auth=None): + client = create_mcp_http_client(headers=headers, timeout=timeout, auth=auth) + client.verify = False + return client class MCPClient: @@ -18,7 +27,14 @@ class MCPClient: async def connect(self, url: str, headers: Optional[dict] = None): async with AsyncExitStack() as exit_stack: try: - self._streams_context = streamablehttp_client(url, headers=headers) + if AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL: + self._streams_context = streamablehttp_client(url, headers=headers) + else: + self._streams_context = streamablehttp_client( + url, + headers=headers, + httpx_client_factory=create_insecure_httpx_client, + ) transport = await exit_stack.enter_async_context(self._streams_context) read_stream, write_stream, _ = transport diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index fe2d7e5dc1..81e07df94e 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -73,6 +73,7 @@ from open_webui.models.models import Models from open_webui.retrieval.utils import get_sources_from_items +from open_webui.utils.sanitize import strip_markdown_code_fences from open_webui.utils.chat import generate_chat_completion from open_webui.utils.task import ( get_task_model_id, @@ -92,6 +93,7 @@ from open_webui.utils.misc import ( prepend_to_first_user_message_content, convert_logit_bias_input_to_json, get_content_from_message, + convert_output_to_messages, ) from open_webui.utils.tools import ( get_tools, @@ -146,6 +148,11 @@ DEFAULT_SOLUTION_TAGS = [("<|begin_of_solution|>", "<|end_of_solution|>")] DEFAULT_CODE_INTERPRETER_TAGS = [("", "")] +def output_id(prefix: str) -> str: + """Generate OR-style ID: prefix + 24-char hex UUID.""" + return f"{prefix}_{uuid4().hex[:24]}" + + def get_citation_source_from_tool_result( tool_name: str, tool_params: dict, tool_result: str, tool_id: str = "" ) -> list[dict]: @@ -1399,6 +1406,30 @@ async def convert_url_images_to_base64(form_data): return form_data +def process_messages_with_output(messages: list[dict]) -> list[dict]: + """ + Process messages with OR-aligned output items for LLM consumption. + + For assistant messages with 'output' field, produces properly formatted + OpenAI-style messages (tool_calls + tool results). Strips 'output' before LLM. + """ + processed = [] + + for message in messages: + if message.get("role") == "assistant" and message.get("output"): + # Use output items for clean OpenAI-format messages + output_messages = convert_output_to_messages(message["output"]) + if output_messages: + processed.extend(output_messages) + continue + + # Strip 'output' field before adding (LLM shouldn't see it) + clean_message = {k: v for k, v in message.items() if k != "output"} + processed.append(clean_message) + + return processed + + async def process_chat_payload(request, form_data, user, metadata, model): # Pipeline Inlet -> Filter Inlet -> Chat Memory -> Chat Web Search -> Chat Image Generation # -> Chat Code Interpreter (Form Data Update) -> (Default) Chat Tools Function Calling @@ -1407,6 +1438,9 @@ async def process_chat_payload(request, form_data, user, metadata, model): form_data = apply_params_to_form_data(form_data, model) log.debug(f"form_data: {form_data}") + # Process messages with OR-aligned output items for clean LLM messages + form_data["messages"] = process_messages_with_output(form_data.get("messages", [])) + system_message = get_system_message(form_data.get("messages", [])) if system_message: # Chat Controls/User Settings try: @@ -1536,12 +1570,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): raise e try: - filter_functions = [ - Functions.get_function_by_id(filter_id) - for filter_id in get_sorted_filter_ids( - request, model, metadata.get("filter_ids", []) - ) - ] + filter_ids = get_sorted_filter_ids(request, model, metadata.get("filter_ids", [])) + filter_functions = Functions.get_functions_by_ids(filter_ids) form_data, flags = await process_filter_functions( request=request, @@ -1808,11 +1838,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Inject builtin tools for native function calling based on enabled features and model capability # Check if builtin_tools capability is enabled for this model (defaults to True if not specified) builtin_tools_enabled = ( - model.get("info", {}) - .get("meta", {}) - .get("capabilities", {}) - .get("builtin_tools", True) - ) + model.get("info", {}).get("meta", {}).get("capabilities") or {} + ).get("builtin_tools", True) if ( metadata.get("params", {}).get("function_calling") == "native" and builtin_tools_enabled @@ -1856,11 +1883,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Check if file context extraction is enabled for this model (default True) file_context_enabled = ( - model.get("info", {}) - .get("meta", {}) - .get("capabilities", {}) - .get("file_context", True) - ) + model.get("info", {}).get("meta", {}).get("capabilities") or {} + ).get("file_context", True) if file_context_enabled: try: @@ -2210,12 +2234,20 @@ async def process_chat_response( title = Chats.get_chat_title_by_id(metadata["chat_id"]) + # Use output from backend if provided (OR-compliant backends) + response_output = response_data.get("output") + await event_emitter( { "type": "chat:completion", "data": { "done": True, "content": content, + **( + {"output": response_output} + if response_output + else {} + ), "title": title, }, } @@ -2228,6 +2260,11 @@ async def process_chat_response( { "role": "assistant", "content": content, + **( + {"output": response_output} + if response_output + else {} + ), }, ) @@ -2481,6 +2518,107 @@ async def process_chat_response( return content.strip() + def serialize_output(output: list) -> str: + """ + Convert OR-aligned output items to HTML for display. + For LLM consumption, use convert_output_to_messages() instead. + """ + content = "" + + # First pass: collect function_call_output items by call_id for lookup + tool_outputs = {} + for item in output: + if item.get("type") == "function_call_output": + tool_outputs[item.get("call_id")] = item + + # Second pass: render items in order + for item in output: + item_type = item.get("type", "") + + if item_type == "message": + for content_part in item.get("content", []): + if content_part.get("type") == "output_text": + text = content_part.get("text", "").strip() + if text: + content = f"{content}{text}\n" + + elif item_type == "function_call": + # Render tool call inline with its result (if available) + if content and not content.endswith("\n"): + content += "\n" + + call_id = item.get("call_id", "") + name = item.get("name", "") + arguments = item.get("arguments", "") + + result_item = tool_outputs.get(call_id) + if result_item: + result_text = "" + for out in result_item.get("output", []): + if out.get("type") == "input_text": + result_text += out.get("text", "") + files = result_item.get("files") + embeds = result_item.get("embeds", "") + + content += f'
\nTool Executed\n
\n' + else: + content += f'
\nExecuting...\n
\n' + + elif item_type == "function_call_output": + # Already handled inline with function_call above + pass + + elif item_type == "reasoning": + reasoning_content = "" + for content_part in item.get("content", []): + if content_part.get("type") == "output_text": + reasoning_content = content_part.get("text", "").strip() + + duration = item.get("duration") + status = item.get("status", "in_progress") + + if content and not content.endswith("\n"): + content += "\n" + + display = html.escape( + "\n".join( + (f"> {line}" if not line.startswith(">") else line) + for line in reasoning_content.splitlines() + ) + ) + + if status == "completed" or duration is not None: + content = f'{content}
\nThought for {duration or 0} seconds\n{display}\n
\n' + else: + content = f'{content}
\nThinking…\n{display}\n
\n' + + elif item_type == "open_webui:code_interpreter": + code = item.get("code", "") + output_val = item.get("output") + lang = item.get("lang", "") + + content_stripped, original_whitespace = ( + split_content_and_whitespace(content) + ) + if is_opening_code_block(content_stripped): + content = ( + content_stripped.rstrip("`").rstrip() + + original_whitespace + ) + else: + content = content_stripped + original_whitespace + + if content and not content.endswith("\n"): + content += "\n" + + if output_val: + output_escaped = html.escape(json.dumps(output_val)) + content = f'{content}
\nAnalyzed\n```{lang}\n{code}\n```\n
\n' + else: + content = f'{content}
\nAnalyzing...\n```{lang}\n{code}\n```\n
\n' + + return content.strip() + def convert_content_blocks_to_messages(content_blocks, raw=False): messages = [] @@ -2521,6 +2659,127 @@ async def process_chat_response( return messages + def convert_content_blocks_to_output(content_blocks): + """ + Convert content_blocks to Open Responses-aligned output items. + See: https://openresponses.org/specification + """ + output_items = [] + + def next_id(prefix): + return f"{prefix}_{uuid4().hex[:24]}" + + for block in content_blocks: + block_type = block.get("type", "") + # Use backend-provided ID if available, fallback to generated + block_id = block.get("id") + + if block_type == "text": + text_content = block.get("content", "").strip() + if text_content: + output_items.append( + { + "type": "message", + "id": block_id or next_id("msg"), + "status": "completed", + "role": "assistant", + "content": [ + {"type": "output_text", "text": text_content} + ], + } + ) + + elif block_type == "tool_calls": + tool_calls = block.get("content", []) + results = block.get("results", []) + + # Emit function_call items + for tool_call in tool_calls: + call_id = tool_call.get("id", "") + func = tool_call.get("function", {}) + output_items.append( + { + "type": "function_call", + "id": call_id + or next_id( + "fc" + ), # Use call_id as item id if available + "call_id": call_id, + "name": func.get("name", ""), + "arguments": func.get("arguments", "{}"), + "status": "completed" if results else "in_progress", + } + ) + + # Emit function_call_output items + for result in results: + output_items.append( + { + "type": "function_call_output", + "id": result.get("id") or next_id("fco"), + "call_id": result.get("tool_call_id", ""), + "output": [ + { + "type": "input_text", + "text": result.get("content", ""), + } + ], + "status": "completed", + **( + {"files": result.get("files")} + if result.get("files") + else {} + ), + **( + {"embeds": result.get("embeds")} + if result.get("embeds") + else {} + ), + } + ) + + elif block_type == "reasoning": + reasoning_content = block.get("content", "").strip() + duration = block.get("duration") + output_items.append( + { + "type": "reasoning", + "id": block_id or next_id("r"), + "status": ( + "completed" + if duration is not None + else "in_progress" + ), + "content": ( + [{"type": "output_text", "text": reasoning_content}] + if reasoning_content + else None + ), + "summary": None, + } + ) + + elif block_type == "code_interpreter": + code = block.get("content", "") + output_val = block.get("output") + attrs = block.get("attributes", {}) + output_items.append( + { + "type": "open_webui:code_interpreter", + "id": block_id or next_id("ci"), + "status": ( + "completed" + if output_val is not None + else "in_progress" + ), + "lang": attrs.get("lang", ""), + "code": code, + "output": output_val, + } + ) + + return output_items + def tag_content_handler(content_type, tags, content, content_blocks): end_flag = False @@ -2718,6 +2977,23 @@ async def process_chat_response( else last_assistant_message if last_assistant_message else "" ) + # Initialize output: use existing from message if continuing, else create new + existing_output = message.get("output") if message else None + if existing_output: + output = existing_output + else: + # Always create an initial message item (even if content is empty) + output = [ + { + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": content}], + } + ] + + # Keep content_blocks for backward compatibility during transition content_blocks = [ { "type": "text", @@ -3129,13 +3405,15 @@ async def process_chat_response( if ENABLE_REALTIME_CHAT_SAVE: # Save message in the database + output = convert_content_blocks_to_output( + content_blocks + ) Chats.upsert_message_to_chat_by_id_and_message_id( metadata["chat_id"], metadata["message_id"], { - "content": serialize_content_blocks( - content_blocks - ), + "content": serialize_output(output), + "output": output, }, ) else: @@ -3220,11 +3498,13 @@ async def process_chat_response( } ) + output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", "data": { - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, }, } ) @@ -3399,11 +3679,13 @@ async def process_chat_response( ) tool_call_sources.clear() + output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", "data": { - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, }, } ) @@ -3445,11 +3727,13 @@ async def process_chat_response( and retries < MAX_RETRIES ): + output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", "data": { - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, }, } ) @@ -3461,6 +3745,9 @@ async def process_chat_response( try: if content_blocks[-1]["attributes"].get("type") == "code": code = content_blocks[-1]["content"] + # Strip markdown fences if model included them + code = strip_markdown_code_fences(code) + if CODE_INTERPRETER_BLOCKED_MODULES: blocking_code = textwrap.dedent( f""" @@ -3576,11 +3863,13 @@ async def process_chat_response( } ) + output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", "data": { - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, }, } ) @@ -3617,9 +3906,11 @@ async def process_chat_response( break title = Chats.get_chat_title_by_id(metadata["chat_id"]) + output = convert_content_blocks_to_output(content_blocks) data = { "done": True, - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, "title": title, } @@ -3629,7 +3920,8 @@ async def process_chat_response( metadata["chat_id"], metadata["message_id"], { - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, }, ) @@ -3663,11 +3955,13 @@ async def process_chat_response( if not ENABLE_REALTIME_CHAT_SAVE: # Save message in the database + output = convert_content_blocks_to_output(content_blocks) Chats.upsert_message_to_chat_by_id_and_message_id( metadata["chat_id"], metadata["message_id"], { - "content": serialize_content_blocks(content_blocks), + "content": serialize_output(output), + "output": output, }, ) diff --git a/backend/open_webui/utils/misc.py b/backend/open_webui/utils/misc.py index e293f3d257..b931476ca9 100644 --- a/backend/open_webui/utils/misc.py +++ b/backend/open_webui/utils/misc.py @@ -128,6 +128,86 @@ def get_content_from_message(message: dict) -> Optional[str]: return None +def convert_output_to_messages(output: list) -> list[dict]: + """ + Convert OR-aligned output items to OpenAI-format messages for LLM consumption. + + This is the inverse of convert_content_blocks_to_output() in middleware.py. + """ + if not output or not isinstance(output, list): + return [] + + messages = [] + pending_tool_calls = [] + pending_content = [] + + for item in output: + item_type = item.get("type", "") + + if item_type == "message": + # Extract text from output_text content parts + content_parts = item.get("content", []) + text = "" + for part in content_parts: + if part.get("type") == "output_text": + text += part.get("text", "") + if text: + pending_content.append(text) + + elif item_type == "function_call": + # Collect tool calls to batch into assistant message + pending_tool_calls.append({ + "id": item.get("call_id", ""), + "type": "function", + "function": { + "name": item.get("name", ""), + "arguments": item.get("arguments", "{}"), + } + }) + + elif item_type == "function_call_output": + # Flush any pending content/tool_calls before adding tool result + if pending_content or pending_tool_calls: + messages.append({ + "role": "assistant", + "content": "\n".join(pending_content) if pending_content else "", + **({"tool_calls": pending_tool_calls} if pending_tool_calls else {}), + }) + pending_content = [] + pending_tool_calls = [] + + # Extract text from output content parts + output_parts = item.get("output", []) + content = "" + for part in output_parts: + if part.get("type") == "input_text": + content += part.get("text", "") + + messages.append({ + "role": "tool", + "tool_call_id": item.get("call_id", ""), + "content": content, + }) + + elif item_type == "reasoning": + # Skip reasoning blocks for LLM messages + pass + + elif item_type.startswith("open_webui:"): + # Skip extension types + pass + + # Flush remaining content/tool_calls + if pending_content or pending_tool_calls: + messages.append({ + "role": "assistant", + "content": "\n".join(pending_content) if pending_content else "", + **({"tool_calls": pending_tool_calls} if pending_tool_calls else {}), + }) + + return messages + + def get_last_user_message(messages: list[dict]) -> Optional[str]: message = get_last_user_message_item(messages) if message is None: diff --git a/backend/open_webui/utils/payload.py b/backend/open_webui/utils/payload.py index 458687b371..5094c910ca 100644 --- a/backend/open_webui/utils/payload.py +++ b/backend/open_webui/utils/payload.py @@ -287,7 +287,11 @@ def convert_payload_openai_to_ollama(openai_payload: dict) -> dict: Returns: dict: A modified payload compatible with the Ollama API. """ - openai_payload = copy.deepcopy(openai_payload) + # Shallow copy metadata separately (may contain non-picklable objects) + metadata = openai_payload.get("metadata") + openai_payload = copy.deepcopy({k: v for k, v in openai_payload.items() if k != "metadata"}) + if metadata is not None: + openai_payload["metadata"] = dict(metadata) ollama_payload = {} # Mapping basic model and message details diff --git a/backend/open_webui/utils/plugin.py b/backend/open_webui/utils/plugin.py index 965b0a688f..79a1c0f0dc 100644 --- a/backend/open_webui/utils/plugin.py +++ b/backend/open_webui/utils/plugin.py @@ -6,6 +6,7 @@ from importlib import util import types import tempfile import logging +from typing import Any from open_webui.env import PIP_OPTIONS, PIP_PACKAGE_INDEX_OPTIONS, OFFLINE_MODE from open_webui.models.functions import Functions @@ -14,6 +15,143 @@ from open_webui.models.tools import Tools log = logging.getLogger(__name__) +def resolve_valves_schema_options( + valves_class: type, schema: dict, user: Any = None +) -> dict: + """ + Resolve dynamic options in a Valves schema. + + For properties with `input.options`, this function handles two cases: + - List: Used directly as dropdown options + - String: Treated as method name, called to get options dynamically + + Usage in Valves: + class UserValves(BaseModel): + # Static options + priority: str = Field( + default="medium", + json_schema_extra={ + "input": { + "type": "select", + "options": ["low", "medium", "high"] + } + } + ) + + # Dynamic options (method name) + model: str = Field( + default="", + json_schema_extra={ + "input": { + "type": "select", + "options": "get_model_options" + } + } + ) + + @classmethod + def get_model_options(cls, __user__=None) -> list[dict]: + return [{"value": "gpt-4", "label": "GPT-4"}] + + Args: + valves_class: The Valves or UserValves Pydantic model class + schema: The JSON schema dict from valves_class.schema() + user: Optional user object passed to methods that accept __user__ + + Returns: + Modified schema dict with resolved options + """ + if not schema or "properties" not in schema: + return schema + + # Make a copy to avoid mutating the original + schema = dict(schema) + schema["properties"] = dict(schema.get("properties", {})) + + for prop_name, prop_schema in list(schema["properties"].items()): + # Get the original field info from the Pydantic model + if not hasattr(valves_class, "model_fields"): + continue + + field_info = valves_class.model_fields.get(prop_name) + if not field_info: + continue + + # Check json_schema_extra for options + json_schema_extra = field_info.json_schema_extra + if not json_schema_extra or not isinstance(json_schema_extra, dict): + continue + + input_config = json_schema_extra.get("input") + if not input_config or not isinstance(input_config, dict): + continue + + options = input_config.get("options") + if options is None: + continue + + resolved_options = None + + # Case 1: options is already a list - use directly + if isinstance(options, list): + resolved_options = options + + # Case 2: options is a string - treat as method name + elif isinstance(options, str) and options: + method = getattr(valves_class, options, None) + if method is None or not callable(method): + log.warning( + f"options '{options}' not found or not callable on {valves_class.__name__}" + ) + continue + + try: + import inspect + + sig = inspect.signature(method) + params = sig.parameters + + # Prepare kwargs based on what the method accepts + kwargs = {} + if "__user__" in params and user is not None: + kwargs["__user__"] = ( + user.model_dump() if hasattr(user, "model_dump") else user + ) + if "user" in params and user is not None: + kwargs["user"] = ( + user.model_dump() if hasattr(user, "model_dump") else user + ) + + resolved_options = method(**kwargs) if kwargs else method() + + # Validate return type + if not isinstance(resolved_options, list): + log.warning( + f"Method '{options}' did not return a list for {prop_name}" + ) + continue + + except Exception as e: + log.warning(f"Failed to resolve options for {prop_name}: {e}") + continue + else: + # Invalid options type - skip + continue + + # Update the schema with resolved options + schema["properties"][prop_name] = dict(prop_schema) + if "input" not in schema["properties"][prop_name]: + schema["properties"][prop_name]["input"] = {"type": "select"} + else: + schema["properties"][prop_name]["input"] = dict( + schema["properties"][prop_name].get("input", {}) + ) + schema["properties"][prop_name]["input"]["options"] = resolved_options + + return schema + + + def extract_frontmatter(content): """ Extract frontmatter as a dictionary from the provided content string. diff --git a/backend/open_webui/utils/redis.py b/backend/open_webui/utils/redis.py index 2040633a7b..fcc4879ba3 100644 --- a/backend/open_webui/utils/redis.py +++ b/backend/open_webui/utils/redis.py @@ -1,5 +1,7 @@ import inspect from urllib.parse import urlparse +import asyncio +import time import logging @@ -12,6 +14,7 @@ from open_webui.env import ( REDIS_SENTINEL_MAX_RETRY_COUNT, REDIS_SENTINEL_PORT, REDIS_URL, + REDIS_RECONNECT_DELAY, ) log = logging.getLogger(__name__) @@ -63,6 +66,8 @@ class SentinelRedisProxy: i + 1, REDIS_SENTINEL_MAX_RETRY_COUNT, ) + if REDIS_RECONNECT_DELAY: + time.sleep(REDIS_RECONNECT_DELAY / 1000) continue log.error( "Redis operation failed after %s retries: %s", @@ -94,6 +99,8 @@ class SentinelRedisProxy: i + 1, REDIS_SENTINEL_MAX_RETRY_COUNT, ) + if REDIS_RECONNECT_DELAY: + await asyncio.sleep(REDIS_RECONNECT_DELAY / 1000) continue log.error( "Redis operation failed after %s retries: %s", @@ -122,6 +129,8 @@ class SentinelRedisProxy: i + 1, REDIS_SENTINEL_MAX_RETRY_COUNT, ) + if REDIS_RECONNECT_DELAY: + time.sleep(REDIS_RECONNECT_DELAY / 1000) continue log.error( "Redis operation failed after %s retries: %s", diff --git a/backend/open_webui/utils/sanitize.py b/backend/open_webui/utils/sanitize.py new file mode 100644 index 0000000000..5755a9be48 --- /dev/null +++ b/backend/open_webui/utils/sanitize.py @@ -0,0 +1,21 @@ +import re + + +def strip_markdown_code_fences(code: str) -> str: + """ + Strip markdown code fences if present. + + This is a defensive, non-breaking change — if the code doesn't + contain fences, it passes through unchanged. + + Handles patterns like: + - ```python + - ```py + - ``` + """ + code = code.strip() + # Remove opening fence (```python, ```py, ``` etc.) + code = re.sub(r"^```\w*\n?", "", code) + # Remove closing fence + code = re.sub(r"\n?```\s*$", "", code) + return code.strip() diff --git a/backend/open_webui/utils/task.py b/backend/open_webui/utils/task.py index ecedd595a7..4acb2c3718 100644 --- a/backend/open_webui/utils/task.py +++ b/backend/open_webui/utils/task.py @@ -69,6 +69,7 @@ def prompt_template(template: str, user: Optional[Any] = None) -> str: USER_VARIABLES = { "name": str(user.get("name")), + "email": str(user.get("email")), "location": str(user_info.get("location")), "bio": str(user.get("bio")), "gender": str(user.get("gender")), @@ -92,6 +93,9 @@ def prompt_template(template: str, user: Optional[Any] = None) -> str: template = template.replace("{{CURRENT_WEEKDAY}}", formatted_weekday) template = template.replace("{{USER_NAME}}", USER_VARIABLES.get("name", "Unknown")) + template = template.replace( + "{{USER_EMAIL}}", USER_VARIABLES.get("email", "Unknown") + ) template = template.replace("{{USER_BIO}}", USER_VARIABLES.get("bio", "Unknown")) template = template.replace( "{{USER_GENDER}}", USER_VARIABLES.get("gender", "Unknown") diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 77f985ebbc..b2f95a8297 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -398,74 +398,87 @@ def get_builtin_tools( # Helper to get model capabilities (defaults to True if not specified) def get_model_capability(name: str, default: bool = True) -> bool: return ( - model.get("info", {}) - .get("meta", {}) - .get("capabilities", {}) + (model.get("info", {}).get("meta", {}).get("capabilities") or {}) .get(name, default) ) - # Time utilities - always available for date calculations - builtin_functions.extend([get_current_timestamp, calculate_timestamp]) + # Helper to check if a builtin tool category is enabled via meta.builtinTools + # Defaults to True if not specified (backward compatible) + def is_builtin_tool_enabled(category: str) -> bool: + builtin_tools = model.get("info", {}).get("meta", {}).get("builtinTools", {}) + return builtin_tools.get(category, True) + + # Time utilities - available for date calculations + if is_builtin_tool_enabled("time"): + builtin_functions.extend([get_current_timestamp, calculate_timestamp]) # Knowledge base tools - conditional injection based on model knowledge # If model has attached knowledge (any type), only provide query_knowledge_files # Otherwise, provide all KB browsing tools model_knowledge = model.get("info", {}).get("meta", {}).get("knowledge", []) - if model_knowledge: - # Model has attached knowledge - only allow semantic search within it - builtin_functions.append(query_knowledge_files) - else: - # No model knowledge - allow full KB browsing - builtin_functions.extend( - [ - list_knowledge_bases, - search_knowledge_bases, - query_knowledge_bases, - search_knowledge_files, - query_knowledge_files, - view_knowledge_file, - ] - ) + if is_builtin_tool_enabled("knowledge"): + if model_knowledge: + # Model has attached knowledge - only allow semantic search within it + builtin_functions.append(query_knowledge_files) + else: + # No model knowledge - allow full KB browsing + builtin_functions.extend( + [ + list_knowledge_bases, + search_knowledge_bases, + query_knowledge_bases, + search_knowledge_files, + query_knowledge_files, + view_knowledge_file, + ] + ) # Chats tools - search and fetch user's chat history - builtin_functions.extend([search_chats, view_chat]) + if is_builtin_tool_enabled("chats"): + builtin_functions.extend([search_chats, view_chat]) - # Add memory tools if enabled for this chat - if features.get("memory"): + # Add memory tools if builtin category enabled AND enabled for this chat + if is_builtin_tool_enabled("memory") and features.get("memory"): builtin_functions.extend([search_memories, add_memory, replace_memory_content]) - # Add web search tools if enabled globally AND model has web_search capability - if getattr( - request.app.state.config, "ENABLE_WEB_SEARCH", False - ) and get_model_capability("web_search"): + # 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 get_model_capability("web_search") + ): builtin_functions.extend([search_web, fetch_url]) - # Add image generation/edit tools if enabled globally AND model has image_generation capability - if getattr( - request.app.state.config, "ENABLE_IMAGE_GENERATION", False - ) and get_model_capability("image_generation"): + # 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 get_model_capability("image_generation") + ): builtin_functions.append(generate_image) - if getattr( - request.app.state.config, "ENABLE_IMAGE_EDIT", False - ) and get_model_capability("image_generation"): + if ( + is_builtin_tool_enabled("image_generation") + and getattr(request.app.state.config, "ENABLE_IMAGE_EDIT", False) + and get_model_capability("image_generation") + ): builtin_functions.append(edit_image) - # Add code interpreter tool if enabled globally AND model has code_interpreter capability - # Supports both pyodide (via frontend event call) and jupyter engines + # Add code interpreter tool if builtin category enabled AND enabled globally AND model has code_interpreter capability if ( - getattr(request.app.state.config, "ENABLE_CODE_INTERPRETER", True) + is_builtin_tool_enabled("code_interpreter") + and getattr(request.app.state.config, "ENABLE_CODE_INTERPRETER", True) and get_model_capability("code_interpreter") ): builtin_functions.append(execute_code) - # Notes tools - search, view, create, and update user's notes (if notes enabled globally) - if getattr(request.app.state.config, "ENABLE_NOTES", False): + # Notes tools - search, view, create, and update user's notes (if builtin category enabled AND notes enabled globally) + if is_builtin_tool_enabled("notes") and getattr(request.app.state.config, "ENABLE_NOTES", False): builtin_functions.extend( [search_notes, view_note, write_note, replace_note_content] ) - # Channels tools - search channels and messages (if channels enabled globally) - if getattr(request.app.state.config, "ENABLE_CHANNELS", False): + # Channels tools - search channels and messages (if builtin category enabled AND channels enabled globally) + if is_builtin_tool_enabled("channels") and getattr(request.app.state.config, "ENABLE_CHANNELS", False): builtin_functions.extend( [ search_channels, diff --git a/backend/requirements-min.txt b/backend/requirements-min.txt index 115d1bc92f..c4daedc446 100644 --- a/backend/requirements-min.txt +++ b/backend/requirements-min.txt @@ -4,7 +4,7 @@ fastapi==0.128.0 uvicorn[standard]==0.40.0 pydantic==2.12.5 -python-multipart==0.0.21 +python-multipart==0.0.22 itsdangerous==2.2.0 python-socketio==5.16.0 @@ -16,20 +16,20 @@ PyJWT[crypto]==2.10.1 authlib==1.6.6 requests==2.32.5 -aiohttp==3.13.2 +aiohttp==3.13.2 # do not update to 3.13.3 - broken async-timeout aiocache aiofiles -starlette-compress==1.6.1 +starlette-compress==1.7.0 httpx[socks,http2,zstd,cli,brotli]==0.28.1 starsessions[redis]==2.2.1 -sqlalchemy==2.0.45 -alembic==1.17.2 -peewee==3.18.3 +sqlalchemy==2.0.46 +alembic==1.18.3 +peewee==3.19.0 peewee-migrate==1.14.3 -pycrdt==0.12.44 +pycrdt==0.12.45 redis APScheduler==3.11.2 @@ -38,17 +38,17 @@ RestrictedPython==8.1 loguru==0.7.3 asgiref==3.11.0 -mcp==1.25.0 +mcp==1.26.0 openai -langchain==1.2.0 +langchain==1.2.7 langchain-community==0.4.1 langchain-classic==1.0.1 langchain-text-splitters==1.1.0 fake-useragent==2.2.0 -chromadb==1.4.0 -black==25.12.0 +chromadb==1.4.1 +black==26.1.0 pydub chardet==5.2.0 diff --git a/backend/requirements.txt b/backend/requirements.txt index 51f0a8a1ae..c9fdd5d0d7 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -1,7 +1,7 @@ fastapi==0.128.0 uvicorn[standard]==0.40.0 pydantic==2.12.5 -python-multipart==0.0.21 +python-multipart==0.0.22 itsdangerous==2.2.0 python-socketio==5.16.0 @@ -13,21 +13,21 @@ PyJWT[crypto]==2.10.1 authlib==1.6.6 requests==2.32.5 -aiohttp==3.13.2 +aiohttp==3.13.2 # do not update to 3.13.3 - broken async-timeout aiocache aiofiles -starlette-compress==1.6.1 +starlette-compress==1.7.0 httpx[socks,http2,zstd,cli,brotli]==0.28.1 starsessions[redis]==2.2.1 python-mimeparse==2.0.0 -sqlalchemy==2.0.45 -alembic==1.17.2 -peewee==3.18.3 +sqlalchemy==2.0.46 +alembic==1.18.3 +peewee==3.19.0 peewee-migrate==1.14.3 -pycrdt==0.12.44 +pycrdt==0.12.45 redis APScheduler==3.11.2 @@ -38,39 +38,39 @@ asgiref==3.11.0 # AI libraries tiktoken -mcp==1.25.0 +mcp==1.26.0 openai anthropic -google-genai==1.56.0 +google-genai==1.60.0 -langchain==1.2.0 +langchain==1.2.7 langchain-community==0.4.1 langchain-classic==1.0.1 langchain-text-splitters==1.1.0 fake-useragent==2.2.0 -chromadb==1.4.0 +chromadb==1.4.1 weaviate-client==4.19.2 opensearch-py==3.1.0 -transformers==4.57.3 -sentence-transformers==5.2.0 +transformers==4.57.6 +sentence-transformers==5.2.2 accelerate pyarrow==20.0.0 # fix: pin pyarrow version to 20 for rpi compatibility #15897 -einops==0.8.1 +einops==0.8.2 ftfy==6.3.1 chardet==5.2.0 -pypdf==6.5.0 +pypdf==6.6.2 fpdf2==2.8.5 -pymdown-extensions==10.20 +pymdown-extensions==10.20.1 docx2txt==0.9 python-pptx==1.0.2 -unstructured==0.18.24 +unstructured==0.18.31 msoffcrypto-tool==5.4.2 nltk==3.9.2 -Markdown==3.10 +Markdown==3.10.1 pypandoc==1.16.2 pandas==2.3.3 openpyxl==3.1.5 @@ -82,15 +82,15 @@ sentencepiece soundfile==0.13.1 pillow==12.1.0 -opencv-python-headless==4.12.0.88 +opencv-python-headless==4.13.0.90 rapidocr-onnxruntime==1.4.4 rank-bm25==0.2.2 onnxruntime==1.23.2 faster-whisper==1.2.1 -black==25.12.0 -youtube-transcript-api==1.2.3 +black==26.1.0 +youtube-transcript-api==1.2.4 pytube==15.0.0 pydub @@ -98,7 +98,7 @@ ddgs==9.10.0 azure-ai-documentintelligence==1.0.2 azure-identity==1.25.1 -azure-storage-blob==12.27.1 +azure-storage-blob==12.28.0 azure-search-documents==11.6.0 ## Google Drive @@ -107,7 +107,7 @@ google-auth-httplib2 google-auth-oauthlib googleapis-common-protos==1.72.0 -google-cloud-storage==3.7.0 +google-cloud-storage==3.8.0 ## Databases pymongo @@ -115,14 +115,14 @@ psycopg2-binary==2.9.11 pgvector==0.4.2 PyMySQL==1.1.2 -boto3==1.42.21 +boto3==1.42.38 -pymilvus==2.6.6 +pymilvus==2.6.8 qdrant-client==1.16.2 -playwright==1.57.0 # Caution: version must match docker-compose.playwright.yaml - Update the docker-compose.yaml if necessary +playwright==1.58.0 # Caution: version must match docker-compose.playwright.yaml - Update the docker-compose.yaml if necessary elasticsearch==9.2.1 pinecone==6.0.2 -oracledb==3.4.1 +oracledb==3.4.2 av==14.0.1 # Caution: Set due to FATAL FIPS SELFTEST FAILURE, see discussion https://github.com/open-webui/open-webui/discussions/15720 @@ -138,7 +138,7 @@ pytest-docker~=3.2.5 ldap3==2.9.1 ## Firecrawl -firecrawl-py==4.12.0 +firecrawl-py==4.14.0 ## Trace opentelemetry-api==1.39.1 diff --git a/docker-compose.playwright.yaml b/docker-compose.playwright.yaml index e00a28df58..167c2501d6 100644 --- a/docker-compose.playwright.yaml +++ b/docker-compose.playwright.yaml @@ -1,8 +1,8 @@ services: playwright: - image: mcr.microsoft.com/playwright:v1.57.0-noble # Version must match requirements.txt + image: mcr.microsoft.com/playwright:v1.58.0-noble # Version must match requirements.txt container_name: playwright - command: npx -y playwright@1.57.0 run-server --port 3000 --host 0.0.0.0 + command: npx -y playwright@1.58.0 run-server --port 3000 --host 0.0.0.0 open-webui: environment: diff --git a/pyproject.toml b/pyproject.toml index 893e1d018a..1f46eaa068 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,7 +9,7 @@ dependencies = [ "fastapi==0.128.0", "uvicorn[standard]==0.40.0", "pydantic==2.12.5", - "python-multipart==0.0.21", + "python-multipart==0.0.22", "itsdangerous==2.2.0", "python-socketio==5.16.0", @@ -21,21 +21,21 @@ dependencies = [ "authlib==1.6.6", "requests==2.32.5", - "aiohttp==3.13.2", + "aiohttp==3.13.2" # do not update to 3.13.3 - broken, "async-timeout", "aiocache", "aiofiles", - "starlette-compress==1.6.1", + "starlette-compress==1.7.0", "httpx[socks,http2,zstd,cli,brotli]==0.28.1", "starsessions[redis]==2.2.1", "python-mimeparse==2.0.0", - "sqlalchemy==2.0.45", - "alembic==1.17.2", - "peewee==3.18.3", + "sqlalchemy==2.0.46", + "alembic==1.18.3", + "peewee==3.19.0", "peewee-migrate==1.14.3", - "pycrdt==0.12.44", + "pycrdt==0.12.45", "redis", "APScheduler==3.11.2", @@ -45,40 +45,40 @@ dependencies = [ "asgiref==3.11.0", "tiktoken", - "mcp==1.25.0", + "mcp==1.26.0", "openai", "anthropic", - "google-genai==1.56.0", + "google-genai==1.60.0", - "langchain==1.2.0", + "langchain==1.2.7", "langchain-community==0.4.1", "langchain-classic==1.0.1", "langchain-text-splitters==1.1.0", "fake-useragent==2.2.0", - "chromadb==1.4.0", + "chromadb==1.4.1", "opensearch-py==3.1.0", "PyMySQL==1.1.2", - "boto3==1.42.21", + "boto3==1.42.38", - "transformers==4.57.3", - "sentence-transformers==5.2.0", + "transformers==4.57.6", + "sentence-transformers==5.2.2", "accelerate", "pyarrow==20.0.0", # fix: pin pyarrow version to 20 for rpi compatibility #15897 - "einops==0.8.1", + "einops==0.8.2", "ftfy==6.3.1", "chardet==5.2.0", - "pypdf==6.5.0", + "pypdf==6.6.2", "fpdf2==2.8.5", - "pymdown-extensions==10.20", + "pymdown-extensions==10.20.1", "docx2txt==0.9", "python-pptx==1.0.2", - "unstructured==0.18.24", + "unstructured==0.18.31", "msoffcrypto-tool==5.4.2", "nltk==3.9.2", - "Markdown==3.10", + "Markdown==3.10.1", "pypandoc==1.16.2", "pandas==2.3.3", "openpyxl==3.1.5", @@ -91,15 +91,15 @@ dependencies = [ "azure-ai-documentintelligence==1.0.2", "pillow==12.1.0", - "opencv-python-headless==4.12.0.88", + "opencv-python-headless==4.13.0.90", "rapidocr-onnxruntime==1.4.4", "rank-bm25==0.2.2", "onnxruntime==1.23.2", "faster-whisper==1.2.1", - "black==25.12.0", - "youtube-transcript-api==1.2.3", + "black==26.1.0", + "youtube-transcript-api==1.2.4", "pytube==15.0.0", "pydub", @@ -110,10 +110,10 @@ dependencies = [ "google-auth-oauthlib", "googleapis-common-protos==1.72.0", - "google-cloud-storage==3.7.0", + "google-cloud-storage==3.8.0", "azure-identity==1.25.1", - "azure-storage-blob==12.27.1", + "azure-storage-blob==12.28.0", "ldap3==2.9.1", ] @@ -145,18 +145,18 @@ all = [ "docker~=7.1.0", "pytest~=8.3.2", "pytest-docker~=3.2.5", - "playwright==1.57.0", # Caution: version must match docker-compose.playwright.yaml - Update the docker-compose.yaml if necessary + "playwright==1.58.0", # Caution: version must match docker-compose.playwright.yaml - Update the docker-compose.yaml if necessary "elasticsearch==9.2.1", "qdrant-client==1.16.2", "weaviate-client==4.19.2", - "pymilvus==2.6.6", + "pymilvus==2.6.8", "pinecone==6.0.2", - "oracledb==3.4.1", + "oracledb==3.4.2", "colbert-ai==0.2.22", - "firecrawl-py==4.12.0", + "firecrawl-py==4.14.0", "azure-search-documents==11.6.0", ] diff --git a/src/lib/apis/chats/index.ts b/src/lib/apis/chats/index.ts index b33072e890..bfe7a16386 100644 --- a/src/lib/apis/chats/index.ts +++ b/src/lib/apis/chats/index.ts @@ -255,6 +255,55 @@ export const getArchivedChatList = async ( })); }; +export const getSharedChatList = async ( + token: string = '', + page: number = 1, + filter?: object +) => { + let error = null; + + const searchParams = new URLSearchParams(); + searchParams.append('page', `${page}`); + + if (filter) { + Object.entries(filter).forEach(([key, value]) => { + if (value !== undefined && value !== null) { + searchParams.append(key, value.toString()); + } + }); + } + + const res = await fetch(`${WEBUI_API_BASE_URL}/chats/shared?${searchParams.toString()}`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + ...(token && { authorization: `Bearer ${token}` }) + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .then((json) => { + return json; + }) + .catch((err) => { + error = err; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res.map((chat) => ({ + ...chat, + time_range: getTimeRange(chat.updated_at) + })); +}; + export const getAllChats = async (token: string) => { let error = null; diff --git a/src/lib/apis/files/index.ts b/src/lib/apis/files/index.ts index 44af669fa1..15785b354d 100644 --- a/src/lib/apis/files/index.ts +++ b/src/lib/apis/files/index.ts @@ -175,6 +175,44 @@ export const getFiles = async (token: string = '') => { return res; }; +export const searchFiles = async ( + token: string, + filename: string = '*', + skip: number = 0, + limit: number = 50 +) => { + let error = null; + + const searchParams = new URLSearchParams(); + searchParams.append('filename', filename); + searchParams.append('skip', String(skip)); + searchParams.append('limit', String(limit)); + + const res = await fetch(`${WEBUI_API_BASE_URL}/files/search?${searchParams.toString()}`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return []; + }); + + if (error) { + throw error; + } + + return res; +}; + export const getFileById = async (token: string, id: string) => { let error = null; diff --git a/src/lib/apis/images/index.ts b/src/lib/apis/images/index.ts index a58d16085f..86b8a90ed1 100644 --- a/src/lib/apis/images/index.ts +++ b/src/lib/apis/images/index.ts @@ -217,7 +217,63 @@ export const imageGenerations = async (token: string = '', prompt: string) => { .catch((err) => { console.error(err); if ('detail' in err) { - error = err.detail; + if (Array.isArray(err.detail)) { + error = err.detail.map((e: { msg?: string }) => e.msg || JSON.stringify(e)).join(', '); + } else { + error = err.detail; + } + } else { + error = 'Server connection failed'; + } + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const imageEdits = async ( + token: string = '', + images: string | string[], + prompt: string, + model?: string, + size?: string, + n?: number +) => { + let error = null; + + const res = await fetch(`${IMAGES_API_BASE_URL}/edit`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + ...(token && { authorization: `Bearer ${token}` }) + }, + body: JSON.stringify({ + form_data: { + image: images, + prompt, + ...(model && { model }), + ...(size && { size }), + ...(n && { n }) + } + }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + console.error(err); + if ('detail' in err) { + if (Array.isArray(err.detail)) { + error = err.detail.map((e: { msg?: string }) => e.msg || JSON.stringify(e)).join(', '); + } else { + error = err.detail; + } } else { error = 'Server connection failed'; } diff --git a/src/lib/apis/prompts/index.ts b/src/lib/apis/prompts/index.ts index 4129ea62aa..e9cd6e8481 100644 --- a/src/lib/apis/prompts/index.ts +++ b/src/lib/apis/prompts/index.ts @@ -1,10 +1,48 @@ import { WEBUI_API_BASE_URL } from '$lib/constants'; type PromptItem = { + id?: string; // Prompt ID command: string; - title: string; + name: string; // Changed from title content: string; + data?: object | null; + meta?: object | null; access_control?: null | object; + version_id?: string | null; // Active version + commit_message?: string | null; // For history tracking + is_production?: boolean; // Whether to set new version as production +}; + +type PromptHistoryItem = { + id: string; + prompt_id: string; + parent_id: string | null; + snapshot: { + name: string; + content: string; + command: string; + data: object; + meta: object; + access_control: object | null; + }; + user_id: string; + commit_message: string | null; + created_at: number; + user?: { + id: string; + name: string; + email: string; + }; +}; + +type PromptDiff = { + from_id: string; + to_id: string; + from_snapshot: object; + to_snapshot: object; + content_diff: string[]; + name_changed: boolean; + access_control_changed: boolean; }; export const createNewPrompt = async (token: string, prompt: PromptItem) => { @@ -19,7 +57,7 @@ export const createNewPrompt = async (token: string, prompt: PromptItem) => { }, body: JSON.stringify({ ...prompt, - command: `/${prompt.command}` + command: prompt.command.startsWith('/') ? prompt.command.slice(1) : prompt.command }) }) .then(async (res) => { @@ -70,7 +108,95 @@ export const getPrompts = async (token: string = '') => { return res; }; +export const getPromptTags = async (token: string = '') => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/tags`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getPromptItems = async ( + token: string = '', + query: string | null, + viewOption: string | null, + selectedTag: string | null, + orderBy: string | null, + direction: string | null, + page: number +) => { + let error = null; + + const searchParams = new URLSearchParams(); + if (query) { + searchParams.append('query', query); + } + if (viewOption) { + searchParams.append('view_option', viewOption); + } + if (selectedTag) { + searchParams.append('tag', selectedTag); + } + if (orderBy) { + searchParams.append('order_by', orderBy); + } + if (direction) { + searchParams.append('direction', direction); + } + if (page) { + searchParams.append('page', page.toString()); + } + + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/list?${searchParams.toString()}`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .then((json) => { + return json; + }) + .catch((err) => { + error = err; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const getPromptList = async (token: string = '') => { + let error = null; const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/list`, { @@ -104,6 +230,8 @@ export const getPromptList = async (token: string = '') => { export const getPromptByCommand = async (token: string, command: string) => { let error = null; + command = command.charAt(0) === '/' ? command.slice(1) : command; + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/command/${command}`, { method: 'GET', headers: { @@ -133,20 +261,16 @@ export const getPromptByCommand = async (token: string, command: string) => { return res; }; -export const updatePromptByCommand = async (token: string, prompt: PromptItem) => { +export const getPromptById = async (token: string, promptId: string) => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/command/${prompt.command}/update`, { - method: 'POST', + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/id/${promptId}`, { + method: 'GET', headers: { Accept: 'application/json', 'Content-Type': 'application/json', authorization: `Bearer ${token}` - }, - body: JSON.stringify({ - ...prompt, - command: `/${prompt.command}` - }) + } }) .then(async (res) => { if (!res.ok) throw await res.json(); @@ -169,12 +293,113 @@ export const updatePromptByCommand = async (token: string, prompt: PromptItem) = return res; }; -export const deletePromptByCommand = async (token: string, command: string) => { +export const updatePromptById = async (token: string, prompt: PromptItem) => { let error = null; - command = command.charAt(0) === '/' ? command.slice(1) : command; + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/id/${prompt.id}/update`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify(prompt) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .then((json) => { + return json; + }) + .catch((err) => { + error = err.detail; - const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/command/${command}/delete`, { + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const updatePromptMetadata = async ( + token: string, + promptId: string, + name: string, + command: string, + tags: string[] = [] +) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/id/${promptId}/update/meta`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify({ name, command, tags }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const setProductionPromptVersion = async ( + token: string, + promptId: string, + version_id: string +) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/id/${promptId}/update/version`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ + version_id: version_id + }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + console.log(err); + error = err.detail; + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const deletePromptById = async (token: string, promptId: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/id/${promptId}/delete`, { method: 'DELETE', headers: { Accept: 'application/json', @@ -202,3 +427,188 @@ export const deletePromptByCommand = async (token: string, command: string) => { return res; }; + +//////////////////////////// +// Prompt History APIs +//////////////////////////// + +export const getPromptHistory = async ( + token: string, + promptId: string, + page: number = 0 +): Promise => { + let error = null; + + const res = await fetch( + `${WEBUI_API_BASE_URL}/prompts/id/${promptId}/history?page=${page}`, + { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const deletePromptHistoryVersion = async ( + token: string, + promptId: string, + historyId: string +): Promise => { + let error = null; + + const res = await fetch( + `${WEBUI_API_BASE_URL}/prompts/id/${promptId}/history/${historyId}`, + { + method: 'DELETE', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return false; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getPromptHistoryEntry = async ( + token: string, + promptId: string, + historyId: string +): Promise => { + let error = null; + + const res = await fetch( + `${WEBUI_API_BASE_URL}/prompts/id/${promptId}/history/${historyId}`, + { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const restorePromptFromHistory = async ( + token: string, + promptId: string, + historyId: string, + commitMessage?: string +) => { + let error = null; + + const res = await fetch( + `${WEBUI_API_BASE_URL}/prompts/id/${promptId}/history/${historyId}/restore`, + { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + }, + body: JSON.stringify({ + commit_message: commitMessage + }) + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getPromptDiff = async ( + token: string, + promptId: string, + fromId: string, + toId: string +): Promise => { + let error = null; + + const res = await fetch( + `${WEBUI_API_BASE_URL}/prompts/id/${promptId}/history/diff?from_id=${fromId}&to_id=${toId}`, + { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + diff --git a/src/lib/components/admin/Evaluations/Feedbacks.svelte b/src/lib/components/admin/Evaluations/Feedbacks.svelte index 0ec5678f0e..ccf1735431 100644 --- a/src/lib/components/admin/Evaluations/Feedbacks.svelte +++ b/src/lib/components/admin/Evaluations/Feedbacks.svelte @@ -302,9 +302,11 @@
{#if feedback.data?.sibling_model_ids} -
+ +
{feedback.data?.model_id}
+
@@ -320,11 +322,11 @@
{:else} -
+ +
{feedback.data?.model_id}
+
{/if}
diff --git a/src/lib/components/admin/Evaluations/Leaderboard.svelte b/src/lib/components/admin/Evaluations/Leaderboard.svelte index e16e62c98c..0dd501a153 100644 --- a/src/lib/components/admin/Evaluations/Leaderboard.svelte +++ b/src/lib/components/admin/Evaluations/Leaderboard.svelte @@ -182,7 +182,9 @@ alt={model.name} class="size-5 rounded-full object-cover" /> - {model.name} + + {model.name} +
diff --git a/src/lib/components/admin/Evaluations/LeaderboardModal.svelte b/src/lib/components/admin/Evaluations/LeaderboardModal.svelte index fc3ec6eb10..6730e739d7 100644 --- a/src/lib/components/admin/Evaluations/LeaderboardModal.svelte +++ b/src/lib/components/admin/Evaluations/LeaderboardModal.svelte @@ -4,6 +4,7 @@ import { getModelHistory } from '$lib/apis/evaluations'; import ModelActivityChart from './ModelActivityChart.svelte'; import XMark from '$lib/components/icons/XMark.svelte'; + import Tooltip from '$lib/components/common/Tooltip.svelte'; export let show = false; export let model = null; @@ -60,9 +61,11 @@ {#if model}
-
- {model.name} -
+ +
+ {model.name} +
+
diff --git a/src/lib/components/admin/Functions.svelte b/src/lib/components/admin/Functions.svelte index 67a1fbbdfd..48c1863e74 100644 --- a/src/lib/components/admin/Functions.svelte +++ b/src/lib/components/admin/Functions.svelte @@ -4,7 +4,7 @@ const { saveAs } = fileSaver; import { WEBUI_NAME, config, functions as _functions, models, settings, user } from '$lib/stores'; - import { onMount, getContext, tick } from 'svelte'; + import { onMount, getContext, tick, onDestroy } from 'svelte'; import { goto } from '$app/navigation'; import { @@ -53,6 +53,7 @@ let viewOption = ''; let query = ''; + let searchDebounceTimer: ReturnType; let selectedTag = ''; let selectedType = ''; @@ -70,12 +71,14 @@ let functions = null; let filteredItems = []; - $: if ( - functions && - query !== undefined && - selectedType !== undefined && - viewOption !== undefined - ) { + $: if (query !== undefined) { + clearTimeout(searchDebounceTimer); + searchDebounceTimer = setTimeout(() => { + setFilteredItems(); + }, 300); + } + + $: if (functions && selectedType !== undefined && viewOption !== undefined) { setFilteredItems(); } @@ -86,7 +89,10 @@ (selectedType !== '' ? f.type === selectedType : true) && (query === '' || f.name.toLowerCase().includes(query.toLowerCase()) || - f.id.toLowerCase().includes(query.toLowerCase())) && + f.id.toLowerCase().includes(query.toLowerCase()) || + (f.user?.name || '').toLowerCase().includes(query.toLowerCase()) || + (f.user?.email || '').toLowerCase().includes(query.toLowerCase()) || + (f.user?.username || '').toLowerCase().includes(query.toLowerCase())) && (viewOption === '' || (viewOption === 'created' && f.user_id === $user?.id) || (viewOption === 'shared' && f.user_id !== $user?.id)) @@ -236,6 +242,10 @@ window.removeEventListener('blur-sm', onBlur); }; }); + + onDestroy(() => { + clearTimeout(searchDebounceTimer); + }); diff --git a/src/lib/components/admin/Settings/Documents.svelte b/src/lib/components/admin/Settings/Documents.svelte index fcc4e5d027..af18376bae 100644 --- a/src/lib/components/admin/Settings/Documents.svelte +++ b/src/lib/components/admin/Settings/Documents.svelte @@ -362,6 +362,30 @@
+ +
+
+
+ + {$i18n.t('PDF Loader Mode')} + +
+
+ +
+
+
{:else if RAGConfig.CONTENT_EXTRACTION_ENGINE === 'datalab_marker'}
searchValue === '' || m.name.toLowerCase().includes(searchValue.toLowerCase())) + .filter((m) => { + if (viewOption === 'enabled') return m?.is_active ?? true; + if (viewOption === 'disabled') return !(m?.is_active ?? true); + if (viewOption === 'visible') return !(m?.meta?.hidden ?? false); + if (viewOption === 'hidden') return m?.meta?.hidden === true; + return true; // All + }) .sort((a, b) => { - // // Check if either model is inactive and push them to the bottom - // if ((a.is_active ?? true) !== (b.is_active ?? true)) { - // return (b.is_active ?? true) - (a.is_active ?? true); - // } - // If both models' active states are the same, sort alphabetically return (a?.name ?? a?.id ?? '').localeCompare(b?.name ?? b?.id ?? ''); }); } let searchValue = ''; + const enableAllHandler = async () => { + const modelsToEnable = filteredModels.filter((m) => !(m.is_active ?? true)); + // Optimistic UI update + modelsToEnable.forEach((m) => (m.is_active = true)); + models = models; + // Sync with server + await Promise.all(modelsToEnable.map((model) => toggleModelById(localStorage.token, model.id))); + }; + + const disableAllHandler = async () => { + const modelsToDisable = filteredModels.filter((m) => m.is_active ?? true); + // Optimistic UI update + modelsToDisable.forEach((m) => (m.is_active = false)); + models = models; + // Sync with server + await Promise.all(modelsToDisable.map((model) => toggleModelById(localStorage.token, model.id))); + }; + const downloadModels = async (models) => { let blob = new Blob([JSON.stringify(models)], { type: 'application/json' @@ -275,41 +301,48 @@ {#if selectedModelId === null}
-
- {$i18n.t('Models')} - {filteredModels.length} +
+
+ {$i18n.t('Models')} +
+ +
+ {filteredModels.length} +
-
- - - +
+ - - - +
+
-
+
+
@@ -333,156 +366,215 @@ {/if}
-
-
- {#if models.length > 0} - {#each filteredModels as model, modelIdx (`${model.id}-${modelIdx}`)} -
+
+
+ +
+ +
+ + + -
- {#if shiftKey} - + + +
+ + { + enableAllHandler(); + }} + > + +
{$i18n.t('Enable All')}
+
+ + { + disableAllHandler(); + }} + > + +
{$i18n.t('Disable All')}
+
+
+
+ +
+ +
+ {#if filteredModels.length > 0} + {#each filteredModels as model, modelIdx (`${model.id}-${modelIdx}`)} +
+ +
+ {#if shiftKey} + + + + {:else} - - {:else} - - { - exportModelHandler(model); - }} - hideHandler={() => { - hideModelHandler(model); - }} - pinModelHandler={() => { - pinModelHandler(model.id); - }} - copyLinkHandler={() => { - copyLinkHandler(model); - }} - cloneHandler={() => { - cloneHandler(model); - }} - onClose={() => {}} - > - - + + -
- - { - toggleModelHandler(model); - }} - /> - -
- {/if} +
+ + { + toggleModelHandler(model); + }} + /> + +
+ {/if} +
+
+ {/each} + {:else} +
+
+
😕
+
{$i18n.t('No models found')}
+
+ {$i18n.t('Try adjusting your search or filter to find what you are looking for.')} +
- {/each} - {:else} -
-
- {$i18n.t('No models found')} -
-
- {/if} + {/if} +
{#if $user?.role === 'admin'} diff --git a/src/lib/components/admin/Settings/Models/AdminViewSelector.svelte b/src/lib/components/admin/Settings/Models/AdminViewSelector.svelte new file mode 100644 index 0000000000..5778ba21ab --- /dev/null +++ b/src/lib/components/admin/Settings/Models/AdminViewSelector.svelte @@ -0,0 +1,63 @@ + + + item.value === value)} + {items} + onSelectedChange={(selectedItem) => { + value = selectedItem.value; + onChange(value); + }} +> + + + + + + + {#each items as item} + + {item.label} + + {#if value === item.value} +
+ +
+ {/if} +
+ {/each} +
+
diff --git a/src/lib/components/admin/Settings/Models/ModelList.svelte b/src/lib/components/admin/Settings/Models/ModelList.svelte index cc86e52e5f..d501a485d4 100644 --- a/src/lib/components/admin/Settings/Models/ModelList.svelte +++ b/src/lib/components/admin/Settings/Models/ModelList.svelte @@ -50,7 +50,7 @@
-
+
{#if $models.find((model) => model.id === modelId)} {$models.find((model) => model.id === modelId).name} {:else} diff --git a/src/lib/components/admin/Settings/WebSearch.svelte b/src/lib/components/admin/Settings/WebSearch.svelte index e91a110f81..587b494f0d 100644 --- a/src/lib/components/admin/Settings/WebSearch.svelte +++ b/src/lib/components/admin/Settings/WebSearch.svelte @@ -7,6 +7,7 @@ import { toast } from 'svelte-sonner'; import SensitiveInput from '$lib/components/common/SensitiveInput.svelte'; import Tooltip from '$lib/components/common/Tooltip.svelte'; + import Textarea from '$lib/components/common/Textarea.svelte'; const i18n = getContext('i18n'); @@ -35,7 +36,8 @@ 'perplexity', 'sougou', 'firecrawl', - 'external' + 'external', + 'yandex' ]; let webLoaderEngines = ['playwright', 'firecrawl', 'tavily', 'external']; @@ -735,6 +737,53 @@ />
+ {:else if webConfig.WEB_SEARCH_ENGINE === 'yandex'} +
+
+
+ {$i18n.t('Yandex Web Search URL')} +
+ +
+
+ +
+
+
+ +
+
+ {$i18n.t('Yandex Web Search API Key')} +
+ + +
+ +
+
{$i18n.t('Yandex Web Search config')}
+ + +