mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-05 02:41:34 +00:00
Merge branch 'open-webui:dev' into dev
This commit is contained in:
commit
9876b60fa8
232 changed files with 99164 additions and 7463 deletions
11
.env.example
11
.env.example
|
|
@ -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'
|
||||
|
|
|
|||
22
.github/workflows/docker.yaml
vendored
22
.github/workflows/docker.yaml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
13
.github/workflows/release.yml
vendored
13
.github/workflows/release.yml
vendored
|
|
@ -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: |
|
||||
|
|
|
|||
134
CHANGELOG.md
134
CHANGELOG.md
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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'])
|
||||
|
|
|
|||
|
|
@ -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'])
|
||||
|
|
|
|||
|
|
@ -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'],
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
765
backend/open_webui/retrieval/vector/dbs/valkey.py
Normal file
765
backend/open_webui/retrieval/vector/dbs/valkey.py
Normal 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))
|
||||
|
|
@ -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}')
|
||||
|
||||
|
|
|
|||
|
|
@ -14,3 +14,4 @@ class VectorType(StrEnum):
|
|||
S3VECTOR = 's3vector'
|
||||
WEAVIATE = 'weaviate'
|
||||
OPENGAUSS = 'opengauss'
|
||||
VALKEY = 'valkey'
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
74
backend/open_webui/retrieval/web/linkup.py
Normal file
74
backend/open_webui/retrieval/web/linkup.py
Normal 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)}')
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {},
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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}',
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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'):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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'],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
|
|
|
|||
|
|
@ -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"]}'
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue