Compare commits

..

No commits in common. "main" and "v0.9.6" have entirely different histories.
main ... v0.9.6

872 changed files with 102244 additions and 235579 deletions

View file

@ -18,6 +18,3 @@ uploads
**/*.db **/*.db
_test _test
backend/data/* backend/data/*
.venv
.git

View file

@ -13,21 +13,6 @@ OPENAI_API_KEY=''
# CORS_ALLOW_ORIGIN='http://localhost:5173;http://localhost:8080' # CORS_ALLOW_ORIGIN='http://localhost:5173;http://localhost:8080'
CORS_ALLOW_ORIGIN='*' CORS_ALLOW_ORIGIN='*'
# Set to false to keep memory tools enabled without adding memory context to the system context.
ENABLE_MEMORY_SYSTEM_CONTEXT=true
# Set to true to add compact row/column stats to parsed CSV retrieval context.
ENABLE_RAG_CSV_SUMMARY=false
# Set to true to preserve backing file records, storage blobs, and per-file vectors when files are removed from knowledge bases.
ENABLE_KNOWLEDGE_FILE_RETENTION=false
# Comma-separated chunk metadata keys to expose to the model alongside retrieved content.
RAG_SOURCE_METADATA_KEYS=''
# Set to false to disable workspace Tools and Functions.
ENABLE_PLUGINS=true
# For production you should set this to match the proxy configuration (127.0.0.1) # For production you should set this to match the proxy configuration (127.0.0.1)
FORWARDED_ALLOW_IPS='*' FORWARDED_ALLOW_IPS='*'

2
.github/FUNDING.yml vendored
View file

@ -1 +1 @@
github: open-webui github: tjbck

View file

@ -1,35 +1,41 @@
name: Bug Report name: Bug Report
description: Tell us what broke in Open WebUI. description: Create a detailed bug report to help us improve Open WebUI.
title: 'issue: ' title: 'issue: '
labels: ['bug', 'triage'] labels: ['bug', 'triage']
assignees: []
body: body:
- type: markdown - type: markdown
attributes: attributes:
value: | value: |
# Bug Report # Bug Report
Use this for real, reproducible bugs. A clear issue is the most useful contribution: include the affected workflow, the expected result, the actual result, and the details needed for someone else to reproduce it. ## Important Notes
Before submitting, search open and closed [Issues](https://github.com/open-webui/open-webui/issues) and [Discussions](https://github.com/open-webui/open-webui/discussions). The issue may already be reported or fixed on `dev`. - **Before submitting a bug report**: Please check the [Issues](https://github.com/open-webui/open-webui/issues) and [Discussions](https://github.com/open-webui/open-webui/discussions) sections to see if a similar issue has already been reported. If unsure, start a discussion first, as this helps us efficiently focus on improving the project. Duplicates may be closed without notice. **Please search for existing issues AND discussions. No matter open or closed.**
**Test on the latest release AND on `dev`, right before you submit this report, not last week.** A huge share of reports are for bugs already fixed on `dev`, sometimes weeks earlier, because the reporter only tested an old version and never rechecked. Reports that don't reproduce on latest and on current `dev` at submission time will be closed without further discussion, no exceptions. - Check for opened, **but also for (recently) CLOSED issues** as the issue you are trying to report **might already have been fixed on the dev branch!**
Please do not open a code pull request for this report unless a maintainer asks for one, or the change is only i18n/localization. If you want to share code as reference, include it here as a local diff or patch. Actionable reproduction details are the most useful next step. - **Respectful collaboration**: Open WebUI is a volunteer-driven project with a single maintainer and contributors who also have full-time jobs. Please be constructive and respectful in your communication.
Security vulnerabilities must not be reported publicly. Use the [GitHub security page](https://github.com/open-webui/open-webui/security) instead. - **Contributing**: If you encounter an issue, consider submitting a pull request or forking the project. We prioritize preventing contributor burnout to maintain Open WebUI's quality.
- **Bug Reproducibility**: If a bug cannot be reproduced using a `:main` or `:dev` Docker setup or with `pip install` on Python 3.11, community assistance may be required. In such cases, we will move it to the "[Issues](https://github.com/open-webui/open-webui/discussions/categories/issues)" Discussions section. Your help is appreciated!
- **Scope**: If you want to report a SECURITY VULNERABILITY, then do so through our [GitHub security page](https://github.com/open-webui/open-webui/security).
- type: checkboxes - type: checkboxes
id: issue-check id: issue-check
attributes: attributes:
label: Before Submitting label: Check Existing Issues
description: Confirm that you’ve checked for existing reports before submitting a new one.
options: options:
- label: I searched open and closed issues and discussions for an existing report. - label: I have searched for any existing and/or related issues.
required: true required: true
- label: I reproduced this bug on the latest release AND on the current `dev` branch, right before submitting this report. I did not just check an old version or rely on a check from days ago. - label: I have searched for any existing and/or related discussions.
required: true required: true
- label: I understand that maintainers want a well-written issue before any code pull request. - label: I have also searched in the CLOSED issues AND CLOSED discussions and found no related items (your issue might already be addressed on the development branch!).
required: true required: true
- label: This is not a security vulnerability. - label: I am using the latest version of Open WebUI.
required: true required: true
- type: dropdown - type: dropdown
@ -38,9 +44,9 @@ body:
label: Installation Method label: Installation Method
description: How did you install Open WebUI? description: How did you install Open WebUI?
options: options:
- Docker
- Pip Install
- Git Clone - Git Clone
- Pip Install
- Docker
- Other - Other
validations: validations:
required: true required: true
@ -49,51 +55,67 @@ body:
id: open-webui-version id: open-webui-version
attributes: attributes:
label: Open WebUI Version label: Open WebUI Version
description: Specify the version, commit, or image tag. description: Specify the version (e.g., v0.6.26)
placeholder: v0.11.0, dev commit SHA, or Docker tag
validations: validations:
required: true required: true
- type: input
id: ollama-version
attributes:
label: Ollama Version (if applicable)
description: Specify the version (e.g., v0.2.0, or v0.1.32-rc1)
validations:
required: false
- type: input - type: input
id: operating-system id: operating-system
attributes: attributes:
label: Operating System label: Operating System
description: Specify the OS and version. description: Specify the OS (e.g., Windows 10, macOS Sonoma, Ubuntu 22.04, Debian 12)
placeholder: Windows 11, macOS Tahoe, Ubuntu 26.04, Debian 13
validations: validations:
required: true required: true
- type: input - type: input
id: browser id: browser
attributes: attributes:
label: Browser label: Browser (if applicable)
description: If the bug appears in the browser, include browser and version. description: Specify the browser/version (e.g., Chrome 100.0, Firefox 98.0)
placeholder: Chrome 151.0, Firefox 153.0.3
validations: validations:
required: false required: false
- type: input - type: checkboxes
id: ollama-version id: confirmation
attributes: attributes:
label: Ollama Version label: Confirmation
description: Include this if Ollama is involved. description: Ensure the following prerequisites have been met.
placeholder: v0.32.5 options:
validations: - label: I have read and followed all instructions in `README.md`.
required: false required: true
- label: I am using the latest version of **both** Open WebUI and Ollama.
- type: textarea required: true
id: summary - label: I have included the browser console logs.
attributes: required: true
label: Summary - label: I have included the Docker container logs.
description: What is wrong, in a few sentences? required: true
validations: - label: I have **provided every relevant configuration, setting, and environment variable used in my setup.**
required: true required: true
- label: I have clearly **listed every relevant configuration, custom setting, environment variable, and command-line option that influences my setup** (such as Docker Compose overrides, .env values, browser settings, authentication configurations, etc).
required: true
- label: |
I have documented **step-by-step reproduction instructions that are precise, sequential, and leave nothing to interpretation**. My steps:
- Start with the initial platform/version/OS and dependencies used,
- Specify exact install/launch/configure commands,
- List URLs visited, user input (incl. example values/emails/passwords if needed),
- Describe all options and toggles enabled or changed,
- Include any files or environmental changes,
- Identify the expected and actual result at each stage,
- Ensure any reasonably skilled user can follow and hit the same issue.
required: true
- type: textarea - type: textarea
id: expected-behavior id: expected-behavior
attributes: attributes:
label: Expected Behavior label: Expected Behavior
description: What should have happened? description: Describe what should have happened.
validations: validations:
required: true required: true
@ -101,7 +123,7 @@ body:
id: actual-behavior id: actual-behavior
attributes: attributes:
label: Actual Behavior label: Actual Behavior
description: What actually happened? description: Describe what actually happened.
validations: validations:
required: true required: true
@ -109,21 +131,32 @@ body:
id: reproduction-steps id: reproduction-steps
attributes: attributes:
label: Steps to Reproduce label: Steps to Reproduce
description: Include the exact commands, settings, URLs, model/provider setup, and user actions needed to hit the bug. description: |
Please provide a **very detailed, step-by-step guide** to reproduce the issue. Your instructions should be so clear and precise that anyone can follow them without guesswork. Include every relevant detail—settings, configuration options, exact commands used, values entered, and any prerequisites or environment variables.
**If full reproduction steps and all relevant settings are not provided, your issue may not be addressed.**
**If your steps to reproduction are incomplete, lacking detail or not reproducible, your issue can not be addressed.**
placeholder: | placeholder: |
1. Start Open WebUI with ... Example (include every detail):
2. Configure ... 1. Start with a clean Ubuntu 22.04 install.
3. Open ... 2. Install Docker v24.0.5 and start the service.
4. Click ... 3. Clone the Open WebUI repo (git clone ...).
5. See ... 4. Use the Docker Compose file without modifications.
5. Open browser Chrome 115.0 in incognito mode.
6. Go to http://localhost:8080 and log in with user "test@example.com".
7. Set the language to "English" and theme to "Dark".
8. Attempt to connect to Ollama at "http://localhost:11434".
9. Observe that the error message "Connection refused" appears at the top right.
Please list each step carefully and include all relevant configuration, settings, and options.
validations: validations:
required: true required: true
- type: textarea - type: textarea
id: logs-screenshots id: logs-screenshots
attributes: attributes:
label: Logs, Screenshots, and Config label: Logs & Screenshots
description: Include relevant browser console logs, server/container logs, screenshots, and configuration. If something does not apply, say so. description: Include relevant logs, errors, or screenshots to help diagnose the issue.
placeholder: 'Attach logs from the browser console, Docker logs, or error messages.'
validations: validations:
required: true required: true
@ -131,6 +164,13 @@ body:
id: additional-info id: additional-info
attributes: attributes:
label: Additional Information label: Additional Information
description: Anything else that might help us understand the report. description: Provide any extra details that may assist in understanding the issue.
validations: validations:
required: false required: false
- type: markdown
attributes:
value: |
## Note
**If the bug report is incomplete, does not follow instructions or is lacking details it may not be addressed.** Ensure that you've followed all the **README.md** and **troubleshooting.md** guidelines, and provide all necessary information for us to reproduce the issue.
Thank you for contributing to Open WebUI!

View file

@ -1,5 +1 @@
blank_issues_enabled: false blank_issues_enabled: false
contact_links:
- name: 🔒 Report a Security Vulnerability
url: https://github.com/open-webui/open-webui/security
about: Do NOT open a public issue for security vulnerabilities, suspected vulnerabilities, or any security-related concern. Please review our Security Policy and report privately via the "Report a vulnerability" button so it can be handled as a private advisory.

View file

@ -1,73 +1,76 @@
name: Feature Request name: Feature Request
description: Describe what you would like Open WebUI to support. description: Suggest a new feature or improvement
title: 'feat: ' title: 'feat: '
labels: ['triage'] labels: ['triage']
body: body:
- type: markdown - type: markdown
attributes: attributes:
value: | value: |
# Feature Request ## Before Submitting
Describe the requested behavior, the problem it solves, and any examples, mockups, screenshots, or workflows that clarify the request. A clear issue or discussion is the most useful contribution. Please check **open AND closed** [Issues](https://github.com/open-webui/open-webui/issues) and [Discussions](https://github.com/open-webui/open-webui/discussions) for similar requests. If you find one, add your input there instead.
Search open and closed [Issues](https://github.com/open-webui/open-webui/issues) and [Discussions](https://github.com/open-webui/open-webui/discussions) before submitting. If the request needs broad product, UX, architecture, compatibility, or maintenance discussion, please start in [Discussions](https://github.com/open-webui/open-webui/discussions) so the community can weigh in. ### Scope Guidelines
Please do not open a code pull request for this request unless a maintainer asks for one, or the change is only i18n/localization. If you want to share code as reference, include it here as a local diff or patch. Clear product context is the most useful next step. Feature requests that require significant implementation effort should be posted in the **Ideas** section of [Discussions](https://github.com/open-webui/open-webui/discussions) instead. We will move oversized feature requests to Discussions to keep the Issues tab focused on actionable items.
Security vulnerabilities must not be reported publicly. Use the [GitHub security page](https://github.com/open-webui/open-webui/security) instead. If your request might impact the broader community, please open a Discussion first so others can weigh in on the design.
### Be Respectful
Open WebUI is a volunteer-driven project maintained by a small team. We value constructive, positive communication. Please be mindful of maintainers' time and energy.
### Contributing
If you encounter an issue, we encourage you to submit a pull request or fork the project. We actively work to prevent contributor burnout and maintain project quality.
### Reproducibility
If a bug cannot be reproduced with a `:main` or `:dev` Docker setup, or a `pip install` with Python 3.11, it may be moved to the "Issues" section in Discussions for community assistance.
- type: checkboxes - type: checkboxes
id: existing-request id: existing-issue
attributes: attributes:
label: Before Submitting label: Check Existing Issues
description: Confirm you have searched for similar requests.
options: options:
- label: I searched open and closed issues and discussions for an existing request. - label: I have searched all existing **open AND closed** issues and discussions and found none comparable to my request.
required: true required: true
- label: I checked whether this already exists on the `dev` branch or latest source.
required: true - type: checkboxes
- label: I understand that maintainers want a well-written issue or discussion before any code pull request. id: feature-scope
required: true attributes:
- label: This request is not a security vulnerability. label: Verify Feature Scope
description: Confirm this request belongs in Issues rather than Discussions.
options:
- label: I believe this feature request is appropriately scoped for the Issues section as described above.
required: true required: true
- type: textarea - type: textarea
id: problem-description id: problem-description
attributes: attributes:
label: Problem label: Problem Description
description: What is missing, frustrating, confusing, or unnecessarily hard today? description: Is this related to a problem? Describe the pain point clearly.
placeholder: "I'm trying to..., but..." placeholder: "e.g., I'm frustrated when..."
validations: validations:
required: true required: true
- type: textarea - type: textarea
id: desired-behavior id: solution-description
attributes: attributes:
label: Desired Behavior label: Proposed Solution
description: What would you like to happen instead? description: Describe what you would like to happen.
placeholder: "I would like Open WebUI to..."
validations: validations:
required: true required: true
- type: textarea
id: why-it-matters
attributes:
label: Why This Matters
description: Who benefits, and what workflow does this unlock or improve?
validations:
required: true
- type: textarea
id: examples
attributes:
label: Examples or References
description: Add mockups, screenshots, links, prompts, workflows, or examples from other tools.
validations:
required: false
- type: textarea - type: textarea
id: alternatives-considered id: alternatives-considered
attributes: attributes:
label: Alternatives or Workarounds label: Alternatives Considered
description: What have you tried instead, if anything? description: Describe any alternative solutions or workarounds you have considered.
validations:
required: false - type: textarea
id: additional-context
attributes:
label: Additional Context
description: Add any other context, mockups, or screenshots about the feature request.

View file

@ -1,95 +1,104 @@
<!-- <!--
Important checks for contributors: ⚠️ CRITICAL CHECKS FOR CONTRIBUTORS (READ, DON'T DELETE) ⚠️
1. DO NOT OPEN A CODE PULL REQUEST unless a maintainer explicitly asked you to, or the change is strictly limited to i18n/localization. 1. Target the `dev` branch. PRs targeting `main` will be automatically closed.
2. Target the `dev` branch. PRs targeting `main` will be closed. 2. Do NOT delete the CLA section at the bottom. It is required for the bot to accept your PR.
3. Do not delete the Contributor License Agreement section at the bottom. The CLA bot requires it.
--> -->
# Pull Request # Pull Request Checklist
**Do not open a code pull request unless a maintainer has explicitly requested it or the change is limited to i18n/localization.** ### Note to first-time contributors: Please open a discussion post in [Discussions](https://github.com/open-webui/open-webui/discussions) to discuss your idea/fix with the community before creating a pull request, and describe your changes before submitting a pull request.
The most useful way to help is to give us a clear understanding of the problem: report reproducible bugs in [Issues](https://github.com/open-webui/open-webui/issues) and share proposals in [Discussions](https://github.com/open-webui/open-webui/discussions). We use that context to evaluate solutions and refine the implementation internally, accounting for the broader codebase and ongoing work. External implementations usually require substantial reworking to fit the project's standards, and coordinating those revisions usually takes more effort than developing the solution internally. Please follow this process before investing time in a pull request. PRs opened outside these guidelines are generally closed without review. This is to ensure large feature PRs are discussed with the community first, before starting work on it. If the community does not want this feature or it is not relevant for Open WebUI as a project, it can be identified in the discussion before working on the feature and submitting the PR.
## Maintainer Request <!--
### ⚠️ Important: Your PR is a contribution, not a guarantee of merge.
Link the maintainer's request for this PR, or state that the change is limited to i18n/localization. The most impactful way to contribute to Open WebUI is through well-written bug reports, detailed feature discussions, and thoughtful ideas. These directly shape the project. If you do open a pull request, please know that Open WebUI is held to the highest standard of code quality, consistency, and architectural coherence, and every line merged becomes something the core team must own, maintain, and support indefinitely. Submitted code may be refactored, rewritten, or used as inspiration for a different implementation. This is not a reflection of your work's quality. It is how we ensure that a small team can deeply understand and evolve every part of the codebase.
-->
## Checklist **Before submitting, make sure you've checked the following:**
- [ ] I have read and I understand the [contribution policy](https://docs.openwebui.com/contributing/#submit-code). - [ ] **Linked Issue/Discussion:** This PR references an existing [Issue](https://github.com/open-webui/open-webui/issues) or [Discussion](https://github.com/open-webui/open-webui/discussions) — `Closes #___` / `Relates to #___`. If one does not exist, create one first. PRs without a linked issue or discussion may be closed without review.
- [ ] This PR targets the `dev` branch. - [ ] **Target branch:** The pull request targets the `dev` branch. **PRs targeting `main` will be immediately closed.**
- [ ] This PR links to a well-described, confirmed Issue or active Discussion: `Closes #___` / `Relates to #___`. - [ ] **Description:** A concise description of the changes is provided below.
- [ ] A maintainer explicitly asked me to open this PR, or this PR only updates i18n/localization. - [ ] **Changelog:** A changelog entry following [Keep a Changelog](https://keepachangelog.com/) format is included at the bottom.
- [ ] The change is one logical unit with no unrelated commits. - [ ] **Documentation:** Relevant documentation has been added or updated in the [Open WebUI Docs Repository](https://github.com/open-webui/docs).
- [ ] I matched nearby code patterns and avoided unnecessary new settings, abstractions, or dependencies. - [ ] **Dependencies:** Any new or updated dependencies are explained, tested, and documented.
- [ ] I manually tested the changed workflow and any nearby behavior that could be affected. - [ ] **Testing:** Manual tests have been performed to verify the fix/feature works correctly and does not introduce regressions. Screenshots or recordings are included where applicable.
- [ ] I have not added or rewritten automated tests, fixtures, snapshots, or testing infrastructure unless a maintainer explicitly requested them. - [ ] **No Unchecked AI Code:** This PR is either human-written or has undergone thorough human review AND manual testing. Unreviewed AI-generated PRs may be closed immediately.
- [ ] I updated relevant docs, including the [Open WebUI Docs Repository](https://github.com/open-webui/docs), if needed. - [ ] **Self-Review:** A self-review of the code has been performed, ensuring adherence to project coding standards.
- [ ] I added screenshots for UI changes, and a recording when motion or interaction matters. - [ ] **Architecture:** Smart defaults are preferred over new settings. Local state is used for ephemeral UI logic. Major architectural or UX changes have been discussed first.
- [ ] I reviewed any AI-generated code before submitting it. - [ ] **Git Hygiene:** The PR is atomic (one logical change), rebased on `dev`, and contains no unrelated commits.
- [ ] The PR title uses one of the prefixes listed below. - [ ] **Title Prefix:** The PR title uses one of the following prefixes:
- **BREAKING CHANGE**: Changes affecting backward compatibility
- **build**: Build system or dependency changes
- **ci**: CI/CD workflow changes
- **chore**: Refactoring, cleanup, or non-functional changes
- **docs**: Documentation additions or updates
- **feat**: New features or enhancements
- **fix**: Bug fixes or corrections
- **i18n**: Internationalization or localization changes
- **perf**: Performance improvements
- **refactor**: Code restructuring
- **style**: Formatting changes (whitespace, semicolons, etc.)
- **test**: Test additions or corrections
- **WIP**: Work in progress
## Title Prefix # Changelog Entry
Use one of the following prefixes: ### Description
- **BREAKING CHANGE**: Changes affecting backward compatibility - [Describe the changes, including motivation and impact]
- **build**: Build system or dependency changes
- **ci**: CI/CD workflow changes
- **chore**: Refactoring, cleanup, or non-functional changes
- **docs**: Documentation additions or updates
- **feat**: New features or enhancements
- **fix**: Bug fixes or corrections
- **i18n**: Internationalization or localization changes
- **perf**: Performance improvements
- **refactor**: Code restructuring
## Summary
Describe the change, the problem it solves, and the impact on users.
## Verification
Describe how you reproduced the problem and manually checked the behavior before and after the change. Include exact steps, setup details, and relevant logs, screenshots, or recordings. Report results from relevant existing checks and anything you could not verify.
Do not add or rewrite automated tests unless a maintainer explicitly requests them. Tests that repeat an implementation's assumptions can pass while preserving the same mistake; maintainers determine the regression coverage needed. Do not remove, disable, or weaken existing tests to make the change pass.
## Changelog Entry
### Added ### Added
- - [New features, functionalities, or additions]
### Changed ### Changed
- - [Changes, updates, refactorings, or optimizations]
### Fixed ### Deprecated
- - [Deprecated functionality or features]
### Removed ### Removed
- - [Removed features, files, or functionalities]
### Fixed
- [Bug fixes or corrections]
### Security ### Security
- - [Security-related changes or vulnerability fixes]
### Breaking Changes ### Breaking Changes
- - **BREAKING CHANGE**: [Changes affecting compatibility or functionality]
## Additional Context ---
Add anything maintainers should know before review. ### Additional Information
## Contributor License Agreement - [Any additional context, notes, or references to related issues/commits]
### Screenshots or Videos
- [Attach relevant screenshots or videos demonstrating the changes]
### Contributor License Agreement
<!-- <!--
DO NOT DELETE THIS SECTION. 🚨 DO NOT DELETE THE TEXT BELOW 🚨
Your PR will not be reviewed or merged until you check the box below confirming that you have read and agree to the CLA. Keep the "Contributor License Agreement" confirmation text intact.
Deleting it will trigger the CLA-Bot to INVALIDATE your PR.
Your PR will NOT be reviewed or merged until you check the box below confirming that you have read and agree to the terms of the CLA.
--> -->
- [ ] By submitting this pull request, I confirm that I have read and fully agree to the [Contributor License Agreement (CLA)](https://github.com/open-webui/open-webui/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT), and I am providing my contributions under its terms. - [ ] By submitting this pull request, I confirm that I have read and fully agree to the [Contributor License Agreement (CLA)](https://github.com/open-webui/open-webui/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT), and I am providing my contributions under its terms.
> [!NOTE]
> Deleting the CLA section will lead to immediate closure of your PR and it will not be merged in.

View file

@ -7,10 +7,10 @@ name: Python CI
on: on:
push: push:
branches: [main, dev] branches: [main, dev]
paths: ['backend/**', 'pyproject.toml', 'uv.lock', '.github/workflows/backend.yaml'] paths: ['backend/**', 'pyproject.toml', 'uv.lock']
pull_request: pull_request:
branches: [main, dev] branches: [main, dev]
paths: ['backend/**', 'pyproject.toml', 'uv.lock', '.github/workflows/backend.yaml'] paths: ['backend/**', 'pyproject.toml', 'uv.lock']
concurrency: concurrency:
group: backend-${{ github.ref }} group: backend-${{ github.ref }}
@ -38,6 +38,3 @@ jobs:
- name: Verify formatting - name: Verify formatting
run: ruff format --check . --exclude .venv --exclude venv run: ruff format --check . --exclude .venv --exclude venv
- name: Detect logic errors
run: ruff check --select=F --ignore=F401,F403,F405,F541,F811,F841 --output-format=github .

View file

@ -75,17 +75,6 @@ jobs:
- name: Set up Docker Buildx - name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3 uses: docker/setup-buildx-action@v3
- name: Prepare CI Dockerfile
run: |
awk '
/^FROM --platform=\$BUILDPLATFORM node:/ {
print
print "ENV NODE_OPTIONS=\"--max-old-space-size=12288\""
next
}
{ print }
' Dockerfile > "${RUNNER_TEMP}/Dockerfile"
- name: Log in to the Container registry - name: Log in to the Container registry
uses: docker/login-action@v3 uses: docker/login-action@v3
with: with:
@ -126,7 +115,6 @@ jobs:
id: build id: build
with: with:
context: . context: .
file: ${{ runner.temp }}/Dockerfile
push: true push: true
platforms: ${{ matrix.platform.arch }} platforms: ${{ matrix.platform.arch }}
labels: ${{ steps.meta.outputs.labels }} labels: ${{ steps.meta.outputs.labels }}
@ -243,74 +231,11 @@ jobs:
run: | run: |
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ steps.meta.outputs.version }} docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ steps.meta.outputs.version }}
notify-helm-charts:
runs-on: ubuntu-latest
needs: [merge]
if: ${{ !cancelled() && needs.merge.result == 'success' && (github.ref == 'refs/heads/dev' || startsWith(github.ref, 'refs/tags/v')) }}
steps:
- name: Create Helm charts app token
id: helm-app-token
uses: actions/create-github-app-token@v2
with:
app-id: ${{ secrets.HELM_CHARTS_APP_ID }}
private-key: ${{ secrets.HELM_CHARTS_APP_PRIVATE_KEY }}
owner: ${{ github.repository_owner }}
repositories: helm-charts
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Verify published Open WebUI image
id: image
run: |
set -euo pipefail
image_name="ghcr.io/${GITHUB_REPOSITORY,,}"
ref_name="${GITHUB_REF_NAME}"
if [ "${GITHUB_REF}" = "refs/heads/dev" ]; then
image_tag="dev"
else
image_tag="${ref_name#v}"
fi
docker buildx imagetools inspect "${image_name}:${image_tag}"
echo "tag=${image_tag}" >> "${GITHUB_OUTPUT}"
- name: Dispatch Helm chart automation
uses: actions/github-script@v8
with:
github-token: ${{ steps.helm-app-token.outputs.token }}
script: |
const isDev = context.ref === 'refs/heads/dev';
const eventType = isDev
? 'open-webui-dev-image-published'
: 'open-webui-release-published';
const refName = context.ref.replace('refs/heads/', '').replace('refs/tags/', '');
const appVersion = refName.startsWith('v') ? refName.slice(1) : refName;
const payload = {
image_tag: isDev ? 'dev' : appVersion,
source_ref: context.ref,
source_sha: context.sha,
source_run_id: String(context.runId),
source_repository: context.repo.repo,
};
if (!isDev) {
payload.app_version = appVersion;
}
await github.rest.repos.createDispatchEvent({
owner: context.repo.owner,
repo: 'helm-charts',
event_type: eventType,
client_payload: payload,
});
copy-to-dockerhub: copy-to-dockerhub:
runs-on: ubuntu-latest runs-on: ubuntu-latest
if: ${{ !cancelled() && (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] needs: [merge]
continue-on-error: true
strategy: strategy:
fail-fast: false fail-fast: false
matrix: matrix:
@ -402,16 +327,5 @@ jobs:
[ -z "$TAG" ] && continue [ -z "$TAG" ] && continue
DEST="${DOCKERHUB_IMAGE}:${TAG}" DEST="${DOCKERHUB_IMAGE}:${TAG}"
echo " -> ${DEST}" echo " -> ${DEST}"
for ATTEMPT in 1 2 3; do docker buildx imagetools create -t "${DEST}" "${SOURCE}"
if docker buildx imagetools create -t "${DEST}" "${SOURCE}" && \
docker buildx imagetools inspect "${DEST}"; then
break
fi
if [ "${ATTEMPT}" = "3" ]; then
echo "Failed to copy ${DEST} after ${ATTEMPT} attempts"
exit 1
fi
echo "Copy attempt ${ATTEMPT} for ${DEST} failed, retrying in 15s..."
sleep 15
done
done <<< "${{ steps.tags.outputs.tags }}" done <<< "${{ steps.tags.outputs.tags }}"

View file

@ -43,8 +43,6 @@ jobs:
- name: Production build - name: Production build
run: npm run build run: npm run build
env:
NODE_OPTIONS: --max-old-space-size=8192
# ── Vitest unit tests ──────────────────────────────────────────────────── # ── Vitest unit tests ────────────────────────────────────────────────────
unit-tests: unit-tests:

View file

@ -1,139 +0,0 @@
name: Issue Labeler
on:
issues:
types: [opened, edited]
permissions:
issues: write
jobs:
label-bug-reports:
runs-on: ubuntu-latest
steps:
- name: Add "bug" label to unlabeled bug reports
uses: actions/github-script@v7
with:
script: |
const issue = context.payload.issue;
// Web-form submissions already carry the label from the issue template
if (issue.labels.some((label) => label.name === 'bug')) {
return;
}
const isEdit = context.payload.action === 'edited';
const titleWasEdited = Boolean(context.payload.changes?.title);
if (isEdit && !titleWasEdited) {
return;
}
const title = issue.title ?? '';
const body = issue.body ?? '';
// Freeform bug reports: "issue: ...", "bug: ...", "fix: ...", "[Bug] ...", "issue/UX: ..."
const bugLikeTitle = /^\s*(\[\s*(bug|issue|fix)\b[^\]]*\]|(bug|issue|fix)\s*[:/\-])/i.test(title);
// API/CLI-created issues that reproduce the bug report form structure.
// Only headings distinctive to the bug form (both are required fields there) —
// generic headings like "Expected Behavior" also appear in freeform feature requests.
const bugFormBody = /###\s*(Installation Method|Open WebUI Version)/i.test(body);
if (!bugLikeTitle && !bugFormBody) {
return;
}
if (isEdit) {
const events = await github.paginate(github.rest.issues.listEvents, {
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: issue.number,
per_page: 100
});
const bugLabelWasRemoved = events.some(
(event) => event.event === 'unlabeled' && event.label?.name === 'bug'
);
if (bugLabelWasRemoved) {
return;
}
}
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: issue.number,
labels: ['bug']
});
label-feature-requests:
runs-on: ubuntu-latest
steps:
- name: Add "enhancement" label to unlabeled feature requests
uses: actions/github-script@v7
with:
script: |
const issue = context.payload.issue;
if (issue.labels.some((label) => label.name === 'enhancement')) {
return;
}
// A human (or the bug form) already classified this as a bug;
// do not stack a second, contradictory classification on it.
if (issue.labels.some((label) => label.name === 'bug')) {
return;
}
const isEdit = context.payload.action === 'edited';
const titleWasEdited = Boolean(context.payload.changes?.title);
if (isEdit && !titleWasEdited) {
return;
}
const title = issue.title ?? '';
const body = issue.body ?? '';
// Feature requests: "feat: ...", "feature: ...", "feature request: ...",
// "enhancement: ...", "enh: ...", "[Feature Request] ..." — the feature
// request form titles every submission "feat: ", so form submissions are
// covered by the same pattern.
const featureLikeTitle =
/^\s*(\[\s*(feat|feature|enhancement|enh)\b[^\]]*\]|(feat|feature( request)?|enhancement|enh)\s*[:/\-])/i.test(
title
);
// API/CLI-created issues that reproduce the feature request form structure.
// Only headings distinctive to that form.
const featureFormBody = /###\s*(Proposed Solution|Alternatives Considered)/i.test(body);
if (!featureLikeTitle && !featureFormBody) {
return;
}
if (isEdit) {
const events = await github.paginate(github.rest.issues.listEvents, {
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: issue.number,
per_page: 100
});
const enhancementLabelWasRemoved = events.some(
(event) => event.event === 'unlabeled' && event.label?.name === 'enhancement'
);
if (enhancementLabelWasRemoved) {
return;
}
}
await github.rest.issues.addLabels({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: issue.number,
labels: ['enhancement']
});

View file

@ -1,45 +0,0 @@
# ─────────────────────────────────────────────────────────────────────────────
# Tests — run the open-webui/tests unit suite against a release candidate
# Release pull requests go from dev into main and are titled with the version
# ─────────────────────────────────────────────────────────────────────────────
name: Tests
on:
pull_request:
branches: [main, dev]
types: [opened, synchronize, reopened, edited]
# An edit must not cancel a running suite: the replacement run would skip it and still report green.
concurrency:
group: regression-${{ github.ref }}
cancel-in-progress: ${{ github.event.action != 'edited' }}
jobs:
# Into main only the release branch counts; into dev any version title does, so the
# suite can be exercised outside a release. Expressions have no regex, hence the literal prefixes.
regression:
name: Suite
if: >-
(github.event.pull_request.base.ref == 'dev' ||
github.event.pull_request.head.ref == 'dev') &&
(github.event.action != 'edited' || github.event.changes.title != null) &&
(startsWith(github.event.pull_request.title, '0.') ||
startsWith(github.event.pull_request.title, '1.'))
permissions:
contents: read
uses: open-webui/tests/.github/workflows/regression.yml@main
with:
open-webui-ref: ${{ github.event.pull_request.head.sha }}
# Single check to require in branch protection.
result:
name: Result
needs: [regression]
if: always()
runs-on: ubuntu-latest
timeout-minutes: 5
permissions: {}
steps:
- name: Fail unless the suite passed or was not required
if: needs.regression.result != 'success' && needs.regression.result != 'skipped'
run: exit 1

1
.gitignore vendored
View file

@ -310,4 +310,3 @@ dist
cypress/videos cypress/videos
cypress/screenshots cypress/screenshots
.vscode/settings.json .vscode/settings.json
.cptr

File diff suppressed because it is too large Load diff

View file

@ -20,7 +20,7 @@ Examples of behavior that contribute to a positive and professional community in
- **Respecting others.** Be considerate, listen actively, and engage with empathy toward others' viewpoints and experiences. - **Respecting others.** Be considerate, listen actively, and engage with empathy toward others' viewpoints and experiences.
- **Constructive feedback.** Provide actionable, thoughtful, and respectful feedback that helps improve the project and encourages collaboration. Avoid unproductive negativity or hypercriticism. - **Constructive feedback.** Provide actionable, thoughtful, and respectful feedback that helps improve the project and encourages collaboration. Avoid unproductive negativity or hypercriticism.
- **Recognizing volunteer contributions.** Appreciate that **contributors dedicate their free time and resources selflessly**. Approach them with gratitude and patience. - **Recognizing volunteer contributions.** Appreciate that contributors dedicate their free time and resources selflessly. Approach them with gratitude and patience.
- **Focusing on shared goals.** Collaborate in ways that prioritize the health, success, and sustainability of the community over individual agendas. - **Focusing on shared goals.** Collaborate in ways that prioritize the health, success, and sustainability of the community over individual agendas.
Examples of unacceptable behavior include: Examples of unacceptable behavior include:
@ -32,23 +32,11 @@ Examples of unacceptable behavior include:
- **Entitlement, demand, or aggression toward contributors.** Volunteers are under no obligation to provide immediate or personalized support. Rude or dismissive behavior will not be tolerated. - **Entitlement, demand, or aggression toward contributors.** Volunteers are under no obligation to provide immediate or personalized support. Rude or dismissive behavior will not be tolerated.
- **Unproductive or destructive behavior.** This includes venting frustration as hostility ("tantrums"), hypercriticism, attention-seeking negativity, or anything that distracts from the project's goals. - **Unproductive or destructive behavior.** This includes venting frustration as hostility ("tantrums"), hypercriticism, attention-seeking negativity, or anything that distracts from the project's goals.
- **Spamming and promotional exploitation.** Sharing irrelevant product promotions or self-promotion in the community is not allowed unless it directly contributes value to the discussion. - **Spamming and promotional exploitation.** Sharing irrelevant product promotions or self-promotion in the community is not allowed unless it directly contributes value to the discussion.
- Posting low-effort, hard to read, essay-length AI generated comments or other forms of low-quality, hard to parse content that puts the burden of understanding on the reader.
### How We Develop the Project
Development is led by the maintainers, and code pull requests are reserved for work we explicitly request or exceptional contributions we choose to consider at our discretion. We use actionable reports and concrete use cases to understand problems, then evaluate, revise, and implement the appropriate approach internally. We assess each change against the project's architecture, existing behavior, quality standards, and future direction before settling on an implementation. Resolving a reported problem requires that broader context, and a working external patch usually requires substantial rewriting to meet the project's standards. Reviewing the patch, explaining the required changes, and coordinating successive revisions usually takes more effort than developing the solution internally. Fragmented commit histories, branches that have not been rebased, unresolved conflicts, and lengthy or unverified AI-generated comments add cleanup and discussion that delay the underlying work. Maintainers remain responsible for testing, documenting, supporting, and maintaining every accepted change, so we choose the approach based on the whole product and its ongoing maintenance. Clear reports, reproduction details, and relevant context give us what we need to make those decisions and develop the solution. A polished implementation, clean commit history, or completed checklist does not establish an exception to this process, and opening an issue or discussion is not an invitation to submit a PR. Wait for an explicit maintainer request before investing in a PR; unsolicited submissions are generally closed without review, and requested PRs remain subject to maintainer judgment.
### Feedback and Community Engagement ### Feedback and Community Engagement
Participation should help maintainers understand a concrete problem while respecting the project's priorities and available capacity. Please follow the [issue templates](.github/ISSUE_TEMPLATE) and [pull request policy](.github/pull_request_template.md) before submitting anything. - **Constructive feedback is encouraged, but hostile or entitled behavior will result in immediate action.** If you disagree with elements of the project, we encourage you to offer meaningful improvements or fork the project if necessary. Healthy discussions and technical disagreements are welcome only when handled with professionalism.
- **Respect contributors' time and efforts.** No one is entitled to personalized or on-demand assistance. This is a community built on collaboration and shared effort; demanding or demeaning behavior undermines that trust and will not be allowed.
- **Make reports actionable.** Search existing issues and discussions, check the latest version and whether the problem is already addressed on `dev`, and use the appropriate template. Bug reports should describe a reproducible problem, the affected workflow, expected and actual behavior, and relevant evidence. Feature requests should explain the user-facing need; broader product, UX, architecture, or maintenance questions belong in Discussions. Report security concerns privately through the [security reporting process](https://github.com/open-webui/open-webui/security).
- **Share the problem before investing in code.** Start with an actionable issue or discussion and leave implementation planning to the maintainers. An issue or discussion alone is not an invitation to submit a PR. Please wait for an explicit request before opening one; any exception is at the maintainers' discretion. Implementation notes, local diffs, or patches may be shared as reference in the relevant issue or discussion.
- **Respect maintainers' discretion.** Submitting an issue, proposal, or pull request does not create an obligation to respond, review, implement, or merge it. Maintainers set the project's direction and defer or close submissions based on scope, quality, maintenance cost, or available capacity. Unsolicited pull requests are generally closed without review.
- **Keep discussion focused and concise.** Provide new information when it helps evaluate the problem. Repeated bumps, duplicate submissions, unsolicited direct messages seeking attention, or pressure for timelines place an unnecessary burden on contributors.
- **Respect decisions and boundaries.** Technical disagreement is welcome when expressed professionally. Reopening a declined request or continuing to press for a different outcome without new, relevant information is not constructive. You are free to explore a different direction in your own fork.
Participants are expected to respect maintainers' decisions and the contribution process. Harassment, hostility, or repeated disregard for these boundaries will result in enforcement under this Code of Conduct.
### Zero Tolerance: No Warnings, Immediate Action ### Zero Tolerance: No Warnings, Immediate Action

View file

@ -26,9 +26,6 @@ ARG GID=0
######## WebUI frontend ######## ######## WebUI frontend ########
FROM --platform=$BUILDPLATFORM node:22-alpine3.20 AS build FROM --platform=$BUILDPLATFORM node:22-alpine3.20 AS build
ARG BUILD_HASH ARG BUILD_HASH
ARG USE_SLIM
ARG UID
ARG GID
# Set Node.js options (heap limit Allocation failed - JavaScript heap out of memory) # Set Node.js options (heap limit Allocation failed - JavaScript heap out of memory)
# ENV NODE_OPTIONS="--max-old-space-size=4096" # ENV NODE_OPTIONS="--max-old-space-size=4096"
@ -43,14 +40,7 @@ RUN npm ci --force
COPY . . COPY . .
ENV APP_BUILD_HASH=${BUILD_HASH} ENV APP_BUILD_HASH=${BUILD_HASH}
RUN npm run build && \ RUN npm run build
if [ "$USE_SLIM" = "true" ]; then find build -type f -name '*.map' -delete; fi
# Prepare backend ownership before the final copy so static assets occupy one layer.
# Group 0 write access lets arbitrary OpenShift UIDs update these assets at startup.
RUN chown -R $UID:$GID /app/backend && \
chgrp -R 0 /app/backend/open_webui/static && \
chmod -R g=u /app/backend/open_webui/static
######## WebUI backend ######## ######## WebUI backend ########
FROM python:3.11-slim-bookworm AS base FROM python:3.11-slim-bookworm AS base
@ -133,33 +123,24 @@ RUN echo -n 00000000-0000-0000-0000-000000000000 > $HOME/.cache/chroma/telemetry
# Make sure the user has access to the app and root directory # Make sure the user has access to the app and root directory
RUN chown -R $UID:$GID /app $HOME RUN chown -R $UID:$GID /app $HOME
# Slim cannot bundle a local model server or GPU runtime. # Install common system dependencies
RUN if [ "$USE_SLIM" = "true" ] && { [ "$USE_CUDA" = "true" ] || [ "$USE_OLLAMA" = "true" ]; }; then \
echo "USE_SLIM cannot be combined with USE_CUDA or USE_OLLAMA" >&2; exit 1; fi
# Keep the slim runtime free of local document/audio processing tools.
# Git-based tool requirements require the standard image.
RUN apt-get update && \ RUN apt-get update && \
apt-get install -y --no-install-recommends \ apt-get install -y --no-install-recommends \
curl jq ca-certificates \ git build-essential pandoc gcc netcat-openbsd curl jq \
&& if [ "$USE_SLIM" != "true" ]; then \ libmariadb-dev \
apt-get install -y --no-install-recommends \ python3-dev \
git build-essential pandoc gcc libmariadb-dev ffmpeg libsm6 libxext6; \ ffmpeg libsm6 libxext6 zstd \
fi && if [ "$USE_OLLAMA" = "true" ]; then \ && rm -rf /var/lib/apt/lists/*
apt-get install -y --no-install-recommends zstd; \
fi && rm -rf /var/lib/apt/lists/*
# install python dependencies # install python dependencies
COPY --chown=$UID:$GID ./backend/requirements*.txt ./ COPY --chown=$UID:$GID ./backend/requirements.txt ./requirements.txt
# Set UV_LINK_MODE to copy to prevent 0-byte file corruption in QEMU arm64 cross-builds # Set UV_LINK_MODE to copy to prevent 0-byte file corruption in QEMU arm64 cross-builds
ENV UV_LINK_MODE=copy ENV UV_LINK_MODE=copy
RUN --mount=from=ghcr.io/astral-sh/uv:0.12.10,source=/uv,target=/bin/uv \ RUN set -e; \
set -e; \ pip3 install --no-cache-dir uv; \
if [ "$USE_SLIM" = "true" ]; then \ if [ "$USE_CUDA" = "true" ]; then \
uv pip install --system -r requirements-slim.txt --no-cache-dir; \
elif [ "$USE_CUDA" = "true" ]; then \
# If you use CUDA the whisper and embedding model will be downloaded on first use # If you use CUDA the whisper and embedding model will be downloaded on first use
# fix: pin torch<=2.9.1 - torch 2.10.0 aarch64 wheels cause SIGILL on ARM devices (RPi 4 Cortex-A72) #21349 # fix: pin torch<=2.9.1 - torch 2.10.0 aarch64 wheels cause SIGILL on ARM devices (RPi 4 Cortex-A72) #21349
pip3 install 'torch<=2.9.1' torchvision torchaudio --index-url https://download.pytorch.org/whl/$USE_CUDA_DOCKER_VER --no-cache-dir; \ pip3 install 'torch<=2.9.1' torchvision torchaudio --index-url https://download.pytorch.org/whl/$USE_CUDA_DOCKER_VER --no-cache-dir; \
@ -168,6 +149,7 @@ RUN --mount=from=ghcr.io/astral-sh/uv:0.12.10,source=/uv,target=/bin/uv \
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')"; \ python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')"; \
python -c "import os; from faster_whisper import WhisperModel; WhisperModel(os.environ['WHISPER_MODEL'], device='cpu', compute_type='int8', download_root=os.environ['WHISPER_MODEL_DIR'])"; \ python -c "import os; from faster_whisper import WhisperModel; WhisperModel(os.environ['WHISPER_MODEL'], device='cpu', compute_type='int8', download_root=os.environ['WHISPER_MODEL_DIR'])"; \
python -c "import os; import tiktoken; tiktoken.get_encoding(os.environ['TIKTOKEN_ENCODING_NAME'])"; \ python -c "import os; import tiktoken; tiktoken.get_encoding(os.environ['TIKTOKEN_ENCODING_NAME'])"; \
python -c "import nltk; nltk.download('punkt_tab')"; \
else \ else \
pip3 install 'torch<=2.9.1' torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu --no-cache-dir; \ pip3 install 'torch<=2.9.1' torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu --no-cache-dir; \
uv pip install --system -r requirements.txt --no-cache-dir; \ uv pip install --system -r requirements.txt --no-cache-dir; \
@ -176,17 +158,12 @@ RUN --mount=from=ghcr.io/astral-sh/uv:0.12.10,source=/uv,target=/bin/uv \
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')"; \ python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')"; \
python -c "import os; from faster_whisper import WhisperModel; WhisperModel(os.environ['WHISPER_MODEL'], device='cpu', compute_type='int8', download_root=os.environ['WHISPER_MODEL_DIR'])"; \ python -c "import os; from faster_whisper import WhisperModel; WhisperModel(os.environ['WHISPER_MODEL'], device='cpu', compute_type='int8', download_root=os.environ['WHISPER_MODEL_DIR'])"; \
python -c "import os; import tiktoken; tiktoken.get_encoding(os.environ['TIKTOKEN_ENCODING_NAME'])"; \ python -c "import os; import tiktoken; tiktoken.get_encoding(os.environ['TIKTOKEN_ENCODING_NAME'])"; \
python -c "import nltk; nltk.download('punkt_tab')"; \
fi; \ fi; \
fi; \ fi; \
mkdir -p /app/backend/data; chown -R $UID:$GID /app/backend/data/; \ mkdir -p /app/backend/data; chown -R $UID:$GID /app/backend/data/; \
if [ -d /app/backend/data/cache ]; then chmod -R a+rX /app/backend/data/cache; fi; \
rm -rf /var/lib/apt/lists/*; rm -rf /var/lib/apt/lists/*;
# Optional: PPTX parsing through unstructured may need spaCy's English model.
# Keep this out of the default image to avoid the extra image bloat; deployments
# with read-only site-packages can uncomment it and bake the model in.
# RUN python -m spacy download en_core_web_sm
# Install Ollama if requested # Install Ollama if requested
RUN if [ "$USE_OLLAMA" = "true" ]; then \ RUN if [ "$USE_OLLAMA" = "true" ]; then \
date +%s > /tmp/ollama_build_hash && \ date +%s > /tmp/ollama_build_hash && \
@ -204,8 +181,8 @@ COPY --chown=$UID:$GID --from=build /app/build /app/build
COPY --chown=$UID:$GID --from=build /app/CHANGELOG.md /app/CHANGELOG.md COPY --chown=$UID:$GID --from=build /app/CHANGELOG.md /app/CHANGELOG.md
COPY --chown=$UID:$GID --from=build /app/package.json /app/package.json COPY --chown=$UID:$GID --from=build /app/package.json /app/package.json
# copy backend files with the ownership and static permissions prepared above # copy backend files
COPY --from=build /app/backend . COPY --chown=$UID:$GID ./backend .
EXPOSE 8080 EXPOSE 8080

View file

@ -8,9 +8,11 @@
![GitHub top language](https://img.shields.io/github/languages/top/open-webui/open-webui) ![GitHub top language](https://img.shields.io/github/languages/top/open-webui/open-webui)
![GitHub last commit](https://img.shields.io/github/last-commit/open-webui/open-webui?color=red) ![GitHub last commit](https://img.shields.io/github/last-commit/open-webui/open-webui?color=red)
[![Discord](https://img.shields.io/badge/Discord-Open_WebUI-blue?logo=discord&logoColor=white)](https://discord.gg/5rJgQTnV4s) [![Discord](https://img.shields.io/badge/Discord-Open_WebUI-blue?logo=discord&logoColor=white)](https://discord.gg/5rJgQTnV4s)
[![](https://img.shields.io/static/v1?label=Sponsor&message=%E2%9D%A4&logo=GitHub&color=%23fe8e86)](https://github.com/sponsors/open-webui) [![](https://img.shields.io/static/v1?label=Sponsor&message=%E2%9D%A4&logo=GitHub&color=%23fe8e86)](https://github.com/sponsors/tjbck)
Open WebUI is **a home for AI**, a self-hosted AI platform that's **[extensible](https://docs.openwebui.com/features/extensibility/plugin/)**, **[feature-rich](https://docs.openwebui.com/features/)**, user-friendly, and built to run **[entirely offline](https://openwebui.com/sovereign-ai)**. With support for **Ollama** and **OpenAI-compatible APIs**, it gives you a powerful, provider-agnostic interface for both local and cloud-based models. ![Open WebUI Banner](./banner.png)
**Open WebUI is an [extensible](https://docs.openwebui.com/features/extensibility/plugin), feature-rich, and user-friendly self-hosted AI platform designed to operate entirely offline.** It supports various LLM runners like **Ollama** and **OpenAI-compatible APIs**, with **built-in inference engine** for RAG, making it a **powerful AI deployment solution**.
Passionate about open-source AI? [Join our team →](https://careers.openwebui.com/) Passionate about open-source AI? [Join our team →](https://careers.openwebui.com/)
@ -18,89 +20,65 @@ Passionate about open-source AI? [Join our team →](https://careers.openwebui.c
> [!TIP] > [!TIP]
> **Looking for an [Enterprise Plan](https://docs.openwebui.com/enterprise)?** – **[Speak with Our Sales Team Today!](https://docs.openwebui.com/enterprise)** > **Looking for an [Enterprise Plan](https://docs.openwebui.com/enterprise)?** – **[Speak with Our Sales Team Today!](https://docs.openwebui.com/enterprise)**
>
> Get **enhanced capabilities**, including **custom theming and branding**, **Service Level Agreement (SLA) support**, **Long-Term Support (LTS) versions**, and **more!**
For more information, be sure to check out our [Open WebUI Documentation](https://docs.openwebui.com/). For more information, be sure to check out our [Open WebUI Documentation](https://docs.openwebui.com/).
## Key Features of Open WebUI ⭐ ## Key Features of Open WebUI ⭐
- 🚀 **Effortless Setup**: Install seamlessly via pip, uv, Docker, or Kubernetes (kubectl, kustomize, or helm), with `:ollama` and `:cuda` tagged images available for container deployments. - 🚀 **Effortless Setup**: Install seamlessly using Docker or Kubernetes (kubectl, kustomize or helm) for a hassle-free experience with support for both `:ollama` and `:cuda` tagged images.
- 🤝 **Broad Model & API Integration**: Connect any OpenAI-compatible API alongside local Ollama models. Point the API URL at **LMStudio, GroqCloud, Mistral, OpenRouter, vLLM, and more** to mix and match providers freely. - 🤝 **Ollama/OpenAI API Integration**: Effortlessly integrate OpenAI-compatible APIs for versatile conversations alongside Ollama models. Customize the OpenAI API URL to link with **LMStudio, GroqCloud, Mistral, OpenRouter, and more**.
- 🔐 **Granular RBAC & User Groups**: Administrators define detailed roles, groups, and permissions, giving each user exactly the access they need. Secure by default, with tailored experiences per group. - 🛡️ **Granular Permissions and User Groups**: By allowing administrators to create detailed user roles and permissions, we ensure a secure user environment. This granularity not only enhances security but also allows for customized user experiences, fostering a sense of ownership and responsibility amongst users.
- 🧩 **Plugin Support**: Extend Open WebUI with **Filters**, **Actions**, **Pipes**, **Tools**, and **Skills**. Connect external services through **MCP**, **MCPO**, and **OpenAPI tool servers**. Build custom integrations, rate limits, approval flows, data connections, and more. - 📱 **Responsive Design**: Enjoy a seamless experience across Desktop PC, Laptop, and Mobile devices.
- 🤖 **Models & Agents**: Wrap any base model with custom instructions, tools, and knowledge to build specialized agents. Supports dynamic variables, per-user/group access control, and community preset imports via [Open WebUI Community](https://openwebui.com/). - 📱 **Progressive Web App (PWA) for Mobile**: Enjoy a native app-like experience on your mobile device with our PWA, providing offline access on localhost and a seamless user interface.
- ⚡ **Agentic Execution with [Open Terminal](https://github.com/open-webui/open-terminal)**: Give your agents a terminal and filesystem to carry out multi-step tasks. Let them analyze data, run scripts, fix errors, and produce files directly in chat. Scale to teams with **[Terminals (Enterprise)](https://github.com/open-webui/terminals)** for per-user isolated environments, resource limits, and automatic lifecycle management. - ✒️🔢 **Full Markdown and LaTeX Support**: Elevate your LLM experience with comprehensive Markdown and LaTeX capabilities for enriched interaction.
- 📝 **Notes**: A dedicated workspace for content outside conversations. Draft with a rich editor, use AI to rewrite selected text, and attach notes to any chat for full-context injection. - 🎤📹 **Hands-Free Voice/Video Call**: Experience seamless communication with integrated hands-free voice and video call features using multiple Speech-to-Text providers (Local Whisper, OpenAI, Deepgram, Azure) and Text-to-Speech engines (Azure, ElevenLabs, OpenAI, Transformers, WebAPI), allowing for dynamic and interactive chat environments.
- 📢 **Channels**: Real-time shared spaces where your team and AI models collaborate in one timeline. Tag models to draft or critique, with threads, reactions, pins, and access control. - 🛠️ **Model Builder**: Easily create Ollama models via the Web UI. Create and add custom characters/agents, customize chat elements, and import models effortlessly through [Open WebUI Community](https://openwebui.com/) integration.
- 🧠 **Persistent Memory**: The AI remembers facts about you across conversations, carrying context from one chat to the next. - 🐍 **Native Python Function Calling Tool**: Enhance your LLMs with built-in code editor support in the tools workspace. Bring Your Own Function (BYOF) by simply adding your pure Python functions, enabling seamless integration with LLMs.
- ✅ **Live Workflow & Message Flow**: Watch the AI build and work through checklists in real time. Queue messages while the AI is still responding; they send automatically when it's ready. - 💾 **Persistent Artifact Storage**: Built-in key-value storage API for artifacts, enabling features like journals, trackers, leaderboards, and collaborative tools with both personal and shared data scopes across sessions.
- 📅 **Calendar & AI Scheduling**: Built-in personal and shared calendars with month/week/day views, recurring events, color coding, attendees, and reminders. Models manage your schedule conversationally through native function calling. - 📚 **Local RAG Integration**: Dive into the future of chat interactions with groundbreaking Retrieval Augmented Generation (RAG) support using your choice of 9 vector databases and multiple content extraction engines (Tika, Docling, Document Intelligence, Mistral OCR, PaddleOCR-vl, External loaders). Load documents directly into chat or add files to your document library, effortlessly accessing them using the `#` command before a query.
- ⏱️ **Automations**: Schedule prompts to run on recurring schedules, with runs surfaced on your calendar and each completed run linking back to the chat it produced. - 🔍 **Web Search for RAG**: Perform web searches using 15+ providers including `SearXNG`, `Google PSE`, `Brave Search`, `Kagi`, `Mojeek`, `Tavily`, `Perplexity`, `serpstack`, `serper`, `Serply`, `DuckDuckGo`, `SearchApi`, `SerpApi`, `Bing`, `Jina`, `Exa`, `Sougou`, `Azure AI Search`, and `Ollama Cloud`, injecting results directly into your chat experience.
- 📱 **Responsive Design & PWA**: Seamless experience across desktop, laptop, and mobile, with a Progressive Web App for native app-like feel and offline access on localhost. - 🌐 **Web Browsing Capability**: Seamlessly integrate websites into your chat experience using the `#` command followed by a URL. This feature allows you to incorporate web content directly into your conversations, enhancing the richness and depth of your interactions.
- ✒️🔢 **Full Markdown and LaTeX Support**: Comprehensive Markdown and LaTeX capabilities for enriched interaction. - 🎨 **Image Generation & Editing Integration**: Create and edit images using multiple engines including OpenAI's DALL-E, Gemini, ComfyUI (local), and AUTOMATIC1111 (local), with support for both generation and prompt-based editing workflows.
- 🎤📹 **Hands-Free Voice/Video Call**: Integrated voice and video calls with multiple Speech-to-Text providers (Local Whisper, OpenAI, Deepgram, Azure) and Text-to-Speech engines (Azure, ElevenLabs, OpenAI, Transformers, WebAPI). - ⚙️ **Many Models Conversations**: Effortlessly engage with various models simultaneously, harnessing their unique strengths for optimal responses. Enhance your experience by leveraging a diverse set of models in parallel.
- 💾 **Persistent Artifact Storage**: Built-in key-value storage API for artifacts, enabling journals, trackers, leaderboards, and collaborative tools with personal and shared data scopes. - 🔐 **Role-Based Access Control (RBAC)**: Ensure secure access with restricted permissions; only authorized individuals can access your Ollama, and exclusive model creation/pulling rights are reserved for administrators.
- 📚 **Local RAG Integration**: Retrieval Augmented Generation backed by 9 vector databases and multiple content-extraction engines (Tika, Docling, Document Intelligence, Mistral OCR, PaddleOCR-vl, external loaders). Supports hybrid search (BM25 + vector) with reranking and full-context mode. Load documents into chat or pull them from your library with the `#` command. - 🗄️ **Flexible Database & Storage Options**: Choose from SQLite (with optional encryption), PostgreSQL, or configure cloud storage backends (S3, Google Cloud Storage, Azure Blob Storage) for scalable deployments.
- 🔍 **Web Search for RAG**: Search the web through dozens of providers including `SearXNG`, `Google PSE`, `Brave Search`, `Kagi`, `Mojeek`, `Tavily`, `Perplexity`, `Firecrawl`, `serpstack`, `serper`, `Serply`, `DuckDuckGo`, `SearchApi`, `SerpApi`, `Bing`, `Jina`, `Exa`, `Sougou`, `Azure AI Search`, and `Ollama Cloud`, injecting results directly into the conversation. - 🔍 **Advanced Vector Database Support**: Select from 9 vector database options including ChromaDB, PGVector, Qdrant, Milvus, Elasticsearch, OpenSearch, Pinecone, S3Vector, and Oracle 23ai for optimal RAG performance.
- 🌐 **Web Browsing Capability**: Pull websites into chat with the `#` command followed by a URL, or let the model fetch them on its own when needed. - 🔐 **Enterprise Authentication**: Full support for LDAP/Active Directory integration, SCIM 2.0 automated provisioning, and SSO via trusted headers alongside OAuth providers. Enterprise-grade user and group provisioning through SCIM 2.0 protocol, enabling seamless integration with identity providers like Okta, Azure AD, and Google Workspace for automated user lifecycle management.
- 🎨 **Image Generation & Editing**: Create and edit images with multiple engines including OpenAI DALL·E, Gemini, ComfyUI (local), and AUTOMATIC1111 (local), supporting both generation and prompt-based editing. - ☁️ **Cloud-Native Integration**: Native support for Google Drive and OneDrive/SharePoint file picking, enabling seamless document import from enterprise cloud storage.
- ⚙️ **Multi-Model Conversations**: Engage several models at once, harnessing their individual strengths in parallel for the best possible responses. - 📊 **Production Observability**: Built-in OpenTelemetry support for traces, metrics, and logs, enabling comprehensive monitoring with your existing observability stack.
- 📊 **Usage Analytics & Model Evaluation**: Admin dashboards track message volume, token consumption, and cost across users and models. Evaluate models with a built-in arena, A/B testing, and ELO-based leaderboards. - ⚖️ **Horizontal Scalability**: Redis-backed session management and WebSocket support for multi-worker and multi-node deployments behind load balancers.
- 🗄️ **Flexible Database & Storage**: Choose SQLite (with optional encryption) or PostgreSQL, and store files locally or on S3, Google Cloud Storage, or Azure Blob Storage. - 🌐🌍 **Multilingual Support**: Experience Open WebUI in your preferred language with our internationalization (i18n) support. Join us in expanding our supported languages! We're actively seeking contributors!
- 🧬 **Advanced Vector Database Support**: Pick from 9 vector databases: ChromaDB, PGVector, Qdrant, Milvus, Elasticsearch, OpenSearch, Pinecone, S3Vector, and Oracle 23ai. - 🧩 **Pipelines, Open WebUI Plugin Support**: Seamlessly integrate custom logic and Python libraries into Open WebUI using [Pipelines Plugin Framework](https://github.com/open-webui/pipelines). Launch your Pipelines instance, set the OpenAI URL to the Pipelines URL, and explore endless possibilities. [Examples](https://github.com/open-webui/pipelines/tree/main/examples) include **Function Calling**, User **Rate Limiting** to control access, **Usage Monitoring** with tools like Langfuse, **Live Translation with LibreTranslate** for multilingual support, **Toxic Message Filtering** and much more.
- 🪪 **Enterprise Authentication & Provisioning**: Full LDAP/Active Directory integration, SSO via trusted headers and OAuth providers, and SCIM 2.0 automated provisioning for identity providers like Okta, Azure AD, and Google Workspace. - 🌟 **Continuous Updates**: We are committed to improving Open WebUI with regular updates, fixes, and new features.
- ☁️ **Cloud-Native File Integration**: Native Google Drive and OneDrive/SharePoint file picking for seamless document import from enterprise cloud storage.
- 🔭 **Production Observability**: Built-in OpenTelemetry support for traces, metrics, and logs, plugging into your existing monitoring stack.
- ⚖️ **Horizontal Scalability**: Redis-backed session management and WebSocket support for multi-worker, multi-node deployments behind load balancers.
- 🌐🌍 **Multilingual Support**: Use Open WebUI in your preferred language with i18n support. We're actively seeking contributors to expand language coverage!
- 🌟 **Continuous Updates**: We're committed to improving Open WebUI with regular updates, fixes, and new features.
- 🛡️ **Transparent Security Process**: Security reports are triaged, fixed, and published as open advisories through a documented responsible-disclosure process. See our [Security Policy](https://github.com/open-webui/open-webui/security).
Want to learn more about Open WebUI's features? Check out our [Open WebUI documentation](https://docs.openwebui.com/features) for a comprehensive overview! Want to learn more about Open WebUI's features? Check out our [Open WebUI documentation](https://docs.openwebui.com/features) for a comprehensive overview!
## The Open WebUI Ecosystem 🌐
Open WebUI is the core, surrounded by companion apps and infrastructure that extend what your AI can do, where it can reach, and how you run it:
- 💻 **Open WebUI Computer** ([open-webui/computer](https://github.com/open-webui/computer)): A standalone, mobile-first computer and coding agent that runs on the machine you own. Files, terminal, and git in a browser tab, reachable from your phone. Connect it into Open WebUI as a model, or reach it from Telegram, WhatsApp, and more.
- ⚡ **Open Terminal** and **Terminals (Enterprise)** ([open-webui/open-terminal](https://github.com/open-webui/open-terminal) & [open-webui/terminals](https://github.com/open-webui/terminals)): A self-hosted computing environment that plugs into Open WebUI, giving the AI a place to write code, run it, read output, fix errors, and iterate inside the chat. Terminals gives you per-user isolated containers with separate credentials, resource limits, and network rules. Automatic lifecycle management on Docker or Kubernetes.
- 🔄 **oikb** ([open-webui/oikb](https://github.com/open-webui/oikb)): Feed your Knowledge Bases from 45+ sources (GitHub, Confluence, ServiceNow, Salesforce, Jira, Slack, SharePoint, Notion, and more), keeping the tools your team already uses continuously in sync.
- 🖥️ **Native Desktop App** ([open-webui/desktop](https://github.com/open-webui/desktop)): Run Open WebUI as a native app on macOS, Windows, and Linux. System-wide Spotlight chat bar with screenshot capture, push-to-talk voice, and optional fully-local inference via a built-in llama.cpp engine.
Want to learn more? Check out our [Open WebUI documentation](https://docs.openwebui.com) for more details!
--- ---
We are incredibly grateful for the generous support of our sponsors. Their contributions help us to maintain and improve our project, ensuring we can continue to deliver quality work to our community. Thank you! We are incredibly grateful for the generous support of our sponsors. Their contributions help us to maintain and improve our project, ensuring we can continue to deliver quality work to our community. Thank you!
@ -244,10 +222,6 @@ This project contains code under multiple licenses. The current codebase include
If you have any questions, suggestions, or need assistance, please open an issue or join our If you have any questions, suggestions, or need assistance, please open an issue or join our
[Open WebUI Discord community](https://discord.gg/5rJgQTnV4s) to connect with us! 🤝 [Open WebUI Discord community](https://discord.gg/5rJgQTnV4s) to connect with us! 🤝
## Security 🛡️
If you believe you've found a security vulnerability, or something that shouldn't be disclosed publicly, please [reach out confidentially through our responsible disclosure program on GitHub](https://github.com/open-webui/open-webui/security). We accept reports only through GitHub, not through any other platform. Thank you for helping us keep Open WebUI secure!
## Star History ## Star History
<a href="https://star-history.com/#open-webui/open-webui&Date"> <a href="https://star-history.com/#open-webui/open-webui&Date">

View file

@ -1,3 +1,3 @@
export CORS_ALLOW_ORIGIN="http://localhost:5173;http://localhost:8080" export CORS_ALLOW_ORIGIN="http://localhost:5173;http://localhost:8080"
PORT="${PORT:-8080}" PORT="${PORT:-8080}"
uvicorn open_webui.main:app --port $PORT --host 0.0.0.0 --forwarded-allow-ips "${FORWARDED_ALLOW_IPS:-*}" --ws-per-message-deflate "${UVICORN_WS_PER_MESSAGE_DEFLATE:-true}" --reload uvicorn open_webui.main:app --port $PORT --host 0.0.0.0 --forwarded-allow-ips "${FORWARDED_ALLOW_IPS:-*}" --reload

View file

@ -11,16 +11,12 @@ import uvicorn
app = typer.Typer() app = typer.Typer()
KEY_FILE = Path.cwd() / '.webui_secret_key' KEY_FILE = Path.cwd() / '.webui_secret_key'
DEFAULT_SECRET_KEY_LENGTH = 24
def version_callback(value: bool) -> None: def version_callback(value: bool) -> None:
if value: if value:
from open_webui.env import VERSION from open_webui.env import VERSION
# LICENSE covers this Open WebUI CLI identifier.
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
# https://docs.openwebui.com/license.
typer.echo(f'Open WebUI version: {VERSION}') typer.echo(f'Open WebUI version: {VERSION}')
raise typer.Exit() raise typer.Exit()
@ -41,11 +37,8 @@ def serve(
if os.getenv('WEBUI_SECRET_KEY') is None: if os.getenv('WEBUI_SECRET_KEY') is None:
typer.echo('Loading WEBUI_SECRET_KEY from file, not provided as an environment variable.') typer.echo('Loading WEBUI_SECRET_KEY from file, not provided as an environment variable.')
if not KEY_FILE.exists(): if not KEY_FILE.exists():
key_length = int(os.getenv('WEBUI_SECRET_KEY_LENGTH', DEFAULT_SECRET_KEY_LENGTH))
if key_length < 1:
raise ValueError('WEBUI_SECRET_KEY_LENGTH must be a positive integer')
typer.echo(f'Generating a new secret key and saving it to {KEY_FILE}') typer.echo(f'Generating a new secret key and saving it to {KEY_FILE}')
KEY_FILE.write_bytes(base64.b64encode(random.randbytes(key_length))) KEY_FILE.write_bytes(base64.b64encode(random.randbytes(12)))
typer.echo(f'Loading WEBUI_SECRET_KEY from {KEY_FILE}') typer.echo(f'Loading WEBUI_SECRET_KEY from {KEY_FILE}')
os.environ['WEBUI_SECRET_KEY'] = KEY_FILE.read_text() os.environ['WEBUI_SECRET_KEY'] = KEY_FILE.read_text()
@ -74,7 +67,7 @@ def serve(
os.environ['LD_LIBRARY_PATH'] = ':'.join(LD_LIBRARY_PATH) os.environ['LD_LIBRARY_PATH'] = ':'.join(LD_LIBRARY_PATH)
import open_webui.main # noqa: F401 import open_webui.main # noqa: F401
from open_webui.env import UVICORN_WORKERS, UVICORN_WS_PER_MESSAGE_DEFLATE from open_webui.env import UVICORN_WORKERS # Import the workers setting
# On Windows, uvicorn's default loop factory hardcodes ProactorEventLoop, # On Windows, uvicorn's default loop factory hardcodes ProactorEventLoop,
# which is incompatible with psycopg v3 async. Setting loop='none' lets # which is incompatible with psycopg v3 async. Setting loop='none' lets
@ -87,7 +80,6 @@ def serve(
port=port, port=port,
forwarded_allow_ips='*', forwarded_allow_ips='*',
workers=UVICORN_WORKERS, workers=UVICORN_WORKERS,
ws_per_message_deflate=UVICORN_WS_PER_MESSAGE_DEFLATE,
loop=loop, loop=loop,
) )
@ -98,15 +90,12 @@ def dev(
port: int = 8080, port: int = 8080,
reload: bool = True, reload: bool = True,
): ):
from open_webui.env import UVICORN_WS_PER_MESSAGE_DEFLATE
uvicorn.run( uvicorn.run(
'open_webui.main:app', 'open_webui.main:app',
host=host, host=host,
port=port, port=port,
reload=reload, reload=reload,
forwarded_allow_ips='*', forwarded_allow_ips='*',
ws_per_message_deflate=UVICORN_WS_PER_MESSAGE_DEFLATE,
) )

File diff suppressed because it is too large Load diff

View file

@ -1,27 +1,7 @@
from __future__ import annotations from __future__ import annotations
import errno
from enum import Enum from enum import Enum
_ERRNO_MESSAGES = {
errno.ENAMETOOLONG: 'File name is too long.',
errno.ENOSPC: 'The server is out of storage space.',
errno.EDQUOT: 'Server storage quota exceeded.',
errno.EACCES: 'Server storage is not writable.',
errno.EPERM: 'Server storage is not writable.',
errno.EROFS: 'Server storage is not writable.',
}
def _error_message(err='', fallback='') -> str:
if not err:
return 'Something went wrong :/'
if isinstance(err, OSError) and err.errno in _ERRNO_MESSAGES:
return f'[ERROR: {_ERRNO_MESSAGES[err.errno]}]'
if isinstance(err, Exception):
return f'[ERROR: {fallback}]' if fallback else 'Something went wrong :/'
return f'[ERROR: {err}]'
class MESSAGES(str, Enum): class MESSAGES(str, Enum):
DEFAULT = lambda msg='': f'{msg if msg else ""}' DEFAULT = lambda msg='': f'{msg if msg else ""}'
@ -38,7 +18,7 @@ class ERROR_MESSAGES(str, Enum):
def __str__(self) -> str: def __str__(self) -> str:
return super().__str__() return super().__str__()
DEFAULT = _error_message DEFAULT = lambda err='': f'{"Something went wrong :/" if err == "" else "[ERROR: " + str(err) + "]"}'
ENV_VAR_NOT_FOUND = 'Required environment variable not found. Terminating now.' ENV_VAR_NOT_FOUND = 'Required environment variable not found. Terminating now.'
CREATE_USER_ERROR = 'Oops! Something went wrong while creating your account. Please try again later. If the issue persists, contact support for assistance.' CREATE_USER_ERROR = 'Oops! Something went wrong while creating your account. Please try again later. If the issue persists, contact support for assistance.'
DELETE_USER_ERROR = 'Oops! Something went wrong. We encountered an issue while trying to delete the user. Please give it another shot.' DELETE_USER_ERROR = 'Oops! Something went wrong. We encountered an issue while trying to delete the user. Please give it another shot.'
@ -98,7 +78,7 @@ class ERROR_MESSAGES(str, Enum):
INVALID_URL = 'The URL you provided is invalid. Please double-check and try again.' INVALID_URL = 'The URL you provided is invalid. Please double-check and try again.'
WEB_SEARCH_ERROR = 'Something went wrong while searching the web.' WEB_SEARCH_ERROR = lambda err='': err if err else 'Something went wrong while searching the web.'
OLLAMA_API_DISABLED = 'The Ollama API is disabled. Please enable it to use this feature.' OLLAMA_API_DISABLED = 'The Ollama API is disabled. Please enable it to use this feature.'
@ -117,16 +97,9 @@ class ERROR_MESSAGES(str, Enum):
AUTOMATION_TOO_FREQUENT = lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.' AUTOMATION_TOO_FREQUENT = lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.'
AUTOMATION_INVALID_RRULE = lambda err='': f'Invalid RRULE: {err}' AUTOMATION_INVALID_RRULE = lambda err='': f'Invalid RRULE: {err}'
AUTOMATION_NO_FUTURE_RUNS = 'RRULE has no future occurrences' AUTOMATION_NO_FUTURE_RUNS = 'RRULE has no future occurrences'
AUTOMATION_COUNT_REQUIRES_DTSTART = (
'RRULE with COUNT requires an explicit DTSTART line to anchor the occurrence window'
)
CALENDAR_RRULE_TOO_FREQUENT = 'Recurring events cannot repeat more often than daily'
FEATURE_DISABLED = lambda name='': f'{name} is disabled' FEATURE_DISABLED = lambda name='': f'{name} is disabled'
INPUT_TOO_LONG = lambda size='': f'Input prompt exceeds maximum length of {size}' INPUT_TOO_LONG = lambda size='': f'Input prompt exceeds maximum length of {size}'
# LICENSE covers this Open WebUI error identifier.
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
# https://docs.openwebui.com/license.
SERVER_CONNECTION_ERROR = 'Open WebUI: Server Connection Error' SERVER_CONNECTION_ERROR = 'Open WebUI: Server Connection Error'
REQUIRED_FIELD_EMPTY = lambda name='': f'Required field {name} is empty' REQUIRED_FIELD_EMPTY = lambda name='': f'Required field {name} is empty'
OAUTH_NOT_CONFIGURED = lambda name='': f"Provider '{name}' is not configured" OAUTH_NOT_CONFIGURED = lambda name='': f"Provider '{name}' is not configured"

View file

@ -7,9 +7,7 @@ import pkgutil
import re import re
import shutil import shutil
import sys import sys
import threading
import traceback import traceback
from contextlib import nullcontext
from pathlib import Path from pathlib import Path
from typing import Any, Optional from typing import Any, Optional
from uuid import uuid4 from uuid import uuid4
@ -42,13 +40,12 @@ except ImportError:
print('dotenv not installed, skipping...') print('dotenv not installed, skipping...')
DOCKER = os.getenv('DOCKER', 'False').lower() == 'true' DOCKER = os.getenv('DOCKER', 'False').lower() == 'true'
USE_SLIM = os.getenv('USE_SLIM_DOCKER', 'False').lower() == 'true'
USE_CUDA = os.getenv('USE_CUDA_DOCKER', 'false') USE_CUDA = os.getenv('USE_CUDA_DOCKER', 'false')
DEVICE_TYPE = 'cpu' DEVICE_TYPE = 'cpu'
_cuda_error: Optional[str] = None _cuda_error: Optional[str] = None
if not USE_SLIM and USE_CUDA.lower() == 'true': if USE_CUDA.lower() == 'true':
try: try:
import torch # noqa: E402 import torch # noqa: E402
@ -60,7 +57,7 @@ if not USE_SLIM and USE_CUDA.lower() == 'true':
os.environ['USE_CUDA_DOCKER'] = 'false' os.environ['USE_CUDA_DOCKER'] = 'false'
USE_CUDA = 'false' USE_CUDA = 'false'
if not USE_SLIM and sys.platform == 'darwin' and DEVICE_TYPE == 'cpu': if sys.platform == 'darwin' and DEVICE_TYPE == 'cpu':
try: try:
import torch # noqa: E402 import torch # noqa: E402
@ -69,9 +66,6 @@ if not USE_SLIM and sys.platform == 'darwin' and DEVICE_TYPE == 'cpu':
except Exception: except Exception:
pass pass
# Torch MPS inference is not thread-safe and a concurrent call kills the whole process.
MPS_INFERENCE_LOCK = threading.Lock() if DEVICE_TYPE == 'mps' else nullcontext()
#################################### ####################################
# LOGGING # LOGGING
#################################### ####################################
@ -108,7 +102,6 @@ class JSONFormatter(logging.Formatter):
LOG_FORMAT = os.getenv('LOG_FORMAT', '').lower() LOG_FORMAT = os.getenv('LOG_FORMAT', '').lower()
LOGURU_DIAGNOSE = os.getenv('LOGURU_DIAGNOSE', 'False').lower() == 'true'
GLOBAL_LOG_LEVEL = os.getenv('GLOBAL_LOG_LEVEL', '').upper() GLOBAL_LOG_LEVEL = os.getenv('GLOBAL_LOG_LEVEL', '').upper()
if GLOBAL_LOG_LEVEL in logging.getLevelNamesMapping(): if GLOBAL_LOG_LEVEL in logging.getLevelNamesMapping():
@ -156,11 +149,6 @@ INSTANCE_ID = os.getenv('INSTANCE_ID', str(uuid4()))
ENABLE_DB_MIGRATIONS = os.getenv('ENABLE_DB_MIGRATIONS', 'True').lower() == 'true' ENABLE_DB_MIGRATIONS = os.getenv('ENABLE_DB_MIGRATIONS', 'True').lower() == 'true'
# Swap the JSON encoder/decoder used across the app (HTTP request bodies, JSONResponse
# bodies, upstream provider responses, socket.io payloads) from the stdlib `json` module
# to orjson. Faster, but stricter: see open_webui/utils/json_codec.py for the differences.
ENABLE_ORJSON = os.getenv('ENABLE_ORJSON', 'False').lower() == 'true'
# Function to parse each section # Function to parse each section
def parse_section(section): def parse_section(section):
@ -233,7 +221,7 @@ if FROM_INIT_PY:
# Check if the data directory exists in the package directory # Check if the data directory exists in the package directory
if DATA_DIR.exists() and DATA_DIR != NEW_DATA_DIR: if DATA_DIR.exists() and DATA_DIR != NEW_DATA_DIR:
log.info('Moving %s to %s', DATA_DIR, NEW_DATA_DIR) log.info(f'Moving {DATA_DIR} to {NEW_DATA_DIR}')
for item in DATA_DIR.iterdir(): for item in DATA_DIR.iterdir():
dest = NEW_DATA_DIR / item.name dest = NEW_DATA_DIR / item.name
if item.is_dir(): if item.is_dir():
@ -251,6 +239,8 @@ if FROM_INIT_PY:
STATIC_DIR = Path(os.getenv('STATIC_DIR', OPEN_WEBUI_DIR / 'static')) STATIC_DIR = Path(os.getenv('STATIC_DIR', OPEN_WEBUI_DIR / 'static'))
FONTS_DIR = Path(os.getenv('FONTS_DIR', OPEN_WEBUI_DIR / 'static' / 'fonts'))
FRONTEND_BUILD_DIR = Path(os.getenv('FRONTEND_BUILD_DIR', BASE_DIR / 'build')).resolve() FRONTEND_BUILD_DIR = Path(os.getenv('FRONTEND_BUILD_DIR', BASE_DIR / 'build')).resolve()
if FROM_INIT_PY: if FROM_INIT_PY:
@ -301,7 +291,6 @@ if 'postgres://' in DATABASE_URL:
DATABASE_URL = DATABASE_URL.replace('postgres://', 'postgresql://') DATABASE_URL = DATABASE_URL.replace('postgres://', 'postgresql://')
DATABASE_SCHEMA = os.getenv('DATABASE_SCHEMA', None) DATABASE_SCHEMA = os.getenv('DATABASE_SCHEMA', None)
DATABASE_ENABLE_IAM_TOKEN_AUTH = os.getenv('DATABASE_ENABLE_IAM_TOKEN_AUTH', 'False').lower() == 'true'
_pool_size_raw = os.getenv('DATABASE_POOL_SIZE') _pool_size_raw = os.getenv('DATABASE_POOL_SIZE')
try: try:
@ -357,23 +346,20 @@ DATABASE_SQLITE_PRAGMA_MMAP_SIZE = os.getenv('DATABASE_SQLITE_PRAGMA_MMAP_SIZE',
# truncated. 67108864 ≈ 64 MB. Set to -1 for no limit (SQLite default). # truncated. 67108864 ≈ 64 MB. Set to -1 for no limit (SQLite default).
DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT = os.getenv('DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT', '67108864') DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT = os.getenv('DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT', '67108864')
# Seconds between presence writes per user per worker; keep under the 180s active-user window. 0 disables. DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL = os.getenv('DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL', None)
try: if DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL is not None:
DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL = float(os.getenv('DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL', '60')) try:
except ValueError: DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL = float(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL)
DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL = 60.0 except Exception:
DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL = 0.0
DATABASE_ENABLE_SESSION_SHARING = os.getenv('DATABASE_ENABLE_SESSION_SHARING', 'False').lower() == 'true' DATABASE_ENABLE_SESSION_SHARING = os.getenv('DATABASE_ENABLE_SESSION_SHARING', 'False').lower() == 'true'
ENABLE_PUBLIC_ACTIVE_USERS_COUNT = os.getenv('ENABLE_PUBLIC_ACTIVE_USERS_COUNT', 'True').lower() == 'true' ENABLE_PUBLIC_ACTIVE_USERS_COUNT = os.getenv('ENABLE_PUBLIC_ACTIVE_USERS_COUNT', 'True').lower() == 'true'
RESET_CONFIG_ON_START = os.getenv('RESET_CONFIG_ON_START', 'False').lower() == 'true' RESET_CONFIG_ON_START = os.getenv('RESET_CONFIG_ON_START', 'False').lower() == 'true'
ENABLE_REALTIME_CHAT_SAVE = os.getenv('ENABLE_REALTIME_CHAT_SAVE', 'False').lower() == 'true' ENABLE_REALTIME_CHAT_SAVE = os.getenv('ENABLE_REALTIME_CHAT_SAVE', 'False').lower() == 'true'
ENABLE_QUERIES_CACHE = os.getenv('ENABLE_QUERIES_CACHE', 'False').lower() == 'true' ENABLE_QUERIES_CACHE = os.getenv('ENABLE_QUERIES_CACHE', 'False').lower() == 'true'
ENABLE_ADMIN_CHAT_ACCESS = os.getenv('ENABLE_ADMIN_CHAT_ACCESS', 'True').lower() == 'true'
RAG_SYSTEM_CONTEXT = os.getenv('RAG_SYSTEM_CONTEXT', 'False').lower() == 'true' RAG_SYSTEM_CONTEXT = os.getenv('RAG_SYSTEM_CONTEXT', 'False').lower() == 'true'
# Empty by default: chunk metadata also holds internal bookkeeping (file hashes, collection names, scores).
RAG_SOURCE_METADATA_KEYS = [key.strip() for key in os.getenv('RAG_SOURCE_METADATA_KEYS', '').split(',') if key.strip()]
#################################### ####################################
# REDIS # REDIS
#################################### ####################################
@ -383,19 +369,6 @@ REDIS_CLUSTER = os.getenv('REDIS_CLUSTER', 'False').lower() == 'true'
REDIS_KEY_PREFIX = os.getenv('REDIS_KEY_PREFIX', 'open-webui') REDIS_KEY_PREFIX = os.getenv('REDIS_KEY_PREFIX', 'open-webui')
try:
REDIS_RESPONSE_STREAM_TTL = int(os.getenv('REDIS_RESPONSE_STREAM_TTL', '3600'))
except ValueError:
REDIS_RESPONSE_STREAM_TTL = 3600
# Seconds a task survives without a heartbeat. 0 disables expiry.
try:
REDIS_TASK_TTL = int(os.getenv('REDIS_TASK_TTL', '300'))
if REDIS_TASK_TTL != 0 and REDIS_TASK_TTL < 60:
REDIS_TASK_TTL = 300
except ValueError:
REDIS_TASK_TTL = 300
REDIS_SENTINEL_HOSTS = os.getenv('REDIS_SENTINEL_HOSTS', '') REDIS_SENTINEL_HOSTS = os.getenv('REDIS_SENTINEL_HOSTS', '')
REDIS_SENTINEL_PORT = os.getenv('REDIS_SENTINEL_PORT', '26379') REDIS_SENTINEL_PORT = os.getenv('REDIS_SENTINEL_PORT', '26379')
@ -415,12 +388,6 @@ try:
except ValueError: except ValueError:
REDIS_SOCKET_CONNECT_TIMEOUT = None REDIS_SOCKET_CONNECT_TIMEOUT = None
REDIS_SOCKET_TIMEOUT = os.getenv('REDIS_SOCKET_TIMEOUT', '')
try:
REDIS_SOCKET_TIMEOUT = float(REDIS_SOCKET_TIMEOUT)
except ValueError:
REDIS_SOCKET_TIMEOUT = None
# Whether to enable TCP SO_KEEPALIVE on Redis client sockets. Opt-in: # Whether to enable TCP SO_KEEPALIVE on Redis client sockets. Opt-in:
# defaults to off so behavior is unchanged for existing deployments. When # defaults to off so behavior is unchanged for existing deployments. When
# enabled, the kernel sends TCP keepalive probes on idle connections so # enabled, the kernel sends TCP keepalive probes on idle connections so
@ -463,9 +430,6 @@ try:
except (ValueError, TypeError): except (ValueError, TypeError):
UVICORN_WORKERS = 1 UVICORN_WORKERS = 1
# tiny delta-stream frames make per-frame websocket compression CPU-bound, allow opting out (true/false)
UVICORN_WS_PER_MESSAGE_DEFLATE = os.getenv('UVICORN_WS_PER_MESSAGE_DEFLATE', 'True').lower() == 'true'
#################################### ####################################
# WEBSOCKET SUPPORT # WEBSOCKET SUPPORT
#################################### ####################################
@ -479,16 +443,17 @@ WEBSOCKET_REDIS_OPTIONS = os.getenv('WEBSOCKET_REDIS_OPTIONS', '')
if WEBSOCKET_REDIS_OPTIONS == '': if WEBSOCKET_REDIS_OPTIONS == '':
WEBSOCKET_REDIS_OPTIONS = {'socket_timeout': None}
if REDIS_SOCKET_CONNECT_TIMEOUT: if REDIS_SOCKET_CONNECT_TIMEOUT:
WEBSOCKET_REDIS_OPTIONS['socket_connect_timeout'] = REDIS_SOCKET_CONNECT_TIMEOUT WEBSOCKET_REDIS_OPTIONS = {'socket_connect_timeout': REDIS_SOCKET_CONNECT_TIMEOUT}
else:
log.debug('No WEBSOCKET_REDIS_OPTIONS provided, defaulting to None')
WEBSOCKET_REDIS_OPTIONS = None
else: else:
try: try:
WEBSOCKET_REDIS_OPTIONS = json.loads(WEBSOCKET_REDIS_OPTIONS) WEBSOCKET_REDIS_OPTIONS = json.loads(WEBSOCKET_REDIS_OPTIONS)
WEBSOCKET_REDIS_OPTIONS.setdefault('socket_timeout', None)
except Exception: except Exception:
log.warning('Invalid WEBSOCKET_REDIS_OPTIONS, defaulting to socket_timeout=None') log.warning('Invalid WEBSOCKET_REDIS_OPTIONS, defaulting to None')
WEBSOCKET_REDIS_OPTIONS = {'socket_timeout': None} WEBSOCKET_REDIS_OPTIONS = None
WEBSOCKET_REDIS_URL = os.getenv('WEBSOCKET_REDIS_URL', REDIS_URL) WEBSOCKET_REDIS_URL = os.getenv('WEBSOCKET_REDIS_URL', REDIS_URL)
WEBSOCKET_REDIS_CLUSTER = os.getenv('WEBSOCKET_REDIS_CLUSTER', str(REDIS_CLUSTER)).lower() == 'true' WEBSOCKET_REDIS_CLUSTER = os.getenv('WEBSOCKET_REDIS_CLUSTER', str(REDIS_CLUSTER)).lower() == 'true'
@ -522,15 +487,6 @@ try:
except ValueError: except ValueError:
WEBSOCKET_SERVER_PING_INTERVAL = 25 WEBSOCKET_SERVER_PING_INTERVAL = 25
WEBSOCKET_HEARTBEAT_INTERVAL = os.getenv('WEBSOCKET_HEARTBEAT_INTERVAL', '')
if WEBSOCKET_HEARTBEAT_INTERVAL == '':
WEBSOCKET_HEARTBEAT_INTERVAL = None
else:
try:
WEBSOCKET_HEARTBEAT_INTERVAL = min(max(int(WEBSOCKET_HEARTBEAT_INTERVAL), 5), 90)
except ValueError:
WEBSOCKET_HEARTBEAT_INTERVAL = 30
WEBSOCKET_EVENT_CALLER_TIMEOUT = os.getenv('WEBSOCKET_EVENT_CALLER_TIMEOUT', '') WEBSOCKET_EVENT_CALLER_TIMEOUT = os.getenv('WEBSOCKET_EVENT_CALLER_TIMEOUT', '')
if WEBSOCKET_EVENT_CALLER_TIMEOUT == '': if WEBSOCKET_EVENT_CALLER_TIMEOUT == '':
@ -542,112 +498,20 @@ else:
WEBSOCKET_EVENT_CALLER_TIMEOUT = 300 WEBSOCKET_EVENT_CALLER_TIMEOUT = 300
import ssl as _ssl
# Dedicated env var for a custom CA bundle file path. When set, this is
# used as the default CA bundle for all outbound HTTPS connections that
# have SSL verification enabled (i.e. when their per-connection SSL env
# var is ``"True"``). Per-connection overrides (setting the SSL env var
# to a path directly) take precedence over this global fallback.
#
# This follows the industry convention of ``SSL_CERT_FILE`` / ``REQUESTS_CA_BUNDLE``
# but is scoped to Open WebUI to avoid interfering with system-level settings.
AIOHTTP_CLIENT_SSL_CERT_FILE = os.getenv('AIOHTTP_CLIENT_SSL_CERT_FILE', '').strip()
def _build_ssl_context_from_file(path: str) -> '_ssl.SSLContext | None':
"""Create an SSLContext from a CA bundle file, or None if invalid."""
if not path:
return None
if not os.path.isfile(path):
log.warning(
'SSL CA bundle path does not exist: %r, ignoring',
path,
)
return None
ctx = _ssl.create_default_context(cafile=path)
log.info('Using custom SSL CA bundle: %s', path)
return ctx
# Pre-built SSLContext from the dedicated env var (cached once at startup).
_GLOBAL_SSL_CONTEXT = _build_ssl_context_from_file(AIOHTTP_CLIENT_SSL_CERT_FILE)
def _parse_ssl_env(value: str) -> 'bool | _ssl.SSLContext':
"""Parse an SSL env var into a bool or SSLContext.
- ``"true"`` → uses ``AIOHTTP_CLIENT_SSL_CERT_FILE`` context if set,
otherwise ``True`` (default SSL verification via certifi)
- ``"false"`` → ``False`` (no verification)
- ``"/path/to/ca-bundle.crt"`` → ``SSLContext`` loading that CA file
(takes precedence over ``AIOHTTP_CLIENT_SSL_CERT_FILE``)
This allows users with corporate or internal CAs to point Open WebUI
at a custom CA bundle without disabling verification entirely.
"""
lower = value.strip().lower()
if lower == 'true':
# Use the global dedicated CA bundle if configured, otherwise default
return _GLOBAL_SSL_CONTEXT if _GLOBAL_SSL_CONTEXT is not None else True
if lower == 'false':
return False
# Treat as a file path to a CA bundle (per-connection override)
ctx = _build_ssl_context_from_file(value.strip())
if ctx is not None:
return ctx
# Path was invalid — fall back to default
return _GLOBAL_SSL_CONTEXT if _GLOBAL_SSL_CONTEXT is not None else True
REQUESTS_VERIFY = os.getenv('REQUESTS_VERIFY', 'True').lower() == 'true' REQUESTS_VERIFY = os.getenv('REQUESTS_VERIFY', 'True').lower() == 'true'
TAVILY_API_BASE_URL = os.getenv('TAVILY_API_BASE_URL', 'https://api.tavily.com').rstrip('/')
_aiohttp_timeout_raw = os.getenv('AIOHTTP_CLIENT_TIMEOUT', '') _aiohttp_timeout_raw = os.getenv('AIOHTTP_CLIENT_TIMEOUT', '')
try: try:
AIOHTTP_CLIENT_TIMEOUT = int(_aiohttp_timeout_raw) if _aiohttp_timeout_raw else None AIOHTTP_CLIENT_TIMEOUT = int(_aiohttp_timeout_raw) if _aiohttp_timeout_raw else None
except (ValueError, TypeError): except (ValueError, TypeError):
AIOHTTP_CLIENT_TIMEOUT = 300 AIOHTTP_CLIENT_TIMEOUT = 300
# Optional between-chunks idle cap for streaming aiohttp requests.
AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT = os.getenv('AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT', '')
if AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT == '':
AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT = None
else:
try:
AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT = int(AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT)
except (ValueError, TypeError):
AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT = None
if AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT is not None and AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT <= 0: AIOHTTP_CLIENT_SESSION_SSL = os.getenv('AIOHTTP_CLIENT_SESSION_SSL', 'True').lower() == 'true'
AIOHTTP_CLIENT_STREAM_IDLE_TIMEOUT = None
# SSL verification for general outbound requests (OpenAI, OAuth, etc.).
# Accepts "True", "False", or a path to a CA bundle file.
# When "True", falls back to AIOHTTP_CLIENT_SSL_CERT_FILE if set.
AIOHTTP_CLIENT_SESSION_SSL = _parse_ssl_env(os.getenv('AIOHTTP_CLIENT_SESSION_SSL', 'True'))
SEARXNG_CLIENT_CERT_FILE = os.getenv('SEARXNG_CLIENT_CERT_FILE', '').strip()
SEARXNG_CLIENT_KEY_FILE = os.getenv('SEARXNG_CLIENT_KEY_FILE', '').strip()
# When False (default), outbound HTTP requests do not follow 3xx redirects. # When False (default), outbound HTTP requests do not follow 3xx redirects.
AIOHTTP_CLIENT_ALLOW_REDIRECTS = os.getenv('AIOHTTP_CLIENT_ALLOW_REDIRECTS', 'False').lower() == 'true' AIOHTTP_CLIENT_ALLOW_REDIRECTS = os.getenv('AIOHTTP_CLIENT_ALLOW_REDIRECTS', 'False').lower() == 'true'
# Opt-in c-ares DNS resolution (aiodns). Off by default: c-ares breaks name
# resolution in some environments (#28013, #28215). Must run before any
# TCPConnector is constructed.
AIOHTTP_CLIENT_ASYNC_DNS_RESOLVER = os.getenv('AIOHTTP_CLIENT_ASYNC_DNS_RESOLVER', 'False').lower() == 'true'
if not AIOHTTP_CLIENT_ASYNC_DNS_RESOLVER:
import aiohttp
aiohttp.DefaultResolver = aiohttp.resolver.ThreadedResolver # for plugin code
aiohttp.resolver.DefaultResolver = aiohttp.resolver.ThreadedResolver
aiohttp.connector.DefaultResolver = aiohttp.resolver.ThreadedResolver
# Optional User-Agent override for outbound web-loader fetches. When set, # Optional User-Agent override for outbound web-loader fetches. When set,
# SafeWebBaseLoader sends this value instead of the default python-requests UA # SafeWebBaseLoader sends this value instead of the default python-requests UA
# which is aggressively blocked by Cloudflare, Wikipedia, and similar services. # which is aggressively blocked by Cloudflare, Wikipedia, and similar services.
@ -668,20 +532,8 @@ try:
except (ValueError, TypeError): except (ValueError, TypeError):
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER_DATA = 10 AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER_DATA = 10
AIOHTTP_FILE_STREAM_CHUNK_SIZE = os.getenv('AIOHTTP_FILE_STREAM_CHUNK_SIZE', str(1024 * 1024))
try:
AIOHTTP_FILE_STREAM_CHUNK_SIZE = int(AIOHTTP_FILE_STREAM_CHUNK_SIZE)
except Exception:
AIOHTTP_FILE_STREAM_CHUNK_SIZE = 1024 * 1024
if AIOHTTP_FILE_STREAM_CHUNK_SIZE <= 0: AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL = os.getenv('AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL', 'True').lower() == 'true'
AIOHTTP_FILE_STREAM_CHUNK_SIZE = 1024 * 1024
# SSL verification for tool server connections specifically.
# Accepts "True", "False", or a path to a CA bundle file.
# When "True", falls back to AIOHTTP_CLIENT_SSL_CERT_FILE if set.
AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL = _parse_ssl_env(os.getenv('AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL', 'True'))
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER = os.getenv('AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER', '') AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER = os.getenv('AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER', '')
@ -764,8 +616,6 @@ WEBUI_SECRET_KEY = os.getenv(
os.getenv('WEBUI_JWT_SECRET_KEY', ''), os.getenv('WEBUI_JWT_SECRET_KEY', ''),
) )
ENABLE_VALVE_ENCRYPTION = os.getenv('ENABLE_VALVE_ENCRYPTION', 'False').lower() == 'true'
WEBUI_SESSION_COOKIE_SAME_SITE = os.getenv('WEBUI_SESSION_COOKIE_SAME_SITE', 'lax') WEBUI_SESSION_COOKIE_SAME_SITE = os.getenv('WEBUI_SESSION_COOKIE_SAME_SITE', 'lax')
WEBUI_SESSION_COOKIE_SECURE = os.getenv('WEBUI_SESSION_COOKIE_SECURE', 'false').lower() == 'true' WEBUI_SESSION_COOKIE_SECURE = os.getenv('WEBUI_SESSION_COOKIE_SECURE', 'false').lower() == 'true'
WEBUI_AUTH_COOKIE_SAME_SITE = os.getenv('WEBUI_AUTH_COOKIE_SAME_SITE', WEBUI_SESSION_COOKIE_SAME_SITE) WEBUI_AUTH_COOKIE_SAME_SITE = os.getenv('WEBUI_AUTH_COOKIE_SAME_SITE', WEBUI_SESSION_COOKIE_SAME_SITE)
@ -812,7 +662,6 @@ WEBUI_AUTH_TRUSTED_ROLE_HEADER = os.getenv('WEBUI_AUTH_TRUSTED_ROLE_HEADER', Non
CUSTOM_API_KEY_HEADER = os.getenv('CUSTOM_API_KEY_HEADER', 'x-api-key') CUSTOM_API_KEY_HEADER = os.getenv('CUSTOM_API_KEY_HEADER', 'x-api-key')
ENABLE_PASSWORD_VALIDATION = os.getenv('ENABLE_PASSWORD_VALIDATION', 'False').lower() == 'true' ENABLE_PASSWORD_VALIDATION = os.getenv('ENABLE_PASSWORD_VALIDATION', 'False').lower() == 'true'
PASSWORD_HASH_ALGORITHM = os.getenv('PASSWORD_HASH_ALGORITHM', 'bcrypt').lower()
PASSWORD_VALIDATION_REGEX_PATTERN = os.getenv( PASSWORD_VALIDATION_REGEX_PATTERN = os.getenv(
'PASSWORD_VALIDATION_REGEX_PATTERN', 'PASSWORD_VALIDATION_REGEX_PATTERN',
r'^(?=.*[a-z])(?=.*[A-Z])(?=.*\d)(?=.*[^\w\s]).{8,}$', r'^(?=.*[a-z])(?=.*[A-Z])(?=.*\d)(?=.*[^\w\s]).{8,}$',
@ -838,23 +687,10 @@ BYPASS_RETRIEVAL_ACCESS_CONTROL = os.getenv('BYPASS_RETRIEVAL_ACCESS_CONTROL', '
# denied — closing the legacy unscoped namespace. # denied — closing the legacy unscoped namespace.
ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS = os.getenv('ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS', 'False').lower() == 'true' ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS = os.getenv('ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS', 'False').lower() == 'true'
# Falls back to the upload size limit, because a document cannot legitimately carry more metadata
# than the file itself is allowed to be. Left unbounded, a small archive that expands enormously
# during extraction can exhaust memory. RAG_FILE_MAX_SIZE is in MB.
RAG_METADATA_MAX_VALUE_CHARS = (
int(os.getenv('RAG_METADATA_MAX_VALUE_CHARS'))
if os.getenv('RAG_METADATA_MAX_VALUE_CHARS')
else ((int(os.getenv('RAG_FILE_MAX_SIZE', '0')) or 0) * 1024 * 1024 or None)
)
MINERU_MAX_MARKDOWN_BYTES = (
int(os.getenv('MINERU_MAX_MARKDOWN_BYTES')) if os.getenv('MINERU_MAX_MARKDOWN_BYTES') else None
)
# When enabled, skips pydub-based preprocessing (format conversion, compression, # When enabled, skips pydub-based preprocessing (format conversion, compression,
# and chunked splitting) before sending files to processing engines. Useful when # and chunked splitting) before sending files to processing engines. Useful when
# the upstream provider handles these steps or when ffmpeg is unavailable. # the upstream provider handles these steps or when ffmpeg is unavailable.
BYPASS_PYDUB_PREPROCESSING = USE_SLIM or os.getenv('BYPASS_PYDUB_PREPROCESSING', 'False').lower() == 'true' BYPASS_PYDUB_PREPROCESSING = os.getenv('BYPASS_PYDUB_PREPROCESSING', 'False').lower() == 'true'
# When disabled (default), the OpenAI catch-all proxy endpoint (/{path:path}) # When disabled (default), the OpenAI catch-all proxy endpoint (/{path:path})
# is blocked. Enable only if you need direct passthrough to upstream OpenAI- # is blocked. Enable only if you need direct passthrough to upstream OpenAI-
@ -881,18 +717,6 @@ OAUTH_MAX_SESSIONS_PER_USER = int(os.getenv('OAUTH_MAX_SESSIONS_PER_USER', '10')
# Token Exchange Configuration # Token Exchange Configuration
# Allows external apps to exchange OAuth tokens for OpenWebUI tokens # Allows external apps to exchange OAuth tokens for OpenWebUI tokens
ENABLE_OAUTH_TOKEN_EXCHANGE = os.getenv('ENABLE_OAUTH_TOKEN_EXCHANGE', 'False').lower() == 'true' ENABLE_OAUTH_TOKEN_EXCHANGE = os.getenv('ENABLE_OAUTH_TOKEN_EXCHANGE', 'False').lower() == 'true'
_oauth_token_exchange_rate_limit = (os.getenv('OAUTH_TOKEN_EXCHANGE_RATE_LIMIT') or '').strip()
OAUTH_TOKEN_EXCHANGE_RATE_LIMIT = (
int(_oauth_token_exchange_rate_limit)
if _oauth_token_exchange_rate_limit and _oauth_token_exchange_rate_limit.lower() != 'none'
else None
)
OAUTH_TOKEN_EXCHANGE_RATE_LIMIT_WINDOW = int(os.getenv('OAUTH_TOKEN_EXCHANGE_RATE_LIMIT_WINDOW', str(60 * 3)))
OAUTH_TOKEN_EXCHANGE_TRUSTED_CLIENT_IDS = [
client_id.strip()
for client_id in os.getenv('OAUTH_TOKEN_EXCHANGE_TRUSTED_CLIENT_IDS', '').split(',')
if client_id.strip()
]
# Back-Channel Logout Configuration # Back-Channel Logout Configuration
# When enabled, exposes POST /oauth/backchannel-logout for IdP-initiated logout # When enabled, exposes POST /oauth/backchannel-logout for IdP-initiated logout
@ -944,18 +768,10 @@ if LICENSE_PUBLIC_KEY:
# WEBUI Identity # WEBUI Identity
#################################### ####################################
# LICENSE covers this Open WebUI branding surface, including name, logo,
# visual, textual, symbolic identifiers, metadata, and surrounding UI.
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
# https://docs.openwebui.com/license.
WEBUI_NAME = os.getenv('WEBUI_NAME', 'Open WebUI') WEBUI_NAME = os.getenv('WEBUI_NAME', 'Open WebUI')
if WEBUI_NAME != 'Open WebUI': if WEBUI_NAME != 'Open WebUI':
WEBUI_NAME += ' (Open WebUI)' WEBUI_NAME += ' (Open WebUI)'
# LICENSE covers this Open WebUI branding surface, including this favicon
# and any visual, textual, or symbolic identifiers it preserves.
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
# https://docs.openwebui.com/license.
WEBUI_FAVICON_URL = 'https://openwebui.com/favicon.png' WEBUI_FAVICON_URL = 'https://openwebui.com/favicon.png'
WEBUI_BUILD_HASH = os.getenv('WEBUI_BUILD_HASH', 'dev-build') WEBUI_BUILD_HASH = os.getenv('WEBUI_BUILD_HASH', 'dev-build')
TRUSTED_SIGNATURE_KEY = os.getenv('TRUSTED_SIGNATURE_KEY', '') TRUSTED_SIGNATURE_KEY = os.getenv('TRUSTED_SIGNATURE_KEY', '')
@ -994,7 +810,6 @@ FORWARD_USER_INFO_HEADER_USER_NAME = os.getenv('FORWARD_USER_INFO_HEADER_USER_NA
FORWARD_USER_INFO_HEADER_USER_ID = os.getenv('FORWARD_USER_INFO_HEADER_USER_ID', 'X-OpenWebUI-User-Id') FORWARD_USER_INFO_HEADER_USER_ID = os.getenv('FORWARD_USER_INFO_HEADER_USER_ID', 'X-OpenWebUI-User-Id')
FORWARD_USER_INFO_HEADER_USER_EMAIL = os.getenv('FORWARD_USER_INFO_HEADER_USER_EMAIL', 'X-OpenWebUI-User-Email') FORWARD_USER_INFO_HEADER_USER_EMAIL = os.getenv('FORWARD_USER_INFO_HEADER_USER_EMAIL', 'X-OpenWebUI-User-Email')
FORWARD_USER_INFO_HEADER_USER_ROLE = os.getenv('FORWARD_USER_INFO_HEADER_USER_ROLE', 'X-OpenWebUI-User-Role') FORWARD_USER_INFO_HEADER_USER_ROLE = os.getenv('FORWARD_USER_INFO_HEADER_USER_ROLE', 'X-OpenWebUI-User-Role')
FORWARD_USER_INFO_HEADER_AUTH_TYPE = os.getenv('FORWARD_USER_INFO_HEADER_AUTH_TYPE', 'X-OpenWebUI-Auth-Type')
FORWARD_SESSION_INFO_HEADER_MESSAGE_ID = os.getenv('FORWARD_SESSION_INFO_HEADER_MESSAGE_ID', 'X-OpenWebUI-Message-Id') 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') FORWARD_SESSION_INFO_HEADER_CHAT_ID = os.getenv('FORWARD_SESSION_INFO_HEADER_CHAT_ID', 'X-OpenWebUI-Chat-Id')
@ -1013,10 +828,6 @@ except ValueError:
# Progressive Web App # Progressive Web App
#################################### ####################################
# LICENSE covers this install-time Open WebUI branding surface, including
# names, logos, manifests, metadata, and surrounding UI.
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
# https://docs.openwebui.com/license.
EXTERNAL_PWA_MANIFEST_URL = os.getenv('EXTERNAL_PWA_MANIFEST_URL', None) EXTERNAL_PWA_MANIFEST_URL = os.getenv('EXTERNAL_PWA_MANIFEST_URL', None)
#################################### ####################################
@ -1051,14 +862,6 @@ else:
ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION = ( ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION = (
os.getenv('ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION', 'False').lower() == 'true' os.getenv('ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION', 'False').lower() == 'true'
) )
ENABLE_API_OUTLET_FILTERS = os.getenv('ENABLE_API_OUTLET_FILTERS', 'True').lower() == 'true'
# Opt in to CPython's in-place string append optimization for streamed responses.
# Off by default for a staged rollout. Only a host already out of memory can lose
# text here; the default path (a full copy per chunk) raises there too.
ENABLE_CHAT_RESPONSE_STREAM_INPLACE_APPEND = (
os.getenv('ENABLE_CHAT_RESPONSE_STREAM_INPLACE_APPEND', 'False').lower() == 'true'
)
# When enabled, uses a hardcoded extension-to-MIME dictionary as a last-resort # When enabled, uses a hardcoded extension-to-MIME dictionary as a last-resort
# fallback when both mimetypes.guess_type() and file.meta.content_type fail to # fallback when both mimetypes.guess_type() and file.meta.content_type fail to
@ -1159,34 +962,10 @@ SENTENCE_TRANSFORMERS_CROSS_ENCODER_SIGMOID_ACTIVATION_FUNCTION = (
os.getenv('SENTENCE_TRANSFORMERS_CROSS_ENCODER_SIGMOID_ACTIVATION_FUNCTION', 'True').lower() == 'true' os.getenv('SENTENCE_TRANSFORMERS_CROSS_ENCODER_SIGMOID_ACTIVATION_FUNCTION', 'True').lower() == 'true'
) )
####################################
# KNOWLEDGE TOOLS
####################################
def _int_env(name: str, default: int) -> int:
try:
return max(int(os.getenv(name) or default), 1)
except (ValueError, TypeError):
return default
# Total output of a single kb_exec call, whatever the command.
KB_EXEC_MAX_OUTPUT_CHARS = _int_env('KB_EXEC_MAX_OUTPUT_CHARS', 30_000)
# Files a single kb_exec grep may scan before it asks for a narrower scope.
KB_EXEC_MAX_GREP_FILES = _int_env('KB_EXEC_MAX_GREP_FILES', 200)
# Matching lines returned by kb_exec grep and grep_knowledge_files.
KNOWLEDGE_GREP_MAX_MATCHES = _int_env('KNOWLEDGE_GREP_MAX_MATCHES', 50)
# Characters returned by view_file / view_knowledge_file.
VIEW_FILE_MAX_CHARS = _int_env('VIEW_FILE_MAX_CHARS', 100_000)
VIEW_FILE_DEFAULT_MAX_CHARS = _int_env('VIEW_FILE_DEFAULT_MAX_CHARS', 10_000)
#################################### ####################################
# TOOLS/FUNCTIONS PIP OPTIONS # TOOLS/FUNCTIONS PIP OPTIONS
#################################### ####################################
ENABLE_PLUGINS = os.getenv('ENABLE_PLUGINS', 'True').lower() == 'true'
ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS = ( ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS = (
os.getenv('ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS', 'True').lower() == 'true' os.getenv('ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS', 'True').lower() == 'true'
) )
@ -1206,12 +985,6 @@ if OFFLINE_MODE:
os.environ['HF_HUB_OFFLINE'] = '1' os.environ['HF_HUB_OFFLINE'] = '1'
ENABLE_VERSION_UPDATE_CHECK = False ENABLE_VERSION_UPDATE_CHECK = False
####################################
# Pyodide file persistence
####################################
ENABLE_PYODIDE_FILE_PERSISTENCE = os.getenv('ENABLE_PYODIDE_FILE_PERSISTENCE', 'false').lower() == 'true'
#################################### ####################################
# Audit logging # Audit logging
#################################### ####################################
@ -1240,19 +1013,15 @@ except ValueError:
MAX_BODY_LOG_SIZE = 2048 MAX_BODY_LOG_SIZE = 2048
# Comma separated list for urls to exclude from audit # Comma separated list for urls to exclude from audit
AUDIT_EXCLUDED_PATHS = [ AUDIT_EXCLUDED_PATHS = os.getenv('AUDIT_EXCLUDED_PATHS', '/chats,/chat,/folders').split(',')
path AUDIT_EXCLUDED_PATHS = [path.strip() for path in AUDIT_EXCLUDED_PATHS]
for path in ( AUDIT_EXCLUDED_PATHS = [path.lstrip('/') for path in AUDIT_EXCLUDED_PATHS]
path.strip().lstrip('/') for path in os.getenv('AUDIT_EXCLUDED_PATHS', '/chats,/chat,/folders').split(',')
)
if path
]
# Comma separated list of urls to include in audit (whitelist mode) # Comma separated list of urls to include in audit (whitelist mode)
# When set, only these paths are audited and AUDIT_EXCLUDED_PATHS is ignored # When set, only these paths are audited and AUDIT_EXCLUDED_PATHS is ignored
AUDIT_INCLUDED_PATHS = [ AUDIT_INCLUDED_PATHS = os.getenv('AUDIT_INCLUDED_PATHS', '').split(',')
path for path in (path.strip().lstrip('/') for path in os.getenv('AUDIT_INCLUDED_PATHS', '').split(',')) if path AUDIT_INCLUDED_PATHS = [path.strip() for path in AUDIT_INCLUDED_PATHS]
] AUDIT_INCLUDED_PATHS = [path.lstrip('/') for path in AUDIT_INCLUDED_PATHS if path]
# When enabled, GET requests are also audited (disabled by default to avoid log noise) # When enabled, GET requests are also audited (disabled by default to avoid log noise)
ENABLE_AUDIT_GET_REQUESTS = os.getenv('ENABLE_AUDIT_GET_REQUESTS', 'False').lower() == 'true' ENABLE_AUDIT_GET_REQUESTS = os.getenv('ENABLE_AUDIT_GET_REQUESTS', 'False').lower() == 'true'

File diff suppressed because it is too large Load diff

View file

@ -1,5 +1,6 @@
import asyncio import asyncio
import inspect import inspect
import json
import logging import logging
import sys import sys
from typing import AsyncGenerator, Generator, Iterator from typing import AsyncGenerator, Generator, Iterator
@ -19,7 +20,7 @@ from starlette.responses import Response, StreamingResponse
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
from open_webui.constants import ERROR_MESSAGES from open_webui.constants import ERROR_MESSAGES
from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_PLUGINS, GLOBAL_LOG_LEVEL from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, GLOBAL_LOG_LEVEL
from open_webui.models.functions import Functions from open_webui.models.functions import Functions
from open_webui.models.models import Models from open_webui.models.models import Models
from open_webui.models.users import UserModel from open_webui.models.users import UserModel
@ -28,7 +29,6 @@ from open_webui.socket.main import (
get_event_emitter, get_event_emitter,
) )
from open_webui.utils.access_control import check_model_access from open_webui.utils.access_control import check_model_access
from open_webui.utils.json_codec import JSONCodec
from open_webui.utils.misc import ( from open_webui.utils.misc import (
add_or_update_system_message, add_or_update_system_message,
get_last_user_message, get_last_user_message,
@ -69,9 +69,6 @@ async def get_function_module_by_id(request: Request, pipe_id: str):
async def get_function_models(request): async def get_function_models(request):
if not ENABLE_PLUGINS:
return []
pipes = await Functions.get_functions_by_type('pipe', active_only=True) pipes = await Functions.get_functions_by_type('pipe', active_only=True)
pipe_models = [] pipe_models = []
@ -100,7 +97,7 @@ async def get_function_models(request):
log.exception(e) log.exception(e)
sub_pipes = [] sub_pipes = []
log.debug("get_function_models: function '%s' is a manifold of %s", pipe.id, sub_pipes) log.debug(f"get_function_models: function '{pipe.id}' is a manifold of {sub_pipes}")
for p in sub_pipes: for p in sub_pipes:
sub_pipe_id = f'{pipe.id}.{p["id"]}' sub_pipe_id = f'{pipe.id}.{p["id"]}'
@ -126,10 +123,7 @@ async def get_function_models(request):
pipe_flag = {'type': 'pipe'} pipe_flag = {'type': 'pipe'}
log.debug( log.debug(
"get_function_models: function '%s' is a single pipe { 'id': %s, 'name': %s }", f"get_function_models: function '{pipe.id}' is a single pipe {{ 'id': {pipe.id}, 'name': {pipe.name} }}"
pipe.id,
pipe.id,
pipe.name,
) )
pipe_models.append( pipe_models.append(
@ -150,10 +144,7 @@ async def get_function_models(request):
return pipe_models return pipe_models
async def generate_function_chat_completion(request, form_data, user, models: dict | None = None): async def generate_function_chat_completion(request, form_data, user, models: dict = {}):
if models is None:
models = {}
async def execute_pipe(pipe, params): async def execute_pipe(pipe, params):
if inspect.iscoroutinefunction(pipe): if inspect.iscoroutinefunction(pipe):
return await pipe(**params) return await pipe(**params)
@ -173,7 +164,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
line = line.model_dump_json() line = line.model_dump_json()
line = f'data: {line}' line = f'data: {line}'
if isinstance(line, dict): if isinstance(line, dict):
line = f'data: {JSONCodec.dumps(line)}' line = f'data: {json.dumps(line)}'
try: try:
line = line.decode('utf-8') line = line.decode('utf-8')
@ -184,7 +175,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
return f'{line}\n\n' return f'{line}\n\n'
else: else:
line = openai_chat_chunk_message_template(form_data['model'], line) line = openai_chat_chunk_message_template(form_data['model'], line)
return f'data: {JSONCodec.dumps(line)}\n\n' return f'data: {json.dumps(line)}\n\n'
def get_pipe_id(form_data: dict) -> str: def get_pipe_id(form_data: dict) -> str:
pipe_id = form_data['model'] pipe_id = form_data['model']
@ -212,13 +203,6 @@ async def generate_function_chat_completion(request, form_data, user, models: di
return params return params
# Set server-side by utils/chat.py, never by client input. Mirrors the routers.
bypass_system_prompt = getattr(request.state, 'bypass_system_prompt', False)
# Copy so the base-model substitution below doesn't leak into the caller's
# payload, which the tool-call continuation re-submits. Mirrors the routers.
form_data = {**form_data}
model_id = form_data.get('model') model_id = form_data.get('model')
model_info = await Models.get_model_by_id(model_id) model_info = await Models.get_model_by_id(model_id)
@ -294,8 +278,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
if params: if params:
system = params.pop('system', None) system = params.pop('system', None)
form_data = apply_model_params_to_body_openai(params, form_data) form_data = apply_model_params_to_body_openai(params, form_data)
if not bypass_system_prompt: form_data = await apply_system_prompt_to_body(system, form_data, metadata, user)
form_data = await apply_system_prompt_to_body(system, form_data, metadata, user)
pipe_id = get_pipe_id(form_data) pipe_id = get_pipe_id(form_data)
function_module = await get_function_module_by_id(request, pipe_id) function_module = await get_function_module_by_id(request, pipe_id)
@ -315,17 +298,17 @@ async def generate_function_chat_completion(request, form_data, user, models: di
yield data yield data
return return
if isinstance(res, dict): if isinstance(res, dict):
yield f'data: {JSONCodec.dumps(res)}\n\n' yield f'data: {json.dumps(res)}\n\n'
return return
except Exception as e: except Exception as e:
log.error(f'Error: {e}') log.error(f'Error: {e}')
yield f'data: {JSONCodec.dumps({"error": {"detail": str(e)}})}\n\n' yield f'data: {json.dumps({"error": {"detail": str(e)}})}\n\n'
return return
if isinstance(res, str): if isinstance(res, str):
message = openai_chat_chunk_message_template(form_data['model'], res) message = openai_chat_chunk_message_template(form_data['model'], res)
yield f'data: {JSONCodec.dumps(message)}\n\n' yield f'data: {json.dumps(message)}\n\n'
if isinstance(res, Iterator): if isinstance(res, Iterator):
for line in res: for line in res:
@ -337,7 +320,7 @@ async def generate_function_chat_completion(request, form_data, user, models: di
finish_message = openai_chat_chunk_message_template(form_data['model'], '') finish_message = openai_chat_chunk_message_template(form_data['model'], '')
finish_message['choices'][0]['finish_reason'] = 'stop' finish_message['choices'][0]['finish_reason'] = 'stop'
yield f'data: {JSONCodec.dumps(finish_message)}\n\n' yield f'data: {json.dumps(finish_message)}\n\n'
yield 'data: [DONE]' yield 'data: [DONE]'
return StreamingResponse(stream_content(), media_type='text/event-stream') return StreamingResponse(stream_content(), media_type='text/event-stream')

View file

@ -0,0 +1,265 @@
"""Database-backed configuration with environment variable defaults."""
from __future__ import annotations
import asyncio
import json
import logging
from datetime import datetime
from functools import reduce
from typing import Any, Optional, Union
import redis
from open_webui.internal.db import Base, get_async_db, get_db
from open_webui.utils.redis import get_redis_connection
from sqlalchemy import JSON, Column, DateTime, Integer, func, select
log = logging.getLogger(__name__)
# ── Model ────────────────────────────────────────────────────────────────────
class ConfigTable(Base):
__tablename__ = 'config'
id = Column(Integer, primary_key=True)
data = Column(JSON, nullable=False)
version = Column(Integer, nullable=False, default=0)
created_at = Column(DateTime, nullable=False, server_default=func.now())
updated_at = Column(DateTime, nullable=True, onupdate=func.now())
# ── Blob ─────────────────────────────────────────────────────────────────────
class ConfigState:
"""In-memory mirror of the single-row config JSON blob."""
__slots__ = ('_data',)
def __init__(self) -> None:
self._data: dict[str, Any] = {}
@property
def snapshot(self) -> dict:
return self._data
def read(self, path: str) -> Any:
return reduce(
lambda n, k: n.get(k) if isinstance(n, dict) else None,
path.split('.'),
self._data,
)
def write(self, path: str, value: Any) -> None:
keys = path.split('.')
reduce(lambda d, k: d.setdefault(k, {}), keys[:-1], self._data)[keys[-1]] = value
def replace(self, data: dict) -> None:
self._data = data
def load(self) -> dict:
with get_db() as db:
row = db.query(ConfigTable).order_by(ConfigTable.id.desc()).first()
self._data = row.data if row else {'version': 0, 'ui': {}}
return self._data
def persist(self, data: dict | None = None) -> None:
if data is not None:
self._data = data
with get_db() as db:
row = db.query(ConfigTable).first()
if row is None:
db.add(ConfigTable(data=self._data, version=0))
else:
row.data, row.updated_at = self._data, datetime.now()
db.add(row)
db.commit()
async def persist_async(self, data: dict | None = None) -> None:
if data is not None:
self._data = data
async with get_async_db() as db:
result = await db.execute(select(ConfigTable).limit(1))
row = result.scalars().first()
if row is None:
db.add(ConfigTable(data=self._data, version=0))
else:
row.data, row.updated_at = self._data, datetime.now()
db.add(row)
await db.commit()
def clear(self) -> None:
with get_db() as db:
db.query(ConfigTable).delete()
db.commit()
async def clear_async(self) -> None:
from sqlalchemy import delete as sa_delete
async with get_async_db() as db:
await db.execute(sa_delete(ConfigTable))
await db.commit()
STATE = ConfigState()
# ── ConfigVar ──────────────────────────────────────────────────────────────────
_persist_enabled: bool = True
_oauth_persist_enabled: bool = False
_all_configs: list[ConfigVar] = []
def initialize(*, enable_persistent: bool = True, enable_oauth_persistent: bool = False) -> dict:
global _persist_enabled, _oauth_persist_enabled
_persist_enabled = enable_persistent
_oauth_persist_enabled = enable_oauth_persistent
return STATE.load()
class ConfigVar:
__slots__ = ('env_name', 'config_path', 'env_value', 'config_value', 'value')
def __init__(self, env_name: str, config_path: str, env_value: Any) -> None:
self.env_name = env_name
self.config_path = config_path
self.env_value = env_value
self.config_value = STATE.read(config_path)
if self.config_value is not None and _persist_enabled:
if config_path.startswith('oauth.') and not _oauth_persist_enabled:
log.info("Skipping DB value for '%s' (OAuth persistence disabled)", env_name)
self.value = env_value
else:
log.info("'%s' loaded from database", env_name)
self.value = self.config_value
else:
self.value = env_value
_all_configs.append(self)
def __str__(self) -> str:
return str(self.value)
def __repr__(self) -> str:
return f'<ConfigVar {self.env_name}={self.value!r}>'
@property
def __dict__(self): # type: ignore[override]
raise TypeError(f"ConfigVar('{self.env_name}') cannot be cast to dict; use .value")
def __getattribute__(self, item: str):
if item == '__dict__':
raise TypeError('ConfigVar cannot be cast to dict; use .value')
return super().__getattribute__(item)
def refresh(self) -> None:
current = STATE.read(self.config_path)
if current is not None:
self.value = current
log.info('Refreshed %s → %s', self.env_name, self.value)
def commit(self) -> None:
log.info("Persisting '%s'", self.env_name)
STATE.write(self.config_path, self.value)
self.config_value = self.value
STATE.persist()
async def commit_async(self) -> None:
log.info("Persisting '%s'", self.env_name)
STATE.write(self.config_path, self.value)
self.config_value = self.value
await STATE.persist_async()
# ── AppConfig ──────────────────────────────────────────────────────────
class AppConfig:
"""Attribute-style container for ConfigVars with optional Redis sync."""
def __init__(
self,
*,
redis_url: Optional[str] = None,
redis_sentinels: Optional[list] = None,
redis_cluster: bool = False,
redis_key_prefix: str = 'open-webui',
) -> None:
super().__setattr__('_entries', {})
super().__setattr__('_key_prefix', redis_key_prefix)
# If sentinels weren't explicitly provided, read from env.
if redis_sentinels is None:
from open_webui.env import REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT
from open_webui.utils.redis import get_sentinels_from_env
redis_sentinels = get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT)
rc: Union[redis.Redis, redis.cluster.RedisCluster, None] = None
if redis_url:
rc = get_redis_connection(redis_url, redis_sentinels or [], redis_cluster, decode_responses=True)
super().__setattr__('_rc', rc)
def __setattr__(self, name: str, value: Any) -> None:
entries: dict = super().__getattribute__('_entries')
if isinstance(value, ConfigVar):
entries[name] = value
return
entries[name].value = value
try:
asyncio.get_running_loop().create_task(self._write_async(name))
except RuntimeError:
entries[name].commit()
rc = super().__getattribute__('_rc')
if rc and _persist_enabled:
prefix = super().__getattribute__('_key_prefix')
try:
rc.set(f'{prefix}:config:{name}', json.dumps(entries[name].value))
except Exception as exc:
log.error("Redis write failed for '%s': %s", name, exc)
async def _write_async(self, name: str) -> None:
try:
await self._entries[name].commit_async()
except Exception as exc:
log.error("Async persist failed for '%s': %s", name, exc)
def __getattr__(self, name: str) -> Any:
entries = super().__getattribute__('_entries')
if name not in entries:
raise AttributeError(f"No config key '{name}'")
rc = super().__getattribute__('_rc')
if rc and _persist_enabled:
prefix = super().__getattribute__('_key_prefix')
try:
raw = rc.get(f'{prefix}:config:{name}')
if raw is not None:
decoded = json.loads(raw)
if entries[name].value != decoded:
entries[name].value = decoded
log.info("Updated '%s' from Redis", name)
except Exception as exc:
log.error("Redis read failed for '%s': %s", name, exc)
return entries[name].value
def _sync_to_redis(self) -> None:
rc = super().__getattribute__('_rc')
if not rc or not _persist_enabled:
return
prefix = super().__getattribute__('_key_prefix')
for name, s in super().__getattribute__('_entries').items():
try:
rc.set(f'{prefix}:config:{name}', json.dumps(s.value))
except Exception as exc:
log.error("Redis sync failed for '%s': %s", name, exc)

View file

@ -1,16 +1,14 @@
from __future__ import annotations from __future__ import annotations
import json
import logging import logging
import os import os
import re
import sys import sys
from contextlib import asynccontextmanager, contextmanager from contextlib import asynccontextmanager, contextmanager
from datetime import datetime, timedelta, timezone
from typing import Any, Optional from typing import Any, Optional
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
from open_webui.env import ( from open_webui.env import (
DATABASE_ENABLE_IAM_TOKEN_AUTH,
DATABASE_ENABLE_SESSION_SHARING, DATABASE_ENABLE_SESSION_SHARING,
DATABASE_ENABLE_SQLITE_WAL, DATABASE_ENABLE_SQLITE_WAL,
DATABASE_POOL_MAX_OVERFLOW, DATABASE_POOL_MAX_OVERFLOW,
@ -27,11 +25,8 @@ from open_webui.env import (
DATABASE_URL, DATABASE_URL,
ENABLE_DB_MIGRATIONS, ENABLE_DB_MIGRATIONS,
OPEN_WEBUI_DIR, OPEN_WEBUI_DIR,
USE_SLIM,
) )
from open_webui.utils.json_codec import JSONCodec
from sqlalchemy import Dialect, MetaData, create_engine, event, types from sqlalchemy import Dialect, MetaData, create_engine, event, types
from sqlalchemy.engine.url import make_url
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import Session, scoped_session, sessionmaker from sqlalchemy.orm import Session, scoped_session, sessionmaker
@ -126,34 +121,23 @@ class JSONField(types.TypeDecorator): # TEXT-backed JSON storage
"""Store arbitrary Python objects as JSON-encoded TEXT. """Store arbitrary Python objects as JSON-encoded TEXT.
Used instead of native JSON columns for portability across SQLite and Used instead of native JSON columns for portability across SQLite and
PostgreSQL. Values are serialized with ``JSONCodec.dumps`` on write and PostgreSQL. Values are serialized with ``json.dumps`` on write and
deserialized with ``JSONCodec.loads`` on read. deserialized with ``json.loads`` on read.
""" """
impl = types.UnicodeText impl = types.UnicodeText
cache_ok = True cache_ok = True
def process_bind_param(self, value: _T | None, dialect: Dialect) -> Any: def process_bind_param(self, value: _T | None, dialect: Dialect) -> Any:
return JSONCodec.dumps(value) if value is not None else None return json.dumps(value) if value is not None else None
def process_result_value(self, value: _T | None, dialect: Dialect) -> Any: def process_result_value(self, value: _T | None, dialect: Dialect) -> Any:
return JSONCodec.loads(value) if value is not None else None return json.loads(value) if value is not None else None
def copy(self, **kwargs: Any) -> Self: def copy(self, **kwargs: Any) -> Self:
return JSONField(length=self.impl.length) return JSONField(length=self.impl.length)
if USE_SLIM:
if make_url(DATABASE_URL).get_backend_name() not in ('sqlite', 'postgresql', 'postgres'):
raise ValueError(
'Slim requires SQLite or PostgreSQL for DATABASE_URL. Use the standard image for other databases.'
)
if DATABASE_ENABLE_IAM_TOKEN_AUTH:
raise ValueError(
'AWS RDS IAM authentication requires the standard image. Slim supports PostgreSQL database credentials.'
)
# Normalize SSL params from the URL once; the sync engine needs them # Normalize SSL params from the URL once; the sync engine needs them
# reattached in canonical libpq form for psycopg2. # reattached in canonical libpq form for psycopg2.
_url_without_ssl, _ssl_dict = extract_ssl_params_from_url(DATABASE_URL) _url_without_ssl, _ssl_dict = extract_ssl_params_from_url(DATABASE_URL)
@ -162,77 +146,6 @@ _url_without_ssl, _ssl_dict = extract_ssl_params_from_url(DATABASE_URL)
SQLALCHEMY_DATABASE_URL = reattach_ssl_params_to_url(_url_without_ssl, _ssl_dict) if _ssl_dict else DATABASE_URL SQLALCHEMY_DATABASE_URL = reattach_ssl_params_to_url(_url_without_ssl, _ssl_dict) if _ssl_dict else DATABASE_URL
class RDSIAMTokenAuth:
_refresh_after = timedelta(minutes=14)
def __init__(self, database_url: str) -> None:
url = make_url(database_url)
if not url.drivername.startswith(('postgresql', 'postgres')):
raise ValueError('DATABASE_ENABLE_IAM_TOKEN_AUTH is only supported for PostgreSQL databases')
if not url.host or not url.username:
raise ValueError('DATABASE_ENABLE_IAM_TOKEN_AUTH requires a database host and user')
self.host = url.host
self.port = url.port or 5432
self.username = url.username
self._client = None
self._token: str | None = None
self._expires_at = datetime.min.replace(tzinfo=timezone.utc)
@property
def client(self):
if self._client is None:
import boto3
self._client = boto3.client('rds')
return self._client
def get_password(self) -> str:
now = datetime.now(timezone.utc)
if self._token and now < self._expires_at:
return self._token
self._token = self.client.generate_db_auth_token(
DBHostname=self.host,
Port=self.port,
DBUsername=self.username,
)
self._expires_at = now + self._refresh_after
log.info('AWS RDS IAM database token refreshed; next refresh after %s', self._expires_at.isoformat())
return self._token
_rds_iam_token_auth = RDSIAMTokenAuth(SQLALCHEMY_DATABASE_URL) if DATABASE_ENABLE_IAM_TOKEN_AUTH else None
def _set_iam_token_password(dialect, conn_rec, cargs, cparams):
if _rds_iam_token_auth is not None:
cparams['password'] = _rds_iam_token_auth.get_password()
def enable_iam_token_auth(connectable) -> None:
if _rds_iam_token_auth is None:
return
engine = getattr(connectable, 'sync_engine', connectable)
url = engine.url
auth = _rds_iam_token_auth
# The token is bound to one host/port/user pair; leave other databases on their own credentials.
if (url.host, url.port or 5432, url.username) != (auth.host, auth.port, auth.username):
log.warning(
'AWS RDS IAM token auth not applied to %s: the token is issued for %s@%s:%s, '
'so this connection uses the password from its own URL',
url.render_as_string(hide_password=True),
auth.username,
auth.host,
auth.port,
)
return
if not event.contains(engine, 'do_connect', _set_iam_token_password):
event.listen(engine, 'do_connect', _set_iam_token_password)
def _make_async_url(url: str) -> str: def _make_async_url(url: str) -> str:
"""Convert a sync database URL to its async driver equivalent. """Convert a sync database URL to its async driver equivalent.
@ -259,27 +172,6 @@ def _make_async_url(url: str) -> str:
return url return url
def _json_codec_kwargs(kwargs: dict) -> dict:
"""Default an engine to JSONCodec for native ``JSON`` columns.
Unlike ``JSONField``, those serialize through the engine, which otherwise uses
stdlib ``json``. With ``ENABLE_ORJSON`` off JSONCodec is stdlib ``json`` anyway.
"""
kwargs.setdefault('json_serializer', JSONCodec.dumps)
kwargs.setdefault('json_deserializer', JSONCodec.loads)
return kwargs
def _create_engine(*args, **kwargs):
"""``create_engine`` with the app JSON codec wired in."""
return create_engine(*args, **_json_codec_kwargs(kwargs))
def _create_async_engine(*args, **kwargs):
"""``create_async_engine`` with the app JSON codec wired in."""
return create_async_engine(*args, **_json_codec_kwargs(kwargs))
# ============================================================ # ============================================================
# SYNC ENGINE (used only for: startup migrations, config loading, # SYNC ENGINE (used only for: startup migrations, config loading,
# Alembic, peewee migration, health checks) # Alembic, peewee migration, health checks)
@ -308,7 +200,7 @@ if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'):
# in the native sqlcipher3 C library. Use NullPool by default for safety, # in the native sqlcipher3 C library. Use NullPool by default for safety,
# or QueuePool if DATABASE_POOL_SIZE is explicitly configured. # or QueuePool if DATABASE_POOL_SIZE is explicitly configured.
if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0: if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0:
engine = _create_engine( engine = create_engine(
'sqlite://', 'sqlite://',
creator=create_sqlcipher_connection, creator=create_sqlcipher_connection,
pool_size=DATABASE_POOL_SIZE, pool_size=DATABASE_POOL_SIZE,
@ -320,7 +212,7 @@ if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'):
echo=False, echo=False,
) )
else: else:
engine = _create_engine( engine = create_engine(
'sqlite://', 'sqlite://',
creator=create_sqlcipher_connection, creator=create_sqlcipher_connection,
poolclass=NullPool, poolclass=NullPool,
@ -330,49 +222,10 @@ if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'):
log.info('Connected to encrypted SQLite database using SQLCipher') log.info('Connected to encrypted SQLite database using SQLCipher')
elif 'sqlite' in SQLALCHEMY_DATABASE_URL: elif 'sqlite' in SQLALCHEMY_DATABASE_URL:
engine = _create_engine(SQLALCHEMY_DATABASE_URL, connect_args={'check_same_thread': False}) engine = create_engine(SQLALCHEMY_DATABASE_URL, connect_args={'check_same_thread': False})
def _apply_sqlite_pragmas(dbapi_connection): def _apply_sqlite_pragmas(dbapi_connection):
"""Apply all configured SQLite PRAGMAs to a raw DBAPI connection.""" """Apply all configured SQLite PRAGMAs to a raw DBAPI connection."""
# SQLite LIKE folds ASCII only; SQLAlchemy SQLite ILIKE compiles to lower(x) LIKE lower(?).
compiled_patterns = {}
def like(pattern, value, escape=None):
if pattern is None or value is None:
return None
pattern = str(pattern).lower()
escape = str(escape).lower() if escape is not None else None
key = (pattern, escape)
compiled = compiled_patterns.get(key)
if compiled is False:
return False
if compiled is None:
regex = []
escaped = False
for char in pattern:
if escape and not escaped and char == escape:
escaped = True
continue
regex.append(
'.*' if not escaped and char == '%' else '.' if not escaped and char == '_' else re.escape(char)
)
escaped = False
if escaped:
compiled = False
if len(compiled_patterns) >= 512:
compiled_patterns.clear()
compiled_patterns[key] = compiled
return False
compiled = re.compile(''.join(regex), re.DOTALL)
if len(compiled_patterns) >= 512:
compiled_patterns.clear()
compiled_patterns[key] = compiled
return compiled.fullmatch(str(value).lower()) is not None
dbapi_connection.create_function('like', 2, like, deterministic=True)
dbapi_connection.create_function('like', 3, like, deterministic=True)
cursor = dbapi_connection.cursor() cursor = dbapi_connection.cursor()
if DATABASE_ENABLE_SQLITE_WAL: if DATABASE_ENABLE_SQLITE_WAL:
cursor.execute('PRAGMA journal_mode=WAL') cursor.execute('PRAGMA journal_mode=WAL')
@ -401,7 +254,7 @@ elif 'sqlite' in SQLALCHEMY_DATABASE_URL:
else: else:
if isinstance(DATABASE_POOL_SIZE, int): if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0: if DATABASE_POOL_SIZE > 0:
engine = _create_engine( engine = create_engine(
SQLALCHEMY_DATABASE_URL, SQLALCHEMY_DATABASE_URL,
pool_size=DATABASE_POOL_SIZE, pool_size=DATABASE_POOL_SIZE,
max_overflow=DATABASE_POOL_MAX_OVERFLOW, max_overflow=DATABASE_POOL_MAX_OVERFLOW,
@ -411,11 +264,9 @@ else:
poolclass=QueuePool, poolclass=QueuePool,
) )
else: else:
engine = _create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool) engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool)
else: else:
engine = _create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True) engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True)
enable_iam_token_auth(engine)
# Sync session — used ONLY for startup config loading (config.py runs at import time) # Sync session — used ONLY for startup config loading (config.py runs at import time)
@ -457,15 +308,14 @@ if sys.platform == 'win32' and _is_postgres_url(DATABASE_URL):
if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL: if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
# Generous default — async coroutines + no session sharing = high connection demand. # Generous default — async coroutines + no session sharing = high connection demand.
# No pool_pre_ping: a local SQLite file cannot drop connections, and the
# ping costs a worker-thread hop plus a SELECT 1 on every checkout.
_sqlite_pool_size = DATABASE_POOL_SIZE if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0 else 512 _sqlite_pool_size = DATABASE_POOL_SIZE if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0 else 512
async_engine = _create_async_engine( async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL, ASYNC_SQLALCHEMY_DATABASE_URL,
connect_args={'check_same_thread': False}, connect_args={'check_same_thread': False},
pool_size=_sqlite_pool_size, pool_size=_sqlite_pool_size,
pool_timeout=DATABASE_POOL_TIMEOUT, pool_timeout=DATABASE_POOL_TIMEOUT,
pool_recycle=DATABASE_POOL_RECYCLE, pool_recycle=DATABASE_POOL_RECYCLE,
pool_pre_ping=True,
) )
@event.listens_for(async_engine.sync_engine, 'connect') @event.listens_for(async_engine.sync_engine, 'connect')
@ -474,7 +324,7 @@ if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
else: else:
if isinstance(DATABASE_POOL_SIZE, int): if isinstance(DATABASE_POOL_SIZE, int):
if DATABASE_POOL_SIZE > 0: if DATABASE_POOL_SIZE > 0:
async_engine = _create_async_engine( async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL, ASYNC_SQLALCHEMY_DATABASE_URL,
pool_size=DATABASE_POOL_SIZE, pool_size=DATABASE_POOL_SIZE,
max_overflow=DATABASE_POOL_MAX_OVERFLOW, max_overflow=DATABASE_POOL_MAX_OVERFLOW,
@ -483,19 +333,17 @@ else:
pool_pre_ping=True, pool_pre_ping=True,
) )
else: else:
async_engine = _create_async_engine( async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL, ASYNC_SQLALCHEMY_DATABASE_URL,
pool_pre_ping=True, pool_pre_ping=True,
poolclass=NullPool, poolclass=NullPool,
) )
else: else:
async_engine = _create_async_engine( async_engine = create_async_engine(
ASYNC_SQLALCHEMY_DATABASE_URL, ASYNC_SQLALCHEMY_DATABASE_URL,
pool_pre_ping=True, pool_pre_ping=True,
) )
enable_iam_token_auth(async_engine)
AsyncSessionLocal = async_sessionmaker( AsyncSessionLocal = async_sessionmaker(
bind=async_engine, bind=async_engine,

File diff suppressed because it is too large Load diff

View file

@ -6,11 +6,9 @@ import logging.config
import logging import logging
import alembic.context import alembic.context
from open_webui.env import DATABASE_PASSWORD, DATABASE_URL, LOG_FORMAT from open_webui.env import DATABASE_PASSWORD, DATABASE_URL, LOG_FORMAT
from open_webui.internal.db import enable_iam_token_auth, extract_ssl_params_from_url, reattach_ssl_params_to_url 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.auths import Auth
from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401 from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401
from open_webui.models.chat_messages import ChatMessage # noqa: F401
from open_webui.models.chats import Chat # noqa: F401
from sqlalchemy import create_engine, engine_from_config, pool from sqlalchemy import create_engine, engine_from_config, pool
alembic_config = alembic.context.config alembic_config = alembic.context.config
@ -70,7 +68,6 @@ def _get_engine_connectable():
def run_migrations_online() -> None: def run_migrations_online() -> None:
"""Execute migrations against a live database connection.""" """Execute migrations against a live database connection."""
live_connectable = _get_engine_connectable() live_connectable = _get_engine_connectable()
enable_iam_token_auth(live_connectable)
with live_connectable.connect() as live_connection: with live_connectable.connect() as live_connection:
alembic.context.configure( alembic.context.configure(
connection=live_connection, connection=live_connection,

View file

@ -1,28 +0,0 @@
"""Add group_member user_id index
Revision ID: 1ce6ade7d93b
Revises: f0bd01a18a3d
Create Date: 2026-07-31 03:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
revision = '1ce6ade7d93b'
down_revision = 'f0bd01a18a3d'
branch_labels = None
depends_on = None
def upgrade():
conn = op.get_bind()
inspector = sa.inspect(conn)
existing_indexes = {idx['name'] for idx in inspector.get_indexes('group_member')}
if 'ix_group_member_user_id_group_id' not in existing_indexes:
op.create_index('ix_group_member_user_id_group_id', 'group_member', ['user_id', 'group_id'])
def downgrade():
op.drop_index('ix_group_member_user_id_group_id', table_name='group_member')

View file

@ -49,7 +49,7 @@ def upgrade():
# Step 3: Migrate data from 'old_chat' to 'chat' (only if old_chat exists) # Step 3: Migrate data from 'old_chat' to 'chat' (only if old_chat exists)
# Re-check columns after potential rename above # Re-check columns after potential rename above
current_cols = {c['name'] for c in sa.inspect(conn).get_columns('chat')} current_cols = {c['name'] for c in inspector.get_columns('chat')}
if 'old_chat' in current_cols: if 'old_chat' in current_cols:
chat_table = table( chat_table = table(
'chat', 'chat',
@ -76,12 +76,8 @@ def upgrade():
def downgrade(): def downgrade():
conn = op.get_bind()
columns = {col['name'] for col in sa.inspect(conn).get_columns('chat')}
# Step 1: Add 'old_chat' column back as Text # Step 1: Add 'old_chat' column back as Text
if 'old_chat' not in columns: op.add_column('chat', sa.Column('old_chat', sa.Text(), nullable=True))
op.add_column('chat', sa.Column('old_chat', sa.Text(), nullable=True))
# Step 2: Convert 'chat' JSON data back to text and store in 'old_chat' # Step 2: Convert 'chat' JSON data back to text and store in 'old_chat'
chat_table = table( chat_table = table(
@ -91,14 +87,14 @@ def downgrade():
sa.Column('old_chat', sa.Text()), sa.Column('old_chat', sa.Text()),
) )
if 'chat' in columns: connection = op.get_bind()
results = conn.execute(select(chat_table.c.id, chat_table.c.chat)) results = connection.execute(select(chat_table.c.id, chat_table.c.chat))
for row in results: for row in results:
text_data = json.dumps(row.chat) if row.chat is not None else None text_data = json.dumps(row.chat) if row.chat is not None else None
conn.execute(sa.update(chat_table).where(chat_table.c.id == row.id).values(old_chat=text_data)) connection.execute(sa.update(chat_table).where(chat_table.c.id == row.id).values(old_chat=text_data))
# Step 3: Remove the new 'chat' JSON column # Step 3: Remove the new 'chat' JSON column
op.drop_column('chat', 'chat') op.drop_column('chat', 'chat')
# Step 4: Rename 'old_chat' back to 'chat' # Step 4: Rename 'old_chat' back to 'chat'
op.alter_column('chat', 'old_chat', new_column_name='chat', existing_type=sa.Text()) op.alter_column('chat', 'old_chat', new_column_name='chat', existing_type=sa.Text())

View file

@ -1,584 +0,0 @@
"""reshape config to per key rows
Revision ID: 3ff2c63645b8
Revises: 461111b60977
Create Date: 2026-06-17 00:50:51.477073
"""
import json
import time
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision: str = '3ff2c63645b8'
down_revision: Union[str, None] = '461111b60977'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
# Maps every dot-notation blob path to its legacy env/config key name.
# Built from the legacy persistent config declarations in config.py.
BLOB_PATH_TO_KEY = {
'audio.stt.allowed_extensions': 'AUDIO_STT_ALLOWED_EXTENSIONS',
'audio.stt.azure.api_key': 'AUDIO_STT_AZURE_API_KEY',
'audio.stt.azure.base_url': 'AUDIO_STT_AZURE_BASE_URL',
'audio.stt.azure.locales': 'AUDIO_STT_AZURE_LOCALES',
'audio.stt.azure.max_speakers': 'AUDIO_STT_AZURE_MAX_SPEAKERS',
'audio.stt.azure.region': 'AUDIO_STT_AZURE_REGION',
'audio.stt.deepgram.api_key': 'DEEPGRAM_API_KEY',
'audio.stt.engine': 'AUDIO_STT_ENGINE',
'audio.stt.mistral.api_base_url': 'AUDIO_STT_MISTRAL_API_BASE_URL',
'audio.stt.mistral.api_key': 'AUDIO_STT_MISTRAL_API_KEY',
'audio.stt.mistral.use_chat_completions': 'AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS',
'audio.stt.model': 'AUDIO_STT_MODEL',
'audio.stt.openai.api_base_url': 'AUDIO_STT_OPENAI_API_BASE_URL',
'audio.stt.openai.api_key': 'AUDIO_STT_OPENAI_API_KEY',
'audio.stt.supported_content_types': 'AUDIO_STT_SUPPORTED_CONTENT_TYPES',
'audio.stt.whisper_model': 'WHISPER_MODEL',
'audio.tts.api_key': 'AUDIO_TTS_API_KEY',
'audio.tts.azure.speech_base_url': 'AUDIO_TTS_AZURE_SPEECH_BASE_URL',
'audio.tts.azure.speech_output_format': 'AUDIO_TTS_AZURE_SPEECH_OUTPUT_FORMAT',
'audio.tts.azure.speech_region': 'AUDIO_TTS_AZURE_SPEECH_REGION',
'audio.tts.engine': 'AUDIO_TTS_ENGINE',
'audio.tts.mistral.api_base_url': 'AUDIO_TTS_MISTRAL_API_BASE_URL',
'audio.tts.mistral.api_key': 'AUDIO_TTS_MISTRAL_API_KEY',
'audio.tts.model': 'AUDIO_TTS_MODEL',
'audio.tts.openai.api_base_url': 'AUDIO_TTS_OPENAI_API_BASE_URL',
'audio.tts.openai.api_key': 'AUDIO_TTS_OPENAI_API_KEY',
'audio.tts.openai.params': 'AUDIO_TTS_OPENAI_PARAMS',
'audio.tts.split_on': 'AUDIO_TTS_SPLIT_ON',
'audio.tts.voice': 'AUDIO_TTS_VOICE',
'auth.admin.email': 'ADMIN_EMAIL',
'auth.admin.show': 'SHOW_ADMIN_DETAILS',
'auth.api_key.allowed_endpoints': 'API_KEYS_ALLOWED_ENDPOINTS',
'auth.api_key.endpoint_restrictions': 'ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS',
'auth.enable_api_keys': 'ENABLE_API_KEYS',
'auth.jwt_expiry': 'JWT_EXPIRES_IN',
'automations.enable': 'ENABLE_AUTOMATIONS',
'automations.max_count': 'AUTOMATION_MAX_COUNT',
'automations.min_interval': 'AUTOMATION_MIN_INTERVAL',
'calendar.enable': 'ENABLE_CALENDAR',
'channels.enable': 'ENABLE_CHANNELS',
'code_execution.enable': 'ENABLE_CODE_EXECUTION',
'code_execution.engine': 'CODE_EXECUTION_ENGINE',
'code_execution.jupyter.auth': 'CODE_EXECUTION_JUPYTER_AUTH',
'code_execution.jupyter.auth_password': 'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD',
'code_execution.jupyter.auth_token': 'CODE_EXECUTION_JUPYTER_AUTH_TOKEN',
'code_execution.jupyter.timeout': 'CODE_EXECUTION_JUPYTER_TIMEOUT',
'code_execution.jupyter.url': 'CODE_EXECUTION_JUPYTER_URL',
'code_interpreter.enable': 'ENABLE_CODE_INTERPRETER',
'code_interpreter.engine': 'CODE_INTERPRETER_ENGINE',
'code_interpreter.jupyter.auth': 'CODE_INTERPRETER_JUPYTER_AUTH',
'code_interpreter.jupyter.auth_password': 'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD',
'code_interpreter.jupyter.auth_token': 'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN',
'code_interpreter.jupyter.timeout': 'CODE_INTERPRETER_JUPYTER_TIMEOUT',
'code_interpreter.jupyter.url': 'CODE_INTERPRETER_JUPYTER_URL',
'code_interpreter.prompt_template': 'CODE_INTERPRETER_PROMPT_TEMPLATE',
'direct.enable': 'ENABLE_DIRECT_CONNECTIONS',
'evaluation.arena.enable': 'ENABLE_EVALUATION_ARENA_MODELS',
'evaluation.arena.models': 'EVALUATION_ARENA_MODELS',
'file.image_compression_height': 'FILE_IMAGE_COMPRESSION_HEIGHT',
'file.image_compression_width': 'FILE_IMAGE_COMPRESSION_WIDTH',
'folders.enable': 'ENABLE_FOLDERS',
'folders.max_file_count': 'FOLDER_MAX_FILE_COUNT',
'google_drive.api_key': 'GOOGLE_DRIVE_API_KEY',
'google_drive.client_id': 'GOOGLE_DRIVE_CLIENT_ID',
'google_drive.enable': 'ENABLE_GOOGLE_DRIVE_INTEGRATION',
'image_generation.automatic1111.api_auth': 'AUTOMATIC1111_API_AUTH',
'image_generation.automatic1111.api_params': 'AUTOMATIC1111_PARAMS',
'image_generation.automatic1111.base_url': 'AUTOMATIC1111_BASE_URL',
'image_generation.comfyui.api_key': 'COMFYUI_API_KEY',
'image_generation.comfyui.base_url': 'COMFYUI_BASE_URL',
'image_generation.comfyui.nodes': 'COMFYUI_WORKFLOW_NODES',
'image_generation.comfyui.workflow': 'COMFYUI_WORKFLOW',
'image_generation.enable': 'ENABLE_IMAGE_GENERATION',
'image_generation.engine': 'IMAGE_GENERATION_ENGINE',
'image_generation.gemini.api_base_url': 'IMAGES_GEMINI_API_BASE_URL',
'image_generation.gemini.api_key': 'IMAGES_GEMINI_API_KEY',
'image_generation.gemini.endpoint_method': 'IMAGES_GEMINI_ENDPOINT_METHOD',
'image_generation.model': 'IMAGE_GENERATION_MODEL',
'image_generation.openai.api_base_url': 'IMAGES_OPENAI_API_BASE_URL',
'image_generation.openai.api_key': 'IMAGES_OPENAI_API_KEY',
'image_generation.openai.api_version': 'IMAGES_OPENAI_API_VERSION',
'image_generation.openai.params': 'IMAGES_OPENAI_API_PARAMS',
'image_generation.prompt.enable': 'ENABLE_IMAGE_PROMPT_GENERATION',
'image_generation.size': 'IMAGE_SIZE',
'image_generation.steps': 'IMAGE_STEPS',
'images.edit.comfyui.api_key': 'IMAGES_EDIT_COMFYUI_API_KEY',
'images.edit.comfyui.base_url': 'IMAGES_EDIT_COMFYUI_BASE_URL',
'images.edit.comfyui.nodes': 'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES',
'images.edit.comfyui.workflow': 'IMAGES_EDIT_COMFYUI_WORKFLOW',
'images.edit.enable': 'ENABLE_IMAGE_EDIT',
'images.edit.engine': 'IMAGE_EDIT_ENGINE',
'images.edit.gemini.api_base_url': 'IMAGES_EDIT_GEMINI_API_BASE_URL',
'images.edit.gemini.api_key': 'IMAGES_EDIT_GEMINI_API_KEY',
'images.edit.model': 'IMAGE_EDIT_MODEL',
'images.edit.openai.api_base_url': 'IMAGES_EDIT_OPENAI_API_BASE_URL',
'images.edit.openai.api_key': 'IMAGES_EDIT_OPENAI_API_KEY',
'images.edit.openai.api_version': 'IMAGES_EDIT_OPENAI_API_VERSION',
'images.edit.size': 'IMAGE_EDIT_SIZE',
'ldap.enable': 'ENABLE_LDAP',
'ldap.group.enable_creation': 'ENABLE_LDAP_GROUP_CREATION',
'ldap.group.enable_management': 'ENABLE_LDAP_GROUP_MANAGEMENT',
'ldap.server.app_dn': 'LDAP_APP_DN',
'ldap.server.app_password': 'LDAP_APP_PASSWORD',
'ldap.server.attribute_for_groups': 'LDAP_ATTRIBUTE_FOR_GROUPS',
'ldap.server.attribute_for_mail': 'LDAP_ATTRIBUTE_FOR_MAIL',
'ldap.server.attribute_for_username': 'LDAP_ATTRIBUTE_FOR_USERNAME',
'ldap.server.ca_cert_file': 'LDAP_CA_CERT_FILE',
'ldap.server.ciphers': 'LDAP_CIPHERS',
'ldap.server.host': 'LDAP_SERVER_HOST',
'ldap.server.label': 'LDAP_SERVER_LABEL',
'ldap.server.port': 'LDAP_SERVER_PORT',
'ldap.server.search_filter': 'LDAP_SEARCH_FILTER',
'ldap.server.use_tls': 'LDAP_USE_TLS',
'ldap.server.users_dn': 'LDAP_SEARCH_BASE',
'ldap.server.validate_cert': 'LDAP_VALIDATE_CERT',
'memories.enable': 'ENABLE_MEMORIES',
'models.base_models_cache': 'ENABLE_BASE_MODELS_CACHE',
'models.default_metadata': 'DEFAULT_MODEL_METADATA',
'models.default_params': 'DEFAULT_MODEL_PARAMS',
'notes.enable': 'ENABLE_NOTES',
# OAuth — direct paths
'oauth.admin_roles': 'OAUTH_ADMIN_ROLES',
'oauth.allowed_domains': 'OAUTH_ALLOWED_DOMAINS',
'oauth.allowed_roles': 'OAUTH_ALLOWED_ROLES',
'oauth.audience': 'OAUTH_AUDIENCE',
'oauth.auto_redirect': 'OAUTH_AUTO_REDIRECT',
'oauth.blocked_groups': 'OAUTH_BLOCKED_GROUPS',
'oauth.client.timeout': 'OAUTH_CLIENT_TIMEOUT',
'oauth.enable_group_creation': 'ENABLE_OAUTH_GROUP_CREATION',
'oauth.enable_group_mapping': 'ENABLE_OAUTH_GROUP_MANAGEMENT',
'oauth.enable_role_mapping': 'ENABLE_OAUTH_ROLE_MANAGEMENT',
'oauth.enable_signup': 'ENABLE_OAUTH_SIGNUP',
'oauth.group_default_share': 'OAUTH_GROUP_DEFAULT_SHARE',
'oauth.merge_accounts_by_email': 'OAUTH_MERGE_ACCOUNTS_BY_EMAIL',
'oauth.refresh_token_include_scope': 'OAUTH_REFRESH_TOKEN_INCLUDE_SCOPE',
'oauth.roles_claim': 'OAUTH_ROLES_CLAIM',
'oauth.update_email_on_login': 'OAUTH_UPDATE_EMAIL_ON_LOGIN',
'oauth.update_name_on_login': 'OAUTH_UPDATE_NAME_ON_LOGIN',
'oauth.update_picture_on_login': 'OAUTH_UPDATE_PICTURE_ON_LOGIN',
# OAuth — generic provider paths
'oauth.client_id': 'OAUTH_CLIENT_ID',
'oauth.client_secret': 'OAUTH_CLIENT_SECRET',
'oauth.code_challenge_method': 'OAUTH_CODE_CHALLENGE_METHOD',
'oauth.email_claim': 'OAUTH_EMAIL_CLAIM',
'oauth.end_session_endpoint': 'OPENID_END_SESSION_ENDPOINT',
'oauth.group_claim': 'OAUTH_GROUP_CLAIM',
'oauth.picture_claim': 'OAUTH_PICTURE_CLAIM',
'oauth.provider_name': 'OAUTH_PROVIDER_NAME',
'oauth.provider_url': 'OPENID_PROVIDER_URL',
'oauth.redirect_uri': 'OPENID_REDIRECT_URI',
'oauth.scopes': 'OAUTH_SCOPES',
'oauth.sub_claim': 'OAUTH_SUB_CLAIM',
'oauth.timeout': 'OAUTH_TIMEOUT',
'oauth.token_endpoint_auth_method': 'OAUTH_TOKEN_ENDPOINT_AUTH_METHOD',
'oauth.username_claim': 'OAUTH_USERNAME_CLAIM',
# OAuth — OIDC nested paths (flattened)
'oauth.oidc.avatar_claim': 'OAUTH_PICTURE_CLAIM',
'oauth.oidc.client_id': 'OAUTH_CLIENT_ID',
'oauth.oidc.client_secret': 'OAUTH_CLIENT_SECRET',
'oauth.oidc.code_challenge_method': 'OAUTH_CODE_CHALLENGE_METHOD',
'oauth.oidc.email_claim': 'OAUTH_EMAIL_CLAIM',
'oauth.oidc.end_session_endpoint': 'OPENID_END_SESSION_ENDPOINT',
'oauth.oidc.group_claim': 'OAUTH_GROUP_CLAIM', # renamed from OAUTH_GROUPS_CLAIM
'oauth.oidc.oauth_timeout': 'OAUTH_TIMEOUT',
'oauth.oidc.provider_name': 'OAUTH_PROVIDER_NAME',
'oauth.oidc.provider_url': 'OPENID_PROVIDER_URL',
'oauth.oidc.redirect_uri': 'OPENID_REDIRECT_URI',
'oauth.oidc.scopes': 'OAUTH_SCOPES',
'oauth.oidc.sub_claim': 'OAUTH_SUB_CLAIM',
'oauth.oidc.token_endpoint_auth_method': 'OAUTH_TOKEN_ENDPOINT_AUTH_METHOD',
'oauth.oidc.username_claim': 'OAUTH_USERNAME_CLAIM',
# OAuth — provider-specific
'oauth.feishu.client_id': 'FEISHU_CLIENT_ID',
'oauth.feishu.client_secret': 'FEISHU_CLIENT_SECRET',
'oauth.feishu.redirect_uri': 'FEISHU_REDIRECT_URI',
'oauth.feishu.scope': 'FEISHU_OAUTH_SCOPE',
'oauth.github.client_id': 'GITHUB_CLIENT_ID',
'oauth.github.client_secret': 'GITHUB_CLIENT_SECRET',
'oauth.github.redirect_uri': 'GITHUB_CLIENT_REDIRECT_URI',
'oauth.github.scope': 'GITHUB_CLIENT_SCOPE',
'oauth.google.client_id': 'GOOGLE_CLIENT_ID',
'oauth.google.client_secret': 'GOOGLE_CLIENT_SECRET',
'oauth.google.redirect_uri': 'GOOGLE_REDIRECT_URI',
'oauth.google.scope': 'GOOGLE_OAUTH_SCOPE',
'oauth.microsoft.client_id': 'MICROSOFT_CLIENT_ID',
'oauth.microsoft.client_secret': 'MICROSOFT_CLIENT_SECRET',
'oauth.microsoft.login_base_url': 'MICROSOFT_CLIENT_LOGIN_BASE_URL',
'oauth.microsoft.picture_url': 'MICROSOFT_CLIENT_PICTURE_URL',
'oauth.microsoft.redirect_uri': 'MICROSOFT_REDIRECT_URI',
'oauth.microsoft.scope': 'MICROSOFT_OAUTH_SCOPE',
'oauth.microsoft.tenant_id': 'MICROSOFT_CLIENT_TENANT_ID',
# Ollama / OpenAI
'ollama.api_configs': 'OLLAMA_API_CONFIGS',
'ollama.base_urls': 'OLLAMA_BASE_URLS',
'ollama.enable': 'ENABLE_OLLAMA_API',
'onedrive.enable': 'ENABLE_ONEDRIVE_INTEGRATION',
'onedrive.sharepoint_tenant_id': 'ONEDRIVE_SHAREPOINT_TENANT_ID',
'onedrive.sharepoint_url': 'ONEDRIVE_SHAREPOINT_URL',
'openai.api_base_urls': 'OPENAI_API_BASE_URLS',
'openai.api_configs': 'OPENAI_API_CONFIGS',
'openai.api_keys': 'OPENAI_API_KEYS',
'openai.enable': 'ENABLE_OPENAI_API',
# RAG
'rag.content_extraction_engine': 'CONTENT_EXTRACTION_ENGINE',
'rag.datalab_marker_use_llm': 'DATALAB_MARKER_USE_LLM',
'rag.mistral_ocr_api_base_url': 'MISTRAL_OCR_API_BASE_URL',
'rag.azure_openai.api_key': 'RAG_AZURE_OPENAI_API_KEY',
'rag.azure_openai.api_version': 'RAG_AZURE_OPENAI_API_VERSION',
'rag.azure_openai.base_url': 'RAG_AZURE_OPENAI_BASE_URL',
'rag.bypass_embedding_and_retrieval': 'BYPASS_EMBEDDING_AND_RETRIEVAL',
'rag.chunk_min_size_target': 'CHUNK_MIN_SIZE_TARGET',
'rag.chunk_overlap': 'CHUNK_OVERLAP',
'rag.chunk_size': 'CHUNK_SIZE',
'rag.datalab_marker_additional_config': 'DATALAB_MARKER_ADDITIONAL_CONFIG',
'rag.datalab_marker_api_base_url': 'DATALAB_MARKER_API_BASE_URL',
'rag.datalab_marker_api_key': 'DATALAB_MARKER_API_KEY',
'rag.datalab_marker_disable_image_extraction': 'DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION',
'rag.datalab_marker_force_ocr': 'DATALAB_MARKER_FORCE_OCR',
'rag.datalab_marker_format_lines': 'DATALAB_MARKER_FORMAT_LINES',
'rag.datalab_marker_output_format': 'DATALAB_MARKER_OUTPUT_FORMAT',
'rag.datalab_marker_paginate': 'DATALAB_MARKER_PAGINATE',
'rag.datalab_marker_skip_cache': 'DATALAB_MARKER_SKIP_CACHE',
'rag.datalab_marker_strip_existing_ocr': 'DATALAB_MARKER_STRIP_EXISTING_OCR',
'rag.docling_api_key': 'DOCLING_API_KEY',
'rag.docling_params': 'DOCLING_PARAMS',
'rag.docling_server_url': 'DOCLING_SERVER_URL',
'rag.document_intelligence_endpoint': 'DOCUMENT_INTELLIGENCE_ENDPOINT',
'rag.document_intelligence_key': 'DOCUMENT_INTELLIGENCE_KEY',
'rag.document_intelligence_model': 'DOCUMENT_INTELLIGENCE_MODEL',
'rag.embedding_batch_size': 'RAG_EMBEDDING_BATCH_SIZE',
'rag.embedding_concurrent_requests': 'RAG_EMBEDDING_CONCURRENT_REQUESTS',
'rag.embedding_engine': 'RAG_EMBEDDING_ENGINE',
'rag.embedding_model': 'RAG_EMBEDDING_MODEL',
'rag.enable_async_embedding': 'ENABLE_ASYNC_EMBEDDING',
'rag.enable_hybrid_search': 'ENABLE_RAG_HYBRID_SEARCH',
'rag.enable_hybrid_search_enriched_texts': 'ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS',
'rag.enable_markdown_header_text_splitter': 'ENABLE_MARKDOWN_HEADER_TEXT_SPLITTER',
'rag.external_document_loader_api_key': 'EXTERNAL_DOCUMENT_LOADER_API_KEY',
'rag.external_document_loader_url': 'EXTERNAL_DOCUMENT_LOADER_URL',
'rag.external_reranker_api_key': 'RAG_EXTERNAL_RERANKER_API_KEY',
'rag.external_reranker_timeout': 'RAG_EXTERNAL_RERANKER_TIMEOUT',
'rag.external_reranker_url': 'RAG_EXTERNAL_RERANKER_URL',
'rag.file.allowed_extensions': 'RAG_ALLOWED_FILE_EXTENSIONS',
'rag.file.max_count': 'RAG_FILE_MAX_COUNT',
'rag.file.max_size': 'RAG_FILE_MAX_SIZE',
'rag.full_context': 'RAG_FULL_CONTEXT',
'rag.hybrid_bm25_weight': 'RAG_HYBRID_BM25_WEIGHT',
'rag.mineru_api_key': 'MINERU_API_KEY',
'rag.mineru_api_mode': 'MINERU_API_MODE',
'rag.mineru_api_timeout': 'MINERU_API_TIMEOUT',
'rag.mineru_api_url': 'MINERU_API_URL',
'rag.mineru_file_extensions': 'MINERU_FILE_EXTENSIONS',
'rag.mineru_params': 'MINERU_PARAMS',
'rag.mistral_ocr_api_key': 'MISTRAL_OCR_API_KEY',
'rag.ollama.key': 'RAG_OLLAMA_API_KEY',
'rag.ollama.url': 'RAG_OLLAMA_BASE_URL',
'rag.openai_api_base_url': 'RAG_OPENAI_API_BASE_URL',
'rag.openai_api_key': 'RAG_OPENAI_API_KEY',
'rag.paddleocr_vl_base_url': 'PADDLEOCR_VL_BASE_URL',
'rag.paddleocr_vl_token': 'PADDLEOCR_VL_TOKEN',
'rag.pdf_extract_images': 'PDF_EXTRACT_IMAGES',
'rag.pdf_loader_mode': 'PDF_LOADER_MODE',
'rag.relevance_threshold': 'RAG_RELEVANCE_THRESHOLD',
'rag.reranking_batch_size': 'RAG_RERANKING_BATCH_SIZE',
'rag.reranking_engine': 'RAG_RERANKING_ENGINE',
'rag.reranking_model': 'RAG_RERANKING_MODEL',
'rag.template': 'RAG_TEMPLATE',
'rag.text_splitter': 'RAG_TEXT_SPLITTER',
'rag.tika_server_url': 'TIKA_SERVER_URL',
'rag.tiktoken_encoding_name': 'TIKTOKEN_ENCODING_NAME',
'rag.top_k': 'RAG_TOP_K',
'rag.top_k_reranker': 'RAG_TOP_K_RERANKER',
# RAG — Web
'rag.web.fetch.max_content_length': 'WEB_FETCH_MAX_CONTENT_LENGTH',
'rag.web.loader.concurrent_requests': 'WEB_LOADER_CONCURRENT_REQUESTS',
'rag.web.loader.engine': 'WEB_LOADER_ENGINE',
'rag.web.loader.external_web_loader_api_key': 'EXTERNAL_WEB_LOADER_API_KEY',
'rag.web.loader.external_web_loader_url': 'EXTERNAL_WEB_LOADER_URL',
'rag.web.loader.firecrawl_api_key': 'FIRECRAWL_API_KEY',
'rag.web.loader.firecrawl_api_url': 'FIRECRAWL_API_BASE_URL',
'rag.web.loader.firecrawl_timeout': 'FIRECRAWL_TIMEOUT',
'rag.web.loader.playwright_timeout': 'PLAYWRIGHT_TIMEOUT',
'rag.web.loader.playwright_ws_url': 'PLAYWRIGHT_WS_URL',
'rag.web.loader.ssl_verification': 'ENABLE_WEB_LOADER_SSL_VERIFICATION',
'rag.web.loader.timeout': 'WEB_LOADER_TIMEOUT',
'rag.web.search.azure_ai_search_api_key': 'AZURE_AI_SEARCH_API_KEY',
'rag.web.search.azure_ai_search_endpoint': 'AZURE_AI_SEARCH_ENDPOINT',
'rag.web.search.azure_ai_search_index_name': 'AZURE_AI_SEARCH_INDEX_NAME',
'rag.web.search.bing_search_v7_endpoint': 'BING_SEARCH_V7_ENDPOINT',
'rag.web.search.bing_search_v7_subscription_key': 'BING_SEARCH_V7_SUBSCRIPTION_KEY',
'rag.web.search.bocha_search_api_key': 'BOCHA_SEARCH_API_KEY',
'rag.web.search.brave_search_api_key': 'BRAVE_SEARCH_API_KEY',
'rag.web.search.brave_search_context_tokens': 'BRAVE_SEARCH_CONTEXT_TOKENS',
'rag.web.search.bypass_embedding_and_retrieval': 'BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL',
'rag.web.search.bypass_web_loader': 'BYPASS_WEB_SEARCH_WEB_LOADER',
'rag.web.search.concurrent_requests': 'WEB_SEARCH_CONCURRENT_REQUESTS',
'rag.web.search.ddgs_backend': 'DDGS_BACKEND',
'rag.web.search.domain.filter_list': 'WEB_SEARCH_DOMAIN_FILTER_LIST',
'rag.web.search.enable': 'ENABLE_WEB_SEARCH',
'rag.web.search.engine': 'WEB_SEARCH_ENGINE',
'rag.web.search.exa_api_key': 'EXA_API_KEY',
'rag.web.search.external_web_search_api_key': 'EXTERNAL_WEB_SEARCH_API_KEY',
'rag.web.search.external_web_search_url': 'EXTERNAL_WEB_SEARCH_URL',
'rag.web.search.google_pse_api_key': 'GOOGLE_PSE_API_KEY',
'rag.web.search.google_pse_engine_id': 'GOOGLE_PSE_ENGINE_ID',
'rag.web.search.jina_api_base_url': 'JINA_API_BASE_URL',
'rag.web.search.jina_api_key': 'JINA_API_KEY',
'rag.web.search.kagi_search_api_key': 'KAGI_SEARCH_API_KEY',
'rag.web.search.linkup_api_key': 'LINKUP_API_KEY',
'rag.web.search.linkup_search_params': 'LINKUP_SEARCH_PARAMS',
'rag.web.search.mojeek_search_api_key': 'MOJEEK_SEARCH_API_KEY',
'rag.web.search.ollama_cloud_api_key': 'OLLAMA_CLOUD_WEB_SEARCH_API_KEY',
'rag.web.search.perplexity_api_key': 'PERPLEXITY_API_KEY',
'rag.web.search.perplexity_model': 'PERPLEXITY_MODEL',
'rag.web.search.perplexity_search_api_url': 'PERPLEXITY_SEARCH_API_URL',
'rag.web.search.perplexity_search_context_usage': 'PERPLEXITY_SEARCH_CONTEXT_USAGE',
'rag.web.search.result_count': 'WEB_SEARCH_RESULT_COUNT',
'rag.web.search.searchapi_api_key': 'SEARCHAPI_API_KEY',
'rag.web.search.searchapi_engine': 'SEARCHAPI_ENGINE',
'rag.web.search.searxng_language': 'SEARXNG_LANGUAGE',
'rag.web.search.searxng_query_url': 'SEARXNG_QUERY_URL',
'rag.web.search.serpapi_api_key': 'SERPAPI_API_KEY',
'rag.web.search.serpapi_engine': 'SERPAPI_ENGINE',
'rag.web.search.serper_api_key': 'SERPER_API_KEY',
'rag.web.search.serply_api_key': 'SERPLY_API_KEY',
'rag.web.search.serpstack_api_key': 'SERPSTACK_API_KEY',
'rag.web.search.serpstack_https': 'SERPSTACK_HTTPS',
'rag.web.search.sougou_api_sid': 'SOUGOU_API_SID',
'rag.web.search.sougou_api_sk': 'SOUGOU_API_SK',
'rag.web.search.tavily_api_key': 'TAVILY_API_KEY',
'rag.web.search.tavily_extract_depth': 'TAVILY_EXTRACT_DEPTH',
'rag.web.search.trust_env': 'WEB_SEARCH_TRUST_ENV',
'rag.web.search.yacy_password': 'YACY_PASSWORD',
'rag.web.search.yacy_query_url': 'YACY_QUERY_URL',
'rag.web.search.yacy_username': 'YACY_USERNAME',
'rag.web.search.yandex_web_search_api_key': 'YANDEX_WEB_SEARCH_API_KEY',
'rag.web.search.yandex_web_search_config': 'YANDEX_WEB_SEARCH_CONFIG',
'rag.web.search.yandex_web_search_url': 'YANDEX_WEB_SEARCH_URL',
'rag.web.search.youcom_api_key': 'YOUCOM_API_KEY',
'rag.youtube_loader_language': 'YOUTUBE_LOADER_LANGUAGE',
'rag.youtube_loader_proxy_url': 'YOUTUBE_LOADER_PROXY_URL',
# Tasks
'task.autocomplete.enable': 'ENABLE_AUTOCOMPLETE_GENERATION',
'task.autocomplete.input_max_length': 'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH',
'task.autocomplete.prompt_template': 'AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE',
'task.follow_up.enable': 'ENABLE_FOLLOW_UP_GENERATION',
'task.follow_up.prompt_template': 'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE',
'task.image.prompt_template': 'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE',
'task.model.default': 'TASK_MODEL',
'task.model.external': 'TASK_MODEL_EXTERNAL',
'task.query.prompt_template': 'QUERY_GENERATION_PROMPT_TEMPLATE',
'task.query.retrieval.enable': 'ENABLE_RETRIEVAL_QUERY_GENERATION',
'task.query.search.enable': 'ENABLE_SEARCH_QUERY_GENERATION',
'task.tags.enable': 'ENABLE_TAGS_GENERATION',
'task.tags.prompt_template': 'TAGS_GENERATION_PROMPT_TEMPLATE',
'task.title.enable': 'ENABLE_TITLE_GENERATION',
'task.title.prompt_template': 'TITLE_GENERATION_PROMPT_TEMPLATE',
'task.tools.prompt_template': 'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE',
'task.voice.prompt.enable': 'ENABLE_VOICE_MODE_PROMPT',
'task.voice.prompt_template': 'VOICE_MODE_PROMPT_TEMPLATE',
# Misc
'terminal_server.connections': 'TERMINAL_SERVER_CONNECTIONS',
'tool_server.connections': 'TOOL_SERVER_CONNECTIONS',
'ui.banners': 'WEBUI_BANNERS',
'ui.default_group_id': 'DEFAULT_GROUP_ID',
'ui.default_locale': 'DEFAULT_LOCALE',
'ui.default_models': 'DEFAULT_MODELS',
'ui.default_pinned_models': 'DEFAULT_PINNED_MODELS',
'ui.default_user_role': 'DEFAULT_USER_ROLE',
'ui.enable_community_sharing': 'ENABLE_COMMUNITY_SHARING',
'ui.enable_login_form': 'ENABLE_LOGIN_FORM',
'ui.enable_message_rating': 'ENABLE_MESSAGE_RATING',
'ui.enable_password_change_form': 'ENABLE_PASSWORD_CHANGE_FORM',
'ui.enable_signup': 'ENABLE_SIGNUP',
'ui.enable_user_webhooks': 'ENABLE_USER_WEBHOOKS',
'ui.model_order_list': 'MODEL_ORDER_LIST',
'ui.pending_user_overlay_content': 'PENDING_USER_OVERLAY_CONTENT',
'ui.pending_user_overlay_title': 'PENDING_USER_OVERLAY_TITLE',
'ui.prompt_suggestions': 'DEFAULT_PROMPT_SUGGESTIONS',
'ui.watermark': 'RESPONSE_WATERMARK',
'user.permissions': 'USER_PERMISSIONS',
'users.enable_status': 'ENABLE_USER_STATUS',
'webhook_url': 'WEBHOOK_URL',
'webui.url': 'WEBUI_URL',
}
STORAGE_KEY_REWRITES = {
'oauth.refresh_token_include_scope': 'oauth.refresh_token.include_scope',
'rag.openai_api_base_url': 'rag.openai.api_base_url',
'rag.openai_api_key': 'rag.openai.api_key',
'rag.ollama.url': 'rag.ollama.base_url',
'rag.ollama.key': 'rag.ollama.api_key',
'oauth.oidc.avatar_claim': 'oauth.picture_claim',
'oauth.oidc.client_id': 'oauth.client_id',
'oauth.oidc.client_secret': 'oauth.client_secret',
'oauth.oidc.code_challenge_method': 'oauth.code_challenge_method',
'oauth.oidc.email_claim': 'oauth.email_claim',
'oauth.oidc.end_session_endpoint': 'oauth.end_session_endpoint',
'oauth.oidc.group_claim': 'oauth.group_claim',
'oauth.oidc.oauth_timeout': 'oauth.timeout',
'oauth.oidc.provider_name': 'oauth.provider_name',
'oauth.oidc.provider_url': 'oauth.provider_url',
'oauth.oidc.redirect_uri': 'oauth.redirect_uri',
'oauth.oidc.scopes': 'oauth.scopes',
'oauth.oidc.sub_claim': 'oauth.sub_claim',
'oauth.oidc.token_endpoint_auth_method': 'oauth.token_endpoint_auth_method',
'oauth.oidc.username_claim': 'oauth.username_claim',
}
LEGACY_KEY_TO_STORAGE_KEY = {
legacy_key: STORAGE_KEY_REWRITES.get(blob_path, blob_path) for blob_path, legacy_key in BLOB_PATH_TO_KEY.items()
}
def _walk_blob(data: dict, prefix: str = '') -> dict:
"""Recursively walk a nested config blob, preserving known config values.
Some config values are intentionally dictionaries, e.g. OPENAI_API_CONFIGS
and OLLAMA_API_CONFIGS. Once the current path is a known config key, keep
that value intact instead of flattening its internals into orphaned rows.
"""
result = {}
for key, value in data.items():
path = f'{prefix}{key}' if not prefix else f'{prefix}.{key}'
if path in BLOB_PATH_TO_KEY or path in LEGACY_KEY_TO_STORAGE_KEY:
result[path] = value
elif isinstance(value, dict):
result.update(_walk_blob(value, path))
else:
result[path] = value
return result
def upgrade() -> None:
"""Reshape config from single-row JSON blob to per-key rows."""
conn = op.get_bind()
inspector = sa.inspect(conn)
table_names = set(inspector.get_table_names())
config_columns = (
{column['name'] for column in inspector.get_columns('config')} if 'config' in table_names else set()
)
has_old_config = {'id', 'data'}.issubset(config_columns)
has_new_config = {'key', 'value'}.issubset(config_columns)
# Ad-hoc table reference for reading the old schema
old_config = sa.table(
'config',
sa.column('id', sa.Integer),
sa.column('data', sa.JSON),
)
# 1. Read existing blob
blob_data = {}
if has_old_config:
try:
result = conn.execute(sa.select(old_config.c.data).order_by(old_config.c.id.desc()).limit(1))
row = result.fetchone()
if row and row[0]:
raw = row[0]
blob_data = json.loads(raw) if isinstance(raw, str) else raw
except Exception:
pass # Table might be partially migrated or empty
# 2. Preserve old blob table for rollback/inspection, then create per-key table.
if has_old_config:
if 'config_old' in table_names:
op.drop_table('config_old')
op.rename_table('config', 'config_old')
# 3. Create new per-key table
new_config = (
sa.table(
'config',
sa.column('key', sa.Text),
sa.column('value', sa.JSON()),
sa.column('updated_at', sa.BigInteger),
)
if has_new_config
else op.create_table(
'config',
sa.Column('key', sa.Text(), primary_key=True),
sa.Column('value', sa.JSON(), nullable=False),
sa.Column('updated_at', sa.BigInteger(), nullable=True),
)
)
# 4. Flatten blob and insert per-key rows
if blob_data:
flat = _walk_blob(blob_data)
# Keep stable dot-notation paths as the database keys.
# Known legacy env-style keys are rewritten to their dotted keys; unknown
# keys are still copied so custom/future config is not silently lost.
rows = {}
for blob_path, value in flat.items():
if blob_path in BLOB_PATH_TO_KEY:
storage_key = STORAGE_KEY_REWRITES.get(blob_path, blob_path)
elif blob_path in LEGACY_KEY_TO_STORAGE_KEY:
storage_key = LEGACY_KEY_TO_STORAGE_KEY[blob_path]
else:
storage_key = STORAGE_KEY_REWRITES.get(blob_path, blob_path)
if storage_key not in rows:
rows[storage_key] = value
# Batch insert via SQLAlchemy table reference
if rows:
now = int(time.time())
op.bulk_insert(
new_config,
[{'key': k, 'value': v, 'updated_at': now} for k, v in rows.items()],
)
def downgrade() -> None:
"""Restore preserved old single-row config table when available."""
conn = op.get_bind()
inspector = sa.inspect(conn)
table_names = set(inspector.get_table_names())
if 'config_old' in table_names:
if 'config' in table_names:
op.drop_table('config')
op.rename_table('config_old', 'config')
return
config_columns = (
{column['name'] for column in inspector.get_columns('config')} if 'config' in table_names else set()
)
has_per_key_config = {'key', 'value'}.issubset(config_columns)
blob_data = {}
if has_per_key_config:
config = sa.table(
'config',
sa.column('key', sa.Text),
sa.column('value', sa.JSON),
)
for key, value in conn.execute(sa.select(config.c.key, config.c.value)):
blob_data[key] = json.loads(value) if isinstance(value, str) else value
op.drop_table('config')
if 'config' in table_names and not has_per_key_config:
return
old_config = op.create_table(
'config',
sa.Column('id', sa.Integer(), primary_key=True),
sa.Column('data', sa.JSON(), nullable=False),
sa.Column('version', sa.Integer(), nullable=False, server_default='0'),
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.Column('updated_at', sa.DateTime(), nullable=True),
)
if blob_data:
op.bulk_insert(old_config, [{'data': blob_data, 'version': 0}])

View file

@ -1,40 +0,0 @@
"""add memory path and meta
Revision ID: 42e2978c7933
Revises: 7b3f2a9c1d4e
Create Date: 2026-06-29 05:35:50.565887
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = '42e2978c7933'
down_revision: Union[str, None] = '7b3f2a9c1d4e'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {column['name'] for column in inspector.get_columns('memory')}
if 'path' not in columns:
op.add_column('memory', sa.Column('path', sa.Text(), nullable=True))
if 'meta' not in columns:
op.add_column('memory', sa.Column('meta', sa.JSON(), nullable=True))
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {column['name'] for column in inspector.get_columns('memory')}
if 'meta' in columns:
op.drop_column('memory', 'meta')
if 'path' in columns:
op.drop_column('memory', 'path')

View file

@ -1,37 +0,0 @@
"""add context summary to chat message
Revision ID: 4c5ce3d2f27f
Revises: 3ff2c63645b8
Create Date: 2026-06-18 23:48:08.310063
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '4c5ce3d2f27f'
down_revision: Union[str, None] = '3ff2c63645b8'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {column['name'] for column in inspector.get_columns('chat_message')}
if 'context_summary' not in columns:
op.add_column('chat_message', sa.Column('context_summary', sa.Text(), nullable=True))
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {column['name'] for column in inspector.get_columns('chat_message')}
if 'context_summary' in columns:
op.drop_column('chat_message', 'context_summary')

View file

@ -1,36 +0,0 @@
"""Add memory (id, user_id) covering index
Revision ID: 55f1302ac17c
Revises: b0018471bbbe
Create Date: 2026-07-24 00:00:00.000000
"""
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
revision: str = '55f1302ac17c'
down_revision: Union[str, None] = 'b0018471bbbe'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
indexes = {index['name'] for index in inspector.get_indexes('memory')}
if 'ix_memory_id_user_id' not in indexes:
op.create_index('ix_memory_id_user_id', 'memory', ['id', 'user_id'])
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
indexes = {index['name'] for index in inspector.get_indexes('memory')}
if 'ix_memory_id_user_id' in indexes:
op.drop_index('ix_memory_id_user_id', table_name='memory')

View file

@ -1,64 +0,0 @@
"""repair double encoded user oauth
Revision ID: 6d09d1bf1f23
Revises: 1ce6ade7d93b
Create Date: 2026-08-10 23:20:20.374826
"""
import json
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
import open_webui.internal.db
# revision identifiers, used by Alembic.
revision: str = '6d09d1bf1f23'
down_revision: Union[str, None] = '1ce6ade7d93b'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
_user = sa.table(
'user',
sa.column('id', sa.Text),
sa.column('oauth', sa.JSON),
)
def _decode_json_object(value: str) -> dict | None:
try:
decoded = json.loads(value)
except Exception:
return None
return decoded if isinstance(decoded, dict) else None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
if 'user' not in inspector.get_table_names():
return
user_columns = {c['name'] for c in inspector.get_columns('user')}
if 'oauth' not in user_columns:
return
rows = conn.execute(sa.select(_user.c.id, _user.c.oauth).where(_user.c.oauth.is_not(None))).fetchall()
for uid, oauth in rows:
if not isinstance(oauth, str):
continue
decoded = _decode_json_object(oauth)
if decoded is None:
continue
conn.execute(sa.update(_user).where(_user.c.id == uid).values(oauth=decoded))
def downgrade() -> None:
pass

View file

@ -1,44 +0,0 @@
"""add memory type
Revision ID: 7b3f2a9c1d4e
Revises: 4c5ce3d2f27f
Create Date: 2026-06-25 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = '7b3f2a9c1d4e'
down_revision: Union[str, None] = '4c5ce3d2f27f'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {column['name'] for column in inspector.get_columns('memory')}
indexes = {index['name'] for index in inspector.get_indexes('memory')}
if 'type' not in columns:
op.add_column('memory', sa.Column('type', sa.String(), server_default='context', nullable=False))
if 'ix_memory_type' not in indexes:
op.create_index('ix_memory_type', 'memory', ['type'])
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {column['name'] for column in inspector.get_columns('memory')}
indexes = {index['name'] for index in inspector.get_indexes('memory')}
if 'ix_memory_type' in indexes:
op.drop_index('ix_memory_type', table_name='memory')
if 'type' in columns:
op.drop_column('memory', 'type')

View file

@ -1,25 +0,0 @@
"""add chat message meta
Revision ID: 856c5b02fb54
Revises: 42e2978c7933
Create Date: 2026-07-16 01:39:39.291935
"""
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
revision: str = '856c5b02fb54'
down_revision: Union[str, None] = '42e2978c7933'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.add_column('chat_message', sa.Column('meta', sa.JSON(), nullable=True))
def downgrade() -> None:
op.drop_column('chat_message', 'meta')

View file

@ -1,54 +0,0 @@
"""add automation folder id
Revision ID: 959eaac8f909
Revises: 55f1302ac17c
Create Date: 2026-07-26 19:19:31.345756
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = '959eaac8f909'
down_revision: str | None = '55f1302ac17c'
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
if context.is_offline_mode():
op.add_column('automation', sa.Column('folder_id', sa.Text(), nullable=True))
op.create_index('ix_automation_user_folder', 'automation', ['user_id', 'folder_id'])
return
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {col['name'] for col in inspector.get_columns('automation')}
indexes = {index['name'] for index in inspector.get_indexes('automation')}
if 'folder_id' not in columns:
op.add_column('automation', sa.Column('folder_id', sa.Text(), nullable=True))
if 'ix_automation_user_folder' not in indexes:
op.create_index('ix_automation_user_folder', 'automation', ['user_id', 'folder_id'])
def downgrade() -> None:
if context.is_offline_mode():
op.drop_index('ix_automation_user_folder', table_name='automation')
op.drop_column('automation', 'folder_id')
return
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {col['name'] for col in inspector.get_columns('automation')}
indexes = {index['name'] for index in inspector.get_indexes('automation')}
if 'ix_automation_user_folder' in indexes:
op.drop_index('ix_automation_user_folder', table_name='automation')
if 'folder_id' in columns:
op.drop_column('automation', 'folder_id')

View file

@ -1,219 +0,0 @@
"""add current_message_id to chat
Revision ID: 9a1b2c3d4e5f
Revises: 856c5b02fb54
Create Date: 2026-07-23 00:00:00.000000
"""
import json
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
revision: str = '9a1b2c3d4e5f'
down_revision: Union[str, None] = '856c5b02fb54'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
BATCH_SIZE = 150
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = [col['name'] for col in inspector.get_columns('chat')]
if 'current_message_id' not in columns:
op.add_column('chat', sa.Column('current_message_id', sa.Text(), nullable=True))
chat = sa.table(
'chat',
sa.column('id', sa.String()),
sa.column('chat', sa.Text()),
sa.column('current_message_id', sa.Text()),
)
chat_message = sa.table(
'chat_message',
sa.column('id', sa.Text()),
sa.column('chat_id', sa.Text()),
sa.column('parent_id', sa.Text()),
sa.column('created_at', sa.BigInteger()),
)
has_chat_message = 'chat_message' in inspector.get_table_names()
result = conn.execute(
sa.select(chat.c.id, chat.c.chat, chat.c.current_message_id).execution_options(
yield_per=BATCH_SIZE,
stream_results=True,
)
)
while True:
rows = result.fetchmany(BATCH_SIZE)
if not rows:
break
batch_chat_ids: list[str] = []
candidates_by_chat: dict[str, list[str]] = {}
current_by_chat: dict[str, str | None] = {}
json_messages_by_chat: dict[str, dict[str, dict]] = {}
for row in rows:
values = row._mapping
chat_id = values['id']
prefix = f'{chat_id}-'
batch_chat_ids.append(chat_id)
current_by_chat[chat_id] = values['current_message_id']
chat_data = {}
if isinstance(values['chat'], dict):
chat_data = values['chat']
elif isinstance(values['chat'], str):
try:
parsed = json.loads(values['chat'])
chat_data = parsed if isinstance(parsed, dict) else {}
except (TypeError, ValueError, json.JSONDecodeError):
pass
history = chat_data.get('history') if isinstance(chat_data.get('history'), dict) else {}
candidates_by_chat[chat_id] = []
for candidate in (
values['current_message_id'],
history.get('currentId'),
chat_data.get('currentId'),
chat_data.get('branchPointMessageId'),
):
if not isinstance(candidate, str) or not candidate:
continue
candidate = candidate[len(prefix) :] if candidate.startswith(prefix) else candidate
if candidate not in candidates_by_chat[chat_id]:
candidates_by_chat[chat_id].append(candidate)
messages = history.get('messages') if isinstance(history.get('messages'), dict) else {}
if not messages and isinstance(chat_data.get('messages'), list):
messages = {
message['id']: message
for message in chat_data['messages']
if isinstance(message, dict) and message.get('id')
}
if messages:
json_messages_by_chat[chat_id] = {
message_id: {
'parent_id': message.get('parentId') if isinstance(message, dict) else None,
'created_at': message.get('timestamp', 0) if isinstance(message, dict) else 0,
}
for message_id, message in messages.items()
}
resolved: dict[str, str] = {}
if has_chat_message:
candidate_ids = {
f'{chat_id}-{candidate}'
for chat_id, candidates in candidates_by_chat.items()
for candidate in candidates
}
if candidate_ids:
valid_by_chat: dict[str, set[str]] = {}
for row in conn.execute(
sa.select(chat_message.c.chat_id, chat_message.c.id).where(
chat_message.c.chat_id.in_(batch_chat_ids),
chat_message.c.id.in_(candidate_ids),
)
):
values = row._mapping
chat_id = values['chat_id']
prefix = f'{chat_id}-'
message_id = values['id']
if message_id and message_id.startswith(prefix):
message_id = message_id[len(prefix) :]
if message_id:
valid_by_chat.setdefault(chat_id, set()).add(message_id)
for chat_id, candidates in candidates_by_chat.items():
valid_ids = valid_by_chat.get(chat_id, set())
for candidate in candidates:
if candidate in valid_ids:
resolved[chat_id] = candidate
break
unresolved_chat_ids = [chat_id for chat_id in batch_chat_ids if chat_id not in resolved]
messages_by_chat: dict[str, dict[str, dict]] = {}
if unresolved_chat_ids:
for row in conn.execute(
sa.select(
chat_message.c.chat_id,
chat_message.c.id,
chat_message.c.parent_id,
chat_message.c.created_at,
).where(chat_message.c.chat_id.in_(unresolved_chat_ids))
):
values = row._mapping
chat_id = values['chat_id']
prefix = f'{chat_id}-'
message_id = values['id']
if message_id and message_id.startswith(prefix):
message_id = message_id[len(prefix) :]
if not message_id:
continue
parent_id = values['parent_id']
if parent_id and parent_id.startswith(prefix):
parent_id = parent_id[len(prefix) :]
messages_by_chat.setdefault(chat_id, {})[message_id] = {
'parent_id': parent_id,
'created_at': values['created_at'] or 0,
}
for chat_id, messages in messages_by_chat.items():
parent_ids = {
message['parent_id'] for message in messages.values() if message.get('parent_id') in messages
}
leaf_ids = [message_id for message_id in messages if message_id not in parent_ids]
resolved[chat_id] = max(
leaf_ids or list(messages),
key=lambda message_id: messages[message_id].get('created_at') or 0,
)
for chat_id in batch_chat_ids:
if chat_id in resolved:
continue
messages = json_messages_by_chat.get(chat_id, {})
valid_candidate = next(
(candidate for candidate in candidates_by_chat[chat_id] if candidate in messages),
None,
)
if valid_candidate:
resolved[chat_id] = valid_candidate
elif messages:
parent_ids = {
message['parent_id'] for message in messages.values() if message.get('parent_id') in messages
}
leaf_ids = [message_id for message_id in messages if message_id not in parent_ids]
resolved[chat_id] = max(
leaf_ids or list(messages),
key=lambda message_id: messages[message_id].get('created_at') or 0,
)
updates = [
{'chat_id': chat_id, 'current_message_id': message_id}
for chat_id, message_id in resolved.items()
if message_id and message_id != current_by_chat.get(chat_id)
]
if updates:
conn.execute(
sa.update(chat)
.where(chat.c.id == sa.bindparam('update_chat_id'))
.values(current_message_id=sa.bindparam('update_current_message_id')),
[
{
'update_chat_id': row['chat_id'],
'update_current_message_id': row['current_message_id'],
}
for row in updates
],
)
def downgrade() -> None:
op.drop_column('chat', 'current_message_id')

View file

@ -1,32 +0,0 @@
"""add user variables
Revision ID: b0018471bbbe
Revises: c49178636c78
Create Date: 2026-07-24 01:21:46.457057
"""
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision: str = 'b0018471bbbe'
down_revision: Union[str, None] = 'c49178636c78'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = [col['name'] for col in inspector.get_columns('user')]
if 'variables' not in columns:
op.add_column('user', sa.Column('variables', sa.JSON(), nullable=True))
def downgrade() -> None:
op.drop_column('user', 'variables')

View file

@ -72,6 +72,7 @@ def _convert_column_to_json(table: str, column: str):
dialect = conn.dialect.name dialect = conn.dialect.name
t = sa.table(table, sa.column('id', sa.Text), sa.column(column, sa.Text)) t = sa.table(table, sa.column('id', sa.Text), sa.column(column, sa.Text))
t_json = sa.column(f'{column}_json', sa.JSON)
# SQLite cannot ALTER COLUMN → must recreate column # SQLite cannot ALTER COLUMN → must recreate column
if dialect == 'sqlite': if dialect == 'sqlite':
@ -89,9 +90,9 @@ def _convert_column_to_json(table: str, column: str):
parsed = None parsed = None
conn.execute( conn.execute(
sa.update(sa.table(table, sa.column('id'), sa.column(f'{column}_json', sa.JSON))) sa.update(sa.table(table, sa.column('id'), t_json))
.where(sa.column('id') == uid) .where(sa.column('id') == uid)
.values({f'{column}_json': parsed}) .values({f'{column}_json': json.dumps(parsed) if parsed else None})
) )
op.drop_column(table, column) op.drop_column(table, column)
@ -111,7 +112,8 @@ def _convert_column_to_text(table: str, column: str):
conn = op.get_bind() conn = op.get_bind()
dialect = conn.dialect.name dialect = conn.dialect.name
t = sa.table(table, sa.column('id', sa.Text), sa.column(column, sa.JSON)) t = sa.table(table, sa.column('id', sa.Text), sa.column(column))
t_text = sa.column(f'{column}_text', sa.Text)
if dialect == 'sqlite': if dialect == 'sqlite':
op.add_column(table, sa.Column(f'{column}_text', sa.Text(), nullable=True)) op.add_column(table, sa.Column(f'{column}_text', sa.Text(), nullable=True))
@ -120,9 +122,9 @@ def _convert_column_to_text(table: str, column: str):
for uid, raw in rows: for uid, raw in rows:
conn.execute( conn.execute(
sa.update(sa.table(table, sa.column('id'), sa.column(f'{column}_text', sa.Text))) sa.update(sa.table(table, sa.column('id'), t_text))
.where(sa.column('id') == uid) .where(sa.column('id') == uid)
.values({f'{column}_text': json.dumps(raw) if raw is not None else None}) .values({f'{column}_text': json.dumps(raw) if raw else None})
) )
op.drop_column(table, column) op.drop_column(table, column)
@ -185,7 +187,9 @@ def upgrade() -> None:
for uid, oauth_sub in rows: for uid, oauth_sub in rows:
if oauth_sub: 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={provider: {'sub': sub}})) conn.execute(
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) # ── Migrate api_key column → api_key table (only if old column still exists)
if 'api_key' in user_columns: if 'api_key' in user_columns:
@ -224,7 +228,7 @@ def downgrade() -> None:
for uid, oauth in rows: for uid, oauth in rows:
try: try:
data = oauth if isinstance(oauth, dict) else json.loads(oauth) data = json.loads(oauth)
provider = list(data.keys())[0] provider = list(data.keys())[0]
sub = data[provider].get('sub') sub = data[provider].get('sub')
oauth_sub = f'{provider}@{sub}' oauth_sub = f'{provider}@{sub}'

View file

@ -1,32 +0,0 @@
"""add chat variables
Revision ID: c49178636c78
Revises: 9a1b2c3d4e5f
Create Date: 2026-07-23 23:33:45.497453
"""
from typing import Sequence, Union
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision: str = 'c49178636c78'
down_revision: Union[str, None] = '9a1b2c3d4e5f'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = [col['name'] for col in inspector.get_columns('chat')]
if 'variables' not in columns:
op.add_column('chat', sa.Column('variables', sa.JSON(), nullable=True))
def downgrade() -> None:
op.drop_column('chat', 'variables')

View file

@ -1,70 +0,0 @@
"""add chat timer_at and chat list, unread and timer indexes
Revision ID: d4c1a8e37b62
Revises: 6d09d1bf1f23
Create Date: 2026-08-23 18:05:12.441907
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision: str = 'd4c1a8e37b62'
down_revision: str | None = '6d09d1bf1f23'
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.add_column('chat', sa.Column('timer_at', sa.BigInteger(), nullable=True))
op.create_index(
'timer_at_idx',
'chat',
['timer_at'],
sqlite_where=sa.text('timer_at IS NOT NULL'),
postgresql_where=sa.text('timer_at IS NOT NULL'),
)
# Timers created before this migration carry their due time in meta only, and would never fire.
chat = sa.table(
'chat', sa.column('id', sa.String), sa.column('meta', sa.JSON), sa.column('timer_at', sa.BigInteger)
)
conn = op.get_bind()
pending = conn.execute(
sa.select(chat.c.id, chat.c.meta)
.where(chat.c.meta['type'].as_string() == 'timer')
.where(chat.c.meta['status'].as_string() == 'pending')
).all()
for chat_id, meta in pending:
try: # imported chats can carry any meta, and a non-numeric due time must not abort the migration
due_at = int(meta.get('timer_at'))
except (TypeError, ValueError):
continue
conn.execute(chat.update().where(chat.c.id == chat_id).values(timer_at=due_at))
op.create_index('user_id_updated_at_id_idx', 'chat', ['user_id', sa.text('updated_at DESC'), 'id'])
op.create_index(
'user_id_timer_at_idx',
'chat',
['user_id', 'timer_at'],
sqlite_where=sa.text('timer_at IS NOT NULL'),
postgresql_where=sa.text('timer_at IS NOT NULL'),
)
op.create_index(
'user_id_folder_unread_idx',
'chat',
['user_id', 'folder_id', 'archived', 'updated_at', 'last_read_at', 'id'],
)
op.create_index('chat_message_chat_role_done_idx', 'chat_message', ['chat_id', 'role', 'done'])
def downgrade() -> None:
op.drop_index('chat_message_chat_role_done_idx', table_name='chat_message')
op.drop_index('user_id_folder_unread_idx', table_name='chat')
op.drop_index('user_id_timer_at_idx', table_name='chat')
op.drop_index('user_id_updated_at_id_idx', table_name='chat')
op.drop_index('timer_at_idx', table_name='chat')
op.drop_column('chat', 'timer_at')

View file

@ -1,84 +0,0 @@
"""add unique normalized user email index
Revision ID: f0bd01a18a3d
Revises: 959eaac8f909
Create Date: 2026-07-27 04:41:12.708743
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = 'f0bd01a18a3d'
down_revision: str | None = '959eaac8f909'
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
INDEX_NAME = 'uq_user_email_lower'
EMAIL_IS_NOT_NULL = sa.text('email IS NOT NULL')
LOWER_EMAIL = sa.text('lower(email)')
def _index_exists() -> bool:
conn = op.get_bind()
inspector = sa.inspect(conn)
return INDEX_NAME in {index['name'] for index in inspector.get_indexes('user')}
def _duplicate_emails() -> list:
conn = op.get_bind()
return conn.execute(
sa.text(
"""
SELECT lower(email) AS email, count(*) AS duplicate_count
FROM "user"
WHERE email IS NOT NULL
GROUP BY lower(email)
HAVING count(*) > 1
ORDER BY lower(email)
"""
)
).fetchall()
def _create_index() -> None:
op.create_index(
INDEX_NAME,
'user',
[LOWER_EMAIL],
unique=True,
postgresql_where=EMAIL_IS_NOT_NULL,
sqlite_where=EMAIL_IS_NOT_NULL,
)
def upgrade() -> None:
if context.is_offline_mode():
_create_index()
return
if _index_exists():
return
duplicates = _duplicate_emails()
if duplicates:
details = ', '.join(f'{row.email} (x{row.duplicate_count})' for row in duplicates)
raise RuntimeError(
'Cannot add unique normalized user email index because duplicate emails exist: '
f'{details}. Merge or remove the duplicate users and rerun migrations.'
)
_create_index()
def downgrade() -> None:
if context.is_offline_mode():
op.drop_index(INDEX_NAME, table_name='user')
return
if _index_exists():
op.drop_index(INDEX_NAME, table_name='user')

View file

@ -11,11 +11,6 @@ from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
PRINCIPAL_TYPE_ANYONE = 'anyone'
PRINCIPAL_TYPE_GROUP = 'group'
PRINCIPAL_TYPE_USER = 'user'
WILDCARD_PRINCIPAL_ID = '*'
#################### ####################
# AccessGrant DB Schema # AccessGrant DB Schema
@ -28,7 +23,7 @@ class AccessGrant(Base):
id = Column(Text, primary_key=True) id = Column(Text, primary_key=True)
resource_type = Column(Text, nullable=False) # "knowledge", "model", "prompt", "tool", "note", "channel", "file" resource_type = Column(Text, nullable=False) # "knowledge", "model", "prompt", "tool", "note", "channel", "file"
resource_id = Column(Text, nullable=False) resource_id = Column(Text, nullable=False)
principal_type = Column(Text, nullable=False) # "user", "group", or "anyone" principal_type = Column(Text, nullable=False) # "user" or "group"
principal_id = Column(Text, nullable=False) # user_id, group_id, or "*" (wildcard for public) principal_id = Column(Text, nullable=False) # user_id, group_id, or "*" (wildcard for public)
permission = Column(Text, nullable=False) # "read" or "write" permission = Column(Text, nullable=False) # "read" or "write"
created_at = Column(BigInteger, nullable=False) created_at = Column(BigInteger, nullable=False)
@ -168,14 +163,12 @@ def normalize_access_grants(access_grants: Optional[list]) -> list[dict]:
principal_id = grant.get('principal_id') principal_id = grant.get('principal_id')
permission = grant.get('permission') permission = grant.get('permission')
if principal_type not in (PRINCIPAL_TYPE_USER, PRINCIPAL_TYPE_GROUP, PRINCIPAL_TYPE_ANYONE): if principal_type not in ('user', 'group'):
continue continue
if permission not in ('read', 'write'): if permission not in ('read', 'write'):
continue continue
if not isinstance(principal_id, str) or not principal_id: if not isinstance(principal_id, str) or not principal_id:
continue continue
if principal_type == PRINCIPAL_TYPE_ANYONE and (principal_id != WILDCARD_PRINCIPAL_ID or permission != 'read'):
continue
key = (principal_type, principal_id, permission) key = (principal_type, principal_id, permission)
deduped[key] = { deduped[key] = {
@ -193,11 +186,7 @@ def has_public_read_access_grant(access_grants: Optional[list]) -> bool:
Returns True when a direct grant list includes wildcard public-read. Returns True when a direct grant list includes wildcard public-read.
""" """
for grant in normalize_access_grants(access_grants): for grant in normalize_access_grants(access_grants):
if ( if grant['principal_type'] == 'user' and grant['principal_id'] == '*' and grant['permission'] == 'read':
grant['principal_type'] == PRINCIPAL_TYPE_USER
and grant['principal_id'] == WILDCARD_PRINCIPAL_ID
and grant['permission'] == 'read'
):
return True return True
return False return False
@ -207,25 +196,7 @@ def has_public_write_access_grant(access_grants: Optional[list]) -> bool:
Returns True when a direct grant list includes wildcard public-write. Returns True when a direct grant list includes wildcard public-write.
""" """
for grant in normalize_access_grants(access_grants): for grant in normalize_access_grants(access_grants):
if ( if grant['principal_type'] == 'user' and grant['principal_id'] == '*' and grant['permission'] == 'write':
grant['principal_type'] == PRINCIPAL_TYPE_USER
and grant['principal_id'] == WILDCARD_PRINCIPAL_ID
and grant['permission'] == 'write'
):
return True
return False
def has_anyone_read_access_grant(access_grants: Optional[list]) -> bool:
"""
Returns True when a direct grant list includes no-auth anyone-read.
"""
for grant in normalize_access_grants(access_grants):
if (
grant['principal_type'] == PRINCIPAL_TYPE_ANYONE
and grant['principal_id'] == WILDCARD_PRINCIPAL_ID
and grant['permission'] == 'read'
):
return True return True
return False return False
@ -235,7 +206,7 @@ def has_user_access_grant(access_grants: Optional[list]) -> bool:
Returns True when a direct grant list includes any non-wildcard user grant. Returns True when a direct grant list includes any non-wildcard user grant.
""" """
for grant in normalize_access_grants(access_grants): for grant in normalize_access_grants(access_grants):
if grant['principal_type'] == PRINCIPAL_TYPE_USER and grant['principal_id'] != WILDCARD_PRINCIPAL_ID: if grant['principal_type'] == 'user' and grant['principal_id'] != '*':
return True return True
return False return False
@ -252,27 +223,12 @@ def strip_user_access_grants(access_grants: Optional[list]) -> list:
for grant in access_grants for grant in access_grants
if not ( if not (
(grant.get('principal_type') if isinstance(grant, dict) else getattr(grant, 'principal_type', None)) (grant.get('principal_type') if isinstance(grant, dict) else getattr(grant, 'principal_type', None))
== PRINCIPAL_TYPE_USER == 'user'
and (grant.get('principal_id') if isinstance(grant, dict) else getattr(grant, 'principal_id', None)) and (grant.get('principal_id') if isinstance(grant, dict) else getattr(grant, 'principal_id', None)) != '*'
!= WILDCARD_PRINCIPAL_ID
) )
] ]
def strip_anyone_access_grants(access_grants: Optional[list]) -> list:
"""
Remove no-auth anyone grants from the list.
"""
if not access_grants:
return []
return [
grant
for grant in access_grants
if (grant.get('principal_type') if isinstance(grant, dict) else getattr(grant, 'principal_type', None))
!= PRINCIPAL_TYPE_ANYONE
]
def grants_to_access_control(grants: list) -> Optional[dict]: def grants_to_access_control(grants: list) -> Optional[dict]:
""" """
Convert a list of grant objects (AccessGrantModel or AccessGrantResponse) Convert a list of grant objects (AccessGrantModel or AccessGrantResponse)
@ -360,6 +316,7 @@ class AccessGrantsTable:
) )
db.add(grant) db.add(grant)
await db.commit() await db.commit()
await db.refresh(grant)
return AccessGrantModel.model_validate(grant) return AccessGrantModel.model_validate(grant)
async def revoke_access( async def revoke_access(
@ -537,28 +494,6 @@ class AccessGrantsTable:
result_dict[g.resource_id].append(AccessGrantModel.model_validate(g)) result_dict[g.resource_id].append(AccessGrantModel.model_validate(g))
return result_dict return result_dict
async def has_anyone_access(
self,
resource_type: str,
resource_id: str,
permission: str = 'read',
db: Optional[AsyncSession] = None,
) -> bool:
"""Check for a no-auth anyone:* grant. Callers must opt in explicitly."""
async with get_async_db_context(db) as db:
result = await db.execute(
select(AccessGrant)
.filter(
AccessGrant.resource_type == resource_type,
AccessGrant.resource_id == resource_id,
AccessGrant.principal_type == PRINCIPAL_TYPE_ANYONE,
AccessGrant.principal_id == WILDCARD_PRINCIPAL_ID,
AccessGrant.permission == permission,
)
.limit(1)
)
return result.scalars().first() is not None
async def has_access( async def has_access(
self, self,
user_id: str, user_id: str,
@ -839,8 +774,7 @@ class AccessGrantsTable:
): ):
""" """
Filter for items where user has read BUT NOT write access. Filter for items where user has read BUT NOT write access.
A public (user:*) read grant counts as read access, so publicly shared Public items are NOT considered read_only.
read-only items are listed rather than being reachable only by direct link.
Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself, Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself,
so it remains synchronous. The caller is responsible for executing the query so it remains synchronous. The caller is responsible for executing the query
@ -851,6 +785,7 @@ class AccessGrantsTable:
from sqlalchemy import exists as sa_exists from sqlalchemy import exists as sa_exists
# Has read grant (not public)
read_grant_exists = ( read_grant_exists = (
select(AccessGrant.id) select(AccessGrant.id)
.where( .where(
@ -858,10 +793,6 @@ class AccessGrantsTable:
AccessGrant.resource_id == DocumentModel.id, AccessGrant.resource_id == DocumentModel.id,
AccessGrant.permission == 'read', AccessGrant.permission == 'read',
or_( or_(
and_(
AccessGrant.principal_type == 'user',
AccessGrant.principal_id == '*',
),
*( *(
[ [
and_( and_(
@ -888,6 +819,7 @@ class AccessGrantsTable:
.exists() .exists()
) )
# Does NOT have write grant
write_grant_exists = ( write_grant_exists = (
select(AccessGrant.id) select(AccessGrant.id)
.where( .where(
@ -895,10 +827,6 @@ class AccessGrantsTable:
AccessGrant.resource_id == DocumentModel.id, AccessGrant.resource_id == DocumentModel.id,
AccessGrant.permission == 'write', AccessGrant.permission == 'write',
or_( or_(
and_(
AccessGrant.principal_type == 'user',
AccessGrant.principal_id == '*',
),
*( *(
[ [
and_( and_(
@ -925,7 +853,21 @@ class AccessGrantsTable:
.exists() .exists()
) )
conditions = [read_grant_exists, ~write_grant_exists] # Is NOT public
public_grant_exists = (
select(AccessGrant.id)
.where(
AccessGrant.resource_type == resource_type,
AccessGrant.resource_id == DocumentModel.id,
AccessGrant.permission == 'read',
AccessGrant.principal_type == 'user',
AccessGrant.principal_id == '*',
)
.correlate(DocumentModel)
.exists()
)
conditions = [read_grant_exists, ~write_grant_exists, ~public_grant_exists]
# Not owner # Not owner
if user_id: if user_id:

View file

@ -6,22 +6,15 @@ import logging
import uuid import uuid
from typing import Optional from typing import Optional
import bcrypt
from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users
from open_webui.utils.validate import validate_image_url from open_webui.utils.validate import validate_profile_image_url
from pydantic import BaseModel, field_validator from pydantic import BaseModel, field_validator
from sqlalchemy import Boolean, Column, String, Text, delete, select, update from sqlalchemy import Boolean, Column, String, Text, delete, select, update
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
# Pre-computed hash verified on signin paths that lack a real credential
# (unknown user, inactive account) so response timing cannot reveal
# whether an account exists (CWE-208).
PLACEHOLDER_HASH = bcrypt.hashpw(b'placeholder', bcrypt.gensalt()).decode('utf-8')
class Auth(Base): # credential ↔ user linkage class Auth(Base): # credential ↔ user linkage
"""Maps a user ID to an email/password pair with an active flag.""" """Maps a user ID to an email/password pair with an active flag."""
@ -87,7 +80,7 @@ class SignupForm(BaseModel):
@classmethod @classmethod
def check_profile_image_url(cls, v: str | None) -> str | None: def check_profile_image_url(cls, v: str | None) -> str | None:
if v is not None: if v is not None:
return validate_image_url(v) return validate_profile_image_url(v)
return v return v
@ -125,20 +118,18 @@ class AuthsTable:
) )
session.add(credential) session.add(credential)
try: created_user = await Users.insert_new_user(
created_user = await Users.insert_new_user( new_id,
new_id, name,
name, email,
email, profile_image_url,
profile_image_url, role,
role, oauth=oauth,
oauth=oauth, db=session,
db=session, )
) # persist both records and reload generated defaults
await session.commit() await session.commit()
except IntegrityError: await session.refresh(credential)
await session.rollback()
raise
return created_user if credential and created_user else None return created_user if credential and created_user else None
async def authenticate_user( async def authenticate_user(
@ -151,15 +142,13 @@ class AuthsTable:
log.info('authenticate_user: %s', email) log.info('authenticate_user: %s', email)
resolved = await Users.get_user_by_email(email, db=db) resolved = await Users.get_user_by_email(email, db=db)
if not resolved: if not resolved:
await verify_password(PLACEHOLDER_HASH)
return return
# load the credential row and verify the password hash # load the credential row and verify the password hash
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
credential = await session.get(Auth, resolved.id) credential = await session.get(Auth, resolved.id)
if not credential or not credential.active: if not credential or not credential.active:
await verify_password(PLACEHOLDER_HASH)
return return
if not await verify_password(credential.password): if not verify_password(credential.password):
return return
return resolved return resolved

View file

@ -1,10 +1,9 @@
import logging import logging
import time import time
from typing import Literal, Optional from typing import Optional
from uuid import uuid4 from uuid import uuid4
from open_webui.internal.db import Base, get_async_db_context from open_webui.internal.db import Base, get_async_db_context
from open_webui.utils.misc import json_text_variants
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from sqlalchemy import JSON, BigInteger, Boolean, Column, Index, String, Text, cast, delete, func, or_, select, update from sqlalchemy import JSON, BigInteger, Boolean, Column, Index, String, Text, cast, delete, func, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -22,7 +21,6 @@ class Automation(Base):
id = Column(Text, primary_key=True) id = Column(Text, primary_key=True)
user_id = Column(Text, nullable=False) user_id = Column(Text, nullable=False)
folder_id = Column(Text, nullable=True)
name = Column(Text, nullable=False) name = Column(Text, nullable=False)
data = Column(JSON, nullable=False) # {prompt, model_id, rrule} data = Column(JSON, nullable=False) # {prompt, model_id, rrule}
meta = Column(JSON, nullable=True) meta = Column(JSON, nullable=True)
@ -33,10 +31,7 @@ class Automation(Base):
created_at = Column(BigInteger, nullable=False) created_at = Column(BigInteger, nullable=False)
updated_at = Column(BigInteger, nullable=False) updated_at = Column(BigInteger, nullable=False)
__table_args__ = ( __table_args__ = (Index('ix_automation_next_run', 'next_run_at'),)
Index('ix_automation_next_run', 'next_run_at'),
Index('ix_automation_user_folder', 'user_id', 'folder_id'),
)
class AutomationRun(Base): class AutomationRun(Base):
@ -65,17 +60,11 @@ class AutomationTerminalConfig(BaseModel):
cwd: Optional[str] = None cwd: Optional[str] = None
class AutomationTarget(BaseModel):
type: Literal['chat', 'channel'] = 'chat'
channel_id: Optional[str] = None
class AutomationData(BaseModel): class AutomationData(BaseModel):
prompt: str prompt: str
model_id: str model_id: str
rrule: str rrule: str
terminal: Optional[AutomationTerminalConfig] = None terminal: Optional[AutomationTerminalConfig] = None
target: Optional[AutomationTarget] = None
class AutomationModel(BaseModel): class AutomationModel(BaseModel):
@ -83,7 +72,6 @@ class AutomationModel(BaseModel):
id: str id: str
user_id: str user_id: str
folder_id: Optional[str] = None
name: str name: str
data: dict data: dict
meta: Optional[dict] = None meta: Optional[dict] = None
@ -108,7 +96,6 @@ class AutomationRunModel(BaseModel):
class AutomationForm(BaseModel): class AutomationForm(BaseModel):
name: str name: str
folder_id: Optional[str] = None
data: AutomationData data: AutomationData
meta: Optional[dict] = None meta: Optional[dict] = None
is_active: Optional[bool] = True is_active: Optional[bool] = True
@ -142,7 +129,6 @@ class AutomationTable:
row = Automation( row = Automation(
id=str(uuid4()), id=str(uuid4()),
user_id=user_id, user_id=user_id,
folder_id=form.folder_id,
name=form.name, name=form.name,
data=form.data.model_dump(), data=form.data.model_dump(),
meta=form.meta, meta=form.meta,
@ -153,6 +139,7 @@ class AutomationTable:
) )
db.add(row) db.add(row)
await db.commit() await db.commit()
await db.refresh(row)
return AutomationModel.model_validate(row) return AutomationModel.model_validate(row)
async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int: async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int:
@ -178,7 +165,6 @@ class AutomationTable:
user_id: str, user_id: str,
query: Optional[str] = None, query: Optional[str] = None,
status: Optional[str] = None, status: Optional[str] = None,
folder_id: Optional[str] = None,
skip: int = 0, skip: int = 0,
limit: int = 30, limit: int = 30,
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
@ -186,16 +172,13 @@ class AutomationTable:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
stmt = select(Automation).filter_by(user_id=user_id) stmt = select(Automation).filter_by(user_id=user_id)
if folder_id:
stmt = stmt.filter(Automation.folder_id == folder_id)
if query: if query:
# Search the name column and the prompt inside the JSON data. search = f'%{query}%'
data_text = cast(Automation.data, String) # Search in name and prompt inside JSON data
stmt = stmt.filter( stmt = stmt.filter(
or_( or_(
Automation.name.ilike(f'%{query}%'), Automation.name.ilike(search),
*(data_text.ilike(f'%{variant}%') for variant in json_text_variants(query)), cast(Automation.data, String).ilike(search),
) )
) )
@ -234,7 +217,6 @@ class AutomationTable:
if not row: if not row:
return None return None
row.name = form.name row.name = form.name
row.folder_id = form.folder_id
row.data = form.data.model_dump() row.data = form.data.model_dump()
row.meta = form.meta row.meta = form.meta
if form.is_active is not None: if form.is_active is not None:
@ -242,25 +224,9 @@ class AutomationTable:
row.next_run_at = next_run_at row.next_run_at = next_run_at
row.updated_at = int(time.time_ns()) row.updated_at = int(time.time_ns())
await db.commit() await db.commit()
await db.refresh(row)
return AutomationModel.model_validate(row) return AutomationModel.model_validate(row)
async def clear_folder_ids(
self,
user_id: str,
folder_ids: list[str],
db: Optional[AsyncSession] = None,
) -> int:
if not folder_ids:
return 0
async with get_async_db_context(db) as db:
result = await db.execute(
update(Automation)
.where(Automation.user_id == user_id, Automation.folder_id.in_(folder_ids))
.values(folder_id=None, updated_at=int(time.time_ns()))
)
await db.commit()
return result.rowcount or 0
async def toggle( async def toggle(
self, self,
id: str, id: str,
@ -275,6 +241,7 @@ class AutomationTable:
row.next_run_at = next_run_at if row.is_active else None row.next_run_at = next_run_at if row.is_active else None
row.updated_at = int(time.time_ns()) row.updated_at = int(time.time_ns())
await db.commit() await db.commit()
await db.refresh(row)
return AutomationModel.model_validate(row) return AutomationModel.model_validate(row)
async def delete(self, id: str, db: Optional[AsyncSession] = None) -> bool: async def delete(self, id: str, db: Optional[AsyncSession] = None) -> bool:
@ -312,7 +279,6 @@ class AutomationTable:
rows = result.scalars().all() rows = result.scalars().all()
from open_webui.utils.automations import next_run_ns from open_webui.utils.automations import next_run_ns
from open_webui.utils.recurrence import RecurrenceEvaluationTimeout
# Batch-fetch user timezones so rescheduling respects each # Batch-fetch user timezones so rescheduling respects each
# user's local timezone instead of falling back to server time. # user's local timezone instead of falling back to server time.
@ -324,20 +290,13 @@ class AutomationTable:
tz_result = await db.execute(select(User.id, User.timezone).where(User.id.in_(user_ids))) tz_result = await db.execute(select(User.id, User.timezone).where(User.id.in_(user_ids)))
timezone_by_user_id = {uid: tz for uid, tz in tz_result.all()} timezone_by_user_id = {uid: tz for uid, tz in tz_result.all()}
claimed = []
for row in rows: for row in rows:
try:
next_run_at = await next_run_ns(row.data.get('rrule', ''), tz=timezone_by_user_id.get(row.user_id))
except RecurrenceEvaluationTimeout:
log.warning('Skipping automation %s: recurrence evaluation timed out', row.id)
continue
row.last_run_at = now_ns row.last_run_at = now_ns
row.next_run_at = next_run_at row.next_run_at = next_run_ns(row.data.get('rrule', ''), tz=timezone_by_user_id.get(row.user_id))
claimed.append(row)
await db.commit() await db.commit()
return [AutomationModel.model_validate(r) for r in claimed] return [AutomationModel.model_validate(r) for r in rows]
#################### ####################
@ -365,6 +324,7 @@ class AutomationRunTable:
) )
db.add(row) db.add(row)
await db.commit() await db.commit()
await db.refresh(row)
return AutomationRunModel.model_validate(row) return AutomationRunModel.model_validate(row)
async def get_latest(self, automation_id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationRunModel]: async def get_latest(self, automation_id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationRunModel]:

View file

@ -4,7 +4,6 @@ from typing import Optional
from uuid import uuid4 from uuid import uuid4
from open_webui.internal.db import Base, get_async_db_context from open_webui.internal.db import Base, get_async_db_context
from open_webui.constants import ERROR_MESSAGES
from open_webui.models.access_grants import AccessGrantModel, AccessGrants from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups from open_webui.models.groups import Groups
from open_webui.models.users import User, UserModel, UserResponse from open_webui.models.users import User, UserModel, UserResponse
@ -27,7 +26,6 @@ from sqlalchemy import (
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
MIN_CALENDAR_RRULE_INTERVAL_SECONDS = 24 * 60 * 60
#################### ####################
@ -179,20 +177,6 @@ class CalendarUpdateForm(BaseModel):
access_grants: Optional[list[dict]] = None access_grants: Optional[list[dict]] = None
async def validate_calendar_rrule(value: Optional[str]) -> None:
if value:
from open_webui.utils.recurrence import rrule_interval_seconds
try:
interval = await rrule_interval_seconds(value)
except ValueError:
raise
except Exception as e:
raise ValueError(ERROR_MESSAGES.AUTOMATION_INVALID_RRULE(e)) from e
if interval is not None and interval < MIN_CALENDAR_RRULE_INTERVAL_SECONDS:
raise ValueError(ERROR_MESSAGES.CALENDAR_RRULE_TOO_FREQUENT)
class CalendarEventForm(BaseModel): class CalendarEventForm(BaseModel):
calendar_id: str calendar_id: str
title: str title: str
@ -257,11 +241,11 @@ class CalendarTable:
access_grants: Optional[list[AccessGrantModel]] = None, access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
) -> CalendarModel: ) -> CalendarModel:
calendar_model = CalendarModel.model_validate(cal) cal_data = CalendarModel.model_validate(cal).model_dump(exclude={'access_grants'})
calendar_model.access_grants = ( cal_data['access_grants'] = (
access_grants if access_grants is not None else await self._get_access_grants(calendar_model.id, db=db) access_grants if access_grants is not None else await self._get_access_grants(cal_data['id'], db=db)
) )
return calendar_model return CalendarModel.model_validate(cal_data)
async def get_or_create_defaults(self, user_id: str, db: Optional[AsyncSession] = None) -> list[CalendarModel]: async def get_or_create_defaults(self, user_id: str, db: Optional[AsyncSession] = None) -> list[CalendarModel]:
"""Return user's calendars, creating 'Personal' default if none exist.""" """Return user's calendars, creating 'Personal' default if none exist."""
@ -447,7 +431,6 @@ class CalendarEventTable:
async def insert_new_event( async def insert_new_event(
self, user_id: str, form_data: CalendarEventForm, db: Optional[AsyncSession] = None self, user_id: str, form_data: CalendarEventForm, db: Optional[AsyncSession] = None
) -> Optional[CalendarEventModel]: ) -> Optional[CalendarEventModel]:
await validate_calendar_rrule(form_data.rrule)
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
now = int(time.time_ns()) now = int(time.time_ns())
event = CalendarEvent( event = CalendarEvent(
@ -517,12 +500,9 @@ class CalendarEventTable:
# Filter to requested calendars only # Filter to requested calendars only
accessible_cal_ids = [c for c in accessible_cal_ids if c in calendar_ids] accessible_cal_ids = [c for c in accessible_cal_ids if c in calendar_ids]
# Also get event IDs where the user is an attendee, excluding invites they declined # Also get event IDs where user is an attendee
attendee_event_ids_result = await db.execute( attendee_event_ids_result = await db.execute(
select(CalendarEventAttendee.event_id).filter( select(CalendarEventAttendee.event_id).filter(CalendarEventAttendee.user_id == user_id)
CalendarEventAttendee.user_id == user_id,
CalendarEventAttendee.status != 'declined',
)
) )
attendee_event_ids = [r[0] for r in attendee_event_ids_result.all()] attendee_event_ids = [r[0] for r in attendee_event_ids_result.all()]
@ -550,8 +530,7 @@ class CalendarEventTable:
& (CalendarEvent.start_at < end) & (CalendarEvent.start_at < end)
& or_( & or_(
CalendarEvent.end_at.is_(None) & (CalendarEvent.start_at >= start), CalendarEvent.end_at.is_(None) & (CalendarEvent.start_at >= start),
CalendarEvent.end_at.isnot(None) CalendarEvent.end_at.isnot(None) & (CalendarEvent.end_at > start),
& ((CalendarEvent.end_at > start) | (CalendarEvent.start_at >= start)),
) )
), ),
# Recurring: fetch all (expansion in Python) # Recurring: fetch all (expansion in Python)
@ -678,7 +657,6 @@ class CalendarEventTable:
async def update_event_by_id( async def update_event_by_id(
self, id: str, form_data: CalendarEventUpdateForm, db: Optional[AsyncSession] = None self, id: str, form_data: CalendarEventUpdateForm, db: Optional[AsyncSession] = None
) -> Optional[CalendarEventModel]: ) -> Optional[CalendarEventModel]:
await validate_calendar_rrule(form_data.rrule)
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute(select(CalendarEvent).filter(CalendarEvent.id == id)) result = await db.execute(select(CalendarEvent).filter(CalendarEvent.id == id))
event = result.scalars().first() event = result.scalars().first()
@ -753,10 +731,10 @@ class CalendarEventTable:
events = [] events = []
for event, tz in rows: for event, tz in rows:
model = CalendarEventModel.model_validate(event) model = CalendarEventModel.model_validate(event)
# meta is user-writable and this poll is shared by every user. # Determine per-event alert window
alert_minutes = (model.meta or {}).get('alert_minutes') alert_minutes = None
if not isinstance(alert_minutes, (int, float)): if model.meta and 'alert_minutes' in model.meta:
alert_minutes = None alert_minutes = model.meta['alert_minutes']
if alert_minutes is not None: if alert_minutes is not None:
if alert_minutes < 0: if alert_minutes < 0:
@ -786,32 +764,22 @@ class CalendarEventAttendeeTable:
async def set_attendees( async def set_attendees(
self, event_id: str, attendees: list[dict], db: Optional[AsyncSession] = None self, event_id: str, attendees: list[dict], db: Optional[AsyncSession] = None
) -> list[CalendarEventAttendeeModel]: ) -> list[CalendarEventAttendeeModel]:
"""Replace all attendees for an event ({user_id, meta?} per dict). """Replace all attendees for an event.
RSVP status is the attendee's alone to set (via update_rsvp): an existing Each dict in attendees: {user_id: str, status?: str, meta?: dict}
attendee keeps their status, a newly added one starts 'pending'. A
caller-supplied status is ignored so an organiser cannot set it for others.
""" """
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
existing_status = {
row.user_id: row.status
for row in (
await db.execute(select(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == event_id))
).scalars()
}
# Remove existing # Remove existing
await db.execute(delete(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == event_id)) await db.execute(delete(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == event_id))
now = int(time.time_ns()) now = int(time.time_ns())
models = [] models = []
for att in attendees: for att in attendees:
user_id = att['user_id']
row = CalendarEventAttendee( row = CalendarEventAttendee(
id=str(uuid4()), id=str(uuid4()),
event_id=event_id, event_id=event_id,
user_id=user_id, user_id=att['user_id'],
status=existing_status.get(user_id, 'pending'), status=att.get('status', 'pending'),
meta=att.get('meta'), meta=att.get('meta'),
created_at=now, created_at=now,
updated_at=now, updated_at=now,

View file

@ -1,3 +1,4 @@
import json
import secrets import secrets
import time import time
import uuid import uuid
@ -9,8 +10,7 @@ from open_webui.models.access_grants import (
AccessGrants, AccessGrants,
) )
from open_webui.models.groups import Groups from open_webui.models.groups import Groups
from open_webui.models.users import User from open_webui.utils.validate import validate_profile_image_url
from open_webui.utils.validate import validate_image_url
from pydantic import BaseModel, ConfigDict, Field, field_validator from pydantic import BaseModel, ConfigDict, Field, field_validator
from sqlalchemy import ( from sqlalchemy import (
JSON, JSON,
@ -253,7 +253,7 @@ class ChannelWebhookForm(BaseModel):
def check_profile_image_url(cls, v: Optional[str]) -> Optional[str]: def check_profile_image_url(cls, v: Optional[str]) -> Optional[str]:
if v is None: if v is None:
return v return v
return validate_image_url(v) return validate_profile_image_url(v)
class ChannelTable: class ChannelTable:
@ -266,11 +266,11 @@ class ChannelTable:
access_grants: Optional[list[AccessGrantModel]] = None, access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
) -> ChannelModel: ) -> ChannelModel:
channel_model = ChannelModel.model_validate(channel) channel_data = ChannelModel.model_validate(channel).model_dump(exclude={'access_grants'})
channel_model.access_grants = ( channel_data['access_grants'] = (
access_grants if access_grants is not None else await self._get_access_grants(channel_model.id, db=db) access_grants if access_grants is not None else await self._get_access_grants(channel_data['id'], db=db)
) )
return channel_model return ChannelModel.model_validate(channel_data)
async def _collect_unique_user_ids( async def _collect_unique_user_ids(
self, self,
@ -438,17 +438,17 @@ class ChannelTable:
match_count = func.sum( match_count = func.sum(
case( case(
(User.id.in_(unique_user_ids), 1), (ChannelMember.user_id.in_(unique_user_ids), 1),
else_=0, else_=0,
) )
) )
subquery = ( subquery = (
select(ChannelMember.channel_id) select(ChannelMember.channel_id)
.join(User, User.id == ChannelMember.user_id)
.group_by(ChannelMember.channel_id) .group_by(ChannelMember.channel_id)
# Match the exact set of accounts that still exist. # 1. Channel must have exactly len(user_ids) members
.having(func.count(User.id) == len(unique_user_ids)) .having(func.count(ChannelMember.user_id) == len(unique_user_ids))
# 2. All those members must be in unique_user_ids
.having(match_count == len(unique_user_ids)) .having(match_count == len(unique_user_ids))
.subquery() .subquery()
) )
@ -869,6 +869,7 @@ class ChannelTable:
result = ChannelFile(**channel_file.model_dump()) result = ChannelFile(**channel_file.model_dump())
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
if result: if result:
return ChannelFileModel.model_validate(result) return ChannelFileModel.model_validate(result)
else: else:

View file

@ -1,14 +1,10 @@
import json
import time import time
import uuid import uuid
from collections import Counter
from datetime import datetime, timedelta
from typing import Any, Optional from typing import Any, Optional
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
from sqlalchemy import select, delete, func, cast, Integer, distinct
from sqlalchemy.ext.asyncio import AsyncSession
from open_webui.internal.db import Base, get_async_db_context from open_webui.internal.db import Base, get_async_db_context
from open_webui.utils.response import merge_usage, normalize_usage from open_webui.utils.response import normalize_usage
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from sqlalchemy import ( from sqlalchemy import (
JSON, JSON,
@ -49,17 +45,6 @@ def _normalize_timestamp(timestamp: int) -> float:
return timestamp return timestamp
def _timezone(tz: Optional[str]) -> ZoneInfo:
try:
return ZoneInfo(tz or 'UTC')
except ZoneInfoNotFoundError:
return ZoneInfo('UTC')
def _date_key(timestamp: int, tz: ZoneInfo) -> str:
return datetime.fromtimestamp(_normalize_timestamp(timestamp), tz=tz).strftime('%Y-%m-%d')
def get_usage(data: dict) -> Optional[dict]: def get_usage(data: dict) -> Optional[dict]:
"""Extract and normalize usage from message data.""" """Extract and normalize usage from message data."""
usage = data.get('usage') or (data.get('info') or {}).get('usage') usage = data.get('usage') or (data.get('info') or {}).get('usage')
@ -85,40 +70,6 @@ def _token_columns(dialect: str):
) )
def _extract_tool_names(value: Any) -> list[str]:
names: list[str] = []
def add(name: Any):
if isinstance(name, str):
cleaned = name.strip()
if cleaned and len(cleaned) <= 128:
names.append(cleaned)
def walk(item: Any):
if isinstance(item, list):
for child in item:
walk(child)
return
if not isinstance(item, dict):
return
item_type = str(item.get('type') or '')
looks_like_tool = 'tool' in item_type or item_type in {'function_call', 'function_call_output'}
if looks_like_tool:
add(item.get('name') or item.get('tool_name'))
function = item.get('function')
if isinstance(function, dict):
add(function.get('name'))
for key in ('tool_calls', 'tools', 'output', 'meta'):
if key in item:
walk(item.get(key))
walk(value)
return names
#################### ####################
# ChatMessage DB Schema # ChatMessage DB Schema
#################### ####################
@ -147,7 +98,6 @@ class ChatMessage(Base):
files = Column(JSON, nullable=True) files = Column(JSON, nullable=True)
sources = Column(JSON, nullable=True) sources = Column(JSON, nullable=True)
embeds = Column(JSON, nullable=True) embeds = Column(JSON, nullable=True)
meta = Column(JSON, nullable=True)
# Status # Status
done = Column(Boolean, default=True) done = Column(Boolean, default=True)
@ -157,9 +107,6 @@ class ChatMessage(Base):
# Usage (tokens, timing, etc.) # Usage (tokens, timing, etc.)
usage = Column(JSON, nullable=True) usage = Column(JSON, nullable=True)
# Context compaction checkpoint
context_summary = Column(Text, nullable=True)
# Timestamps # Timestamps
created_at = Column(BigInteger, index=True) created_at = Column(BigInteger, index=True)
updated_at = Column(BigInteger) updated_at = Column(BigInteger)
@ -168,7 +115,6 @@ class ChatMessage(Base):
Index('chat_message_chat_parent_idx', 'chat_id', 'parent_id'), Index('chat_message_chat_parent_idx', 'chat_id', 'parent_id'),
Index('chat_message_model_created_idx', 'model_id', 'created_at'), Index('chat_message_model_created_idx', 'model_id', 'created_at'),
Index('chat_message_user_created_idx', 'user_id', 'created_at'), Index('chat_message_user_created_idx', 'user_id', 'created_at'),
Index('chat_message_chat_role_done_idx', 'chat_id', 'role', 'done'), # unfinished-assistant probe
) )
@ -191,12 +137,10 @@ class ChatMessageModel(BaseModel):
files: Optional[list] = None files: Optional[list] = None
sources: Optional[list] = None sources: Optional[list] = None
embeds: Optional[list] = None embeds: Optional[list] = None
meta: Optional[dict] = None
done: bool = True done: bool = True
status_history: Optional[list] = None status_history: Optional[list] = None
error: Optional[dict | str] = None error: Optional[dict | str] = None
usage: Optional[dict] = None usage: Optional[dict] = None
context_summary: Optional[str] = None
created_at: int created_at: int
updated_at: int updated_at: int
@ -207,66 +151,6 @@ class ChatMessageModel(BaseModel):
class ChatMessageTable: class ChatMessageTable:
@staticmethod
def _apply_message_data(message: ChatMessage, data: dict, now: int) -> None:
"""Overwrite only the fields the payload carries."""
if 'role' in data:
message.role = data['role']
if 'parent_id' in data or 'parentId' in data:
message.parent_id = data.get('parent_id') or data.get('parentId')
if 'content' in data:
message.content = data.get('content')
if 'output' in data:
message.output = data.get('output')
if 'model_id' in data or 'model' in data:
message.model_id = data.get('model_id') or data.get('model')
if 'files' in data:
message.files = data.get('files')
if 'sources' in data:
message.sources = data.get('sources')
if 'embeds' in data:
message.embeds = data.get('embeds')
if 'meta' in data:
message.meta = data.get('meta')
if 'done' in data:
message.done = data['done']
if 'status_history' in data or 'statusHistory' in data:
message.status_history = data.get('status_history') or data.get('statusHistory')
if 'error' in data:
message.error = data.get('error')
if 'context_summary' in data or 'contextSummary' in data:
message.context_summary = data.get('context_summary') or data.get('contextSummary')
usage = get_usage(data)
if usage:
existing_usage = normalize_usage(message.usage)
message.usage = existing_usage if usage == existing_usage else merge_usage(existing_usage, usage)
message.updated_at = now
@staticmethod
def _build_message(composite_id: str, chat_id: str, user_id: str, data: dict, now: int) -> ChatMessage:
return ChatMessage(
id=composite_id,
chat_id=chat_id,
user_id=user_id,
role=data.get('role', 'user'),
parent_id=data.get('parent_id') or data.get('parentId'),
content=data.get('content'),
output=data.get('output'),
model_id=data.get('model_id') or data.get('model'),
files=data.get('files'),
sources=data.get('sources'),
embeds=data.get('embeds'),
meta=data.get('meta'),
done=data.get('done', True),
status_history=data.get('status_history') or data.get('statusHistory'),
error=data.get('error'),
usage=get_usage(data),
context_summary=data.get('context_summary') or data.get('contextSummary'),
created_at=data.get('timestamp', now),
updated_at=now,
)
async def upsert_message( async def upsert_message(
self, self,
message_id: str, message_id: str,
@ -278,67 +162,80 @@ class ChatMessageTable:
"""Insert or update a chat message.""" """Insert or update a chat message."""
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
now = int(time.time()) now = int(time.time())
timestamp = data.get('timestamp', now)
# Use composite ID: {chat_id}-{message_id} # Use composite ID: {chat_id}-{message_id}
composite_id = f'{chat_id}-{message_id}' composite_id = f'{chat_id}-{message_id}'
message = await db.get(ChatMessage, composite_id) existing = await db.get(ChatMessage, composite_id)
if message: if existing:
self._apply_message_data(message, data, now) # Update existing
if 'role' in data:
existing.role = data['role']
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')
if 'output' in data:
existing.output = data.get('output')
if 'model_id' in data or 'model' in data:
existing.model_id = data.get('model_id') or data.get('model')
if 'files' in data:
existing.files = data.get('files')
if 'sources' in data:
existing.sources = data.get('sources')
if 'embeds' in data:
existing.embeds = data.get('embeds')
if 'done' in data:
existing.done = data.get('done', True)
if 'status_history' in data or 'statusHistory' in data:
existing.status_history = data.get('status_history') or data.get('statusHistory')
if 'error' in data:
existing.error = data.get('error')
# Extract and normalize usage
usage = get_usage(data)
if usage:
# Deep-merge: preserve existing keys not present in new data
# This prevents background tasks (follow-ups, title, tags)
# from accidentally clearing the primary response's token counts
existing.usage = {**(existing.usage or {}), **usage}
existing.updated_at = now
await db.commit()
await db.refresh(existing)
return ChatMessageModel.model_validate(existing)
else: else:
message = self._build_message(composite_id, chat_id, user_id, data, now) # Insert new
# Extract and normalize usage
usage = get_usage(data)
message = ChatMessage(
id=composite_id,
chat_id=chat_id,
user_id=user_id,
role=data.get('role', 'user'),
parent_id=data.get('parent_id') or data.get('parentId'),
content=data.get('content'),
output=data.get('output'),
model_id=data.get('model_id') or data.get('model'),
files=data.get('files'),
sources=data.get('sources'),
embeds=data.get('embeds'),
done=data.get('done', True),
status_history=data.get('status_history') or data.get('statusHistory'),
error=data.get('error'),
usage=usage,
created_at=timestamp,
updated_at=now,
)
db.add(message) db.add(message)
await db.commit()
await db.commit() await db.refresh(message)
return ChatMessageModel.model_validate(message) return ChatMessageModel.model_validate(message)
async def upsert_messages(
self,
chat_id: str,
user_id: str,
messages: dict[str, dict],
db: AsyncSession | None = None,
) -> None:
"""Insert or update the given messages of one chat."""
if not messages:
return
async with get_async_db_context(db) as db:
now = int(time.time())
result = await db.execute(
select(ChatMessage).filter(ChatMessage.id.in_([f'{chat_id}-{message_id}' for message_id in messages]))
)
existing_by_id = {row.id: row for row in result.scalars().all()}
for message_id, data in messages.items():
composite_id = f'{chat_id}-{message_id}'
message = existing_by_id.get(composite_id)
if message:
self._apply_message_data(message, data, now)
else:
db.add(self._build_message(composite_id, chat_id, user_id, data, now))
await db.commit()
async def get_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatMessageModel]: async def get_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatMessageModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
message = await db.get(ChatMessage, id) message = await db.get(ChatMessage, id)
return ChatMessageModel.model_validate(message) if message else None return ChatMessageModel.model_validate(message) if message else None
async def has_unfinished_assistant_by_chat_id(
self,
chat_id: str,
db: Optional[AsyncSession] = None,
) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChatMessage.id)
.where(ChatMessage.chat_id == chat_id)
.where(ChatMessage.role == 'assistant')
.where(ChatMessage.done.is_(False))
.limit(1)
)
return result.scalar_one_or_none() is not None
async def get_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[ChatMessageModel]: async def get_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[ChatMessageModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute( result = await db.execute(
@ -352,7 +249,6 @@ class ChatMessageTable:
'parent_id': 'parentId', 'parent_id': 'parentId',
'model_id': 'model', 'model_id': 'model',
'status_history': 'statusHistory', 'status_history': 'statusHistory',
'context_summary': 'contextSummary',
'created_at': 'timestamp', 'created_at': 'timestamp',
} }
# DB-internal columns excluded from the reconstructed message dict. # DB-internal columns excluded from the reconstructed message dict.
@ -544,44 +440,6 @@ class ChatMessageTable:
result = await db.execute(stmt) result = await db.execute(stmt)
return {row.model_id: row.count for row in result.all()} return {row.model_id: row.count for row in result.all()}
async def get_unique_counts_by_model(
self,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
group_id: Optional[str] = None,
db: Optional[AsyncSession] = None,
) -> dict[str, dict]:
"""Count distinct users and chats per model."""
async with get_async_db_context(db) as db:
from open_webui.models.groups import GroupMember
stmt = select(
ChatMessage.model_id,
func.count(distinct(ChatMessage.user_id)).label('unique_users'),
func.count(distinct(ChatMessage.chat_id)).label('unique_chats'),
).filter(
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
)
if start_date:
stmt = stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
stmt = stmt.filter(ChatMessage.created_at <= end_date)
if group_id:
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
stmt = stmt.group_by(ChatMessage.model_id)
result = await db.execute(stmt)
return {
row.model_id: {
'unique_users': row.unique_users,
'unique_chats': row.unique_chats,
}
for row in result.all()
}
async def get_token_usage_by_model( async def get_token_usage_by_model(
self, self,
start_date: Optional[int] = None, start_date: Optional[int] = None,
@ -680,233 +538,6 @@ class ChatMessageTable:
for row in result.all() for row in result.all()
} }
async def get_user_usage_summary(
self,
user_id: str,
start_date: Optional[int] = None,
end_date: Optional[int] = None,
include_active_days: bool = True,
timezone: Optional[str] = None,
db: Optional[AsyncSession] = None,
) -> dict:
async with get_async_db_context(db) as db:
bind = await db.connection()
dialect = bind.dialect.name
input_tokens, output_tokens = _token_columns(dialect)
messages_stmt = select(ChatMessage.role, func.count(ChatMessage.id).label('count')).filter(
ChatMessage.user_id == user_id,
)
token_stmt = select(
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
).filter(
ChatMessage.user_id == user_id,
ChatMessage.role == 'assistant',
ChatMessage.usage.isnot(None),
)
models_stmt = select(func.count(distinct(ChatMessage.model_id)).label('models_used')).filter(
ChatMessage.user_id == user_id,
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
)
if start_date:
messages_stmt = messages_stmt.filter(ChatMessage.created_at >= start_date)
token_stmt = token_stmt.filter(ChatMessage.created_at >= start_date)
models_stmt = models_stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
messages_stmt = messages_stmt.filter(ChatMessage.created_at <= end_date)
token_stmt = token_stmt.filter(ChatMessage.created_at <= end_date)
models_stmt = models_stmt.filter(ChatMessage.created_at <= end_date)
messages_result = await db.execute(messages_stmt.group_by(ChatMessage.role))
message_counts = {row.role: row.count for row in messages_result.all()}
token_result = (await db.execute(token_stmt)).one()
models_used = (await db.execute(models_stmt)).scalar() or 0
active_days = set()
if include_active_days:
tz = _timezone(timezone)
day_stmt = select(ChatMessage.created_at).filter(ChatMessage.user_id == user_id)
if start_date:
day_stmt = day_stmt.filter(ChatMessage.created_at >= start_date)
if end_date:
day_stmt = day_stmt.filter(ChatMessage.created_at <= end_date)
day_result = await db.execute(day_stmt)
active_days = {_date_key(row.created_at, tz) for row in day_result.all()}
input_total = int(token_result.input_tokens or 0)
output_total = int(token_result.output_tokens or 0)
return {
'messages': sum(message_counts.values()),
'user_messages': message_counts.get('user', 0),
'assistant_messages': message_counts.get('assistant', 0),
'input_tokens': input_total,
'output_tokens': output_total,
'total_tokens': input_total + output_total,
'models_used': int(models_used),
'active_days': len(active_days),
}
async def get_user_first_message_created_at(
self,
user_id: str,
db: Optional[AsyncSession] = None,
) -> Optional[int]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(func.min(ChatMessage.created_at)).filter(
ChatMessage.user_id == user_id,
ChatMessage.created_at.isnot(None),
)
)
value = result.scalar()
return int(value) if value else None
async def get_user_daily_usage(
self,
user_id: str,
start_date: int,
end_date: int,
timezone: Optional[str] = None,
db: Optional[AsyncSession] = None,
) -> list[dict]:
async with get_async_db_context(db) as db:
tz = _timezone(timezone)
bind = await db.connection()
dialect = bind.dialect.name
input_tokens, output_tokens = _token_columns(dialect)
stmt = select(
ChatMessage.created_at,
ChatMessage.chat_id,
ChatMessage.role,
ChatMessage.model_id,
ChatMessage.usage,
input_tokens.label('input_tokens'),
output_tokens.label('output_tokens'),
).filter(
ChatMessage.user_id == user_id,
ChatMessage.created_at >= start_date,
ChatMessage.created_at <= end_date,
)
result = await db.execute(stmt)
daily: dict[str, dict] = {}
for row in result.all():
date = _date_key(row.created_at, tz)
entry = daily.setdefault(
date,
{
'date': date,
'messages': 0,
'chat_ids': set(),
'tokens': 0,
'models': Counter(),
},
)
entry['messages'] += 1
entry['chat_ids'].add(row.chat_id)
if row.role == 'assistant' and row.model_id:
entry['models'][row.model_id] += 1
if row.usage:
entry['tokens'] += int(row.input_tokens or 0) + int(row.output_tokens or 0)
current = datetime.fromtimestamp(_normalize_timestamp(start_date), tz=tz).replace(
hour=0, minute=0, second=0, microsecond=0
)
end_dt = datetime.fromtimestamp(_normalize_timestamp(end_date), tz=tz).replace(
hour=0, minute=0, second=0, microsecond=0
)
while current <= end_dt:
date = current.strftime('%Y-%m-%d')
daily.setdefault(
date,
{'date': date, 'messages': 0, 'chat_ids': set(), 'tokens': 0, 'models': Counter()},
)
current += timedelta(days=1)
return [
{
'date': item['date'],
'messages': item['messages'],
'chats': len(item['chat_ids']),
'tokens': item['tokens'],
'models': dict(item['models']),
}
for item in sorted(daily.values(), key=lambda x: x['date'])
]
async def get_user_top_models(
self,
user_id: str,
start_date: int,
end_date: int,
limit: int = 5,
db: Optional[AsyncSession] = None,
) -> list[dict]:
async with get_async_db_context(db) as db:
bind = await db.connection()
dialect = bind.dialect.name
input_tokens, output_tokens = _token_columns(dialect)
stmt = (
select(
ChatMessage.model_id,
func.count(ChatMessage.id).label('messages'),
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
)
.filter(
ChatMessage.user_id == user_id,
ChatMessage.role == 'assistant',
ChatMessage.model_id.isnot(None),
ChatMessage.created_at >= start_date,
ChatMessage.created_at <= end_date,
)
.group_by(ChatMessage.model_id)
.order_by(func.count(ChatMessage.id).desc())
.limit(limit)
)
result = await db.execute(stmt)
return [
{
'model_id': row.model_id,
'messages': row.messages,
'input_tokens': int(row.input_tokens or 0),
'output_tokens': int(row.output_tokens or 0),
'total_tokens': int(row.input_tokens or 0) + int(row.output_tokens or 0),
}
for row in result.all()
]
async def get_user_top_tools(
self,
user_id: str,
start_date: int,
end_date: int,
limit: int = 5,
db: Optional[AsyncSession] = None,
) -> list[dict]:
async with get_async_db_context(db) as db:
stmt = select(ChatMessage.output, ChatMessage.meta).filter(
ChatMessage.user_id == user_id,
ChatMessage.created_at >= start_date,
ChatMessage.created_at <= end_date,
)
result = await db.execute(stmt)
counts: Counter[str] = Counter()
for output, meta in result.all():
for name in _extract_tool_names(output):
counts[name] += 1
for name in _extract_tool_names(meta):
counts[name] += 1
return [{'name': name, 'count': count} for name, count in counts.most_common(limit)]
async def get_message_count_by_user( async def get_message_count_by_user(
self, self,
start_date: Optional[int] = None, start_date: Optional[int] = None,

File diff suppressed because it is too large Load diff

View file

@ -1,382 +0,0 @@
"""Database-backed configuration with per-key storage.
Replaces the old single-row JSON blob machinery with a simple per-key model
mirroring cptr's Config.
Each config key is stored as its own row: key TEXT PK, value JSON.
Reads are direct DB lookups. Writes are explicit awaited upserts that raise on
failure (no more fire-and-forget create_task).
"""
from __future__ import annotations
import logging
import time
from typing import Any, ClassVar
from fastapi.encoders import jsonable_encoder
from open_webui.internal.db import Base, get_async_db
from sqlalchemy import JSON, BigInteger, Column, Text, delete, select
log = logging.getLogger(__name__)
API_CONFIG_KEYS = ('openai.api_configs', 'ollama.api_configs')
DICT_CONFIG_KEY_ALIASES = {
'openai.api_configs': ('OPENAI_API_CONFIGS',),
'ollama.api_configs': ('OLLAMA_API_CONFIGS',),
'rag.mineru_params': ('MINERU_PARAMS',),
'rag.docling_params': ('DOCLING_PARAMS',),
'web.search.linkup_search_params': ('LINKUP_SEARCH_PARAMS',),
'image_generation.automatic1111.api_params': ('AUTOMATIC1111_PARAMS',),
'image_generation.openai.params': ('IMAGES_OPENAI_API_PARAMS',),
'audio.tts.openai.params': ('AUDIO_TTS_OPENAI_PARAMS',),
'models.default_metadata': ('DEFAULT_MODEL_METADATA',),
'models.default_params': ('DEFAULT_MODEL_PARAMS',),
'task.model.params': ('TASK_MODEL_PARAMS',),
'ui.default_interface_settings': ('DEFAULT_INTERFACE_SETTINGS',),
'user.permissions': ('USER_PERMISSIONS',),
}
DICT_CONFIG_KEYS = tuple(DICT_CONFIG_KEY_ALIASES)
API_CONFIG_FIELDS = (
'enable',
'key',
'prefix_id',
'tags',
'model_ids',
'connection_type',
'provider',
'auth_type',
'headers',
'azure',
'api_type',
'api_version',
'extra_params',
'passthrough_params',
)
def _split_api_config_fragment(fragment: str) -> tuple[str, list[str]] | None:
if not fragment:
return None
first, _, rest = fragment.partition('.')
if first.isdigit() and rest:
return first, rest.split('.')
match: tuple[int, str] | None = None
for field in API_CONFIG_FIELDS:
marker = f'.{field}'
marker_index = fragment.rfind(marker)
if marker_index != -1 and (match is None or marker_index > match[0]):
match = (marker_index, field)
if match:
marker_index, field = match
connection_key = fragment[:marker_index]
field_path = fragment[marker_index + 1 :]
if connection_key:
return connection_key, field_path.split('.')
return None
def _assign_path(target: dict, path: list[str], value: Any) -> None:
current = target
for part in path[:-1]:
next_value = current.get(part)
if not isinstance(next_value, dict):
next_value = {}
current[part] = next_value
current = next_value
current[path[-1]] = value
def _json_value(value: Any) -> Any:
return jsonable_encoder(value)
# ── Model ────────────────────────────────────────────────────────────────────
class Config(Base):
"""Per-key config storage. Each row is one config key."""
__tablename__ = 'config'
key = Column(Text, primary_key=True)
value = Column(JSON, nullable=False)
updated_at = Column(BigInteger, nullable=True)
DEFAULTS: ClassVar[dict[str, Any]] = {}
PERSISTENT_ENABLED: ClassVar[bool] = True
OAUTH_PERSISTENT_ENABLED: ClassVar[bool] = False
# ── Class methods ────────────────────────────────────────
@classmethod
def configure(
cls,
*,
defaults: dict[str, Any] | None = None,
enable_persistent: bool = True,
enable_oauth_persistent: bool = False,
) -> None:
cls.DEFAULTS = dict(defaults or {})
cls.PERSISTENT_ENABLED = enable_persistent
cls.OAUTH_PERSISTENT_ENABLED = enable_oauth_persistent
@classmethod
def default_value(cls, key: str, default: Any = None) -> Any:
return cls.DEFAULTS.get(key, default)
@classmethod
def persistent_enabled_for(cls, key: str) -> bool:
if not cls.PERSISTENT_ENABLED:
return False
if key.startswith('oauth.') and not cls.OAUTH_PERSISTENT_ENABLED:
return False
return True
@staticmethod
async def get(key: str, default: Any = None) -> Any:
"""Get a config value by key. Returns default if not set."""
if not Config.persistent_enabled_for(key):
return Config.default_value(key, default)
async with get_async_db() as db:
row = await db.get(Config, key)
return row.value if row else Config.default_value(key, default)
@staticmethod
async def get_many(*keys: str) -> dict:
"""Get multiple config values. Returns {key: value} for keys that exist."""
disabled_values = {
key: Config.default_value(key)
for key in keys
if not Config.persistent_enabled_for(key) and key in Config.DEFAULTS
}
enabled_keys = {key for key in keys if Config.persistent_enabled_for(key)}
if not enabled_keys:
return disabled_values
async with get_async_db() as db:
result = await db.execute(select(Config).where(Config.key.in_(enabled_keys)))
values = {row.key: row.value for row in result.scalars().all()}
return {
key: values.get(key, Config.default_value(key))
for key in keys
if key in values or key in Config.DEFAULTS or key in disabled_values
}
@staticmethod
async def get_namespace(namespace: str) -> dict:
"""Get all config keys under a dotted namespace."""
default_values = {
key: value
for key, value in Config.DEFAULTS.items()
if key.startswith(f'{namespace}.') and not Config.persistent_enabled_for(key)
}
if not Config.PERSISTENT_ENABLED:
return default_values
async with get_async_db() as db:
result = await db.execute(select(Config).where(Config.key.like(f'{namespace}.%')))
values = {row.key: row.value for row in result.scalars().all()}
values.update(default_values)
return values
@staticmethod
async def get_all() -> dict:
"""Get all config as {key: value}."""
if not Config.PERSISTENT_ENABLED:
return dict(Config.DEFAULTS)
async with get_async_db() as db:
result = await db.execute(select(Config))
values = {row.key: row.value for row in result.scalars().all()}
if not Config.OAUTH_PERSISTENT_ENABLED:
values.update({key: value for key, value in Config.DEFAULTS.items() if key.startswith('oauth.')})
return values
@staticmethod
async def upsert(updates: dict) -> None:
"""Upsert multiple config key-value pairs. Raises on failure."""
persistent_updates = {}
for key, value in updates.items():
value = _json_value(value)
if Config.persistent_enabled_for(key):
persistent_updates[key] = value
else:
Config.DEFAULTS[key] = value
if not persistent_updates:
return
async with get_async_db() as db:
now = int(time.time())
for key, value in persistent_updates.items():
existing = await db.get(Config, key)
if existing:
existing.value = value
existing.updated_at = now
else:
db.add(Config(key=key, value=value, updated_at=now))
await db.commit()
@staticmethod
async def delete(key: str) -> bool:
"""Delete a config key. Returns True if it existed."""
async with get_async_db() as db:
row = await db.get(Config, key)
if row:
await db.delete(row)
await db.commit()
return True
return False
@staticmethod
async def clear() -> None:
"""Delete all config rows."""
async with get_async_db() as db:
await db.execute(delete(Config))
await db.commit()
@staticmethod
async def seed_defaults(defaults: dict) -> None:
"""Insert keys that don't yet exist in the DB.
Called at startup to ensure all known config keys have values.
Existing DB values take precedence over defaults.
"""
async with get_async_db() as db:
result = await db.execute(select(Config.key))
existing_keys = {row[0] for row in result.all()}
now = int(time.time())
new_count = 0
for key, value in defaults.items():
# Skip keys the DB is not authoritative for (e.g. oauth.* while
# ENABLE_OAUTH_PERSISTENT_CONFIG is off), matching the read paths.
if not Config.persistent_enabled_for(key):
continue
if key not in existing_keys:
value = _json_value(value)
db.add(Config(key=key, value=value, updated_at=now))
existing_keys.add(key)
new_count += 1
if new_count:
await db.commit()
log.info('Seeded %d new config defaults', new_count)
@staticmethod
async def rename_prefix(old_prefix: str, new_prefix: str) -> None:
"""Move persisted config keys from one dotted prefix to another."""
if not Config.PERSISTENT_ENABLED:
return
async with get_async_db() as db:
result = await db.execute(select(Config).where(Config.key.like(f'{old_prefix}.%')))
rows = result.scalars().all()
if not rows:
return
now = int(time.time())
moved_count = 0
deleted_count = 0
for row in rows:
new_key = f'{new_prefix}.{row.key.removeprefix(f"{old_prefix}.")}'
existing = await db.get(Config, new_key)
if existing is None:
db.add(Config(key=new_key, value=row.value, updated_at=now))
moved_count += 1
else:
deleted_count += 1
await db.delete(row)
await db.commit()
log.info(
'Renamed %d config keys from %s.* to %s.*; deleted %d old duplicates',
moved_count,
old_prefix,
new_prefix,
deleted_count,
)
@staticmethod
async def repair_config_rows() -> None:
"""Repair known legacy config row shapes."""
if not Config.PERSISTENT_ENABLED:
return
async with get_async_db() as db:
repaired_keys: list[str] = []
orphan_keys: list[str] = []
default_model_keys: list[str] = []
now = int(time.time())
for config_key, aliases in DICT_CONFIG_KEY_ALIASES.items():
prefixes = (config_key, *aliases)
rows = []
for key_prefix in prefixes:
result = await db.execute(select(Config).where(Config.key.like(f'{key_prefix}.%')))
rows.extend(result.scalars().all())
if not rows:
continue
existing = await db.get(Config, config_key)
repaired = existing.value if existing and isinstance(existing.value, dict) else {}
repaired_any = False
for row in rows:
fragment = None
for key_prefix in prefixes:
prefix = f'{key_prefix}.'
if row.key.startswith(prefix):
fragment = row.key.removeprefix(prefix)
break
if fragment is None:
continue
if config_key in API_CONFIG_KEYS:
split = _split_api_config_fragment(fragment)
if not split:
continue
object_key, field_path = split
else:
object_key, field_path = None, fragment.split('.')
target = repaired
if object_key is not None:
target = repaired.setdefault(object_key, {})
if not isinstance(target, dict):
continue
_assign_path(target, field_path, row.value)
orphan_keys.append(row.key)
repaired_any = True
if not repaired_any:
continue
if existing:
existing.value = repaired
existing.updated_at = now
else:
db.add(Config(key=config_key, value=repaired, updated_at=now))
repaired_keys.append(config_key)
if orphan_keys:
await db.execute(delete(Config).where(Config.key.in_(orphan_keys)))
for key in ('ui.default_models', 'ui.default_pinned_models'):
row = await db.get(Config, key)
if not row or not isinstance(row.value, list):
continue
row.value = ','.join(model_id for model_id in (str(item).strip() for item in row.value) if model_id)
row.updated_at = now
default_model_keys.append(key)
if repaired_keys or orphan_keys or default_model_keys:
await db.commit()
if repaired_keys or orphan_keys:
log.info('Repaired flattened dict config rows for %s', ', '.join(repaired_keys))
if default_model_keys:
log.info('Repaired default model config rows for %s', ', '.join(default_model_keys))

View file

@ -133,12 +133,6 @@ class ModelHistoryEntry(BaseModel):
lost: int lost: int
class ModelHistoryCounts(BaseModel):
date: str
won: int = 0
lost: int = 0
class ModelHistoryResponse(BaseModel): class ModelHistoryResponse(BaseModel):
model_id: str model_id: str
history: list[ModelHistoryEntry] history: list[ModelHistoryEntry]
@ -165,6 +159,7 @@ class FeedbackTable:
result = Feedback(**feedback.model_dump()) result = Feedback(**feedback.model_dump())
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
if result: if result:
return FeedbackModel.model_validate(result) return FeedbackModel.model_validate(result)
else: else:
@ -221,15 +216,12 @@ class FeedbackTable:
) -> FeedbackListResponse: ) -> FeedbackListResponse:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
stmt = select(Feedback, User).join(User, Feedback.user_id == User.id) stmt = select(Feedback, User).join(User, Feedback.user_id == User.id)
count_stmt = select(func.count(Feedback.id)).select_from(Feedback).join(User, Feedback.user_id == User.id)
if filter: if filter:
# Apply model_id filter (exact match) # Apply model_id filter (exact match)
model_id = filter.get('model_id') model_id = filter.get('model_id')
if model_id: if model_id:
model_id_filter = Feedback.data['model_id'].as_string() == model_id stmt = stmt.filter(Feedback.data['model_id'].as_string() == model_id)
stmt = stmt.filter(model_id_filter)
count_stmt = count_stmt.filter(model_id_filter)
order_by = filter.get('order_by') order_by = filter.get('order_by')
direction = filter.get('direction') direction = filter.get('direction')
@ -258,9 +250,9 @@ class FeedbackTable:
else: else:
stmt = stmt.order_by(Feedback.created_at.desc()) stmt = stmt.order_by(Feedback.created_at.desc())
# Count before pagination without wrapping the ordered item query. # Count BEFORE pagination
count_result = await db.execute(count_stmt) count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar() or 0 total = count_result.scalar()
if skip: if skip:
stmt = stmt.offset(skip) stmt = stmt.offset(skip)
@ -383,45 +375,6 @@ class FeedbackTable:
return result return result
async def get_model_feedback_counts_by_day(
self,
model_id: str,
start_date: Optional[int] = None,
db: Optional[AsyncSession] = None,
) -> list[ModelHistoryCounts]:
"""Get aggregated feedback counts per day for a model, preserving all matching days."""
from collections import defaultdict
from datetime import datetime
async with get_async_db_context(db) as db:
stmt = select(Feedback.created_at, Feedback.data).filter(Feedback.data['model_id'].as_string() == model_id)
if start_date is not None:
stmt = stmt.filter(Feedback.created_at >= start_date)
result = await db.execute(stmt.order_by(Feedback.created_at.asc()))
rows = result.all()
daily_counts = defaultdict(lambda: {'won': 0, 'lost': 0})
for created_at, data in rows:
if not data:
continue
rating_str = str(data.get('rating', ''))
if rating_str not in ('1', '-1'):
continue
date_str = datetime.fromtimestamp(created_at).strftime('%Y-%m-%d')
if rating_str == '1':
daily_counts[date_str]['won'] += 1
else:
daily_counts[date_str]['lost'] += 1
return [
ModelHistoryCounts(date=date_str, won=counts['won'], lost=counts['lost'])
for date_str, counts in sorted(daily_counts.items())
]
async def get_feedbacks_by_type(self, type: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]: async def get_feedbacks_by_type(self, type: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute(select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc())) result = await db.execute(select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()))

View file

@ -142,6 +142,7 @@ class FilesTable:
result = File(**file.model_dump()) result = File(**file.model_dump())
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
if result: if result:
return FileModel.model_validate(result) return FileModel.model_validate(result)
else: else:
@ -200,18 +201,6 @@ class FilesTable:
result = await db.execute(select(File)) result = await db.execute(select(File))
return [FileModel.model_validate(file) for file in result.scalars().all()] return [FileModel.model_validate(file) for file in result.scalars().all()]
async def count_files_by_user_id(
self,
user_id: str | None = None,
db: AsyncSession | None = None,
) -> int:
async with get_async_db_context(db) as db:
stmt = select(func.count(File.id))
if user_id:
stmt = stmt.filter_by(user_id=user_id)
result = await db.execute(stmt)
return result.scalar() or 0
async def check_access_by_user_id(self, id, user_id, permission='write', db: AsyncSession | None = None) -> bool: async def check_access_by_user_id(self, id, user_id, permission='write', db: AsyncSession | None = None) -> bool:
file = await self.get_file_by_id(id, db=db) file = await self.get_file_by_id(id, db=db)
if not file: if not file:

View file

@ -6,7 +6,7 @@ from typing import Optional
from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.internal.db import Base, JSONField, get_async_db_context
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from sqlalchemy import JSON, BigInteger, Boolean, Column, Text, delete, func, select, or_, and_ from sqlalchemy import JSON, BigInteger, Boolean, Column, Text, delete, func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -58,21 +58,6 @@ class FolderNameIdResponse(BaseModel):
meta: Optional[FolderMetadataResponse] = None meta: Optional[FolderMetadataResponse] = None
parent_id: Optional[str] = None parent_id: Optional[str] = None
is_expanded: bool = False is_expanded: bool = False
unread_count: int = 0
created_at: int
updated_at: int
class SharedFolderResponse(BaseModel):
id: str
name: str
parent_id: Optional[str] = None
user_id: str
owner_name: Optional[str] = None
permission: str = 'read'
access_grants: list = []
is_expanded: bool = False
meta: Optional[dict] = None
created_at: int created_at: int
updated_at: int updated_at: int
@ -145,71 +130,16 @@ class FolderTable:
except Exception: except Exception:
return None return None
async def get_folder_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FolderModel]:
"""Fetch folder by ID only (no user_id filter). Used for shared access."""
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id))
folder = result.scalars().first()
if not folder:
return None
return FolderModel.model_validate(folder)
except Exception:
return None
async def get_folders_by_ids(self, ids: list[str], db: AsyncSession | None = None) -> list[FolderModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter(Folder.id.in_(ids)).order_by(Folder.updated_at.desc()))
return [FolderModel.model_validate(folder) for folder in result.scalars().all()]
async def get_shared_folder_ids_for_user(
self, user_id: str, user_group_ids: set[str], db: Optional[AsyncSession] = None
) -> dict[str, str]:
"""
Returns {folder_id: highest_permission} for all folders shared with user.
Checks direct user grants, group grants, and public (user:*) grants.
"""
from open_webui.models.access_grants import AccessGrant
async with get_async_db_context(db) as db:
conditions = [
and_(AccessGrant.principal_type == 'user', AccessGrant.principal_id == '*'),
and_(AccessGrant.principal_type == 'user', AccessGrant.principal_id == user_id),
]
if user_group_ids:
conditions.append(
and_(AccessGrant.principal_type == 'group', AccessGrant.principal_id.in_(user_group_ids))
)
result = await db.execute(
select(AccessGrant).filter(
AccessGrant.resource_type == 'folder',
or_(*conditions),
)
)
grants = result.scalars().all()
# Build {folder_id: highest_permission} ('write' > 'read')
folder_perms = {}
for g in grants:
existing = folder_perms.get(g.resource_id)
if existing != 'write':
folder_perms[g.resource_id] = g.permission
return folder_perms
async def get_children_folders_by_id_and_user_id( async def get_children_folders_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[list[FolderModel]]: ) -> Optional[list[FolderModel]]:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
folders = [] folders = []
seen_ids = {id}
async def get_children(folder): async def get_children(folder):
children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db) children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
for child in children: for child in children:
if child.id in seen_ids:
continue
seen_ids.add(child.id)
await get_children(child) await get_children(child)
folders.append(child) folders.append(child)
@ -239,9 +169,7 @@ class FolderTable:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
# Check if folder exists # Check if folder exists
result = await db.execute( result = await db.execute(
select(Folder) select(Folder).filter_by(parent_id=parent_id, user_id=user_id).filter(Folder.name.ilike(name))
.filter_by(parent_id=parent_id, user_id=user_id)
.filter(func.lower(Folder.name) == func.lower(name))
) )
folder = result.scalars().first() folder = result.scalars().first()
@ -257,32 +185,9 @@ class FolderTable:
self, parent_id: Optional[str], user_id: str, db: Optional[AsyncSession] = None self, parent_id: Optional[str], user_id: str, db: Optional[AsyncSession] = None
) -> list[FolderModel]: ) -> list[FolderModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute( result = await db.execute(select(Folder).filter_by(parent_id=parent_id, user_id=user_id))
select(Folder).filter_by(parent_id=parent_id, user_id=user_id).order_by(Folder.updated_at.desc())
)
return [FolderModel.model_validate(folder) for folder in result.scalars().all()] return [FolderModel.model_validate(folder) for folder in result.scalars().all()]
async def get_folder_ids_by_id_and_user_id_in_subtree(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> list[str]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
folder = result.scalars().first()
if not folder:
return []
folder_ids = {folder.id}
folders = [FolderModel.model_validate(folder)]
while folders:
current_folder = folders.pop()
children = await self.get_folders_by_parent_id_and_user_id(current_folder.id, user_id, db=db)
for child in children:
if child.id not in folder_ids:
folder_ids.add(child.id)
folders.append(child)
return list(folder_ids)
async def update_folder_parent_id_by_id_and_user_id( async def update_folder_parent_id_by_id_and_user_id(
self, self,
id: str, id: str,
@ -391,15 +296,11 @@ class FolderTable:
return folder_ids return folder_ids
folder_ids.append(folder.id) folder_ids.append(folder.id)
seen_ids = {folder.id}
# Delete all children folders # Delete all children folders
async def delete_children(folder): async def delete_children(folder):
folder_children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db) folder_children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
for folder_child in folder_children: for folder_child in folder_children:
if folder_child.id in seen_ids:
continue
seen_ids.add(folder_child.id)
await delete_children(folder_child) await delete_children(folder_child)
folder_ids.append(folder_child.id) folder_ids.append(folder_child.id)

View file

@ -7,8 +7,7 @@ import time
# local imports # local imports
from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.users import User, UserResponse, Users, UserSettings from open_webui.models.users import UserModel, UserResponse, Users
from open_webui.utils.valves import decrypt_valves, encrypt_valves
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Boolean, Column, Index, String, Text, delete, select, update from sqlalchemy import BigInteger, Boolean, Column, Index, String, Text, delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -42,7 +41,7 @@ class FunctionMeta(BaseModel):
class FunctionModel(BaseModel): class FunctionModel(BaseModel):
id: str id: str
user_id: str | None = None # may be null for legacy/malformed records user_id: str
name: str name: str
type: str type: str
content: str content: str
@ -58,7 +57,7 @@ class FunctionModel(BaseModel):
# --- form / schema definitions --- # --- form / schema definitions ---
class FunctionWithValvesModel(BaseModel): class FunctionWithValvesModel(BaseModel):
id: str id: str
user_id: str | None = None # may be null for legacy/malformed records user_id: str
name: str name: str
type: str type: str
content: str content: str
@ -79,7 +78,7 @@ class FunctionWithValvesModel(BaseModel):
class FunctionResponse(BaseModel): class FunctionResponse(BaseModel):
id: str id: str
user_id: str | None = None # may be null for legacy/malformed records user_id: str
type: str type: str
name: str name: str
meta: FunctionMeta meta: FunctionMeta
@ -129,6 +128,7 @@ class FunctionsTable:
result = Function(**function.model_dump()) result = Function(**function.model_dump())
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
if result: if result:
return FunctionModel.model_validate(result) return FunctionModel.model_validate(result)
else: else:
@ -143,8 +143,7 @@ class FunctionsTable:
functions: list[FunctionWithValvesModel], functions: list[FunctionWithValvesModel],
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> list[FunctionWithValvesModel]: ) -> list[FunctionWithValvesModel]:
# Synchronize functions by updating existing ones, inserting new ones, # Synchronize functions for a user by updating existing ones, inserting new ones, and removing those that are no longer present.
# and removing those that are no longer present.
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
# Get existing functions # Get existing functions
@ -157,15 +156,24 @@ class FunctionsTable:
# Update or insert functions # Update or insert functions
for func in functions: for func in functions:
func_data = func.model_dump()
func_data['valves'] = encrypt_valves(func_data['valves']) if func_data.get('valves') else None
func_data['user_id'] = user_id
func_data['updated_at'] = int(time.time())
if func.id in existing_ids: if func.id in existing_ids:
await db.execute(update(Function).filter_by(id=func.id).values(**func_data)) await db.execute(
update(Function)
.filter_by(id=func.id)
.values(
**func.model_dump(),
user_id=user_id,
updated_at=int(time.time()),
)
)
else: else:
new_func = Function(**func_data) new_func = Function(
**{
**func.model_dump(),
'user_id': user_id,
'updated_at': int(time.time()),
}
)
db.add(new_func) db.add(new_func)
# Remove functions that are no longer present # Remove functions that are no longer present
@ -219,15 +227,7 @@ class FunctionsTable:
functions = result.scalars().all() functions = result.scalars().all()
if include_valves: if include_valves:
return [ return [FunctionWithValvesModel.model_validate(function) for function in functions]
FunctionWithValvesModel.model_validate(
{
**FunctionModel.model_validate(function).model_dump(),
'valves': decrypt_valves(function.valves),
}
)
for function in functions
]
else: else:
return [FunctionModel.model_validate(function) for function in functions] return [FunctionModel.model_validate(function) for function in functions]
@ -274,18 +274,6 @@ class FunctionsTable:
result = await db.execute(select(Function).filter_by(type='filter', is_active=True, is_global=True)) result = await db.execute(select(Function).filter_by(type='filter', is_active=True, is_global=True))
return [FunctionModel.model_validate(function) for function in result.scalars().all()] return [FunctionModel.model_validate(function) for function in result.scalars().all()]
async def get_active_function_ids_by_type(
self, type: str, db: AsyncSession | None = None
) -> list[tuple[str, bool]]:
"""Return (id, is_global) for active functions without fetching plugin source."""
async with get_async_db_context(db) as db:
result = await db.execute(select(Function.id, Function.is_global).filter_by(type=type, is_active=True))
return [(id, bool(is_global)) for id, is_global in result.all()]
async def get_active_filter_ids(self, db: AsyncSession | None = None) -> list[tuple[str, bool]]:
"""Return (id, is_global) for active filters without fetching plugin source."""
return await self.get_active_function_ids_by_type('filter', db=db)
async def get_global_action_functions(self, db: AsyncSession | None = None) -> list[FunctionModel]: async def get_global_action_functions(self, db: AsyncSession | None = None) -> list[FunctionModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter_by(type='action', is_active=True, is_global=True)) result = await db.execute(select(Function).filter_by(type='action', is_active=True, is_global=True))
@ -294,8 +282,8 @@ class FunctionsTable:
async def get_function_valves_by_id(self, id: str, db: AsyncSession | None = None) -> dict | None: async def get_function_valves_by_id(self, id: str, db: AsyncSession | None = None) -> dict | None:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
try: try:
result = await db.execute(select(Function.valves).filter_by(id=id)) function = await db.get(Function, id)
return decrypt_valves(result.scalar_one_or_none()) return function.valves if function.valves else {}
except Exception as e: except Exception as e:
log.exception(f'Error getting function valves by id {id}: {e}') log.exception(f'Error getting function valves by id {id}: {e}')
return None return None
@ -311,7 +299,8 @@ class FunctionsTable:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids))) result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids)))
return {id: decrypt_valves(valves) for id, valves in result.all()} functions = result.all()
return {f.id: (f.valves if f.valves else {}) for f in functions}
except Exception as e: except Exception as e:
log.exception(f'Error batch-fetching function valves: {e}') log.exception(f'Error batch-fetching function valves: {e}')
return {} return {}
@ -322,9 +311,10 @@ class FunctionsTable:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
try: try:
function = await db.get(Function, id) function = await db.get(Function, id)
function.valves = encrypt_valves(valves) function.valves = valves
function.updated_at = int(time.time()) function.updated_at = int(time.time())
await db.commit() await db.commit()
await db.refresh(function)
return FunctionModel.model_validate(function) return FunctionModel.model_validate(function)
except Exception: except Exception:
return None return None
@ -344,6 +334,7 @@ class FunctionsTable:
function.updated_at = int(time.time()) function.updated_at = int(time.time())
await db.commit() await db.commit()
await db.refresh(function)
return FunctionModel.model_validate(function) return FunctionModel.model_validate(function)
else: else:
return None return None
@ -355,11 +346,8 @@ class FunctionsTable:
self, id: str, user_id: str, db: AsyncSession | None = None self, id: str, user_id: str, db: AsyncSession | None = None
) -> dict | None: ) -> dict | None:
try: try:
async with get_async_db_context(db) as db: user = await Users.get_user_by_id(user_id, db=db)
result = await db.execute(select(User.settings).filter_by(id=user_id)) user_settings = user.settings.model_dump() if user.settings else {}
settings = result.scalar_one_or_none()
user_settings = UserSettings(**settings).model_dump() if settings else {}
# Check if user has "functions" and "valves" settings # Check if user has "functions" and "valves" settings
if 'functions' not in user_settings: if 'functions' not in user_settings:
@ -367,8 +355,8 @@ class FunctionsTable:
if 'valves' not in user_settings['functions']: if 'valves' not in user_settings['functions']:
user_settings['functions']['valves'] = {} user_settings['functions']['valves'] = {}
return decrypt_valves(user_settings['functions']['valves'].get(id)) return user_settings['functions']['valves'].get(id, {})
except Exception: except Exception as e:
log.exception(f'Error getting user values by id {id} and user id {user_id}') log.exception(f'Error getting user values by id {id} and user id {user_id}')
return None return None
@ -385,12 +373,12 @@ class FunctionsTable:
if 'valves' not in user_settings['functions']: if 'valves' not in user_settings['functions']:
user_settings['functions']['valves'] = {} user_settings['functions']['valves'] = {}
user_settings['functions']['valves'][id] = encrypt_valves(valves) user_settings['functions']['valves'][id] = valves
# Update the user settings in the database # Update the user settings in the database
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
return valves return user_settings['functions']['valves'][id]
except Exception as e: except Exception as e:
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}') log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
return None return None

View file

@ -1,3 +1,4 @@
import json
import logging import logging
import time import time
import uuid import uuid
@ -12,7 +13,6 @@ from sqlalchemy import (
BigInteger, BigInteger,
Column, Column,
ForeignKey, ForeignKey,
Index,
String, String,
Text, Text,
and_, and_,
@ -72,8 +72,6 @@ class GroupModel(BaseModel):
class GroupMember(Base): class GroupMember(Base):
__tablename__ = 'group_member' __tablename__ = 'group_member'
# The table's (group_id, user_id) unique constraint cannot serve user_id lookups.
__table_args__ = (Index('ix_group_member_user_id_group_id', 'user_id', 'group_id'),)
id = Column(Text, unique=True, primary_key=True) id = Column(Text, unique=True, primary_key=True)
group_id = Column( group_id = Column(

View file

@ -1,9 +1,9 @@
import json
import logging import logging
import time import time
import uuid import uuid
from typing import Optional from typing import Optional
from open_webui.config import RAG_FILE_CONTENT_SEARCH_MAX_CHARS
from open_webui.internal.db import Base, JSONField, get_async_db_context 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.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.files import ( from open_webui.models.files import (
@ -31,13 +31,9 @@ from sqlalchemy import (
update, update,
) )
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import defer
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
# Columns the knowledge base list may be ordered by; anything else falls back to the default.
KNOWLEDGE_SORTABLE_FIELDS = {'name', 'created_at', 'updated_at'}
#################### ####################
# Knowledge DB Schema # Knowledge DB Schema
# Let what was gathered here outlast the one who gathered it, # Let what was gathered here outlast the one who gathered it,
@ -151,7 +147,6 @@ class KnowledgeDirectoryForm(BaseModel):
#################### ####################
class KnowledgeUserModel(KnowledgeModel): class KnowledgeUserModel(KnowledgeModel):
user: Optional[UserResponse] = None user: Optional[UserResponse] = None
file_count: int | None = None
class KnowledgeResponse(KnowledgeModel): class KnowledgeResponse(KnowledgeModel):
@ -194,11 +189,11 @@ class KnowledgeTable:
access_grants: Optional[list[AccessGrantModel]] = None, access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
) -> KnowledgeModel: ) -> KnowledgeModel:
knowledge_model = KnowledgeModel.model_validate(knowledge) knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump(exclude={'access_grants'})
knowledge_model.access_grants = ( knowledge_data['access_grants'] = (
access_grants if access_grants is not None else await self._get_access_grants(knowledge_model.id, db=db) access_grants if access_grants is not None else await self._get_access_grants(knowledge_data['id'], db=db)
) )
return knowledge_model return KnowledgeModel.model_validate(knowledge_data)
async def insert_new_knowledge( async def insert_new_knowledge(
self, user_id: str, form_data: KnowledgeForm, db: Optional[AsyncSession] = None self, user_id: str, form_data: KnowledgeForm, db: Optional[AsyncSession] = None
@ -291,17 +286,6 @@ class KnowledgeTable:
elif view_option == 'shared': elif view_option == 'shared':
stmt = stmt.filter(Knowledge.user_id != user_id) stmt = stmt.filter(Knowledge.user_id != user_id)
source = filter.get('source')
if source == 'external':
stmt = stmt.filter(Knowledge.meta['source'].as_string() == 'external')
elif source == 'local':
stmt = stmt.filter(
or_(
Knowledge.meta.is_(None),
Knowledge.meta['source'].as_string() != 'external',
)
)
stmt = AccessGrants.has_permission_filter( stmt = AccessGrants.has_permission_filter(
db=db, db=db,
query=stmt, query=stmt,
@ -311,17 +295,7 @@ class KnowledgeTable:
permission='read', permission='read',
) )
order_by = (filter or {}).get('order_by') stmt = stmt.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc())
direction = (filter or {}).get('direction')
if order_by in KNOWLEDGE_SORTABLE_FIELDS:
column = getattr(Knowledge, order_by)
if (direction or 'desc').lower() == 'asc':
stmt = stmt.order_by(column.asc(), Knowledge.id.asc())
else:
stmt = stmt.order_by(column.desc(), Knowledge.id.asc())
else:
stmt = stmt.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc())
count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar() total = count_result.scalar()
@ -335,14 +309,6 @@ class KnowledgeTable:
knowledge_ids = [kb.id for kb, _ in items] knowledge_ids = [kb.id for kb, _ in items]
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db) grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
file_counts = {}
if knowledge_ids:
file_count_result = await db.execute(
select(KnowledgeFile.knowledge_id, func.count(KnowledgeFile.id))
.where(KnowledgeFile.knowledge_id.in_(knowledge_ids))
.group_by(KnowledgeFile.knowledge_id)
)
file_counts = dict(file_count_result.all())
knowledge_bases = [] knowledge_bases = []
for knowledge_base, user in items: for knowledge_base, user in items:
@ -357,7 +323,6 @@ class KnowledgeTable:
) )
).model_dump(), ).model_dump(),
'user': (UserModel.model_validate(user).model_dump() if user else None), 'user': (UserModel.model_validate(user).model_dump() if user else None),
'file_count': file_counts.get(knowledge_base.id, 0),
} }
) )
) )
@ -404,7 +369,6 @@ class KnowledgeTable:
# to avoid PostgreSQL "invalid memory alloc request # to avoid PostgreSQL "invalid memory alloc request
# size" on large extracted-content rows (#24670). # size" on large extracted-content rows (#24670).
content_text = File.data['content'].as_string() content_text = File.data['content'].as_string()
content_text = func.substr(content_text, 1, RAG_FILE_CONTENT_SEARCH_MAX_CHARS)
search_filter = or_( search_filter = or_(
File.filename.ilike(f'%{q}%'), File.filename.ilike(f'%{q}%'),
content_text.ilike(f'%{q}%'), content_text.ilike(f'%{q}%'),
@ -441,7 +405,6 @@ class KnowledgeTable:
if limit: if limit:
stmt = stmt.limit(limit) stmt = stmt.limit(limit)
stmt = stmt.options(defer(File.data))
result = await db.execute(stmt) result = await db.execute(stmt)
rows = result.all() rows = result.all()
@ -449,13 +412,7 @@ class KnowledgeTable:
for file, user, knowledge in rows: for file, user, knowledge in rows:
items.append( items.append(
FileUserResponse( FileUserResponse(
id=file.id, **FileModel.model_validate(file).model_dump(),
user_id=file.user_id,
hash=file.hash,
filename=file.filename,
meta=file.meta,
created_at=file.created_at,
updated_at=file.updated_at,
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
collection=(await self._to_knowledge_model(knowledge, db=db)).model_dump(), collection=(await self._to_knowledge_model(knowledge, db=db)).model_dump(),
) )
@ -467,22 +424,14 @@ class KnowledgeTable:
print('search_knowledge_files error:', e) print('search_knowledge_files error:', e)
return KnowledgeFileListResponse(items=[], total=0) return KnowledgeFileListResponse(items=[], total=0)
async def check_access_by_user_id( async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool:
self,
id,
user_id,
permission='write',
db: Optional[AsyncSession] = None,
user_group_ids: set[str] | None = None,
) -> bool:
knowledge = await self.get_knowledge_by_id(id, db=db) knowledge = await self.get_knowledge_by_id(id, db=db)
if not knowledge: if not knowledge:
return False return False
if knowledge.user_id == user_id: if knowledge.user_id == user_id:
return True return True
if user_group_ids is None: user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db) user_group_ids = {group.id for group in user_groups}
user_group_ids = {group.id for group in user_groups}
return await AccessGrants.has_access( return await AccessGrants.has_access(
user_id=user_id, user_id=user_id,
resource_type='knowledge', resource_type='knowledge',
@ -492,6 +441,28 @@ class KnowledgeTable:
db=db, db=db,
) )
async def get_knowledge_bases_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
) -> list[KnowledgeUserModel]:
knowledge_bases = await self.get_knowledge_bases(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
result = []
for knowledge_base in knowledge_bases:
if knowledge_base.user_id == user_id:
result.append(knowledge_base)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge_base.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
):
result.append(knowledge_base)
return result
async def get_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]: async def get_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
@ -501,6 +472,29 @@ class KnowledgeTable:
except Exception: except Exception:
return None return None
async def get_knowledge_by_id_and_user_id(
self, id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeModel]:
knowledge = await self.get_knowledge_by_id(id, db=db)
if not knowledge:
return None
if knowledge.user_id == user_id:
return knowledge
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
if await AccessGrants.has_access(
user_id=user_id,
resource_type='knowledge',
resource_id=knowledge.id,
permission='write',
user_group_ids=user_group_ids,
db=db,
):
return knowledge
return None
async def get_knowledges_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[KnowledgeModel]: async def get_knowledges_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[KnowledgeModel]:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
@ -560,7 +554,6 @@ class KnowledgeTable:
# to avoid PostgreSQL memory allocation failures on # to avoid PostgreSQL memory allocation failures on
# large content (#24670). # large content (#24670).
content_text = File.data['content'].as_string() content_text = File.data['content'].as_string()
content_text = func.substr(content_text, 1, RAG_FILE_CONTENT_SEARCH_MAX_CHARS)
stmt = stmt.filter( stmt = stmt.filter(
or_( or_(
File.filename.ilike(f'%{query_key}%'), File.filename.ilike(f'%{query_key}%'),
@ -599,23 +592,17 @@ class KnowledgeTable:
if limit: if limit:
stmt = stmt.limit(limit) stmt = stmt.limit(limit)
stmt = stmt.options(defer(File.data))
result = await db.execute(stmt) result = await db.execute(stmt)
items = result.all() items = result.all()
files = [ files = []
FileUserResponse( for file, user in items:
id=file.id, files.append(
user_id=file.user_id, FileUserResponse(
hash=file.hash, **FileModel.model_validate(file).model_dump(),
filename=file.filename, user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
meta=file.meta, )
created_at=file.created_at,
updated_at=file.updated_at,
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
) )
for file, user in items
]
return KnowledgeFileListResponse( return KnowledgeFileListResponse(
items=files, items=files,
@ -625,7 +612,6 @@ class KnowledgeTable:
db=db, db=db,
), ),
breadcrumbs=await self.get_directory_breadcrumbs( breadcrumbs=await self.get_directory_breadcrumbs(
knowledge_id,
filter.get('directory_id') if filter else None, filter.get('directory_id') if filter else None,
db=db, db=db,
), ),
@ -651,25 +637,9 @@ class KnowledgeTable:
async def get_file_metadatas_by_id( async def get_file_metadatas_by_id(
self, knowledge_id: str, db: Optional[AsyncSession] = None self, knowledge_id: str, db: Optional[AsyncSession] = None
) -> list[FileMetadataResponse]: ) -> list[FileMetadataResponse]:
"""Column-only listing: File.data holds each file's full extracted
text, which metadata views must never load."""
try: try:
async with get_async_db_context(db) as db: files = await self.get_files_by_id(knowledge_id, db=db)
result = await db.execute( return [FileMetadataResponse(**file.model_dump()) for file in files]
select(File.id, File.hash, File.meta, File.created_at, File.updated_at)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
.filter(KnowledgeFile.knowledge_id == knowledge_id)
)
return [
FileMetadataResponse(
id=row.id,
hash=row.hash,
meta=row.meta,
created_at=row.created_at,
updated_at=row.updated_at,
)
for row in result.all()
]
except Exception: except Exception:
return [] return []
@ -776,8 +746,8 @@ class KnowledgeTable:
log.exception(e) log.exception(e)
return None return None
async def update_knowledge_meta_by_id( async def update_knowledge_data_by_id(
self, id: str, meta: dict, db: Optional[AsyncSession] = None self, id: str, data: dict, db: Optional[AsyncSession] = None
) -> Optional[KnowledgeModel]: ) -> Optional[KnowledgeModel]:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
@ -785,7 +755,7 @@ class KnowledgeTable:
update(Knowledge) update(Knowledge)
.filter_by(id=id) .filter_by(id=id)
.values( .values(
meta=meta, data=data,
updated_at=int(time.time()), updated_at=int(time.time()),
) )
) )
@ -909,7 +879,6 @@ class KnowledgeTable:
async def get_directory_breadcrumbs( async def get_directory_breadcrumbs(
self, self,
knowledge_id: str,
directory_id: Optional[str], directory_id: Optional[str],
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
) -> list[KnowledgeDirectoryModel]: ) -> list[KnowledgeDirectoryModel]:
@ -924,10 +893,7 @@ class KnowledgeTable:
while current_id and current_id not in seen: while current_id and current_id not in seen:
seen.add(current_id) seen.add(current_id)
# Scoped by knowledge base so a caller-supplied id cannot walk another one's tree. result = await db.execute(select(KnowledgeDirectory).filter_by(id=current_id))
result = await db.execute(
select(KnowledgeDirectory).filter_by(id=current_id, knowledge_id=knowledge_id)
)
directory = result.scalars().first() directory = result.scalars().first()
if not directory: if not directory:
break break
@ -1074,26 +1040,6 @@ class KnowledgeTable:
for child_id in child_ids: for child_id in child_ids:
await self._delete_files_in_subtree(child_id, db=db) await self._delete_files_in_subtree(child_id, db=db)
async def get_files_by_id_and_directory_id(
self,
knowledge_id: str,
directory_id: str,
db: Optional[AsyncSession] = None,
) -> list[FileModel]:
"""Get all files in a directory and its subdirectories."""
async with get_async_db_context(db) as db:
directory_ids = [directory_id]
for parent_id in directory_ids:
result = await db.execute(select(KnowledgeDirectory.id).filter_by(parent_id=parent_id))
directory_ids.extend(result.scalars().all())
result = await db.execute(
select(File)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
.filter(KnowledgeFile.knowledge_id == knowledge_id)
.filter(KnowledgeFile.directory_id.in_(directory_ids))
)
return [FileModel.model_validate(file) for file in result.scalars().all()]
async def move_file_to_directory( async def move_file_to_directory(
self, self,
knowledge_id: str, knowledge_id: str,

View file

@ -4,11 +4,11 @@ from __future__ import annotations
import time import time
import uuid import uuid
from typing import Literal from typing import Optional
from open_webui.internal.db import Base, get_async_db_context from open_webui.internal.db import Base, get_async_db_context
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from sqlalchemy import JSON, BigInteger, Column, Index, String, Text, delete, select from sqlalchemy import BigInteger, Column, String, Text, delete, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -16,14 +16,10 @@ class Memory(Base): # user memory store
"""Stores user-created memory entries linked to a vector collection.""" """Stores user-created memory entries linked to a vector collection."""
__tablename__ = 'memory' __tablename__ = 'memory'
__table_args__ = (Index('ix_memory_id_user_id', 'id', 'user_id'),)
id = Column(String, primary_key=True, unique=True) id = Column(String, primary_key=True, unique=True)
user_id = Column(String, index=True) user_id = Column(String, index=True)
type = Column(String, default='context', server_default='context', index=True)
path = Column(Text, nullable=True)
content = Column(Text) # free-form text learned from conversation content = Column(Text) # free-form text learned from conversation
meta = Column(JSON, nullable=True)
updated_at = Column(BigInteger) # epoch seconds updated_at = Column(BigInteger) # epoch seconds
created_at = Column(BigInteger) # epoch seconds created_at = Column(BigInteger) # epoch seconds
@ -33,27 +29,17 @@ class MemoryModel(BaseModel):
id: str id: str
user_id: str user_id: str
type: Literal['user', 'context'] = 'context'
path: str | None = None
content: str content: str
meta: dict | None = None
updated_at: int # timestamp in epoch updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True) # allows ORM mapping model_config = ConfigDict(from_attributes=True) # allows ORM mapping
class MemoriesTable: class MemoriesTable:
@staticmethod
def normalize_memory_type(memory_type: str | None = None) -> str:
return 'user' if memory_type == 'user' else 'context'
async def insert_new_memory( async def insert_new_memory(
self, self,
user_id: str, user_id: str,
content: str, content: str,
memory_type: str | None = None,
path: str | None = None,
meta: dict | None = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> MemoryModel | None: ) -> MemoryModel | None:
"""Persist a new memory entry and return the created model.""" """Persist a new memory entry and return the created model."""
@ -62,26 +48,20 @@ class MemoriesTable:
record = Memory( record = Memory(
id=str(uuid.uuid4()), id=str(uuid.uuid4()),
user_id=user_id, user_id=user_id,
type=self.normalize_memory_type(memory_type),
path=path,
content=content, content=content,
meta=meta,
created_at=now, created_at=now,
updated_at=now, updated_at=now,
) )
db.add(record) db.add(record)
await db.commit() await db.commit()
await db.refresh(record)
return MemoryModel.model_validate(record) if record else None return MemoryModel.model_validate(record) if record else None
async def update_memory_by_id_and_user_id( async def update_memory_by_id_and_user_id(
self, self,
id: str, id: str,
user_id: str, user_id: str,
content: str | None, content: str,
memory_type: str | None = None,
path: str | None = None,
update_path: bool = False,
meta: dict | None = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> MemoryModel | None: ) -> MemoryModel | None:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
@ -90,17 +70,11 @@ class MemoriesTable:
if not memory or memory.user_id != user_id: if not memory or memory.user_id != user_id:
return None return None
if content is not None: memory.content = content
memory.content = content
if memory_type is not None:
memory.type = self.normalize_memory_type(memory_type)
if update_path:
memory.path = path
if meta is not None:
memory.meta = {**(memory.meta or {}), **meta}
memory.updated_at = int(time.time()) memory.updated_at = int(time.time())
await db.commit() await db.commit()
await db.refresh(memory)
return MemoryModel.model_validate(memory) return MemoryModel.model_validate(memory)
except Exception: except Exception:
return None return None
@ -165,104 +139,5 @@ class MemoriesTable:
except Exception: except Exception:
return False return False
async def apply_memory_operations(
self,
user_id: str,
operations: list[dict],
db: AsyncSession | None = None,
) -> list[dict]:
now = int(time.time())
results: list[dict] = []
async with get_async_db_context(db) as db:
for operation in operations:
action = operation.get('action')
if action == 'add':
content = operation.get('content', '').strip()
memory_type = self.normalize_memory_type(operation.get('type'))
path = operation.get('path')
result = await db.execute(
select(Memory).filter_by(user_id=user_id, content=content, type=memory_type, path=path)
)
existing = result.scalars().first()
if existing:
results.append(
{
'action': action,
'status': 'skipped',
'memory': MemoryModel.model_validate(existing),
'reason': 'duplicate',
}
)
continue
memory = Memory(
id=str(uuid.uuid4()),
user_id=user_id,
type=memory_type,
path=path,
content=content,
meta=operation.get('meta'),
created_at=now,
updated_at=now,
)
db.add(memory)
await db.flush()
results.append(
{'action': action, 'status': 'created', 'memory': MemoryModel.model_validate(memory)}
)
elif action == 'replace':
memory_id = operation.get('id')
content = operation.get('content', '').strip()
memory = await db.get(Memory, memory_id)
if not memory or memory.user_id != user_id:
raise ValueError(f'Memory not found: {memory_id}')
memory.content = content
if operation.get('type') is not None:
memory.type = self.normalize_memory_type(operation.get('type'))
if 'path' in operation:
memory.path = operation.get('path')
if operation.get('meta') is not None:
memory.meta = {**(memory.meta or {}), **operation.get('meta')}
memory.updated_at = now
await db.flush()
results.append(
{'action': action, 'status': 'updated', 'memory': MemoryModel.model_validate(memory)}
)
elif action == 'move':
memory_id = operation.get('id')
memory = await db.get(Memory, memory_id)
if not memory or memory.user_id != user_id:
raise ValueError(f'Memory not found: {memory_id}')
memory.path = operation.get('path')
if operation.get('meta') is not None:
memory.meta = {**(memory.meta or {}), **operation.get('meta')}
memory.updated_at = now
await db.flush()
results.append(
{'action': action, 'status': 'updated', 'memory': MemoryModel.model_validate(memory)}
)
elif action == 'remove':
memory_id = operation.get('id')
memory = await db.get(Memory, memory_id)
if not memory or memory.user_id != user_id:
raise ValueError(f'Memory not found: {memory_id}')
await db.delete(memory)
results.append({'action': action, 'status': 'deleted', 'id': memory_id})
else:
raise ValueError(f'Unsupported memory operation: {action}')
await db.commit()
return results
Memories = MemoriesTable() # user memory registry Memories = MemoriesTable() # user memory registry

View file

@ -1,3 +1,4 @@
import json
import time import time
import uuid import uuid
from typing import Optional from typing import Optional
@ -327,8 +328,7 @@ class MessageTable:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
message = await db.get(Message, parent_id) message = await db.get(Message, parent_id)
# Thread parent must belong to the requested channel; never disclose a foreign-channel message. if not message:
if not message or message.channel_id != channel_id:
return [] return []
result = await db.execute( result = await db.execute(
@ -500,71 +500,6 @@ class MessageTable:
return [Reactions(**reaction) for reaction in reactions.values()] return [Reactions(**reaction) for reaction in reactions.values()]
async def get_reactions_by_message_ids(
self, ids: list[str], db: Optional[AsyncSession] = None
) -> dict[str, list[Reactions]]:
"""Batch-fetch reactions for multiple messages in a single query.
Returns a dict mapping each message_id to its list of Reactions.
Messages with no reactions map to an empty list.
"""
if not ids:
return {}
async with get_async_db_context(db) as db:
result = await db.execute(
select(MessageReaction, User)
.join(User, MessageReaction.user_id == User.id)
.filter(MessageReaction.message_id.in_(ids))
)
rows = result.all()
# Group by (message_id, reaction_name)
grouped: dict[str, dict[str, dict]] = {mid: {} for mid in ids}
for reaction, user in rows:
mid = reaction.message_id
if mid not in grouped:
grouped[mid] = {}
if reaction.name not in grouped[mid]:
grouped[mid][reaction.name] = {
'name': reaction.name,
'users': [],
'count': 0,
}
grouped[mid][reaction.name]['users'].append(
{
'id': user.id,
'name': user.name,
}
)
grouped[mid][reaction.name]['count'] += 1
return {mid: [Reactions(**r) for r in reactions.values()] for mid, reactions in grouped.items()}
async def get_thread_reply_counts_by_message_ids(
self, ids: list[str], db: Optional[AsyncSession] = None
) -> dict[str, tuple[int, int | None]]:
"""Batch-fetch reply counts and latest reply timestamps for multiple parent messages.
Returns a dict mapping each parent message_id to a
(reply_count, latest_reply_created_at) tuple.
Messages with no replies are omitted from the result.
"""
if not ids:
return {}
async with get_async_db_context(db) as db:
result = await db.execute(
select(
Message.parent_id,
func.count(Message.id),
func.max(Message.created_at),
)
.filter(Message.parent_id.in_(ids))
.group_by(Message.parent_id)
)
return {row[0]: (row[1], row[2]) for row in result.all()}
async def remove_reaction_by_id_and_user_id_and_name( async def remove_reaction_by_id_and_user_id_and_name(
self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
) -> bool: ) -> bool:

View file

@ -1,65 +1,25 @@
from __future__ import annotations from __future__ import annotations
import json
import logging import logging
import time import time
from copy import deepcopy from typing import Optional
from typing import Any
from open_webui.internal.db import Base, JSONField, get_async_db_context 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.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups from open_webui.models.groups import Groups
from open_webui.models.users import User, UserModel, UserResponse, Users from open_webui.models.users import User, UserModel, UserResponse, Users
from open_webui.utils.misc import json_text_variants from open_webui.utils.validate import validate_profile_image_url
from open_webui.utils.validate import validate_image_url from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic import BaseModel, ConfigDict, Field, ValidationInfo, field_validator, model_validator
from sqlalchemy import BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, update 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 from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
# Track invalid profile_image_url values we've already warned about so we
def normalize_model_tags(tags: Any) -> list[dict[str, str]]: # don't flood the logs on every DB read (the validator fires per-row).
if not isinstance(tags, list): _warned_profile_urls: set[str] = set()
return []
normalized = []
for tag in tags:
name = tag.get('name') if isinstance(tag, dict) else tag
if isinstance(name, str) and name.strip():
normalized.append({'name': name.strip()})
return normalized
def strip_extracted_content_from_model_knowledge(knowledge: Any) -> Any:
"""Drop duplicated extracted text from ModelMeta.knowledge."""
if not isinstance(knowledge, list):
return knowledge
sanitized = []
for item in knowledge:
if not isinstance(item, dict):
sanitized.append(item)
continue
next_item = item
data = item.get('data')
if isinstance(data, dict) and 'content' in data:
next_item = deepcopy(item)
next_item.get('data', {}).pop('content', None)
file = next_item.get('file')
file_data = file.get('data') if isinstance(file, dict) else None
if isinstance(file_data, dict) and 'content' in file_data:
if next_item is item:
next_item = deepcopy(item)
file = next_item.get('file')
file_data = file.get('data') if isinstance(file, dict) else None
file_data.pop('content', None)
sanitized.append(next_item)
return sanitized
# --- Models DB Schema --- # --- Models DB Schema ---
@ -75,36 +35,40 @@ class ModelMeta(BaseModel):
"""Metadata for a workspace model entry (profile, description, tags, capabilities).""" """Metadata for a workspace model entry (profile, description, tags, capabilities)."""
profile_image_url: str | None = None profile_image_url: str | None = None
background_image_url: str | None = None
description: str | None = Field(default=None, description='User-facing description of the model.') description: str | None = Field(default=None, description='User-facing description of the model.')
i18n: dict[str, Any] | None = None
capabilities: dict | None = None capabilities: dict | None = None
knowledge: list[Any] | None = None
model_config = ConfigDict(extra='allow') model_config = ConfigDict(extra='allow')
@field_validator('profile_image_url', 'background_image_url', mode='before') @field_validator('profile_image_url', mode='before')
@classmethod @classmethod
def check_image_url(cls, v: str | None, info: ValidationInfo) -> str | None: def check_profile_image_url(cls, v: str | None) -> str | None:
if v is None: if v is None:
return v return v
try: try:
return validate_image_url(v, file_only=info.field_name == 'background_image_url') return validate_profile_image_url(v)
except ValueError: except ValueError:
if info.field_name == 'background_image_url': if v not in _warned_profile_urls:
raise _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 return None
@field_validator('knowledge', mode='before')
@classmethod
def strip_knowledge_content(cls, v):
return strip_extracted_content_from_model_knowledge(v)
@model_validator(mode='before') @model_validator(mode='before')
@classmethod @classmethod
def normalize_tags(cls, data): def normalize_tags(cls, data):
if isinstance(data, dict) and 'tags' in data: if isinstance(data, dict) and 'tags' in data:
data['tags'] = normalize_model_tags(data['tags']) raw_tags = data['tags']
if isinstance(raw_tags, list):
normalized = []
for tag in raw_tags:
if isinstance(tag, str):
normalized.append({'name': tag})
elif isinstance(tag, dict) and 'name' in tag:
normalized.append(tag)
data['tags'] = normalized
return data return data
@ -169,12 +133,12 @@ class ModelAccessListResponse(BaseModel):
class ModelForm(BaseModel): class ModelForm(BaseModel):
model_config = ConfigDict(extra='ignore') model_config = ConfigDict(extra='ignore')
id: str = Field(pattern=r'^\S+$') id: str
base_model_id: str | None = None base_model_id: str | None = None
name: str name: str
meta: ModelMeta meta: ModelMeta
params: ModelParams params: ModelParams
access_grants: list[dict] | None = None access_grants: list[dict | None] = None
is_active: bool = True is_active: bool = True
@ -185,22 +149,14 @@ class ModelsTable:
async def _to_model_model( async def _to_model_model(
self, self,
model: Model, model: Model,
access_grants: list[AccessGrantModel] | None = None, access_grants: list[AccessGrantModel | None] = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> ModelModel: ) -> ModelModel:
if isinstance(model.meta, dict): model_data = ModelModel.model_validate(model).model_dump(exclude={'access_grants'})
knowledge = model.meta.get('knowledge') model_data['access_grants'] = (
stripped_knowledge = strip_extracted_content_from_model_knowledge(knowledge) access_grants if access_grants is not None else await self._get_access_grants(model_data['id'], db=db)
if stripped_knowledge != knowledge:
model.meta = {**model.meta, 'knowledge': stripped_knowledge}
if db is not None:
await db.commit()
model_model = ModelModel.model_validate(model)
model_model.access_grants = (
access_grants if access_grants is not None else await self._get_access_grants(model_model.id, db=db)
) )
return model_model return ModelModel.model_validate(model_data)
async def insert_new_model( async def insert_new_model(
self, form_data: ModelForm, user_id: str, db: AsyncSession | None = None self, form_data: ModelForm, user_id: str, db: AsyncSession | None = None
@ -217,6 +173,7 @@ class ModelsTable:
) )
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db) await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db)
if result: if result:
@ -241,24 +198,9 @@ class ModelsTable:
log.error('Skipping model %r during get_all_models due to error: %s', model.id, exc) log.error('Skipping model %r during get_all_models due to error: %s', model.id, exc)
return models return models
async def get_models( async def get_models(self, db: AsyncSession | None = None) -> list[ModelUserResponse]:
self, writable_by_user_id: str | None = None, db: AsyncSession | None = None, ids: list[str] | None = None
) -> list[ModelUserResponse]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
stmt = select(Model).filter(Model.base_model_id != None) result = await db.execute(select(Model).filter(Model.base_model_id != None))
if ids is not None:
stmt = stmt.filter(Model.id.in_(ids))
if writable_by_user_id:
user_group_ids = {
group.id for group in await Groups.get_groups_by_member_id(writable_by_user_id, db=db)
}
stmt = self._has_permission(
db, stmt, {'user_id': writable_by_user_id, 'group_ids': user_group_ids}, permission='write'
)
result = await db.execute(stmt)
all_models = result.scalars().all() all_models = result.scalars().all()
user_ids = list(set(model.user_id for model in all_models)) user_ids = list(set(model.user_id for model in all_models))
@ -287,46 +229,10 @@ class ModelsTable:
) )
return models return models
async def get_model_owner_ids_by_file_id( async def get_base_models(self, db: AsyncSession | None = None) -> list[ModelModel]:
self, file_id: str, db: AsyncSession | None = None, include_background: bool = False
) -> dict[str, str]:
"""Return model IDs mapped to owner IDs for models referencing the file."""
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
# File ids are server-generated uuids, so the text match can only over-match. result = await db.execute(select(Model).filter(Model.base_model_id == None))
result = await db.execute(
select(Model.id, Model.user_id, Model.meta).filter(
Model.base_model_id.is_not(None), cast(Model.meta, String).like(f'%{file_id}%')
)
)
return {
model_id: user_id
for model_id, user_id, meta in result.all()
if any(
isinstance(item, dict) and item.get('type') == 'file' and item.get('id') == file_id
for item in meta.get('knowledge') or []
)
or (include_background and meta.get('background_image_url') == f'/api/v1/files/{file_id}/content')
}
@staticmethod
def _meta_has_tag(meta: dict | None, tag: str) -> bool:
if not meta:
return False
for raw_tag in meta.get('tags', []):
name = raw_tag.get('name') if isinstance(raw_tag, dict) else str(raw_tag)
if name == tag:
return True
return False
async def get_base_models(self, tag: str | None = None, db: AsyncSession | None = None) -> list[ModelModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Model).filter(Model.base_model_id.is_(None)))
all_models = result.scalars().all() all_models = result.scalars().all()
if tag:
all_models = [model for model in all_models if self._meta_has_tag(model.meta, tag)]
model_ids = [model.id for model in all_models] model_ids = [model.id for model in all_models]
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
return [ return [
@ -334,6 +240,28 @@ class ModelsTable:
for model in all_models for model in all_models
] ]
async def get_models_by_user_id(
self, user_id: str, permission: str = 'write', db: AsyncSession | None = None
) -> list[ModelUserResponse]:
models = await self.get_models(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
result = []
for model in models:
if model.user_id == user_id:
result.append(model)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='model',
resource_id=model.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
):
result.append(model)
return result
def _has_permission(self, db, query, filter: dict, permission: str = 'read'): def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
return AccessGrants.has_permission_filter( return AccessGrants.has_permission_filter(
db=db, db=db,
@ -385,14 +313,20 @@ class ModelsTable:
tag = filter.get('tag') tag = filter.get('tag')
if tag: if tag:
if db.bind.dialect.name == 'sqlite' and not tag.isascii(): # SQLite stores JSON text via json.dumps(ensure_ascii=True),
# SQLite's LOWER() is ASCII-only, so match non-ASCII tags exact-case. # so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
meta_text = cast(Model.meta, String) # stores literal Unicode. Use the right pattern for each.
variants = json_text_variants(tag) if db.bind.dialect.name == 'sqlite':
if tag.isascii():
meta_text = func.lower(cast(Model.meta, String))
pattern = f'%{json.dumps(tag.lower())}%'
else:
meta_text = cast(Model.meta, String)
pattern = f'%{json.dumps(tag)}%'
else: else:
meta_text = func.lower(cast(Model.meta, String)) meta_text = func.lower(cast(Model.meta, String))
variants = json_text_variants(tag.lower()) pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
stmt = stmt.filter(or_(*(meta_text.like(f'%"{variant}"%') for variant in variants))) stmt = stmt.filter(meta_text.like(pattern))
order_by = filter.get('order_by') order_by = filter.get('order_by')
direction = filter.get('direction') direction = filter.get('direction')
@ -448,13 +382,11 @@ class ModelsTable:
return ModelListResponse(items=models, total=total) return ModelListResponse(items=models, total=total)
async def get_model_meta_by_id( async def get_model_meta_by_id(self, id: str, db: AsyncSession | None = None) -> tuple[dict, int | None]:
self, id: str, db: AsyncSession | None = None """Return (meta, updated_at) for a model, skipping access grant resolution."""
) -> tuple[dict, str, int | None] | None:
"""Return (meta, user_id, updated_at) for a model, skipping access grant resolution."""
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute(select(Model.meta, Model.user_id, Model.updated_at).filter_by(id=id)) result = await db.execute(select(Model.meta, Model.updated_at).filter_by(id=id))
return result.first() return result.first()
except Exception: except Exception:
return None return None
@ -463,14 +395,11 @@ class ModelsTable:
self, self,
user_id: str, user_id: str,
is_admin: bool = False, is_admin: bool = False,
is_base_model: bool = False,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> set[str]: ) -> set[str]:
"""Extract unique tag names from model meta, querying only the meta column.""" """Extract unique tag names from model meta, querying only the meta column."""
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
stmt = select(Model.meta).filter( stmt = select(Model.meta).filter(Model.base_model_id != None)
Model.base_model_id.is_(None) if is_base_model else Model.base_model_id.is_not(None)
)
if not is_admin: if not is_admin:
user_groups = await Groups.get_groups_by_member_id(user_id, db=db) user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
@ -536,6 +465,7 @@ class ModelsTable:
model.is_active = not model.is_active model.is_active = not model.is_active
model.updated_at = int(time.time()) model.updated_at = int(time.time())
await db.commit() await db.commit()
await db.refresh(model)
return await self._to_model_model(model, db=db) return await self._to_model_model(model, db=db)
except Exception: except Exception:
@ -562,12 +492,13 @@ class ModelsTable:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
result = await db.execute(select(Model).filter_by(id=id)) result = await db.execute(select(Model).filter_by(id=id))
model = result.scalars().first() model_obj = result.scalars().first()
if not model: if not model_obj:
return None return None
model.updated_at = int(time.time()) model_obj.updated_at = int(time.time())
await db.commit() await db.commit()
return await self._to_model_model(model, db=db) await db.refresh(model_obj)
return await self._to_model_model(model_obj, db=db)
except Exception as e: except Exception as e:
log.exception(f'Failed to update the model updated_at by id {id}: {e}') log.exception(f'Failed to update the model updated_at by id {id}: {e}')
return None return None
@ -612,16 +543,25 @@ class ModelsTable:
# Update or insert models # Update or insert models
for model in models: for model in models:
model_data = {
**model.model_dump(exclude={'access_grants'}),
'user_id': user_id,
'updated_at': int(time.time()),
}
if model.id in existing_ids: if model.id in existing_ids:
await db.execute(update(Model).filter_by(id=model.id).values(**model_data)) await db.execute(
update(Model)
.filter_by(id=model.id)
.values(
**model.model_dump(exclude={'access_grants'}),
user_id=user_id,
updated_at=int(time.time()),
)
)
else: else:
db.add(Model(**model_data)) new_model = Model(
**{
**model.model_dump(exclude={'access_grants'}),
'user_id': user_id,
'updated_at': int(time.time()),
}
)
db.add(new_model)
await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db) await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db)
# Remove models that are no longer present # Remove models that are no longer present

View file

@ -1,3 +1,4 @@
import json
import time import time
import uuid import uuid
from functools import lru_cache from functools import lru_cache
@ -7,8 +8,7 @@ from open_webui.internal.db import Base, get_async_db_context
from open_webui.models.access_grants import AccessGrantModel, AccessGrants from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups from open_webui.models.groups import Groups
from open_webui.models.users import User, UserModel, UserResponse, Users from open_webui.models.users import User, UserModel, UserResponse, Users
from open_webui.utils.json_codec import JSONCodec from pydantic import BaseModel, ConfigDict, Field
from pydantic import BaseModel, ConfigDict, Field, field_validator
from sqlalchemy import JSON, BigInteger, Boolean, Column, ForeignKey, Text, delete, func, or_, select, update from sqlalchemy import JSON, BigInteger, Boolean, Column, ForeignKey, Text, delete, func, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -31,32 +31,6 @@ class Note(Base):
updated_at = Column(BigInteger) updated_at = Column(BigInteger)
def sanitize_note_data(data: Optional[dict]) -> Optional[dict]:
"""Sanitize malformed note.data so content.md is always markdown text."""
if data is None:
return None
if not isinstance(data, dict):
return {'content': {'md': str(data)}}
content = data.get('content')
if not isinstance(content, dict) or 'md' not in content or isinstance(content.get('md'), str):
return data
md = content.get('md') if content.get('md') is not None else ''
if isinstance(md, (dict, list)):
md = f'```json\n{JSONCodec.dumps(md, indent=2, ensure_ascii=False)}\n```'
else:
md = str(md)
return {
**data,
'content': {
**content,
'md': md,
},
}
class NoteModel(BaseModel): class NoteModel(BaseModel):
model_config = ConfigDict(from_attributes=True) model_config = ConfigDict(from_attributes=True)
@ -73,11 +47,6 @@ class NoteModel(BaseModel):
created_at: int # timestamp in epoch created_at: int # timestamp in epoch
updated_at: int # timestamp in epoch updated_at: int # timestamp in epoch
@field_validator('data', mode='before')
@classmethod
def sanitize_data(cls, data):
return sanitize_note_data(data)
class PinnedNote(Base): class PinnedNote(Base):
__tablename__ = 'pinned_note' __tablename__ = 'pinned_note'
@ -99,11 +68,6 @@ class NoteForm(BaseModel):
meta: Optional[dict] = None meta: Optional[dict] = None
access_grants: Optional[list[dict]] = None access_grants: Optional[list[dict]] = None
@field_validator('data', mode='before')
@classmethod
def sanitize_data(cls, data):
return sanitize_note_data(data)
class NoteUpdateForm(BaseModel): class NoteUpdateForm(BaseModel):
title: Optional[str] = None title: Optional[str] = None
@ -111,11 +75,6 @@ class NoteUpdateForm(BaseModel):
meta: Optional[dict] = None meta: Optional[dict] = None
access_grants: Optional[list[dict]] = None access_grants: Optional[list[dict]] = None
@field_validator('data', mode='before')
@classmethod
def sanitize_data(cls, data):
return sanitize_note_data(data)
class NoteUserResponse(NoteModel): class NoteUserResponse(NoteModel):
user: Optional[UserResponse] = None user: Optional[UserResponse] = None
@ -147,12 +106,11 @@ class NoteTable:
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
) -> NoteModel: ) -> NoteModel:
# We exclude access_grants to inject them # We exclude access_grants to inject them
note_model = NoteModel.model_validate(note) note_data = NoteModel.model_validate(note).model_dump(exclude={'access_grants'})
note_model.data = note_model.data or {} note_data['access_grants'] = (
note_model.access_grants = ( access_grants if access_grants is not None else await self._get_access_grants(note_data['id'], db=db)
access_grants if access_grants is not None else await self._get_access_grants(note_model.id, db=db)
) )
return note_model return NoteModel.model_validate(note_data)
def _has_permission(self, db, query, filter: dict, permission: str = 'read'): def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
return AccessGrants.has_permission_filter( return AccessGrants.has_permission_filter(
@ -346,17 +304,13 @@ class NoteTable:
return None return None
form_data = form_data.model_dump(exclude_unset=True) form_data = form_data.model_dump(exclude_unset=True)
note.data = sanitize_note_data(note.data) or {}
if 'title' in form_data: if 'title' in form_data:
note.title = form_data['title'] note.title = form_data['title']
if 'data' in form_data: if 'data' in form_data:
note.data = {**(note.data or {}), **(form_data['data'] or {})} note.data = {**note.data, **form_data['data']}
if 'meta' in form_data: if 'meta' in form_data:
note.meta = {**(note.meta or {}), **(form_data['meta'] or {})} note.meta = {**note.meta, **form_data['meta']}
if not db.is_modified(note) and 'access_grants' not in form_data:
return await self._to_note_model(note, db=db)
if 'access_grants' in form_data: if 'access_grants' in form_data:
await AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db) await AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db)

View file

@ -1,5 +1,6 @@
import base64 import base64
import hashlib import hashlib
import json
import logging import logging
import time import time
import uuid import uuid
@ -8,7 +9,6 @@ from typing import List, Optional
from cryptography.fernet import Fernet from cryptography.fernet import Fernet
from open_webui.env import OAUTH_SESSION_TOKEN_ENCRYPTION_KEY from open_webui.env import OAUTH_SESSION_TOKEN_ENCRYPTION_KEY
from open_webui.internal.db import Base, get_async_db_context from open_webui.internal.db import Base, get_async_db_context
from open_webui.utils.json_codec import JSONCodec
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Column, Index, String, Text, delete, select, update from sqlalchemy import BigInteger, Column, Index, String, Text, delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -85,7 +85,7 @@ class OAuthSessionTable:
def _encrypt_token(self, token) -> str: def _encrypt_token(self, token) -> str:
"""Encrypt OAuth tokens for storage""" """Encrypt OAuth tokens for storage"""
try: try:
token_json = JSONCodec.dumps(token) token_json = json.dumps(token)
encrypted = self.fernet.encrypt(token_json.encode()).decode() encrypted = self.fernet.encrypt(token_json.encode()).decode()
return encrypted return encrypted
except Exception as e: except Exception as e:
@ -96,7 +96,7 @@ class OAuthSessionTable:
"""Decrypt OAuth tokens from storage""" """Decrypt OAuth tokens from storage"""
try: try:
decrypted = self.fernet.decrypt(token.encode()).decode() decrypted = self.fernet.decrypt(token.encode()).decode()
return JSONCodec.loads(decrypted) return json.loads(decrypted)
except Exception as e: except Exception as e:
log.error(f'Error decrypting tokens: {type(e).__name__}: {e}') log.error(f'Error decrypting tokens: {type(e).__name__}: {e}')
raise raise
@ -128,6 +128,7 @@ class OAuthSessionTable:
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
if result: if result:
# Make a copy of the model data before closing session # Make a copy of the model data before closing session

View file

@ -1,6 +1,7 @@
"""Prompt history model for version tracking.""" """Prompt history model for version tracking."""
import difflib import difflib
import json
import time import time
import uuid import uuid
from typing import Optional from typing import Optional
@ -69,6 +70,7 @@ class PromptHistoryTable:
) )
db.add(history) db.add(history)
await db.commit() await db.commit()
await db.refresh(history)
return PromptHistoryModel.model_validate(history) return PromptHistoryModel.model_validate(history)
async def get_history_by_prompt_id( async def get_history_by_prompt_id(

View file

@ -2,6 +2,7 @@
from __future__ import annotations from __future__ import annotations
import json
import logging import logging
import time import time
import uuid import uuid
@ -14,7 +15,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups from open_webui.models.groups import Groups
from open_webui.models.prompt_history import PromptHistories from open_webui.models.prompt_history import PromptHistories
from open_webui.models.users import User, UserModel, UserResponse, Users from open_webui.models.users import User, UserModel, UserResponse, Users
from open_webui.utils.misc import json_text_variants
from pydantic import BaseModel, ConfigDict, Field from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, text, update from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -47,7 +47,7 @@ class PromptModel(BaseModel):
content: str content: str
data: dict | None = None data: dict | None = None
meta: dict | None = None meta: dict | None = None
tags: list[str] | None = None tags: list[str | None] = None
is_active: bool | None = True is_active: bool | None = True
version_id: str | None = None version_id: str | None = None
created_at: int | None = None created_at: int | None = None
@ -86,8 +86,8 @@ class PromptForm(BaseModel):
content: str content: str
data: dict | None = None data: dict | None = None
meta: dict | None = None meta: dict | None = None
tags: list[str] | None = None tags: list[str | None] = None
access_grants: list[dict] | None = None access_grants: list[dict | None] = None
version_id: str | None = None # Active version version_id: str | None = None # Active version
commit_message: str | None = None # For history tracking commit_message: str | None = None # For history tracking
is_production: bool | None = True # Whether to set new version as production is_production: bool | None = True # Whether to set new version as production
@ -100,14 +100,14 @@ class PromptsTable:
async def _to_prompt_model( async def _to_prompt_model(
self, self,
prompt: Prompt, prompt: Prompt,
access_grants: list[AccessGrantModel] | None = None, access_grants: list[AccessGrantModel | None] = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> PromptModel: ) -> PromptModel:
prompt_model = PromptModel.model_validate(prompt) prompt_data = PromptModel.model_validate(prompt).model_dump(exclude={'access_grants'})
prompt_model.access_grants = ( prompt_data['access_grants'] = (
access_grants if access_grants is not None else await self._get_access_grants(prompt_model.id, db=db) access_grants if access_grants is not None else await self._get_access_grants(prompt_data['id'], db=db)
) )
return prompt_model return PromptModel.model_validate(prompt_data)
async def insert_new_prompt( async def insert_new_prompt(
self, user_id: str, form_data: PromptForm, db: AsyncSession | None = None self, user_id: str, form_data: PromptForm, db: AsyncSession | None = None
@ -132,6 +132,7 @@ class PromptsTable:
) )
session.add(record) session.add(record)
await session.commit() await session.commit()
await session.refresh(record) # populate generated defaults
await AccessGrants.set_access_grants( await AccessGrants.set_access_grants(
'prompt', 'prompt',
@ -168,6 +169,7 @@ class PromptsTable:
if history_entry: if history_entry:
record.version_id = history_entry.id record.version_id = history_entry.id
await session.commit() await session.commit()
await session.refresh(record) # re-read version_id
return await self._to_prompt_model(record, db=session) return await self._to_prompt_model(record, db=session)
except Exception as e: except Exception as e:
@ -334,19 +336,17 @@ class PromptsTable:
tag_lower = tag.lower() tag_lower = tag.lower()
if dialect_name == 'sqlite': if dialect_name == 'sqlite':
tag_lower = tag.replace('\\', '\\\\').replace('%', '\\%').replace('_', '\\_')
tag_clause = text( tag_clause = text(
"EXISTS (SELECT 1 FROM json_each(prompt.tags) t WHERE t.value LIKE :tag_val ESCAPE '\\')" 'EXISTS (SELECT 1 FROM json_each(prompt.tags) t WHERE LOWER(t.value) = :tag_val)'
) )
elif dialect_name == 'postgresql': elif dialect_name == 'postgresql':
tag_clause = text( tag_clause = text(
'EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)' 'EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)'
) )
else: else:
# Fallback for dialects with no JSON array function: LIKE on the text. # Fallback: LIKE on serialised JSON text (ASCII-safe only)
tags_text = func.lower(cast(Prompt.tags, String)) tag_clause = func.lower(cast(Prompt.tags, String)).like(
tag_clause = or_( f'%{json.dumps(tag_lower, ensure_ascii=False)}%'
*(tags_text.like(f'%"{variant}"%') for variant in json_text_variants(tag_lower))
) )
tag_lower = None tag_lower = None
@ -506,16 +506,14 @@ class PromptsTable:
) )
# Update prompt fields # Update prompt fields
prompt.name = form_data.name
prompt.command = form_data.command prompt.command = form_data.command
prompt.content = form_data.content
prompt.data = form_data.data or prompt.data
prompt.meta = form_data.meta or prompt.meta
if form_data.is_production: if form_data.tags is not None:
prompt.name = form_data.name prompt.tags = form_data.tags
prompt.content = form_data.content
prompt.data = form_data.data or prompt.data
prompt.meta = form_data.meta or prompt.meta
if form_data.tags is not None:
prompt.tags = form_data.tags
if form_data.access_grants is not None: if form_data.access_grants is not None:
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=session) await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=session)
@ -533,7 +531,7 @@ class PromptsTable:
'command': prompt.command, 'command': prompt.command,
'data': form_data.data or {}, 'data': form_data.data or {},
'meta': form_data.meta or {}, 'meta': form_data.meta or {},
'tags': form_data.tags if form_data.tags is not None else (prompt.tags or []), 'tags': prompt.tags or [],
'access_grants': [grant.model_dump() for grant in current_access_grants], 'access_grants': [grant.model_dump() for grant in current_access_grants],
} }
@ -560,7 +558,7 @@ class PromptsTable:
prompt_id: str, prompt_id: str,
name: str, name: str,
command: str, command: str,
tags: list[str] | None = None, tags: list[str | None] = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> PromptModel | None: ) -> PromptModel | None:
"""Update only name, command, and tags (no history created).""" """Update only name, command, and tags (no history created)."""
@ -639,6 +637,7 @@ class PromptsTable:
prompt.is_active = not prompt.is_active prompt.is_active = not prompt.is_active
prompt.updated_at = int(time.time()) prompt.updated_at = int(time.time())
await session.commit() await session.commit()
await session.refresh(prompt)
return await self._to_prompt_model(prompt, db=session) return await self._to_prompt_model(prompt, db=session)
return None return None
except Exception: except Exception:

View file

@ -201,15 +201,5 @@ class SharedChatsTable:
except Exception: except Exception:
return False return False
async def delete_all_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
"""Delete all shared chats created by a user."""
try:
async with get_async_db_context(db) as db:
await db.execute(delete(SharedChat).filter_by(user_id=user_id))
await db.commit()
return True
except Exception:
return False
SharedChats = SharedChatsTable() SharedChats = SharedChatsTable()

View file

@ -33,7 +33,6 @@ class Skill(Base):
class SkillMeta(BaseModel): class SkillMeta(BaseModel):
i18n: dict[str, dict[str, str]] | None = None
tags: Optional[list[str]] = [] tags: Optional[list[str]] = []
@ -114,11 +113,11 @@ class SkillsTable:
access_grants: Optional[list[AccessGrantModel]] = None, access_grants: Optional[list[AccessGrantModel]] = None,
db: Optional[AsyncSession] = None, db: Optional[AsyncSession] = None,
) -> SkillModel: ) -> SkillModel:
skill_model = SkillModel.model_validate(skill) skill_data = SkillModel.model_validate(skill).model_dump(exclude={'access_grants'})
skill_model.access_grants = ( skill_data['access_grants'] = (
access_grants if access_grants is not None else await self._get_access_grants(skill_model.id, db=db) access_grants if access_grants is not None else await self._get_access_grants(skill_data['id'], db=db)
) )
return skill_model return SkillModel.model_validate(skill_data)
async def insert_new_skill( async def insert_new_skill(
self, self,
@ -138,6 +137,7 @@ class SkillsTable:
) )
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('skill', result.id, form_data.access_grants, db=db) await AccessGrants.set_access_grants('skill', result.id, form_data.access_grants, db=db)
if result: if result:
return await self._to_skill_model(result, db=db) return await self._to_skill_model(result, db=db)
@ -164,30 +164,9 @@ class SkillsTable:
except Exception: except Exception:
return None return None
async def get_skills( async def get_skills(self, db: Optional[AsyncSession] = None) -> list[SkillUserModel]:
self,
user_id: str | None = None,
ids: list[str] | None = None,
db: AsyncSession | None = None,
) -> list[SkillUserModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
stmt = select(Skill).order_by(Skill.updated_at.desc()) result = await db.execute(select(Skill).order_by(Skill.updated_at.desc()))
if ids is not None:
stmt = stmt.filter(Skill.id.in_(ids))
if user_id is not None:
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)}
stmt = AccessGrants.has_permission_filter(
db=db,
query=stmt,
DocumentModel=Skill,
filter={'user_id': user_id, 'group_ids': user_group_ids},
resource_type='skill',
permission='read',
)
result = await db.execute(stmt)
all_skills = result.scalars().all() all_skills = result.scalars().all()
user_ids = list(set(skill.user_id for skill in all_skills)) user_ids = list(set(skill.user_id for skill in all_skills))
@ -216,6 +195,28 @@ class SkillsTable:
) )
return skills return skills
async def get_skills_by_user_id(
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
) -> list[SkillUserModel]:
skills = await self.get_skills(db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups}
result = []
for skill in skills:
if skill.user_id == user_id:
result.append(skill)
elif await AccessGrants.has_access(
user_id=user_id,
resource_type='skill',
resource_id=skill.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
):
result.append(skill)
return result
async def search_skills( async def search_skills(
self, self,
user_id: str, user_id: str,
@ -258,26 +259,7 @@ class SkillsTable:
permission='read', permission='read',
) )
order_by = filter.get('order_by') stmt = stmt.order_by(Skill.updated_at.desc())
direction = filter.get('direction')
if order_by == 'name':
if direction == 'asc':
stmt = stmt.order_by(Skill.name.asc())
else:
stmt = stmt.order_by(Skill.name.desc())
elif order_by == 'created_at':
if direction == 'asc':
stmt = stmt.order_by(Skill.created_at.asc())
else:
stmt = stmt.order_by(Skill.created_at.desc())
elif order_by == 'updated_at':
if direction == 'asc':
stmt = stmt.order_by(Skill.updated_at.asc())
else:
stmt = stmt.order_by(Skill.updated_at.desc())
else:
stmt = stmt.order_by(Skill.updated_at.desc())
# Count BEFORE pagination # Count BEFORE pagination
count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
@ -325,8 +307,8 @@ class SkillsTable:
if access_grants is not None: if access_grants is not None:
await AccessGrants.set_access_grants('skill', id, access_grants, db=db) await AccessGrants.set_access_grants('skill', id, access_grants, db=db)
# populate_existing: the Core update above bypasses any identity-map copy skill = await db.get(Skill, id)
skill = await db.get(Skill, id, populate_existing=True) await db.refresh(skill)
return await self._to_skill_model(skill, db=db) return await self._to_skill_model(skill, db=db)
except Exception: except Exception:
return None return None
@ -342,6 +324,7 @@ class SkillsTable:
skill.is_active = not skill.is_active skill.is_active = not skill.is_active
skill.updated_at = int(time.time()) skill.updated_at = int(time.time())
await db.commit() await db.commit()
await db.refresh(skill)
return await self._to_skill_model(skill, db=db) return await self._to_skill_model(skill, db=db)
except Exception: except Exception:

View file

@ -63,6 +63,7 @@ class TagTable:
record = Tag(id=tag_id, user_id=user_id, name=name) record = Tag(id=tag_id, user_id=user_id, name=name)
db.add(record) db.add(record)
await db.commit() await db.commit()
await db.refresh(record)
return TagModel.model_validate(record) if record else None return TagModel.model_validate(record) if record else None
except Exception as e: except Exception as e:
log.exception('Error inserting tag %r: %s', name, e) log.exception('Error inserting tag %r: %s', name, e)
@ -97,7 +98,7 @@ class TagTable:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
id = name.replace(' ', '_').lower() id = name.replace(' ', '_').lower()
result = await db.execute(delete(Tag).filter_by(id=id, user_id=user_id)) result = await db.execute(delete(Tag).filter_by(id=id, user_id=user_id))
log.debug('res: %s', result.rowcount) log.debug(f'res: {result.rowcount}')
await db.commit() await db.commit()
return True return True
except Exception as e: except Exception as e:

View file

@ -10,7 +10,6 @@ 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.access_grants import AccessGrantModel, AccessGrants
from open_webui.models.groups import Groups from open_webui.models.groups import Groups
from open_webui.models.users import UserResponse, Users from open_webui.models.users import UserResponse, Users
from open_webui.utils.valves import decrypt_valves, encrypt_valves
from pydantic import BaseModel, ConfigDict, Field from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import BigInteger, Column, String, Text, delete, select, update from sqlalchemy import BigInteger, Column, String, Text, delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -34,18 +33,15 @@ class Tool(Base): # database table definition
class ToolMeta(BaseModel): class ToolMeta(BaseModel):
i18n: dict[str, dict[str, str]] | None = None
description: str | None = None description: str | None = None
manifest: dict | None = {} manifest: dict | None = {}
has_user_valves: bool = False
class ToolModel(BaseModel): class ToolModel(BaseModel):
id: str id: str
user_id: str | None = None # may be null for legacy/malformed records user_id: str
name: str name: str
# None when listed with defer_content=True (source skipped for listings) content: str
content: str | None = None
specs: list[dict] specs: list[dict]
meta: ToolMeta meta: ToolMeta
access_grants: list[AccessGrantModel] = Field(default_factory=list) access_grants: list[AccessGrantModel] = Field(default_factory=list)
@ -67,7 +63,7 @@ class ToolUserModel(ToolModel):
class ToolResponse(BaseModel): class ToolResponse(BaseModel):
id: str id: str
user_id: str | None = None # may be null for legacy/malformed records user_id: str
name: str name: str
meta: ToolMeta meta: ToolMeta
access_grants: list[AccessGrantModel] = Field(default_factory=list) access_grants: list[AccessGrantModel] = Field(default_factory=list)
@ -90,7 +86,7 @@ class ToolForm(BaseModel):
name: str name: str
content: str content: str
meta: ToolMeta meta: ToolMeta
access_grants: list[dict] | None = None access_grants: list[dict | None] = None
class ToolValves(BaseModel): class ToolValves(BaseModel):
@ -104,14 +100,14 @@ class ToolsTable:
async def _to_tool_model( async def _to_tool_model(
self, self,
tool: Tool, tool: Tool,
access_grants: list[AccessGrantModel] | None = None, access_grants: list[AccessGrantModel | None] = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> ToolModel: ) -> ToolModel:
tool_model = ToolModel.model_validate(tool) tool_data = ToolModel.model_validate(tool).model_dump(exclude={'access_grants'})
tool_model.access_grants = ( tool_data['access_grants'] = (
access_grants if access_grants is not None else await self._get_access_grants(tool_model.id, db=db) access_grants if access_grants is not None else await self._get_access_grants(tool_data['id'], db=db)
) )
return tool_model return ToolModel.model_validate(tool_data)
async def insert_new_tool( async def insert_new_tool(
self, self,
@ -133,6 +129,7 @@ class ToolsTable:
) )
db.add(result) db.add(result)
await db.commit() await db.commit()
await db.refresh(result)
await AccessGrants.set_access_grants('tool', result.id, form_data.access_grants, db=db) await AccessGrants.set_access_grants('tool', result.id, form_data.access_grants, db=db)
if result: if result:
return await self._to_tool_model(result, db=db) return await self._to_tool_model(result, db=db)
@ -170,37 +167,13 @@ class ToolsTable:
for tool in tools for tool in tools
} }
async def get_tools( async def get_tools(self, defer_content: bool = False, db: AsyncSession | None = None) -> list[ToolUserModel]:
self,
defer_content: bool = False,
db: AsyncSession | None = None,
user_id: str | None = None,
user_group_ids: set[str] | None = None,
permission: str = 'read',
) -> list[ToolUserModel]:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
# Skip Tool.content (plugin source, potentially large) via a stmt = select(Tool).order_by(Tool.updated_at.desc())
# column select; Row attributes satisfy from_attributes. if defer_content:
stmt = ( stmt = stmt
select(Tool.id, Tool.user_id, Tool.name, Tool.specs, Tool.meta, Tool.updated_at, Tool.created_at)
if defer_content
else select(Tool)
).order_by(Tool.updated_at.desc())
if user_id is not None:
if user_group_ids is None:
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)}
stmt = AccessGrants.has_permission_filter(
db=db,
query=stmt,
DocumentModel=Tool,
filter={'user_id': user_id, 'group_ids': user_group_ids},
resource_type='tool',
permission=permission,
)
result = await db.execute(stmt) result = await db.execute(stmt)
all_tools = result.all() if defer_content else result.scalars().all() all_tools = result.scalars().all()
user_ids = list(set(tool.user_id for tool in all_tools)) user_ids = list(set(tool.user_id for tool in all_tools))
tool_ids = [tool.id for tool in all_tools] tool_ids = [tool.id for tool in all_tools]
@ -235,22 +208,31 @@ class ToolsTable:
defer_content: bool = False, defer_content: bool = False,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> list[ToolUserModel]: ) -> list[ToolUserModel]:
tools = await self.get_tools(defer_content=defer_content, db=db)
user_groups = await Groups.get_groups_by_member_id(user_id, db=db) user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
user_group_ids = {group.id for group in user_groups} user_group_ids = {group.id for group in user_groups}
return await self.get_tools(
defer_content=defer_content, result = []
db=db, for tool in tools:
user_id=user_id, if tool.user_id == user_id:
user_group_ids=user_group_ids, result.append(tool)
permission=permission, elif await AccessGrants.has_access(
) user_id=user_id,
resource_type='tool',
resource_id=tool.id,
permission=permission,
user_group_ids=user_group_ids,
db=db,
):
result.append(tool)
return result
async def get_tool_valves_by_id(self, id: str, db: AsyncSession | None = None) -> dict | None: async def get_tool_valves_by_id(self, id: str, db: AsyncSession | None = None) -> dict | None:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
tool = await db.get(Tool, id) tool = await db.get(Tool, id)
return decrypt_valves(tool.valves if tool else None) return tool.valves if tool.valves else {}
except Exception: except Exception as e:
log.exception(f'Error getting tool valves by id {id}') log.exception(f'Error getting tool valves by id {id}')
return None return None
@ -259,9 +241,7 @@ class ToolsTable:
) -> ToolValves | None: ) -> ToolValves | None:
try: try:
async with get_async_db_context(db) as db: async with get_async_db_context(db) as db:
await db.execute( await db.execute(update(Tool).filter_by(id=id).values(valves=valves, updated_at=int(time.time())))
update(Tool).filter_by(id=id).values(valves=encrypt_valves(valves), updated_at=int(time.time()))
)
await db.commit() await db.commit()
return await self.get_tool_by_id(id, db=db) return await self.get_tool_by_id(id, db=db)
except Exception: except Exception:
@ -280,7 +260,7 @@ class ToolsTable:
if 'valves' not in user_settings['tools']: if 'valves' not in user_settings['tools']:
user_settings['tools']['valves'] = {} user_settings['tools']['valves'] = {}
return decrypt_valves(user_settings['tools']['valves'].get(id)) return user_settings['tools']['valves'].get(id, {})
except Exception as e: except Exception as e:
log.exception(f'Error getting user values by id {id} and user_id {user_id}: {e}') log.exception(f'Error getting user values by id {id} and user_id {user_id}: {e}')
return None return None
@ -298,12 +278,12 @@ class ToolsTable:
if 'valves' not in user_settings['tools']: if 'valves' not in user_settings['tools']:
user_settings['tools']['valves'] = {} user_settings['tools']['valves'] = {}
user_settings['tools']['valves'][id] = encrypt_valves(valves) user_settings['tools']['valves'][id] = valves
# Update the user settings in the database # Update the user settings in the database
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db) await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
return valves return user_settings['tools']['valves'][id]
except Exception as e: except Exception as e:
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}') log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
return None return None
@ -317,8 +297,8 @@ class ToolsTable:
if access_grants is not None: if access_grants is not None:
await AccessGrants.set_access_grants('tool', id, access_grants, db=db) await AccessGrants.set_access_grants('tool', id, access_grants, db=db)
# populate_existing: the Core update above bypasses any identity-map copy tool = await db.get(Tool, id)
tool = await db.get(Tool, id, populate_existing=True) await db.refresh(tool)
return await self._to_tool_model(tool, db=db) return await self._to_tool_model(tool, db=db)
except Exception: except Exception:
return None return None

View file

@ -4,18 +4,12 @@ from __future__ import annotations
import datetime import datetime
import time import time
from typing import Literal, Optional from typing import Optional
from open_webui.env import DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL 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.internal.db import Base, JSONField, get_async_db_context
from open_webui.utils.misc import throttle from open_webui.utils.misc import throttle
from open_webui.utils.validate import validate_image_url from open_webui.utils.validate import validate_profile_image_url
from pydantic import ( from pydantic import BaseModel, ConfigDict, field_validator, model_validator
BaseModel,
ConfigDict,
Field,
field_validator,
model_validator,
)
from sqlalchemy import ( from sqlalchemy import (
JSON, JSON,
BigInteger, BigInteger,
@ -33,6 +27,7 @@ from sqlalchemy import (
select, select,
update, update,
) )
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
#################### ####################
@ -42,97 +37,6 @@ from sqlalchemy.ext.asyncio import AsyncSession
#################### ####################
class InterfaceTitleSettings(BaseModel):
model_config = ConfigDict(extra='forbid')
auto: bool | None = None
class InterfaceImageCompressionSize(BaseModel):
model_config = ConfigDict(extra='forbid')
width: int | float | Literal[''] | None = None
height: int | float | Literal[''] | None = None
class InterfaceFloatingActionButton(BaseModel):
model_config = ConfigDict(extra='forbid')
id: str
label: str
input: bool
prompt: str
class InterfaceSettings(BaseModel):
"""Fields owned by the Interface settings panel; not the entire user UI dict."""
model_config = ConfigDict(extra='forbid')
autoTags: bool | None = None
autoFollowUps: bool | None = None
highContrastMode: bool | None = None
detectArtifacts: bool | None = None
responseAutoCopy: bool | None = None
showUsername: bool | None = None
showUpdateToast: bool | None = None
showChangelog: bool | None = None
showEmojiInCall: bool | None = None
voiceInterruption: bool | None = None
displayMultiModelResponsesInTabs: bool | None = None
chatFadeStreamingText: bool | None = None
richTextInput: bool | None = None
showFormattingToolbar: bool | None = None
insertPromptAsRichText: bool | None = None
promptAutocomplete: bool | None = None
insertSuggestionPrompt: bool | None = None
keepFollowUpPrompts: bool | None = None
insertFollowUpPrompt: bool | None = None
regenerateMenu: bool | None = None
enableMessageQueue: bool | None = None
largeTextAsFile: bool | None = None
copyFormatted: bool | None = None
collapseCodeBlocks: bool | None = None
renderMarkdownInUserMessages: bool | None = None
renderMarkdownInAssistantMessages: bool | None = None
expandDetails: bool | None = None
chatHoverPreview: bool | None = None
renderMarkdownInPreviews: bool | None = None
chatBubble: bool | None = None
widescreenMode: bool | None = None
splitLargeChunks: bool | None = None
scrollOnBranchChange: bool | None = None
scrollOnResponseGeneration: bool | None = None
showFilesOnTerminalSelect: bool | None = None
temporaryChatByDefault: bool | None = None
userLocation: bool | None = None
showChatTitleInTab: bool | None = None
iframeSandboxAllowScripts: bool | None = None
iframeSandboxAllowSameOrigin: bool | None = None
iframeSandboxAllowForms: bool | None = None
iframeSandboxAllowDownloads: bool | None = None
terminalPreviewAllowSameOrigin: bool | None = None
stylizedPdfExport: bool | None = None
hapticFeedback: bool | None = None
ctrlEnterToSend: bool | None = None
showFloatingActionButtons: bool | None = None
imageCompression: bool | None = None
imageCompressionInChannels: bool | None = None
landingPageMode: Literal['', 'chat'] | None = None
chatDirection: Literal['LTR', 'RTL', 'auto'] | None = None
terminalFileDisplay: Literal['sidebar', 'inline'] | None = None
defaultUploadContext: Literal['full', 'focused'] | None = None
webSearch: Literal['always'] | None = None
models: list[str] | None = None
backgroundImageUrl: str | None = None
fontFamily: str | None = None
textScale: float | None = None
title: InterfaceTitleSettings | None = None
imageCompressionSize: InterfaceImageCompressionSize | None = None
floatingActionButtons: list[InterfaceFloatingActionButton] | None = None
class UserSettings(BaseModel): class UserSettings(BaseModel):
ui: dict | None = {} ui: dict | None = {}
model_config = ConfigDict(extra='allow') model_config = ConfigDict(extra='allow')
@ -165,7 +69,6 @@ class User(Base): # identity & profile
# Metadata # Metadata
info = Column(JSON, nullable=True) info = Column(JSON, nullable=True)
variables = Column(JSON, nullable=True)
settings = Column(JSON, nullable=True) settings = Column(JSON, nullable=True)
oauth = Column(JSON, nullable=True) oauth = Column(JSON, nullable=True)
scim = Column(JSON, nullable=True) scim = Column(JSON, nullable=True)
@ -202,7 +105,6 @@ class UserModel(BaseModel):
status_expires_at: int | None = None status_expires_at: int | None = None
info: dict | None = None info: dict | None = None
variables: dict = Field(default_factory=dict, exclude=True)
settings: UserSettings | None = None settings: UserSettings | None = None
oauth: dict | None = None oauth: dict | None = None
@ -224,11 +126,6 @@ class UserModel(BaseModel):
self.profile_image_url = self.profile_image_url or _DEFAULT_PROFILE_IMAGE_URL.format(user_id=self.id) self.profile_image_url = self.profile_image_url or _DEFAULT_PROFILE_IMAGE_URL.format(user_id=self.id)
return self return self
@field_validator('variables', mode='before')
@classmethod
def normalize_variables(cls, value):
return value if isinstance(value, dict) else {}
class UserStatusModel(UserModel): class UserStatusModel(UserModel):
is_active: bool = False is_active: bool = False
@ -277,7 +174,7 @@ class UpdateProfileForm(BaseModel):
@field_validator('profile_image_url') @field_validator('profile_image_url')
@classmethod @classmethod
def check_profile_image_url(cls, v: str) -> str: def check_profile_image_url(cls, v: str) -> str:
return validate_image_url(v) return validate_profile_image_url(v)
class UserGroupIdsModel(UserModel): class UserGroupIdsModel(UserModel):
@ -367,7 +264,7 @@ class UserUpdateForm(BaseModel):
def check_profile_image_url(cls, v: str | None) -> str | None: def check_profile_image_url(cls, v: str | None) -> str | None:
if v is None: if v is None:
return v return v
return validate_image_url(v) return validate_profile_image_url(v)
class UsersTable: class UsersTable:
@ -382,11 +279,6 @@ class UsersTable:
oauth: dict | None = None, oauth: dict | None = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> UserModel | None: ) -> UserModel | None:
try:
profile_image_url = validate_image_url(profile_image_url)
except ValueError:
profile_image_url = '/user.png'
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
user = UserModel( user = UserModel(
**{ **{
@ -405,6 +297,7 @@ class UsersTable:
result = User(**user.model_dump()) result = User(**user.model_dump())
session.add(result) session.add(result)
await session.commit() await session.commit()
await session.refresh(result)
return user if result else None return user if result else None
# database read methods # database read methods
@ -456,17 +349,16 @@ class UsersTable:
sub: str, sub: str,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> UserModel | None: ) -> UserModel | None:
"""Look up a user by OAuth provider + subject claim.""" """Look up a user by OAuth provider + subject claim (dialect-aware JSON filter)."""
sub = str(sub)
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
# Subscript, never contains(): on a JSON column contains() degrades to a substring LIKE. dialect = session.bind.dialect.name
sub_expr = User.oauth[provider]['sub'].as_string() query = select(User)
query = select(User).where(sub_expr == sub) if dialect == 'sqlite':
# SQLite preserves JSON numeric type here; Postgres ->> already compares numeric JSON as text. oauth_match = User.oauth.contains({provider: {'sub': sub}})
if session.get_bind().dialect.name == 'sqlite' and sub.isdecimal(): query = query.where(oauth_match)
sub_int = int(sub) elif dialect == 'postgresql':
if str(sub_int) == sub and sub_int <= 2**63 - 1: oauth_match = User.oauth[provider].cast(JSONB)['sub'].astext == sub
query = select(User).where(or_(sub_expr == sub, sub_expr == sub_int)) query = query.where(oauth_match)
row = (await session.execute(query)).scalars().first() row = (await session.execute(query)).scalars().first()
return UserModel.model_validate(row) if row else None return UserModel.model_validate(row) if row else None
@ -476,76 +368,27 @@ class UsersTable:
external_id: str, external_id: str,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> UserModel | None: ) -> UserModel | None:
"""Look up a user by SCIM provider + external ID.""" """Look up a user by SCIM provider + external ID (dialect-aware JSON filter)."""
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
# Subscript, never contains(): on a JSON column contains() degrades to a substring LIKE. dialect = session.bind.dialect.name
query = select(User).where(User.scim[provider]['external_id'].as_string() == external_id) 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() row = (await session.execute(query)).scalars().first()
return UserModel.model_validate(row) if row else None return UserModel.model_validate(row) if row else None
async def get_scim_users(
self,
filter: dict | None = None,
sort: dict | None = None,
skip: int | None = None,
limit: int | None = None,
db: AsyncSession | None = None,
) -> dict:
async with get_async_db_context(db) as session:
stmt = select(User).where(or_(User.oauth.cast(String) != 'null', User.scim.cast(String) != 'null'))
if filter:
user_id = filter.get('id')
if user_id:
stmt = stmt.where(User.id == user_id)
email = filter.get('email')
if email:
stmt = stmt.where(func.lower(User.email) == email.lower())
order_by = sort.get('order_by') if sort else None
direction = sort.get('direction') if sort else None
if order_by == 'created_at':
stmt = stmt.order_by(User.created_at.asc() if direction == 'asc' else User.created_at.desc())
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
total = count_result.scalar()
if skip is not None:
stmt = stmt.offset(skip)
if limit is not None:
stmt = stmt.limit(limit)
result = await session.execute(stmt)
users = result.scalars().all()
return {
'users': [UserModel.model_validate(user) for user in users],
'total': total,
}
async def get_scim_user_by_id(
self,
id: str,
db: AsyncSession | None = None,
) -> UserModel | None:
async with get_async_db_context(db) as session:
stmt = select(User).where(
User.id == id,
or_(User.oauth.cast(String) != 'null', User.scim.cast(String) != 'null'),
)
user = (await session.execute(stmt)).scalars().first()
return UserModel.model_validate(user) if user else None
async def get_users( async def get_users(
self, self,
filter: dict | None = None, filter: dict | None = None,
sort: dict | None = None,
skip: int | None = None, skip: int | None = None,
limit: int | None = None, limit: int | None = None,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> dict: ) -> dict:
"""Paginated user listing with optional filters and sort.""" """Paginated user listing with optional filters for role, group, and channel."""
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
# Deferred imports to avoid circular dependencies # Deferred imports to avoid circular dependencies
from open_webui.models.channels import ChannelMember from open_webui.models.channels import ChannelMember
@ -606,63 +449,64 @@ class UsersTable:
if exclude_roles: if exclude_roles:
stmt = stmt.filter(~User.role.in_(exclude_roles)) stmt = stmt.filter(~User.role.in_(exclude_roles))
order_by = sort.get('order_by') if sort else None order_by = filter.get('order_by')
direction = sort.get('direction') if sort else None direction = filter.get('direction')
if order_by and order_by.startswith('group_id:'): if order_by and order_by.startswith('group_id:'):
group_id = order_by.split(':', 1)[1] group_id = order_by.split(':', 1)[1]
# Subquery that checks if the user belongs to the group # Subquery that checks if the user belongs to the group
membership_exists = exists( membership_exists = exists(
select(GroupMember.id).where( select(GroupMember.id).where(
GroupMember.user_id == User.id, GroupMember.user_id == User.id,
GroupMember.group_id == group_id, GroupMember.group_id == group_id,
)
) )
)
# CASE: user in group → 1, user not in group → 0 # CASE: user in group → 1, user not in group → 0
group_sort = case((membership_exists, 1), else_=0) group_sort = case((membership_exists, 1), else_=0)
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(group_sort.asc(), User.name.asc()) stmt = stmt.order_by(group_sort.asc(), User.name.asc())
else: else:
stmt = stmt.order_by(group_sort.desc(), User.name.asc()) stmt = stmt.order_by(group_sort.desc(), User.name.asc())
elif order_by == 'name': elif order_by == 'name':
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(User.name.asc()) stmt = stmt.order_by(User.name.asc())
else: else:
stmt = stmt.order_by(User.name.desc()) stmt = stmt.order_by(User.name.desc())
elif order_by == 'email': elif order_by == 'email':
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(User.email.asc()) stmt = stmt.order_by(User.email.asc())
else: else:
stmt = stmt.order_by(User.email.desc()) stmt = stmt.order_by(User.email.desc())
elif order_by == 'created_at': elif order_by == 'created_at':
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(User.created_at.asc()) stmt = stmt.order_by(User.created_at.asc())
else: else:
stmt = stmt.order_by(User.created_at.desc()) stmt = stmt.order_by(User.created_at.desc())
elif order_by == 'last_active_at': elif order_by == 'last_active_at':
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(User.last_active_at.asc()) stmt = stmt.order_by(User.last_active_at.asc())
else: else:
stmt = stmt.order_by(User.last_active_at.desc()) stmt = stmt.order_by(User.last_active_at.desc())
elif order_by == 'updated_at': elif order_by == 'updated_at':
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(User.updated_at.asc()) stmt = stmt.order_by(User.updated_at.asc())
else: else:
stmt = stmt.order_by(User.updated_at.desc()) stmt = stmt.order_by(User.updated_at.desc())
elif order_by == 'role': elif order_by == 'role':
if direction == 'asc': if direction == 'asc':
stmt = stmt.order_by(User.role.asc()) stmt = stmt.order_by(User.role.asc())
else: else:
stmt = stmt.order_by(User.role.desc()) stmt = stmt.order_by(User.role.desc())
elif not filter:
else:
stmt = stmt.order_by(User.created_at.desc()) stmt = stmt.order_by(User.created_at.desc())
# Count BEFORE pagination # Count BEFORE pagination
@ -717,6 +561,13 @@ class UsersTable:
row = (await session.execute(stmt)).scalars().first() row = (await session.execute(stmt)).scalars().first()
return UserModel.model_validate(row) if row else None 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:
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 def get_num_users_active_today(self, db: AsyncSession | None = None) -> int | None:
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
current_timestamp = int(time.time()) current_timestamp = int(time.time())
@ -733,6 +584,7 @@ class UsersTable:
return None return None
user.role = role user.role = role
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
async def update_user_status_by_id( async def update_user_status_by_id(
@ -745,6 +597,7 @@ class UsersTable:
for key, value in form_data.model_dump(exclude_none=True).items(): for key, value in form_data.model_dump(exclude_none=True).items():
setattr(user, key, value) setattr(user, key, value)
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
async def update_user_profile_image_url_by_id( async def update_user_profile_image_url_by_id(
@ -753,17 +606,13 @@ class UsersTable:
profile_image_url: str, profile_image_url: str,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> UserModel | None: ) -> UserModel | None:
try:
profile_image_url = validate_image_url(profile_image_url)
except ValueError:
profile_image_url = '/user.png'
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
user = await session.get(User, id) user = await session.get(User, id)
if user is None: if user is None:
return None return None
user.profile_image_url = profile_image_url user.profile_image_url = profile_image_url
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
@throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL) @throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL)
@ -781,19 +630,17 @@ class UsersTable:
if not user: if not user:
return None return None
oauth = dict(user.oauth or {}) oauth = dict(user.oauth or {})
provider_oauth = oauth.get(provider) oauth[provider] = {'sub': sub}
provider_oauth = dict(provider_oauth) if isinstance(provider_oauth, dict) else {}
provider_oauth['sub'] = str(sub)
oauth[provider] = provider_oauth
user.oauth = oauth user.oauth = oauth
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
async def update_user_scim_by_id( async def update_user_scim_by_id(
self, self,
id: str, id: str,
provider: str, provider: str,
external_id: str | None, external_id: str,
db: AsyncSession | None = None, db: AsyncSession | None = None,
) -> UserModel | None: ) -> UserModel | None:
"""Update or insert a SCIM provider/external_id pair into the user's scim JSON field.""" """Update or insert a SCIM provider/external_id pair into the user's scim JSON field."""
@ -803,10 +650,9 @@ class UsersTable:
return None return None
scim = dict(user.scim or {}) scim = dict(user.scim or {})
scim[provider] = {'external_id': external_id} scim[provider] = {'external_id': external_id}
if scim != user.scim: user.scim = scim
user.scim = scim
user.updated_at = int(time.time())
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
async def update_user_by_id(self, id: str, updated: dict, db: AsyncSession | None = None) -> UserModel | None: async def update_user_by_id(self, id: str, updated: dict, db: AsyncSession | None = None) -> UserModel | None:
@ -817,6 +663,7 @@ class UsersTable:
for key, value in updated.items(): for key, value in updated.items():
setattr(user, key, value) setattr(user, key, value)
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
# settings update helper # settings update helper
@ -828,20 +675,10 @@ class UsersTable:
if not user: if not user:
return None return None
user_settings = dict(user.settings or {}) user_settings = dict(user.settings or {})
updated = dict(updated)
ui_settings = updated.pop('ui', None)
user_settings.update(updated) user_settings.update(updated)
if ui_settings is not None:
# UI updates are field-level patches: omission keeps a value; null resets it.
current_ui_settings = dict(user_settings.get('ui') or {})
for key, value in ui_settings.items():
if value is None:
current_ui_settings.pop(key, None)
else:
current_ui_settings[key] = value
user_settings['ui'] = current_ui_settings
user.settings = user_settings user.settings = user_settings
await session.commit() await session.commit()
await session.refresh(user)
return UserModel.model_validate(user) return UserModel.model_validate(user)
async def delete_user_by_id(self, id: str, db: AsyncSession | None = None) -> bool: async def delete_user_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
@ -888,8 +725,8 @@ class UsersTable:
async def get_valid_user_ids(self, user_ids: list[str], db: AsyncSession | None = None) -> list[str]: 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 session: async with get_async_db_context(db) as session:
result = await session.execute(select(User.id).where(User.id.in_(user_ids))) result = await session.execute(select(User).where(User.id.in_(user_ids)))
return list(result.scalars().all()) return [u.id for u in result.scalars().all()]
async def get_super_admin_user(self, db: AsyncSession | None = None) -> UserModel | None: async def get_super_admin_user(self, db: AsyncSession | None = None) -> UserModel | None:
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
@ -915,11 +752,11 @@ class UsersTable:
async def is_user_active(self, user_id: str, db: AsyncSession | None = None) -> bool: async def is_user_active(self, user_id: str, db: AsyncSession | None = None) -> bool:
async with get_async_db_context(db) as session: async with get_async_db_context(db) as session:
last_active_at = await session.scalar(select(User.last_active_at).where(User.id == user_id)) user = await session.get(User, user_id)
if last_active_at: if user and user.last_active_at:
# Consider user active if last_active_at within the last 3 minutes # Consider user active if last_active_at within the last 3 minutes
three_minutes_ago = int(time.time()) - 180 three_minutes_ago = int(time.time()) - 180
return last_active_at >= three_minutes_ago return user.last_active_at >= three_minutes_ago
return False return False

View file

@ -1,379 +0,0 @@
import asyncio
import logging
import re
import time
from typing import Any, Optional
from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
from open_webui.models.config import Config
from open_webui.models.knowledge import KnowledgeModel
log = logging.getLogger(__name__)
EXTERNAL_KNOWLEDGE_CONNECTIONS_CONFIG_KEY = 'external_knowledge.connections'
IDENTIFIER_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]*$')
async def _get_external_connection(connection_id: str) -> Optional[dict]:
connections = await Config.get(EXTERNAL_KNOWLEDGE_CONNECTIONS_CONFIG_KEY, []) or []
return next((connection for connection in connections if connection.get('id') == connection_id), None)
def _get_path(data: Any, path: Optional[str], default=None):
if not path:
return default
value = data
for part in path.split('.'):
if isinstance(value, dict):
value = value.get(part, default)
else:
return default
return value
def _normalize_result(result: dict, mapping: dict, knowledge: KnowledgeModel, distance: Optional[float] = None) -> dict:
content = _get_path(result, mapping.get('content_field', 'content'), '')
title = _get_path(result, mapping.get('title_field', 'title'), None)
source = _get_path(result, mapping.get('source_field', 'source'), None)
url = _get_path(result, mapping.get('url_field', 'url'), None)
document_id = _get_path(result, mapping.get('document_id_field', 'document_id'), None)
page = _get_path(result, mapping.get('page_field', 'page'), None)
metadata = _get_path(result, mapping.get('metadata_field', 'metadata'), {}) or {}
score = _get_path(result, mapping.get('score_field', 'score'), distance)
if not isinstance(metadata, dict):
metadata = {'external_metadata': metadata}
source_name = source or title or metadata.get('source') or metadata.get('name') or knowledge.name
metadata.update(
{
'name': title or source_name,
'source': source_name,
'url': url,
'file_id': document_id or f'external-{knowledge.id}',
'knowledge_id': knowledge.id,
'knowledge_name': knowledge.name,
'external': True,
}
)
if page is not None:
metadata['page'] = page
if document_id is not None:
metadata['document_id'] = document_id
return {
'content': content,
'metadata': metadata,
'distance': score,
}
def _source_config(knowledge: KnowledgeModel) -> dict:
external = (knowledge.meta or {}).get('external', {})
source = external.get('source') or {}
return source.get('config') or {}
def _root_field(path: Optional[str]) -> Optional[str]:
if not path:
return None
return path.split('.')[0]
def _safe_identifier(value: str, label: str) -> str:
if not value or not IDENTIFIER_RE.match(value):
raise RuntimeError(f'Invalid {label}')
return value
async def _retrieve_qdrant(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]:
try:
from qdrant_client import QdrantClient
except ImportError as exc:
raise RuntimeError('qdrant-client is not installed') from exc
if not embedding_function:
raise RuntimeError('Embedding function is not configured')
config = connection.get('config') or {}
external = (knowledge.meta or {}).get('external', {})
source = external.get('source') or {}
collection_name = source.get('name')
if not collection_name:
raise RuntimeError('External source collection is not configured')
source_config = _source_config(knowledge)
vector_field = source_config.get('vector_field') or None
vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX)
def _search():
client = QdrantClient(
url=connection.get('endpoint'),
api_key=(auth_config or {}).get('api_key'),
timeout=config.get('timeout') or 30,
)
return client.query_points(
collection_name=collection_name,
query=vector,
using=vector_field,
limit=count,
)
response = await asyncio.to_thread(_search)
mapping = {
'content_field': source_config.get('content_field') or 'payload.text',
'metadata_field': source_config.get('metadata_field') or 'payload.metadata',
'document_id_field': source_config.get('document_id_field') or 'id',
'score_field': 'score',
}
normalized = []
for point in response.points:
normalized.append(_normalize_result(point.model_dump(), mapping, knowledge, distance=point.score))
return normalized
async def _retrieve_milvus(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]:
try:
from pymilvus import MilvusClient
except ImportError as exc:
raise RuntimeError('pymilvus is not installed') from exc
if not embedding_function:
raise RuntimeError('Embedding function is not configured')
config = connection.get('config') or {}
external = (knowledge.meta or {}).get('external', {})
source = external.get('source') or {}
collection_name = source.get('name')
if not collection_name:
raise RuntimeError('Milvus collection is not configured')
source_config = _source_config(knowledge)
vector_field = source_config.get('vector_field') or 'vector'
content_field = source_config.get('content_field') or 'data.text'
metadata_field = source_config.get('metadata_field') or 'metadata'
vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX)
def _search():
client_kwargs = {
'uri': connection.get('endpoint'),
}
token = (auth_config or {}).get('api_key') or (auth_config or {}).get('token')
if token:
client_kwargs['token'] = token
if config.get('db_name'):
client_kwargs['db_name'] = config.get('db_name')
client = MilvusClient(**client_kwargs)
output_fields = {
field
for field in (
_root_field(content_field),
_root_field(metadata_field),
_root_field(source_config.get('document_id_field')),
)
if field and field != vector_field
}
kwargs = {
'collection_name': collection_name,
'data': [vector],
'anns_field': vector_field,
'limit': count,
'output_fields': list(output_fields),
}
return client.search(**kwargs)
response = await asyncio.to_thread(_search)
mapping = {
'content_field': content_field,
'metadata_field': metadata_field,
'document_id_field': source_config.get('document_id_field') or 'id',
'score_field': 'distance',
}
normalized = []
for hit in response[0] if response else []:
item = dict(hit)
entity = item.get('entity') or {}
result = {
**entity,
'id': item.get('id') or entity.get('id'),
'distance': item.get('distance'),
}
normalized.append(_normalize_result(result, mapping, knowledge, distance=item.get('distance')))
return normalized
async def _retrieve_pgvector(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]:
try:
import psycopg
from pgvector.psycopg import register_vector
from psycopg.rows import dict_row
except ImportError as exc:
raise RuntimeError('psycopg and pgvector are required for pgvector retrieval') from exc
if not embedding_function:
raise RuntimeError('Embedding function is not configured')
config = connection.get('config') or {}
external = (knowledge.meta or {}).get('external', {})
source = external.get('source') or {}
collection_name = source.get('name')
if not collection_name:
raise RuntimeError('pgvector collection is not configured')
source_config = _source_config(knowledge)
table_name = source_config.get('table_name') or 'document_chunk'
collection_field = source_config.get('collection_field') or 'collection_name'
content_field = source_config.get('content_field') or 'text'
vector_field = source_config.get('vector_field') or 'vector'
metadata_field = source_config.get('metadata_field') or 'vmetadata'
document_id_field = source_config.get('document_id_field') or 'id'
vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX)
def _search():
from psycopg import sql
table_identifier = sql.SQL('.').join(
sql.Identifier(_safe_identifier(part, 'table name')) for part in table_name.split('.')
)
collection_identifier = sql.Identifier(_safe_identifier(collection_field, 'collection field'))
content_identifier = sql.Identifier(_safe_identifier(content_field, 'content field'))
vector_identifier = sql.Identifier(_safe_identifier(vector_field, 'vector field'))
document_id_identifier = sql.Identifier(_safe_identifier(document_id_field, 'document id field'))
metadata_sql = (
sql.Identifier(_safe_identifier(metadata_field, 'metadata field'))
if metadata_field
else sql.SQL("'{}'::jsonb")
)
with psycopg.connect(
connection.get('endpoint'),
row_factory=dict_row,
connect_timeout=config.get('timeout') or 30,
) as conn:
register_vector(conn)
with conn.cursor() as cur:
cur.execute(
sql.SQL(
"""
SELECT {document_id} AS id,
{content} AS content,
{metadata} AS metadata,
{vector_column} <=> %s AS distance
FROM {table_name}
WHERE {collection} = %s
ORDER BY distance ASC
LIMIT %s
"""
).format(
document_id=document_id_identifier,
content=content_identifier,
metadata=metadata_sql,
vector_column=vector_identifier,
table_name=table_identifier,
collection=collection_identifier,
),
(vector, collection_name, count),
)
return cur.fetchall()
rows = await asyncio.to_thread(_search)
mapping = {
'content_field': 'content',
'metadata_field': 'metadata',
'document_id_field': 'id',
'score_field': 'distance',
}
return [_normalize_result(row, mapping, knowledge, distance=row.get('distance')) for row in rows]
async def retrieve_external_knowledge(
request,
knowledge: KnowledgeModel,
queries: list[str],
count: int,
user=None,
) -> dict:
external = (knowledge.meta or {}).get('external', {})
connection_id = external.get('connection_id')
if not connection_id:
raise RuntimeError('External knowledge connection is not configured')
connection = await _get_external_connection(connection_id)
if not connection:
raise RuntimeError('External knowledge connection not found')
return await retrieve_external_knowledge_for_connection(request, knowledge, connection, queries, count, user=user)
async def retrieve_external_knowledge_for_connection(
request,
knowledge: KnowledgeModel,
connection: dict,
queries: list[str],
count: int,
user=None,
) -> dict:
auth_config = connection.get('auth_config') or {}
if not connection.get('enabled', True):
raise RuntimeError('External knowledge connection is disabled')
started_at = time.monotonic()
chunks = []
provider = (connection.get('provider') or '').lower()
for query in queries:
if provider == 'qdrant':
chunks.extend(
await _retrieve_qdrant(
connection,
auth_config,
knowledge,
query,
count,
getattr(request.app.state, 'EMBEDDING_FUNCTION', None),
)
)
elif provider == 'milvus':
chunks.extend(
await _retrieve_milvus(
connection,
auth_config,
knowledge,
query,
count,
getattr(request.app.state, 'EMBEDDING_FUNCTION', None),
)
)
elif provider == 'pgvector':
chunks.extend(
await _retrieve_pgvector(
connection,
auth_config,
knowledge,
query,
count,
getattr(request.app.state, 'EMBEDDING_FUNCTION', None),
)
)
else:
raise RuntimeError(f'Unsupported external knowledge provider: {connection.get("provider")}')
chunks = chunks[:count]
log.info(
'external_knowledge_retrieval knowledge_id=%s connection_id=%s provider=%s user_id=%s latency_ms=%s result_count=%s',
knowledge.id,
connection.get('id'),
connection.get('provider'),
getattr(user, 'id', None),
round((time.monotonic() - started_at) * 1000),
len(chunks),
)
return {
'documents': [[chunk['content'] for chunk in chunks]],
'metadatas': [[chunk['metadata'] for chunk in chunks]],
'distances': [[chunk['distance'] for chunk in chunks]],
}

View file

@ -7,7 +7,6 @@ from typing import List, Optional
import requests import requests
from fastapi import HTTPException, status from fastapi import HTTPException, status
from langchain_core.documents import Document from langchain_core.documents import Document
from open_webui.utils.json_codec import JSONCodec
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -65,6 +64,25 @@ class DatalabMarkerLoader:
} }
return mime_map.get(ext, 'application/octet-stream') return mime_map.get(ext, 'application/octet-stream')
def check_marker_request_status(self, request_id: str) -> dict:
url = f'{self.api_base_url}/{request_id}'
headers = {'X-Api-Key': self.api_key}
try:
response = requests.get(url, headers=headers)
response.raise_for_status()
result = response.json()
log.info(f'Marker API status check for request {request_id}: {result}')
return result
except requests.HTTPError as e:
log.error(f'Error checking Marker request status: {e}')
raise HTTPException(
status.HTTP_502_BAD_GATEWAY,
detail=f'Failed to check Marker request: {e}',
)
except ValueError as e:
log.error(f'Invalid JSON checking Marker request: {e}')
raise HTTPException(status.HTTP_502_BAD_GATEWAY, detail=f'Invalid JSON: {e}')
def load(self) -> List[Document]: def load(self) -> List[Document]:
filename = os.path.basename(self.file_path) filename = os.path.basename(self.file_path)
mime_type = self._get_mime_type(filename) mime_type = self._get_mime_type(filename)
@ -85,10 +103,7 @@ class DatalabMarkerLoader:
form_data['additional_config'] = self.additional_config form_data['additional_config'] = self.additional_config
log.info( log.info(
"Datalab Marker POST request parameters: {'filename': '%s', 'mime_type': '%s', **%s}", f"Datalab Marker POST request parameters: {{'filename': '{filename}', 'mime_type': '{mime_type}', **{form_data}}}"
filename,
mime_type,
form_data,
) )
try: try:
@ -152,7 +167,7 @@ class DatalabMarkerLoader:
'total_cost', 'total_cost',
) )
} }
log.info('Marker processing completed successfully: %s', json.dumps(summary, indent=2)) log.info(f'Marker processing completed successfully: {json.dumps(summary, indent=2)}')
break break
if status_val == 'failed' or success_val is False: if status_val == 'failed' or success_val is False:
@ -219,7 +234,7 @@ class DatalabMarkerLoader:
try: try:
with open(output_path, 'w', encoding='utf-8') as f: with open(output_path, 'w', encoding='utf-8') as f:
f.write(full_text) f.write(full_text)
log.info('Saved Marker output to: %s', output_path) log.info(f'Saved Marker output to: {output_path}')
except Exception as e: except Exception as e:
log.warning(f'Failed to write marker output to disk: {e}') log.warning(f'Failed to write marker output to disk: {e}')
@ -234,11 +249,11 @@ class DatalabMarkerLoader:
images = final_result.get('images', {}) images = final_result.get('images', {})
if images: if images:
metadata['image_count'] = len(images) metadata['image_count'] = len(images)
metadata['images'] = JSONCodec.dumps(list(images.keys())) metadata['images'] = json.dumps(list(images.keys()))
for k, v in metadata.items(): for k, v in metadata.items():
if isinstance(v, (dict, list)): if isinstance(v, (dict, list)):
metadata[k] = JSONCodec.dumps(v) metadata[k] = json.dumps(v)
elif v is None: elif v is None:
metadata[k] = '' metadata[k] = ''

View file

@ -6,7 +6,7 @@ from urllib.parse import quote
import requests import requests
from langchain_core.document_loaders import BaseLoader from langchain_core.document_loaders import BaseLoader
from langchain_core.documents import Document from langchain_core.documents import Document
from open_webui.utils.headers import include_user_info_headers, parse_custom_headers from open_webui.utils.headers import include_user_info_headers
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -19,9 +19,6 @@ class ExternalDocumentLoader(BaseLoader):
api_key: str, api_key: str,
mime_type=None, mime_type=None,
user=None, user=None,
user_groups=None,
headers=None,
metadata=None,
**kwargs, **kwargs,
) -> None: ) -> None:
self.url = url self.url = url
@ -31,9 +28,6 @@ class ExternalDocumentLoader(BaseLoader):
self.mime_type = mime_type self.mime_type = mime_type
self.user = user self.user = user
self.user_groups = user_groups
self.headers = headers
self.metadata = metadata
def load(self) -> List[Document]: def load(self) -> List[Document]:
with open(self.file_path, 'rb') as f: with open(self.file_path, 'rb') as f:
@ -51,8 +45,6 @@ class ExternalDocumentLoader(BaseLoader):
except Exception: except Exception:
pass pass
headers.update(parse_custom_headers(self.headers, self.user, self.metadata, user_groups=self.user_groups))
if self.user is not None: if self.user is not None:
headers = include_user_info_headers(headers, self.user) headers = include_user_info_headers(headers, self.user)

View file

@ -30,9 +30,6 @@ class ExternalWebLoader(BaseLoader):
response = requests.post( response = requests.post(
self.external_url, self.external_url,
headers={ headers={
# LICENSE covers this Open WebUI user-agent identifier.
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
# https://docs.openwebui.com/license.
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) External Web Loader', 'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) External Web Loader',
'Authorization': f'Bearer {self.external_api_key}', 'Authorization': f'Bearer {self.external_api_key}',
}, },

View file

@ -1,121 +0,0 @@
from importlib import import_module
from pathlib import Path
from bs4 import BeautifulSoup
from langchain_core.documents import Document
class TextLoader:
def __init__(self, file_path, encoding=None):
self.file_path = str(file_path)
self.encoding = encoding
def load(self) -> list[Document]:
try:
text = Path(self.file_path).read_text(encoding=self.encoding)
except Exception as e:
raise RuntimeError(f'Error loading {self.file_path}') from e
return [
Document(
page_content=text,
metadata={'source': self.file_path},
)
]
class HTMLLoader(TextLoader):
def load(self) -> list[Document]:
with open(self.file_path, encoding=self.encoding) as file:
soup = BeautifulSoup(file, 'lxml')
return [
Document(
page_content=soup.get_text(),
metadata={'source': self.file_path, 'title': str(soup.title.string) if soup.title else ''},
)
]
class DocxLoader(TextLoader):
def load(self) -> list[Document]:
import docx2txt
return [
Document(
page_content=docx2txt.process(Path(self.file_path).expanduser()),
metadata={'source': self.file_path},
)
]
class UnstructuredLoader:
def __init__(self, file_path, file_format, mode='single', **kwargs):
# Match the optional-package check; format dependencies are loaded when parsing.
import_module('unstructured')
self.file_path = file_path
self.file_format = file_format
self.mode = mode
self.kwargs = kwargs
def load(self) -> list[Document]:
file_format = self.file_format
if file_format in ('doc', 'ppt', 'pptx'):
from unstructured.file_utils.filetype import detect_filetype
legacy_format = 'doc' if file_format == 'doc' else 'ppt'
try:
import_module('magic')
except ImportError:
is_legacy = Path(self.file_path).suffix == f'.{legacy_format}'
else:
is_legacy = detect_filetype(self.file_path).name.lower() == legacy_format
file_format = legacy_format if is_legacy else legacy_format + 'x'
elif file_format == 'msg':
from unstructured.file_utils.filetype import detect_filetype
detected = detect_filetype(self.file_path).name
if detected not in ('EML', 'MSG'):
raise ValueError(f'Unsupported email file type: {detected}')
file_format = 'email' if detected == 'EML' else 'msg'
module = import_module(f'unstructured.partition.{file_format}')
elements = getattr(module, f'partition_{file_format}')(filename=self.file_path, **self.kwargs)
metadata = {'source': str(self.file_path)}
if self.mode == 'elements':
return [
Document(
page_content=str(element),
metadata={
**metadata,
**element.metadata.to_dict(),
'category': element.category,
'element_id': element.id,
},
)
for element in elements
]
return [Document(page_content='\n\n'.join(map(str, elements)), metadata=metadata)]
class DocumentIntelligenceLoader:
def __init__(self, file_path, api_endpoint, api_key=None, azure_credential=None, api_model='prebuilt-layout'):
if (api_key is None) == (azure_credential is None):
raise ValueError('Provide exactly one of api_key or azure_credential.')
self.file_path = file_path
self.api_endpoint = api_endpoint
self.api_key = api_key
self.azure_credential = azure_credential
self.api_model = api_model
def load(self) -> list[Document]:
from azure.ai.documentintelligence import DocumentIntelligenceClient
from azure.core.credentials import AzureKeyCredential
credential = self.azure_credential if self.azure_credential is not None else AzureKeyCredential(self.api_key)
with DocumentIntelligenceClient(self.api_endpoint, credential) as client, open(self.file_path, 'rb') as file:
result = client.begin_analyze_document(
self.api_model,
body=file,
content_type='application/octet-stream',
output_content_format='markdown',
).result()
return [Document(page_content=result.content, metadata=result.as_dict())]

View file

@ -1,37 +1,28 @@
import asyncio import asyncio
import csv import json
import logging import logging
import os
import sys import sys
import zipfile
import ftfy import ftfy
import requests import requests
from fastapi import HTTPException
from azure.identity import DefaultAzureCredential from azure.identity import DefaultAzureCredential
from langchain_core.documents import Document from langchain_community.document_loaders import (
from open_webui.env import ( AzureAIDocumentIntelligenceLoader,
AIOHTTP_CLIENT_SESSION_SSL, BSHTMLLoader,
GLOBAL_LOG_LEVEL, CSVLoader,
USE_SLIM, Docx2txtLoader,
MINERU_MAX_MARKDOWN_BYTES, OutlookMessageLoader,
REQUESTS_VERIFY, PyPDFLoader,
TextLoader,
YoutubeLoader,
) )
from langchain_core.documents import Document
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, GLOBAL_LOG_LEVEL, REQUESTS_VERIFY
from open_webui.retrieval.loaders.datalab_marker import DatalabMarkerLoader from open_webui.retrieval.loaders.datalab_marker import DatalabMarkerLoader
from open_webui.retrieval.loaders.external_document import ExternalDocumentLoader from open_webui.retrieval.loaders.external_document import ExternalDocumentLoader
from open_webui.retrieval.loaders.local import (
DocumentIntelligenceLoader,
DocxLoader,
HTMLLoader,
TextLoader,
UnstructuredLoader,
)
from open_webui.retrieval.loaders.mineru import MinerULoader from open_webui.retrieval.loaders.mineru import MinerULoader
from open_webui.retrieval.loaders.mistral import MistralLoader from open_webui.retrieval.loaders.mistral import MistralLoader
from open_webui.retrieval.loaders.paddleocr_vl import PADDLEOCR_VL_SUPPORTED_EXTENSIONS, PaddleOCRVLLoader from open_webui.retrieval.loaders.paddleocr_vl import PaddleOCRVLLoader
from open_webui.retrieval.loaders.pdf import PDFLoader
from open_webui.utils.headers import get_user_groups_for_custom_headers
from open_webui.utils.json_codec import JSONCodec
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -52,7 +43,6 @@ known_source_ext = [
'h', 'h',
'c', 'c',
'cs', 'cs',
'ino',
'sql', 'sql',
'log', 'log',
'ini', 'ini',
@ -91,18 +81,8 @@ known_source_ext = [
'yaml', 'yaml',
'yml', 'yml',
'toml', 'toml',
'svg',
] ]
known_archive_ext = {'docx', 'epub', 'odt', 'pptx', 'xlsx'}
known_archive_content_types = {
'application/epub+zip',
'application/vnd.oasis.opendocument.text',
'application/vnd.openxmlformats-officedocument.presentationml.presentation',
'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet',
'application/vnd.openxmlformats-officedocument.wordprocessingml.document',
}
class ExcelLoader: class ExcelLoader:
"""Fallback Excel loader using pandas when unstructured is not installed.""" """Fallback Excel loader using pandas when unstructured is not installed."""
@ -126,66 +106,6 @@ class ExcelLoader:
] ]
def get_csv_summary(filename: str, file_path: str, encoding: str) -> str | None:
try:
with open(file_path, newline='', encoding=encoding) as f:
sample = f.read(4096)
f.seek(0)
try:
dialect = csv.Sniffer().sniff(sample)
except csv.Error:
dialect = csv.excel
total_rows = 0
max_columns = 0
headers = []
for row in csv.reader(f, dialect):
total_rows += 1
max_columns = max(max_columns, len(row))
if total_rows == 1:
headers = [header.lstrip('\ufeff') for header in row]
except Exception:
return None
if total_rows == 0:
return None
return (
f'Table: {total_rows} rows incl. header; '
f'{max(total_rows - 1, 0)} data rows; '
f'{max_columns} columns: {", ".join(headers)}.'
)
class CSVLoaderWithSummary:
def __init__(self, file_path: str, filename: str, encoding: str):
self.file_path = file_path
self.filename = filename
self.encoding = encoding
def load(self) -> list[Document]:
docs = []
try:
with open(self.file_path, newline='', encoding=self.encoding) as file:
for index, row in enumerate(csv.DictReader(file)):
fields = []
for key, value in row.items():
if isinstance(value, str):
value = value.strip()
elif isinstance(value, list):
value = ','.join(v.strip() for v in value)
fields.append(f'{key.strip() if key is not None else key}: {value}')
content = '\n'.join(fields)
docs.append(Document(page_content=content, metadata={'source': self.file_path, 'row': index}))
except Exception as e:
raise RuntimeError(f'Error loading {self.file_path}') from e
if os.getenv('ENABLE_RAG_CSV_SUMMARY', 'False').lower() == 'true':
summary = get_csv_summary(self.filename, self.file_path, self.encoding)
if summary:
docs.insert(0, Document(page_content=summary, metadata={'source': self.file_path, 'row': -1}))
return docs
class PptxLoader: class PptxLoader:
"""Fallback PowerPoint loader using python-pptx when unstructured is not installed.""" """Fallback PowerPoint loader using python-pptx when unstructured is not installed."""
@ -213,11 +133,10 @@ class PptxLoader:
class TikaLoader: class TikaLoader:
def __init__(self, url, file_path, mime_type=None, extract_images=None, server_version='3'): def __init__(self, url, file_path, mime_type=None, extract_images=None):
self.url = url self.url = url
self.file_path = file_path self.file_path = file_path
self.mime_type = mime_type self.mime_type = mime_type
self.server_version = str(server_version or '3')
self.extract_images = extract_images self.extract_images = extract_images
@ -233,15 +152,16 @@ class TikaLoader:
if self.extract_images == True: if self.extract_images == True:
headers['X-Tika-PDFextractInlineImages'] = 'true' headers['X-Tika-PDFextractInlineImages'] = 'true'
endpoint_path = 'tika/json/md' if self.server_version == '4' else 'tika/text' endpoint = self.url
content_key = 'tk:content' if self.server_version == '4' else 'X-TIKA:content' if not endpoint.endswith('/'):
endpoint = f'{self.url.rstrip("/")}/{endpoint_path}' endpoint += '/'
endpoint += 'tika/text'
r = requests.put(endpoint, data=data, headers=headers, verify=REQUESTS_VERIFY) r = requests.put(endpoint, data=data, headers=headers, verify=REQUESTS_VERIFY)
if r.ok: if r.ok:
raw_metadata = r.json() raw_metadata = r.json()
text = raw_metadata.get(content_key, '<No text content found>').strip() text = raw_metadata.get('X-TIKA:content', '<No text content found>').strip()
if 'Content-Type' in raw_metadata: if 'Content-Type' in raw_metadata:
headers['Content-Type'] = raw_metadata['Content-Type'] headers['Content-Type'] = raw_metadata['Content-Type']
@ -263,7 +183,6 @@ class DoclingLoader:
self.params = params or {} self.params = params or {}
def load(self) -> list[Document]: def load(self) -> list[Document]:
page_break_marker = '\f'
with open(self.file_path, 'rb') as f: with open(self.file_path, 'rb') as f:
headers = {} headers = {}
if self.api_key: if self.api_key:
@ -280,10 +199,6 @@ class DoclingLoader:
}, },
data={ data={
'image_export_mode': 'placeholder', 'image_export_mode': 'placeholder',
'md_page_break_placeholder': page_break_marker,
# Keep Docling params as user-provided form values. Encoding nested
# values here would make Open WebUI responsible for Docling's API
# quirks and could break when Docling changes its form contract.
**self.params, **self.params,
}, },
headers=headers, headers=headers,
@ -291,29 +206,10 @@ class DoclingLoader:
) )
if r.ok: if r.ok:
result = r.json() result = r.json()
# Docling reports failed and skipped conversions inside HTTP 200 responses.
conversion_status = result.get('status')
if conversion_status in ['failure', 'skipped']:
error_details = (
'; '.join(filter(None, (error.get('error_message') for error in result.get('errors', []))))
or 'no error message provided'
)
raise Exception(f'Error calling Docling: conversion status {conversion_status} - {error_details}')
document_data = result.get('document', {}) document_data = result.get('document', {})
md_content = document_data.get('md_content') or '' text = document_data.get('md_content', '<No text content found>')
text = md_content or '<No text content found>'
metadata = {'Content-Type': self.mime_type} if self.mime_type else {} metadata = {'Content-Type': self.mime_type} if self.mime_type else {}
if page_break_marker in md_content:
documents = [
Document(page_content=page.strip(), metadata={**metadata, 'page': page_idx})
for page_idx, page in enumerate(md_content.split(page_break_marker))
if page.strip()
]
if documents:
log.debug('Docling extracted text: %s', text)
return documents
log.debug('Docling extracted text: %s', text) log.debug('Docling extracted text: %s', text)
return [Document(page_content=text, metadata=metadata)] return [Document(page_content=text, metadata=metadata)]
@ -333,18 +229,12 @@ class Loader:
def __init__(self, engine: str = '', **kwargs): def __init__(self, engine: str = '', **kwargs):
self.engine = engine self.engine = engine
self.user = kwargs.get('user', None) self.user = kwargs.get('user', None)
self.user_groups = kwargs.get('user_groups', None)
self.metadata = kwargs.get('metadata', {})
self.kwargs = kwargs self.kwargs = kwargs
def load(self, filename: str, file_content_type: str, file_path: str) -> list[Document]: def load(self, filename: str, file_content_type: str, file_path: str) -> list[Document]:
loader = self._get_loader(filename, file_content_type, file_path) loader = self._get_loader(filename, file_content_type, file_path)
docs = loader.load() docs = loader.load()
# ftfy's auto mode unescapes entities on every line before the first literal '<', rewriting the document. return [Document(page_content=ftfy.fix_text(doc.page_content), metadata=doc.metadata) for doc in docs]
return [
Document(page_content=ftfy.fix_text(doc.page_content, unescape_html=False), metadata=doc.metadata)
for doc in docs
]
async def aload(self, filename: str, file_content_type: str, file_path: str) -> list[Document]: async def aload(self, filename: str, file_content_type: str, file_path: str) -> list[Document]:
""" """
@ -356,13 +246,6 @@ class Loader:
loop for the entire parse — minutes for large PDFs. This offloads loop for the entire parse — minutes for large PDFs. This offloads
the work to a worker thread so the loop stays responsive. the work to a worker thread so the loop stays responsive.
""" """
# Group lookup is async-only, so it must happen before `load`
# is offloaded to a thread without a running event loop.
if self.engine == 'external' and self.user_groups is None:
self.user_groups = await get_user_groups_for_custom_headers(
self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_HEADERS'), self.user
)
return await asyncio.to_thread(self.load, filename, file_content_type, file_path) return await asyncio.to_thread(self.load, filename, file_content_type, file_path)
def _is_text_file(self, file_ext: str, file_content_type: str) -> bool: def _is_text_file(self, file_ext: str, file_content_type: str) -> bool:
@ -402,20 +285,13 @@ class Loader:
try: try:
raw.decode('utf-8') raw.decode('utf-8')
return 'utf-8' return 'utf-8'
except UnicodeDecodeError as e: except UnicodeDecodeError:
first_non_utf8 = e.start pass
# Use chardet as a hint, not as ground truth # Use chardet as a hint, not as ground truth
import chardet import chardet
# chardet is pure Python (~1.3s/MB), so sample around the first bad byte detected = chardet.detect(raw)
window = 256 * 1024
sample_start = max(0, first_non_utf8 - window // 2)
sample = raw[sample_start : sample_start + window]
detected = chardet.detect(sample)
# A stray byte can sit far from the real payload, leaving the sample with nothing to read
if len(sample.translate(None, delete=bytes(range(128)))) < 64 and len(sample) < len(raw):
detected = chardet.detect(raw)
detected_enc = (detected.get('encoding') or '').lower().replace('-', '').replace('_', '') detected_enc = (detected.get('encoding') or '').lower().replace('-', '').replace('_', '')
# Map chardet's detected encoding to the correct superset codec. # Map chardet's detected encoding to the correct superset codec.
@ -517,27 +393,6 @@ class Loader:
def _get_loader(self, filename: str, file_content_type: str, file_path: str): def _get_loader(self, filename: str, file_content_type: str, file_path: str):
file_ext = filename.split('.')[-1].lower() file_ext = filename.split('.')[-1].lower()
if file_ext in known_archive_ext or file_content_type in known_archive_content_types:
max_file_size = self.kwargs.get('FILE_MAX_SIZE')
try:
max_file_size_bytes = int(max_file_size) * 1024 * 1024 if max_file_size else 100 * 1024 * 1024
except (TypeError, ValueError):
max_file_size_bytes = 100 * 1024 * 1024
if max_file_size_bytes > 0:
try:
with zipfile.ZipFile(file_path) as archive:
uncompressed_size = sum(entry.file_size for entry in archive.infolist())
except (zipfile.BadZipFile, OSError):
pass
else:
max_bytes = min(
max(10 * 1024 * 1024, os.path.getsize(file_path) * 100),
max_file_size_bytes,
)
if uncompressed_size > max_bytes:
raise ValueError('Document archive is too large after decompression')
if ( if (
self.engine == 'external' self.engine == 'external'
and self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_URL') and self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_URL')
@ -549,13 +404,6 @@ class Loader:
api_key=self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_API_KEY'), api_key=self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_API_KEY'),
mime_type=file_content_type, mime_type=file_content_type,
user=self.user, user=self.user,
user_groups=self.user_groups,
headers=self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_HEADERS'),
metadata={
**self.metadata,
'file_name': filename,
'file_content_type': file_content_type,
},
) )
elif self.engine == 'tika' and self.kwargs.get('TIKA_SERVER_URL'): elif self.engine == 'tika' and self.kwargs.get('TIKA_SERVER_URL'):
if self._is_text_file(file_ext, file_content_type): if self._is_text_file(file_ext, file_content_type):
@ -564,7 +412,6 @@ class Loader:
loader = TikaLoader( loader = TikaLoader(
url=self.kwargs.get('TIKA_SERVER_URL'), url=self.kwargs.get('TIKA_SERVER_URL'),
file_path=file_path, file_path=file_path,
server_version=self.kwargs.get('TIKA_SERVER_VERSION'),
extract_images=self.kwargs.get('PDF_EXTRACT_IMAGES'), extract_images=self.kwargs.get('PDF_EXTRACT_IMAGES'),
) )
elif ( elif (
@ -618,8 +465,8 @@ class Loader:
params = self.kwargs.get('DOCLING_PARAMS', {}) params = self.kwargs.get('DOCLING_PARAMS', {})
if not isinstance(params, dict): if not isinstance(params, dict):
try: try:
params = JSONCodec.loads(params) params = json.loads(params)
except JSONCodec.JSONDecodeError: except json.JSONDecodeError:
log.error('Invalid DOCLING_PARAMS format, expected JSON object') log.error('Invalid DOCLING_PARAMS format, expected JSON object')
params = {} params = {}
@ -644,14 +491,14 @@ class Loader:
) )
): ):
if self.kwargs.get('DOCUMENT_INTELLIGENCE_KEY') != '': if self.kwargs.get('DOCUMENT_INTELLIGENCE_KEY') != '':
loader = DocumentIntelligenceLoader( loader = AzureAIDocumentIntelligenceLoader(
file_path=file_path, file_path=file_path,
api_endpoint=self.kwargs.get('DOCUMENT_INTELLIGENCE_ENDPOINT'), api_endpoint=self.kwargs.get('DOCUMENT_INTELLIGENCE_ENDPOINT'),
api_key=self.kwargs.get('DOCUMENT_INTELLIGENCE_KEY'), api_key=self.kwargs.get('DOCUMENT_INTELLIGENCE_KEY'),
api_model=self.kwargs.get('DOCUMENT_INTELLIGENCE_MODEL'), api_model=self.kwargs.get('DOCUMENT_INTELLIGENCE_MODEL'),
) )
else: else:
loader = DocumentIntelligenceLoader( loader = AzureAIDocumentIntelligenceLoader(
file_path=file_path, file_path=file_path,
api_endpoint=self.kwargs.get('DOCUMENT_INTELLIGENCE_ENDPOINT'), api_endpoint=self.kwargs.get('DOCUMENT_INTELLIGENCE_ENDPOINT'),
azure_credential=DefaultAzureCredential(), azure_credential=DefaultAzureCredential(),
@ -664,6 +511,7 @@ class Loader:
mineru_timeout = int(mineru_timeout) mineru_timeout = int(mineru_timeout)
except ValueError: except ValueError:
mineru_timeout = 300 mineru_timeout = 300
loader = MinerULoader( loader = MinerULoader(
file_path=file_path, file_path=file_path,
api_mode=self.kwargs.get('MINERU_API_MODE', 'local'), api_mode=self.kwargs.get('MINERU_API_MODE', 'local'),
@ -671,7 +519,6 @@ class Loader:
api_key=self.kwargs.get('MINERU_API_KEY', ''), api_key=self.kwargs.get('MINERU_API_KEY', ''),
params=self.kwargs.get('MINERU_PARAMS', {}), params=self.kwargs.get('MINERU_PARAMS', {}),
timeout=mineru_timeout, timeout=mineru_timeout,
max_markdown_bytes=MINERU_MAX_MARKDOWN_BYTES,
) )
elif ( elif (
self.engine == 'mistral_ocr' self.engine == 'mistral_ocr'
@ -682,49 +529,27 @@ class Loader:
base_url=self.kwargs.get('MISTRAL_OCR_API_BASE_URL'), base_url=self.kwargs.get('MISTRAL_OCR_API_BASE_URL'),
api_key=self.kwargs.get('MISTRAL_OCR_API_KEY'), api_key=self.kwargs.get('MISTRAL_OCR_API_KEY'),
file_path=file_path, file_path=file_path,
use_base64=self.kwargs.get('MISTRAL_OCR_USE_BASE64', False),
user=self.user,
) )
elif ( elif self.engine == 'paddleocr_vl' and self.kwargs.get('PADDLEOCR_VL_TOKEN') != '':
self.engine == 'paddleocr_vl'
and self.kwargs.get('PADDLEOCR_VL_BASE_URL')
and self.kwargs.get('PADDLEOCR_VL_TOKEN')
and file_ext in PADDLEOCR_VL_SUPPORTED_EXTENSIONS
):
loader = PaddleOCRVLLoader( loader = PaddleOCRVLLoader(
api_url=self.kwargs.get('PADDLEOCR_VL_BASE_URL'), api_url=self.kwargs.get('PADDLEOCR_VL_BASE_URL'),
token=self.kwargs.get('PADDLEOCR_VL_TOKEN'), token=self.kwargs.get('PADDLEOCR_VL_TOKEN'),
file_path=file_path, file_path=file_path,
) )
else: else:
if USE_SLIM:
if file_ext == 'csv':
return CSVLoaderWithSummary(file_path, filename, self._detect_text_encoding(file_path))
if file_ext in ['htm', 'html']:
return HTMLLoader(file_path, encoding=self._detect_text_encoding(file_path))
if file_ext in ['txt', 'md', 'markdown', 'rst', 'xml'] or self._is_text_file(
file_ext, file_content_type
):
return TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
raise HTTPException(
503,
'This file type requires an external document extractor in slim. Configure one that supports it.',
)
if file_ext == 'pdf': if file_ext == 'pdf':
loader = PDFLoader( loader = PyPDFLoader(
file_path, file_path,
extract_images=self.kwargs.get('PDF_EXTRACT_IMAGES'), extract_images=self.kwargs.get('PDF_EXTRACT_IMAGES'),
mode=self.kwargs.get('PDF_LOADER_MODE', 'page'), mode=self.kwargs.get('PDF_LOADER_MODE', 'page'),
) )
elif file_ext == 'csv': elif file_ext == 'csv':
loader = CSVLoaderWithSummary( loader = CSVLoader(file_path, encoding=self._detect_text_encoding(file_path))
file_path,
filename,
self._detect_text_encoding(file_path),
)
elif file_ext == 'rst': elif file_ext == 'rst':
try: try:
loader = UnstructuredLoader(file_path, 'rst', mode='elements') from langchain_community.document_loaders import UnstructuredRSTLoader
loader = UnstructuredRSTLoader(file_path, mode='elements')
except ImportError: except ImportError:
log.warning( log.warning(
"The 'unstructured' package is not installed. " "The 'unstructured' package is not installed. "
@ -734,7 +559,9 @@ class Loader:
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path)) loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_ext == 'xml': elif file_ext == 'xml':
try: try:
loader = UnstructuredLoader(file_path, 'xml') from langchain_community.document_loaders import UnstructuredXMLLoader
loader = UnstructuredXMLLoader(file_path)
except ImportError: except ImportError:
log.warning( log.warning(
"The 'unstructured' package is not installed. " "The 'unstructured' package is not installed. "
@ -743,12 +570,14 @@ class Loader:
) )
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path)) loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_ext in ['htm', 'html']: elif file_ext in ['htm', 'html']:
loader = HTMLLoader(file_path, encoding='unicode_escape') loader = BSHTMLLoader(file_path, open_encoding='unicode_escape')
elif file_ext == 'md': elif file_ext == 'md':
loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path)) loader = TextLoader(file_path, encoding=self._detect_text_encoding(file_path))
elif file_content_type == 'application/epub+zip': elif file_content_type == 'application/epub+zip':
try: try:
loader = UnstructuredLoader(file_path, 'epub') from langchain_community.document_loaders import UnstructuredEPubLoader
loader = UnstructuredEPubLoader(file_path)
except ImportError: except ImportError:
raise ValueError( raise ValueError(
"Processing .epub files requires the 'unstructured' package. " "Processing .epub files requires the 'unstructured' package. "
@ -758,10 +587,12 @@ class Loader:
file_content_type == 'application/vnd.openxmlformats-officedocument.wordprocessingml.document' file_content_type == 'application/vnd.openxmlformats-officedocument.wordprocessingml.document'
or file_ext == 'docx' or file_ext == 'docx'
): ):
loader = DocxLoader(file_path) loader = Docx2txtLoader(file_path)
elif file_ext == 'doc' or file_content_type == 'application/msword': elif file_ext == 'doc' or file_content_type == 'application/msword':
try: try:
loader = UnstructuredLoader(file_path, 'doc') from langchain_community.document_loaders import UnstructuredWordDocumentLoader
loader = UnstructuredWordDocumentLoader(file_path)
except ImportError: except ImportError:
raise ValueError( raise ValueError(
"Processing .doc files requires the 'unstructured' package. " "Processing .doc files requires the 'unstructured' package. "
@ -772,7 +603,9 @@ class Loader:
'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet', 'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet',
] or file_ext in ['xls', 'xlsx']: ] or file_ext in ['xls', 'xlsx']:
try: try:
loader = UnstructuredLoader(file_path, 'xlsx') from langchain_community.document_loaders import UnstructuredExcelLoader
loader = UnstructuredExcelLoader(file_path)
except ImportError: except ImportError:
log.warning( log.warning(
"The 'unstructured' package is not installed. " "The 'unstructured' package is not installed. "
@ -785,7 +618,9 @@ class Loader:
'application/vnd.openxmlformats-officedocument.presentationml.presentation', 'application/vnd.openxmlformats-officedocument.presentationml.presentation',
] or file_ext in ['ppt', 'pptx']: ] or file_ext in ['ppt', 'pptx']:
try: try:
loader = UnstructuredLoader(file_path, 'ppt' if file_ext == 'ppt' else 'pptx') from langchain_community.document_loaders import UnstructuredPowerPointLoader
loader = UnstructuredPowerPointLoader(file_path)
except ImportError: except ImportError:
log.warning( log.warning(
"The 'unstructured' package is not installed. " "The 'unstructured' package is not installed. "
@ -794,17 +629,12 @@ class Loader:
) )
loader = PptxLoader(file_path) loader = PptxLoader(file_path)
elif file_ext == 'msg': elif file_ext == 'msg':
try: loader = OutlookMessageLoader(file_path)
# unstructured parses .msg via python-oxmsg; avoids extract_msg's beautifulsoup4<4.14 conflict
loader = UnstructuredLoader(file_path, 'msg', process_attachments=False)
except ImportError:
raise ValueError(
"Processing .msg files requires the 'unstructured' package. "
'Install it with: pip install unstructured'
)
elif file_ext == 'odt': elif file_ext == 'odt':
try: try:
loader = UnstructuredLoader(file_path, 'odt') from langchain_community.document_loaders import UnstructuredODTLoader
loader = UnstructuredODTLoader(file_path)
except ImportError: except ImportError:
raise ValueError( raise ValueError(
"Processing .odt files requires the 'unstructured' package. " "Processing .odt files requires the 'unstructured' package. "

View file

@ -1,109 +0,0 @@
import logging
import time
from collections.abc import Iterator
from typing import Any
from urllib.parse import urlparse
import requests
from langchain_core.document_loaders import BaseLoader
from langchain_core.documents import Document
log = logging.getLogger(__name__)
MICROSOFT_BROWSE_RETRY_STATUS_CODES = {202, 429, 500, 502, 503, 504}
MICROSOFT_BROWSE_MAX_RETRIES = 2
class MicrosoftWebIQLoader(BaseLoader):
def __init__(
self,
urls: str | list[str],
api_base_url: str,
api_key: str,
language: str = 'en',
verify_ssl: bool = True,
timeout: Any = None,
continue_on_failure: bool = True,
) -> None:
self.urls = urls if isinstance(urls, list) else [urls]
self.api_base_url = api_base_url.rstrip('/')
self.api_key = api_key
self.language = language
self.verify_ssl = verify_ssl
self.timeout = timeout
self.continue_on_failure = continue_on_failure
def lazy_load(self) -> Iterator[Document]:
for url in self.urls:
try:
doc = self._browse_url(url)
if doc is not None:
yield doc
except Exception as e:
if self.continue_on_failure:
log.warning(f'Error browsing {url} with Microsoft Web IQ: {e}')
else:
raise e
def _browse_url(self, url: str) -> Document | None:
headers = {
'host': urlparse(self.api_base_url).netloc or 'api.microsoft.ai',
'x-apikey': self.api_key,
'content-type': 'application/json',
}
payload = {
'url': url,
'contentFormat': 'markdown',
'liveCrawl': 'fallback',
'renderDynamicPages': True,
'language': self.language,
}
try:
request_timeout = float(self.timeout)
except (TypeError, ValueError):
request_timeout = 60
request_timeout = request_timeout if request_timeout > 0 else 60
data: dict[str, Any] = {}
for attempt in range(MICROSOFT_BROWSE_MAX_RETRIES + 1):
response = requests.post(
f'{self.api_base_url}/browse',
json=payload,
headers=headers,
timeout=request_timeout,
verify=self.verify_ssl,
)
if response.status_code in MICROSOFT_BROWSE_RETRY_STATUS_CODES and attempt < MICROSOFT_BROWSE_MAX_RETRIES:
try:
body = response.json()
except Exception:
body = {}
retry_after = body.get('retryAfter') if isinstance(body, dict) else None
retry_after = retry_after or response.headers.get('Retry-After')
try:
delay = min(10.0, max(0.0, float(str(retry_after).rstrip('s'))))
except (TypeError, ValueError):
delay = min(8.0, float(2**attempt))
log.warning(
'Microsoft Browse %s returned HTTP %s; retrying in %.1fs',
url,
response.status_code,
delay,
)
time.sleep(delay)
continue
response.raise_for_status()
data = response.json()
break
content = data.get('content') or ''
if not isinstance(content, str) or not content.strip():
return None
metadata = {'source': data.get('url') or url}
if data.get('title'):
metadata['title'] = data['title']
return Document(page_content=content, metadata=metadata)

View file

@ -28,22 +28,20 @@ class MinerULoader:
api_key: str = '', api_key: str = '',
params: dict = None, params: dict = None,
timeout: Optional[int] = 300, timeout: Optional[int] = 300,
max_markdown_bytes: Optional[int] = None,
): ):
self.file_path = file_path self.file_path = file_path
self.api_mode = api_mode.lower() self.api_mode = api_mode.lower()
self.api_url = api_url.rstrip('/') self.api_url = api_url.rstrip('/')
self.api_key = api_key self.api_key = api_key
self.timeout = timeout self.timeout = timeout
self.max_markdown_bytes = max_markdown_bytes
# Parse params dict with defaults # Parse params dict with defaults
self.params = params or {} self.params = params or {}
self.enable_ocr = self.params.get('enable_ocr', False) self.enable_ocr = params.get('enable_ocr', False)
self.enable_formula = self.params.get('enable_formula', True) self.enable_formula = params.get('enable_formula', True)
self.enable_table = self.params.get('enable_table', True) self.enable_table = params.get('enable_table', True)
self.language = self.params.get('language', 'en') self.language = params.get('language', 'en')
self.model_version = self.params.get('model_version', 'pipeline') self.model_version = params.get('model_version', 'pipeline')
self.page_ranges = self.params.pop('page_ranges', '') self.page_ranges = self.params.pop('page_ranges', '')
@ -74,7 +72,7 @@ class MinerULoader:
Load document using Local API (synchronous). Load document using Local API (synchronous).
Posts file to /file_parse endpoint and gets immediate response. Posts file to /file_parse endpoint and gets immediate response.
""" """
log.info('Using MinerU Local API at %s', self.api_url) log.info(f'Using MinerU Local API at {self.api_url}')
filename = os.path.basename(self.file_path) filename = os.path.basename(self.file_path)
@ -97,8 +95,8 @@ class MinerULoader:
with open(self.file_path, 'rb') as f: with open(self.file_path, 'rb') as f:
files = {'files': (filename, f, 'application/octet-stream')} files = {'files': (filename, f, 'application/octet-stream')}
log.info('Sending file to MinerU Local API: %s', filename) log.info(f'Sending file to MinerU Local API: {filename}')
log.debug('Local API parameters: %s', form_data) log.debug(f'Local API parameters: {form_data}')
response = requests.post( response = requests.post(
f'{self.api_url}/file_parse', f'{self.api_url}/file_parse',
@ -163,7 +161,7 @@ class MinerULoader:
detail='MinerU returned empty markdown content', detail='MinerU returned empty markdown content',
) )
log.info('Successfully parsed document with MinerU Local API: %s', filename) log.info(f'Successfully parsed document with MinerU Local API: {filename}')
# Create metadata # Create metadata
metadata = { metadata = {
@ -180,7 +178,7 @@ class MinerULoader:
Load document using Cloud API (asynchronous). Load document using Cloud API (asynchronous).
Uses batch upload endpoint to avoid need for public file URLs. Uses batch upload endpoint to avoid need for public file URLs.
""" """
log.info('Using MinerU Cloud API at %s', self.api_url) log.info(f'Using MinerU Cloud API at {self.api_url}')
filename = os.path.basename(self.file_path) filename = os.path.basename(self.file_path)
@ -196,7 +194,7 @@ class MinerULoader:
# Step 4: Download and extract markdown from ZIP # Step 4: Download and extract markdown from ZIP
markdown_content = self._download_and_extract_zip(result['full_zip_url'], filename) markdown_content = self._download_and_extract_zip(result['full_zip_url'], filename)
log.info('Successfully parsed document with MinerU Cloud API: %s', filename) log.info(f'Successfully parsed document with MinerU Cloud API: {filename}')
# Create metadata # Create metadata
metadata = { metadata = {
@ -232,8 +230,8 @@ class MinerULoader:
if self.page_ranges: if self.page_ranges:
request_body['files'][0]['page_ranges'] = self.page_ranges request_body['files'][0]['page_ranges'] = self.page_ranges
log.info('Requesting upload URL for: %s', filename) log.info(f'Requesting upload URL for: {filename}')
log.debug('Cloud API request body: %s', request_body) log.debug(f'Cloud API request body: {request_body}')
try: try:
response = requests.post( response = requests.post(
@ -284,7 +282,7 @@ class MinerULoader:
) )
upload_url = file_urls[0] upload_url = file_urls[0]
log.info('Received upload URL for batch: %s', batch_id) log.info(f'Received upload URL for batch: {batch_id}')
return batch_id, upload_url return batch_id, upload_url
@ -334,7 +332,7 @@ class MinerULoader:
max_iterations = 300 # 10 minutes max (2 seconds per iteration) max_iterations = 300 # 10 minutes max (2 seconds per iteration)
poll_interval = 2 # seconds poll_interval = 2 # seconds
log.info('Polling batch status: %s', batch_id) log.info(f'Polling batch status: {batch_id}')
for iteration in range(max_iterations): for iteration in range(max_iterations):
try: try:
@ -393,7 +391,7 @@ class MinerULoader:
state = file_result.get('state') state = file_result.get('state')
if state == 'done': if state == 'done':
log.info('Processing complete for %s', filename) log.info(f'Processing complete for {filename}')
return file_result return file_result
elif state == 'failed': elif state == 'failed':
error_msg = file_result.get('err_msg', 'Unknown error') error_msg = file_result.get('err_msg', 'Unknown error')
@ -404,7 +402,7 @@ class MinerULoader:
elif state in ['waiting-file', 'pending', 'running', 'converting']: elif state in ['waiting-file', 'pending', 'running', 'converting']:
# Still processing # Still processing
if iteration % 10 == 0: # Log every 20 seconds if iteration % 10 == 0: # Log every 20 seconds
log.info('Processing status: %s (iteration %s/%s)', state, iteration + 1, max_iterations) log.info(f'Processing status: {state} (iteration {iteration + 1}/{max_iterations})')
time.sleep(poll_interval) time.sleep(poll_interval)
else: else:
log.warning(f'Unknown state: {state}') log.warning(f'Unknown state: {state}')
@ -421,7 +419,7 @@ class MinerULoader:
Download ZIP file from CDN and extract markdown content. Download ZIP file from CDN and extract markdown content.
Returns the markdown content as a string. Returns the markdown content as a string.
""" """
log.info('Downloading results from: %s', zip_url) log.info(f'Downloading results from: {zip_url}')
try: try:
response = requests.get(zip_url, timeout=60) response = requests.get(zip_url, timeout=60)
@ -437,77 +435,67 @@ class MinerULoader:
detail=f'Error downloading results: {str(e)}', detail=f'Error downloading results: {str(e)}',
) )
# Save ZIP to temporary file before reading. # Save ZIP to temporary file and extract
tmp_zip_path = None
markdown_content = None
try: try:
with tempfile.NamedTemporaryFile(delete=False, suffix='.zip') as tmp_zip: with tempfile.NamedTemporaryFile(delete=False, suffix='.zip') as tmp_zip:
tmp_zip.write(response.content) tmp_zip.write(response.content)
tmp_zip_path = tmp_zip.name tmp_zip_path = tmp_zip.name
with zipfile.ZipFile(tmp_zip_path, 'r') as zip_ref: with tempfile.TemporaryDirectory() as tmp_dir:
members = zip_ref.infolist() # Extract ZIP
all_files = [member.filename for member in members] with zipfile.ZipFile(tmp_zip_path, 'r') as zip_ref:
md_members = [member for member in members if member.filename.endswith('.md')] zip_ref.extractall(tmp_dir)
read_errors = []
for member in md_members: # Find markdown file - search recursively for any .md file
log.info('Found markdown file in ZIP: %s', member.filename) markdown_content = None
try: found_md_path = None
with zip_ref.open(member, 'r') as f:
if self.max_markdown_bytes is None: # First, list all files in the ZIP for debugging
content = f.read() all_files = []
else: for root, dirs, files in os.walk(tmp_dir):
content = f.read(self.max_markdown_bytes + 1) for file in files:
if len(content) > self.max_markdown_bytes: full_path = os.path.join(root, file)
raise HTTPException( all_files.append(full_path)
status.HTTP_502_BAD_GATEWAY, # Look for any .md file
detail=f'Markdown file in results ZIP is too large: {member.filename}', if file.endswith('.md'):
) found_md_path = full_path
markdown_content = content.decode('utf-8') log.info(f'Found markdown file at: {full_path}')
except UnicodeDecodeError as e: try:
read_errors.append(f'{member.filename}: {e}') with open(full_path, 'r', encoding='utf-8') as f:
log.warning(f'Failed to decode {member.filename}: {e}') markdown_content = f.read()
continue if markdown_content: # Use the first non-empty markdown file
except HTTPException: break
raise except Exception as e:
except Exception as e: log.warning(f'Failed to read {full_path}: {e}')
read_errors.append(f'{member.filename}: {e}')
log.warning(f'Failed to read {member.filename}: {e}')
continue
if markdown_content: if markdown_content:
break break
if markdown_content is None: if markdown_content is None:
log.error(f'Available files in ZIP: {all_files}') log.error(f'Available files in ZIP: {all_files}')
if read_errors: # Try to provide more helpful error message
error_msg = f"Found .md files but couldn't read them: {read_errors}" md_files = [f for f in all_files if f.endswith('.md')]
if md_files:
error_msg = f"Found .md files but couldn't read them: {md_files}"
else: else:
error_msg = f'No .md files found in ZIP. Available files: {all_files}' error_msg = f'No .md files found in ZIP. Available files: {all_files}'
raise HTTPException( raise HTTPException(
status.HTTP_502_BAD_GATEWAY, status.HTTP_502_BAD_GATEWAY,
detail=error_msg, detail=error_msg,
) )
# Clean up temporary ZIP file
os.unlink(tmp_zip_path)
except zipfile.BadZipFile as e: except zipfile.BadZipFile as e:
raise HTTPException( raise HTTPException(
status.HTTP_502_BAD_GATEWAY, status.HTTP_502_BAD_GATEWAY,
detail=f'Invalid ZIP file received: {e}', detail=f'Invalid ZIP file received: {e}',
) )
except HTTPException:
raise
except Exception as e: except Exception as e:
raise HTTPException( raise HTTPException(
status.HTTP_500_INTERNAL_SERVER_ERROR, status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f'Error extracting ZIP: {str(e)}', detail=f'Error extracting ZIP: {str(e)}',
) )
finally:
if tmp_zip_path:
try:
os.unlink(tmp_zip_path)
except FileNotFoundError:
pass
except Exception as e:
log.warning(f'Failed to remove temporary ZIP file {tmp_zip_path}: {e}')
if not markdown_content: if not markdown_content:
raise HTTPException( raise HTTPException(
@ -515,5 +503,5 @@ class MinerULoader:
detail='Extracted markdown content is empty', detail='Extracted markdown content is empty',
) )
log.info('Successfully extracted markdown content (%s characters)', len(markdown_content)) log.info(f'Successfully extracted markdown content ({len(markdown_content)} characters)')
return markdown_content return markdown_content

View file

@ -1,14 +1,15 @@
import base64 import asyncio
import logging import logging
import os import os
import sys import sys
import time import time
from typing import Any, Dict, List, Optional from contextlib import asynccontextmanager
from typing import Any, Dict, List
import aiohttp
import requests import requests
from langchain_core.documents import Document from langchain_core.documents import Document
from open_webui.env import ENABLE_FORWARD_USER_INFO_HEADERS, GLOBAL_LOG_LEVEL from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, GLOBAL_LOG_LEVEL
from open_webui.utils.headers import include_user_info_headers
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -16,12 +17,15 @@ log = logging.getLogger(__name__)
class MistralLoader: class MistralLoader:
""" """
Enhanced Mistral OCR loader. Enhanced Mistral OCR loader with both sync and async support.
Loads documents by processing them through the Mistral OCR API. Loads documents by processing them through the Mistral OCR API.
Performance Optimizations: Performance Optimizations:
- Differentiated timeouts for different operations - Differentiated timeouts for different operations
- Intelligent retry logic with exponential backoff - Intelligent retry logic with exponential backoff
- Memory-efficient file streaming for large files
- Connection pooling and keepalive optimization
- Semaphore-based concurrency control for batch processing
- Enhanced error handling with retryable error classification - Enhanced error handling with retryable error classification
""" """
@ -33,8 +37,6 @@ class MistralLoader:
timeout: int = 300, # 5 minutes default timeout: int = 300, # 5 minutes default
max_retries: int = 3, max_retries: int = 3,
enable_debug_logging: bool = False, enable_debug_logging: bool = False,
use_base64: bool = False,
user: Optional[Any] = None,
): ):
""" """
Initializes the loader with enhanced features. Initializes the loader with enhanced features.
@ -45,9 +47,6 @@ class MistralLoader:
timeout: Request timeout in seconds. timeout: Request timeout in seconds.
max_retries: Maximum number of retry attempts. max_retries: Maximum number of retry attempts.
enable_debug_logging: Enable detailed debug logs. enable_debug_logging: Enable detailed debug logs.
use_base64: Send the document as a data URL instead of uploading it first.
user: The requesting user, forwarded to Mistral via user-info headers
when ENABLE_FORWARD_USER_INFO_HEADERS is enabled.
""" """
if not api_key: if not api_key:
raise ValueError('API key cannot be empty.') raise ValueError('API key cannot be empty.')
@ -57,10 +56,9 @@ class MistralLoader:
self.base_url = base_url.rstrip('/') if base_url else 'https://api.mistral.ai/v1' self.base_url = base_url.rstrip('/') if base_url else 'https://api.mistral.ai/v1'
self.api_key = api_key self.api_key = api_key
self.file_path = file_path self.file_path = file_path
self.timeout = timeout
self.max_retries = max_retries self.max_retries = max_retries
self.debug = enable_debug_logging self.debug = enable_debug_logging
self.use_base64 = use_base64
self.user = user
# PERFORMANCE OPTIMIZATION: Differentiated timeouts for different operations # PERFORMANCE OPTIMIZATION: Differentiated timeouts for different operations
# This prevents long-running OCR operations from affecting quick operations # This prevents long-running OCR operations from affecting quick operations
@ -80,8 +78,6 @@ class MistralLoader:
'Authorization': f'Bearer {self.api_key}', 'Authorization': f'Bearer {self.api_key}',
'User-Agent': 'OpenWebUI-MistralLoader/2.0', # Helps API provider track usage 'User-Agent': 'OpenWebUI-MistralLoader/2.0', # Helps API provider track usage
} }
if self.user is not None and ENABLE_FORWARD_USER_INFO_HEADERS:
self.headers = include_user_info_headers(self.headers, self.user)
def _debug_log(self, message: str, *args) -> None: def _debug_log(self, message: str, *args) -> None:
""" """
@ -111,6 +107,32 @@ class MistralLoader:
log.error(f'JSON decode error: {json_err} - Response: {response.text}') log.error(f'JSON decode error: {json_err} - Response: {response.text}')
raise # Re-raise after logging raise # Re-raise after logging
async def _handle_response_async(self, response: aiohttp.ClientResponse) -> Dict[str, Any]:
"""Async version of response handling with better error info."""
try:
response.raise_for_status()
# Check content type
content_type = response.headers.get('content-type', '')
if 'application/json' not in content_type:
if response.status == 204:
return {}
text = await response.text()
raise ValueError(f'Unexpected content type: {content_type}, body: {text[:200]}...')
return await response.json()
except aiohttp.ClientResponseError as e:
error_text = await response.text() if response else 'No response'
log.error(f'HTTP {e.status}: {e.message} - Response: {error_text[:500]}')
raise
except aiohttp.ClientError as e:
log.error(f'Client error: {e}')
raise
except Exception as e:
log.error(f'Unexpected error processing response: {e}')
raise
def _is_retryable_error(self, error: Exception) -> bool: def _is_retryable_error(self, error: Exception) -> bool:
""" """
ENHANCEMENT: Intelligent error classification for retry logic. ENHANCEMENT: Intelligent error classification for retry logic.
@ -140,6 +162,10 @@ class MistralLoader:
status_code = error.response.status_code status_code = error.response.status_code
return status_code >= 500 or status_code == 429 return status_code >= 500 or status_code == 429
return False return False
if isinstance(error, (aiohttp.ClientConnectionError, aiohttp.ServerTimeoutError)):
return True # Async network/timeout errors are retryable
if isinstance(error, aiohttp.ClientResponseError):
return error.status >= 500 or error.status == 429
return False # All other errors are non-retryable return False # All other errors are non-retryable
def _retry_request_sync(self, request_func, *args, **kwargs): def _retry_request_sync(self, request_func, *args, **kwargs):
@ -166,11 +192,32 @@ class MistralLoader:
) )
time.sleep(wait_time) time.sleep(wait_time)
async def _retry_request_async(self, request_func, *args, **kwargs):
"""
ENHANCEMENT: Async retry logic with intelligent error classification.
Async version of retry logic that doesn't block the event loop during
wait periods. Uses the same exponential backoff strategy as sync version.
"""
for attempt in range(self.max_retries):
try:
return await request_func(*args, **kwargs)
except Exception as e:
if attempt == self.max_retries - 1 or not self._is_retryable_error(e):
raise
# PERFORMANCE OPTIMIZATION: Non-blocking exponential backoff
wait_time = min((2**attempt) + 0.5, 30) # Cap at 30 seconds
log.warning(
f'Retryable error (attempt {attempt + 1}/{self.max_retries}): {e}. Retrying in {wait_time}s...'
)
await asyncio.sleep(wait_time) # Non-blocking wait
def _upload_file(self) -> str: def _upload_file(self) -> str:
""" """
PERFORMANCE OPTIMIZATION: Enhanced file upload with streaming consideration. PERFORMANCE OPTIMIZATION: Enhanced file upload with streaming consideration.
Uploads the file to Mistral for OCR processing. Uploads the file to Mistral for OCR processing (sync version).
Uses context manager for file handling to ensure proper resource cleanup. Uses context manager for file handling to ensure proper resource cleanup.
Although streaming is not enabled for this endpoint, the file is opened Although streaming is not enabled for this endpoint, the file is opened
in a context manager to minimize memory usage duration. in a context manager to minimize memory usage duration.
@ -203,15 +250,57 @@ class MistralLoader:
file_id = response_data.get('id') file_id = response_data.get('id')
if not file_id: if not file_id:
raise ValueError('File ID not found in upload response.') raise ValueError('File ID not found in upload response.')
log.info('File uploaded successfully. File ID: %s', file_id) log.info(f'File uploaded successfully. File ID: {file_id}')
return file_id return file_id
except Exception as e: except Exception as e:
log.error(f'Failed to upload file: {e}') log.error(f'Failed to upload file: {e}')
raise raise
async def _upload_file_async(self, session: aiohttp.ClientSession) -> str:
"""Async file upload with streaming for better memory efficiency."""
url = f'{self.base_url}/files'
async def upload_request():
# Create multipart writer for streaming upload
writer = aiohttp.MultipartWriter('form-data')
# Add purpose field
purpose_part = writer.append('ocr')
purpose_part.set_content_disposition('form-data', name='purpose')
# Add file part with streaming
file_part = writer.append_payload(
aiohttp.streams.FilePayload(
self.file_path,
filename=self.file_name,
content_type='application/pdf',
)
)
file_part.set_content_disposition('form-data', name='file', filename=self.file_name)
self._debug_log(f'Uploading file: {self.file_name} ({self.file_size:,} bytes)')
async with session.post(
url,
data=writer,
headers=self.headers,
timeout=aiohttp.ClientTimeout(total=self.upload_timeout),
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as response:
return await self._handle_response_async(response)
response_data = await self._retry_request_async(upload_request)
file_id = response_data.get('id')
if not file_id:
raise ValueError('File ID not found in upload response.')
log.info(f'File uploaded successfully. File ID: {file_id}')
return file_id
def _get_signed_url(self, file_id: str) -> str: def _get_signed_url(self, file_id: str) -> str:
"""Retrieves a temporary signed URL for the uploaded file.""" """Retrieves a temporary signed URL for the uploaded file (sync version)."""
log.info('Getting signed URL for file ID: %s', file_id) log.info(f'Getting signed URL for file ID: {file_id}')
url = f'{self.base_url}/files/{file_id}/url' url = f'{self.base_url}/files/{file_id}/url'
params = {'expiry': 1} params = {'expiry': 1}
signed_url_headers = {**self.headers, 'Accept': 'application/json'} signed_url_headers = {**self.headers, 'Accept': 'application/json'}
@ -231,8 +320,35 @@ class MistralLoader:
log.error(f'Failed to get signed URL: {e}') log.error(f'Failed to get signed URL: {e}')
raise raise
async def _get_signed_url_async(self, session: aiohttp.ClientSession, file_id: str) -> str:
"""Async signed URL retrieval."""
url = f'{self.base_url}/files/{file_id}/url'
params = {'expiry': 1}
headers = {**self.headers, 'Accept': 'application/json'}
async def url_request():
self._debug_log(f'Getting signed URL for file ID: {file_id}')
async with session.get(
url,
headers=headers,
params=params,
timeout=aiohttp.ClientTimeout(total=self.url_timeout),
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as response:
return await self._handle_response_async(response)
response_data = await self._retry_request_async(url_request)
signed_url = response_data.get('url')
if not signed_url:
raise ValueError('Signed URL not found in response.')
self._debug_log('Signed URL received successfully')
return signed_url
def _process_ocr(self, signed_url: str) -> Dict[str, Any]: def _process_ocr(self, signed_url: str) -> Dict[str, Any]:
"""Sends the signed URL to the OCR endpoint for processing.""" """Sends the signed URL to the OCR endpoint for processing (sync version)."""
log.info('Processing OCR via Mistral API') log.info('Processing OCR via Mistral API')
url = f'{self.base_url}/ocr' url = f'{self.base_url}/ocr'
ocr_headers = { ocr_headers = {
@ -262,24 +378,108 @@ class MistralLoader:
log.error(f'Failed during OCR processing: {e}') log.error(f'Failed during OCR processing: {e}')
raise raise
def _get_file_data_url(self) -> str: async def _process_ocr_async(self, session: aiohttp.ClientSession, signed_url: str) -> Dict[str, Any]:
with open(self.file_path, 'rb') as f: """Async OCR processing with timing metrics."""
encoded_file = base64.b64encode(f.read()).decode('utf-8') url = f'{self.base_url}/ocr'
return f'data:application/pdf;base64,{encoded_file}'
headers = {
**self.headers,
'Content-Type': 'application/json',
'Accept': 'application/json',
}
payload = {
'model': 'mistral-ocr-latest',
'document': {
'type': 'document_url',
'document_url': signed_url,
},
'include_image_base64': False,
}
async def ocr_request():
log.info('Starting OCR processing via Mistral API')
start_time = time.time()
async with session.post(
url,
json=payload,
headers=headers,
timeout=aiohttp.ClientTimeout(total=self.ocr_timeout),
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as response:
ocr_response = await self._handle_response_async(response)
processing_time = time.time() - start_time
log.info(f'OCR processing completed in {processing_time:.2f}s')
return ocr_response
return await self._retry_request_async(ocr_request)
def _delete_file(self, file_id: str) -> None: def _delete_file(self, file_id: str) -> None:
"""Deletes the file from Mistral storage.""" """Deletes the file from Mistral storage (sync version)."""
log.info('Deleting uploaded file ID: %s', file_id) log.info(f'Deleting uploaded file ID: {file_id}')
url = f'{self.base_url}/files/{file_id}' url = f'{self.base_url}/files/{file_id}'
try: try:
response = requests.delete(url, headers=self.headers, timeout=self.cleanup_timeout) response = requests.delete(url, headers=self.headers, timeout=self.cleanup_timeout)
delete_response = self._handle_response(response) delete_response = self._handle_response(response)
log.info('File deleted successfully: %s', delete_response) log.info(f'File deleted successfully: {delete_response}')
except Exception as e: except Exception as e:
# Log error but don't necessarily halt execution if deletion fails # Log error but don't necessarily halt execution if deletion fails
log.error(f'Failed to delete file ID {file_id}: {e}') log.error(f'Failed to delete file ID {file_id}: {e}')
async def _delete_file_async(self, session: aiohttp.ClientSession, file_id: str) -> None:
"""Async file deletion with error tolerance."""
try:
async def delete_request():
self._debug_log(f'Deleting file ID: {file_id}')
async with session.delete(
url=f'{self.base_url}/files/{file_id}',
headers=self.headers,
timeout=aiohttp.ClientTimeout(total=self.cleanup_timeout),
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as response:
return await self._handle_response_async(response)
await self._retry_request_async(delete_request)
self._debug_log(f'File {file_id} deleted successfully')
except Exception as e:
# Don't fail the entire process if cleanup fails
log.warning(f'Failed to delete file ID {file_id}: {e}')
@asynccontextmanager
async def _get_session(self):
"""Context manager for HTTP session with optimized settings."""
connector = aiohttp.TCPConnector(
limit=20, # Increased total connection limit for better throughput
limit_per_host=10, # Increased per-host limit for API endpoints
ttl_dns_cache=600, # Longer DNS cache TTL (10 minutes)
use_dns_cache=True,
keepalive_timeout=60, # Increased keepalive for connection reuse
enable_cleanup_closed=True,
force_close=False, # Allow connection reuse
resolver=aiohttp.AsyncResolver(), # Use async DNS resolver
)
timeout = aiohttp.ClientTimeout(
total=self.timeout,
connect=30, # Connection timeout
sock_read=60, # Socket read timeout
)
async with aiohttp.ClientSession(
connector=connector,
timeout=timeout,
headers={'User-Agent': 'OpenWebUI-MistralLoader/2.0'},
raise_for_status=False, # We handle status codes manually
trust_env=True,
) as session:
yield session
def _process_results(self, ocr_response: Dict[str, Any]) -> List[Document]: def _process_results(self, ocr_response: Dict[str, Any]) -> List[Document]:
"""Process OCR results into Document objects with enhanced metadata and memory efficiency.""" """Process OCR results into Document objects with enhanced metadata and memory efficiency."""
pages_data = ocr_response.get('pages') pages_data = ocr_response.get('pages')
@ -304,7 +504,7 @@ class MistralLoader:
if page_content is None or page_index is None: if page_content is None or page_index is None:
skipped_pages += 1 skipped_pages += 1
self._debug_log( self._debug_log(
"Skipping page due to missing 'markdown' or 'index'. Data keys: %s", list(page_data.keys()) f"Skipping page due to missing 'markdown' or 'index'. Data keys: {list(page_data.keys())}"
) )
continue continue
@ -316,7 +516,7 @@ class MistralLoader:
if not cleaned_content: if not cleaned_content:
skipped_pages += 1 skipped_pages += 1
self._debug_log('Skipping empty page %s', page_index) self._debug_log(f'Skipping empty page {page_index}')
continue continue
# Create document with optimized metadata # Create document with optimized metadata
@ -336,7 +536,7 @@ class MistralLoader:
) )
if skipped_pages > 0: if skipped_pages > 0:
log.info('Processed %s pages, skipped %s empty/invalid pages', len(documents), skipped_pages) log.info(f'Processed {len(documents)} pages, skipped {skipped_pages} empty/invalid pages')
if not documents: if not documents:
# Case where pages existed but none had valid markdown/index # Case where pages existed but none had valid markdown/index
@ -357,6 +557,7 @@ class MistralLoader:
def load(self) -> List[Document]: def load(self) -> List[Document]:
""" """
Executes the full OCR workflow: upload, get URL, process OCR, delete file. Executes the full OCR workflow: upload, get URL, process OCR, delete file.
Synchronous version for backward compatibility.
Returns: Returns:
A list of Document objects, one for each page processed. A list of Document objects, one for each page processed.
@ -365,12 +566,6 @@ class MistralLoader:
start_time = time.time() start_time = time.time()
try: try:
if self.use_base64:
documents = self._process_results(self._process_ocr(self._get_file_data_url()))
total_time = time.time() - start_time
log.info('Sync OCR workflow completed in %.2fs, produced %s documents', total_time, len(documents))
return documents
# 1. Upload file # 1. Upload file
file_id = self._upload_file() file_id = self._upload_file()
@ -384,7 +579,7 @@ class MistralLoader:
documents = self._process_results(ocr_response) documents = self._process_results(ocr_response)
total_time = time.time() - start_time total_time = time.time() - start_time
log.info('Sync OCR workflow completed in %.2fs, produced %s documents', total_time, len(documents)) log.info(f'Sync OCR workflow completed in {total_time:.2f}s, produced {len(documents)} documents')
return documents return documents
@ -409,3 +604,118 @@ class MistralLoader:
except Exception as del_e: except Exception as del_e:
# Log deletion error, but don't overwrite original error if one occurred # Log deletion error, but don't overwrite original error if one occurred
log.error(f'Cleanup error: Could not delete file ID {file_id}. Reason: {del_e}') log.error(f'Cleanup error: Could not delete file ID {file_id}. Reason: {del_e}')
async def load_async(self) -> List[Document]:
"""
Asynchronous OCR workflow execution with optimized performance.
Returns:
A list of Document objects, one for each page processed.
"""
file_id = None
start_time = time.time()
try:
async with self._get_session() as session:
# 1. Upload file with streaming
file_id = await self._upload_file_async(session)
# 2. Get signed URL
signed_url = await self._get_signed_url_async(session, file_id)
# 3. Process OCR
ocr_response = await self._process_ocr_async(session, signed_url)
# 4. Process results
documents = self._process_results(ocr_response)
total_time = time.time() - start_time
log.info(f'Async OCR workflow completed in {total_time:.2f}s, produced {len(documents)} documents')
return documents
except Exception as e:
total_time = time.time() - start_time
log.error(f'Async OCR workflow failed after {total_time:.2f}s: {e}')
return [
Document(
page_content=f'Error during OCR processing: {e}',
metadata={
'error': 'processing_failed',
'file_name': self.file_name,
},
)
]
finally:
# 5. Cleanup - always attempt file deletion
if file_id:
try:
async with self._get_session() as session:
await self._delete_file_async(session, file_id)
except Exception as cleanup_error:
log.error(f'Cleanup failed for file ID {file_id}: {cleanup_error}')
@staticmethod
async def load_multiple_async(
loaders: List['MistralLoader'],
max_concurrent: int = 5, # Limit concurrent requests
) -> List[List[Document]]:
"""
Process multiple files concurrently with controlled concurrency.
Args:
loaders: List of MistralLoader instances
max_concurrent: Maximum number of concurrent requests
Returns:
List of document lists, one for each loader
"""
if not loaders:
return []
log.info(f'Starting concurrent processing of {len(loaders)} files with max {max_concurrent} concurrent')
start_time = time.time()
# Use semaphore to control concurrency
semaphore = asyncio.Semaphore(max_concurrent)
async def process_with_semaphore(loader: 'MistralLoader') -> List[Document]:
async with semaphore:
return await loader.load_async()
# Process all files with controlled concurrency
tasks = [process_with_semaphore(loader) for loader in loaders]
results = await asyncio.gather(*tasks, return_exceptions=True)
# Handle any exceptions in results
processed_results = []
for i, result in enumerate(results):
if isinstance(result, Exception):
log.error(f'File {i} failed: {result}')
processed_results.append(
[
Document(
page_content=f'Error processing file: {result}',
metadata={
'error': 'batch_processing_failed',
'file_index': i,
},
)
]
)
else:
processed_results.append(result)
# MONITORING: Log comprehensive batch processing statistics
total_time = time.time() - start_time
total_docs = sum(len(docs) for docs in processed_results)
success_count = sum(1 for result in results if not isinstance(result, Exception))
failure_count = len(results) - success_count
log.info(
f'Batch processing completed in {total_time:.2f}s: '
f'{success_count} files succeeded, {failure_count} files failed, '
f'produced {total_docs} total documents'
)
return processed_results

View file

@ -11,9 +11,6 @@ from open_webui.env import GLOBAL_LOG_LEVEL
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
PADDLEOCR_VL_IMAGE_EXTENSIONS = ['png', 'jpg', 'jpeg', 'bmp', 'tiff', 'webp']
PADDLEOCR_VL_SUPPORTED_EXTENSIONS = ['pdf'] + PADDLEOCR_VL_IMAGE_EXTENSIONS
class PaddleOCRVLLoader: class PaddleOCRVLLoader:
"""Loader that uses PaddleOCR-vl API to extract text from PDF/images.""" """Loader that uses PaddleOCR-vl API to extract text from PDF/images."""
@ -35,7 +32,7 @@ class PaddleOCRVLLoader:
self.file_name = os.path.basename(file_path) self.file_name = os.path.basename(file_path)
def load(self) -> List[Document]: def load(self) -> List[Document]:
log.info('Processing with PaddleOCR-vl: %s', self.file_path) log.info(f'Processing with PaddleOCR-vl: {self.file_path}')
try: try:
with open(self.file_path, 'rb') as file: with open(self.file_path, 'rb') as file:
@ -49,7 +46,8 @@ class PaddleOCRVLLoader:
# Detect fileType based on file extension # Detect fileType based on file extension
ext = self.file_path.lower().split('.')[-1] ext = self.file_path.lower().split('.')[-1]
file_type = 1 if ext in PADDLEOCR_VL_IMAGE_EXTENSIONS else 0 image_extensions = ['png', 'jpg', 'jpeg', 'bmp', 'tiff', 'webp']
file_type = 1 if ext in image_extensions else 0
payload = { payload = {
'file': file_data, 'file': file_data,
@ -96,7 +94,7 @@ class PaddleOCRVLLoader:
) )
if skipped_pages > 0: if skipped_pages > 0:
log.info('PaddleOCR-vl: Processed %s pages, skipped %s empty pages.', len(documents), skipped_pages) log.info(f'PaddleOCR-vl: Processed {len(documents)} pages, skipped {skipped_pages} empty pages.')
if not documents: if not documents:
log.warning('No valid text content found by PaddleOCR-vl.') log.warning('No valid text content found by PaddleOCR-vl.')

View file

@ -1,101 +0,0 @@
import datetime as dt
import io
import logging
from pathlib import Path
from langchain_core.document_loaders import BaseLoader
from langchain_core.documents import Document
log = logging.getLogger(__name__)
class PDFLoader(BaseLoader):
def __init__(self, file_path, *, extract_images=False, mode='page'):
if mode not in ('single', 'page'):
raise ValueError("PDF mode must be 'single' or 'page'")
self.file_path = str(Path(file_path).expanduser())
self.extract_images = extract_images
self.mode = mode
self.ocr = None
def lazy_load(self):
from pypdf import PdfReader
with open(self.file_path, 'rb') as file:
reader = PdfReader(file)
metadata = {'producer': 'PyPDF', 'creator': 'PyPDF', 'creationdate': ''}
for key, value in (reader.metadata or {}).items():
key = key.removeprefix('/').lower()
value = value if type(value) in (str, int) else str(value)
if key in ('creationdate', 'moddate') and isinstance(value, str):
try:
value = dt.datetime.strptime(value.replace("'", ''), 'D:%Y%m%d%H%M%S%z').isoformat()
except ValueError:
pass
metadata[key] = (
value.strip()
if isinstance(value, str) and key not in ('creationdate', 'moddate', 'page_count', 'file_path')
else value
)
metadata.update(source=self.file_path, total_pages=len(reader.pages))
labels = reader.page_labels if self.mode == 'page' else None
texts = []
for index, page in enumerate(reader.pages):
text = page.extract_text()
if self.extract_images:
image_text = self._extract_images(page)
if image_text:
text = self._merge_image_text(text, image_text)
text = text.strip()
if self.mode == 'page':
yield Document(page_content=text, metadata={**metadata, 'page': index, 'page_label': labels[index]})
else:
texts.append(text)
if self.mode == 'single':
yield Document(page_content='\n\f'.join(texts), metadata=metadata)
@staticmethod
def _merge_image_text(text, image_text):
# Insert before the final paragraphs/footer where possible, matching existing chunks.
position, separator = len(text), '\n\n'
for _ in range(2):
for delimiter in ('\n\n\n', '\n\n'):
found = text.rfind(delimiter, 0, position)
if found >= 0:
position, separator = found, delimiter
break
else:
break
return text[:position] + separator + image_text + text[position:]
def _extract_images(self, page):
import numpy as np
from PIL import Image, UnidentifiedImageError
if '/Resources' not in page or '/XObject' not in page['/Resources']:
return ''
texts = []
xobjects = page['/Resources']['/XObject']
for name in xobjects:
try:
stream = xobjects[name]
if stream.get('/Subtype') != '/Image':
continue
try:
# Encoded images, including CMYK JPEGs, can go straight to Pillow.
image = Image.open(io.BytesIO(stream.get_data()))
except UnidentifiedImageError:
image = stream.decode_as_image()
pixels = np.array(image.convert('RGB'))
except Exception as e:
log.warning('Skipping unreadable PDF image %s: %s', name, e)
continue
if self.ocr is None:
from rapidocr import RapidOCR
self.ocr = RapidOCR()
result = self.ocr(pixels)
if result and result.txts:
texts.append('\n'.join(result.txts).strip())
return '\n\n' + '\n'.join(filter(None, texts)) + '\n\n' if any(texts) else ''

View file

@ -4,7 +4,6 @@ from typing import Iterator, List, Literal, Union
import requests import requests
from langchain_core.document_loaders import BaseLoader from langchain_core.document_loaders import BaseLoader
from langchain_core.documents import Document from langchain_core.documents import Document
from open_webui.env import TAVILY_API_BASE_URL
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -49,7 +48,7 @@ class TavilyLoader(BaseLoader):
self.urls = urls if isinstance(urls, list) else [urls] self.urls = urls if isinstance(urls, list) else [urls]
self.extract_depth = extract_depth self.extract_depth = extract_depth
self.continue_on_failure = continue_on_failure self.continue_on_failure = continue_on_failure
self.api_url = f'{TAVILY_API_BASE_URL}/extract' self.api_url = 'https://api.tavily.com/extract'
def lazy_load(self) -> Iterator[Document]: def lazy_load(self) -> Iterator[Document]:
"""Extract and yield documents from the URLs using Tavily Extract API.""" """Extract and yield documents from the URLs using Tavily Extract API."""

View file

@ -18,32 +18,6 @@ ALLOWED_NETLOCS = {
} }
class YoutubeTranscriptError(Exception):
"""A YouTube transcript could not be retrieved."""
def _transcript_error_message(error: Exception, video_id: str) -> str:
name = type(error).__name__
if name in {'RequestBlocked', 'IpBlocked'}:
return (
f'YouTube blocked the transcript request for {video_id} from this server. '
'This usually means the server address is rate limited or belongs to a cloud '
'provider. A proxy for these requests can be configured under Admin Settings, '
'Web Search, Youtube Proxy URL.'
)
if name == 'TranscriptsDisabled':
return f'Transcripts are disabled for the YouTube video {video_id}.'
if name == 'AgeRestricted':
return f'The YouTube video {video_id} is age restricted, so its transcript cannot be retrieved.'
if name in {'VideoUnavailable', 'VideoUnplayable', 'InvalidVideoId'}:
return f'The YouTube video {video_id} is unavailable.'
if name == 'PoTokenRequired':
return f'YouTube requires additional verification to return the transcript for {video_id}.'
return f'Could not retrieve a transcript for the YouTube video {video_id}.'
def _parse_video_id(url: str) -> Optional[str]: def _parse_video_id(url: str) -> Optional[str]:
"""Parse a YouTube URL and return the video ID if valid, otherwise None.""" """Parse a YouTube URL and return the video ID if valid, otherwise None."""
parsed_url = urlparse(url) parsed_url = urlparse(url)
@ -116,7 +90,7 @@ class YoutubeLoader:
if self.proxy_url: if self.proxy_url:
youtube_proxies = GenericProxyConfig(http_url=self.proxy_url, https_url=self.proxy_url) youtube_proxies = GenericProxyConfig(http_url=self.proxy_url, https_url=self.proxy_url)
log.debug('Using proxy URL: %s...', self.proxy_url[:14]) log.debug(f'Using proxy URL: {self.proxy_url[:14]}...')
else: else:
youtube_proxies = None youtube_proxies = None
@ -124,31 +98,31 @@ class YoutubeLoader:
try: try:
transcript_list = transcript_api.list(self.video_id) transcript_list = transcript_api.list(self.video_id)
except Exception as e: except Exception as e:
log.warning('Loading YouTube transcript failed: %s', e) log.warning(f'Loading YouTube transcript failed: {e}')
raise YoutubeTranscriptError(_transcript_error_message(e, self.video_id)) from e return []
# Try each language in order of priority # Try each language in order of priority
for lang in self.language: for lang in self.language:
try: try:
transcript = transcript_list.find_transcript([lang]) transcript = transcript_list.find_transcript([lang])
if transcript.is_generated: if transcript.is_generated:
log.debug("Found generated transcript for language '%s'", lang) log.debug(f"Found generated transcript for language '{lang}'")
try: try:
transcript = transcript_list.find_manually_created_transcript([lang]) transcript = transcript_list.find_manually_created_transcript([lang])
log.debug("Found manual transcript for language '%s'", lang) log.debug(f"Found manual transcript for language '{lang}'")
except NoTranscriptFound: except NoTranscriptFound:
log.debug("No manual transcript found for language '%s', using generated", lang) log.debug(f"No manual transcript found for language '{lang}', using generated")
pass pass
log.debug("Found transcript for language '%s'", lang) log.debug(f"Found transcript for language '{lang}'")
try: try:
transcript_pieces: List[Dict[str, Any]] = transcript.fetch() transcript_pieces: List[Dict[str, Any]] = transcript.fetch()
except ParseError: except ParseError:
log.debug("Empty or invalid transcript for language '%s'", lang) log.debug(f"Empty or invalid transcript for language '{lang}'")
continue continue
if not transcript_pieces: if not transcript_pieces:
log.debug("Empty transcript for language '%s'", lang) log.debug(f"Empty transcript for language '{lang}'")
continue continue
transcript_text = ' '.join( transcript_text = ' '.join(
@ -161,20 +135,18 @@ class YoutubeLoader:
) )
return [Document(page_content=transcript_text, metadata=self._metadata)] return [Document(page_content=transcript_text, metadata=self._metadata)]
except NoTranscriptFound: except NoTranscriptFound:
log.debug("No transcript found for language '%s'", lang) log.debug(f"No transcript found for language '{lang}'")
continue continue
except Exception as e: except Exception as e:
log.info("Error finding transcript for language '%s'", lang) log.info(f"Error finding transcript for language '{lang}'")
raise YoutubeTranscriptError(_transcript_error_message(e, self.video_id)) from e raise e
# If we get here, all languages failed # If we get here, all languages failed
languages_tried = ', '.join(self.language) languages_tried = ', '.join(self.language)
log.warning( log.warning(
f'No transcript found for any of the specified languages: {languages_tried}. Verify if the video has transcripts, add more languages if needed.' f'No transcript found for any of the specified languages: {languages_tried}. Verify if the video has transcripts, add more languages if needed.'
) )
raise YoutubeTranscriptError( raise NoTranscriptFound(self.video_id, self.language, list(transcript_list))
f'No transcript found for the YouTube video {self.video_id} in these languages: {languages_tried}.'
)
async def aload(self) -> Generator[Document, None, None]: async def aload(self) -> Generator[Document, None, None]:
"""Asynchronously load YouTube transcripts into `Document` objects.""" """Asynchronously load YouTube transcripts into `Document` objects."""

View file

@ -12,7 +12,7 @@ log = logging.getLogger(__name__)
class ColBERT(BaseReranker): class ColBERT(BaseReranker):
def __init__(self, name, **kwargs) -> None: def __init__(self, name, **kwargs) -> None:
log.info('ColBERT: Loading model %s', name) log.info('ColBERT: Loading model', name)
self.device = 'cuda' if torch.cuda.is_available() else 'cpu' self.device = 'cuda' if torch.cuda.is_available() else 'cpu'
DOCKER = kwargs.get('env') == 'docker' DOCKER = kwargs.get('env') == 'docker'

View file

@ -35,8 +35,8 @@ class ExternalReranker(BaseReranker):
} }
try: try:
log.info('ExternalReranker:predict:model %s', self.model) log.info(f'ExternalReranker:predict:model {self.model}')
log.info('ExternalReranker:predict:query %s', query) log.info(f'ExternalReranker:predict:query {query}')
headers = { headers = {
'Content-Type': 'application/json', 'Content-Type': 'application/json',

File diff suppressed because it is too large Load diff

View file

@ -15,8 +15,9 @@ transparently dispatches each call to a worker thread via
`asyncio.to_thread`. Async callers can `await ASYNC_VECTOR_DB_CLIENT.x(...)` `asyncio.to_thread`. Async callers can `await ASYNC_VECTOR_DB_CLIENT.x(...)`
in place of `VECTOR_DB_CLIENT.x(...)` and the loop stays responsive. in place of `VECTOR_DB_CLIENT.x(...)` and the loop stays responsive.
Client initialization and calls run in the worker thread. Synchronous callers The original `VECTOR_DB_CLIENT` is unchanged, so callers already running
already inside `run_in_threadpool` use `get_vector_db_client()` directly. inside `run_in_threadpool` (e.g. `save_docs_to_vector_db`) are not
affected.
Thread-safety expectations Thread-safety expectations
-------------------------- --------------------------
@ -54,7 +55,7 @@ from __future__ import annotations
import asyncio import asyncio
from typing import Dict, List, Optional, Union from typing import Dict, List, Optional, Union
from open_webui.retrieval.vector.factory import get_vector_db_client from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
from open_webui.retrieval.vector.main import ( from open_webui.retrieval.vector.main import (
GetResult, GetResult,
SearchResult, SearchResult,
@ -72,30 +73,26 @@ class AsyncVectorDBClient:
typically swallowed by surrounding ``try/except``). typically swallowed by surrounding ``try/except``).
""" """
def __init__(self, sync_client: Optional[VectorDBBase] = None) -> None: def __init__(self, sync_client: VectorDBBase) -> None:
self._sync = sync_client self._sync = sync_client
@property @property
def sync(self) -> VectorDBBase: def sync(self) -> VectorDBBase:
"""Escape hatch for code that must call the sync client directly """Escape hatch for code that must call the sync client directly
(e.g. already inside a worker thread).""" (e.g. already inside a worker thread)."""
return self._sync if self._sync is not None else get_vector_db_client() return self._sync
@property
def supports_hybrid_search(self) -> bool:
return type(self.sync).hybrid_search is not VectorDBBase.hybrid_search
async def has_collection(self, collection_name: str) -> bool: async def has_collection(self, collection_name: str) -> bool:
return await asyncio.to_thread(lambda: self.sync.has_collection(collection_name)) return await asyncio.to_thread(self._sync.has_collection, collection_name)
async def delete_collection(self, collection_name: str) -> None: async def delete_collection(self, collection_name: str) -> None:
return await asyncio.to_thread(lambda: self.sync.delete_collection(collection_name)) return await asyncio.to_thread(self._sync.delete_collection, collection_name)
async def insert(self, collection_name: str, items: List[VectorItem]) -> None: async def insert(self, collection_name: str, items: List[VectorItem]) -> None:
return await asyncio.to_thread(lambda: self.sync.insert(collection_name, items)) return await asyncio.to_thread(self._sync.insert, collection_name, items)
async def upsert(self, collection_name: str, items: List[VectorItem]) -> None: async def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
return await asyncio.to_thread(lambda: self.sync.upsert(collection_name, items)) return await asyncio.to_thread(self._sync.upsert, collection_name, items)
async def search( async def search(
self, self,
@ -104,20 +101,7 @@ class AsyncVectorDBClient:
filter: Optional[Dict] = None, filter: Optional[Dict] = None,
limit: int = 10, limit: int = 10,
) -> Optional[SearchResult]: ) -> Optional[SearchResult]:
return await asyncio.to_thread(lambda: self.sync.search(collection_name, vectors, filter, limit)) return await asyncio.to_thread(self._sync.search, collection_name, vectors, filter, limit)
async def hybrid_search(
self,
collection_name: str,
query: str,
vectors: List[List[Union[float, int]]],
filter: Optional[Dict] = None,
limit: int = 10,
hybrid_bm25_weight: float = 0.5,
) -> Optional[SearchResult]:
return await asyncio.to_thread(
lambda: self.sync.hybrid_search(collection_name, query, vectors, filter, limit, hybrid_bm25_weight)
)
async def query( async def query(
self, self,
@ -125,10 +109,10 @@ class AsyncVectorDBClient:
filter: Dict, filter: Dict,
limit: Optional[int] = None, limit: Optional[int] = None,
) -> Optional[GetResult]: ) -> Optional[GetResult]:
return await asyncio.to_thread(lambda: self.sync.query(collection_name, filter, limit)) return await asyncio.to_thread(self._sync.query, collection_name, filter, limit)
async def get(self, collection_name: str) -> Optional[GetResult]: async def get(self, collection_name: str) -> Optional[GetResult]:
return await asyncio.to_thread(lambda: self.sync.get(collection_name)) return await asyncio.to_thread(self._sync.get, collection_name)
async def delete( async def delete(
self, self,
@ -136,10 +120,10 @@ class AsyncVectorDBClient:
ids: Optional[List[str]] = None, ids: Optional[List[str]] = None,
filter: Optional[Dict] = None, filter: Optional[Dict] = None,
) -> None: ) -> None:
return await asyncio.to_thread(lambda: self.sync.delete(collection_name, ids, filter)) return await asyncio.to_thread(self._sync.delete, collection_name, ids, filter)
async def reset(self) -> None: async def reset(self) -> None:
return await asyncio.to_thread(lambda: self.sync.reset()) return await asyncio.to_thread(self._sync.reset)
ASYNC_VECTOR_DB_CLIENT = AsyncVectorDBClient() ASYNC_VECTOR_DB_CLIENT = AsyncVectorDBClient(VECTOR_DB_CLIENT)

View file

@ -3,7 +3,6 @@ from typing import Optional
import chromadb import chromadb
from chromadb import Settings from chromadb import Settings
from chromadb.errors import NotFoundError
from chromadb.utils.batch_utils import create_batches from chromadb.utils.batch_utils import create_batches
from open_webui.config import ( from open_webui.config import (
CHROMA_CLIENT_AUTH_CREDENTIALS, CHROMA_CLIENT_AUTH_CREDENTIALS,
@ -16,8 +15,6 @@ from open_webui.config import (
CHROMA_HTTP_SSL, CHROMA_HTTP_SSL,
CHROMA_TENANT, CHROMA_TENANT,
) )
from open_webui.env import USE_SLIM
from fastapi import HTTPException
from open_webui.retrieval.vector.main import ( from open_webui.retrieval.vector.main import (
GetResult, GetResult,
SearchResult, SearchResult,
@ -31,8 +28,6 @@ log = logging.getLogger(__name__)
class ChromaClient(VectorDBBase): class ChromaClient(VectorDBBase):
def __init__(self): def __init__(self):
if USE_SLIM and not CHROMA_HTTP_HOST:
raise HTTPException(503, 'Configure CHROMA_HTTP_HOST: embedded Chroma is unavailable in slim.')
settings_dict = { settings_dict = {
'allow_reset': True, 'allow_reset': True,
'anonymized_telemetry': False, 'anonymized_telemetry': False,
@ -61,11 +56,9 @@ class ChromaClient(VectorDBBase):
) )
def has_collection(self, collection_name: str) -> bool: def has_collection(self, collection_name: str) -> bool:
try: # Check if the collection exists based on the collection name.
self.client.get_collection(name=collection_name, embedding_function=None) collection_names = self.client.list_collections()
return True return collection_name in collection_names
except NotFoundError:
return False
def delete_collection(self, collection_name: str): def delete_collection(self, collection_name: str):
# Delete the collection based on the collection name. # Delete the collection based on the collection name.
@ -80,7 +73,7 @@ class ChromaClient(VectorDBBase):
) -> Optional[SearchResult]: ) -> Optional[SearchResult]:
# Search for the nearest neighbor items based on the vectors and return 'limit' number of results. # Search for the nearest neighbor items based on the vectors and return 'limit' number of results.
try: try:
collection = self.client.get_collection(name=collection_name, embedding_function=None) collection = self.client.get_collection(name=collection_name)
if collection: if collection:
result = collection.query( result = collection.query(
query_embeddings=vectors, query_embeddings=vectors,
@ -109,7 +102,7 @@ class ChromaClient(VectorDBBase):
def query(self, collection_name: str, filter: dict, limit: Optional[int] = None) -> Optional[GetResult]: def query(self, collection_name: str, filter: dict, limit: Optional[int] = None) -> Optional[GetResult]:
# Query the items from the collection based on the filter. # Query the items from the collection based on the filter.
try: try:
collection = self.client.get_collection(name=collection_name, embedding_function=None) collection = self.client.get_collection(name=collection_name)
if collection: if collection:
result = collection.get( result = collection.get(
where=filter, where=filter,
@ -129,7 +122,7 @@ class ChromaClient(VectorDBBase):
def get(self, collection_name: str) -> Optional[GetResult]: def get(self, collection_name: str) -> Optional[GetResult]:
# Get all the items in the collection. # Get all the items in the collection.
collection = self.client.get_collection(name=collection_name, embedding_function=None) collection = self.client.get_collection(name=collection_name)
if collection: if collection:
result = collection.get() result = collection.get()
return GetResult( return GetResult(
@ -143,9 +136,7 @@ class ChromaClient(VectorDBBase):
def insert(self, collection_name: str, items: list[VectorItem]): def insert(self, collection_name: str, items: list[VectorItem]):
# Insert the items into the collection, if the collection does not exist, it will be created. # Insert the items into the collection, if the collection does not exist, it will be created.
collection = self.client.get_or_create_collection( collection = self.client.get_or_create_collection(name=collection_name, metadata={'hnsw:space': 'cosine'})
name=collection_name, metadata={'hnsw:space': 'cosine'}, embedding_function=None
)
ids = [item['id'] for item in items] ids = [item['id'] for item in items]
documents = [item['text'] for item in items] documents = [item['text'] for item in items]
@ -163,9 +154,7 @@ class ChromaClient(VectorDBBase):
def upsert(self, collection_name: str, items: list[VectorItem]): def upsert(self, collection_name: str, items: list[VectorItem]):
# Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created. # Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
collection = self.client.get_or_create_collection( collection = self.client.get_or_create_collection(name=collection_name, metadata={'hnsw:space': 'cosine'})
name=collection_name, metadata={'hnsw:space': 'cosine'}, embedding_function=None
)
ids = [item['id'] for item in items] ids = [item['id'] for item in items]
documents = [item['text'] for item in items] documents = [item['text'] for item in items]
@ -182,7 +171,7 @@ class ChromaClient(VectorDBBase):
): ):
# Delete the items from the collection based on the ids. # Delete the items from the collection based on the ids.
try: try:
collection = self.client.get_collection(name=collection_name, embedding_function=None) collection = self.client.get_collection(name=collection_name)
if collection: if collection:
if ids: if ids:
collection.delete(ids=ids) collection.delete(ids=ids)
@ -190,7 +179,7 @@ class ChromaClient(VectorDBBase):
collection.delete(where=filter) collection.delete(where=filter)
except Exception as e: except Exception as e:
# If collection doesn't exist, that's fine - nothing to delete # If collection doesn't exist, that's fine - nothing to delete
log.debug('Attempted to delete from non-existent collection %s. Ignoring.', collection_name) log.debug(f'Attempted to delete from non-existent collection {collection_name}. Ignoring.')
pass pass
def reset(self): def reset(self):

View file

@ -3,7 +3,7 @@ NOTE: This vector database integration is community-supported and maintained on
""" """
import ssl import ssl
from typing import Any, Optional from typing import Optional
from elasticsearch import BadRequestError, Elasticsearch from elasticsearch import BadRequestError, Elasticsearch
from elasticsearch.helpers import bulk, scan from elasticsearch.helpers import bulk, scan
@ -23,13 +23,7 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata from open_webui.retrieval.vector.utils import process_metadata
def _metadata_filter(key: str, op: str, value: Any) -> dict:
if op == '$in':
return {'terms': {f'metadata.{key}': value}}
return {'term': {f'metadata.{key}': value}}
class ElasticsearchClient(VectorDBBase): class ElasticsearchClient(VectorDBBase):
@ -167,16 +161,12 @@ class ElasticsearchClient(VectorDBBase):
filter: Optional[dict] = None, filter: Optional[dict] = None,
limit: int = 10, limit: int = 10,
) -> Optional[SearchResult]: ) -> Optional[SearchResult]:
filters = [{'term': {'collection': collection_name}}]
if filter:
filters.extend(_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter))
query = { query = {
'size': limit, 'size': limit,
'_source': ['text', 'metadata'], '_source': ['text', 'metadata'],
'query': { 'query': {
'script_score': { 'script_score': {
'query': {'bool': {'filter': filters}}, 'query': {'bool': {'filter': [{'term': {'collection': collection_name}}]}},
'script': { 'script': {
'source': "cosineSimilarity(params.vector, 'vector') + 1.0", 'source': "cosineSimilarity(params.vector, 'vector') + 1.0",
'params': {'vector': vectors[0]}, # Assuming single query vector 'params': {'vector': vectors[0]}, # Assuming single query vector

View file

@ -5,6 +5,7 @@ NOTE: This vector database integration is community-supported and maintained on
from __future__ import annotations from __future__ import annotations
import array import array
import json
import logging import logging
import math import math
import re import re
@ -29,7 +30,6 @@ from open_webui.retrieval.vector.main import (
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import process_metadata from open_webui.retrieval.vector.utils import process_metadata
from open_webui.utils.json_codec import JSONCodec
from sqlalchemy import create_engine from sqlalchemy import create_engine
from sqlalchemy.pool import NullPool, QueuePool from sqlalchemy.pool import NullPool, QueuePool
@ -72,7 +72,7 @@ def _safe_json(v: Any) -> Dict[str, Any]:
return {} return {}
if isinstance(v, str): if isinstance(v, str):
try: try:
j = JSONCodec.loads(v) j = json.loads(v)
return j if isinstance(j, dict) else {} return j if isinstance(j, dict) else {}
except Exception: except Exception:
return {} return {}
@ -324,7 +324,7 @@ class MariaDBVectorClient(VectorDBBase):
emb, emb,
collection_name, collection_name,
item.get('text'), item.get('text'),
JSONCodec.dumps(meta), json.dumps(meta),
) )
) )
cur.executemany(sql, params) cur.executemany(sql, params)
@ -367,7 +367,7 @@ class MariaDBVectorClient(VectorDBBase):
emb, emb,
collection_name, collection_name,
item.get('text'), item.get('text'),
JSONCodec.dumps(meta), json.dumps(meta),
) )
) )
cur.executemany(sql, params) cur.executemany(sql, params)

View file

@ -2,9 +2,9 @@
NOTE: This vector database integration is community-supported and maintained on a best-effort basis. NOTE: This vector database integration is community-supported and maintained on a best-effort basis.
""" """
import json
import logging import logging
import re from typing import Optional
from typing import Any, Optional
from open_webui.config import ( from open_webui.config import (
MILVUS_DB, MILVUS_DB,
@ -24,49 +24,12 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata from open_webui.retrieval.vector.utils import process_metadata
from open_webui.utils.json_codec import JSONCodec from pymilvus import Collection, DataType, FieldSchema, connections
from pymilvus import DataType
from pymilvus import MilvusClient as Client from pymilvus import MilvusClient as Client
from pymilvus.exceptions import MilvusException
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
# Milvus caps stored text length (here the chunk lives under the JSON `data`
# field). Clamp long chunks before insert so one oversized chunk can't fail the
# whole batch and leave the file with zero embeddings.
MILVUS_TEXT_MAX_LENGTH = 65535
_SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$')
def _escape_milvus_string(value: str) -> str:
if not isinstance(value, str):
raise TypeError(f'Expected str, got {type(value).__name__}')
return value.replace('\\', '\\\\').replace("'", "\\'")
def _milvus_literal(value: Any) -> str:
if isinstance(value, str):
return f"'{_escape_milvus_string(value)}'"
if isinstance(value, bool):
return str(value).lower()
if isinstance(value, (int, float)):
return str(value)
raise TypeError(f'Unsupported Milvus filter value type: {type(value).__name__}')
def _metadata_exprs(filter: Optional[dict]) -> list[str]:
exprs = []
for key, op, value in iter_filter_conditions(filter):
if not isinstance(key, str) or not _SAFE_METADATA_KEY_RE.fullmatch(key):
raise ValueError(f'Invalid Milvus metadata filter key: {key!r}')
if op == '$in':
items = [f"metadata['{key}'] == {_milvus_literal(item)}" for item in value]
exprs.append(f'({" or ".join(items)})' if items else 'false')
else:
exprs.append(f"metadata['{key}'] == {_milvus_literal(value)}")
return exprs
class MilvusClient(VectorDBBase): class MilvusClient(VectorDBBase):
def __init__(self): def __init__(self):
@ -156,7 +119,7 @@ class MilvusClient(VectorDBBase):
index_type = MILVUS_INDEX_TYPE.upper() index_type = MILVUS_INDEX_TYPE.upper()
metric_type = MILVUS_METRIC_TYPE.upper() metric_type = MILVUS_METRIC_TYPE.upper()
log.info('Using Milvus index type: %s, metric type: %s', index_type, metric_type) log.info(f'Using Milvus index type: {index_type}, metric type: {metric_type}')
index_creation_params = {} index_creation_params = {}
if index_type == 'HNSW': if index_type == 'HNSW':
@ -164,18 +127,18 @@ class MilvusClient(VectorDBBase):
'M': MILVUS_HNSW_M, 'M': MILVUS_HNSW_M,
'efConstruction': MILVUS_HNSW_EFCONSTRUCTION, 'efConstruction': MILVUS_HNSW_EFCONSTRUCTION,
} }
log.info('HNSW params: %s', index_creation_params) log.info(f'HNSW params: {index_creation_params}')
elif index_type == 'IVF_FLAT': elif index_type == 'IVF_FLAT':
index_creation_params = {'nlist': MILVUS_IVF_FLAT_NLIST} index_creation_params = {'nlist': MILVUS_IVF_FLAT_NLIST}
log.info('IVF_FLAT params: %s', index_creation_params) log.info(f'IVF_FLAT params: {index_creation_params}')
elif index_type == 'DISKANN': elif index_type == 'DISKANN':
index_creation_params = { index_creation_params = {
'max_degree': MILVUS_DISKANN_MAX_DEGREE, 'max_degree': MILVUS_DISKANN_MAX_DEGREE,
'search_list_size': MILVUS_DISKANN_SEARCH_LIST_SIZE, 'search_list_size': MILVUS_DISKANN_SEARCH_LIST_SIZE,
} }
log.info('DISKANN params: %s', index_creation_params) log.info(f'DISKANN params: {index_creation_params}')
elif index_type in ['FLAT', 'AUTOINDEX']: elif index_type in ['FLAT', 'AUTOINDEX']:
log.info('Using %s index with no specific build-time params.', index_type) log.info(f'Using {index_type} index with no specific build-time params.')
else: else:
log.warning( log.warning(
f"Unsupported MILVUS_INDEX_TYPE: '{index_type}'. " f"Unsupported MILVUS_INDEX_TYPE: '{index_type}'. "
@ -198,11 +161,7 @@ class MilvusClient(VectorDBBase):
index_params=index_params, index_params=index_params,
) )
log.info( log.info(
"Successfully created collection '%s_%s' with index type '%s' and metric '%s'.", f"Successfully created collection '{self.collection_prefix}_{collection_name}' with index type '{index_type}' and metric '{metric_type}'."
self.collection_prefix,
collection_name,
index_type,
metric_type,
) )
def has_collection(self, collection_name: str) -> bool: def has_collection(self, collection_name: str) -> bool:
@ -224,9 +183,6 @@ class MilvusClient(VectorDBBase):
) -> Optional[SearchResult]: ) -> Optional[SearchResult]:
# Search for the nearest neighbor items based on the vectors and return 'limit' number of results. # Search for the nearest neighbor items based on the vectors and return 'limit' number of results.
collection_name = collection_name.replace('-', '_') collection_name = collection_name.replace('-', '_')
kwargs = {}
if filter:
kwargs['filter'] = ' and '.join(_metadata_exprs(filter))
# For some index types like IVF_FLAT, search params like nprobe can be set. # For some index types like IVF_FLAT, search params like nprobe can be set.
# Example: search_params = {"nprobe": 10} if using IVF_FLAT # Example: search_params = {"nprobe": 10} if using IVF_FLAT
# For simplicity, not adding configurable search_params here, but could be extended. # For simplicity, not adding configurable search_params here, but could be extended.
@ -235,12 +191,13 @@ class MilvusClient(VectorDBBase):
data=vectors, data=vectors,
limit=limit, limit=limit,
output_fields=['data', 'metadata'], output_fields=['data', 'metadata'],
**kwargs,
# search_params=search_params # Potentially add later if needed # search_params=search_params # Potentially add later if needed
) )
return self._result_to_search_result(result) return self._result_to_search_result(result)
def query(self, collection_name: str, filter: dict, limit: int = -1): def query(self, collection_name: str, filter: dict, limit: int = -1):
connections.connect(uri=MILVUS_URI, token=MILVUS_TOKEN, db_name=MILVUS_DB)
collection_name = collection_name.replace('-', '_') collection_name = collection_name.replace('-', '_')
if not self.has_collection(collection_name): if not self.has_collection(collection_name):
log.warning(f'Query attempted on non-existent collection: {self.collection_prefix}_{collection_name}') log.warning(f'Query attempted on non-existent collection: {self.collection_prefix}_{collection_name}')
@ -255,20 +212,16 @@ class MilvusClient(VectorDBBase):
filter_string = ' && '.join(filter_expressions) filter_string = ' && '.join(filter_expressions)
self.client.load_collection(collection_name=f'{self.collection_prefix}_{collection_name}') collection = Collection(f'{self.collection_prefix}_{collection_name}')
collection.load()
try: try:
log.info( log.info(
"Querying collection %s_%s with filter: '%s', limit: %s", f"Querying collection {self.collection_prefix}_{collection_name} with filter: '{filter_string}', limit: {limit}"
self.collection_prefix,
collection_name,
filter_string,
limit,
) )
iterator = self.client.query_iterator( iterator = collection.query_iterator(
collection_name=f'{self.collection_prefix}_{collection_name}', expr=filter_string,
filter=filter_string,
output_fields=[ output_fields=[
'id', 'id',
'data', 'data',
@ -285,7 +238,7 @@ class MilvusClient(VectorDBBase):
break break
all_results.extend(batch) all_results.extend(batch)
log.debug('Total results from query: %s', len(all_results)) log.debug(f'Total results from query: {len(all_results)}')
return self._result_to_get_result([all_results] if all_results else [[]]) return self._result_to_get_result([all_results] if all_results else [[]])
except Exception as e: except Exception as e:
@ -308,7 +261,7 @@ class MilvusClient(VectorDBBase):
# Insert the items into the collection, if the collection does not exist, it will be created. # Insert the items into the collection, if the collection does not exist, it will be created.
collection_name = collection_name.replace('-', '_') collection_name = collection_name.replace('-', '_')
if not self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}'): if not self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}'):
log.info('Collection %s_%s does not exist. Creating now.', self.collection_prefix, collection_name) log.info(f'Collection {self.collection_prefix}_{collection_name} does not exist. Creating now.')
if not items: if not items:
log.error( log.error(
f'Cannot create collection {self.collection_prefix}_{collection_name} without items to determine dimension.' f'Cannot create collection {self.collection_prefix}_{collection_name} without items to determine dimension.'
@ -316,37 +269,25 @@ class MilvusClient(VectorDBBase):
raise ValueError('Cannot create Milvus collection without items to determine vector dimension.') raise ValueError('Cannot create Milvus collection without items to determine vector dimension.')
self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector'])) self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector']))
log.info('Inserting %s items into collection %s_%s.', len(items), self.collection_prefix, collection_name) log.info(f'Inserting {len(items)} items into collection {self.collection_prefix}_{collection_name}.')
data = [] return self.client.insert(
for item in items: collection_name=f'{self.collection_prefix}_{collection_name}',
text = item['text'] or '' data=[
if len(text) > MILVUS_TEXT_MAX_LENGTH:
log.warning(f'Milvus: truncating text id={item["id"]} {len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars')
text = text[:MILVUS_TEXT_MAX_LENGTH]
data.append(
{ {
'id': item['id'], 'id': item['id'],
'vector': item['vector'], 'vector': item['vector'],
'data': {'text': text}, 'data': {'text': item['text']},
'metadata': process_metadata(item['metadata']), 'metadata': process_metadata(item['metadata']),
} }
) for item in items
try: ],
return self.client.insert( )
collection_name=f'{self.collection_prefix}_{collection_name}',
data=data,
)
except MilvusException as e:
log.error(f'Milvus insert failed for {self.collection_prefix}_{collection_name} ({len(items)} items): {e}')
raise
def upsert(self, collection_name: str, items: list[VectorItem]): def upsert(self, collection_name: str, items: list[VectorItem]):
# Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created. # Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
collection_name = collection_name.replace('-', '_') collection_name = collection_name.replace('-', '_')
if not self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}'): if not self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}'):
log.info( log.info(f'Collection {self.collection_prefix}_{collection_name} does not exist for upsert. Creating now.')
'Collection %s_%s does not exist for upsert. Creating now.', self.collection_prefix, collection_name
)
if not items: if not items:
log.error( log.error(
f'Cannot create collection {self.collection_prefix}_{collection_name} for upsert without items to determine dimension.' f'Cannot create collection {self.collection_prefix}_{collection_name} for upsert without items to determine dimension.'
@ -356,29 +297,19 @@ class MilvusClient(VectorDBBase):
) )
self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector'])) self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector']))
log.info('Upserting %s items into collection %s_%s.', len(items), self.collection_prefix, collection_name) log.info(f'Upserting {len(items)} items into collection {self.collection_prefix}_{collection_name}.')
data = [] return self.client.upsert(
for item in items: collection_name=f'{self.collection_prefix}_{collection_name}',
text = item['text'] or '' data=[
if len(text) > MILVUS_TEXT_MAX_LENGTH:
log.warning(f'Milvus: truncating text id={item["id"]} {len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars')
text = text[:MILVUS_TEXT_MAX_LENGTH]
data.append(
{ {
'id': item['id'], 'id': item['id'],
'vector': item['vector'], 'vector': item['vector'],
'data': {'text': text}, 'data': {'text': item['text']},
'metadata': process_metadata(item['metadata']), 'metadata': process_metadata(item['metadata']),
} }
) for item in items
try: ],
return self.client.upsert( )
collection_name=f'{self.collection_prefix}_{collection_name}',
data=data,
)
except MilvusException as e:
log.error(f'Milvus upsert failed for {self.collection_prefix}_{collection_name} ({len(items)} items): {e}')
raise
def delete( def delete(
self, self,
@ -393,20 +324,15 @@ class MilvusClient(VectorDBBase):
return None return None
if ids: if ids:
log.info('Deleting items by IDs from %s_%s. IDs: %s', self.collection_prefix, collection_name, ids) log.info(f'Deleting items by IDs from {self.collection_prefix}_{collection_name}. IDs: {ids}')
return self.client.delete( return self.client.delete(
collection_name=f'{self.collection_prefix}_{collection_name}', collection_name=f'{self.collection_prefix}_{collection_name}',
ids=ids, ids=ids,
) )
elif filter: elif filter:
filter_string = ' && '.join( filter_string = ' && '.join([f'metadata["{key}"] == {json.dumps(value)}' for key, value in filter.items()])
[f'metadata["{key}"] == {JSONCodec.dumps(value)}' for key, value in filter.items()]
)
log.info( log.info(
'Deleting items by filter from %s_%s. Filter: %s', f'Deleting items by filter from {self.collection_prefix}_{collection_name}. Filter: {filter_string}'
self.collection_prefix,
collection_name,
filter_string,
) )
return self.client.delete( return self.client.delete(
collection_name=f'{self.collection_prefix}_{collection_name}', collection_name=f'{self.collection_prefix}_{collection_name}',
@ -428,7 +354,7 @@ class MilvusClient(VectorDBBase):
try: try:
self.client.drop_collection(collection_name=collection_name_full) self.client.drop_collection(collection_name=collection_name_full)
deleted_collections.append(collection_name_full) deleted_collections.append(collection_name_full)
log.info('Deleted collection: %s', collection_name_full) log.info(f'Deleted collection: {collection_name_full}')
except Exception as e: except Exception as e:
log.error(f'Error deleting collection {collection_name_full}: {e}') log.error(f'Error deleting collection {collection_name_full}: {e}')
log.info('Milvus reset complete. Deleted collections: %s', deleted_collections) log.info(f'Milvus reset complete. Deleted collections: {deleted_collections}')

View file

@ -17,25 +17,24 @@ from open_webui.config import (
MILVUS_TOKEN, MILVUS_TOKEN,
MILVUS_URI, MILVUS_URI,
) )
from open_webui.retrieval.vector.dbs.milvus import _metadata_exprs
from open_webui.retrieval.vector.main import ( from open_webui.retrieval.vector.main import (
GetResult, GetResult,
SearchResult, SearchResult,
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import process_metadata from pymilvus import (
from pymilvus import DataType Collection,
from pymilvus import MilvusClient as Client CollectionSchema,
from pymilvus.exceptions import MilvusException DataType,
FieldSchema,
connections,
utility,
)
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
RESOURCE_ID_FIELD = 'resource_id' RESOURCE_ID_FIELD = 'resource_id'
# Milvus VARCHAR hard cap for the `text` field (see _create_shared_collection).
# Chunks longer than this are truncated before insert so one oversized chunk
# can't fail the whole batch (and leave the file with zero embeddings).
MILVUS_TEXT_MAX_LENGTH = 65535
# Milvus expressions are SQL-like strings with no parameterized-query API; # Milvus expressions are SQL-like strings with no parameterized-query API;
# values get interpolated into single-quoted literals. Reject anything that # values get interpolated into single-quoted literals. Reject anything that
@ -66,7 +65,12 @@ class MilvusClient(VectorDBBase):
def __init__(self): def __init__(self):
# Milvus collection names can only contain numbers, letters, and underscores. # Milvus collection names can only contain numbers, letters, and underscores.
self.collection_prefix = MILVUS_COLLECTION_PREFIX.replace('-', '_') self.collection_prefix = MILVUS_COLLECTION_PREFIX.replace('-', '_')
self.client = Client(uri=MILVUS_URI, token=MILVUS_TOKEN, db_name=MILVUS_DB) connections.connect(
alias='default',
uri=MILVUS_URI,
token=MILVUS_TOKEN,
db_name=MILVUS_DB,
)
# Main collection types for multi-tenancy # Main collection types for multi-tenancy
self.MEMORY_COLLECTION = f'{self.collection_prefix}_memories' self.MEMORY_COLLECTION = f'{self.collection_prefix}_memories'
@ -107,66 +111,53 @@ class MilvusClient(VectorDBBase):
return self.KNOWLEDGE_COLLECTION, resource_id return self.KNOWLEDGE_COLLECTION, resource_id
def _create_shared_collection(self, mt_collection_name: str, dimension: int): def _create_shared_collection(self, mt_collection_name: str, dimension: int):
schema = self.client.create_schema(auto_id=False, description='Shared collection for multi-tenancy') fields = [
schema.add_field(field_name='id', datatype=DataType.VARCHAR, is_primary=True, max_length=36) FieldSchema(
schema.add_field(field_name='vector', datatype=DataType.FLOAT_VECTOR, dim=dimension) name='id',
schema.add_field(field_name='text', datatype=DataType.VARCHAR, max_length=MILVUS_TEXT_MAX_LENGTH) dtype=DataType.VARCHAR,
schema.add_field(field_name='metadata', datatype=DataType.JSON) is_primary=True,
schema.add_field(field_name=RESOURCE_ID_FIELD, datatype=DataType.VARCHAR, max_length=255) auto_id=False,
max_length=36,
),
FieldSchema(name='vector', dtype=DataType.FLOAT_VECTOR, dim=dimension),
FieldSchema(name='text', dtype=DataType.VARCHAR, max_length=65535),
FieldSchema(name='metadata', dtype=DataType.JSON),
FieldSchema(name=RESOURCE_ID_FIELD, dtype=DataType.VARCHAR, max_length=255),
]
schema = CollectionSchema(fields, 'Shared collection for multi-tenancy')
collection = Collection(mt_collection_name, schema)
index_build_params = {} index_params = {
'metric_type': MILVUS_METRIC_TYPE,
'index_type': MILVUS_INDEX_TYPE,
'params': {},
}
if MILVUS_INDEX_TYPE == 'HNSW': if MILVUS_INDEX_TYPE == 'HNSW':
index_build_params = { index_params['params'] = {
'M': MILVUS_HNSW_M, 'M': MILVUS_HNSW_M,
'efConstruction': MILVUS_HNSW_EFCONSTRUCTION, 'efConstruction': MILVUS_HNSW_EFCONSTRUCTION,
} }
elif MILVUS_INDEX_TYPE == 'IVF_FLAT': elif MILVUS_INDEX_TYPE == 'IVF_FLAT':
index_build_params = {'nlist': MILVUS_IVF_FLAT_NLIST} index_params['params'] = {'nlist': MILVUS_IVF_FLAT_NLIST}
vector_index = self.client.prepare_index_params( collection.create_index('vector', index_params)
field_name='vector', collection.create_index(RESOURCE_ID_FIELD)
index_type=MILVUS_INDEX_TYPE, log.info(f'Created shared collection: {mt_collection_name}')
metric_type=MILVUS_METRIC_TYPE, return collection
params=index_build_params,
)
self.client.create_collection(collection_name=mt_collection_name, schema=schema)
self.client.create_index(collection_name=mt_collection_name, index_params=vector_index)
try:
# A Milvus server auto-selects the scalar index type from a parameterless call.
self.client.create_index(
collection_name=mt_collection_name,
index_params=self.client.prepare_index_params(field_name=RESOURCE_ID_FIELD),
)
except MilvusException:
try:
self.client.create_index(
collection_name=mt_collection_name,
index_params=self.client.prepare_index_params(field_name=RESOURCE_ID_FIELD, index_type='INVERTED'),
)
except MilvusException as e:
# The index only accelerates resource_id filters; never fail
# collection creation over it.
log.warning(f'Could not create {RESOURCE_ID_FIELD} index on {mt_collection_name}: {e}')
log.info('Created shared collection: %s', mt_collection_name)
def _ensure_collection(self, mt_collection_name: str, dimension: int): def _ensure_collection(self, mt_collection_name: str, dimension: int):
if not self.client.has_collection(mt_collection_name): if not utility.has_collection(mt_collection_name):
self._create_shared_collection(mt_collection_name, dimension) self._create_shared_collection(mt_collection_name, dimension)
def has_collection(self, collection_name: str) -> bool: def has_collection(self, collection_name: str) -> bool:
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id) _validate_resource_id(resource_id)
if not self.client.has_collection(mt_collection): if not utility.has_collection(mt_collection):
return False return False
self.client.load_collection(mt_collection) collection = Collection(mt_collection)
res = self.client.query( collection.load()
collection_name=mt_collection, res = collection.query(expr=f"{RESOURCE_ID_FIELD} == '{resource_id}'", limit=1)
filter=f"{RESOURCE_ID_FIELD} == '{resource_id}'",
output_fields=['id'],
limit=1,
)
return len(res) > 0 return len(res) > 0
def upsert(self, collection_name: str, items: List[VectorItem]): def upsert(self, collection_name: str, items: List[VectorItem]):
@ -176,35 +167,19 @@ class MilvusClient(VectorDBBase):
_validate_resource_id(resource_id) _validate_resource_id(resource_id)
dimension = len(items[0]['vector']) dimension = len(items[0]['vector'])
self._ensure_collection(mt_collection, dimension) self._ensure_collection(mt_collection, dimension)
collection = Collection(mt_collection)
entities = [] entities = [
for item in items: {
text = item['text'] or '' 'id': item['id'],
if len(text) > MILVUS_TEXT_MAX_LENGTH: 'vector': item['vector'],
log.warning( 'text': item['text'],
f'Milvus: truncating text id={item["id"]} ' 'metadata': item['metadata'],
f'{len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars ' RESOURCE_ID_FIELD: resource_id,
f'(collection={mt_collection}, resource_id={resource_id})' }
) for item in items
text = text[:MILVUS_TEXT_MAX_LENGTH] ]
entities.append( collection.insert(entities)
{
'id': item['id'],
'vector': item['vector'],
'text': text,
'metadata': process_metadata(item['metadata']),
RESOURCE_ID_FIELD: resource_id,
}
)
try:
self.client.insert(collection_name=mt_collection, data=entities)
except MilvusException as e:
log.error(
f'Milvus insert failed (collection={mt_collection}, '
f'resource_id={resource_id}, items={len(entities)}): {e}'
)
raise
def search( def search(
self, self,
@ -218,19 +193,19 @@ class MilvusClient(VectorDBBase):
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id) _validate_resource_id(resource_id)
if not self.client.has_collection(mt_collection): if not utility.has_collection(mt_collection):
return None return None
self.client.load_collection(mt_collection) collection = Collection(mt_collection)
collection.load()
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'", *_metadata_exprs(filter)] search_params = {'metric_type': MILVUS_METRIC_TYPE, 'params': {}}
results = self.client.search( results = collection.search(
collection_name=mt_collection,
data=vectors, data=vectors,
anns_field='vector', anns_field='vector',
search_params={'metric_type': MILVUS_METRIC_TYPE, 'params': {}}, param=search_params,
limit=limit, limit=limit,
filter=' and '.join(expr), expr=f"{RESOURCE_ID_FIELD} == '{resource_id}'",
output_fields=['id', 'text', 'metadata'], output_fields=['id', 'text', 'metadata'],
) )
@ -238,11 +213,10 @@ class MilvusClient(VectorDBBase):
for hits in results: for hits in results:
batch_ids, batch_docs, batch_metadatas, batch_dists = [], [], [], [] batch_ids, batch_docs, batch_metadatas, batch_dists = [], [], [], []
for hit in hits: for hit in hits:
entity = hit.get('entity', {}) batch_ids.append(hit.entity.get('id'))
batch_ids.append(entity.get('id')) batch_docs.append(hit.entity.get('text'))
batch_docs.append(entity.get('text')) batch_metadatas.append(hit.entity.get('metadata'))
batch_metadatas.append(entity.get('metadata')) batch_dists.append(hit.distance)
batch_dists.append(hit.get('distance'))
ids.append(batch_ids) ids.append(batch_ids)
documents.append(batch_docs) documents.append(batch_docs)
metadatas.append(batch_metadatas) metadatas.append(batch_metadatas)
@ -258,9 +232,11 @@ class MilvusClient(VectorDBBase):
): ):
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id) _validate_resource_id(resource_id)
if not self.client.has_collection(mt_collection): if not utility.has_collection(mt_collection):
return return
collection = Collection(mt_collection)
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"] expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
if ids: if ids:
# Milvus expects a string list for 'in' operator # Milvus expects a string list for 'in' operator
@ -272,28 +248,30 @@ class MilvusClient(VectorDBBase):
_validate_metadata_key(key) _validate_metadata_key(key)
expr.append(f"metadata['{key}'] == '{_escape_milvus_string(str(value))}'") expr.append(f"metadata['{key}'] == '{_escape_milvus_string(str(value))}'")
self.client.delete(collection_name=mt_collection, filter=' and '.join(expr)) collection.delete(' and '.join(expr))
def reset(self): def reset(self):
for collection_name in self.shared_collections: for collection_name in self.shared_collections:
if self.client.has_collection(collection_name): if utility.has_collection(collection_name):
self.client.drop_collection(collection_name) utility.drop_collection(collection_name)
def delete_collection(self, collection_name: str): def delete_collection(self, collection_name: str):
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id) _validate_resource_id(resource_id)
if not self.client.has_collection(mt_collection): if not utility.has_collection(mt_collection):
return return
self.client.delete(collection_name=mt_collection, filter=f"{RESOURCE_ID_FIELD} == '{resource_id}'") collection = Collection(mt_collection)
collection.delete(f"{RESOURCE_ID_FIELD} == '{resource_id}'")
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]: 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) mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
_validate_resource_id(resource_id) _validate_resource_id(resource_id)
if not self.client.has_collection(mt_collection): if not utility.has_collection(mt_collection):
return None return None
self.client.load_collection(mt_collection) collection = Collection(mt_collection)
collection.load()
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"] expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
if filter: if filter:
@ -308,9 +286,8 @@ class MilvusClient(VectorDBBase):
else: else:
raise TypeError(f'Unsupported Milvus filter value type for key {key!r}: {type(value).__name__}') raise TypeError(f'Unsupported Milvus filter value type for key {key!r}: {type(value).__name__}')
iterator = self.client.query_iterator( iterator = collection.query_iterator(
collection_name=mt_collection, expr=' and '.join(expr),
filter=' and '.join(expr),
output_fields=['id', 'text', 'metadata'], output_fields=['id', 'text', 'metadata'],
limit=limit if limit else -1, limit=limit if limit else -1,
) )

View file

@ -2,6 +2,7 @@
NOTE: This vector database integration is community-supported and maintained on a best-effort basis. NOTE: This vector database integration is community-supported and maintained on a best-effort basis.
""" """
import json
import logging import logging
import re import re
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
@ -62,18 +63,20 @@ from open_webui.config import (
OPENGAUSS_POOL_SIZE, OPENGAUSS_POOL_SIZE,
OPENGAUSS_POOL_TIMEOUT, OPENGAUSS_POOL_TIMEOUT,
) )
from open_webui.env import SRC_LOG_LEVELS
from open_webui.retrieval.vector.main import ( from open_webui.retrieval.vector.main import (
GetResult, GetResult,
SearchResult, SearchResult,
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata from open_webui.retrieval.vector.utils import process_metadata
VECTOR_LENGTH = OPENGAUSS_INITIALIZE_MAX_VECTOR_LENGTH VECTOR_LENGTH = OPENGAUSS_INITIALIZE_MAX_VECTOR_LENGTH
Base = declarative_base() Base = declarative_base()
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
log.setLevel(SRC_LOG_LEVELS['RAG'])
class DocumentChunk(Base): class DocumentChunk(Base):
@ -86,12 +89,6 @@ class DocumentChunk(Base):
vmetadata = Column(MutableDict.as_mutable(JSONB), nullable=True) vmetadata = Column(MutableDict.as_mutable(JSONB), nullable=True)
def _metadata_clause(key: str, op: str, value: Any):
if op == '$in':
return DocumentChunk.vmetadata[key].astext.in_([str(v) for v in value])
return DocumentChunk.vmetadata[key].astext == str(value)
class OpenGaussClient(VectorDBBase): class OpenGaussClient(VectorDBBase):
def __init__(self) -> None: def __init__(self) -> None:
if not OPENGAUSS_DB_URL: if not OPENGAUSS_DB_URL:
@ -185,7 +182,7 @@ class OpenGaussClient(VectorDBBase):
new_items.append(new_chunk) new_items.append(new_chunk)
self.session.bulk_save_objects(new_items) self.session.bulk_save_objects(new_items)
self.session.commit() self.session.commit()
log.info("Inserting %s items into collection '%s'.", len(new_items), collection_name) log.info(f"Inserting {len(new_items)} items into collection '{collection_name}'.")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Failed to insert data: {e}') log.exception(f'Failed to insert data: {e}')
@ -211,7 +208,7 @@ class OpenGaussClient(VectorDBBase):
) )
self.session.add(new_chunk) self.session.add(new_chunk)
self.session.commit() self.session.commit()
log.info("Inserting/updating %s items in collection '%s'.", len(items), collection_name) log.info(f"Inserting/updating {len(items)} items in collection '{collection_name}'.")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Failed to insert or update data.: {e}') log.exception(f'Failed to insert or update data.: {e}')
@ -248,15 +245,10 @@ class OpenGaussClient(VectorDBBase):
DocumentChunk.vmetadata, DocumentChunk.vmetadata,
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label('distance'), (DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label('distance'),
] ]
where_clauses = [DocumentChunk.collection_name == collection_name]
if filter:
where_clauses.extend(
_metadata_clause(key, op, value) for key, op, value in iter_filter_conditions(filter)
)
subq = ( subq = (
select(*result_fields) select(*result_fields)
.where(*where_clauses) .where(DocumentChunk.collection_name == collection_name)
.order_by(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)) .order_by(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector))
) )
if limit is not None: if limit is not None:
@ -311,7 +303,6 @@ class OpenGaussClient(VectorDBBase):
results = query.all() results = query.all()
if not results: if not results:
self.session.rollback()
return None return None
ids = [[result.id for result in results]] ids = [[result.id for result in results]]
@ -334,7 +325,6 @@ class OpenGaussClient(VectorDBBase):
results = query.all() results = query.all()
if not results: if not results:
self.session.rollback()
return None return None
ids = [[result.id for result in results]] ids = [[result.id for result in results]]
@ -363,7 +353,7 @@ class OpenGaussClient(VectorDBBase):
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value)) query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
deleted = query.delete(synchronize_session=False) deleted = query.delete(synchronize_session=False)
self.session.commit() self.session.commit()
log.info("Deleted %s items from collection '%s'", deleted, collection_name) log.info(f"Deleted {deleted} items from collection '{collection_name}'")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Failed to delete data: {e}') log.exception(f'Failed to delete data: {e}')
@ -373,7 +363,7 @@ class OpenGaussClient(VectorDBBase):
try: try:
deleted = self.session.query(DocumentChunk).delete() deleted = self.session.query(DocumentChunk).delete()
self.session.commit() self.session.commit()
log.info('Reset completed. Deleted %s items', deleted) log.info(f'Reset completed. Deleted {deleted} items')
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Reset failed: {e}') log.exception(f'Reset failed: {e}')
@ -397,4 +387,4 @@ class OpenGaussClient(VectorDBBase):
def delete_collection(self, collection_name: str) -> None: def delete_collection(self, collection_name: str) -> None:
self.delete(collection_name) self.delete(collection_name)
log.info("Collection '%s' has been deleted", collection_name) log.info(f"Collection '{collection_name}' has been deleted")

View file

@ -2,7 +2,7 @@
NOTE: This vector database integration is community-supported and maintained on a best-effort basis. NOTE: This vector database integration is community-supported and maintained on a best-effort basis.
""" """
from typing import Any, Optional from typing import Optional
from open_webui.config import ( from open_webui.config import (
OPENSEARCH_CERT_VERIFY, OPENSEARCH_CERT_VERIFY,
@ -17,17 +17,11 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata from open_webui.retrieval.vector.utils import process_metadata
from opensearchpy import OpenSearch from opensearchpy import OpenSearch
from opensearchpy.helpers import bulk from opensearchpy.helpers import bulk
def _metadata_filter(key: str, op: str, value: Any) -> dict:
if op == '$in':
return {'terms': {f'metadata.{key}.keyword': value}}
return {'term': {f'metadata.{key}.keyword': value}}
class OpenSearchClient(VectorDBBase): class OpenSearchClient(VectorDBBase):
def __init__(self): def __init__(self):
self.index_prefix = 'open_webui' self.index_prefix = 'open_webui'
@ -127,8 +121,6 @@ class OpenSearchClient(VectorDBBase):
filter: Optional[dict] = None, filter: Optional[dict] = None,
limit: int = 10, limit: int = 10,
) -> Optional[SearchResult]: ) -> Optional[SearchResult]:
filter_clauses = [_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter)]
try: try:
if not self.has_collection(collection_name): if not self.has_collection(collection_name):
return None return None
@ -138,7 +130,7 @@ class OpenSearchClient(VectorDBBase):
'_source': ['text', 'metadata'], '_source': ['text', 'metadata'],
'query': { 'query': {
'script_score': { 'script_score': {
'query': {'bool': {'filter': filter_clauses}} if filter_clauses else {'match_all': {}}, 'query': {'match_all': {}},
'script': { 'script': {
'source': '(cosineSimilarity(params.query_value, doc[params.field]) + 1.0) / 2.0', 'source': '(cosineSimilarity(params.query_value, doc[params.field]) + 1.0) / 2.0',
'params': { 'params': {

View file

@ -32,7 +32,6 @@ import array
import json import json
import logging import logging
import os import os
import re
import threading import threading
import time import time
from decimal import Decimal from decimal import Decimal
@ -57,29 +56,8 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
from open_webui.utils.json_codec import JSONCodec
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
_SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$')
def _metadata_where(filter: Optional[dict]) -> tuple[str, dict[str, Any]]:
clause = ''
params: dict[str, Any] = {}
for i, (key, op, value) in enumerate(iter_filter_conditions(filter)):
if not isinstance(key, str) or not _SAFE_METADATA_KEY_RE.fullmatch(key):
raise ValueError(f'Invalid Oracle metadata filter key: {key!r}')
json_value = f"JSON_VALUE(dc.vmetadata, '$.{key}' RETURNING VARCHAR2(4096))"
if op == '$in':
names = [f'value_{i}_{j}' for j, _ in enumerate(value)]
clause += f' AND {json_value} IN ({", ".join(f":{name}" for name in names)})' if names else ' AND 1 = 0'
params.update({name: str(item) for name, item in zip(names, value)})
else:
name = f'value_{i}'
clause += f' AND {json_value} = :{name}'
params[name] = str(value)
return clause, params
class Oracle23aiClient(VectorDBBase): class Oracle23aiClient(VectorDBBase):
@ -115,10 +93,10 @@ class Oracle23aiClient(VectorDBBase):
self._create_dbcs_pool() self._create_dbcs_pool()
dsn = ORACLE_DB_DSN dsn = ORACLE_DB_DSN
log.info('Creating Connection Pool [%s:**@%s]', ORACLE_DB_USER, dsn) log.info(f'Creating Connection Pool [{ORACLE_DB_USER}:**@{dsn}]')
with self.get_connection() as connection: with self.get_connection() as connection:
log.info('Connection version: %s', connection.version) log.info(f'Connection version: {connection.version}')
self._initialize_database(connection) self._initialize_database(connection)
log.info('Oracle Vector Search initialization complete.') log.info('Oracle Vector Search initialization complete.')
@ -180,7 +158,7 @@ class Oracle23aiClient(VectorDBBase):
if attempt < max_retries - 1: if attempt < max_retries - 1:
wait_time = 2**attempt wait_time = 2**attempt
log.info('Retrying in %s seconds...', wait_time) log.info(f'Retrying in {wait_time} seconds...')
time.sleep(wait_time) time.sleep(wait_time)
else: else:
raise raise
@ -205,7 +183,7 @@ class Oracle23aiClient(VectorDBBase):
thread = threading.Thread(target=_monitor, daemon=True) thread = threading.Thread(target=_monitor, daemon=True)
thread.start() thread.start()
log.info('Started DB health monitor every %s seconds.', interval_seconds) log.info(f'Started DB health monitor every {interval_seconds} seconds.')
def _reconnect_pool(self): def _reconnect_pool(self):
""" """
@ -400,7 +378,7 @@ class Oracle23aiClient(VectorDBBase):
Returns: Returns:
str: JSON representation of metadata str: JSON representation of metadata
""" """
return json.dumps(process_metadata(metadata), default=self._decimal_handler) if metadata else '{}' return json.dumps(metadata, default=self._decimal_handler) if metadata else '{}'
def _json_to_metadata(self, json_str: str) -> Dict: def _json_to_metadata(self, json_str: str) -> Dict:
""" """
@ -412,7 +390,7 @@ class Oracle23aiClient(VectorDBBase):
Returns: Returns:
Dict: Metadata dictionary Dict: Metadata dictionary
""" """
return JSONCodec.loads(json_str) if json_str else {} return json.loads(json_str) if json_str else {}
def insert(self, collection_name: str, items: List[VectorItem]) -> None: def insert(self, collection_name: str, items: List[VectorItem]) -> None:
""" """
@ -433,7 +411,7 @@ class Oracle23aiClient(VectorDBBase):
... ] ... ]
>>> client.insert("my_collection", items) >>> client.insert("my_collection", items)
""" """
log.info("Inserting %s items into collection '%s'.", len(items), collection_name) log.info(f"Inserting {len(items)} items into collection '{collection_name}'.")
with self.get_connection() as connection: with self.get_connection() as connection:
try: try:
@ -458,7 +436,7 @@ class Oracle23aiClient(VectorDBBase):
) )
connection.commit() connection.commit()
log.info("Successfully inserted %s items into collection '%s'.", len(items), collection_name) log.info(f"Successfully inserted {len(items)} items into collection '{collection_name}'.")
except Exception as e: except Exception as e:
connection.rollback() connection.rollback()
@ -487,7 +465,7 @@ class Oracle23aiClient(VectorDBBase):
... ] ... ]
>>> client.upsert("my_collection", items) >>> client.upsert("my_collection", items)
""" """
log.info("Upserting %s items into collection '%s'.", len(items), collection_name) log.info(f"Upserting {len(items)} items into collection '{collection_name}'.")
with self.get_connection() as connection: with self.get_connection() as connection:
try: try:
@ -526,7 +504,7 @@ class Oracle23aiClient(VectorDBBase):
) )
connection.commit() connection.commit()
log.info("Successfully upserted %s items into collection '%s'.", len(items), collection_name) log.info(f"Successfully upserted {len(items)} items into collection '{collection_name}'.")
except Exception as e: except Exception as e:
connection.rollback() connection.rollback()
@ -562,7 +540,7 @@ class Oracle23aiClient(VectorDBBase):
... for i, (id, dist) in enumerate(zip(results.ids[0], results.distances[0])): ... for i, (id, dist) in enumerate(zip(results.ids[0], results.distances[0])):
... log.info(f"Match {i+1}: id={id}, distance={dist}") ... log.info(f"Match {i+1}: id={id}, distance={dist}")
""" """
log.info("Searching items from collection '%s' with limit %s.", collection_name, limit) log.info(f"Searching items from collection '{collection_name}' with limit {limit}.")
try: try:
if not vectors: if not vectors:
@ -570,7 +548,6 @@ class Oracle23aiClient(VectorDBBase):
return None return None
num_queries = len(vectors) num_queries = len(vectors)
filter_clause, filter_params = _metadata_where(filter)
ids = [[] for _ in range(num_queries)] ids = [[] for _ in range(num_queries)]
distances = [[] for _ in range(num_queries)] distances = [[] for _ in range(num_queries)]
@ -583,12 +560,12 @@ class Oracle23aiClient(VectorDBBase):
vector_blob = self._vector_to_blob(vector) vector_blob = self._vector_to_blob(vector)
cursor.execute( cursor.execute(
f""" """
SELECT dc.id, dc.text, SELECT dc.id, dc.text,
JSON_SERIALIZE(dc.vmetadata RETURNING VARCHAR2(4096)) as vmetadata, JSON_SERIALIZE(dc.vmetadata RETURNING VARCHAR2(4096)) as vmetadata,
VECTOR_DISTANCE(dc.vector, :query_vector, COSINE) as distance VECTOR_DISTANCE(dc.vector, :query_vector, COSINE) as distance
FROM document_chunk dc FROM document_chunk dc
WHERE dc.collection_name = :collection_name{filter_clause} WHERE dc.collection_name = :collection_name
ORDER BY VECTOR_DISTANCE(dc.vector, :query_vector, COSINE) ORDER BY VECTOR_DISTANCE(dc.vector, :query_vector, COSINE)
FETCH APPROX FIRST :limit ROWS ONLY FETCH APPROX FIRST :limit ROWS ONLY
""", """,
@ -596,7 +573,6 @@ class Oracle23aiClient(VectorDBBase):
'query_vector': vector_blob, 'query_vector': vector_blob,
'collection_name': collection_name, 'collection_name': collection_name,
'limit': limit, 'limit': limit,
**filter_params,
}, },
) )
@ -610,7 +586,7 @@ class Oracle23aiClient(VectorDBBase):
metadatas[qid].append(self._json_to_metadata(metadata_str)) metadatas[qid].append(self._json_to_metadata(metadata_str))
distances[qid].append(float(row[3])) distances[qid].append(float(row[3]))
log.info('Search completed. Found %s total results.', sum(len(ids[i]) for i in range(num_queries))) log.info(f'Search completed. Found {sum(len(ids[i]) for i in range(num_queries))} total results.')
return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas) return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas)
@ -639,7 +615,7 @@ class Oracle23aiClient(VectorDBBase):
>>> if results: >>> if results:
... print(f"Found {len(results.ids[0])} matching documents") ... print(f"Found {len(results.ids[0])} matching documents")
""" """
log.info("Querying items from collection '%s' with filters.", collection_name) log.info(f"Querying items from collection '{collection_name}' with filters.")
try: try:
limit = limit or 100 limit = limit or 100
@ -679,7 +655,7 @@ class Oracle23aiClient(VectorDBBase):
] ]
] ]
log.info('Query completed. Found %s results.', len(results)) log.info(f'Query completed. Found {len(results)} results.')
return GetResult(ids=ids, documents=documents, metadatas=metadatas) return GetResult(ids=ids, documents=documents, metadatas=metadatas)
@ -770,7 +746,7 @@ class Oracle23aiClient(VectorDBBase):
>>> # Or delete by metadata filter >>> # Or delete by metadata filter
>>> client.delete("my_collection", filter={"source": "deprecated_source"}) >>> client.delete("my_collection", filter={"source": "deprecated_source"})
""" """
log.info("Deleting items from collection '%s'.", collection_name) log.info(f"Deleting items from collection '{collection_name}'.")
try: try:
query = 'DELETE FROM document_chunk WHERE collection_name = :collection_name' query = 'DELETE FROM document_chunk WHERE collection_name = :collection_name'
@ -795,7 +771,7 @@ class Oracle23aiClient(VectorDBBase):
deleted = cursor.rowcount deleted = cursor.rowcount
connection.commit() connection.commit()
log.info("Deleted %s items from collection '%s'.", deleted, collection_name) log.info(f"Deleted {deleted} items from collection '{collection_name}'.")
except Exception as e: except Exception as e:
log.exception(f'Error during delete: {e}') log.exception(f'Error during delete: {e}')
@ -823,7 +799,7 @@ class Oracle23aiClient(VectorDBBase):
deleted = cursor.rowcount deleted = cursor.rowcount
connection.commit() connection.commit()
log.info("Reset complete. Deleted %s items from 'document_chunk' table.", deleted) log.info(f"Reset complete. Deleted {deleted} items from 'document_chunk' table.")
except Exception as e: except Exception as e:
log.exception(f'Error during reset: {e}') log.exception(f'Error during reset: {e}')
@ -898,7 +874,7 @@ class Oracle23aiClient(VectorDBBase):
>>> client = Oracle23aiClient() >>> client = Oracle23aiClient()
>>> client.delete_collection("obsolete_collection") >>> client.delete_collection("obsolete_collection")
""" """
log.info("Deleting collection '%s'.", collection_name) log.info(f"Deleting collection '{collection_name}'.")
try: try:
with self.get_connection() as connection: with self.get_connection() as connection:
@ -914,7 +890,7 @@ class Oracle23aiClient(VectorDBBase):
deleted = cursor.rowcount deleted = cursor.rowcount
connection.commit() connection.commit()
log.info("Collection '%s' deleted. Removed %s items.", collection_name, deleted) log.info(f"Collection '{collection_name}' deleted. Removed {deleted} items.")
except Exception as e: except Exception as e:
log.exception(f"Error deleting collection '{collection_name}': {e}") log.exception(f"Error deleting collection '{collection_name}': {e}")

View file

@ -1,3 +1,4 @@
import json
import logging import logging
from typing import Any, Dict, List, Optional, Tuple from typing import Any, Dict, List, Optional, Tuple
@ -8,7 +9,6 @@ from open_webui.config import (
PGVECTOR_HNSW_M, PGVECTOR_HNSW_M,
PGVECTOR_INDEX_METHOD, PGVECTOR_INDEX_METHOD,
PGVECTOR_INITIALIZE_MAX_VECTOR_LENGTH, PGVECTOR_INITIALIZE_MAX_VECTOR_LENGTH,
PGVECTOR_ITERATIVE_SCAN,
PGVECTOR_IVFFLAT_LISTS, PGVECTOR_IVFFLAT_LISTS,
PGVECTOR_PGCRYPTO, PGVECTOR_PGCRYPTO,
PGVECTOR_PGCRYPTO_KEY, PGVECTOR_PGCRYPTO_KEY,
@ -18,15 +18,13 @@ from open_webui.config import (
PGVECTOR_POOL_TIMEOUT, PGVECTOR_POOL_TIMEOUT,
PGVECTOR_USE_HALFVEC, PGVECTOR_USE_HALFVEC,
) )
from open_webui.internal.db import ScopedSession, enable_iam_token_auth
from open_webui.retrieval.vector.main import ( from open_webui.retrieval.vector.main import (
GetResult, GetResult,
SearchResult, SearchResult,
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import merge_hybrid_search_results, process_metadata from open_webui.retrieval.vector.utils import process_metadata
from open_webui.utils.json_codec import JSONCodec
from open_webui.utils.misc import sanitize_text_for_db from open_webui.utils.misc import sanitize_text_for_db
from pgvector.sqlalchemy import HALFVEC, Vector from pgvector.sqlalchemy import HALFVEC, Vector
from sqlalchemy import ( from sqlalchemy import (
@ -89,6 +87,8 @@ class PgvectorClient(VectorDBBase):
def __init__(self) -> None: def __init__(self) -> None:
# if no pgvector uri, use the existing database connection # if no pgvector uri, use the existing database connection
if not PGVECTOR_DB_URL: if not PGVECTOR_DB_URL:
from open_webui.internal.db import ScopedSession
self.session = ScopedSession self.session = ScopedSession
else: else:
if isinstance(PGVECTOR_POOL_SIZE, int): if isinstance(PGVECTOR_POOL_SIZE, int):
@ -107,7 +107,6 @@ class PgvectorClient(VectorDBBase):
else: else:
engine = create_engine(PGVECTOR_DB_URL, pool_pre_ping=True) engine = create_engine(PGVECTOR_DB_URL, pool_pre_ping=True)
enable_iam_token_auth(engine)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
self.session = scoped_session(SessionLocal) self.session = scoped_session(SessionLocal)
@ -154,8 +153,6 @@ class PgvectorClient(VectorDBBase):
index_method, index_options = self._vector_index_configuration() index_method, index_options = self._vector_index_configuration()
self._ensure_vector_index(index_method, index_options) self._ensure_vector_index(index_method, index_options)
self._ensure_text_search_index()
self.iterative_scan_sql = self._iterative_scan_setting(index_method)
self.session.execute( self.session.execute(
text( text(
@ -225,9 +222,6 @@ class PgvectorClient(VectorDBBase):
) )
if not existing_index_def: if not existing_index_def:
if index_method == 'ivfflat' and not self._has_enough_ivfflat_training_rows():
return
index_sql = ( index_sql = (
f'CREATE INDEX IF NOT EXISTS {index_name} ' f'CREATE INDEX IF NOT EXISTS {index_name} '
f'ON document_chunk USING {index_method} (vector {VECTOR_OPCLASS})' f'ON document_chunk USING {index_method} (vector {VECTOR_OPCLASS})'
@ -242,51 +236,6 @@ class PgvectorClient(VectorDBBase):
f' {index_options}' if index_options else '', f' {index_options}' if index_options else '',
) )
def _has_enough_ivfflat_training_rows(self) -> bool:
# ivfflat samples 50 rows per list to place its centroids, so recall stays poor until the table holds that many
min_training_rows = 50 * PGVECTOR_IVFFLAT_LISTS
row_count = self.session.execute(
text('SELECT count(*) FROM (SELECT 1 FROM document_chunk LIMIT :min_training_rows) AS sample'),
{'min_training_rows': min_training_rows},
).scalar()
if row_count < min_training_rows:
log.info(
"Deferring vector index 'idx_document_chunk_vector' until document_chunk holds %s rows to cluster on, "
'it has %s. Searches run as an exact scan until then.',
min_training_rows,
row_count,
)
return False
return True
def _iterative_scan_setting(self, index_method: str) -> Optional[str]:
if PGVECTOR_ITERATIVE_SCAN == 'off':
return None
version = self.session.execute(text("SELECT extversion FROM pg_extension WHERE extname = 'vector'")).scalar()
version_parts = [int(part) for part in (version or '').split('.') if part.isdigit()]
if version_parts[:2] < [0, 8]:
log.info('Iterative scan needs pgvector 0.8 or newer, the server has %s.', version or 'none')
return None
# ivfflat only accepts relaxed_order
mode = 'relaxed_order' if index_method == 'ivfflat' else PGVECTOR_ITERATIVE_SCAN
return f'SET LOCAL {index_method}.iterative_scan = {mode}'
def _ensure_text_search_index(self) -> None:
if PGVECTOR_PGCRYPTO:
return
self.session.execute(
text("""
CREATE INDEX IF NOT EXISTS idx_document_chunk_text_search
ON document_chunk
USING GIN (to_tsvector('simple', coalesce(text, '')));
""")
)
log.info("Ensured text search index 'idx_document_chunk_text_search'.")
def check_vector_length(self) -> None: def check_vector_length(self) -> None:
""" """
Check if the VECTOR_LENGTH matches the existing vector column dimension in the database. Check if the VECTOR_LENGTH matches the existing vector column dimension in the database.
@ -340,7 +289,7 @@ class PgvectorClient(VectorDBBase):
# Use raw SQL for BYTEA/pgcrypto # Use raw SQL for BYTEA/pgcrypto
# Ensure metadata is converted to its JSON text representation # Ensure metadata is converted to its JSON text representation
# Sanitize to strip null bytes / surrogates that PostgreSQL cannot store # Sanitize to strip null bytes / surrogates that PostgreSQL cannot store
json_metadata = sanitize_text_for_db(JSONCodec.dumps(item['metadata'])) json_metadata = sanitize_text_for_db(json.dumps(item['metadata']))
item_text = sanitize_text_for_db(item['text']) item_text = sanitize_text_for_db(item['text'])
self.session.execute( self.session.execute(
text(""" text("""
@ -363,7 +312,7 @@ class PgvectorClient(VectorDBBase):
}, },
) )
self.session.commit() self.session.commit()
log.info("Encrypted & inserted %s into '%s'", len(items), collection_name) log.info(f"Encrypted & inserted {len(items)} into '{collection_name}'")
else: else:
new_items = [] new_items = []
@ -379,7 +328,7 @@ class PgvectorClient(VectorDBBase):
new_items.append(new_chunk) new_items.append(new_chunk)
self.session.bulk_save_objects(new_items) self.session.bulk_save_objects(new_items)
self.session.commit() self.session.commit()
log.info("Inserted %s items into collection '%s'.", len(new_items), collection_name) log.info(f"Inserted {len(new_items)} items into collection '{collection_name}'.")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Error during insert: {e}') log.exception(f'Error during insert: {e}')
@ -391,7 +340,7 @@ class PgvectorClient(VectorDBBase):
for item in items: for item in items:
vector = self.adjust_vector_length(item['vector']) vector = self.adjust_vector_length(item['vector'])
# Sanitize to strip null bytes / surrogates that PostgreSQL cannot store # Sanitize to strip null bytes / surrogates that PostgreSQL cannot store
json_metadata = sanitize_text_for_db(JSONCodec.dumps(item['metadata'])) json_metadata = sanitize_text_for_db(json.dumps(item['metadata']))
item_text = sanitize_text_for_db(item['text']) item_text = sanitize_text_for_db(item['text'])
self.session.execute( self.session.execute(
text(""" text("""
@ -418,7 +367,7 @@ class PgvectorClient(VectorDBBase):
}, },
) )
self.session.commit() self.session.commit()
log.info("Encrypted & upserted %s into '%s'", len(items), collection_name) log.info(f"Encrypted & upserted {len(items)} into '{collection_name}'")
else: else:
for item in items: for item in items:
vector = self.adjust_vector_length(item['vector']) vector = self.adjust_vector_length(item['vector'])
@ -438,7 +387,7 @@ class PgvectorClient(VectorDBBase):
) )
self.session.add(new_chunk) self.session.add(new_chunk)
self.session.commit() self.session.commit()
log.info("Upserted %s items into collection '%s'.", len(items), collection_name) log.info(f"Upserted {len(items)} items into collection '{collection_name}'.")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Error during upsert: {e}') log.exception(f'Error during upsert: {e}')
@ -540,9 +489,6 @@ class PgvectorClient(VectorDBBase):
.order_by(query_vectors.c.qid, subq.c.distance) .order_by(query_vectors.c.qid, subq.c.distance)
) )
if self.iterative_scan_sql:
self.session.execute(text(self.iterative_scan_sql))
result_proxy = self.session.execute(stmt) result_proxy = self.session.execute(stmt)
results = result_proxy.all() results = result_proxy.all()
@ -552,7 +498,6 @@ class PgvectorClient(VectorDBBase):
metadatas = [[] for _ in range(num_queries)] metadatas = [[] for _ in range(num_queries)]
if not results: if not results:
self.session.rollback()
return SearchResult( return SearchResult(
ids=ids, ids=ids,
distances=distances, distances=distances,
@ -576,71 +521,6 @@ class PgvectorClient(VectorDBBase):
log.exception(f'Error during search: {e}') log.exception(f'Error during search: {e}')
return None return None
def hybrid_search(
self,
collection_name: str,
query: str,
vectors: List[List[float]],
filter: Optional[Dict[str, Any]] = None,
limit: int = 10,
hybrid_bm25_weight: float = 0.5,
) -> Optional[SearchResult]:
if PGVECTOR_PGCRYPTO or filter:
return None
try:
limit = max(1, limit)
vectors = [self.adjust_vector_length(vector) for vector in vectors] if vectors else []
num_queries = len(vectors) if vectors else 1
bm25_weight = min(max(hybrid_bm25_weight, 0.0), 1.0)
vector_weight = 1.0 - bm25_weight
vector_result = None
if vector_weight > 0 and vectors:
vector_result = self.search(collection_name=collection_name, vectors=vectors, limit=limit)
fts_results = []
if bm25_weight > 0 and query and query.strip():
fts_rows = self.session.execute(
text("""
WITH fts_query AS (
SELECT plainto_tsquery('simple', :query) AS query
)
SELECT
document_chunk.id AS id,
document_chunk.text AS text,
document_chunk.vmetadata AS vmetadata,
ts_rank_cd(
to_tsvector('simple', coalesce(document_chunk.text, '')),
fts_query.query
) AS rank
FROM document_chunk, fts_query
WHERE document_chunk.collection_name = :collection_name
AND to_tsvector('simple', coalesce(document_chunk.text, '')) @@ fts_query.query
ORDER BY rank DESC
LIMIT :limit
"""),
{
'collection_name': collection_name,
'query': query,
'limit': limit,
},
)
fts_results = [dict(row) for row in fts_rows.mappings().all()]
self.session.rollback()
return merge_hybrid_search_results(
vector_result=vector_result,
fts_results=fts_results,
num_queries=num_queries,
limit=limit,
hybrid_bm25_weight=hybrid_bm25_weight,
)
except Exception as e:
self.session.rollback()
log.exception(f'Error during hybrid search: {e}')
return None
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]: def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]:
try: try:
if PGVECTOR_PGCRYPTO: if PGVECTOR_PGCRYPTO:
@ -672,7 +552,6 @@ class PgvectorClient(VectorDBBase):
results = query.all() results = query.all()
if not results: if not results:
self.session.rollback()
return None return None
ids = [[result.id for result in results]] ids = [[result.id for result in results]]
@ -712,7 +591,6 @@ class PgvectorClient(VectorDBBase):
results = query.all() results = query.all()
if not results: if not results:
self.session.rollback()
return None return None
ids = [[result.id for result in results]] ids = [[result.id for result in results]]
@ -755,7 +633,7 @@ class PgvectorClient(VectorDBBase):
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value)) query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
deleted = query.delete(synchronize_session=False) deleted = query.delete(synchronize_session=False)
self.session.commit() self.session.commit()
log.info("Deleted %s items from collection '%s'.", deleted, collection_name) log.info(f"Deleted {deleted} items from collection '{collection_name}'.")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Error during delete: {e}') log.exception(f'Error during delete: {e}')
@ -765,7 +643,7 @@ class PgvectorClient(VectorDBBase):
try: try:
deleted = self.session.query(DocumentChunk).delete() deleted = self.session.query(DocumentChunk).delete()
self.session.commit() self.session.commit()
log.info("Reset complete. Deleted %s items from 'document_chunk' table.", deleted) log.info(f"Reset complete. Deleted {deleted} items from 'document_chunk' table.")
except Exception as e: except Exception as e:
self.session.rollback() self.session.rollback()
log.exception(f'Error during reset: {e}') log.exception(f'Error during reset: {e}')
@ -789,4 +667,4 @@ class PgvectorClient(VectorDBBase):
def delete_collection(self, collection_name: str) -> None: def delete_collection(self, collection_name: str) -> None:
self.delete(collection_name) self.delete(collection_name)
log.info("Collection '%s' deleted.", collection_name) log.info(f"Collection '{collection_name}' deleted.")

View file

@ -35,7 +35,7 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import normalize_filter, process_metadata from open_webui.retrieval.vector.utils import process_metadata
NO_LIMIT = 10000 # Reasonable limit to avoid overwhelming the system NO_LIMIT = 10000 # Reasonable limit to avoid overwhelming the system
BATCH_SIZE = 100 # Recommended batch size for Pinecone operations BATCH_SIZE = 100 # Recommended batch size for Pinecone operations
@ -106,16 +106,16 @@ class PineconeClient(VectorDBBase):
try: try:
# Check if index exists # Check if index exists
if self.index_name not in self.client.list_indexes().names(): if self.index_name not in self.client.list_indexes().names():
log.info("Creating Pinecone index '%s'...", self.index_name) log.info(f"Creating Pinecone index '{self.index_name}'...")
self.client.create_index( self.client.create_index(
name=self.index_name, name=self.index_name,
dimension=self.dimension, dimension=self.dimension,
metric=self.metric, metric=self.metric,
spec=ServerlessSpec(cloud=self.cloud, region=self.environment), spec=ServerlessSpec(cloud=self.cloud, region=self.environment),
) )
log.info("Successfully created Pinecone index '%s'", self.index_name) log.info(f"Successfully created Pinecone index '{self.index_name}'")
else: else:
log.info("Using existing Pinecone index '%s'", self.index_name) log.info(f"Using existing Pinecone index '{self.index_name}'")
# Connect to the index # Connect to the index
self.index = self.client.Index( self.index = self.client.Index(
@ -245,7 +245,7 @@ class PineconeClient(VectorDBBase):
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name) collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
try: try:
self.index.delete(filter={'collection_name': collection_name_with_prefix}) self.index.delete(filter={'collection_name': collection_name_with_prefix})
log.info("Collection '%s' deleted (all vectors removed).", collection_name_with_prefix) log.info(f"Collection '{collection_name_with_prefix}' deleted (all vectors removed).")
except Exception as e: except Exception as e:
log.warning(f"Failed to delete collection '{collection_name_with_prefix}': {e}") log.warning(f"Failed to delete collection '{collection_name_with_prefix}': {e}")
raise raise
@ -274,9 +274,9 @@ class PineconeClient(VectorDBBase):
log.error(f'Error inserting batch: {e}') log.error(f'Error inserting batch: {e}')
raise raise
elapsed = time.time() - start_time elapsed = time.time() - start_time
log.debug('Insert of %s vectors took %.2f seconds', len(points), elapsed) log.debug(f'Insert of {len(points)} vectors took {elapsed:.2f} seconds')
log.info( log.info(
"Successfully inserted %s vectors in parallel batches into '%s'", len(points), collection_name_with_prefix f"Successfully inserted {len(points)} vectors in parallel batches into '{collection_name_with_prefix}'"
) )
def upsert(self, collection_name: str, items: List[VectorItem]) -> None: def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
@ -303,9 +303,9 @@ class PineconeClient(VectorDBBase):
log.error(f'Error upserting batch: {e}') log.error(f'Error upserting batch: {e}')
raise raise
elapsed = time.time() - start_time elapsed = time.time() - start_time
log.debug('Upsert of %s vectors took %.2f seconds', len(points), elapsed) log.debug(f'Upsert of {len(points)} vectors took {elapsed:.2f} seconds')
log.info( log.info(
"Successfully upserted %s vectors in parallel batches into '%s'", len(points), collection_name_with_prefix f"Successfully upserted {len(points)} vectors in parallel batches into '{collection_name_with_prefix}'"
) )
async def insert_async(self, collection_name: str, items: List[VectorItem]) -> None: async def insert_async(self, collection_name: str, items: List[VectorItem]) -> None:
@ -326,9 +326,7 @@ class PineconeClient(VectorDBBase):
if isinstance(result, Exception): if isinstance(result, Exception):
log.error(f'Error in async insert batch: {result}') log.error(f'Error in async insert batch: {result}')
raise result raise result
log.info( log.info(f"Successfully async inserted {len(points)} vectors in batches into '{collection_name_with_prefix}'")
"Successfully async inserted %s vectors in batches into '%s'", len(points), collection_name_with_prefix
)
async def upsert_async(self, collection_name: str, items: List[VectorItem]) -> None: async def upsert_async(self, collection_name: str, items: List[VectorItem]) -> None:
"""Async version of upsert using asyncio and run_in_executor for improved performance.""" """Async version of upsert using asyncio and run_in_executor for improved performance."""
@ -348,9 +346,7 @@ class PineconeClient(VectorDBBase):
if isinstance(result, Exception): if isinstance(result, Exception):
log.error(f'Error in async upsert batch: {result}') log.error(f'Error in async upsert batch: {result}')
raise result raise result
log.info( log.info(f"Successfully async upserted {len(points)} vectors in batches into '{collection_name_with_prefix}'")
"Successfully async upserted %s vectors in batches into '%s'", len(points), collection_name_with_prefix
)
def search( def search(
self, self,
@ -372,15 +368,13 @@ class PineconeClient(VectorDBBase):
try: try:
# Search using the first vector (assuming this is the intended behavior) # Search using the first vector (assuming this is the intended behavior)
query_vector = vectors[0] query_vector = vectors[0]
pinecone_filter = normalize_filter(filter)
pinecone_filter['collection_name'] = collection_name_with_prefix
# Perform the search # Perform the search
query_response = self.index.query( query_response = self.index.query(
vector=query_vector, vector=query_vector,
top_k=limit, top_k=limit,
include_metadata=True, include_metadata=True,
filter=pinecone_filter, filter={'collection_name': collection_name_with_prefix},
) )
matches = getattr(query_response, 'matches', []) or [] matches = getattr(query_response, 'matches', []) or []
@ -480,10 +474,8 @@ class PineconeClient(VectorDBBase):
# Note: When deleting by ID, we can't filter by collection_name # Note: When deleting by ID, we can't filter by collection_name
# This is a limitation of Pinecone - be careful with ID uniqueness # This is a limitation of Pinecone - be careful with ID uniqueness
self.index.delete(ids=batch_ids) self.index.delete(ids=batch_ids)
log.debug( log.debug(f"Deleted batch of {len(batch_ids)} vectors by ID from '{collection_name_with_prefix}'")
"Deleted batch of %s vectors by ID from '%s'", len(batch_ids), collection_name_with_prefix log.info(f"Successfully deleted {len(ids)} vectors by ID from '{collection_name_with_prefix}'")
)
log.info("Successfully deleted %s vectors by ID from '%s'", len(ids), collection_name_with_prefix)
elif filter: elif filter:
# Combine user filter with collection_name # Combine user filter with collection_name
@ -492,7 +484,7 @@ class PineconeClient(VectorDBBase):
pinecone_filter.update(filter) pinecone_filter.update(filter)
# Delete by metadata filter # Delete by metadata filter
self.index.delete(filter=pinecone_filter) self.index.delete(filter=pinecone_filter)
log.info("Successfully deleted vectors by filter from '%s'", collection_name_with_prefix) log.info(f"Successfully deleted vectors by filter from '{collection_name_with_prefix}'")
else: else:
log.warning('No ids or filter provided for delete operation') log.warning('No ids or filter provided for delete operation')

View file

@ -3,7 +3,7 @@ NOTE: This vector database integration is community-supported and maintained on
""" """
import logging import logging
from typing import Any, Optional from typing import Optional
from urllib.parse import urlparse from urllib.parse import urlparse
from open_webui.config import ( from open_webui.config import (
@ -22,7 +22,6 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
from qdrant_client import QdrantClient as Qclient from qdrant_client import QdrantClient as Qclient
from qdrant_client.http.models import PointStruct from qdrant_client.http.models import PointStruct
from qdrant_client.models import models from qdrant_client.models import models
@ -32,11 +31,6 @@ NO_LIMIT = 999999999
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
def _metadata_filter(key: str, op: str, value: Any) -> models.FieldCondition:
match = models.MatchAny(any=value) if op == '$in' else models.MatchValue(value=value)
return models.FieldCondition(key=f'metadata.{key}', match=match)
class QdrantClient(VectorDBBase): class QdrantClient(VectorDBBase):
def __init__(self): def __init__(self):
self.collection_prefix = QDRANT_COLLECTION_PREFIX self.collection_prefix = QDRANT_COLLECTION_PREFIX
@ -125,7 +119,7 @@ class QdrantClient(VectorDBBase):
on_disk=self.QDRANT_ON_DISK, on_disk=self.QDRANT_ON_DISK,
), ),
) )
log.info('collection %s successfully created!', collection_name_with_prefix) log.info(f'collection {collection_name_with_prefix} successfully created!')
def _create_collection_if_not_exists(self, collection_name, dimension): def _create_collection_if_not_exists(self, collection_name, dimension):
if not self.has_collection(collection_name=collection_name): if not self.has_collection(collection_name=collection_name):
@ -136,7 +130,7 @@ class QdrantClient(VectorDBBase):
PointStruct( PointStruct(
id=item['id'], id=item['id'],
vector=item['vector'], vector=item['vector'],
payload={'text': item['text'], 'metadata': process_metadata(item['metadata'])}, payload={'text': item['text'], 'metadata': item['metadata']},
) )
for item in items for item in items
] ]
@ -158,13 +152,10 @@ class QdrantClient(VectorDBBase):
if limit is None: if limit is None:
limit = NO_LIMIT # otherwise qdrant would set limit to 10! limit = NO_LIMIT # otherwise qdrant would set limit to 10!
conditions = [_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter)]
query_filter = models.Filter(must=conditions) if conditions else None
query_response = self.client.query_points( query_response = self.client.query_points(
collection_name=f'{self.collection_prefix}_{collection_name}', collection_name=f'{self.collection_prefix}_{collection_name}',
query=vectors[0], query=vectors[0],
limit=limit, limit=limit,
query_filter=query_filter,
) )
get_result = self._result_to_get_result(query_response.points) get_result = self._result_to_get_result(query_response.points)
return SearchResult( return SearchResult(

View file

@ -23,7 +23,6 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
from qdrant_client import QdrantClient as Qclient from qdrant_client import QdrantClient as Qclient
from qdrant_client.http.exceptions import UnexpectedResponse from qdrant_client.http.exceptions import UnexpectedResponse
from qdrant_client.http.models import PointStruct from qdrant_client.http.models import PointStruct
@ -40,9 +39,8 @@ def _tenant_filter(tenant_id: str) -> models.FieldCondition:
return models.FieldCondition(key=TENANT_ID_FIELD, match=models.MatchValue(value=tenant_id)) return models.FieldCondition(key=TENANT_ID_FIELD, match=models.MatchValue(value=tenant_id))
def _metadata_filter(key: str, op: str, value: Any) -> models.FieldCondition: def _metadata_filter(key: str, value: Any) -> models.FieldCondition:
match = models.MatchAny(any=value) if op == '$in' else models.MatchValue(value=value) return models.FieldCondition(key=f'metadata.{key}', match=models.MatchValue(value=value))
return models.FieldCondition(key=f'metadata.{key}', match=match)
class QdrantClient(VectorDBBase): class QdrantClient(VectorDBBase):
@ -150,7 +148,7 @@ class QdrantClient(VectorDBBase):
m=0, m=0,
), ),
) )
log.info('Multi-tenant collection %s created with dimension %s!', mt_collection_name, dimension) log.info(f'Multi-tenant collection {mt_collection_name} created with dimension {dimension}!')
self.client.create_payload_index( self.client.create_payload_index(
collection_name=mt_collection_name, collection_name=mt_collection_name,
@ -182,7 +180,7 @@ class QdrantClient(VectorDBBase):
vector=item['vector'], vector=item['vector'],
payload={ payload={
'text': item['text'], 'text': item['text'],
'metadata': process_metadata(item['metadata']), 'metadata': item['metadata'],
TENANT_ID_FIELD: tenant_id, TENANT_ID_FIELD: tenant_id,
}, },
) )
@ -226,7 +224,7 @@ class QdrantClient(VectorDBBase):
mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name) mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name)
if not self.client.collection_exists(collection_name=mt_collection): if not self.client.collection_exists(collection_name=mt_collection):
log.debug("Collection %s doesn't exist, nothing to delete", mt_collection) log.debug(f"Collection {mt_collection} doesn't exist, nothing to delete")
return None return None
must_conditions = [_tenant_filter(tenant_id)] must_conditions = [_tenant_filter(tenant_id)]
@ -236,7 +234,7 @@ class QdrantClient(VectorDBBase):
# whose payload omits an id (e.g. memories), leaving orphaned vectors. # whose payload omits an id (e.g. memories), leaving orphaned vectors.
must_conditions.append(models.HasIdCondition(has_id=ids)) must_conditions.append(models.HasIdCondition(has_id=ids))
elif filter: elif filter:
must_conditions += [_metadata_filter(k, '$eq', v) for k, v in filter.items()] must_conditions += [_metadata_filter(k, v) for k, v in filter.items()]
return self.client.delete( return self.client.delete(
collection_name=mt_collection, collection_name=mt_collection,
@ -257,17 +255,15 @@ class QdrantClient(VectorDBBase):
return None return None
mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name) mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name)
if not self.client.collection_exists(collection_name=mt_collection): if not self.client.collection_exists(collection_name=mt_collection):
log.debug("Collection %s doesn't exist, search returns None", mt_collection) log.debug(f"Collection {mt_collection} doesn't exist, search returns None")
return None return None
conditions = [_tenant_filter(tenant_id)] tenant_filter = _tenant_filter(tenant_id)
if filter:
conditions.extend(_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter))
query_response = self.client.query_points( query_response = self.client.query_points(
collection_name=mt_collection, collection_name=mt_collection,
query=vectors[0], query=vectors[0],
limit=limit, limit=limit,
query_filter=models.Filter(must=conditions), query_filter=models.Filter(must=[tenant_filter]),
) )
get_result = self._result_to_get_result(query_response.points) get_result = self._result_to_get_result(query_response.points)
return SearchResult( return SearchResult(
@ -285,12 +281,12 @@ class QdrantClient(VectorDBBase):
return None return None
mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name) mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name)
if not self.client.collection_exists(collection_name=mt_collection): if not self.client.collection_exists(collection_name=mt_collection):
log.debug("Collection %s doesn't exist, query returns None", mt_collection) log.debug(f"Collection {mt_collection} doesn't exist, query returns None")
return None return None
if limit is None: if limit is None:
limit = NO_LIMIT limit = NO_LIMIT
tenant_filter = _tenant_filter(tenant_id) tenant_filter = _tenant_filter(tenant_id)
field_conditions = [_metadata_filter(k, '$eq', v) for k, v in filter.items()] field_conditions = [_metadata_filter(k, v) for k, v in filter.items()]
combined_filter = models.Filter(must=[tenant_filter, *field_conditions]) combined_filter = models.Filter(must=[tenant_filter, *field_conditions])
points = self.client.scroll( points = self.client.scroll(
collection_name=mt_collection, collection_name=mt_collection,
@ -307,7 +303,7 @@ class QdrantClient(VectorDBBase):
return None return None
mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name) mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name)
if not self.client.collection_exists(collection_name=mt_collection): if not self.client.collection_exists(collection_name=mt_collection):
log.debug("Collection %s doesn't exist, get returns None", mt_collection) log.debug(f"Collection {mt_collection} doesn't exist, get returns None")
return None return None
tenant_filter = _tenant_filter(tenant_id) tenant_filter = _tenant_filter(tenant_id)
points = self.client.scroll( points = self.client.scroll(
@ -354,7 +350,7 @@ class QdrantClient(VectorDBBase):
return None return None
mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name) mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name)
if not self.client.collection_exists(collection_name=mt_collection): if not self.client.collection_exists(collection_name=mt_collection):
log.debug("Collection %s doesn't exist, nothing to delete", mt_collection) log.debug(f"Collection {mt_collection} doesn't exist, nothing to delete")
return None return None
self.client.delete( self.client.delete(
collection_name=mt_collection, collection_name=mt_collection,

View file

@ -13,7 +13,7 @@ from open_webui.retrieval.vector.main import (
VectorDBBase, VectorDBBase,
VectorItem, VectorItem,
) )
from open_webui.retrieval.vector.utils import metadata_matches_filter, normalize_filter, process_metadata from open_webui.retrieval.vector.utils import process_metadata
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -36,7 +36,7 @@ class S3VectorClient(VectorDBBase):
if self.bucket_name and self.region: if self.bucket_name and self.region:
try: try:
self.client = boto3.client('s3vectors', region_name=self.region) self.client = boto3.client('s3vectors', region_name=self.region)
log.info("S3Vector client initialized for bucket '%s' in region '%s'", self.bucket_name, self.region) log.info(f"S3Vector client initialized for bucket '{self.bucket_name}' in region '{self.region}'")
except Exception as e: except Exception as e:
log.error(f'Failed to initialize S3Vector client: {e}') log.error(f'Failed to initialize S3Vector client: {e}')
self.client = None self.client = None
@ -54,7 +54,7 @@ class S3VectorClient(VectorDBBase):
Create a new index in the S3 vector bucket for the given collection if it does not exist. Create a new index in the S3 vector bucket for the given collection if it does not exist.
""" """
if self.has_collection(index_name): if self.has_collection(index_name):
log.debug("Index '%s' already exists, skipping creation", index_name) log.debug(f"Index '{index_name}' already exists, skipping creation")
return return
try: try:
@ -70,9 +70,7 @@ class S3VectorClient(VectorDBBase):
] ]
}, },
) )
log.info( log.info(f'Created S3 index: {index_name} (dim={dimension}, type={data_type}, metric={distance_metric})')
'Created S3 index: %s (dim=%s, type=%s, metric=%s)', index_name, dimension, data_type, distance_metric
)
except Exception as e: except Exception as e:
log.error(f"Error creating S3 index '{index_name}': {e}") log.error(f"Error creating S3 index '{index_name}': {e}")
raise raise
@ -139,9 +137,9 @@ class S3VectorClient(VectorDBBase):
return return
try: try:
log.info("Deleting collection '%s'", collection_name) log.info(f"Deleting collection '{collection_name}'")
self.client.delete_index(vectorBucketName=self.bucket_name, indexName=collection_name) self.client.delete_index(vectorBucketName=self.bucket_name, indexName=collection_name)
log.info("Successfully deleted collection '%s'", collection_name) log.info(f"Successfully deleted collection '{collection_name}'")
except Exception as e: except Exception as e:
log.error(f"Error deleting collection '{collection_name}': {e}") log.error(f"Error deleting collection '{collection_name}': {e}")
raise raise
@ -158,7 +156,7 @@ class S3VectorClient(VectorDBBase):
try: try:
if not self.has_collection(collection_name): if not self.has_collection(collection_name):
log.info("Index '%s' does not exist. Creating index.", collection_name) log.info(f"Index '{collection_name}' does not exist. Creating index.")
self._create_index( self._create_index(
index_name=collection_name, index_name=collection_name,
dimension=dimension, dimension=dimension,
@ -204,11 +202,9 @@ class S3VectorClient(VectorDBBase):
indexName=collection_name, indexName=collection_name,
vectors=batch, vectors=batch,
) )
log.info( log.info(f"Inserted batch {i // batch_size + 1}: {len(batch)} vectors into index '{collection_name}'.")
"Inserted batch %s: %s vectors into index '%s'.", i // batch_size + 1, len(batch), collection_name
)
log.info("Completed insertion of %s vectors into index '%s'.", len(vectors), collection_name) log.info(f"Completed insertion of {len(vectors)} vectors into index '{collection_name}'.")
except Exception as e: except Exception as e:
log.error(f'Error inserting vectors: {e}') log.error(f'Error inserting vectors: {e}')
raise raise
@ -222,11 +218,11 @@ class S3VectorClient(VectorDBBase):
return return
dimension = len(items[0]['vector']) dimension = len(items[0]['vector'])
log.info('Upsert dimension: %s', dimension) log.info(f'Upsert dimension: {dimension}')
try: try:
if not self.has_collection(collection_name): if not self.has_collection(collection_name):
log.info("Index '%s' does not exist. Creating index for upsert.", collection_name) log.info(f"Index '{collection_name}' does not exist. Creating index for upsert.")
self._create_index( self._create_index(
index_name=collection_name, index_name=collection_name,
dimension=dimension, dimension=dimension,
@ -268,14 +264,10 @@ class S3VectorClient(VectorDBBase):
batch = vectors[i : i + batch_size] batch = vectors[i : i + batch_size]
if i == 0: # Log sample info for first batch only if i == 0: # Log sample info for first batch only
log.info( log.info(
'Upserting batch 1: %s vectors. First vector sample: key=%s, data_type=%s, data_len=%s', f'Upserting batch 1: {len(batch)} vectors. First vector sample: key={batch[0]["key"]}, data_type={type(batch[0]["data"]["float32"])}, data_len={len(batch[0]["data"]["float32"])}'
len(batch),
batch[0]['key'],
type(batch[0]['data']['float32']),
len(batch[0]['data']['float32']),
) )
else: else:
log.info('Upserting batch %s: %s vectors.', i // batch_size + 1, len(batch)) log.info(f'Upserting batch {i // batch_size + 1}: {len(batch)} vectors.')
self.client.put_vectors( self.client.put_vectors(
vectorBucketName=self.bucket_name, vectorBucketName=self.bucket_name,
@ -283,7 +275,7 @@ class S3VectorClient(VectorDBBase):
vectors=batch, vectors=batch,
) )
log.info("Completed upsert of %s vectors into index '%s'.", len(vectors), collection_name) log.info(f"Completed upsert of {len(vectors)} vectors into index '{collection_name}'.")
except Exception as e: except Exception as e:
log.error(f'Error upserting vectors: {e}') log.error(f'Error upserting vectors: {e}')
raise raise
@ -308,8 +300,7 @@ class S3VectorClient(VectorDBBase):
return None return None
try: try:
log.info("Searching collection '%s' with %s query vectors, limit=%s", collection_name, len(vectors), limit) log.info(f"Searching collection '{collection_name}' with {len(vectors)} query vectors, limit={limit}")
vector_filter = normalize_filter(filter)
# Initialize result lists # Initialize result lists
all_ids = [] all_ids = []
@ -319,23 +310,20 @@ class S3VectorClient(VectorDBBase):
# Process each query vector # Process each query vector
for i, query_vector in enumerate(vectors): for i, query_vector in enumerate(vectors):
log.debug('Processing query vector %s/%s', i + 1, len(vectors)) log.debug(f'Processing query vector {i + 1}/{len(vectors)}')
# Prepare the query vector in S3 Vector format # Prepare the query vector in S3 Vector format
query_vector_dict = {'float32': [float(x) for x in query_vector]} query_vector_dict = {'float32': [float(x) for x in query_vector]}
request_params = { # Call S3 Vector query API
'vectorBucketName': self.bucket_name, response = self.client.query_vectors(
'indexName': collection_name, vectorBucketName=self.bucket_name,
'topK': limit, indexName=collection_name,
'queryVector': query_vector_dict, topK=limit,
'returnMetadata': True, queryVector=query_vector_dict,
'returnDistance': True, returnMetadata=True,
} returnDistance=True,
if vector_filter: )
request_params['filter'] = vector_filter
response = self.client.query_vectors(**request_params)
# Process results for this query # Process results for this query
query_ids = [] query_ids = []
@ -350,9 +338,6 @@ class S3VectorClient(VectorDBBase):
vector_metadata = vector.get('metadata', {}) vector_metadata = vector.get('metadata', {})
vector_distance = vector.get('distance', 0.0) vector_distance = vector.get('distance', 0.0)
if vector_filter and not metadata_matches_filter(vector_metadata, vector_filter):
continue
# Extract document text from metadata # Extract document text from metadata
document_text = '' document_text = ''
if isinstance(vector_metadata, dict): if isinstance(vector_metadata, dict):
@ -377,7 +362,7 @@ class S3VectorClient(VectorDBBase):
all_metadatas.append(query_metadatas) all_metadatas.append(query_metadatas)
all_distances.append(query_distances) all_distances.append(query_distances)
log.info('Search completed. Found results for %s queries', len(all_ids)) log.info(f'Search completed. Found results for {len(all_ids)} queries')
# Return SearchResult format # Return SearchResult format
return SearchResult( return SearchResult(
@ -417,7 +402,7 @@ class S3VectorClient(VectorDBBase):
return self.get(collection_name) return self.get(collection_name)
try: try:
log.info("Querying collection '%s' with filter: %s", collection_name, filter) log.info(f"Querying collection '{collection_name}' with filter: {filter}")
# For S3 Vector, we need to use list_vectors and then filter results # For S3 Vector, we need to use list_vectors and then filter results
# Since S3 Vector may not support complex server-side filtering, # Since S3 Vector may not support complex server-side filtering,
@ -452,7 +437,7 @@ class S3VectorClient(VectorDBBase):
if limit and len(filtered_ids) >= limit: if limit and len(filtered_ids) >= limit:
break break
log.info('Filter applied: %s vectors match out of %s total', len(filtered_ids), len(all_ids)) log.info(f'Filter applied: {len(filtered_ids)} vectors match out of {len(all_ids)} total')
# Return GetResult format # Return GetResult format
if filtered_ids: if filtered_ids:
@ -487,7 +472,7 @@ class S3VectorClient(VectorDBBase):
return GetResult(ids=[[]], documents=[[]], metadatas=[[]]) return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
try: try:
log.info("Retrieving all vectors from collection '%s'", collection_name) log.info(f"Retrieving all vectors from collection '{collection_name}'")
# Initialize result lists # Initialize result lists
all_ids = [] all_ids = []
@ -536,7 +521,7 @@ class S3VectorClient(VectorDBBase):
) )
# Log the actual content for debugging # Log the actual content for debugging
log.debug('Document text preview (first 200 chars): %s', str(document_text)[:200]) log.debug(f'Document text preview (first 200 chars): {str(document_text)[:200]}')
else: else:
document_text = vector_id document_text = vector_id
@ -549,7 +534,7 @@ class S3VectorClient(VectorDBBase):
if not next_token: if not next_token:
break break
log.info("Retrieved %s vectors from collection '%s'", len(all_ids), collection_name) log.info(f"Retrieved {len(all_ids)} vectors from collection '{collection_name}'")
# Return in GetResult format # Return in GetResult format
# The Open WebUI GetResult expects lists of lists, so we wrap each list # The Open WebUI GetResult expects lists of lists, so we wrap each list
@ -591,17 +576,17 @@ class S3VectorClient(VectorDBBase):
try: try:
if ids: if ids:
# Delete by specific vector IDs/keys # Delete by specific vector IDs/keys
log.info("Deleting %s vectors by IDs from collection '%s'", len(ids), collection_name) log.info(f"Deleting {len(ids)} vectors by IDs from collection '{collection_name}'")
self.client.delete_vectors( self.client.delete_vectors(
vectorBucketName=self.bucket_name, vectorBucketName=self.bucket_name,
indexName=collection_name, indexName=collection_name,
keys=ids, keys=ids,
) )
log.info("Deleted %s vectors from index '%s'", len(ids), collection_name) log.info(f"Deleted {len(ids)} vectors from index '{collection_name}'")
elif filter: elif filter:
# Handle filter-based deletion # Handle filter-based deletion
log.info("Deleting vectors by filter from collection '%s': %s", collection_name, filter) log.info(f"Deleting vectors by filter from collection '{collection_name}': {filter}")
# If this is a knowledge collection and we have a file_id filter, # If this is a knowledge collection and we have a file_id filter,
# also clean up the corresponding file-specific collection # also clean up the corresponding file-specific collection
@ -610,8 +595,7 @@ class S3VectorClient(VectorDBBase):
file_collection_name = f'file-{file_id}' file_collection_name = f'file-{file_id}'
if self.has_collection(file_collection_name): if self.has_collection(file_collection_name):
log.info( log.info(
"Found related file-specific collection '%s', deleting it to prevent duplicates", f"Found related file-specific collection '{file_collection_name}', deleting it to prevent duplicates"
file_collection_name,
) )
self.delete_collection(file_collection_name) self.delete_collection(file_collection_name)
@ -620,7 +604,7 @@ class S3VectorClient(VectorDBBase):
query_result = self.query(collection_name, filter) query_result = self.query(collection_name, filter)
if query_result and query_result.ids and query_result.ids[0]: if query_result and query_result.ids and query_result.ids[0]:
matching_ids = query_result.ids[0] matching_ids = query_result.ids[0]
log.info('Found %s vectors matching filter, deleting them', len(matching_ids)) log.info(f'Found {len(matching_ids)} vectors matching filter, deleting them')
# Delete the matching vectors by ID # Delete the matching vectors by ID
self.client.delete_vectors( self.client.delete_vectors(
@ -628,7 +612,7 @@ class S3VectorClient(VectorDBBase):
indexName=collection_name, indexName=collection_name,
keys=matching_ids, keys=matching_ids,
) )
log.info("Deleted %s vectors from index '%s' using filter", len(matching_ids), collection_name) log.info(f"Deleted {len(matching_ids)} vectors from index '{collection_name}' using filter")
else: else:
log.warning('No vectors found matching the filter criteria') log.warning('No vectors found matching the filter criteria')
else: else:
@ -661,11 +645,11 @@ class S3VectorClient(VectorDBBase):
try: try:
self.client.delete_index(vectorBucketName=self.bucket_name, indexName=index_name) self.client.delete_index(vectorBucketName=self.bucket_name, indexName=index_name)
deleted_count += 1 deleted_count += 1
log.info('Deleted index: %s', index_name) log.info(f'Deleted index: {index_name}')
except Exception as e: except Exception as e:
log.error(f"Error deleting index '{index_name}': {e}") log.error(f"Error deleting index '{index_name}': {e}")
log.info('Reset completed: deleted %s indexes', deleted_count) log.info(f'Reset completed: deleted {deleted_count} indexes')
except Exception as e: except Exception as e:
log.error(f'Error during reset: {e}') log.error(f'Error during reset: {e}')

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