diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 01e18ffbde..85a9aedad1 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -573,6 +573,20 @@ ENABLE_OAUTH_GROUP_CREATION = PersistentConfig( ) +oauth_group_default_share = ( + os.environ.get("OAUTH_GROUP_DEFAULT_SHARE", "true").strip().lower() +) +OAUTH_GROUP_DEFAULT_SHARE = PersistentConfig( + "OAUTH_GROUP_DEFAULT_SHARE", + "oauth.group_default_share", + ( + "members" + if oauth_group_default_share == "members" + else oauth_group_default_share == "true" + ), +) + + OAUTH_BLOCKED_GROUPS = PersistentConfig( "OAUTH_BLOCKED_GROUPS", "oauth.blocked_groups", @@ -1678,6 +1692,10 @@ ENABLE_ADMIN_CHAT_ACCESS = ( os.environ.get("ENABLE_ADMIN_CHAT_ACCESS", "True").lower() == "true" ) +ENABLE_ADMIN_ANALYTICS = ( + os.environ.get("ENABLE_ADMIN_ANALYTICS", "True").lower() == "true" +) + ENABLE_COMMUNITY_SHARING = PersistentConfig( "ENABLE_COMMUNITY_SHARING", "ui.enable_community_sharing", @@ -2915,6 +2933,12 @@ ENABLE_ASYNC_EMBEDDING = PersistentConfig( os.environ.get("ENABLE_ASYNC_EMBEDDING", "True").lower() == "true", ) +RAG_EMBEDDING_CONCURRENT_REQUESTS = PersistentConfig( + "RAG_EMBEDDING_CONCURRENT_REQUESTS", + "rag.embedding_concurrent_requests", + int(os.getenv("RAG_EMBEDDING_CONCURRENT_REQUESTS", "0")), +) + RAG_EMBEDDING_QUERY_PREFIX = os.environ.get("RAG_EMBEDDING_QUERY_PREFIX", None) RAG_EMBEDDING_CONTENT_PREFIX = os.environ.get("RAG_EMBEDDING_CONTENT_PREFIX", None) @@ -3497,6 +3521,12 @@ YANDEX_WEB_SEARCH_CONFIG = PersistentConfig( os.environ.get("YANDEX_WEB_SEARCH_CONFIG", ""), ) +YOUCOM_API_KEY = PersistentConfig( + "YOUCOM_API_KEY", + "rag.web.search.youcom_api_key", + os.environ.get("YOUCOM_API_KEY", ""), +) + #################################### # Images #################################### diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index 86fe787899..65f95b2755 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -557,6 +557,10 @@ OAUTH_SESSION_TOKEN_ENCRYPTION_KEY = os.environ.get( "OAUTH_SESSION_TOKEN_ENCRYPTION_KEY", WEBUI_SECRET_KEY ) +# Maximum number of concurrent OAuth sessions per user per provider +# This prevents unbounded session growth while allowing multi-device usage +OAUTH_MAX_SESSIONS_PER_USER = int(os.environ.get("OAUTH_MAX_SESSIONS_PER_USER", "10")) + # Token Exchange Configuration # Allows external apps to exchange OAuth tokens for OpenWebUI tokens ENABLE_OAUTH_TOKEN_EXCHANGE = ( @@ -978,6 +982,11 @@ OTEL_LOGS_OTLP_SPAN_EXPORTER = os.environ.get( # TOOLS/FUNCTIONS PIP OPTIONS #################################### +ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS = ( + os.environ.get("ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS", "True").lower() + == "true" +) + PIP_OPTIONS = os.getenv("PIP_OPTIONS", "").split() PIP_PACKAGE_INDEX_OPTIONS = os.getenv("PIP_PACKAGE_INDEX_OPTIONS", "").split() diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index c1e4234647..8c99801148 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -240,6 +240,7 @@ from open_webui.config import ( RAG_EMBEDDING_ENGINE, RAG_EMBEDDING_BATCH_SIZE, ENABLE_ASYNC_EMBEDDING, + RAG_EMBEDDING_CONCURRENT_REQUESTS, RAG_TOP_K, RAG_TOP_K_RERANKER, RAG_RELEVANCE_THRESHOLD, @@ -361,6 +362,7 @@ from open_webui.config import ( YANDEX_WEB_SEARCH_URL, YANDEX_WEB_SEARCH_API_KEY, YANDEX_WEB_SEARCH_CONFIG, + YOUCOM_API_KEY, # WebUI WEBUI_AUTH, WEBUI_NAME, @@ -434,6 +436,7 @@ from open_webui.config import ( RESPONSE_WATERMARK, # Admin ENABLE_ADMIN_CHAT_ACCESS, + ENABLE_ADMIN_ANALYTICS, BYPASS_ADMIN_ACCESS_CONTROL, ENABLE_ADMIN_EXPORT, # Tasks @@ -984,6 +987,7 @@ app.state.config.RAG_EMBEDDING_ENGINE = RAG_EMBEDDING_ENGINE app.state.config.RAG_EMBEDDING_MODEL = RAG_EMBEDDING_MODEL app.state.config.RAG_EMBEDDING_BATCH_SIZE = RAG_EMBEDDING_BATCH_SIZE app.state.config.ENABLE_ASYNC_EMBEDDING = ENABLE_ASYNC_EMBEDDING +app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS = RAG_EMBEDDING_CONCURRENT_REQUESTS app.state.config.RAG_RERANKING_ENGINE = RAG_RERANKING_ENGINE app.state.config.RAG_RERANKING_MODEL = RAG_RERANKING_MODEL @@ -1069,6 +1073,7 @@ 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.YOUCOM_API_KEY = YOUCOM_API_KEY app.state.config.PLAYWRIGHT_WS_URL = PLAYWRIGHT_WS_URL @@ -1137,6 +1142,7 @@ app.state.EMBEDDING_FUNCTION = get_embedding_function( else None ), enable_async=app.state.config.ENABLE_ASYNC_EMBEDDING, + concurrent_requests=app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS, ) app.state.RERANKING_FUNCTION = get_reranking_function( @@ -1446,6 +1452,16 @@ async def check_url(request: Request, call_next): scheme="Bearer", credentials=request.cookies.get("token") ) + # Fallback to x-api-key header for Anthropic Messages API routes + if request.state.token is None and request.headers.get("x-api-key"): + request_path = request.url.path + if request_path in ("/api/message", "/api/v1/messages"): + from fastapi.security import HTTPAuthorizationCredentials + + request.state.token = HTTPAuthorizationCredentials( + scheme="Bearer", credentials=request.headers.get("x-api-key") + ) + request.state.enable_api_keys = app.state.config.ENABLE_API_KEYS response = await call_next(request) process_time = int(time.time()) - start_time @@ -1519,7 +1535,8 @@ app.include_router(functions.router, prefix="/api/v1/functions", tags=["function app.include_router( evaluations.router, prefix="/api/v1/evaluations", tags=["evaluations"] ) -app.include_router(analytics.router, prefix="/api/v1/analytics", tags=["analytics"]) +if ENABLE_ADMIN_ANALYTICS: + app.include_router(analytics.router, prefix="/api/v1/analytics", tags=["analytics"]) app.include_router(utils.router, prefix="/api/v1/utils", tags=["utils"]) # SCIM 2.0 API for identity management @@ -1749,9 +1766,12 @@ async def chat_completion( "local:" ): # temporary chats are not stored - # Verify chat ownership - chat = Chats.get_chat_by_id_and_user_id(metadata["chat_id"], user.id) - if chat is None and user.role != "admin": # admins can access any chat + # Verify chat ownership — lightweight EXISTS check avoids + # deserializing the full chat JSON blob just to confirm the row exists + if ( + not Chats.is_chat_owner(metadata["chat_id"], user.id) + and user.role != "admin" + ): # admins can access any chat raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.DEFAULT(), @@ -1897,6 +1917,68 @@ generate_chat_completions = chat_completion generate_chat_completion = chat_completion +################################## +# +# Anthropic Messages API Compatible Endpoint +# +################################## + + +from open_webui.utils.anthropic import ( + convert_anthropic_to_openai_payload, + convert_openai_to_anthropic_response, + openai_stream_to_anthropic_stream, +) + + +@app.post("/api/message") +@app.post("/api/v1/messages") # Anthropic Messages API compatible endpoint +async def generate_messages( + request: Request, + form_data: dict, + user=Depends(get_verified_user), +): + """ + Anthropic Messages API compatible endpoint. + + Accepts the Anthropic Messages API format, converts internally to OpenAI + Chat Completions format, routes through the existing chat completion + pipeline, then converts the response back to Anthropic Messages format. + + Supports both streaming and non-streaming requests. + All models configured in Open WebUI are accessible via this endpoint. + + Authentication: Supports both standard Authorization header and + Anthropic's x-api-key header (via middleware translation). + """ + # Convert Anthropic payload to OpenAI format + requested_model = form_data.get("model", "") + + openai_payload = convert_anthropic_to_openai_payload(form_data) + + # Route through the existing chat_completion handler + response = await chat_completion(request, openai_payload, user) + + # Convert response back to Anthropic format + if isinstance(response, StreamingResponse): + # Streaming response: wrap the generator to convert SSE format + return StreamingResponse( + openai_stream_to_anthropic_stream( + response.body_iterator, model=requested_model + ), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + }, + ) + elif isinstance(response, dict): + return convert_openai_to_anthropic_response(response, model=requested_model) + else: + # Passthrough for error responses (JSONResponse, PlainTextResponse, etc.) + return response + + @app.post("/api/chat/completed") async def chat_completed( request: Request, form_data: dict, user=Depends(get_verified_user) @@ -2046,6 +2128,7 @@ async def get_app_config(request: Request): "enable_user_status": app.state.config.ENABLE_USER_STATUS, "enable_admin_export": ENABLE_ADMIN_EXPORT, "enable_admin_chat_access": ENABLE_ADMIN_CHAT_ACCESS, + "enable_admin_analytics": ENABLE_ADMIN_ANALYTICS, "enable_google_drive_integration": app.state.config.ENABLE_GOOGLE_DRIVE_INTEGRATION, "enable_onedrive_integration": app.state.config.ENABLE_ONEDRIVE_INTEGRATION, "enable_memories": app.state.config.ENABLE_MEMORIES, diff --git a/backend/open_webui/models/access_grants.py b/backend/open_webui/models/access_grants.py index dd5a344b46..227621becd 100644 --- a/backend/open_webui/models/access_grants.py +++ b/backend/open_webui/models/access_grants.py @@ -402,7 +402,7 @@ class AccessGrantsTable: results = [] for grant_dict in normalized_grants: grant = AccessGrant( - id=grant_dict["id"], + id=str(uuid.uuid4()), resource_type=resource_type, resource_id=resource_id, principal_type=grant_dict["principal_type"], @@ -456,6 +456,31 @@ class AccessGrantsTable: ) return [AccessGrantModel.model_validate(g) for g in grants] + def get_grants_by_resources( + self, + resource_type: str, + resource_ids: list[str], + db: Optional[Session] = None, + ) -> dict[str, list[AccessGrantModel]]: + """Batch-fetch grants for multiple resources. Returns {resource_id: [grants]}.""" + if not resource_ids: + return {} + with get_db_context(db) as db: + grants = ( + db.query(AccessGrant) + .filter( + AccessGrant.resource_type == resource_type, + AccessGrant.resource_id.in_(resource_ids), + ) + .all() + ) + result: dict[str, list[AccessGrantModel]] = { + rid: [] for rid in resource_ids + } + for g in grants: + result[g.resource_id].append(AccessGrantModel.model_validate(g)) + return result + def has_access( self, user_id: str, diff --git a/backend/open_webui/models/channels.py b/backend/open_webui/models/channels.py index 8a55da9345..e212789a44 100644 --- a/backend/open_webui/models/channels.py +++ b/backend/open_webui/models/channels.py @@ -261,13 +261,19 @@ class ChannelTable: return AccessGrants.get_grants_by_resource("channel", channel_id, db=db) def _to_channel_model( - self, channel: Channel, db: Optional[Session] = None + self, + channel: Channel, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, ) -> ChannelModel: channel_data = ChannelModel.model_validate(channel).model_dump( exclude={"access_grants"} ) - access_grants = self._get_access_grants(channel_data["id"], db=db) - channel_data["access_grants"] = access_grants + channel_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(channel_data["id"], db=db) + ) return ChannelModel.model_validate(channel_data) def _collect_unique_user_ids( @@ -368,7 +374,18 @@ class ChannelTable: def get_channels(self, db: Optional[Session] = None) -> list[ChannelModel]: with get_db_context(db) as db: channels = db.query(Channel).all() - return [self._to_channel_model(channel, db=db) for channel in channels] + channel_ids = [channel.id for channel in channels] + grants_map = AccessGrants.get_grants_by_resources( + "channel", channel_ids, db=db + ) + return [ + self._to_channel_model( + channel, + access_grants=grants_map.get(channel.id, []), + db=db, + ) + for channel in channels + ] def _has_permission(self, db, query, filter: dict, permission: str = "read"): return AccessGrants.has_permission_filter( @@ -417,7 +434,14 @@ class ChannelTable: standard_channels = query.all() all_channels = membership_channels + standard_channels - return [self._to_channel_model(c, db=db) for c in all_channels] + channel_ids = [c.id for c in all_channels] + grants_map = AccessGrants.get_grants_by_resources( + "channel", channel_ids, db=db + ) + return [ + self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) + for c in all_channels + ] def get_dm_channel_by_user_ids( self, user_ids: list[str], db: Optional[Session] = None @@ -724,7 +748,17 @@ class ChannelTable: ) channel_ids = [cf.channel_id for cf in channel_files] channels = db.query(Channel).filter(Channel.id.in_(channel_ids)).all() - return [self._to_channel_model(channel, db=db) for channel in channels] + grants_map = AccessGrants.get_grants_by_resources( + "channel", channel_ids, db=db + ) + return [ + self._to_channel_model( + channel, + access_grants=grants_map.get(channel.id, []), + db=db, + ) + for channel in channels + ] def get_channels_by_file_id_and_user_id( self, file_id: str, user_id: str, db: Optional[Session] = None diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index 1418abd62d..5025cf7fca 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -456,11 +456,11 @@ class ChatTable: return ChatModel.model_validate(chat) def get_chat_title_by_id(self, id: str) -> Optional[str]: - chat = self.get_chat_by_id(id) - if chat is None: - return None - - return chat.chat.get("title", "New Chat") + with get_db_context() as db: + result = db.query(Chat.title).filter_by(id=id).first() + if result is None: + return None + return result[0] or "New Chat" def get_messages_map_by_chat_id(self, id: str) -> Optional[dict]: chat = self.get_chat_by_id(id) @@ -489,6 +489,7 @@ class ChatTable: if isinstance(message.get("content"), str): message["content"] = sanitize_text_for_db(message["content"]) + user_id = chat.user_id chat = chat.chat history = chat.get("history", {}) @@ -509,7 +510,7 @@ class ChatTable: ChatMessages.upsert_message( message_id=message_id, chat_id=id, - user_id=self.get_chat_by_id(id).user_id, + user_id=user_id, data=history["messages"][message_id], ) except Exception as e: @@ -713,7 +714,7 @@ class ChatTable: skip: int = 0, limit: int = 50, db: Optional[Session] = None, - ) -> list[ChatModel]: + ) -> list[ChatTitleIdResponse]: with get_db_context(db) as db: query = db.query(Chat).filter_by(user_id=user_id, archived=True) @@ -739,13 +740,27 @@ class ChatTable: else: query = query.order_by(Chat.updated_at.desc()) + query = query.with_entities( + Chat.id, Chat.title, Chat.updated_at, Chat.created_at + ) + 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] + return [ + ChatTitleIdResponse.model_validate( + { + "id": chat[0], + "title": chat[1], + "updated_at": chat[2], + "created_at": chat[3], + } + ) + for chat in all_chats + ] def get_shared_chat_list_by_user_id( self, @@ -754,7 +769,7 @@ class ChatTable: skip: int = 0, limit: int = 50, db: Optional[Session] = None, - ) -> list[ChatModel]: + ) -> list[SharedChatResponse]: with get_db_context(db) as db: query = ( @@ -784,13 +799,34 @@ class ChatTable: else: query = query.order_by(Chat.updated_at.desc()) + # Select only the columns needed for SharedChatResponse + # to avoid loading the heavy chat JSON blob + query = query.with_entities( + Chat.id, + Chat.title, + Chat.share_id, + Chat.updated_at, + Chat.created_at, + ) + 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] + return [ + SharedChatResponse.model_validate( + { + "id": chat[0], + "title": chat[1], + "share_id": chat[2], + "updated_at": chat[3], + "created_at": chat[4], + } + ) + for chat in all_chats + ] def get_chat_list_by_user_id( self, @@ -938,6 +974,37 @@ class ChatTable: except Exception: return None + def is_chat_owner( + self, id: str, user_id: str, db: Optional[Session] = None + ) -> bool: + """ + Lightweight ownership check — uses EXISTS subquery instead of loading + the full Chat row (which includes the potentially large JSON blob). + """ + try: + with get_db_context(db) as db: + return db.query( + exists().where(and_(Chat.id == id, Chat.user_id == user_id)) + ).scalar() + except Exception: + return False + + def get_chat_folder_id( + self, id: str, user_id: str, db: Optional[Session] = None + ) -> Optional[str]: + """ + Fetch only the folder_id column for a chat, without loading the full + JSON blob. Returns None if chat doesn't exist or doesn't belong to user. + """ + try: + with get_db_context(db) as db: + result = ( + db.query(Chat.folder_id).filter_by(id=id, user_id=user_id).first() + ) + return result[0] if result else None + except Exception: + return None + def get_chats( self, skip: int = 0, limit: int = 50, db: Optional[Session] = None ) -> list[ChatModel]: @@ -997,14 +1064,25 @@ class ChatTable: def get_pinned_chats_by_user_id( self, user_id: str, db: Optional[Session] = None - ) -> list[ChatModel]: + ) -> list[ChatTitleIdResponse]: with get_db_context(db) as db: all_chats = ( db.query(Chat) .filter_by(user_id=user_id, pinned=True, archived=False) .order_by(Chat.updated_at.desc()) + .with_entities(Chat.id, Chat.title, Chat.updated_at, Chat.created_at) ) - return [ChatModel.model_validate(chat) for chat in all_chats] + return [ + ChatTitleIdResponse.model_validate( + { + "id": chat[0], + "title": chat[1], + "updated_at": chat[2], + "created_at": chat[3], + } + ) + for chat in all_chats + ] def get_archived_chats_by_user_id( self, user_id: str, db: Optional[Session] = None diff --git a/backend/open_webui/models/knowledge.py b/backend/open_webui/models/knowledge.py index 1d21d5d910..4e5e208d8b 100644 --- a/backend/open_webui/models/knowledge.py +++ b/backend/open_webui/models/knowledge.py @@ -144,13 +144,18 @@ class KnowledgeTable: return AccessGrants.get_grants_by_resource("knowledge", knowledge_id, db=db) def _to_knowledge_model( - self, knowledge: Knowledge, db: Optional[Session] = None + self, + knowledge: Knowledge, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, ) -> KnowledgeModel: knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump( exclude={"access_grants"} ) - knowledge_data["access_grants"] = self._get_access_grants( - knowledge_data["id"], db=db + knowledge_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(knowledge_data["id"], db=db) ) return KnowledgeModel.model_validate(knowledge_data) @@ -192,9 +197,13 @@ class KnowledgeTable: db.query(Knowledge).order_by(Knowledge.updated_at.desc()).all() ) user_ids = list(set(knowledge.user_id for knowledge in all_knowledge)) + knowledge_ids = [knowledge.id for knowledge in all_knowledge] users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} + grants_map = AccessGrants.get_grants_by_resources( + "knowledge", knowledge_ids, db=db + ) knowledge_bases = [] for knowledge in all_knowledge: @@ -202,7 +211,11 @@ class KnowledgeTable: knowledge_bases.append( KnowledgeUserModel.model_validate( { - **self._to_knowledge_model(knowledge, db=db).model_dump(), + **self._to_knowledge_model( + knowledge, + access_grants=grants_map.get(knowledge.id, []), + db=db, + ).model_dump(), "user": user.model_dump() if user else None, } ) @@ -261,13 +274,20 @@ class KnowledgeTable: items = query.all() + knowledge_ids = [kb.id for kb, _ in items] + grants_map = AccessGrants.get_grants_by_resources( + "knowledge", knowledge_ids, db=db + ) + knowledge_bases = [] for knowledge_base, user in items: knowledge_bases.append( KnowledgeUserModel.model_validate( { **self._to_knowledge_model( - knowledge_base, db=db + knowledge_base, + access_grants=grants_map.get(knowledge_base.id, []), + db=db, ).model_dump(), "user": ( UserModel.model_validate(user).model_dump() @@ -440,8 +460,16 @@ class KnowledgeTable: .filter(KnowledgeFile.file_id == file_id) .all() ) + knowledge_ids = [k.id for k in knowledges] + grants_map = AccessGrants.get_grants_by_resources( + "knowledge", knowledge_ids, db=db + ) return [ - self._to_knowledge_model(knowledge, db=db) + self._to_knowledge_model( + knowledge, + access_grants=grants_map.get(knowledge.id, []), + db=db, + ) for knowledge in knowledges ] except Exception: diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index cfece00e35..3bffe4ddcf 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -144,11 +144,20 @@ class ModelsTable: ) -> list[AccessGrantModel]: return AccessGrants.get_grants_by_resource("model", model_id, db=db) - def _to_model_model(self, model: Model, db: Optional[Session] = None) -> ModelModel: + def _to_model_model( + self, + model: Model, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, + ) -> ModelModel: model_data = ModelModel.model_validate(model).model_dump( exclude={"access_grants"} ) - model_data["access_grants"] = self._get_access_grants(model_data["id"], db=db) + model_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(model_data["id"], db=db) + ) return ModelModel.model_validate(model_data) def insert_new_model( @@ -181,8 +190,14 @@ class ModelsTable: def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]: with get_db_context(db) as db: + all_models = db.query(Model).all() + model_ids = [model.id for model in all_models] + grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db) return [ - self._to_model_model(model, db=db) for model in db.query(Model).all() + self._to_model_model( + model, access_grants=grants_map.get(model.id, []), db=db + ) + for model in all_models ] def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]: @@ -190,9 +205,11 @@ class ModelsTable: all_models = db.query(Model).filter(Model.base_model_id != None).all() user_ids = list(set(model.user_id for model in all_models)) + model_ids = [model.id for model in all_models] users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} + grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db) models = [] for model in all_models: @@ -200,7 +217,11 @@ class ModelsTable: models.append( ModelUserResponse.model_validate( { - **self._to_model_model(model, db=db).model_dump(), + **self._to_model_model( + model, + access_grants=grants_map.get(model.id, []), + db=db, + ).model_dump(), "user": user.model_dump() if user else None, } ) @@ -209,9 +230,14 @@ class ModelsTable: def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]: with get_db_context(db) as db: + all_models = db.query(Model).filter(Model.base_model_id == None).all() + model_ids = [model.id for model in all_models] + grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db) return [ - self._to_model_model(model, db=db) - for model in db.query(Model).filter(Model.base_model_id == None).all() + self._to_model_model( + model, access_grants=grants_map.get(model.id, []), db=db + ) + for model in all_models ] def get_models_by_user_id( @@ -325,11 +351,18 @@ class ModelsTable: items = query.all() + model_ids = [model.id for model, _ in items] + grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db) + models = [] for model, user in items: models.append( ModelUserResponse( - **self._to_model_model(model, db=db).model_dump(), + **self._to_model_model( + model, + access_grants=grants_map.get(model.id, []), + db=db, + ).model_dump(), user=( UserResponse(**UserModel.model_validate(user).model_dump()) if user @@ -356,7 +389,18 @@ class ModelsTable: try: with get_db_context(db) as db: models = db.query(Model).filter(Model.id.in_(ids)).all() - return [self._to_model_model(model, db=db) for model in models] + model_ids = [model.id for model in models] + grants_map = AccessGrants.get_grants_by_resources( + "model", model_ids, db=db + ) + return [ + self._to_model_model( + model, + access_grants=grants_map.get(model.id, []), + db=db, + ) + for model in models + ] except Exception: return [] @@ -465,9 +509,18 @@ class ModelsTable: db.commit() + all_models = db.query(Model).all() + model_ids = [model.id for model in all_models] + grants_map = AccessGrants.get_grants_by_resources( + "model", model_ids, db=db + ) return [ - self._to_model_model(model, db=db) - for model in db.query(Model).all() + self._to_model_model( + model, + access_grants=grants_map.get(model.id, []), + db=db, + ) + for model in all_models ] except Exception as e: log.exception(f"Error syncing models for user {user_id}: {e}") diff --git a/backend/open_webui/models/notes.py b/backend/open_webui/models/notes.py index d17c749d1c..ff8a3ac635 100644 --- a/backend/open_webui/models/notes.py +++ b/backend/open_webui/models/notes.py @@ -93,9 +93,18 @@ class NoteTable: ) -> list[AccessGrantModel]: return AccessGrants.get_grants_by_resource("note", note_id, db=db) - def _to_note_model(self, note: Note, db: Optional[Session] = None) -> NoteModel: + def _to_note_model( + self, + note: Note, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, + ) -> NoteModel: note_data = NoteModel.model_validate(note).model_dump(exclude={"access_grants"}) - note_data["access_grants"] = self._get_access_grants(note_data["id"], db=db) + note_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(note_data["id"], db=db) + ) return NoteModel.model_validate(note_data) def _has_permission(self, db, query, filter: dict, permission: str = "read"): @@ -142,7 +151,14 @@ class NoteTable: if limit is not None: query = query.limit(limit) notes = query.all() - return [self._to_note_model(note, db=db) for note in notes] + note_ids = [note.id for note in notes] + grants_map = AccessGrants.get_grants_by_resources("note", note_ids, db=db) + return [ + self._to_note_model( + note, access_grants=grants_map.get(note.id, []), db=db + ) + for note in notes + ] def search_notes( self, @@ -227,11 +243,18 @@ class NoteTable: items = query.all() + note_ids = [note.id for note, _ in items] + grants_map = AccessGrants.get_grants_by_resources("note", note_ids, db=db) + notes = [] for note, user in items: notes.append( NoteUserResponse( - **self._to_note_model(note, db=db).model_dump(), + **self._to_note_model( + note, + access_grants=grants_map.get(note.id, []), + db=db, + ).model_dump(), user=( UserResponse(**UserModel.model_validate(user).model_dump()) if user @@ -266,7 +289,14 @@ class NoteTable: query = query.limit(limit) notes = query.all() - return [self._to_note_model(note, db=db) for note in notes] + note_ids = [note.id for note in notes] + grants_map = AccessGrants.get_grants_by_resources("note", note_ids, db=db) + return [ + self._to_note_model( + note, access_grants=grants_map.get(note.id, []), db=db + ) + for note in notes + ] def get_note_by_id( self, id: str, db: Optional[Session] = None diff --git a/backend/open_webui/models/oauth_sessions.py b/backend/open_webui/models/oauth_sessions.py index 538937483f..fbcd763f34 100644 --- a/backend/open_webui/models/oauth_sessions.py +++ b/backend/open_webui/models/oauth_sessions.py @@ -188,6 +188,7 @@ class OAuthSessionTable: session = ( db.query(OAuthSession) .filter_by(provider=provider, user_id=user_id) + .order_by(OAuthSession.created_at.desc()) .first() ) if session: diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index 3ab7a496ab..e32621f4e5 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -97,12 +97,19 @@ class PromptsTable: return AccessGrants.get_grants_by_resource("prompt", prompt_id, db=db) def _to_prompt_model( - self, prompt: Prompt, db: Optional[Session] = None + self, + prompt: Prompt, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, ) -> PromptModel: prompt_data = PromptModel.model_validate(prompt).model_dump( exclude={"access_grants"} ) - prompt_data["access_grants"] = self._get_access_grants(prompt_data["id"], db=db) + prompt_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(prompt_data["id"], db=db) + ) return PromptModel.model_validate(prompt_data) def insert_new_prompt( @@ -206,9 +213,13 @@ class PromptsTable: ) user_ids = list(set(prompt.user_id for prompt in all_prompts)) + prompt_ids = [prompt.id for prompt in all_prompts] users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} + grants_map = AccessGrants.get_grants_by_resources( + "prompt", prompt_ids, db=db + ) prompts = [] for prompt in all_prompts: @@ -216,7 +227,11 @@ class PromptsTable: prompts.append( PromptUserResponse.model_validate( { - **self._to_prompt_model(prompt, db=db).model_dump(), + **self._to_prompt_model( + prompt, + access_grants=grants_map.get(prompt.id, []), + db=db, + ).model_dump(), "user": user.model_dump() if user else None, } ) @@ -259,7 +274,6 @@ class PromptsTable: # 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") @@ -330,11 +344,20 @@ class PromptsTable: items = query.all() + prompt_ids = [prompt.id for prompt, _ in items] + grants_map = AccessGrants.get_grants_by_resources( + "prompt", prompt_ids, db=db + ) + prompts = [] for prompt, user in items: prompts.append( PromptUserResponse( - **self._to_prompt_model(prompt, db=db).model_dump(), + **self._to_prompt_model( + prompt, + access_grants=grants_map.get(prompt.id, []), + db=db, + ).model_dump(), user=( UserResponse(**UserModel.model_validate(user).model_dump()) if user @@ -562,43 +585,24 @@ class PromptsTable: except Exception: return None - 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: - prompt = db.query(Prompt).filter_by(command=command).first() - if prompt: - PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) - AccessGrants.revoke_all_access("prompt", 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 delete_prompt_by_id(self, prompt_id: str, db: Optional[Session] = None) -> bool: - """Soft delete a prompt by setting is_active to False.""" + def toggle_prompt_active( + self, prompt_id: str, db: Optional[Session] = None + ) -> Optional[PromptModel]: + """Toggle the is_active flag on a prompt.""" 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) - AccessGrants.revoke_all_access("prompt", prompt.id, db=db) - - prompt.is_active = False + prompt.is_active = not prompt.is_active prompt.updated_at = int(time.time()) db.commit() - return True - return False + db.refresh(prompt) + return self._to_prompt_model(prompt, db=db) + return None except Exception: - return False + return None - def hard_delete_prompt_by_command( + def delete_prompt_by_command( self, command: str, db: Optional[Session] = None ) -> bool: """Permanently delete a prompt and its history.""" @@ -609,8 +613,23 @@ class PromptsTable: PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) AccessGrants.revoke_all_access("prompt", prompt.id, db=db) - # Delete prompt - db.query(Prompt).filter_by(command=command).delete() + db.delete(prompt) + db.commit() + return True + return False + except Exception: + return False + + def delete_prompt_by_id(self, prompt_id: 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(id=prompt_id).first() + if prompt: + PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + AccessGrants.revoke_all_access("prompt", prompt.id, db=db) + + db.delete(prompt) db.commit() return True return False diff --git a/backend/open_webui/models/skills.py b/backend/open_webui/models/skills.py index 1262830153..6bd5affce8 100644 --- a/backend/open_webui/models/skills.py +++ b/backend/open_webui/models/skills.py @@ -110,11 +110,20 @@ class SkillsTable: ) -> list[AccessGrantModel]: return AccessGrants.get_grants_by_resource("skill", skill_id, db=db) - def _to_skill_model(self, skill: Skill, db: Optional[Session] = None) -> SkillModel: + def _to_skill_model( + self, + skill: Skill, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, + ) -> SkillModel: skill_data = SkillModel.model_validate(skill).model_dump( exclude={"access_grants"} ) - skill_data["access_grants"] = self._get_access_grants(skill_data["id"], db=db) + skill_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(skill_data["id"], db=db) + ) return SkillModel.model_validate(skill_data) def insert_new_skill( @@ -172,9 +181,11 @@ class SkillsTable: all_skills = db.query(Skill).order_by(Skill.updated_at.desc()).all() user_ids = list(set(skill.user_id for skill in all_skills)) + skill_ids = [skill.id for skill in all_skills] users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} + grants_map = AccessGrants.get_grants_by_resources("skill", skill_ids, db=db) skills = [] for skill in all_skills: @@ -182,7 +193,11 @@ class SkillsTable: skills.append( SkillUserModel.model_validate( { - **self._to_skill_model(skill, db=db).model_dump(), + **self._to_skill_model( + skill, + access_grants=grants_map.get(skill.id, []), + db=db, + ).model_dump(), "user": user.model_dump() if user else None, } ) @@ -267,11 +282,20 @@ class SkillsTable: items = query.all() + skill_ids = [skill.id for skill, _ in items] + grants_map = AccessGrants.get_grants_by_resources( + "skill", skill_ids, db=db + ) + skills = [] for skill, user in items: skills.append( SkillUserResponse( - **self._to_skill_model(skill, db=db).model_dump(), + **self._to_skill_model( + skill, + access_grants=grants_map.get(skill.id, []), + db=db, + ).model_dump(), user=( UserResponse( **UserModel.model_validate(user).model_dump() diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index eaac4c385d..62fe71abee 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -2,7 +2,7 @@ import logging import time from typing import Optional -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, defer from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.users import Users, UserResponse from open_webui.models.groups import Groups @@ -100,9 +100,18 @@ class ToolsTable: ) -> list[AccessGrantModel]: return AccessGrants.get_grants_by_resource("tool", tool_id, db=db) - def _to_tool_model(self, tool: Tool, db: Optional[Session] = None) -> ToolModel: + def _to_tool_model( + self, + tool: Tool, + access_grants: Optional[list[AccessGrantModel]] = None, + db: Optional[Session] = None, + ) -> ToolModel: tool_data = ToolModel.model_validate(tool).model_dump(exclude={"access_grants"}) - tool_data["access_grants"] = self._get_access_grants(tool_data["id"], db=db) + tool_data["access_grants"] = ( + access_grants + if access_grants is not None + else self._get_access_grants(tool_data["id"], db=db) + ) return ToolModel.model_validate(tool_data) def insert_new_tool( @@ -147,14 +156,21 @@ class ToolsTable: except Exception: return None - def get_tools(self, db: Optional[Session] = None) -> list[ToolUserModel]: + def get_tools( + self, defer_content: bool = False, db: Optional[Session] = None + ) -> list[ToolUserModel]: with get_db_context(db) as db: - all_tools = db.query(Tool).order_by(Tool.updated_at.desc()).all() + query = db.query(Tool).order_by(Tool.updated_at.desc()) + if defer_content: + query = query.options(defer(Tool.content), defer(Tool.specs)) + all_tools = query.all() user_ids = list(set(tool.user_id for tool in all_tools)) + tool_ids = [tool.id for tool in all_tools] users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} + grants_map = AccessGrants.get_grants_by_resources("tool", tool_ids, db=db) tools = [] for tool in all_tools: @@ -162,7 +178,11 @@ class ToolsTable: tools.append( ToolUserModel.model_validate( { - **self._to_tool_model(tool, db=db).model_dump(), + **self._to_tool_model( + tool, + access_grants=grants_map.get(tool.id, []), + db=db, + ).model_dump(), "user": user.model_dump() if user else None, } ) @@ -170,9 +190,9 @@ class ToolsTable: return tools def get_tools_by_user_id( - self, user_id: str, permission: str = "write", db: Optional[Session] = None + self, user_id: str, permission: str = "write", defer_content: bool = False, db: Optional[Session] = None ) -> list[ToolUserModel]: - tools = self.get_tools(db=db) + tools = self.get_tools(defer_content=defer_content, db=db) user_group_ids = { group.id for group in Groups.get_groups_by_member_id(user_id, db=db) } diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 96fd9d3f89..d328ba51a8 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -292,7 +292,7 @@ async def query_doc_with_hybrid_search( # retrieve only min(k, k_reranker) items, sort and cut by distance if k < k_reranker if k < k_reranker: sorted_items = sorted( - zip(distances, metadatas, documents), key=lambda x: x[0], reverse=True + zip(distances, documents, metadatas), key=lambda x: x[0], reverse=True ) sorted_items = sorted_items[:k] @@ -803,6 +803,7 @@ def get_embedding_function( embedding_batch_size, azure_api_version=None, enable_async=True, + concurrent_requests=0, ) -> Awaitable: if embedding_engine == "": # Sentence transformers: CPU-bound sync operation @@ -844,11 +845,25 @@ def get_embedding_function( log.debug( f"generate_multiple_async: Processing {len(batches)} batches in parallel" ) - # Execute all batches in parallel - tasks = [ - embedding_function(batch, prefix=prefix, user=user) - for batch in batches - ] + # Use semaphore to limit concurrent embedding API requests + # 0 = unlimited (no semaphore) + if concurrent_requests: + semaphore = asyncio.Semaphore(concurrent_requests) + + async def generate_batch_with_semaphore(batch): + async with semaphore: + return await embedding_function( + batch, prefix=prefix, user=user + ) + + tasks = [ + generate_batch_with_semaphore(batch) for batch in batches + ] + else: + tasks = [ + embedding_function(batch, prefix=prefix, user=user) + for batch in batches + ] batch_results = await asyncio.gather(*tasks) else: log.debug( diff --git a/backend/open_webui/retrieval/web/ydc.py b/backend/open_webui/retrieval/web/ydc.py new file mode 100644 index 0000000000..21d725a895 --- /dev/null +++ b/backend/open_webui/retrieval/web/ydc.py @@ -0,0 +1,73 @@ +import logging +from typing import Optional, List + +import requests +from open_webui.retrieval.web.main import SearchResult, get_filtered_results + +log = logging.getLogger(__name__) + + +def search_youcom( + api_key: str, + query: str, + count: int, + filter_list: Optional[List[str]] = None, + language: str = "EN", +) -> List[SearchResult]: + """Search using You.com's YDC Index API and return the results as a list of SearchResult objects. + + Args: + api_key (str): A You.com API key + query (str): The query to search for + count (int): Maximum number of results to return + filter_list (list[str], optional): Domain filter list + language (str): Language code for search results (default: "EN") + """ + url = "https://ydc-index.io/v1/search" + headers = { + "Accept": "application/json", + "X-API-KEY": api_key, + } + params = { + "query": query, + "count": count, + "language": language, + } + + response = requests.get(url, headers=headers, params=params) + response.raise_for_status() + + json_response = response.json() + results = json_response.get("results", {}).get("web", []) + + if filter_list: + results = get_filtered_results(results, filter_list) + + return [ + SearchResult( + link=result["url"], + title=result.get("title"), + snippet=_build_snippet(result), + ) + for result in results[:count] + ] + + +def _build_snippet(result: dict) -> str: + """Combine the description and snippets list into a single string. + + The You.com API returns a short ``description`` plus a ``snippets`` + list with richer passages. Merging them gives downstream retrieval + (embedding, BM25, bypass-loader context) the most content to work with. + """ + parts: list[str] = [] + + description = result.get("description") + if description: + parts.append(description) + + snippets = result.get("snippets") + if snippets and isinstance(snippets, list): + parts.extend(snippets) + + return "\n\n".join(parts) diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 8ada514cc2..9b246b297b 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -700,8 +700,9 @@ async def signup_handler( Returns the newly created UserModel. Raises HTTPException on failure. """ - has_users = Users.has_users(db=db) - role = "admin" if not has_users else request.app.state.config.DEFAULT_USER_ROLE + # Insert with default role first to avoid TOCTOU race on first signup. + # If has_users() is checked before insert, concurrent requests during + # first-user registration can all see an empty table and each get admin. hashed = get_password_hash(password) user = Auths.insert_new_auth( @@ -709,12 +710,19 @@ async def signup_handler( password=hashed, name=name, profile_image_url=profile_image_url, - role=role, + role=request.app.state.config.DEFAULT_USER_ROLE, db=db, ) if not user: raise HTTPException(500, detail=ERROR_MESSAGES.CREATE_USER_ERROR) + # Atomically check if this is the only user *after* the insert. + # Only the single user present at this point should become admin. + if Users.get_num_users(db=db) == 1: + Users.update_user_role_by_id(user.id, "admin", db=db) + user = Users.get_user_by_id(user.id, db=db) + request.app.state.config.ENABLE_SIGNUP = False + if request.app.state.config.WEBHOOK_URL: await post_webhook( request.app.state.WEBUI_NAME, @@ -727,10 +735,6 @@ async def signup_handler( }, ) - if not has_users: - # Disable signup after the first user is created - request.app.state.config.ENABLE_SIGNUP = False - apply_default_group_assignment( request.app.state.config.DEFAULT_GROUP_ID, user.id, diff --git a/backend/open_webui/routers/chats.py b/backend/open_webui/routers/chats.py index 69e47123f0..ff3f6c78c7 100644 --- a/backend/open_webui/routers/chats.py +++ b/backend/open_webui/routers/chats.py @@ -723,10 +723,7 @@ async def get_chat_list_by_folder_id( async def get_user_pinned_chats( user=Depends(get_verified_user), db: Session = Depends(get_session) ): - return [ - ChatTitleIdResponse(**chat.model_dump()) - for chat in Chats.get_pinned_chats_by_user_id(user.id, db=db) - ] + return Chats.get_pinned_chats_by_user_id(user.id, db=db) ############################ @@ -821,18 +818,13 @@ async def get_archived_session_user_chat_list( if direction: filter["direction"] = direction - chat_list = [ - ChatTitleIdResponse(**chat.model_dump()) - for chat in Chats.get_archived_chat_list_by_user_id( - user.id, - filter=filter, - skip=skip, - limit=limit, - db=db, - ) - ] - - return chat_list + return Chats.get_archived_chat_list_by_user_id( + user.id, + filter=filter, + skip=skip, + limit=limit, + db=db, + ) ############################ @@ -887,18 +879,13 @@ async def get_shared_session_user_chat_list( 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 + return Chats.get_shared_chat_list_by_user_id( + user.id, + filter=filter, + skip=skip, + limit=limit, + db=db, + ) ############################ diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 4a12db5cd9..9022ca66ff 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -47,6 +47,7 @@ from open_webui.routers.audio import transcribe from open_webui.storage.provider import Storage +from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.misc import strict_match_mime_type from pydantic import BaseModel @@ -362,7 +363,7 @@ async def list_files( content: bool = Query(True), db: Session = Depends(get_session), ): - if user.role == "admin": + if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL: files = Files.get_files(db=db) else: files = Files.get_files_by_user_id(user.id, db=db) @@ -398,8 +399,10 @@ async def search_files( Search for files by filename with support for wildcard patterns. Uses SQL-based filtering with pagination for better performance. """ - # Determine user_id: null for admin (search all), user.id for regular users - user_id = None if user.role == "admin" else user.id + # Determine user_id: null for admin with bypass (search all), user.id otherwise + user_id = ( + None if (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) else user.id + ) # Use optimized database query with pagination files = Files.search_files( @@ -689,6 +692,8 @@ async def get_file_content_by_id( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) + except HTTPException as e: + raise e except Exception as e: log.exception(e) log.error("Error getting file content") @@ -740,6 +745,8 @@ async def get_html_file_content_by_id( status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) + except HTTPException as e: + raise e except Exception as e: log.exception(e) log.error("Error getting file content") diff --git a/backend/open_webui/routers/notes.py b/backend/open_webui/routers/notes.py index 8d1a66c4af..41bb65f55a 100644 --- a/backend/open_webui/routers/notes.py +++ b/backend/open_webui/routers/notes.py @@ -36,6 +36,14 @@ log = logging.getLogger(__name__) router = APIRouter() + +def _truncate_note_data(data: Optional[dict], max_length: int = 1000) -> Optional[dict]: + if not data: + return data + md = (data.get("content") or {}).get("md") or "" + return {"content": {"md": md[:max_length]}} + + ############################ # GetNotes ############################ @@ -82,6 +90,7 @@ async def get_notes( NoteUserResponse( **{ **note.model_dump(), + "data": _truncate_note_data(note.data), "user": UserResponse(**users[note.user_id].model_dump()), } ) @@ -135,7 +144,10 @@ async def search_notes( filter["user_id"] = user.id - return Notes.search_notes(user.id, filter, skip=skip, limit=limit, db=db) + result = Notes.search_notes(user.id, filter, skip=skip, limit=limit, db=db) + for note in result.items: + note.data = _truncate_note_data(note.data) + return result ############################ diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 9fe421b4e9..767aa791f1 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -57,6 +57,7 @@ from open_webui.utils.misc import ( from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.headers import include_user_info_headers +from open_webui.utils.anthropic import is_anthropic_url, get_anthropic_models log = logging.getLogger(__name__) @@ -91,6 +92,12 @@ async def send_get_request(url, key=None, user: UserModel = None): return None +async def get_models_request(url, key=None, user: UserModel = None): + if is_anthropic_url(url): + return await get_anthropic_models(url, key, user=user) + return await send_get_request(f"{url}/models", key, user=user) + + def openai_reasoning_model_handler(payload): """ Handle reasoning model specific parameters @@ -365,13 +372,7 @@ async def get_all_models_responses(request: Request, user: UserModel) -> list: request_tasks = [] for idx, url in enumerate(api_base_urls): if (str(idx) not in api_configs) and (url not in api_configs): # Legacy support - request_tasks.append( - send_get_request( - f"{url}/models", - api_keys[idx], - user=user, - ) - ) + request_tasks.append(get_models_request(url, api_keys[idx], user=user)) else: api_config = api_configs.get( str(idx), @@ -384,11 +385,7 @@ async def get_all_models_responses(request: Request, user: UserModel) -> list: if enable: if len(model_ids) == 0: request_tasks.append( - send_get_request( - f"{url}/models", - api_keys[idx], - user=user, - ) + get_models_request(url, api_keys[idx], user=user) ) else: model_list = { @@ -594,6 +591,10 @@ async def get_models( "data": api_config.get("model_ids", []) or [], "object": "list", } + elif is_anthropic_url(url): + models = await get_anthropic_models(url, key, user=user) + if models is None: + raise Exception("Failed to connect to Anthropic API") else: async with session.get( f"{url}/models", @@ -602,7 +603,6 @@ async def get_models( ssl=AIOHTTP_CLIENT_SESSION_SSL, ) as r: if r.status != 200: - # Extract response error details if available error_detail = f"HTTP Error: {r.status}" try: res = await r.json() @@ -614,9 +614,7 @@ async def get_models( response_data = await r.json() - # Check if we're calling OpenAI API based on the URL if "api.openai.com" in url: - # Filter models according to the specified conditions response_data["data"] = [ model for model in response_data.get("data", []) @@ -707,6 +705,15 @@ async def verify_connection( ) return response_data + elif is_anthropic_url(url): + result = await get_anthropic_models(url, key) + if result is None: + raise HTTPException( + status_code=500, detail="Failed to connect to Anthropic API" + ) + if "error" in result: + raise HTTPException(status_code=500, detail=result["error"]) + return result else: async with session.get( f"{url}/models", @@ -1181,7 +1188,10 @@ async def embeddings(request: Request, form_data: dict, user): request, url, key, api_config, user=user ) try: - session = aiohttp.ClientSession(trust_env=True) + session = aiohttp.ClientSession( + trust_env=True, + timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT), + ) r = await session.request( method="POST", url=f"{url}/embeddings", @@ -1408,7 +1418,10 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)): else: request_url = f"{url}/{path}" - session = aiohttp.ClientSession(trust_env=True) + session = aiohttp.ClientSession( + trust_env=True, + timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT), + ) r = await session.request( method=request.method, url=request_url, diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py index 2491578959..9653571fbb 100644 --- a/backend/open_webui/routers/prompts.py +++ b/backend/open_webui/routers/prompts.py @@ -497,6 +497,48 @@ async def update_prompt_access_by_id( return Prompts.get_prompt_by_id(prompt_id, db=db) +############################ +# TogglePromptActiveById +############################ + + +@router.post("/id/{prompt_id}/toggle", response_model=Optional[PromptModel]) +async def toggle_prompt_active( + 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 AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) + and user.role != "admin" + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + result = Prompts.toggle_prompt_active(prompt.id, db=db) + if result: + return result + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT(), + ) + + ############################ # DeletePromptById ############################ diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 6c95d3e606..ae894608c5 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -77,6 +77,7 @@ 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.web.ydc import search_youcom from open_webui.retrieval.utils import ( get_content_from_url, @@ -270,6 +271,7 @@ async def get_status(request: Request): "RAG_RERANKING_MODEL": request.app.state.config.RAG_RERANKING_MODEL, "RAG_EMBEDDING_BATCH_SIZE": request.app.state.config.RAG_EMBEDDING_BATCH_SIZE, "ENABLE_ASYNC_EMBEDDING": request.app.state.config.ENABLE_ASYNC_EMBEDDING, + "RAG_EMBEDDING_CONCURRENT_REQUESTS": request.app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS, } @@ -281,6 +283,7 @@ async def get_embedding_config(request: Request, user=Depends(get_admin_user)): "RAG_EMBEDDING_MODEL": request.app.state.config.RAG_EMBEDDING_MODEL, "RAG_EMBEDDING_BATCH_SIZE": request.app.state.config.RAG_EMBEDDING_BATCH_SIZE, "ENABLE_ASYNC_EMBEDDING": request.app.state.config.ENABLE_ASYNC_EMBEDDING, + "RAG_EMBEDDING_CONCURRENT_REQUESTS": request.app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS, "openai_config": { "url": request.app.state.config.RAG_OPENAI_API_BASE_URL, "key": request.app.state.config.RAG_OPENAI_API_KEY, @@ -321,6 +324,7 @@ class EmbeddingModelUpdateForm(BaseModel): RAG_EMBEDDING_MODEL: str RAG_EMBEDDING_BATCH_SIZE: Optional[int] = 1 ENABLE_ASYNC_EMBEDDING: Optional[bool] = True + RAG_EMBEDDING_CONCURRENT_REQUESTS: Optional[int] = 0 def unload_embedding_model(request: Request): @@ -355,6 +359,9 @@ async def update_embedding_config( request.app.state.config.ENABLE_ASYNC_EMBEDDING = ( form_data.ENABLE_ASYNC_EMBEDDING ) + request.app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS = ( + form_data.RAG_EMBEDDING_CONCURRENT_REQUESTS + ) if request.app.state.config.RAG_EMBEDDING_ENGINE in [ "ollama", @@ -422,6 +429,7 @@ async def update_embedding_config( else None ), enable_async=request.app.state.config.ENABLE_ASYNC_EMBEDDING, + concurrent_requests=request.app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS, ) return { @@ -430,6 +438,7 @@ async def update_embedding_config( "RAG_EMBEDDING_MODEL": request.app.state.config.RAG_EMBEDDING_MODEL, "RAG_EMBEDDING_BATCH_SIZE": request.app.state.config.RAG_EMBEDDING_BATCH_SIZE, "ENABLE_ASYNC_EMBEDDING": request.app.state.config.ENABLE_ASYNC_EMBEDDING, + "RAG_EMBEDDING_CONCURRENT_REQUESTS": request.app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS, "openai_config": { "url": request.app.state.config.RAG_OPENAI_API_BASE_URL, "key": request.app.state.config.RAG_OPENAI_API_KEY, @@ -585,6 +594,7 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)): "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, + "YOUCOM_API_KEY": request.app.state.config.YOUCOM_API_KEY, }, } @@ -651,6 +661,7 @@ class WebConfig(BaseModel): YANDEX_WEB_SEARCH_URL: Optional[str] = None YANDEX_WEB_SEARCH_API_KEY: Optional[str] = None YANDEX_WEB_SEARCH_CONFIG: Optional[str] = None + YOUCOM_API_KEY: Optional[str] = None class ConfigForm(BaseModel): @@ -1219,6 +1230,7 @@ async def update_rag_config( request.app.state.config.YANDEX_WEB_SEARCH_CONFIG = ( form_data.web.YANDEX_WEB_SEARCH_CONFIG ) + request.app.state.config.YOUCOM_API_KEY = form_data.web.YOUCOM_API_KEY return { "status": True, @@ -1348,6 +1360,7 @@ async def update_rag_config( "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, + "YOUCOM_API_KEY": request.app.state.config.YOUCOM_API_KEY, }, } @@ -1605,6 +1618,7 @@ def save_docs_to_vector_db( else None ), enable_async=request.app.state.config.ENABLE_ASYNC_EMBEDDING, + concurrent_requests=request.app.state.config.RAG_EMBEDDING_CONCURRENT_REQUESTS, ) # Run async embedding in sync context using the main event loop @@ -1762,6 +1776,7 @@ def process_file( DOCLING_API_KEY=request.app.state.config.DOCLING_API_KEY, DOCLING_PARAMS=request.app.state.config.DOCLING_PARAMS, PDF_EXTRACT_IMAGES=request.app.state.config.PDF_EXTRACT_IMAGES, + PDF_LOADER_MODE=request.app.state.config.PDF_LOADER_MODE, DOCUMENT_INTELLIGENCE_ENDPOINT=request.app.state.config.DOCUMENT_INTELLIGENCE_ENDPOINT, DOCUMENT_INTELLIGENCE_KEY=request.app.state.config.DOCUMENT_INTELLIGENCE_KEY, DOCUMENT_INTELLIGENCE_MODEL=request.app.state.config.DOCUMENT_INTELLIGENCE_MODEL, @@ -1951,6 +1966,9 @@ async def process_web( request: Request, form_data: ProcessUrlForm, process: bool = Query(True, description="Whether to process and save the content"), + overwrite: bool = Query( + True, description="Whether to overwrite existing collection" + ), user=Depends(get_verified_user), ): try: @@ -1970,7 +1988,7 @@ async def process_web( request, docs, collection_name, - overwrite=True, + overwrite=overwrite, user=user, ) else: @@ -2306,6 +2324,13 @@ def search_web( request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST, user=user, ) + elif engine == "youcom": + return search_youcom( + request.app.state.config.YOUCOM_API_KEY, + query, + request.app.state.config.WEB_SEARCH_RESULT_COUNT, + request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST, + ) else: raise Exception("No search engine API key found in environment variables") @@ -2345,7 +2370,7 @@ async def process_web_search( # Limited concurrency with semaphore semaphore = asyncio.Semaphore(concurrent_limit) - async def search_with_limit(query): + async def search_query_with_semaphore(query): async with semaphore: return await run_in_threadpool( search_web, @@ -2355,7 +2380,9 @@ async def process_web_search( user, ) - search_tasks = [search_with_limit(query) for query in form_data.queries] + search_tasks = [ + search_query_with_semaphore(query) for query in form_data.queries + ] else: # Unlimited parallel execution (previous behavior) search_tasks = [ diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index 0c16eb99bd..4d56a7e97f 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -523,13 +523,17 @@ async def get_schemas(): @router.get("/Users", response_model=SCIMListResponse) async def get_users( request: Request, - startIndex: int = Query(1, ge=1), - count: int = Query(20, ge=1, le=100), + startIndex: int = Query(1), + count: int = Query(20), filter: Optional[str] = None, _: bool = Depends(get_scim_auth), db: Session = Depends(get_session), ): """List SCIM Users""" + # Clamp per SCIM 2.0 spec (RFC 7644 §3.4.2.4): + # startIndex < 1 SHALL be treated as 1; count < 0 SHALL be treated as 0. + startIndex = max(1, startIndex) + count = max(0, min(100, count)) skip = startIndex - 1 limit = count @@ -794,13 +798,18 @@ async def delete_user( @router.get("/Groups", response_model=SCIMListResponse) async def get_groups( request: Request, - startIndex: int = Query(1, ge=1), - count: int = Query(20, ge=1, le=100), + startIndex: int = Query(1), + count: int = Query(20), filter: Optional[str] = None, _: bool = Depends(get_scim_auth), db: Session = Depends(get_session), ): """List SCIM Groups""" + # Clamp per SCIM 2.0 spec (RFC 7644 §3.4.2.4): + # startIndex < 1 SHALL be treated as 1; count < 0 SHALL be treated as 0. + startIndex = max(1, startIndex) + count = max(0, min(100, count)) + # Get all groups groups_list = Groups.get_all_groups(db=db) diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index fab5039909..b2d35ccc6c 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -64,13 +64,13 @@ async def get_tools( tools = [] # Local Tools - for tool in Tools.get_tools(db=db): - tool_module = get_tool_module(request, tool.id) + for tool in Tools.get_tools(defer_content=True, db=db): + tool_module = request.app.state.TOOLS.get(tool.id) if hasattr(request.app.state, 'TOOLS') else None tools.append( ToolUserResponse( **{ **tool.model_dump(), - "has_user_valves": hasattr(tool_module, "UserValves"), + "has_user_valves": hasattr(tool_module, "UserValves") if tool_module else False, } ) ) @@ -196,27 +196,35 @@ async def get_tool_list( user=Depends(get_verified_user), db: Session = Depends(get_session) ): if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL: - tools = Tools.get_tools(db=db) + tools = Tools.get_tools(defer_content=True, db=db) else: - tools = Tools.get_tools_by_user_id(user.id, "read", db=db) + tools = Tools.get_tools_by_user_id(user.id, "read", defer_content=True, db=db) - return [ - ToolAccessResponse( - **tool.model_dump(), - write_access=( - (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) - or user.id == tool.user_id - or AccessGrants.has_access( - user_id=user.id, - resource_type="tool", - resource_id=tool.id, - permission="write", - db=db, + user_group_ids = { + group.id for group in Groups.get_groups_by_member_id(user.id, db=db) + } + + result = [] + for tool in tools: + has_write = ( + (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) + or user.id == tool.user_id + or any( + g.permission == "write" + and ( + (g.principal_type == "user" and (g.principal_id == user.id or g.principal_id == "*")) + or (g.principal_type == "group" and g.principal_id in user_group_ids) ) - ), + for g in tool.access_grants + ) ) - for tool in tools - ] + result.append( + ToolAccessResponse( + **tool.model_dump(), + write_access=has_write, + ) + ) + return result ############################ diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 78df66b8dc..c1d7e5c29e 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -511,6 +511,8 @@ async def channel_events(sid, data): async def ydoc_document_join(sid, data): """Handle user joining a document""" user = SESSION_POOL.get(sid) + if not user: + return try: document_id = data["document_id"] @@ -683,11 +685,13 @@ async def yjs_document_update(sid, data): skip_sid=sid, ) + user = SESSION_POOL.get(sid) + if not user: + return + async def debounced_save(): await asyncio.sleep(0.5) - await document_save_handler( - document_id, data.get("data", {}), SESSION_POOL.get(sid) - ) + await document_save_handler(document_id, data.get("data", {}), user) if data.get("data"): await create_task(REDIS, debounced_save(), document_id) diff --git a/backend/open_webui/test/apps/webui/routers/test_prompts.py b/backend/open_webui/test/apps/webui/routers/test_prompts.py deleted file mode 100644 index d91bf77dc5..0000000000 --- a/backend/open_webui/test/apps/webui/routers/test_prompts.py +++ /dev/null @@ -1,91 +0,0 @@ -from test.util.abstract_integration_test import AbstractPostgresTest -from test.util.mock_user import mock_webui_user - - -class TestPrompts(AbstractPostgresTest): - BASE_PATH = "/api/v1/prompts" - - def test_prompts(self): - # Get all prompts - with mock_webui_user(id="2"): - response = self.fast_api_client.get(self.create_url("/")) - assert response.status_code == 200 - assert len(response.json()) == 0 - - # Create a two new prompts - with mock_webui_user(id="2"): - response = self.fast_api_client.post( - self.create_url("/create"), - json={ - "command": "/my-command", - "title": "Hello World", - "content": "description", - }, - ) - assert response.status_code == 200 - with mock_webui_user(id="3"): - response = self.fast_api_client.post( - self.create_url("/create"), - json={ - "command": "/my-command2", - "title": "Hello World 2", - "content": "description 2", - }, - ) - assert response.status_code == 200 - - # Get all prompts - with mock_webui_user(id="2"): - response = self.fast_api_client.get(self.create_url("/")) - assert response.status_code == 200 - assert len(response.json()) == 2 - - # Get prompt by command - with mock_webui_user(id="2"): - response = self.fast_api_client.get(self.create_url("/command/my-command")) - assert response.status_code == 200 - data = response.json() - assert data["command"] == "/my-command" - assert data["title"] == "Hello World" - assert data["content"] == "description" - assert data["user_id"] == "2" - - # Update prompt - with mock_webui_user(id="2"): - response = self.fast_api_client.post( - self.create_url("/command/my-command2/update"), - json={ - "command": "irrelevant for request", - "title": "Hello World Updated", - "content": "description Updated", - }, - ) - assert response.status_code == 200 - data = response.json() - assert data["command"] == "/my-command2" - assert data["title"] == "Hello World Updated" - assert data["content"] == "description Updated" - assert data["user_id"] == "3" - - # Get prompt by command - with mock_webui_user(id="2"): - response = self.fast_api_client.get(self.create_url("/command/my-command2")) - assert response.status_code == 200 - data = response.json() - assert data["command"] == "/my-command2" - assert data["title"] == "Hello World Updated" - assert data["content"] == "description Updated" - assert data["user_id"] == "3" - - # Delete prompt - with mock_webui_user(id="2"): - response = self.fast_api_client.delete( - self.create_url("/command/my-command/delete") - ) - assert response.status_code == 200 - - # Get all prompts - with mock_webui_user(id="2"): - response = self.fast_api_client.get(self.create_url("/")) - assert response.status_code == 200 - assert len(response.json()) == 1 diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 31c79484ab..330baf1318 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -36,6 +36,8 @@ 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.models.memories import Memories +from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT from open_webui.utils.sanitize import sanitize_code log = logging.getLogger(__name__) @@ -634,6 +636,79 @@ async def replace_memory_content( return json.dumps({"error": str(e)}) +async def delete_memory( + memory_id: str, + __request__: Request = None, + __user__: dict = None, +) -> str: + """ + Delete a memory by its ID. + + :param memory_id: The ID of the memory to delete + :return: Confirmation that the memory was deleted + """ + if __request__ is None: + return json.dumps({"error": "Request context not available"}) + + try: + user = UserModel(**__user__) if __user__ else None + + result = Memories.delete_memory_by_id_and_user_id(memory_id, user.id) + + if result: + VECTOR_DB_CLIENT.delete( + collection_name=f"user-memory-{user.id}", ids=[memory_id] + ) + return json.dumps( + {"status": "success", "message": f"Memory {memory_id} deleted"}, + ensure_ascii=False, + ) + else: + return json.dumps({"error": "Memory not found or access denied"}) + except Exception as e: + log.exception(f"delete_memory error: {e}") + return json.dumps({"error": str(e)}) + + +async def list_memories( + __request__: Request = None, + __user__: dict = None, +) -> str: + """ + List all stored memories for the user. + + :return: JSON list of all memories with id, content, and dates + """ + if __request__ is None: + return json.dumps({"error": "Request context not available"}) + + try: + user = UserModel(**__user__) if __user__ else None + + memories = Memories.get_memories_by_user_id(user.id) + + if memories: + result = [ + { + "id": m.id, + "content": m.content, + "created_at": time.strftime( + "%Y-%m-%d %H:%M", time.localtime(m.created_at) + ), + "updated_at": time.strftime( + "%Y-%m-%d %H:%M", time.localtime(m.updated_at) + ), + } + for m in memories + ] + return json.dumps(result, ensure_ascii=False) + else: + return json.dumps([]) + except Exception as e: + log.exception(f"list_memories error: {e}") + return json.dumps({"error": str(e)}) + + # ============================================================================= # NOTES TOOLS # ============================================================================= diff --git a/backend/open_webui/utils/anthropic.py b/backend/open_webui/utils/anthropic.py new file mode 100644 index 0000000000..736f7238bc --- /dev/null +++ b/backend/open_webui/utils/anthropic.py @@ -0,0 +1,534 @@ +import json +import logging + +import aiohttp + +from open_webui.env import ( + AIOHTTP_CLIENT_SESSION_SSL, + AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST, + ENABLE_FORWARD_USER_INFO_HEADERS, +) +from open_webui.models.users import UserModel +from open_webui.utils.headers import include_user_info_headers + +log = logging.getLogger(__name__) + + +def is_anthropic_url(url: str) -> bool: + """Check if the URL is an Anthropic API endpoint.""" + return "api.anthropic.com" in url + + +async def get_anthropic_models(url: str, key: str, user: UserModel = None) -> dict: + """ + Fetch models from Anthropic's /v1/models endpoint with pagination. + Normalizes the response to OpenAI format. + """ + timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST) + all_models = [] + after_id = None + + try: + async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session: + headers = { + "x-api-key": key, + "anthropic-version": "2023-06-01", + } + + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) + + while True: + params = {"limit": 1000} + if after_id: + params["after_id"] = after_id + + async with session.get( + f"{url}/models", + headers=headers, + params=params, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as response: + if response.status != 200: + error_detail = f"HTTP Error: {response.status}" + try: + res = await response.json() + if "error" in res: + error_detail = f"External Error: {res['error']}" + except Exception: + pass + return {"object": "list", "data": [], "error": error_detail} + + data = await response.json() + + for model in data.get("data", []): + all_models.append( + { + "id": model.get("id"), + "object": "model", + "created": 0, + "owned_by": "anthropic", + "name": model.get("display_name", model.get("id")), + } + ) + + if not data.get("has_more", False): + break + after_id = data.get("last_id") + + except Exception as e: + log.error(f"Anthropic connection error: {e}") + return None + + return {"object": "list", "data": all_models} + + +############################## +# +# Anthropic Messages API Conversion Utilities +# +############################## + + +def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict: + """ + Convert an Anthropic Messages API request to OpenAI Chat Completions format. + + Anthropic format: + {model, messages: [{role, content}], system, max_tokens, ...} + OpenAI format: + {model, messages: [{role, content}], max_tokens, ...} + """ + openai_payload = {} + + # Model + openai_payload["model"] = anthropic_payload.get("model", "") + + # Build messages list + messages = [] + + # System prompt (Anthropic has it as top-level, OpenAI as a system message) + system = anthropic_payload.get("system") + if system: + if isinstance(system, str): + messages.append({"role": "system", "content": system}) + elif isinstance(system, list): + # Anthropic supports system as list of content blocks + text_parts = [] + for block in system: + if isinstance(block, dict) and block.get("type") == "text": + text_parts.append(block.get("text", "")) + elif isinstance(block, str): + text_parts.append(block) + messages.append({"role": "system", "content": "\n".join(text_parts)}) + + # Convert messages + for msg in anthropic_payload.get("messages", []): + role = msg.get("role", "user") + content = msg.get("content") + + if isinstance(content, str): + messages.append({"role": role, "content": content}) + elif isinstance(content, list): + # Convert Anthropic content blocks to OpenAI format + openai_content = [] + tool_calls = [] + + for block in content: + block_type = block.get("type", "text") + + if block_type == "text": + openai_content.append( + { + "type": "text", + "text": block.get("text", ""), + } + ) + elif block_type == "image": + source = block.get("source", {}) + if source.get("type") == "base64": + media_type = source.get("media_type", "image/png") + data = source.get("data", "") + openai_content.append( + { + "type": "image_url", + "image_url": { + "url": f"data:{media_type};base64,{data}", + }, + } + ) + elif source.get("type") == "url": + openai_content.append( + { + "type": "image_url", + "image_url": {"url": source.get("url", "")}, + } + ) + elif block_type == "tool_use": + tool_calls.append( + { + "id": block.get("id", ""), + "type": "function", + "function": { + "name": block.get("name", ""), + "arguments": ( + json.dumps(block.get("input", {})) + if isinstance(block.get("input"), dict) + else str(block.get("input", "{}")) + ), + }, + } + ) + elif block_type == "tool_result": + # Tool results become separate tool messages in OpenAI format + tool_content = block.get("content", "") + if isinstance(tool_content, list): + tool_text_parts = [] + for tc in tool_content: + if isinstance(tc, dict) and tc.get("type") == "text": + tool_text_parts.append(tc.get("text", "")) + tool_content = "\n".join(tool_text_parts) + + # Propagate error status if present + if block.get("is_error"): + tool_content = f"Error: {tool_content}" + + messages.append( + { + "role": "tool", + "tool_call_id": block.get("tool_use_id", ""), + "content": tool_content, + } + ) + + # Build the message + if tool_calls: + # Assistant message with tool calls + msg_dict = {"role": role} + if openai_content: + # If there's only text, flatten it + if len(openai_content) == 1 and openai_content[0]["type"] == "text": + msg_dict["content"] = openai_content[0]["text"] + else: + msg_dict["content"] = openai_content + else: + msg_dict["content"] = "" + msg_dict["tool_calls"] = tool_calls + messages.append(msg_dict) + elif openai_content: + # If there's only a single text block, flatten it to a string + if len(openai_content) == 1 and openai_content[0]["type"] == "text": + messages.append( + {"role": role, "content": openai_content[0]["text"]} + ) + else: + messages.append({"role": role, "content": openai_content}) + else: + messages.append({"role": role, "content": str(content) if content else ""}) + + openai_payload["messages"] = messages + + # max_tokens + if "max_tokens" in anthropic_payload: + openai_payload["max_tokens"] = anthropic_payload["max_tokens"] + + # Common parameters + for param in ("temperature", "top_p", "stop_sequences", "stream"): + if param in anthropic_payload: + if param == "stop_sequences": + openai_payload["stop"] = anthropic_payload[param] + else: + openai_payload[param] = anthropic_payload[param] + + # Tools conversion: Anthropic → OpenAI + if "tools" in anthropic_payload: + openai_tools = [] + for tool in anthropic_payload["tools"]: + openai_tools.append( + { + "type": "function", + "function": { + "name": tool.get("name", ""), + "description": tool.get("description", ""), + "parameters": tool.get("input_schema", {}), + }, + } + ) + openai_payload["tools"] = openai_tools + + # tool_choice + if "tool_choice" in anthropic_payload: + tc = anthropic_payload["tool_choice"] + if isinstance(tc, dict): + tc_type = tc.get("type", "auto") + if tc_type == "auto": + openai_payload["tool_choice"] = "auto" + elif tc_type == "any": + openai_payload["tool_choice"] = "required" + elif tc_type == "tool": + openai_payload["tool_choice"] = { + "type": "function", + "function": {"name": tc.get("name", "")}, + } + + return openai_payload + + +def convert_openai_to_anthropic_response( + openai_response: dict, model: str = "" +) -> dict: + """ + Convert a non-streaming OpenAI Chat Completions response to Anthropic Messages format. + """ + import uuid as _uuid + + choice = {} + if openai_response.get("choices"): + choice = openai_response["choices"][0] + + message = choice.get("message", {}) + finish_reason = choice.get("finish_reason", "stop") + + # Map finish_reason to stop_reason + stop_reason_map = { + "stop": "end_turn", + "length": "max_tokens", + "tool_calls": "tool_use", + "content_filter": "end_turn", + } + stop_reason = stop_reason_map.get(finish_reason, "end_turn") + + # Build content blocks + content = [] + msg_content = message.get("content") + if msg_content: + content.append({"type": "text", "text": msg_content}) + + # Tool calls → tool_use blocks + tool_calls = message.get("tool_calls", []) + for tc in tool_calls: + func = tc.get("function", {}) + try: + tool_input = json.loads(func.get("arguments", "{}")) + except (json.JSONDecodeError, TypeError): + tool_input = {} + content.append( + { + "type": "tool_use", + "id": tc.get("id", f"toolu_{_uuid.uuid4().hex[:24]}"), + "name": func.get("name", ""), + "input": tool_input, + } + ) + + # Usage + openai_usage = openai_response.get("usage", {}) + usage = { + "input_tokens": openai_usage.get("prompt_tokens", 0), + "output_tokens": openai_usage.get("completion_tokens", 0), + } + + return { + "id": openai_response.get("id", f"msg_{_uuid.uuid4().hex[:24]}"), + "type": "message", + "role": "assistant", + "content": content, + "model": model or openai_response.get("model", ""), + "stop_reason": stop_reason, + "stop_sequence": None, + "usage": usage, + } + + +async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str = ""): + """ + Convert an OpenAI SSE streaming response to Anthropic Messages SSE format. + + OpenAI sends: data: {"choices": [{"delta": {"content": "..."}}]} + Anthropic sends: event: content_block_delta\\ndata: {"type": "content_block_delta", ...} + + Handles text content, tool calls, and mixed content with proper + multi-block indexing as required by Anthropic's streaming protocol. + """ + import uuid as _uuid + + msg_id = f"msg_{_uuid.uuid4().hex[:24]}" + input_tokens = 0 + output_tokens = 0 + stop_reason = "end_turn" + + # Track content blocks with a running index. + # Each text block or tool_use block gets its own index. + current_block_index = 0 + text_block_open = False + + # Track tool call state: maps OpenAI tool_call index -> Anthropic block index + # This allows handling multiple concurrent tool calls. + tool_call_blocks = {} # {openai_tc_index: anthropic_block_index} + tool_call_started = {} # {openai_tc_index: bool} + + # Emit message_start + message_start = { + "type": "message_start", + "message": { + "id": msg_id, + "type": "message", + "role": "assistant", + "content": [], + "model": model, + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 0, "output_tokens": 0}, + }, + } + yield f"event: message_start\ndata: {json.dumps(message_start)}\n\n".encode() + + try: + async for chunk in openai_stream_generator: + if isinstance(chunk, bytes): + chunk = chunk.decode("utf-8", errors="ignore") + + for line in chunk.strip().split("\n"): + line = line.strip() + + if not line or not line.startswith("data:"): + continue + + data_str = line[5:].strip() + if data_str == "[DONE]": + continue + if data_str == "{}": + continue + + try: + data = json.loads(data_str) + except (json.JSONDecodeError, TypeError): + continue + + choices = data.get("choices", []) + if not choices: + # Check for usage in the final chunk + if data.get("usage"): + input_tokens = data["usage"].get("prompt_tokens", input_tokens) + output_tokens = data["usage"].get( + "completion_tokens", output_tokens + ) + continue + + delta = choices[0].get("delta", {}) + finish_reason = choices[0].get("finish_reason") + + # Update usage if present + if data.get("usage"): + input_tokens = data["usage"].get("prompt_tokens", input_tokens) + output_tokens = data["usage"].get( + "completion_tokens", output_tokens + ) + + # --- Handle text content --- + content = delta.get("content") + if content is not None: + if not text_block_open: + # Start a new text content block + block_start = { + "type": "content_block_start", + "index": current_block_index, + "content_block": {"type": "text", "text": ""}, + } + yield f"event: content_block_start\ndata: {json.dumps(block_start)}\n\n".encode() + text_block_open = True + + # Send text delta + block_delta = { + "type": "content_block_delta", + "index": current_block_index, + "delta": {"type": "text_delta", "text": content}, + } + yield f"event: content_block_delta\ndata: {json.dumps(block_delta)}\n\n".encode() + + # --- Handle tool calls --- + tool_calls = delta.get("tool_calls") + if tool_calls: + # Close text block if one is open (text comes before tools) + if text_block_open: + block_stop = { + "type": "content_block_stop", + "index": current_block_index, + } + yield f"event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n".encode() + text_block_open = False + current_block_index += 1 + + for tc in tool_calls: + tc_index = tc.get("index", 0) + + if tc_index not in tool_call_started: + # First time seeing this tool call — emit content_block_start + tool_call_blocks[tc_index] = current_block_index + tool_call_started[tc_index] = True + + # Extract tool call ID and name from the first chunk + tc_id = tc.get("id", f"toolu_{_uuid.uuid4().hex[:24]}") + tc_name = tc.get("function", {}).get("name", "") + + block_start = { + "type": "content_block_start", + "index": current_block_index, + "content_block": { + "type": "tool_use", + "id": tc_id, + "name": tc_name, + "input": {}, + }, + } + yield f"event: content_block_start\ndata: {json.dumps(block_start)}\n\n".encode() + current_block_index += 1 + + # Emit argument chunks as input_json_delta + args_chunk = tc.get("function", {}).get("arguments", "") + if args_chunk: + block_delta = { + "type": "content_block_delta", + "index": tool_call_blocks[tc_index], + "delta": { + "type": "input_json_delta", + "partial_json": args_chunk, + }, + } + yield f"event: content_block_delta\ndata: {json.dumps(block_delta)}\n\n".encode() + + # --- Handle finish reason --- + if finish_reason is not None: + stop_reason_map = { + "stop": "end_turn", + "length": "max_tokens", + "tool_calls": "tool_use", + } + stop_reason = stop_reason_map.get(finish_reason, "end_turn") + + except Exception as e: + log.error(f"Error in Anthropic stream conversion: {e}") + + # Close any open text block + if text_block_open: + block_stop = {"type": "content_block_stop", "index": current_block_index} + yield f"event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n".encode() + + # Close any open tool call blocks + for tc_index, block_index in tool_call_blocks.items(): + block_stop = {"type": "content_block_stop", "index": block_index} + yield f"event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n".encode() + + # Emit message_delta with stop reason + message_delta = { + "type": "message_delta", + "delta": { + "stop_reason": stop_reason, + "stop_sequence": None, + }, + "usage": {"output_tokens": output_tokens}, + } + yield f"event: message_delta\ndata: {json.dumps(message_delta)}\n\n".encode() + + # Emit message_stop + yield f"event: message_stop\ndata: {json.dumps({'type': 'message_stop'})}\n\n".encode() diff --git a/backend/open_webui/utils/auth.py b/backend/open_webui/utils/auth.py index 27af3631e4..a12c6db881 100644 --- a/backend/open_webui/utils/auth.py +++ b/backend/open_webui/utils/auth.py @@ -290,6 +290,10 @@ async def get_current_user( if token is None and "token" in request.cookies: token = request.cookies.get("token") + # Fallback to request.state.token (set by middleware, e.g. for x-api-key) + if token is None and hasattr(request.state, "token") and request.state.token: + token = request.state.token.credentials + if token is None: raise HTTPException(status_code=401, detail="Not authenticated") diff --git a/backend/open_webui/utils/mcp/client.py b/backend/open_webui/utils/mcp/client.py index 33803648be..1004536e4d 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -9,14 +9,27 @@ 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 + """Create an httpx AsyncClient with SSL verification disabled. + + Note: verify=False must be passed at construction time because httpx + configures the SSL context during __init__. Setting client.verify = False + after construction does not affect the underlying transport's SSL context. + """ + kwargs = { + "follow_redirects": True, + "verify": False, + } + if timeout is not None: + kwargs["timeout"] = timeout + if headers is not None: + kwargs["headers"] = headers + if auth is not None: + kwargs["auth"] = auth + return httpx.AsyncClient(**kwargs) class MCPClient: diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index ec7af7733b..218deed17e 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -171,7 +171,10 @@ def get_citation_source_from_tool_result( Returns a list of sources (usually one, but query_knowledge_files may return multiple). """ try: - tool_result = json.loads(tool_result) + try: + tool_result = json.loads(tool_result) + except (json.JSONDecodeError, TypeError): + pass # keep tool_result as-is (e.g. fetch_url returns plain text) if isinstance(tool_result, dict) and "error" in tool_result: return [] @@ -232,6 +235,25 @@ def get_citation_source_from_tool_result( } ] + elif tool_name == "fetch_url": + url = tool_params.get("url", "") + content = tool_result if isinstance(tool_result, str) else str(tool_result) + snippet = content[:500] + ("..." if len(content) > 500 else "") + + return [ + { + "source": {"name": url or "fetch_url", "id": url or "fetch_url"}, + "document": [snippet], + "metadata": [ + { + "source": url, + "name": url, + "url": url, + } + ], + } + ] + elif tool_name == "query_knowledge_files": chunks = tool_result @@ -2056,11 +2078,12 @@ async def process_chat_payload(request, form_data, user, metadata, model): # Folder "Project" handling # Check if the request has chat_id and is inside of a folder + # Uses lightweight column query — only fetches folder_id, not the full chat JSON blob chat_id = metadata.get("chat_id", None) if chat_id and user: - chat = Chats.get_chat_by_id_and_user_id(chat_id, user.id) - if chat and chat.folder_id: - folder = Folders.get_folder_by_id_and_user_id(chat.folder_id, user.id) + folder_id = Chats.get_chat_folder_id(chat_id, user.id) + if folder_id: + folder = Folders.get_folder_by_id_and_user_id(folder_id, user.id) if folder and folder.data: if "system_prompt" in folder.data: @@ -4102,6 +4125,7 @@ async def streaming_chat_response_handler(response, ctx): tool_function_name in [ "search_web", + "fetch_url", "view_knowledge_file", "query_knowledge_files", ] diff --git a/backend/open_webui/utils/misc.py b/backend/open_webui/utils/misc.py index a63d425ff5..447e334227 100644 --- a/backend/open_webui/utils/misc.py +++ b/backend/open_webui/utils/misc.py @@ -91,14 +91,22 @@ def get_message_list(messages_map, message_id): # Reconstruct the chain by following the parentId links message_list = [] + visited_message_ids = set() while current_message: - message_list.insert( - 0, current_message - ) # Insert the message at the beginning of the list + message_id = current_message.get("id") + if message_id in visited_message_ids: + # Cycle detected, break to prevent infinite loop + break + + if message_id is not None: + visited_message_ids.add(message_id) + + message_list.append(current_message) parent_id = current_message.get("parentId") # Use .get() for safety current_message = messages_map.get(parent_id) if parent_id else None + message_list.reverse() return message_list diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index c849eb25a8..b23c5c90a3 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -41,6 +41,7 @@ from open_webui.config import ( ENABLE_OAUTH_ROLE_MANAGEMENT, ENABLE_OAUTH_GROUP_MANAGEMENT, ENABLE_OAUTH_GROUP_CREATION, + OAUTH_GROUP_DEFAULT_SHARE, OAUTH_BLOCKED_GROUPS, OAUTH_GROUPS_SEPARATOR, OAUTH_ROLES_SEPARATOR, @@ -69,6 +70,7 @@ from open_webui.env import ( ENABLE_OAUTH_ID_TOKEN_COOKIE, ENABLE_OAUTH_EMAIL_FALLBACK, OAUTH_CLIENT_INFO_ENCRYPTION_KEY, + OAUTH_MAX_SESSIONS_PER_USER, ) from open_webui.utils.misc import parse_duration from open_webui.utils.auth import get_password_hash, create_token @@ -113,6 +115,7 @@ auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL = OAUTH_MERGE_ACCOUNTS_BY_EMAI auth_manager_config.ENABLE_OAUTH_ROLE_MANAGEMENT = ENABLE_OAUTH_ROLE_MANAGEMENT auth_manager_config.ENABLE_OAUTH_GROUP_MANAGEMENT = ENABLE_OAUTH_GROUP_MANAGEMENT auth_manager_config.ENABLE_OAUTH_GROUP_CREATION = ENABLE_OAUTH_GROUP_CREATION +auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE = OAUTH_GROUP_DEFAULT_SHARE auth_manager_config.OAUTH_BLOCKED_GROUPS = OAUTH_BLOCKED_GROUPS auth_manager_config.OAUTH_ROLES_CLAIM = OAUTH_ROLES_CLAIM auth_manager_config.OAUTH_SUB_CLAIM = OAUTH_SUB_CLAIM @@ -1245,7 +1248,11 @@ class OAuthManager: name=group_name, description=f"Group '{group_name}' created automatically via OAuth.", permissions=default_permissions, # Use default permissions from function args - user_ids=[], # Start with no users, user will be added later by subsequent logic + data={ + "config": { + "share": auth_manager_config.OAUTH_GROUP_DEFAULT_SHARE + } + }, ) # Use determined creator ID (admin or fallback to current user) created_group = Groups.insert_new_group( @@ -1679,11 +1686,18 @@ class OAuthManager: if "expires_in" in token and "expires_at" not in token: token["expires_at"] = datetime.now().timestamp() + token["expires_in"] - # Clean up any existing sessions for this user/provider first + # Enforce max concurrent sessions per user/provider to prevent + # unbounded growth while allowing multi-device usage sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db) - for session in sessions: - if session.provider == provider: - OAuthSessions.delete_session_by_id(session.id, db=db) + provider_sessions = sorted( + [session for session in sessions if session.provider == provider], + key=lambda session: session.created_at, + reverse=True, + ) + # Keep the newest sessions up to the limit, prune the rest + if len(provider_sessions) >= OAUTH_MAX_SESSIONS_PER_USER: + for old_session in provider_sessions[OAUTH_MAX_SESSIONS_PER_USER - 1 :]: + OAuthSessions.delete_session_by_id(old_session.id, db=db) session = OAuthSessions.create_session( user_id=user.id, diff --git a/backend/open_webui/utils/plugin.py b/backend/open_webui/utils/plugin.py index 2dd49fb8ff..a2d7b9ad11 100644 --- a/backend/open_webui/utils/plugin.py +++ b/backend/open_webui/utils/plugin.py @@ -8,7 +8,12 @@ import tempfile import logging from typing import Any -from open_webui.env import PIP_OPTIONS, PIP_PACKAGE_INDEX_OPTIONS, OFFLINE_MODE +from open_webui.env import ( + PIP_OPTIONS, + PIP_PACKAGE_INDEX_OPTIONS, + OFFLINE_MODE, + ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS, +) from open_webui.models.functions import Functions from open_webui.models.tools import Tools @@ -401,6 +406,12 @@ def get_function_module_from_cache(request, function_id, load_from_db=True): def install_frontmatter_requirements(requirements: str): + if not ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS: + log.info( + "ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS is disabled, skipping installation of requirements." + ) + return + if OFFLINE_MODE: log.info("Offline mode enabled, skipping installation of requirements.") return diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index aeba213392..310fa999c7 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -60,6 +60,8 @@ from open_webui.tools.builtin import ( search_memories, add_memory, replace_memory_content, + delete_memory, + list_memories, get_current_timestamp, calculate_timestamp, search_notes, @@ -471,7 +473,15 @@ def get_builtin_tools( # 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]) + builtin_functions.extend( + [ + search_memories, + add_memory, + replace_memory_content, + delete_memory, + list_memories, + ] + ) # Add web search tools if builtin category enabled AND enabled globally AND model has web_search capability if ( diff --git a/src/lib/apis/prompts/index.ts b/src/lib/apis/prompts/index.ts index 1fd311c76f..5db7dc3540 100644 --- a/src/lib/apis/prompts/index.ts +++ b/src/lib/apis/prompts/index.ts @@ -395,6 +395,34 @@ export const setProductionPromptVersion = async ( return res; }; +export const togglePromptById = async (token: string, promptId: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/prompts/id/${promptId}/toggle`, { + method: 'POST', + 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 deletePromptById = async (token: string, promptId: string) => { let error = null; diff --git a/src/lib/components/AddConnectionModal.svelte b/src/lib/components/AddConnectionModal.svelte index a455627e11..b8aa7f5e64 100644 --- a/src/lib/components/AddConnectionModal.svelte +++ b/src/lib/components/AddConnectionModal.svelte @@ -309,8 +309,21 @@ bind:value={url} placeholder={$i18n.t('API Base URL')} autocomplete="off" + list={ollama ? undefined : 'suggestions'} required /> + + {#if !ollama} + + {/if} diff --git a/src/lib/components/admin/Functions.svelte b/src/lib/components/admin/Functions.svelte index 48c1863e74..57dba96af2 100644 --- a/src/lib/components/admin/Functions.svelte +++ b/src/lib/components/admin/Functions.svelte @@ -83,7 +83,7 @@ } const setFilteredItems = () => { - filteredItems = functions + filteredItems = (functions ?? []) .filter( (f) => (selectedType !== '' ? f.type === selectedType : true) && @@ -681,7 +681,8 @@ } toast.success($i18n.t('Functions imported successfully')); - functions.set(await getFunctions(localStorage.token)); + functions = await getFunctionList(localStorage.token); + _functions.set(await getFunctions(localStorage.token)); models.set( await getModels( localStorage.token, @@ -690,6 +691,8 @@ true ) ); + importFiles = null; + functionsImportInputElement.value = ''; }; reader.readAsText(importFiles[0]); diff --git a/src/lib/components/admin/Settings/Connections/OllamaConnection.svelte b/src/lib/components/admin/Settings/Connections/OllamaConnection.svelte index b6d33eb49f..036deadd05 100644 --- a/src/lib/components/admin/Settings/Connections/OllamaConnection.svelte +++ b/src/lib/components/admin/Settings/Connections/OllamaConnection.svelte @@ -3,6 +3,7 @@ const i18n = getContext('i18n'); import Tooltip from '$lib/components/common/Tooltip.svelte'; + import Switch from '$lib/components/common/Switch.svelte'; import SensitiveInput from '$lib/components/common/SensitiveInput.svelte'; import AddConnectionModal from '$lib/components/AddConnectionModal.svelte'; import ConfirmDialog from '$lib/components/common/ConfirmDialog.svelte'; @@ -75,7 +76,7 @@ /> -