Merge branch 'open-webui:dev' into dev

This commit is contained in:
Kevin Rohn 2026-06-08 09:45:59 +02:00 • committed by GitHub
commit 9876b60fa8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
232 changed files with 99164 additions and 7463 deletions

View file

@ -19,4 +19,13 @@ FORWARDED_ALLOW_IPS='*'
# DO NOT TRACK
SCARF_NO_ANALYTICS=true
DO_NOT_TRACK=true
ANONYMIZED_TELEMETRY=false
ANONYMIZED_TELEMETRY=false
# Valkey Vector Store (requires VECTOR_DB=valkey)
# VALKEY_URL='valkey://localhost:6379'
# VALKEY_COLLECTION_PREFIX='open_webui'
# VALKEY_INDEX_TYPE='HNSW'
# VALKEY_DISTANCE_METRIC='COSINE'
# VALKEY_HNSW_M='16'
# VALKEY_HNSW_EF_CONSTRUCTION='200'
# VALKEY_HNSW_EF_RUNTIME='10'

View file

@ -143,6 +143,7 @@ jobs:
merge:
runs-on: ubuntu-latest
needs: [build]
if: ${{ !cancelled() }}
permissions:
contents: read
packages: write
@ -168,16 +169,32 @@ jobs:
echo "FULL_IMAGE_NAME=${REGISTRY}/${GITHUB_REPOSITORY,,}" >> ${GITHUB_ENV}
- name: Download digests
id: download
uses: actions/download-artifact@v5
with:
pattern: digests-${{ matrix.variant.name }}-*
path: /tmp/digests
merge-multiple: true
continue-on-error: true
- name: Check digests
id: check
run: |
count=$(find /tmp/digests -type f 2>/dev/null | wc -l | tr -d ' ')
echo "digest_count=$count" >> $GITHUB_OUTPUT
if [ "$count" -lt 2 ]; then
echo "::warning::${{ matrix.variant.name }}: found $count digest(s), need 2 (one per arch). Skipping merge."
echo "skip=true" >> $GITHUB_OUTPUT
else
echo "skip=false" >> $GITHUB_OUTPUT
fi
- name: Set up Docker Buildx
if: steps.check.outputs.skip != 'true'
uses: docker/setup-buildx-action@v3
- name: Log in to the Container registry
if: steps.check.outputs.skip != 'true'
uses: docker/login-action@v3
with:
registry: ${{ env.REGISTRY }}
@ -185,6 +202,7 @@ jobs:
password: ${{ secrets.GITHUB_TOKEN }}
- name: Extract metadata for Docker images
if: steps.check.outputs.skip != 'true'
id: meta
uses: docker/metadata-action@v5
with:
@ -201,6 +219,7 @@ jobs:
${{ matrix.variant.suffix != '' && format('suffix={0},onlatest=true', matrix.variant.suffix) || '' }}
- name: Create manifest list and push
if: steps.check.outputs.skip != 'true'
working-directory: /tmp/digests
run: |
docker buildx imagetools create \
@ -208,12 +227,13 @@ jobs:
$(printf '${{ env.FULL_IMAGE_NAME }}@sha256:%s ' *)
- name: Inspect image
if: steps.check.outputs.skip != 'true'
run: |
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ steps.meta.outputs.version }}
copy-to-dockerhub:
runs-on: ubuntu-latest
if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')
if: ${{ !cancelled() && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')) }}
needs: [merge]
continue-on-error: true
strategy:

View file

@ -19,6 +19,7 @@ jobs:
timeout-minutes: 10
permissions:
contents: write
actions: write
steps:
- uses: actions/checkout@v5
@ -43,9 +44,13 @@ jobs:
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
gh release create "v${{ steps.pkg.outputs.version }}" \
--title "v${{ steps.pkg.outputs.version }}" \
--notes-file /tmp/release-notes.md
if gh release view "v${{ steps.pkg.outputs.version }}" &>/dev/null; then
echo "Release v${{ steps.pkg.outputs.version }} already exists — skipping creation"
else
gh release create "v${{ steps.pkg.outputs.version }}" \
--title "v${{ steps.pkg.outputs.version }}" \
--notes-file /tmp/release-notes.md
fi
- name: Archive source
uses: actions/upload-artifact@v4
@ -57,7 +62,7 @@ jobs:
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
- name: Dispatch Docker build
- name: Trigger Docker build
uses: actions/github-script@v8
with:
script: |

View file

@ -5,6 +5,140 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [0.9.6] - 2026-06-01
### Added
- 📦 **Official knowledge base sync tool.** A new companion tool from Open WebUI, oikb, keeps a knowledge base in sync with a local directory, GitHub repo, S3 bucket, Confluence space, or any of more than 40 other sources, uploading only new and changed files using the incremental sync support added in this release. [oikb](https://github.com/open-webui/oikb)
- 📂 **Smart directory sync for knowledge bases.** Local directories can now be synced into a knowledge base in one action: file checksums are compared against what's already stored, and only added or modified files are uploaded while removed files and orphaned subdirectories are cleaned up, with the directory structure mirrored automatically and per-file progress shown throughout. [#19190](https://github.com/open-webui/open-webui/issues/19190), [#19394](https://github.com/open-webui/open-webui/issues/19394), [Commit](https://github.com/open-webui/open-webui/commit/60c9db1cb81d021589cb49bee8744a799b51211f), [Commit](https://github.com/open-webui/open-webui/commit/73bdf86766d3f44467cdec436786fe089481baf3), [Commit](https://github.com/open-webui/open-webui/commit/9835b3f1dd7aa9c92ff08a634a0f97cc2a046e42), [Commit](https://github.com/open-webui/open-webui/commit/97252fa609440573251d8e75090f347ed1e51e1d), [Commit](https://github.com/open-webui/open-webui/commit/8f2d346e10c47b57bf6b5a6aa02d453488a88b89), [Commit](https://github.com/open-webui/open-webui/commit/1527eb6e01d441225c979d12d78b7625edf1086f)
- 🗂️ **Knowledge base folders.** Files inside a knowledge base can now be organized into nested folders, with breadcrumb navigation that makes it much easier to manage and find content in large collections. [Commit](https://github.com/open-webui/open-webui/commit/c2cbc47ca76ebfa21e1d36279dbd859283cbbfb1), [Commit](https://github.com/open-webui/open-webui/commit/ab0ee858b73738c5b77d1d455e66e66bc61df1c4), [Commit](https://github.com/open-webui/open-webui/commit/2ad327a4dce64a8eedf5eb41c73b8520344fdd82), [Commit](https://github.com/open-webui/open-webui/commit/171150c1e12b5713f77a5ea0dd19c6581c7bef2e), [Commit](https://github.com/open-webui/open-webui/commit/e7d2ddbb1d14c03a52797751d24a98132aac0cd8), [Commit](https://github.com/open-webui/open-webui/commit/32a417bbf66ed28a5ca8a759ad3a87b2a2aa358a)
- 🧰 **Filesystem tool for knowledge bases.** A new built-in tool, enabled via the "ENABLE_KB_EXEC" environment variable, lets AI models browse and search knowledge base contents using familiar filesystem commands such as 'ls', 'cat', 'grep', 'find', 'head', 'tail', and 'sed', including pipes between them. [Commit](https://github.com/open-webui/open-webui/commit/5b125c24d4eae925d3287efa626595bab29d5c33), [Commit](https://github.com/open-webui/open-webui/commit/ecec86dd32aa396dd5a383e9be5b0f16c4346c48), [Commit](https://github.com/open-webui/open-webui/commit/9ef579ce4be5b82f9281b085a621fed5b04e0d26), [Commit](https://github.com/open-webui/open-webui/commit/2f642754ac6b8c430228a7fc289357a6c6652235), [Commit](https://github.com/open-webui/open-webui/commit/3b00e5721a60518ea317c4fd61a6d0d182961a2f), [Commit](https://github.com/open-webui/open-webui/commit/4e78b355efe0611f27fddef591c69f7c1259a9f6), [Commit](https://github.com/open-webui/open-webui/commit/74f95a9b0d6b89b241701422170d54e21e701baa), [Commit](https://github.com/open-webui/open-webui/commit/cc16e06c32de48bdf464db5814483a49f3527b63), [Commit](https://github.com/open-webui/open-webui/commit/1ea54c3217f487990b6a4c0f0d8ea675a479b6e1)
- ✏️ **File renaming in knowledge bases.** Files inside a knowledge base can now be renamed directly from the workspace, with the new name reflected wherever the file is referenced. [Commit](https://github.com/open-webui/open-webui/commit/3127f1b46255626ea92eb0b598e45078287303f1)
- 😀 **Emoji picker in message input.** A new emoji button in the rich text formatting toolbar lets you browse and insert emojis directly into your messages. [#24704](https://github.com/open-webui/open-webui/pull/24704)
- 🪄 **Per-chat skills toggle.** Skills can now be turned on or off for a conversation directly from the chat Integrations menu, the same way tools and capabilities already work, instead of only through the model preset. [#25036](https://github.com/open-webui/open-webui/issues/25036), [#25037](https://github.com/open-webui/open-webui/pull/25037)
- 🔎 **Access preview for users and groups.** Administrators can now preview exactly which models, knowledge bases, and tools a given user or group can access, making it easier to audit and verify permission setups. [Commit](https://github.com/open-webui/open-webui/commit/9c14740ffb009d550dcbd5d6c599dac57053f112)
- 📄 **Configurable knowledge base file page size.** Administrators can now request a larger page size when listing a knowledge base's files through the API, reducing the number of requests needed to retrieve large collections instead of paging through fixed increments of 30. [#25148](https://github.com/open-webui/open-webui/issues/25148), [Commit](https://github.com/open-webui/open-webui/commit/a4d1b3e9378a61a11dd9822a8dbb525f39753081)
- 🔃 **Persistent processing indicator for knowledge files.** Files still being processed in a knowledge base now keep showing a processing indicator across page reloads, so you can tell what's still ingesting after navigating away and back. [#25031](https://github.com/open-webui/open-webui/issues/25031), [Commit](https://github.com/open-webui/open-webui/commit/ad9f2eeb15a448147d5e47e6a391cabb9aefd0ea)
- 📑 **MinerU file type configuration.** Administrators can now configure which file types are processed by the MinerU document loader, via the new "MINERU_FILE_EXTENSIONS" setting, extending it beyond PDF to formats like DOCX, PPTX, and XLSX. [Commit](https://github.com/open-webui/open-webui/commit/d4030a8aa5d48c2a1cb06c461566844aca2530ab)
- 📃 **Legacy Word document support.** Older ".doc" Word files can now have their text extracted by the default document extraction engine, in addition to the modern ".docx" format. [Commit](https://github.com/open-webui/open-webui/commit/9e3e24e304f8ff210494380b46f614a2984cafb3)
- 📁 **Create subfolders from the folder header.** Chat folders can now have subfolders created directly from the folder header in the chat view, not just from the sidebar. [Commit](https://github.com/open-webui/open-webui/commit/1f0948bcbef2af73b155535ea27762c522260afc)
- ⚡ **Faster initial page loads.** The configuration endpoint that loads on every page visit no longer runs an unnecessary user-count query, making the initial application load lighter on the database, especially on instances with many users. [Commit](https://github.com/open-webui/open-webui/commit/0adc090dcbe636c9c9645d8ec7f6b89ea514870b)
- 🚀 **Faster tool-enabled chat completions.** Chat completions that use multiple tools now start faster because the tools they reference are fetched from the database in a single batch query instead of one query per tool. [#24808](https://github.com/open-webui/open-webui/pull/24808), [Commit](https://github.com/open-webui/open-webui/commit/cc94a90b4d4d690bc7cb9f7124f2d6e552973970)
- 🏎️ **More responsive web search under load.** Web search through SearXNG, Google PSE, Brave, Serper, and Serpstack now uses non-blocking network calls, so the server stays responsive to other users while a search is in flight, and concurrent multi-query searches complete faster. [Commit](https://github.com/open-webui/open-webui/commit/b94245d2ee191e8ef118bf9de1ff7539503bfec9)
- 🐎 **Lighter Ollama backend connections.** Requests to Ollama backends now reuse a shared connection pool instead of opening a fresh session each time, reducing TCP and TLS handshake overhead for installs that poll Ollama frequently or have multiple backends configured. [Commit](https://github.com/open-webui/open-webui/commit/5d9a09a88a9094ebfcd249340be3aaee544b34d0)
- 💽 **Fewer redundant model-list writes.** On multi-instance deployments backed by Redis, the model list is no longer rewritten when it hasn't changed, cutting a major source of redundant writes. [#25469](https://github.com/open-webui/open-webui/issues/25469), [#25474](https://github.com/open-webui/open-webui/pull/25474), [Commit](https://github.com/open-webui/open-webui/commit/fd76b51ab2ad4c4192f2c98153e47888712c2009)
- 📉 **Faster websocket disconnect cleanup.** Disconnecting from a collaborative session no longer triggers a scan across the entire Redis keyspace, using a per-session index instead, which keeps disconnects cheap on large deployments. [#25466](https://github.com/open-webui/open-webui/issues/25466), [Commit](https://github.com/open-webui/open-webui/commit/c7de057a4a54cc80f366ad949417e78fae756d4c)
- 📝 **Frontmatter auto-fill for tools, functions, and skills.** Opening a tool, function, or skill editor now auto-fills the name, id, and description fields from the file's frontmatter, saving you from re-entering metadata already declared in the source. [#24649](https://github.com/open-webui/open-webui/pull/24649), [Commit](https://github.com/open-webui/open-webui/commit/ef975649b26d3e7cd589c49be2fc77cee80fad8f)
- 🪪 **More user placeholders in custom headers.** Custom-header templates for direct connections and tool servers now support "{{USER_EMAIL}}" and "{{USER_ROLE}}" alongside the existing user and session placeholders. [Commit](https://github.com/open-webui/open-webui/commit/ed73ef3d8df988b0e9646b82df5b1a453202ef8d)
- ⏱️ **Configurable MCP connection timeout.** The timeout for the initial handshake with an MCP tool server is now configurable via the new "MCP_INITIALIZE_TIMEOUT" setting, so servers that are slow to start or expose many tools can finish connecting instead of timing out. [#25011](https://github.com/open-webui/open-webui/pull/25011), [Commit](https://github.com/open-webui/open-webui/commit/4297c02b121180e239a61c483ac8477cc557d4ef)
- 📐 **Profile image size limit.** Administrators can now cap the size of inline profile images via the new "PROFILE_IMAGE_MAX_DATA_URI_SIZE" setting, bounding how much database and cache space inline avatars and model icons can consume. [#25468](https://github.com/open-webui/open-webui/issues/25468), [#25476](https://github.com/open-webui/open-webui/pull/25476)
- 🎫 **Wildcard OAuth role mapping.** Administrators can now set "\*" in the allowed OAuth roles to grant the user role to any authenticated OAuth user, instead of having to enumerate every accepted role. [#25062](https://github.com/open-webui/open-webui/pull/25062), [Commit](https://github.com/open-webui/open-webui/commit/07cbc91a8eba3a9a3b39588b2ae5de916930af70)
- 📊 **Paginated feedback history.** The feedback and evaluation history list is now paginated, keeping it responsive for instances that have accumulated large numbers of feedback entries. [Commit](https://github.com/open-webui/open-webui/commit/160a6694e4bd66fc42e6361947516e8e15414ef3)
- 🔘 **Bulk enable or disable automations.** Automations can now be enabled or disabled in bulk from an actions menu on the automations page, instead of toggling each one individually. [Commit](https://github.com/open-webui/open-webui/commit/675e9bee5af8b9fb390fcb235a6c2c4766e78c75)
- ➡️ **Optional auto-redirect to single sign-on.** Administrators can now enable "OAUTH_AUTO_REDIRECT" so that, on deployments with a single sign-on provider and no other login methods, users are sent straight to the provider instead of seeing a login page first. [#25067](https://github.com/open-webui/open-webui/pull/25067), [Commit](https://github.com/open-webui/open-webui/commit/d64ef1803d2f5cedb7b3a308151d08ff2cd2b8e1)
- ☁️ **Azure AI Foundry v1 with Entra ID.** Open WebUI now supports Azure AI Foundry's OpenAI v1 endpoint together with Microsoft Entra ID authentication, so these connections work without manual workarounds. [#24761](https://github.com/open-webui/open-webui/issues/24761), [#24985](https://github.com/open-webui/open-webui/pull/24985), [Commit](https://github.com/open-webui/open-webui/commit/eb4eebc3ce1042cb0d393bf890c1895db6e08b19)
- 🌎 **Linkup web search provider.** Administrators can now select Linkup as the web search provider from the admin settings, with options to configure the API key and search depth. [#24752](https://github.com/open-webui/open-webui/pull/24752), [Commit](https://github.com/open-webui/open-webui/commit/56c0d00e13c74d665124ec9f1cae78e83ca60b6a)
- 🧊 **Valkey vector database support.** Valkey can now be used as the vector database backend, configurable through new "VALKEY_URL" and related settings including index type, distance metric, and HNSW tuning. [#24769](https://github.com/open-webui/open-webui/pull/24769), [Commit](https://github.com/open-webui/open-webui/commit/c0f1aa291938bbef62db8d4013dbc498b0abea15)
- 🔄 **General improvements.** Various improvements were implemented across the application to enhance performance, stability, and security.
- 🌐 **Translation updates.** Translations for Spanish (Spain), Swedish, German, Korean, Catalan, Russian, Irish, Simplified Chinese, Traditional Chinese, Finnish, Polish, Turkish, and Malay were enhanced and expanded.
### Fixed
- 🛡️ **Security Advisory**: This release includes security and access-control fixes. We recommend updating production deployments at your earliest convenience. Not all security fixes in this version may be enumerated in the fixed section — some may be withheld for a short time to give administrators time to upgrade. [Advisories](https://github.com/open-webui/open-webui/security)
- 🛡️ **Tool server permission enforcement.** The per-user permission for inline tool servers is now enforced on chat-completion requests, so users without that permission can no longer bypass the admin setting by supplying tool servers directly in their requests. [Commit](https://github.com/open-webui/open-webui/commit/5cc1eb517094e3507915664817285f1e6e37a16d)
- 🔒 **Knowledge base access check in search tool.** The built-in knowledge search tool now verifies that the caller can access a knowledge base before searching it by id, preventing users from reading the contents of knowledge bases they have not been granted access to. [#25113](https://github.com/open-webui/open-webui/pull/25113)
- 🗄️ **Cross-user access to retrieval collections.** Resolving the documents used for retrieval now verifies the caller's access to each referenced file and rejects client-supplied collection names, preventing a crafted request from pulling another user's files or vector collections into its context. [Commit](https://github.com/open-webui/open-webui/commit/ee47c9c833f8889f3abe99d2266c56e8f8d40230)
- 🔣 **Collection name validation.** Vector collection names are now rejected unless they contain only safe characters, preventing malformed names from reaching the vector store or breaking out of a database query expression. [#24982](https://github.com/open-webui/open-webui/pull/24982)
- 🚫 **Unscoped retrieval collections denied by default.** Retrieval requests for collection names that don't correspond to a known file, memory, web-search, or knowledge base are now denied for non-admins by default, with a new "ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS" setting to restore the previous behavior if needed. [Commit](https://github.com/open-webui/open-webui/commit/c93f071700520a60833b7f9616c292b11f210880)
- 📜 **Prompt history authorization.** Comparing, deleting, and restoring prompt versions now verify the history entry belongs to the prompt you're authorized for, preventing access to or modification of another prompt's version history. [#25056](https://github.com/open-webui/open-webui/pull/25056)
- 🚦 **Code interpreter permission on the legacy path.** The legacy code-execution path now enforces the same permission and capability checks as the current one, so users without the code interpreter permission can no longer trigger code execution through it. [#24724](https://github.com/open-webui/open-webui/pull/24724)
- 🧱 **API key endpoint restriction bypass.** The endpoint allow-list that limits which paths an API key may reach is now matched against the routed request path directly, preventing a crafted request from slipping past the restriction. [#25123](https://github.com/open-webui/open-webui/pull/25123)
- 🚧 **System prompt bypass via request parameter.** The flag that skips a model's configured system prompt can no longer be set by external clients through a request parameter, so admin-configured system prompts can't be bypassed from the API. [#25156](https://github.com/open-webui/open-webui/pull/25156)
- 🚪 **Terminal proxy path traversal.** The terminal proxy now fully decodes request paths before validating them, blocking multi-encoded payloads that could otherwise escape the intended path. [#25157](https://github.com/open-webui/open-webui/pull/25157)
- 🪤 **Cache file path traversal.** The cache file server now requires an exact directory boundary match, closing a gap where a sibling directory whose name began with the cache directory's name could be used to serve files from outside it. [#25086](https://github.com/open-webui/open-webui/pull/25086)
- 🔀 **Ollama backend selection access check.** Requests can no longer target an arbitrary Ollama backend by index; a caller-supplied backend selector is now verified against the backends that actually serve the requested model. [Commit](https://github.com/open-webui/open-webui/commit/7139797be04030b9d1016782f1dbb251c5fe68bb)
- 🔓 **Cross-user file exfiltration via image URLs.** When a chat message references a file by id in an "image_url" field, the server now resolves that file only for its owner, an administrator, or a user with an explicit read grant, preventing other authenticated users from extracting a file's contents by routing it through the model. [#24625](https://github.com/open-webui/open-webui/pull/24625), [Commit](https://github.com/open-webui/open-webui/commit/c75fe8e74b72617c51282cc3ea0a2e8d9cdd9140)
- 📌 **Chat file attachment access checks.** Attaching files to a chat now links only files the caller can read, preventing a user from associating another user's file with their chat to access its contents. [#25054](https://github.com/open-webui/open-webui/pull/25054)
- 🧾 **Model knowledge file ownership checks.** Creating or updating a model now verifies that any knowledge files attached to it are files the editor can access, preventing another user's files from being attached to a model. [#25055](https://github.com/open-webui/open-webui/pull/25055), [Commit](https://github.com/open-webui/open-webui/commit/27fb20c13a4bf8501f7a485abe0654eb5880980d)
- 📅 **Calendar event move authorization.** Updating a calendar event to move it into a different calendar now requires write access on the destination calendar, preventing users from injecting events into calendars they cannot write to. [#24764](https://github.com/open-webui/open-webui/pull/24764)
- 📣 **Channel chat access control.** Generating a response in a channel context now verifies the caller's access to that channel and scopes the included messages, preventing access to channels or messages the user isn't permitted to see. [#24725](https://github.com/open-webui/open-webui/pull/24725)
- 🕸️ **Web loader SSRF gating with Playwright.** When the Playwright-based web loader is in use, page navigations and redirects are now validated the same way as the default loader, closing a gap where the Playwright path could reach internal or otherwise blocked URLs. [#24756](https://github.com/open-webui/open-webui/pull/24756)
- 🛂 **DNS rebinding protection for URL fetches.** The IP address validated for an outbound URL fetch is now the same one used for the actual connection, closing a DNS rebinding window where an attacker-controlled hostname could resolve to a public IP during the safety check and then to a private IP when the connection was opened. [#24759](https://github.com/open-webui/open-webui/pull/24759)
- 🪞 **OAuth profile picture redirect handling.** The OAuth profile picture fetch now follows redirects only when administrators have explicitly allowed it, closing a window where a redirect from an externally validated URL could be used to reach internal addresses. [#24809](https://github.com/open-webui/open-webui/pull/24809)
- 🧼 **Model profile image script injection.** Model profile images are now validated on save and only served inline when they are a known-safe image type, preventing a crafted SVG profile image from running scripts in other users' browsers, while existing legacy images that fail validation are cleared gracefully instead of breaking the model list. [#25060](https://github.com/open-webui/open-webui/pull/25060), [#25173](https://github.com/open-webui/open-webui/pull/25173)
- 🧯 **Diagram rendering script injection.** Mermaid diagrams rendered in chat are now sanitized before display, preventing a crafted diagram from running scripts in the viewer's browser. [#25219](https://github.com/open-webui/open-webui/pull/25219)
- 🔐 **Shared-chat file write protection.** Access to a file through a shared chat now only grants read access, so users who can read a shared chat can no longer modify or delete files attached to it. [#24755](https://github.com/open-webui/open-webui/pull/24755)
- 🔏 **Cross-origin embed prompt control.** When Open WebUI is embedded in an iframe on a different origin, the embedding page can now only drive the chat input or submit prompts if the user has explicitly opted in via the "iframe Sandbox Allow Same Origin" setting, preventing untrusted host pages from triggering confirmation dialogs or controlling the chat. [#24767](https://github.com/open-webui/open-webui/pull/24767), [Commit](https://github.com/open-webui/open-webui/commit/eb3076c1b02d90c2ce6e6d3beb08a37987c740ec)
- 🗂️ **Chat folder ownership checks.** Creating a chat or updating a chat's folder now verifies the referenced folder belongs to the current user, preventing chats from being associated with folders owned by other people. [#24588](https://github.com/open-webui/open-webui/pull/24588)
- 🧩 **Chat recovery from corrupted history.** Chats whose internal message graph was left in a malformed state by a failed regeneration now open and load correctly, with missing roles, parent references, and current-message pointers reconstructed automatically instead of breaking the chat. [#24424](https://github.com/open-webui/open-webui/issues/24424), [#24157](https://github.com/open-webui/open-webui/issues/24157), [#20474](https://github.com/open-webui/open-webui/issues/20474), [#24799](https://github.com/open-webui/open-webui/pull/24799), [Commit](https://github.com/open-webui/open-webui/commit/d310a0777c4c48ec772bdc9a510005d5e91b09c7)
- 📨 **Imported chats with folders appear correctly.** Importing grouped chats no longer leaves them invisible when a referenced folder is missing; such chats now appear in the chat list instead of being silently orphaned. [#24910](https://github.com/open-webui/open-webui/issues/24910), [Commit](https://github.com/open-webui/open-webui/commit/7f7cd210186cb6a67e028c17974eb210fc7ba9fd)
- 🎟️ **MCP tool server sessions stay connected.** OAuth-authenticated MCP tool server sessions are no longer mistakenly refreshed and deleted by the single sign-on session handler, so those connections stay active. [#24618](https://github.com/open-webui/open-webui/issues/24618), [Commit](https://github.com/open-webui/open-webui/commit/c8eb8edca4174ec68fabd071d6b08c0bc07f8117)
- 🤝 **MCP OAuth scope discovery.** The OAuth flow for MCP tool servers now reads the scopes a server advertises through its Protected Resource Metadata, so connecting to servers that declare their own scopes succeeds. [#24730](https://github.com/open-webui/open-webui/issues/24730), [#24690](https://github.com/open-webui/open-webui/pull/24690)
- 🔍 **Web search reliability.** Web search again fetches page content reliably with the default web loader engine, a new "USER_AGENT" environment variable lets administrators set a real browser user-agent so fetches aren't blocked by Cloudflare, Wikipedia, and other bot-detection systems, and the startup script no longer fails to launch when these new environment variables are unset. [#24560](https://github.com/open-webui/open-webui/issues/24560), [#24793](https://github.com/open-webui/open-webui/issues/24793), [#24683](https://github.com/open-webui/open-webui/pull/24683), [Commit](https://github.com/open-webui/open-webui/commit/f60733758272c4532cd032e518c9cd73f648043a)
- 🔥 **Firecrawl web search results.** Web search using Firecrawl now returns results correctly regardless of which response format the Firecrawl version uses. [#24712](https://github.com/open-webui/open-webui/pull/24712)
- 🦅 **Kagi web search.** Web search using Kagi works again after its API endpoint and request method were updated to match Kagi's current API. [#25015](https://github.com/open-webui/open-webui/pull/25015)
- 🔢 **Bracketed numbers in code blocks.** Numbers in square brackets such as "[0]" inside code blocks are no longer stripped out as if they were source citations, so code displays and copies correctly. [#24948](https://github.com/open-webui/open-webui/issues/24948), [Commit](https://github.com/open-webui/open-webui/commit/e90a618f4555cea2c024bb60c1332ca04eed96da)
- 🔌 **API chat completions reliability.** Direct calls to the chat completions API no longer fail with an internal error when no chat session identifier is supplied. [#24553](https://github.com/open-webui/open-webui/issues/24553), [#25235](https://github.com/open-webui/open-webui/issues/25235), [Commit](https://github.com/open-webui/open-webui/commit/bc244fdc90504824b76654880898bf3f6649c299), [Commit](https://github.com/open-webui/open-webui/commit/f16b5c446027eae1bd767617bac2fdf54b24d6fc)
- 🖼️ **ComfyUI image generation and editing.** Generating and editing images via a ComfyUI backend now works again, including when ComfyUI is hosted on a private or internal network where URL validation was previously blocking the admin-configured endpoint. [#24565](https://github.com/open-webui/open-webui/issues/24565), [Commit](https://github.com/open-webui/open-webui/commit/7dcd932ad7cb5feab007737c6e0281def5cd47fc), [Commit](https://github.com/open-webui/open-webui/commit/8aa2a42dc7887512e697ff85d044220945ced40f)
- 🖌️ **Image generation with non-standard response headers.** Image generation now works with backends that return valid JSON without a standard content-type header, instead of rejecting the response. [#24838](https://github.com/open-webui/open-webui/pull/24838)
- 🐘 **Knowledge search on large documents.** Searching knowledge bases on PostgreSQL no longer fails when scanning across documents with very large extracted text content. [#24670](https://github.com/open-webui/open-webui/issues/24670), [Commit](https://github.com/open-webui/open-webui/commit/d74ee34d9128295faa116c919bd1dcca77744975)
- 💬 **Chat title generation.** Automatically generated chat titles now use the model currently selected in the dropdown for the active chat and fall back to the model from the active message branch otherwise, and a clear message is shown if no model is available instead of an unhelpful error. [#24604](https://github.com/open-webui/open-webui/issues/24604), [#24745](https://github.com/open-webui/open-webui/issues/24745), [Commit](https://github.com/open-webui/open-webui/commit/e5c8f8110a88739e011d8ab23e5cb11b4768469f), [Commit](https://github.com/open-webui/open-webui/commit/3c5e7968f0b5130890e17f1f2a2a52297cb145a6)
- 🧮 **Message search and analytics consistency.** Edits, deletions, and branch changes made in a chat are now reflected in message search results and analytics counts instead of leaving stale entries behind. [#25205](https://github.com/open-webui/open-webui/pull/25205), [Commit](https://github.com/open-webui/open-webui/commit/aa06200f789a4c3fc54ea9a141cd347d76cda4c5)
- 🩹 **Graceful handling of in-chat task failures.** When web search query generation, image prompt generation, or a tool call fails or references a missing tool, the chat now falls back or surfaces a clear error instead of breaking partway through the response. [#25038](https://github.com/open-webui/open-webui/issues/25038), [#25144](https://github.com/open-webui/open-webui/issues/25144), [Commit](https://github.com/open-webui/open-webui/commit/b64fd988f02b8a4295193892b24b00bdf635a4dc)
- 🎛️ **Filter changes to message output.** Filter functions that modify a message's structured output after generation now have those changes saved and displayed, instead of being discarded when only the output, not the text content, was changed. [#24884](https://github.com/open-webui/open-webui/pull/24884)
- ⏩ **Titles and tags reflect filtered output.** Outlet filters now run before automatic title, tag, and follow-up generation, so those are based on the final filtered message instead of the unfiltered version. [#24717](https://github.com/open-webui/open-webui/pull/24717)
- 💾 **Action-replaced message content persists.** Message content replaced by an action function through its event emitter is now kept when the chat is saved, instead of reverting to the original after a page reload. [#24585](https://github.com/open-webui/open-webui/issues/24585), [#25485](https://github.com/open-webui/open-webui/pull/25485)
- 🏷️ **Skill mentions in messages.** Mentioning a skill in a message now keeps the skill's name as readable text instead of removing it, and selecting a skill without typing anything no longer causes an error on providers that reject empty messages. [#24929](https://github.com/open-webui/open-webui/issues/24929), [Commit](https://github.com/open-webui/open-webui/commit/01810e32ad51305ca3247e8f83da95f6e6260f0c)
- 🧹 **Usage timer cleanup on send failure.** The background usage-stats timer started during message generation is now always cleared, even when sending a message fails, preventing leaked timers from accumulating over a session. [#25478](https://github.com/open-webui/open-webui/pull/25478)
- 🗑️ **Background tasks stop when a chat is removed.** Deleting or archiving a chat now cancels any in-flight generation or title and tag tasks for it, instead of leaving orphaned background work running. [#25050](https://github.com/open-webui/open-webui/pull/25050), [Commit](https://github.com/open-webui/open-webui/commit/778dba1d6b8dc3962163fa7bf9d802d9a07fba26)
- ⌨️ **Responsive knowledge file search.** Searching for knowledge files in the chat picker and model knowledge selector now matches on file names by default instead of scanning the full extracted text of every document on each keystroke, keeping the search responsive on large deployments, with content search available as an explicit opt-in. [#25082](https://github.com/open-webui/open-webui/issues/25082), [#25119](https://github.com/open-webui/open-webui/pull/25119), [Commit](https://github.com/open-webui/open-webui/commit/591e0aafa1d5e21dcd9fa1279dada2076968b9cb)
- 📥 **Document processing with empty embeddings.** Saving documents to the vector database no longer crashes when an embedding step returns no vectors, allowing the process to continue instead of failing the whole upload. [#25166](https://github.com/open-webui/open-webui/issues/25166)
- 🔤 **Non-UTF-8 text and CSV uploads.** Text and CSV files saved in legacy encodings, including Latin-1, Windows-1252, and Chinese encodings such as GB18030, are now detected and loaded correctly instead of being rejected as binary or failing with an empty-content error. [#25172](https://github.com/open-webui/open-webui/issues/25172), [#24973](https://github.com/open-webui/open-webui/issues/24973), [Commit](https://github.com/open-webui/open-webui/commit/6f0277db52d005420480abb0702d421525d6ea8b), [Commit](https://github.com/open-webui/open-webui/commit/1bbb2b933d4cee70a5102eca4109e77b8625f8fc)
- 🧽 **Null bytes in nested data no longer break saves.** Data containing null bytes nested inside structured fields is now sanitized correctly before being written, preventing database errors that the previous check failed to catch. [#25018](https://github.com/open-webui/open-webui/pull/25018), [Commit](https://github.com/open-webui/open-webui/commit/e3ab4bd212e44c39f439ec4ff5df7c3dbd046895)
- 🧠 **Clear error when no embedding model is configured.** Using knowledge or retrieval features without a loaded embedding model now returns a clear setup error explaining what to configure, instead of failing with a cryptic crash. [Commit](https://github.com/open-webui/open-webui/commit/55ca719bbf76306ed647424c728f730d0bef21f9)
- 🧲 **Memory search quality.** Memory searches now apply the configured embedding query prefix, so retrieval works correctly with embedding models that require one for queries. [#24921](https://github.com/open-webui/open-webui/pull/24921), [Commit](https://github.com/open-webui/open-webui/commit/ce4dca47cb19a6582fd8a550806c89ab297c038d)
- 📚 **Knowledge tool context overflow.** The built-in tool that lists a model's knowledge no longer dumps every file in every knowledge base into the model's context; it now returns summaries by default and paginates file listings only for a requested knowledge base. [#25105](https://github.com/open-webui/open-webui/pull/25105), [Commit](https://github.com/open-webui/open-webui/commit/0e73f7af099e1f9437c55a5c2d16741c97dbd57f)
- ⏳ **Terminal session stability.** The terminal proxy no longer hangs when one direction of the connection closes before the other, so terminal sessions shut down cleanly instead of stalling. [#25464](https://github.com/open-webui/open-webui/issues/25464), [#25479](https://github.com/open-webui/open-webui/pull/25479)
- 🧷 **Tool call continuity with strict providers.** Chats that contain incomplete tool calls or orphaned tool results no longer fail to continue when sent to providers that strictly validate tool pairings, such as Anthropic and AWS Bedrock Converse. [#24758](https://github.com/open-webui/open-webui/issues/24758), [#24940](https://github.com/open-webui/open-webui/issues/24940), [#24798](https://github.com/open-webui/open-webui/pull/24798), [Commit](https://github.com/open-webui/open-webui/commit/cfa6908d579e1f7f202a321289029996130b8411)
- 🛑 **Stream termination for pipe functions.** Streamed responses from pipe functions now always send the standard end-of-stream marker, so chat clients and external integrations reliably detect when a response is complete instead of waiting on streams that already finished. [#24763](https://github.com/open-webui/open-webui/pull/24763)
- 🔊 **Non-blocking text-to-speech transcoding.** Converting text-to-speech audio to MP3 no longer blocks the server's event loop, so other requests stay responsive even while a TTS response is being transcoded. [#24876](https://github.com/open-webui/open-webui/pull/24876)
- 🎚️ **Default text-to-speech voice.** Text-to-speech requests now honor the voice specified in the request and fall back to the configured default only when none is given, instead of always using the admin default or failing. [#15143](https://github.com/open-webui/open-webui/issues/15143), [#25035](https://github.com/open-webui/open-webui/issues/25035), [Commit](https://github.com/open-webui/open-webui/commit/f16b5c446027eae1bd767617bac2fdf54b24d6fc), [Commit](https://github.com/open-webui/open-webui/commit/750604a11d4adcb5ae568b9fc93010031ad394f9)
- 🪝 **Reliable knowledge base file linking.** Files uploaded to a knowledge collection are now linked on the server as part of the upload itself, so they remain attached to the collection even if you navigate away or close the page before processing finishes. [#24807](https://github.com/open-webui/open-webui/issues/24807), [Commit](https://github.com/open-webui/open-webui/commit/d0b17f056911ec73df2b686a0d92bb789391db27)
- ☁️ **Azure connections on custom hostnames.** Connections marked as the Azure provider now use the Azure code path even when the endpoint does not contain "azure" in its hostname, fixing custom Azure deployments served from non-standard domains. [#24882](https://github.com/open-webui/open-webui/pull/24882), [Commit](https://github.com/open-webui/open-webui/commit/c8f851bd2de3d127f1e818fc116c27174f96d3f1)
- 🗓️ **Clearing calendar event fields.** Removing the description or location from a calendar event now saves correctly instead of silently keeping the previous value. [#25026](https://github.com/open-webui/open-webui/issues/25026), [Commit](https://github.com/open-webui/open-webui/commit/91810f1c4e93d4f559bdf61cd085cadc288a732d), [Commit](https://github.com/open-webui/open-webui/commit/78b1637a035d71099262412e5dee3e4d65c7fb2f)
- 💭 **Advanced parameter settings.** Custom reasoning tags and custom model parameters are now saved correctly instead of being dropped, and the presence penalty and repeat penalty no longer save the frequency penalty's value instead of their own. [#25183](https://github.com/open-webui/open-webui/pull/25183), [#25200](https://github.com/open-webui/open-webui/pull/25200), [#25204](https://github.com/open-webui/open-webui/pull/25204)
- 📏 **Long username display.** Long usernames no longer overflow their containers in the admin user list, user modals, and sidebar. [#25185](https://github.com/open-webui/open-webui/pull/25185)
- 🎯 **All skills selectable in the model editor.** The model editor's skills selector now lists every skill you have access to, with a search box for large lists, instead of showing only the first 30 with no way to reach the rest. [#24873](https://github.com/open-webui/open-webui/issues/24873), [Commit](https://github.com/open-webui/open-webui/commit/936d5f2676dfbd2ba763b30af33a8557dcda9b05)
- 🔔 **Accurate knowledge upload feedback.** Dragging files into a knowledge base no longer shows an upload notification before the upload has actually been processed. [#25484](https://github.com/open-webui/open-webui/pull/25484)
- ♿ **High-contrast timestamp readability.** The user message timestamp now uses the correct colors in high-contrast mode instead of inverted ones, keeping it readable. [#25461](https://github.com/open-webui/open-webui/pull/25461)
- ♿ **Keyboard and screen reader access to menus.** The integrations, more-options, and user menus are now real buttons with labels and keyboard support, so they can be opened with the keyboard and announced by screen readers. [Commit](https://github.com/open-webui/open-webui/commit/346dab3d8f909fc321a49ea2be633ea5c4c4a349)
- 🖱️ **Focus-loss handling in editors.** Workspace and admin editors for models, tools, functions, and skills again respond correctly when the browser window loses focus, after the wrong event name was being listened for. [#25459](https://github.com/open-webui/open-webui/pull/25459)
- 🛟 **Resilience to corrupted local storage.** Corrupted data in the browser's local storage no longer crashes the interface; affected settings and dismissed-banner state now fall back to safe defaults. [#25481](https://github.com/open-webui/open-webui/pull/25481)
- 📶 **Quieter reconnection notifications.** Brief connection interruptions, such as backgrounding a mobile tab, no longer flash a "connection lost" warning, and the "reconnected" message only appears if a disconnect was actually shown. [Commit](https://github.com/open-webui/open-webui/commit/77c8c54b1ea07189e21313f9cd401f3b956bb93a)
- 🍎 **Safari PDF handling.** PDF processing now works in Safari, which doesn't support the stream iteration the previous code relied on. [#25151](https://github.com/open-webui/open-webui/issues/25151), [#25473](https://github.com/open-webui/open-webui/pull/25473)
- 🎙️ **Voice mode mute shortcut listing.** The keyboard shortcut for muting voice mode now appears in the keyboard shortcuts help modal. [#25193](https://github.com/open-webui/open-webui/pull/25193)
- 📎 **Document attachments in channel model replies.** Tagging a model in a channel thread now forwards uploaded non-image documents such as PDFs and DOCX files into the model's context, so document summarization and comparison workflows that already worked in direct chat now work in channels too. [#24896](https://github.com/open-webui/open-webui/issues/24896), [#24898](https://github.com/open-webui/open-webui/pull/24898), [Commit](https://github.com/open-webui/open-webui/commit/7e9d41d664d7065a92ac93606d18e88c45eb36b6)
- 🙈 **Hidden models in channel mentions.** Models marked as hidden no longer appear in the channel message-input model mention selector, matching how hidden models are excluded elsewhere in the interface. [#24892](https://github.com/open-webui/open-webui/pull/24892)
- 🧵 **Channel thread and pinned message stability.** Opening a channel thread or the pinned messages view no longer fails to render when a message or its data is missing. [#25209](https://github.com/open-webui/open-webui/pull/25209)
- 📺 **YouTube short link transcripts.** Pasting a "youtu.be" short link into a chat now loads the video transcript correctly instead of failing with an empty-content error. [#24856](https://github.com/open-webui/open-webui/issues/24856), [Commit](https://github.com/open-webui/open-webui/commit/1e36a206008c3abce5ccdc98d5750e30a7345b98)
- 🙉 **Hidden models in default-model and automation pickers.** The admin pickers for default models and default pinned models, and the automation model dropdown, now filter out hidden models, consistent with how hidden models are treated elsewhere. [#24869](https://github.com/open-webui/open-webui/issues/24869), [Commit](https://github.com/open-webui/open-webui/commit/1fa3050f069a72de1daaaa29233b545c4deaea51), [Commit](https://github.com/open-webui/open-webui/commit/4705c2d98812c1189c2ae4960399bc8888d4861c)
- 🔊 **Speech-to-text SSL setting honored.** Speech-to-text requests now respect the "AIOHTTP_CLIENT_SESSION_SSL" setting, so administrators using self-signed certificates or custom SSL configurations can use STT engines that were previously failing TLS verification. [#24568](https://github.com/open-webui/open-webui/issues/24568), [#24857](https://github.com/open-webui/open-webui/pull/24857), [Commit](https://github.com/open-webui/open-webui/commit/2ca91ceeeca6a6a21e5a6ad68ea3ab8c9d9f6deb), [Commit](https://github.com/open-webui/open-webui/commit/94b66b17972e0ad77954df4b81a6bb86a7a6a04b)
- 🔗 **Placeholders in MCP connection headers.** Custom header templates configured on MCP server connections now have their "{{USER_ID}}", "{{USER_NAME}}", "{{USER_EMAIL}}", "{{USER_ROLE}}", "{{CHAT_ID}}", and "{{MESSAGE_ID}}" placeholders interpolated at request time, matching how custom headers already work for direct connections and tool servers. [#24822](https://github.com/open-webui/open-webui/pull/24822)
- 🪟 **Bing search CLI smoke test.** Running the Bing web-search module from the command line for a quick connectivity check no longer raises an error about missing arguments. [#24765](https://github.com/open-webui/open-webui/issues/24765), [#24768](https://github.com/open-webui/open-webui/pull/24768)
- 🩺 **Database health check recovery.** After a transient database connection error, the health check endpoint now recovers automatically instead of staying permanently broken on the affected worker. [Commit](https://github.com/open-webui/open-webui/commit/0b81520e072bb4ed15532b8e153b43dd3243feaf)
- 🥾 **Startup on non-Unicode consoles.** Open WebUI no longer crashes at startup when the console can't encode the banner's box-drawing characters, such as on Windows or with redirected or headless output, falling back to a plain-text banner instead. [#24965](https://github.com/open-webui/open-webui/issues/24965), [#25482](https://github.com/open-webui/open-webui/pull/25482)
- 🆕 **First admin signup after a reset.** Creating the first administrator account is no longer blocked by a previously stored signup setting, so a fresh or reset instance can always be bootstrapped. [#24821](https://github.com/open-webui/open-webui/pull/24821)
- 🪵 **JSON exception logging.** With JSON log formatting enabled, exceptions are now recorded correctly with a structured type, message, and stacktrace instead of being dropped, and a logging failure can no longer crash the application. [#25135](https://github.com/open-webui/open-webui/issues/25135), [Commit](https://github.com/open-webui/open-webui/commit/79bf3d28d88e78b0136ef3c3c9f8e7bb85d3cea9)
- 🧭 **Workspace skills permission.** Users granted only the "workspace.skills" permission can now see the workspace entry in the sidebar and are correctly routed to the skills page from the workspace index. [#24729](https://github.com/open-webui/open-webui/pull/24729)
- 🔁 **Resilient database migrations.** Database migrations now skip tables, indexes, and columns that already exist and add missing primary keys to legacy tables, so upgrades succeed even when parts of the schema were manually or partially created beforehand. [Commit](https://github.com/open-webui/open-webui/commit/81f611fb73c726cfcfc20ee57a926a543f07e95f), [Commit](https://github.com/open-webui/open-webui/commit/f0e88dadc8502ea05b1a00f3155bee7d1cf32249), [Commit](https://github.com/open-webui/open-webui/commit/bd9f82d5a681ee94bde44325033447500a0c76bf), [Commit](https://github.com/open-webui/open-webui/commit/459b1c3fda2ec3579fbd2ab408d0fdeb07c96b99), [Commit](https://github.com/open-webui/open-webui/commit/6df09a4039d181f4e2d41324e93cd36c6fb27dfd), [Commit](https://github.com/open-webui/open-webui/commit/95840e307a66429168775de389b329487a4311c4), [Commit](https://github.com/open-webui/open-webui/commit/6b1df94bf933af92f5c7d093a2d92e2e50d6fdd0), [Commit](https://github.com/open-webui/open-webui/commit/98d3b2308564e2112290773ea240b90419ebb48d), [Commit](https://github.com/open-webui/open-webui/commit/1b9d22e324181b96511a15cbacc40fdbb439ad79), [Commit](https://github.com/open-webui/open-webui/commit/dc0f8ae6f2372b6e5dd42f3520813b09c54146c0), [Commit](https://github.com/open-webui/open-webui/commit/ee3b14233a642ff1bc38aef03a2418afcf5d6c1d), [Commit](https://github.com/open-webui/open-webui/commit/1004dad2749bcc369d511315a808d4bb5277854c), [Commit](https://github.com/open-webui/open-webui/commit/db2b3d7fd86cc7d76424671ead3c62a3f6302fe7), [Commit](https://github.com/open-webui/open-webui/commit/2e1b671e8db49f47d69fd59ece2b58e109ec8b84), [Commit](https://github.com/open-webui/open-webui/commit/9a8969ca93f8c9109d956518a085be3e2afacc07), [Commit](https://github.com/open-webui/open-webui/commit/9717ada92fdb3761a217d309f322ee9fbe418d0d), [Commit](https://github.com/open-webui/open-webui/commit/d7cfc1e46a8f3e5c8f6758c12ef641f41fb51a39), [Commit](https://github.com/open-webui/open-webui/commit/9263b7568eaa83f43e4a118c98490c4f7aaba570), [Commit](https://github.com/open-webui/open-webui/commit/73d2065227e651cf90a2cfade721516815743ab4), [#24722](https://github.com/open-webui/open-webui/pull/24722)
### Changed
- ⚠️ **Database Migrations**: This release includes database schema changes; we strongly recommend backing up your database and all associated data before upgrading in production environments. If you are running a multi-worker, multi-server, or load-balanced deployment, all instances must be updated simultaneously, rolling updates are not supported and will cause application failures due to schema incompatibility.
- ⚙️ **Tool-call iteration cap renamed and raised.** The environment variable that limits how many tool calls a single chat response may make is now "CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS", with its default raised from 30 to 256 and a new "-1" value for unlimited; the previous "CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES" name continues to work as a fallback, and chats that hit the cap now show a clear error in-chat instead of stopping silently. [#24918](https://github.com/open-webui/open-webui/pull/24918), [Commit](https://github.com/open-webui/open-webui/commit/2b99945d2726bdf5aed7b1712f9e3b7b622671df)
- 🔐 **Reduced public "/api/config" exposure.** The "/api/config" response no longer includes several feature flags ("enable_api_keys", "enable_password_change_form", "enable_version_update_check", "enable_public_active_users_count", "enable_easter_eggs") for unauthenticated callers, reducing information disclosure to anonymous visitors. [Commit](https://github.com/open-webui/open-webui/commit/245e0ee029e9e10617f62953a9e6f67dd00ecf81), [Commit](https://github.com/open-webui/open-webui/commit/ae06e199d5d3f65296a978cc35079bdacba596d2)
- 🔑 **"WEBUI_SECRET_KEY" is now a hard requirement even for unsupported deployments.** Deployments that start the backend in an explicitly unsupported way (such as invoking uvicorn directly) without setting "WEBUI_SECRET_KEY" will now refuse to start instead of falling back to an empty key; the supported start methods (start.sh, start_windows.bat, and "open-webui serve") still set or auto-generate it automatically, so standard deployments are unaffected. Direct Uvicorn startup is not supported. [#25218](https://github.com/open-webui/open-webui/pull/25218)
## [0.9.5] - 2026-05-09
### Added

View file

@ -199,9 +199,7 @@ if frontend_loader.exists():
logging.error(f'An error occurred: {e}')
####################################
# STORAGE PROVIDER
####################################
# --- Storage Provider ---
STORAGE_PROVIDER = os.getenv('STORAGE_PROVIDER', 'local') # defaults to local, s3
STORAGE_LOCAL_CACHE = os.getenv('STORAGE_LOCAL_CACHE', 'true').lower() == 'true'
@ -948,6 +946,15 @@ log.info(f'VECTOR_DB: {VECTOR_DB}')
S3_VECTOR_BUCKET_NAME = os.getenv('S3_VECTOR_BUCKET_NAME', None)
S3_VECTOR_REGION = os.getenv('S3_VECTOR_REGION', None)
# Valkey Vector Store
VALKEY_URL = os.getenv('VALKEY_URL', '')
VALKEY_COLLECTION_PREFIX = os.getenv('VALKEY_COLLECTION_PREFIX', 'open_webui')
VALKEY_INDEX_TYPE = os.getenv('VALKEY_INDEX_TYPE', 'HNSW').upper()
VALKEY_DISTANCE_METRIC = os.getenv('VALKEY_DISTANCE_METRIC', 'COSINE').upper()
VALKEY_HNSW_M = int(os.getenv('VALKEY_HNSW_M', '16'))
VALKEY_HNSW_EF_CONSTRUCTION = int(os.getenv('VALKEY_HNSW_EF_CONSTRUCTION', '200'))
VALKEY_HNSW_EF_RUNTIME = int(os.getenv('VALKEY_HNSW_EF_RUNTIME', '10'))
####################################
# Information Retrieval (RAG)
####################################
@ -1111,6 +1118,12 @@ MINERU_PARAMS = ConfigVar(
mineru_params,
)
MINERU_FILE_EXTENSIONS = ConfigVar(
'MINERU_FILE_EXTENSIONS',
'rag.mineru_file_extensions',
[ext.strip() for ext in os.getenv('MINERU_FILE_EXTENSIONS', 'pdf').split(',') if ext.strip()],
)
EXTERNAL_DOCUMENT_LOADER_URL = ConfigVar(
'EXTERNAL_DOCUMENT_LOADER_URL',
'rag.external_document_loader_url',
@ -1914,6 +1927,24 @@ YOUCOM_API_KEY = ConfigVar(
os.getenv('YOUCOM_API_KEY', ''),
)
LINKUP_API_KEY = ConfigVar(
'LINKUP_API_KEY',
'rag.web.search.linkup_api_key',
os.getenv('LINKUP_API_KEY', ''),
)
linkup_search_params = os.getenv('LINKUP_SEARCH_PARAMS', '')
try:
linkup_search_params = json.loads(linkup_search_params)
except json.JSONDecodeError:
linkup_search_params = {}
LINKUP_SEARCH_PARAMS = ConfigVar(
'LINKUP_SEARCH_PARAMS',
'rag.web.search.linkup_search_params',
linkup_search_params,
)
####################################
# Images
####################################
@ -3430,6 +3461,12 @@ ENABLE_OAUTH_SIGNUP = ConfigVar(
os.getenv('ENABLE_OAUTH_SIGNUP', 'False').lower() == 'true',
)
OAUTH_AUTO_REDIRECT = ConfigVar(
'OAUTH_AUTO_REDIRECT',
'oauth.auto_redirect',
os.getenv('OAUTH_AUTO_REDIRECT', 'False').lower() == 'true',
)
OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE = ConfigVar(
'OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE',
'oauth.refresh_token_include_scope',

View file

@ -16,8 +16,6 @@ import markdown
from bs4 import BeautifulSoup
from cryptography.hazmat.primitives import serialization
from open_webui.constants import ERROR_MESSAGES
####################################
# Load .env file
####################################
@ -547,6 +545,15 @@ else:
except Exception:
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER = AIOHTTP_CLIENT_TIMEOUT
# Timeout (in seconds) for the MCP session.initialize() handshake.
# The handshake performs a list-tools round-trip and can take tens of
# seconds on cold-start servers or servers exposing many tools.
MCP_INITIALIZE_TIMEOUT = os.getenv('MCP_INITIALIZE_TIMEOUT', '10')
try:
MCP_INITIALIZE_TIMEOUT = int(MCP_INITIALIZE_TIMEOUT)
except (ValueError, TypeError):
MCP_INITIALIZE_TIMEOUT = 10
####################################
# AIOHTTP Connection Pool
@ -603,9 +610,10 @@ ENABLE_SIGNUP_PASSWORD_CONFIRMATION = os.getenv('ENABLE_SIGNUP_PASSWORD_CONFIRMA
####################################
# WEBUI_JWT_SECRET_KEY is deprecated; use WEBUI_SECRET_KEY instead.
# No hardcoded fallback by design: the supported start scripts set/auto-generate it; unset is rejected below.
WEBUI_SECRET_KEY = os.getenv(
'WEBUI_SECRET_KEY',
os.getenv('WEBUI_JWT_SECRET_KEY', 't0p-s3cr3t'),
os.getenv('WEBUI_JWT_SECRET_KEY', ''),
)
WEBUI_SESSION_COOKIE_SAME_SITE = os.getenv('WEBUI_SESSION_COOKIE_SAME_SITE', 'lax')
@ -620,7 +628,14 @@ WEBUI_AUTH_COOKIE_SECURE = (
)
if WEBUI_AUTH and WEBUI_SECRET_KEY == '':
raise ValueError(ERROR_MESSAGES.ENV_VAR_NOT_FOUND)
raise SystemExit(
'WEBUI_SECRET_KEY is not set. It is a hard requirement when authentication is enabled.\n'
'The supported start methods set or auto-generate it for you: use start.sh (Linux/macOS), '
'start_windows.bat (Windows), or `open-webui serve`.\n'
'If you start the backend another way (e.g. invoking uvicorn directly, which is unsupported), '
'you must set WEBUI_SECRET_KEY yourself to a long random value.\n'
'See https://docs.openwebui.com/reference/env-configuration#webui_secret_key'
)
ENABLE_COMPRESSION_MIDDLEWARE = os.getenv('ENABLE_COMPRESSION_MIDDLEWARE', 'True').lower() == 'true'
@ -664,6 +679,13 @@ PASSWORD_VALIDATION_HINT = os.getenv('PASSWORD_VALIDATION_HINT', '')
BYPASS_MODEL_ACCESS_CONTROL = os.getenv('BYPASS_MODEL_ACCESS_CONTROL', 'False').lower() == 'true'
BYPASS_RETRIEVAL_ACCESS_CONTROL = os.getenv('BYPASS_RETRIEVAL_ACCESS_CONTROL', 'False').lower() == 'true'
# When True, collection names that do not match any known file-*, user-memory-*,
# web-search-*, or knowledge-base collection are allowed through access control
# for non-admin users. When False (default), unknown collection names are
# denied — closing the legacy unscoped namespace.
ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS = os.getenv('ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS', 'False').lower() == 'true'
# When enabled, skips pydub-based preprocessing (format conversion, compression,
# and chunked splitting) before sending files to processing engines. Useful when
@ -773,6 +795,11 @@ PROFILE_IMAGE_ALLOWED_MIME_TYPES = frozenset(
if t.strip()
)
# Max stored length (bytes) of a data:image profile URI; bounds Postgres/Redis
# bloat from inline avatars and model icons. Unset (default) disables the cap.
_profile_image_max_data_uri_size = os.getenv('PROFILE_IMAGE_MAX_DATA_URI_SIZE', '').strip()
PROFILE_IMAGE_MAX_DATA_URI_SIZE = int(_profile_image_max_data_uri_size) if _profile_image_max_data_uri_size else None
####################################
# Forward Headers
####################################
@ -786,6 +813,17 @@ FORWARD_USER_INFO_HEADER_USER_ROLE = os.getenv('FORWARD_USER_INFO_HEADER_USER_RO
FORWARD_SESSION_INFO_HEADER_MESSAGE_ID = os.getenv('FORWARD_SESSION_INFO_HEADER_MESSAGE_ID', 'X-OpenWebUI-Message-Id')
FORWARD_SESSION_INFO_HEADER_CHAT_ID = os.getenv('FORWARD_SESSION_INFO_HEADER_CHAT_ID', 'X-OpenWebUI-Chat-Id')
# If set while ENABLE_FORWARD_USER_INFO_HEADERS is True, send one signed HS256 JWT
# (FORWARD_USER_INFO_HEADER_JWT) instead of separate X-OpenWebUI-User-* headers.
FORWARD_USER_INFO_HEADER_JWT_SECRET = (os.environ.get('FORWARD_USER_INFO_HEADER_JWT_SECRET') or '').strip() or None
FORWARD_USER_INFO_HEADER_JWT = os.environ.get('FORWARD_USER_INFO_HEADER_JWT', 'X-OpenWebUI-User-Jwt')
try:
FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS = int(
os.environ.get('FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS', '300')
)
except ValueError:
FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS = 300
####################################
# Progressive Web App
####################################
@ -845,15 +883,24 @@ else:
CHAT_RESPONSE_STREAM_DELTA_CHUNK_SIZE = 1
CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES = os.getenv('CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES', '30')
# Maximum tool-call iterations per chat response. Set to -1 for unlimited.
# The old CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES name is accepted as a fallback.
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS = os.getenv(
'CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS',
os.getenv('CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES', '256'),
)
if CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES == '':
CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES = 30
if CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS == '':
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS = 256
else:
try:
CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES = int(CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES)
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS = int(CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS)
except Exception:
CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES = 30
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS = 256
# -1 means unlimited (no cap).
if CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS == -1:
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS = None
# WARNING: Experimental. Only enable if your upstream Responses API endpoint

View file

@ -318,11 +318,10 @@ async def generate_function_chat_completion(request, form_data, user, models: di
async for line in res:
yield process_line(form_data, line)
if isinstance(res, str) or isinstance(res, Generator):
finish_message = openai_chat_chunk_message_template(form_data['model'], '')
finish_message['choices'][0]['finish_reason'] = 'stop'
yield f'data: {json.dumps(finish_message)}\n\n'
yield 'data: [DONE]'
finish_message = openai_chat_chunk_message_template(form_data['model'], '')
finish_message['choices'][0]['finish_reason'] = 'stop'
yield f'data: {json.dumps(finish_message)}\n\n'
yield 'data: [DONE]'
return StreamingResponse(stream_content(), media_type='text/event-stream')
else:

View file

@ -193,6 +193,13 @@ class AppConfig:
super().__setattr__('_entries', {})
super().__setattr__('_key_prefix', redis_key_prefix)
# If sentinels weren't explicitly provided, read from env.
if redis_sentinels is None:
from open_webui.env import REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT
from open_webui.utils.redis import get_sentinels_from_env
redis_sentinels = get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT)
rc: Union[redis.Redis, redis.cluster.RedisCluster, None] = None
if redis_url:
rc = get_redis_connection(redis_url, redis_sentinels or [], redis_cluster, decode_responses=True)

View file

@ -117,26 +117,25 @@ extract_ssl_mode_from_url = extract_ssl_params_from_url
reattach_ssl_mode_to_url = reattach_ssl_params_to_url
class JSONField(types.TypeDecorator):
impl = types.Text
class JSONField(types.TypeDecorator): # TEXT-backed JSON storage
"""Store arbitrary Python objects as JSON-encoded TEXT.
Used instead of native JSON columns for portability across SQLite and
PostgreSQL. Values are serialized with ``json.dumps`` on write and
deserialized with ``json.loads`` on read.
"""
impl = types.UnicodeText
cache_ok = True
def process_bind_param(self, value: _T | None, dialect: Dialect) -> Any:
return json.dumps(value)
return json.dumps(value) if value is not None else None
def process_result_value(self, value: _T | None, dialect: Dialect) -> Any:
if value is not None:
return json.loads(value)
return json.loads(value) if value is not None else None
def copy(self, **kw: Any) -> Self:
return JSONField(self.impl.length)
def db_value(self, value):
return json.dumps(value)
def python_value(self, value):
if value is not None:
return json.loads(value)
def copy(self, **kwargs: Any) -> Self:
return JSONField(length=self.impl.length)
# Normalize SSL params from the URL once; the sync engine needs them

View file

@ -281,6 +281,7 @@ from open_webui.config import (
MINERU_API_MODE,
MINERU_API_TIMEOUT,
MINERU_API_URL,
MINERU_FILE_EXTENSIONS,
MINERU_PARAMS,
MISTRAL_OCR_API_BASE_URL,
MISTRAL_OCR_API_KEY,
@ -288,6 +289,7 @@ from open_webui.config import (
MOJEEK_SEARCH_API_KEY,
OAUTH_ADMIN_ROLES,
OAUTH_ALLOWED_ROLES,
OAUTH_AUTO_REDIRECT,
OAUTH_EMAIL_CLAIM,
OAUTH_PICTURE_CLAIM,
OAUTH_PROVIDERS,
@ -452,6 +454,8 @@ from open_webui.config import (
YANDEX_WEB_SEARCH_CONFIG,
YANDEX_WEB_SEARCH_URL,
YOUCOM_API_KEY,
LINKUP_API_KEY,
LINKUP_SEARCH_PARAMS,
YOUTUBE_LOADER_LANGUAGE,
YOUTUBE_LOADER_PROXY_URL,
AppConfig,
@ -487,10 +491,9 @@ from open_webui.env import (
LICENSE_KEY,
LOG_FORMAT,
MAX_BODY_LOG_SIZE,
# Redis
REDIS_CLUSTER,
REDIS_KEY_PREFIX,
REDIS_SENTINEL_HOSTS,
REDIS_SENTINEL_PORT,
REDIS_URL,
RESET_CONFIG_ON_START,
SAFE_MODE,
@ -509,8 +512,11 @@ from open_webui.env import (
WEBUI_SESSION_COOKIE_SECURE,
)
from open_webui.internal.db import ScopedSession, engine, get_async_session
from open_webui.models.access_grants import AccessGrants
from open_webui.models.channels import Channels
from open_webui.models.chats import ChatForm, Chats
from open_webui.models.functions import Functions
from open_webui.models.messages import Messages
from open_webui.models.models import Models
from open_webui.models.users import UserModel, Users
from open_webui.routers import (
@ -619,7 +625,7 @@ from open_webui.utils.oauth import (
resolve_oauth_client_info,
)
from open_webui.utils.plugin import install_tool_and_function_dependencies
from open_webui.utils.redis import get_redis_connection, get_sentinels_from_env
from open_webui.utils.redis import get_redis_client, get_redis_connection
from open_webui.utils.security_headers import SecurityHeadersMiddleware
from open_webui.utils.session_pool import get_session
from open_webui.utils.tools import set_terminal_servers, set_tool_servers
@ -648,7 +654,7 @@ class SPAStaticFiles(StaticFiles):
if LOG_FORMAT != 'json':
print(rf"""
banner = rf"""
██████╗ ██████╗ ███████╗███╗ ██╗ ██╗ ██╗███████╗██████╗ ██╗ ██╗██╗
██╔═══██╗██╔══██╗██╔════╝████╗ ██║ ██║ ██║██╔════╝██╔══██╗██║ ██║██║
██║ ██║██████╔╝█████╗ ██╔██╗ ██║ ██║ █╗ ██║█████╗ ██████╔╝██║ ██║██║
@ -660,7 +666,12 @@ if LOG_FORMAT != 'json':
v{VERSION} - building the best AI user interface.
{f'Commit: {WEBUI_BUILD_HASH}' if WEBUI_BUILD_HASH != 'dev-build' else ''}
https://github.com/open-webui/open-webui
""")
"""
try:
print(banner)
except UnicodeEncodeError:
# Stdout can't encode the box-drawing banner (Windows cp1252, redirected/headless stdout); fall back to ASCII.
print(f'Open WebUI v{VERSION} - building the best AI user interface.\nhttps://github.com/open-webui/open-webui')
@asynccontextmanager
@ -692,12 +703,7 @@ async def lifespan(app: FastAPI):
log.info('Installing external dependencies of functions and tools...')
await install_tool_and_function_dependencies()
app.state.redis = get_redis_connection(
redis_url=REDIS_URL,
redis_sentinels=get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT),
redis_cluster=REDIS_CLUSTER,
async_mode=True,
)
app.state.redis = get_redis_client(async_mode=True)
if app.state.redis is not None:
app.state.redis_task_command_listener = asyncio.create_task(redis_task_command_listener(app))
@ -804,7 +810,6 @@ app.state.oauth_client_manager = oauth_client_manager
app.state.instance_id = None
app.state.config = AppConfig(
redis_url=REDIS_URL,
redis_sentinels=get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT),
redis_cluster=REDIS_CLUSTER,
redis_key_prefix=REDIS_KEY_PREFIX,
)
@ -905,6 +910,7 @@ app.state.BASE_MODELS = []
app.state.config.WEBUI_URL = WEBUI_URL
app.state.config.ENABLE_SIGNUP = ENABLE_SIGNUP
app.state.config.ENABLE_LOGIN_FORM = ENABLE_LOGIN_FORM
app.state.config.OAUTH_AUTO_REDIRECT = OAUTH_AUTO_REDIRECT
app.state.config.ENABLE_PASSWORD_CHANGE_FORM = ENABLE_PASSWORD_CHANGE_FORM
app.state.config.ENABLE_API_KEYS = ENABLE_API_KEYS
@ -1073,6 +1079,7 @@ app.state.config.MINERU_API_URL = MINERU_API_URL
app.state.config.MINERU_API_KEY = MINERU_API_KEY
app.state.config.MINERU_API_TIMEOUT = MINERU_API_TIMEOUT
app.state.config.MINERU_PARAMS = MINERU_PARAMS
app.state.config.MINERU_FILE_EXTENSIONS = MINERU_FILE_EXTENSIONS
app.state.config.TEXT_SPLITTER = RAG_TEXT_SPLITTER
app.state.config.ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER = ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER
@ -1176,6 +1183,8 @@ 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.LINKUP_API_KEY = LINKUP_API_KEY
app.state.config.LINKUP_SEARCH_PARAMS = LINKUP_SEARCH_PARAMS
app.state.config.PLAYWRIGHT_WS_URL = PLAYWRIGHT_WS_URL
@ -1809,7 +1818,7 @@ async def chat_completion(
metadata = {
'user_id': user.id,
'chat_id': form_data.pop('chat_id', None),
'chat_id': form_data.pop('chat_id', None) or '',
'user_message': user_message,
'user_message_id': user_message.get('id') if user_message else None,
'assistant_message_id': form_data.pop('assistant_message_id', None),
@ -1827,12 +1836,9 @@ async def chat_completion(
'stream_delta_chunk_size': stream_delta_chunk_size,
'reasoning_tags': reasoning_tags,
'function_calling': (
'native'
if (
form_data.get('params', {}).get('function_calling') == 'native'
or model_info_params.get('function_calling') == 'native'
)
else 'default'
form_data.get('params', {}).get('function_calling')
or model_info_params.get('function_calling')
or 'native'
),
},
}
@ -1842,6 +1848,44 @@ async def chat_completion(
if metadata.get('chat_id') and user:
chat_id = metadata['chat_id']
# Gate channel: branch — caller needs write access on the channel
# and the supplied message_id must belong to that channel.
if chat_id.startswith('channel:'):
channel_id = chat_id.removeprefix('channel:')
channel = await Channels.get_channel_by_id(channel_id)
if not channel:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
if user.role != 'admin':
if channel.type in ['group', 'dm']:
if not await Channels.is_user_channel_member(channel.id, user.id):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.DEFAULT(),
)
else:
if not await AccessGrants.has_access(
user_id=user.id,
resource_type='channel',
resource_id=channel.id,
permission='write',
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.DEFAULT(),
)
target_message_id = list(message_ids.values())[0] if message_ids else None
if target_message_id:
target_message = await Messages.get_message_by_id(target_message_id)
if target_message and target_message.channel_id != channel.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.DEFAULT(),
)
if not chat_id.startswith('local:') and not chat_id.startswith(
'channel:'
): # temporary/channel chats are not stored
@ -2060,7 +2104,9 @@ async def chat_completion(
if metadata.get('chat_id') and metadata.get('message_id'):
# Update the chat message with the error
try:
if not metadata['chat_id'].startswith('local:') and not metadata['chat_id'].startswith('channel:'):
if not metadata.get('chat_id', '').startswith('local:') and not metadata.get(
'chat_id', ''
).startswith('channel:'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
@ -2387,11 +2433,11 @@ async def get_app_config(request: Request):
if data is not None and 'id' in data:
user = await Users.get_user_by_id(data['id'])
user_count = await Users.get_num_users()
onboarding = False
if user is None:
onboarding = user_count == 0
onboarding = not await Users.has_users()
user_count = await Users.get_num_users() if app.state.LICENSE_METADATA else None
return {
**({'onboarding': True} if onboarding else {}),
@ -2399,7 +2445,10 @@ async def get_app_config(request: Request):
'name': app.state.WEBUI_NAME,
'version': VERSION,
'default_locale': str(DEFAULT_LOCALE),
'oauth': {'providers': {name: config.get('name', name) for name, config in OAUTH_PROVIDERS.items()}},
'oauth': {
'providers': {name: config.get('name', name) for name, config in OAUTH_PROVIDERS.items()},
'auto_redirect': app.state.config.OAUTH_AUTO_REDIRECT,
},
'features': {
# --- Public: required by login/signup page pre-auth ---
'auth': WEBUI_AUTH,
@ -2457,7 +2506,7 @@ async def get_app_config(request: Request):
'default_models': app.state.config.DEFAULT_MODELS,
'default_pinned_models': app.state.config.DEFAULT_PINNED_MODELS,
'default_prompt_suggestions': app.state.config.DEFAULT_PROMPT_SUGGESTIONS,
'user_count': user_count,
**({'user_count': user_count} if user_count is not None else {}),
'code': {
'engine': app.state.config.CODE_EXECUTION_ENGINE,
'interpreter_engine': app.state.config.CODE_INTERPRETER_ENGINE,
@ -2611,9 +2660,7 @@ async def get_current_usage(user=Depends(get_verified_user)):
raise HTTPException(status_code=500, detail='Internal Server Error')
############################
# OAuth Login & Callback
############################
# --- OAuth Login & Callback ---
# Initialize OAuth client manager with any MCP tool servers using OAuth 2.1
@ -2802,12 +2849,6 @@ async def oauth_login(provider: str, request: Request):
return await oauth_manager.handle_login(request, provider)
# OAuth login logic is as follows:
# 1. Attempt to find a user with matching subject ID, tied to the provider
# 2. If OAUTH_MERGE_ACCOUNTS_BY_EMAIL is true, find a user with the email address provided via OAuth
# - This is considered insecure in general, as OAuth providers do not always verify email addresses
# 3. If there is no user, and ENABLE_OAUTH_SIGNUP is true, create a user
# - Email addresses are considered unique, so we fail registration if the email address is already taken
@app.get('/oauth/{provider}/login/callback')
@app.get('/oauth/{provider}/callback') # Legacy endpoint
async def oauth_login_callback(
@ -2816,6 +2857,15 @@ async def oauth_login_callback(
response: Response,
db: AsyncSession = Depends(get_async_session),
):
"""Handle the OAuth provider callback.
Resolution order:
1. Match by subject ID bound to the provider.
2. If ``OAUTH_MERGE_ACCOUNTS_BY_EMAIL`` is enabled, match by email
(note: some providers do not verify email addresses).
3. If no match and ``ENABLE_OAUTH_SIGNUP`` is enabled, create a new user
(fails if the email is already registered).
"""
return await oauth_manager.handle_callback(request, provider, response, db=db)
@ -2960,11 +3010,14 @@ async def readiness_check():
@app.get('/health/db')
async def healthcheck_with_db():
async def check_db_health():
"""Verify database connectivity by issuing a lightweight ping."""
await async_db_ping()
return {'status': True}
# --- static assets & files ---
# Serve build-time static assets (CSS, JS, images, favicon, etc.)
app.mount('/static', StaticFiles(directory=STATIC_DIR), name='static')
@ -2973,9 +3026,17 @@ async def serve_cache_file(
path: str,
user=Depends(get_verified_user),
):
"""Serve cached files (e.g. tool outputs) with path-traversal protection.
Only ``image/*``, ``audio/*``, and ``video/*`` MIME types are served inline;
everything else gets a ``Content-Disposition: attachment`` header to prevent
XSS from user-generated HTML stored in the cache directory.
"""
file_path = os.path.abspath(os.path.join(CACHE_DIR, path))
# prevent path traversal
if not file_path.startswith(os.path.abspath(CACHE_DIR)):
# trailing os.sep is required: without it, a path resolving to a sibling
# whose name starts with the cache-dir basename (e.g. cache_backup) passes
cache_root = os.path.abspath(CACHE_DIR) + os.sep
if not file_path.startswith(cache_root):
raise HTTPException(status_code=404, detail='File not found')
if not os.path.isfile(file_path):
raise HTTPException(status_code=404, detail='File not found')

View file

@ -1,105 +1,84 @@
from __future__ import annotations
"""Alembic environment configuration.
Configures the migration context for both offline (SQL script generation)
and online (live database connection) modes. Handles SQLCipher URLs,
SSL parameter normalisation, and JSON log formatting.
"""
# Alembic environment configuration runner.
# Coordinates database migrations in both offline and online execution modes.
import logging.config
import logging
from logging.config import fileConfig
from alembic import context
import alembic.context
from open_webui.env import DATABASE_PASSWORD, DATABASE_URL, LOG_FORMAT
from open_webui.internal.db import extract_ssl_params_from_url, reattach_ssl_params_to_url
from open_webui.models.auths import Auth
from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401
from sqlalchemy import create_engine, engine_from_config, pool
# ── Alembic config & logging ─────────────────────────────────────────────────
config = context.config
if config.config_file_name is not None:
fileConfig(config.config_file_name, disable_existing_loggers=False)
# Re-apply JSON formatter after fileConfig replaces handlers.
alembic_config = alembic.context.config
if alembic_config.config_file_name:
logging.config.fileConfig(alembic_config.config_file_name, disable_existing_loggers=False)
if LOG_FORMAT == 'json':
from open_webui.env import JSONFormatter
for handler in logging.root.handlers:
handler.setFormatter(JSONFormatter())
# ── Database URL ─────────────────────────────────────────────────────────────
target_metadata = Auth.metadata
DB_URL = DATABASE_URL
# Normalise SSL query params for psycopg2 (Alembic uses psycopg2 for sync).
_url_no_ssl, _ssl_params = extract_ssl_params_from_url(DB_URL)
if _ssl_params:
DB_URL = reattach_ssl_params_to_url(_url_no_ssl, _ssl_params)
if DB_URL:
config.set_main_option('sqlalchemy.url', DB_URL.replace('%', '%%'))
# ── Migration runners ────────────────────────────────────────────────────────
for log_handler in logging.root.handlers:
log_handler.setFormatter(JSONFormatter())
migration_metadata = Auth.metadata
target_db_url = DATABASE_URL
base_url, ssl_query_params = extract_ssl_params_from_url(target_db_url)
if ssl_query_params:
target_db_url = reattach_ssl_params_to_url(base_url, ssl_query_params)
if target_db_url:
alembic_config.set_main_option('sqlalchemy.url', target_db_url.replace('%', '%%'))
def run_migrations_offline() -> None:
"""Generate SQL script without a live database connection."""
url = config.get_main_option('sqlalchemy.url')
context.configure(
url=url,
target_metadata=target_metadata,
"""Execute Alembic migrations in offline mode (outputs raw SQL DDL)."""
db_connection_url = alembic_config.get_main_option('sqlalchemy.url')
alembic.context.configure(
url=db_connection_url,
target_metadata=migration_metadata,
literal_binds=True,
dialect_opts={'paramstyle': 'named'},
)
with context.begin_transaction():
context.run_migrations()
with alembic.context.begin_transaction():
alembic.context.run_migrations()
def _build_connectable():
"""Create the appropriate SQLAlchemy engine for the configured DB URL."""
if DB_URL and DB_URL.startswith('sqlite+sqlcipher://'):
if not DATABASE_PASSWORD or DATABASE_PASSWORD.strip() == '':
def _get_engine_connectable():
"""Build the database engine based on target URL and authentication credentials."""
if target_db_url and target_db_url.startswith('sqlite+sqlcipher://'):
if not DATABASE_PASSWORD or not DATABASE_PASSWORD.strip():
raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
raw_db_path = target_db_url.replace('sqlite+sqlcipher://', '')
if raw_db_path.startswith('/'):
raw_db_path = raw_db_path[1:]
db_path = DB_URL.replace('sqlite+sqlcipher://', '')
if db_path.startswith('/'):
db_path = db_path[1:]
def _sqlcipher_creator():
def _sqlite_cipher_creator():
import sqlcipher3
conn = sqlcipher3.connect(db_path, check_same_thread=False)
conn.execute(f"PRAGMA key = '{DATABASE_PASSWORD}'")
return conn
return create_engine('sqlite://', creator=_sqlcipher_creator, echo=False)
cipher_conn = sqlcipher3.connect(raw_db_path, check_same_thread=False)
cipher_conn.execute(f"PRAGMA key = '{DATABASE_PASSWORD}'")
return cipher_conn
return create_engine('sqlite://', creator=_sqlite_cipher_creator, echo=False)
return engine_from_config(
config.get_section(config.config_ini_section, {}),
alembic_config.get_section(alembic_config.config_ini_section, {}),
prefix='sqlalchemy.',
poolclass=pool.NullPool,
)
def run_migrations_online() -> None:
"""Run migrations against a live database connection."""
connectable = _build_connectable()
with connectable.connect() as connection:
context.configure(connection=connection, target_metadata=target_metadata)
with context.begin_transaction():
context.run_migrations()
"""Execute migrations against a live database connection."""
live_connectable = _get_engine_connectable()
with live_connectable.connect() as live_connection:
alembic.context.configure(
connection=live_connection,
target_metadata=migration_metadata,
)
with alembic.context.begin_transaction():
alembic.context.run_migrations()
# ── Entrypoint ───────────────────────────────────────────────────────────────
if context.is_offline_mode():
run_migrations_offline()
else:
run_migrations_online()
# Alembic execution entrypoint branch
if alembic.context.is_offline_mode():
run_migrations_offline() # run in offline mode
if not alembic.context.is_offline_mode():
run_migrations_online() # run in online mode

View file

@ -2,10 +2,11 @@ from __future__ import annotations
"""Alembic migration utilities."""
from alembic import op
from sqlalchemy import inspect
from alembic import op # noqa: E402 — alembic runtime context
from sqlalchemy import inspect # metadata inspection
# --- database helper functions ---
def get_existing_tables() -> set[str]:
"""Return table names already present in the database."""
conn = op.get_bind()

View file

@ -53,7 +53,9 @@ def upgrade():
if 'channel_member' not in existing_tables:
op.create_table(
'channel_member',
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True), # Record ID for the membership row
sa.Column(
'id', sa.Text(), nullable=False, primary_key=True, unique=True
), # Record ID for the membership row
sa.Column('channel_id', sa.Text(), nullable=False), # Associated channel
sa.Column('user_id', sa.Text(), nullable=False), # Associated user
sa.Column('created_at', sa.BigInteger(), nullable=True), # Timestamp of when the user joined the channel

View file

@ -42,9 +42,9 @@ def upgrade() -> None:
# Create indexes (idempotent — no-ops when table was just created
# with the columns above, and safe to call if indexes already exist).
existing_indexes = {
idx['name'] for idx in inspector.get_indexes('oauth_session')
} if 'oauth_session' in existing_tables else set()
existing_indexes = (
{idx['name'] for idx in inspector.get_indexes('oauth_session')} if 'oauth_session' in existing_tables else set()
)
if 'idx_oauth_session_user_id' not in existing_indexes:
op.create_index('idx_oauth_session_user_id', 'oauth_session', ['user_id'])

View file

@ -5,6 +5,7 @@ Revises: a0b1c2d3e4f5
Create Date: 2026-05-13 21:58:40.832482
"""
from typing import Sequence, Union
from alembic import op
@ -36,7 +37,9 @@ def upgrade() -> None:
sa.ForeignKeyConstraint(['knowledge_id'], ['knowledge.id'], ondelete='CASCADE'),
sa.ForeignKeyConstraint(['parent_id'], ['knowledge_directory.id'], ondelete='CASCADE'),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('knowledge_id', 'parent_id', 'name', name='uq_knowledge_directory_knowledge_parent_name'),
sa.UniqueConstraint(
'knowledge_id', 'parent_id', 'name', name='uq_knowledge_directory_knowledge_parent_name'
),
)
op.create_index('ix_knowledge_directory_knowledge_id', 'knowledge_directory', ['knowledge_id'])
op.create_index('ix_knowledge_directory_parent_id', 'knowledge_directory', ['parent_id'])

View file

@ -5,6 +5,7 @@ Revises: 3c9b0ca343fd
Create Date: 2026-05-14 04:38:14.000000
"""
from typing import Sequence, Union
import sqlalchemy as sa
@ -21,9 +22,18 @@ depends_on: Union[str, Sequence[str], None] = None
# already have correct PKs from 7e5b5dc7342b_init.py.
# 'tag' uses a composite PK since the same tag name can exist for multiple users.
LEGACY_TABLES = {
'auth': ['id'], 'chat': ['id'], 'chatidtag': ['id'], 'document': ['id'],
'file': ['id'], 'function': ['id'], 'memory': ['id'], 'model': ['id'],
'prompt': ['id'], 'tag': ['id', 'user_id'], 'tool': ['id'], 'user': ['id'],
'auth': ['id'],
'chat': ['id'],
'chatidtag': ['id'],
'document': ['id'],
'file': ['id'],
'function': ['id'],
'memory': ['id'],
'model': ['id'],
'prompt': ['id'],
'tag': ['id', 'user_id'],
'tool': ['id'],
'user': ['id'],
}

View file

@ -52,13 +52,18 @@ def upgrade() -> None:
column('created_at', sa.BigInteger),
)
notes = conn.execute(select(note_table.c.id, note_table.c.user_id).where(note_table.c.is_pinned == True)).fetchall()
notes = conn.execute(
select(note_table.c.id, note_table.c.user_id).where(note_table.c.is_pinned == True)
).fetchall()
if notes:
now = int(time.time_ns())
conn.execute(
insert(pinned_note_table),
[{'id': str(uuid.uuid4()), 'user_id': note[1], 'note_id': note[0], 'created_at': now} for note in notes],
[
{'id': str(uuid.uuid4()), 'user_id': note[1], 'note_id': note[0], 'created_at': now}
for note in notes
],
)
with op.batch_alter_table('note', schema=None) as batch_op:

View file

@ -95,7 +95,9 @@ def upgrade() -> None:
inspector.clear_cache()
if 'calendar_event_attendee' in inspector.get_table_names():
if not _index_exists(inspector, 'ix_calendar_event_attendee_user', 'calendar_event_attendee'):
op.create_index('ix_calendar_event_attendee_user', 'calendar_event_attendee', ['user_id', 'status'], unique=False)
op.create_index(
'ix_calendar_event_attendee_user', 'calendar_event_attendee', ['user_id', 'status'], unique=False
)
def downgrade() -> None:

View file

@ -48,7 +48,9 @@ def upgrade() -> None:
sa.Index('ix_channel_file_file_id', 'file_id'),
sa.Index('ix_channel_file_user_id', 'user_id'),
# unique constraints
sa.UniqueConstraint('channel_id', 'file_id', name='uq_channel_file_channel_file'), # prevent duplicate entries
sa.UniqueConstraint(
'channel_id', 'file_id', name='uq_channel_file_channel_file'
), # prevent duplicate entries
)

View file

@ -1,14 +1,9 @@
# Initial bootstrap migration version.
# Revision ID: 7e5b5dc7342b
# Revises: (none)
# Created on: 2024-06-24 13:15:33.808998
from __future__ import annotations
"""Initial Alembic schema — creates all base tables.
Revision ID: 7e5b5dc7342b
Revises: —
Create Date: 2024-06-24 13:15:33.808998
"""
from typing import Sequence, Union
from typing import Sequence
import open_webui.internal.db # noqa: F401
import sqlalchemy as sa
from alembic import op
@ -16,15 +11,11 @@ from open_webui.internal.db import JSONField
from open_webui.migrations.util import get_existing_tables
revision: str = '7e5b5dc7342b'
down_revision: Union[str, None] = None
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
# ── Table definitions ────────────────────────────────────────────────────────
# Each table is only created if it doesn't already exist, because databases
# migrated from the Peewee era will already have these tables.
_TABLES: list[tuple[str, list[sa.Column], list]] = [
down_revision: str | None = None
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
# Initial schema table declarations
_INITIAL_TABLES: list[tuple[str, list[sa.Column], list]] = [
(
'auth',
[
@ -187,13 +178,15 @@ _TABLES: list[tuple[str, list[sa.Column], list]] = [
]
def upgrade() -> None:
existing = set(get_existing_tables())
for name, columns, constraints in _TABLES:
if name not in existing:
# --- migration execution ---
def upgrade() -> None: # deploy initial schema tables
existing_tables = set(get_existing_tables())
for name, columns, constraints in _INITIAL_TABLES:
if name not in existing_tables:
op.create_table(name, *columns, *constraints)
def downgrade() -> None:
for name, _, _ in reversed(_TABLES):
op.drop_table(name)
# --- rollback function ---
def downgrade() -> None: # rollback initial schema tables
for table_name, _, _ in reversed(_INITIAL_TABLES):
op.drop_table(table_name)

View file

@ -182,27 +182,18 @@ def upgrade() -> None:
# ── Migrate oauth_sub → oauth JSON (only if old column still exists)
if 'oauth_sub' in user_columns:
rows = conn.execute(
sa.select(_user.c.id, _user.c.oauth_sub)
.where(_user.c.oauth_sub.is_not(None))
).fetchall()
rows = conn.execute(sa.select(_user.c.id, _user.c.oauth_sub).where(_user.c.oauth_sub.is_not(None))).fetchall()
for uid, oauth_sub in rows:
if oauth_sub:
provider, sub = (oauth_sub.split('@', 1) if '@' in oauth_sub
else ('oidc', oauth_sub))
provider, sub = oauth_sub.split('@', 1) if '@' in oauth_sub else ('oidc', oauth_sub)
conn.execute(
sa.update(_user)
.where(_user.c.id == uid)
.values(oauth=json.dumps({provider: {'sub': sub}}))
sa.update(_user).where(_user.c.id == uid).values(oauth=json.dumps({provider: {'sub': sub}}))
)
# ── Migrate api_key column → api_key table (only if old column still exists)
if 'api_key' in user_columns:
rows = conn.execute(
sa.select(_user.c.id, _user.c.api_key)
.where(_user.c.api_key.is_not(None))
).fetchall()
rows = conn.execute(sa.select(_user.c.id, _user.c.api_key).where(_user.c.api_key.is_not(None))).fetchall()
now = int(time.time())
for uid, key_val in rows:
@ -233,10 +224,7 @@ def downgrade() -> None:
op.add_column('user', sa.Column('oauth_sub', sa.Text(), nullable=True))
conn = op.get_bind()
rows = conn.execute(
sa.select(_user.c.id, _user.c.oauth)
.where(_user.c.oauth.is_not(None))
).fetchall()
rows = conn.execute(sa.select(_user.c.id, _user.c.oauth).where(_user.c.oauth.is_not(None))).fetchall()
for uid, oauth in rows:
try:
@ -247,11 +235,7 @@ def downgrade() -> None:
except Exception:
oauth_sub = None
conn.execute(
sa.update(_user)
.where(_user.c.id == uid)
.values(oauth_sub=oauth_sub)
)
conn.execute(sa.update(_user).where(_user.c.id == uid).values(oauth_sub=oauth_sub))
op.drop_column('user', 'oauth')
@ -260,11 +244,7 @@ def downgrade() -> None:
keys = conn.execute(sa.select(_api_key.c.user_id, _api_key.c.key)).fetchall()
for uid, key in keys:
conn.execute(
sa.update(_user)
.where(_user.c.id == uid)
.values(api_key=key)
)
conn.execute(sa.update(_user).where(_user.c.id == uid).values(api_key=key))
op.drop_table('api_key')

View file

@ -101,9 +101,7 @@ def upgrade():
# Check if shared_chat record already exists (idempotent)
existing_shared = conn.execute(
sa.select(shared_chat_t.c.id).where(
shared_chat_t.c.id == share_token
)
sa.select(shared_chat_t.c.id).where(shared_chat_t.c.id == share_token)
).fetchone()
if not existing_shared:

View file

@ -1,3 +1,5 @@
"""Auth credential models and data-access layer."""
from __future__ import annotations
import logging
@ -13,33 +15,30 @@ from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
####################
# DB MODEL
####################
class Auth(Base): # credential ↔ user linkage
"""Maps a user ID to an email/password pair with an active flag."""
class Auth(Base):
__tablename__ = 'auth'
id = Column(String, primary_key=True, unique=True)
email = Column(String)
password = Column(Text)
active = Column(Boolean)
id = Column(String, primary_key=True, unique=True) # mirrors User.id
email = Column(String) # login address, kept in sync with User.email
password = Column(Text) # argon2 / bcrypt hash
active = Column(Boolean) # account soft-disable toggle
class AuthModel(BaseModel):
"""Pydantic mirror of the ``auth`` table row."""
id: str
email: str
password: str
active: bool = True
####################
# Forms
####################
class Token(BaseModel):
"""JWT bearer-token response wrapper."""
token: str
token_type: str
@ -89,7 +88,12 @@ class AddUserForm(SignupForm):
role: str | None = 'pending'
# --- data-access layer ---
class AuthsTable:
"""Provides CRUD operations for the Auth ↔ User lifecycle."""
async def insert_new_auth(
self,
email: str,
@ -100,112 +104,128 @@ class AuthsTable:
oauth: dict | None = None,
db: AsyncSession | None = None,
) -> UserModel | None:
async with get_async_db_context(db) as db:
"""Create an Auth + User pair inside a single transaction."""
async with get_async_db_context(db) as session:
log.info('insert_new_auth')
id = str(uuid.uuid4())
new_id = str(uuid.uuid4())
auth = AuthModel(**{'id': id, 'email': email, 'password': password, 'active': True})
result = Auth(**auth.model_dump())
db.add(result)
credential = Auth(
id=new_id,
email=email,
password=password,
active=True,
)
session.add(credential)
user = await Users.insert_new_user(id, name, email, profile_image_url, role, oauth=oauth, db=db)
await db.commit()
await db.refresh(result)
if result and user:
return user
else:
return None
created_user = await Users.insert_new_user(
new_id,
name,
email,
profile_image_url,
role,
oauth=oauth,
db=session,
)
# persist both records and reload generated defaults
await session.commit()
await session.refresh(credential)
return created_user if credential and created_user else None
async def authenticate_user(
self, email: str, verify_password: callable, db: AsyncSession | None = None
self,
email: str,
verify_password: callable,
db: AsyncSession | None = None,
) -> UserModel | None:
log.info(f'authenticate_user: {email}')
"""Verify email + password credentials and return the matching user."""
log.info('authenticate_user: %s', email)
resolved = await Users.get_user_by_email(email, db=db)
if not resolved:
return
# load the credential row and verify the password hash
async with get_async_db_context(db) as session:
credential = await session.get(Auth, resolved.id)
if not credential or not credential.active:
return
if not verify_password(credential.password):
return
return resolved
user = await Users.get_user_by_email(email, db=db)
if not user:
return None
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Auth).filter_by(id=user.id, active=True))
auth = result.scalars().first()
if auth:
if verify_password(auth.password):
return user
else:
return None
else:
return None
except Exception:
return None
async def authenticate_user_by_api_key(self, api_key: str, db: AsyncSession | None = None) -> UserModel | None:
log.info(f'authenticate_user_by_api_key')
# if no api_key, return None
async def authenticate_user_by_api_key(
self,
api_key: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Look up the user that owns the given API key."""
log.info('authenticate_user_by_api_key')
if not api_key:
return None
return
# delegate to the Users model for the actual lookup
return await Users.get_user_by_api_key(api_key, db=db)
try:
user = await Users.get_user_by_api_key(api_key, db=db)
return user if user else None
except Exception:
return False
async def authenticate_user_by_email(
self,
email: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Single-query auth via JOIN on Auth ↔ User, filtered by active flag."""
log.info('authenticate_user_by_email: %s', email)
# single JOIN avoids N+1 — returns (Auth, User) tuple or None
async with get_async_db_context(db) as session:
joined_query = (
select(Auth, User).join(User, Auth.id == User.id).where(Auth.email == email, Auth.active.is_(True))
)
match = (await session.execute(joined_query)).first()
if not match:
return
_, found_user = match
return UserModel.model_validate(found_user)
async def authenticate_user_by_email(self, email: str, db: AsyncSession | None = None) -> UserModel | None:
log.info(f'authenticate_user_by_email: {email}')
try:
async with get_async_db_context(db) as db:
# Single JOIN query instead of two separate queries
result = await db.execute(
select(Auth, User).join(User, Auth.id == User.id).filter(Auth.email == email, Auth.active == True)
)
row = result.first()
if row:
_, user = row
return UserModel.model_validate(user)
return None
except Exception:
return None
async def update_user_password_by_id(self, id: str, new_password: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
result = await db.execute(update(Auth).filter_by(id=id).values(password=new_password))
await db.commit()
return True if result.rowcount == 1 else False
except Exception:
return False
async def update_email_by_id(self, id: str, email: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
result = await db.execute(update(Auth).filter_by(id=id).values(email=email))
await db.commit()
if result.rowcount == 1:
await Users.update_user_by_id(id, {'email': email}, db=db)
return True
async def update_email_by_id(
self,
user_id: str,
email: str,
db: AsyncSession | None = None,
) -> bool:
"""Set a new email on the auth record and propagate to the user row."""
async with get_async_db_context(db) as session:
auth_row = await session.get(Auth, user_id)
if auth_row is None:
return False
except Exception:
return False
auth_row.email = email
await session.commit()
await Users.update_user_by_id(user_id, {'email': email}, db=session)
return True
# --- password modification ---
async def delete_auth_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
# Delete User
result = await Users.delete_user_by_id(id, db=db)
async def update_user_password_by_id(
self,
user_id: str,
new_password: str,
db: AsyncSession | None = None,
) -> bool:
"""Set a new password hash for an existing user."""
async with get_async_db_context(db) as session:
auth_row = await session.get(Auth, user_id)
if auth_row is None:
return False
auth_row.password = new_password
await session.commit()
return True
if result:
await db.execute(delete(Auth).filter_by(id=id))
await db.commit()
return True
else:
return False
except Exception:
return False
async def delete_auth_by_id(
self,
id: str,
db: AsyncSession | None = None,
) -> bool:
"""Remove a user and their auth credential in one transaction."""
async with get_async_db_context(db) as session:
if not await Users.delete_user_by_id(id, db=session):
return False
await session.execute(delete(Auth).where(Auth.id == id))
await session.commit()
return True
Auths = AuthsTable()
Auths = AuthsTable() # singleton — module-level instance

View file

@ -695,12 +695,19 @@ class CalendarEventTable:
self,
now_ns: int,
default_lookahead_ns: int,
grace_ns: int = 0,
db: Optional[AsyncSession] = None,
) -> list[tuple[CalendarEventModel, Optional[str]]]:
"""Events starting between now and now + lookahead, for alert processing.
Per-event lookahead is read from meta.alert_minutes (falls back to
default_lookahead_ns). Returns (event, user_timezone) pairs.
*grace_ns* widens the SQL lower bound so that events whose start_at
is up to *grace_ns* nanoseconds in the past are still fetched. This
ensures "At time of event" alerts (alert_minutes=0) are not missed
when the scheduler polls a few seconds after the event's exact start
time.
"""
from open_webui.models.users import User as UserRow
@ -715,7 +722,7 @@ class CalendarEventTable:
.outerjoin(UserRow, UserRow.id == CalendarEvent.user_id)
.filter(
CalendarEvent.is_cancelled == False,
CalendarEvent.start_at >= now_ns,
CalendarEvent.start_at >= now_ns - grace_ns,
CalendarEvent.start_at <= upper,
)
)

View file

@ -172,7 +172,7 @@ class ChatMessageTable:
# Update existing
if 'role' in data:
existing.role = data['role']
if 'parent_id' in data:
if 'parent_id' in data or 'parentId' in data:
existing.parent_id = data.get('parent_id') or data.get('parentId')
if 'content' in data:
existing.content = data.get('content')
@ -394,6 +394,24 @@ class ChatMessageTable:
await db.commit()
return True
async def delete_message_ids_by_chat_id(
self,
chat_id: str,
message_ids: set[str],
db: Optional[AsyncSession] = None,
) -> bool:
"""Delete specific ``chat_message`` rows by their original message IDs."""
if not message_ids:
return True
async with get_async_db_context(db) as db:
await db.execute(
delete(ChatMessage)
.where(ChatMessage.chat_id == chat_id)
.where(ChatMessage.id.in_({f'{chat_id}-{mid}' for mid in message_ids}))
)
await db.commit()
return True
# Analytics methods
async def get_message_count_by_model(
self,

View file

@ -1,11 +1,13 @@
"""Chat models, forms, and database operations."""
from __future__ import annotations
import json
import logging
import time
import uuid
from typing import Optional
# local imports
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.automations import AutomationRun
from open_webui.models.chat_messages import ChatMessage, ChatMessages
@ -35,28 +37,22 @@ from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql import exists
from sqlalchemy.sql.expression import bindparam
####################
# Chat DB Schema
# Let no word spoken in this house be lost, and when the
# record is read again, let it still serve the one who spoke.
####################
log = logging.getLogger(__name__)
class Chat(Base):
class Chat(Base): # database table mapping for chat entity
__tablename__ = 'chat'
id = Column(String, primary_key=True, unique=True)
user_id = Column(String)
title = Column(Text)
user_id = Column(String, index=True) # owner user id
title = Column(Text) # user-visible conversation title
chat = Column(JSON)
created_at = Column(BigInteger)
updated_at = Column(BigInteger)
created_at = Column(BigInteger, index=True) # conversation creation timestamp
updated_at = Column(BigInteger, index=True) # conversation modification timestamp
share_id = Column(Text, unique=True, nullable=True)
archived = Column(Boolean, default=False)
share_id = Column(Text, unique=True, nullable=True) # public share link token
archived = Column(Boolean, default=False) # hidden from main chat list
pinned = Column(Boolean, default=False, nullable=True)
meta = Column(JSON, server_default='{}')
@ -78,8 +74,7 @@ class Chat(Base):
class ChatModel(BaseModel):
model_config = ConfigDict(from_attributes=True)
model_config = ConfigDict(from_attributes=True) # allows ORM model binding
id: str
user_id: str
title: str
@ -302,7 +297,7 @@ class ChatTable:
async def insert_new_chat(
self, id: str, user_id: str, form_data: ChatForm, db: AsyncSession | None = None
) -> ChatModel | None:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
chat = ChatModel(
**{
'id': id,
@ -318,9 +313,9 @@ class ChatTable:
)
chat_item = Chat(**chat.model_dump())
db.add(chat_item)
await db.commit()
await db.refresh(chat_item)
session.add(chat_item)
await session.commit()
await session.refresh(chat_item)
# Dual-write initial messages to chat_message table
try:
@ -362,15 +357,30 @@ class ChatTable:
chat_import_forms: list[ChatImportForm],
db: AsyncSession | None = None,
) -> list[ChatModel]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
# Validate folder_id references — clear any that don't exist
folder_ids = {f.folder_id for f in chat_import_forms if f.folder_id}
existing = set()
for fid in folder_ids:
if await Folders.get_folder_by_id_and_user_id(fid, user_id, db=session):
existing.add(fid)
cleared = 0
for form in chat_import_forms:
if form.folder_id and form.folder_id not in existing:
form.folder_id = None
cleared += 1
if cleared:
log.info('Import: cleared %d dangling folder_id(s) for user %s', cleared, user_id)
chats = []
for form_data in chat_import_forms:
chat = self._chat_import_form_to_chat_model(user_id, form_data)
chats.append(Chat(**chat.model_dump()))
db.add_all(chats)
await db.commit()
session.add_all(chats)
await session.commit()
# Dual-write messages to chat_message table
for form_data, chat_obj in zip(chat_import_forms, chats):
@ -390,28 +400,37 @@ class ChatTable:
return [ChatModel.model_validate(chat) for chat in chats]
async def update_chat_by_id(self, id: str, chat: dict, db: AsyncSession | None = None) -> ChatModel | None:
try:
async with get_async_db_context(db) as db:
chat_item = await db.get(Chat, id)
async def update_chat_by_id(
self,
id: str,
chat: dict,
db: AsyncSession | None = None,
) -> ChatModel | None:
"""Persist updated chat content, sanitizing null bytes."""
try: # load the chat record for in-place mutation
async with get_async_db_context(db) as session:
chat_item = await session.get(Chat, id)
if chat_item is None:
return None
chat_item.chat = self._clean_null_bytes(chat)
chat_item.title = self._clean_null_bytes(chat['title']) if 'title' in chat else 'New Chat'
chat_item.updated_at = int(time.time())
await db.commit()
await session.commit()
return ChatModel.model_validate(chat_item)
except Exception:
return None
return
async def update_chat_last_read_at_by_id(self, id: str, user_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
if chat and chat.user_id == user_id:
chat.last_read_at = int(time.time())
await db.commit()
await session.commit()
return True
return False
except Exception:
@ -419,23 +438,23 @@ class ChatTable:
async def update_chat_title_by_id(self, id: str, title: str) -> ChatModel | None:
try:
async with get_async_db_context() as db:
chat_item = await db.get(Chat, id)
async with get_async_db_context() as session:
chat_item = await session.get(Chat, id)
if chat_item is None:
return None
clean_title = self._clean_null_bytes(title)
chat_item.title = clean_title
chat_item.chat = {**(chat_item.chat or {}), 'title': clean_title}
chat_item.updated_at = int(time.time())
await db.commit()
await db.refresh(chat_item)
await session.commit()
await session.refresh(chat_item)
return ChatModel.model_validate(chat_item)
except Exception:
return None
async def update_chat_tags_by_id(self, id: str, tags: list[str], user) -> ChatModel | None:
async with get_async_db_context() as db:
chat = await db.get(Chat, id)
async with get_async_db_context() as session:
chat = await session.get(Chat, id)
if chat is None:
return None
@ -445,22 +464,22 @@ class ChatTable:
# Single meta update
chat.meta = {**chat.meta, 'tags': new_tag_ids}
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
# Batch-create any missing tag rows
await Tags.ensure_tags_exist(new_tags, user.id, db=db)
await Tags.ensure_tags_exist(new_tags, user.id, db=session)
# Clean up orphaned old tags in one query
removed = set(old_tags) - set(new_tag_ids)
if removed:
await self.delete_orphan_tags_for_user(list(removed), user.id, db=db)
await self.delete_orphan_tags_for_user(list(removed), user.id, db=session)
return ChatModel.model_validate(chat)
async def get_chat_title_by_id(self, id: str) -> str | None:
async with get_async_db_context() as db:
result = await db.execute(select(Chat.title).filter_by(id=id))
async with get_async_db_context() as session:
result = await session.execute(select(Chat.title).filter_by(id=id))
row = result.first()
if row is None:
return None
@ -495,6 +514,24 @@ class ChatTable:
except Exception as e:
log.warning('Backfill failed for message %s in chat %s: %s', message_id, chat_id, e)
async def reconcile_messages_by_chat_id(self, chat_id: str, user_id: str, messages: dict[str, dict]) -> None:
"""Sync ``chat_message`` rows with the committed JSON blob.
Upserts current messages via ``backfill_messages_by_chat_id``
and deletes orphaned rows whose message_id no longer appears
in the blob. Best-effort: errors are logged but never raised.
"""
try:
await self.backfill_messages_by_chat_id(chat_id, user_id, messages)
existing_map = await ChatMessages.get_messages_map_by_chat_id(chat_id)
if existing_map is not None:
orphaned_ids = set(existing_map.keys()) - set(messages.keys())
if orphaned_ids:
await ChatMessages.delete_message_ids_by_chat_id(chat_id, orphaned_ids)
except Exception as e:
log.warning('Failed to reconcile chat_message rows for chat %s: %s', chat_id, e)
async def get_messages_map_by_chat_id(self, id: str) -> dict | None:
"""Message map for walking history (see ``get_message_list``).
@ -614,8 +651,8 @@ class ChatTable:
return await self.update_chat_by_id(id, chat)
async def add_message_files_by_id_and_message_id(self, id: str, message_id: str, files: list[dict]) -> list[dict]:
async with get_async_db_context() as db:
chat = await self.get_chat_by_id(id, db=db)
async with get_async_db_context() as session:
chat = await self.get_chat_by_id(id, db=session)
if chat is None:
return None
@ -630,46 +667,49 @@ class ChatTable:
history['messages'][message_id]['files'] = message_files
chat['history'] = history
await self.update_chat_by_id(id, chat, db=db)
await self.update_chat_by_id(id, chat, db=session)
return message_files
async def insert_shared_chat_by_chat_id(self, chat_id: str, db: AsyncSession | None = None) -> ChatModel | None:
"""Create a shared snapshot for a chat. Returns the original chat with share_id set."""
from open_webui.models.shared_chats import SharedChats
async with get_async_db_context(db) as db:
chat = await db.get(Chat, chat_id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, chat_id)
if not chat:
return None
# If already shared, just update the existing snapshot
if chat.share_id:
return await self.update_shared_chat_by_chat_id(chat_id, db=db)
return await self.update_shared_chat_by_chat_id(chat_id, db=session)
shared = await SharedChats.create(chat_id, chat.user_id, db=db)
shared = await SharedChats.create(chat_id, chat.user_id, db=session)
if not shared:
return None
# Set share_id on the original chat
chat.share_id = shared.id
await db.commit()
await db.refresh(chat)
return ChatModel.model_validate(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat) # return the updated original
async def update_shared_chat_by_chat_id(self, chat_id: str, db: AsyncSession | None = None) -> ChatModel | None:
"""Re-snapshot the shared chat with current chat data."""
# refresh helper
async def update_shared_chat_by_chat_id(
self,
chat_id: str,
db: AsyncSession | None = None,
) -> ChatModel | None:
"""Refresh the shared snapshot with current chat content."""
from open_webui.models.shared_chats import SharedChats
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, chat_id)
if not chat or not chat.share_id:
return await self.insert_shared_chat_by_chat_id(chat_id, db=db)
await SharedChats.update(chat.share_id, db=db)
return ChatModel.model_validate(chat)
except Exception:
return None
async with get_async_db_context(db) as session:
record = await session.get(Chat, chat_id)
if not record or not record.share_id:
return await self.insert_shared_chat_by_chat_id(chat_id, db=session)
await SharedChats.update(record.share_id, db=session)
return ChatModel.model_validate(record)
# unreachable — context manager above always returns
return
async def delete_shared_chat_by_chat_id(self, chat_id: str, db: AsyncSession | None = None) -> bool:
"""Delete shared snapshot for a chat."""
@ -682,9 +722,9 @@ class ChatTable:
async def unarchive_all_chats_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(update(Chat).filter_by(user_id=user_id).values(archived=False))
await db.commit()
async with get_async_db_context(db) as session:
await session.execute(update(Chat).filter_by(user_id=user_id).values(archived=False))
await session.commit()
return True
except Exception:
return False
@ -693,45 +733,45 @@ class ChatTable:
self, id: str, share_id: str | None, db: AsyncSession | None = None
) -> ChatModel | None:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
chat.share_id = share_id
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat)
except Exception:
return None
async def toggle_chat_pinned_by_id(self, id: str, db: AsyncSession | None = None) -> ChatModel | None:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
chat.pinned = not chat.pinned
chat.updated_at = int(time.time())
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat)
except Exception:
return None
async def toggle_chat_archive_by_id(self, id: str, db: AsyncSession | None = None) -> ChatModel | None:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
chat.archived = not chat.archived
chat.folder_id = None
chat.updated_at = int(time.time())
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat)
except Exception:
return None
async def archive_all_chats_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(update(Chat).filter_by(user_id=user_id).values(archived=True))
await db.commit()
async with get_async_db_context(db) as session:
await session.execute(update(Chat).filter_by(user_id=user_id).values(archived=True))
await session.commit()
return True
except Exception:
return False
@ -744,7 +784,7 @@ class ChatTable:
limit: int = 50,
db: AsyncSession | None = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at).filter_by(
user_id=user_id, archived=True
)
@ -775,7 +815,7 @@ class ChatTable:
if limit:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.all()
return [
ChatTitleIdResponse.model_validate(
@ -811,7 +851,7 @@ class ChatTable:
limit: int = 50,
db: AsyncSession | None = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(
user_id=user_id
)
@ -841,7 +881,7 @@ class ChatTable:
if limit:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.all()
return [
ChatTitleIdResponse.model_validate(
@ -866,7 +906,7 @@ class ChatTable:
limit: int | None = None,
db: AsyncSession | None = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(
user_id=user_id
)
@ -887,7 +927,7 @@ class ChatTable:
if limit:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.all()
return [
@ -910,23 +950,29 @@ class ChatTable:
limit: int = 50,
db: AsyncSession | None = None,
) -> list[ChatModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
async with get_async_db_context(db) as session:
result = await session.execute(
select(Chat).filter(Chat.id.in_(chat_ids)).filter_by(archived=False).order_by(Chat.updated_at.desc())
)
all_chats = result.scalars().all()
return [ChatModel.model_validate(chat) for chat in all_chats]
async def get_chat_by_id(self, id: str, db: AsyncSession | None = None) -> ChatModel | None:
# retrieve conversation
async def get_chat_by_id(
self,
id: str,
db: AsyncSession | None = None,
) -> ChatModel | None:
"""Fetch a chat by PK, auto-sanitizing null bytes on read."""
try:
async with get_async_db_context(db) as db:
chat_item = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat_item = await session.get(Chat, id)
if chat_item is None:
return None
if self._sanitize_chat_row(chat_item):
await db.commit()
await db.refresh(chat_item)
await session.commit()
await session.refresh(chat_item)
return ChatModel.model_validate(chat_item)
except Exception:
@ -957,8 +1003,8 @@ class ChatTable:
self, id: str, user_id: str, db: AsyncSession | None = None
) -> ChatModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Chat).filter_by(id=id, user_id=user_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(Chat).filter_by(id=id, user_id=user_id))
chat = result.scalars().first()
return ChatModel.model_validate(chat) if chat else None
except Exception:
@ -970,8 +1016,8 @@ class ChatTable:
the full Chat row (which includes the potentially large JSON blob).
"""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(exists().where(and_(Chat.id == id, Chat.user_id == user_id))))
async with get_async_db_context(db) as session:
result = await session.execute(select(exists().where(and_(Chat.id == id, Chat.user_id == user_id))))
return result.scalar()
except Exception:
return False
@ -982,19 +1028,20 @@ class ChatTable:
JSON blob. Returns None if chat doesn't exist or doesn't belong to user.
"""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Chat.folder_id).filter_by(id=id, user_id=user_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(Chat.folder_id).filter_by(id=id, user_id=user_id))
row = result.first()
return row[0] if row else None
except Exception:
return None
async def get_chats(self, skip: int = 0, limit: int = 50, db: AsyncSession | None = None) -> list[ChatModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Chat).order_by(Chat.updated_at.desc()))
async with get_async_db_context(db) as session:
result = await session.execute(select(Chat).order_by(Chat.updated_at.desc()))
all_chats = result.scalars().all()
return [ChatModel.model_validate(chat) for chat in all_chats]
# list user conversations
async def get_chats_by_user_id(
self,
user_id: str,
@ -1003,7 +1050,7 @@ class ChatTable:
limit: int | None = None,
db: AsyncSession | None = None,
) -> ChatListResponse:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat).filter_by(user_id=user_id)
if filter:
@ -1025,7 +1072,7 @@ class ChatTable:
else:
stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip is not None:
@ -1033,7 +1080,7 @@ class ChatTable:
if limit is not None:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.scalars().all()
return ChatListResponse(
@ -1043,11 +1090,12 @@ class ChatTable:
}
)
# list pinned chats
async def get_pinned_chats_by_user_id(
self, user_id: str, db: AsyncSession | None = None
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
async with get_async_db_context(db) as session:
result = await session.execute(
select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at)
.filter_by(user_id=user_id, pinned=True, archived=False)
.order_by(Chat.updated_at.desc())
@ -1067,12 +1115,13 @@ class ChatTable:
]
async def get_archived_chats_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> list[ChatModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
async with get_async_db_context(db) as session:
result = await session.execute(
select(Chat).filter_by(user_id=user_id, archived=True).order_by(Chat.updated_at.desc())
)
return [ChatModel.model_validate(chat) for chat in result.scalars().all()]
# search user conversations
async def get_chats_by_user_id_and_search_text(
self,
user_id: str,
@ -1138,7 +1187,7 @@ class ChatTable:
search_text = ' '.join(search_text_words)
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat).filter(Chat.user_id == user_id)
if is_archived is not None:
@ -1161,7 +1210,7 @@ class ChatTable:
stmt = stmt.order_by(Chat.updated_at.desc(), Chat.id)
# Check if the database dialect is either 'sqlite' or 'postgresql'
bind = await db.connection()
bind = await session.connection()
dialect_name = bind.dialect.name
if dialect_name == 'sqlite':
# SQLite case: using JSON1 extension for JSON searching
@ -1259,7 +1308,7 @@ class ChatTable:
# Perform pagination at the SQL level
stmt = stmt.offset(skip).limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.scalars().all()
log.info(f'The number of chats: {len(all_chats)}')
@ -1275,7 +1324,7 @@ class ChatTable:
limit: int = 60,
db: AsyncSession | None = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = (
select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at)
.filter_by(folder_id=folder_id, user_id=user_id)
@ -1289,7 +1338,7 @@ class ChatTable:
if limit:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.all()
return [
ChatTitleIdResponse.model_validate(
@ -1307,7 +1356,7 @@ class ChatTable:
async def get_chats_by_folder_ids_and_user_id(
self, folder_ids: list[str], user_id: str, db: AsyncSession | None = None
) -> list[ChatModel]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = (
select(Chat)
.filter(Chat.folder_id.in_(folder_ids), Chat.user_id == user_id)
@ -1316,7 +1365,7 @@ class ChatTable:
.order_by(Chat.updated_at.desc())
)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.scalars().all()
return [ChatModel.model_validate(chat) for chat in all_chats]
@ -1324,13 +1373,13 @@ class ChatTable:
self, id: str, user_id: str, folder_id: str, db: AsyncSession | None = None
) -> ChatModel | None:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
chat.folder_id = folder_id
chat.updated_at = int(time.time())
chat.pinned = False
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat)
except Exception:
return None
@ -1338,12 +1387,12 @@ class ChatTable:
async def get_chat_tags_by_id_and_user_id(
self, id: str, user_id: str, db: AsyncSession | None = None
) -> list[TagModel]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat.meta).where(Chat.id == id)
result = await db.execute(stmt)
result = await session.execute(stmt)
meta = result.scalar_one_or_none()
tag_ids = (meta or {}).get('tags', [])
return await Tags.get_tags_by_ids_and_user_id(tag_ids, user_id, db=db)
return await Tags.get_tags_by_ids_and_user_id(tag_ids, user_id, db=session)
async def get_chat_list_by_user_id_and_tag_name(
self,
@ -1353,13 +1402,13 @@ class ChatTable:
limit: int = 50,
db: AsyncSession | None = None,
) -> list[ChatTitleIdResponse]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(Chat.id, Chat.title, Chat.updated_at, Chat.created_at, Chat.last_read_at).filter_by(
user_id=user_id
)
tag_id = tag_name.replace(' ', '_').lower()
bind = await db.connection()
bind = await session.connection()
dialect_name = bind.dialect.name
log.info(f'DB dialect name: {dialect_name}')
if dialect_name == 'sqlite':
@ -1380,7 +1429,7 @@ class ChatTable:
if limit:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
all_chats = result.all()
return [
ChatTitleIdResponse.model_validate(
@ -1401,15 +1450,15 @@ class ChatTable:
tag_id = tag_name.replace(' ', '_').lower()
await Tags.ensure_tags_exist([tag_name], user_id, db=db)
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
if tag_id not in chat.meta.get('tags', []):
chat.meta = {
**chat.meta,
'tags': list(set(chat.meta.get('tags', []) + [tag_id])),
}
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat)
except Exception:
return None
@ -1417,11 +1466,11 @@ class ChatTable:
async def count_chats_by_tag_name_and_user_id(
self, tag_name: str, user_id: str, db: AsyncSession | None = None
) -> int:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
stmt = select(func.count(Chat.id)).filter_by(user_id=user_id, archived=False)
tag_id = tag_name.replace(' ', '_').lower()
bind = await db.connection()
bind = await session.connection()
dialect_name = bind.dialect.name
if dialect_name == 'sqlite':
stmt = stmt.filter(
@ -1434,7 +1483,7 @@ class ChatTable:
else:
raise NotImplementedError(f'Unsupported dialect: {dialect_name}')
result = await db.execute(stmt)
result = await session.execute(stmt)
return result.scalar()
async def delete_orphan_tags_for_user(
@ -1454,19 +1503,19 @@ class ChatTable:
"""
if not tag_ids:
return
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
orphans = []
for tag_id in tag_ids:
count = await self.count_chats_by_tag_name_and_user_id(tag_id, user_id, db=db)
count = await self.count_chats_by_tag_name_and_user_id(tag_id, user_id, db=session)
if count <= threshold:
orphans.append(tag_id)
await Tags.delete_tags_by_ids_and_user_id(orphans, user_id, db=db)
await Tags.delete_tags_by_ids_and_user_id(orphans, user_id, db=session)
async def count_chats_by_folder_id_and_user_id(
self, folder_id: str, user_id: str, db: AsyncSession | None = None
) -> int:
async with get_async_db_context(db) as db:
result = await db.execute(select(func.count(Chat.id)).filter_by(user_id=user_id, folder_id=folder_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(func.count(Chat.id)).filter_by(user_id=user_id, folder_id=folder_id))
count = result.scalar()
log.info(f"Count of chats for folder '{folder_id}': {count}")
@ -1476,8 +1525,8 @@ class ChatTable:
self, id: str, user_id: str, tag_name: str, db: AsyncSession | None = None
) -> bool:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
async with get_async_db_context(db) as session:
chat = await session.get(Chat, id)
tags = chat.meta.get('tags', [])
tag_id = tag_name.replace(' ', '_').lower()
@ -1486,65 +1535,51 @@ class ChatTable:
**chat.meta,
'tags': list(set(tags)),
}
await db.commit()
return True
except Exception:
return False
async def delete_all_tags_by_id_and_user_id(self, id: str, user_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
chat = await db.get(Chat, id)
chat.meta = {
**chat.meta,
'tags': [],
}
await db.commit()
await session.commit()
return True
except Exception:
return False
async def delete_chat_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
await db.execute(delete(ChatMessage).filter_by(chat_id=id))
await db.execute(delete(Chat).filter_by(id=id))
await db.commit()
async with get_async_db_context(db) as session:
await session.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
await session.execute(delete(ChatMessage).filter_by(chat_id=id))
await session.execute(delete(Chat).filter_by(id=id))
await session.commit()
return True and await self.delete_shared_chat_by_chat_id(id, db=db)
return True and await self.delete_shared_chat_by_chat_id(id, db=session)
except Exception:
return False
async def delete_chat_by_id_and_user_id(self, id: str, user_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
await db.execute(delete(ChatMessage).filter_by(chat_id=id))
await db.execute(delete(Chat).filter_by(id=id, user_id=user_id))
await db.commit()
async with get_async_db_context(db) as session:
await session.execute(update(AutomationRun).filter_by(chat_id=id).values(chat_id=None))
await session.execute(delete(ChatMessage).filter_by(chat_id=id))
await session.execute(delete(Chat).filter_by(id=id, user_id=user_id))
await session.commit()
return True and await self.delete_shared_chat_by_chat_id(id, db=db)
return True and await self.delete_shared_chat_by_chat_id(id, db=session)
except Exception:
return False
async def delete_chats_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await self.delete_shared_chats_by_user_id(user_id, db=db)
async with get_async_db_context(db) as session:
await self.delete_shared_chats_by_user_id(user_id, db=session)
chat_id_subquery = select(Chat.id).filter_by(user_id=user_id).scalar_subquery()
await db.execute(
await session.execute(
update(AutomationRun)
.filter(AutomationRun.chat_id.in_(select(Chat.id).filter_by(user_id=user_id)))
.values(chat_id=None)
)
await db.execute(
await session.execute(
delete(ChatMessage).filter(ChatMessage.chat_id.in_(select(Chat.id).filter_by(user_id=user_id)))
)
await db.execute(delete(Chat).filter_by(user_id=user_id))
await db.commit()
await session.execute(delete(Chat).filter_by(user_id=user_id))
await session.commit()
return True
except Exception:
@ -1554,14 +1589,14 @@ class ChatTable:
self, user_id: str, folder_id: str, db: AsyncSession | None = None
) -> bool:
try:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
chat_ids_stmt = select(Chat.id).filter_by(user_id=user_id, folder_id=folder_id)
await db.execute(
await session.execute(
update(AutomationRun).filter(AutomationRun.chat_id.in_(chat_ids_stmt)).values(chat_id=None)
)
await db.execute(delete(ChatMessage).filter(ChatMessage.chat_id.in_(chat_ids_stmt)))
await db.execute(delete(Chat).filter_by(user_id=user_id, folder_id=folder_id))
await db.commit()
await session.execute(delete(ChatMessage).filter(ChatMessage.chat_id.in_(chat_ids_stmt)))
await session.execute(delete(Chat).filter_by(user_id=user_id, folder_id=folder_id))
await session.commit()
return True
except Exception:
@ -1575,11 +1610,11 @@ class ChatTable:
db: AsyncSession | None = None,
) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(
async with get_async_db_context(db) as session:
await session.execute(
update(Chat).filter_by(user_id=user_id, folder_id=folder_id).values(folder_id=new_folder_id)
)
await db.commit()
await session.commit()
return True
except Exception:
@ -1591,13 +1626,13 @@ class ChatTable:
from open_webui.models.shared_chats import SharedChats
try:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
# Delete shared_chat rows for this user's chats
await db.execute(delete(SharedChatTable).filter_by(user_id=user_id))
await session.execute(delete(SharedChatTable).filter_by(user_id=user_id))
# Clear share_id on all of this user's chats
await db.execute(update(Chat).filter_by(user_id=user_id).values(share_id=None))
await db.commit()
await session.execute(update(Chat).filter_by(user_id=user_id).values(share_id=None))
await session.commit()
return True
except Exception:
@ -1622,8 +1657,29 @@ class ChatTable:
if not file_ids:
return None
# Only link files the caller can read; blocks forging a chat_file row to another user's file.
from open_webui.models.files import Files
from open_webui.models.users import Users
from open_webui.utils.access_control.files import has_access_to_file
user = await Users.get_user_by_id(user_id, db=db)
accessible_file_ids = []
for file_id in file_ids:
file = await Files.get_file_by_id(file_id, db=db)
if not file:
continue
if (
file.user_id == user_id
or (user and user.role == 'admin')
or (user and await has_access_to_file(file_id, 'read', user, db=db))
):
accessible_file_ids.append(file_id)
file_ids = accessible_file_ids
if not file_ids:
return None
try:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
now = int(time.time())
chat_files = [
@ -1641,8 +1697,8 @@ class ChatTable:
results = [ChatFile(**chat_file.model_dump()) for chat_file in chat_files]
db.add_all(results)
await db.commit()
session.add_all(results)
await session.commit()
return chat_files
except Exception:
@ -1651,8 +1707,8 @@ class ChatTable:
async def get_chat_files_by_chat_id_and_message_id(
self, chat_id: str, message_id: str, db: AsyncSession | None = None
) -> list[ChatFileModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
async with get_async_db_context(db) as session:
result = await session.execute(
select(ChatFile).filter_by(chat_id=chat_id, message_id=message_id).order_by(ChatFile.created_at.asc())
)
all_chat_files = result.scalars().all()
@ -1660,17 +1716,17 @@ class ChatTable:
async def delete_chat_file(self, chat_id: str, file_id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(delete(ChatFile).filter_by(chat_id=chat_id, file_id=file_id))
await db.commit()
async with get_async_db_context(db) as session:
await session.execute(delete(ChatFile).filter_by(chat_id=chat_id, file_id=file_id))
await session.commit()
return True
except Exception:
return False
async def get_shared_chat_ids_by_file_id(self, file_id: str, db: AsyncSession | None = None) -> list[str]:
"""Return IDs of chats that contain this file and have an active share link."""
async with get_async_db_context(db) as db:
result = await db.execute(
async with get_async_db_context(db) as session:
result = await session.execute(
select(Chat.id)
.join(ChatFile, Chat.id == ChatFile.chat_id)
.filter(ChatFile.file_id == file_id, Chat.share_id.isnot(None))
@ -1680,25 +1736,25 @@ class ChatTable:
async def update_chat_tasks_by_id(self, id: str, tasks: list[dict]) -> ChatModel | None:
"""Update the tasks list on a chat."""
try:
async with get_async_db_context() as db:
chat = await db.get(Chat, id)
async with get_async_db_context() as session:
chat = await session.get(Chat, id)
if chat is None:
return None
chat.tasks = tasks
await db.commit()
await db.refresh(chat)
await session.commit()
await session.refresh(chat)
return ChatModel.model_validate(chat)
except Exception:
return None
async def get_chat_tasks_by_id(self, id: str) -> list[dict]:
"""Read the tasks list from a chat (lightweight column query)."""
async with get_async_db_context() as db:
result = await db.execute(select(Chat.tasks).filter_by(id=id))
async with get_async_db_context() as session:
result = await session.execute(select(Chat.tasks).filter_by(id=id))
row = result.first()
if row is None or row[0] is None:
return []
return row[0]
Chats = ChatTable()
Chats = ChatTable() # singleton chats repository

View file

@ -380,10 +380,39 @@ class FeedbackTable:
result = await db.execute(select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()))
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
async def get_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async def get_feedbacks_by_user_id(
self,
user_id: str,
skip: int = 0,
limit: int = 30,
db: Optional[AsyncSession] = None,
) -> FeedbackListResponse:
async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc()))
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
stmt = (
select(Feedback, User)
.join(User, Feedback.user_id == User.id)
.filter(Feedback.user_id == user_id)
.order_by(Feedback.updated_at.desc())
)
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip:
stmt = stmt.offset(skip)
if limit:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
items = result.all()
feedbacks = []
for feedback, user in items:
feedback_model = FeedbackModel.model_validate(feedback)
user_model = UserResponse.model_validate(user)
feedbacks.append(FeedbackUserResponse(**feedback_model.model_dump(), user=user_model))
return FeedbackListResponse(items=feedbacks, total=total)
async def update_feedback_by_id(
self, id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None

View file

@ -1,9 +1,11 @@
"""File upload models, forms, and database operations."""
from __future__ import annotations
import logging
import time
from typing import Optional
# local imports
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.utils.misc import sanitize_metadata
from pydantic import BaseModel, ConfigDict, model_validator
@ -12,26 +14,20 @@ from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
####################
# Files DB Schema
# What is written here bears witness. Let the testimony
# remain as it was given, and let none tamper with it.
####################
class File(Base):
class File(Base): # uploaded file record
__tablename__ = 'file'
id = Column(String, primary_key=True, unique=True)
user_id = Column(String)
user_id = Column(String, index=True) # owner user id
hash = Column(Text, nullable=True)
filename = Column(Text)
filename = Column(Text) # original upload filename
path = Column(Text, nullable=True)
data = Column(JSON, nullable=True)
meta = Column(JSON, nullable=True)
created_at = Column(BigInteger)
created_at = Column(BigInteger, index=True) # upload timestamp
updated_at = Column(BigInteger)
@ -52,11 +48,7 @@ class FileModel(BaseModel):
updated_at: int | None # timestamp in epoch
####################
# Forms
####################
# --- metadata structures ---
class FileMeta(BaseModel):
name: str | None = None
content_type: str | None = None
@ -157,16 +149,20 @@ class FilesTable:
return None
except Exception as e:
log.exception(f'Error inserting a new file: {e}')
return None
return None # insertion failed
async def get_file_by_id(self, id: str, db: AsyncSession | None = None) -> FileModel | None:
async def get_file_by_id(
self,
id: str,
db: AsyncSession | None = None,
) -> FileModel | None:
"""Look up a file by its primary key."""
try:
async with get_async_db_context(db) as db:
try:
file = await db.get(File, id)
return FileModel.model_validate(file) if file else None
except Exception:
file = await db.get(File, id)
if not file:
return None
return FileModel.model_validate(file)
except Exception:
return None
@ -397,6 +393,44 @@ class FilesTable:
except Exception:
return None
async def get_pending_files_for_knowledge(
self, knowledge_id: str, db: AsyncSession | None = None
) -> list[FileModelResponse]:
"""Return files still being processed for this knowledge base.
These are files uploaded with ``meta.data.knowledge_id`` set, whose
``data.status`` is still ``pending`` or ``processing``, and which
have not yet been added to the ``knowledge_file`` join table.
The JSON subscript syntax (``Column['key']['subkey'].as_string()``)
is supported by both SQLite (``json_extract``) and PostgreSQL
(``->>``/``->``).
"""
async with get_async_db_context(db) as db:
try:
# Lazy import to avoid circular dependency
from open_webui.models.knowledge import KnowledgeFile
# Subquery: file IDs already linked to this knowledge base
linked_ids = (
select(KnowledgeFile.file_id).filter(KnowledgeFile.knowledge_id == knowledge_id).correlate(None)
)
stmt = (
select(File)
.filter(
File.meta['data']['knowledge_id'].as_string() == knowledge_id,
File.data['status'].as_string().in_(['pending', 'processing']),
File.id.notin_(linked_ids),
)
.order_by(File.created_at.desc())
)
result = await db.execute(stmt)
return [FileModelResponse.model_validate(f, from_attributes=True) for f in result.scalars().all()]
except Exception as e:
log.warning(f'Error fetching pending files for knowledge {knowledge_id}: {e}')
return []
async def delete_file_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
async with get_async_db_context(db) as db:
try:
@ -418,4 +452,4 @@ class FilesTable:
return False
Files = FilesTable()
Files = FilesTable() # singleton files repository

View file

@ -1,9 +1,11 @@
"""Function (filter/action/pipe) models, forms, and database operations."""
from __future__ import annotations
import logging
import time
from typing import Optional
# local imports
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import UserModel, UserResponse, Users
from pydantic import BaseModel, ConfigDict
@ -12,29 +14,23 @@ from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
####################
# Functions DB Schema
# Each function here is a promise made. Let no promise
# go unkept, and let none be called who cannot answer.
####################
class Function(Base):
class Function(Base): # database table mapping
__tablename__ = 'function'
id = Column(String, primary_key=True, unique=True)
user_id = Column(String)
name = Column(Text)
type = Column(Text)
content = Column(Text)
meta = Column(JSONField)
valves = Column(JSONField)
is_active = Column(Boolean)
is_global = Column(Boolean)
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
user_id = Column(String, index=True) # creator user id
name = Column(Text, nullable=False) # function identifier
type = Column(Text, nullable=False) # function type (pipe, filter, etc.)
content = Column(Text, nullable=True) # Python source code
meta = Column(JSONField, nullable=True) # function metadata
valves = Column(JSONField, nullable=True) # function configuration valves
is_active = Column(Boolean, default=False) # function activation status
is_global = Column(Boolean) # if True, applied to every chat automatically
updated_at = Column(BigInteger) # epoch seconds
created_at = Column(BigInteger) # epoch seconds
__table_args__ = (Index('is_global_idx', 'is_global'),)
__table_args__ = (Index('is_global_idx', 'is_global'),) # speed up global-function lookups
class FunctionMeta(BaseModel):
@ -55,9 +51,10 @@ class FunctionModel(BaseModel):
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
model_config = ConfigDict(from_attributes=True) # allows ORM model binding
# --- form / schema definitions ---
class FunctionWithValvesModel(BaseModel):
id: str
user_id: str
@ -430,4 +427,4 @@ class FunctionsTable:
return False
Functions = FunctionsTable()
Functions = FunctionsTable() # singleton functions engine

View file

@ -360,18 +360,21 @@ class KnowledgeTable:
)
# Apply filename / content search
# Use ->> (as_string) instead of CAST(-> AS TEXT) to avoid
# PostgreSQL "invalid memory alloc request size" on large
# extracted-content rows (#24670).
content_text = File.data['content'].as_string()
search_filter = None
if filter:
q = filter.get('query')
if q:
search_filter = or_(
File.filename.ilike(f'%{q}%'),
content_text.ilike(f'%{q}%'),
)
if filter.get('include_content'):
# Use ->> (as_string) instead of CAST(-> AS TEXT)
# to avoid PostgreSQL "invalid memory alloc request
# size" on large extracted-content rows (#24670).
content_text = File.data['content'].as_string()
search_filter = or_(
File.filename.ilike(f'%{q}%'),
content_text.ilike(f'%{q}%'),
)
else:
search_filter = File.filename.ilike(f'%{q}%')
stmt = stmt.filter(search_filter)
# Order by file changes
@ -546,16 +549,19 @@ class KnowledgeTable:
if filter:
query_key = filter.get('query')
if query_key:
# Use ->> (as_string) instead of CAST(-> AS TEXT) to
# avoid PostgreSQL memory allocation failures on large
# content (#24670).
content_text = File.data['content'].as_string()
stmt = stmt.filter(
or_(
File.filename.ilike(f'%{query_key}%'),
content_text.ilike(f'%{query_key}%'),
if filter.get('include_content'):
# Use ->> (as_string) instead of CAST(-> AS TEXT)
# to avoid PostgreSQL memory allocation failures on
# large content (#24670).
content_text = File.data['content'].as_string()
stmt = stmt.filter(
or_(
File.filename.ilike(f'%{query_key}%'),
content_text.ilike(f'%{query_key}%'),
)
)
)
else:
stmt = stmt.filter(File.filename.ilike(f'%{query_key}%'))
view_option = filter.get('view_option')
if view_option == 'created':
@ -692,11 +698,18 @@ class KnowledgeTable:
except Exception:
return False
async def reset_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]:
async def reset_knowledge_by_id(
self, id: str, include_directories: bool = True, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeModel]:
try:
async with get_async_db_context(db) as db:
# Delete all knowledge_file entries for this knowledge_id
await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=id))
# Delete all directories if requested
if include_directories:
await db.execute(delete(KnowledgeDirectory).filter_by(knowledge_id=id))
await db.commit()
# Update the knowledge entry's updated_at timestamp
@ -814,9 +827,7 @@ class KnowledgeTable:
) -> list[KnowledgeDirectoryModel]:
"""List directories at a given level (parent_id=None for root)."""
async with get_async_db_context(db) as db:
stmt = select(KnowledgeDirectory).filter(
KnowledgeDirectory.knowledge_id == knowledge_id
)
stmt = select(KnowledgeDirectory).filter(KnowledgeDirectory.knowledge_id == knowledge_id)
if parent_id:
stmt = stmt.filter(KnowledgeDirectory.parent_id == parent_id)
else:
@ -854,10 +865,7 @@ class KnowledgeTable:
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
.filter(KnowledgeFile.knowledge_id == knowledge_id)
)
return [
(FileModel.model_validate(file), dir_id)
for file, dir_id in result.all()
]
return [(FileModel.model_validate(file), dir_id) for file, dir_id in result.all()]
except Exception:
return []
@ -865,9 +873,7 @@ class KnowledgeTable:
self, directory_id: str, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeDirectoryModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(KnowledgeDirectory).filter_by(id=directory_id)
)
result = await db.execute(select(KnowledgeDirectory).filter_by(id=directory_id))
directory = result.scalars().first()
return KnowledgeDirectoryModel.model_validate(directory) if directory else None
@ -887,9 +893,7 @@ class KnowledgeTable:
while current_id and current_id not in seen:
seen.add(current_id)
result = await db.execute(
select(KnowledgeDirectory).filter_by(id=current_id)
)
result = await db.execute(select(KnowledgeDirectory).filter_by(id=current_id))
directory = result.scalars().first()
if not directory:
break
@ -908,9 +912,7 @@ class KnowledgeTable:
async with get_async_db_context(db) as db:
try:
await db.execute(
update(KnowledgeDirectory)
.filter_by(id=directory_id)
.values(name=name, updated_at=int(time.time()))
update(KnowledgeDirectory).filter_by(id=directory_id).values(name=name, updated_at=int(time.time()))
)
await db.commit()
return await self.get_directory_by_id(directory_id, db=db)
@ -936,9 +938,7 @@ class KnowledgeTable:
if current == directory_id:
return None # Would create a cycle
seen.add(current)
result = await db.execute(
select(KnowledgeDirectory.parent_id).filter_by(id=current)
)
result = await db.execute(select(KnowledgeDirectory.parent_id).filter_by(id=current))
row = result.first()
current = row[0] if row else None
@ -986,9 +986,7 @@ class KnowledgeTable:
async with get_async_db_context(db) as db:
try:
# Get the directory to find its parent
result = await db.execute(
select(KnowledgeDirectory).filter_by(id=directory_id)
)
result = await db.execute(select(KnowledgeDirectory).filter_by(id=directory_id))
directory = result.scalars().first()
if not directory:
return False
@ -998,9 +996,7 @@ class KnowledgeTable:
if move_files_to_parent:
# Move files in this directory to its parent (or root)
await db.execute(
update(KnowledgeFile)
.filter_by(directory_id=directory_id)
.values(directory_id=parent_id)
update(KnowledgeFile).filter_by(directory_id=directory_id).values(directory_id=parent_id)
)
# Recursively move files from all subdirectories too
await self._move_files_from_subtree(directory_id, parent_id, db=db)
@ -1009,9 +1005,7 @@ class KnowledgeTable:
await self._delete_files_in_subtree(directory_id, db=db)
# CASCADE on parent_id will handle deleting subdirectories
await db.execute(
delete(KnowledgeDirectory).filter_by(id=directory_id)
)
await db.execute(delete(KnowledgeDirectory).filter_by(id=directory_id))
await db.commit()
return True
except Exception as e:
@ -1025,16 +1019,12 @@ class KnowledgeTable:
db: AsyncSession,
) -> None:
"""Recursively move all files from a directory subtree to the target."""
result = await db.execute(
select(KnowledgeDirectory.id).filter_by(parent_id=directory_id)
)
result = await db.execute(select(KnowledgeDirectory.id).filter_by(parent_id=directory_id))
child_ids = [row[0] for row in result.all()]
for child_id in child_ids:
await db.execute(
update(KnowledgeFile)
.filter_by(directory_id=child_id)
.values(directory_id=target_directory_id)
update(KnowledgeFile).filter_by(directory_id=child_id).values(directory_id=target_directory_id)
)
await self._move_files_from_subtree(child_id, target_directory_id, db=db)
@ -1044,12 +1034,8 @@ class KnowledgeTable:
db: AsyncSession,
) -> None:
"""Recursively delete all files from a directory subtree."""
await db.execute(
delete(KnowledgeFile).filter_by(directory_id=directory_id)
)
result = await db.execute(
select(KnowledgeDirectory.id).filter_by(parent_id=directory_id)
)
await db.execute(delete(KnowledgeFile).filter_by(directory_id=directory_id))
result = await db.execute(select(KnowledgeDirectory.id).filter_by(parent_id=directory_id))
child_ids = [row[0] for row in result.all()]
for child_id in child_ids:
await self._delete_files_in_subtree(child_id, db=db)

View file

@ -1,3 +1,5 @@
"""Long-term memory storage for per-user context recall."""
from __future__ import annotations
import time
@ -9,36 +11,28 @@ from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, String, Text, delete, select
from sqlalchemy.ext.asyncio import AsyncSession
####################
# Memory DB Schema
# What was learned at cost should not need to be paid
# for again. Let the memory hold.
####################
class Memory(Base): # user memory store
"""Stores user-created memory entries linked to a vector collection."""
class Memory(Base):
__tablename__ = 'memory'
id = Column(String, primary_key=True, unique=True)
user_id = Column(String, index=True)
content = Column(Text)
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
content = Column(Text) # free-form text learned from conversation
updated_at = Column(BigInteger) # epoch seconds
created_at = Column(BigInteger) # epoch seconds
class MemoryModel(BaseModel):
"""Pydantic mirror of the Memory table row."""
id: str
user_id: str
content: str
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
####################
# Forms
####################
model_config = ConfigDict(from_attributes=True) # allows ORM mapping
class MemoriesTable:
@ -48,26 +42,20 @@ class MemoriesTable:
content: str,
db: AsyncSession | None = None,
) -> MemoryModel | None:
"""Persist a new memory entry and return the created model."""
async with get_async_db_context(db) as db:
id = str(uuid.uuid4())
memory = MemoryModel(
**{
'id': id,
'user_id': user_id,
'content': content,
'created_at': int(time.time()),
'updated_at': int(time.time()),
}
now = int(time.time())
record = Memory(
id=str(uuid.uuid4()),
user_id=user_id,
content=content,
created_at=now,
updated_at=now,
)
result = Memory(**memory.model_dump())
db.add(result)
db.add(record)
await db.commit()
await db.refresh(result)
if result:
return MemoryModel.model_validate(result)
else:
return None
await db.refresh(record)
return MemoryModel.model_validate(record) if record else None
async def update_memory_by_id_and_user_id(
self,
@ -143,15 +131,13 @@ class MemoriesTable:
try:
memory = await db.get(Memory, id)
if not memory or memory.user_id != user_id:
return None
return False
# Delete the memory
await db.delete(memory)
await db.commit()
return True
except Exception:
return False
Memories = MemoriesTable()
Memories = MemoriesTable() # user memory registry

View file

@ -9,40 +9,53 @@ from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups
from open_webui.models.users import User, UserModel, UserResponse, Users
from pydantic import BaseModel, ConfigDict, Field, model_validator
from open_webui.utils.validate import validate_profile_image_url
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from sqlalchemy import BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, update
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
####################
# Models DB Schema
# A misconfigured model wastes the time of everyone
# who trusts it. Let what is set here be set with care.
####################
# Track invalid profile_image_url values we've already warned about so we
# don't flood the logs on every DB read (the validator fires per-row).
_warned_profile_urls: set[str] = set()
# --- Models DB Schema ---
# ModelParams is a model for the data stored in the params field of the Model table
class ModelParams(BaseModel):
"""Parameters for model inference (temperature, top_p, etc.)."""
model_config = ConfigDict(extra='allow')
pass
# ModelMeta is a model for the data stored in the meta field of the Model table
class ModelMeta(BaseModel):
profile_image_url: str | None = '/static/favicon.png'
description: str | None = None
"""
User-facing description of the model.
"""
"""Metadata for a workspace model entry (profile, description, tags, capabilities)."""
profile_image_url: str | None = None
description: str | None = Field(default=None, description='User-facing description of the model.')
capabilities: dict | None = None
model_config = ConfigDict(extra='allow')
@field_validator('profile_image_url', mode='before')
@classmethod
def check_profile_image_url(cls, v: str | None) -> str | None:
if v is None:
return v
try:
return validate_profile_image_url(v)
except ValueError:
if v not in _warned_profile_urls:
_warned_profile_urls.add(v)
log.warning(
'Clearing invalid profile_image_url stored in DB (likely a legacy SVG data-URI): %.80s…',
v,
)
return None
@model_validator(mode='before')
@classmethod
def normalize_tags(cls, data):
@ -60,38 +73,19 @@ class ModelMeta(BaseModel):
class Model(Base):
"""Workspace model entry — wraps an upstream LLM with custom params and metadata."""
__tablename__ = 'model'
id = Column(Text, primary_key=True, unique=True)
"""
The model's id as used in the API. If set to an existing model, it will override the model.
"""
user_id = Column(Text)
base_model_id = Column(Text, nullable=True)
"""
An optional pointer to the actual model that should be used when proxying requests.
"""
name = Column(Text)
"""
The human-readable display name of the model.
"""
params = Column(JSONField)
"""
Holds a JSON encoded blob of parameters, see `ModelParams`.
"""
meta = Column(JSONField)
"""
Holds a JSON encoded blob of metadata, see `ModelMeta`.
"""
is_active = Column(Boolean, default=True)
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
id = Column(Text, primary_key=True, unique=True) # API model identifier; overrides built-in when matching
user_id = Column(Text) # owner
base_model_id = Column(Text, nullable=True) # actual upstream model for proxied requests
name = Column(Text) # human-readable display name
params = Column(JSONField) # see ModelParams
meta = Column(JSONField) # see ModelMeta
is_active = Column(Boolean, default=True) # soft-disable toggle
updated_at = Column(BigInteger) # epoch seconds
created_at = Column(BigInteger) # epoch seconds
class ModelModel(BaseModel):
@ -109,12 +103,9 @@ class ModelModel(BaseModel):
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
####################
# Forms
####################
model_config = ConfigDict(
from_attributes=True,
)
class ModelUserResponse(ModelModel):
@ -199,10 +190,13 @@ class ModelsTable:
all_models = result.scalars().all()
model_ids = [model.id for model in all_models]
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)
for model in all_models
]
models: list[ModelModel] = []
for model in all_models:
try:
models.append(await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db))
except Exception as exc:
log.error('Skipping model %r during get_all_models due to error: %s', model.id, exc)
return models
async def get_models(self, db: AsyncSession | None = None) -> list[ModelUserResponse]:
async with get_async_db_context(db) as db:
@ -595,4 +589,4 @@ class ModelsTable:
return []
Models = ModelsTable()
Models = ModelsTable() # singleton model registry

View file

@ -151,13 +151,19 @@ class PromptHistoryTable:
self,
from_id: str,
to_id: str,
prompt_id: str,
db: Optional[AsyncSession] = None,
) -> Optional[dict]:
"""Compute diff between two history entries."""
async with get_async_db_context(db) as db:
result_from = await db.execute(select(PromptHistory).filter(PromptHistory.id == from_id))
# Bind both entries to the authorized prompt; an unbound id reads another prompt's snapshot.
result_from = await db.execute(
select(PromptHistory).filter(PromptHistory.id == from_id, PromptHistory.prompt_id == prompt_id)
)
from_entry = result_from.scalars().first()
result_to = await db.execute(select(PromptHistory).filter(PromptHistory.id == to_id))
result_to = await db.execute(
select(PromptHistory).filter(PromptHistory.id == to_id, PromptHistory.prompt_id == prompt_id)
)
to_entry = result_to.scalars().first()
if not from_entry or not to_entry:
@ -203,11 +209,13 @@ class PromptHistoryTable:
async def delete_history_entry(
self,
history_id: str,
prompt_id: str,
db: Optional[AsyncSession] = None,
) -> bool:
"""Delete a history entry and reparent its children to grandparent."""
async with get_async_db_context(db) as db:
result = await db.execute(select(PromptHistory).filter_by(id=history_id))
# Bind to the authorized prompt; an unbound id deletes another prompt's history.
result = await db.execute(select(PromptHistory).filter_by(id=history_id, prompt_id=prompt_id))
entry = result.scalars().first()
if not entry:
return False

View file

@ -1,10 +1,15 @@
"""Prompt template models, forms, and database operations."""
from __future__ import annotations
import json
import logging
import time
import uuid
from typing import Optional
log = logging.getLogger(__name__)
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups
@ -14,23 +19,19 @@ from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession
####################
# Prompts DB Schema
# Every word here was weighed before it was set down.
# Let the weight not be wasted when it is spoken aloud.
####################
class Prompt(Base): # versioned template
"""Slash-command prompt with history tracking and access control."""
class Prompt(Base):
__tablename__ = 'prompt'
id = Column(Text, primary_key=True)
command = Column(String, unique=True, index=True)
user_id = Column(String)
user_id = Column(String, index=True) # owner user id
name = Column(Text)
content = Column(Text)
data = Column(JSON, nullable=True)
meta = Column(JSON, nullable=True)
content = Column(Text) # the prompt template body
data = Column(JSON, nullable=True) # structured prompt parameters
meta = Column(JSON, nullable=True) # freeform metadata (description, etc.)
tags = Column(JSON, nullable=True)
is_active = Column(Boolean, default=True)
version_id = Column(Text, nullable=True) # Points to active history entry
@ -53,10 +54,10 @@ class PromptModel(BaseModel):
updated_at: int | None = None
access_grants: list[AccessGrantModel] = Field(default_factory=list)
model_config = ConfigDict(from_attributes=True)
model_config = ConfigDict(from_attributes=True) # allows ORM model binding
####################
# --- form / schema definitions ---
# Forms
####################
@ -114,101 +115,112 @@ class PromptsTable:
now = int(time.time())
prompt_id = str(uuid.uuid4())
prompt = PromptModel(
id=prompt_id,
user_id=user_id,
command=form_data.command,
name=form_data.name,
content=form_data.content,
data=form_data.data or {},
meta=form_data.meta or {},
tags=form_data.tags or [],
access_grants=[],
is_active=True,
created_at=now,
updated_at=now,
)
async with get_async_db_context(db) as session:
try:
record = Prompt(
id=prompt_id,
user_id=user_id,
command=form_data.command,
name=form_data.name,
content=form_data.content,
data=form_data.data or {},
meta=form_data.meta or {},
tags=form_data.tags or [],
is_active=True,
created_at=now,
updated_at=now,
)
session.add(record)
await session.commit()
await session.refresh(record) # populate generated defaults
try:
async with get_async_db_context(db) as db:
result = Prompt(**prompt.model_dump(exclude={'access_grants'}))
db.add(result)
await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
await AccessGrants.set_access_grants(
'prompt',
prompt_id,
form_data.access_grants,
db=session,
) # persist sharing rules
if result:
current_access_grants = await self._get_access_grants(prompt_id, db=db)
snapshot = {
'name': form_data.name,
'content': form_data.content,
'command': form_data.command,
'data': form_data.data or {},
'meta': form_data.meta or {},
'tags': form_data.tags or [],
'access_grants': [grant.model_dump() for grant in current_access_grants],
}
history_entry = await PromptHistories.create_history_entry(
prompt_id=prompt_id,
snapshot=snapshot,
user_id=user_id,
parent_id=None, # Initial commit has no parent
commit_message=form_data.commit_message or 'Initial version',
db=db,
)
# Set the initial version as the production version
if history_entry:
result.version_id = history_entry.id
await db.commit()
await db.refresh(result)
return await self._to_prompt_model(result, db=db)
else:
if not record: # shouldn't happen, but guard anyway
return None
except Exception:
return None
# Build the initial version snapshot.
grants = await self._get_access_grants(prompt_id, db=session)
snapshot = {
'name': form_data.name,
'content': form_data.content,
'command': form_data.command,
'data': form_data.data or {},
'meta': form_data.meta or {},
'tags': form_data.tags or [],
'access_grants': [g.model_dump() for g in grants],
}
history_entry = await PromptHistories.create_history_entry(
prompt_id=prompt_id,
snapshot=snapshot,
user_id=user_id,
parent_id=None,
commit_message=form_data.commit_message or 'Initial version',
db=session,
) # creates the first version entry
# Pin the initial history entry as the production version.
if history_entry:
record.version_id = history_entry.id
await session.commit()
await session.refresh(record) # re-read version_id
return await self._to_prompt_model(record, db=session)
except Exception as e:
log.exception('Error creating prompt: %s', e)
return None
async def get_prompt_by_id(self, prompt_id: str, db: AsyncSession | None = None) -> PromptModel | None:
"""Get prompt by UUID."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if prompt:
return await self._to_prompt_model(prompt, db=db)
return None
except Exception:
return None
async with get_async_db_context(db) as session:
result = await session.execute(
select(Prompt).filter_by(id=prompt_id),
)
prompt = result.scalars().first() # None when not found
if not prompt:
return None
return await self._to_prompt_model(prompt, db=session)
except Exception: # connection / integrity error
return
async def get_prompt_by_command(self, command: str, db: AsyncSession | None = None) -> PromptModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(command=command))
prompt = result.scalars().first()
if prompt:
return await self._to_prompt_model(prompt, db=db)
return None
except Exception:
return None
"""Look up a prompt by its unique slash-command string."""
async with get_async_db_context(db) as session:
match = (await session.execute(select(Prompt).where(Prompt.command == command))).scalars().first()
if match is None:
return
return await self._to_prompt_model(match, db=session)
# --- context manager always returns above ---
return
async def get_prompts(self, db: AsyncSession | None = None) -> list[PromptUserResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc())
"""Return all active prompts ordered by most recently updated."""
async with get_async_db_context(db) as session:
active = (
(
await session.execute(
select(Prompt).where(Prompt.is_active.is_(True)).order_by(Prompt.updated_at.desc())
)
)
.scalars()
.all()
)
all_prompts = result.scalars().all()
user_ids = list(set(prompt.user_id for prompt in all_prompts))
prompt_ids = [prompt.id for prompt in all_prompts]
user_ids = list(set(p.user_id for p in active))
prompt_ids = [p.id for p in active]
users = await 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 = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
users = await Users.get_users_by_user_ids(user_ids, db=session) if user_ids else []
users_dict = {u.id: u for u in users}
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=session)
prompts = []
for prompt in all_prompts:
for prompt in active:
user = users_dict.get(prompt.user_id)
prompts.append(
PromptUserResponse.model_validate(
@ -217,7 +229,7 @@ class PromptsTable:
await self._to_prompt_model(
prompt,
access_grants=grants_map.get(prompt.id, []),
db=db,
db=session,
)
).model_dump(),
'user': user.model_dump() if user else None,
@ -230,8 +242,8 @@ class PromptsTable:
async def get_prompts_by_user_id(
self, user_id: str, permission: str = 'write', db: AsyncSession | None = None
) -> list[PromptUserResponse]:
async with get_async_db_context(db) as db:
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
async with get_async_db_context(db) as session:
user_groups = await Groups.get_groups_by_member_id(user_id, db=session)
user_group_ids = [group.id for group in user_groups]
query = select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc())
@ -244,7 +256,7 @@ class PromptsTable:
permission=permission,
)
result = await db.execute(query)
result = await session.execute(query)
accessible_prompts = result.scalars().all()
if not accessible_prompts:
@ -253,9 +265,9 @@ class PromptsTable:
prompt_ids = [p.id for p in accessible_prompts]
owner_ids = list({p.user_id for p in accessible_prompts})
users = await Users.get_users_by_user_ids(owner_ids, db=db)
users = await Users.get_users_by_user_ids(owner_ids, db=session)
users_dict = {u.id: u for u in users}
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=session)
results = []
for prompt in accessible_prompts:
@ -284,7 +296,7 @@ class PromptsTable:
limit: int = 30,
db: AsyncSession | None = None,
) -> PromptListResponse:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
# Join with User table for user filtering and sorting
query = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
@ -319,7 +331,7 @@ class PromptsTable:
tag = filter.get('tag')
if tag:
bind = await db.connection()
bind = await session.connection()
dialect_name = bind.dialect.name
tag_lower = tag.lower()
@ -367,7 +379,7 @@ class PromptsTable:
query = query.order_by(Prompt.updated_at.desc())
# Count BEFORE pagination
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
count_result = await session.execute(select(func.count()).select_from(query.subquery()))
total = count_result.scalar()
if skip:
@ -375,11 +387,11 @@ class PromptsTable:
if limit:
query = query.limit(limit)
result = await db.execute(query)
result = await session.execute(query)
items = result.all()
prompt_ids = [prompt.id for prompt, _ in items]
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=session)
prompts = []
for prompt, user in items:
@ -405,16 +417,18 @@ class PromptsTable:
user_id: str,
db: AsyncSession | None = None,
) -> PromptModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(command=command))
if not command:
return None
try: # database transaction
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(command=command))
prompt = result.scalars().first()
if not prompt:
return None
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db)
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=session)
parent_id = latest_history.id if latest_history else None
current_access_grants = await self._get_access_grants(prompt.id, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=session)
# Check if content changed to decide on history creation
content_changed = (
@ -430,10 +444,10 @@ class PromptsTable:
prompt.meta = form_data.meta or prompt.meta
prompt.updated_at = int(time.time())
if form_data.access_grants is not None:
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=db)
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=session)
current_access_grants = await self._get_access_grants(prompt.id, db=session)
await db.commit()
await session.commit()
# Create history entry only if content changed
if content_changed:
@ -458,9 +472,9 @@ class PromptsTable:
# Set as production if flag is True (default)
if form_data.is_production and history_entry:
prompt.version_id = history_entry.id
await db.commit()
await session.commit()
return await self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=session)
except Exception:
return None
@ -472,15 +486,15 @@ class PromptsTable:
db: AsyncSession | None = None,
) -> PromptModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if not prompt:
return None
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db)
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=session)
parent_id = latest_history.id if latest_history else None
current_access_grants = await self._get_access_grants(prompt.id, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=session)
# Check if content changed to decide on history creation
content_changed = (
@ -502,12 +516,12 @@ class PromptsTable:
prompt.tags = form_data.tags
if form_data.access_grants is not None:
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
current_access_grants = await self._get_access_grants(prompt.id, db=db)
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=session)
current_access_grants = await self._get_access_grants(prompt.id, db=session)
prompt.updated_at = int(time.time())
await db.commit()
await session.commit()
# Create history entry only if content changed
if content_changed:
@ -533,9 +547,9 @@ class PromptsTable:
# Set as production if flag is True (default)
if form_data.is_production and history_entry:
prompt.version_id = history_entry.id
await db.commit()
await session.commit()
return await self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=session)
except Exception:
return None
@ -549,8 +563,8 @@ class PromptsTable:
) -> PromptModel | None:
"""Update only name, command, and tags (no history created)."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if not prompt:
return None
@ -562,9 +576,9 @@ class PromptsTable:
prompt.tags = tags
prompt.updated_at = int(time.time())
await db.commit()
await session.commit()
return await self._to_prompt_model(prompt, db=db)
return await self._to_prompt_model(prompt, db=session)
except Exception:
return None
@ -576,15 +590,16 @@ class PromptsTable:
) -> PromptModel | None:
"""Set the active version of a prompt and restore content from that version's snapshot."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if not prompt:
return None
history_entry = await PromptHistories.get_history_entry_by_id(version_id, db=db)
history_entry = await PromptHistories.get_history_entry_by_id(version_id, db=session)
if not history_entry:
# Reject a version_id from another prompt; restoring it would copy a foreign snapshot in.
if not history_entry or history_entry.prompt_id != prompt_id:
return None
# Restore prompt content from the snapshot
@ -599,24 +614,31 @@ class PromptsTable:
prompt.version_id = version_id
prompt.updated_at = int(time.time())
await db.commit()
await session.commit()
return await self._to_prompt_model(prompt, db=db)
except Exception:
return await self._to_prompt_model(prompt, db=session)
except Exception as e: # connection error
log.error(f'Failed to restore prompt version: {e}')
return None # restoration failed
async def toggle_prompt_active(
self,
prompt_id: str,
db: AsyncSession | None = None,
) -> PromptModel | None:
"""Flip the is_active flag on a prompt."""
if not prompt_id:
return None
async def toggle_prompt_active(self, prompt_id: str, db: AsyncSession | None = None) -> PromptModel | None:
"""Toggle the is_active flag on a prompt."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
try: # activation state toggle
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if prompt:
prompt.is_active = not prompt.is_active
prompt.updated_at = int(time.time())
await db.commit()
await db.refresh(prompt)
return await self._to_prompt_model(prompt, db=db)
await session.commit()
await session.refresh(prompt)
return await self._to_prompt_model(prompt, db=session)
return None
except Exception:
return None
@ -624,15 +646,15 @@ class PromptsTable:
async def delete_prompt_by_command(self, command: str, db: AsyncSession | None = None) -> bool:
"""Permanently delete a prompt and its history."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(command=command))
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(command=command))
prompt = result.scalars().first()
if prompt:
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
await AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=session)
await AccessGrants.revoke_all_access('prompt', prompt.id, db=session)
await db.delete(prompt)
await db.commit()
await session.delete(prompt)
await session.commit()
return True
return False
except Exception:
@ -641,24 +663,25 @@ class PromptsTable:
async def delete_prompt_by_id(self, prompt_id: str, db: AsyncSession | None = None) -> bool:
"""Permanently delete a prompt and its history."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt).filter_by(id=prompt_id))
prompt = result.scalars().first()
if prompt:
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
await AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=session)
await AccessGrants.revoke_all_access('prompt', prompt.id, db=session)
await db.delete(prompt)
await db.commit()
await session.delete(prompt)
await session.commit()
return True
return False
except Exception:
return False
except Exception as err:
log.error(f'Failed to delete prompt: {err}')
return False # deletion failed
async def get_tags(self, db: AsyncSession | None = None) -> list[str]:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Prompt.tags).filter(Prompt.is_active == True))
async with get_async_db_context(db) as session:
result = await session.execute(select(Prompt.tags).filter(Prompt.is_active == True))
tags = set()
for (tag_list,) in result.all():
if tag_list:
@ -671,8 +694,8 @@ class PromptsTable:
async def get_tags_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> list[str]:
try:
async with get_async_db_context(db) as db:
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
async with get_async_db_context(db) as session:
user_groups = await Groups.get_groups_by_member_id(user_id, db=session)
user_group_ids = [group.id for group in user_groups]
query = select(Prompt.tags).filter(Prompt.is_active == True)
@ -685,7 +708,7 @@ class PromptsTable:
permission='read',
)
result = await db.execute(query)
result = await session.execute(query)
tags = set()
for (tag_list,) in result.all():
if tag_list:
@ -697,4 +720,4 @@ class PromptsTable:
return []
Prompts = PromptsTable()
Prompts = PromptsTable() # singleton prompts registry

View file

@ -1,10 +1,12 @@
"""Tag models and database operations."""
from __future__ import annotations
import logging
import time
import uuid
from typing import Optional
# local imports
from open_webui.internal.db import Base, JSONField, get_async_db_context
from pydantic import BaseModel, ConfigDict
from sqlalchemy import JSON, BigInteger, Column, Index, PrimaryKeyConstraint, String, delete, select
@ -13,16 +15,11 @@ from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
####################
# Tag DB Schema
# To name a thing is to claim it. The creator has
# already named everything stored in this table.
####################
class Tag(Base):
class Tag(Base): # database table mapping for tag entity
__tablename__ = 'tag'
id = Column(String)
name = Column(String)
user_id = Column(String)
name = Column(String, index=True) # tag label
user_id = Column(String, index=True) # user identifier
meta = Column(JSON, nullable=True)
__table_args__ = (
@ -39,10 +36,10 @@ class TagModel(BaseModel):
name: str
user_id: str
meta: dict | None = None
model_config = ConfigDict(from_attributes=True)
model_config = ConfigDict(from_attributes=True) # allows ORM model binding
####################
# --- tag schema forms ---
# Forms
####################
@ -53,22 +50,24 @@ class TagChatIdForm(BaseModel):
class TagTable:
async def insert_new_tag(self, name: str, user_id: str, db: AsyncSession | None = None) -> TagModel | None:
async def insert_new_tag(
self,
name: str,
user_id: str,
db: AsyncSession | None = None,
) -> TagModel | None:
"""Create a new tag, deriving the id from the name."""
async with get_async_db_context(db) as db:
id = name.replace(' ', '_').lower()
tag = TagModel(**{'id': id, 'user_id': user_id, 'name': name})
tag_id = name.replace(' ', '_').lower()
try:
result = Tag(**tag.model_dump())
db.add(result)
record = Tag(id=tag_id, user_id=user_id, name=name)
db.add(record)
await db.commit()
await db.refresh(result)
if result:
return TagModel.model_validate(result)
else:
return None
await db.refresh(record)
return TagModel.model_validate(record) if record else None
except Exception as e:
log.exception(f'Error inserting a new tag: {e}')
return None
log.exception('Error inserting tag %r: %s', name, e)
return None # insertion failed
async def get_tag_by_name_and_user_id(
self, name: str, user_id: str, db: AsyncSession | None = None
@ -137,4 +136,4 @@ class TagTable:
await db.commit()
Tags = TagTable()
Tags = TagTable() # singleton tag repository

View file

@ -1,9 +1,11 @@
"""Tool models, forms, and database operations."""
from __future__ import annotations
import logging
import time
from typing import Optional
# local imports
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups
@ -14,26 +16,20 @@ from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
####################
# Tools DB Schema
# A tool that fails silently is worse than one that
# refuses outright. Let each one here be honest in its work.
####################
class Tool(Base):
class Tool(Base): # database table definition
__tablename__ = 'tool'
id = Column(String, primary_key=True, unique=True)
user_id = Column(String)
name = Column(Text)
content = Column(Text)
specs = Column(JSONField)
meta = Column(JSONField)
valves = Column(JSONField)
user_id = Column(String, index=True) # owner user id
name = Column(Text) # human-readable label
content = Column(Text) # Python source code
specs = Column(JSONField) # OpenAPI-style function specs
meta = Column(JSONField) # description, manifest, etc.
valves = Column(JSONField) # admin-configurable runtime parameters
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
updated_at = Column(BigInteger, nullable=False) # modification timestamp
created_at = Column(BigInteger, index=True) # creation timestamp
class ToolMeta(BaseModel):
@ -53,10 +49,10 @@ class ToolModel(BaseModel):
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
model_config = ConfigDict(from_attributes=True) # enables ORM mapping
####################
# --- tool request forms ---
# Forms
####################
@ -141,16 +137,36 @@ class ToolsTable:
return None
except Exception as e:
log.exception(f'Error creating a new tool: {e}')
return None
return None # creation failed
async def get_tool_by_id(self, id: str, db: AsyncSession | None = None) -> ToolModel | None:
try:
async with get_async_db_context(db) as db:
tool = await db.get(Tool, id)
return await self._to_tool_model(tool, db=db) if tool else None
async def get_tool_by_id(
self,
id: str,
db: AsyncSession | None = None,
) -> ToolModel | None:
"""Fetch a single tool by primary key, including access grants."""
try: # single PK lookup + access grants
async with get_async_db_context(db) as session:
tool = await session.get(Tool, id)
if not tool:
return None
return await self._to_tool_model(tool, db=session)
except Exception:
return None
async def get_tools_by_ids(self, tool_ids: list[str], db: AsyncSession | None = None) -> dict[str, ToolModel]:
"""Batch-fetch multiple tools by ID, returning a dict keyed by tool ID."""
if not tool_ids:
return {}
async with get_async_db_context(db) as db:
result = await db.execute(select(Tool).where(Tool.id.in_(tool_ids)))
tools = result.scalars().all()
grants_map = await AccessGrants.get_grants_by_resources('tool', [tool.id for tool in tools], db=db)
return {
tool.id: await self._to_tool_model(tool, access_grants=grants_map.get(tool.id, []), db=db)
for tool in tools
}
async def get_tools(self, defer_content: bool = False, db: AsyncSession | None = None) -> list[ToolUserModel]:
async with get_async_db_context(db) as db:
stmt = select(Tool).order_by(Tool.updated_at.desc())
@ -299,4 +315,4 @@ class ToolsTable:
return False
Tools = ToolsTable()
Tools = ToolsTable() # singleton tool registry

View file

@ -1,9 +1,10 @@
"""User models, Pydantic schemas, and database access layer."""
from __future__ import annotations
import datetime
import time
from typing import Optional
from open_webui.env import DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.utils.misc import throttle
@ -42,40 +43,45 @@ class UserSettings(BaseModel):
pass
class User(Base):
__tablename__ = 'user'
class User(Base): # identity & profile
"""One row per registered account — profile, role, and settings."""
id = Column(String, primary_key=True, unique=True)
email = Column(String)
username = Column(String(50), nullable=True)
role = Column(String)
__tablename__: str = 'user' # Identity & Credentials
id = Column(String, primary_key=True, unique=True) # unique user id
email = Column(String, unique=True) # user email address
username = Column(String(50), nullable=True) # custom handle
role = Column(String, default='pending') # permissions role
name = Column(String, nullable=False) # display name
name = Column(String)
profile_image_url = Column(Text)
# Profile
profile_image_url = Column(Text) # data-uri, path, or external URL
profile_banner_image_url = Column(Text, nullable=True)
bio = Column(Text, nullable=True)
gender = Column(Text, nullable=True)
date_of_birth = Column(Date, nullable=True)
timezone = Column(String, nullable=True)
# Online status
presence_state = Column(String, nullable=True)
status_emoji = Column(String, nullable=True)
status_message = Column(Text, nullable=True)
status_expires_at = Column(BigInteger, nullable=True)
# Metadata
info = Column(JSON, nullable=True)
settings = Column(JSON, nullable=True)
oauth = Column(JSON, nullable=True)
scim = Column(JSON, nullable=True)
# Timestamps (epoch seconds)
last_active_at = Column(BigInteger)
updated_at = Column(BigInteger)
created_at = Column(BigInteger)
_DEFAULT_PROFILE_IMAGE_URL = '/api/v1/users/{user_id}/profile/image'
class UserModel(BaseModel):
id: str
@ -108,12 +114,16 @@ class UserModel(BaseModel):
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
model_config = ConfigDict(
from_attributes=True,
)
# validation schema logic
# --- model validators ---
@model_validator(mode='after')
def set_profile_image_url(self):
if not self.profile_image_url:
self.profile_image_url = f'/api/v1/users/{self.id}/profile/image'
def _ensure_profile_image(self) -> 'UserModel':
"""Assign a generated avatar when no profile image is provided."""
self.profile_image_url = self.profile_image_url or _DEFAULT_PROFILE_IMAGE_URL.format(user_id=self.id)
return self
@ -269,7 +279,7 @@ class UsersTable:
oauth: dict | None = None,
db: AsyncSession | None = None,
) -> UserModel | None:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
user = UserModel(
**{
'id': id,
@ -285,79 +295,91 @@ class UsersTable:
}
)
result = User(**user.model_dump())
db.add(result)
await db.commit()
await db.refresh(result)
if result:
return user
else:
return None
session.add(result)
await session.commit()
await session.refresh(result)
return user if result else None
async def get_user_by_id(self, id: str, db: AsyncSession | None = None) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
# database read methods
# --- read / lookup operations ---
async def get_user_by_id(
self,
id: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Fetch a single user by primary key."""
async with get_async_db_context(db) as session:
user = await session.get(User, id)
return UserModel.model_validate(user) if user else None
async def get_user_by_api_key(self, api_key: str, db: AsyncSession | None = None) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(
select(User).join(ApiKey, User.id == ApiKey.user_id).filter(ApiKey.key == api_key)
)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
# api key auth helper
async def get_user_by_api_key(
self,
api_key: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Resolve a user from their API key via a JOIN on the api_key table."""
async with get_async_db_context(db) as session:
result = await session.execute(
select(User).join(ApiKey, User.id == ApiKey.user_id).where(ApiKey.key == api_key),
)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
async def get_user_by_email(self, email: str, db: AsyncSession | None = None) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter(func.lower(User.email) == email.lower()))
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
async def get_user_by_email(
self,
email: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Case-insensitive email lookup using SQL lower()."""
async with get_async_db_context(db) as session:
email_filter = func.lower(User.email) == email.lower()
query = select(User).where(email_filter)
match = (await session.execute(query)).scalars().first()
if match is None:
return
return UserModel.model_validate(match)
# --- context manager above always returns ---
return
async def get_user_by_oauth_sub(self, provider: str, sub: str, db: AsyncSession | None = None) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
dialect_name = db.bind.dialect.name
stmt = select(User)
if dialect_name == 'sqlite':
stmt = stmt.filter(User.oauth.contains({provider: {'sub': sub}}))
elif dialect_name == 'postgresql':
stmt = stmt.filter(User.oauth[provider].cast(JSONB)['sub'].astext == sub)
result = await db.execute(stmt)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception as e:
# You may want to log the exception here
return None
# --- oauth & integrations ---
async def get_user_by_oauth_sub(
self,
provider: str,
sub: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""Look up a user by OAuth provider + subject claim (dialect-aware JSON filter)."""
async with get_async_db_context(db) as session:
dialect = session.bind.dialect.name
query = select(User)
if dialect == 'sqlite':
oauth_match = User.oauth.contains({provider: {'sub': sub}})
query = query.where(oauth_match)
elif dialect == 'postgresql':
oauth_match = User.oauth[provider].cast(JSONB)['sub'].astext == sub
query = query.where(oauth_match)
row = (await session.execute(query)).scalars().first()
return UserModel.model_validate(row) if row else None
async def get_user_by_scim_external_id(
self, provider: str, external_id: str, db: AsyncSession | None = None
self,
provider: str,
external_id: str,
db: AsyncSession | None = None,
) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
dialect_name = db.bind.dialect.name
stmt = select(User)
if dialect_name == 'sqlite':
stmt = stmt.filter(User.scim.contains({provider: {'external_id': external_id}}))
elif dialect_name == 'postgresql':
stmt = stmt.filter(User.scim[provider].cast(JSONB)['external_id'].astext == external_id)
result = await db.execute(stmt)
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
"""Look up a user by SCIM provider + external ID (dialect-aware JSON filter)."""
async with get_async_db_context(db) as session:
dialect = session.bind.dialect.name
query = select(User)
if dialect == 'sqlite':
scim_match = User.scim.contains({provider: {'external_id': external_id}})
query = query.where(scim_match)
elif dialect == 'postgresql':
scim_match = User.scim[provider].cast(JSONB)['external_id'].astext == external_id
query = query.where(scim_match)
row = (await session.execute(query)).scalars().first()
return UserModel.model_validate(row) if row else None
async def get_users(
self,
@ -366,8 +388,9 @@ class UsersTable:
limit: int | None = None,
db: AsyncSession | None = None,
) -> dict:
async with get_async_db_context(db) as db:
# Import here to avoid circular imports
"""Paginated user listing with optional filters for role, group, and channel."""
async with get_async_db_context(db) as session:
# Deferred imports to avoid circular dependencies
from open_webui.models.channels import ChannelMember
from open_webui.models.groups import GroupMember
@ -487,7 +510,7 @@ class UsersTable:
stmt = stmt.order_by(User.created_at.desc())
# Count BEFORE pagination
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
# correct pagination logic
@ -496,7 +519,7 @@ class UsersTable:
if limit is not None:
stmt = stmt.limit(limit)
result = await db.execute(stmt)
result = await session.execute(stmt)
users = result.scalars().all()
return {
'users': [UserModel.model_validate(user) for user in users],
@ -504,150 +527,114 @@ class UsersTable:
}
async def get_users_by_group_id(self, group_id: str, db: AsyncSession | None = None) -> list[UserModel]:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
from open_webui.models.groups import GroupMember
result = await db.execute(
result = await session.execute(
select(User).join(GroupMember, User.id == GroupMember.user_id).filter(GroupMember.group_id == group_id)
)
users = result.scalars().all()
return [UserModel.model_validate(user) for user in users]
async def get_users_by_user_ids(self, user_ids: list[str], db: AsyncSession | None = None) -> list[UserStatusModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter(User.id.in_(user_ids)))
async with get_async_db_context(db) as session:
result = await session.execute(select(User).filter(User.id.in_(user_ids)))
users = result.scalars().all()
return [UserModel.model_validate(user) for user in users]
# count registered accounts
async def get_num_users(self, db: AsyncSession | None = None) -> int | None:
async with get_async_db_context(db) as db:
result = await db.execute(select(func.count()).select_from(User))
async with get_async_db_context(db) as session:
result = await session.execute(select(func.count()).select_from(User))
return result.scalar()
# check user existence
async def has_users(self, db: AsyncSession | None = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(exists(select(User))))
async with get_async_db_context(db) as session:
result = await session.execute(select(exists(select(User))))
return result.scalar()
async def get_first_user(self, db: AsyncSession | None = None) -> UserModel:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).order_by(User.created_at).limit(1))
user = result.scalars().first()
return UserModel.model_validate(user) if user else None
except Exception:
return None
async def get_first_user(self, db: AsyncSession | None = None) -> UserModel | None:
"""Return the earliest-created user (bootstrap admin detection)."""
async with get_async_db_context(db) as session:
stmt = select(User).order_by(User.created_at).limit(1)
row = (await session.execute(stmt)).scalars().first()
return UserModel.model_validate(row) if row else None
async def get_user_webhook_url_by_id(self, id: str, db: AsyncSession | None = None) -> str | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if user.settings is None:
return None
else:
return user.settings.get('ui', {}).get('notifications', {}).get('webhook_url', None)
except Exception:
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if user and user.settings:
return user.settings.get('ui', {}).get('notifications', {}).get('webhook_url', None)
return None
async def get_num_users_active_today(self, db: AsyncSession | None = None) -> int | None:
async with get_async_db_context(db) as db:
current_timestamp = int(datetime.datetime.now().timestamp())
async with get_async_db_context(db) as session:
current_timestamp = int(time.time())
today_midnight_timestamp = current_timestamp - (current_timestamp % 86400)
result = await db.execute(
select(func.count()).select_from(User).filter(User.last_active_at > today_midnight_timestamp)
result = await session.execute(
select(func.count()).select_from(User).where(User.last_active_at > today_midnight_timestamp)
)
return result.scalar()
async def update_user_role_by_id(self, id: str, role: str, db: AsyncSession | None = None) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
user.role = role
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception:
return None
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if not user:
return None
user.role = role
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
async def update_user_status_by_id(
self, id: str, form_data: UserStatus, db: AsyncSession | None = None
) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
for key, value in form_data.model_dump(exclude_none=True).items():
setattr(user, key, value)
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception:
return None
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if not user:
return None
for key, value in form_data.model_dump(exclude_none=True).items():
setattr(user, key, value)
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
async def update_user_profile_image_url_by_id(
self, id: str, profile_image_url: str, db: AsyncSession | None = None
self,
id: str,
profile_image_url: str,
db: AsyncSession | None = None,
) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
user.profile_image_url = profile_image_url
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception:
return None
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if user is None:
return None
user.profile_image_url = profile_image_url
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
@throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL)
async def update_last_active_by_id(self, id: str, db: AsyncSession | None = None) -> None:
try:
async with get_async_db_context(db) as db:
await db.execute(update(User).filter_by(id=id).values(last_active_at=int(time.time())))
await db.commit()
except Exception:
pass
async with get_async_db_context(db) as session:
await session.execute(update(User).where(User.id == id).values(last_active_at=int(time.time())))
await session.commit()
async def update_user_oauth_by_id(
self, id: str, provider: str, sub: str, db: AsyncSession | None = None
) -> UserModel | None:
"""
Update or insert an OAuth provider/sub pair into the user's oauth JSON field.
Example resulting structure:
{
"google": { "sub": "123" },
"github": { "sub": "abc" }
}
"""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
# Load existing oauth JSON or create empty
oauth = user.oauth or {}
# Update or insert provider entry
oauth[provider] = {'sub': sub}
# Persist updated JSON
await db.execute(update(User).filter_by(id=id).values(oauth=oauth))
await db.commit()
return UserModel.model_validate(user)
except Exception:
return None
"""Update or insert an OAuth provider/sub pair into the user's oauth JSON field."""
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if not user:
return None
oauth = dict(user.oauth or {})
oauth[provider] = {'sub': sub}
user.oauth = oauth
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
async def update_user_scim_by_id(
self,
@ -656,157 +643,102 @@ class UsersTable:
external_id: str,
db: AsyncSession | None = None,
) -> UserModel | None:
"""
Update or insert a SCIM provider/external_id pair into the user's scim JSON field.
Example resulting structure:
{
"microsoft": { "external_id": "abc" },
"okta": { "external_id": "def" }
}
"""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
scim = user.scim or {}
scim[provider] = {'external_id': external_id}
await db.execute(update(User).filter_by(id=id).values(scim=scim))
await db.commit()
return UserModel.model_validate(user)
except Exception:
return None
"""Update or insert a SCIM provider/external_id pair into the user's scim JSON field."""
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if not user:
return None
scim = dict(user.scim or {})
scim[provider] = {'external_id': external_id}
user.scim = scim
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
async def update_user_by_id(self, id: str, updated: dict, db: AsyncSession | None = None) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
for key, value in updated.items():
setattr(user, key, value)
await db.commit()
await db.refresh(user)
return UserModel.model_validate(user)
except Exception as e:
print(e)
return None
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if not user:
return None
for key, value in updated.items():
setattr(user, key, value)
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
# settings update helper
async def update_user_settings_by_id(
self, id: str, updated: dict, db: AsyncSession | None = None
) -> UserModel | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
if not user:
return None
user_settings = user.settings
if user_settings is None:
user_settings = {}
user_settings.update(updated)
await db.execute(update(User).filter_by(id=id).values(settings=user_settings))
await db.commit()
result = await db.execute(select(User).filter_by(id=id))
user = result.scalars().first()
return UserModel.model_validate(user)
except Exception:
return None
async with get_async_db_context(db) as session:
user = await session.get(User, id)
if not user:
return None
user_settings = dict(user.settings or {})
user_settings.update(updated)
user.settings = user_settings
await session.commit()
await session.refresh(user)
return UserModel.model_validate(user)
async def delete_user_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
try:
from open_webui.models.chats import Chats
from open_webui.models.groups import Groups
from open_webui.models.chats import Chats
from open_webui.models.groups import Groups
# Remove User from Groups
await Groups.remove_user_from_all_groups(id)
# Remove User from Groups
await Groups.remove_user_from_all_groups(id)
# Delete User Chats
result = await Chats.delete_chats_by_user_id(id, db=db)
if result:
async with get_async_db_context(db) as db:
# Delete User
await db.execute(delete(User).filter_by(id=id))
await db.commit()
return True
else:
return False
except Exception:
return False
# Delete User Chats
async with get_async_db_context(db) as session:
deleted_chats = await Chats.delete_chats_by_user_id(id, db=session)
if not deleted_chats:
return False # chats deletion failed
await session.execute(delete(User).where(User.id == id))
await session.commit()
return True
async def get_user_api_key_by_id(self, id: str, db: AsyncSession | None = None) -> str | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(ApiKey).filter_by(user_id=id))
api_key = result.scalars().first()
return api_key.key if api_key else None
except Exception:
return None
async with get_async_db_context(db) as session:
api_key = (await session.execute(select(ApiKey).where(ApiKey.user_id == id))).scalars().first()
return api_key.key if api_key else None
async def update_user_api_key_by_id(self, id: str, api_key: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(delete(ApiKey).filter_by(user_id=id))
await db.commit()
now = int(time.time())
new_api_key = ApiKey(
id=f'key_{id}',
user_id=id,
key=api_key,
created_at=now,
updated_at=now,
)
db.add(new_api_key)
await db.commit()
return True
except Exception:
return False
async with get_async_db_context(db) as session:
await session.execute(delete(ApiKey).where(ApiKey.user_id == id))
now_ts = int(time.time())
new_key = ApiKey(
id=f'key_{id}',
user_id=id,
key=api_key,
created_at=now_ts,
updated_at=now_ts,
)
session.add(new_key)
await session.commit()
return True
async def delete_user_api_key_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(delete(ApiKey).filter_by(user_id=id))
await db.commit()
return True
except Exception:
return False
async with get_async_db_context(db) as session:
await session.execute(delete(ApiKey).where(ApiKey.user_id == id))
await session.commit()
return True
async def get_valid_user_ids(self, user_ids: list[str], db: AsyncSession | None = None) -> list[str]:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter(User.id.in_(user_ids)))
users = result.scalars().all()
return [user.id for user in users]
async with get_async_db_context(db) as session:
result = await session.execute(select(User).where(User.id.in_(user_ids)))
return [u.id for u in result.scalars().all()]
async def get_super_admin_user(self, db: AsyncSession | None = None) -> UserModel | None:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(role='admin').limit(1))
user = result.scalars().first()
if user:
return UserModel.model_validate(user)
else:
return None
async with get_async_db_context(db) as session:
row = (await session.execute(select(User).where(User.role == 'admin').limit(1))).scalars().first()
return UserModel.model_validate(row) if row else None
async def get_active_user_count(self, db: AsyncSession | None = None) -> int:
async with get_async_db_context(db) as db:
async with get_async_db_context(db) as session:
# Consider user active if last_active_at within the last 3 minutes
three_minutes_ago = int(time.time()) - 180
result = await db.execute(
select(func.count()).select_from(User).filter(User.last_active_at >= three_minutes_ago)
result = await session.execute(
select(func.count()).select_from(User).where(User.last_active_at >= three_minutes_ago)
)
return result.scalar()
@ -819,9 +751,8 @@ class UsersTable:
return False
async def is_user_active(self, user_id: str, db: AsyncSession | None = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(select(User).filter_by(id=user_id))
user = result.scalars().first()
async with get_async_db_context(db) as session:
user = await session.get(User, user_id)
if user and user.last_active_at:
# Consider user active if last_active_at within the last 3 minutes
three_minutes_ago = int(time.time()) - 180
@ -829,4 +760,4 @@ class UsersTable:
return False
Users = UsersTable()
Users = UsersTable() # singleton user repository

View file

@ -234,7 +234,6 @@ class Loader:
def load(self, filename: str, file_content_type: str, file_path: str) -> list[Document]:
loader = self._get_loader(filename, file_content_type, file_path)
docs = loader.load()
return [Document(page_content=ftfy.fix_text(doc.page_content), metadata=doc.metadata) for doc in docs]
async def aload(self, filename: str, file_content_type: str, file_path: str) -> list[Document]:
@ -257,6 +256,140 @@ class Loader:
and not file_content_type.find('html') >= 0
)
def _detect_text_encoding(self, file_path: str) -> str:
"""Detect the encoding of a text file with CJK-aware fallbacks.
Langchain's ``TextLoader`` uses chardet internally when
``autodetect_encoding=True``, but chardet frequently misidentifies
CJK encodings (e.g. GB18030 detected as GB2312 or even Cyrillic).
This method replaces that by:
1. Trying UTF-8 first (fast path for the vast majority of files).
2. Using chardet as a *hint* to prioritise the right CJK codec
family, but mapping subset names to their superset
(e.g. GB2312 → gb18030).
3. Validating that decoded text actually contains CJK characters,
guarding against codecs that "succeed" but produce garbage.
4. Falling back to latin-1 (always valid, ftfy fixes mojibake later).
"""
try:
with open(file_path, 'rb') as f:
raw = f.read()
except OSError:
return 'utf-8'
if not raw:
return 'utf-8'
# Fast path: most files are UTF-8
try:
raw.decode('utf-8')
return 'utf-8'
except UnicodeDecodeError:
pass
# Use chardet as a hint, not as ground truth
import chardet
detected = chardet.detect(raw)
detected_enc = (detected.get('encoding') or '').lower().replace('-', '').replace('_', '')
# Map chardet's detected encoding to the correct superset codec.
# chardet often reports GB2312 for content that is actually GB18030;
# GB18030 is a strict superset of both GB2312 and GBK.
_ENC_FAMILY = {
'gb2312': 'gb18030',
'gb18030': 'gb18030',
'gbk': 'gb18030',
'big5': 'big5',
'euckr': 'euc-kr',
'eucjp': 'euc-jp',
'iso2022jp': 'euc-jp',
'shiftjis': 'shift_jis',
}
# Build priority list: chardet-hinted codec first, then remaining CJK
base_order = ['gb18030', 'big5', 'euc-kr', 'euc-jp']
hinted = _ENC_FAMILY.get(detected_enc)
if hinted and hinted in base_order:
ordered = [hinted] + [e for e in base_order if e != hinted]
else:
ordered = base_order
for enc in ordered:
try:
text = raw.decode(enc)
if text.strip() and self._has_cjk_characters(text):
log.info(
'Detected encoding %s for %s (chardet guessed %s)',
enc,
file_path,
detected.get('encoding'),
)
return enc
except (UnicodeDecodeError, LookupError):
continue
# If chardet gave a non-CJK answer that isn't in our family map,
# try it directly — it might be a valid Western encoding.
chardet_encoding = detected.get('encoding')
if chardet_encoding:
try:
raw.decode(chardet_encoding)
log.info(
'Using chardet-detected encoding %s for %s',
chardet_encoding,
file_path,
)
return chardet_encoding
except (UnicodeDecodeError, LookupError):
pass
# latin-1 is the ultimate fallback: every byte 0x00–0xFF is valid.
# ftfy.fix_text() (applied downstream) repairs most mojibake that
# results from treating Windows-1252 content as Latin-1.
log.info('Falling back to latin-1 encoding for %s', file_path)
return 'latin-1'
@staticmethod
def _has_cjk_characters(text: str, threshold: float = 0.05) -> bool:
"""Check if decoded text contains a meaningful proportion of CJK characters.
This guards against codecs that technically "succeed" but decode the
bytes into wrong Unicode codepoints (e.g. PUA chars, random symbols).
A genuine CJK document should have at least ``threshold`` fraction of
its non-whitespace characters in CJK Unicode blocks.
"""
if not text:
return False
cjk_count = 0
total = 0
for ch in text:
if ch.isspace():
continue
total += 1
cp = ord(ch)
if (
0x4E00 <= cp <= 0x9FFF # CJK Unified Ideographs
or 0x3400 <= cp <= 0x4DBF # CJK Extension A
or 0x20000 <= cp <= 0x2A6DF # CJK Extension B
or 0x2A700 <= cp <= 0x2B73F # CJK Extension C
or 0x2B740 <= cp <= 0x2B81F # CJK Extension D
or 0xF900 <= cp <= 0xFAFF # CJK Compatibility Ideographs
or 0x3000 <= cp <= 0x303F # CJK Symbols and Punctuation
or 0x3040 <= cp <= 0x309F # Hiragana
or 0x30A0 <= cp <= 0x30FF # Katakana
or 0xAC00 <= cp <= 0xD7AF # Hangul Syllables
or 0xFF00 <= cp <= 0xFFEF # Halfwidth and Fullwidth Forms
):
cjk_count += 1
if total == 0:
return False
return (cjk_count / total) >= threshold
def _get_loader(self, filename: str, file_content_type: str, file_path: str):
file_ext = filename.split('.')[-1].lower()
@ -274,7 +407,7 @@ class Loader:
)
elif self.engine == 'tika' and self.kwargs.get('TIKA_SERVER_URL'):
if self._is_text_file(file_ext, file_content_type):
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
else:
loader = TikaLoader(
url=self.kwargs.get('TIKA_SERVER_URL'),
@ -326,7 +459,7 @@ class Loader:
)
elif self.engine == 'docling' and self.kwargs.get('DOCLING_SERVER_URL'):
if self._is_text_file(file_ext, file_content_type):
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
else:
# Build params for DoclingLoader
params = self.kwargs.get('DOCLING_PARAMS', {})
@ -371,7 +504,7 @@ class Loader:
azure_credential=DefaultAzureCredential(),
api_model=self.kwargs.get('DOCUMENT_INTELLIGENCE_MODEL'),
)
elif self.engine == 'mineru' and file_ext in ['pdf']: # MinerU currently only supports PDF
elif self.engine == 'mineru' and file_ext in self.kwargs.get('MINERU_FILE_EXTENSIONS', ['pdf']):
mineru_timeout = self.kwargs.get('MINERU_API_TIMEOUT', 300)
if mineru_timeout:
try:
@ -413,7 +546,7 @@ class Loader:
mode=self.kwargs.get('PDF_LOADER_MODE', 'page'),
)
elif file_ext == 'csv':
loader = CSVLoader(file_path, autodetect_encoding=True)
loader = CSVLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_ext == 'rst':
try:
from langchain_community.document_loaders import UnstructuredRSTLoader
@ -425,7 +558,7 @@ class Loader:
'Falling back to plain text loading for .rst file. '
'Install it with: pip install unstructured'
)
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_ext == 'xml':
try:
from langchain_community.document_loaders import UnstructuredXMLLoader
@ -437,11 +570,11 @@ class Loader:
'Falling back to plain text loading for .xml file. '
'Install it with: pip install unstructured'
)
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_ext in ['htm', 'html']:
loader = BSHTMLLoader(file_path, open_encoding='unicode_escape')
elif file_ext == 'md':
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_content_type == 'application/epub+zip':
try:
from langchain_community.document_loaders import UnstructuredEPubLoader
@ -457,6 +590,16 @@ class Loader:
or file_ext == 'docx'
):
loader = Docx2txtLoader(file_path)
elif file_ext == 'doc' or file_content_type == 'application/msword':
try:
from langchain_community.document_loaders import UnstructuredWordDocumentLoader
loader = UnstructuredWordDocumentLoader(file_path)
except ImportError:
raise ValueError(
"Processing .doc files requires the 'unstructured' package. "
'Install it with: pip install unstructured'
)
elif file_content_type in [
'application/vnd.ms-excel',
'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet',
@ -500,8 +643,8 @@ class Loader:
'Install it with: pip install unstructured'
)
elif self._is_text_file(file_ext, file_content_type):
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
else:
loader = TextLoader(file_path, autodetect_encoding=True)
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
return loader

View file

@ -98,7 +98,7 @@ class YoutubeLoader:
try:
transcript_list = transcript_api.list(self.video_id)
except Exception as e:
log.exception('Loading YouTube transcript failed')
log.warning(f'Loading YouTube transcript failed: {e}')
return []
# Try each language in order of priority

View file

@ -29,7 +29,9 @@ from open_webui.env import (
AIOHTTP_CLIENT_ALLOW_REDIRECTS,
AIOHTTP_CLIENT_SESSION_SSL,
AIOHTTP_CLIENT_TIMEOUT,
BYPASS_RETRIEVAL_ACCESS_CONTROL,
ENABLE_FORWARD_USER_INFO_HEADERS,
ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS,
OFFLINE_MODE,
)
from open_webui.models.access_grants import AccessGrants
@ -117,6 +119,7 @@ def build_loader_from_config(request):
MINERU_API_KEY=config.MINERU_API_KEY,
MINERU_API_TIMEOUT=config.MINERU_API_TIMEOUT,
MINERU_PARAMS=config.MINERU_PARAMS,
MINERU_FILE_EXTENSIONS=config.MINERU_FILE_EXTENSIONS,
)
@ -175,6 +178,18 @@ def get_content_from_url(request, url: str) -> str:
# Validate URL before making any request (blocks private IPs, non-HTTP, filter list)
validate_url(url)
# YouTube URLs (including youtu.be short links) should go straight to
# YoutubeLoader, which uses youtube-transcript-api and never needs the
# HTTP response body. Probing the URL first is harmful for short URLs:
# youtu.be returns a 303 redirect with Content-Type: application/binary
# when allow_redirects=False, causing the binary-content path to run
# and produce empty docs → HTTP 400.
if is_youtube_url(url):
loader = get_loader(request, url)
docs = loader.load()
content = ' '.join([doc.page_content for doc in docs])
return content, docs
# Streamed GET to check Content-Type without downloading the body.
# allow_redirects=False prevents redirect-based SSRF: validate_url() above is
# called on the originally-submitted URL only; following 3xx redirects without
@ -914,6 +929,13 @@ def get_embedding_function(
concurrent_requests=0,
) -> Awaitable:
if embedding_engine == '':
if embedding_function is None:
raise ValueError(
'No embedding model is loaded. Set RAG_EMBEDDING_MODEL to a valid '
'SentenceTransformer model name, or configure an external '
'RAG_EMBEDDING_ENGINE (ollama, openai, azure_openai).'
)
# Sentence transformers: CPU-bound sync operation
async def async_embedding_function(query, prefix=None, user=None):
return await asyncio.to_thread(
@ -1053,6 +1075,16 @@ def get_reranking_function(reranking_engine, reranking_model, reranking_function
)
# UUIDs, SHA-256 digests, and prefixed variants thereof all fit [A-Za-z0-9_-].
# Anything else cannot be a real Open WebUI collection and could break out of
# a Milvus expression literal.
_SAFE_COLLECTION_NAME_RE = re.compile(r'^[A-Za-z0-9_-]{1,255}$')
def _is_safe_collection_name(name: str) -> bool:
return isinstance(name, str) and bool(_SAFE_COLLECTION_NAME_RE.match(name))
async def filter_accessible_collections(
collection_names: set[str],
user: UserModel,
@ -1062,20 +1094,33 @@ async def filter_accessible_collections(
Return only the collection names the user is allowed to access.
Admins bypass all checks. For non-admins the policy is:
- any name with characters outside [A-Za-z0-9_-] → rejected
- file-* → validated via has_access_to_file
- user-memory-* → must match user's own memory collection
- web-search-* → ephemeral per-query collections, always allowed
- knowledge-bases → always denied (system meta-collection)
- everything else → if the name matches a knowledge base, validated
via Knowledges.check_access_by_user_id; if no
such KB exists, the name is treated as an
ephemeral/legacy collection and allowed
such KB exists, denied by default. When
ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS is True,
the name is treated as a legacy/ephemeral
collection and allowed.
"""
# Applied before the admin bypass — malformed names should never reach the vector store.
safe_names = {n for n in collection_names if _is_safe_collection_name(n)}
rejected = collection_names - safe_names
if rejected:
log.warning(
'filter_accessible_collections: rejected %d collection name(s) with unsafe characters (user_id=%s)',
len(rejected),
getattr(user, 'id', '<unknown>'),
)
if user.role == 'admin':
return collection_names
return safe_names
validated = set()
for name in collection_names:
for name in safe_names:
if name == 'knowledge-bases':
# System meta-collection — never exposed to non-admins.
continue
@ -1094,11 +1139,13 @@ async def filter_accessible_collections(
else:
# May be a knowledge-base ID or a legacy/ephemeral collection.
# If it IS a KB, enforce access control. If no such KB
# exists, treat it as a non-sensitive collection (e.g. legacy
# model knowledge, process_text SHA256 collections) and allow.
# exists, the behaviour depends on
# ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS:
# False (default) — deny (closes the unscoped namespace)
# True — allow (preserves legacy behaviour)
if await Knowledges.check_access_by_user_id(name, user.id, permission=access_type):
validated.add(name)
elif not await Knowledges.get_knowledge_by_id(name):
elif ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS and not await Knowledges.get_knowledge_by_id(name):
# Not a KB at all — legacy/ephemeral collection, allow
validated.add(name)
return validated
@ -1242,11 +1289,27 @@ async def get_sources_from_items(
],
}
else:
# Fallback to collection names
if item.get('legacy'):
collection_names.append(f'{item["id"]}')
else:
collection_names.append(f'file-{item["id"]}')
# Chunked-retrieval fallback — verify read access before
# exposing the file's vector collection (same posture as the
# full-context branch above).
file_id = item.get('id')
if file_id:
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
if item.get('legacy'):
collection_names.append(f'{file_id}')
else:
collection_names.append(f'file-{file_id}')
else:
file_object = await Files.get_file_by_id(file_id)
if file_object and (
user.role == 'admin'
or file_object.user_id == user.id
or await has_access_to_file(file_id, 'read', user)
):
if item.get('legacy'):
collection_names.append(f'{file_id}')
else:
collection_names.append(f'file-{file_id}')
elif item.get('type') == 'collection':
# Manual Full Mode Toggle for Collection
@ -1292,9 +1355,18 @@ async def get_sources_from_items(
'metadatas': [metadatas],
}
else:
# Fallback to collection names
if item.get('legacy'):
collection_names = item.get('collection_names', [])
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
collection_names = item.get('collection_names', [])
else:
# Legacy KB: item.collection_names is client-supplied.
# Validate against the KB's actual files to prevent
# cross-tenant collection name substitution.
files = await Knowledges.get_files_by_id(knowledge_base.id)
owned_names = {f'file-{f.id}' for f in files}
owned_names.add(knowledge_base.id)
valid_names = [n for n in (item.get('collection_names') or []) if n in owned_names]
collection_names = valid_names if valid_names else [knowledge_base.id]
else:
collection_names.append(item['id'])
@ -1305,11 +1377,20 @@ async def get_sources_from_items(
'metadatas': [[doc.get('metadata') for doc in item.get('docs')]],
}
elif item.get('collection_name'):
# Direct Collection Name
collection_names.append(item['collection_name'])
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
collection_names.append(item['collection_name'])
else:
log.debug(
"get_sources_from_items: ignoring untrusted direct collection_name '%s' on item without type",
item.get('collection_name'),
)
elif item.get('collection_names'):
# Collection Names List
collection_names.extend(item['collection_names'])
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
collection_names.extend(item['collection_names'])
else:
log.debug(
'get_sources_from_items: ignoring untrusted direct collection_names on item without type',
)
# If query_result is None
# Fallback to collection names and vector search the collections

View file

@ -3,6 +3,7 @@ NOTE: This vector database integration is community-supported and maintained on
"""
import logging
import re
from typing import Any, Dict, List, Optional, Tuple
from open_webui.config import (
@ -35,6 +36,30 @@ log = logging.getLogger(__name__)
RESOURCE_ID_FIELD = 'resource_id'
# Milvus expressions are SQL-like strings with no parameterized-query API;
# values get interpolated into single-quoted literals. Reject anything that
# can't be a legitimate Open WebUI collection name.
_SAFE_RESOURCE_ID_RE = re.compile(r'^[A-Za-z0-9_-]{1,255}$')
_SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$')
def _validate_resource_id(resource_id: str) -> str:
if not isinstance(resource_id, str) or not _SAFE_RESOURCE_ID_RE.match(resource_id):
raise ValueError(f'Invalid Milvus resource_id (collection name): {resource_id!r}')
return resource_id
def _validate_metadata_key(key: str) -> str:
if not isinstance(key, str) or not _SAFE_METADATA_KEY_RE.match(key):
raise ValueError(f'Invalid Milvus metadata filter key: {key!r}')
return key
def _escape_milvus_string(value: str) -> str:
if not isinstance(value, str):
raise TypeError(f'Expected str for Milvus expression value, got {type(value).__name__}')
return value.replace('\\', '\\\\').replace("'", "\\'")
class MilvusClient(VectorDBBase):
def __init__(self):
@ -126,6 +151,7 @@ class MilvusClient(VectorDBBase):
def has_collection(self, collection_name: str) -> bool:
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id)
if not utility.has_collection(mt_collection):
return False
@ -138,6 +164,7 @@ class MilvusClient(VectorDBBase):
if not items:
return
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id)
dimension = len(items[0]['vector'])
self._ensure_collection(mt_collection, dimension)
collection = Collection(mt_collection)
@ -165,6 +192,7 @@ class MilvusClient(VectorDBBase):
return None
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id)
if not utility.has_collection(mt_collection):
return None
@ -203,21 +231,22 @@ class MilvusClient(VectorDBBase):
filter: Optional[Dict[str, Any]] = None,
):
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id)
if not utility.has_collection(mt_collection):
return
collection = Collection(mt_collection)
# Build expression
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
if ids:
# Milvus expects a string list for 'in' operator
id_list_str = ', '.join([f"'{id_val}'" for id_val in ids])
id_list_str = ', '.join([f"'{_escape_milvus_string(str(id_val))}'" for id_val in ids])
expr.append(f'id in [{id_list_str}]')
if filter:
for key, value in filter.items():
expr.append(f"metadata['{key}'] == '{value}'")
_validate_metadata_key(key)
expr.append(f"metadata['{key}'] == '{_escape_milvus_string(str(value))}'")
collection.delete(' and '.join(expr))
@ -228,6 +257,7 @@ class MilvusClient(VectorDBBase):
def delete_collection(self, collection_name: str):
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id)
if not utility.has_collection(mt_collection):
return
@ -236,6 +266,7 @@ class MilvusClient(VectorDBBase):
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]:
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id)
if not utility.has_collection(mt_collection):
return None
@ -245,10 +276,15 @@ class MilvusClient(VectorDBBase):
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
if filter:
for key, value in filter.items():
_validate_metadata_key(key)
if isinstance(value, str):
expr.append(f"metadata['{key}'] == '{value}'")
else:
expr.append(f"metadata['{key}'] == '{_escape_milvus_string(value)}'")
elif isinstance(value, bool):
expr.append(f"metadata['{key}'] == {str(value).lower()}")
elif isinstance(value, (int, float)):
expr.append(f"metadata['{key}'] == {value}")
else:
raise TypeError(f'Unsupported Milvus filter value type for key {key!r}: {type(value).__name__}')
iterator = collection.query_iterator(
expr=' and '.join(expr),

View file

@ -216,28 +216,23 @@ class QdrantClient(VectorDBBase):
ids: Optional[list[str]] = None,
filter: Optional[dict] = None,
):
# Delete the items from the collection based on the ids.
field_conditions = []
# Delete by point ID: the point ID is the item's id (see _create_points).
# Filtering on metadata.id silently misses points whose payload omits an
# id (e.g. memories), leaving orphaned vectors behind.
if ids:
for id_value in ids:
(
field_conditions.append(
models.FieldCondition(
key='metadata.id',
match=models.MatchValue(value=id_value),
),
),
)
elif filter:
return self.client.delete(
collection_name=f'{self.collection_prefix}_{collection_name}',
points_selector=models.PointIdsList(points=ids),
)
field_conditions = []
if filter:
for key, value in filter.items():
(
field_conditions.append(
models.FieldCondition(
key=f'metadata.{key}',
match=models.MatchValue(value=value),
),
),
field_conditions.append(
models.FieldCondition(
key=f'metadata.{key}',
match=models.MatchValue(value=value),
)
)
return self.client.delete(

View file

@ -228,15 +228,17 @@ class QdrantClient(VectorDBBase):
return None
must_conditions = [_tenant_filter(tenant_id)]
should_conditions = []
if ids:
should_conditions = [_metadata_filter('id', id_value) for id_value in ids]
# Delete by point ID within the tenant. The point ID is the item's id
# (see _create_points); filtering on metadata.id silently misses points
# whose payload omits an id (e.g. memories), leaving orphaned vectors.
must_conditions.append(models.HasIdCondition(has_id=ids))
elif filter:
must_conditions += [_metadata_filter(k, v) for k, v in filter.items()]
return self.client.delete(
collection_name=mt_collection,
points_selector=models.FilterSelector(filter=models.Filter(must=must_conditions, should=should_conditions)),
points_selector=models.FilterSelector(filter=models.Filter(must=must_conditions)),
)
def search(

View file

@ -0,0 +1,765 @@
# NOTE: This vector database integration is community-supported and maintained on a best-effort basis.
# Requires Valkey core >= 9.0.1 with the valkey-search module >= 1.2.0 loaded.
import atexit
import json
import logging
import re
import struct
from urllib.parse import urlparse
from open_webui.config import (
VALKEY_COLLECTION_PREFIX,
VALKEY_DISTANCE_METRIC,
VALKEY_HNSW_EF_CONSTRUCTION,
VALKEY_HNSW_EF_RUNTIME,
VALKEY_HNSW_M,
VALKEY_INDEX_TYPE,
VALKEY_URL,
)
from open_webui.retrieval.vector.main import (
GetResult,
SearchResult,
VectorDBBase,
VectorItem,
)
from open_webui.retrieval.vector.utils import process_metadata
log = logging.getLogger(__name__)
def _import_glide():
"""Lazily import glide_sync so the module can be loaded without valkey-glide-sync installed."""
try:
from glide_sync import (
Batch,
DataType,
DistanceMetricType,
FtCreateOptions,
FtSearchLimit,
FtSearchOptions,
GlideClient,
GlideClientConfiguration,
NodeAddress,
RequestError,
ReturnField,
TagField,
TextField,
VectorAlgorithm,
VectorField,
VectorFieldAttributesFlat,
VectorFieldAttributesHnsw,
VectorType,
)
from glide_sync import (
ft as glide_ft,
)
except ImportError as e:
raise ImportError(
'valkey-glide-sync is required when VECTOR_DB=valkey. Install it with: pip install valkey-glide-sync==2.3.1'
) from e
return {
'Batch': Batch,
'DataType': DataType,
'DistanceMetricType': DistanceMetricType,
'FtCreateOptions': FtCreateOptions,
'FtSearchLimit': FtSearchLimit,
'FtSearchOptions': FtSearchOptions,
'GlideClient': GlideClient,
'GlideClientConfiguration': GlideClientConfiguration,
'NodeAddress': NodeAddress,
'RequestError': RequestError,
'ReturnField': ReturnField,
'TagField': TagField,
'TextField': TextField,
'VectorAlgorithm': VectorAlgorithm,
'VectorField': VectorField,
'VectorFieldAttributesFlat': VectorFieldAttributesFlat,
'VectorFieldAttributesHnsw': VectorFieldAttributesHnsw,
'VectorType': VectorType,
'glide_ft': glide_ft,
}
# valkey-search 1.2.0 requires Valkey core 9.0.1+ per upstream release notes.
# Unlike RediSearch (dialects 1-4), valkey-search only implements DIALECT 2 — GLIDE's
# FtSearchOptions doesn't expose a dialect parameter because it's always dialect 2.
MIN_VALKEY_VERSION = (9, 0, 1)
MIN_SEARCH_MODULE_VERSION = (1, 2, 0)
_VALID_DISTANCE_METRICS = {'COSINE', 'L2', 'IP'}
_NEVER_MATCH_SENTINEL = '__open_webui_valkey_never_match__'
# Compile once at module load — includes `?` which is a single-char wildcard in TAG queries.
_TAG_SPECIAL_RE = re.compile(r'([,.<>{}\[\]"\':;!@#$%^&*()\-+=~?\\/| \t\n\r])')
_SAFE_FIELD_RE = re.compile(r'^[a-zA-Z_][a-zA-Z0-9_]*$')
def _vector_to_bytes(vector: list[float | int]) -> bytes:
"""Pack a list of floats as a float32 little-endian binary blob."""
return struct.pack(f'<{len(vector)}f', *vector)
def _escape_tag_value(value: str) -> str:
"""Escape special characters for RediSearch/Valkey-Search TAG field queries."""
return _TAG_SPECIAL_RE.sub(r'\\\1', str(value))
def _build_filter_expression(filter: dict) -> str:
"""Translate a Chroma-style filter dict into a valkey-search filter expression.
Supports simple equality, $in, $ne, and $eq. Multiple keys are ANDed together.
Raises ValueError on unsupported operators rather than silently matching nothing.
"""
parts = []
for key, value in filter.items():
if not _SAFE_FIELD_RE.match(key):
raise ValueError(
f'Invalid filter field name: {key!r}. '
'Field names must start with a letter or underscore and contain only alphanumerics/underscores.'
)
if isinstance(value, dict):
for op, operand in value.items():
if op == '$in' and isinstance(operand, list):
if not operand:
# Empty $in → match nothing, not "match all".
parts.append(f'@{key}:{{{_NEVER_MATCH_SENTINEL}}}')
continue
escaped = [_escape_tag_value(str(v)) for v in operand]
parts.append(f'@{key}:{{{"|".join(escaped)}}}')
elif op in ('$eq', '$ne'):
prefix = '-' if op == '$ne' else ''
parts.append(f'{prefix}@{key}:{{{_escape_tag_value(str(operand))}}}')
else:
raise ValueError(
f'Unsupported filter operator {op!r} for key {key!r}. Supported operators: $in, $ne, $eq.'
)
else:
parts.append(f'@{key}:{{{_escape_tag_value(str(value))}}}')
return ' '.join(parts)
def _decode(value) -> str:
"""Decode bytes to str; pass through str unchanged."""
if isinstance(value, (bytes, bytearray)):
return value.decode()
return str(value) if value is not None else ''
class ValkeyClient(VectorDBBase):
def __init__(self):
if not VALKEY_URL:
raise ValueError(
'VALKEY_URL is required when VECTOR_DB=valkey. '
'Set it to your Valkey server URL (e.g., valkey://localhost:6379).'
)
# Lazily import glide_sync — only needed when this backend is actually used.
self._g = _import_glide()
# Validate distance metric at init — invalid values pass through to FT.CREATE
# and fail with a cryptic server error.
metric = VALKEY_DISTANCE_METRIC.upper()
if metric not in _VALID_DISTANCE_METRICS:
raise ValueError(
f'Invalid VALKEY_DISTANCE_METRIC={VALKEY_DISTANCE_METRIC!r}. '
f'Must be one of: {", ".join(sorted(_VALID_DISTANCE_METRICS))}.'
)
DistanceMetricType = self._g['DistanceMetricType']
self._distance_metric_map = {
'COSINE': DistanceMetricType.COSINE,
'L2': DistanceMetricType.L2,
'IP': DistanceMetricType.IP,
}
self.collection_prefix = VALKEY_COLLECTION_PREFIX
self.index_type = VALKEY_INDEX_TYPE
self.distance_metric = metric
parsed = urlparse(VALKEY_URL)
host = parsed.hostname or 'localhost'
port = parsed.port or 6379
db = int(parsed.path.lstrip('/') or 0)
GlideClientConfiguration = self._g['GlideClientConfiguration']
NodeAddress = self._g['NodeAddress']
GlideClient = self._g['GlideClient']
config = GlideClientConfiguration(
addresses=[NodeAddress(host=host, port=port)],
database_id=db if db else None,
request_timeout=5000,
client_name='open_webui_vector_store_client',
)
try:
self.client = GlideClient.create(config)
except Exception as e:
raise ConnectionError(f'Failed to connect to Valkey at {host}:{port}: {e}') from e
# Separate client for batch writes — large HSET payloads on the multiplexed
# connection can starve concurrent reads.
batch_config = GlideClientConfiguration(
addresses=[NodeAddress(host=host, port=port)],
database_id=db if db else None,
request_timeout=10000, # 10s — HNSW indexing can take 1-4s per vector
client_name='open_webui_vector_store_batch_client',
)
try:
self.batch_client = GlideClient.create(batch_config)
except Exception as e:
raise ConnectionError(f'Failed to create batch write client for Valkey at {host}:{port}: {e}') from e
try:
self.client.ping()
except Exception as e:
raise ConnectionError(f'Failed to ping Valkey at {host}:{port}: {e}') from e
# Catch misconfigured deployments at startup (e.g., valkey-bundle:9.0.1 ships
# valkey-search 1.0.0 which lacks TEXT fields and filter-only FT.SEARCH).
self._check_core_version()
self._check_search_module()
atexit.register(self.close)
def close(self) -> None:
"""Close both GLIDE clients, flushing in-flight requests."""
try:
self.client.close()
except Exception:
pass
try:
self.batch_client.close()
except Exception:
pass
# ----- version checks ----------------------------------------------------
@staticmethod
def _parse_semver(version_str: str) -> tuple[int, int, int] | None:
if not version_str:
return None
m = re.match(r'^(\d+)\.(\d+)\.(\d+)', version_str)
return (int(m.group(1)), int(m.group(2)), int(m.group(3))) if m else None
@staticmethod
def _format_version(v: tuple[int, int, int]) -> str:
return f'{v[0]}.{v[1]}.{v[2]}'
def _check_core_version(self) -> None:
try:
info_raw = self.client.info()
except Exception as e:
log.warning(f'Could not fetch Valkey INFO for version check, proceeding: {e}')
return
raw = None
text = _decode(info_raw) if info_raw else ''
redis_fallback = None
for line in text.splitlines():
if line.startswith('valkey_version:'):
raw = line.split(':', 1)[1].strip()
break
if line.startswith('redis_version:') and redis_fallback is None:
redis_fallback = line.split(':', 1)[1].strip()
if raw is None:
raw = redis_fallback
version = self._parse_semver(raw) if raw else None
if version is None:
log.warning(
f'Could not determine Valkey core version (raw={raw!r}); proceeding but '
f'minimum {self._format_version(MIN_VALKEY_VERSION)} is required.'
)
elif version < MIN_VALKEY_VERSION:
raise RuntimeError(
f'Valkey core {self._format_version(version)} is below the minimum required version '
f'{self._format_version(MIN_VALKEY_VERSION)}. valkey-search 1.2.0 requires Valkey core '
'9.0.1 or later. Upgrade your server or use valkey-bundle:9.1.0-rc2+.'
)
log.info(f'Valkey core version: {self._format_version(version) if version else "unknown"}')
def _check_search_module(self) -> None:
try:
modules = self.client.custom_command(['MODULE', 'LIST'])
except Exception as e:
log.warning(
f'Could not list modules on the Valkey server ({e}); proceeding but '
f'valkey-search >= {self._format_version(MIN_SEARCH_MODULE_VERSION)} is required.'
)
return
# MODULE LIST returns [{b'name': b'search', b'ver': 66048, ...}]
# ver encoding: major*10000 + minor*100 + patch
search_version: tuple[int, int, int] | None = None
module_present = False
raw_ver = None
for entry in modules or []:
if isinstance(entry, dict):
name = _decode(entry.get(b'name') or entry.get('name') or '')
raw_ver = entry.get(b'ver') or entry.get('ver', 0)
else:
parsed = self._decode_kv_pairs(entry)
name = parsed.get('name', '')
raw_ver = parsed.get('ver', 0)
if name.lower() == 'search':
module_present = True
try:
ver_int = int(raw_ver)
search_version = (ver_int // 10000, (ver_int % 10000) // 100, ver_int % 100)
except (TypeError, ValueError):
search_version = None
break
if not module_present:
raise RuntimeError(
'The valkey-search module is not loaded on the Valkey server. '
f'This backend requires valkey-search >= {self._format_version(MIN_SEARCH_MODULE_VERSION)}. '
'Use valkey-bundle:9.1.0-rc2+ or load libsearch.so via --loadmodule on a Valkey 9.0.1+ server.'
)
if search_version is None:
log.warning(
f'valkey-search module is loaded but version could not be parsed (raw={raw_ver!r}); '
f'proceeding but minimum {self._format_version(MIN_SEARCH_MODULE_VERSION)} is required.'
)
elif search_version < MIN_SEARCH_MODULE_VERSION:
raise RuntimeError(
f'valkey-search {self._format_version(search_version)} is below the minimum required '
f'version {self._format_version(MIN_SEARCH_MODULE_VERSION)}. Earlier versions lack the '
'TEXT field type and filter-only FT.SEARCH support required by this backend. '
'Upgrade to valkey-bundle:9.1.0-rc2+ or load valkey-search 1.2.0+ as a module.'
)
log.info(f'valkey-search version: {self._format_version(search_version) if search_version else "unknown"}')
def _index_name(self, collection_name: str) -> str:
return f'idx:{self.collection_prefix}:{collection_name}'
def _key_prefix(self, collection_name: str) -> str:
return f'{self.collection_prefix}:{collection_name}:'
def _item_key(self, collection_name: str, item_id: str) -> str:
return f'{self.collection_prefix}:{collection_name}:{item_id}'
def _create_index(self, collection_name: str, dimension: int) -> None:
"""Create an FT index for a collection with the given vector dimension."""
g = self._g
index_name = self._index_name(collection_name)
prefix = self._key_prefix(collection_name)
distance_metric = self._distance_metric_map[self.distance_metric]
if self.index_type == 'HNSW':
vector_attrs = g['VectorFieldAttributesHnsw'](
dimensions=dimension,
distance_metric=distance_metric,
type=g['VectorType'].FLOAT32,
number_of_edges=VALKEY_HNSW_M,
vectors_examined_on_construction=VALKEY_HNSW_EF_CONSTRUCTION,
vectors_examined_on_runtime=VALKEY_HNSW_EF_RUNTIME,
)
algo = g['VectorAlgorithm'].HNSW
else:
if self.index_type != 'FLAT':
log.warning(f'Unrecognized VALKEY_INDEX_TYPE={self.index_type!r}; falling back to FLAT.')
vector_attrs = g['VectorFieldAttributesFlat'](
dimensions=dimension,
distance_metric=distance_metric,
type=g['VectorType'].FLOAT32,
)
algo = g['VectorAlgorithm'].FLAT
schema = [
g['VectorField'](name='vector', algorithm=algo, attributes=vector_attrs),
g['TextField'](name='text'),
g['TagField'](name='id'),
g['TextField'](name='metadata_json'),
g['TagField'](name='hash'),
g['TagField'](name='file_id'),
g['TagField'](name='source'),
g['TagField'](name='knowledge_base_id'),
]
options = g['FtCreateOptions'](data_type=g['DataType'].HASH, prefixes=[prefix])
try:
g['glide_ft'].create(self.client, index_name, schema, options)
log.info(
f'Created Valkey index {index_name} with dimension={dimension}, '
f'type={self.index_type}, metric={self.distance_metric}'
)
except g['RequestError'] as e:
if 'already exists' in str(e).lower():
log.debug(f'Index {index_name} already exists, skipping creation.')
else:
raise
def _verify_collection_dimension(self, collection_name: str, dimension: int) -> None:
index_name = self._index_name(collection_name)
try:
info = self._g['glide_ft'].info(self.client, index_name)
except Exception as e:
log.warning(f'Could not FT.INFO {index_name} for dimension check, skipping: {e}')
return
# ft.info response has nested structure: b'attributes' → list of fields,
# each field is [k1, v1, ...] with a nested 'index' sub-list containing 'dimensions'.
existing = None
attrs = None
if isinstance(info, dict):
attrs = info.get(b'attributes') or info.get('attributes')
elif isinstance(info, (list, tuple)):
attrs = self._find_in_kv_pairs(info, 'attributes', case_insensitive=True)
for attr in attrs or []:
if not isinstance(attr, (list, tuple)):
continue
field_type = self._find_in_kv_pairs(attr, 'type', case_insensitive=True)
if _decode(field_type).upper() != 'VECTOR':
continue
index_params = self._find_in_kv_pairs(attr, 'index', case_insensitive=True)
if index_params and isinstance(index_params, (list, tuple)):
dim_raw = self._find_in_kv_pairs(index_params, 'dimensions', case_insensitive=True)
if dim_raw is not None:
try:
existing = int(dim_raw)
except (ValueError, TypeError):
pass
break
if existing is None:
log.warning(
f'Could not determine vector dimension for {index_name} from FT.INFO response, '
'skipping dimension check.'
)
return
if existing != dimension:
raise ValueError(
f'Collection {collection_name!r} was created with dim={existing}, refusing to '
f'insert vectors with dim={dimension}. Recreate the collection (e.g., via '
'VECTOR_DB_CLIENT.delete_collection) if you intend to switch embedding models.'
)
def has_collection(self, collection_name: str) -> bool:
index_name = self._index_name(collection_name)
try:
self._g['glide_ft'].info(self.client, index_name)
return True
except self._g['RequestError'] as e:
msg = str(e).lower()
if 'no such index' in msg or 'unknown index' in msg or 'not found in database' in msg:
return False
log.warning(f'Unexpected FT.INFO response for collection {collection_name}: {e}')
raise
def delete_collection(self, collection_name: str):
index_name = self._index_name(collection_name)
try:
self._g['glide_ft'].dropindex(self.client, index_name)
log.info(f'Dropped index {index_name}')
except self._g['RequestError'] as e:
log.debug(f'Could not drop index {index_name}: {e}')
self._delete_keys_by_prefix(self._key_prefix(collection_name))
def insert(self, collection_name: str, items: list[VectorItem]):
if not items:
return
dimension = len(items[0]['vector'])
if not self.has_collection(collection_name):
self._create_index(collection_name, dimension)
else:
self._verify_collection_dimension(collection_name, dimension)
# Individual HSET rather than Batch.exec() — each command gets its own timeout.
# HNSW indexing can take 1-4s per vector (ef_construction=200), and Batch.exec()
# applies a single timeout to ALL commands, causing all-or-nothing failures on
# large inserts.
for item in items:
metadata = process_metadata(item['metadata']) if item.get('metadata') else {}
mapping = {
'id': item['id'],
'vector': _vector_to_bytes(item['vector']),
'text': item['text'],
'metadata_json': json.dumps(metadata),
# `or ''` prevents indexing literal 'None' as a TAG value, which would
# poison $ne / equality queries.
'hash': str(metadata.get('hash') or ''),
'file_id': str(metadata.get('file_id') or ''),
'source': str(metadata.get('source') or ''),
'knowledge_base_id': str(metadata.get('knowledge_base_id') or ''),
}
self.batch_client.hset(self._item_key(collection_name, item['id']), mapping)
log.debug(f'Inserted {len(items)} items into collection {collection_name}')
def upsert(self, collection_name: str, items: list[VectorItem]):
self.insert(collection_name, items)
def search(
self,
collection_name: str,
vectors: list[list[float | int]],
filter: dict | None = None,
limit: int = 10,
) -> SearchResult | None:
if not vectors:
return None
if not self.has_collection(collection_name):
return None
filter_expr = _build_filter_expression(filter) if filter else ''
query_str = (
f'({filter_expr})=>[KNN {limit} @vector $query_vec]'
if filter_expr
else f'*=>[KNN {limit} @vector $query_vec]'
)
g = self._g
try:
opts = g['FtSearchOptions'](
params={'query_vec': _vector_to_bytes(vectors[0])},
limit=g['FtSearchLimit'](offset=0, count=limit),
)
result = g['glide_ft'].search(self.client, self._index_name(collection_name), query_str, opts)
except g['RequestError'] as e:
log.error(f'Valkey search error on collection {collection_name}: {e}')
return None
return self._parse_glide_search_response(result, include_score=True)
def query(self, collection_name: str, filter: dict, limit: int | None = None) -> GetResult | None:
if not self.has_collection(collection_name):
return None
if not filter:
return self.get(collection_name, limit=limit)
query_str = _build_filter_expression(filter)
if not query_str:
return self.get(collection_name, limit=limit)
# Hard cap when no limit provided — FT.SEARCH requires a finite count.
effective_limit = limit if limit and limit > 0 else 10000
if not (limit and limit > 0):
log.warning(
f'query() called without a limit on collection {collection_name}; '
f'capping at {effective_limit} results. Pass an explicit limit to avoid silent truncation.'
)
g = self._g
try:
opts = g['FtSearchOptions'](
return_fields=[
g['ReturnField'](field_identifier='id'),
g['ReturnField'](field_identifier='text'),
g['ReturnField'](field_identifier='metadata_json'),
],
limit=g['FtSearchLimit'](offset=0, count=effective_limit),
)
result = g['glide_ft'].search(self.client, self._index_name(collection_name), query_str, opts)
except g['RequestError'] as e:
log.error(f'Valkey query error on collection {collection_name}: {e}')
return None
return self._parse_glide_search_response(result, include_score=False)
def get(self, collection_name: str, limit: int | None = None) -> GetResult | None:
if not self.has_collection(collection_name):
return None
# FT.SEARCH "*" wildcard not yet in a tagged valkey-search release (tracked in #957).
# SCAN fallback is acceptable here — get() is not on the hot search path.
prefix = self._key_prefix(collection_name)
ids, documents, metadatas = [], [], []
cursor = '0'
while True:
scan_result = self.client.scan(cursor=cursor, match=f'{prefix}*', count=500)
cursor = _decode(scan_result[0])
keys = scan_result[1]
if keys:
batch = self._g['Batch'](is_atomic=False)
for key in keys:
batch.hgetall(key)
results = self.client.exec(batch, raise_on_error=False) or []
for fields in results:
if not fields:
continue
ids.append(_decode(fields.get(b'id', b'')))
documents.append(_decode(fields.get(b'text', b'')))
try:
metadatas.append(json.loads(_decode(fields.get(b'metadata_json', b'{}'))))
except (json.JSONDecodeError, TypeError):
metadatas.append({})
if limit is not None and limit > 0 and len(ids) >= limit:
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
if cursor == '0':
break
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
def delete(
self,
collection_name: str,
ids: list[str] | None = None,
filter: dict | None = None,
):
if ids:
keys = [self._item_key(collection_name, item_id) for item_id in ids]
try:
self.batch_client.delete(keys)
except self._g['RequestError'] as e:
log.error(f'Valkey delete error on collection {collection_name}: {e}')
return
if not filter:
return
filter_expr = _build_filter_expression(filter)
if not filter_expr:
return
index_name = self._index_name(collection_name)
page_size = 10000
g = self._g
while True:
try:
opts = g['FtSearchOptions'](
return_fields=[g['ReturnField'](field_identifier='id')],
limit=g['FtSearchLimit'](offset=0, count=page_size),
)
result = g['glide_ft'].search(self.client, index_name, filter_expr, opts)
except g['RequestError'] as e:
log.error(f'Valkey delete-by-filter error on collection {collection_name}: {e}')
return
if not result or result[0] == 0:
return
keys_map = result[1] if len(result) > 1 else {}
keys = [_decode(k) for k in keys_map.keys()] if isinstance(keys_map, dict) else []
if not keys:
return
self.batch_client.delete(keys)
if len(keys) < page_size:
return
def reset(self):
glide_ft = self._g['glide_ft']
collections: list[str] = []
try:
indexes = glide_ft.list(self.client) or []
idx_prefix = f'idx:{self.collection_prefix}:'
for idx in indexes:
name = _decode(idx)
if name.startswith(idx_prefix):
collections.append(name[len(idx_prefix) :])
try:
glide_ft.dropindex(self.client, idx)
log.info(f'Dropped index: {name}')
except Exception as e:
log.error(f'Error dropping index {name}: {e}')
except Exception as e:
log.error(f'Error listing indexes during reset: {e}')
for collection in collections:
self._delete_keys_by_prefix(self._key_prefix(collection))
log.info(f'Valkey vector store reset complete (prefix: {self.collection_prefix})')
def _delete_keys_by_prefix(self, prefix: str) -> None:
cursor = '0'
while True:
scan_result = self.client.scan(cursor=cursor, match=f'{prefix}*', count=500)
cursor = _decode(scan_result[0])
keys = scan_result[1]
if keys:
self.batch_client.delete(keys)
if cursor == '0':
break
@staticmethod
def _decode_kv_pairs(fields) -> dict:
"""Decode a flat [k1, v1, k2, v2, ...] wire array into a dict."""
if not fields:
return {}
if len(fields) % 2 != 0:
fields = fields[:-1]
out = {}
for k, v in zip(fields[::2], fields[1::2]):
key = _decode(k)
if isinstance(v, (bytes, bytearray)):
try:
val = v.decode()
except UnicodeDecodeError:
val = v
else:
val = v
out[key] = val
return out
@staticmethod
def _find_in_kv_pairs(pairs, target: str, case_insensitive: bool = False):
"""Look up a value in a flat [k1, v1, k2, v2, ...] array or dict."""
if isinstance(pairs, dict):
needle = target.lower() if case_insensitive else target
for k, v in pairs.items():
key = _decode(k)
if (key.lower() if case_insensitive else key) == needle:
return v
return None
if not isinstance(pairs, (list, tuple)) or len(pairs) < 2:
return None
needle = target.lower() if case_insensitive else target
for j in range(0, len(pairs) - 1, 2):
key = _decode(pairs[j])
if (key.lower() if case_insensitive else key) == needle:
return pairs[j + 1]
return None
def _parse_glide_search_response(self, result, include_score: bool) -> SearchResult | GetResult | None:
"""Parse ft.search response: [total_count, {key: {field: value, ...}, ...}]"""
empty_search = SearchResult(ids=[[]], distances=[[]], documents=[[]], metadatas=[[]])
empty_get = GetResult(ids=[[]], documents=[[]], metadatas=[[]])
if not result or result[0] == 0:
return empty_search if include_score else empty_get
docs_map = result[1] if len(result) > 1 else {}
if not isinstance(docs_map, dict):
return empty_search if include_score else empty_get
ids, documents, metadatas, distances = [], [], [], []
for _key, fields in docs_map.items():
if not isinstance(fields, dict):
continue
ids.append(_decode(fields.get(b'id', b'')))
documents.append(_decode(fields.get(b'text', b'')))
try:
metadatas.append(json.loads(_decode(fields.get(b'metadata_json', b'{}'))))
except (json.JSONDecodeError, TypeError):
metadatas.append({})
if include_score:
try:
raw_score = _decode(fields.get(b'__vector_score', b'0'))
distances.append(self._normalize_score(float(raw_score)))
except (ValueError, TypeError):
distances.append(0.0)
if not include_score:
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
return SearchResult(ids=[ids], distances=[distances], documents=[documents], metadatas=[metadatas])
def _normalize_score(self, score: float) -> float:
"""Convert valkey-search __vector_score (a distance, lower = more similar) to [0, 1] similarity.
All metrics return distance: COSINE/IP in [0, 2] for unit vectors, L2 in [0, ∞).
"""
if self.distance_metric == 'COSINE':
# COSINE distance: 0 (identical) → 2 (opposite). Map to similarity [1, -1], clamp [0, 1].
return max(0.0, min(1.0, 1.0 - score))
if self.distance_metric == 'L2':
# L2 distance: 0 (identical) → ∞.
return 1.0 / (1.0 + score)
# IP: distance = 1 - inner_product
return max(0.0, min(1.0, 1.0 - score))

View file

@ -80,6 +80,10 @@ class Vector:
from open_webui.retrieval.vector.dbs.weaviate import WeaviateClient
return WeaviateClient()
case VectorType.VALKEY:
from open_webui.retrieval.vector.dbs.valkey import ValkeyClient
return ValkeyClient()
case _:
raise ValueError(f'Unsupported vector type: {vector_type}')

View file

@ -14,3 +14,4 @@ class VectorType(StrEnum):
S3VECTOR = 's3vector'
WEAVIATE = 'weaviate'
OPENGAUSS = 'opengauss'
VALKEY = 'valkey'

View file

@ -64,5 +64,12 @@ def main():
args = parser.parse_args()
results = search_bing(args.locale, args.query, args.count, args.filter)
results = search_bing(
os.environ.get('BING_SEARCH_V7_SUBSCRIPTION_KEY', ''),
os.environ.get('BING_SEARCH_V7_ENDPOINT', 'https://api.bing.microsoft.com/v7.0/search'),
args.locale,
args.query,
args.count,
args.filter,
)
pprint(results)

View file

@ -1,20 +1,26 @@
from __future__ import annotations
import asyncio
import logging
import time
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
from open_webui.utils.session_pool import get_session
log = logging.getLogger(__name__)
# Brave free-tier rate limit: 1 request per second.
_RATE_LIMIT_RETRY_DELAY = 1.0
def search_brave(api_key: str, query: str, count: int, filter_list: list[str | None] = None) -> list[SearchResult]:
"""Search using Brave's Search API and return the results as a list of SearchResult objects.
Args:
api_key (str): A Brave Search API key
query (str): The query to search for
async def search_brave(
api_key: str,
query: str,
count: int,
filter_list: list[str | None] | None = None,
) -> list[SearchResult]:
"""Query the Brave Web Search API and return normalised results.
Retries once on HTTP 429 (rate-limit) after a short delay.
"""
url = 'https://api.search.brave.com/res/v1/web/search'
headers = {
@ -24,27 +30,27 @@ def search_brave(api_key: str, query: str, count: int, filter_list: list[str | N
}
params = {'q': query, 'count': count}
response = requests.get(url, headers=headers, params=params)
session = await get_session()
async with session.get(url, headers=headers, params=params) as response:
if response.status == 429:
log.info('Brave Search rate-limited (429); retrying after %.1fs', _RATE_LIMIT_RETRY_DELAY)
await asyncio.sleep(_RATE_LIMIT_RETRY_DELAY)
async with session.get(url, headers=headers, params=params) as retry_resp:
retry_resp.raise_for_status()
payload = await retry_resp.json()
else:
response.raise_for_status()
payload = await response.json()
# Handle 429 rate limiting - Brave free tier allows 1 request/second
# If rate limited, wait 1 second and retry once before failing
if response.status_code == 429:
log.info('Brave Search API rate limited (429), retrying after 1 second...')
time.sleep(1)
response = requests.get(url, headers=headers, params=params)
response.raise_for_status()
json_response = response.json()
results = json_response.get('web', {}).get('results', [])
web_results = payload.get('web', {}).get('results', [])
if filter_list:
results = get_filtered_results(results, filter_list)
web_results = get_filtered_results(web_results, filter_list)
return [
SearchResult(
link=result['url'],
title=result.get('title'),
snippet=result.get('description'),
link=item.get('url', ''),
title=item.get('title'),
snippet=item.get('description'),
)
for result in results[:count]
for item in web_results[:count]
]

View file

@ -197,8 +197,11 @@ def search_firecrawl(
},
timeout=count * 3 + 10,
)
# Firecrawl /search has historically returned both `{"data": [...]}`
# (flat list, what v1 did and what frost19k reported under #23966)
# and `{"data": {"web": [...]}}` (current v2). Accept either.
data = response.get('data') or {}
results = data.get('web') or []
results = data if isinstance(data, list) else (data.get('web') or [])
if filter_list:
from open_webui.retrieval.web.main import get_filtered_results

View file

@ -2,70 +2,65 @@ from __future__ import annotations
import logging
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
from open_webui.utils.session_pool import get_session
log = logging.getLogger(__name__)
def search_google_pse(
async def search_google_pse(
api_key: str,
search_engine_id: str,
query: str,
count: int,
filter_list: list[str | None] = None,
filter_list: list[str | None] | None = None,
referer: str | None = None,
) -> list[SearchResult]:
"""Search using Google's Programmable Search Engine API and return the results as a list of SearchResult objects.
Handles pagination for counts greater than 10.
"""Query Google Programmable Search Engine with automatic pagination.
Args:
api_key (str): A Programmable Search Engine API key
search_engine_id (str): A Programmable Search Engine ID
query (str): The query to search for
count (int): The number of results to return (max 100, as PSE max results per query is 10 and max page is 10)
filter_list (list[str | None], optional): A list of keywords to filter out from results. Defaults to None.
Returns:
list[SearchResult]: A list of SearchResult objects.
The PSE API returns at most 10 results per request, so this function
issues multiple requests when ``count > 10``.
"""
url = 'https://www.googleapis.com/customsearch/v1'
headers = {'Content-Type': 'application/json'}
headers: dict[str, str] = {'Content-Type': 'application/json'}
if referer:
headers['Referer'] = referer
all_results = []
start_index = 1 # Google PSE start parameter is 1-based
all_items: list[dict] = []
start_index = 1 # PSE uses 1-based pagination
while count > 0:
num_results_this_page = min(count, 10) # Google PSE max results per page is 10
session = await get_session()
remaining = count
while remaining > 0:
page_size = min(remaining, 10)
params = {
'cx': search_engine_id,
'q': query,
'key': api_key,
'num': num_results_this_page,
'start': start_index,
'num': str(page_size),
'start': str(start_index),
}
response = requests.request('GET', url, headers=headers, params=params)
response.raise_for_status()
json_response = response.json()
results = json_response.get('items', [])
if results: # check if results are returned. If not, no more pages to fetch.
all_results.extend(results)
count -= len(results) # Decrement count by the number of results fetched in this page.
start_index += 10 # Increment start index for the next page
else:
break # No more results from Google PSE, break the loop
async with session.get(url, headers=headers, params=params) as response:
response.raise_for_status()
payload = await response.json()
items = payload.get('items', [])
if not items:
break
all_items.extend(items)
remaining -= len(items)
start_index += 10
if filter_list:
all_results = get_filtered_results(all_results, filter_list)
all_items = get_filtered_results(all_items, filter_list)
return [
SearchResult(
link=result['link'],
title=result.get('title'),
snippet=result.get('snippet'),
link=item.get('link', ''),
title=item.get('title'),
snippet=item.get('snippet'),
)
for result in all_results
for item in all_items
]

View file

@ -17,21 +17,20 @@ def search_kagi(api_key: str, query: str, count: int, filter_list: Optional[list
query (str): The query to search for
count (int): The number of results to return
"""
url = 'https://kagi.com/api/v0/search'
url = 'https://kagi.com/api/v1/search'
headers = {
'Authorization': f'Bot {api_key}',
'Authorization': f'Bearer {api_key}',
}
params = {'q': query, 'limit': count}
params = {'query': query, 'limit': count}
response = requests.get(url, headers=headers, params=params)
response = requests.post(url, headers=headers, json=params)
response.raise_for_status()
json_response = response.json()
search_results = json_response.get('data', [])
search_results = json_response.get('data', {}).get('search', [])
results = [
SearchResult(link=result['url'], title=result['title'], snippet=result.get('snippet'))
for result in search_results
if result['t'] == 0
]
print(results)

View file

@ -0,0 +1,74 @@
import logging
from typing import Optional
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
log = logging.getLogger(__name__)
DEFAULT_LINKUP_PARAMS = {
'url': 'https://api.linkup.so/v1/search',
'depth': 'standard',
'outputType': 'sourcedAnswer',
}
def search_linkup(
api_key: str,
query: str,
count: int,
filter_list: Optional[list[str]] = None,
params: Optional[dict] = None,
) -> list[SearchResult]:
"""Search using the Linkup Search API.
``params`` is forwarded almost verbatim as the JSON body; only ``q``
and ``maxResults`` are injected automatically. The special key
``url`` (default ``https://api.linkup.so/v1/search``) is popped and
used as the endpoint.
"""
if hasattr(api_key, '__str__'):
api_key = str(api_key)
merged = {**DEFAULT_LINKUP_PARAMS, **(params or {})}
api_url = str(merged.pop('url', DEFAULT_LINKUP_PARAMS['url']))
payload = {**merged, 'q': query, 'maxResults': count}
try:
response = requests.post(
api_url,
headers={
'Authorization': f'Bearer {api_key}',
'Content-Type': 'application/json',
},
json=payload,
timeout=30,
)
response.raise_for_status()
json_response = response.json()
output_type = merged.get('outputType', 'sourcedAnswer')
search_results = (
json_response.get('sources', []) if output_type == 'sourcedAnswer' else json_response.get('results', [])
)
if filter_list:
search_results = get_filtered_results(search_results, filter_list)
return [
SearchResult(
link=r.get('url', ''),
title=r.get('name') or r.get('title'),
snippet=r.get('content') or r.get('text') or r.get('snippet'),
)
for r in search_results
][:count]
except requests.exceptions.RequestException as e:
log.error(f'Linkup API request failed: {e}')
raise Exception(f'Linkup search failed: {str(e)}')
except Exception as e:
log.error(f'Error searching Linkup: {e}')
raise Exception(f'Linkup search error: {str(e)}')

View file

@ -2,6 +2,7 @@ import logging
from typing import Literal, Optional
import requests
from open_webui.env import VERSION
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
MODELS = Literal[
@ -64,6 +65,7 @@ def search_perplexity(
headers = {
'Authorization': f'Bearer {api_key}',
'Content-Type': 'application/json',
'X-Pplx-Integration': f'open-webui/{VERSION}',
}
# Make the API request

View file

@ -2,6 +2,7 @@ import logging
from typing import Literal, Optional
import requests
from open_webui.env import VERSION
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
from open_webui.utils.headers import include_user_info_headers
@ -47,6 +48,7 @@ def search_perplexity_search(
headers = {
'Authorization': f'Bearer {api_key}',
'Content-Type': 'application/json',
'X-Pplx-Integration': f'open-webui/{VERSION}',
}
# Forward user info headers if user is provided

View file

@ -2,87 +2,65 @@ from __future__ import annotations
import logging
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
from open_webui.utils.session_pool import get_session
log = logging.getLogger(__name__)
# SearXNG request headers — identifies the bot to instance operators.
_SEARXNG_HEADERS = {
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) RAG Bot',
'Accept': 'text/html',
'Accept-Encoding': 'gzip, deflate',
'Accept-Language': 'en-US,en;q=0.5',
'Connection': 'keep-alive',
}
def search_searxng( # noqa: PLR0913
async def search_searxng(
query_url: str,
query: str,
count: int,
filter_list: list[str | None] = None,
filter_list: list[str | None] | None = None,
**kwargs,
) -> list[SearchResult]:
"""Query a SearXNG instance and return results sorted by relevance score.
Optional keyword arguments (language, safesearch, time_range, categories)
are forwarded directly as SearXNG query parameters.
"""
Search a SearXNG instance for a given query and return the results as a list of SearchResult objects.
The function allows passing additional parameters such as language or time_range to tailor the search result.
Args:
query_url (str): The base URL of the SearXNG server.
query (str): The search term or question to find in the SearXNG database.
count (int): The maximum number of results to retrieve from the search.
Keyword Args:
language (str): Language filter for the search results; e.g., "all", "en-US", "es". Defaults to "all".
safesearch (int): Safe search filter for safer web results; 0 = off, 1 = moderate, 2 = strict. Defaults to 1 (moderate).
time_range (str): Time range for filtering results by date; e.g., "2023-04-05..today" or "all-time". Defaults to ''.
categories: (list[str | None]): Specific categories within which the search should be performed, defaulting to an empty string if not provided.
Returns:
list[SearchResult]: A list of SearchResults sorted by relevance score in descending order.
Raise:
requests.exceptions.RequestException: If a request error occurs during the search process.
"""
# Default values for optional parameters are provided as empty strings or None when not specified.
language = kwargs.get('language', 'all').strip().rstrip(',')
safesearch = kwargs.get('safesearch', '1')
time_range = kwargs.get('time_range', '')
categories = ''.join(kwargs.get('categories', []))
# Normalise legacy ``<query>``-style URLs by stripping any query string.
if '<query>' in query_url:
query_url = query_url.split('?')[0]
params = {
'q': query,
'format': 'json',
'pageno': 1,
'safesearch': safesearch,
'language': language,
'time_range': time_range,
'categories': categories,
'safesearch': kwargs.get('safesearch', '1'),
'language': kwargs.get('language', 'all').strip().rstrip(','),
'time_range': kwargs.get('time_range', ''),
'categories': ''.join(kwargs.get('categories', [])),
'theme': 'simple',
'image_proxy': 0,
}
# Legacy query format
if '<query>' in query_url:
# Strip all query parameters from the URL
query_url = query_url.split('?')[0]
log.debug('searching %s', query_url)
log.debug(f'searching {query_url}')
session = await get_session()
async with session.get(query_url, headers=_SEARXNG_HEADERS, params=params) as response:
response.raise_for_status()
payload = await response.json()
response = requests.get(
query_url,
headers={
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) RAG Bot',
'Accept': 'text/html',
'Accept-Encoding': 'gzip, deflate',
'Accept-Language': 'en-US,en;q=0.5',
'Connection': 'keep-alive',
},
params=params,
)
response.raise_for_status() # Raise an exception for HTTP errors.
json_response = response.json()
results = json_response.get('results', [])
sorted_results = sorted(results, key=lambda x: x.get('score', 0), reverse=True)
results = sorted(payload.get('results', []), key=lambda x: x.get('score', 0), reverse=True)
if filter_list:
sorted_results = get_filtered_results(sorted_results, filter_list)
results = get_filtered_results(results, filter_list)
return [
SearchResult(link=result['url'], title=result.get('title'), snippet=result.get('content'))
for result in sorted_results[:count]
SearchResult(
link=item.get('url', ''),
title=item.get('title'),
snippet=item.get('content'),
)
for item in results[:count]
]

View file

@ -3,36 +3,39 @@ from __future__ import annotations
import json
import logging
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
from open_webui.utils.session_pool import get_session
log = logging.getLogger(__name__)
def search_serper(api_key: str, query: str, count: int, filter_list: list[str | None] = None) -> list[SearchResult]:
"""Search using serper.dev's API and return the results as a list of SearchResult objects.
async def search_serper(
api_key: str,
query: str,
count: int,
filter_list: list[str | None] | None = None,
) -> list[SearchResult]:
"""Query the serper.dev Google Search API and return normalised results.
Args:
api_key (str): A serper.dev API key
query (str): The query to search for
Results are sorted by their position field before truncation.
"""
url = 'https://google.serper.dev/search'
payload = json.dumps({'q': query})
headers = {'X-API-KEY': api_key, 'Content-Type': 'application/json'}
response = requests.request('POST', url, headers=headers, data=payload)
response.raise_for_status()
session = await get_session()
async with session.post(url, headers=headers, data=json.dumps({'q': query})) as response:
response.raise_for_status()
payload = await response.json()
json_response = response.json()
results = sorted(json_response.get('organic', []), key=lambda x: x.get('position', 0))
organic = sorted(payload.get('organic', []), key=lambda item: item.get('position', 0))
if filter_list:
results = get_filtered_results(results, filter_list)
organic = get_filtered_results(organic, filter_list)
return [
SearchResult(
link=result['link'],
title=result.get('title'),
snippet=result.get('snippet'),
link=item.get('link', ''),
title=item.get('title'),
snippet=item.get('snippet'),
)
for result in results[:count]
for item in organic[:count]
]

View file

@ -2,42 +2,41 @@ from __future__ import annotations
import logging
import requests
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
from open_webui.utils.session_pool import get_session
log = logging.getLogger(__name__)
def search_serpstack(
async def search_serpstack(
api_key: str,
query: str,
count: int,
filter_list: list[str | None] = None,
filter_list: list[str | None] | None = None,
https_enabled: bool = True,
) -> list[SearchResult]:
"""Search using serpstack.com's and return the results as a list of SearchResult objects.
"""Query the serpstack.com API and return normalised results.
Args:
api_key (str): A serpstack.com API key
query (str): The query to search for
https_enabled (bool): Whether to use HTTPS or HTTP for the API request
Uses HTTPS by default; set ``https_enabled=False`` for free-tier HTTP access.
"""
url = f'{"https" if https_enabled else "http"}://api.serpstack.com/search'
scheme = 'https' if https_enabled else 'http'
url = f'{scheme}://api.serpstack.com/search'
params = {'access_key': api_key, 'query': query}
headers = {'Content-Type': 'application/json'}
params = {
'access_key': api_key,
'query': query,
}
session = await get_session()
async with session.get(url, params=params) as response:
response.raise_for_status()
payload = await response.json()
response = requests.request('POST', url, headers=headers, params=params)
response.raise_for_status()
json_response = response.json()
results = sorted(json_response.get('organic_results', []), key=lambda x: x.get('position', 0))
organic = sorted(payload.get('organic_results', []), key=lambda x: x.get('position', 0))
if filter_list:
results = get_filtered_results(results, filter_list)
organic = get_filtered_results(organic, filter_list)
return [
SearchResult(link=result['url'], title=result.get('title'), snippet=result.get('snippet'))
for result in results[:count]
SearchResult(
link=item.get('url', ''),
title=item.get('title'),
snippet=item.get('snippet'),
)
for item in organic[:count]
]

View file

@ -19,9 +19,13 @@ from typing import (
)
import aiohttp
import aiohttp.resolver
import certifi
import requests
import urllib3.connection
import urllib3.connectionpool
import validators
from requests.adapters import HTTPAdapter
from fastapi.concurrency import run_in_threadpool
from langchain_community.document_loaders import PlaywrightURLLoader, WebBaseLoader
from langchain_community.document_loaders.base import BaseLoader
@ -94,7 +98,7 @@ def validate_url(url: Union[str, Sequence[str]]):
# Get IPv4 and IPv6 addresses
ipv4_addresses, ipv6_addresses = resolve_hostname(parsed_url.hostname)
# Check if any of the resolved addresses are private
# This is technically still vulnerable to DNS rebinding attacks, as we don't control WebBaseLoader
# DNS rebinding is mitigated at the connection layer; see _SSRFSafeResolver / _SSRFSafeAdapter
for ip in ipv4_addresses + ipv6_addresses:
addr = ipaddress.ip_address(ip)
if not addr.is_global:
@ -118,6 +122,81 @@ def safe_validate_urls(url: Sequence[str]) -> Sequence[str]:
return valid_urls
def _ssrf_safe_new_conn(self):
"""Resolve DNS, validate all IPs are global, connect to validated IP.
Replaces urllib3's _new_conn so the DNS lookup that feeds the actual TCP
connect is the same one we validate — no second resolution, no rebinding
window.
"""
host = getattr(self, '_dns_host', self.host)
port = self.port
infos = socket.getaddrinfo(host, port, 0, socket.SOCK_STREAM)
if not infos:
raise OSError(f'getaddrinfo for {host!r} returned empty list')
if not ENABLE_RAG_LOCAL_WEB_FETCH:
for _, _, _, _, sa in infos:
if not ipaddress.ip_address(sa[0]).is_global:
raise ValueError(ERROR_MESSAGES.INVALID_URL)
err = None
for fam, typ, proto, _, sa in infos:
sock = None
try:
sock = socket.socket(fam, typ, proto)
if self.timeout is not socket._GLOBAL_DEFAULT_TIMEOUT:
sock.settimeout(self.timeout)
if getattr(self, 'source_address', None):
sock.bind(self.source_address)
for opt in getattr(self, 'socket_options', None) or ():
sock.setsockopt(*opt)
sock.connect(sa)
return sock
except OSError as exc:
err = exc
if sock is not None:
sock.close()
raise err or OSError(f'connect to {host!r}:{port} failed')
class _SafeHTTPConn(urllib3.connection.HTTPConnection):
_new_conn = _ssrf_safe_new_conn
class _SafeHTTPSConn(urllib3.connection.HTTPSConnection):
_new_conn = _ssrf_safe_new_conn
class _SafeHTTPPool(urllib3.connectionpool.HTTPConnectionPool):
ConnectionCls = _SafeHTTPConn
class _SafeHTTPSPool(urllib3.connectionpool.HTTPSConnectionPool):
ConnectionCls = _SafeHTTPSConn
class _SSRFSafeAdapter(HTTPAdapter):
"""requests transport adapter that validates resolved IPs at connect time."""
def init_poolmanager(self, *args, **kwargs):
super().init_poolmanager(*args, **kwargs)
self.poolmanager.pool_classes_by_scheme = {
'http': _SafeHTTPPool,
'https': _SafeHTTPSPool,
}
class _SSRFSafeResolver(aiohttp.resolver.DefaultResolver):
"""aiohttp resolver that rejects non-global IPs unless local fetch is on."""
async def resolve(self, host, port=0, family=socket.AF_INET):
results = await super().resolve(host, port, family)
if not ENABLE_RAG_LOCAL_WEB_FETCH:
for entry in results:
if not ipaddress.ip_address(entry['host']).is_global:
raise ValueError(ERROR_MESSAGES.INVALID_URL)
return results
def extract_metadata(soup, url):
metadata = {'source': url}
if title := soup.find('title'):
@ -421,6 +500,62 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
self.trust_env = trust_env
self.playwright_timeout = playwright_timeout
def _intercept_navigation_sync(self, route, request=None):
req = request or route.request
if req.resource_type != 'document':
route.continue_()
return
try:
validate_url(req.url)
except Exception:
route.abort()
return
if AIOHTTP_CLIENT_ALLOW_REDIRECTS:
resp = route.fetch()
else:
try:
resp = route.fetch(max_redirects=0)
except TypeError:
route.abort()
return
if 300 <= resp.status < 400:
route.abort()
return
route.fulfill(response=resp)
async def _intercept_navigation(self, route, request=None):
req = request or route.request
if req.resource_type != 'document':
await route.continue_()
return
try:
await run_in_threadpool(validate_url, req.url)
except Exception:
await route.abort()
return
if AIOHTTP_CLIENT_ALLOW_REDIRECTS:
resp = await route.fetch()
else:
try:
resp = await route.fetch(max_redirects=0)
except TypeError:
await route.abort()
return
if 300 <= resp.status < 400:
await route.abort()
return
await route.fulfill(response=resp)
def lazy_load(self) -> Iterator[Document]:
"""Safely load URLs synchronously with support for remote browser."""
from playwright.sync_api import sync_playwright
@ -436,6 +571,7 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
try:
self._safe_process_url_sync(url)
page = browser.new_page()
page.route('**/*', self._intercept_navigation_sync)
response = page.goto(url, timeout=self.playwright_timeout)
if response is None:
raise ValueError(f'page.goto() returned None for url {url}')
@ -465,6 +601,7 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
try:
await self._safe_process_url(url)
page = await browser.new_page()
await page.route('**/*', self._intercept_navigation)
response = await page.goto(url, timeout=self.playwright_timeout)
if response is None:
raise ValueError(f'page.goto() returned None for url {url}')
@ -512,8 +649,12 @@ class SafeWebBaseLoader(WebBaseLoader):
'allow_redirects': AIOHTTP_CLIENT_ALLOW_REDIRECTS,
}
self.session.mount('http://', _SSRFSafeAdapter())
self.session.mount('https://', _SSRFSafeAdapter())
async def _fetch(self, url: str, retries: int = 3, cooldown: int = 2, backoff: float = 1.5) -> str:
async with aiohttp.ClientSession(trust_env=self.trust_env) as session:
connector = aiohttp.TCPConnector(resolver=_SSRFSafeResolver())
async with aiohttp.ClientSession(trust_env=self.trust_env, connector=connector) as session:
for i in range(retries):
try:
kwargs: Dict = dict(

File diff suppressed because it is too large Load diff

View file

@ -286,8 +286,9 @@ async def update_password(
session_user=Depends(get_current_user),
db: AsyncSession = Depends(get_async_session),
):
# Trusted-header auth mode delegates passwords to the reverse proxy
if WEBUI_AUTH_TRUSTED_EMAIL_HEADER:
raise HTTPException(400, detail=ERROR_MESSAGES.ACTION_PROHIBITED)
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.ACTION_PROHIBITED)
if session_user:
user = await Auths.authenticate_user(
session_user.email,
@ -578,7 +579,7 @@ async def signin(
if WEBUI_AUTH_TRUSTED_EMAIL_HEADER:
if WEBUI_AUTH_TRUSTED_EMAIL_HEADER not in request.headers:
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_TRUSTED_HEADER)
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_TRUSTED_HEADER)
email = request.headers[WEBUI_AUTH_TRUSTED_EMAIL_HEADER].lower()
name = email
@ -746,9 +747,12 @@ async def signup(
has_users = await Users.has_users(db=db)
if WEBUI_AUTH:
if not request.app.state.config.ENABLE_SIGNUP or not request.app.state.config.ENABLE_LOGIN_FORM:
if has_users or not ENABLE_INITIAL_ADMIN_SIGNUP:
if has_users:
if not request.app.state.config.ENABLE_SIGNUP or not request.app.state.config.ENABLE_LOGIN_FORM:
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
# Don't gate the first admin on ENABLE_SIGNUP: it auto-disables and can persist stale across a DB reset.
elif not request.app.state.config.ENABLE_LOGIN_FORM and not ENABLE_INITIAL_ADMIN_SIGNUP:
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
else:
if has_users:
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)

View file

@ -301,6 +301,12 @@ async def update_event(
await _check_calendar_access(event.calendar_id, user, 'write')
# A new calendar_id in the form moves the event; require write access on the
# destination too, mirroring create_event. Without this, write on the source
# calendar alone is enough to inject an event into any other calendar.
if form_data.calendar_id is not None and form_data.calendar_id != event.calendar_id:
await _check_calendar_access(form_data.calendar_id, user, 'write')
updated = await CalendarEvents.update_event_by_id(event_id, form_data)
if not updated:
raise HTTPException(status_code=500, detail='Failed to update')

View file

@ -911,6 +911,7 @@ async def model_response_handler(request, channel, message, user, db=None):
thread_history = []
images = []
files = []
# Batch fetch all users in a single query (fixes N+1 problem)
user_ids = list({message.user_id for message in thread_messages})
@ -934,9 +935,11 @@ async def model_response_handler(request, channel, message, user, db=None):
if file.get('type', '') == 'image':
images.append(file.get('url', ''))
elif file.get('content_type', '').startswith('image/'):
image = await get_image_base64_from_file_id(file.get('id', ''))
image = await get_image_base64_from_file_id(file.get('id', ''), user=user)
if image:
images.append(image)
elif file.get('id'):
files.append(file)
thread_history_string = '\n\n'.join(thread_history)
system_message = {
@ -994,6 +997,8 @@ async def model_response_handler(request, channel, message, user, db=None):
'session_id': f'channel:{channel.id}',
'background_tasks': {},
}
if files:
form_data['files'] = files
if tool_ids:
form_data['tool_ids'] = tool_ids
if features:

View file

@ -30,6 +30,7 @@ from open_webui.models.folders import Folders
from open_webui.models.shared_chats import SharedChatResponse, SharedChats
from open_webui.models.tags import TagModel, Tags
from open_webui.socket.main import get_event_emitter
from open_webui.tasks import stop_item_tasks
from open_webui.utils.access_control import filter_allowed_access_grants, has_permission
from open_webui.utils.auth import get_admin_user, get_verified_user
from open_webui.utils.middleware import serialize_output
@ -521,17 +522,13 @@ async def get_user_chat_list_by_user_id(
user=Depends(get_admin_user),
db: AsyncSession = Depends(get_async_session),
):
"""List chat summaries for a given user (admin-only endpoint)."""
if not ENABLE_ADMIN_CHAT_ACCESS:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
if page is None:
page = 1
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
effective_page = page if page is not None else 1
limit = 60
skip = (page - 1) * limit
skip = (effective_page - 1) * limit
filter = {}
if query:
@ -762,10 +759,7 @@ async def get_all_user_tags(user=Depends(get_verified_user), db: AsyncSession =
@router.get('/all/db', response_model=list[ChatResponse])
async def get_all_user_chats_in_db(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
if not ENABLE_ADMIN_EXPORT:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
return [ChatResponse(**chat.model_dump()) for chat in await Chats.get_chats(db=db)]
@ -979,14 +973,25 @@ async def update_chat_by_id(
if chat:
updated_chat = {**chat.chat, **form_data.chat}
# Re-derive content from output for assistant messages so that
# frontend edits to output items are always reflected in content.
# serialize_output() is the single source of truth for this conversion.
for msg in updated_chat.get('history', {}).get('messages', {}).values():
# Re-derive content from output for assistant messages so that frontend
# edits to output items are reflected in content. Only when output
# actually changed — otherwise content set independently of output
# (e.g. a `replace` event or an outlet filter footer) would be reverted.
existing_messages = (chat.chat.get('history') or {}).get('messages') or {}
for msg_id, msg in updated_chat.get('history', {}).get('messages', {}).items():
if msg.get('role') == 'assistant' and msg.get('output'):
msg['content'] = serialize_output(msg['output'])
if msg.get('output') != existing_messages.get(msg_id, {}).get('output'):
msg['content'] = serialize_output(msg['output'])
chat = await Chats.update_chat_by_id(id, updated_chat, db=db)
# Reconcile chat_message rows with the committed blob.
# This is the only caller where the frontend pushes a full
# history with potential edits, deletions, or new branches.
messages = (updated_chat.get('history') or {}).get('messages') or {}
if messages:
await Chats.reconcile_messages_by_chat_id(id, user.id, messages)
return ChatResponse(**chat.model_dump())
else:
raise HTTPException(
@ -1116,6 +1121,10 @@ async def delete_chat_by_id(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
# Cancel any in-flight LLM tasks (streaming, title/tags generation)
# before deleting the chat to prevent orphaned requests.
await stop_item_tasks(request.app.state.redis, id)
if user.role == 'admin':
chat = await Chats.get_chat_by_id(id, db=db)
if not chat:
@ -1305,13 +1314,20 @@ async def clone_shared_chat_by_id(
@router.post('/{id}/archive', response_model=ChatResponse | None)
async def archive_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
async def archive_chat_by_id(
request: Request,
id: str,
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
chat = await Chats.toggle_chat_archive_by_id(id, db=db)
tag_ids = chat.meta.get('tags', [])
if chat.archived:
# Cancel any in-flight LLM tasks before archiving
await stop_item_tasks(request.app.state.redis, id)
# Archived chats are excluded from count — clean up orphans
await Chats.delete_orphan_tags_for_user(tag_ids, user.id, db=db)
else:
@ -1323,9 +1339,7 @@ async def archive_chat_by_id(id: str, user=Depends(get_verified_user), db: Async
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
############################
# ShareChatById
############################
# --- Share Chat ---
@router.post('/{id}/share', response_model=ChatResponse | None)
@ -1335,51 +1349,35 @@ async def share_chat_by_id(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
if (user.role != 'admin') and (
not await has_permission(user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS)
if user.role != 'admin' and not await has_permission(
user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if not chat:
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
if chat:
if chat.share_id:
# Re-snapshot existing share
shared = await SharedChats.update(chat.share_id, db=db)
if shared:
# Re-fetch the original chat to return
chat = await Chats.get_chat_by_id(id, db=db)
return ChatResponse(**chat.model_dump())
# If a share already exists, re-snapshot it
if chat.share_id:
shared = await SharedChats.update(chat.share_id, db=db)
if shared:
chat = await Chats.get_chat_by_id(id, db=db)
return ChatResponse(**chat.model_dump())
# Create new share
shared = await SharedChats.create(id, user.id, db=db)
if not shared:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=ERROR_MESSAGES.DEFAULT(),
)
# Set share_id on the original chat
chat = await Chats.update_chat_share_id_by_id(id, shared.id, db=db)
if not chat:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=ERROR_MESSAGES.DEFAULT(),
)
return ChatResponse(**chat.model_dump())
# Create a new share
shared = await SharedChats.create(id, user.id, db=db)
if not shared:
raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, detail=ERROR_MESSAGES.DEFAULT())
else:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
chat = await Chats.update_chat_share_id_by_id(id, shared.id, db=db)
if not chat:
raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, detail=ERROR_MESSAGES.DEFAULT())
return ChatResponse(**chat.model_dump())
############################
# DeleteSharedChatById
############################
# --- Delete Shared Chat ---
@router.delete('/{id}/share', response_model=bool | None)
@ -1387,22 +1385,17 @@ async def delete_shared_chat_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
if not chat.share_id:
return False
if not chat:
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
await SharedChats.delete_by_chat_id(id, db=db)
await Chats.update_chat_share_id_by_id(id, None, db=db)
if not chat.share_id:
return False
# Revoke all access grants for this shared chat
await AccessGrants.set_access_grants('shared_chat', id, [], db=db)
await SharedChats.delete_by_chat_id(id, db=db)
await Chats.update_chat_share_id_by_id(id, None, db=db)
await AccessGrants.set_access_grants('shared_chat', id, [], db=db)
return True
else:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
return True
############################
@ -1583,23 +1576,3 @@ async def delete_tag_by_id_and_tag_name(
return await Tags.get_tags_by_ids_and_user_id(tags, user.id, db=db)
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
############################
# DeleteAllTagsById
############################
@router.delete('/{id}/tags/all', response_model=bool | None)
async def delete_all_tags_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
if chat:
old_tags = chat.meta.get('tags', [])
await Chats.delete_all_tags_by_id_and_user_id(id, user.id, db=db)
await Chats.delete_orphan_tags_for_user(old_tags, user.id, db=db)
return True
else:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)

View file

@ -10,9 +10,7 @@ from open_webui.models.feedbacks import (
FeedbackIdResponse,
FeedbackListResponse,
FeedbackModel,
FeedbackResponse,
Feedbacks,
FeedbackUserResponse,
LeaderboardFeedbackData,
ModelHistoryEntry,
ModelHistoryResponse,
@ -295,12 +293,6 @@ async def get_feedback_model_ids(user=Depends(get_admin_user), db: AsyncSession
return await Feedbacks.get_distinct_model_ids(db=db)
@router.get('/feedbacks/all', response_model=list[FeedbackResponse])
async def get_all_feedbacks(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
feedbacks = await Feedbacks.get_all_feedbacks(db=db)
return feedbacks
@router.get('/feedbacks/all/ids', response_model=list[FeedbackIdResponse])
async def get_all_feedback_ids(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
return await Feedbacks.get_all_feedback_ids(db=db)
@ -324,10 +316,19 @@ async def export_all_feedbacks(
return feedbacks
@router.get('/feedbacks/user', response_model=list[FeedbackUserResponse])
async def get_feedbacks(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
feedbacks = await Feedbacks.get_feedbacks_by_user_id(user.id, db=db)
return feedbacks
PAGE_ITEM_COUNT = 30
@router.get('/feedbacks/user', response_model=FeedbackListResponse)
async def get_user_feedbacks(
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
limit = PAGE_ITEM_COUNT
page = max(1, page)
skip = (page - 1) * limit
return await Feedbacks.get_feedbacks_by_user_id(user.id, skip=skip, limit=limit, db=db)
@router.delete('/feedbacks', response_model=bool)
@ -336,9 +337,6 @@ async def delete_feedbacks(user=Depends(get_verified_user), db: AsyncSession = D
return success
PAGE_ITEM_COUNT = 30
@router.get('/feedbacks/list', response_model=FeedbackListResponse)
async def get_feedbacks(
order_by: Optional[str] = None,

View file

@ -1,4 +1,5 @@
import asyncio
import hashlib
import json
import logging
import os
@ -60,7 +61,11 @@ from open_webui.utils.access_control.files import has_access_to_file
def _is_text_file(file_path: str, chunk_size: int = 8192) -> bool:
"""Check if a file is likely a text file by reading a chunk and validating UTF-8.
"""Check if a file is likely a text file by reading a chunk and decoding it.
Tries UTF-8 first, then falls back to Latin-1 (which accepts every byte
in 0x00–0xFF) so that legacy-encoded files from Windows environments are
not misclassified as binary.
This catches files whose extensions are mis-mapped by mimetypes/browsers
(e.g. TypeScript .ts → video/mp2t) without maintaining an extension whitelist.
@ -74,9 +79,15 @@ def _is_text_file(file_path: str, chunk_size: int = 8192) -> bool:
# Null bytes are a strong indicator of binary content
if b'\x00' in chunk:
return False
chunk.decode('utf-8')
try:
chunk.decode('utf-8')
except UnicodeDecodeError:
# Latin-1 always succeeds (every byte is valid), so this
# effectively just means "the file has no null bytes and is
# therefore likely text, even if not valid UTF-8".
chunk.decode('latin-1')
return True
except (UnicodeDecodeError, Exception):
except Exception:
return False
@ -112,15 +123,12 @@ async def process_uploaded_file(
if _is_text_file(file_path):
content_type = 'text/plain'
stt_supported = getattr(
request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', []
)
stt_supported = getattr(request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', [])
if content_type and strict_match_mime_type(stt_supported, content_type):
# Audio / STT-supported files → transcribe then index
file_path_processed = await asyncio.to_thread(Storage.get_file, file_path)
result = await asyncio.to_thread(
transcribe,
result = await transcribe(
request,
file_path_processed,
file_metadata,
@ -151,17 +159,12 @@ async def process_uploaded_file(
db=db_session,
)
else:
raise Exception(
f'File type {content_type} is not supported for processing'
)
raise Exception(f'File type {content_type} is not supported for processing')
else:
# Documents, or any file when an external engine is configured
if not content_type:
log.info(
f'File type {file.content_type} is not provided, '
'but trying to process anyway'
)
log.info(f'File type {file.content_type} is not provided, but trying to process anyway')
await process_file(
request,
ProcessFileForm(file_id=file_item.id),
@ -169,6 +172,28 @@ async def process_uploaded_file(
db=db_session,
)
# Auto-link to Knowledge Collection when uploaded from one (#24807).
# Mirrors POST /knowledge/{id}/file/add so linking doesn't depend
# on the frontend staying connected after upload.
knowledge_id = file_metadata.get('knowledge_id')
if knowledge_id:
try:
await Knowledges.add_file_to_knowledge_by_id(
knowledge_id=knowledge_id,
file_id=file_item.id,
user_id=user.id,
directory_id=file_metadata.get('directory_id'),
)
await process_file(
request,
ProcessFileForm(file_id=file_item.id, collection_name=knowledge_id),
user=user,
db=db_session,
)
log.info(f'Linked file {file_item.id} to knowledge {knowledge_id}')
except Exception as e:
log.warning(f'Failed to link file {file_item.id} to knowledge {knowledge_id}: {e}')
except Exception as e:
log.error(f'Error processing file: {file_item.id}')
await Files.update_file_data_by_id(
@ -270,6 +295,10 @@ async def upload_file_handler(
},
)
# SHA-256 of raw uploaded bytes for incremental sync diffing.
# If the client pre-computed and sent file_hash, use that.
file_hash = file_metadata.get('file_hash') or hashlib.sha256(contents).hexdigest()
file_item = await Files.insert_new_file(
user.id,
FileForm(
@ -284,6 +313,7 @@ async def upload_file_handler(
'name': name,
'content_type': (file.content_type if isinstance(file.content_type, str) else None),
'size': len(contents),
'file_hash': file_hash,
'data': file_metadata,
},
}

View file

@ -336,29 +336,16 @@ async def preview_group_access(
return {
'group': {'id': group.id, 'name': group.name},
'models': {
'items': [
{'id': m.id, 'name': m.name}
for m in active_models
if m.id in accessible_model_ids
],
'items': [{'id': m.id, 'name': m.name} for m in active_models if m.id in accessible_model_ids],
'total': len(active_models),
},
'knowledge': {
'items': [
{'id': k.id, 'name': k.name}
for k in all_knowledge
if k.id in accessible_knowledge_ids
],
'items': [{'id': k.id, 'name': k.name} for k in all_knowledge if k.id in accessible_knowledge_ids],
'total': len(all_knowledge),
},
'tools': {
'items': [
{'id': t.id, 'name': t.name}
for t in all_tools
if t.id in accessible_tool_ids
],
'items': [{'id': t.id, 'name': t.name} for t in all_tools if t.id in accessible_tool_ids],
'total': len(all_tools),
},
'permissions': group.permissions or {},
}

View file

@ -447,6 +447,7 @@ def _is_same_origin(url: str, base_url: str) -> bool:
and comparing the three origin components eliminates those
attack vectors.
"""
def _default_port(scheme: str) -> int:
return 443 if scheme == 'https' else 80
@ -455,8 +456,7 @@ def _is_same_origin(url: str, base_url: str) -> bool:
return (
parsed.scheme == trusted.scheme
and parsed.hostname == trusted.hostname
and (parsed.port or _default_port(parsed.scheme))
== (trusted.port or _default_port(trusted.scheme))
and (parsed.port or _default_port(parsed.scheme)) == (trusted.port or _default_port(trusted.scheme))
)
@ -626,7 +626,7 @@ async def image_generations(
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
res = await r.json(content_type=None)
images = []
@ -676,7 +676,7 @@ async def image_generations(
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
res = await r.json(content_type=None)
images = []
@ -743,7 +743,8 @@ async def image_generations(
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
image_data, content_type = await get_image_data(
image['url'], headers,
image['url'],
headers,
trusted_base_url=request.app.state.config.COMFYUI_BASE_URL,
)
_, url = await upload_image(
@ -785,7 +786,7 @@ async def image_generations(
headers={'authorization': get_automatic1111_api_auth(request)},
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
res = await r.json()
res = await r.json(content_type=None)
log.debug(f'res: {res}')
images = []
@ -960,7 +961,7 @@ async def image_edits(
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
res = await r.json(content_type=None)
images = []
for image in res['data']:
@ -1015,7 +1016,7 @@ async def image_edits(
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
res = await r.json()
res = await r.json(content_type=None)
images = []
for image in res['candidates']:
@ -1102,7 +1103,8 @@ async def image_edits(
headers = {'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY}'}
image_data, content_type = await get_image_data(
image_url, headers,
image_url,
headers,
trusted_base_url=request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
)
_, url = await upload_image(

View file

@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import io
import logging
import zipfile
@ -12,7 +13,7 @@ from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
from open_webui.constants import ERROR_MESSAGES
from open_webui.internal.db import get_async_session
from open_webui.models.access_grants import AccessGrants
from open_webui.models.files import FileMetadataResponse, FileModel, Files
from open_webui.models.files import FileMetadataResponse, FileModel, FileModelResponse, Files
from open_webui.models.groups import Groups
from open_webui.models.knowledge import (
KnowledgeDirectoryForm,
@ -216,6 +217,7 @@ async def search_knowledge_bases(
@router.get('/search/files', response_model=KnowledgeFileListResponse)
async def search_knowledge_files(
query: str | None = None,
include_content: bool = Query(False, description='Include file content in search (expensive).'),
page: int | None = 1,
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
@ -227,6 +229,8 @@ async def search_knowledge_files(
filter = {}
if query:
filter['query'] = query
if include_content:
filter['include_content'] = True
groups = await Groups.get_groups_by_member_id(user.id, db=db)
if groups:
@ -548,6 +552,71 @@ async def update_knowledge_access_by_id(
)
############################
# GetPendingKnowledgeFiles
############################
@router.get('/{id}/files/pending')
async def get_pending_knowledge_files(
id: str,
stream: bool = Query(False),
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""Return files that are being processed for this knowledge base but not yet linked.
After a file is uploaded with ``knowledge_id`` in its metadata, the backend
processes it in a background task before linking it to the ``knowledge_file``
join table. During this window the file is invisible to the normal file
list endpoint. This endpoint exposes those in-flight files so the frontend
can show them with a processing indicator even after a page reload.
When ``stream=true``, returns an SSE stream that polls every 3 seconds
and emits the current pending file list. Closes when no files remain.
"""
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
if not (
user.role == 'admin'
or knowledge.user_id == user.id
or await AccessGrants.has_access(
user_id=user.id,
resource_type='knowledge',
resource_id=knowledge.id,
permission='read',
db=db,
)
):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
if not stream:
return await Files.get_pending_files_for_knowledge(id, db=db)
async def event_stream(knowledge_id: str):
MAX_POLL_DURATION = 3600 # 1 hour max
for _ in range(MAX_POLL_DURATION // 3):
pending = await Files.get_pending_files_for_knowledge(knowledge_id)
data = [f.model_dump() for f in pending]
yield f'data: {json.dumps(data)}\n\n'
if len(pending) == 0:
break
await asyncio.sleep(3)
return StreamingResponse(
event_stream(id),
media_type='text/event-stream',
)
############################
# GetKnowledgeFilesById
############################
@ -557,11 +626,13 @@ async def update_knowledge_access_by_id(
async def get_knowledge_files_by_id(
id: str,
query: str | None = None,
include_content: bool = Query(False, description='Include file content in search (expensive).'),
view_option: str | None = None,
order_by: str | None = None,
direction: str | None = None,
directory_id: str | None = Query(None, description='Filter by directory ID. Pass empty string for root.'),
page: int | None = 1,
limit: int | None = Query(None, description='Page size (admin only). Defaults to 30.'),
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
@ -590,12 +661,18 @@ async def get_knowledge_files_by_id(
page = max(page, 1)
limit = 30
# Allow admins to configure page size; non-admins always get the default
if user.role == 'admin' and limit is not None:
limit = max(1, limit)
else:
limit = PAGE_ITEM_COUNT
skip = (page - 1) * limit
filter = {}
if query:
filter['query'] = query
if include_content:
filter['include_content'] = True
if view_option:
filter['view_option'] = view_option
if order_by:
@ -846,10 +923,7 @@ async def remove_file_from_knowledge_by_id(
log.debug(e)
pass
# Only the file owner or an admin may permanently delete the underlying
# file. Collaborators with KB write access can unlink a file from the
# knowledge base but must not be able to destroy files they do not own,
# as the same file may be referenced by other KBs and chats.
# Anyone with write permission or higher can delete files
if delete_file and (file.user_id == user.id or user.role == 'admin'):
try:
# Remove the file's collection from vector database
@ -949,7 +1023,10 @@ async def delete_knowledge_by_id(
@router.post('/{id}/reset', response_model=KnowledgeResponse | None)
async def reset_knowledge_by_id(
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
id: str,
include_directories: bool = Query(True),
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
@ -980,10 +1057,189 @@ async def reset_knowledge_by_id(
log.debug(e)
pass
knowledge = await Knowledges.reset_knowledge_by_id(id=id, db=db)
knowledge = await Knowledges.reset_knowledge_by_id(id=id, include_directories=include_directories, db=db)
return knowledge
############################
# SyncKnowledgeDiff
############################
class FileManifestEntry(BaseModel):
filename: str # basename: "readme.md"
path: str # relative dir: "docs/api" or "" for root
checksum: str # SHA-256 of raw bytes
size: int
class SyncDiffForm(BaseModel):
manifest: list[FileManifestEntry]
class SyncDiffResponse(BaseModel):
added: list[dict] # [{filename, path}] — new files
modified: list[dict] # [{filename, path, stale_file_id}] — changed files
deleted: list[dict] # [{file_id, filename}] — files to remove
mkdir: list[str] # directory paths to create
rmdir: list[str] # directory IDs to remove
unmodified_count: int
directory_map: dict[str, str] # existing path → directory ID
@router.post('/{id}/sync/diff', response_model=SyncDiffResponse)
async def sync_knowledge_diff(
id: str,
form_data: SyncDiffForm,
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""
Compare a local file manifest against the knowledge base to determine
which files need uploading, removing, and which directories to create/remove.
"""
await _verify_knowledge_write_access(id, user, db)
# ── Index existing state ──
knowledge_files = await Knowledges.get_files_with_directory_ids(id, db=db)
existing_directories = await Knowledges.get_all_directories(id, db=db)
# Build directory path lookups
directory_path_by_id: dict[str, str] = {}
directory_id_by_path: dict[str, str] = {}
for directory in existing_directories:
segments = [directory.name]
parent_id = directory.parent_id
while parent_id:
parent = next((d for d in existing_directories if d.id == parent_id), None)
if not parent:
break
segments.insert(0, parent.name)
parent_id = parent.parent_id
full_path = '/'.join(segments)
directory_path_by_id[directory.id] = full_path
directory_id_by_path[full_path] = directory.id
# Index existing files by (path, filename) → {file_id, checksum}
indexed_files: dict[tuple[str, str], dict] = {}
for file_model, directory_id in knowledge_files:
file_path = directory_path_by_id.get(directory_id, '') if directory_id else ''
stored_checksum = (file_model.meta or {}).get('file_hash')
indexed_files[(file_path, file_model.filename)] = {
'file_id': file_model.id,
'checksum': stored_checksum,
}
# ── Diff files ──
added: list[dict] = []
modified: list[dict] = []
deleted: list[dict] = []
unmodified_count = 0
manifest_keys: set[tuple[str, str]] = set()
for entry in form_data.manifest:
key = (entry.path, entry.filename)
manifest_keys.add(key)
if key not in indexed_files:
added.append({'filename': entry.filename, 'path': entry.path})
elif indexed_files[key]['checksum'] != entry.checksum:
modified.append(
{
'filename': entry.filename,
'path': entry.path,
'stale_file_id': indexed_files[key]['file_id'],
}
)
else:
unmodified_count += 1
for key, file_info in indexed_files.items():
if key not in manifest_keys:
deleted.append({'file_id': file_info['file_id'], 'filename': key[1]})
# ── Diff directories ──
required_directory_paths: set[str] = set()
for entry in form_data.manifest:
if entry.path:
segments = entry.path.split('/')
for depth in range(len(segments)):
required_directory_paths.add('/'.join(segments[: depth + 1]))
mkdir = sorted([p for p in required_directory_paths if p not in directory_id_by_path], key=lambda p: p.count('/'))
orphaned_directory_paths = set(directory_id_by_path) - required_directory_paths
rmdir = [directory_id_by_path[p] for p in orphaned_directory_paths]
return SyncDiffResponse(
added=added,
modified=modified,
deleted=deleted,
mkdir=mkdir,
rmdir=rmdir,
unmodified_count=unmodified_count,
directory_map=directory_id_by_path,
)
############################
# SyncKnowledgeCleanup
############################
class SyncCleanupForm(BaseModel):
file_ids: list[str] # file IDs to delete
dir_ids: list[str] = [] # directory IDs to rmdir
@router.post('/{id}/sync/cleanup')
async def sync_knowledge_cleanup(
id: str,
form_data: SyncCleanupForm,
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""
Remove stale files and orphaned directories from a knowledge base
after an incremental sync.
"""
await _verify_knowledge_write_access(id, user, db)
# ── Remove deleted files ──
for file_id in form_data.file_ids:
file = await Files.get_file_by_id(file_id, db=db)
if not file:
continue
await Knowledges.remove_file_from_knowledge_by_id(id, file_id, db=db)
try:
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=id, filter={'file_id': file_id})
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=id, filter={'hash': file.hash})
except Exception:
pass
try:
collection_name = f'file-{file_id}'
if await ASYNC_VECTOR_DB_CLIENT.has_collection(collection_name):
await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name)
except Exception:
pass
if file.user_id == user.id or user.role == 'admin':
await Files.delete_file_by_id(file_id, db=db)
try:
await asyncio.to_thread(Storage.delete_file, file.path)
except Exception:
pass
# ── Remove orphaned directories (children before parents) ──
for dir_id in reversed(form_data.dir_ids):
await Knowledges.delete_directory(dir_id, move_files_to_parent=False, db=db)
return {'status': True}
############################
# AddFilesToKnowledge
############################
@ -1046,6 +1302,23 @@ async def add_files_to_knowledge_batch(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
# Filter out files already linked to this knowledge base to prevent
# duplicate embeddings in the vector DB (issue #10679).
new_entries = []
for form in form_data:
if not await Knowledges.has_file(knowledge_id=id, file_id=form.file_id, db=db):
new_entries.append(form)
if not new_entries:
return KnowledgeFilesResponse(
**knowledge.model_dump(),
files=await Knowledges.get_file_metadatas_by_id(knowledge.id, db=db),
)
# Narrow the file list to only new files for processing
new_file_ids = {form.file_id for form in new_entries}
files = [f for f in files if f.id in new_file_ids]
# Process files
try:
result = await process_files_batch(
@ -1060,8 +1333,15 @@ async def add_files_to_knowledge_batch(
# Only add files that were successfully processed
successful_file_ids = [r.file_id for r in result.results if r.status == 'completed']
dir_map = {form.file_id: form.directory_id for form in new_entries}
for file_id in successful_file_ids:
await Knowledges.add_file_to_knowledge_by_id(knowledge_id=id, file_id=file_id, user_id=user.id, db=db)
await Knowledges.add_file_to_knowledge_by_id(
knowledge_id=id,
file_id=file_id,
user_id=user.id,
directory_id=dir_map.get(file_id),
db=db,
)
# If there were any errors, include them in the response
if result.errors:
@ -1152,9 +1432,7 @@ class KnowledgeFileMoveForm(BaseModel):
directory_id: Optional[str] = None
async def _verify_knowledge_write_access(
id: str, user, db: AsyncSession
):
async def _verify_knowledge_write_access(id: str, user, db: AsyncSession):
"""Verify the user has write access to the knowledge base. Returns the knowledge model."""
knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db)
if not knowledge:
@ -1304,4 +1582,3 @@ async def move_file_in_knowledge(
detail='Failed to move file.',
)
return {'status': True}

View file

@ -9,6 +9,7 @@ from open_webui.constants import ERROR_MESSAGES
from open_webui.internal.db import get_async_session
from open_webui.models.memories import Memories, MemoryModel
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
from open_webui.utils.access_control import has_permission
from open_webui.utils.auth import get_verified_user
from pydantic import BaseModel
@ -66,10 +67,12 @@ async def add_memory(
form_data: AddMemoryForm,
user=Depends(get_verified_user),
):
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (insert_new_memory) manage their own short-lived sessions.
# This prevents holding a connection during EMBEDDING_FUNCTION()
# which makes external embedding API calls (1-5+ seconds).
"""Persist a new memory and embed it into the user's vector collection.
Does NOT use ``Depends(get_async_session)`` — database operations manage their
own short-lived sessions so a connection is not held during the external
embedding API call (``EMBEDDING_FUNCTION``), which can take 1-5+ seconds.
"""
if not request.app.state.config.ENABLE_MEMORIES:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -137,7 +140,7 @@ async def query_memory(
if not memories:
raise HTTPException(status_code=404, detail='No memories found for user')
vector = await request.app.state.EMBEDDING_FUNCTION(form_data.content, user=user)
vector = await request.app.state.EMBEDDING_FUNCTION(form_data.content, RAG_EMBEDDING_QUERY_PREFIX, user=user)
results = await ASYNC_VECTOR_DB_CLIENT.search(
collection_name=f'user-memory-{user.id}',

View file

@ -20,7 +20,7 @@ from fastapi import (
from fastapi.responses import RedirectResponse, StreamingResponse
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
from open_webui.constants import ERROR_MESSAGES
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING, PROFILE_IMAGE_ALLOWED_MIME_TYPES
from open_webui.internal.db import get_async_session
from open_webui.models.access_grants import AccessGrants
from open_webui.models.groups import Groups
@ -36,6 +36,7 @@ from open_webui.models.models import (
Models,
)
from open_webui.utils.access_control import filter_allowed_access_grants, has_permission
from open_webui.utils.access_control.files import has_access_to_file
from open_webui.utils.auth import get_admin_user, get_verified_user
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import AsyncSession
@ -77,6 +78,32 @@ def is_valid_model_id(model_id: str) -> bool:
return model_id and len(model_id) <= 256
async def _verify_knowledge_file_access(
knowledge_items: list | None,
user,
db: AsyncSession,
) -> None:
"""Raise 403 if any knowledge item references a file the caller cannot read."""
if not knowledge_items or user.role == 'admin':
return
for item in knowledge_items:
if not isinstance(item, dict) or item.get('type') != 'file':
continue
file_id = item.get('id')
if not file_id:
continue
if not await has_access_to_file(file_id, 'read', user, db=db):
log.warning(
'knowledge file access denied: user %s cannot read file %s',
user.id,
file_id,
)
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
###########################
# GetModels
# Let each model here be judged by what it does and not
@ -198,6 +225,7 @@ async def create_new_model(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""Create a new workspace model entry."""
if user.role != 'admin' and not await has_permission(
user.id, 'workspace.models', request.app.state.config.USER_PERMISSIONS, db=db
):
@ -220,6 +248,12 @@ async def create_new_model(
)
else:
await _verify_knowledge_file_access(
getattr(form_data.meta, 'knowledge', None) if form_data.meta else None,
user,
db,
)
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,
@ -326,6 +360,21 @@ async def import_models(
model_id = model_data.get('id')
if model_id and is_valid_model_id(model_id):
# Defense-in-depth: skip models referencing inaccessible files
try:
await _verify_knowledge_file_access(
(model_data.get('meta') or {}).get('knowledge'),
user,
db,
)
except HTTPException:
log.warning(
'import_models: user %s skipped model %s (knowledge file access denied)',
user.id,
model_id,
)
continue
existing_model = existing_models.get(model_id)
if existing_model:
# Enforce ownership/write-access before allowing overwrite
@ -504,9 +553,19 @@ async def get_model_profile_image(
header, base64_data = profile_image_url.split(',', 1)
image_data = base64.b64decode(base64_data)
image_buffer = io.BytesIO(image_data)
media_type = header.split(';')[0].lstrip('data:')
media_type = header.split(';')[0].lstrip('data:').lower()
headers = {'Content-Disposition': 'inline'}
# only serve known-safe raster types inline; reject SVG/unknown (can run script on our origin)
if media_type not in PROFILE_IMAGE_ALLOWED_MIME_TYPES:
return RedirectResponse(
url='/static/favicon.png',
status_code=status.HTTP_302_FOUND,
)
headers = {
'Content-Disposition': 'inline',
'X-Content-Type-Options': 'nosniff',
}
if updated_at:
headers['ETag'] = f'"{updated_at}"'
@ -584,6 +643,7 @@ async def update_model_by_id(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""Update a workspace model's configuration."""
model = await Models.get_model_by_id(form_data.id, db=db)
if not model:
raise HTTPException(
@ -607,6 +667,12 @@ async def update_model_by_id(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
await _verify_knowledge_file_access(
getattr(form_data.meta, 'knowledge', None) if form_data.meta else None,
user,
db,
)
form_data.access_grants = await filter_allowed_access_grants(
request.app.state.config.USER_PERMISSIONS,
user.id,

File diff suppressed because it is too large Load diff

View file

@ -595,7 +595,7 @@ async def get_models(request: Request, url_idx: int | None = None, user=Depends(
try:
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
if api_config.get('azure', False):
if api_config.get('azure') or api_config.get('provider') == 'azure':
models = {
'data': api_config.get('model_ids', []) or [],
'object': 'list',
@ -681,15 +681,24 @@ async def verify_connection(
try:
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
if api_config.get('azure', False):
if api_config.get('azure') or api_config.get('provider') == 'azure':
# Only set api-key header if not using Azure Entra ID authentication
auth_type = api_config.get('auth_type', 'bearer')
if auth_type not in ('azure_ad', 'microsoft_entra_id'):
headers['api-key'] = key
api_version = api_config.get('api_version', '') or '2023-03-15-preview'
# Azure v1 format: base URL already ends with /openai/v1,
# use standard /models endpoint without api-version.
is_azure_v1 = bool(re.search(r'/openai/v1(?:/|$)', url))
if is_azure_v1:
verify_url = f'{url.rstrip("/")}/models'
else:
api_version = api_config.get('api_version', '') or '2023-03-15-preview'
verify_url = f'{url}/openai/models?api-version={api_version}'
async with session.get(
url=f'{url}/openai/models?api-version={api_version}',
url=verify_url,
headers=headers,
cookies=cookies,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
@ -1045,19 +1054,20 @@ async def generate_chat_completion(
request: Request,
form_data: dict,
user=Depends(get_verified_user),
bypass_system_prompt: bool = False,
):
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
# Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions.
# This prevents holding a connection during the entire LLM call (30-60+ seconds),
# which would exhaust the connection pool under concurrent load.
# bypass_filter is read from request.state to prevent external clients from
# setting it via query parameter (CVE fix). Only internal server-side callers
# (e.g. utils/chat.py) should set request.state.bypass_filter = True.
# bypass_filter and bypass_system_prompt are read from request.state to prevent
# external clients from setting them via query parameter. Only internal
# server-side callers (e.g. utils/chat.py) should set
# request.state.bypass_filter / request.state.bypass_system_prompt = True.
bypass_filter = getattr(request.state, 'bypass_filter', False)
if BYPASS_MODEL_ACCESS_CONTROL:
bypass_filter = True
bypass_system_prompt = getattr(request.state, 'bypass_system_prompt', False)
idx = 0
@ -1151,7 +1161,7 @@ async def generate_chat_completion(
is_responses = api_config.get('api_type') == 'responses'
if api_config.get('azure', False):
if api_config.get('azure') or api_config.get('provider') == 'azure':
# Only set api-key header if not using Azure Entra ID authentication
auth_type = api_config.get('auth_type', 'bearer')
if auth_type not in ('azure_ad', 'microsoft_entra_id'):
@ -1304,11 +1314,32 @@ async def embeddings(request: Request, form_data: dict, user):
streaming = False
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
if api_config.get('azure') or api_config.get('provider') == 'azure':
# Only set api-key header if not using Azure Entra ID authentication
auth_type = api_config.get('auth_type', 'bearer')
if auth_type not in ('azure_ad', 'microsoft_entra_id'):
headers['api-key'] = key
# Azure v1 format: base URL already ends with /openai/v1,
# model stays in the payload, no deployment URL rewriting.
is_azure_v1 = bool(re.search(r'/openai/v1(?:/|$)', url))
if is_azure_v1:
embeddings_url = f'{url.rstrip("/")}/embeddings'
else:
api_version = api_config.get('api_version', '2023-03-15-preview')
model = _sanitize_model_for_url(form_data.get('model', ''))
embeddings_url = f'{url}/openai/deployments/{model}/embeddings?api-version={api_version}'
headers['api-version'] = api_version
else:
embeddings_url = f'{url}/embeddings'
try:
session = await get_session()
r = await session.request(
method='POST',
url=f'{url}/embeddings',
url=embeddings_url,
data=body,
headers=headers,
cookies=cookies,
@ -1408,7 +1439,7 @@ async def responses(
try:
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
if api_config.get('azure', False):
if api_config.get('azure') or api_config.get('provider') == 'azure':
auth_type = api_config.get('auth_type', 'bearer')
if auth_type not in ('azure_ad', 'microsoft_entra_id'):
headers['api-key'] = key
@ -1519,7 +1550,7 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
try:
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
if api_config.get('azure', False):
if api_config.get('azure') or api_config.get('provider') == 'azure':
# Only set api-key header if not using Azure Entra ID authentication
auth_type = api_config.get('auth_type', 'bearer')
if auth_type not in ('azure_ad', 'microsoft_entra_id'):

View file

@ -188,50 +188,6 @@ async def create_new_prompt(
)
############################
# GetPromptByCommand
############################
@router.get('/command/{command}', response_model=PromptAccessResponse | None)
async def get_prompt_by_command(
command: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
):
prompt = await Prompts.get_prompt_by_command(command, db=db)
if prompt:
if (
user.role == 'admin'
or prompt.user_id == user.id
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
permission='read',
db=db,
)
):
return PromptAccessResponse(
**prompt.model_dump(),
write_access=(
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
or user.id == prompt.user_id
or await AccessGrants.has_access(
user_id=user.id,
resource_type='prompt',
resource_id=prompt.id,
permission='write',
db=db,
)
),
)
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.NOT_FOUND,
)
############################
# GetPromptById
############################
@ -289,6 +245,7 @@ async def update_prompt_by_id(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""Update a prompt's content, creating a new history entry if changed."""
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
if not prompt:
@ -699,7 +656,7 @@ async def delete_prompt_history_entry(
detail='Cannot delete the active production version',
)
success = await PromptHistories.delete_history_entry(history_id, db=db)
success = await PromptHistories.delete_history_entry(history_id, prompt.id, db=db)
if not success:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@ -743,7 +700,7 @@ async def get_prompt_diff(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
diff = await PromptHistories.compute_diff(from_id, to_id, db=db)
diff = await PromptHistories.compute_diff(from_id, to_id, prompt.id, db=db)
if not diff:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,

View file

@ -10,7 +10,7 @@ import shutil
import uuid
from datetime import datetime
from pathlib import Path
from typing import Iterator, Optional, Sequence, Union
from typing import Callable, Iterator, Optional, Sequence, Union
import tiktoken
from fastapi import (
@ -107,6 +107,7 @@ from open_webui.retrieval.web.utils import get_web_loader
from open_webui.retrieval.web.yacy import search_yacy
from open_webui.retrieval.web.yandex import search_yandex
from open_webui.retrieval.web.ydc import search_youcom
from open_webui.retrieval.web.linkup import search_linkup
from open_webui.storage.provider import Storage
from open_webui.utils.access_control import has_permission
from open_webui.utils.access_control.files import has_access_to_file
@ -468,6 +469,7 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)):
'MINERU_API_KEY': request.app.state.config.MINERU_API_KEY,
'MINERU_API_TIMEOUT': request.app.state.config.MINERU_API_TIMEOUT,
'MINERU_PARAMS': request.app.state.config.MINERU_PARAMS,
'MINERU_FILE_EXTENSIONS': request.app.state.config.MINERU_FILE_EXTENSIONS,
# Reranking settings
'RAG_RERANKING_MODEL': request.app.state.config.RAG_RERANKING_MODEL,
'RAG_RERANKING_ENGINE': request.app.state.config.RAG_RERANKING_ENGINE,
@ -556,6 +558,8 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)):
'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,
'LINKUP_API_KEY': request.app.state.config.LINKUP_API_KEY,
'LINKUP_SEARCH_PARAMS': request.app.state.config.LINKUP_SEARCH_PARAMS,
},
}
@ -625,6 +629,8 @@ class WebConfig(BaseModel):
YANDEX_WEB_SEARCH_API_KEY: str | None = None
YANDEX_WEB_SEARCH_CONFIG: str | None = None
YOUCOM_API_KEY: str | None = None
LINKUP_API_KEY: str | None = None
LINKUP_SEARCH_PARAMS: dict | None = None
class ConfigForm(BaseModel):
@ -681,6 +687,7 @@ class ConfigForm(BaseModel):
MINERU_API_KEY: str | None = None
MINERU_API_TIMEOUT: str | None = None
MINERU_PARAMS: dict | None = None
MINERU_FILE_EXTENSIONS: list[str] | None = None
# Reranking settings
RAG_RERANKING_MODEL: str | None = None
@ -915,6 +922,11 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend
request.app.state.config.MINERU_PARAMS = (
form_data.MINERU_PARAMS if form_data.MINERU_PARAMS is not None else request.app.state.config.MINERU_PARAMS
)
request.app.state.config.MINERU_FILE_EXTENSIONS = (
form_data.MINERU_FILE_EXTENSIONS
if form_data.MINERU_FILE_EXTENSIONS is not None
else request.app.state.config.MINERU_FILE_EXTENSIONS
)
# Reranking settings
if request.app.state.config.RAG_RERANKING_ENGINE == '':
@ -1125,6 +1137,8 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend
request.app.state.config.YANDEX_WEB_SEARCH_API_KEY = form_data.web.YANDEX_WEB_SEARCH_API_KEY
request.app.state.config.YANDEX_WEB_SEARCH_CONFIG = form_data.web.YANDEX_WEB_SEARCH_CONFIG
request.app.state.config.YOUCOM_API_KEY = form_data.web.YOUCOM_API_KEY
request.app.state.config.LINKUP_API_KEY = form_data.web.LINKUP_API_KEY
request.app.state.config.LINKUP_SEARCH_PARAMS = form_data.web.LINKUP_SEARCH_PARAMS
return {
'status': True,
@ -1260,6 +1274,8 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend
'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,
'LINKUP_API_KEY': request.app.state.config.LINKUP_API_KEY,
'LINKUP_SEARCH_PARAMS': request.app.state.config.LINKUP_SEARCH_PARAMS,
},
}
@ -1294,20 +1310,43 @@ def merge_docs_to_target_size(
Attempts to grow small chunks up to a desired minimum size,
without exceeding the maximum size or crossing source/file
boundaries.
"""
min_chunk_size_target = request.app.state.config.CHUNK_MIN_SIZE_TARGET
max_chunk_size = request.app.state.config.CHUNK_SIZE
if min_chunk_size_target <= 0:
Uses forward merging first (absorb the next chunk), then
backward merging (append into the previous emitted chunk)
for undersized chunks that can't grow forward.
"""
min_size = request.app.state.config.CHUNK_MIN_SIZE_TARGET
max_size = request.app.state.config.CHUNK_SIZE
if min_size <= 0:
return chunks
measure_chunk_size = len
measure: Callable[[str], int] = len
if request.app.state.config.TEXT_SPLITTER == 'token':
encoding = tiktoken.get_encoding(str(request.app.state.config.TIKTOKEN_ENCODING_NAME))
measure_chunk_size = lambda text: len(encoding.encode(text))
measure = lambda text: len(encoding.encode(text))
processed_chunks: list[Document] = []
def _merge_backward(result: list[Document], content: str, chunk: Document) -> bool:
"""Try to append content into the last emitted chunk. Returns True on success."""
if not result:
return False
prev = result[-1]
if not can_merge_chunks(prev, chunk):
return False
merged = f'{prev.page_content}\n\n{content}'
if measure(merged) > max_size:
return False
result[-1] = Document(page_content=merged, metadata={**prev.metadata})
return True
def _emit(result: list[Document], content: str, chunk: Document) -> None:
"""Emit a chunk, trying backward merge first if it's undersized."""
is_undersized = measure(content) < min_size
if is_undersized and _merge_backward(result, content, chunk):
return
result.append(Document(page_content=content, metadata={**chunk.metadata}))
result: list[Document] = []
current_chunk: Document | None = None
current_content: str = ''
@ -1315,37 +1354,27 @@ def merge_docs_to_target_size(
if current_chunk is None:
current_chunk = next_chunk
current_content = next_chunk.page_content
continue # First chunk initialization
continue
proposed_content = f'{current_content}\n\n{next_chunk.page_content}'
can_merge = (
# Forward merge: absorb next chunk into current if undersized and fits
merged_content = f'{current_content}\n\n{next_chunk.page_content}'
can_merge_forward = (
can_merge_chunks(current_chunk, next_chunk)
and measure_chunk_size(current_content) < min_chunk_size_target
and measure_chunk_size(proposed_content) <= max_chunk_size
and measure(current_content) < min_size
and measure(merged_content) <= max_size
)
if can_merge:
current_content = proposed_content
if can_merge_forward:
current_content = merged_content
else:
processed_chunks.append(
Document(
page_content=current_content,
metadata={**current_chunk.metadata},
)
)
_emit(result, current_content, current_chunk)
current_chunk = next_chunk
current_content = next_chunk.page_content
if current_chunk is not None:
processed_chunks.append(
Document(
page_content=current_content,
metadata={**current_chunk.metadata},
)
)
_emit(result, current_content, current_chunk)
return processed_chunks
return result
def save_docs_to_vector_db(
@ -1876,32 +1905,18 @@ async def process_web(
)
def search_web(request: Request, engine: str, query: str, user=None) -> list[SearchResult]:
"""Search the web using a search engine and return the results as a list of SearchResult objects.
Will look for a search engine API key in environment variables in the following order:
- SEARXNG_QUERY_URL
- YACY_QUERY_URL + YACY_USERNAME + YACY_PASSWORD
- GOOGLE_PSE_API_KEY + GOOGLE_PSE_ENGINE_ID
- BRAVE_SEARCH_API_KEY
- KAGI_SEARCH_API_KEY
- MOJEEK_SEARCH_API_KEY
- BOCHA_SEARCH_API_KEY
- SERPSTACK_API_KEY
- SERPER_API_KEY
- SERPLY_API_KEY
- TAVILY_API_KEY
- EXA_API_KEY
- PERPLEXITY_API_KEY
- SOUGOU_API_SID + SOUGOU_API_SK
- SEARCHAPI_API_KEY + SEARCHAPI_ENGINE (by default `google`)
- SERPAPI_API_KEY + SERPAPI_ENGINE (by default `google`)
Args:
query (str): The query to search for
async def search_web(request: Request, engine: str, query: str, user=None) -> list[SearchResult]:
"""Dispatch a web search query to the configured engine and return results.
Providers that have been migrated to async (aiohttp) are awaited natively.
Legacy sync providers are offloaded via ``asyncio.to_thread`` to avoid
blocking the event loop.
"""
# TODO: add playwright to search the web
if engine == 'ollama_cloud':
return search_ollama_cloud(
return await asyncio.to_thread(
search_ollama_cloud,
'https://ollama.com',
request.app.state.config.OLLAMA_CLOUD_WEB_SEARCH_API_KEY,
query,
@ -1910,7 +1925,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
)
elif engine == 'perplexity_search':
if request.app.state.config.PERPLEXITY_API_KEY:
return search_perplexity_search(
return await asyncio.to_thread(
search_perplexity_search,
request.app.state.config.PERPLEXITY_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -1923,7 +1939,7 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
elif engine == 'searxng':
if request.app.state.config.SEARXNG_QUERY_URL:
searxng_kwargs = {'language': request.app.state.config.SEARXNG_LANGUAGE}
return search_searxng(
return await search_searxng(
request.app.state.config.SEARXNG_QUERY_URL,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -1934,7 +1950,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No SEARXNG_QUERY_URL found in environment variables')
elif engine == 'yacy':
if request.app.state.config.YACY_QUERY_URL:
return search_yacy(
return await asyncio.to_thread(
search_yacy,
request.app.state.config.YACY_QUERY_URL,
request.app.state.config.YACY_USERNAME,
request.app.state.config.YACY_PASSWORD,
@ -1946,7 +1963,7 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No YACY_QUERY_URL found in environment variables')
elif engine == 'google_pse':
if request.app.state.config.GOOGLE_PSE_API_KEY and request.app.state.config.GOOGLE_PSE_ENGINE_ID:
return search_google_pse(
return await search_google_pse(
request.app.state.config.GOOGLE_PSE_API_KEY,
request.app.state.config.GOOGLE_PSE_ENGINE_ID,
query,
@ -1958,7 +1975,7 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No GOOGLE_PSE_API_KEY or GOOGLE_PSE_ENGINE_ID found in environment variables')
elif engine == 'brave':
if request.app.state.config.BRAVE_SEARCH_API_KEY:
return search_brave(
return await search_brave(
request.app.state.config.BRAVE_SEARCH_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -1968,7 +1985,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No BRAVE_SEARCH_API_KEY found in environment variables')
elif engine == 'brave_llm_context':
if request.app.state.config.BRAVE_SEARCH_API_KEY:
return search_brave_llm_context(
return await asyncio.to_thread(
search_brave_llm_context,
request.app.state.config.BRAVE_SEARCH_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -1979,7 +1997,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No BRAVE_SEARCH_API_KEY found in environment variables')
elif engine == 'kagi':
if request.app.state.config.KAGI_SEARCH_API_KEY:
return search_kagi(
return await asyncio.to_thread(
search_kagi,
request.app.state.config.KAGI_SEARCH_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -1989,7 +2008,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No KAGI_SEARCH_API_KEY found in environment variables')
elif engine == 'mojeek':
if request.app.state.config.MOJEEK_SEARCH_API_KEY:
return search_mojeek(
return await asyncio.to_thread(
search_mojeek,
request.app.state.config.MOJEEK_SEARCH_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -1999,7 +2019,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No MOJEEK_SEARCH_API_KEY found in environment variables')
elif engine == 'bocha':
if request.app.state.config.BOCHA_SEARCH_API_KEY:
return search_bocha(
return await asyncio.to_thread(
search_bocha,
request.app.state.config.BOCHA_SEARCH_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2009,7 +2030,7 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No BOCHA_SEARCH_API_KEY found in environment variables')
elif engine == 'serpstack':
if request.app.state.config.SERPSTACK_API_KEY:
return search_serpstack(
return await search_serpstack(
request.app.state.config.SERPSTACK_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2020,7 +2041,7 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No SERPSTACK_API_KEY found in environment variables')
elif engine == 'serper':
if request.app.state.config.SERPER_API_KEY:
return search_serper(
return await search_serper(
request.app.state.config.SERPER_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2030,7 +2051,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No SERPER_API_KEY found in environment variables')
elif engine == 'serply':
if request.app.state.config.SERPLY_API_KEY:
return search_serply(
return await asyncio.to_thread(
search_serply,
request.app.state.config.SERPLY_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2039,7 +2061,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
else:
raise Exception('No SERPLY_API_KEY found in environment variables')
elif engine == 'duckduckgo':
return search_duckduckgo(
return await asyncio.to_thread(
search_duckduckgo,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST,
@ -2048,7 +2071,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
)
elif engine == 'tavily':
if request.app.state.config.TAVILY_API_KEY:
return search_tavily(
return await asyncio.to_thread(
search_tavily,
request.app.state.config.TAVILY_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2058,7 +2082,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No TAVILY_API_KEY found in environment variables')
elif engine == 'exa':
if request.app.state.config.EXA_API_KEY:
return search_exa(
return await asyncio.to_thread(
search_exa,
request.app.state.config.EXA_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2068,7 +2093,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No EXA_API_KEY found in environment variables')
elif engine == 'searchapi':
if request.app.state.config.SEARCHAPI_API_KEY:
return search_searchapi(
return await asyncio.to_thread(
search_searchapi,
request.app.state.config.SEARCHAPI_API_KEY,
request.app.state.config.SEARCHAPI_ENGINE,
query,
@ -2079,7 +2105,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception('No SEARCHAPI_API_KEY found in environment variables')
elif engine == 'serpapi':
if request.app.state.config.SERPAPI_API_KEY:
return search_serpapi(
return await asyncio.to_thread(
search_serpapi,
request.app.state.config.SERPAPI_API_KEY,
request.app.state.config.SERPAPI_ENGINE,
query,
@ -2089,14 +2116,16 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
else:
raise Exception('No SERPAPI_API_KEY found in environment variables')
elif engine == 'jina':
return search_jina(
return await asyncio.to_thread(
search_jina,
request.app.state.config.JINA_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
request.app.state.config.JINA_API_BASE_URL,
)
elif engine == 'bing':
return search_bing(
return await asyncio.to_thread(
search_bing,
request.app.state.config.BING_SEARCH_V7_SUBSCRIPTION_KEY,
request.app.state.config.BING_SEARCH_V7_ENDPOINT,
str(DEFAULT_LOCALE),
@ -2110,7 +2139,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
and request.app.state.config.AZURE_AI_SEARCH_ENDPOINT
and request.app.state.config.AZURE_AI_SEARCH_INDEX_NAME
):
return search_azure(
return await asyncio.to_thread(
search_azure,
request.app.state.config.AZURE_AI_SEARCH_API_KEY,
request.app.state.config.AZURE_AI_SEARCH_ENDPOINT,
request.app.state.config.AZURE_AI_SEARCH_INDEX_NAME,
@ -2122,15 +2152,9 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
raise Exception(
'AZURE_AI_SEARCH_API_KEY, AZURE_AI_SEARCH_ENDPOINT, and AZURE_AI_SEARCH_INDEX_NAME are required for Azure AI Search'
)
elif engine == 'exa':
return search_exa(
request.app.state.config.EXA_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST,
)
elif engine == 'perplexity':
return search_perplexity(
return await asyncio.to_thread(
search_perplexity,
request.app.state.config.PERPLEXITY_API_KEY,
query,
request.app.state.config.WEB_SEARCH_RESULT_COUNT,
@ -2140,7 +2164,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
)
elif engine == 'sougou':
if request.app.state.config.SOUGOU_API_SID and request.app.state.config.SOUGOU_API_SK:
return search_sougou(
return await asyncio.to_thread(
search_sougou,
request.app.state.config.SOUGOU_API_SID,
request.app.state.config.SOUGOU_API_SK,
query,
@ -2150,7 +2175,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
else:
raise Exception('No SOUGOU_API_SID or SOUGOU_API_SK found in environment variables')
elif engine == 'firecrawl':
return search_firecrawl(
return await asyncio.to_thread(
search_firecrawl,
request.app.state.config.FIRECRAWL_API_BASE_URL,
request.app.state.config.FIRECRAWL_API_KEY,
query,
@ -2158,7 +2184,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST,
)
elif engine == 'external':
return search_external(
return await asyncio.to_thread(
search_external,
request,
request.app.state.config.EXTERNAL_WEB_SEARCH_URL,
request.app.state.config.EXTERNAL_WEB_SEARCH_API_KEY,
@ -2168,7 +2195,8 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
user=user,
)
elif engine == 'yandex':
return search_yandex(
return await asyncio.to_thread(
search_yandex,
request,
request.app.state.config.YANDEX_WEB_SEARCH_URL,
request.app.state.config.YANDEX_WEB_SEARCH_API_KEY,
@ -2179,12 +2207,25 @@ def search_web(request: Request, engine: str, query: str, user=None) -> list[Sea
user=user,
)
elif engine == 'youcom':
return search_youcom(
return await asyncio.to_thread(
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,
)
elif engine == 'linkup':
if request.app.state.config.LINKUP_API_KEY:
return await asyncio.to_thread(
search_linkup,
api_key=request.app.state.config.LINKUP_API_KEY,
query=query,
count=request.app.state.config.WEB_SEARCH_RESULT_COUNT,
filter_list=request.app.state.config.WEB_SEARCH_DOMAIN_FILTER_LIST,
params=request.app.state.config.LINKUP_SEARCH_PARAMS,
)
else:
raise Exception('No LINKUP_API_KEY found in environment variables')
else:
raise Exception('No search engine API key found in environment variables')
@ -2222,8 +2263,7 @@ async def process_web_search(request: Request, form_data: SearchForm, user=Depen
async def search_query_with_semaphore(query):
async with semaphore:
return await run_in_threadpool(
search_web,
return await search_web(
request,
request.app.state.config.WEB_SEARCH_ENGINE,
query,
@ -2232,10 +2272,9 @@ async def process_web_search(request: Request, form_data: SearchForm, user=Depen
search_tasks = [search_query_with_semaphore(query) for query in form_data.queries]
else:
# Unlimited parallel execution (previous behavior)
# Unlimited parallel execution
search_tasks = [
run_in_threadpool(
search_web,
search_web(
request,
request.app.state.config.WEB_SEARCH_ENGINE,
query,
@ -2257,12 +2296,8 @@ async def process_web_search(request: Request, form_data: SearchForm, user=Depen
log.debug(f'urls: {urls}')
except Exception as e:
log.exception(e)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.WEB_SEARCH_ERROR(e),
)
log.exception('Web search failed')
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.WEB_SEARCH_ERROR(e))
if len(urls) == 0:
raise HTTPException(
@ -2342,11 +2377,8 @@ async def process_web_search(request: Request, form_data: SearchForm, user=Depen
'loaded_count': len(docs),
}
except Exception as e:
log.exception(e)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(e),
)
log.exception('Web search content loading failed')
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT(e))
async def _validate_collection_access(collection_names: list[str], user, access_type: str = 'read') -> None:

View file

@ -35,7 +35,14 @@ def _sanitize_proxy_path(path: str) -> str | None:
Trailing slashes are preserved — many upstream frameworks treat
``/path`` and ``/path/`` differently.
"""
decoded = unquote(path)
# Decode until stable: a single unquote pass leaves %252e%252e as %2e%2e,
# which the upstream then re-decodes into '..', bypassing the check below.
decoded = path
for _ in range(8):
once = unquote(decoded)
if once == decoded:
break
decoded = once
had_trailing_slash = decoded.endswith('/')
normalized = posixpath.normpath(decoded)
# Remove any leading slashes that would reset the base
@ -324,11 +331,20 @@ async def ws_terminal(
except Exception:
pass
await asyncio.gather(
_client_to_upstream(),
_upstream_to_client(),
return_exceptions=True,
)
# End the proxy as soon as either direction finishes (e.g. a
# graceful upstream CLOSE) and cancel the sibling, which would
# otherwise hang on a blocked ws.receive() until the browser leaves.
tasks = [
asyncio.create_task(_client_to_upstream()),
asyncio.create_task(_upstream_to_client()),
]
_done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
for task in pending:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
except Exception as e:
log.exception('Terminal WebSocket proxy error: %s', e)
finally:

View file

@ -329,6 +329,7 @@ async def create_new_tools(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""Create a new tool from user-supplied Python source code."""
if user.role != 'admin' and not (
await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
or await has_permission(
@ -455,6 +456,7 @@ async def update_tools_by_id(
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
"""Update an existing tool's source code and metadata."""
tools = await Tools.get_tool_by_id(id, db=db)
if not tools:
raise HTTPException(

View file

@ -383,21 +383,25 @@ async def get_user_info_by_session_user(user=Depends(get_verified_user), db: Asy
@router.post('/user/info/update', response_model=dict | None)
async def update_user_info_by_session_user(
form_data: dict, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
async def update_user_info_by_session_user( # PATCH-style merge
form_data: dict,
user=Depends(get_verified_user),
db: AsyncSession = Depends(get_async_session),
):
# Merges against the auth-time snapshot of user.info. The previous pre-merge
# refetch only narrowed (did not eliminate) the lost-update window on concurrent
# same-user writes; real safety needs row locking or a version column.
existing_info = user.info or {}
updated = await Users.update_user_by_id(user.id, {'info': {**existing_info, **form_data}}, db=db)
if updated:
return updated.info
else:
"""Merge caller-supplied fields into the current user's info dict.
Uses the auth-time snapshot of ``user.info`` as the merge base. This does
NOT eliminate lost-update races on concurrent same-user writes; real safety
would need row locking or an optimistic-concurrency version column.
"""
merged_info = {**(user.info or {}), **form_data}
updated = await Users.update_user_by_id(user.id, {'info': merged_info}, db=db)
if not updated:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.USER_NOT_FOUND,
)
return updated.info
############################
@ -539,7 +543,7 @@ async def get_user_active_status_by_id(
async def update_user_by_id(
user_id: str,
form_data: UserUpdateForm,
session_user=Depends(get_admin_user),
session_user: UserModel = Depends(get_admin_user),
db: AsyncSession = Depends(get_async_session),
):
# Prevent modification of the primary admin user by other admins
@ -743,27 +747,15 @@ async def get_user_preview(
'user': {'id': target_user.id, 'name': target_user.name},
'groups': [{'id': g.id, 'name': g.name} for g in user_groups],
'models': {
'items': [
{'id': m.id, 'name': m.name}
for m in active_models
if m.id in accessible_model_ids
],
'items': [{'id': m.id, 'name': m.name} for m in active_models if m.id in accessible_model_ids],
'total': len(active_models),
},
'knowledge': {
'items': [
{'id': k.id, 'name': k.name}
for k in all_knowledge
if k.id in accessible_knowledge_ids
],
'items': [{'id': k.id, 'name': k.name} for k in all_knowledge if k.id in accessible_knowledge_ids],
'total': len(all_knowledge),
},
'tools': {
'items': [
{'id': t.id, 'name': t.name}
for t in all_tools
if t.id in accessible_tool_ids
],
'items': [{'id': t.id, 'name': t.name} for t in all_tools if t.id in accessible_tool_ids],
'total': len(all_tools),
},
}

View file

@ -3,7 +3,6 @@ from __future__ import annotations
import logging
import black
import markdown
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from open_webui.config import DATA_DIR, ENABLE_ADMIN_EXPORT
from open_webui.constants import ERROR_MESSAGES
@ -73,15 +72,6 @@ async def execute_code(request: Request, form_data: CodeForm, user=Depends(get_v
)
class MarkdownForm(BaseModel):
md: str
@router.post('/markdown')
async def get_html_from_markdown(form_data: MarkdownForm, user=Depends(get_verified_user)):
return {'html': markdown.markdown(form_data.md)}
class ChatForm(BaseModel):
title: str
messages: list[dict]
@ -104,20 +94,18 @@ async def download_chat_as_pdf(form_data: ChatTitleMessagesForm, user=Depends(ge
@router.get('/db/download')
async def download_db(user=Depends(get_admin_user)):
"""Download the raw SQLite database file (admin-only, SQLite deployments only)."""
if not ENABLE_ADMIN_EXPORT:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
# Lazy import avoids circular dependency at module load time
from open_webui.internal.db import engine
if engine.name != 'sqlite':
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DB_NOT_SQLITE,
)
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DB_NOT_SQLITE)
return FileResponse(
engine.url.database,
str(engine.url.database),
media_type='application/octet-stream',
filename='webui.db',
)

View file

@ -40,8 +40,8 @@ from open_webui.tasks import create_task, stop_item_tasks
from open_webui.utils.access_control import has_permission
from open_webui.utils.auth import decode_token
from open_webui.utils.redis import (
build_sentinel_url,
get_redis_connection,
get_sentinel_url_from_env,
get_sentinels_from_env,
)
from redis import asyncio as aioredis
@ -58,20 +58,20 @@ REDIS = None
SOCKETIO_CORS_ORIGINS = '*' if CORS_ALLOW_ORIGIN == ['*'] else CORS_ALLOW_ORIGIN
if WEBSOCKET_MANAGER == 'redis':
if WEBSOCKET_SENTINEL_HOSTS:
mgr = socketio.AsyncRedisManager(
get_sentinel_url_from_env(WEBSOCKET_REDIS_URL, WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT),
redis_options=WEBSOCKET_REDIS_OPTIONS,
)
else:
mgr = socketio.AsyncRedisManager(WEBSOCKET_REDIS_URL, redis_options=WEBSOCKET_REDIS_OPTIONS)
sentinel_hosts = WEBSOCKET_SENTINEL_HOSTS or ''
ws_redis_url = (
build_sentinel_url(WEBSOCKET_REDIS_URL, sentinel_hosts, WEBSOCKET_SENTINEL_PORT)
if sentinel_hosts
else WEBSOCKET_REDIS_URL
)
redis_manager = socketio.AsyncRedisManager(ws_redis_url, redis_options=WEBSOCKET_REDIS_OPTIONS)
sio = socketio.AsyncServer(
cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
async_mode='asgi',
transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
always_connect=True,
client_manager=mgr,
client_manager=redis_manager,
logger=WEBSOCKET_SERVER_LOGGING,
ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
@ -99,32 +99,31 @@ SESSION_POOL_TIMEOUT = 120 # seconds without heartbeat before session is reaped
if WEBSOCKET_MANAGER == 'redis':
log.debug('Using Redis to manage websockets.')
ws_sentinels = get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT)
REDIS = get_redis_connection(
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT),
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
async_mode=True,
)
redis_sentinels = get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT)
MODELS = RedisDict(
f'{REDIS_KEY_PREFIX}:models',
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=redis_sentinels,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
SESSION_POOL = RedisDict(
f'{REDIS_KEY_PREFIX}:session_pool',
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=redis_sentinels,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
USAGE_POOL = RedisDict(
f'{REDIS_KEY_PREFIX}:usage_pool',
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=redis_sentinels,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
@ -132,7 +131,7 @@ if WEBSOCKET_MANAGER == 'redis':
redis_url=WEBSOCKET_REDIS_URL,
lock_name=f'{REDIS_KEY_PREFIX}:usage_cleanup_lock',
timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
redis_sentinels=redis_sentinels,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
aquire_func = clean_up_lock.aquire_lock
@ -143,7 +142,7 @@ if WEBSOCKET_MANAGER == 'redis':
redis_url=WEBSOCKET_REDIS_URL,
lock_name=f'{REDIS_KEY_PREFIX}:session_cleanup_lock',
timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
redis_sentinels=redis_sentinels,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
session_aquire_func = session_cleanup_lock.aquire_lock
@ -806,7 +805,7 @@ async def yjs_awareness_update(sid, data):
@sio.event
async def disconnect(sid):
async def disconnect(sid, reason=None):
if sid in SESSION_POOL:
user = SESSION_POOL[sid]
del SESSION_POOL[sid]

View file

@ -1,13 +1,21 @@
"""Redis-backed distributed data structures for WebSocket state management."""
from __future__ import annotations
import hashlib
import json
import uuid
from typing import List, Optional, Tuple
import pycrdt as Y
from open_webui.env import REDIS_KEY_PREFIX
from open_webui.utils.redis import get_redis_connection
from open_webui.env import REDIS_KEY_PREFIX
YDOC_KEY_PREFIX = f'{REDIS_KEY_PREFIX}:ydoc:documents'
class RedisLock:
"""Distributed lock backed by a Redis SET with NX/EX semantics."""
def __init__(
self,
redis_url,
@ -45,6 +53,10 @@ class RedisLock:
class RedisDict:
def __init__(self, name, redis_url, redis_sentinels=[], redis_cluster=False):
self.name = name
# Per-process cache of the last payload fingerprint written by set().
# Used to skip redundant HSET round-trips when the model list hasn't
# changed — the dominant Redis write source on busy multi-pod setups.
self._last_signature: str | None = None
self.redis = get_redis_connection(
redis_url,
redis_sentinels,
@ -85,6 +97,18 @@ class RedisDict:
def set(self, mapping: dict):
if not mapping:
self.redis.delete(self.name)
self._last_signature = None
return
# Serialize values once — reused for both the fingerprint and the write.
serialized = {k: json.dumps(v) for k, v in mapping.items()}
# Skip the write when the prepared mapping is identical to the last one
# this process wrote. The check is per-instance (not distributed), but
# still eliminates the majority of redundant writes because each pod
# typically produces the same model list on consecutive refreshes.
signature = hashlib.sha256(json.dumps(serialized, sort_keys=True).encode()).hexdigest()
if signature == self._last_signature:
return
# Fetch existing keys before writing so we know which ones to remove.
@ -96,10 +120,12 @@ class RedisDict:
# HSET first (add/update all new values), then HDEL (remove stale keys).
# We never DELETE the whole hash — this eliminates the race window
# where concurrent readers would see an empty models dict.
self.redis.hset(self.name, mapping={k: json.dumps(v) for k, v in mapping.items()})
self.redis.hset(self.name, mapping=serialized)
if keys_to_remove:
self.redis.hdel(self.name, *keys_to_remove)
self._last_signature = signature
def get(self, key, default=None):
try:
return self[key]
@ -108,6 +134,7 @@ class RedisDict:
def clear(self):
self.redis.delete(self.name)
self._last_signature = None
def update(self, other=None, **kwargs):
if other is not None:
@ -128,7 +155,7 @@ class YdocManager:
def __init__(
self,
redis=None,
redis_key_prefix: str = f'{REDIS_KEY_PREFIX}:ydoc:documents',
redis_key_prefix: str = YDOC_KEY_PREFIX,
):
self._updates = {}
self._users = {}
@ -177,7 +204,7 @@ class YdocManager:
ydoc.apply_update(bytes(update))
self._updates[document_id] = [ydoc.get_update()] + updates[mid:]
async def get_updates(self, document_id: str) -> List[bytes]:
async def get_updates(self, document_id: str) -> list[bytes]:
document_id = document_id.replace(':', '_')
if self._redis:
@ -196,7 +223,7 @@ class YdocManager:
else:
return document_id in self._updates
async def get_users(self, document_id: str) -> List[str]:
async def get_users(self, document_id: str) -> list[str]:
document_id = document_id.replace(':', '_')
if self._redis:
@ -212,6 +239,11 @@ class YdocManager:
if self._redis:
redis_key = f'{self._redis_key_prefix}:{document_id}:users'
await self._redis.sadd(redis_key, user_id)
# Maintain a per-session reverse index so disconnect cleanup
# can look up only the documents this session joined, instead
# of issuing a cluster-wide SCAN over the entire keyspace.
session_key = f'{self._redis_key_prefix}:session:{user_id}:documents'
await self._redis.sadd(session_key, document_id)
else:
if document_id not in self._users:
self._users[document_id] = set()
@ -223,22 +255,31 @@ class YdocManager:
if self._redis:
redis_key = f'{self._redis_key_prefix}:{document_id}:users'
await self._redis.srem(redis_key, user_id)
# Keep the reverse index in sync.
session_key = f'{self._redis_key_prefix}:session:{user_id}:documents'
await self._redis.srem(session_key, document_id)
else:
if document_id in self._users and user_id in self._users[document_id]:
self._users[document_id].remove(user_id)
async def remove_user_from_all_documents(self, user_id: str):
if self._redis:
keys = []
async for key in self._redis.scan_iter(match=f'{self._redis_key_prefix}:*', count=100):
keys.append(key)
for key in keys:
if key.endswith(':users'):
await self._redis.srem(key, user_id)
# Use the per-session reverse index instead of a cluster-wide
# SCAN. This set contains only the document IDs that this
# session actually joined, so the cost is proportional to
# the session's footprint — not the total number of documents.
session_key = f'{self._redis_key_prefix}:session:{user_id}:documents'
document_ids = await self._redis.smembers(session_key)
document_id = key.split(':')[-2]
if len(await self.get_users(document_id)) == 0:
await self.clear_document(document_id)
for document_id in document_ids:
users_key = f'{self._redis_key_prefix}:{document_id}:users'
await self._redis.srem(users_key, user_id)
if len(await self.get_users(document_id)) == 0:
await self.clear_document(document_id)
# Clean up the reverse index itself.
await self._redis.delete(session_key)
else:
for document_id in list(self._users.keys()):

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

View file

@ -50,23 +50,25 @@ MAX_KNOWLEDGE_BASE_SEARCH_ITEMS = 10_000
async def _has_read_access_to_file(
file, user_id: str, user_role: str,
file,
user_id: str,
user_role: str,
model_knowledge: Optional[list[dict]] = None,
) -> bool:
"""Check if a user can read a file via ownership, admin role, model attachment, or access grants."""
if file.user_id == user_id or user_role == 'admin':
return True
if model_knowledge and any(
item.get('type') == 'file' and item.get('id') == file.id
for item in model_knowledge
):
if model_knowledge and any(item.get('type') == 'file' and item.get('id') == file.id for item in model_knowledge):
return True
from open_webui.utils.access_control.files import has_access_to_file
return await has_access_to_file(
file_id=file.id, access_type='read',
file_id=file.id,
access_type='read',
user=UserModel(**{'id': user_id, 'role': user_role}),
)
# =============================================================================
# TIME UTILITIES
# =============================================================================
@ -230,7 +232,7 @@ async def search_web(
max_count = 5 if configured is None else configured
count = max(1, min(count, max_count)) if count is not None else max_count
results = await asyncio.to_thread(_search_web, __request__, engine, query, user)
results = await _search_web(__request__, engine, query, user)
# Limit results
results = results[:count] if results else []
@ -272,7 +274,7 @@ async def fetch_url(
return content
except Exception as e:
log.exception(f'fetch_url error: {e}')
log.warning(f'fetch_url error: {e}')
return json.dumps({'error': str(e)})
@ -1713,6 +1715,21 @@ async def search_knowledge_files(
# No attached knowledge - search all accessible KBs
if knowledge_id:
# search_files_by_id does not enforce knowledge_id ownership; mirror the attached-KB check above.
knowledge = await Knowledges.get_knowledge_by_id(knowledge_id)
if not knowledge or not (
user_role == 'admin'
or knowledge.user_id == user_id
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
permission='read',
user_group_ids=set(user_group_ids),
)
):
return json.dumps({'error': f'Access denied to knowledge base {knowledge_id}'})
result = await Knowledges.search_files_by_id(
knowledge_id=knowledge_id,
user_id=user_id,
@ -2178,14 +2195,24 @@ async def view_knowledge_file(
async def list_knowledge(
knowledge_id: Optional[str] = None,
skip: int = 0,
count: int = 50,
__request__: Request = None,
__user__: dict = None,
__model_knowledge__: Optional[list[dict]] = None,
) -> str:
"""
List all knowledge bases, files, and notes attached to the current model.
List knowledge bases, files, and notes attached to the current model.
Use this first to discover what knowledge is available before querying or reading files.
Without knowledge_id: returns KB summaries (name, description, file_count)
plus standalone files and notes — no file listing inside KBs.
With knowledge_id: includes paginated file listing for that specific KB.
Use skip/count to page through large KBs.
:param knowledge_id: Optional KB ID to get file listing for
:param skip: Number of files to skip for pagination (default: 0)
:param count: Maximum files per page (default: 50, max: 200)
:return: JSON with knowledge_bases, files, and notes attached to this model
"""
if __request__ is None:
@ -2197,6 +2224,22 @@ async def list_knowledge(
if not __model_knowledge__:
return json.dumps({'knowledge_bases': [], 'files': [], 'notes': []})
# Coerce parameters from LLM tool calls (may come as strings)
if isinstance(skip, str):
try:
skip = int(skip)
except ValueError:
skip = 0
if isinstance(count, str):
try:
count = int(count)
except ValueError:
count = 50
if isinstance(knowledge_id, str) and knowledge_id.lower() in ('none', 'null', ''):
knowledge_id = None
count = min(count, 200)
try:
from open_webui.models.access_grants import AccessGrants
from open_webui.models.files import Files
@ -2238,9 +2281,15 @@ async def list_knowledge(
'file_count': file_count,
}
# Include file listing for each KB
if kb_files:
kb_entry['files'] = [{'id': f.id, 'filename': f.filename} for f in kb_files]
# Include file listing only when this KB is targeted
if knowledge_id and knowledge_id == knowledge.id:
if kb_files:
paged_files = kb_files[skip : skip + count]
kb_entry['files'] = [{'id': f.id, 'filename': f.filename} for f in paged_files]
kb_entry['files_skip'] = skip
kb_entry['files_count'] = len(paged_files)
kb_entry['files_total'] = file_count
kb_entry['has_more'] = skip + count < file_count
knowledge_bases.append(kb_entry)

View file

@ -34,9 +34,16 @@ MAX_GREP_MATCHES = 50
def is_regex_pattern(pattern: str) -> bool:
"""Detect if a pattern looks like regex (\|, .*, .+, \d, \w, \s, [...])."""
return ('\|' in pattern or '.*' in pattern or '.+' in pattern
or '.?' in pattern or '\d' in pattern or '\w' in pattern
or '\s' in pattern or bool(re.search(r'\[.+\]', pattern)))
return (
'\|' in pattern
or '.*' in pattern
or '.+' in pattern
or '.?' in pattern
or '\d' in pattern
or '\w' in pattern
or '\s' in pattern
or bool(re.search(r'\[.+\]', pattern))
)
def normalize_regex(pattern: str) -> str:
@ -44,8 +51,7 @@ def normalize_regex(pattern: str) -> str:
return pattern.replace('\\|', '|').replace('\|', '|')
def build_matcher(pattern: str, case_insensitive: bool = False,
use_regex: bool = False) -> tuple:
def build_matcher(pattern: str, case_insensitive: bool = False, use_regex: bool = False) -> tuple:
"""Build a matcher function. Returns (match_fn, error_str_or_None)."""
if not use_regex and is_regex_pattern(pattern):
use_regex = True
@ -157,6 +163,7 @@ async def _build_directory_tree(knowledge_id: str) -> dict:
# Compute full path for each directory
dir_id_to_path = {}
def _get_dir_path(dir_id):
if dir_id in dir_id_to_path:
return dir_id_to_path[dir_id]
@ -165,7 +172,7 @@ async def _build_directory_tree(knowledge_id: str) -> dict:
return ''
if d['parent_id'] and d['parent_id'] in dir_map:
parent_path = _get_dir_path(d['parent_id'])
path = f"{parent_path}/{d['name']}" if parent_path else d['name']
path = f'{parent_path}/{d["name"]}' if parent_path else d['name']
else:
path = d['name']
dir_id_to_path[dir_id] = path
@ -180,16 +187,20 @@ async def _build_directory_tree(knowledge_id: str) -> dict:
files = []
for file_model, directory_id in files_with_dirs:
if directory_id and directory_id in dir_id_to_path:
file_path = f"{dir_id_to_path[directory_id]}/{file_model.filename}"
file_path = f'{dir_id_to_path[directory_id]}/{file_model.filename}'
else:
file_path = file_model.filename
files.append({
'id': file_model.id, 'filename': file_model.filename,
'path': file_path, 'directory_id': directory_id,
'size': file_model.meta.get('size') if file_model.meta else None,
'type': file_model.meta.get('content_type') if file_model.meta else None,
'updated_at': file_model.updated_at,
})
files.append(
{
'id': file_model.id,
'filename': file_model.filename,
'path': file_path,
'directory_id': directory_id,
'size': file_model.meta.get('size') if file_model.meta else None,
'type': file_model.meta.get('content_type') if file_model.meta else None,
'updated_at': file_model.updated_at,
}
)
return {
'dirs': dir_map,
@ -212,10 +223,7 @@ def _get_files_in_dir(tree: dict, dir_id: str | None) -> list[dict]:
def _get_subdirs(tree: dict, parent_id: str | None) -> list[dict]:
"""Get immediate child directories."""
return sorted(
[d for d in tree['dirs'].values() if d['parent_id'] == parent_id],
key=lambda d: d['name']
)
return sorted([d for d in tree['dirs'].values() if d['parent_id'] == parent_id], key=lambda d: d['name'])
def _get_files_under_dir(tree: dict, dir_id: str) -> list[dict]:
@ -237,8 +245,9 @@ def _get_files_under_dir(tree: dict, dir_id: str) -> list[dict]:
# =============================================================================
async def _get_accessible_kb_ids(user: dict, model_knowledge: list[dict] | None,
knowledge_id: str | None = None) -> list[tuple[str, str, str]]:
async def _get_accessible_kb_ids(
user: dict, model_knowledge: list[dict] | None, knowledge_id: str | None = None
) -> list[tuple[str, str, str]]:
"""Get list of (kb_id, kb_name, kb_description) the user can access."""
from open_webui.models.access_grants import AccessGrants
from open_webui.models.groups import Groups
@ -249,11 +258,17 @@ async def _get_accessible_kb_ids(user: dict, model_knowledge: list[dict] | None,
user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)]
async def _has_access(kb):
return (user_role == 'admin' or kb.user_id == user_id
or await AccessGrants.has_access(
user_id=user_id, resource_type='knowledge',
resource_id=kb.id, permission='read',
user_group_ids=set(user_group_ids)))
return (
user_role == 'admin'
or kb.user_id == user_id
or await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=kb.id,
permission='read',
user_group_ids=set(user_group_ids),
)
)
result = []
@ -276,8 +291,10 @@ async def _get_accessible_kb_ids(user: dict, model_knowledge: list[dict] | None,
result.append((kb.id, kb.name, kb.description or ''))
else:
search = await Knowledges.search_knowledge_bases(
user_id, filter={'query': '', 'user_id': user_id, 'group_ids': user_group_ids},
skip=0, limit=50,
user_id,
filter={'query': '', 'user_id': user_id, 'group_ids': user_group_ids},
skip=0,
limit=50,
)
for kb in search.items:
result.append((kb.id, kb.name, kb.description or ''))
@ -285,8 +302,9 @@ async def _get_accessible_kb_ids(user: dict, model_knowledge: list[dict] | None,
return result
async def _get_accessible_files(user: dict, model_knowledge: list[dict] | None,
knowledge_id: str | None = None) -> list[dict]:
async def _get_accessible_files(
user: dict, model_knowledge: list[dict] | None, knowledge_id: str | None = None
) -> list[dict]:
"""Get all files the user can access, with KB metadata and directory_id (no path computation)."""
from open_webui.models.files import Files
from open_webui.models.knowledge import Knowledges
@ -297,15 +315,18 @@ async def _get_accessible_files(user: dict, model_knowledge: list[dict] | None,
for kb_id, kb_name, _ in kb_ids:
kb_files = await Knowledges.get_files_with_directory_ids(kb_id)
for file_model, dir_id in kb_files:
files.append({
'id': file_model.id, 'filename': file_model.filename,
'directory_id': dir_id,
'size': file_model.meta.get('size') if file_model.meta else None,
'type': file_model.meta.get('content_type') if file_model.meta else None,
'updated_at': file_model.updated_at,
'knowledge_id': kb_id,
'knowledge_name': kb_name,
})
files.append(
{
'id': file_model.id,
'filename': file_model.filename,
'directory_id': dir_id,
'size': file_model.meta.get('size') if file_model.meta else None,
'type': file_model.meta.get('content_type') if file_model.meta else None,
'updated_at': file_model.updated_at,
'knowledge_id': kb_id,
'knowledge_name': kb_name,
}
)
# Also handle directly attached files (not in any KB)
if model_knowledge:
@ -316,14 +337,18 @@ async def _get_accessible_files(user: dict, model_knowledge: list[dict] | None,
for fid in attached_file_ids:
f = await Files.get_file_by_id(fid)
if f:
files.append({
'id': f.id, 'filename': f.filename,
'directory_id': None,
'size': f.meta.get('size') if f.meta else None,
'type': f.meta.get('content_type') if f.meta else None,
'updated_at': f.updated_at,
'knowledge_id': None, 'knowledge_name': None,
})
files.append(
{
'id': f.id,
'filename': f.filename,
'directory_id': None,
'size': f.meta.get('size') if f.meta else None,
'type': f.meta.get('content_type') if f.meta else None,
'updated_at': f.updated_at,
'knowledge_id': None,
'knowledge_name': None,
}
)
return files
@ -374,8 +399,14 @@ async def _resolve_file(ref: str, user: dict, model_knowledge: list[dict] | None
if f and f.data:
if f.id not in accessible_ids:
return None
return {'id': f.id, 'filename': f.filename, 'content': f.data.get('content', ''),
'meta': f.meta, 'updated_at': f.updated_at, 'created_at': f.created_at}
return {
'id': f.id,
'filename': f.filename,
'content': f.data.get('content', ''),
'meta': f.meta,
'updated_at': f.updated_at,
'created_at': f.created_at,
}
# Try path match (e.g. "docs/api/auth.md") — lazy dir walk
ref_clean = ref.strip('/')
@ -388,15 +419,20 @@ async def _resolve_file(ref: str, user: dict, model_knowledge: list[dict] | None
if dir_id is None:
continue
# Find file with that name in that directory
matches = [fi for fi in accessible
if fi['filename'] == filename and fi['directory_id'] == dir_id]
matches = [fi for fi in accessible if fi['filename'] == filename and fi['directory_id'] == dir_id]
if len(matches) == 1:
f = await Files.get_file_by_id(matches[0]['id'])
if f and f.data:
return {'id': f.id, 'filename': f.filename, 'content': f.data.get('content', ''),
'meta': f.meta, 'updated_at': f.updated_at, 'created_at': f.created_at,
'knowledge_id': matches[0].get('knowledge_id'),
'knowledge_name': matches[0].get('knowledge_name')}
return {
'id': f.id,
'filename': f.filename,
'content': f.data.get('content', ''),
'meta': f.meta,
'updated_at': f.updated_at,
'created_at': f.created_at,
'knowledge_id': matches[0].get('knowledge_id'),
'knowledge_name': matches[0].get('knowledge_name'),
}
# Try filename match within accessible files
matches = [fi for fi in accessible if fi['filename'] == ref]
@ -404,13 +440,21 @@ async def _resolve_file(ref: str, user: dict, model_knowledge: list[dict] | None
if len(matches) == 1:
f = await Files.get_file_by_id(matches[0]['id'])
if f and f.data:
return {'id': f.id, 'filename': f.filename, 'content': f.data.get('content', ''),
'meta': f.meta, 'updated_at': f.updated_at, 'created_at': f.created_at,
'knowledge_id': matches[0].get('knowledge_id'),
'knowledge_name': matches[0].get('knowledge_name')}
return {
'id': f.id,
'filename': f.filename,
'content': f.data.get('content', ''),
'meta': f.meta,
'updated_at': f.updated_at,
'created_at': f.created_at,
'knowledge_id': matches[0].get('knowledge_id'),
'knowledge_name': matches[0].get('knowledge_name'),
}
elif len(matches) > 1:
return {'error': f'Ambiguous filename "{ref}". Use full path to disambiguate:\n' +
'\n'.join(f" {m['id']} {m['filename']} ({m.get('knowledge_name', 'direct')})" for m in matches)}
return {
'error': f'Ambiguous filename "{ref}". Use full path to disambiguate:\n'
+ '\n'.join(f' {m["id"]} {m["filename"]} ({m.get("knowledge_name", "direct")})' for m in matches)
}
return None
@ -418,6 +462,7 @@ async def _resolve_file(ref: str, user: dict, model_knowledge: list[dict] | None
async def _get_file_content(file_id: str) -> str | None:
"""Get file content by ID."""
from open_webui.models.files import Files
f = await Files.get_file_by_id(file_id)
if f and f.data:
return f.data.get('content', '')
@ -429,8 +474,7 @@ async def _get_file_content(file_id: str) -> str | None:
# =============================================================================
async def _kb_ls(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None) -> str:
async def _kb_ls(args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None) -> str:
"""List files and directories. Supports: ls, ls <path>, ls -a (flat)."""
from open_webui.models.knowledge import Knowledges
@ -506,13 +550,13 @@ def _fmt_size(f: dict) -> str:
def _fmt_date(f: dict) -> str:
if f.get('updated_at'):
from datetime import datetime, timezone
dt = datetime.fromtimestamp(f['updated_at'], tz=timezone.utc)
return dt.strftime('%Y-%m-%d')
return ''
async def _kb_cat(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None) -> str:
async def _kb_cat(args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None) -> str:
"""Read file content. Use -n for line numbers."""
if not args:
return 'Usage: cat [-n] <file_id or filename>'
@ -542,9 +586,9 @@ async def _kb_cat(args: list[str], flags: set[str], user: dict,
return content
async def _kb_head(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None,
piped_input: str | None = None) -> str:
async def _kb_head(
args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None, piped_input: str | None = None
) -> str:
"""First N lines of a file or piped input."""
n, args = _extract_numeric_flag(args)
if n is None:
@ -571,9 +615,9 @@ async def _kb_head(args: list[str], flags: set[str], user: dict,
return result
async def _kb_tail(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None,
piped_input: str | None = None) -> str:
async def _kb_tail(
args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None, piped_input: str | None = None
) -> str:
"""Last N lines of a file or piped input."""
n, args = _extract_numeric_flag(args)
if n is None:
@ -600,9 +644,9 @@ async def _kb_tail(args: list[str], flags: set[str], user: dict,
return result
async def _kb_grep(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None,
piped_input: str | None = None) -> str:
async def _kb_grep(
args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None, piped_input: str | None = None
) -> str:
"""Text search across files or piped input. Supports -E for regex."""
if not args:
return 'Usage: grep [-E] [-i] [-l] [-c] "pattern" [file] [*.ext]'
@ -684,8 +728,7 @@ async def _kb_grep(args: list[str], flags: set[str], user: dict,
accessible = [f for f in accessible if f['filename'].endswith(f'.{ext_filter}')]
if len(accessible) > MAX_GREP_FILES:
return (f'Too many files ({len(accessible)}). '
f'Scope your search: grep "{pattern}" docs/ or grep "{pattern}" *.py')
return f'Too many files ({len(accessible)}). Scope your search: grep "{pattern}" docs/ or grep "{pattern}" *.py'
from open_webui.models.files import Files
@ -717,9 +760,7 @@ async def _kb_grep(args: list[str], flags: set[str], user: dict,
if not count_only and not filenames_only:
for line_num, line_text in file_matches:
if len(results) < MAX_GREP_MATCHES:
results.append(
f'{file_info["id"]} {file_info["filename"]}:{line_num}: {line_text.rstrip()}'
)
results.append(f'{file_info["id"]} {file_info["filename"]}:{line_num}: {line_text.rstrip()}')
if count_only:
if not file_match_counts:
@ -742,8 +783,7 @@ async def _kb_grep(args: list[str], flags: set[str], user: dict,
return output
async def _kb_find(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None) -> str:
async def _kb_find(args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None) -> str:
"""Find files by name/glob pattern, optionally scoped to a directory."""
if not args:
return 'Usage: find "*.md" or find docs/ "*.md"'
@ -783,9 +823,9 @@ async def _kb_find(args: list[str], flags: set[str], user: dict,
return '\n'.join(lines)
async def _kb_wc(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None,
piped_input: str | None = None) -> str:
async def _kb_wc(
args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None, piped_input: str | None = None
) -> str:
"""Word, line, character counts."""
if piped_input is not None:
lines = piped_input.count('\n') + (1 if piped_input and not piped_input.endswith('\n') else 0)
@ -814,8 +854,7 @@ async def _kb_wc(args: list[str], flags: set[str], user: dict,
return f' {lines} {words} {chars} {resolved["filename"]}'
async def _kb_stat(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None) -> str:
async def _kb_stat(args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None) -> str:
"""File metadata."""
if not args:
return 'Usage: stat <file>'
@ -847,10 +886,12 @@ async def _kb_stat(args: list[str], flags: set[str], user: dict,
if resolved.get('created_at'):
from datetime import datetime, timezone
dt = datetime.fromtimestamp(resolved['created_at'], tz=timezone.utc)
out.append(f' Created: {dt.strftime("%Y-%m-%d %H:%M:%S UTC")}')
if resolved.get('updated_at'):
from datetime import datetime, timezone
dt = datetime.fromtimestamp(resolved['updated_at'], tz=timezone.utc)
out.append(f' Updated: {dt.strftime("%Y-%m-%d %H:%M:%S UTC")}')
if resolved.get('knowledge_name'):
@ -859,20 +900,20 @@ async def _kb_stat(args: list[str], flags: set[str], user: dict,
return '\n'.join(out)
async def _kb_sed(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None,
piped_input: str | None = None) -> str:
async def _kb_sed(
args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None, piped_input: str | None = None
) -> str:
"""Extract line range from a file. Usage: sed -n 'M,Np' <file>"""
if piped_input is not None:
# sed on piped input: parse range from args
start, end = 1, None
if 'n' in flags and args:
m = re.match(r"^(\d+),(\d+)p?$", args[0])
m = re.match(r'^(\d+),(\d+)p?$', args[0])
if m:
start, end = int(m.group(1)), int(m.group(2))
args = args[1:]
lines = piped_input.split('\n')
selected = lines[max(0, start - 1):(end or len(lines))]
selected = lines[max(0, start - 1) : (end or len(lines))]
return '\n'.join(selected)
# Parse: sed -n '40,60p' <file>
@ -903,7 +944,7 @@ async def _kb_sed(args: list[str], flags: set[str], user: dict,
lines = resolved['content'].split('\n')
total = len(lines)
selected = lines[max(0, start - 1):end]
selected = lines[max(0, start - 1) : end]
result = '\n'.join(selected)
result += f'\n[lines {start}-{min(end, total)} of {total}]'
return result
@ -914,8 +955,7 @@ async def _kb_sed(args: list[str], flags: set[str], user: dict,
# =============================================================================
async def _kb_tree(args: list[str], flags: set[str], user: dict,
model_knowledge: list[dict] | None) -> str:
async def _kb_tree(args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None) -> str:
"""Show directory tree structure."""
kb_ids = await _get_accessible_kb_ids(user, model_knowledge)
if not kb_ids:

View file

@ -66,7 +66,7 @@ async def has_access_to_file(
# Check if the file is associated with any chats the user has access to
shared_chat_ids = await Chats.get_shared_chat_ids_by_file_id(file_id, db=db)
if shared_chat_ids:
if access_type == 'read' and shared_chat_ids:
accessible_ids = await AccessGrants.get_accessible_resource_ids(
user_id=user.id,
resource_type='shared_chat',

View file

@ -426,7 +426,7 @@ async def get_current_user_by_api_key(request, api_key: str):
allowed_paths = [
path.strip() for path in str(request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',') if path.strip()
]
request_path = request.url.path
request_path = request.scope['path'] # Use raw ASGI path — not spoofable via Host header (CVE-2026-48710)
is_allowed = any(request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths)
if not is_allowed:
raise HTTPException(

View file

@ -500,9 +500,12 @@ async def _check_calendar_alerts(app) -> None:
now_ns = int(time.time_ns())
default_lookahead_ns = CALENDAR_ALERT_LOOKAHEAD_MINUTES * 60 * 1_000_000_000
# Grace window covers one poll cycle + jitter so "At time of event"
# alerts (alert_minutes=0) are not missed.
grace_ns = (SCHEDULER_POLL_INTERVAL + 5) * 1_000_000_000
async with get_async_db() as db:
upcoming = await CalendarEvents.get_upcoming_events(now_ns, default_lookahead_ns, db=db)
upcoming = await CalendarEvents.get_upcoming_events(now_ns, default_lookahead_ns, grace_ns=grace_ns, db=db)
if not upcoming:
return

View file

@ -160,9 +160,11 @@ async def generate_chat_completion(
if BYPASS_MODEL_ACCESS_CONTROL:
bypass_filter = True
# Propagate bypass_filter via request.state so that downstream route
# handlers (openai/ollama) can read it without exposing it as a query param.
# Propagate bypass_filter and bypass_system_prompt via request.state so that
# downstream route handlers (openai/ollama) can read them without exposing
# them as query parameters.
request.state.bypass_filter = bypass_filter
request.state.bypass_system_prompt = bypass_system_prompt
if hasattr(request.state, 'metadata'):
if 'metadata' not in form_data:
@ -279,7 +281,6 @@ async def generate_chat_completion(
request=request,
form_data=form_data,
user=user,
bypass_system_prompt=bypass_system_prompt,
)
if form_data.get('stream'):
response.headers['content-type'] = 'text/event-stream'
@ -295,7 +296,6 @@ async def generate_chat_completion(
request=request,
form_data=form_data,
user=user,
bypass_system_prompt=bypass_system_prompt,
)

View file

@ -22,6 +22,7 @@ from open_webui.models.chats import Chats
from open_webui.models.files import Files
from open_webui.retrieval.web.utils import validate_url
from open_webui.routers.files import upload_file_handler
from open_webui.utils.access_control.files import has_access_to_file
from open_webui.routers.images import (
get_image_data,
upload_image,
@ -50,7 +51,7 @@ _IMAGE_MIME_FALLBACK = {
}
async def get_image_base64_from_url(url: str) -> Optional[str]:
async def get_image_base64_from_url(url: str, user=None) -> Optional[str]:
try:
if url.startswith('http'):
# Validate URL to prevent SSRF attacks against local/private networks.
@ -70,25 +71,9 @@ async def get_image_base64_from_url(url: str) -> Optional[str]:
content_type = response.headers.get('Content-Type', 'image/png')
return f'data:{content_type};base64,{encoded_string}'
else:
file = await Files.get_file_by_id(url)
if not file:
return None
file_path = await asyncio.to_thread(Storage.get_file, file.path)
file_path = Path(file_path)
if file_path.is_file():
with open(file_path, 'rb') as image_file:
encoded_string = base64.b64encode(image_file.read()).decode('utf-8')
content_type = mimetypes.guess_type(file_path.name)[0] or (file.meta or {}).get('content_type')
if not content_type and ENABLE_IMAGE_CONTENT_TYPE_EXTENSION_FALLBACK:
content_type = _IMAGE_MIME_FALLBACK.get(file_path.suffix.lower())
if not content_type:
return None
return f'data:{content_type};base64,{encoded_string}'
else:
return None
# Non-URL string — treat as file_id. Delegate to the canonical
# file-ID resolver which enforces ownership/access checks.
return await get_image_base64_from_file_id(url, user=user)
except Exception as e:
return None
@ -194,11 +179,21 @@ async def get_file_url_from_base64(request, base64_file_string, metadata, user):
return None
async def get_image_base64_from_file_id(id: str) -> Optional[str]:
async def get_image_base64_from_file_id(id: str, user=None) -> Optional[str]:
file = await Files.get_file_by_id(id)
if not file:
return None
# Gate file-by-id resolution by ownership to prevent exfiltration.
# A caller could place another user's file_id in an image_url field;
# without this check the server reads the file from disk, inlines it
# base64 into the LLM request, and the content leaks via OCR/describe.
# Owner, admin, and explicit read-grant holders are allowed.
if user is None:
return None
if file.user_id != user.id and user.role != 'admin' and not await has_access_to_file(file.id, 'read', user):
return None
try:
file_path = await asyncio.to_thread(Storage.get_file, file.path)
file_path = Path(file_path)

View file

@ -1,14 +1,55 @@
import logging
import time
from typing import Any, Optional
from urllib.parse import quote
import jwt
from open_webui.env import (
FORWARD_USER_INFO_HEADER_JWT,
FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS,
FORWARD_USER_INFO_HEADER_JWT_SECRET,
FORWARD_USER_INFO_HEADER_USER_EMAIL,
FORWARD_USER_INFO_HEADER_USER_ID,
FORWARD_USER_INFO_HEADER_USER_NAME,
FORWARD_USER_INFO_HEADER_USER_ROLE,
)
log = logging.getLogger(__name__)
def _mint_forward_user_jwt(user: Any) -> str:
now = int(time.time())
payload = {
'sub': str(user.id),
'email': str(user.email),
'name': str(user.name),
'role': str(user.role),
'iss': 'open-webui',
'iat': now,
'exp': now + FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS,
}
return jwt.encode(payload, FORWARD_USER_INFO_HEADER_JWT_SECRET, algorithm='HS256')
def include_user_info_headers(headers: dict, user: Optional[Any] = None) -> dict:
"""
Forward user identity to external backends: signed JWT in
FORWARD_USER_INFO_HEADER_JWT if FORWARD_USER_INFO_HEADER_JWT_SECRET is set;
otherwise the legacy X-OpenWebUI-User-* headers.
"""
if user is None:
return headers
if FORWARD_USER_INFO_HEADER_JWT_SECRET:
try:
token = _mint_forward_user_jwt(user)
return {**headers, FORWARD_USER_INFO_HEADER_JWT: token}
except Exception:
log.exception(
'Failed to mint %s; falling back to plain user-info headers.',
FORWARD_USER_INFO_HEADER_JWT,
)
def include_user_info_headers(headers, user):
return {
**headers,
FORWARD_USER_INFO_HEADER_USER_NAME: quote(user.name, safe=' '),
@ -28,6 +69,8 @@ def get_custom_headers(custom_headers: dict, user=None, metadata: dict = None) -
'{{MESSAGE_ID}}': metadata.get('message_id', '') or '',
'{{USER_ID}}': (user.id if user else '') or '',
'{{USER_NAME}}': (user.name if user else '') or '',
'{{USER_EMAIL}}': (user.email if user else '') or '',
'{{USER_ROLE}}': (user.role if user else '') or '',
}
parsed_headers = {}

View file

@ -1,6 +1,7 @@
import json
import logging
import sys
import traceback
from typing import TYPE_CHECKING
from loguru import logger
@ -49,22 +50,42 @@ def _json_sink(message: 'Message') -> None:
Used as a Loguru sink when LOG_FORMAT is set to "json".
"""
record = message.record
log_entry = {
'ts': record['time'].strftime('%Y-%m-%dT%H:%M:%S.%f')[:-3] + 'Z',
'level': _LEVEL_MAP.get(record['level'].name, record['level'].name.lower()),
'msg': record['message'],
'caller': f'{record["name"]}:{record["function"]}:{record["line"]}',
}
try:
record = message.record
log_entry = {
'ts': record['time'].isoformat(timespec='milliseconds'),
'level': _LEVEL_MAP.get(record['level'].name, record['level'].name.lower()),
'msg': record['message'],
'caller': f'{record["name"]}:{record["function"]}:{record["line"]}',
}
if record['extra']:
log_entry['extra'] = record['extra']
if record['extra']:
log_entry['extra'] = record['extra']
if record['exception'] is not None:
log_entry['error'] = ''.join(record['exception'].format_exception()).rstrip()
exc = record['exception']
if exc is not None:
log_entry['error'] = {
'type': exc.type.__name__ if exc.type else None,
'message': str(exc.value) if exc.value else None,
'stacktrace': ''.join(traceback.format_exception(exc.type, exc.value, exc.traceback)).rstrip(),
}
sys.stdout.write(json.dumps(log_entry, ensure_ascii=False, default=str) + '\n')
sys.stdout.flush()
sys.stdout.write(json.dumps(log_entry, ensure_ascii=False, default=str) + '\n')
sys.stdout.flush()
except Exception:
# Last-resort fallback: never let a logging failure crash the application.
# Emit a minimal valid JSON line so the structured logging pipeline stays intact.
try:
fallback = {
'ts': message.record['time'].isoformat(timespec='milliseconds'),
'level': 'error',
'msg': f'[logging error] failed to serialize log record: {message}',
}
sys.stdout.write(json.dumps(fallback, ensure_ascii=False, default=str) + '\n')
sys.stdout.flush()
except Exception:
sys.stderr.write(f'[logging error] _json_sink failed: {message}\n')
sys.stderr.flush()
class InterceptHandler(logging.Handler):

View file

@ -11,7 +11,11 @@ from mcp import ClientSession
from mcp.client.auth import OAuthClientProvider, TokenStorage
from mcp.client.streamable_http import streamablehttp_client
from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken
from open_webui.env import AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL, AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER
from open_webui.env import (
AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL,
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER,
MCP_INITIALIZE_TIMEOUT,
)
def _build_httpx_client(headers=None, timeout=None, auth=None, verify=True):
@ -69,7 +73,7 @@ class MCPClient:
self._session_context = ClientSession(read_stream, write_stream) # pylint: disable=W0201
self.session = await exit_stack.enter_async_context(self._session_context)
with anyio.fail_after(10):
with anyio.fail_after(MCP_INITIALIZE_TIMEOUT):
await self.session.initialize()
self.exit_stack = exit_stack.pop_all()
except Exception as e:

View file

@ -18,7 +18,7 @@ from uuid import uuid4
from aiocache import cached
from fastapi import HTTPException, Request
from fastapi.responses import HTMLResponse
from fastapi.responses import HTMLResponse, JSONResponse
from open_webui.config import (
CACHE_DIR,
CODE_INTERPRETER_BLOCKED_MODULES,
@ -30,15 +30,12 @@ from open_webui.config import (
from open_webui.constants import TASKS
from open_webui.env import (
BYPASS_MODEL_ACCESS_CONTROL,
CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES,
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS,
CHAT_RESPONSE_STREAM_DELTA_CHUNK_SIZE,
ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION,
ENABLE_FORWARD_USER_INFO_HEADERS,
ENABLE_QUERIES_CACHE,
ENABLE_REALTIME_CHAT_SAVE,
ENABLE_RESPONSES_API_STATEFUL,
FORWARD_SESSION_INFO_HEADER_CHAT_ID,
FORWARD_SESSION_INFO_HEADER_MESSAGE_ID,
GLOBAL_LOG_LEVEL,
RAG_SYSTEM_CONTEXT,
)
@ -75,7 +72,7 @@ from open_webui.socket.main import (
get_event_call,
get_event_emitter,
)
from open_webui.utils.access_control import has_connection_access
from open_webui.utils.access_control import has_connection_access, has_permission
from open_webui.utils.access_control.files import get_accessible_folder_files
from open_webui.utils.chat import generate_chat_completion
from open_webui.utils.code_interpreter import execute_code_jupyter
@ -89,7 +86,7 @@ from open_webui.utils.filter import (
get_sorted_filter_ids,
process_filter_functions,
)
from open_webui.utils.headers import include_user_info_headers
from open_webui.utils.mcp.client import MCPClient
from open_webui.utils.misc import (
add_or_update_system_message,
@ -121,6 +118,7 @@ from open_webui.utils.task import (
tools_function_calling_generation_template,
)
from open_webui.utils.tools import (
build_tool_server_headers,
get_builtin_tools,
get_terminal_tools,
get_tools,
@ -1334,7 +1332,8 @@ async def chat_completion_tools_handler(
tool_function_name = tool_call.get('name', None)
if tool_function_name not in tools:
return body, {}
log.warning(f'Tool "{tool_function_name}" not found')
return
tool_function_params = tool_call.get('parameters', {})
@ -1522,6 +1521,18 @@ async def chat_web_search_handler(request: Request, form_data: dict, extra_param
user,
)
# generate_queries returns a JSONResponse on error (e.g. model not
# found, chat completion failure). Extract the error detail and
# re-raise so the outer except block falls back to using the raw
# user message as the search query.
if isinstance(res, JSONResponse):
try:
error_body = json.loads(res.body)
detail = error_body.get('detail', 'Query generation failed')
except Exception:
detail = 'Query generation failed'
raise Exception(detail)
response = res['choices'][0]['message']['content']
try:
@ -1853,6 +1864,15 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra
user,
)
# Handle JSONResponse from error paths
if isinstance(res, JSONResponse):
try:
error_body = json.loads(res.body)
detail = error_body.get('detail', 'Image prompt generation failed')
except Exception:
detail = 'Image prompt generation failed'
raise Exception(detail)
response = res['choices'][0]['message']['content']
try:
@ -2096,7 +2116,7 @@ def apply_params_to_form_data(form_data, model):
return form_data
async def convert_url_images_to_base64(form_data):
async def convert_url_images_to_base64(form_data, user=None):
messages = form_data.get('messages', [])
for message in messages:
@ -2117,7 +2137,7 @@ async def convert_url_images_to_base64(form_data):
continue
try:
base64_data = await get_image_base64_from_url(image_url)
base64_data = await get_image_base64_from_url(image_url, user=user)
if base64_data:
new_content.append(
{
@ -2223,21 +2243,76 @@ def extract_skill_ids_from_messages(messages: list[dict]) -> set[str]:
def strip_skill_mentions(messages: list[dict]) -> None:
"""Strip <$skillId|label> mention tags from message content in-place."""
strip_re = re.compile(r'<\$[^>]+>')
"""Replace <$skillId|label> mention tags with the label in message content in-place."""
strip_re = re.compile(r'<\$[^|>]+\|?([^>]*)>')
for message in messages:
content = message.get('content')
if isinstance(content, str) and strip_re.search(content):
message['content'] = strip_re.sub('', content).strip()
message['content'] = strip_re.sub(r'\1', content).strip()
elif isinstance(content, list):
for part in content:
if isinstance(part, dict) and part.get('type') == 'text':
text = part.get('text', '')
if strip_re.search(text):
part['text'] = strip_re.sub('', text).strip()
part['text'] = strip_re.sub(r'\1', text).strip()
async def connect_mcp_server(
request,
server_id: str,
user,
metadata: dict,
extra_params: dict,
) -> tuple[MCPClient, list[dict]] | None:
"""Resolve an MCP server connection, authenticate, and return (client, tool_specs).
Returns None if the server is not found or access is denied.
"""
mcp_server_connection = None
for server_connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
if server_connection.get('type', '') == 'mcp' and server_connection.get('info', {}).get('id') == server_id:
mcp_server_connection = server_connection
break
if not mcp_server_connection:
log.error(f'MCP server with id {server_id} not found')
return None
if not await has_connection_access(user, mcp_server_connection):
log.warning(f'Access denied to MCP server {server_id} for user {user.id}')
return None
headers, _ = await build_tool_server_headers(
mcp_server_connection,
request,
user,
server_id=server_id,
metadata=metadata,
extra_params=extra_params,
)
client = MCPClient()
await client.connect(
url=mcp_server_connection.get('url', ''),
headers=headers if headers else None,
)
function_name_filter_list = mcp_server_connection.get('config', {}).get('function_name_filter_list', '')
if isinstance(function_name_filter_list, str):
function_name_filter_list = function_name_filter_list.split(',')
tool_specs = await client.list_tool_specs()
if function_name_filter_list:
tool_specs = [spec for spec in tool_specs if is_string_allowed(spec['name'], function_name_filter_list)]
return client, tool_specs
async def process_chat_payload(request, form_data, user, metadata, model):
# Ensure chat_id is always a string — external API clients may omit it.
if not isinstance(metadata.get('chat_id'), str):
metadata['chat_id'] = ''
# Pipeline Inlet -> Filter Inlet -> Chat Memory -> Chat Web Search -> Chat Image Generation
# -> Chat Code Interpreter (Form Data Update) -> (Default) Chat Tools Function Calling
# -> Chat Files
@ -2340,7 +2415,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
except Exception:
pass
form_data = await convert_url_images_to_base64(form_data)
form_data = await convert_url_images_to_base64(form_data, user=user)
event_emitter = await get_event_emitter(metadata)
event_caller = await get_event_call(metadata)
@ -2396,7 +2471,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
if 'files' in folder.data:
# Defensive: filter to entries the caller can still read.
allowed_files = await get_accessible_folder_files(folder.data['files'], user)
if metadata.get('params', {}).get('function_calling') != 'native':
if metadata.get('params', {}).get('function_calling') == 'legacy':
form_data['files'] = [
*allowed_files,
*form_data.get('files', []),
@ -2410,7 +2485,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
user_message = get_last_user_message(form_data['messages'])
model_knowledge = model.get('info', {}).get('meta', {}).get('knowledge', False)
if model_knowledge and metadata.get('params', {}).get('function_calling') != 'native':
if model_knowledge and metadata.get('params', {}).get('function_calling') == 'legacy':
await event_emitter(
{
'type': 'status',
@ -2488,17 +2563,17 @@ async def process_chat_payload(request, form_data, user, metadata, model):
if 'memory' in features and features['memory']:
# Skip forced memory injection when native FC is enabled - model can use memory tools
if metadata.get('params', {}).get('function_calling') != 'native':
if metadata.get('params', {}).get('function_calling') == 'legacy':
form_data = await chat_memory_handler(request, form_data, extra_params, user)
if 'web_search' in features and features['web_search']:
# Skip forced RAG web search when native FC is enabled - model can use web_search tool
if metadata.get('params', {}).get('function_calling') != 'native':
if metadata.get('params', {}).get('function_calling') == 'legacy':
form_data = await chat_web_search_handler(request, form_data, extra_params, user)
if 'image_generation' in features and features['image_generation']:
# Skip forced image generation when native FC is enabled - model can use generate_image tool
if metadata.get('params', {}).get('function_calling') != 'native':
if metadata.get('params', {}).get('function_calling') == 'legacy':
form_data = await chat_image_generation_handler(request, form_data, extra_params, user)
if 'code_interpreter' in features and features['code_interpreter']:
@ -2506,7 +2581,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
# Skip XML-tag prompt injection when native FC is enabled —
# execute_code will be injected as a builtin tool instead
if metadata.get('params', {}).get('function_calling') != 'native':
if metadata.get('params', {}).get('function_calling') == 'legacy':
prompt = (
request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE
if request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE != ''
@ -2587,6 +2662,16 @@ async def process_chat_payload(request, form_data, user, metadata, model):
strip_skill_mentions(form_data.get('messages', []))
prompt = get_last_user_message(form_data['messages'])
# Guard against empty user message after skill mention stripping.
# When a user selects a skill ($skill-name) without typing additional text,
# the stripped result is an empty string which causes 400 errors on providers
# that reject empty content blocks (e.g. AWS Bedrock ConverseStream).
if not prompt or not prompt.strip():
fallback = ', '.join(s.name for s in available_skills)
if fallback:
set_last_user_message_content(fallback, form_data['messages'])
prompt = fallback
# TODO: re-enable URL extraction from prompt
# urls = []
# if prompt and len(prompt or "") < 500 and (not files or len(files) == 0):
@ -2642,79 +2727,19 @@ async def process_chat_payload(request, form_data, user, metadata, model):
try:
server_id = tool_id[len('server:mcp:') :]
mcp_server_connection = None
for server_connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
if (
server_connection.get('type', '') == 'mcp'
and server_connection.get('info', {}).get('id') == server_id
):
mcp_server_connection = server_connection
break
if not mcp_server_connection:
log.error(f'MCP server with id {server_id} not found')
result = await connect_mcp_server(
request,
server_id,
user,
metadata,
extra_params,
)
if result is None:
continue
# Check access control for MCP server
if not await has_connection_access(user, mcp_server_connection):
log.warning(f'Access denied to MCP server {server_id} for user {user.id}')
continue
client, tool_specs = result
mcp_clients[server_id] = client
auth_type = mcp_server_connection.get('auth_type', '')
headers = {}
if auth_type == 'bearer':
headers['Authorization'] = f'Bearer {mcp_server_connection.get("key", "")}'
elif auth_type == 'none':
# No authentication
pass
elif auth_type == 'session':
headers['Authorization'] = f'Bearer {request.state.token.credentials}'
elif auth_type == 'system_oauth':
oauth_token = extra_params.get('__oauth_token__', None)
if oauth_token:
headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}'
elif auth_type in ('oauth_2.1', 'oauth_2.1_static'):
try:
splits = server_id.split(':')
server_id = splits[-1] if len(splits) > 1 else server_id
oauth_token = await request.app.state.oauth_client_manager.get_oauth_token(
user.id, f'mcp:{server_id}'
)
if oauth_token:
headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}'
except Exception as e:
log.error(f'Error getting OAuth token: {e}')
oauth_token = None
connection_headers = mcp_server_connection.get('headers', None)
if connection_headers and isinstance(connection_headers, dict):
for key, value in connection_headers.items():
headers[key] = value
# Add user info headers if enabled
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
headers = include_user_info_headers(headers, user)
if metadata and metadata.get('chat_id'):
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
if metadata and metadata.get('message_id'):
headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id')
mcp_clients[server_id] = MCPClient()
await mcp_clients[server_id].connect(
url=mcp_server_connection.get('url', ''),
headers=headers if headers else None,
)
function_name_filter_list = mcp_server_connection.get('config', {}).get(
'function_name_filter_list', ''
)
if isinstance(function_name_filter_list, str):
function_name_filter_list = function_name_filter_list.split(',')
tool_specs = await mcp_clients[server_id].list_tool_specs()
for tool_spec in tool_specs:
async def make_tool_function(client, function_name):
@ -2726,12 +2751,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
return tool_function
if function_name_filter_list:
if not is_string_allowed(tool_spec['name'], function_name_filter_list):
# Skip this function
continue
tool_function = await make_tool_function(mcp_clients[server_id], tool_spec['name'])
tool_function = await make_tool_function(client, tool_spec['name'])
mcp_tools_dict[f'{server_id}_{tool_spec["name"]}'] = {
'spec': {
@ -2740,7 +2760,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
},
'callable': tool_function,
'type': 'mcp',
'client': mcp_clients[server_id],
'client': client,
'direct': False,
}
except Exception as e:
@ -2823,7 +2843,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
builtin_tools_enabled = (model.get('info', {}).get('meta', {}).get('capabilities') or {}).get(
'builtin_tools', True
)
if metadata.get('params', {}).get('function_calling') == 'native' and builtin_tools_enabled:
if metadata.get('params', {}).get('function_calling') != 'legacy' and builtin_tools_enabled:
# Add file context to user messages
chat_id = metadata.get('chat_id')
form_data['messages'] = await add_file_context(form_data.get('messages', []), chat_id, user)
@ -2846,7 +2866,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
# (e.g. pipe functions) can access all tools including MCP and builtins.
metadata['tools'] = tools_dict
if metadata.get('params', {}).get('function_calling') == 'native':
if metadata.get('params', {}).get('function_calling') != 'legacy':
# If the function calling is native, then call the tools function calling handler
form_data['tools'] = [
{'type': 'function', 'function': tool.get('spec', {})} for tool in tools_dict.values()
@ -3022,6 +3042,10 @@ async def get_system_oauth_token(request, user):
from open_webui.models.oauth_sessions import OAuthSessions
sessions = await OAuthSessions.get_sessions_by_user_id(user.id)
# Filter out MCP-provider sessions — their token refresh is handled
# separately by oauth_client_manager. Passing them to the SSO
# oauth_manager causes a failed refresh and session deletion (#24618).
sessions = [s for s in sessions if not (s.provider or '').startswith('mcp:')]
if sessions:
best = max(sessions, key=lambda s: s.updated_at)
oauth_token = await request.app.state.oauth_manager.get_oauth_token(
@ -3046,11 +3070,14 @@ async def background_tasks_handler(ctx):
if (
'chat_id' in metadata
and not metadata['chat_id'].startswith('local:')
and not metadata['chat_id'].startswith('channel:')
and not metadata.get('chat_id', '').startswith('local:')
and not metadata.get('chat_id', '').startswith('channel:')
):
messages_map = await Chats.get_messages_map_by_chat_id(metadata['chat_id'])
message = messages_map.get(metadata['message_id']) if messages_map else None
if not messages_map:
# Chat was deleted while the response was streaming — skip background tasks
return
message = messages_map.get(metadata['message_id'])
message_list = get_message_list(messages_map, metadata['message_id'])
@ -3364,13 +3391,22 @@ async def outlet_filter_handler(ctx):
outlet_message_id = message.get('id')
if outlet_message_id and outlet_message_id in messages_map:
original_message = messages_map[outlet_message_id]
if original_message.get('content') != message.get('content'):
content_changed = original_message.get('content') != message.get('content')
output_changed = message.get('output') and message.get('output') != original_message.get(
'output'
)
if content_changed or output_changed:
# If output was modified, re-derive content from it
new_content = message.get('content', original_message.get('content', ''))
if output_changed:
new_content = serialize_output(message['output'])
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
outlet_message_id,
{
'content': message['content'],
'content': new_content,
'originalContent': original_message.get('content'),
**({'output': message['output']} if output_changed else {}),
},
)
@ -3410,7 +3446,7 @@ async def non_streaming_chat_response_handler(response, ctx):
log.error('Provider returned error (non-streaming): %s', error)
if not metadata['chat_id'].startswith('channel:'):
if not metadata.get('chat_id', '').startswith('channel:'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
@ -3426,7 +3462,7 @@ async def non_streaming_chat_response_handler(response, ctx):
}
)
if 'selected_model_id' in response_data and not metadata['chat_id'].startswith('channel:'):
if 'selected_model_id' in response_data and not metadata.get('chat_id', '').startswith('channel:'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
@ -3449,7 +3485,7 @@ async def non_streaming_chat_response_handler(response, ctx):
title = (
await Chats.get_chat_title_by_id(metadata['chat_id'])
if not metadata['chat_id'].startswith('channel:')
if not metadata.get('chat_id', '').startswith('channel:')
else ''
)
@ -3482,7 +3518,7 @@ async def non_streaming_chat_response_handler(response, ctx):
# Save message in the database
usage = normalize_usage(response_data.get('usage', {}) or {})
if not metadata['chat_id'].startswith('channel:'):
if not metadata.get('chat_id', '').startswith('channel:'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
@ -3511,13 +3547,13 @@ async def non_streaming_chat_response_handler(response, ctx):
},
)
await background_tasks_handler(ctx)
ctx['assistant_message'] = {
'content': content,
'output': response_output,
**({'usage': usage} if usage else {}),
}
await outlet_filter_handler(ctx)
await background_tasks_handler(ctx)
response = build_response_object(response, merge_events_into_response(response_data, events))
except Exception as e:
@ -3836,7 +3872,26 @@ async def streaming_chat_response_handler(response, ctx):
reasoning_tags_param = metadata.get('params', {}).get('reasoning_tags')
DETECT_REASONING_TAGS = reasoning_tags_param is not False
DETECT_CODE_INTERPRETER = metadata.get('features', {}).get('code_interpreter', False)
# Mirror the five gates from utils/tools.py get_builtin_tools so the
# legacy XML-tag path enforces the same authz as native FC.
features = metadata.get('features', {}) or {}
model_capabilities = model.get('info', {}).get('meta', {}).get('capabilities') or {}
builtin_tools_meta = model.get('info', {}).get('meta', {}).get('builtinTools', {})
DETECT_CODE_INTERPRETER = (
bool(features.get('code_interpreter'))
and builtin_tools_meta.get('code_interpreter', True)
and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True)
and model_capabilities.get('code_interpreter', True)
and (
getattr(user, 'role', None) == 'admin'
or await has_permission(
getattr(user, 'id', ''),
'features.code_interpreter',
request.app.state.config.USER_PERMISSIONS,
)
)
)
reasoning_tags = []
if DETECT_REASONING_TAGS:
@ -4348,7 +4403,9 @@ async def streaming_chat_response_handler(response, ctx):
if end:
break
if ENABLE_REALTIME_CHAT_SAVE and not metadata['chat_id'].startswith('channel:'):
if ENABLE_REALTIME_CHAT_SAVE and not metadata.get('chat_id', '').startswith(
'channel:'
):
# Save message in the database
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
@ -4454,7 +4511,7 @@ async def streaming_chat_response_handler(response, ctx):
if response.background:
await response.background()
tool_call_retries = 0
tool_call_iterations = 0
tool_call_sources = [] # Track citation sources from tool results
all_tool_call_sources = [] # Accumulated sources across all iterations
user_message = get_last_user_message(form_data['messages'])
@ -4474,8 +4531,11 @@ async def streaming_chat_response_handler(response, ctx):
get_content_from_message(original_system_message) if original_system_message else None
)
while len(tool_calls) > 0 and tool_call_retries < CHAT_RESPONSE_MAX_TOOL_CALL_RETRIES:
tool_call_retries += 1
while tool_calls and (
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS is None
or tool_call_iterations < CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS
):
tool_call_iterations += 1
response_tool_calls = tool_calls.pop(0)
@ -4586,6 +4646,8 @@ async def streaming_chat_response_handler(response, ctx):
except Exception as e:
tool_result = str(e)
else:
tool_result = f'Error: Tool "{tool_function_name}" not found.'
tool_result, tool_result_files, tool_result_embeds = await process_tool_result(
request,
@ -4848,6 +4910,26 @@ async def streaming_chat_response_handler(response, ctx):
log.debug(e)
break
if (
CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS is not None
and tool_calls
and tool_call_iterations >= CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS
):
log.warning('Tool-call iteration limit reached (%s)', CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS)
error_content = f'Tool-call limit reached ({CHAT_RESPONSE_MAX_TOOL_CALL_ITERATIONS} iterations).'
if not metadata.get('chat_id', '').startswith('channel:'):
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],
metadata['message_id'],
{'error': {'content': error_content}},
)
await event_emitter(
{
'type': 'chat:message:error',
'data': {'error': {'content': error_content}},
}
)
if DETECT_CODE_INTERPRETER:
MAX_RETRIES = 5
retries = 0
@ -5026,7 +5108,7 @@ async def streaming_chat_response_handler(response, ctx):
title = (
await Chats.get_chat_title_by_id(metadata['chat_id'])
if not metadata['chat_id'].startswith('channel:')
if not metadata.get('chat_id', '').startswith('channel:')
else ''
)
data = {
@ -5037,7 +5119,7 @@ async def streaming_chat_response_handler(response, ctx):
**({'usage': usage} if usage else {}),
}
if not metadata['chat_id'].startswith('channel:'):
if not metadata.get('chat_id', '').startswith('channel:'):
if not ENABLE_REALTIME_CHAT_SAVE:
# Save message in the database
await Chats.upsert_message_to_chat_by_id_and_message_id(
@ -5086,13 +5168,13 @@ async def streaming_chat_response_handler(response, ctx):
}
)
await background_tasks_handler(ctx)
ctx['assistant_message'] = {
'content': serialize_output(output),
'output': output,
**({'usage': usage} if usage else {}),
}
await outlet_filter_handler(ctx)
await background_tasks_handler(ctx)
except asyncio.CancelledError:
log.warning('Task was cancelled!')
@ -5108,7 +5190,7 @@ async def streaming_chat_response_handler(response, ctx):
async def save_cancelled_state():
await event_emitter({'type': 'chat:tasks:cancel'})
if not metadata['chat_id'].startswith('channel:'):
if not metadata.get('chat_id', '').startswith('channel:'):
if not ENABLE_REALTIME_CHAT_SAVE:
await Chats.upsert_message_to_chat_by_id_and_message_id(
metadata['chat_id'],

View file

@ -130,6 +130,58 @@ def get_content_from_message(message: dict) -> str | None:
return None
def reconcile_tool_pairs(messages: list[dict]) -> list[dict]:
"""Drop unpaired tool_use / tool_result from a reconstructed conversation.
Stored output can be incomplete — a tool result may be missing (e.g. the
knowledge base was updated mid-chat, or the call was interrupted), or a
tool call may be missing while its result survived. Strict providers
(Anthropic, AWS Bedrock Converse) reject either direction of mismatch.
Well-formed output is unaffected: every id pairs, so nothing is stripped.
"""
completed_tool_call_ids = {
message['tool_call_id'] for message in messages if message.get('role') == 'tool' and message.get('tool_call_id')
}
requested_tool_call_ids = {
tool_call['id']
for message in messages
for tool_call in message.get('tool_calls') or ()
if message.get('role') == 'assistant' and tool_call.get('id')
}
reconciled_messages = []
for message in messages:
role = message.get('role')
# Orphan tool result — no assistant ever claimed this call_id.
if role == 'tool' and message.get('tool_call_id') not in requested_tool_call_ids:
continue
# Non-assistant or no tool_calls — pass through unchanged.
if role != 'assistant' or not message.get('tool_calls'):
reconciled_messages.append(message)
continue
# Keep only tool_calls whose id received a tool-role response.
valid_tool_calls = [
tool_call for tool_call in message['tool_calls'] if tool_call.get('id') in completed_tool_call_ids
]
if valid_tool_calls:
reconciled_messages.append({**message, 'tool_calls': valid_tool_calls})
continue
# All tool_calls were orphans — keep the message only if it
# carries meaningful text or reasoning content.
content = message.get('content', '')
has_meaningful_content = content.strip() if isinstance(content, str) else content
if has_meaningful_content or message.get('reasoning_content'):
reconciled_messages.append({key: value for key, value in message.items() if key != 'tool_calls'})
return reconciled_messages
def convert_output_to_messages(
output: list,
raw: bool = False,
@ -296,7 +348,7 @@ def convert_output_to_messages(
# Flush remaining content/tool_calls
flush_pending()
return messages
return reconcile_tool_pairs(messages)
def get_last_user_message(messages: list[dict]) -> str | None:
@ -653,7 +705,7 @@ def sanitize_data_for_db(obj):
# json.dumps is implemented in C and much faster than a Python-level
# recursive walk over every leaf string.
try:
if '\x00' not in json.dumps(obj, ensure_ascii=False):
if '\\u0000' not in json.dumps(obj, ensure_ascii=False):
return obj
except (TypeError, ValueError):
pass

View file

@ -2,7 +2,6 @@ import asyncio
import copy
import logging
import sys
import time
from aiocache import cached
from fastapi import Request
@ -36,7 +35,7 @@ async def fetch_ollama_models(request: Request, user: UserModel = None):
'id': model['model'],
'name': model['name'],
'object': 'model',
'created': int(time.time()),
'created': 0,
'owned_by': 'ollama',
'ollama': model,
'loaded': 'expires_at' in model,
@ -100,7 +99,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None)
'meta': model['meta'],
},
'object': 'model',
'created': int(time.time()),
'created': 0,
'owned_by': 'arena',
'arena': True,
}
@ -116,7 +115,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None)
'meta': DEFAULT_ARENA_MODEL['meta'],
},
'object': 'model',
'created': int(time.time()),
'created': 0,
'owned_by': 'arena',
'arena': True,
}

View file

@ -293,6 +293,7 @@ class ProtectedResourceMetadata:
resource: str | None = None
authorization_servers: list[str] = field(default_factory=list)
scopes_supported: list[str] = field(default_factory=list)
def get_discovery_urls(self, server_url: str) -> list[str]:
"""Build all candidate OAuth discovery URLs from this metadata and the server URL."""
@ -315,6 +316,7 @@ async def get_protected_resource_metadata(server_url: str) -> ProtectedResourceM
"""
authorization_servers = []
resource = None
scopes = []
try:
async with aiohttp.ClientSession(trust_env=True) as session:
async with session.post(
@ -359,6 +361,10 @@ async def get_protected_resource_metadata(server_url: str) -> ProtectedResourceM
log.debug(f'Discovered resource indicator: {resource}')
servers = resource_metadata.get('authorization_servers', [])
scopes = resource_metadata.get('scopes_supported', [])
if scopes:
log.debug(f'Discovered resource scopes: {scopes}')
if servers:
authorization_servers = servers
log.debug(f'Discovered authorization servers: {servers}')
@ -369,7 +375,9 @@ async def get_protected_resource_metadata(server_url: str) -> ProtectedResourceM
except Exception as e:
log.debug(f'MCP Protected Resource discovery failed: {e}')
return ProtectedResourceMetadata(resource=resource, authorization_servers=authorization_servers)
return ProtectedResourceMetadata(
resource=resource, authorization_servers=authorization_servers, scopes_supported=scopes
)
def _build_well_known_urls(server_url: str) -> list[str]:
@ -553,12 +561,11 @@ async def get_oauth_client_info_with_static_credentials(
log.error(f'Error parsing OAuth metadata from {url}: {e}')
continue
# Let the OAuth provider apply its default scopes.
# We intentionally do NOT join all scopes_supported here — that list
# represents every scope the server *can* grant, not what the client
# should request. Requesting all of them is almost always wrong and
# can break providers like Entra ID that require resource-specific scopes.
scope = None
# Use scopes from the Protected Resource Metadata (RFC 9728) if available.
# Unlike the Authorization Server's scopes_supported (which is a full catalog
# of every scope the server can grant), the PRM scopes_supported represents
# what this specific resource requires — making it safe to request them all.
scope = ' '.join(resource_metadata.scopes_supported) if resource_metadata.scopes_supported else None
# Determine token_endpoint_auth_method
token_endpoint_auth_method = 'client_secret_post'
@ -1077,6 +1084,18 @@ class OAuthManager:
log.warning(f'No OAuth session found for user {user_id}, session {session_id}')
return None
# Guard: MCP-provider sessions must be refreshed by
# oauth_client_manager, not the SSO OAuthManager. If one
# reaches here (e.g. via a stale cookie), bail out early
# instead of attempting a refresh that will fail and delete
# the session (#24618).
if (session.provider or '').startswith('mcp:'):
log.debug(
f'Skipping MCP session {session.id} (provider={session.provider}) '
f'in SSO OAuthManager — handled by oauth_client_manager'
)
return None
if (
force_refresh
or session.expires_at is None
@ -1264,7 +1283,7 @@ class OAuthManager:
if oauth_roles:
matched = False
for allowed_role in oauth_allowed_roles:
if allowed_role in oauth_roles:
if allowed_role == '*' or allowed_role in oauth_roles:
log.debug('Assigned user the user role')
role = 'user'
matched = True
@ -1448,7 +1467,13 @@ class OAuthManager:
'Authorization': f'Bearer {access_token}',
}
async with aiohttp.ClientSession(trust_env=True) as session:
async with session.get(picture_url, **get_kwargs, ssl=AIOHTTP_CLIENT_SESSION_SSL) as resp:
# allow_redirects=False prevents redirect-based SSRF: validate_url() only vetted the initial URL (CVE-2026-45401 cohort).
async with session.get(
picture_url,
**get_kwargs,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS,
) as resp:
if resp.ok:
picture = await resp.read()
base64_encoded_picture = base64.b64encode(picture).decode('utf-8')

View file

@ -1,10 +1,20 @@
"""Redis connection utilities.
Provides connection factory functions for standalone, Sentinel, and Cluster
Redis deployments, with optional async support and automatic connection caching.
"""
from __future__ import annotations
import asyncio
import inspect
import logging
import time
from urllib.parse import urlparse
from typing import Any
from urllib.parse import ParseResult, urlparse
import redis as _redis_sync
import redis
from open_webui.env import (
REDIS_CLUSTER,
REDIS_HEALTH_CHECK_INTERVAL,
@ -19,286 +29,299 @@ from open_webui.env import (
log = logging.getLogger(__name__)
MAX_RETRY_COUNT = REDIS_SENTINEL_MAX_RETRY_COUNT
_ACCEPTED_SCHEMES = frozenset({'redis', 'rediss'})
_SENTINEL_RETRYABLE = (
_redis_sync.exceptions.ConnectionError,
_redis_sync.exceptions.ReadOnlyError,
)
_FACTORY_METHODS = frozenset({'pipeline', 'pubsub', 'monitor', 'client', 'transaction'})
_CONNECTION_POOL: dict[tuple, Any] = {}
# Let not our connections be timed out but deliver them from
# partition. For the cache and the socket and the uptime
# belong to the one who first opened them, now and always.
_CONNECTION_CACHE = {}
class SentinelRedisProxy:
def __init__(self, sentinel, service, *, async_mode: bool = True, **kw):
self._sentinel = sentinel
self._service = service
self._kw = kw
self._async_mode = async_mode
def _master(self):
return self._sentinel.master_for(self._service, **self._kw)
def __getattr__(self, item):
master = self._master()
orig_attr = getattr(master, item)
if not callable(orig_attr):
return orig_attr
FACTORY_METHODS = {'pipeline', 'pubsub', 'monitor', 'client', 'transaction'}
if item in FACTORY_METHODS:
return orig_attr
if self._async_mode:
if inspect.isasyncgenfunction(orig_attr):
def _wrapped_iter(*args, **kwargs):
async def _iter():
for i in range(REDIS_SENTINEL_MAX_RETRY_COUNT):
try:
method = getattr(self._master(), item)
async for value in method(*args, **kwargs):
yield value
return
except (
redis.exceptions.ConnectionError,
redis.exceptions.ReadOnlyError,
) as e:
if i < REDIS_SENTINEL_MAX_RETRY_COUNT - 1:
log.debug(
'Redis sentinel fail-over (%s). Retry %s/%s',
type(e).__name__,
i + 1,
REDIS_SENTINEL_MAX_RETRY_COUNT,
)
if REDIS_RECONNECT_DELAY:
time.sleep(REDIS_RECONNECT_DELAY / 1000)
continue
log.error(
'Redis operation failed after %s retries: %s',
REDIS_SENTINEL_MAX_RETRY_COUNT,
e,
)
raise e from e
return _iter()
return _wrapped_iter
async def _wrapped(*args, **kwargs):
for i in range(REDIS_SENTINEL_MAX_RETRY_COUNT):
try:
method = getattr(self._master(), item)
result = method(*args, **kwargs)
if inspect.iscoroutine(result):
return await result
return result
except (
redis.exceptions.ConnectionError,
redis.exceptions.ReadOnlyError,
) as e:
if i < REDIS_SENTINEL_MAX_RETRY_COUNT - 1:
log.debug(
'Redis sentinel fail-over (%s). Retry %s/%s',
type(e).__name__,
i + 1,
REDIS_SENTINEL_MAX_RETRY_COUNT,
)
if REDIS_RECONNECT_DELAY:
await asyncio.sleep(REDIS_RECONNECT_DELAY / 1000)
continue
log.error(
'Redis operation failed after %s retries: %s',
REDIS_SENTINEL_MAX_RETRY_COUNT,
e,
)
raise e from e
return _wrapped
else:
def _wrapped(*args, **kwargs):
for i in range(REDIS_SENTINEL_MAX_RETRY_COUNT):
try:
method = getattr(self._master(), item)
return method(*args, **kwargs)
except (
redis.exceptions.ConnectionError,
redis.exceptions.ReadOnlyError,
) as e:
if i < REDIS_SENTINEL_MAX_RETRY_COUNT - 1:
log.debug(
'Redis sentinel fail-over (%s). Retry %s/%s',
type(e).__name__,
i + 1,
REDIS_SENTINEL_MAX_RETRY_COUNT,
)
if REDIS_RECONNECT_DELAY:
time.sleep(REDIS_RECONNECT_DELAY / 1000)
continue
log.error(
'Redis operation failed after %s retries: %s',
REDIS_SENTINEL_MAX_RETRY_COUNT,
e,
)
raise e from e
return _wrapped
def parse_redis_service_url(redis_url):
parsed_url = urlparse(redis_url)
if parsed_url.scheme != 'redis' and parsed_url.scheme != 'rediss':
raise ValueError("Invalid Redis URL scheme. Must be 'redis' or 'rediss'.")
def parse_redis_url(url: str) -> dict[str, Any]:
"""Break a ``redis://`` URL into its parts: service, port, db, username, password."""
parts: ParseResult = urlparse(url)
if parts.scheme not in _ACCEPTED_SCHEMES:
raise ValueError(f"Invalid Redis URL scheme '{parts.scheme}'; expected 'redis' or 'rediss'.")
return {
'username': parsed_url.username or None,
'password': parsed_url.password or None,
'service': parsed_url.hostname or 'mymaster',
'port': parsed_url.port or 6379,
'db': int(parsed_url.path.lstrip('/') or 0),
'service': parts.hostname or 'mymaster',
'port': parts.port or 6379,
'db': int(parts.path.lstrip('/') or '0'),
'username': parts.username or None,
'password': parts.password or None,
}
def get_redis_client(async_mode=False):
parse_redis_service_url = parse_redis_url
def get_sentinels_from_env(
hosts_csv: str | None,
port: str | int | None,
) -> list[tuple[str, int]]:
"""Turn a comma-separated host string into ``[(host, port), …]``."""
if not hosts_csv:
return []
resolved_port = int(port) if port else 26379
return [(host.strip(), resolved_port) for host in hosts_csv.split(',') if host.strip()]
def build_sentinel_url(
base_url: str,
hosts_csv: str,
port: str | int,
) -> str:
"""Construct a ``redis+sentinel://`` connection string.
``base_url`` supplies credentials, db index, and master service name.
``hosts_csv`` is a comma-separated list of sentinel hostnames.
"""
cfg = parse_redis_url(base_url)
auth = ''
if cfg['username'] or cfg['password']:
auth = f'{cfg["username"] or ""}:{cfg["password"] or ""}@'
nodes = ','.join(f'{host.strip()}:{port}' for host in hosts_csv.split(',') if host.strip())
return f'redis+sentinel://{auth}{nodes}/{cfg["db"]}/{cfg["service"]}'
def get_redis_client(async_mode: bool = False) -> Any | None:
"""Create a Redis connection using settings from environment variables.
Returns ``None`` when Redis is not configured or the connection fails.
"""
sentinel_list = get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT)
if not REDIS_URL and not sentinel_list:
return None
try:
return get_redis_connection(
redis_url=REDIS_URL,
redis_sentinels=get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT),
REDIS_URL,
redis_sentinels=sentinel_list,
redis_cluster=REDIS_CLUSTER,
async_mode=async_mode,
)
except Exception as e:
log.debug(f'Failed to get Redis client: {e}')
except Exception:
log.debug('Could not establish Redis connection', exc_info=True)
return None
# ---------------------------------------------------------------------------
# Sentinel proxy with automatic failover retry
# ---------------------------------------------------------------------------
class SentinelRedisProxy:
"""Transparent proxy that re-resolves the Sentinel master on connection errors.
Every call (sync or async) is wrapped with retry logic so that transient
Sentinel failovers are handled without caller intervention.
"""
def __init__(
self,
sentinel: Any,
service_name: str,
*,
async_mode: bool = True,
) -> None:
self._sentinel = sentinel
self._service_name = service_name
self._async_mode = async_mode
def __getattr__(self, name: str) -> Any:
"""Proxy attribute access with automatic Sentinel failover retry."""
current_master = self._sentinel.master_for(self._service_name)
original = getattr(current_master, name)
# Non-callable or factory attributes pass through without wrapping.
if not callable(original) or name in _FACTORY_METHODS:
return original
# Select the retry wrapper matching the execution mode.
if not self._async_mode:
return self._wrap_sync(name)
return self._wrap_async(name, original)
def _resolve_master(self) -> Any:
"""Ask Sentinel for the current master connection."""
return self._sentinel.master_for(self._service_name)
def _should_retry(self, attempt: int) -> bool:
return attempt < REDIS_SENTINEL_MAX_RETRY_COUNT - 1
def _log_retry(self, exc: Exception, attempt: int) -> None:
log.debug(
'Sentinel failover (%s) — retry %d/%d',
type(exc).__name__,
attempt + 1,
REDIS_SENTINEL_MAX_RETRY_COUNT,
)
def _log_exhausted(self, exc: Exception) -> None:
log.error(
'Redis operation failed after %d retries: %s',
REDIS_SENTINEL_MAX_RETRY_COUNT,
exc,
)
# -- async wrappers -----------------------------------------------------
def _wrap_async(self, name: str, attr: Any) -> Any:
if inspect.isasyncgenfunction(attr):
return self._wrap_async_gen(name)
return self._wrap_async_call(name)
def _wrap_async_gen(self, name: str) -> Any:
proxy = self
def wrapper(*args: Any, **kwargs: Any) -> Any:
async def _inner():
for attempt in range(REDIS_SENTINEL_MAX_RETRY_COUNT):
try:
method = getattr(proxy._resolve_master(), name)
async for value in method(*args, **kwargs):
yield value
return
except _SENTINEL_RETRYABLE as exc:
if proxy._should_retry(attempt):
proxy._log_retry(exc, attempt)
if REDIS_RECONNECT_DELAY:
time.sleep(REDIS_RECONNECT_DELAY / 1000)
continue
proxy._log_exhausted(exc)
raise
return _inner()
return wrapper
def _wrap_async_call(self, name: str) -> Any:
proxy = self
async def wrapper(*args: Any, **kwargs: Any) -> Any:
for attempt in range(REDIS_SENTINEL_MAX_RETRY_COUNT):
try:
method = getattr(proxy._resolve_master(), name)
result = method(*args, **kwargs)
if inspect.iscoroutine(result):
return await result
return result
except _SENTINEL_RETRYABLE as exc:
if proxy._should_retry(attempt):
proxy._log_retry(exc, attempt)
if REDIS_RECONNECT_DELAY:
await asyncio.sleep(REDIS_RECONNECT_DELAY / 1000)
continue
proxy._log_exhausted(exc)
raise
return wrapper
# -- sync wrapper -------------------------------------------------------
def _wrap_sync(self, name: str) -> Any:
proxy = self
def wrapper(*args: Any, **kwargs: Any) -> Any:
for attempt in range(REDIS_SENTINEL_MAX_RETRY_COUNT):
try:
method = getattr(proxy._resolve_master(), name)
return method(*args, **kwargs)
except _SENTINEL_RETRYABLE as exc:
if proxy._should_retry(attempt):
proxy._log_retry(exc, attempt)
if REDIS_RECONNECT_DELAY:
time.sleep(REDIS_RECONNECT_DELAY / 1000)
continue
proxy._log_exhausted(exc)
raise
return wrapper
# ---------------------------------------------------------------------------
# Connection factory
# ---------------------------------------------------------------------------
def _socket_options() -> dict[str, Any]:
"""Collect optional socket-level kwargs once instead of repeating them."""
opts: dict[str, Any] = {}
if REDIS_SOCKET_CONNECT_TIMEOUT is not None:
opts['socket_connect_timeout'] = REDIS_SOCKET_CONNECT_TIMEOUT
if REDIS_SOCKET_KEEPALIVE:
opts['socket_keepalive'] = True
if REDIS_HEALTH_CHECK_INTERVAL:
opts['health_check_interval'] = REDIS_HEALTH_CHECK_INTERVAL
return opts
def _build_sentinel(
redis_module: Any,
url: str,
sentinels: list[tuple[str, int]],
decode_responses: bool,
async_mode: bool,
) -> SentinelRedisProxy:
"""Create a SentinelRedisProxy from a redis URL and sentinel list."""
cfg = parse_redis_url(url)
sentinel = redis_module.sentinel.Sentinel(
sentinels,
port=cfg['port'],
db=cfg['db'],
username=cfg['username'],
password=cfg['password'],
decode_responses=decode_responses,
socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT,
**{k: v for k, v in _socket_options().items() if k != 'socket_connect_timeout'},
)
return SentinelRedisProxy(sentinel, cfg['service'], async_mode=async_mode)
def get_redis_connection(
redis_url,
redis_sentinels,
redis_cluster=False,
async_mode=False,
decode_responses=True,
):
redis_url: str | None,
redis_sentinels: list[tuple[str, int]] | None = None,
redis_cluster: bool = False,
async_mode: bool = False,
decode_responses: bool = True,
) -> Any | None:
"""Return a cached Redis connection (or create one).
Supports three topologies in order of precedence:
1. **Sentinel** — when ``redis_sentinels`` is non-empty.
2. **Cluster** — when ``redis_cluster`` is True.
3. **Standalone** — plain ``redis://`` connection.
"""
cache_key = (
redis_url,
tuple(redis_sentinels) if redis_sentinels else (),
async_mode,
decode_responses,
)
if cache_key in _CONNECTION_POOL:
return _CONNECTION_POOL[cache_key]
if cache_key in _CONNECTION_CACHE:
return _CONNECTION_CACHE[cache_key]
connection = None
connect_timeout_kwargs = (
{'socket_connect_timeout': REDIS_SOCKET_CONNECT_TIMEOUT} if REDIS_SOCKET_CONNECT_TIMEOUT is not None else {}
)
keepalive_kwargs = {'socket_keepalive': True} if REDIS_SOCKET_KEEPALIVE else {}
health_check_kwargs = {'health_check_interval': REDIS_HEALTH_CHECK_INTERVAL} if REDIS_HEALTH_CHECK_INTERVAL else {}
extra = _socket_options()
connection: Any = None
# Pick the right redis module for sync vs async.
if async_mode:
import redis.asyncio as redis
# If using sentinel in async mode
if redis_sentinels:
redis_config = parse_redis_service_url(redis_url)
sentinel = redis.sentinel.Sentinel(
redis_sentinels,
port=redis_config['port'],
db=redis_config['db'],
username=redis_config['username'],
password=redis_config['password'],
decode_responses=decode_responses,
socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT,
**keepalive_kwargs,
**health_check_kwargs,
)
connection = SentinelRedisProxy(
sentinel,
redis_config['service'],
async_mode=async_mode,
)
elif redis_cluster:
if not redis_url:
raise ValueError('Redis URL must be provided for cluster mode.')
return redis.cluster.RedisCluster.from_url(
redis_url,
decode_responses=decode_responses,
**connect_timeout_kwargs,
**keepalive_kwargs,
**health_check_kwargs,
)
elif redis_url:
connection = redis.from_url(
redis_url,
decode_responses=decode_responses,
**connect_timeout_kwargs,
**keepalive_kwargs,
**health_check_kwargs,
)
import redis.asyncio as redis_mod
else:
import redis
import redis as redis_mod # type: ignore[no-redef]
if redis_sentinels:
redis_config = parse_redis_service_url(redis_url)
sentinel = redis.sentinel.Sentinel(
redis_sentinels,
port=redis_config['port'],
db=redis_config['db'],
username=redis_config['username'],
password=redis_config['password'],
decode_responses=decode_responses,
socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT,
**keepalive_kwargs,
**health_check_kwargs,
)
connection = SentinelRedisProxy(
sentinel,
redis_config['service'],
async_mode=async_mode,
)
elif redis_cluster:
if not redis_url:
raise ValueError('Redis URL must be provided for cluster mode.')
return redis.cluster.RedisCluster.from_url(
redis_url,
decode_responses=decode_responses,
**connect_timeout_kwargs,
**keepalive_kwargs,
**health_check_kwargs,
)
elif redis_url:
connection = redis.Redis.from_url(
redis_url,
decode_responses=decode_responses,
**connect_timeout_kwargs,
**keepalive_kwargs,
**health_check_kwargs,
)
if redis_sentinels:
connection = _build_sentinel(redis_mod, redis_url, redis_sentinels, decode_responses, async_mode)
elif redis_cluster:
if not redis_url:
raise ValueError('Redis URL is required for cluster mode.')
connection = redis_mod.cluster.RedisCluster.from_url(
redis_url,
decode_responses=decode_responses,
**extra,
)
elif redis_url:
factory = getattr(redis_mod, 'from_url', None) or redis_mod.Redis.from_url
connection = factory(redis_url, decode_responses=decode_responses, **extra)
_CONNECTION_CACHE[cache_key] = connection
_CONNECTION_POOL[cache_key] = connection
return connection
def get_sentinels_from_env(sentinel_hosts_env, sentinel_port_env):
if sentinel_hosts_env:
sentinel_hosts = sentinel_hosts_env.split(',')
sentinel_port = int(sentinel_port_env)
return [(host, sentinel_port) for host in sentinel_hosts]
return []
def get_sentinel_url_from_env(redis_url, sentinel_hosts_env, sentinel_port_env):
redis_config = parse_redis_service_url(redis_url)
username = redis_config['username'] or ''
password = redis_config['password'] or ''
auth_part = ''
if username or password:
auth_part = f'{username}:{password}@'
hosts_part = ','.join(f'{host}:{sentinel_port_env}' for host in sentinel_hosts_env.split(','))
return f'redis+sentinel://{auth_part}{hosts_part}/{redis_config["db"]}/{redis_config["service"]}'

View file

@ -101,6 +101,68 @@ from pydantic.fields import FieldInfo
log = logging.getLogger(__name__)
async def build_tool_server_headers(
connection: dict,
request,
user,
server_id: str = '',
metadata: dict | None = None,
extra_params: dict | None = None,
) -> tuple[dict, dict]:
"""Build auth headers and cookies for a tool server connection.
Handles bearer, session, system_oauth, and oauth_2.1 auth types plus
custom header interpolation and user-info forwarding.
Shared by MCP and OpenAPI paths.
Returns (headers, cookies).
"""
extra_params = extra_params or {}
metadata = metadata or {}
auth_type = connection.get('auth_type', 'bearer')
headers = {}
cookies = {}
if auth_type == 'bearer':
headers['Authorization'] = f'Bearer {connection.get("key", "")}'
elif auth_type == 'session':
cookies = request.cookies if hasattr(request, 'cookies') else {}
headers['Authorization'] = f'Bearer {request.state.token.credentials}'
elif auth_type == 'system_oauth':
cookies = request.cookies if hasattr(request, 'cookies') else {}
oauth_token = extra_params.get('__oauth_token__', None)
if oauth_token:
headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}'
elif auth_type in ('oauth_2.1', 'oauth_2.1_static'):
try:
splits = server_id.split(':')
oauth_server_id = splits[-1] if len(splits) > 1 else server_id
connection_type = connection.get('type', 'openapi')
oauth_token = await request.app.state.oauth_client_manager.get_oauth_token(
user.id, f'{connection_type}:{oauth_server_id}'
)
if oauth_token:
headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}'
except Exception as e:
log.error(f'Error getting OAuth token: {e}')
# Interpolate template vars in custom connection headers
connection_headers = connection.get('headers', None)
if connection_headers and isinstance(connection_headers, dict):
headers.update(get_custom_headers(connection_headers, user, metadata))
# Add user info headers if enabled
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
headers = include_user_info_headers(headers, user)
if metadata.get('chat_id'):
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata['chat_id']
if metadata.get('message_id'):
headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata['message_id']
return headers, cookies
# Let no function be called without need, and let what
# it yields justify the cost of running it.
async def get_async_tool_function_and_apply_extra_params(
@ -166,8 +228,11 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
# Get user's group memberships for access control checks
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
# Batch-fetch all DB tools in one query instead of one per tool_id
tool_models = await Tools.get_tools_by_ids(tool_ids)
for tool_id in tool_ids:
tool = await Tools.get_tool_by_id(tool_id)
tool = tool_models.get(tool_id)
if tool:
# Check access control for local tools
if (
@ -309,41 +374,16 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
# Skip this function
continue
auth_type = tool_server_connection.get('auth_type', 'bearer')
cookies = {}
headers = {
'Content-Type': 'application/json',
}
if auth_type == 'bearer':
headers['Authorization'] = f'Bearer {tool_server_connection.get("key", "")}'
elif auth_type == 'none':
# No authentication
pass
elif auth_type == 'session':
cookies = request.cookies
headers['Authorization'] = f'Bearer {request.state.token.credentials}'
elif auth_type == 'system_oauth':
cookies = request.cookies
oauth_token = extra_params.get('__oauth_token__', None)
if oauth_token:
headers['Authorization'] = f'Bearer {oauth_token.get("access_token", "")}'
connection_headers = tool_server_connection.get('headers', None)
if connection_headers and isinstance(connection_headers, dict):
metadata = extra_params.get('__metadata__', {})
custom_headers = get_custom_headers(connection_headers, user, metadata)
headers.update(custom_headers)
# Add user info headers if enabled
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
headers = include_user_info_headers(headers, user)
metadata = extra_params.get('__metadata__', {})
if metadata and metadata.get('chat_id'):
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
if metadata and metadata.get('message_id'):
headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id')
metadata = extra_params.get('__metadata__', {})
headers, cookies = await build_tool_server_headers(
tool_server_connection,
request,
user,
server_id=server_id,
metadata=metadata,
extra_params=extra_params,
)
headers.setdefault('Content-Type', 'application/json')
async def make_tool_function(function_name, tool_server_data, headers):
async def tool_function(**kwargs):
@ -448,7 +488,9 @@ async def get_builtin_tools(
builtin_functions.append(query_knowledge_bases)
builtin_functions.append(search_knowledge_bases)
elif model_knowledge:
builtin_functions.extend([list_knowledge, search_knowledge_files, grep_knowledge_files, query_knowledge_files])
builtin_functions.extend(
[list_knowledge, search_knowledge_files, grep_knowledge_files, query_knowledge_files]
)
knowledge_types = {item.get('type') for item in model_knowledge}
if 'file' in knowledge_types or 'collection' in knowledge_types:
@ -456,11 +498,17 @@ async def get_builtin_tools(
if 'note' in knowledge_types:
builtin_functions.append(view_note)
else:
builtin_functions.extend([
list_knowledge_bases, search_knowledge_bases, query_knowledge_bases,
grep_knowledge_files, search_knowledge_files, query_knowledge_files,
view_knowledge_file,
])
builtin_functions.extend(
[
list_knowledge_bases,
search_knowledge_bases,
query_knowledge_bases,
grep_knowledge_files,
search_knowledge_files,
query_knowledge_files,
view_knowledge_file,
]
)
# Chats tools - search and fetch user's chat history
if is_builtin_tool_enabled('chats'):
@ -1477,7 +1525,7 @@ async def execute_tool_server(
except Exception as err:
error = str(err)
log.exception(f'API Request Error: {error}')
log.warning(f'API Request Error: {error}')
return ({'error': error}, None)

View file

@ -3,7 +3,10 @@
import re
from urllib.parse import urlparse
from open_webui.env import PROFILE_IMAGE_ALLOWED_MIME_TYPES
from open_webui.env import (
PROFILE_IMAGE_ALLOWED_MIME_TYPES,
PROFILE_IMAGE_MAX_DATA_URI_SIZE,
)
_USER_PROFILE_IMAGE_RE = re.compile(r'^/api/v1/users/[^/?#]+/profile/image$')
@ -40,6 +43,7 @@ def validate_profile_image_url(url: str) -> str:
- SVG data URIs (can contain embedded scripts)
- Arbitrary relative paths (prevents authenticated GET triggers)
- Scheme-relative URLs (``//host/path``)
- data URIs larger than PROFILE_IMAGE_MAX_DATA_URI_SIZE bytes
"""
if not url:
return url
@ -70,6 +74,10 @@ def validate_profile_image_url(url: str) -> str:
# The regex enforces the ;base64, boundary and is case-insensitive
# per the data-URI / MIME-type specs.
if _SAFE_DATA_URI_RE.match(url):
if PROFILE_IMAGE_MAX_DATA_URI_SIZE and len(url) > PROFILE_IMAGE_MAX_DATA_URI_SIZE:
raise ValueError(
f'Invalid profile image URL: data URI exceeds the {PROFILE_IMAGE_MAX_DATA_URI_SIZE}-byte limit.'
)
return url
raise ValueError(

View file

@ -3,7 +3,13 @@ import logging
import aiohttp
from open_webui.config import WEBUI_FAVICON_URL
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT, VERSION
from open_webui.env import (
AIOHTTP_CLIENT_ALLOW_REDIRECTS,
AIOHTTP_CLIENT_SESSION_SSL,
AIOHTTP_CLIENT_TIMEOUT,
VERSION,
)
from open_webui.retrieval.web.utils import validate_url
log = logging.getLogger(__name__)
@ -13,6 +19,10 @@ log = logging.getLogger(__name__)
async def post_webhook(name: str, url: str, message: str, event_data: dict) -> bool:
try:
log.debug(f'post_webhook: {url}, {message}, {event_data}')
# Block private-IP / loopback / cloud-metadata targets — the URL is
# caller-controlled (user notification settings under
# ENABLE_USER_WEBHOOKS, automation notification triggers).
validate_url(url)
payload = {}
# Slack and Google Chat Webhooks
@ -53,7 +63,12 @@ async def post_webhook(name: str, url: str, message: str, event_data: dict) -> b
async with aiohttp.ClientSession(
trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)
) as session:
async with session.post(url, json=payload, ssl=AIOHTTP_CLIENT_SESSION_SSL) as r:
async with session.post(
url,
json=payload,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS,
) as r:
r_text = await r.text()
r.raise_for_status()
log.debug(f'r.text: {r_text}')

Some files were not shown because too many files have changed in this diff Show more