mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-10 03:27:57 +00:00
Compare commits
No commits in common. "main" and "v0.9.6" have entirely different histories.
872 changed files with 102244 additions and 235579 deletions
|
|
@ -18,6 +18,3 @@ uploads
|
||||||
**/*.db
|
**/*.db
|
||||||
_test
|
_test
|
||||||
backend/data/*
|
backend/data/*
|
||||||
|
|
||||||
.venv
|
|
||||||
.git
|
|
||||||
|
|
|
||||||
15
.env.example
15
.env.example
|
|
@ -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
2
.github/FUNDING.yml
vendored
|
|
@ -1 +1 @@
|
||||||
github: open-webui
|
github: tjbck
|
||||||
|
|
|
||||||
136
.github/ISSUE_TEMPLATE/bug_report.yaml
vendored
136
.github/ISSUE_TEMPLATE/bug_report.yaml
vendored
|
|
@ -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!
|
||||||
|
|
|
||||||
4
.github/ISSUE_TEMPLATE/config.yml
vendored
4
.github/ISSUE_TEMPLATE/config.yml
vendored
|
|
@ -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.
|
|
||||||
|
|
|
||||||
85
.github/ISSUE_TEMPLATE/feature_request.yaml
vendored
85
.github/ISSUE_TEMPLATE/feature_request.yaml
vendored
|
|
@ -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.
|
||||||
|
|
|
||||||
125
.github/pull_request_template.md
vendored
125
.github/pull_request_template.md
vendored
|
|
@ -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.
|
||||||
|
|
|
||||||
7
.github/workflows/backend.yaml
vendored
7
.github/workflows/backend.yaml
vendored
|
|
@ -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 .
|
|
||||||
|
|
|
||||||
90
.github/workflows/docker.yaml
vendored
90
.github/workflows/docker.yaml
vendored
|
|
@ -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 }}"
|
||||||
|
|
|
||||||
2
.github/workflows/frontend.yaml
vendored
2
.github/workflows/frontend.yaml
vendored
|
|
@ -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:
|
||||||
|
|
|
||||||
139
.github/workflows/issue-label.yaml
vendored
139
.github/workflows/issue-label.yaml
vendored
|
|
@ -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']
|
|
||||||
});
|
|
||||||
45
.github/workflows/regression.yaml
vendored
45
.github/workflows/regression.yaml
vendored
|
|
@ -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
1
.gitignore
vendored
|
|
@ -310,4 +310,3 @@ dist
|
||||||
cypress/videos
|
cypress/videos
|
||||||
cypress/screenshots
|
cypress/screenshots
|
||||||
.vscode/settings.json
|
.vscode/settings.json
|
||||||
.cptr
|
|
||||||
|
|
|
||||||
1186
CHANGELOG.md
1186
CHANGELOG.md
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
53
Dockerfile
53
Dockerfile
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
88
README.md
88
README.md
|
|
@ -8,9 +8,11 @@
|
||||||

|

|
||||||

|

|
||||||
[](https://discord.gg/5rJgQTnV4s)
|
[](https://discord.gg/5rJgQTnV4s)
|
||||||
[](https://github.com/sponsors/open-webui)
|
[](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 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">
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
@ -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')
|
||||||
|
|
|
||||||
265
backend/open_webui/internal/config.py
Normal file
265
backend/open_webui/internal/config.py
Normal 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)
|
||||||
|
|
@ -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
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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())
|
||||||
|
|
|
||||||
|
|
@ -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}])
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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}'
|
||||||
|
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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')
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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]:
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
@ -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))
|
|
||||||
|
|
@ -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()))
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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]],
|
|
||||||
}
|
|
||||||
|
|
@ -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] = ''
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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}',
|
||||||
},
|
},
|
||||||
|
|
|
||||||
|
|
@ -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())]
|
|
||||||
|
|
@ -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. "
|
||||||
|
|
|
||||||
|
|
@ -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)
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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.')
|
||||||
|
|
|
||||||
|
|
@ -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 ''
|
|
||||||
|
|
@ -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."""
|
||||||
|
|
|
||||||
|
|
@ -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."""
|
||||||
|
|
|
||||||
|
|
@ -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'
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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):
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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}')
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
|
||||||
|
|
@ -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': {
|
||||||
|
|
|
||||||
|
|
@ -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}")
|
||||||
|
|
|
||||||
|
|
@ -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.")
|
||||||
|
|
|
||||||
|
|
@ -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')
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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
Loading…
Add table
Reference in a new issue